Refactor test cases for improved readability and consistency
- Added blank lines for better separation of test cases in multiple test files. - Reformatted event handling in tests for clarity and consistency. - Ensured consistent use of multi-line formatting for dictionary arguments in event handling. - Improved assertions and test descriptions for better understanding. - Updated test cases across various modules including test_stream_state, test_stream_utils, test_summarization, test_thread_selector, test_tool_error_handler, test_tui_widgets, test_ui_runtime, and test_wechat_channel.
This commit is contained in:
+4
-2
@@ -10,8 +10,10 @@ NVIDIA_API_KEY= # build.nvidia.com
|
||||
SILICONFLOW_API_KEY= # siliconflow.cn
|
||||
OPENROUTER_API_KEY= # openrouter.ai
|
||||
ZHIPU_API_KEY= # open.bigmodel.cn
|
||||
CUSTOM_API_KEY= # Your custom OpenAI-compatible endpoint
|
||||
CUSTOM_BASE_URL= # Your custom API base URL (optional)
|
||||
CUSTOM_OPENAI_API_KEY= # Third-party OpenAI-compatible endpoint
|
||||
CUSTOM_OPENAI_BASE_URL= # OpenAI-compatible base URL (optional)
|
||||
CUSTOM_ANTHROPIC_API_KEY= # Third-party Anthropic-compatible endpoint
|
||||
CUSTOM_ANTHROPIC_BASE_URL= # Anthropic-compatible base URL
|
||||
|
||||
# Local models (optional)
|
||||
OLLAMA_BASE_URL= # http://localhost:11434 (default)
|
||||
|
||||
@@ -283,6 +283,7 @@ def _get_default_middleware():
|
||||
]
|
||||
if cfg.enable_ask_user and not cfg.auto_approve:
|
||||
from .middleware.ask_user import AskUserMiddleware
|
||||
|
||||
mw.insert(0, AskUserMiddleware())
|
||||
return mw
|
||||
|
||||
@@ -348,6 +349,7 @@ def create_cli_agent(workspace_dir: str | None = None, checkpointer=None, config
|
||||
|
||||
if checkpointer is None:
|
||||
from langgraph.checkpoint.memory import InMemorySaver # type: ignore[import-untyped]
|
||||
|
||||
checkpointer = InMemorySaver()
|
||||
|
||||
# When no explicit workspace_dir is provided, apply config.default_workdir
|
||||
@@ -396,6 +398,7 @@ def create_cli_agent(workspace_dir: str | None = None, checkpointer=None, config
|
||||
]
|
||||
if cfg.enable_ask_user and not cfg.auto_approve:
|
||||
from .middleware.ask_user import AskUserMiddleware
|
||||
|
||||
mw.insert(0, AskUserMiddleware())
|
||||
|
||||
# Re-load MCP tools from current config (picks up /mcp add changes)
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
"""Enable `python -m EvoScientist` execution."""
|
||||
|
||||
from EvoScientist.cli import main
|
||||
|
||||
main()
|
||||
|
||||
+46
-26
@@ -19,27 +19,37 @@ from deepagents.backends.protocol import (
|
||||
# System path prefixes that should never appear in virtual paths.
|
||||
# If the agent hallucinates an absolute system path, we block it.
|
||||
_SYSTEM_PATH_PREFIXES = (
|
||||
"/Users/", "/home/", "/tmp/", "/var/", "/etc/",
|
||||
"/opt/", "/usr/", "/bin/", "/sbin/", "/dev/",
|
||||
"/proc/", "/sys/", "/root/",
|
||||
"/Users/",
|
||||
"/home/",
|
||||
"/tmp/",
|
||||
"/var/",
|
||||
"/etc/",
|
||||
"/opt/",
|
||||
"/usr/",
|
||||
"/bin/",
|
||||
"/sbin/",
|
||||
"/dev/",
|
||||
"/proc/",
|
||||
"/sys/",
|
||||
"/root/",
|
||||
)
|
||||
|
||||
# Dangerous patterns that could escape the workspace
|
||||
BLOCKED_PATTERNS = [
|
||||
r'~/', # home directory
|
||||
r'\bcd\s+/', # cd to absolute path
|
||||
r'\brm\s+-rf\s+/', # rm -rf with absolute path
|
||||
r"~/", # home directory
|
||||
r"\bcd\s+/", # cd to absolute path
|
||||
r"\brm\s+-rf\s+/", # rm -rf with absolute path
|
||||
]
|
||||
|
||||
# Dangerous commands that should never be executed
|
||||
BLOCKED_COMMANDS = [
|
||||
'sudo',
|
||||
'chmod',
|
||||
'chown',
|
||||
'mkfs',
|
||||
'dd',
|
||||
'shutdown',
|
||||
'reboot',
|
||||
"sudo",
|
||||
"chmod",
|
||||
"chown",
|
||||
"mkfs",
|
||||
"dd",
|
||||
"shutdown",
|
||||
"reboot",
|
||||
]
|
||||
|
||||
|
||||
@@ -50,7 +60,7 @@ def _split_shell_commands(command: str) -> list[str]:
|
||||
"""
|
||||
base_commands: list[str] = []
|
||||
# Split by sequential operators first
|
||||
for segment in re.split(r'\s*(?:&&|\|\||;)\s*', command):
|
||||
for segment in re.split(r"\s*(?:&&|\|\||;)\s*", command):
|
||||
# Then split by pipe
|
||||
for pipe_seg in segment.split("|"):
|
||||
pipe_seg = pipe_seg.strip()
|
||||
@@ -140,7 +150,7 @@ def convert_virtual_paths_in_command(
|
||||
path = match.group(0)
|
||||
|
||||
# Skip content that looks like a URL
|
||||
if '://' in command[max(0, match.start() - 10):match.end() + 10]:
|
||||
if "://" in command[max(0, match.start() - 10) : match.end() + 10]:
|
||||
return path
|
||||
|
||||
# Fix hallucinated system absolute paths that reference the workspace.
|
||||
@@ -152,17 +162,17 @@ def convert_virtual_paths_in_command(
|
||||
marker = f"/{workspace_name}/"
|
||||
idx = path.find(marker)
|
||||
if idx != -1:
|
||||
relative = path[idx + len(marker):]
|
||||
relative = path[idx + len(marker) :]
|
||||
return "./" + relative if relative else "."
|
||||
elif path.endswith(f"/{workspace_name}"):
|
||||
return "."
|
||||
break # Matched system prefix but no workspace → fall through
|
||||
|
||||
# Convert virtual path
|
||||
if path == '/':
|
||||
return '.'
|
||||
if path == "/":
|
||||
return "."
|
||||
else:
|
||||
return '.' + path
|
||||
return "." + path
|
||||
|
||||
# Match pattern: paths starting with / (but not URLs)
|
||||
pattern = r'(?<=\s)/[^\s;|&<>\'"`]*|^/[^\s;|&<>\'"`]*'
|
||||
@@ -209,8 +219,12 @@ class MergedReadOnlyBackend(BackendProtocol):
|
||||
"""
|
||||
|
||||
def __init__(self, primary_dir: str, secondary_dir: str):
|
||||
self._primary = ReadOnlyFilesystemBackend(root_dir=primary_dir, virtual_mode=True)
|
||||
self._secondary = ReadOnlyFilesystemBackend(root_dir=secondary_dir, virtual_mode=True)
|
||||
self._primary = ReadOnlyFilesystemBackend(
|
||||
root_dir=primary_dir, virtual_mode=True
|
||||
)
|
||||
self._secondary = ReadOnlyFilesystemBackend(
|
||||
root_dir=secondary_dir, virtual_mode=True
|
||||
)
|
||||
|
||||
# -- read: try primary first, fall back to secondary --
|
||||
|
||||
@@ -233,7 +247,9 @@ class MergedReadOnlyBackend(BackendProtocol):
|
||||
|
||||
# -- grep_raw: search both, deduplicate --
|
||||
|
||||
def grep_raw(self, pattern: str, path: str | None = None, glob: str | None = None) -> list:
|
||||
def grep_raw(
|
||||
self, pattern: str, path: str | None = None, glob: str | None = None
|
||||
) -> list:
|
||||
results = self._secondary.grep_raw(pattern, path, glob)
|
||||
try:
|
||||
results += self._primary.grep_raw(pattern, path, glob)
|
||||
@@ -244,9 +260,13 @@ class MergedReadOnlyBackend(BackendProtocol):
|
||||
# -- glob_info: merge both --
|
||||
|
||||
def glob_info(self, pattern: str, path: str = "/") -> list:
|
||||
secondary = {item["path"]: item for item in self._secondary.glob_info(pattern, path)}
|
||||
secondary = {
|
||||
item["path"]: item for item in self._secondary.glob_info(pattern, path)
|
||||
}
|
||||
try:
|
||||
primary = {item["path"]: item for item in self._primary.glob_info(pattern, path)}
|
||||
primary = {
|
||||
item["path"]: item for item in self._primary.glob_info(pattern, path)
|
||||
}
|
||||
secondary.update(primary)
|
||||
except Exception:
|
||||
pass
|
||||
@@ -349,7 +369,7 @@ class CustomSandboxBackend(LocalShellBackend):
|
||||
# Auto-strip /<ws_name>/ prefix to prevent nesting
|
||||
ws_prefix = f"/{ws_name}/"
|
||||
if key.startswith(ws_prefix):
|
||||
key = key[len(ws_prefix) - 1:] # "/<ws>/main.py" → "/main.py"
|
||||
key = key[len(ws_prefix) - 1 :] # "/<ws>/main.py" → "/main.py"
|
||||
elif key == f"/{ws_name}":
|
||||
key = "/"
|
||||
|
||||
@@ -359,7 +379,7 @@ class CustomSandboxBackend(LocalShellBackend):
|
||||
# Try to extract path after "<ws_name>/"
|
||||
idx = key.find(ws_prefix)
|
||||
if idx != -1:
|
||||
key = "/" + key[idx + len(ws_prefix):]
|
||||
key = "/" + key[idx + len(ws_prefix) :]
|
||||
elif key.endswith(f"/{ws_name}"):
|
||||
key = "/"
|
||||
else:
|
||||
|
||||
@@ -7,7 +7,12 @@ This module provides an extensible interface for different messaging channels
|
||||
from .base import Channel, RawIncoming, IncomingMessage, OutgoingMessage, chunk_text
|
||||
from .bus import MessageBus, InboundMessage, OutboundMessage
|
||||
from .capabilities import ChannelCapabilities
|
||||
from .channel_manager import ChannelManager, register_channel, create_channel, available_channels
|
||||
from .channel_manager import (
|
||||
ChannelManager,
|
||||
register_channel,
|
||||
create_channel,
|
||||
available_channels,
|
||||
)
|
||||
from .consumer import InboundConsumer
|
||||
from .formatter import UnifiedFormatter
|
||||
from .middleware import TypingManager
|
||||
|
||||
+104
-60
@@ -27,6 +27,7 @@ _logger = logging.getLogger(__name__)
|
||||
|
||||
# ── Text chunking ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def chunk_text(text: str, limit: int) -> list[str]:
|
||||
"""Split text into chunks that respect logical boundaries.
|
||||
|
||||
@@ -173,7 +174,9 @@ async def download_attachment(
|
||||
local_path = media_path(f"{prefix}{safe_name}")
|
||||
|
||||
async with httpx.AsyncClient(proxy=proxy) as client:
|
||||
async with client.stream("GET", url, headers=headers or {}, timeout=30) as resp:
|
||||
async with client.stream(
|
||||
"GET", url, headers=headers or {}, timeout=30
|
||||
) as resp:
|
||||
if resp.status_code != 200:
|
||||
return None, f"[attachment: {filename} - download failed]"
|
||||
|
||||
@@ -203,6 +206,7 @@ async def download_attachment(
|
||||
_logger.warning(f"Failed to download attachment: {e}")
|
||||
return None, f"[attachment: {filename} - download failed]"
|
||||
|
||||
|
||||
# Deprecated aliases — use InboundMessage / OutboundMessage instead.
|
||||
IncomingMessage = InboundMessage
|
||||
OutgoingMessage = OutboundMessage
|
||||
@@ -260,13 +264,17 @@ class Channel(ChannelPlugin, ABC):
|
||||
|
||||
# Auto-configure formatter from capabilities
|
||||
self._formatter = UnifiedFormatter.for_channel(self.capabilities.format_type)
|
||||
self._queue: asyncio.Queue[InboundMessage] = asyncio.Queue(maxsize=queue_maxsize)
|
||||
self._queue: asyncio.Queue[InboundMessage] = asyncio.Queue(
|
||||
maxsize=queue_maxsize
|
||||
)
|
||||
self._running = False
|
||||
|
||||
# Typing indicator — delegated to TypingManager
|
||||
from .middleware import TypingManager
|
||||
|
||||
self._typing_manager = TypingManager(
|
||||
self._send_typing_action, interval=self._typing_interval,
|
||||
self._send_typing_action,
|
||||
interval=self._typing_interval,
|
||||
)
|
||||
# Keep legacy dict reference for any subclass that touches it directly
|
||||
self._typing_tasks = self._typing_manager._tasks
|
||||
@@ -300,6 +308,7 @@ class Channel(ChannelPlugin, ABC):
|
||||
|
||||
# Retry configuration (auto-resolved from channel name)
|
||||
from .retry import RetryConfig, DEFAULT_RETRY, RETRY_PRESETS
|
||||
|
||||
self._retry_config: RetryConfig = RETRY_PRESETS.get(self.name, DEFAULT_RETRY)
|
||||
|
||||
# Per-chat send locks to prevent message reordering.
|
||||
@@ -321,9 +330,13 @@ class Channel(ChannelPlugin, ABC):
|
||||
5. MentionGatingMiddleware — filter by mention policy
|
||||
"""
|
||||
from .middleware import (
|
||||
DedupMiddleware, AllowListMiddleware,
|
||||
PairingMiddleware, GroupHistoryMiddleware, MentionGatingMiddleware,
|
||||
DedupMiddleware,
|
||||
AllowListMiddleware,
|
||||
PairingMiddleware,
|
||||
GroupHistoryMiddleware,
|
||||
MentionGatingMiddleware,
|
||||
)
|
||||
|
||||
middlewares = []
|
||||
middlewares.append(DedupMiddleware())
|
||||
# AllowList
|
||||
@@ -333,29 +346,37 @@ class Channel(ChannelPlugin, ABC):
|
||||
allowed_senders = set(allowed_senders)
|
||||
if allowed_channels and not isinstance(allowed_channels, set):
|
||||
allowed_channels = set(allowed_channels)
|
||||
middlewares.append(AllowListMiddleware(
|
||||
allowed_senders=allowed_senders,
|
||||
allowed_channels=allowed_channels,
|
||||
dm_policy=self.dm_policy,
|
||||
))
|
||||
middlewares.append(
|
||||
AllowListMiddleware(
|
||||
allowed_senders=allowed_senders,
|
||||
allowed_channels=allowed_channels,
|
||||
dm_policy=self.dm_policy,
|
||||
)
|
||||
)
|
||||
# Pairing
|
||||
if self.dm_policy == "pairing":
|
||||
|
||||
async def _send_pair(chat_id, text):
|
||||
await self._send_chunk(chat_id, text, text, None, {})
|
||||
middlewares.append(PairingMiddleware(
|
||||
channel_name=self.name,
|
||||
send_response_fn=_send_pair,
|
||||
dm_policy=self.dm_policy,
|
||||
))
|
||||
|
||||
middlewares.append(
|
||||
PairingMiddleware(
|
||||
channel_name=self.name,
|
||||
send_response_fn=_send_pair,
|
||||
dm_policy=self.dm_policy,
|
||||
)
|
||||
)
|
||||
# GroupHistory
|
||||
if self.capabilities.groups:
|
||||
middlewares.append(GroupHistoryMiddleware())
|
||||
# MentionGating
|
||||
if self.capabilities.mentions:
|
||||
middlewares.append(MentionGatingMiddleware(
|
||||
require_mention=self.require_mention,
|
||||
strip_fn=self._strip_mention,
|
||||
))
|
||||
middlewares.append(
|
||||
MentionGatingMiddleware(
|
||||
require_mention=self.require_mention,
|
||||
strip_fn=self._strip_mention,
|
||||
)
|
||||
)
|
||||
return middlewares
|
||||
|
||||
@abstractmethod
|
||||
@@ -433,7 +454,8 @@ class Channel(ChannelPlugin, ABC):
|
||||
# to prevent unbounded growth.
|
||||
if len(self._send_locks) > self._send_locks_max:
|
||||
to_evict = [
|
||||
k for k, lock in self._send_locks.items()
|
||||
k
|
||||
for k, lock in self._send_locks.items()
|
||||
if not lock.locked() and k != chat_id
|
||||
]
|
||||
for k in to_evict:
|
||||
@@ -475,9 +497,7 @@ class Channel(ChannelPlugin, ABC):
|
||||
)
|
||||
)
|
||||
except Exception as chunk_err:
|
||||
_logger.error(
|
||||
f"{self.name} chunk {i} send error: {chunk_err}"
|
||||
)
|
||||
_logger.error(f"{self.name} chunk {i} send error: {chunk_err}")
|
||||
had_error = True
|
||||
return not had_error
|
||||
except Exception as e:
|
||||
@@ -494,7 +514,9 @@ class Channel(ChannelPlugin, ABC):
|
||||
return reply_to if chunk_index == 0 else None
|
||||
|
||||
def _prepare_chunks(
|
||||
self, content: str, limit: int,
|
||||
self,
|
||||
content: str,
|
||||
limit: int,
|
||||
) -> list[tuple[str, str]]:
|
||||
"""Build ``(formatted, raw)`` pairs, re-splitting when formatting
|
||||
expands a chunk beyond *limit*.
|
||||
@@ -549,8 +571,12 @@ class Channel(ChannelPlugin, ABC):
|
||||
|
||||
@abstractmethod
|
||||
async def _send_chunk(
|
||||
self, chat_id: str, formatted_text: str, raw_text: str,
|
||||
reply_to: str | None, metadata: dict,
|
||||
self,
|
||||
chat_id: str,
|
||||
formatted_text: str,
|
||||
raw_text: str,
|
||||
reply_to: str | None,
|
||||
metadata: dict,
|
||||
) -> None:
|
||||
"""Send a single text chunk. Platform-specific implementation."""
|
||||
...
|
||||
@@ -558,7 +584,10 @@ class Channel(ChannelPlugin, ABC):
|
||||
_format_fallback_patterns: tuple[str, ...] = ("parse", "invalid")
|
||||
|
||||
async def _send_with_format_fallback(
|
||||
self, send_fn: CallableABC[[str], Awaitable], formatted: str, raw: str,
|
||||
self,
|
||||
send_fn: CallableABC[[str], Awaitable],
|
||||
formatted: str,
|
||||
raw: str,
|
||||
) -> None:
|
||||
"""Try *send_fn(formatted)*; on format-related errors retry with *raw*.
|
||||
|
||||
@@ -645,7 +674,8 @@ class Channel(ChannelPlugin, ABC):
|
||||
Delegates to :func:`download_attachment`.
|
||||
"""
|
||||
return await download_attachment(
|
||||
url, filename,
|
||||
url,
|
||||
filename,
|
||||
channel_name=self.name,
|
||||
headers=headers,
|
||||
file_size=file_size,
|
||||
@@ -797,11 +827,15 @@ class Channel(ChannelPlugin, ABC):
|
||||
|
||||
# ── ACK reaction ─────────────────────────────────────────────────
|
||||
|
||||
async def _send_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
async def _send_ack_reaction(
|
||||
self, chat_id: str, message_id: str, emoji: str = "👀"
|
||||
) -> None:
|
||||
"""Send an acknowledgment reaction to a message. Override in subclasses that support reactions."""
|
||||
pass # Default no-op; channels override if they support reactions
|
||||
|
||||
async def _remove_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
async def _remove_ack_reaction(
|
||||
self, chat_id: str, message_id: str, emoji: str = "👀"
|
||||
) -> None:
|
||||
"""Remove the ack reaction after replying. Override in subclasses."""
|
||||
pass
|
||||
|
||||
@@ -837,7 +871,8 @@ class Channel(ChannelPlugin, ABC):
|
||||
|
||||
if loop is not None and loop.is_running():
|
||||
future = asyncio.run_coroutine_threadsafe(
|
||||
self._build_inbound_async(raw), loop,
|
||||
self._build_inbound_async(raw),
|
||||
loop,
|
||||
)
|
||||
return future.result()
|
||||
else:
|
||||
@@ -863,10 +898,16 @@ class Channel(ChannelPlugin, ABC):
|
||||
meta = dict(raw.metadata)
|
||||
meta.setdefault("chat_id", raw.chat_id)
|
||||
return InboundMessage(
|
||||
channel=self.name, sender_id=raw.sender_id, chat_id=raw.chat_id,
|
||||
content=content or "[media only]", timestamp=raw.timestamp,
|
||||
message_id=raw.message_id, media=raw.media_files, metadata=meta,
|
||||
is_group=raw.is_group, was_mentioned=raw.was_mentioned,
|
||||
channel=self.name,
|
||||
sender_id=raw.sender_id,
|
||||
chat_id=raw.chat_id,
|
||||
content=content or "[media only]",
|
||||
timestamp=raw.timestamp,
|
||||
message_id=raw.message_id,
|
||||
media=raw.media_files,
|
||||
metadata=meta,
|
||||
is_group=raw.is_group,
|
||||
was_mentioned=raw.was_mentioned,
|
||||
)
|
||||
|
||||
async def _enqueue_raw(self, raw: RawIncoming) -> None:
|
||||
@@ -921,22 +962,16 @@ class Channel(ChannelPlugin, ABC):
|
||||
self.initial_debounce + (msg_count - 1) * self.debounce_step,
|
||||
self.max_debounce,
|
||||
)
|
||||
_logger.debug(
|
||||
f"Debounce for {sender}: {wait:.1f}s (message #{msg_count})"
|
||||
)
|
||||
_logger.debug(f"Debounce for {sender}: {wait:.1f}s (message #{msg_count})")
|
||||
|
||||
async def debounce_callback(_s=sender, _w=wait):
|
||||
await asyncio.sleep(_w)
|
||||
try:
|
||||
await self._process_buffered_messages(_s)
|
||||
except Exception as e:
|
||||
_logger.error(
|
||||
f"{self.name} debounce flush error for {_s}: {e}"
|
||||
)
|
||||
_logger.error(f"{self.name} debounce flush error for {_s}: {e}")
|
||||
|
||||
self._debounce_tasks[sender] = asyncio.create_task(
|
||||
debounce_callback()
|
||||
)
|
||||
self._debounce_tasks[sender] = asyncio.create_task(debounce_callback())
|
||||
|
||||
async def _process_buffered_messages(self, sender: str) -> None:
|
||||
"""Flush buffered messages for *sender* and publish to bus."""
|
||||
@@ -954,9 +989,7 @@ class Channel(ChannelPlugin, ABC):
|
||||
return
|
||||
|
||||
merged_content = "\n".join(messages)
|
||||
_logger.info(
|
||||
f"Processing {len(messages)} merged message(s) from {sender}"
|
||||
)
|
||||
_logger.info(f"Processing {len(messages)} merged message(s) from {sender}")
|
||||
|
||||
if self._bus:
|
||||
chat_id = (metadata or {}).get("chat_id", sender)
|
||||
@@ -974,27 +1007,40 @@ class Channel(ChannelPlugin, ABC):
|
||||
await self._bus.publish_inbound(inbound)
|
||||
|
||||
async def _send_status_message(
|
||||
self, sender: str, content: str, metadata: dict | None = None,
|
||||
self,
|
||||
sender: str,
|
||||
content: str,
|
||||
metadata: dict | None = None,
|
||||
) -> None:
|
||||
"""Send a status/intermediate message to the channel."""
|
||||
chat_id = (metadata or {}).get("chat_id", sender)
|
||||
await self.send(OutboundMessage(
|
||||
channel=self.name,
|
||||
chat_id=str(chat_id),
|
||||
content=content,
|
||||
metadata=metadata or {},
|
||||
))
|
||||
await self.send(
|
||||
OutboundMessage(
|
||||
channel=self.name,
|
||||
chat_id=str(chat_id),
|
||||
content=content,
|
||||
metadata=metadata or {},
|
||||
)
|
||||
)
|
||||
|
||||
async def send_thinking_message(
|
||||
self, sender: str, thinking: str, metadata: dict | None = None,
|
||||
self,
|
||||
sender: str,
|
||||
thinking: str,
|
||||
metadata: dict | None = None,
|
||||
) -> None:
|
||||
"""Send a thinking intermediate message to the channel."""
|
||||
if not self.send_thinking:
|
||||
return
|
||||
await self._send_status_message(sender, f"\U0001f9e0\n{thinking}\n\u23f3", metadata)
|
||||
await self._send_status_message(
|
||||
sender, f"\U0001f9e0\n{thinking}\n\u23f3", metadata
|
||||
)
|
||||
|
||||
async def send_todo_message(
|
||||
self, sender: str, content: str, metadata: dict | None = None,
|
||||
self,
|
||||
sender: str,
|
||||
content: str,
|
||||
metadata: dict | None = None,
|
||||
) -> None:
|
||||
"""Send a todo list intermediate message to the channel."""
|
||||
await self._send_status_message(sender, content, metadata)
|
||||
@@ -1029,9 +1075,7 @@ class Channel(ChannelPlugin, ABC):
|
||||
self._running = should_reconnect
|
||||
|
||||
if self._running:
|
||||
_logger.info(
|
||||
f"Reconnecting {self.name} in {backoff:.1f}s..."
|
||||
)
|
||||
_logger.info(f"Reconnecting {self.name} in {backoff:.1f}s...")
|
||||
await asyncio.sleep(backoff)
|
||||
backoff = min(backoff * 2, max_backoff)
|
||||
|
||||
|
||||
@@ -51,7 +51,9 @@ class MessageBus:
|
||||
# ── subscriber routing ──
|
||||
|
||||
def subscribe_outbound(
|
||||
self, channel: str, callback: OutboundCallback,
|
||||
self,
|
||||
channel: str,
|
||||
callback: OutboundCallback,
|
||||
) -> None:
|
||||
"""Register a callback for outbound messages targeting *channel*."""
|
||||
if channel not in self._outbound_subscribers:
|
||||
@@ -67,7 +69,8 @@ class MessageBus:
|
||||
while self._running:
|
||||
try:
|
||||
msg = await asyncio.wait_for(
|
||||
self.outbound.get(), timeout=1.0,
|
||||
self.outbound.get(),
|
||||
timeout=1.0,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
@@ -79,9 +82,7 @@ class MessageBus:
|
||||
try:
|
||||
await callback(msg)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error dispatching to {msg.channel}: {e}"
|
||||
)
|
||||
logger.error(f"Error dispatching to {msg.channel}: {e}")
|
||||
|
||||
def stop(self) -> None:
|
||||
"""Stop the dispatcher loop."""
|
||||
|
||||
@@ -27,35 +27,35 @@ class ChannelCapabilities:
|
||||
max_file_size: int = 20 * 1024 * 1024 # 20 MB
|
||||
|
||||
# ── Interaction capabilities ────────────────────────────────────
|
||||
streaming: bool = False # edit-in-place streaming output
|
||||
threading: bool = False # message threads / topics
|
||||
reactions: bool = False # emoji reactions on messages
|
||||
typing: bool = False # typing indicator API
|
||||
streaming: bool = False # edit-in-place streaming output
|
||||
threading: bool = False # message threads / topics
|
||||
reactions: bool = False # emoji reactions on messages
|
||||
typing: bool = False # typing indicator API
|
||||
inline_buttons: bool = False # inline keyboard / action buttons
|
||||
|
||||
# ── Media capabilities ──────────────────────────────────────────
|
||||
media_send: bool = False # can send files/images
|
||||
media_receive: bool = False # can receive files/images
|
||||
voice: bool = False # platform has voice/audio messages that arrive as downloadable files (receive only, not bot sending)
|
||||
stickers: bool = False # supports sticker receive (not bot sending)
|
||||
location: bool = False # supports location message receive (not bot sending)
|
||||
video: bool = False # video messages
|
||||
media_send: bool = False # can send files/images
|
||||
media_receive: bool = False # can receive files/images
|
||||
voice: bool = False # platform has voice/audio messages that arrive as downloadable files (receive only, not bot sending)
|
||||
stickers: bool = False # supports sticker receive (not bot sending)
|
||||
location: bool = False # supports location message receive (not bot sending)
|
||||
video: bool = False # video messages
|
||||
|
||||
# ── Group features ──────────────────────────────────────────────
|
||||
groups: bool = False # group chat support
|
||||
mentions: bool = False # @mention detection
|
||||
groups: bool = False # group chat support
|
||||
mentions: bool = False # @mention detection
|
||||
|
||||
# ── Rich text ───────────────────────────────────────────────────
|
||||
markdown: bool = False # supports Markdown rendering
|
||||
html: bool = False # supports HTML rendering
|
||||
markdown: bool = False # supports Markdown rendering
|
||||
html: bool = False # supports HTML rendering
|
||||
|
||||
# ── Extended capabilities ────────────────────────────────────────
|
||||
chat_types: tuple[str, ...] = () # ("direct", "group", "channel", "thread")
|
||||
edit: bool = False # message editing after send
|
||||
unsend: bool = False # message recall / unsend
|
||||
block_streaming: bool = False # block edit-in-place streaming
|
||||
native_commands: bool = False # platform-native slash commands
|
||||
polls: bool = False # poll / vote messages
|
||||
chat_types: tuple[str, ...] = () # ("direct", "group", "channel", "thread")
|
||||
edit: bool = False # message editing after send
|
||||
unsend: bool = False # message recall / unsend
|
||||
block_streaming: bool = False # block edit-in-place streaming
|
||||
native_commands: bool = False # platform-native slash commands
|
||||
polls: bool = False # poll / vote messages
|
||||
|
||||
def supports(self, feature: str) -> bool:
|
||||
"""Check if a feature is supported by name."""
|
||||
@@ -70,7 +70,7 @@ TELEGRAM = ChannelCapabilities(
|
||||
format_type="html",
|
||||
max_text_length=4000,
|
||||
streaming=False, # could edit messages, but not implemented yet
|
||||
threading=False, # topics exist but not used yet
|
||||
threading=False, # topics exist but not used yet
|
||||
reactions=True,
|
||||
typing=True,
|
||||
media_send=True,
|
||||
@@ -209,12 +209,12 @@ EMAIL = ChannelCapabilities(
|
||||
IMESSAGE = ChannelCapabilities(
|
||||
format_type="plain",
|
||||
max_text_length=999_999,
|
||||
typing=False, # Apple does not expose typing indicator API
|
||||
typing=False, # Apple does not expose typing indicator API
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
voice=True,
|
||||
groups=True,
|
||||
mentions=False, # iMessage has no @mention concept
|
||||
reactions=False, # imsg CLI cannot send tapback reactions
|
||||
mentions=False, # iMessage has no @mention concept
|
||||
reactions=False, # imsg CLI cannot send tapback reactions
|
||||
chat_types=("direct", "group"),
|
||||
)
|
||||
|
||||
@@ -33,6 +33,7 @@ logger = logging.getLogger(__name__)
|
||||
# Account management (formerly account.py)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChannelAccountSnapshot:
|
||||
"""Point-in-time snapshot of a single account's connection state."""
|
||||
@@ -125,13 +126,16 @@ class AccountManager:
|
||||
try:
|
||||
account_config = config
|
||||
if plugin.config_adapter is not None and config is not None:
|
||||
account_config = plugin.config_adapter.resolve_account(config, account_id)
|
||||
account_config = plugin.config_adapter.resolve_account(
|
||||
config, account_id
|
||||
)
|
||||
|
||||
await plugin.start(account_config, account_id=account_id)
|
||||
state.status = "running"
|
||||
state.started_at = time.monotonic()
|
||||
state.snapshot = ChannelAccountSnapshot(
|
||||
account_id=account_id, channel=channel_id,
|
||||
account_id=account_id,
|
||||
channel=channel_id,
|
||||
)
|
||||
state.snapshot.mark_connected()
|
||||
logger.info(f"Account {key} started")
|
||||
@@ -194,7 +198,8 @@ class AccountManager:
|
||||
|
||||
for account_id in adapter.list_account_ids(config):
|
||||
if adapter.is_enabled(
|
||||
adapter.resolve_account(config, account_id), config,
|
||||
adapter.resolve_account(config, account_id),
|
||||
config,
|
||||
):
|
||||
try:
|
||||
await self.start_account(channel_id, account_id, config)
|
||||
@@ -217,23 +222,26 @@ class AccountManager:
|
||||
logger.error(f"Failed to stop account {cid}:{aid}: {e}")
|
||||
|
||||
def get_state(
|
||||
self, channel_id: str, account_id: str,
|
||||
self,
|
||||
channel_id: str,
|
||||
account_id: str,
|
||||
) -> AccountState | None:
|
||||
"""Get the runtime state for a specific account."""
|
||||
return self._states.get(self._key(channel_id, account_id))
|
||||
|
||||
def list_accounts(
|
||||
self, channel_id: str | None = None,
|
||||
self,
|
||||
channel_id: str | None = None,
|
||||
) -> list[AccountState]:
|
||||
"""List account states, optionally filtered by channel."""
|
||||
if channel_id is None:
|
||||
return list(self._states.values())
|
||||
return [
|
||||
s for s in self._states.values() if s.channel_id == channel_id
|
||||
]
|
||||
return [s for s in self._states.values() if s.channel_id == channel_id]
|
||||
|
||||
def get_snapshot(
|
||||
self, channel_id: str, account_id: str,
|
||||
self,
|
||||
channel_id: str,
|
||||
account_id: str,
|
||||
) -> ChannelAccountSnapshot | None:
|
||||
"""Get the connection snapshot for a specific account."""
|
||||
state = self._states.get(self._key(channel_id, account_id))
|
||||
@@ -286,6 +294,7 @@ def build_outbound_pipeline(
|
||||
|
||||
# ── Per-channel health tracking ──────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChannelHealth:
|
||||
"""Tracks send success / failure metrics for a single channel."""
|
||||
@@ -299,6 +308,7 @@ class ChannelHealth:
|
||||
|
||||
# ── Minimal HTTP health-check server ────────────────────────────────
|
||||
|
||||
|
||||
class _HealthServer:
|
||||
"""Zero-dependency HTTP health-check endpoint using ``asyncio.start_server``.
|
||||
|
||||
@@ -318,7 +328,9 @@ class _HealthServer:
|
||||
async def start(self) -> None:
|
||||
self._start_time = time.monotonic()
|
||||
self._server = await asyncio.start_server(
|
||||
self._handle_connection, "0.0.0.0", self._port,
|
||||
self._handle_connection,
|
||||
"0.0.0.0",
|
||||
self._port,
|
||||
)
|
||||
addrs = [s.getsockname() for s in self._server.sockets]
|
||||
logger.info(f"Health server listening on {addrs}")
|
||||
@@ -449,8 +461,7 @@ def create_channel(name: str, config) -> Channel:
|
||||
factory = _CHANNEL_REGISTRY.get(name)
|
||||
if not factory:
|
||||
raise ValueError(
|
||||
f"Unknown channel type: {name}. "
|
||||
f"Available: {list(_CHANNEL_REGISTRY.keys())}"
|
||||
f"Unknown channel type: {name}. Available: {list(_CHANNEL_REGISTRY.keys())}"
|
||||
)
|
||||
return factory(config)
|
||||
|
||||
@@ -504,6 +515,7 @@ def _ensure_channels_registered(types: list[str] | None = None) -> None:
|
||||
|
||||
# ── Shared webhook server ─────────────────────────────────────────
|
||||
|
||||
|
||||
class SharedWebhookServer:
|
||||
"""Single aiohttp server that hosts routes from multiple HTTP channels.
|
||||
|
||||
@@ -601,7 +613,9 @@ class ChannelManager:
|
||||
bus = MessageBus()
|
||||
shared_webhook_port = getattr(config, "shared_webhook_port", 0) or 0
|
||||
manager = cls(bus, shared_webhook_port=shared_webhook_port)
|
||||
types = [t.strip() for t in (config.channel_enabled or "").split(",") if t.strip()]
|
||||
types = [
|
||||
t.strip() for t in (config.channel_enabled or "").split(",") if t.strip()
|
||||
]
|
||||
if not types:
|
||||
raise ValueError("No channels enabled")
|
||||
_ensure_channels_registered(types)
|
||||
@@ -666,9 +680,7 @@ class ChannelManager:
|
||||
# Start shared webhook server before individual channels
|
||||
await self._setup_shared_webhook()
|
||||
|
||||
self._dispatch_task = asyncio.create_task(
|
||||
self._dispatch_outbound()
|
||||
)
|
||||
self._dispatch_task = asyncio.create_task(self._dispatch_outbound())
|
||||
|
||||
now = datetime.now()
|
||||
for name, channel in self._channels.items():
|
||||
@@ -777,8 +789,7 @@ class ChannelManager:
|
||||
channel._shared_webhook_server = True # type: ignore[attr-defined]
|
||||
all_routes.extend(routes)
|
||||
logger.debug(
|
||||
f"Shared webhook: collected {len(routes)} route(s) "
|
||||
f"from '{name}'"
|
||||
f"Shared webhook: collected {len(routes)} route(s) from '{name}'"
|
||||
)
|
||||
|
||||
if not all_routes:
|
||||
@@ -791,7 +802,9 @@ class ChannelManager:
|
||||
await self._shared_webhook_server.start(all_routes)
|
||||
|
||||
def register_health_provider(
|
||||
self, name: str, provider: Callable[[], dict],
|
||||
self,
|
||||
name: str,
|
||||
provider: Callable[[], dict],
|
||||
) -> None:
|
||||
"""Register a callable that returns extra data for ``/healthz``."""
|
||||
self._health_providers[name] = provider
|
||||
@@ -804,7 +817,8 @@ class ChannelManager:
|
||||
while True:
|
||||
try:
|
||||
msg: OutboundMessage = await asyncio.wait_for(
|
||||
self.bus.consume_outbound(), timeout=1.0,
|
||||
self.bus.consume_outbound(),
|
||||
timeout=1.0,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
@@ -848,9 +862,7 @@ class ChannelManager:
|
||||
)
|
||||
delivery_failed = True
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error sending media to {msg.channel}: {e}"
|
||||
)
|
||||
logger.error(f"Error sending media to {msg.channel}: {e}")
|
||||
delivery_failed = True
|
||||
|
||||
if delivery_failed:
|
||||
@@ -862,9 +874,7 @@ class ChannelManager:
|
||||
health.consecutive_failures = 0
|
||||
health.total_successes += 1
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Error sending to {msg.channel}: {e}"
|
||||
)
|
||||
logger.error(f"Error sending to {msg.channel}: {e}")
|
||||
health = self._health.get(msg.channel)
|
||||
if health is not None:
|
||||
health.consecutive_failures += 1
|
||||
|
||||
@@ -44,7 +44,9 @@ class SingleAccountConfigAdapter:
|
||||
return ["default"]
|
||||
|
||||
def resolve_account(
|
||||
self, config: Any, account_id: str | None = None,
|
||||
self,
|
||||
config: Any,
|
||||
account_id: str | None = None,
|
||||
) -> Any:
|
||||
return config
|
||||
|
||||
@@ -59,10 +61,7 @@ class SingleAccountConfigAdapter:
|
||||
return bool(account)
|
||||
# dataclass / object — check that at least one field is truthy
|
||||
if hasattr(account, "__dataclass_fields__"):
|
||||
return any(
|
||||
getattr(account, f, None)
|
||||
for f in account.__dataclass_fields__
|
||||
)
|
||||
return any(getattr(account, f, None) for f in account.__dataclass_fields__)
|
||||
return True
|
||||
|
||||
|
||||
@@ -101,7 +100,9 @@ class MultiAccountConfigAdapter:
|
||||
return list(self._get_accounts_map(config).keys())
|
||||
|
||||
def resolve_account(
|
||||
self, config: Any, account_id: str | None = None,
|
||||
self,
|
||||
config: Any,
|
||||
account_id: str | None = None,
|
||||
) -> Any:
|
||||
accounts = self._get_accounts_map(config)
|
||||
if account_id is None:
|
||||
|
||||
@@ -28,7 +28,9 @@ _MAX_CHAT_LOCKS = 10_000
|
||||
_MAX_SESSIONS = 10_000
|
||||
_MAX_HITL_ROUNDS = 50
|
||||
_HITL_APPROVAL_TIMEOUT = 120.0 # seconds to wait for HITL approval reply
|
||||
_ASK_USER_TIMEOUT = 300.0 # seconds to wait for ask_user reply (longer for thinking time)
|
||||
_ASK_USER_TIMEOUT = (
|
||||
300.0 # seconds to wait for ask_user reply (longer for thinking time)
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -86,6 +88,7 @@ def _should_auto_approve(action_requests: list[dict]) -> bool:
|
||||
|
||||
try:
|
||||
from ..config.settings import load_config
|
||||
|
||||
cfg = load_config()
|
||||
except Exception:
|
||||
return False # fail-closed
|
||||
@@ -95,14 +98,19 @@ def _should_auto_approve(action_requests: list[dict]) -> bool:
|
||||
|
||||
shell_allow_list = (
|
||||
[s.strip() for s in cfg.shell_allow_list.split(",") if s.strip()]
|
||||
if cfg.shell_allow_list else []
|
||||
if cfg.shell_allow_list
|
||||
else []
|
||||
)
|
||||
|
||||
for req in action_requests:
|
||||
name = req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "")
|
||||
name = (
|
||||
req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "")
|
||||
)
|
||||
if name != "execute":
|
||||
continue
|
||||
args = req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {})
|
||||
args = (
|
||||
req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {})
|
||||
)
|
||||
command = args.get("command", "") if isinstance(args, dict) else ""
|
||||
cmd = command.strip()
|
||||
if not any(cmd.startswith(prefix) for prefix in shell_allow_list):
|
||||
@@ -114,8 +122,12 @@ def _format_approval_prompt(action_requests: list[dict]) -> str:
|
||||
"""Format an approval prompt as a text message for channel users."""
|
||||
lines = ["\u26a0\ufe0f Approval Required\n"]
|
||||
for i, req in enumerate(action_requests, 1):
|
||||
name = req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "")
|
||||
args = req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {})
|
||||
name = (
|
||||
req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "")
|
||||
)
|
||||
args = (
|
||||
req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {})
|
||||
)
|
||||
if isinstance(args, dict):
|
||||
command = args.get("command", args.get("path", ""))
|
||||
else:
|
||||
@@ -148,6 +160,7 @@ def _parse_approval_reply(text: str) -> str | None:
|
||||
@dataclass
|
||||
class _PendingInterrupt:
|
||||
"""Stored state for a pending HITL interrupt awaiting channel user reply."""
|
||||
|
||||
thread_id: str
|
||||
action_requests: list
|
||||
event: asyncio.Event # set when user replies
|
||||
@@ -157,6 +170,7 @@ class _PendingInterrupt:
|
||||
@dataclass
|
||||
class _PendingAskUserReply:
|
||||
"""Stored state for a pending ask_user question awaiting channel user reply."""
|
||||
|
||||
event: asyncio.Event # set when user replies
|
||||
reply: str | None = None # raw reply text
|
||||
|
||||
@@ -221,7 +235,9 @@ class InboundConsumer:
|
||||
self._on_message_received = on_message_received
|
||||
self._on_streaming_event = on_streaming_event
|
||||
self._on_message_sent = on_message_sent
|
||||
self._sessions: OrderedDict[str, str] = OrderedDict() # sender_id -> thread_id (LRU)
|
||||
self._sessions: OrderedDict[str, str] = (
|
||||
OrderedDict()
|
||||
) # sender_id -> thread_id (LRU)
|
||||
|
||||
# Per-chat locks: same chat is processed serially (bounded)
|
||||
self._chat_locks: dict[str, asyncio.Lock] = {}
|
||||
@@ -282,14 +298,14 @@ class InboundConsumer:
|
||||
"""
|
||||
self._stopping = False
|
||||
self._workers = [
|
||||
asyncio.create_task(self._worker(i))
|
||||
for i in range(self._max_concurrent)
|
||||
asyncio.create_task(self._worker(i)) for i in range(self._max_concurrent)
|
||||
]
|
||||
try:
|
||||
while not self._stopping:
|
||||
try:
|
||||
msg = await asyncio.wait_for(
|
||||
self.bus.consume_inbound(), timeout=1.0,
|
||||
self.bus.consume_inbound(),
|
||||
timeout=1.0,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
@@ -318,7 +334,8 @@ class InboundConsumer:
|
||||
# Wait for workers to finish, then force-cancel stragglers
|
||||
if self._workers:
|
||||
done, still_running = await asyncio.wait(
|
||||
self._workers, timeout=self._drain_timeout,
|
||||
self._workers,
|
||||
timeout=self._drain_timeout,
|
||||
)
|
||||
for task in still_running:
|
||||
task.cancel()
|
||||
@@ -417,8 +434,12 @@ class InboundConsumer:
|
||||
|
||||
async for event in _timeout_aiter(
|
||||
stream_agent_events(
|
||||
self.agent, stream_input, thread_id,
|
||||
media=msg.media or None if isinstance(stream_input, str) else None,
|
||||
self.agent,
|
||||
stream_input,
|
||||
thread_id,
|
||||
media=msg.media or None
|
||||
if isinstance(stream_input, str)
|
||||
else None,
|
||||
),
|
||||
self._inference_timeout,
|
||||
):
|
||||
@@ -475,7 +496,9 @@ class InboundConsumer:
|
||||
full_thinking = "".join(thinking_buffer)
|
||||
if full_thinking:
|
||||
await channel.send_thinking_message(
|
||||
msg.sender_id, full_thinking, msg.metadata,
|
||||
msg.sender_id,
|
||||
full_thinking,
|
||||
msg.metadata,
|
||||
)
|
||||
|
||||
# No interrupt — normal completion
|
||||
@@ -499,9 +522,12 @@ class InboundConsumer:
|
||||
# ask_user: send questions to channel user, collect answers
|
||||
if interrupt_data.get("type") == "ask_user":
|
||||
result = await self._resolve_ask_user(
|
||||
msg, interrupt_data, session_key,
|
||||
msg,
|
||||
interrupt_data,
|
||||
session_key,
|
||||
)
|
||||
from langgraph.types import Command # type: ignore[import-untyped]
|
||||
|
||||
stream_input = Command(resume=result)
|
||||
continue
|
||||
|
||||
@@ -512,23 +538,31 @@ class InboundConsumer:
|
||||
# Session auto-approve (user previously chose "Approve all")
|
||||
if session_key in self._auto_approve_sessions:
|
||||
from langgraph.types import Command # type: ignore[import-untyped]
|
||||
stream_input = Command(resume={"decisions": [{"type": "approve"} for _ in range(n)]})
|
||||
|
||||
stream_input = Command(
|
||||
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
|
||||
)
|
||||
continue
|
||||
|
||||
# Config auto-approve (auto_approve, non-execute, allow_list)
|
||||
if _should_auto_approve(action_reqs):
|
||||
from langgraph.types import Command # type: ignore[import-untyped]
|
||||
stream_input = Command(resume={"decisions": [{"type": "approve"} for _ in range(n)]})
|
||||
|
||||
stream_input = Command(
|
||||
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
|
||||
)
|
||||
continue
|
||||
|
||||
# Needs user approval — send prompt to channel
|
||||
prompt_text = _format_approval_prompt(action_reqs)
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content=prompt_text,
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content=prompt_text,
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
|
||||
# Wait for user reply
|
||||
pending = _PendingInterrupt(
|
||||
@@ -552,19 +586,24 @@ class InboundConsumer:
|
||||
decision = pending.decision or "approve"
|
||||
|
||||
if decision == "reject":
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="Tool execution rejected.",
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="Tool execution rejected.",
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
return
|
||||
|
||||
if decision == "auto":
|
||||
self._auto_approve_sessions.add(session_key)
|
||||
|
||||
from langgraph.types import Command # type: ignore[import-untyped]
|
||||
stream_input = Command(resume={"decisions": [{"type": "approve"} for _ in range(n)]})
|
||||
|
||||
stream_input = Command(
|
||||
resume={"decisions": [{"type": "approve"} for _ in range(n)]}
|
||||
)
|
||||
# continue to next HITL round
|
||||
|
||||
except asyncio.TimeoutError:
|
||||
@@ -573,22 +612,26 @@ class InboundConsumer:
|
||||
f"Inference timeout ({self._inference_timeout}s idle) "
|
||||
f"for {msg.sender_id} in {session_key}"
|
||||
)
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="Sorry, the response timed out. Please try again.",
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="Sorry, the response timed out. Please try again.",
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
self._metrics.total_failures += 1
|
||||
logger.error(f"Agent error: {e}")
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="Sorry, something went wrong. Please try again later.",
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="Sorry, something went wrong. Please try again later.",
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
finally:
|
||||
if channel:
|
||||
await channel.stop_typing(msg.chat_id)
|
||||
@@ -623,7 +666,9 @@ class InboundConsumer:
|
||||
# ── ask_user helpers ──
|
||||
|
||||
async def _wait_for_ask_user_reply(
|
||||
self, session_key: str, timeout: float,
|
||||
self,
|
||||
session_key: str,
|
||||
timeout: float,
|
||||
) -> str | None:
|
||||
"""Register a pending ask_user slot and wait for the user to reply.
|
||||
|
||||
@@ -684,38 +729,37 @@ class InboundConsumer:
|
||||
lines.append(f" {letter}. {label}")
|
||||
other_letter = chr(ord("A") + len(choices))
|
||||
lines.append(f" {other_letter}. Other")
|
||||
letters = "/".join(
|
||||
chr(ord("A") + k) for k in range(len(choices) + 1)
|
||||
)
|
||||
lines.append(
|
||||
f"\nReply with a letter ({letters}), or 'cancel'."
|
||||
)
|
||||
letters = "/".join(chr(ord("A") + k) for k in range(len(choices) + 1))
|
||||
lines.append(f"\nReply with a letter ({letters}), or 'cancel'.")
|
||||
else:
|
||||
skip_hint = " Leave empty to skip." if not required else ""
|
||||
lines.append(
|
||||
f"\nReply with your answer, or 'cancel'.{skip_hint}"
|
||||
)
|
||||
lines.append(f"\nReply with your answer, or 'cancel'.{skip_hint}")
|
||||
|
||||
# -- Send question --
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="\n".join(lines),
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="\n".join(lines),
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
|
||||
# -- Wait for user reply --
|
||||
reply = await self._wait_for_ask_user_reply(
|
||||
session_key, _ASK_USER_TIMEOUT,
|
||||
session_key,
|
||||
_ASK_USER_TIMEOUT,
|
||||
)
|
||||
|
||||
if not reply:
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="\u23f0 Response timed out.",
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="\u23f0 Response timed out.",
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
return {"status": "cancelled"}
|
||||
|
||||
raw = reply.strip()
|
||||
@@ -728,22 +772,27 @@ class InboundConsumer:
|
||||
other_letter = chr(ord("A") + len(choices))
|
||||
if len(raw) == 1 and raw.upper() == other_letter:
|
||||
# "Other" selected — ask for free-form input
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="Please type your answer:",
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
other_reply = await self._wait_for_ask_user_reply(
|
||||
session_key, _ASK_USER_TIMEOUT,
|
||||
)
|
||||
if not other_reply:
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="\u23f0 Response timed out.",
|
||||
content="Please type your answer:",
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
)
|
||||
)
|
||||
other_reply = await self._wait_for_ask_user_reply(
|
||||
session_key,
|
||||
_ASK_USER_TIMEOUT,
|
||||
)
|
||||
if not other_reply:
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="\u23f0 Response timed out.",
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
return {"status": "cancelled"}
|
||||
if other_reply.strip().lower() == "cancel":
|
||||
return {"status": "cancelled"}
|
||||
@@ -766,5 +815,5 @@ class InboundConsumer:
|
||||
def _evict_chat_locks(self) -> None:
|
||||
"""Remove chat locks that are not currently held."""
|
||||
stale = [k for k, lock in self._chat_locks.items() if not lock.locked()]
|
||||
for k in stale[:max(1, len(stale) // 2)]:
|
||||
for k in stale[: max(1, len(stale) // 2)]:
|
||||
del self._chat_locks[k]
|
||||
|
||||
@@ -18,12 +18,14 @@ __all__ = ["DingTalkChannel", "DingTalkConfig"]
|
||||
def create_from_config(config) -> DingTalkChannel:
|
||||
allowed = _parse_csv(config.dingtalk_allowed_senders)
|
||||
proxy = config.dingtalk_proxy if config.dingtalk_proxy else None
|
||||
return DingTalkChannel(DingTalkConfig(
|
||||
client_id=config.dingtalk_client_id,
|
||||
client_secret=config.dingtalk_client_secret,
|
||||
allowed_senders=allowed,
|
||||
proxy=proxy,
|
||||
))
|
||||
return DingTalkChannel(
|
||||
DingTalkConfig(
|
||||
client_id=config.dingtalk_client_id,
|
||||
client_secret=config.dingtalk_client_secret,
|
||||
allowed_senders=allowed,
|
||||
proxy=proxy,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
register_channel("dingtalk", create_from_config)
|
||||
|
||||
@@ -43,6 +43,7 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
||||
|
||||
async def start(self) -> None:
|
||||
import httpx
|
||||
|
||||
if not self.config.client_id or not self.config.client_secret:
|
||||
raise ChannelError("DingTalk client_id and client_secret are required")
|
||||
self._http_client = httpx.AsyncClient(timeout=15, proxy=self.config.proxy)
|
||||
@@ -54,10 +55,13 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
||||
# ── TokenMixin ────────────────────────────────────────────────
|
||||
|
||||
async def _fetch_token(self) -> tuple[str, int]:
|
||||
data = await self._api_post(TOKEN_URL, {
|
||||
"appKey": self.config.client_id,
|
||||
"appSecret": self.config.client_secret,
|
||||
})
|
||||
data = await self._api_post(
|
||||
TOKEN_URL,
|
||||
{
|
||||
"appKey": self.config.client_id,
|
||||
"appSecret": self.config.client_secret,
|
||||
},
|
||||
)
|
||||
token = data.get("accessToken")
|
||||
if not token:
|
||||
raise ChannelError(f"DingTalk auth error: {data}")
|
||||
@@ -87,12 +91,17 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
||||
# ── WebSocketMixin ────────────────────────────────────────────
|
||||
|
||||
async def _get_ws_url(self) -> str:
|
||||
resp = await self._http_client.post(GATEWAY_URL, json={
|
||||
"clientId": self.config.client_id,
|
||||
"clientSecret": self.config.client_secret,
|
||||
"subscriptions": [{"type": "CALLBACK", "topic": "/v1.0/im/bot/messages/get"}],
|
||||
"ua": "dingtalk-sdk-python/v0.24.3-union",
|
||||
})
|
||||
resp = await self._http_client.post(
|
||||
GATEWAY_URL,
|
||||
json={
|
||||
"clientId": self.config.client_id,
|
||||
"clientSecret": self.config.client_secret,
|
||||
"subscriptions": [
|
||||
{"type": "CALLBACK", "topic": "/v1.0/im/bot/messages/get"}
|
||||
],
|
||||
"ua": "dingtalk-sdk-python/v0.24.3-union",
|
||||
},
|
||||
)
|
||||
data = resp.json()
|
||||
endpoint, ticket = data.get("endpoint"), data.get("ticket")
|
||||
if not endpoint or not ticket:
|
||||
@@ -107,11 +116,25 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
||||
|
||||
# System ping
|
||||
if data.get("type") == "SYSTEM" and headers.get("topic") == "ping":
|
||||
await self._ws_send_json({"code": 200, "headers": headers, "message": "OK", "data": data.get("data", "")})
|
||||
await self._ws_send_json(
|
||||
{
|
||||
"code": 200,
|
||||
"headers": headers,
|
||||
"message": "OK",
|
||||
"data": data.get("data", ""),
|
||||
}
|
||||
)
|
||||
return
|
||||
|
||||
# ACK
|
||||
await self._ws_send_json({"code": 200, "headers": {"contentType": "application/json", "messageId": msg_id}, "message": "OK", "data": "{}"})
|
||||
await self._ws_send_json(
|
||||
{
|
||||
"code": 200,
|
||||
"headers": {"contentType": "application/json", "messageId": msg_id},
|
||||
"message": "OK",
|
||||
"data": "{}",
|
||||
}
|
||||
)
|
||||
|
||||
if data.get("type") != "CALLBACK":
|
||||
return
|
||||
@@ -119,7 +142,9 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
||||
payload = data.get("data", "{}")
|
||||
payload = json.loads(payload) if isinstance(payload, str) else payload
|
||||
text_obj = payload.get("text", {})
|
||||
content = (text_obj.get("content", "") if isinstance(text_obj, dict) else str(text_obj)).strip()
|
||||
content = (
|
||||
text_obj.get("content", "") if isinstance(text_obj, dict) else str(text_obj)
|
||||
).strip()
|
||||
if not content:
|
||||
raw_content = payload.get("content", "")
|
||||
content = raw_content.strip() if isinstance(raw_content, str) else ""
|
||||
@@ -133,12 +158,21 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
||||
# "fileContent"/"imageContent" key.
|
||||
raw_content_obj = payload.get("content")
|
||||
if isinstance(raw_content_obj, dict) and raw_content_obj not in [
|
||||
payload.get(k) for k in ("imageContent", "fileContent", "videoContent", "audioContent")
|
||||
payload.get(k)
|
||||
for k in ("imageContent", "fileContent", "videoContent", "audioContent")
|
||||
]:
|
||||
msg_type = payload.get("msgtype") or payload.get("msgType") or ""
|
||||
media_label = msg_type or "file"
|
||||
file_size = raw_content_obj.get("fileSize") or raw_content_obj.get("downloadSize") or 0
|
||||
file_name = raw_content_obj.get("fileName") or raw_content_obj.get("name") or f"dingtalk_{msg_type}"
|
||||
file_size = (
|
||||
raw_content_obj.get("fileSize")
|
||||
or raw_content_obj.get("downloadSize")
|
||||
or 0
|
||||
)
|
||||
file_name = (
|
||||
raw_content_obj.get("fileName")
|
||||
or raw_content_obj.get("name")
|
||||
or f"dingtalk_{msg_type}"
|
||||
)
|
||||
download_code = raw_content_obj.get("downloadCode") or ""
|
||||
download_url = raw_content_obj.get("downloadUrl") or ""
|
||||
# downloadCode is NOT a URL — resolve it via DingTalk API first
|
||||
@@ -155,7 +189,8 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
||||
except Exception:
|
||||
dl_headers = None
|
||||
local, ann = await self._download_attachment(
|
||||
download_url, f"dingtalk_{file_name}",
|
||||
download_url,
|
||||
f"dingtalk_{file_name}",
|
||||
headers=dl_headers,
|
||||
file_size=int(file_size) if file_size else None,
|
||||
)
|
||||
@@ -183,7 +218,11 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
||||
download_url = download_code
|
||||
# DingTalk audioContent is voice messages
|
||||
media_label = "voice" if att_key == "audioContent" else att_key
|
||||
if download_url and (self.config.include_attachments if hasattr(self.config, 'include_attachments') else True):
|
||||
if download_url and (
|
||||
self.config.include_attachments
|
||||
if hasattr(self.config, "include_attachments")
|
||||
else True
|
||||
):
|
||||
# DingTalk download URLs require access token
|
||||
try:
|
||||
dl_token = await self._ensure_token()
|
||||
@@ -191,7 +230,8 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
||||
except Exception:
|
||||
dl_headers = None
|
||||
local, ann = await self._download_attachment(
|
||||
download_url, f"dingtalk_{file_name}",
|
||||
download_url,
|
||||
f"dingtalk_{file_name}",
|
||||
headers=dl_headers,
|
||||
file_size=int(file_size) if file_size else None,
|
||||
)
|
||||
@@ -231,17 +271,32 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
||||
break
|
||||
|
||||
try:
|
||||
ts = datetime.fromtimestamp(int(create_time) / 1000) if create_time else datetime.now()
|
||||
ts = (
|
||||
datetime.fromtimestamp(int(create_time) / 1000)
|
||||
if create_time
|
||||
else datetime.now()
|
||||
)
|
||||
except (ValueError, TypeError, OSError):
|
||||
ts = datetime.now()
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=sender_id, chat_id=chat_id, text=content, timestamp=ts,
|
||||
message_id=msg_id, is_group=is_group, was_mentioned=was_mentioned,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
metadata={"chat_id": chat_id, "sender_nick": payload.get("senderNick", ""), "backend": "dingtalk"},
|
||||
))
|
||||
await self._enqueue_raw(
|
||||
RawIncoming(
|
||||
sender_id=sender_id,
|
||||
chat_id=chat_id,
|
||||
text=content,
|
||||
timestamp=ts,
|
||||
message_id=msg_id,
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"sender_nick": payload.get("senderNick", ""),
|
||||
"backend": "dingtalk",
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
# _send_typing_action: inherited no-op (DingTalk has no typing API)
|
||||
# _format_chunk: inherited from base (UnifiedFormatter)
|
||||
@@ -250,12 +305,16 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
||||
|
||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||
token = await self._ensure_token()
|
||||
data = await self._api_post(SEND_URL, {
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleMarkdown",
|
||||
"msgParam": json.dumps({"text": raw_text, "title": "EvoScientist"}),
|
||||
}, headers={"x-acs-dingtalk-access-token": token})
|
||||
data = await self._api_post(
|
||||
SEND_URL,
|
||||
{
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleMarkdown",
|
||||
"msgParam": json.dumps({"text": raw_text, "title": "EvoScientist"}),
|
||||
},
|
||||
headers={"x-acs-dingtalk-access-token": token},
|
||||
)
|
||||
return data
|
||||
|
||||
# ── Media send ────────────────────────────────────────────────
|
||||
@@ -284,53 +343,76 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin):
|
||||
# Try uploading image to get media_id for native image message
|
||||
media_id = await self._upload_dingtalk_media(token, file_path, "image")
|
||||
if media_id:
|
||||
await self._api_post(MEDIA_SEND_URL, {
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleImageMsg",
|
||||
"msgParam": json.dumps({"photoURL": media_id}),
|
||||
}, headers=headers)
|
||||
await self._api_post(
|
||||
MEDIA_SEND_URL,
|
||||
{
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleImageMsg",
|
||||
"msgParam": json.dumps({"photoURL": media_id}),
|
||||
},
|
||||
headers=headers,
|
||||
)
|
||||
else:
|
||||
# Fallback to markdown with file path
|
||||
await self._api_post(MEDIA_SEND_URL, {
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleMarkdown",
|
||||
"msgParam": json.dumps({
|
||||
"text": f"" + (f"\n{caption}" if caption else ""),
|
||||
"title": caption or "Image",
|
||||
}),
|
||||
}, headers=headers)
|
||||
await self._api_post(
|
||||
MEDIA_SEND_URL,
|
||||
{
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleMarkdown",
|
||||
"msgParam": json.dumps(
|
||||
{
|
||||
"text": f""
|
||||
+ (f"\n{caption}" if caption else ""),
|
||||
"title": caption or "Image",
|
||||
}
|
||||
),
|
||||
},
|
||||
headers=headers,
|
||||
)
|
||||
else:
|
||||
# Non-image: send as markdown with filename
|
||||
name = Path(file_path).name
|
||||
text = f"[文件] {name}" + (f"\n{caption}" if caption else "")
|
||||
await self._api_post(MEDIA_SEND_URL, {
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleMarkdown",
|
||||
"msgParam": json.dumps({"text": text, "title": name}),
|
||||
}, headers=headers)
|
||||
await self._api_post(
|
||||
MEDIA_SEND_URL,
|
||||
{
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleMarkdown",
|
||||
"msgParam": json.dumps({"text": text, "title": name}),
|
||||
},
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
if caption and ext in self._IMAGE_EXTS:
|
||||
# Send caption separately for image messages
|
||||
await self._api_post(MEDIA_SEND_URL, {
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleMarkdown",
|
||||
"msgParam": json.dumps({"text": caption, "title": "Caption"}),
|
||||
}, headers=headers)
|
||||
await self._api_post(
|
||||
MEDIA_SEND_URL,
|
||||
{
|
||||
"robotCode": self.config.client_id,
|
||||
"userIds": [chat_id],
|
||||
"msgKey": "sampleMarkdown",
|
||||
"msgParam": json.dumps({"text": caption, "title": "Caption"}),
|
||||
},
|
||||
headers=headers,
|
||||
)
|
||||
return True
|
||||
|
||||
async def _upload_dingtalk_media(
|
||||
self, token: str, file_path: str, media_type: str = "image",
|
||||
self,
|
||||
token: str,
|
||||
file_path: str,
|
||||
media_type: str = "image",
|
||||
) -> str | None:
|
||||
"""Upload a file to DingTalk media API and return the media_id."""
|
||||
try:
|
||||
url = f"{MEDIA_UPLOAD_URL}?access_token={token}&type={media_type}"
|
||||
with open(file_path, "rb") as f:
|
||||
resp = await self._http_client.post(
|
||||
url, files={"media": (Path(file_path).name, f)},
|
||||
url,
|
||||
files={"media": (Path(file_path).name, f)},
|
||||
)
|
||||
data = resp.json()
|
||||
return data.get("media_id")
|
||||
|
||||
@@ -8,12 +8,14 @@ def create_from_config(config) -> DiscordChannel:
|
||||
allowed = _parse_csv(config.discord_allowed_senders)
|
||||
channels = _parse_csv(config.discord_allowed_channels)
|
||||
proxy = config.discord_proxy if config.discord_proxy else None
|
||||
return DiscordChannel(DiscordConfig(
|
||||
bot_token=config.discord_bot_token,
|
||||
allowed_senders=allowed,
|
||||
allowed_channels=channels,
|
||||
proxy=proxy,
|
||||
))
|
||||
return DiscordChannel(
|
||||
DiscordConfig(
|
||||
bot_token=config.discord_bot_token,
|
||||
allowed_senders=allowed,
|
||||
allowed_channels=channels,
|
||||
proxy=proxy,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
register_channel("discord", create_from_config)
|
||||
|
||||
@@ -127,7 +127,9 @@ class DiscordChannel(Channel):
|
||||
|
||||
# ── ACK Reactions ───────────────────────────────────────────────
|
||||
|
||||
async def _send_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
async def _send_ack_reaction(
|
||||
self, chat_id: str, message_id: str, emoji: str = "👀"
|
||||
) -> None:
|
||||
msg = self._message_cache.get(message_id)
|
||||
if msg:
|
||||
try:
|
||||
@@ -135,7 +137,9 @@ class DiscordChannel(Channel):
|
||||
except Exception as e:
|
||||
logger.debug(f"Discord ACK reaction failed: {e}")
|
||||
|
||||
async def _remove_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
async def _remove_ack_reaction(
|
||||
self, chat_id: str, message_id: str, emoji: str = "👀"
|
||||
) -> None:
|
||||
msg = self._message_cache.get(message_id)
|
||||
if msg and self._client and self._client.user:
|
||||
try:
|
||||
@@ -155,7 +159,6 @@ class DiscordChannel(Channel):
|
||||
|
||||
# ── Send ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||
import discord
|
||||
|
||||
@@ -168,7 +171,8 @@ class DiscordChannel(Channel):
|
||||
if reply_to:
|
||||
try:
|
||||
ref = discord.MessageReference(
|
||||
message_id=int(reply_to), channel_id=target_id,
|
||||
message_id=int(reply_to),
|
||||
channel_id=target_id,
|
||||
)
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
@@ -179,8 +183,11 @@ class DiscordChannel(Channel):
|
||||
await self._send_with_format_fallback(_send, formatted_text, raw_text)
|
||||
|
||||
async def _send_media_impl(
|
||||
self, recipient: str, file_path: str,
|
||||
caption: str = "", metadata: dict | None = None,
|
||||
self,
|
||||
recipient: str,
|
||||
file_path: str,
|
||||
caption: str = "",
|
||||
metadata: dict | None = None,
|
||||
) -> bool:
|
||||
import discord
|
||||
|
||||
@@ -222,7 +229,8 @@ class DiscordChannel(Channel):
|
||||
if self.config.include_attachments and message.attachments:
|
||||
for attachment in message.attachments:
|
||||
too_large = self._check_attachment_size(
|
||||
attachment.size or 0, attachment.filename,
|
||||
attachment.size or 0,
|
||||
attachment.filename,
|
||||
)
|
||||
if too_large:
|
||||
annotations.append(too_large)
|
||||
@@ -235,7 +243,9 @@ class DiscordChannel(Channel):
|
||||
annotations.append(f"[attachment: {file_path}]")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to download Discord attachment: {e}")
|
||||
annotations.append(f"[attachment: {attachment.filename} - download failed]")
|
||||
annotations.append(
|
||||
f"[attachment: {attachment.filename} - download failed]"
|
||||
)
|
||||
|
||||
# Detect thread context
|
||||
thread_id = ""
|
||||
@@ -245,11 +255,17 @@ class DiscordChannel(Channel):
|
||||
thread_id = channel_id # the thread IS the channel
|
||||
parent_channel_id = str(message.channel.parent.id)
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=user_id, chat_id=parent_channel_id, text=text,
|
||||
media_files=media_paths, content_annotations=annotations,
|
||||
timestamp=message.created_at or datetime.now(),
|
||||
message_id=str(message.id),
|
||||
metadata={"chat_id": parent_channel_id, "thread_id": thread_id},
|
||||
is_group=not is_dm, was_mentioned=was_mentioned,
|
||||
))
|
||||
await self._enqueue_raw(
|
||||
RawIncoming(
|
||||
sender_id=user_id,
|
||||
chat_id=parent_channel_id,
|
||||
text=text,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=message.created_at or datetime.now(),
|
||||
message_id=str(message.id),
|
||||
metadata={"chat_id": parent_channel_id, "thread_id": thread_id},
|
||||
is_group=not is_dm,
|
||||
was_mentioned=was_mentioned,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -5,7 +5,9 @@ import logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def validate_discord_token(token: str, proxy: str | None = None) -> tuple[bool, str]:
|
||||
async def validate_discord_token(
|
||||
token: str, proxy: str | None = None
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate a Discord bot token via the REST API.
|
||||
|
||||
Returns:
|
||||
|
||||
@@ -17,25 +17,27 @@ __all__ = ["EmailChannel", "EmailConfig"]
|
||||
|
||||
def create_from_config(config) -> EmailChannel:
|
||||
allowed = _parse_csv(config.email_allowed_senders)
|
||||
return EmailChannel(EmailConfig(
|
||||
imap_host=config.email_imap_host,
|
||||
imap_port=config.email_imap_port,
|
||||
imap_username=config.email_imap_username,
|
||||
imap_password=config.email_imap_password,
|
||||
imap_mailbox=config.email_imap_mailbox,
|
||||
imap_use_ssl=config.email_imap_use_ssl,
|
||||
smtp_host=config.email_smtp_host,
|
||||
smtp_port=config.email_smtp_port,
|
||||
smtp_username=config.email_smtp_username,
|
||||
smtp_password=config.email_smtp_password,
|
||||
smtp_starttls=config.email_smtp_use_tls,
|
||||
from_address=config.email_from_address,
|
||||
poll_interval=config.email_poll_interval,
|
||||
mark_seen=config.email_mark_seen,
|
||||
max_body_chars=config.email_max_body_chars,
|
||||
subject_prefix=config.email_subject_prefix,
|
||||
allowed_senders=allowed,
|
||||
))
|
||||
return EmailChannel(
|
||||
EmailConfig(
|
||||
imap_host=config.email_imap_host,
|
||||
imap_port=config.email_imap_port,
|
||||
imap_username=config.email_imap_username,
|
||||
imap_password=config.email_imap_password,
|
||||
imap_mailbox=config.email_imap_mailbox,
|
||||
imap_use_ssl=config.email_imap_use_ssl,
|
||||
smtp_host=config.email_smtp_host,
|
||||
smtp_port=config.email_smtp_port,
|
||||
smtp_username=config.email_smtp_username,
|
||||
smtp_password=config.email_smtp_password,
|
||||
smtp_starttls=config.email_smtp_use_tls,
|
||||
from_address=config.email_from_address,
|
||||
poll_interval=config.email_poll_interval,
|
||||
mark_seen=config.email_mark_seen,
|
||||
max_body_chars=config.email_max_body_chars,
|
||||
subject_prefix=config.email_subject_prefix,
|
||||
allowed_senders=allowed,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
register_channel("email", create_from_config)
|
||||
|
||||
@@ -56,7 +56,9 @@ class EmailConfig(BaseChannelConfig):
|
||||
smtp_port: int = 587
|
||||
smtp_username: str = ""
|
||||
smtp_password: str = ""
|
||||
smtp_starttls: bool = True # True=STARTTLS (port 587), False=implicit SSL (port 465)
|
||||
smtp_starttls: bool = (
|
||||
True # True=STARTTLS (port 587), False=implicit SSL (port 465)
|
||||
)
|
||||
from_address: str = ""
|
||||
poll_interval: int = 30
|
||||
mark_seen: bool = True
|
||||
@@ -84,7 +86,9 @@ class EmailChannel(Channel, PollingMixin):
|
||||
loop = asyncio.get_running_loop()
|
||||
await loop.run_in_executor(None, self._connect_imap)
|
||||
self._running = True
|
||||
logger.info(f"Email channel started (IMAP: {cfg.imap_host}, poll {cfg.poll_interval}s)")
|
||||
logger.info(
|
||||
f"Email channel started (IMAP: {cfg.imap_host}, poll {cfg.poll_interval}s)"
|
||||
)
|
||||
await self._start_polling()
|
||||
|
||||
async def _cleanup(self) -> None:
|
||||
@@ -102,7 +106,11 @@ class EmailChannel(Channel, PollingMixin):
|
||||
cfg = self.config
|
||||
try:
|
||||
if cfg.imap_use_ssl:
|
||||
self._imap = imaplib.IMAP4_SSL(cfg.imap_host, cfg.imap_port, ssl_context=ssl.create_default_context())
|
||||
self._imap = imaplib.IMAP4_SSL(
|
||||
cfg.imap_host,
|
||||
cfg.imap_port,
|
||||
ssl_context=ssl.create_default_context(),
|
||||
)
|
||||
else:
|
||||
self._imap = imaplib.IMAP4(cfg.imap_host, cfg.imap_port)
|
||||
self._imap.login(cfg.imap_username, cfg.imap_password)
|
||||
@@ -140,7 +148,7 @@ class EmailChannel(Channel, PollingMixin):
|
||||
from_name, from_addr = parseaddr(msg.get("From", ""))
|
||||
body = self._extract_body(msg)
|
||||
if len(body) > self.config.max_body_chars:
|
||||
body = body[:self.config.max_body_chars] + "\n[...truncated]"
|
||||
body = body[: self.config.max_body_chars] + "\n[...truncated]"
|
||||
# Extract attachments and inline images
|
||||
attachments = []
|
||||
if msg.is_multipart():
|
||||
@@ -168,22 +176,44 @@ class EmailChannel(Channel, PollingMixin):
|
||||
payload_data = part.get_payload(decode=True)
|
||||
if payload_data:
|
||||
from ..base import MAX_ATTACHMENT_BYTES, MEDIA_DIR
|
||||
|
||||
if len(payload_data) > MAX_ATTACHMENT_BYTES:
|
||||
attachments.append({"annotation": f"[attachment: {filename} - too large ({len(payload_data)} bytes)]"})
|
||||
attachments.append(
|
||||
{
|
||||
"annotation": f"[attachment: {filename} - too large ({len(payload_data)} bytes)]"
|
||||
}
|
||||
)
|
||||
else:
|
||||
MEDIA_DIR.mkdir(parents=True, exist_ok=True)
|
||||
local_path = MEDIA_DIR / f"email_{mid.decode()}_{filename}"
|
||||
local_path = (
|
||||
MEDIA_DIR / f"email_{mid.decode()}_{filename}"
|
||||
)
|
||||
local_path.write_bytes(payload_data)
|
||||
label = "inline-image" if is_inline_image else "attachment"
|
||||
attachments.append({"path": str(local_path), "annotation": f"[{label}: {local_path}]"})
|
||||
label = (
|
||||
"inline-image"
|
||||
if is_inline_image
|
||||
else "attachment"
|
||||
)
|
||||
attachments.append(
|
||||
{
|
||||
"path": str(local_path),
|
||||
"annotation": f"[{label}: {local_path}]",
|
||||
}
|
||||
)
|
||||
if self.config.mark_seen:
|
||||
self._imap.store(mid, "+FLAGS", "\\Seen")
|
||||
results.append({
|
||||
"from_addr": from_addr, "from_name": _decode_hdr(from_name),
|
||||
"subject": _decode_hdr(msg.get("Subject", "")), "body": body,
|
||||
"message_id": msg.get("Message-ID", ""), "date": msg.get("Date", ""),
|
||||
"references": msg.get("References", ""), "attachments": attachments,
|
||||
})
|
||||
results.append(
|
||||
{
|
||||
"from_addr": from_addr,
|
||||
"from_name": _decode_hdr(from_name),
|
||||
"subject": _decode_hdr(msg.get("Subject", "")),
|
||||
"body": body,
|
||||
"message_id": msg.get("Message-ID", ""),
|
||||
"date": msg.get("Date", ""),
|
||||
"references": msg.get("References", ""),
|
||||
"attachments": attachments,
|
||||
}
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"IMAP fetch: {e}")
|
||||
return results
|
||||
@@ -224,14 +254,24 @@ class EmailChannel(Channel, PollingMixin):
|
||||
media_paths.append(att["path"])
|
||||
if att.get("annotation"):
|
||||
annotations.append(att["annotation"])
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=m["from_addr"], chat_id=m["from_addr"], text=text, timestamp=ts,
|
||||
message_id=m["message_id"],
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
metadata={"chat_id": m["from_addr"], "subject": subject,
|
||||
"original_message_id": m["message_id"], "references": m["references"], "backend": "email"},
|
||||
))
|
||||
await self._enqueue_raw(
|
||||
RawIncoming(
|
||||
sender_id=m["from_addr"],
|
||||
chat_id=m["from_addr"],
|
||||
text=text,
|
||||
timestamp=ts,
|
||||
message_id=m["message_id"],
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
metadata={
|
||||
"chat_id": m["from_addr"],
|
||||
"subject": subject,
|
||||
"original_message_id": m["message_id"],
|
||||
"references": m["references"],
|
||||
"backend": "email",
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
# ── Send ──────────────────────────────────────────────────────
|
||||
|
||||
@@ -254,8 +294,10 @@ class EmailChannel(Channel, PollingMixin):
|
||||
srv.starttls()
|
||||
else:
|
||||
srv = smtplib.SMTP_SSL(
|
||||
cfg.smtp_host, cfg.smtp_port,
|
||||
context=ssl.create_default_context(), timeout=30,
|
||||
cfg.smtp_host,
|
||||
cfg.smtp_port,
|
||||
context=ssl.create_default_context(),
|
||||
timeout=30,
|
||||
)
|
||||
srv.login(cfg.smtp_username, cfg.smtp_password)
|
||||
yield srv
|
||||
@@ -273,16 +315,27 @@ class EmailChannel(Channel, PollingMixin):
|
||||
loop = asyncio.get_running_loop()
|
||||
try:
|
||||
await loop.run_in_executor(
|
||||
None, self._smtp_send_html, chat_id, formatted_text, raw_text, metadata or {},
|
||||
None,
|
||||
self._smtp_send_html,
|
||||
chat_id,
|
||||
formatted_text,
|
||||
raw_text,
|
||||
metadata or {},
|
||||
)
|
||||
except Exception as e:
|
||||
err_str = str(e).lower()
|
||||
# Only fall back to plain text for format-related errors, not server rejections
|
||||
if any(code in err_str for code in ("550", "553", "554", "auth", "rejected")):
|
||||
if any(
|
||||
code in err_str for code in ("550", "553", "554", "auth", "rejected")
|
||||
):
|
||||
raise
|
||||
logger.warning(f"HTML email failed ({e}), falling back to plain text")
|
||||
await loop.run_in_executor(
|
||||
None, self._smtp_send, chat_id, raw_text, metadata or {},
|
||||
None,
|
||||
self._smtp_send,
|
||||
chat_id,
|
||||
raw_text,
|
||||
metadata or {},
|
||||
)
|
||||
|
||||
def _smtp_send(self, to: str, content: str, meta: dict) -> None:
|
||||
@@ -291,7 +344,11 @@ class EmailChannel(Channel, PollingMixin):
|
||||
logger.debug(f"SMTP plain send: from={from_addr} to={to}")
|
||||
msg = EmailMessage()
|
||||
orig_subj = meta.get("subject", "")
|
||||
msg["Subject"] = f"{cfg.subject_prefix}{orig_subj}" if orig_subj and not orig_subj.lower().startswith("re:") else (orig_subj or "EvoScientist Reply")
|
||||
msg["Subject"] = (
|
||||
f"{cfg.subject_prefix}{orig_subj}"
|
||||
if orig_subj and not orig_subj.lower().startswith("re:")
|
||||
else (orig_subj or "EvoScientist Reply")
|
||||
)
|
||||
msg["From"] = from_addr
|
||||
msg["To"] = to
|
||||
orig_id = meta.get("original_message_id", "")
|
||||
@@ -306,14 +363,20 @@ class EmailChannel(Channel, PollingMixin):
|
||||
logger.error(f"SMTP send failed: from={from_addr} to={to}")
|
||||
raise RuntimeError("SMTP send failed") from e
|
||||
|
||||
def _smtp_send_html(self, to: str, html_content: str, plain_content: str, meta: dict) -> None:
|
||||
def _smtp_send_html(
|
||||
self, to: str, html_content: str, plain_content: str, meta: dict
|
||||
) -> None:
|
||||
"""Send an email with both HTML and plain-text parts."""
|
||||
cfg = self.config
|
||||
from_addr = cfg.from_address or cfg.smtp_username
|
||||
logger.debug(f"SMTP HTML send: from={from_addr} to={to}")
|
||||
msg = MIMEMultipart("alternative")
|
||||
orig_subj = meta.get("subject", "")
|
||||
msg["Subject"] = f"{cfg.subject_prefix}{orig_subj}" if orig_subj and not orig_subj.lower().startswith("re:") else (orig_subj or "EvoScientist Reply")
|
||||
msg["Subject"] = (
|
||||
f"{cfg.subject_prefix}{orig_subj}"
|
||||
if orig_subj and not orig_subj.lower().startswith("re:")
|
||||
else (orig_subj or "EvoScientist Reply")
|
||||
)
|
||||
msg["From"] = from_addr
|
||||
msg["To"] = to
|
||||
orig_id = meta.get("original_message_id", "")
|
||||
@@ -341,18 +404,29 @@ class EmailChannel(Channel, PollingMixin):
|
||||
"""Send a file as an email attachment via SMTP."""
|
||||
loop = asyncio.get_running_loop()
|
||||
await loop.run_in_executor(
|
||||
None, self._smtp_send_attachment, recipient, file_path, caption, metadata or {},
|
||||
None,
|
||||
self._smtp_send_attachment,
|
||||
recipient,
|
||||
file_path,
|
||||
caption,
|
||||
metadata or {},
|
||||
)
|
||||
return True
|
||||
|
||||
def _smtp_send_attachment(self, to: str, file_path: str, caption: str, meta: dict) -> None:
|
||||
def _smtp_send_attachment(
|
||||
self, to: str, file_path: str, caption: str, meta: dict
|
||||
) -> None:
|
||||
"""Send an email with a file attachment."""
|
||||
cfg = self.config
|
||||
from_addr = cfg.from_address or cfg.smtp_username
|
||||
logger.debug(f"SMTP attachment send: from={from_addr} to={to} file={file_path}")
|
||||
msg = MIMEMultipart()
|
||||
orig_subj = meta.get("subject", "")
|
||||
msg["Subject"] = f"{cfg.subject_prefix}{orig_subj}" if orig_subj and not orig_subj.lower().startswith("re:") else (orig_subj or "EvoScientist Reply")
|
||||
msg["Subject"] = (
|
||||
f"{cfg.subject_prefix}{orig_subj}"
|
||||
if orig_subj and not orig_subj.lower().startswith("re:")
|
||||
else (orig_subj or "EvoScientist Reply")
|
||||
)
|
||||
msg["From"] = from_addr
|
||||
msg["To"] = to
|
||||
orig_id = meta.get("original_message_id", "")
|
||||
|
||||
@@ -9,7 +9,10 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def validate_email_imap(
|
||||
host: str, port: int, username: str, password: str,
|
||||
host: str,
|
||||
port: int,
|
||||
username: str,
|
||||
password: str,
|
||||
use_ssl: bool = True,
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate IMAP credentials.
|
||||
@@ -21,6 +24,7 @@ async def validate_email_imap(
|
||||
return False, "host, username, and password are required"
|
||||
|
||||
import asyncio
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
|
||||
def _check():
|
||||
@@ -42,7 +46,10 @@ async def validate_email_imap(
|
||||
|
||||
|
||||
async def validate_email_smtp(
|
||||
host: str, port: int, username: str, password: str,
|
||||
host: str,
|
||||
port: int,
|
||||
username: str,
|
||||
password: str,
|
||||
use_tls: bool = True,
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate SMTP credentials.
|
||||
@@ -54,6 +61,7 @@ async def validate_email_smtp(
|
||||
return False, "host, username, and password are required"
|
||||
|
||||
import asyncio
|
||||
|
||||
loop = asyncio.get_event_loop()
|
||||
|
||||
def _check():
|
||||
|
||||
@@ -7,16 +7,18 @@ __all__ = ["FeishuChannel", "FeishuConfig"]
|
||||
def create_from_config(config) -> FeishuChannel:
|
||||
allowed = _parse_csv(config.feishu_allowed_senders)
|
||||
proxy = config.feishu_proxy if config.feishu_proxy else None
|
||||
return FeishuChannel(FeishuConfig(
|
||||
app_id=config.feishu_app_id,
|
||||
app_secret=config.feishu_app_secret,
|
||||
verification_token=config.feishu_verification_token,
|
||||
encrypt_key=config.feishu_encrypt_key,
|
||||
webhook_port=config.feishu_webhook_port,
|
||||
allowed_senders=allowed,
|
||||
feishu_domain=config.feishu_domain,
|
||||
proxy=proxy,
|
||||
))
|
||||
return FeishuChannel(
|
||||
FeishuConfig(
|
||||
app_id=config.feishu_app_id,
|
||||
app_secret=config.feishu_app_secret,
|
||||
verification_token=config.feishu_verification_token,
|
||||
encrypt_key=config.feishu_encrypt_key,
|
||||
webhook_port=config.feishu_webhook_port,
|
||||
allowed_senders=allowed,
|
||||
feishu_domain=config.feishu_domain,
|
||||
proxy=proxy,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
register_channel("feishu", create_from_config)
|
||||
|
||||
@@ -50,43 +50,59 @@ def _parse_inline_text(text: str) -> list[dict]:
|
||||
elements: list[dict] = []
|
||||
# Pattern order matters: code first (protect content), then bold, strikethrough, link, italic
|
||||
pattern = re.compile(
|
||||
r"`([^`]+)`" # inline code
|
||||
r"|\*\*(.+?)\*\*" # bold
|
||||
r"|~~(.+?)~~" # strikethrough
|
||||
r"|\[([^\]]+)\]\(([^)]+)\)" # link
|
||||
r"|_(.+?)_" # italic
|
||||
r"`([^`]+)`" # inline code
|
||||
r"|\*\*(.+?)\*\*" # bold
|
||||
r"|~~(.+?)~~" # strikethrough
|
||||
r"|\[([^\]]+)\]\(([^)]+)\)" # link
|
||||
r"|_(.+?)_" # italic
|
||||
)
|
||||
pos = 0
|
||||
for m in pattern.finditer(text):
|
||||
# Plain text before this match
|
||||
if m.start() > pos:
|
||||
elements.append({"tag": "text", "text": text[pos:m.start()]})
|
||||
elements.append({"tag": "text", "text": text[pos : m.start()]})
|
||||
|
||||
if m.group(1) is not None:
|
||||
# inline code → code_block would be block-level; use text with style
|
||||
elements.append({
|
||||
"tag": "text", "text": m.group(1),
|
||||
"style": ["code_block"],
|
||||
})
|
||||
elements.append(
|
||||
{
|
||||
"tag": "text",
|
||||
"text": m.group(1),
|
||||
"style": ["code_block"],
|
||||
}
|
||||
)
|
||||
elif m.group(2) is not None:
|
||||
elements.append({
|
||||
"tag": "text", "text": m.group(2),
|
||||
"style": ["bold"],
|
||||
})
|
||||
elements.append(
|
||||
{
|
||||
"tag": "text",
|
||||
"text": m.group(2),
|
||||
"style": ["bold"],
|
||||
}
|
||||
)
|
||||
elif m.group(3) is not None:
|
||||
elements.append({
|
||||
"tag": "text", "text": m.group(3),
|
||||
"style": ["strikethrough"],
|
||||
})
|
||||
elements.append(
|
||||
{
|
||||
"tag": "text",
|
||||
"text": m.group(3),
|
||||
"style": ["strikethrough"],
|
||||
}
|
||||
)
|
||||
elif m.group(4) is not None:
|
||||
elements.append({
|
||||
"tag": "a", "text": m.group(4), "href": m.group(5),
|
||||
})
|
||||
elements.append(
|
||||
{
|
||||
"tag": "a",
|
||||
"text": m.group(4),
|
||||
"href": m.group(5),
|
||||
}
|
||||
)
|
||||
elif m.group(6) is not None:
|
||||
elements.append({
|
||||
"tag": "text", "text": m.group(6),
|
||||
"style": ["italic"],
|
||||
})
|
||||
elements.append(
|
||||
{
|
||||
"tag": "text",
|
||||
"text": m.group(6),
|
||||
"style": ["italic"],
|
||||
}
|
||||
)
|
||||
pos = m.end()
|
||||
|
||||
# Remaining plain text
|
||||
@@ -161,11 +177,15 @@ def _markdown_to_feishu_post(text: str) -> dict | None:
|
||||
else:
|
||||
# End of code block
|
||||
code_text = "\n".join(code_lines)
|
||||
paragraphs.append([{
|
||||
"tag": "code_block",
|
||||
"language": code_lang or "plain",
|
||||
"text": code_text,
|
||||
}])
|
||||
paragraphs.append(
|
||||
[
|
||||
{
|
||||
"tag": "code_block",
|
||||
"language": code_lang or "plain",
|
||||
"text": code_text,
|
||||
}
|
||||
]
|
||||
)
|
||||
in_code_block = False
|
||||
code_lines = []
|
||||
code_lang = ""
|
||||
@@ -193,11 +213,15 @@ def _markdown_to_feishu_post(text: str) -> dict | None:
|
||||
# Flush remaining
|
||||
if in_code_block and code_lines:
|
||||
code_text = "\n".join(code_lines)
|
||||
paragraphs.append([{
|
||||
"tag": "code_block",
|
||||
"language": code_lang or "plain",
|
||||
"text": code_text,
|
||||
}])
|
||||
paragraphs.append(
|
||||
[
|
||||
{
|
||||
"tag": "code_block",
|
||||
"language": code_lang or "plain",
|
||||
"text": code_text,
|
||||
}
|
||||
]
|
||||
)
|
||||
elif current_paragraph:
|
||||
paragraphs.append(current_paragraph)
|
||||
|
||||
@@ -225,12 +249,12 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
name = "feishu"
|
||||
_ready_attrs = ("_http_client", "_access_token")
|
||||
_non_retryable_patterns = (
|
||||
"app_access_token is empty", # invalid credentials
|
||||
"10003", # invalid app_id
|
||||
"10014", # invalid app_secret
|
||||
"99991401", # permission denied
|
||||
"99991663", # no permission
|
||||
"99991672", # feature not enabled
|
||||
"app_access_token is empty", # invalid credentials
|
||||
"10003", # invalid app_id
|
||||
"10014", # invalid app_secret
|
||||
"99991401", # permission denied
|
||||
"99991663", # no permission
|
||||
"99991672", # feature not enabled
|
||||
)
|
||||
_rate_limit_patterns = ("99991400", "rate limit", "频率限制")
|
||||
_rate_limit_delay = 2.0
|
||||
@@ -263,9 +287,7 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
raise ChannelError(f"Failed to get Feishu access token: {e}")
|
||||
|
||||
if data.get("code") != 0:
|
||||
raise ChannelError(
|
||||
f"Feishu auth error: {data.get('msg', 'unknown')}"
|
||||
)
|
||||
raise ChannelError(f"Feishu auth error: {data.get('msg', 'unknown')}")
|
||||
return data["tenant_access_token"], data.get("expire", 7200)
|
||||
|
||||
# ── Lifecycle ─────────────────────────────────────────────────
|
||||
@@ -293,8 +315,7 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
|
||||
self._running = True
|
||||
logger.info(
|
||||
f"Feishu channel started "
|
||||
f"(webhook on port {self.config.webhook_port})"
|
||||
f"Feishu channel started (webhook on port {self.config.webhook_port})"
|
||||
)
|
||||
|
||||
async def _cleanup(self) -> None:
|
||||
@@ -320,7 +341,12 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
return False
|
||||
|
||||
async def _send_chunk(
|
||||
self, chat_id, formatted_text, raw_text, reply_to, metadata,
|
||||
self,
|
||||
chat_id,
|
||||
formatted_text,
|
||||
raw_text,
|
||||
reply_to,
|
||||
metadata,
|
||||
):
|
||||
token = await self._ensure_token()
|
||||
headers = {"Authorization": f"Bearer {token}"}
|
||||
@@ -329,13 +355,15 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
# If reply_to is set, try the reply API first
|
||||
if reply_to:
|
||||
reply_url = (
|
||||
f"{self.config.feishu_domain}"
|
||||
f"/open-apis/im/v1/messages/{reply_to}/reply"
|
||||
f"{self.config.feishu_domain}/open-apis/im/v1/messages/{reply_to}/reply"
|
||||
)
|
||||
if post_content is not None:
|
||||
body = {"msg_type": "post", "content": json.dumps(post_content)}
|
||||
else:
|
||||
body = {"msg_type": "text", "content": json.dumps({"text": formatted_text})}
|
||||
body = {
|
||||
"msg_type": "text",
|
||||
"content": json.dumps({"text": formatted_text}),
|
||||
}
|
||||
if await self._feishu_send(reply_url, body, headers):
|
||||
return
|
||||
|
||||
@@ -369,7 +397,10 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
_IMAGE_EXTENSIONS = {".jpg", ".jpeg", ".png", ".gif", ".bmp", ".webp"}
|
||||
|
||||
async def _download_media(
|
||||
self, message_id: str, file_key: str, msg_type: str,
|
||||
self,
|
||||
message_id: str,
|
||||
file_key: str,
|
||||
msg_type: str,
|
||||
) -> str | None:
|
||||
"""Download an image or file attachment from Feishu.
|
||||
|
||||
@@ -386,9 +417,7 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
try:
|
||||
resp = await self._http_client.get(url, headers=headers, timeout=30)
|
||||
if resp.status_code != 200:
|
||||
logger.warning(
|
||||
f"Feishu media download failed: HTTP {resp.status_code}"
|
||||
)
|
||||
logger.warning(f"Feishu media download failed: HTTP {resp.status_code}")
|
||||
return None
|
||||
|
||||
# Check attachment size before writing to disk
|
||||
@@ -402,10 +431,9 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
from ..base import MAX_ATTACHMENT_BYTES
|
||||
|
||||
if len(resp.content) > MAX_ATTACHMENT_BYTES:
|
||||
logger.warning(
|
||||
f"Feishu media too large: {len(resp.content)} bytes"
|
||||
)
|
||||
logger.warning(f"Feishu media too large: {len(resp.content)} bytes")
|
||||
return None
|
||||
|
||||
# Determine extension from Content-Type or default
|
||||
@@ -426,13 +454,19 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
return None
|
||||
|
||||
async def _upload_feishu_resource(
|
||||
self, url: str, headers: dict, file_path: str,
|
||||
field_name: str, extra_data: dict,
|
||||
self,
|
||||
url: str,
|
||||
headers: dict,
|
||||
file_path: str,
|
||||
field_name: str,
|
||||
extra_data: dict,
|
||||
) -> dict | None:
|
||||
"""Upload a file to Feishu API. Returns response data or None on failure."""
|
||||
with open(file_path, "rb") as f:
|
||||
resp = await self._http_client.post(
|
||||
url, headers=headers, data=extra_data,
|
||||
url,
|
||||
headers=headers,
|
||||
data=extra_data,
|
||||
files={field_name: (Path(file_path).name, f)},
|
||||
)
|
||||
data = resp.json()
|
||||
@@ -465,7 +499,11 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
if is_image:
|
||||
upload_url = f"{self.config.feishu_domain}/open-apis/im/v1/images"
|
||||
data = await self._upload_feishu_resource(
|
||||
upload_url, headers, file_path, "image", {"image_type": "message"},
|
||||
upload_url,
|
||||
headers,
|
||||
file_path,
|
||||
"image",
|
||||
{"image_type": "message"},
|
||||
)
|
||||
if not data:
|
||||
return False
|
||||
@@ -477,7 +515,10 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
else:
|
||||
upload_url = f"{self.config.feishu_domain}/open-apis/im/v1/files"
|
||||
data = await self._upload_feishu_resource(
|
||||
upload_url, headers, file_path, "file",
|
||||
upload_url,
|
||||
headers,
|
||||
file_path,
|
||||
"file",
|
||||
{"file_type": "stream", "file_name": path.name},
|
||||
)
|
||||
if not data:
|
||||
@@ -504,7 +545,9 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
|
||||
# ── ACK reaction ───────────────────────────────────────────────
|
||||
|
||||
async def _send_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "THUMBSUP") -> None:
|
||||
async def _send_ack_reaction(
|
||||
self, chat_id: str, message_id: str, emoji: str = "THUMBSUP"
|
||||
) -> None:
|
||||
"""Send an acknowledgment reaction via Feishu Open API."""
|
||||
try:
|
||||
token = await self._ensure_token()
|
||||
@@ -517,7 +560,9 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
except Exception as e:
|
||||
logger.debug(f"Feishu ack reaction failed: {e}")
|
||||
|
||||
async def _remove_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "THUMBSUP") -> None:
|
||||
async def _remove_ack_reaction(
|
||||
self, chat_id: str, message_id: str, emoji: str = "THUMBSUP"
|
||||
) -> None:
|
||||
"""Remove ACK reaction via Feishu Open API.
|
||||
|
||||
Feishu's DELETE /reactions endpoint requires the reaction_id, which
|
||||
@@ -633,11 +678,7 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
"""Handle im.message.receive_v1 event (v2 schema)."""
|
||||
sender_info = event.get("sender", {})
|
||||
sender_id_info = sender_info.get("sender_id", {})
|
||||
sender_id = (
|
||||
sender_id_info.get("open_id")
|
||||
or sender_id_info.get("user_id")
|
||||
or ""
|
||||
)
|
||||
sender_id = sender_id_info.get("open_id") or sender_id_info.get("user_id") or ""
|
||||
sender_type = sender_info.get("sender_type", "")
|
||||
|
||||
# Skip bot's own messages
|
||||
@@ -739,27 +780,31 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
# Parse timestamp (milliseconds)
|
||||
create_time = message.get("create_time", "")
|
||||
try:
|
||||
timestamp = datetime.fromtimestamp(
|
||||
int(create_time) / 1000
|
||||
) if create_time else datetime.now()
|
||||
timestamp = (
|
||||
datetime.fromtimestamp(int(create_time) / 1000)
|
||||
if create_time
|
||||
else datetime.now()
|
||||
)
|
||||
except (ValueError, TypeError, OSError):
|
||||
timestamp = datetime.now()
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=sender_id,
|
||||
chat_id=chat_id,
|
||||
text=text,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=timestamp,
|
||||
message_id=message_id,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"chat_type": message.get("chat_type", ""),
|
||||
},
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
))
|
||||
await self._enqueue_raw(
|
||||
RawIncoming(
|
||||
sender_id=sender_id,
|
||||
chat_id=chat_id,
|
||||
text=text,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=timestamp,
|
||||
message_id=message_id,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"chat_type": message.get("chat_type", ""),
|
||||
},
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
)
|
||||
)
|
||||
|
||||
async def _on_message_v1(self, event: dict) -> None:
|
||||
"""Handle v1 schema message event (legacy)."""
|
||||
@@ -782,19 +827,21 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin):
|
||||
chat_id = event.get("open_chat_id", "")
|
||||
message_id = event.get("open_message_id", "")
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=sender_id,
|
||||
chat_id=chat_id,
|
||||
text=text,
|
||||
timestamp=datetime.now(),
|
||||
message_id=message_id,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"chat_type": event.get("chat_type", ""),
|
||||
},
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
))
|
||||
await self._enqueue_raw(
|
||||
RawIncoming(
|
||||
sender_id=sender_id,
|
||||
chat_id=chat_id,
|
||||
text=text,
|
||||
timestamp=datetime.now(),
|
||||
message_id=message_id,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"chat_type": event.get("chat_type", ""),
|
||||
},
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _extract_post_text(content: dict) -> str:
|
||||
|
||||
@@ -98,10 +98,12 @@ def convert_markdown(
|
||||
|
||||
return text
|
||||
|
||||
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
# Shared helpers
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def _escape_html(text: str) -> str:
|
||||
return text.replace("&", "&").replace("<", "<").replace(">", ">")
|
||||
|
||||
@@ -114,6 +116,7 @@ def _noop_escape(text: str) -> str:
|
||||
# HTML profile (Telegram, Email, Teams)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def _html_code_block(lang: str, code: str) -> str:
|
||||
escaped = _escape_html(code)
|
||||
if lang:
|
||||
@@ -147,6 +150,7 @@ _HTML_INLINE_RULES: list[InlineRule] = [
|
||||
# Slack mrkdwn profile
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def _slack_code_block(lang: str, code: str) -> str:
|
||||
return f"```\n{code}```"
|
||||
|
||||
@@ -168,6 +172,7 @@ _SLACK_INLINE_RULES: list[InlineRule] = [
|
||||
# Discord profile (mostly passthrough, headings → bold)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def _discord_code_block(lang: str, code: str) -> str:
|
||||
return f"```{lang}\n{code}```"
|
||||
|
||||
@@ -185,6 +190,7 @@ _DISCORD_INLINE_RULES: list[InlineRule] = [
|
||||
# Plain text profile (strip all formatting)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def _plain_code_block(lang: str, code: str) -> str:
|
||||
return code
|
||||
|
||||
@@ -207,6 +213,7 @@ _PLAIN_INLINE_RULES: list[InlineRule] = [
|
||||
# Markdown passthrough profile (Feishu, DingTalk, WeCom)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
def _md_code_block(lang: str, code: str) -> str:
|
||||
return f"```{lang}\n{code}```"
|
||||
|
||||
@@ -222,6 +229,7 @@ _MD_INLINE_RULES: list[InlineRule] = [] # passthrough — already Markdown
|
||||
# Unified Formatter
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class UnifiedFormatter:
|
||||
"""Converts internal Markdown to a target platform format.
|
||||
|
||||
|
||||
@@ -32,7 +32,7 @@ class _IMessageAllowListMiddleware:
|
||||
does not cover.
|
||||
"""
|
||||
|
||||
def __init__(self, channel: 'IMessageChannelRpc'):
|
||||
def __init__(self, channel: "IMessageChannelRpc"):
|
||||
self._channel = channel
|
||||
|
||||
async def process_inbound(self, raw, context):
|
||||
@@ -80,6 +80,7 @@ class IMessageChannelRpc(Channel):
|
||||
iMessage doesn't need MentionGating (always sets was_mentioned=True).
|
||||
"""
|
||||
from ..middleware import DedupMiddleware, GroupHistoryMiddleware
|
||||
|
||||
middlewares = []
|
||||
middlewares.append(DedupMiddleware())
|
||||
middlewares.append(_IMessageAllowListMiddleware(self))
|
||||
@@ -150,6 +151,7 @@ class IMessageChannelRpc(Channel):
|
||||
fname = att_path.name
|
||||
# Check file size before copying
|
||||
from ..base import MAX_ATTACHMENT_BYTES
|
||||
|
||||
if att_path.stat().st_size > MAX_ATTACHMENT_BYTES:
|
||||
annotations.append(
|
||||
f"[{media_label}: {fname} - too large "
|
||||
@@ -159,12 +161,15 @@ class IMessageChannelRpc(Channel):
|
||||
local = self._media_path(f"imsg_{fname}")
|
||||
try:
|
||||
import shutil
|
||||
|
||||
shutil.copy2(str(att_path), str(local))
|
||||
media_paths.append(str(local))
|
||||
annotations.append(f"[{media_label}: {local}]")
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to copy iMessage attachment: {e}")
|
||||
annotations.append(f"[{media_label}: {fname} - copy failed]")
|
||||
annotations.append(
|
||||
f"[{media_label}: {fname} - copy failed]"
|
||||
)
|
||||
else:
|
||||
annotations.append(f"[{media_label}: {file_path} - not found]")
|
||||
|
||||
@@ -173,18 +178,20 @@ class IMessageChannelRpc(Channel):
|
||||
|
||||
is_group = message.get("is_group", False)
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=sender,
|
||||
chat_id=str(metadata.get("chat_id", sender)),
|
||||
text=text,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=timestamp,
|
||||
message_id=str(message.get("id", "")),
|
||||
metadata=metadata,
|
||||
is_group=is_group,
|
||||
was_mentioned=True, # iMessage has no mention concept
|
||||
))
|
||||
await self._enqueue_raw(
|
||||
RawIncoming(
|
||||
sender_id=sender,
|
||||
chat_id=str(metadata.get("chat_id", sender)),
|
||||
text=text,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=timestamp,
|
||||
message_id=str(message.get("id", "")),
|
||||
metadata=metadata,
|
||||
is_group=is_group,
|
||||
was_mentioned=True, # iMessage has no mention concept
|
||||
)
|
||||
)
|
||||
|
||||
# ── Sender filtering ──────────────────────────────────────────
|
||||
|
||||
@@ -363,7 +370,6 @@ class IMessageChannelRpc(Channel):
|
||||
"""iMessage uses plain text; no formatting conversion needed."""
|
||||
return text
|
||||
|
||||
|
||||
def _extract_retry_after(self, exc: Exception) -> float | None:
|
||||
"""iMessage-specific retry logic.
|
||||
|
||||
|
||||
@@ -42,7 +42,8 @@ async def get_cli_version(cli_path: str) -> str | None:
|
||||
"""
|
||||
try:
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
cli_path, "--version",
|
||||
cli_path,
|
||||
"--version",
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
)
|
||||
@@ -63,7 +64,9 @@ async def check_rpc_support(cli_path: str) -> bool:
|
||||
"""
|
||||
try:
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
cli_path, "rpc", "--help",
|
||||
cli_path,
|
||||
"rpc",
|
||||
"--help",
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
)
|
||||
|
||||
@@ -16,6 +16,7 @@ logger = logging.getLogger(__name__)
|
||||
@dataclass
|
||||
class RpcError:
|
||||
"""RPC error response."""
|
||||
|
||||
code: int | None = None
|
||||
message: str | None = None
|
||||
data: Any = None
|
||||
@@ -24,6 +25,7 @@ class RpcError:
|
||||
@dataclass
|
||||
class RpcNotification:
|
||||
"""RPC notification (no id, server-initiated)."""
|
||||
|
||||
method: str
|
||||
params: Any = None
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ from typing import Union
|
||||
|
||||
class IMessageService(Enum):
|
||||
"""iMessage service type."""
|
||||
|
||||
IMESSAGE = "imessage"
|
||||
SMS = "sms"
|
||||
AUTO = "auto"
|
||||
@@ -20,6 +21,7 @@ class IMessageService(Enum):
|
||||
@dataclass
|
||||
class ChatIdTarget:
|
||||
"""Target by chat ID."""
|
||||
|
||||
kind: str = "chat_id"
|
||||
chat_id: int = 0
|
||||
|
||||
@@ -27,6 +29,7 @@ class ChatIdTarget:
|
||||
@dataclass
|
||||
class ChatGuidTarget:
|
||||
"""Target by chat GUID."""
|
||||
|
||||
kind: str = "chat_guid"
|
||||
chat_guid: str = ""
|
||||
|
||||
@@ -34,6 +37,7 @@ class ChatGuidTarget:
|
||||
@dataclass
|
||||
class ChatIdentifierTarget:
|
||||
"""Target by chat identifier."""
|
||||
|
||||
kind: str = "chat_identifier"
|
||||
chat_identifier: str = ""
|
||||
|
||||
@@ -41,6 +45,7 @@ class ChatIdentifierTarget:
|
||||
@dataclass
|
||||
class HandleTarget:
|
||||
"""Target by handle (phone/email)."""
|
||||
|
||||
kind: str = "handle"
|
||||
to: str = ""
|
||||
service: IMessageService = IMessageService.AUTO
|
||||
@@ -111,22 +116,22 @@ def normalize_handle(raw: str) -> str:
|
||||
# Strip service prefixes
|
||||
for prefix, _ in SERVICE_PREFIXES:
|
||||
if lowered.startswith(prefix):
|
||||
return normalize_handle(trimmed[len(prefix):])
|
||||
return normalize_handle(trimmed[len(prefix) :])
|
||||
|
||||
# Normalize chat_id/chat_guid/chat_identifier prefixes
|
||||
for prefix in CHAT_ID_PREFIXES:
|
||||
if lowered.startswith(prefix):
|
||||
value = trimmed[len(prefix):].strip()
|
||||
value = trimmed[len(prefix) :].strip()
|
||||
return f"chat_id:{value}"
|
||||
|
||||
for prefix in CHAT_GUID_PREFIXES:
|
||||
if lowered.startswith(prefix):
|
||||
value = trimmed[len(prefix):].strip()
|
||||
value = trimmed[len(prefix) :].strip()
|
||||
return f"chat_guid:{value}"
|
||||
|
||||
for prefix in CHAT_IDENTIFIER_PREFIXES:
|
||||
if lowered.startswith(prefix):
|
||||
value = trimmed[len(prefix):].strip()
|
||||
value = trimmed[len(prefix) :].strip()
|
||||
return f"chat_identifier:{value}"
|
||||
|
||||
# Email - lowercase
|
||||
@@ -172,7 +177,7 @@ def parse_target(raw: str) -> IMessageTarget:
|
||||
# Check service prefixes first
|
||||
for prefix, service in SERVICE_PREFIXES:
|
||||
if lower.startswith(prefix):
|
||||
remainder = trimmed[len(prefix):].strip()
|
||||
remainder = trimmed[len(prefix) :].strip()
|
||||
if not remainder:
|
||||
raise ValueError(f"{prefix} target is required")
|
||||
|
||||
@@ -181,7 +186,9 @@ def parse_target(raw: str) -> IMessageTarget:
|
||||
# Check if remainder is a chat target
|
||||
is_chat = any(
|
||||
remainder_lower.startswith(p)
|
||||
for p in CHAT_ID_PREFIXES + CHAT_GUID_PREFIXES + CHAT_IDENTIFIER_PREFIXES
|
||||
for p in CHAT_ID_PREFIXES
|
||||
+ CHAT_GUID_PREFIXES
|
||||
+ CHAT_IDENTIFIER_PREFIXES
|
||||
)
|
||||
if is_chat:
|
||||
return parse_target(remainder)
|
||||
@@ -191,7 +198,7 @@ def parse_target(raw: str) -> IMessageTarget:
|
||||
# Check chat_id prefixes
|
||||
for prefix in CHAT_ID_PREFIXES:
|
||||
if lower.startswith(prefix):
|
||||
value = trimmed[len(prefix):].strip()
|
||||
value = trimmed[len(prefix) :].strip()
|
||||
try:
|
||||
chat_id = int(value)
|
||||
return ChatIdTarget(chat_id=chat_id)
|
||||
@@ -201,7 +208,7 @@ def parse_target(raw: str) -> IMessageTarget:
|
||||
# Check chat_guid prefixes
|
||||
for prefix in CHAT_GUID_PREFIXES:
|
||||
if lower.startswith(prefix):
|
||||
value = trimmed[len(prefix):].strip()
|
||||
value = trimmed[len(prefix) :].strip()
|
||||
if not value:
|
||||
raise ValueError("chat_guid is required")
|
||||
return ChatGuidTarget(chat_guid=value)
|
||||
@@ -209,7 +216,7 @@ def parse_target(raw: str) -> IMessageTarget:
|
||||
# Check chat_identifier prefixes
|
||||
for prefix in CHAT_IDENTIFIER_PREFIXES:
|
||||
if lower.startswith(prefix):
|
||||
value = trimmed[len(prefix):].strip()
|
||||
value = trimmed[len(prefix) :].strip()
|
||||
if not value:
|
||||
raise ValueError("chat_identifier is required")
|
||||
return ChatIdentifierTarget(chat_identifier=value)
|
||||
|
||||
@@ -28,6 +28,7 @@ _logger = logging.getLogger(__name__)
|
||||
|
||||
# ── Task cancellation helper ─────────────────────────────────────────
|
||||
|
||||
|
||||
async def _cancel_task(task: asyncio.Task) -> None:
|
||||
"""Cancel an asyncio task and await its completion.
|
||||
|
||||
@@ -129,6 +130,7 @@ class DedupCache:
|
||||
|
||||
# ── Group history buffer ─────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class HistoryEntry:
|
||||
sender_id: str
|
||||
@@ -178,6 +180,7 @@ class GroupHistoryBuffer:
|
||||
|
||||
# ── Typing indicator manager ─────────────────────────────────────────
|
||||
|
||||
|
||||
class TypingManager:
|
||||
"""Manages background typing-indicator loops per chat_id.
|
||||
|
||||
@@ -229,6 +232,7 @@ class TypingManager:
|
||||
|
||||
# ── Pairing manager ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class PairingRequest:
|
||||
sender_id: str
|
||||
@@ -309,7 +313,9 @@ class PairingManager:
|
||||
|
||||
def _cleanup_expired(self):
|
||||
now = time.monotonic()
|
||||
expired = [c for c, r in self._pending.items() if now - r.created_at > self.CODE_EXPIRY]
|
||||
expired = [
|
||||
c for c, r in self._pending.items() if now - r.created_at > self.CODE_EXPIRY
|
||||
]
|
||||
for c in expired:
|
||||
del self._pending[c]
|
||||
|
||||
@@ -321,11 +327,14 @@ class PairingManager:
|
||||
|
||||
# ── Inbound middleware base ──────────────────────────────────────────
|
||||
|
||||
|
||||
class InboundMiddleware:
|
||||
"""Base class for inbound message processing middleware."""
|
||||
|
||||
async def process_inbound(
|
||||
self, raw: RawIncoming, context: dict[str, Any],
|
||||
self,
|
||||
raw: RawIncoming,
|
||||
context: dict[str, Any],
|
||||
) -> RawIncoming | None:
|
||||
"""Process an inbound raw message.
|
||||
|
||||
@@ -339,7 +348,9 @@ class OutboundMiddlewareBase:
|
||||
"""Base class for outbound message processing middleware."""
|
||||
|
||||
async def process_outbound(
|
||||
self, message: OutboundMessage, context: dict[str, Any],
|
||||
self,
|
||||
message: OutboundMessage,
|
||||
context: dict[str, Any],
|
||||
) -> OutboundMessage | None:
|
||||
"""Process an outbound message.
|
||||
|
||||
@@ -351,6 +362,7 @@ class OutboundMiddlewareBase:
|
||||
|
||||
# ── Dedup ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class DedupMiddleware(InboundMiddleware):
|
||||
"""Message deduplication using a bounded TTL cache."""
|
||||
|
||||
@@ -361,11 +373,15 @@ class DedupMiddleware(InboundMiddleware):
|
||||
ttl_seconds: float = 3600.0,
|
||||
) -> None:
|
||||
self._cache = DedupCache(
|
||||
max_size=max_size, trim_to=trim_to, ttl_seconds=ttl_seconds,
|
||||
max_size=max_size,
|
||||
trim_to=trim_to,
|
||||
ttl_seconds=ttl_seconds,
|
||||
)
|
||||
|
||||
async def process_inbound(
|
||||
self, raw: RawIncoming, context: dict[str, Any],
|
||||
self,
|
||||
raw: RawIncoming,
|
||||
context: dict[str, Any],
|
||||
) -> RawIncoming | None:
|
||||
if raw.message_id and self._cache.is_duplicate(raw.message_id):
|
||||
_logger.debug(f"Dedup: skipping duplicate message {raw.message_id}")
|
||||
@@ -375,6 +391,7 @@ class DedupMiddleware(InboundMiddleware):
|
||||
|
||||
# ── Debounce ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class DebounceMiddleware:
|
||||
"""Per-sender message batching with configurable timing.
|
||||
|
||||
@@ -471,6 +488,7 @@ class DebounceMiddleware:
|
||||
|
||||
# ── Chunking ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class ChunkingMiddleware(OutboundMiddlewareBase):
|
||||
"""Auto-split messages respecting format expansion.
|
||||
|
||||
@@ -480,6 +498,7 @@ class ChunkingMiddleware(OutboundMiddlewareBase):
|
||||
|
||||
def __init__(self, capabilities: Any) -> None:
|
||||
from .capabilities import ChannelCapabilities
|
||||
|
||||
self._capabilities: ChannelCapabilities = capabilities
|
||||
|
||||
def prepare_chunks(
|
||||
@@ -516,6 +535,7 @@ class ChunkingMiddleware(OutboundMiddlewareBase):
|
||||
|
||||
# ── Formatting ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class FormattingMiddleware(OutboundMiddlewareBase):
|
||||
"""Markdown -> channel format conversion.
|
||||
|
||||
@@ -525,6 +545,7 @@ class FormattingMiddleware(OutboundMiddlewareBase):
|
||||
def __init__(self, capabilities: Any) -> None:
|
||||
from .formatter import UnifiedFormatter
|
||||
from .capabilities import ChannelCapabilities
|
||||
|
||||
caps: ChannelCapabilities = capabilities
|
||||
self._formatter = UnifiedFormatter.for_channel(caps.format_type)
|
||||
|
||||
@@ -533,7 +554,9 @@ class FormattingMiddleware(OutboundMiddlewareBase):
|
||||
return self._formatter.format(text)
|
||||
|
||||
async def process_outbound(
|
||||
self, message: OutboundMessage, context: dict[str, Any],
|
||||
self,
|
||||
message: OutboundMessage,
|
||||
context: dict[str, Any],
|
||||
) -> OutboundMessage | None:
|
||||
formatted = self._formatter.format(message.content)
|
||||
return dataclasses.replace(message, content=formatted)
|
||||
@@ -541,6 +564,7 @@ class FormattingMiddleware(OutboundMiddlewareBase):
|
||||
|
||||
# ── Retry ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class RetryMiddleware:
|
||||
"""Exponential backoff send retry.
|
||||
|
||||
@@ -549,6 +573,7 @@ class RetryMiddleware:
|
||||
|
||||
def __init__(self, channel_name: str = "unknown") -> None:
|
||||
from .retry import DEFAULT_RETRY, RETRY_PRESETS
|
||||
|
||||
self._config = RETRY_PRESETS.get(channel_name, DEFAULT_RETRY)
|
||||
self._channel_name = channel_name
|
||||
|
||||
@@ -576,6 +601,7 @@ class RetryMiddleware:
|
||||
|
||||
# ── Typing ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TypingMiddleware:
|
||||
"""Typing indicator management.
|
||||
|
||||
@@ -601,6 +627,7 @@ class TypingMiddleware:
|
||||
|
||||
# ── ACK Reaction ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class AckReactionMiddleware:
|
||||
"""ACK emoji reaction with configurable scope.
|
||||
|
||||
@@ -661,6 +688,7 @@ class AckReactionMiddleware:
|
||||
|
||||
# ── Mention Gating ───────────────────────────────────────────────────
|
||||
|
||||
|
||||
class MentionGatingMiddleware(InboundMiddleware):
|
||||
"""Filter messages based on mention policy.
|
||||
|
||||
@@ -679,7 +707,9 @@ class MentionGatingMiddleware(InboundMiddleware):
|
||||
self._strip_fn = strip_fn
|
||||
|
||||
async def process_inbound(
|
||||
self, raw: RawIncoming, context: dict[str, Any],
|
||||
self,
|
||||
raw: RawIncoming,
|
||||
context: dict[str, Any],
|
||||
) -> RawIncoming | None:
|
||||
if not self._should_process(raw):
|
||||
return None
|
||||
@@ -701,6 +731,7 @@ class MentionGatingMiddleware(InboundMiddleware):
|
||||
|
||||
# ── AllowList ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class AllowListMiddleware(InboundMiddleware):
|
||||
"""Sender and channel allow-list enforcement."""
|
||||
|
||||
@@ -715,7 +746,9 @@ class AllowListMiddleware(InboundMiddleware):
|
||||
self.dm_policy = dm_policy
|
||||
|
||||
async def process_inbound(
|
||||
self, raw: RawIncoming, context: dict[str, Any],
|
||||
self,
|
||||
raw: RawIncoming,
|
||||
context: dict[str, Any],
|
||||
) -> RawIncoming | None:
|
||||
# Channel allow-list
|
||||
if self.allowed_channels and str(raw.chat_id) not in self.allowed_channels:
|
||||
@@ -747,6 +780,7 @@ class AllowListMiddleware(InboundMiddleware):
|
||||
|
||||
# ── Group History ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class GroupHistoryMiddleware(InboundMiddleware):
|
||||
"""Buffer non-mentioned group messages, inject as context when mentioned."""
|
||||
|
||||
@@ -756,11 +790,14 @@ class GroupHistoryMiddleware(InboundMiddleware):
|
||||
max_age_seconds: int = 3600,
|
||||
) -> None:
|
||||
self._buffer = GroupHistoryBuffer(
|
||||
max_per_chat=max_per_chat, max_age_seconds=max_age_seconds,
|
||||
max_per_chat=max_per_chat,
|
||||
max_age_seconds=max_age_seconds,
|
||||
)
|
||||
|
||||
async def process_inbound(
|
||||
self, raw: RawIncoming, context: dict[str, Any],
|
||||
self,
|
||||
raw: RawIncoming,
|
||||
context: dict[str, Any],
|
||||
) -> RawIncoming | None:
|
||||
if not raw.is_group:
|
||||
return raw
|
||||
@@ -786,7 +823,9 @@ class GroupHistoryMiddleware(InboundMiddleware):
|
||||
if history_context:
|
||||
raw = dataclasses.replace(
|
||||
raw,
|
||||
text=history_context + "\n\n[Current message - respond to this]\n" + raw.text,
|
||||
text=history_context
|
||||
+ "\n\n[Current message - respond to this]\n"
|
||||
+ raw.text,
|
||||
)
|
||||
self._buffer.clear(raw.chat_id)
|
||||
return raw
|
||||
@@ -794,6 +833,7 @@ class GroupHistoryMiddleware(InboundMiddleware):
|
||||
|
||||
# ── Pairing ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class PairingMiddleware(InboundMiddleware):
|
||||
"""DM pairing flow management.
|
||||
|
||||
@@ -814,7 +854,9 @@ class PairingMiddleware(InboundMiddleware):
|
||||
self._background_tasks: set[asyncio.Task] = set()
|
||||
|
||||
async def process_inbound(
|
||||
self, raw: RawIncoming, context: dict[str, Any],
|
||||
self,
|
||||
raw: RawIncoming,
|
||||
context: dict[str, Any],
|
||||
) -> RawIncoming | None:
|
||||
if raw.is_group:
|
||||
return raw # pairing only applies to DMs
|
||||
|
||||
@@ -26,6 +26,7 @@ logger = logging.getLogger(__name__)
|
||||
# Token refresh mixin (shared by Webhook & WebSocket channels)
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class TokenMixin:
|
||||
"""Mixin for channels that need OAuth-style token management.
|
||||
|
||||
@@ -48,7 +49,9 @@ class TokenMixin:
|
||||
token, expire = await self._fetch_token()
|
||||
self._access_token = token
|
||||
self._token_expires = time.monotonic() + expire - 300
|
||||
logger.debug(f"{getattr(self, 'name', '?')} token refreshed, expires in {expire}s")
|
||||
logger.debug(
|
||||
f"{getattr(self, 'name', '?')} token refreshed, expires in {expire}s"
|
||||
)
|
||||
|
||||
async def _ensure_token(self) -> str:
|
||||
if not self._access_token or time.monotonic() >= self._token_expires:
|
||||
@@ -60,6 +63,7 @@ class TokenMixin:
|
||||
# Webhook + REST mixin
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class WebhookMixin:
|
||||
"""Mixin for channels that use an HTTP webhook server for inbound
|
||||
and REST API for outbound.
|
||||
@@ -129,7 +133,9 @@ class WebhookMixin:
|
||||
await self._http_client.aclose()
|
||||
self._http_client = None
|
||||
|
||||
async def _api_post(self, url: str, body: dict, headers: dict | None = None) -> dict:
|
||||
async def _api_post(
|
||||
self, url: str, body: dict, headers: dict | None = None
|
||||
) -> dict:
|
||||
"""POST JSON to API, return parsed response. Raises on HTTP error."""
|
||||
resp = await self._http_client.post(url, json=body, headers=headers)
|
||||
data = resp.json()
|
||||
@@ -144,6 +150,7 @@ class WebhookMixin:
|
||||
# WebSocket mixin
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class WebSocketMixin:
|
||||
"""Mixin for channels that receive messages via WebSocket.
|
||||
|
||||
@@ -193,14 +200,22 @@ class WebSocketMixin:
|
||||
# Resolve proxy: channel config > environment variable
|
||||
proxy = getattr(getattr(self, "config", None), "proxy", None)
|
||||
if not proxy:
|
||||
proxy = (os.environ.get("https_proxy")
|
||||
or os.environ.get("HTTPS_PROXY")
|
||||
or os.environ.get("http_proxy")
|
||||
or os.environ.get("HTTP_PROXY")
|
||||
or None)
|
||||
logger.debug(f"{getattr(self, 'name', '?')} WS connecting to {ws_url[:60]}... proxy={proxy}")
|
||||
proxy = (
|
||||
os.environ.get("https_proxy")
|
||||
or os.environ.get("HTTPS_PROXY")
|
||||
or os.environ.get("http_proxy")
|
||||
or os.environ.get("HTTP_PROXY")
|
||||
or None
|
||||
)
|
||||
logger.debug(
|
||||
f"{getattr(self, 'name', '?')} WS connecting to {ws_url[:60]}... proxy={proxy}"
|
||||
)
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.ws_connect(ws_url, proxy=proxy, timeout=aiohttp.ClientWSTimeout(ws_close=30)) as ws:
|
||||
async with session.ws_connect(
|
||||
ws_url,
|
||||
proxy=proxy,
|
||||
timeout=aiohttp.ClientWSTimeout(ws_close=30),
|
||||
) as ws:
|
||||
logger.info(f"{getattr(self, 'name', '?')} WebSocket connected")
|
||||
self._ws_session = ws
|
||||
await self._on_ws_connected(ws)
|
||||
@@ -233,7 +248,9 @@ class WebSocketMixin:
|
||||
self._ws_session = None
|
||||
|
||||
if getattr(self, "_running", False):
|
||||
logger.info(f"{getattr(self, 'name', '?')} reconnecting in {self._ws_reconnect_delay}s...")
|
||||
logger.info(
|
||||
f"{getattr(self, 'name', '?')} reconnecting in {self._ws_reconnect_delay}s..."
|
||||
)
|
||||
await asyncio.sleep(self._ws_reconnect_delay)
|
||||
|
||||
async def _ws_heartbeat_loop(self, ws) -> None:
|
||||
@@ -274,6 +291,7 @@ class WebSocketMixin:
|
||||
# Polling mixin
|
||||
# ═════════════════════════════════════════════════════════════════════
|
||||
|
||||
|
||||
class PollingMixin:
|
||||
"""Mixin for channels that poll for new messages.
|
||||
|
||||
|
||||
@@ -18,6 +18,7 @@ from .capabilities import ChannelCapabilities
|
||||
|
||||
# ── Channel metadata ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChannelMeta:
|
||||
"""Channel metadata for registry and UI."""
|
||||
@@ -31,6 +32,7 @@ class ChannelMeta:
|
||||
|
||||
# ── Adapter Protocols (slots) ────────────────────────────────────────
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ConfigAdapter(Protocol):
|
||||
"""Account configuration management."""
|
||||
@@ -45,7 +47,9 @@ class ConfigAdapter(Protocol):
|
||||
class SecurityAdapter(Protocol):
|
||||
"""DM policy and security warnings."""
|
||||
|
||||
def resolve_dm_policy(self, ctx: Any) -> str: ... # "open" | "allowlist" | "pairing"
|
||||
def resolve_dm_policy(
|
||||
self, ctx: Any
|
||||
) -> str: ... # "open" | "allowlist" | "pairing"
|
||||
def collect_warnings(self, ctx: Any) -> list[str]: ...
|
||||
|
||||
|
||||
@@ -142,6 +146,7 @@ class OnboardingAdapter(Protocol):
|
||||
|
||||
# ── Reload policy ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReloadPolicy:
|
||||
"""Declares which config prefixes trigger a channel reload."""
|
||||
@@ -152,6 +157,7 @@ class ReloadPolicy:
|
||||
|
||||
# ── ChannelPlugin ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class ChannelPlugin:
|
||||
"""Declarative channel plugin with optional adapter slots.
|
||||
|
||||
@@ -192,7 +198,9 @@ class ChannelPlugin:
|
||||
# Provide default SingleAccountConfigAdapter if not overridden
|
||||
if self.config_adapter is None:
|
||||
from .config import SingleAccountConfigAdapter
|
||||
|
||||
self.config_adapter = SingleAccountConfigAdapter()
|
||||
|
||||
security: SecurityAdapter | None = None
|
||||
groups: GroupAdapter | None = None
|
||||
mentions: MentionAdapter | None = None
|
||||
@@ -219,8 +227,18 @@ class ChannelPlugin:
|
||||
def filled_slots(self) -> list[str]:
|
||||
"""Return names of adapter slots that are not None."""
|
||||
slot_names = [
|
||||
"config_adapter", "security", "groups", "mentions", "outbound",
|
||||
"threading", "streaming", "directory", "status", "heartbeat",
|
||||
"actions", "pairing", "onboarding",
|
||||
"config_adapter",
|
||||
"security",
|
||||
"groups",
|
||||
"mentions",
|
||||
"outbound",
|
||||
"threading",
|
||||
"streaming",
|
||||
"directory",
|
||||
"status",
|
||||
"heartbeat",
|
||||
"actions",
|
||||
"pairing",
|
||||
"onboarding",
|
||||
]
|
||||
return [s for s in slot_names if getattr(self, s, None) is not None]
|
||||
|
||||
@@ -16,11 +16,13 @@ __all__ = ["QQChannel", "QQConfig"]
|
||||
|
||||
def create_from_config(config) -> QQChannel:
|
||||
allowed = _parse_csv(config.qq_allowed_senders)
|
||||
return QQChannel(QQConfig(
|
||||
app_id=config.qq_app_id,
|
||||
app_secret=config.qq_app_secret,
|
||||
allowed_senders=allowed,
|
||||
))
|
||||
return QQChannel(
|
||||
QQConfig(
|
||||
app_id=config.qq_app_id,
|
||||
app_secret=config.qq_app_secret,
|
||||
allowed_senders=allowed,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
register_channel("qq", create_from_config)
|
||||
|
||||
@@ -86,7 +86,9 @@ class QQChannel(Channel):
|
||||
|
||||
async def _run_bot(self) -> None:
|
||||
try:
|
||||
await self._client.start(appid=self.config.app_id, secret=self.config.app_secret)
|
||||
await self._client.start(
|
||||
appid=self.config.app_id, secret=self.config.app_secret
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"QQ auth failed: {e}")
|
||||
self._running = False
|
||||
@@ -130,7 +132,8 @@ class QQChannel(Channel):
|
||||
content_type = getattr(att, "content_type", "") or ""
|
||||
if url:
|
||||
local, ann = await self._download_attachment(
|
||||
url, f"qq_{filename}",
|
||||
url,
|
||||
f"qq_{filename}",
|
||||
)
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
@@ -142,23 +145,25 @@ class QQChannel(Channel):
|
||||
if not content and not media_paths and not annotations:
|
||||
return
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=sender_id,
|
||||
chat_id=chat_id,
|
||||
text=content,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=datetime.now(),
|
||||
message_id=message.id,
|
||||
is_group=(msg_type == "group"),
|
||||
was_mentioned=True,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"msg_type": msg_type,
|
||||
"event_id": message.id,
|
||||
"backend": "qq",
|
||||
},
|
||||
))
|
||||
await self._enqueue_raw(
|
||||
RawIncoming(
|
||||
sender_id=sender_id,
|
||||
chat_id=chat_id,
|
||||
text=content,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=datetime.now(),
|
||||
message_id=message.id,
|
||||
is_group=(msg_type == "group"),
|
||||
was_mentioned=True,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"msg_type": msg_type,
|
||||
"event_id": message.id,
|
||||
"backend": "qq",
|
||||
},
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Error handling QQ message: {e}")
|
||||
|
||||
@@ -185,13 +190,19 @@ class QQChannel(Channel):
|
||||
seq = self._next_msg_seq(msg_id)
|
||||
if msg_type == "group":
|
||||
await self._client.api.post_group_message(
|
||||
group_openid=chat_id, msg_type=0,
|
||||
content=raw_text, msg_id=msg_id, msg_seq=seq,
|
||||
group_openid=chat_id,
|
||||
msg_type=0,
|
||||
content=raw_text,
|
||||
msg_id=msg_id,
|
||||
msg_seq=seq,
|
||||
)
|
||||
else:
|
||||
await self._client.api.post_c2c_message(
|
||||
openid=chat_id, msg_type=0,
|
||||
content=raw_text, msg_id=msg_id, msg_seq=seq,
|
||||
openid=chat_id,
|
||||
msg_type=0,
|
||||
content=raw_text,
|
||||
msg_id=msg_id,
|
||||
msg_seq=seq,
|
||||
)
|
||||
|
||||
# _send_typing_action: inherited no-op (QQ Bot API has no typing indicator)
|
||||
@@ -200,9 +211,20 @@ class QQChannel(Channel):
|
||||
|
||||
# qq-botpy file_type constants: 1=image, 2=video, 3=audio
|
||||
_FILE_TYPE_MAP = {
|
||||
".jpg": 1, ".jpeg": 1, ".png": 1, ".gif": 1, ".webp": 1, ".bmp": 1,
|
||||
".mp4": 2, ".mov": 2, ".avi": 2,
|
||||
".mp3": 3, ".ogg": 3, ".m4a": 3, ".wav": 3, ".silk": 3,
|
||||
".jpg": 1,
|
||||
".jpeg": 1,
|
||||
".png": 1,
|
||||
".gif": 1,
|
||||
".webp": 1,
|
||||
".bmp": 1,
|
||||
".mp4": 2,
|
||||
".mov": 2,
|
||||
".avi": 2,
|
||||
".mp3": 3,
|
||||
".ogg": 3,
|
||||
".m4a": 3,
|
||||
".wav": 3,
|
||||
".silk": 3,
|
||||
}
|
||||
|
||||
async def _send_media_impl(
|
||||
@@ -221,6 +243,7 @@ class QQChannel(Channel):
|
||||
raise ChannelError("QQ client not initialized")
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
chat_id = self._resolve_media_chat_id(recipient, metadata)
|
||||
msg_type = (metadata or {}).get("msg_type", "c2c")
|
||||
ext = Path(file_path).suffix.lower()
|
||||
|
||||
@@ -92,13 +92,15 @@ async def retry_async(
|
||||
delay = max(config.min_delay_s, min(jittered, config.max_delay_s))
|
||||
|
||||
if on_retry is not None:
|
||||
on_retry(RetryInfo(
|
||||
attempt=attempt,
|
||||
max_attempts=config.attempts,
|
||||
delay_s=delay,
|
||||
error=exc,
|
||||
label=label,
|
||||
))
|
||||
on_retry(
|
||||
RetryInfo(
|
||||
attempt=attempt,
|
||||
max_attempts=config.attempts,
|
||||
delay_s=delay,
|
||||
error=exc,
|
||||
label=label,
|
||||
)
|
||||
)
|
||||
|
||||
await asyncio.sleep(delay)
|
||||
|
||||
|
||||
@@ -15,13 +15,15 @@ __all__ = ["SignalChannel", "SignalConfig"]
|
||||
|
||||
def create_from_config(config) -> SignalChannel:
|
||||
allowed = _parse_csv(config.signal_allowed_senders)
|
||||
return SignalChannel(SignalConfig(
|
||||
phone_number=config.signal_phone_number,
|
||||
cli_path=config.signal_cli_path,
|
||||
config_dir=config.signal_config_dir or None,
|
||||
rpc_port=config.signal_rpc_port,
|
||||
allowed_senders=allowed,
|
||||
))
|
||||
return SignalChannel(
|
||||
SignalConfig(
|
||||
phone_number=config.signal_phone_number,
|
||||
cli_path=config.signal_cli_path,
|
||||
config_dir=config.signal_config_dir or None,
|
||||
rpc_port=config.signal_rpc_port,
|
||||
allowed_senders=allowed,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
register_channel("signal", create_from_config)
|
||||
|
||||
@@ -109,13 +109,21 @@ class SignalChannel(Channel):
|
||||
cmd = [self.config.cli_path, "-u", self.config.phone_number]
|
||||
if self.config.config_dir:
|
||||
cmd.extend(["--config", self.config.config_dir])
|
||||
cmd.extend(["daemon", "--tcp",
|
||||
f"localhost:{self.config.rpc_port}", "--no-receive-stdout"])
|
||||
cmd.extend(
|
||||
[
|
||||
"daemon",
|
||||
"--tcp",
|
||||
f"localhost:{self.config.rpc_port}",
|
||||
"--no-receive-stdout",
|
||||
]
|
||||
)
|
||||
|
||||
logger.info(f"Starting signal-cli daemon: {' '.join(cmd)}")
|
||||
try:
|
||||
self._daemon_proc = subprocess.Popen(
|
||||
cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL,
|
||||
cmd,
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
)
|
||||
except FileNotFoundError:
|
||||
raise ChannelError(
|
||||
@@ -128,7 +136,8 @@ class SignalChannel(Channel):
|
||||
await asyncio.sleep(1)
|
||||
try:
|
||||
reader, writer = await asyncio.open_connection(
|
||||
"localhost", self.config.rpc_port,
|
||||
"localhost",
|
||||
self.config.rpc_port,
|
||||
)
|
||||
writer.close()
|
||||
await writer.wait_closed()
|
||||
@@ -143,7 +152,8 @@ class SignalChannel(Channel):
|
||||
"""Connect to signal-cli JSON RPC socket."""
|
||||
try:
|
||||
self._reader, self._writer = await asyncio.open_connection(
|
||||
"localhost", self.config.rpc_port,
|
||||
"localhost",
|
||||
self.config.rpc_port,
|
||||
)
|
||||
except Exception as e:
|
||||
raise ChannelError(f"Cannot connect to signal-cli: {e}")
|
||||
@@ -199,7 +209,10 @@ class SignalChannel(Channel):
|
||||
timestamp = envelope.get("timestamp", 0)
|
||||
|
||||
# Ignore messages from self
|
||||
if source_number == self.config.phone_number or source == self.config.phone_number:
|
||||
if (
|
||||
source_number == self.config.phone_number
|
||||
or source == self.config.phone_number
|
||||
):
|
||||
logger.debug("Ignoring message from self")
|
||||
return
|
||||
|
||||
@@ -209,12 +222,20 @@ class SignalChannel(Channel):
|
||||
text = data_msg.get("message", "")
|
||||
group_info = data_msg.get("groupInfo", {})
|
||||
is_group = bool(group_info)
|
||||
chat_id = group_info.get("groupId", source_number) if is_group else source_number
|
||||
chat_id = (
|
||||
group_info.get("groupId", source_number) if is_group else source_number
|
||||
)
|
||||
msg_ts = data_msg.get("timestamp", timestamp)
|
||||
|
||||
media_paths: list[str] = []
|
||||
annotations: list[str] = []
|
||||
_VOICE_TYPES = {"audio/aac", "audio/ogg", "audio/mp4", "audio/mpeg", "audio/opus"}
|
||||
_VOICE_TYPES = {
|
||||
"audio/aac",
|
||||
"audio/ogg",
|
||||
"audio/mp4",
|
||||
"audio/mpeg",
|
||||
"audio/opus",
|
||||
}
|
||||
attachments = data_msg.get("attachments", [])
|
||||
for att in attachments:
|
||||
att_size = att.get("size", 0)
|
||||
@@ -225,19 +246,26 @@ class SignalChannel(Channel):
|
||||
media_label = "voice" if is_voice else "attachment"
|
||||
if att_file:
|
||||
from pathlib import Path as _Path
|
||||
|
||||
att_path = _Path(att_file)
|
||||
if att_path.exists():
|
||||
from ..base import MAX_ATTACHMENT_BYTES
|
||||
|
||||
if att_path.stat().st_size > MAX_ATTACHMENT_BYTES:
|
||||
annotations.append(f"[{media_label}: {att_name} - too large ({att_path.stat().st_size} bytes)]")
|
||||
annotations.append(
|
||||
f"[{media_label}: {att_name} - too large ({att_path.stat().st_size} bytes)]"
|
||||
)
|
||||
else:
|
||||
local = self._media_path(f"signal_{att_name}")
|
||||
import shutil
|
||||
|
||||
shutil.copy2(str(att_path), str(local))
|
||||
media_paths.append(str(local))
|
||||
annotations.append(f"[{media_label}: {local}]")
|
||||
else:
|
||||
annotations.append(f"[{media_label}: {att_name} - file not found]")
|
||||
annotations.append(
|
||||
f"[{media_label}: {att_name} - file not found]"
|
||||
)
|
||||
elif att_size:
|
||||
too_large = self._check_attachment_size(att_size, att_name)
|
||||
if too_large:
|
||||
@@ -261,31 +289,40 @@ class SignalChannel(Channel):
|
||||
if is_group:
|
||||
mentions = data_msg.get("mentions", [])
|
||||
for m in mentions:
|
||||
if m.get("uuid") == self.config.phone_number or m.get("number") == self.config.phone_number:
|
||||
if (
|
||||
m.get("uuid") == self.config.phone_number
|
||||
or m.get("number") == self.config.phone_number
|
||||
):
|
||||
was_mentioned = True
|
||||
break
|
||||
|
||||
# Cache message_id → sender for reaction targetAuthor
|
||||
self._cache_msg_sender(str(msg_ts), source_number)
|
||||
|
||||
logger.info("Signal message from %s: %s", source_number, text[:50] if text else "[media]")
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=source_number,
|
||||
chat_id=chat_id,
|
||||
text=text,
|
||||
content_annotations=annotations,
|
||||
media_files=media_paths,
|
||||
timestamp=ts,
|
||||
message_id=str(msg_ts),
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"source_name": source_name,
|
||||
"sender_id": source_number,
|
||||
"backend": "signal",
|
||||
},
|
||||
))
|
||||
logger.info(
|
||||
"Signal message from %s: %s",
|
||||
source_number,
|
||||
text[:50] if text else "[media]",
|
||||
)
|
||||
await self._enqueue_raw(
|
||||
RawIncoming(
|
||||
sender_id=source_number,
|
||||
chat_id=chat_id,
|
||||
text=text,
|
||||
content_annotations=annotations,
|
||||
media_files=media_paths,
|
||||
timestamp=ts,
|
||||
message_id=str(msg_ts),
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"source_name": source_name,
|
||||
"sender_id": source_number,
|
||||
"backend": "signal",
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
# ── Typing indicator ────────────────────────────────────────────
|
||||
|
||||
@@ -313,7 +350,9 @@ class SignalChannel(Channel):
|
||||
self._msg_senders[message_id] = sender
|
||||
self._msg_senders_order.append(message_id)
|
||||
|
||||
async def _send_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
async def _send_ack_reaction(
|
||||
self, chat_id: str, message_id: str, emoji: str = "👀"
|
||||
) -> None:
|
||||
"""Send an acknowledgment reaction via signal-cli sendReaction."""
|
||||
target_author = self._msg_senders.get(message_id, "")
|
||||
if not target_author:
|
||||
@@ -333,7 +372,9 @@ class SignalChannel(Channel):
|
||||
except Exception as e:
|
||||
logger.debug(f"Signal ack reaction failed: {e}")
|
||||
|
||||
async def _remove_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
async def _remove_ack_reaction(
|
||||
self, chat_id: str, message_id: str, emoji: str = "👀"
|
||||
) -> None:
|
||||
"""Remove ACK reaction via signal-cli sendReaction --remove."""
|
||||
target_author = self._msg_senders.get(message_id, "")
|
||||
if not target_author:
|
||||
@@ -369,7 +410,9 @@ class SignalChannel(Channel):
|
||||
def _is_ready(self) -> bool:
|
||||
return self._writer is not None and not self._writer.is_closing()
|
||||
|
||||
async def _rpc_call(self, method: str, params: dict, timeout: float = 10.0) -> dict | None:
|
||||
async def _rpc_call(
|
||||
self, method: str, params: dict, timeout: float = 10.0
|
||||
) -> dict | None:
|
||||
"""Send a JSON RPC call to signal-cli and wait for the response."""
|
||||
if not self._writer:
|
||||
return None
|
||||
@@ -400,7 +443,12 @@ class SignalChannel(Channel):
|
||||
return None
|
||||
|
||||
async def _send_chunk(
|
||||
self, chat_id, formatted_text, raw_text, reply_to, metadata,
|
||||
self,
|
||||
chat_id,
|
||||
formatted_text,
|
||||
raw_text,
|
||||
reply_to,
|
||||
metadata,
|
||||
):
|
||||
# Determine if group or individual
|
||||
params: dict[str, Any] = {
|
||||
@@ -429,7 +477,7 @@ class SignalChannel(Channel):
|
||||
# Remove phone number if directly mentioned as text
|
||||
text = re.sub(rf"@?{re.escape(phone)}\s*", "", text).strip()
|
||||
# Remove Unicode Object Replacement Character used as mention placeholder
|
||||
text = text.replace("\uFFFC", "").strip()
|
||||
text = text.replace("\ufffc", "").strip()
|
||||
return text
|
||||
|
||||
# ── Media send ────────────────────────────────────────────────
|
||||
|
||||
@@ -23,10 +23,14 @@ async def validate_signal(
|
||||
|
||||
# Check signal-cli binary
|
||||
loop = asyncio.get_event_loop()
|
||||
|
||||
def _check():
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[cli_path, "--version"], capture_output=True, text=True, timeout=5,
|
||||
[cli_path, "--version"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=5,
|
||||
)
|
||||
if result.returncode == 0:
|
||||
return True, f"signal-cli {result.stdout.strip()}"
|
||||
|
||||
@@ -8,13 +8,15 @@ def create_from_config(config) -> SlackChannel:
|
||||
allowed = _parse_csv(config.slack_allowed_senders)
|
||||
channels = _parse_csv(config.slack_allowed_channels)
|
||||
proxy = config.slack_proxy if config.slack_proxy else None
|
||||
return SlackChannel(SlackConfig(
|
||||
bot_token=config.slack_bot_token,
|
||||
app_token=config.slack_app_token,
|
||||
allowed_senders=allowed,
|
||||
allowed_channels=channels,
|
||||
proxy=proxy,
|
||||
))
|
||||
return SlackChannel(
|
||||
SlackConfig(
|
||||
bot_token=config.slack_bot_token,
|
||||
app_token=config.slack_app_token,
|
||||
allowed_senders=allowed,
|
||||
allowed_channels=channels,
|
||||
proxy=proxy,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
register_channel("slack", create_from_config)
|
||||
|
||||
@@ -39,8 +39,7 @@ class SlackChannel(Channel):
|
||||
raise ChannelError("Slack bot token is required")
|
||||
if not self.config.app_token:
|
||||
raise ChannelError(
|
||||
"Slack app token is required for Socket Mode "
|
||||
"(starts with xapp-)"
|
||||
"Slack app token is required for Socket Mode (starts with xapp-)"
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -62,7 +61,8 @@ class SlackChannel(Channel):
|
||||
# Get bot user ID for filtering own messages
|
||||
try:
|
||||
auth = await asyncio.wait_for(
|
||||
self._web_client.auth_test(), timeout=15,
|
||||
self._web_client.auth_test(),
|
||||
timeout=15,
|
||||
)
|
||||
self._bot_user_id = auth["user_id"]
|
||||
except asyncio.TimeoutError:
|
||||
@@ -93,19 +93,22 @@ class SlackChannel(Channel):
|
||||
if event_type == "message" and "subtype" not in event:
|
||||
is_dm = event.get("channel_type") == "im"
|
||||
await self._on_message(
|
||||
event, is_group=not is_dm, was_mentioned=is_dm,
|
||||
event,
|
||||
is_group=not is_dm,
|
||||
was_mentioned=is_dm,
|
||||
)
|
||||
elif event_type == "app_mention":
|
||||
await self._on_message(
|
||||
event, is_group=True, was_mentioned=True,
|
||||
event,
|
||||
is_group=True,
|
||||
was_mentioned=True,
|
||||
)
|
||||
|
||||
self._socket_client.socket_mode_request_listeners.append(
|
||||
_event_handler
|
||||
)
|
||||
self._socket_client.socket_mode_request_listeners.append(_event_handler)
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
self._socket_client.connect(), timeout=30,
|
||||
self._socket_client.connect(),
|
||||
timeout=30,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
raise ChannelError(
|
||||
@@ -159,7 +162,6 @@ class SlackChannel(Channel):
|
||||
|
||||
# ── Send (template method overrides) ──────────────────────────
|
||||
|
||||
|
||||
async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata):
|
||||
kwargs = dict(channel=chat_id)
|
||||
# Always route to thread if thread_ts is present in metadata,
|
||||
@@ -195,22 +197,30 @@ class SlackChannel(Channel):
|
||||
|
||||
# ── ACK Reactions ───────────────────────────────────────────────
|
||||
|
||||
async def _send_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "eyes") -> None:
|
||||
async def _send_ack_reaction(
|
||||
self, chat_id: str, message_id: str, emoji: str = "eyes"
|
||||
) -> None:
|
||||
"""Add an emoji reaction to acknowledge receipt."""
|
||||
if self._web_client and message_id:
|
||||
try:
|
||||
await self._web_client.reactions_add(
|
||||
channel=chat_id, timestamp=message_id, name=emoji,
|
||||
channel=chat_id,
|
||||
timestamp=message_id,
|
||||
name=emoji,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Slack ACK reaction failed: {e}")
|
||||
|
||||
async def _remove_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "eyes") -> None:
|
||||
async def _remove_ack_reaction(
|
||||
self, chat_id: str, message_id: str, emoji: str = "eyes"
|
||||
) -> None:
|
||||
"""Remove the ACK reaction after replying."""
|
||||
if self._web_client and message_id:
|
||||
try:
|
||||
await self._web_client.reactions_remove(
|
||||
channel=chat_id, timestamp=message_id, name=emoji,
|
||||
channel=chat_id,
|
||||
timestamp=message_id,
|
||||
name=emoji,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.debug(f"Slack remove ACK reaction failed: {e}")
|
||||
@@ -253,11 +263,10 @@ class SlackChannel(Channel):
|
||||
"url_private"
|
||||
)
|
||||
if url and self._web_client:
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.config.bot_token}"
|
||||
}
|
||||
headers = {"Authorization": f"Bearer {self.config.bot_token}"}
|
||||
local_path, annotation = await self._download_attachment(
|
||||
url, f"{file_info.get('id', 'unknown')}_{filename}",
|
||||
url,
|
||||
f"{file_info.get('id', 'unknown')}_{filename}",
|
||||
headers=headers,
|
||||
file_size=file_size,
|
||||
)
|
||||
@@ -273,18 +282,20 @@ class SlackChannel(Channel):
|
||||
except (ValueError, TypeError):
|
||||
timestamp = datetime.now()
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=user_id,
|
||||
chat_id=channel_id,
|
||||
text=text,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=timestamp,
|
||||
message_id=ts,
|
||||
metadata={"chat_id": channel_id, "thread_ts": thread_ts},
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
))
|
||||
await self._enqueue_raw(
|
||||
RawIncoming(
|
||||
sender_id=user_id,
|
||||
chat_id=channel_id,
|
||||
text=text,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=timestamp,
|
||||
message_id=ts,
|
||||
metadata={"chat_id": channel_id, "thread_ts": thread_ts},
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
)
|
||||
)
|
||||
logger.info(
|
||||
f"Slack message queued: sender={user_id}, "
|
||||
f"channel={channel_id}, content={text[:50]}"
|
||||
|
||||
@@ -26,13 +26,15 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def standalone_outbound_dispatcher(
|
||||
bus: MessageBus, channel: Channel,
|
||||
bus: MessageBus,
|
||||
channel: Channel,
|
||||
) -> None:
|
||||
"""Consume outbound messages from the bus and send via channel."""
|
||||
while True:
|
||||
try:
|
||||
msg: OutboundMessage = await asyncio.wait_for(
|
||||
bus.consume_outbound(), timeout=1.0,
|
||||
bus.consume_outbound(),
|
||||
timeout=1.0,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
@@ -47,8 +49,10 @@ async def standalone_outbound_dispatcher(
|
||||
|
||||
|
||||
async def _async_main(
|
||||
channel: Channel, bus: MessageBus,
|
||||
use_agent: bool, send_thinking: bool,
|
||||
channel: Channel,
|
||||
bus: MessageBus,
|
||||
use_agent: bool,
|
||||
send_thinking: bool,
|
||||
) -> None:
|
||||
"""Async entry point — gather channel, dispatcher and optional consumer."""
|
||||
from .channel_manager import ChannelManager
|
||||
@@ -72,6 +76,7 @@ async def _async_main(
|
||||
if use_agent:
|
||||
logger.info("Loading EvoScientist agent...")
|
||||
from ..EvoScientist import create_cli_agent
|
||||
|
||||
agent = create_cli_agent()
|
||||
logger.info("Agent loaded")
|
||||
|
||||
@@ -114,15 +119,19 @@ async def _async_main(
|
||||
loop = asyncio.get_event_loop()
|
||||
for sig in (signal.SIGINT, signal.SIGTERM):
|
||||
loop.add_signal_handler(
|
||||
sig, lambda s=sig: asyncio.create_task(_graceful_shutdown()),
|
||||
sig,
|
||||
lambda s=sig: asyncio.create_task(_graceful_shutdown()),
|
||||
)
|
||||
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
|
||||
def run_standalone(
|
||||
channel: Channel, bus: MessageBus, *,
|
||||
use_agent: bool = False, send_thinking: bool = False,
|
||||
channel: Channel,
|
||||
bus: MessageBus,
|
||||
*,
|
||||
use_agent: bool = False,
|
||||
send_thinking: bool = False,
|
||||
) -> None:
|
||||
"""Synchronous entry point that spins up the standalone runner.
|
||||
|
||||
|
||||
@@ -7,11 +7,13 @@ __all__ = ["TelegramChannel", "TelegramConfig"]
|
||||
def create_from_config(config) -> TelegramChannel:
|
||||
allowed = _parse_csv(config.telegram_allowed_senders)
|
||||
proxy = config.telegram_proxy if config.telegram_proxy else None
|
||||
return TelegramChannel(TelegramConfig(
|
||||
bot_token=config.telegram_bot_token,
|
||||
allowed_senders=allowed,
|
||||
proxy=proxy,
|
||||
))
|
||||
return TelegramChannel(
|
||||
TelegramConfig(
|
||||
bot_token=config.telegram_bot_token,
|
||||
allowed_senders=allowed,
|
||||
proxy=proxy,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
register_channel("telegram", create_from_config)
|
||||
|
||||
@@ -5,7 +5,14 @@ from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from ..base import Channel, RawIncoming, ChannelError, IMAGE_EXTS, VIDEO_EXTS, AUDIO_EXTS
|
||||
from ..base import (
|
||||
Channel,
|
||||
RawIncoming,
|
||||
ChannelError,
|
||||
IMAGE_EXTS,
|
||||
VIDEO_EXTS,
|
||||
AUDIO_EXTS,
|
||||
)
|
||||
from ..capabilities import TELEGRAM as TELEGRAM_CAPS
|
||||
from ..config import BaseChannelConfig
|
||||
|
||||
@@ -52,7 +59,9 @@ class TelegramChannel(Channel):
|
||||
|
||||
builder = ApplicationBuilder().token(self.config.bot_token)
|
||||
if self.config.proxy:
|
||||
builder = builder.proxy(self.config.proxy).get_updates_proxy(self.config.proxy)
|
||||
builder = builder.proxy(self.config.proxy).get_updates_proxy(
|
||||
self.config.proxy
|
||||
)
|
||||
self._app = builder.build()
|
||||
|
||||
# Accept text and media message types
|
||||
@@ -96,7 +105,8 @@ class TelegramChannel(Channel):
|
||||
"""Send typing action via Telegram Bot API."""
|
||||
if self._app:
|
||||
await self._app.bot.send_chat_action(
|
||||
chat_id=int(chat_id), action="typing",
|
||||
chat_id=int(chat_id),
|
||||
action="typing",
|
||||
)
|
||||
|
||||
# ── Send (template method overrides) ──────────────────────────
|
||||
@@ -106,7 +116,8 @@ class TelegramChannel(Channel):
|
||||
|
||||
async def _send(text):
|
||||
await self._app.bot.send_message(
|
||||
chat_id=int(chat_id), text=text,
|
||||
chat_id=int(chat_id),
|
||||
text=text,
|
||||
parse_mode="HTML" if text == formatted_text else None,
|
||||
reply_to_message_id=reply_id,
|
||||
)
|
||||
@@ -133,22 +144,29 @@ class TelegramChannel(Channel):
|
||||
for exts, (method, param) in self._MEDIA_SENDERS.items():
|
||||
if ext in exts:
|
||||
await getattr(self._app.bot, method)(
|
||||
chat_id=chat_id, caption=cap, **{param: file_path},
|
||||
chat_id=chat_id,
|
||||
caption=cap,
|
||||
**{param: file_path},
|
||||
)
|
||||
return True
|
||||
await self._app.bot.send_document(
|
||||
chat_id=chat_id, document=file_path, caption=cap,
|
||||
chat_id=chat_id,
|
||||
document=file_path,
|
||||
caption=cap,
|
||||
)
|
||||
return True
|
||||
|
||||
def _get_bot_identifier(self) -> str | None:
|
||||
return self._bot_username or None
|
||||
|
||||
async def _send_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
async def _send_ack_reaction(
|
||||
self, chat_id: str, message_id: str, emoji: str = "👀"
|
||||
) -> None:
|
||||
"""Send an acknowledgment reaction via Telegram."""
|
||||
if self._app:
|
||||
try:
|
||||
from telegram import ReactionTypeEmoji
|
||||
|
||||
await self._app.bot.set_message_reaction(
|
||||
chat_id=int(chat_id),
|
||||
message_id=int(message_id),
|
||||
@@ -157,7 +175,9 @@ class TelegramChannel(Channel):
|
||||
except Exception as e:
|
||||
logger.debug(f"Telegram ACK reaction failed: {e}")
|
||||
|
||||
async def _remove_ack_reaction(self, chat_id: str, message_id: str, emoji: str = "👀") -> None:
|
||||
async def _remove_ack_reaction(
|
||||
self, chat_id: str, message_id: str, emoji: str = "👀"
|
||||
) -> None:
|
||||
"""Remove the ack reaction by setting empty reaction list."""
|
||||
if self._app:
|
||||
try:
|
||||
@@ -222,12 +242,10 @@ class TelegramChannel(Channel):
|
||||
# Location is not a downloadable file — handle separately
|
||||
if message.location and not media_file:
|
||||
loc = message.location
|
||||
annotations.append(
|
||||
f"[位置] ({loc.latitude}, {loc.longitude})"
|
||||
)
|
||||
annotations.append(f"[位置] ({loc.latitude}, {loc.longitude})")
|
||||
|
||||
if media_file and self._app:
|
||||
file_size = getattr(media_file, 'file_size', 0) or 0
|
||||
file_size = getattr(media_file, "file_size", 0) or 0
|
||||
too_large = self._check_attachment_size(file_size, media_type)
|
||||
if too_large:
|
||||
annotations.append(too_large)
|
||||
@@ -238,47 +256,53 @@ class TelegramChannel(Channel):
|
||||
)
|
||||
ext = self._get_extension(
|
||||
media_type,
|
||||
getattr(media_file, 'mime_type', None),
|
||||
)
|
||||
file_path = self._media_path(
|
||||
f"{media_file.file_id[:16]}{ext}"
|
||||
getattr(media_file, "mime_type", None),
|
||||
)
|
||||
file_path = self._media_path(f"{media_file.file_id[:16]}{ext}")
|
||||
await file.download_to_drive(str(file_path))
|
||||
|
||||
media_paths.append(str(file_path))
|
||||
annotations.append(f"[{media_type}: {file_path}]")
|
||||
logger.debug(
|
||||
f"Downloaded {media_type} to {file_path}"
|
||||
)
|
||||
logger.debug(f"Downloaded {media_type} to {file_path}")
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to download media: {e}")
|
||||
annotations.append(
|
||||
f"[{media_type}: download failed]"
|
||||
)
|
||||
annotations.append(f"[{media_type}: download failed]")
|
||||
|
||||
text_content = "\n".join(content_parts) if content_parts else ""
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=user_id,
|
||||
chat_id=chat_id,
|
||||
text=text_content,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=message.date or datetime.now(),
|
||||
message_id=str(message.message_id),
|
||||
metadata={"chat_id": chat_id},
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
))
|
||||
await self._enqueue_raw(
|
||||
RawIncoming(
|
||||
sender_id=user_id,
|
||||
chat_id=chat_id,
|
||||
text=text_content,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=message.date or datetime.now(),
|
||||
message_id=str(message.message_id),
|
||||
metadata={"chat_id": chat_id},
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
)
|
||||
)
|
||||
|
||||
_MIME_TO_EXT = {
|
||||
"image/jpeg": ".jpg", "image/png": ".png", "image/gif": ".gif",
|
||||
"image/webp": ".webp", "audio/ogg": ".ogg", "audio/mpeg": ".mp3",
|
||||
"audio/mp4": ".m4a", "video/mp4": ".mp4", "video/quicktime": ".mov",
|
||||
"image/jpeg": ".jpg",
|
||||
"image/png": ".png",
|
||||
"image/gif": ".gif",
|
||||
"image/webp": ".webp",
|
||||
"audio/ogg": ".ogg",
|
||||
"audio/mpeg": ".mp3",
|
||||
"audio/mp4": ".m4a",
|
||||
"video/mp4": ".mp4",
|
||||
"video/quicktime": ".mov",
|
||||
}
|
||||
_TYPE_TO_EXT = {
|
||||
"image": ".jpg", "voice": ".ogg", "audio": ".mp3",
|
||||
"video": ".mp4", "file": "", "sticker": ".webp",
|
||||
"image": ".jpg",
|
||||
"voice": ".ogg",
|
||||
"audio": ".mp3",
|
||||
"video": ".mp4",
|
||||
"file": "",
|
||||
"sticker": ".webp",
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -5,7 +5,9 @@ import logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def validate_telegram_token(token: str, proxy: str | None = None) -> tuple[bool, str]:
|
||||
async def validate_telegram_token(
|
||||
token: str, proxy: str | None = None
|
||||
) -> tuple[bool, str]:
|
||||
"""Validate a Telegram bot token via the getMe API.
|
||||
|
||||
Returns:
|
||||
|
||||
@@ -42,6 +42,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
# ── Markdown → plain text (fallback for WeChat text messages) ────
|
||||
|
||||
|
||||
def _strip_markdown(text: str) -> str:
|
||||
"""Strip Markdown formatting for plain-text WeChat messages."""
|
||||
# Remove code blocks
|
||||
@@ -65,9 +66,11 @@ def _strip_markdown(text: str) -> str:
|
||||
|
||||
# ── Config dataclasses ───────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class WeComConfig(BaseChannelConfig):
|
||||
"""Configuration for WeCom (企业微信) backend."""
|
||||
|
||||
corp_id: str = ""
|
||||
agent_id: str = ""
|
||||
secret: str = ""
|
||||
@@ -79,6 +82,7 @@ class WeComConfig(BaseChannelConfig):
|
||||
@dataclass
|
||||
class WeChatMPConfig(BaseChannelConfig):
|
||||
"""Configuration for WeChat Official Account (公众号) backend."""
|
||||
|
||||
app_id: str = ""
|
||||
app_secret: str = ""
|
||||
token: str = ""
|
||||
@@ -88,6 +92,7 @@ class WeChatMPConfig(BaseChannelConfig):
|
||||
|
||||
# ── Unified WeChat Channel ───────────────────────────────────────
|
||||
|
||||
|
||||
class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
capabilities = WECHAT_CAPS
|
||||
"""Unified WeChat channel supporting WeCom and Official Account backends.
|
||||
@@ -119,7 +124,9 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
self._site = None
|
||||
self._http_client = None
|
||||
self._crypto = None # WeChatCrypto instance (optional)
|
||||
self._typing_message_ids: dict[str, list[str]] = {} # chat_id → [msgid, ...] for typing recall
|
||||
self._typing_message_ids: dict[
|
||||
str, list[str]
|
||||
] = {} # chat_id → [msgid, ...] for typing recall
|
||||
|
||||
# ── Lifecycle ─────────────────────────────────────────────────
|
||||
|
||||
@@ -143,6 +150,7 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
self._validate_config()
|
||||
|
||||
import httpx
|
||||
|
||||
self._http_client = httpx.AsyncClient(
|
||||
timeout=15,
|
||||
proxy=self._get_proxy(),
|
||||
@@ -151,6 +159,7 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
# Set up message encryption if configured
|
||||
if self.config.encoding_aes_key and self.config.token:
|
||||
from .crypto import WeChatCrypto
|
||||
|
||||
app_id = self._get_app_id()
|
||||
self._crypto = WeChatCrypto(
|
||||
token=self.config.token,
|
||||
@@ -169,7 +178,9 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
self._runner = web.AppRunner(app)
|
||||
await self._runner.setup()
|
||||
self._site = web.TCPSite(
|
||||
self._runner, "0.0.0.0", self.config.webhook_port,
|
||||
self._runner,
|
||||
"0.0.0.0",
|
||||
self.config.webhook_port,
|
||||
)
|
||||
await self._site.start()
|
||||
|
||||
@@ -272,7 +283,9 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
"""
|
||||
from aiohttp import web
|
||||
|
||||
signature = request.query.get("msg_signature") or request.query.get("signature", "")
|
||||
signature = request.query.get("msg_signature") or request.query.get(
|
||||
"signature", ""
|
||||
)
|
||||
timestamp = request.query.get("timestamp", "")
|
||||
nonce = request.query.get("nonce", "")
|
||||
echostr = request.query.get("echostr", "")
|
||||
@@ -363,7 +376,9 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
msg_id = xml_data.get("MsgId", "")
|
||||
create_time = xml_data.get("CreateTime", "")
|
||||
|
||||
logger.info(f"WeChat message received: type={msg_type}, from={from_user}, id={msg_id}, keys={list(xml_data.keys())}")
|
||||
logger.info(
|
||||
f"WeChat message received: type={msg_type}, from={from_user}, id={msg_id}, keys={list(xml_data.keys())}"
|
||||
)
|
||||
|
||||
if not from_user:
|
||||
return
|
||||
@@ -400,7 +415,8 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
media_id = xml_data.get("MediaId", "")
|
||||
if pic_url:
|
||||
local, ann = await self._download_attachment(
|
||||
pic_url, f"wechat_{msg_id}.jpg",
|
||||
pic_url,
|
||||
f"wechat_{msg_id}.jpg",
|
||||
)
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
@@ -408,7 +424,8 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
annotations.append(ann)
|
||||
elif media_id:
|
||||
local, ann = await self._download_wechat_media(
|
||||
media_id, f"wechat_image_{msg_id}",
|
||||
media_id,
|
||||
f"wechat_image_{msg_id}",
|
||||
)
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
@@ -420,7 +437,9 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
recognition = xml_data.get("Recognition", "")
|
||||
media_id = xml_data.get("MediaId", "")
|
||||
if media_id:
|
||||
local, ann = await self._download_wechat_media(media_id, f"wechat_voice_{msg_id}")
|
||||
local, ann = await self._download_wechat_media(
|
||||
media_id, f"wechat_voice_{msg_id}"
|
||||
)
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
if ann:
|
||||
@@ -433,7 +452,9 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
elif msg_type in ("video", "shortvideo"):
|
||||
media_id = xml_data.get("MediaId", "")
|
||||
if media_id:
|
||||
local, ann = await self._download_wechat_media(media_id, f"wechat_{msg_type}_{msg_id}")
|
||||
local, ann = await self._download_wechat_media(
|
||||
media_id, f"wechat_{msg_type}_{msg_id}"
|
||||
)
|
||||
if local:
|
||||
media_paths.append(local)
|
||||
if ann:
|
||||
@@ -447,11 +468,16 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
text = f"[位置] {label} ({lat}, {lon})"
|
||||
elif msg_type == "file":
|
||||
media_id = xml_data.get("MediaId", "")
|
||||
file_name = xml_data.get("FileName", "") or xml_data.get("Title", f"wechat_file_{msg_id}")
|
||||
logger.info(f"WeChat file message: name={file_name}, media_id={media_id!r}, keys={list(xml_data.keys())}")
|
||||
file_name = xml_data.get("FileName", "") or xml_data.get(
|
||||
"Title", f"wechat_file_{msg_id}"
|
||||
)
|
||||
logger.info(
|
||||
f"WeChat file message: name={file_name}, media_id={media_id!r}, keys={list(xml_data.keys())}"
|
||||
)
|
||||
if media_id:
|
||||
local, ann = await self._download_wechat_media(
|
||||
media_id, f"wechat_file_{msg_id}_{file_name}",
|
||||
media_id,
|
||||
f"wechat_file_{msg_id}_{file_name}",
|
||||
)
|
||||
logger.info(f"WeChat file download result: local={local}, ann={ann}")
|
||||
if local:
|
||||
@@ -489,28 +515,32 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
|
||||
# Parse timestamp
|
||||
try:
|
||||
timestamp = datetime.fromtimestamp(
|
||||
int(create_time)
|
||||
) if create_time else datetime.now()
|
||||
timestamp = (
|
||||
datetime.fromtimestamp(int(create_time))
|
||||
if create_time
|
||||
else datetime.now()
|
||||
)
|
||||
except (ValueError, TypeError, OSError):
|
||||
timestamp = datetime.now()
|
||||
|
||||
await self._enqueue_raw(RawIncoming(
|
||||
sender_id=from_user,
|
||||
chat_id=chat_id,
|
||||
text=text,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=timestamp,
|
||||
message_id=msg_id,
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"to_user": to_user,
|
||||
"backend": self._backend,
|
||||
},
|
||||
))
|
||||
await self._enqueue_raw(
|
||||
RawIncoming(
|
||||
sender_id=from_user,
|
||||
chat_id=chat_id,
|
||||
text=text,
|
||||
media_files=media_paths,
|
||||
content_annotations=annotations,
|
||||
timestamp=timestamp,
|
||||
message_id=msg_id,
|
||||
is_group=is_group,
|
||||
was_mentioned=was_mentioned,
|
||||
metadata={
|
||||
"chat_id": chat_id,
|
||||
"to_user": to_user,
|
||||
"backend": self._backend,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
# ── Send (template method overrides) ──────────────────────────
|
||||
|
||||
@@ -521,7 +551,12 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
return _strip_markdown(text)
|
||||
|
||||
async def _send_chunk(
|
||||
self, chat_id, formatted_text, raw_text, reply_to, metadata,
|
||||
self,
|
||||
chat_id,
|
||||
formatted_text,
|
||||
raw_text,
|
||||
reply_to,
|
||||
metadata,
|
||||
):
|
||||
token = await self._ensure_token()
|
||||
|
||||
@@ -548,13 +583,13 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
# ── WeCom send ────────────────────────────────────────────────
|
||||
|
||||
async def _wecom_send_text(
|
||||
self, token: str, user_id: str, text: str,
|
||||
self,
|
||||
token: str,
|
||||
user_id: str,
|
||||
text: str,
|
||||
) -> None:
|
||||
"""Send a text message via WeCom API."""
|
||||
url = (
|
||||
f"https://qyapi.weixin.qq.com/cgi-bin/message/send"
|
||||
f"?access_token={token}"
|
||||
)
|
||||
url = f"https://qyapi.weixin.qq.com/cgi-bin/message/send?access_token={token}"
|
||||
body = {
|
||||
"touser": user_id,
|
||||
"msgtype": "text",
|
||||
@@ -564,7 +599,10 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
await self._post_api(url, body)
|
||||
|
||||
async def _wecom_send_markdown(
|
||||
self, token: str, user_id: str, text: str,
|
||||
self,
|
||||
token: str,
|
||||
user_id: str,
|
||||
text: str,
|
||||
) -> None:
|
||||
"""Send a markdown message via WeCom API.
|
||||
|
||||
@@ -572,10 +610,7 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
(no code blocks, no images). Falls back to text if the
|
||||
message is too complex.
|
||||
"""
|
||||
url = (
|
||||
f"https://qyapi.weixin.qq.com/cgi-bin/message/send"
|
||||
f"?access_token={token}"
|
||||
)
|
||||
url = f"https://qyapi.weixin.qq.com/cgi-bin/message/send?access_token={token}"
|
||||
body = {
|
||||
"touser": user_id,
|
||||
"msgtype": "markdown",
|
||||
@@ -587,13 +622,13 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
# ── WeCom group send ────────────────────────────────────────────
|
||||
|
||||
async def _wecom_send_group_text(
|
||||
self, token: str, chatid: str, text: str,
|
||||
self,
|
||||
token: str,
|
||||
chatid: str,
|
||||
text: str,
|
||||
) -> None:
|
||||
"""Send a text message to a WeCom group chat."""
|
||||
url = (
|
||||
f"https://qyapi.weixin.qq.com/cgi-bin/appchat/send"
|
||||
f"?access_token={token}"
|
||||
)
|
||||
url = f"https://qyapi.weixin.qq.com/cgi-bin/appchat/send?access_token={token}"
|
||||
body = {
|
||||
"chatid": chatid,
|
||||
"msgtype": "text",
|
||||
@@ -602,13 +637,13 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
await self._post_api(url, body)
|
||||
|
||||
async def _wecom_send_group_markdown(
|
||||
self, token: str, chatid: str, text: str,
|
||||
self,
|
||||
token: str,
|
||||
chatid: str,
|
||||
text: str,
|
||||
) -> None:
|
||||
"""Send a markdown message to a WeCom group chat."""
|
||||
url = (
|
||||
f"https://qyapi.weixin.qq.com/cgi-bin/appchat/send"
|
||||
f"?access_token={token}"
|
||||
)
|
||||
url = f"https://qyapi.weixin.qq.com/cgi-bin/appchat/send?access_token={token}"
|
||||
body = {
|
||||
"chatid": chatid,
|
||||
"msgtype": "markdown",
|
||||
@@ -619,7 +654,10 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
# ── MP send ───────────────────────────────────────────────────
|
||||
|
||||
async def _mp_send_text(
|
||||
self, token: str, openid: str, text: str,
|
||||
self,
|
||||
token: str,
|
||||
openid: str,
|
||||
text: str,
|
||||
) -> None:
|
||||
"""Send a text message via WeChat MP customer service API."""
|
||||
url = (
|
||||
@@ -689,7 +727,9 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
# MP doesn't support file via customer service API;
|
||||
# send caption as text instead
|
||||
if caption:
|
||||
await self._mp_send_text(token, chat_id, f"[文件] {path.name}\n{caption}")
|
||||
await self._mp_send_text(
|
||||
token, chat_id, f"[文件] {path.name}\n{caption}"
|
||||
)
|
||||
return True
|
||||
body = {
|
||||
"touser": chat_id,
|
||||
@@ -712,7 +752,9 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
return True
|
||||
|
||||
async def _upload_media(
|
||||
self, token: str, file_path: str,
|
||||
self,
|
||||
token: str,
|
||||
file_path: str,
|
||||
) -> str | None:
|
||||
"""Upload a media file and return the media_id."""
|
||||
path = Path(file_path)
|
||||
@@ -739,9 +781,7 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
)
|
||||
data = resp.json()
|
||||
if data.get("errcode", 0) != 0 and "media_id" not in data:
|
||||
logger.error(
|
||||
f"WeChat media upload failed: {data.get('errmsg')}"
|
||||
)
|
||||
logger.error(f"WeChat media upload failed: {data.get('errmsg')}")
|
||||
return None
|
||||
return data.get("media_id")
|
||||
except Exception as e:
|
||||
@@ -751,7 +791,9 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
# ── Media download helper ────────────────────────────────────
|
||||
|
||||
async def _download_wechat_media(
|
||||
self, media_id: str, filename: str,
|
||||
self,
|
||||
media_id: str,
|
||||
filename: str,
|
||||
) -> tuple[str | None, str | None]:
|
||||
"""Download media by media_id via WeChat/WeCom media API."""
|
||||
token = await self._ensure_token()
|
||||
@@ -790,14 +832,11 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
data = resp.json()
|
||||
if data.get("errcode", 0) != 0:
|
||||
raise RuntimeError(
|
||||
f"WeChat API error after retry: "
|
||||
f"{data.get('errmsg')}"
|
||||
f"WeChat API error after retry: {data.get('errmsg')}"
|
||||
)
|
||||
return data
|
||||
else:
|
||||
raise RuntimeError(
|
||||
f"WeChat API error ({errcode}): {errmsg}"
|
||||
)
|
||||
raise RuntimeError(f"WeChat API error ({errcode}): {errmsg}")
|
||||
|
||||
return data
|
||||
|
||||
@@ -861,5 +900,3 @@ class WeChatChannel(Channel, WebhookMixin, TokenMixin):
|
||||
except Exception:
|
||||
pass
|
||||
await super().stop_typing(chat_id)
|
||||
|
||||
|
||||
|
||||
@@ -21,6 +21,7 @@ import xml.etree.ElementTree as ET
|
||||
# (but we'll use a pure-Python fallback if not available)
|
||||
try:
|
||||
from Crypto.Cipher import AES
|
||||
|
||||
_HAS_PYCRYPTO = True
|
||||
except ImportError:
|
||||
_HAS_PYCRYPTO = False
|
||||
@@ -50,9 +51,8 @@ def _aes_decrypt(key: bytes, iv: bytes, ciphertext: bytes) -> bytes:
|
||||
# We'll try pyaes as a fallback
|
||||
try:
|
||||
import pyaes
|
||||
decrypter = pyaes.Decrypter(
|
||||
pyaes.AESModeOfOperationCBC(key, iv=iv)
|
||||
)
|
||||
|
||||
decrypter = pyaes.Decrypter(pyaes.AESModeOfOperationCBC(key, iv=iv))
|
||||
decrypted = decrypter.feed(ciphertext)
|
||||
decrypted += decrypter.feed()
|
||||
return decrypted
|
||||
@@ -71,9 +71,8 @@ def _aes_encrypt(key: bytes, iv: bytes, plaintext: bytes) -> bytes:
|
||||
else:
|
||||
try:
|
||||
import pyaes
|
||||
encrypter = pyaes.Encrypter(
|
||||
pyaes.AESModeOfOperationCBC(key, iv=iv)
|
||||
)
|
||||
|
||||
encrypter = pyaes.Encrypter(pyaes.AESModeOfOperationCBC(key, iv=iv))
|
||||
encrypted = encrypter.feed(plaintext)
|
||||
encrypted += encrypter.feed()
|
||||
return encrypted
|
||||
@@ -106,7 +105,10 @@ class WeChatCrypto:
|
||||
self.iv = self.aes_key[:16]
|
||||
|
||||
def verify_signature(
|
||||
self, signature: str, timestamp: str, nonce: str,
|
||||
self,
|
||||
signature: str,
|
||||
timestamp: str,
|
||||
nonce: str,
|
||||
encrypt: str = "",
|
||||
) -> bool:
|
||||
"""Verify the callback signature.
|
||||
@@ -129,8 +131,8 @@ class WeChatCrypto:
|
||||
# plaintext layout:
|
||||
# 16 bytes random + 4 bytes msg_len (big-endian) + msg + app_id
|
||||
msg_len = struct.unpack("!I", plaintext[16:20])[0]
|
||||
msg = plaintext[20:20 + msg_len].decode("utf-8")
|
||||
from_app_id = plaintext[20 + msg_len:].decode("utf-8")
|
||||
msg = plaintext[20 : 20 + msg_len].decode("utf-8")
|
||||
from_app_id = plaintext[20 + msg_len :].decode("utf-8")
|
||||
return msg, from_app_id
|
||||
|
||||
def encrypt(self, reply_msg: str) -> str:
|
||||
@@ -143,6 +145,7 @@ class WeChatCrypto:
|
||||
|
||||
# Random 16 bytes + msg_len (4 bytes big-endian) + msg + app_id
|
||||
import os
|
||||
|
||||
random_bytes = os.urandom(16)
|
||||
msg_len = struct.pack("!I", len(msg_bytes))
|
||||
plaintext = random_bytes + msg_len + msg_bytes + app_id_bytes
|
||||
@@ -152,7 +155,10 @@ class WeChatCrypto:
|
||||
return base64.b64encode(ciphertext).decode("utf-8")
|
||||
|
||||
def generate_signature(
|
||||
self, encrypt: str, timestamp: str, nonce: str,
|
||||
self,
|
||||
encrypt: str,
|
||||
timestamp: str,
|
||||
nonce: str,
|
||||
) -> str:
|
||||
"""Generate the msg_signature for an encrypted reply."""
|
||||
parts = sorted([self.token, timestamp, nonce, encrypt])
|
||||
|
||||
@@ -50,6 +50,7 @@ class VerifyServer:
|
||||
|
||||
if encoding_aes_key and token and app_id:
|
||||
from .crypto import WeChatCrypto
|
||||
|
||||
self._crypto = WeChatCrypto(
|
||||
token=token,
|
||||
encoding_aes_key=encoding_aes_key,
|
||||
@@ -110,9 +111,8 @@ class VerifyServer:
|
||||
"""
|
||||
from aiohttp import web
|
||||
|
||||
signature = (
|
||||
request.query.get("msg_signature")
|
||||
or request.query.get("signature", "")
|
||||
signature = request.query.get("msg_signature") or request.query.get(
|
||||
"signature", ""
|
||||
)
|
||||
timestamp = request.query.get("timestamp", "")
|
||||
nonce = request.query.get("nonce", "")
|
||||
@@ -130,7 +130,10 @@ class VerifyServer:
|
||||
# Attempt 1: Encrypted mode with full signature verification
|
||||
if self._crypto and request.query.get("msg_signature"):
|
||||
sig_ok = self._crypto.verify_signature(
|
||||
signature, timestamp, nonce, echostr,
|
||||
signature,
|
||||
timestamp,
|
||||
nonce,
|
||||
echostr,
|
||||
)
|
||||
if sig_ok:
|
||||
try:
|
||||
@@ -172,4 +175,5 @@ class VerifyServer:
|
||||
async def _handle_post(self, request) -> web.Response:
|
||||
"""Handle POST — just acknowledge during verification phase."""
|
||||
from aiohttp import web
|
||||
|
||||
return web.Response(text="success")
|
||||
|
||||
@@ -30,7 +30,9 @@ def main():
|
||||
import warnings
|
||||
|
||||
warnings.filterwarnings("ignore", message=".*not known to support tools.*")
|
||||
warnings.filterwarnings("ignore", message=".*type is unknown and inference may fail.*")
|
||||
warnings.filterwarnings(
|
||||
"ignore", message=".*type is unknown and inference may fail.*"
|
||||
)
|
||||
from .commands import _configure_logging
|
||||
|
||||
_configure_logging()
|
||||
|
||||
@@ -9,7 +9,9 @@ app = typer.Typer(
|
||||
)
|
||||
|
||||
# Config subcommand group
|
||||
config_app = typer.Typer(help="Configuration management commands", invoke_without_command=True)
|
||||
config_app = typer.Typer(
|
||||
help="Configuration management commands", invoke_without_command=True
|
||||
)
|
||||
app.add_typer(config_app, name="config")
|
||||
|
||||
# MCP subcommand group
|
||||
|
||||
@@ -14,8 +14,12 @@ def _shorten_path(path: str) -> str:
|
||||
try:
|
||||
cwd = os.getcwd()
|
||||
if path.startswith(cwd):
|
||||
rel = path[len(cwd):].lstrip(os.sep)
|
||||
return os.path.join(os.path.basename(cwd), rel) if rel else os.path.basename(cwd)
|
||||
rel = path[len(cwd) :].lstrip(os.sep)
|
||||
return (
|
||||
os.path.join(os.path.basename(cwd), rel)
|
||||
if rel
|
||||
else os.path.basename(cwd)
|
||||
)
|
||||
return path
|
||||
except Exception:
|
||||
return path
|
||||
@@ -25,6 +29,7 @@ def _deduplicate_run_name(name: str, runs_dir: Path | None = None) -> str:
|
||||
"""Return *name* if available, otherwise *name_1*, *name_2*, etc."""
|
||||
if runs_dir is None:
|
||||
from ..paths import RUNS_DIR
|
||||
|
||||
runs_dir = RUNS_DIR
|
||||
if not (runs_dir / name).exists():
|
||||
return name
|
||||
@@ -44,6 +49,7 @@ def _create_session_workspace(name: str | None = None) -> str:
|
||||
"""
|
||||
if name:
|
||||
from ..paths import RUNS_DIR
|
||||
|
||||
session_id = _deduplicate_run_name(name, RUNS_DIR)
|
||||
else:
|
||||
session_id = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
@@ -63,4 +69,7 @@ def _load_agent(workspace_dir: str | None = None, checkpointer=None, config=None
|
||||
``create_cli_agent`` to avoid double config loading.
|
||||
"""
|
||||
from ..EvoScientist import create_cli_agent
|
||||
return create_cli_agent(workspace_dir=workspace_dir, checkpointer=checkpointer, config=config)
|
||||
|
||||
return create_cli_agent(
|
||||
workspace_dir=workspace_dir, checkpointer=checkpointer, config=config
|
||||
)
|
||||
|
||||
+50
-33
@@ -31,9 +31,11 @@ _channel_logger = logging.getLogger(__name__)
|
||||
# Queue bridge: bus thread ⇄ main CLI thread
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class ChannelMessage:
|
||||
"""A message from a channel, enqueued for the main CLI thread."""
|
||||
|
||||
msg_id: str
|
||||
content: str
|
||||
sender: str
|
||||
@@ -41,7 +43,7 @@ class ChannelMessage:
|
||||
metadata: dict | None = None
|
||||
# Filled by the bus consumer so the main thread can send callbacks
|
||||
channel_ref: Any = None # Channel instance (for thinking / todo / file)
|
||||
bus_ref: Any = None # MessageBus (for publishing outbound)
|
||||
bus_ref: Any = None # MessageBus (for publishing outbound)
|
||||
chat_id: str = ""
|
||||
message_id: str | None = None
|
||||
|
||||
@@ -91,7 +93,9 @@ _hitl_lock = threading.Lock()
|
||||
_hitl_auto_approve: set[str] = set() # "channel:chat_id" keys with auto-approve
|
||||
|
||||
_HITL_APPROVAL_TIMEOUT = 120.0 # seconds to wait for HITL approval reply
|
||||
_ASK_USER_TIMEOUT = 300.0 # seconds to wait for ask_user reply (longer for thinking time)
|
||||
_ASK_USER_TIMEOUT = (
|
||||
300.0 # seconds to wait for ask_user reply (longer for thinking time)
|
||||
)
|
||||
|
||||
|
||||
def _register_hitl_wait(channel_type: str, chat_id: str) -> threading.Event:
|
||||
@@ -152,12 +156,14 @@ def channel_ask_user_prompt(
|
||||
def _send(content: str) -> bool:
|
||||
try:
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
msg.bus_ref.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel_type,
|
||||
chat_id=msg.chat_id,
|
||||
content=content,
|
||||
metadata=msg.metadata,
|
||||
)),
|
||||
msg.bus_ref.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel_type,
|
||||
chat_id=msg.chat_id,
|
||||
content=content,
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
),
|
||||
bus_loop,
|
||||
).result(timeout=15)
|
||||
return True
|
||||
@@ -192,7 +198,9 @@ def channel_ask_user_prompt(
|
||||
lines.append(f" {letter}. {label}")
|
||||
other_letter = chr(ord("A") + len(choices))
|
||||
lines.append(f" {other_letter}. Other")
|
||||
lines.append(f"\nReply with a letter ({'/'.join(chr(ord('A') + k) for k in range(len(choices) + 1))}), or 'cancel'.")
|
||||
lines.append(
|
||||
f"\nReply with a letter ({'/'.join(chr(ord('A') + k) for k in range(len(choices) + 1))}), or 'cancel'."
|
||||
)
|
||||
else:
|
||||
skip_hint = " Leave empty to skip." if not required else ""
|
||||
lines.append(f"\nReply with your answer, or 'cancel'.{skip_hint}")
|
||||
@@ -275,12 +283,14 @@ def channel_hitl_prompt(
|
||||
"""Send a message to the channel user. Returns True on success."""
|
||||
try:
|
||||
asyncio.run_coroutine_threadsafe(
|
||||
msg.bus_ref.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel_type,
|
||||
chat_id=msg.chat_id,
|
||||
content=content,
|
||||
metadata=msg.metadata,
|
||||
)),
|
||||
msg.bus_ref.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel_type,
|
||||
chat_id=msg.chat_id,
|
||||
content=content,
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
),
|
||||
bus_loop,
|
||||
).result(timeout=15)
|
||||
return True
|
||||
@@ -311,7 +321,8 @@ def channel_hitl_prompt(
|
||||
return [{"type": "approve"} for _ in action_requests]
|
||||
|
||||
feedback = (
|
||||
"Action rejected." if decision == "reject"
|
||||
"Action rejected."
|
||||
if decision == "reject"
|
||||
else "Unrecognized reply. Action rejected."
|
||||
)
|
||||
_send(feedback)
|
||||
@@ -325,7 +336,7 @@ def channel_hitl_prompt(
|
||||
_manager: Optional[Any] = None # ChannelManager
|
||||
_bus_loop: Optional[asyncio.AbstractEventLoop] = None
|
||||
_bus_thread: Optional[threading.Thread] = None
|
||||
_cli_agent: Any = None # shared agent reference (same as CLI)
|
||||
_cli_agent: Any = None # shared agent reference (same as CLI)
|
||||
_cli_thread_id: Optional[str] = None # shared thread_id (same conversation)
|
||||
|
||||
|
||||
@@ -353,7 +364,8 @@ def _channels_stop(channel_type: str | None = None) -> None:
|
||||
if _bus_loop and _manager:
|
||||
try:
|
||||
future = asyncio.run_coroutine_threadsafe(
|
||||
_manager.stop_all(), _bus_loop,
|
||||
_manager.stop_all(),
|
||||
_bus_loop,
|
||||
)
|
||||
future.result(timeout=10)
|
||||
except Exception as e:
|
||||
@@ -373,7 +385,8 @@ def _channels_stop(channel_type: str | None = None) -> None:
|
||||
if _manager and _bus_loop:
|
||||
try:
|
||||
future = asyncio.run_coroutine_threadsafe(
|
||||
_manager.remove_channel(channel_type), _bus_loop,
|
||||
_manager.remove_channel(channel_type),
|
||||
_bus_loop,
|
||||
)
|
||||
future.result(timeout=5)
|
||||
except Exception as e:
|
||||
@@ -419,9 +432,7 @@ def _start_channels_bus_mode(
|
||||
_bus_loop = loop
|
||||
|
||||
async def _run():
|
||||
consumer = asyncio.create_task(
|
||||
_bus_inbound_consumer(mgr.bus, mgr)
|
||||
)
|
||||
consumer = asyncio.create_task(_bus_inbound_consumer(mgr.bus, mgr))
|
||||
try:
|
||||
await mgr.start_all()
|
||||
finally:
|
||||
@@ -509,8 +520,7 @@ async def _handle_bus_message(bus, manager, msg) -> None:
|
||||
from ..channels.bus.events import OutboundMessage
|
||||
|
||||
_channel_logger.info(
|
||||
f"[bus] Received from {msg.channel}:{msg.sender_id}: "
|
||||
f"{msg.content[:60]}..."
|
||||
f"[bus] Received from {msg.channel}:{msg.sender_id}: {msg.content[:60]}..."
|
||||
)
|
||||
manager.record_message(msg.channel, "received")
|
||||
|
||||
@@ -543,13 +553,15 @@ async def _handle_bus_message(bus, manager, msg) -> None:
|
||||
|
||||
# Publish the response back through the bus → channel
|
||||
try:
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content=response,
|
||||
reply_to=msg.message_id or None,
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
await bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content=response,
|
||||
reply_to=msg.message_id or None,
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
manager.record_message(msg.channel, "sent")
|
||||
except Exception as e:
|
||||
_channel_logger.error(f"[bus] Outbound error: {e}")
|
||||
@@ -581,7 +593,9 @@ def _print_channel_panel(channels: list[tuple[str, bool, str]]) -> None:
|
||||
|
||||
body = Text("\n").join(lines)
|
||||
border = "green" if all_ok else "yellow"
|
||||
console.print(Panel(body, title="[bold]Channels[/bold]", border_style=border, expand=False))
|
||||
console.print(
|
||||
Panel(body, title="[bold]Channels[/bold]", border_style=border, expand=False)
|
||||
)
|
||||
console.print()
|
||||
|
||||
|
||||
@@ -602,6 +616,7 @@ def _cmd_channel(
|
||||
global _cli_agent, _cli_thread_id
|
||||
|
||||
from ..config import load_config
|
||||
|
||||
app_config = load_config()
|
||||
|
||||
channel_type = args.strip().lower() if args and args.strip() else ""
|
||||
@@ -634,7 +649,9 @@ def _cmd_channel(
|
||||
channel_type = app_config.channel_enabled
|
||||
if not channel_type:
|
||||
console.print("[yellow]No channel configured.[/yellow]")
|
||||
console.print("[dim]Run[/dim] evosci onboard [dim]or specify:[/dim] /channel telegram\n")
|
||||
console.print(
|
||||
"[dim]Run[/dim] evosci onboard [dim]or specify:[/dim] /channel telegram\n"
|
||||
)
|
||||
return
|
||||
|
||||
requested = [t.strip() for t in channel_type.split(",") if t.strip()]
|
||||
|
||||
@@ -69,7 +69,8 @@ def copy_selection_to_clipboard(app: App) -> None:
|
||||
except (AttributeError, TypeError, ValueError, IndexError) as exc:
|
||||
logger.debug(
|
||||
"Failed to get selection from %s: %s",
|
||||
type(widget).__name__, exc,
|
||||
type(widget).__name__,
|
||||
exc,
|
||||
)
|
||||
continue
|
||||
if not result:
|
||||
@@ -88,6 +89,7 @@ def copy_selection_to_clipboard(app: App) -> None:
|
||||
|
||||
try:
|
||||
import pyperclip
|
||||
|
||||
copy_methods.insert(0, pyperclip.copy)
|
||||
except ImportError:
|
||||
pass
|
||||
@@ -104,7 +106,9 @@ def copy_selection_to_clipboard(app: App) -> None:
|
||||
markup=False,
|
||||
)
|
||||
except (OSError, RuntimeError, TypeError) as exc:
|
||||
logger.debug("Clipboard method %s failed: %s", getattr(fn, "__name__", repr(fn)), exc)
|
||||
logger.debug(
|
||||
"Clipboard method %s failed: %s", getattr(fn, "__name__", repr(fn)), exc
|
||||
)
|
||||
continue
|
||||
else:
|
||||
return
|
||||
|
||||
+186
-52
@@ -17,7 +17,12 @@ from ..stream.display import console
|
||||
from ..paths import ensure_dirs, set_workspace_root
|
||||
from ._app import app, config_app, mcp_app, channel_app
|
||||
from ._constants import build_metadata
|
||||
from .agent import _deduplicate_run_name, _create_session_workspace, _load_agent, _shorten_path
|
||||
from .agent import (
|
||||
_deduplicate_run_name,
|
||||
_create_session_workspace,
|
||||
_load_agent,
|
||||
_shorten_path,
|
||||
)
|
||||
from .channel import (
|
||||
ChannelMessage,
|
||||
_channels_stop,
|
||||
@@ -42,12 +47,11 @@ from .interactive import cmd_interactive, cmd_run
|
||||
# Onboard command
|
||||
# =============================================================================
|
||||
|
||||
|
||||
@app.command()
|
||||
def onboard(
|
||||
skip_validation: bool = typer.Option(
|
||||
False,
|
||||
"--skip-validation",
|
||||
help="Skip API key validation during setup"
|
||||
False, "--skip-validation", help="Skip API key validation during setup"
|
||||
),
|
||||
):
|
||||
"""Interactive setup wizard for EvoScientist
|
||||
@@ -56,6 +60,7 @@ def onboard(
|
||||
workspace settings, and agent parameters.
|
||||
"""
|
||||
from ..config import run_onboard
|
||||
|
||||
run_onboard(skip_validation=skip_validation)
|
||||
|
||||
|
||||
@@ -63,6 +68,7 @@ def onboard(
|
||||
# Channel setup command
|
||||
# =============================================================================
|
||||
|
||||
|
||||
@channel_app.command("setup")
|
||||
def channel_setup():
|
||||
"""Interactive channel configuration wizard.
|
||||
@@ -71,6 +77,7 @@ def channel_setup():
|
||||
(Telegram, Discord, or iMessage).
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
try:
|
||||
asyncio.get_event_loop()
|
||||
except RuntimeError:
|
||||
@@ -111,10 +118,14 @@ class CompactResult:
|
||||
"""
|
||||
|
||||
__slots__ = (
|
||||
"status", "message",
|
||||
"messages_compacted", "messages_kept",
|
||||
"tokens_before", "tokens_after",
|
||||
"tokens_summarized", "tokens_summary",
|
||||
"status",
|
||||
"message",
|
||||
"messages_compacted",
|
||||
"messages_kept",
|
||||
"tokens_before",
|
||||
"tokens_after",
|
||||
"tokens_summarized",
|
||||
"tokens_summary",
|
||||
"pct_decrease",
|
||||
)
|
||||
|
||||
@@ -164,7 +175,12 @@ def render_compact_result(result: CompactResult): # -> rich.text.Text
|
||||
output.append(" tokens, within retention budget", style="dim")
|
||||
elif result.message:
|
||||
# Extract reason from message (e.g. "no messages")
|
||||
output.append(f" — {result.message.split('—')[-1].strip()}" if "—" in result.message else "", style="dim")
|
||||
output.append(
|
||||
f" — {result.message.split('—')[-1].strip()}"
|
||||
if "—" in result.message
|
||||
else "",
|
||||
style="dim",
|
||||
)
|
||||
return output
|
||||
|
||||
if result.status == "error":
|
||||
@@ -222,7 +238,9 @@ async def compact_conversation(agent: Any, thread_id: str | None) -> CompactResu
|
||||
|
||||
messages = state_snapshot.values.get("messages", [])
|
||||
if not messages:
|
||||
return CompactResult("noop", "Nothing to compact — no messages in conversation.")
|
||||
return CompactResult(
|
||||
"noop", "Nothing to compact — no messages in conversation."
|
||||
)
|
||||
|
||||
from ..EvoScientist import _ensure_chat_model, _get_default_backend
|
||||
from deepagents.middleware.summarization import (
|
||||
@@ -234,7 +252,9 @@ async def compact_conversation(agent: Any, thread_id: str | None) -> CompactResu
|
||||
try:
|
||||
model = _ensure_chat_model()
|
||||
except Exception as exc:
|
||||
return CompactResult("error", f"Compaction requires a working model configuration: {exc}")
|
||||
return CompactResult(
|
||||
"error", f"Compaction requires a working model configuration: {exc}"
|
||||
)
|
||||
|
||||
backend = _get_default_backend()
|
||||
|
||||
@@ -364,8 +384,7 @@ def _serve_process_message(
|
||||
from .channel import _bus_loop
|
||||
|
||||
console.print(
|
||||
f"[dim][{msg.channel_type}] {msg.sender}: "
|
||||
f"{escape(msg.content[:80])}[/dim]"
|
||||
f"[dim][{msg.channel_type}] {msg.sender}: {escape(msg.content[:80])}[/dim]"
|
||||
)
|
||||
|
||||
# -- channel callback helpers (same pattern as interactive.py) --
|
||||
@@ -450,12 +469,25 @@ def _serve_process_message(
|
||||
# Serve command (headless mode)
|
||||
# =============================================================================
|
||||
|
||||
|
||||
@app.command()
|
||||
def serve(
|
||||
no_thinking: bool = typer.Option(False, "--no-thinking", help="Disable thinking relay to channels"),
|
||||
workdir: Optional[str] = typer.Option(None, "--workdir", help="Override workspace directory"),
|
||||
auto_approve: bool = typer.Option(False, "--auto-approve", help="Auto-approve all tool executions without prompting"),
|
||||
ask_user: bool = typer.Option(False, "--ask-user", help="Enable agent to ask clarifying questions about your research preferences"),
|
||||
no_thinking: bool = typer.Option(
|
||||
False, "--no-thinking", help="Disable thinking relay to channels"
|
||||
),
|
||||
workdir: Optional[str] = typer.Option(
|
||||
None, "--workdir", help="Override workspace directory"
|
||||
),
|
||||
auto_approve: bool = typer.Option(
|
||||
False,
|
||||
"--auto-approve",
|
||||
help="Auto-approve all tool executions without prompting",
|
||||
),
|
||||
ask_user: bool = typer.Option(
|
||||
False,
|
||||
"--ask-user",
|
||||
help="Enable agent to ask clarifying questions about your research preferences",
|
||||
),
|
||||
):
|
||||
"""Run EvoScientist in headless mode -- channels only, no interactive prompt.
|
||||
|
||||
@@ -463,9 +495,11 @@ def serve(
|
||||
Press Ctrl+C to shut down.
|
||||
"""
|
||||
import nest_asyncio # type: ignore[import-untyped]
|
||||
|
||||
nest_asyncio.apply()
|
||||
|
||||
from dotenv import load_dotenv, find_dotenv # type: ignore[import-untyped]
|
||||
|
||||
load_dotenv(find_dotenv(), override=True)
|
||||
|
||||
from ..config import get_effective_config, apply_config_to_env
|
||||
@@ -483,9 +517,11 @@ def serve(
|
||||
if config.provider == "anthropic" and config.anthropic_auth_mode == "oauth":
|
||||
try:
|
||||
from ..ccproxy_manager import maybe_start_ccproxy, stop_ccproxy
|
||||
|
||||
_ccproxy_proc_serve = maybe_start_ccproxy(config)
|
||||
if _ccproxy_proc_serve:
|
||||
import atexit
|
||||
|
||||
atexit.register(stop_ccproxy, _ccproxy_proc_serve)
|
||||
except RuntimeError as exc:
|
||||
console.print(f"[red]{exc}[/red]")
|
||||
@@ -510,6 +546,7 @@ def serve(
|
||||
console.print("[dim]Loading agent...[/dim]")
|
||||
agent = _load_agent(workspace_dir=ws, config=config)
|
||||
from ..sessions import generate_thread_id
|
||||
|
||||
tid = generate_thread_id()
|
||||
|
||||
_start_channels_bus_mode(
|
||||
@@ -549,6 +586,7 @@ def serve(
|
||||
# Config commands
|
||||
# =============================================================================
|
||||
|
||||
|
||||
@config_app.callback(invoke_without_command=True)
|
||||
def config_callback(ctx: typer.Context):
|
||||
"""Configuration management commands"""
|
||||
@@ -656,6 +694,7 @@ def config_path():
|
||||
# MCP commands
|
||||
# =============================================================================
|
||||
|
||||
|
||||
@mcp_app.callback(invoke_without_command=True)
|
||||
def mcp_callback(ctx: typer.Context):
|
||||
"""MCP server management commands"""
|
||||
@@ -682,7 +721,9 @@ def mcp_config(
|
||||
"""
|
||||
status = _show_mcp_config(name or "", show_blank_line=False)
|
||||
if status == "empty":
|
||||
console.print("[dim]Add one with:[/dim] EvoSci mcp add <name> <transport> <command-or-url> [args...]")
|
||||
console.print(
|
||||
"[dim]Add one with:[/dim] EvoSci mcp add <name> <transport> <command-or-url> [args...]"
|
||||
)
|
||||
return
|
||||
if status == "missing":
|
||||
raise typer.Exit(1)
|
||||
@@ -692,13 +733,30 @@ def mcp_config(
|
||||
def mcp_add(
|
||||
name: str = typer.Argument(..., help="Server name"),
|
||||
target: str = typer.Argument(..., help="Command (stdio) or URL (http/sse)"),
|
||||
args: Optional[list[str]] = typer.Argument(None, help="Extra args for stdio command"),
|
||||
transport: Optional[str] = typer.Option(None, "--transport", "-T", help="Transport type (default: auto-detect)"),
|
||||
tools: Optional[str] = typer.Option(None, "--tools", "-t", help="Comma-separated tool allowlist (supports wildcards: *_exa, read_*)"),
|
||||
expose_to: Optional[str] = typer.Option(None, "--expose-to", "-e", help="Comma-separated target agents"),
|
||||
header: Optional[list[str]] = typer.Option(None, "--header", "-H", help="HTTP header as Key:Value (repeatable)"),
|
||||
env: Optional[list[str]] = typer.Option(None, "--env", help="Env var as KEY=VALUE for stdio (repeatable)"),
|
||||
env_ref: Optional[list[str]] = typer.Option(None, "--env-ref", help="Env var name as ${NAME} runtime ref (repeatable)"),
|
||||
args: Optional[list[str]] = typer.Argument(
|
||||
None, help="Extra args for stdio command"
|
||||
),
|
||||
transport: Optional[str] = typer.Option(
|
||||
None, "--transport", "-T", help="Transport type (default: auto-detect)"
|
||||
),
|
||||
tools: Optional[str] = typer.Option(
|
||||
None,
|
||||
"--tools",
|
||||
"-t",
|
||||
help="Comma-separated tool allowlist (supports wildcards: *_exa, read_*)",
|
||||
),
|
||||
expose_to: Optional[str] = typer.Option(
|
||||
None, "--expose-to", "-e", help="Comma-separated target agents"
|
||||
),
|
||||
header: Optional[list[str]] = typer.Option(
|
||||
None, "--header", "-H", help="HTTP header as Key:Value (repeatable)"
|
||||
),
|
||||
env: Optional[list[str]] = typer.Option(
|
||||
None, "--env", help="Env var as KEY=VALUE for stdio (repeatable)"
|
||||
),
|
||||
env_ref: Optional[list[str]] = typer.Option(
|
||||
None, "--env-ref", help="Env var name as ${NAME} runtime ref (repeatable)"
|
||||
),
|
||||
):
|
||||
"""Add an MCP server to user config
|
||||
|
||||
@@ -716,11 +774,11 @@ def mcp_add(
|
||||
|
||||
# Merge env and env_ref into a single dict
|
||||
env_dict: dict[str, str] = {}
|
||||
for e in (env or []):
|
||||
for e in env or []:
|
||||
if "=" in e:
|
||||
k, v = e.split("=", 1)
|
||||
env_dict[k.strip()] = v.strip()
|
||||
for ref in (env_ref or []):
|
||||
for ref in env_ref or []:
|
||||
env_dict[ref] = "${" + ref + "}"
|
||||
|
||||
kwargs = build_mcp_add_kwargs(
|
||||
@@ -729,8 +787,16 @@ def mcp_add(
|
||||
extra_args=list(args) if args else None,
|
||||
transport=transport,
|
||||
tools=[t.strip() for t in tools.split(",") if t.strip()] if tools else None,
|
||||
expose_to=[a.strip() for a in expose_to.split(",") if a.strip()] if expose_to else None,
|
||||
headers={k.strip(): v.strip() for h in (header or []) for k, v in [h.split(":", 1)] if ":" in h} or None,
|
||||
expose_to=[a.strip() for a in expose_to.split(",") if a.strip()]
|
||||
if expose_to
|
||||
else None,
|
||||
headers={
|
||||
k.strip(): v.strip()
|
||||
for h in (header or [])
|
||||
for k, v in [h.split(":", 1)]
|
||||
if ":" in h
|
||||
}
|
||||
or None,
|
||||
env=env_dict or None,
|
||||
)
|
||||
|
||||
@@ -741,13 +807,33 @@ def mcp_add(
|
||||
@mcp_app.command("edit")
|
||||
def mcp_edit(
|
||||
name: str = typer.Argument(..., help="Server name to edit"),
|
||||
transport: Optional[str] = typer.Option(None, "--transport", help="New transport type"),
|
||||
command: Optional[str] = typer.Option(None, "--command", help="New command (stdio)"),
|
||||
url: Optional[str] = typer.Option(None, "--url", help="New URL (http/sse/websocket)"),
|
||||
tools: Optional[str] = typer.Option(None, "--tools", "-t", help="Comma-separated tool allowlist, supports wildcards ('none' to clear)"),
|
||||
expose_to: Optional[str] = typer.Option(None, "--expose-to", "-e", help="Comma-separated target agents ('none' to clear)"),
|
||||
header: Optional[list[str]] = typer.Option(None, "--header", "-H", help="HTTP header as Key:Value (repeatable)"),
|
||||
env: Optional[list[str]] = typer.Option(None, "--env", help="Env var as KEY=VALUE for stdio (repeatable)"),
|
||||
transport: Optional[str] = typer.Option(
|
||||
None, "--transport", help="New transport type"
|
||||
),
|
||||
command: Optional[str] = typer.Option(
|
||||
None, "--command", help="New command (stdio)"
|
||||
),
|
||||
url: Optional[str] = typer.Option(
|
||||
None, "--url", help="New URL (http/sse/websocket)"
|
||||
),
|
||||
tools: Optional[str] = typer.Option(
|
||||
None,
|
||||
"--tools",
|
||||
"-t",
|
||||
help="Comma-separated tool allowlist, supports wildcards ('none' to clear)",
|
||||
),
|
||||
expose_to: Optional[str] = typer.Option(
|
||||
None,
|
||||
"--expose-to",
|
||||
"-e",
|
||||
help="Comma-separated target agents ('none' to clear)",
|
||||
),
|
||||
header: Optional[list[str]] = typer.Option(
|
||||
None, "--header", "-H", help="HTTP header as Key:Value (repeatable)"
|
||||
),
|
||||
env: Optional[list[str]] = typer.Option(
|
||||
None, "--env", help="Env var as KEY=VALUE for stdio (repeatable)"
|
||||
),
|
||||
):
|
||||
"""Edit an existing MCP server in user config
|
||||
|
||||
@@ -787,6 +873,7 @@ def mcp_remove(
|
||||
# Main callback (default behavior)
|
||||
# =============================================================================
|
||||
|
||||
|
||||
@app.callback(invoke_without_command=True)
|
||||
def _main_callback(
|
||||
ctx: typer.Context,
|
||||
@@ -802,13 +889,31 @@ def _main_callback(
|
||||
"--name",
|
||||
help="Name for this run (used as directory name instead of timestamp; requires --mode run)",
|
||||
),
|
||||
prompt: Optional[str] = typer.Option(None, "-p", "--prompt", help="Query to execute (single-shot mode)"),
|
||||
thread_id: Optional[str] = typer.Option(None, "--thread-id", help="Thread ID for conversation persistence"),
|
||||
workdir: Optional[str] = typer.Option(None, "--workdir", help="Override workspace directory for this session"),
|
||||
use_cwd: bool = typer.Option(False, "--use-cwd", help="Use current working directory as workspace"),
|
||||
no_thinking: bool = typer.Option(False, "--no-thinking", help="Disable thinking display"),
|
||||
auto_approve: bool = typer.Option(False, "--auto-approve", help="Auto-approve all tool executions without prompting"),
|
||||
ask_user: bool = typer.Option(False, "--ask-user", help="Enable agent to ask clarifying questions about your research preferences"),
|
||||
prompt: Optional[str] = typer.Option(
|
||||
None, "-p", "--prompt", help="Query to execute (single-shot mode)"
|
||||
),
|
||||
thread_id: Optional[str] = typer.Option(
|
||||
None, "--thread-id", help="Thread ID for conversation persistence"
|
||||
),
|
||||
workdir: Optional[str] = typer.Option(
|
||||
None, "--workdir", help="Override workspace directory for this session"
|
||||
),
|
||||
use_cwd: bool = typer.Option(
|
||||
False, "--use-cwd", help="Use current working directory as workspace"
|
||||
),
|
||||
no_thinking: bool = typer.Option(
|
||||
False, "--no-thinking", help="Disable thinking display"
|
||||
),
|
||||
auto_approve: bool = typer.Option(
|
||||
False,
|
||||
"--auto-approve",
|
||||
help="Auto-approve all tool executions without prompting",
|
||||
),
|
||||
ask_user: bool = typer.Option(
|
||||
False,
|
||||
"--ask-user",
|
||||
help="Enable agent to ask clarifying questions about your research preferences",
|
||||
),
|
||||
auth_mode: Optional[str] = typer.Option(
|
||||
None,
|
||||
"--auth-mode",
|
||||
@@ -826,6 +931,7 @@ def _main_callback(
|
||||
return
|
||||
|
||||
from dotenv import load_dotenv, find_dotenv # type: ignore[import-untyped]
|
||||
|
||||
# find_dotenv() traverses up the directory tree to locate .env
|
||||
load_dotenv(find_dotenv(), override=True)
|
||||
|
||||
@@ -859,9 +965,11 @@ def _main_callback(
|
||||
if config.provider == "anthropic" and config.anthropic_auth_mode == "oauth":
|
||||
try:
|
||||
from ..ccproxy_manager import maybe_start_ccproxy, stop_ccproxy
|
||||
|
||||
_ccproxy_proc = maybe_start_ccproxy(config)
|
||||
if _ccproxy_proc:
|
||||
import atexit
|
||||
|
||||
atexit.register(stop_ccproxy, _ccproxy_proc)
|
||||
except RuntimeError as exc:
|
||||
console.print(f"[red]{exc}[/red]")
|
||||
@@ -875,7 +983,9 @@ def _main_callback(
|
||||
raise typer.BadParameter("Use either --workdir or --use-cwd, not both.")
|
||||
|
||||
if mode and (workdir or use_cwd):
|
||||
raise typer.BadParameter("--mode cannot be combined with --workdir or --use-cwd")
|
||||
raise typer.BadParameter(
|
||||
"--mode cannot be combined with --workdir or --use-cwd"
|
||||
)
|
||||
|
||||
if mode and mode not in ("run", "daemon"):
|
||||
raise typer.BadParameter("--mode must be 'run' or 'daemon'")
|
||||
@@ -883,16 +993,23 @@ def _main_callback(
|
||||
raise typer.BadParameter("--ui must be 'tui' or 'cli'")
|
||||
|
||||
# --name only makes sense in run mode
|
||||
if name and not (mode == "run" or (not mode and not workdir and not use_cwd and config.default_mode == "run")):
|
||||
if name and not (
|
||||
mode == "run"
|
||||
or (not mode and not workdir and not use_cwd and config.default_mode == "run")
|
||||
):
|
||||
raise typer.BadParameter("--name can only be used with --mode run")
|
||||
|
||||
# Sanitize run name: allow alphanumeric, hyphens, underscores
|
||||
if name:
|
||||
if not re.fullmatch(r"[A-Za-z0-9_-]+", name):
|
||||
raise typer.BadParameter("--name may only contain letters, digits, hyphens, and underscores")
|
||||
raise typer.BadParameter(
|
||||
"--name may only contain letters, digits, hyphens, and underscores"
|
||||
)
|
||||
|
||||
# Resolve effective mode from config (CLI mode already applied via overrides)
|
||||
effective_mode: str | None = None # None means explicit --workdir/--use-cwd was used
|
||||
effective_mode: str | None = (
|
||||
None # None means explicit --workdir/--use-cwd was used
|
||||
)
|
||||
|
||||
# Resolve workspace directory for this session
|
||||
# Priority: --workdir > --mode (explicit) > default_workdir > default_mode > cwd
|
||||
@@ -914,7 +1031,11 @@ def _main_callback(
|
||||
set_workspace_root(workspace_root)
|
||||
if effective_mode == "run":
|
||||
runs_dir = Path(workspace_root, "runs")
|
||||
session_id = _deduplicate_run_name(name, runs_dir) if name else datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
session_id = (
|
||||
_deduplicate_run_name(name, runs_dir)
|
||||
if name
|
||||
else datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
)
|
||||
workspace_dir = os.path.join(runs_dir, session_id)
|
||||
os.makedirs(workspace_dir, exist_ok=True)
|
||||
workspace_fixed = False
|
||||
@@ -928,7 +1049,11 @@ def _main_callback(
|
||||
effective_mode = config.default_mode
|
||||
if effective_mode == "run":
|
||||
runs_dir = Path(workspace_root, "runs")
|
||||
session_id = _deduplicate_run_name(name, runs_dir) if name else datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
session_id = (
|
||||
_deduplicate_run_name(name, runs_dir)
|
||||
if name
|
||||
else datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
)
|
||||
workspace_dir = os.path.join(runs_dir, session_id)
|
||||
os.makedirs(workspace_dir, exist_ok=True)
|
||||
workspace_fixed = False
|
||||
@@ -957,7 +1082,11 @@ def _main_callback(
|
||||
async def _single_shot():
|
||||
async with get_checkpointer() as checkpointer:
|
||||
console.print("[dim]Loading agent...[/dim]")
|
||||
agent = _load_agent(workspace_dir=workspace_dir, checkpointer=checkpointer, config=config)
|
||||
agent = _load_agent(
|
||||
workspace_dir=workspace_dir,
|
||||
checkpointer=checkpointer,
|
||||
config=config,
|
||||
)
|
||||
tid = thread_id or generate_thread_id()
|
||||
cmd_run(
|
||||
agent,
|
||||
@@ -970,6 +1099,7 @@ def _main_callback(
|
||||
)
|
||||
|
||||
import nest_asyncio # type: ignore[import-untyped]
|
||||
|
||||
nest_asyncio.apply()
|
||||
asyncio.get_event_loop().run_until_complete(_single_shot())
|
||||
else:
|
||||
@@ -1000,12 +1130,16 @@ def _configure_logging():
|
||||
if record.levelno == logging.WARNING:
|
||||
# Use Rich console to print dim warning
|
||||
msg = record.getMessage()
|
||||
console.print(f"[dim yellow]\u26a0\ufe0f Warning:[/dim yellow] [dim]{escape(msg)}[/dim]")
|
||||
console.print(
|
||||
f"[dim yellow]\u26a0\ufe0f Warning:[/dim yellow] [dim]{escape(msg)}[/dim]"
|
||||
)
|
||||
else:
|
||||
super().emit(record)
|
||||
|
||||
# Configure root logger to use our handler for WARNING and above
|
||||
handler = DimWarningHandler(console=console, show_time=False, show_path=False, show_level=False)
|
||||
handler = DimWarningHandler(
|
||||
console=console, show_time=False, show_path=False, show_level=False
|
||||
)
|
||||
handler.setLevel(logging.WARNING)
|
||||
|
||||
# Apply to root logger (catches all loggers including deepagents)
|
||||
|
||||
+158
-63
@@ -55,6 +55,7 @@ _channel_logger = logging.getLogger(__name__)
|
||||
# Banner
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def print_banner(
|
||||
thread_id: str,
|
||||
workspace_dir: str | None = None,
|
||||
@@ -85,9 +86,14 @@ def print_banner(
|
||||
info.append(value, style="magenta")
|
||||
# Directory line
|
||||
import os
|
||||
|
||||
effective_dir = workspace_dir or os.getcwd()
|
||||
home = os.path.expanduser("~")
|
||||
dir_display = effective_dir.replace(home, "~", 1) if effective_dir.startswith(home) else effective_dir
|
||||
dir_display = (
|
||||
effective_dir.replace(home, "~", 1)
|
||||
if effective_dir.startswith(home)
|
||||
else effective_dir
|
||||
)
|
||||
info.append("\n ", style="dim")
|
||||
info.append("Directory: ", style="dim")
|
||||
info.append(dir_display, style="magenta")
|
||||
@@ -116,26 +122,30 @@ _SLASH_COMMANDS = [
|
||||
("/exit", "Quit EvoScientist"),
|
||||
]
|
||||
|
||||
_COMPLETION_STYLE = PtStyle.from_dict({
|
||||
"completion-menu": "bg:default noreverse nounderline noitalic",
|
||||
"completion-menu.completion": "bg:default #888888 noreverse",
|
||||
"completion-menu.completion.current": "bg:default default bold noreverse",
|
||||
"completion-menu.meta.completion": "bg:default #888888 noreverse",
|
||||
"completion-menu.meta.completion.current": "bg:default default bold noreverse",
|
||||
"scrollbar.background": "bg:default",
|
||||
"scrollbar.button": "bg:default",
|
||||
})
|
||||
_COMPLETION_STYLE = PtStyle.from_dict(
|
||||
{
|
||||
"completion-menu": "bg:default noreverse nounderline noitalic",
|
||||
"completion-menu.completion": "bg:default #888888 noreverse",
|
||||
"completion-menu.completion.current": "bg:default default bold noreverse",
|
||||
"completion-menu.meta.completion": "bg:default #888888 noreverse",
|
||||
"completion-menu.meta.completion.current": "bg:default default bold noreverse",
|
||||
"scrollbar.background": "bg:default",
|
||||
"scrollbar.button": "bg:default",
|
||||
}
|
||||
)
|
||||
|
||||
# Style for questionary pickers — matches _COMPLETION_STYLE visual language:
|
||||
# gray (#888888) for non-selected, bold for selected, no background changes.
|
||||
_PICKER_STYLE = PtStyle.from_dict({
|
||||
"questionmark": "#888888",
|
||||
"question": "",
|
||||
"pointer": "bold",
|
||||
"highlighted": "bold",
|
||||
"text": "#888888",
|
||||
"answer": "bold",
|
||||
})
|
||||
_PICKER_STYLE = PtStyle.from_dict(
|
||||
{
|
||||
"questionmark": "#888888",
|
||||
"question": "",
|
||||
"pointer": "bold",
|
||||
"highlighted": "bold",
|
||||
"text": "#888888",
|
||||
"answer": "bold",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class SlashCommandCompleter(Completer):
|
||||
@@ -191,11 +201,13 @@ def cmd_interactive(
|
||||
ui_backend: UI backend ('cli' or 'tui')
|
||||
"""
|
||||
import nest_asyncio
|
||||
|
||||
nest_asyncio.apply()
|
||||
|
||||
resolved_ui_backend = resolve_ui_backend(ui_backend, warn_fallback=True)
|
||||
if resolved_ui_backend == "tui":
|
||||
from functools import partial
|
||||
|
||||
load_agent = partial(_load_agent, config=config)
|
||||
run_textual_interactive(
|
||||
show_thinking=show_thinking,
|
||||
@@ -213,6 +225,7 @@ def cmd_interactive(
|
||||
return
|
||||
|
||||
from .. import paths
|
||||
|
||||
memory_dir = str(paths.MEMORY_DIR)
|
||||
|
||||
from ..config.settings import get_config_dir
|
||||
@@ -252,7 +265,9 @@ def cmd_interactive(
|
||||
if len(similar) == 1:
|
||||
return similar[0]
|
||||
if len(similar) > 1:
|
||||
console.print(f"[yellow]Ambiguous thread ID '{escape(tid)}'. Matches:[/yellow]")
|
||||
console.print(
|
||||
f"[yellow]Ambiguous thread ID '{escape(tid)}'. Matches:[/yellow]"
|
||||
)
|
||||
for s in similar:
|
||||
console.print(f" [cyan]{s}[/cyan]")
|
||||
return None
|
||||
@@ -262,7 +277,9 @@ def cmd_interactive(
|
||||
async def _cmd_threads():
|
||||
"""Handle /threads command — show recent sessions."""
|
||||
threads = await list_threads(
|
||||
limit=0, include_message_count=True, include_preview=True,
|
||||
limit=0,
|
||||
include_message_count=True,
|
||||
include_preview=True,
|
||||
)
|
||||
if not threads:
|
||||
console.print("[yellow]No saved sessions.[/yellow]")
|
||||
@@ -285,7 +302,9 @@ def cmd_interactive(
|
||||
)
|
||||
console.print()
|
||||
console.print(table)
|
||||
console.print("[dim] /resume[/dim] to continue a session [dim]/delete <id>[/dim] to remove [dim]/new[/dim] to start fresh")
|
||||
console.print(
|
||||
"[dim] /resume[/dim] to continue a session [dim]/delete <id>[/dim] to remove [dim]/new[/dim] to start fresh"
|
||||
)
|
||||
console.print()
|
||||
|
||||
async def _render_history(thread_id: str):
|
||||
@@ -308,24 +327,32 @@ def cmd_interactive(
|
||||
content = getattr(msg, "content", "") or ""
|
||||
# content can be a list of blocks (multimodal) — extract text
|
||||
if isinstance(content, list):
|
||||
parts = [b.get("text", "") for b in content if isinstance(b, dict) and b.get("type") == "text"]
|
||||
parts = [
|
||||
b.get("text", "")
|
||||
for b in content
|
||||
if isinstance(b, dict) and b.get("type") == "text"
|
||||
]
|
||||
content = " ".join(parts) if parts else ""
|
||||
|
||||
if msg_type == "human":
|
||||
console.print(Text.assemble(
|
||||
("\u276f ", "bold blue"),
|
||||
(_truncate(content), ""),
|
||||
))
|
||||
console.print(
|
||||
Text.assemble(
|
||||
("\u276f ", "bold blue"),
|
||||
(_truncate(content), ""),
|
||||
)
|
||||
)
|
||||
elif msg_type == "ai":
|
||||
tool_calls = getattr(msg, "tool_calls", None) or []
|
||||
if content:
|
||||
console.print(Text(_truncate(content), style="dim"))
|
||||
if tool_calls:
|
||||
names = [tc.get("name", "?") for tc in tool_calls]
|
||||
console.print(Text(
|
||||
f" \u25b6 {', '.join(names)}",
|
||||
style="dim italic",
|
||||
))
|
||||
console.print(
|
||||
Text(
|
||||
f" \u25b6 {', '.join(names)}",
|
||||
style="dim italic",
|
||||
)
|
||||
)
|
||||
# Skip tool messages — they are verbose and not useful in replay
|
||||
|
||||
console.print("[dim]── End of history ──[/dim]")
|
||||
@@ -336,7 +363,9 @@ def cmd_interactive(
|
||||
if not arg:
|
||||
# Show interactive session picker with conversation previews
|
||||
threads = await list_threads(
|
||||
limit=0, include_message_count=True, include_preview=True,
|
||||
limit=0,
|
||||
include_message_count=True,
|
||||
include_preview=True,
|
||||
)
|
||||
if not threads:
|
||||
console.print("[yellow]No sessions to resume.[/yellow]")
|
||||
@@ -397,14 +426,20 @@ def cmd_interactive(
|
||||
if ws:
|
||||
state["workspace_dir"] = ws
|
||||
console.print("[dim]Loading session...[/dim]")
|
||||
state["agent"] = _load_agent(workspace_dir=state["workspace_dir"], checkpointer=checkpointer, config=config)
|
||||
state["agent"] = _load_agent(
|
||||
workspace_dir=state["workspace_dir"],
|
||||
checkpointer=checkpointer,
|
||||
config=config,
|
||||
)
|
||||
# Sync shared refs if channel is running
|
||||
if _channels_is_running():
|
||||
_ch_mod._cli_agent = state["agent"]
|
||||
_ch_mod._cli_thread_id = state["thread_id"]
|
||||
console.print(f"[green]Resumed session:[/green] [yellow]{resolved}[/yellow]")
|
||||
if state["workspace_dir"]:
|
||||
console.print(f"[dim]Workspace:[/dim] [cyan]{_shorten_path(state['workspace_dir'])}[/cyan]")
|
||||
console.print(
|
||||
f"[dim]Workspace:[/dim] [cyan]{_shorten_path(state['workspace_dir'])}[/cyan]"
|
||||
)
|
||||
console.print()
|
||||
await _render_history(resolved)
|
||||
|
||||
@@ -440,7 +475,11 @@ def cmd_interactive(
|
||||
state["workspace_dir"] = ws
|
||||
|
||||
console.print("[dim]Loading agent...[/dim]")
|
||||
state["agent"] = _load_agent(workspace_dir=state["workspace_dir"], checkpointer=checkpointer, config=config)
|
||||
state["agent"] = _load_agent(
|
||||
workspace_dir=state["workspace_dir"],
|
||||
checkpointer=checkpointer,
|
||||
config=config,
|
||||
)
|
||||
|
||||
# Print banner
|
||||
if state["resumed"]:
|
||||
@@ -453,7 +492,9 @@ def cmd_interactive(
|
||||
provider,
|
||||
state["ui_backend"],
|
||||
)
|
||||
console.print(f"[green]Resumed session [yellow]{state['thread_id']}[/yellow][/green]\n")
|
||||
console.print(
|
||||
f"[green]Resumed session [yellow]{state['thread_id']}[/yellow][/green]\n"
|
||||
)
|
||||
else:
|
||||
print_banner(
|
||||
state["thread_id"],
|
||||
@@ -504,7 +545,9 @@ def cmd_interactive(
|
||||
if not loop:
|
||||
return
|
||||
try:
|
||||
asyncio.run_coroutine_threadsafe(coro, loop).result(timeout=timeout)
|
||||
asyncio.run_coroutine_threadsafe(coro, loop).result(
|
||||
timeout=timeout
|
||||
)
|
||||
except Exception as e:
|
||||
_channel_logger.debug(f"{label} send failed: {e}")
|
||||
|
||||
@@ -512,23 +555,37 @@ def cmd_interactive(
|
||||
ch = msg.channel_ref
|
||||
if ch and ch.send_thinking:
|
||||
_send_to_channel(
|
||||
ch.send_thinking_message(sender=msg.chat_id, thinking=thinking, metadata=msg.metadata),
|
||||
ch.send_thinking_message(
|
||||
sender=msg.chat_id,
|
||||
thinking=thinking,
|
||||
metadata=msg.metadata,
|
||||
),
|
||||
"Thinking",
|
||||
)
|
||||
|
||||
def _send_todo_to_channel(items: list[dict]) -> None:
|
||||
from ..channels.consumer import _format_todo_list
|
||||
|
||||
if msg.channel_ref:
|
||||
_send_to_channel(
|
||||
msg.channel_ref.send_todo_message(sender=msg.chat_id, content=_format_todo_list(items), metadata=msg.metadata),
|
||||
msg.channel_ref.send_todo_message(
|
||||
sender=msg.chat_id,
|
||||
content=_format_todo_list(items),
|
||||
metadata=msg.metadata,
|
||||
),
|
||||
"Todo",
|
||||
)
|
||||
|
||||
def _send_media_to_channel(file_path: str) -> None:
|
||||
if msg.channel_ref:
|
||||
_send_to_channel(
|
||||
msg.channel_ref.send_media(recipient=msg.chat_id, file_path=file_path, metadata=msg.metadata),
|
||||
"Media", timeout=30,
|
||||
msg.channel_ref.send_media(
|
||||
recipient=msg.chat_id,
|
||||
file_path=file_path,
|
||||
metadata=msg.metadata,
|
||||
),
|
||||
"Media",
|
||||
timeout=30,
|
||||
)
|
||||
|
||||
def _channel_hitl_prompt(action_requests: list) -> list[dict] | None:
|
||||
@@ -586,8 +643,13 @@ def cmd_interactive(
|
||||
|
||||
# Auto-start channel if enabled in config
|
||||
from ..config import load_config
|
||||
|
||||
_channel_cfg = load_config()
|
||||
if _channel_cfg and _channel_cfg.channel_enabled and not _channels_is_running():
|
||||
if (
|
||||
_channel_cfg
|
||||
and _channel_cfg.channel_enabled
|
||||
and not _channels_is_running()
|
||||
):
|
||||
_auto_start_channel(
|
||||
state["agent"],
|
||||
state["thread_id"],
|
||||
@@ -596,7 +658,9 @@ def cmd_interactive(
|
||||
)
|
||||
|
||||
# Slogan — after channels, right before user input
|
||||
console.print(Text(f" {random.choice(WELCOME_SLOGANS)}", style="dim italic"))
|
||||
console.print(
|
||||
Text(f" {random.choice(WELCOME_SLOGANS)}", style="dim italic")
|
||||
)
|
||||
console.print()
|
||||
|
||||
try:
|
||||
@@ -604,7 +668,7 @@ def cmd_interactive(
|
||||
while state["running"]:
|
||||
try:
|
||||
user_input = await session.prompt_async(
|
||||
HTML('<ansiblue><b>\u276f</b></ansiblue> ')
|
||||
HTML("<ansiblue><b>\u276f</b></ansiblue> ")
|
||||
)
|
||||
user_input = user_input.strip()
|
||||
|
||||
@@ -627,39 +691,57 @@ def cmd_interactive(
|
||||
continue
|
||||
|
||||
if user_input.lower().startswith("/resume"):
|
||||
arg = user_input[len("/resume"):].strip()
|
||||
arg = user_input[len("/resume") :].strip()
|
||||
await _cmd_resume(arg, checkpointer)
|
||||
continue
|
||||
|
||||
if user_input.lower().startswith("/delete"):
|
||||
arg = user_input[len("/delete"):].strip()
|
||||
arg = user_input[len("/delete") :].strip()
|
||||
await _cmd_delete(arg)
|
||||
continue
|
||||
|
||||
if user_input.lower() == "/new":
|
||||
# New session: new thread; workspace only changes if not fixed
|
||||
if not workspace_fixed:
|
||||
state["workspace_dir"] = _create_session_workspace(run_name)
|
||||
state["workspace_dir"] = _create_session_workspace(
|
||||
run_name
|
||||
)
|
||||
console.print("[dim]Loading new session...[/dim]")
|
||||
state["agent"] = _load_agent(workspace_dir=state["workspace_dir"], checkpointer=checkpointer, config=config)
|
||||
state["agent"] = _load_agent(
|
||||
workspace_dir=state["workspace_dir"],
|
||||
checkpointer=checkpointer,
|
||||
config=config,
|
||||
)
|
||||
state["thread_id"] = generate_thread_id()
|
||||
state["resumed"] = False
|
||||
# Sync channel refs so the queue checker uses the new agent
|
||||
if _channels_is_running():
|
||||
_ch_mod._cli_agent = state["agent"]
|
||||
_ch_mod._cli_thread_id = state["thread_id"]
|
||||
console.print(f"[green]New session:[/green] [yellow]{state['thread_id']}[/yellow]")
|
||||
console.print(
|
||||
f"[green]New session:[/green] [yellow]{state['thread_id']}[/yellow]"
|
||||
)
|
||||
if state["workspace_dir"]:
|
||||
console.print(f"[dim]Workspace:[/dim] [cyan]{_shorten_path(state['workspace_dir'])}[/cyan]\n")
|
||||
console.print(
|
||||
f"[dim]Workspace:[/dim] [cyan]{_shorten_path(state['workspace_dir'])}[/cyan]\n"
|
||||
)
|
||||
continue
|
||||
|
||||
if user_input.lower() == "/current":
|
||||
console.print(f"[dim]Thread:[/dim] [yellow]{state['thread_id']}[/yellow]")
|
||||
console.print(
|
||||
f"[dim]Thread:[/dim] [yellow]{state['thread_id']}[/yellow]"
|
||||
)
|
||||
if state["workspace_dir"]:
|
||||
console.print(f"[dim]Workspace:[/dim] [cyan]{_shorten_path(state['workspace_dir'])}[/cyan]")
|
||||
console.print(
|
||||
f"[dim]Workspace:[/dim] [cyan]{_shorten_path(state['workspace_dir'])}[/cyan]"
|
||||
)
|
||||
if memory_dir:
|
||||
console.print(f"[dim]Memory dir:[/dim] [cyan]{_shorten_path(memory_dir)}[/cyan]")
|
||||
console.print(f"[dim]UI:[/dim] [cyan]{state['ui_backend']}[/cyan]")
|
||||
console.print(
|
||||
f"[dim]Memory dir:[/dim] [cyan]{_shorten_path(memory_dir)}[/cyan]"
|
||||
)
|
||||
console.print(
|
||||
f"[dim]UI:[/dim] [cyan]{state['ui_backend']}[/cyan]"
|
||||
)
|
||||
console.print()
|
||||
continue
|
||||
|
||||
@@ -668,23 +750,23 @@ def cmd_interactive(
|
||||
continue
|
||||
|
||||
if user_input.lower().startswith("/install-skill"):
|
||||
source = user_input[len("/install-skill"):].strip()
|
||||
source = user_input[len("/install-skill") :].strip()
|
||||
_cmd_install_skill(source)
|
||||
continue
|
||||
|
||||
if user_input.lower().startswith("/uninstall-skill"):
|
||||
name = user_input[len("/uninstall-skill"):].strip()
|
||||
name = user_input[len("/uninstall-skill") :].strip()
|
||||
_cmd_uninstall_skill(name)
|
||||
continue
|
||||
|
||||
if user_input.lower().startswith("/mcp"):
|
||||
_cmd_mcp(user_input[len("/mcp"):])
|
||||
_cmd_mcp(user_input[len("/mcp") :])
|
||||
continue
|
||||
|
||||
if user_input.lower().startswith("/channel"):
|
||||
args = user_input[len("/channel"):].strip()
|
||||
args = user_input[len("/channel") :].strip()
|
||||
if args.lower().startswith("stop"):
|
||||
stop_arg = args[len("stop"):].strip()
|
||||
stop_arg = args[len("stop") :].strip()
|
||||
_cmd_channel_stop(stop_arg or None)
|
||||
else:
|
||||
_cmd_channel(
|
||||
@@ -696,8 +778,14 @@ def cmd_interactive(
|
||||
continue
|
||||
|
||||
if user_input.lower() == "/compact":
|
||||
from .commands import compact_conversation, render_compact_result
|
||||
with console.status("[cyan]Compacting conversation...[/cyan]"):
|
||||
from .commands import (
|
||||
compact_conversation,
|
||||
render_compact_result,
|
||||
)
|
||||
|
||||
with console.status(
|
||||
"[cyan]Compacting conversation...[/cyan]"
|
||||
):
|
||||
result = await compact_conversation(
|
||||
agent=state["agent"],
|
||||
thread_id=state["thread_id"],
|
||||
@@ -731,9 +819,14 @@ def cmd_interactive(
|
||||
break
|
||||
except Exception as e:
|
||||
error_msg = str(e)
|
||||
if "authentication" in error_msg.lower() or "api_key" in error_msg.lower():
|
||||
if (
|
||||
"authentication" in error_msg.lower()
|
||||
or "api_key" in error_msg.lower()
|
||||
):
|
||||
console.print("[red]Error: API key not configured.[/red]")
|
||||
console.print("[dim]Run [bold]EvoSci onboard[/bold] to set up your API key.[/dim]")
|
||||
console.print(
|
||||
"[dim]Run [bold]EvoSci onboard[/bold] to set up your API key.[/dim]"
|
||||
)
|
||||
state["running"] = False
|
||||
break
|
||||
else:
|
||||
@@ -799,7 +892,9 @@ def cmd_run(
|
||||
error_msg = str(e)
|
||||
if "authentication" in error_msg.lower() or "api_key" in error_msg.lower():
|
||||
console.print("[red]Error: API key not configured.[/red]")
|
||||
console.print("[dim]Run [bold]EvoSci onboard[/bold] to set up your API key.[/dim]")
|
||||
console.print(
|
||||
"[dim]Run [bold]EvoSci onboard[/bold] to set up your API key.[/dim]"
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
else:
|
||||
console.print(f"[red]Error: {e}[/red]")
|
||||
|
||||
+30
-10
@@ -16,7 +16,9 @@ def _mcp_list_servers() -> None:
|
||||
|
||||
if not config:
|
||||
console.print("[dim]No MCP servers configured.[/dim]")
|
||||
console.print("[dim]Add one with:[/dim] /mcp add <name> <command-or-url> [args...]")
|
||||
console.print(
|
||||
"[dim]Add one with:[/dim] /mcp add <name> <command-or-url> [args...]"
|
||||
)
|
||||
console.print()
|
||||
return
|
||||
|
||||
@@ -51,7 +53,9 @@ def _mcp_add_server_from_kwargs(
|
||||
|
||||
try:
|
||||
entry = add_mcp_server(**kwargs)
|
||||
console.print(f"[green]Added MCP server:[/green] [cyan]{kwargs['name']}[/cyan] ({entry['transport']})")
|
||||
console.print(
|
||||
f"[green]Added MCP server:[/green] [cyan]{kwargs['name']}[/cyan] ({entry['transport']})"
|
||||
)
|
||||
if show_reload_hint:
|
||||
console.print("[dim]Reload with /new to apply.[/dim]")
|
||||
return True
|
||||
@@ -70,7 +74,9 @@ def _mcp_edit_server_fields(
|
||||
from ..mcp import edit_mcp_server
|
||||
|
||||
if not fields:
|
||||
console.print("[red]No fields to edit. Use --transport, --command, --url, --tools, --expose-to, etc.[/red]")
|
||||
console.print(
|
||||
"[red]No fields to edit. Use --transport, --command, --url, --tools, --expose-to, etc.[/red]"
|
||||
)
|
||||
return False
|
||||
|
||||
try:
|
||||
@@ -185,20 +191,30 @@ def _cmd_mcp_add(args_str: str) -> None:
|
||||
if not args_str.strip():
|
||||
console.print("[bold]Usage:[/bold] /mcp add <name> <command-or-url> [args...]")
|
||||
console.print()
|
||||
console.print("[dim]Transport is auto-detected: URLs \u2192 http, commands \u2192 stdio[/dim]")
|
||||
console.print(
|
||||
"[dim]Transport is auto-detected: URLs \u2192 http, commands \u2192 stdio[/dim]"
|
||||
)
|
||||
console.print()
|
||||
console.print("[bold]Examples:[/bold]")
|
||||
console.print(" /mcp add sequential-thinking npx -y @modelcontextprotocol/server-sequential-thinking")
|
||||
console.print(
|
||||
" /mcp add sequential-thinking npx -y @modelcontextprotocol/server-sequential-thinking"
|
||||
)
|
||||
console.print(" /mcp add docs-langchain https://docs.langchain.com/mcp")
|
||||
console.print(" /mcp add my-sse http://localhost:9090/sse --transport sse --expose-to research-agent")
|
||||
console.print(
|
||||
" /mcp add my-sse http://localhost:9090/sse --transport sse --expose-to research-agent"
|
||||
)
|
||||
console.print()
|
||||
console.print("[dim]Options:[/dim]")
|
||||
console.print(" --transport T Transport type (default: auto-detect)")
|
||||
console.print(" --tools t1,t2 Tool allowlist (supports wildcards: *_exa, read_*)")
|
||||
console.print(
|
||||
" --tools t1,t2 Tool allowlist (supports wildcards: *_exa, read_*)"
|
||||
)
|
||||
console.print(" --expose-to a1,a2 Target agents (default: main)")
|
||||
console.print(" --header Key:Value HTTP header (repeatable)")
|
||||
console.print(" --env KEY=VALUE Env var for stdio (repeatable)")
|
||||
console.print(" --env-ref KEY Env var as runtime ${KEY} reference (repeatable)")
|
||||
console.print(
|
||||
" --env-ref KEY Env var as runtime ${KEY} reference (repeatable)"
|
||||
)
|
||||
console.print()
|
||||
return
|
||||
|
||||
@@ -219,8 +235,12 @@ def _cmd_mcp_edit(args_str: str) -> None:
|
||||
if not args_str.strip():
|
||||
console.print("[bold]Usage:[/bold] /mcp edit <name> --<field> <value> ...")
|
||||
console.print()
|
||||
console.print("[dim]Fields:[/dim] --transport, --command, --url, --args, --tools, --expose-to, --header, --env")
|
||||
console.print("[dim]Use[/dim] --tools none [dim]or[/dim] --expose-to none [dim]to clear a field.[/dim]")
|
||||
console.print(
|
||||
"[dim]Fields:[/dim] --transport, --command, --url, --args, --tools, --expose-to, --header, --env"
|
||||
)
|
||||
console.print(
|
||||
"[dim]Use[/dim] --tools none [dim]or[/dim] --expose-to none [dim]to clear a field.[/dim]"
|
||||
)
|
||||
console.print()
|
||||
console.print("[bold]Examples:[/bold]")
|
||||
console.print(" /mcp edit filesystem --expose-to main,code-agent")
|
||||
|
||||
@@ -14,7 +14,9 @@ def _cmd_list_skills() -> None:
|
||||
if not skills:
|
||||
console.print("[dim]No skills available.[/dim]")
|
||||
console.print("[dim]Install with:[/dim] /install-skill <path-or-url>")
|
||||
console.print(f"[dim]Skills directory:[/dim] [cyan]{_shorten_path(str(USER_SKILLS_DIR))}[/cyan]")
|
||||
console.print(
|
||||
f"[dim]Skills directory:[/dim] [cyan]{_shorten_path(str(USER_SKILLS_DIR))}[/cyan]"
|
||||
)
|
||||
console.print()
|
||||
return
|
||||
|
||||
@@ -34,7 +36,9 @@ def _cmd_list_skills() -> None:
|
||||
for skill in system_skills:
|
||||
console.print(f" [cyan]{skill.name}[/cyan] - {skill.description}")
|
||||
|
||||
console.print(f"\n[dim]User skills folder:[/dim] [green]{_shorten_path(str(USER_SKILLS_DIR))}[/green]")
|
||||
console.print(
|
||||
f"\n[dim]User skills folder:[/dim] [green]{_shorten_path(str(USER_SKILLS_DIR))}[/green]"
|
||||
)
|
||||
console.print()
|
||||
|
||||
|
||||
@@ -46,7 +50,9 @@ def _cmd_install_skill(source: str) -> None:
|
||||
console.print("[red]Usage:[/red] /install-skill <path-or-url>")
|
||||
console.print("[dim]Examples:[/dim]")
|
||||
console.print(" /install-skill ./my-skill")
|
||||
console.print(" /install-skill https://github.com/user/repo/tree/main/skill-name")
|
||||
console.print(
|
||||
" /install-skill https://github.com/user/repo/tree/main/skill-name"
|
||||
)
|
||||
console.print(" /install-skill user/repo@skill-name")
|
||||
console.print()
|
||||
return
|
||||
@@ -59,8 +65,12 @@ def _cmd_install_skill(source: str) -> None:
|
||||
# Batch install — multiple skills
|
||||
for item in result.get("installed", []):
|
||||
console.print(f"[green]Installed:[/green] {item['name']}")
|
||||
console.print(f" [dim]Description:[/dim] {item.get('description', '(none)')}")
|
||||
console.print(f" [dim]Path:[/dim] [cyan]{_shorten_path(item['path'])}[/cyan]")
|
||||
console.print(
|
||||
f" [dim]Description:[/dim] {item.get('description', '(none)')}"
|
||||
)
|
||||
console.print(
|
||||
f" [dim]Path:[/dim] [cyan]{_shorten_path(item['path'])}[/cyan]"
|
||||
)
|
||||
for item in result.get("failed", []):
|
||||
console.print(f"[red]Failed:[/red] {item['name']} — {item['error']}")
|
||||
installed_count = len(result.get("installed", []))
|
||||
|
||||
+366
-241
File diff suppressed because it is too large
Load Diff
@@ -26,6 +26,7 @@ def normalize_ui_backend(value: str | None) -> str:
|
||||
def _has_textual_support() -> bool:
|
||||
try:
|
||||
import textual # noqa: F401
|
||||
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
@@ -37,14 +38,16 @@ def resolve_ui_backend(value: str | None, *, warn_fallback: bool = False) -> str
|
||||
if requested == "tui" and not _has_textual_support():
|
||||
if warn_fallback:
|
||||
console.print(
|
||||
'[yellow]TUI is unavailable (missing textual package). '
|
||||
'Falling back to CLI.[/yellow]'
|
||||
"[yellow]TUI is unavailable (missing textual package). "
|
||||
"Falling back to CLI.[/yellow]"
|
||||
)
|
||||
return DEFAULT_UI_BACKEND
|
||||
return requested
|
||||
|
||||
|
||||
def get_backend(name: str | None, *, warn_fallback: bool = False) -> StreamingTUIBackend:
|
||||
def get_backend(
|
||||
name: str | None, *, warn_fallback: bool = False
|
||||
) -> StreamingTUIBackend:
|
||||
"""Instantiate a streaming backend by name.
|
||||
|
||||
Note: The Textual TUI is now a full interactive app (tui_interactive.py),
|
||||
|
||||
@@ -105,7 +105,11 @@ class ApprovalWidget(Widget):
|
||||
self._option_widgets = []
|
||||
count = len(self._action_requests)
|
||||
if count == 1:
|
||||
name = self._action_requests[0].get("name", "") if isinstance(self._action_requests[0], dict) else getattr(self._action_requests[0], "name", "")
|
||||
name = (
|
||||
self._action_requests[0].get("name", "")
|
||||
if isinstance(self._action_requests[0], dict)
|
||||
else getattr(self._action_requests[0], "name", "")
|
||||
)
|
||||
title = f">>> {name} Requires Approval <<<"
|
||||
else:
|
||||
title = f">>> {count} Tool Calls Require Approval <<<"
|
||||
@@ -113,8 +117,16 @@ class ApprovalWidget(Widget):
|
||||
|
||||
# Show each action request as a compact line
|
||||
for req in self._action_requests:
|
||||
name = req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "")
|
||||
args = req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {})
|
||||
name = (
|
||||
req.get("name", "")
|
||||
if isinstance(req, dict)
|
||||
else getattr(req, "name", "")
|
||||
)
|
||||
args = (
|
||||
req.get("args", {})
|
||||
if isinstance(req, dict)
|
||||
else getattr(req, "args", {})
|
||||
)
|
||||
if isinstance(args, dict):
|
||||
command = args.get("command", args.get("path", ""))
|
||||
else:
|
||||
@@ -159,7 +171,9 @@ class ApprovalWidget(Widget):
|
||||
"3. Auto-approve for this session (a)",
|
||||
]
|
||||
|
||||
for i, (text, widget) in enumerate(zip(options, self._option_widgets, strict=True)):
|
||||
for i, (text, widget) in enumerate(
|
||||
zip(options, self._option_widgets, strict=True)
|
||||
):
|
||||
cursor = "▸ " if i == self._selected else " "
|
||||
widget.update(f"{cursor}{text}")
|
||||
widget.remove_class("approval-option-selected")
|
||||
|
||||
@@ -137,7 +137,9 @@ class AskUserWidget(Widget):
|
||||
if total == 1:
|
||||
title = ">>> Quick check-in from EvoScientist <<<"
|
||||
else:
|
||||
title = ">>> Question 1/{} — Quick check-in from EvoScientist <<<".format(total)
|
||||
title = ">>> Question 1/{} — Quick check-in from EvoScientist <<<".format(
|
||||
total
|
||||
)
|
||||
|
||||
self._title_w = Static(title, classes="ask-title")
|
||||
yield self._title_w
|
||||
@@ -212,7 +214,9 @@ class AskUserWidget(Widget):
|
||||
)
|
||||
|
||||
# Question text
|
||||
suffix = " [dim](required)[/dim]" if self._required else " [dim](optional)[/dim]"
|
||||
suffix = (
|
||||
" [dim](required)[/dim]" if self._required else " [dim](optional)[/dim]"
|
||||
)
|
||||
if self._question_w:
|
||||
self._question_w.update(
|
||||
f"[bold]{index + 1}. {escape_markup(q_text)}[/bold]{suffix}"
|
||||
|
||||
@@ -122,7 +122,9 @@ class SubAgentWidget(Vertical):
|
||||
line = Text()
|
||||
if self._is_active:
|
||||
char = _SPINNER_FRAMES[self._frame]
|
||||
line.append(f"\u250c \u25b6 {self._display_name()} {char}", style="bold cyan")
|
||||
line.append(
|
||||
f"\u250c \u25b6 {self._display_name()} {char}", style="bold cyan"
|
||||
)
|
||||
else:
|
||||
line.append(f"\u2713 {self._display_name()}", style="bold green")
|
||||
line.append(f" ({self._tool_count} tools)", style="dim")
|
||||
@@ -225,11 +227,21 @@ class SubAgentWidget(Vertical):
|
||||
# Determine how many completed slots are available
|
||||
running_visible = self._running_ids[-_MAX_VISIBLE_RUNNING:]
|
||||
completed_slots = max(0, _MAX_VISIBLE_COMPLETED - len(running_visible))
|
||||
completed_visible = self._completed_ids[-completed_slots:] if completed_slots else []
|
||||
completed_hidden = self._completed_ids[:-completed_slots] if completed_slots and len(self._completed_ids) > completed_slots else (self._completed_ids if not completed_slots else [])
|
||||
completed_visible = (
|
||||
self._completed_ids[-completed_slots:] if completed_slots else []
|
||||
)
|
||||
completed_hidden = (
|
||||
self._completed_ids[:-completed_slots]
|
||||
if completed_slots and len(self._completed_ids) > completed_slots
|
||||
else (self._completed_ids if not completed_slots else [])
|
||||
)
|
||||
|
||||
# Running tools to hide
|
||||
running_hidden = self._running_ids[:-_MAX_VISIBLE_RUNNING] if len(self._running_ids) > _MAX_VISIBLE_RUNNING else []
|
||||
running_hidden = (
|
||||
self._running_ids[:-_MAX_VISIBLE_RUNNING]
|
||||
if len(self._running_ids) > _MAX_VISIBLE_RUNNING
|
||||
else []
|
||||
)
|
||||
|
||||
# Apply visibility
|
||||
visible_keys = set(completed_visible) | set(running_visible)
|
||||
@@ -269,7 +281,9 @@ class SubAgentWidget(Vertical):
|
||||
if hidden_running_count > 0:
|
||||
if total_hidden > 0:
|
||||
line.append(" | ", style="dim")
|
||||
line.append(f"\u25cf {hidden_running_count} more running...", style="dim yellow")
|
||||
line.append(
|
||||
f"\u25cf {hidden_running_count} more running...", style="dim yellow"
|
||||
)
|
||||
summary_w.update(line)
|
||||
summary_w.add_class("--visible")
|
||||
else:
|
||||
|
||||
@@ -82,7 +82,7 @@ class SummarizationWidget(Static):
|
||||
title = f"Context Summarized ({self._char_count_label()})"
|
||||
first_line = self._content.strip().split("\n")[0].strip()
|
||||
if len(first_line) > _MAX_COLLAPSED_CHARS:
|
||||
first_line = first_line[:_MAX_COLLAPSED_CHARS - 3] + "\u2026"
|
||||
first_line = first_line[: _MAX_COLLAPSED_CHARS - 3] + "\u2026"
|
||||
preview = Text(first_line, style="dim italic")
|
||||
preview.append(" [click to expand]", style="dim italic")
|
||||
body = preview
|
||||
@@ -92,17 +92,15 @@ class SummarizationWidget(Static):
|
||||
if len(display) > _MAX_EXPANDED_CHARS:
|
||||
half = _MAX_EXPANDED_CHARS // 2
|
||||
display = (
|
||||
display[:half]
|
||||
+ "\n\n... (truncated) ...\n\n"
|
||||
+ display[-half:]
|
||||
display[:half] + "\n\n... (truncated) ...\n\n" + display[-half:]
|
||||
)
|
||||
body = Text(display, style="dim italic") if display else Text(
|
||||
"(empty)", style="dim"
|
||||
body = (
|
||||
Text(display, style="dim italic")
|
||||
if display
|
||||
else Text("(empty)", style="dim")
|
||||
)
|
||||
|
||||
self.update(
|
||||
Panel(body, title=title, border_style="#f59e0b", padding=(0, 1))
|
||||
)
|
||||
self.update(Panel(body, title=title, border_style="#f59e0b", padding=(0, 1)))
|
||||
|
||||
def append_text(self, text: str) -> None:
|
||||
"""Append a chunk of summarization text (streaming)."""
|
||||
|
||||
@@ -90,8 +90,12 @@ class ThinkingWidget(Static):
|
||||
display = self._content.rstrip()
|
||||
if len(display) > _MAX_EXPANDED_CHARS:
|
||||
half = _MAX_EXPANDED_CHARS // 2
|
||||
display = display[:half] + "\n\n... (truncated) ...\n\n" + display[-half:]
|
||||
body = Text(display, style="dim") if display else Text("(empty)", style="dim")
|
||||
display = (
|
||||
display[:half] + "\n\n... (truncated) ...\n\n" + display[-half:]
|
||||
)
|
||||
body = (
|
||||
Text(display, style="dim") if display else Text("(empty)", style="dim")
|
||||
)
|
||||
|
||||
self.update(Panel(body, title=title, border_style="blue", padding=(0, 1)))
|
||||
|
||||
|
||||
@@ -26,6 +26,7 @@ if TYPE_CHECKING:
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def build_row_text(
|
||||
thread: dict,
|
||||
*,
|
||||
@@ -67,6 +68,7 @@ def build_row_text(
|
||||
# Widget
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ThreadPickerWidget(Widget):
|
||||
"""Inline thread picker — mounts in chat, keyboard-driven.
|
||||
|
||||
@@ -164,7 +166,9 @@ class ThreadPickerWidget(Widget):
|
||||
def _update_rows(self) -> None:
|
||||
for i, (thread, widget) in enumerate(zip(self._threads, self._row_widgets)):
|
||||
is_current = thread["thread_id"] == self._current_thread
|
||||
text = build_row_text(thread, selected=(i == self._selected), current=is_current)
|
||||
text = build_row_text(
|
||||
thread, selected=(i == self._selected), current=is_current
|
||||
)
|
||||
widget.update(text)
|
||||
widget.remove_class("picker-row-selected")
|
||||
if i == self._selected:
|
||||
|
||||
@@ -51,9 +51,7 @@ class TodoWidget(Static):
|
||||
if i > 0:
|
||||
lines.append("\n")
|
||||
status = str(item.get("status", "todo")).lower()
|
||||
content = str(
|
||||
item.get("content", item.get("task", item.get("title", "")))
|
||||
)
|
||||
content = str(item.get("content", item.get("task", item.get("title", ""))))
|
||||
|
||||
if status in ("done", "completed", "complete"):
|
||||
symbol = "\u2713"
|
||||
@@ -68,10 +66,12 @@ class TodoWidget(Static):
|
||||
lines.append(f"{symbol} ", style=style)
|
||||
lines.append(content, style=style)
|
||||
|
||||
self.update(Panel(
|
||||
lines,
|
||||
title="Task List",
|
||||
title_align="center",
|
||||
border_style="cyan",
|
||||
padding=(0, 1),
|
||||
))
|
||||
self.update(
|
||||
Panel(
|
||||
lines,
|
||||
title="Task List",
|
||||
title_align="center",
|
||||
border_style="cyan",
|
||||
padding=(0, 1),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -146,7 +146,9 @@ class ToolCallWidget(Vertical):
|
||||
|
||||
def _should_collapse(self) -> bool:
|
||||
lines = self._result_content.strip().split("\n")
|
||||
return len(lines) > _COLLAPSE_LINES or len(self._result_content) > _COLLAPSE_CHARS
|
||||
return (
|
||||
len(lines) > _COLLAPSE_LINES or len(self._result_content) > _COLLAPSE_CHARS
|
||||
)
|
||||
|
||||
def set_success(self, content: str) -> None:
|
||||
"""Mark tool call as successfully completed."""
|
||||
@@ -189,7 +191,11 @@ class ToolCallWidget(Vertical):
|
||||
if not self._result_content.strip():
|
||||
return
|
||||
# Diff rendering for edit_file (never truncates — collapses instead)
|
||||
if self._tool_name == "edit_file" and self._status == "success" and self._tool_args:
|
||||
if (
|
||||
self._tool_name == "edit_file"
|
||||
and self._status == "success"
|
||||
and self._tool_args
|
||||
):
|
||||
old_str = self._tool_args.get("old_string", "")
|
||||
new_str = self._tool_args.get("new_string", "")
|
||||
path = self._tool_args.get("path", self._tool_args.get("file_path", ""))
|
||||
|
||||
@@ -43,5 +43,6 @@ __all__ = [
|
||||
def __getattr__(name: str):
|
||||
if name == "run_onboard":
|
||||
from .onboard import run_onboard
|
||||
|
||||
return run_onboard
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
@@ -532,7 +532,14 @@ def _step_provider(config: EvoScientistConfig) -> str:
|
||||
value="zhipu-code",
|
||||
),
|
||||
Choice(title="Ollama (local models)", value="ollama"),
|
||||
Choice(title="Other (OpenAI-compatible)", value="custom"),
|
||||
Choice(
|
||||
title="OpenAI-compatible (third-party OpenAI endpoint)",
|
||||
value="custom-openai",
|
||||
),
|
||||
Choice(
|
||||
title="Claude-compatible (third-party Anthropic endpoint)",
|
||||
value="custom-anthropic",
|
||||
),
|
||||
]
|
||||
|
||||
# Set default based on current config
|
||||
@@ -592,9 +599,15 @@ def _provider_key_info(config: EvoScientistConfig, provider: str):
|
||||
config.zhipu_api_key or os.environ.get("ZHIPU_API_KEY", ""),
|
||||
validate_zhipu_key,
|
||||
),
|
||||
"custom": (
|
||||
"Custom",
|
||||
config.custom_api_key or os.environ.get("CUSTOM_API_KEY", ""),
|
||||
"custom-openai": (
|
||||
"OpenAI-compatible",
|
||||
config.custom_openai_api_key or os.environ.get("CUSTOM_OPENAI_API_KEY", ""),
|
||||
None,
|
||||
),
|
||||
"custom-anthropic": (
|
||||
"Custom Anthropic",
|
||||
config.custom_anthropic_api_key
|
||||
or os.environ.get("CUSTOM_ANTHROPIC_API_KEY", ""),
|
||||
None,
|
||||
),
|
||||
"ollama": ("Ollama", "__no_key__", None),
|
||||
@@ -682,13 +695,15 @@ def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str:
|
||||
if not is_ccproxy_available():
|
||||
console.print(
|
||||
" [dim]OAuth via ccproxy not available. "
|
||||
"Install with: pip install \"evoscientist[oauth]\"[/dim]"
|
||||
'Install with: pip install "evoscientist[oauth]"[/dim]'
|
||||
)
|
||||
return "api_key"
|
||||
|
||||
choices = [
|
||||
Choice(title="API Key (direct Anthropic access)", value="api_key"),
|
||||
Choice(title="Claude Code OAuth (via ccproxy — no API key needed)", value="oauth"),
|
||||
Choice(
|
||||
title="Claude Code OAuth (via ccproxy — no API key needed)", value="oauth"
|
||||
),
|
||||
]
|
||||
|
||||
current = config.anthropic_auth_mode
|
||||
@@ -768,16 +783,17 @@ def _step_provider_api_key(
|
||||
)
|
||||
|
||||
|
||||
def _step_base_url(config: EvoScientistConfig) -> str:
|
||||
def _step_base_url(config: EvoScientistConfig, current_value: str | None = None) -> str:
|
||||
"""Prompt for custom provider base URL.
|
||||
|
||||
Args:
|
||||
config: Current configuration.
|
||||
current_value: Current base URL value (if None, defaults to empty).
|
||||
|
||||
Returns:
|
||||
Base URL string.
|
||||
"""
|
||||
current = config.custom_base_url
|
||||
current = current_value if current_value is not None else ""
|
||||
hint = f"Current: {current}" if current else ""
|
||||
default = current if current else ""
|
||||
|
||||
@@ -1472,7 +1488,9 @@ def _step_mcp_servers() -> list[str]:
|
||||
console.print(f" [dim]Installing {pip_pkg}...[/dim]")
|
||||
if not _install_pip_package(pip_pkg):
|
||||
_print_step_result(
|
||||
"MCP", f"{name} — {_pip_install_hint()} {pip_pkg} failed", success=False
|
||||
"MCP",
|
||||
f"{name} — {_pip_install_hint()} {pip_pkg} failed",
|
||||
success=False,
|
||||
)
|
||||
continue
|
||||
|
||||
@@ -1803,7 +1821,11 @@ def _step_channels(config: EvoScientistConfig) -> dict[str, object]:
|
||||
console.print(" [yellow]✗ Required package not installed.[/yellow]")
|
||||
# Determine packages to install
|
||||
_pip_pkgs = _CHANNEL_PIP_DEPS.get(pip_extra, []) if pip_extra else []
|
||||
_pkg_display = " ".join(f'"{p}"' for p in _pip_pkgs) if _pip_pkgs else f'"evoscientist[{pip_extra}]"'
|
||||
_pkg_display = (
|
||||
" ".join(f'"{p}"' for p in _pip_pkgs)
|
||||
if _pip_pkgs
|
||||
else f'"evoscientist[{pip_extra}]"'
|
||||
)
|
||||
install_now = questionary.confirm(
|
||||
f"Install {_pkg_display} now?",
|
||||
default=True,
|
||||
@@ -1813,9 +1835,7 @@ def _step_channels(config: EvoScientistConfig) -> dict[str, object]:
|
||||
if install_now is None:
|
||||
raise KeyboardInterrupt()
|
||||
if install_now:
|
||||
console.print(
|
||||
f" [dim]Installing {_pkg_display}...[/dim]"
|
||||
)
|
||||
console.print(f" [dim]Installing {_pkg_display}...[/dim]")
|
||||
if _pip_pkgs:
|
||||
_ok = all(_install_pip_package(p) for p in _pip_pkgs)
|
||||
else:
|
||||
@@ -1827,14 +1847,16 @@ def _step_channels(config: EvoScientistConfig) -> dict[str, object]:
|
||||
console.print(" [green]✓ Installed successfully.[/green]")
|
||||
_pkg_ready = True
|
||||
except ImportError:
|
||||
console.print(" [red]✗ Package installed but import failed.[/red]")
|
||||
console.print(
|
||||
" [red]✗ Package installed but import failed.[/red]"
|
||||
)
|
||||
console.print(
|
||||
" [dim]Try restarting and running:[/dim] evosci channel setup"
|
||||
)
|
||||
else:
|
||||
console.print(" [red]✗ Installation failed.[/red]")
|
||||
console.print(
|
||||
f' [dim]Run manually:[/dim] {_pip_install_hint()} {_pkg_display}'
|
||||
f" [dim]Run manually:[/dim] {_pip_install_hint()} {_pkg_display}"
|
||||
)
|
||||
if not _pkg_ready:
|
||||
continue
|
||||
@@ -2154,11 +2176,20 @@ def run_onboard(skip_validation: bool = False) -> bool:
|
||||
provider = _step_provider(config)
|
||||
config.provider = provider
|
||||
|
||||
# Step 2a: Base URL (custom or ollama provider)
|
||||
# Step 2a: Base URL (custom-openai, custom-anthropic, or ollama provider)
|
||||
ollama_detected_models: list[str] = []
|
||||
if provider == "custom":
|
||||
base_url = _step_base_url(config)
|
||||
config.custom_base_url = base_url
|
||||
if provider == "custom-openai":
|
||||
current_base_url = config.custom_openai_base_url or os.environ.get(
|
||||
"CUSTOM_OPENAI_BASE_URL", ""
|
||||
)
|
||||
base_url = _step_base_url(config, current_value=current_base_url)
|
||||
config.custom_openai_base_url = base_url
|
||||
elif provider == "custom-anthropic":
|
||||
current_base_url = config.custom_anthropic_base_url or os.environ.get(
|
||||
"CUSTOM_ANTHROPIC_BASE_URL", ""
|
||||
)
|
||||
base_url = _step_base_url(config, current_value=current_base_url)
|
||||
config.custom_anthropic_base_url = base_url
|
||||
elif provider == "ollama":
|
||||
ollama_url, ollama_detected_models = _step_ollama_base_url(config)
|
||||
config.ollama_base_url = ollama_url
|
||||
@@ -2178,11 +2209,11 @@ def run_onboard(skip_validation: bool = False) -> bool:
|
||||
"openrouter": "openrouter_api_key",
|
||||
"zhipu": "zhipu_api_key",
|
||||
"zhipu-code": "zhipu_api_key",
|
||||
"custom": "custom_api_key",
|
||||
"custom-openai": "custom_openai_api_key",
|
||||
"custom-anthropic": "custom_anthropic_api_key",
|
||||
}
|
||||
_skip_api_key = (
|
||||
provider == "ollama"
|
||||
or (provider == "anthropic" and config.anthropic_auth_mode == "oauth")
|
||||
_skip_api_key = provider == "ollama" or (
|
||||
provider == "anthropic" and config.anthropic_auth_mode == "oauth"
|
||||
)
|
||||
if not _skip_api_key:
|
||||
new_key = _step_provider_api_key(config, provider, skip_validation)
|
||||
|
||||
@@ -68,8 +68,10 @@ class EvoScientistConfig:
|
||||
siliconflow_api_key: str = ""
|
||||
openrouter_api_key: str = ""
|
||||
zhipu_api_key: str = ""
|
||||
custom_api_key: str = ""
|
||||
custom_base_url: str = ""
|
||||
custom_openai_api_key: str = ""
|
||||
custom_openai_base_url: str = ""
|
||||
custom_anthropic_api_key: str = ""
|
||||
custom_anthropic_base_url: str = ""
|
||||
ollama_base_url: str = ""
|
||||
tavily_api_key: str = ""
|
||||
|
||||
@@ -339,8 +341,10 @@ _ENV_MAPPINGS = {
|
||||
"siliconflow_api_key": "SILICONFLOW_API_KEY",
|
||||
"openrouter_api_key": "OPENROUTER_API_KEY",
|
||||
"zhipu_api_key": "ZHIPU_API_KEY",
|
||||
"custom_api_key": "CUSTOM_API_KEY",
|
||||
"custom_base_url": "CUSTOM_BASE_URL",
|
||||
"custom_openai_api_key": "CUSTOM_OPENAI_API_KEY",
|
||||
"custom_openai_base_url": "CUSTOM_OPENAI_BASE_URL",
|
||||
"custom_anthropic_api_key": "CUSTOM_ANTHROPIC_API_KEY",
|
||||
"custom_anthropic_base_url": "CUSTOM_ANTHROPIC_BASE_URL",
|
||||
"ollama_base_url": "OLLAMA_BASE_URL",
|
||||
"tavily_api_key": "TAVILY_API_KEY",
|
||||
"default_mode": "EVOSCIENTIST_DEFAULT_MODE",
|
||||
@@ -416,10 +420,18 @@ def apply_config_to_env(config: EvoScientistConfig) -> None:
|
||||
os.environ["OPENROUTER_API_KEY"] = config.openrouter_api_key
|
||||
if config.zhipu_api_key and not os.environ.get("ZHIPU_API_KEY"):
|
||||
os.environ["ZHIPU_API_KEY"] = config.zhipu_api_key
|
||||
if config.custom_api_key and not os.environ.get("CUSTOM_API_KEY"):
|
||||
os.environ["CUSTOM_API_KEY"] = config.custom_api_key
|
||||
if config.custom_base_url and not os.environ.get("CUSTOM_BASE_URL"):
|
||||
os.environ["CUSTOM_BASE_URL"] = config.custom_base_url
|
||||
if config.custom_openai_api_key and not os.environ.get("CUSTOM_OPENAI_API_KEY"):
|
||||
os.environ["CUSTOM_OPENAI_API_KEY"] = config.custom_openai_api_key
|
||||
if config.custom_openai_base_url and not os.environ.get("CUSTOM_OPENAI_BASE_URL"):
|
||||
os.environ["CUSTOM_OPENAI_BASE_URL"] = config.custom_openai_base_url
|
||||
if config.custom_anthropic_api_key and not os.environ.get(
|
||||
"CUSTOM_ANTHROPIC_API_KEY"
|
||||
):
|
||||
os.environ["CUSTOM_ANTHROPIC_API_KEY"] = config.custom_anthropic_api_key
|
||||
if config.custom_anthropic_base_url and not os.environ.get(
|
||||
"CUSTOM_ANTHROPIC_BASE_URL"
|
||||
):
|
||||
os.environ["CUSTOM_ANTHROPIC_BASE_URL"] = config.custom_anthropic_base_url
|
||||
if config.ollama_base_url and not os.environ.get("OLLAMA_BASE_URL"):
|
||||
os.environ["OLLAMA_BASE_URL"] = config.ollama_base_url
|
||||
if config.tavily_api_key and not os.environ.get("TAVILY_API_KEY"):
|
||||
|
||||
+53
-15
@@ -13,6 +13,7 @@ from typing import Any
|
||||
|
||||
from langchain.chat_models import init_chat_model
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Patch: langchain-anthropic (>=1.3.4) calls .model_dump() on
|
||||
# context_management / container objects returned by the Anthropic SDK.
|
||||
@@ -37,14 +38,16 @@ def _patch_anthropic_proxy_compat() -> None:
|
||||
val = getattr(obj, attr, None)
|
||||
if isinstance(val, dict):
|
||||
d = val.copy()
|
||||
setattr(obj, attr,
|
||||
_types.SimpleNamespace(model_dump=lambda **kw: d))
|
||||
setattr(
|
||||
obj, attr, _types.SimpleNamespace(model_dump=lambda **kw: d)
|
||||
)
|
||||
return _orig(self, event, *args, **kwargs)
|
||||
|
||||
_CA._make_message_chunk_from_anthropic_event = _safe
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
_patch_anthropic_proxy_compat()
|
||||
|
||||
_SILICONFLOW_BASE_URL = "https://api.siliconflow.cn/v1"
|
||||
@@ -59,18 +62,31 @@ _THIRD_PARTY_PROVIDERS: dict[str, tuple[str | None, str]] = {
|
||||
"openrouter": (_OPENROUTER_BASE_URL, "OPENROUTER_API_KEY"),
|
||||
"zhipu": (_ZHIPU_BASE_URL, "ZHIPU_API_KEY"),
|
||||
"zhipu-code": (_ZHIPU_CODE_BASE_URL, "ZHIPU_API_KEY"),
|
||||
"custom": (None, "CUSTOM_API_KEY"), # base_url from CUSTOM_BASE_URL env
|
||||
"custom-openai": (
|
||||
None,
|
||||
"CUSTOM_OPENAI_API_KEY",
|
||||
), # base_url from CUSTOM_OPENAI_BASE_URL env
|
||||
}
|
||||
|
||||
# Model registry: list of (short_name, model_id, provider)
|
||||
# Allows same short_name across different providers.
|
||||
_MODEL_ENTRIES: list[tuple[str, str, str]] = [
|
||||
# Custom Anthropic (third-party Claude-compatible endpoints, 3 defaults)
|
||||
# Listed BEFORE native anthropic so MODELS dict defaults to native provider
|
||||
("claude-sonnet-4-6", "claude-sonnet-4-6", "custom-anthropic"),
|
||||
("claude-sonnet-4-5", "claude-sonnet-4-5", "custom-anthropic"),
|
||||
("claude-haiku-4-5", "claude-haiku-4-5", "custom-anthropic"),
|
||||
# Custom OpenAI (third-party OpenAI-compatible endpoints, 3 defaults)
|
||||
# Listed BEFORE native openai so MODELS dict defaults to native provider
|
||||
("gpt-5.4", "gpt-5.4", "custom-openai"),
|
||||
("gpt-5.3-codex", "gpt-5.3-codex", "custom-openai"),
|
||||
("gpt-5-mini", "gpt-5-mini", "custom-openai"),
|
||||
# Anthropic (ordered by capability)
|
||||
("claude-opus-4-6", "claude-opus-4-6", "anthropic"),
|
||||
("claude-sonnet-4-6", "claude-sonnet-4-6", "anthropic"),
|
||||
("claude-opus-4-5", "claude-opus-4-5-20251101", "anthropic"),
|
||||
("claude-sonnet-4-5", "claude-sonnet-4-5-20250929", "anthropic"),
|
||||
("claude-haiku-4-5", "claude-haiku-4-5-20251001", "anthropic"),
|
||||
("claude-opus-4-5", "claude-opus-4-5", "anthropic"),
|
||||
("claude-sonnet-4-5", "claude-sonnet-4-5", "anthropic"),
|
||||
("claude-haiku-4-5", "claude-haiku-4-5", "anthropic"),
|
||||
# OpenAI
|
||||
("gpt-5.4", "gpt-5.4-2026-03-05", "openai"),
|
||||
("gpt-5.3-codex", "gpt-5.3-codex", "openai"),
|
||||
@@ -82,7 +98,11 @@ _MODEL_ENTRIES: list[tuple[str, str, str]] = [
|
||||
("gpt-5-nano", "gpt-5-nano-2025-08-07", "openai"),
|
||||
# Google GenAI
|
||||
("gemini-3.1-pro", "gemini-3.1-pro-preview", "google-genai"),
|
||||
("gemini-3.1-pro-customtools", "gemini-3.1-pro-preview-customtools", "google-genai"),
|
||||
(
|
||||
"gemini-3.1-pro-customtools",
|
||||
"gemini-3.1-pro-preview-customtools",
|
||||
"google-genai",
|
||||
),
|
||||
("gemini-3.1-flash-lite", "gemini-3.1-flash-lite-preview", "google-genai"),
|
||||
("gemini-3-flash", "gemini-3-flash-preview", "google-genai"),
|
||||
("gemini-2.5-flash", "gemini-2.5-flash", "google-genai"),
|
||||
@@ -106,6 +126,7 @@ _MODEL_ENTRIES: list[tuple[str, str, str]] = [
|
||||
("kimi-k2.5", "Pro/moonshotai/Kimi-K2.5", "siliconflow"),
|
||||
("glm-4.7", "Pro/zai-org/GLM-4.7", "siliconflow"),
|
||||
# OpenRouter
|
||||
("gpt-5.4", "openai/gpt-5.4", "openrouter"),
|
||||
("minimax-m2.5", "minimax/minimax-m2.5", "openrouter"),
|
||||
("grok-4.1-fast", "x-ai/grok-4.1-fast", "openrouter"),
|
||||
("qwen3.5-122b", "qwen/qwen3.5-122b-a10b", "openrouter"),
|
||||
@@ -137,11 +158,7 @@ def get_models_for_provider(provider: str) -> list[tuple[str, str]]:
|
||||
Returns:
|
||||
List of (short_name, model_id) tuples for the provider.
|
||||
"""
|
||||
return [
|
||||
(name, model_id)
|
||||
for name, model_id, p in _MODEL_ENTRIES
|
||||
if p == provider
|
||||
]
|
||||
return [(name, model_id) for name, model_id, p in _MODEL_ENTRIES if p == provider]
|
||||
|
||||
|
||||
def _apply_auto_config(
|
||||
@@ -159,7 +176,7 @@ def _apply_auto_config(
|
||||
if provider == "anthropic" and "thinking" not in kwargs:
|
||||
base_url = os.environ.get("ANTHROPIC_BASE_URL", "")
|
||||
_is_proxy = "127.0.0.1" in base_url or "localhost" in base_url
|
||||
if _is_proxy:
|
||||
if is_third_party or _is_proxy:
|
||||
# ccproxy manages thinking internally; don't set it here
|
||||
# to avoid 422 errors with thinking content blocks in history
|
||||
pass
|
||||
@@ -249,8 +266,15 @@ def get_chat_model(
|
||||
# Third-party providers → route through OpenAI provider with base_url
|
||||
elif provider in _THIRD_PARTY_PROVIDERS:
|
||||
base_url_default, api_key_env = _THIRD_PARTY_PROVIDERS[provider]
|
||||
if provider == "custom":
|
||||
base_url = os.environ.get("CUSTOM_BASE_URL", "")
|
||||
if provider == "custom-openai":
|
||||
base_url = os.environ.get("CUSTOM_OPENAI_BASE_URL", "")
|
||||
if not base_url:
|
||||
raise ValueError(
|
||||
"CUSTOM_OPENAI_BASE_URL environment variable is required when using "
|
||||
"the 'custom-openai' provider. Please set it to your "
|
||||
"OpenAI-compatible API endpoint URL (e.g. https://api.openai.com/v1)."
|
||||
)
|
||||
base_url = base_url.rstrip("/")
|
||||
else:
|
||||
base_url = base_url_default
|
||||
if base_url:
|
||||
@@ -263,6 +287,20 @@ def get_chat_model(
|
||||
if provider == "siliconflow":
|
||||
kwargs.setdefault("extra_body", {})["enable_thinking"] = False
|
||||
provider = "openai"
|
||||
elif provider == "custom-anthropic":
|
||||
base_url = os.environ.get("CUSTOM_ANTHROPIC_BASE_URL", "")
|
||||
if not base_url:
|
||||
raise ValueError(
|
||||
"CUSTOM_ANTHROPIC_BASE_URL environment variable is required when using "
|
||||
"the 'custom-anthropic' provider. Please set it to your "
|
||||
"Anthropic-compatible API endpoint URL (e.g. https://api.anthropic.com)."
|
||||
)
|
||||
kwargs["base_url"] = base_url.rstrip("/")
|
||||
api_key = os.environ.get("CUSTOM_ANTHROPIC_API_KEY", "")
|
||||
if api_key:
|
||||
kwargs["api_key"] = api_key
|
||||
_is_third_party = True # skip thinking in _apply_auto_config
|
||||
provider = "anthropic"
|
||||
elif provider == "ollama":
|
||||
base_url = os.environ.get("OLLAMA_BASE_URL", "")
|
||||
if base_url:
|
||||
|
||||
@@ -555,8 +555,7 @@ def _filter_tools(tools: list, allowed_names: list[str] | None) -> list:
|
||||
|
||||
# Check if any pattern contains wildcard characters
|
||||
has_wildcards = any(
|
||||
any(char in pattern for char in "*?[]")
|
||||
for pattern in allowed_names
|
||||
any(char in pattern for char in "*?[]") for pattern in allowed_names
|
||||
)
|
||||
|
||||
if not has_wildcards:
|
||||
|
||||
@@ -62,23 +62,31 @@ class EvoMemoryState(AgentState):
|
||||
|
||||
evo_memory_content: NotRequired[Annotated[str, PrivateStateAttr]]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Structured extraction schemas
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class UserProfile(BaseModel):
|
||||
"""Extracted user profile information."""
|
||||
|
||||
name: str | None = Field(None, description="User's name")
|
||||
role: str | None = Field(None, description="User's role (e.g. researcher, student)")
|
||||
institution: str | None = Field(None, description="User's institution or organization")
|
||||
institution: str | None = Field(
|
||||
None, description="User's institution or organization"
|
||||
)
|
||||
language: str | None = Field(None, description="User's preferred language")
|
||||
|
||||
|
||||
class ResearchPreferences(BaseModel):
|
||||
"""Extracted research preference information."""
|
||||
|
||||
primary_domain: str | None = Field(None, description="Primary research domain")
|
||||
sub_fields: str | None = Field(None, description="Research sub-fields")
|
||||
preferred_frameworks: str | None = Field(None, description="Preferred software frameworks")
|
||||
preferred_frameworks: str | None = Field(
|
||||
None, description="Preferred software frameworks"
|
||||
)
|
||||
preferred_models: str | None = Field(None, description="Preferred AI/ML models")
|
||||
hardware: str | None = Field(None, description="Available hardware (GPUs, etc.)")
|
||||
constraints: str | None = Field(None, description="Resource or time constraints")
|
||||
@@ -86,6 +94,7 @@ class ResearchPreferences(BaseModel):
|
||||
|
||||
class ExperimentConclusion(BaseModel):
|
||||
"""Extracted experiment conclusion (only when a complete experiment was run)."""
|
||||
|
||||
title: str = Field(description="Experiment name")
|
||||
question: str | None = Field(None, description="Research question")
|
||||
method: str | None = Field(None, description="Method summary")
|
||||
@@ -99,10 +108,19 @@ class ExtractedMemory(BaseModel):
|
||||
|
||||
Only fields with genuinely new information should be populated.
|
||||
"""
|
||||
user_profile: UserProfile | None = Field(None, description="New user profile information")
|
||||
research_preferences: ResearchPreferences | None = Field(None, description="New research preferences")
|
||||
experiment_conclusion: ExperimentConclusion | None = Field(None, description="Completed experiment conclusion")
|
||||
learned_preferences: list[str] | None = Field(None, description="New preferences or habits observed")
|
||||
|
||||
user_profile: UserProfile | None = Field(
|
||||
None, description="New user profile information"
|
||||
)
|
||||
research_preferences: ResearchPreferences | None = Field(
|
||||
None, description="New research preferences"
|
||||
)
|
||||
experiment_conclusion: ExperimentConclusion | None = Field(
|
||||
None, description="Completed experiment conclusion"
|
||||
)
|
||||
learned_preferences: list[str] | None = Field(
|
||||
None, description="New preferences or habits observed"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -289,6 +307,7 @@ def _normalize_item(value: str) -> str:
|
||||
# Helper: merge extracted JSON into MEMORY.md markdown
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _merge_memory(existing_md: str, extracted: dict[str, Any]) -> str:
|
||||
"""Merge extracted fields into the existing MEMORY.md content.
|
||||
|
||||
@@ -338,6 +357,7 @@ def _merge_memory(existing_md: str, extracted: dict[str, Any]) -> str:
|
||||
should_add_exp = bool(exp and isinstance(exp, dict) and exp.get("title"))
|
||||
if should_add_exp:
|
||||
from datetime import datetime
|
||||
|
||||
date_str = datetime.now().strftime("%Y-%m-%d")
|
||||
title = str(exp.get("title", "Untitled")).strip()
|
||||
entry = f"\n### [{date_str}] {title}\n"
|
||||
@@ -353,9 +373,17 @@ def _merge_memory(existing_md: str, extracted: dict[str, Any]) -> str:
|
||||
if exp_start is not None and exp_end is not None:
|
||||
exp_section = result[exp_start:exp_end]
|
||||
exp_lines = [
|
||||
line for line in exp_section.splitlines() if "(No experiments yet)" not in line
|
||||
line
|
||||
for line in exp_section.splitlines()
|
||||
if "(No experiments yet)" not in line
|
||||
]
|
||||
result = result[:exp_start] + "\n" + "\n".join(exp_lines).strip("\n") + "\n" + result[exp_end:]
|
||||
result = (
|
||||
result[:exp_start]
|
||||
+ "\n"
|
||||
+ "\n".join(exp_lines).strip("\n")
|
||||
+ "\n"
|
||||
+ result[exp_end:]
|
||||
)
|
||||
|
||||
# De-duplicate by title if already present
|
||||
if re.search(rf"### \[[0-9-]+\] {re.escape(title)}\b", result):
|
||||
@@ -382,7 +410,8 @@ def _merge_memory(existing_md: str, extracted: dict[str, Any]) -> str:
|
||||
if start is not None and end is not None:
|
||||
section = result[start:end]
|
||||
section_lines = [
|
||||
line for line in section.splitlines()
|
||||
line
|
||||
for line in section.splitlines()
|
||||
if line.strip() and line.strip() not in {"- (none yet)", "- (none)"}
|
||||
]
|
||||
existing_items = {
|
||||
@@ -412,6 +441,7 @@ def _merge_memory(existing_md: str, extracted: dict[str, Any]) -> str:
|
||||
# Middleware
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class EvoMemoryMiddleware(AgentMiddleware):
|
||||
"""Middleware that injects and auto-extracts long-term memory.
|
||||
|
||||
@@ -496,7 +526,11 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
"""Read MEMORY.md content (raw bytes → str)."""
|
||||
try:
|
||||
responses = backend.download_files([self._memory_path])
|
||||
if responses and responses[0].content is not None and responses[0].error is None:
|
||||
if (
|
||||
responses
|
||||
and responses[0].content is not None
|
||||
and responses[0].error is None
|
||||
):
|
||||
return responses[0].content.decode("utf-8")
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.debug("Failed to read memory at %s: %s", self._memory_path, e)
|
||||
@@ -505,13 +539,19 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
async def _aread_memory(self, backend: BackendProtocol) -> str:
|
||||
try:
|
||||
responses = await backend.adownload_files([self._memory_path])
|
||||
if responses and responses[0].content is not None and responses[0].error is None:
|
||||
if (
|
||||
responses
|
||||
and responses[0].content is not None
|
||||
and responses[0].error is None
|
||||
):
|
||||
return responses[0].content.decode("utf-8")
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.debug("Failed to read memory at %s: %s", self._memory_path, e)
|
||||
return ""
|
||||
|
||||
def _write_memory(self, backend: BackendProtocol, old_content: str, new_content: str) -> None:
|
||||
def _write_memory(
|
||||
self, backend: BackendProtocol, old_content: str, new_content: str
|
||||
) -> None:
|
||||
"""Write updated MEMORY.md (edit if exists, write if new)."""
|
||||
try:
|
||||
if old_content:
|
||||
@@ -523,10 +563,14 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("Exception writing memory: %s", e)
|
||||
|
||||
async def _awrite_memory(self, backend: BackendProtocol, old_content: str, new_content: str) -> None:
|
||||
async def _awrite_memory(
|
||||
self, backend: BackendProtocol, old_content: str, new_content: str
|
||||
) -> None:
|
||||
try:
|
||||
if old_content:
|
||||
result = await backend.aedit(self._memory_path, old_content, new_content)
|
||||
result = await backend.aedit(
|
||||
self._memory_path, old_content, new_content
|
||||
)
|
||||
else:
|
||||
result = await backend.awrite(self._memory_path, new_content)
|
||||
if result and result.error:
|
||||
@@ -610,26 +654,34 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
# Fallback for non-Pydantic or unusual model classes
|
||||
return model.bind(**{k: v for k, v in updates.items() if v is not None})
|
||||
|
||||
def _extract(self, model: BaseChatModel, memory: str, messages: list[AnyMessage]) -> dict[str, Any]:
|
||||
def _extract(
|
||||
self, model: BaseChatModel, memory: str, messages: list[AnyMessage]
|
||||
) -> dict[str, Any]:
|
||||
"""Run LLM extraction on recent messages using structured output."""
|
||||
prompt = self._build_extraction_prompt(memory, messages)
|
||||
try:
|
||||
plain_model = self._disable_thinking(model)
|
||||
so_kwargs = self._structured_output_kwargs(plain_model)
|
||||
structured_model = plain_model.with_structured_output(ExtractedMemory, **so_kwargs)
|
||||
structured_model = plain_model.with_structured_output(
|
||||
ExtractedMemory, **so_kwargs
|
||||
)
|
||||
result = structured_model.invoke(prompt)
|
||||
return result.model_dump(exclude_none=True)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("Memory extraction failed: %s", e)
|
||||
return {}
|
||||
|
||||
async def _aextract(self, model: BaseChatModel, memory: str, messages: list[AnyMessage]) -> dict[str, Any]:
|
||||
async def _aextract(
|
||||
self, model: BaseChatModel, memory: str, messages: list[AnyMessage]
|
||||
) -> dict[str, Any]:
|
||||
"""Async: Run LLM extraction on recent messages using structured output."""
|
||||
prompt = self._build_extraction_prompt(memory, messages)
|
||||
try:
|
||||
plain_model = self._disable_thinking(model)
|
||||
so_kwargs = self._structured_output_kwargs(plain_model)
|
||||
structured_model = plain_model.with_structured_output(ExtractedMemory, **so_kwargs)
|
||||
structured_model = plain_model.with_structured_output(
|
||||
ExtractedMemory, **so_kwargs
|
||||
)
|
||||
result = await structured_model.ainvoke(prompt)
|
||||
return result.model_dump(exclude_none=True)
|
||||
except Exception as e: # noqa: BLE001
|
||||
@@ -660,6 +712,7 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
memory_content = "(No memory saved yet. Create `/memory/MEMORY.md` when you learn important information.)"
|
||||
|
||||
from deepagents.middleware._utils import append_to_system_message
|
||||
|
||||
injection = MEMORY_INJECTION_TEMPLATE.format(memory_content=memory_content)
|
||||
new_system = append_to_system_message(request.system_message, injection)
|
||||
return request.override(system_message=new_system)
|
||||
@@ -748,6 +801,7 @@ class EvoMemoryMiddleware(AgentMiddleware):
|
||||
# Factory
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def create_memory_middleware(
|
||||
memory_dir: str | None = None,
|
||||
extraction_model: BaseChatModel | None = None,
|
||||
|
||||
+10
-2
@@ -34,12 +34,20 @@ def set_workspace_root(path: str | Path) -> None:
|
||||
env-var value; all others are re-derived from the new root.
|
||||
Also resets ``_active_workspace`` to the new root as a safe default.
|
||||
"""
|
||||
global WORKSPACE_ROOT, RUNS_DIR, MEMORY_DIR, USER_SKILLS_DIR, MEDIA_DIR, _active_workspace
|
||||
global \
|
||||
WORKSPACE_ROOT, \
|
||||
RUNS_DIR, \
|
||||
MEMORY_DIR, \
|
||||
USER_SKILLS_DIR, \
|
||||
MEDIA_DIR, \
|
||||
_active_workspace
|
||||
WORKSPACE_ROOT = Path(path).resolve()
|
||||
_active_workspace = WORKSPACE_ROOT
|
||||
RUNS_DIR = _env_path("EVOSCIENTIST_RUNS_DIR") or (WORKSPACE_ROOT / "runs")
|
||||
MEMORY_DIR = _env_path("EVOSCIENTIST_MEMORY_DIR") or (WORKSPACE_ROOT / "memory")
|
||||
USER_SKILLS_DIR = _env_path("EVOSCIENTIST_SKILLS_DIR") or (WORKSPACE_ROOT / "skills")
|
||||
USER_SKILLS_DIR = _env_path("EVOSCIENTIST_SKILLS_DIR") or (
|
||||
WORKSPACE_ROOT / "skills"
|
||||
)
|
||||
MEDIA_DIR = _env_path("EVOSCIENTIST_MEDIA_DIR") or (WORKSPACE_ROOT / "media")
|
||||
|
||||
|
||||
|
||||
@@ -330,6 +330,7 @@ Finding one with context [1]. Another insight [2].
|
||||
# Combined exports
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def get_system_prompt() -> str:
|
||||
"""Generate the complete system prompt.
|
||||
|
||||
|
||||
@@ -38,6 +38,7 @@ AGENT_NAME = "EvoScientist"
|
||||
# Paths & ID generation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def get_db_path() -> Path:
|
||||
"""Return ``~/.config/evoscientist/sessions.db``, creating parents."""
|
||||
db_dir = Path.home() / ".config" / "evoscientist"
|
||||
@@ -54,6 +55,7 @@ def generate_thread_id() -> str:
|
||||
# Checkpointer context manager
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def get_checkpointer() -> AsyncIterator[AsyncSqliteSaver]:
|
||||
"""Yield an ``AsyncSqliteSaver`` connected to the global sessions DB."""
|
||||
@@ -65,6 +67,7 @@ async def get_checkpointer() -> AsyncIterator[AsyncSqliteSaver]:
|
||||
# Internal helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def _table_exists(conn: aiosqlite.Connection, table: str) -> bool:
|
||||
query = "SELECT 1 FROM sqlite_master WHERE type='table' AND name=?"
|
||||
async with conn.execute(query, (table,)) as cur:
|
||||
@@ -159,6 +162,7 @@ def _format_relative_time(iso_ts: str | None) -> str:
|
||||
# Thread CRUD
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def list_threads(
|
||||
limit: int = 20,
|
||||
include_message_count: bool = False,
|
||||
|
||||
@@ -19,6 +19,7 @@ import sys
|
||||
# Charset detection (simplified from upstream config.py)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _detect_unicode_support() -> bool:
|
||||
"""Check if the terminal supports Unicode glyphs."""
|
||||
encoding = getattr(sys.stdout, "encoding", "") or ""
|
||||
@@ -30,8 +31,8 @@ def _detect_unicode_support() -> bool:
|
||||
|
||||
# Module-level glyph constants
|
||||
_UNICODE = _detect_unicode_support()
|
||||
GUTTER_BAR = "\u258c" if _UNICODE else "|" # ▌ or |
|
||||
BOX_VERTICAL = "\u2502" if _UNICODE else "|" # │ or |
|
||||
GUTTER_BAR = "\u258c" if _UNICODE else "|" # ▌ or |
|
||||
BOX_VERTICAL = "\u2502" if _UNICODE else "|" # │ or |
|
||||
BOX_DOUBLE_HORIZ = "\u2550" if _UNICODE else "=" # ═ or =
|
||||
|
||||
|
||||
@@ -39,6 +40,7 @@ BOX_DOUBLE_HORIZ = "\u2550" if _UNICODE else "=" # ═ or =
|
||||
# Markup escaping
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _escape_markup(text: str) -> str:
|
||||
"""Escape Rich markup characters in text.
|
||||
|
||||
@@ -51,6 +53,7 @@ def _escape_markup(text: str) -> str:
|
||||
# Diff formatting (produces Rich markup string)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _build_stats_text(additions: int, deletions: int) -> str:
|
||||
"""Build a ``+N -M`` stats string with Rich markup."""
|
||||
parts: list[str] = []
|
||||
@@ -102,7 +105,9 @@ def format_diff_rich(
|
||||
# Title header (═══ title ═══)
|
||||
h = BOX_DOUBLE_HORIZ
|
||||
if title:
|
||||
formatted.append(f"[bold cyan]{h}{h}{h} {_escape_markup(title)} {h}{h}{h}[/bold cyan]")
|
||||
formatted.append(
|
||||
f"[bold cyan]{h}{h}{h} {_escape_markup(title)} {h}{h}{h}[/bold cyan]"
|
||||
)
|
||||
formatted.append("")
|
||||
|
||||
# Stats header
|
||||
@@ -115,9 +120,7 @@ def format_diff_rich(
|
||||
|
||||
for line in lines:
|
||||
if max_lines is not None and line_count >= max_lines:
|
||||
formatted.append(
|
||||
f"\n[dim]... ({len(lines) - line_count} more lines)[/dim]"
|
||||
)
|
||||
formatted.append(f"\n[dim]... ({len(lines) - line_count} more lines)[/dim]")
|
||||
break
|
||||
|
||||
# Skip file headers
|
||||
@@ -147,9 +150,7 @@ def format_diff_rich(
|
||||
new_num += 1
|
||||
line_count += 1
|
||||
elif line.startswith(" "):
|
||||
formatted.append(
|
||||
f"[dim]{BOX_VERTICAL}{old_num:>{width}}[/dim] {escaped}"
|
||||
)
|
||||
formatted.append(f"[dim]{BOX_VERTICAL}{old_num:>{width}}[/dim] {escaped}")
|
||||
old_num += 1
|
||||
new_num += 1
|
||||
line_count += 1
|
||||
@@ -168,6 +169,7 @@ def format_diff_rich(
|
||||
# High-level helper: build diff from edit_file tool args
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def build_edit_diff(
|
||||
file_path: str,
|
||||
old_string: str,
|
||||
@@ -191,14 +193,16 @@ def build_edit_diff(
|
||||
if not old_string and not new_string:
|
||||
return None
|
||||
|
||||
diff_lines = list(difflib.unified_diff(
|
||||
old_string.splitlines(),
|
||||
new_string.splitlines(),
|
||||
fromfile=file_path,
|
||||
tofile=file_path,
|
||||
lineterm="",
|
||||
n=3,
|
||||
))
|
||||
diff_lines = list(
|
||||
difflib.unified_diff(
|
||||
old_string.splitlines(),
|
||||
new_string.splitlines(),
|
||||
fromfile=file_path,
|
||||
tofile=file_path,
|
||||
lineterm="",
|
||||
n=3,
|
||||
)
|
||||
)
|
||||
|
||||
if not diff_lines:
|
||||
return None
|
||||
|
||||
+253
-128
@@ -20,7 +20,13 @@ from rich.text import Text # type: ignore[import-untyped]
|
||||
|
||||
from ..paths import resolve_virtual_path
|
||||
from .formatter import ToolResultFormatter
|
||||
from .state import StreamState, SubAgentState, _build_todo_stats, _parse_todo_items, _INTERNAL_TOOLS
|
||||
from .state import (
|
||||
StreamState,
|
||||
SubAgentState,
|
||||
_build_todo_stats,
|
||||
_parse_todo_items,
|
||||
_INTERNAL_TOOLS,
|
||||
)
|
||||
from .utils import DisplayLimits, ToolStatus, format_tool_compact, is_success
|
||||
from .diff_format import build_edit_diff
|
||||
from .events import stream_agent_events
|
||||
@@ -33,8 +39,8 @@ from .events import stream_agent_events
|
||||
_MEDIA_EXTENSIONS = {".png", ".jpg", ".jpeg", ".gif", ".webp", ".bmp", ".svg", ".pdf"}
|
||||
|
||||
console = Console(
|
||||
legacy_windows=(sys.platform == 'win32'),
|
||||
no_color=os.getenv('NO_COLOR') is not None,
|
||||
legacy_windows=(sys.platform == "win32"),
|
||||
no_color=os.getenv("NO_COLOR") is not None,
|
||||
)
|
||||
|
||||
formatter = ToolResultFormatter()
|
||||
@@ -44,6 +50,7 @@ formatter = ToolResultFormatter()
|
||||
# Todo formatting
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _format_single_todo(item: dict) -> Text:
|
||||
"""Format a single todo item with status symbol."""
|
||||
status = str(item.get("status", "todo")).lower()
|
||||
@@ -77,6 +84,7 @@ def _format_single_todo(item: dict) -> Text:
|
||||
# Tool result formatting
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def format_tool_result_compact(
|
||||
_name: str,
|
||||
content: str,
|
||||
@@ -148,12 +156,13 @@ def format_tool_result_compact(
|
||||
# Tool call line rendering
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _render_tool_call_line(tc: dict, tr: dict | None) -> Text:
|
||||
"""Render a single tool call line with status indicator."""
|
||||
is_task = tc.get('name', '').lower() == 'task'
|
||||
is_task = tc.get("name", "").lower() == "task"
|
||||
|
||||
if tr is not None:
|
||||
content = tr.get('content', '')
|
||||
content = tr.get("content", "")
|
||||
if is_success(content):
|
||||
style = "bold green"
|
||||
indicator = "\u2713" if is_task else ToolStatus.SUCCESS.value
|
||||
@@ -165,19 +174,19 @@ def _render_tool_call_line(tc: dict, tr: dict | None) -> Text:
|
||||
indicator = "\u25b6" if is_task else ToolStatus.RUNNING.value
|
||||
|
||||
# Try to get display name from args first
|
||||
tool_compact = format_tool_compact(tc['name'], tc.get('args'))
|
||||
tool_compact = format_tool_compact(tc["name"], tc.get("args"))
|
||||
|
||||
# If args were empty and we have a result, try to infer memory operations from result
|
||||
tool_name = tc.get('name', '').lower()
|
||||
if tool_name in ('write_file', 'edit_file') and tr is not None:
|
||||
result_content = tr.get('content', '')
|
||||
if '/MEMORY.md' in result_content or 'MEMORY.md' in result_content:
|
||||
tool_name = tc.get("name", "").lower()
|
||||
if tool_name in ("write_file", "edit_file") and tr is not None:
|
||||
result_content = tr.get("content", "")
|
||||
if "/MEMORY.md" in result_content or "MEMORY.md" in result_content:
|
||||
tool_compact = "Updating memory"
|
||||
elif tool_name == 'read_file' and tr is not None:
|
||||
result_content = tr.get('content', '')
|
||||
elif tool_name == "read_file" and tr is not None:
|
||||
result_content = tr.get("content", "")
|
||||
# read_file result doesn't contain path, check if args is empty and result looks like memory
|
||||
args = tc.get('args') or {}
|
||||
if not args.get('path') and '# EvoScientist Memory' in result_content:
|
||||
args = tc.get("args") or {}
|
||||
if not args.get("path") and "# EvoScientist Memory" in result_content:
|
||||
tool_compact = "Reading memory"
|
||||
|
||||
tool_text = Text()
|
||||
@@ -190,7 +199,8 @@ def _render_tool_call_line(tc: dict, tr: dict | None) -> Text:
|
||||
# Sub-agent section rendering
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _render_subagent_section(sa: 'SubAgentState', compact: bool = False) -> list:
|
||||
|
||||
def _render_subagent_section(sa: "SubAgentState", compact: bool = False) -> list:
|
||||
"""Render a sub-agent's activity as a bordered section.
|
||||
|
||||
Args:
|
||||
@@ -241,8 +251,8 @@ def _render_subagent_section(sa: 'SubAgentState', compact: bool = False) -> list
|
||||
return elements
|
||||
|
||||
# --- Full mode: bordered section for Live streaming ---
|
||||
MAX_SA_VISIBLE = 3 # max completed tools shown
|
||||
MAX_SA_RUNNING = 2 # max running tools shown
|
||||
MAX_SA_VISIBLE = 3 # max completed tools shown
|
||||
MAX_SA_RUNNING = 2 # max running tools shown
|
||||
|
||||
# Header
|
||||
header = Text()
|
||||
@@ -255,7 +265,11 @@ def _render_subagent_section(sa: 'SubAgentState', compact: bool = False) -> list
|
||||
|
||||
# Completed tools — collapse older ones into a summary
|
||||
slots = max(0, MAX_SA_VISIBLE - len(pending))
|
||||
hidden = completed[:-slots] if slots and len(completed) > slots else (completed if not slots else [])
|
||||
hidden = (
|
||||
completed[:-slots]
|
||||
if slots and len(completed) > slots
|
||||
else (completed if not slots else [])
|
||||
)
|
||||
visible = completed[-slots:] if slots else []
|
||||
|
||||
if hidden:
|
||||
@@ -288,7 +302,9 @@ def _render_subagent_section(sa: 'SubAgentState', compact: bool = False) -> list
|
||||
hidden_running = len(pending) - MAX_SA_RUNNING
|
||||
if hidden_running > 0:
|
||||
run_summary = Text("\u2502 ", style=BORDER)
|
||||
run_summary.append(f"\u25cf {hidden_running} more running...", style="dim yellow")
|
||||
run_summary.append(
|
||||
f"\u25cf {hidden_running} more running...", style="dim yellow"
|
||||
)
|
||||
elements.append(run_summary)
|
||||
pending = pending[-MAX_SA_RUNNING:]
|
||||
|
||||
@@ -317,6 +333,7 @@ def _render_subagent_section(sa: 'SubAgentState', compact: bool = False) -> list
|
||||
# Todo panel
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _render_todo_panel(todo_items: list[dict]) -> Panel:
|
||||
"""Render a bordered Task List panel from todo items.
|
||||
|
||||
@@ -355,6 +372,7 @@ def _render_todo_panel(todo_items: list[dict]) -> Panel:
|
||||
# Streaming display layout
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def create_streaming_display(
|
||||
thinking_text: str = "",
|
||||
response_text: str = "",
|
||||
@@ -392,7 +410,7 @@ def create_streaming_display(
|
||||
return Group(*elements)
|
||||
|
||||
# Thinking panel
|
||||
_show_thinking = (final_show_thinking if is_final else show_thinking)
|
||||
_show_thinking = final_show_thinking if is_final else show_thinking
|
||||
if _show_thinking and thinking_text:
|
||||
thinking_title = "Thinking"
|
||||
display_thinking = thinking_text.rstrip()
|
||||
@@ -400,18 +418,26 @@ def create_streaming_display(
|
||||
# Final frame: middle-elision truncation
|
||||
if len(display_thinking) > final_thinking_max_length:
|
||||
half = final_thinking_max_length // 2
|
||||
display_thinking = display_thinking[:half] + "\n\n... (truncated) ...\n\n" + display_thinking[-half:]
|
||||
display_thinking = (
|
||||
display_thinking[:half]
|
||||
+ "\n\n... (truncated) ...\n\n"
|
||||
+ display_thinking[-half:]
|
||||
)
|
||||
else:
|
||||
if is_thinking:
|
||||
thinking_title += " ..."
|
||||
if len(display_thinking) > DisplayLimits.THINKING_STREAM:
|
||||
display_thinking = "..." + display_thinking[-DisplayLimits.THINKING_STREAM:]
|
||||
elements.append(Panel(
|
||||
Text(display_thinking, style="dim"),
|
||||
title=thinking_title,
|
||||
border_style="blue",
|
||||
padding=(0, 1),
|
||||
))
|
||||
display_thinking = (
|
||||
"..." + display_thinking[-DisplayLimits.THINKING_STREAM :]
|
||||
)
|
||||
elements.append(
|
||||
Panel(
|
||||
Text(display_thinking, style="dim"),
|
||||
title=thinking_title,
|
||||
border_style="blue",
|
||||
padding=(0, 1),
|
||||
)
|
||||
)
|
||||
|
||||
# Summarization panel (context was compressed by LangGraph middleware)
|
||||
if summarization_text:
|
||||
@@ -420,12 +446,14 @@ def create_streaming_display(
|
||||
char_label = f"{n / 1000:.1f}k chars" if n >= 1000 else f"{n:,} chars"
|
||||
if n > 300:
|
||||
summary_display = summary_display[:300] + " ..."
|
||||
elements.append(Panel(
|
||||
Text(summary_display, style="dim italic"),
|
||||
title=f"Context Summarized ({char_label})",
|
||||
border_style="#f59e0b",
|
||||
padding=(0, 1),
|
||||
))
|
||||
elements.append(
|
||||
Panel(
|
||||
Text(summary_display, style="dim italic"),
|
||||
title=f"Context Summarized ({char_label})",
|
||||
border_style="#f59e0b",
|
||||
padding=(0, 1),
|
||||
)
|
||||
)
|
||||
|
||||
# Tool calls and results paired display
|
||||
# Collapse older completed tools to prevent overflow in Live mode
|
||||
@@ -435,22 +463,22 @@ def create_streaming_display(
|
||||
|
||||
if tool_calls:
|
||||
# Split into categories
|
||||
completed_regular = [] # completed non-task tools
|
||||
task_tools = [] # task tools (always visible)
|
||||
running_regular = [] # running non-task tools
|
||||
completed_regular = [] # completed non-task tools
|
||||
task_tools = [] # task tools (always visible)
|
||||
running_regular = [] # running non-task tools
|
||||
|
||||
for i, tc in enumerate(tool_calls):
|
||||
has_result = i < len(tool_results)
|
||||
tr = tool_results[i] if has_result else None
|
||||
is_task = tc.get('name') == 'task'
|
||||
is_task = tc.get("name") == "task"
|
||||
|
||||
# Skip internal middleware tools
|
||||
if tc.get('name') in _INTERNAL_TOOLS:
|
||||
if tc.get("name") in _INTERNAL_TOOLS:
|
||||
continue
|
||||
|
||||
if is_task:
|
||||
# Skip task calls with empty args (still streaming)
|
||||
if tc.get('args'):
|
||||
if tc.get("args"):
|
||||
task_tools.append((tc, tr))
|
||||
elif has_result:
|
||||
completed_regular.append((tc, tr))
|
||||
@@ -463,22 +491,26 @@ def create_streaming_display(
|
||||
|
||||
for tc, tr in completed_regular:
|
||||
elements.append(_render_tool_call_line(tc, tr))
|
||||
content = tr.get('content', '') if tr else ''
|
||||
if tr and (not is_success(content) or tc.get('name') == 'edit_file'):
|
||||
content = tr.get("content", "") if tr else ""
|
||||
if tr and (not is_success(content) or tc.get("name") == "edit_file"):
|
||||
result_elements = format_tool_result_compact(
|
||||
tr['name'], content, max_lines=10,
|
||||
tool_args=tc.get('args'),
|
||||
tr["name"],
|
||||
content,
|
||||
max_lines=10,
|
||||
tool_args=tc.get("args"),
|
||||
)
|
||||
elements.extend(result_elements)
|
||||
|
||||
# Task tools with compact sub-agent summaries
|
||||
for tc, tr in task_tools:
|
||||
elements.append(_render_tool_call_line(tc, tr))
|
||||
sa_name = tc.get('args', {}).get('subagent_type', '')
|
||||
task_desc = tc.get('args', {}).get('description', '')
|
||||
sa_name = tc.get("args", {}).get("subagent_type", "")
|
||||
task_desc = tc.get("args", {}).get("description", "")
|
||||
matched_sa = None
|
||||
for sa in subagents:
|
||||
if sa.name == sa_name or (task_desc and task_desc in (sa.description or '')):
|
||||
if sa.name == sa_name or (
|
||||
task_desc and task_desc in (sa.description or "")
|
||||
):
|
||||
matched_sa = sa
|
||||
break
|
||||
if matched_sa:
|
||||
@@ -494,11 +526,15 @@ def create_streaming_display(
|
||||
# Streaming mode: collapse older tools, show spinners
|
||||
# --- Completed regular tools (collapsible) ---
|
||||
slots = max(0, MAX_VISIBLE_TOOLS - len(running_regular))
|
||||
hidden = completed_regular[:-slots] if slots and len(completed_regular) > slots else (completed_regular if not slots else [])
|
||||
hidden = (
|
||||
completed_regular[:-slots]
|
||||
if slots and len(completed_regular) > slots
|
||||
else (completed_regular if not slots else [])
|
||||
)
|
||||
visible = completed_regular[-slots:] if slots else []
|
||||
|
||||
if hidden:
|
||||
ok = sum(1 for _, tr in hidden if is_success(tr.get('content', '')))
|
||||
ok = sum(1 for _, tr in hidden if is_success(tr.get("content", "")))
|
||||
fail = len(hidden) - ok
|
||||
summary = Text()
|
||||
summary.append(f"\u2713 {ok} completed", style="dim green")
|
||||
@@ -508,11 +544,13 @@ def create_streaming_display(
|
||||
|
||||
for tc, tr in visible:
|
||||
elements.append(_render_tool_call_line(tc, tr))
|
||||
content = tr.get('content', '') if tr else ''
|
||||
if tr and (not is_success(content) or tc.get('name') == 'edit_file'):
|
||||
content = tr.get("content", "") if tr else ""
|
||||
if tr and (not is_success(content) or tc.get("name") == "edit_file"):
|
||||
result_elements = format_tool_result_compact(
|
||||
tr['name'], content, max_lines=5,
|
||||
tool_args=tc.get('args'),
|
||||
tr["name"],
|
||||
content,
|
||||
max_lines=5,
|
||||
tool_args=tc.get("args"),
|
||||
)
|
||||
elements.extend(result_elements)
|
||||
|
||||
@@ -520,7 +558,9 @@ def create_streaming_display(
|
||||
hidden_running = len(running_regular) - MAX_VISIBLE_RUNNING
|
||||
if hidden_running > 0:
|
||||
summary = Text()
|
||||
summary.append(f"\u25cf {hidden_running} more running...", style="dim yellow")
|
||||
summary.append(
|
||||
f"\u25cf {hidden_running} more running...", style="dim yellow"
|
||||
)
|
||||
elements.append(summary)
|
||||
running_regular = running_regular[-MAX_VISIBLE_RUNNING:]
|
||||
|
||||
@@ -535,7 +575,7 @@ def create_streaming_display(
|
||||
_n_visible = 0
|
||||
_n_visible_done = 0
|
||||
for i, tc in enumerate(tool_calls):
|
||||
if tc.get('name') in _INTERNAL_TOOLS:
|
||||
if tc.get("name") in _INTERNAL_TOOLS:
|
||||
continue
|
||||
_n_visible += 1
|
||||
if i < len(tool_results):
|
||||
@@ -598,11 +638,18 @@ def create_streaming_display(
|
||||
elements.extend(_render_subagent_section(sa, compact=not sa.is_active))
|
||||
|
||||
# Processing state after tool execution
|
||||
if is_processing and not is_thinking and not is_responding and not response_text:
|
||||
if (
|
||||
is_processing
|
||||
and not is_thinking
|
||||
and not is_responding
|
||||
and not response_text
|
||||
):
|
||||
# Check if any sub-agent is active
|
||||
any_active = any(sa.is_active for sa in subagents)
|
||||
if not any_active:
|
||||
elements.append(Spinner("dots", text=" Analyzing results...", style="cyan"))
|
||||
elements.append(
|
||||
Spinner("dots", text=" Analyzing results...", style="cyan")
|
||||
)
|
||||
|
||||
# Stream response in real-time as tokens arrive (all tools done)
|
||||
if response_text and all_done:
|
||||
@@ -618,6 +665,7 @@ def create_streaming_display(
|
||||
# Final results display
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def display_final_results(
|
||||
state: StreamState,
|
||||
thinking_max_length: int = DisplayLimits.THINKING_FINAL,
|
||||
@@ -629,22 +677,30 @@ def display_final_results(
|
||||
display_thinking = state.thinking_text.rstrip()
|
||||
if len(display_thinking) > thinking_max_length:
|
||||
half = thinking_max_length // 2
|
||||
display_thinking = display_thinking[:half] + "\n\n... (truncated) ...\n\n" + display_thinking[-half:]
|
||||
console.print(Panel(
|
||||
Text(display_thinking, style="dim"),
|
||||
title="Thinking",
|
||||
border_style="blue",
|
||||
))
|
||||
display_thinking = (
|
||||
display_thinking[:half]
|
||||
+ "\n\n... (truncated) ...\n\n"
|
||||
+ display_thinking[-half:]
|
||||
)
|
||||
console.print(
|
||||
Panel(
|
||||
Text(display_thinking, style="dim"),
|
||||
title="Thinking",
|
||||
border_style="blue",
|
||||
)
|
||||
)
|
||||
|
||||
if state.summarization_text:
|
||||
summary_display = state.summarization_text.rstrip()
|
||||
if len(summary_display) > 500:
|
||||
summary_display = summary_display[:500] + " ..."
|
||||
console.print(Panel(
|
||||
Text(summary_display, style="dim italic"),
|
||||
title="Context Summarized",
|
||||
border_style="#f59e0b",
|
||||
))
|
||||
console.print(
|
||||
Panel(
|
||||
Text(summary_display, style="dim italic"),
|
||||
title="Context Summarized",
|
||||
border_style="#f59e0b",
|
||||
)
|
||||
)
|
||||
|
||||
if show_tools and state.tool_calls:
|
||||
shown_sa_names: set[str] = set()
|
||||
@@ -652,9 +708,9 @@ def display_final_results(
|
||||
for i, tc in enumerate(state.tool_calls):
|
||||
has_result = i < len(state.tool_results)
|
||||
tr = state.tool_results[i] if has_result else None
|
||||
content = tr.get('content', '') if tr is not None else ''
|
||||
tool_name = tc.get('name', '')
|
||||
is_task = tool_name.lower() == 'task'
|
||||
content = tr.get("content", "") if tr is not None else ""
|
||||
tool_name = tc.get("name", "")
|
||||
is_task = tool_name.lower() == "task"
|
||||
|
||||
# Skip internal middleware tools
|
||||
if tool_name in _INTERNAL_TOOLS:
|
||||
@@ -663,11 +719,13 @@ def display_final_results(
|
||||
# Task tools: show delegation line + compact sub-agent summary
|
||||
if is_task:
|
||||
console.print(_render_tool_call_line(tc, tr))
|
||||
sa_name = tc.get('args', {}).get('subagent_type', '')
|
||||
task_desc = tc.get('args', {}).get('description', '')
|
||||
sa_name = tc.get("args", {}).get("subagent_type", "")
|
||||
task_desc = tc.get("args", {}).get("description", "")
|
||||
matched_sa = None
|
||||
for sa in state.subagents:
|
||||
if sa.name == sa_name or (task_desc and task_desc in (sa.description or '')):
|
||||
if sa.name == sa_name or (
|
||||
task_desc and task_desc in (sa.description or "")
|
||||
):
|
||||
matched_sa = sa
|
||||
break
|
||||
if matched_sa:
|
||||
@@ -680,10 +738,10 @@ def display_final_results(
|
||||
console.print(_render_tool_call_line(tc, tr))
|
||||
if has_result and tr is not None:
|
||||
result_elements = format_tool_result_compact(
|
||||
tr['name'],
|
||||
tr["name"],
|
||||
content,
|
||||
max_lines=10,
|
||||
tool_args=tc.get('args'),
|
||||
tool_args=tc.get("args"),
|
||||
)
|
||||
for elem in result_elements:
|
||||
console.print(elem)
|
||||
@@ -764,6 +822,7 @@ def _resolve_hitl_approval(
|
||||
|
||||
# Config-level auto-approve
|
||||
from ..config.settings import load_config
|
||||
|
||||
cfg = load_config()
|
||||
if cfg.auto_approve:
|
||||
return [{"type": "approve"} for _ in action_requests]
|
||||
@@ -771,13 +830,18 @@ def _resolve_hitl_approval(
|
||||
# Per-tool auto-approval: only execute needs manual approval
|
||||
shell_allow_list = (
|
||||
[s.strip() for s in cfg.shell_allow_list.split(",") if s.strip()]
|
||||
if cfg.shell_allow_list else []
|
||||
if cfg.shell_allow_list
|
||||
else []
|
||||
)
|
||||
|
||||
needs_prompt = False
|
||||
for req in action_requests:
|
||||
name = req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "")
|
||||
args = req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {})
|
||||
name = (
|
||||
req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "")
|
||||
)
|
||||
args = (
|
||||
req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {})
|
||||
)
|
||||
|
||||
if name != "execute":
|
||||
continue # Non-execute tools auto-approve
|
||||
@@ -807,21 +871,29 @@ def _prompt_hitl_approval(action_requests: list) -> list[dict] | None:
|
||||
console.print()
|
||||
panel_text = Text()
|
||||
for i, req in enumerate(action_requests):
|
||||
name = req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "")
|
||||
args = req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {})
|
||||
name = (
|
||||
req.get("name", "") if isinstance(req, dict) else getattr(req, "name", "")
|
||||
)
|
||||
args = (
|
||||
req.get("args", {}) if isinstance(req, dict) else getattr(req, "args", {})
|
||||
)
|
||||
desc = format_tool_compact(name, args if isinstance(args, dict) else {})
|
||||
if panel_text.plain:
|
||||
panel_text.append("\n")
|
||||
panel_text.append(f" {i + 1}. {desc}", style="yellow")
|
||||
panel_text.append("\n\n")
|
||||
panel_text.append(" [1] Approve [2] Reject [3] Approve all (session)", style="dim")
|
||||
panel_text.append(
|
||||
" [1] Approve [2] Reject [3] Approve all (session)", style="dim"
|
||||
)
|
||||
|
||||
console.print(Panel(
|
||||
panel_text,
|
||||
title="Approval Required",
|
||||
border_style="yellow",
|
||||
padding=(0, 1),
|
||||
))
|
||||
console.print(
|
||||
Panel(
|
||||
panel_text,
|
||||
title="Approval Required",
|
||||
border_style="yellow",
|
||||
padding=(0, 1),
|
||||
)
|
||||
)
|
||||
|
||||
try:
|
||||
choice = input(" Choose [1/2/3, Enter=Approve]: ").strip() or "1"
|
||||
@@ -843,6 +915,7 @@ def _prompt_hitl_approval(action_requests: list) -> list[dict] | None:
|
||||
# Async-to-sync bridge
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _create_event_loop() -> asyncio.AbstractEventLoop:
|
||||
"""Create and set the event loop for asyncio.
|
||||
|
||||
@@ -853,6 +926,7 @@ def _create_event_loop() -> asyncio.AbstractEventLoop:
|
||||
asyncio.set_event_loop(loop)
|
||||
return loop
|
||||
|
||||
|
||||
def _get_event_loop() -> asyncio.AbstractEventLoop:
|
||||
"""Get the event loop for asyncio.
|
||||
|
||||
@@ -866,6 +940,7 @@ def _get_event_loop() -> asyncio.AbstractEventLoop:
|
||||
loop = _create_event_loop()
|
||||
return loop
|
||||
|
||||
|
||||
def _resolve_ask_user_prompt(ask_user_data: dict) -> dict:
|
||||
"""Interactive console Q&A for ask_user events.
|
||||
|
||||
@@ -880,11 +955,13 @@ def _resolve_ask_user_prompt(ask_user_data: dict) -> dict:
|
||||
return {"answers": [], "status": "answered"}
|
||||
|
||||
console.print()
|
||||
console.print(Panel(
|
||||
Text("Quick check-in from EvoScientist", style="bold"),
|
||||
border_style="cyan",
|
||||
padding=(0, 1),
|
||||
))
|
||||
console.print(
|
||||
Panel(
|
||||
Text("Quick check-in from EvoScientist", style="bold"),
|
||||
border_style="cyan",
|
||||
padding=(0, 1),
|
||||
)
|
||||
)
|
||||
console.print()
|
||||
|
||||
answers: list[str] = []
|
||||
@@ -903,12 +980,18 @@ def _resolve_ask_user_prompt(ask_user_data: dict) -> dict:
|
||||
letter = chr(ord("A") + j)
|
||||
console.print(Text(f" {letter}. {label}", style="dim"))
|
||||
other_letter = chr(ord("A") + len(choices))
|
||||
console.print(Text(f" {other_letter}. Other (type your answer)", style="dim"))
|
||||
console.print(
|
||||
Text(f" {other_letter}. Other (type your answer)", style="dim")
|
||||
)
|
||||
|
||||
letters = "/".join(chr(ord("A") + k) for k in range(len(choices) + 1))
|
||||
raw = pt_prompt(HTML(f" <b><style fg='#1565c0'>Choice [{letters}]:</style></b> ")).strip()
|
||||
raw = pt_prompt(
|
||||
HTML(f" <b><style fg='#1565c0'>Choice [{letters}]:</style></b> ")
|
||||
).strip()
|
||||
if raw.upper() == other_letter:
|
||||
raw = pt_prompt(HTML(" <b><style fg='#42a5f5'>> Your answer:</style></b> ")).strip()
|
||||
raw = pt_prompt(
|
||||
HTML(" <b><style fg='#42a5f5'>> Your answer:</style></b> ")
|
||||
).strip()
|
||||
answers.append(raw)
|
||||
elif len(raw) == 1 and raw.upper().isalpha():
|
||||
idx = ord(raw.upper()) - ord("A")
|
||||
@@ -919,7 +1002,9 @@ def _resolve_ask_user_prompt(ask_user_data: dict) -> dict:
|
||||
else:
|
||||
answers.append(raw)
|
||||
else:
|
||||
raw = pt_prompt(HTML(" <b><style fg='#42a5f5'>> Answer:</style></b> ")).strip()
|
||||
raw = pt_prompt(
|
||||
HTML(" <b><style fg='#42a5f5'>> Answer:</style></b> ")
|
||||
).strip()
|
||||
answers.append(raw)
|
||||
console.print()
|
||||
except (EOFError, KeyboardInterrupt):
|
||||
@@ -970,6 +1055,7 @@ def _run_streaming(
|
||||
The final response text.
|
||||
"""
|
||||
import nest_asyncio
|
||||
|
||||
nest_asyncio.apply()
|
||||
|
||||
state = _state if _state is not None else StreamState()
|
||||
@@ -981,36 +1067,49 @@ def _run_streaming(
|
||||
|
||||
async def _consume() -> None:
|
||||
nonlocal _thinking_sent, _todo_sent
|
||||
async for event in stream_agent_events(agent, message, thread_id, metadata=metadata):
|
||||
async for event in stream_agent_events(
|
||||
agent, message, thread_id, metadata=metadata
|
||||
):
|
||||
event_type = state.handle_event(event)
|
||||
|
||||
# Send thinking to channel when transitioning away from thinking
|
||||
if (on_thinking and not _thinking_sent
|
||||
and state.thinking_text
|
||||
and event_type != "thinking"
|
||||
and len(state.thinking_text) >= _MIN_THINKING_LEN):
|
||||
if (
|
||||
on_thinking
|
||||
and not _thinking_sent
|
||||
and state.thinking_text
|
||||
and event_type != "thinking"
|
||||
and len(state.thinking_text) >= _MIN_THINKING_LEN
|
||||
):
|
||||
on_thinking(state.thinking_text.rstrip())
|
||||
_thinking_sent = True
|
||||
|
||||
# Send todo list to channel on first write_todos tool_call
|
||||
if (on_todo and not _todo_sent
|
||||
and event_type == "tool_call"
|
||||
and event.get("name") == "write_todos"
|
||||
and state.todo_items):
|
||||
if (
|
||||
on_todo
|
||||
and not _todo_sent
|
||||
and event_type == "tool_call"
|
||||
and event.get("name") == "write_todos"
|
||||
and state.todo_items
|
||||
):
|
||||
# Flush thinking before todo if not sent yet
|
||||
if (on_thinking and not _thinking_sent
|
||||
and state.thinking_text
|
||||
and len(state.thinking_text) >= _MIN_THINKING_LEN):
|
||||
if (
|
||||
on_thinking
|
||||
and not _thinking_sent
|
||||
and state.thinking_text
|
||||
and len(state.thinking_text) >= _MIN_THINKING_LEN
|
||||
):
|
||||
on_thinking(state.thinking_text.rstrip())
|
||||
_thinking_sent = True
|
||||
on_todo(state.todo_items)
|
||||
_todo_sent = True
|
||||
|
||||
# Send media file to channel when write_file succeeds
|
||||
if (on_file_write
|
||||
and event_type == "tool_result"
|
||||
and event.get("name") == "write_file"
|
||||
and event.get("success")):
|
||||
if (
|
||||
on_file_write
|
||||
and event_type == "tool_result"
|
||||
and event.get("name") == "write_file"
|
||||
and event.get("success")
|
||||
):
|
||||
wf_path = ""
|
||||
for tc in reversed(state.tool_calls):
|
||||
if tc.get("name") == "write_file":
|
||||
@@ -1027,14 +1126,18 @@ def _run_streaming(
|
||||
on_file_write(real_path)
|
||||
|
||||
# Send media file to channel when read_file returns an image
|
||||
if (on_file_write
|
||||
and event_type == "tool_result"
|
||||
and event.get("name") == "read_file"
|
||||
and event.get("success")):
|
||||
if (
|
||||
on_file_write
|
||||
and event_type == "tool_result"
|
||||
and event.get("name") == "read_file"
|
||||
and event.get("success")
|
||||
):
|
||||
rf_path = ""
|
||||
for tc in reversed(state.tool_calls):
|
||||
if tc.get("name") == "read_file":
|
||||
p = tc.get("args", {}).get("file_path", "") or tc.get("args", {}).get("path", "")
|
||||
p = tc.get("args", {}).get("file_path", "") or tc.get(
|
||||
"args", {}
|
||||
).get("path", "")
|
||||
if p and p not in _media_sent:
|
||||
rf_path = p
|
||||
break
|
||||
@@ -1048,13 +1151,20 @@ def _run_streaming(
|
||||
_media_sent.add(rf_path)
|
||||
on_file_write(real_path)
|
||||
|
||||
live.update(create_streaming_display(
|
||||
**state.get_display_args(),
|
||||
show_thinking=show_thinking,
|
||||
response_markdown=state.get_response_markdown(),
|
||||
))
|
||||
live.update(
|
||||
create_streaming_display(
|
||||
**state.get_display_args(),
|
||||
show_thinking=show_thinking,
|
||||
response_markdown=state.get_response_markdown(),
|
||||
)
|
||||
)
|
||||
|
||||
with Live(console=console, auto_refresh=False, transient=False, vertical_overflow="visible") as live:
|
||||
with Live(
|
||||
console=console,
|
||||
auto_refresh=False,
|
||||
transient=False,
|
||||
vertical_overflow="visible",
|
||||
) as live:
|
||||
live.update(create_streaming_display(is_waiting=True))
|
||||
try:
|
||||
loop = _get_event_loop()
|
||||
@@ -1081,7 +1191,10 @@ def _run_streaming(
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
# Render clean final frame before Live exits (no spinners, expanded tools)
|
||||
if state.pending_interrupt is not None or state.pending_ask_user is not None:
|
||||
if (
|
||||
state.pending_interrupt is not None
|
||||
or state.pending_ask_user is not None
|
||||
):
|
||||
# Interrupted: render current state (not final) so it
|
||||
# looks continuous when prompt appears.
|
||||
final_display = create_streaming_display(
|
||||
@@ -1123,6 +1236,7 @@ def _run_streaming(
|
||||
else:
|
||||
result = _resolve_ask_user_prompt(state.pending_ask_user)
|
||||
from langgraph.types import Command # type: ignore[import-untyped]
|
||||
|
||||
state.pending_ask_user = None
|
||||
return _run_streaming(
|
||||
agent=agent,
|
||||
@@ -1144,10 +1258,12 @@ def _run_streaming(
|
||||
# HITL: check for pending interrupt and handle approval
|
||||
if state.pending_interrupt is not None and _hitl_depth < _MAX_HITL_ITERATIONS:
|
||||
decisions = _resolve_hitl_approval(
|
||||
state.pending_interrupt, prompt_fn=hitl_prompt_fn,
|
||||
state.pending_interrupt,
|
||||
prompt_fn=hitl_prompt_fn,
|
||||
)
|
||||
if decisions is not None:
|
||||
from langgraph.types import Command # type: ignore[import-untyped]
|
||||
|
||||
state.pending_interrupt = None
|
||||
return _run_streaming(
|
||||
agent=agent,
|
||||
@@ -1181,6 +1297,7 @@ def _run_streaming(
|
||||
# Thread-safe static streaming (for background channels)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def _astream_to_console(
|
||||
agent: Any,
|
||||
message: str,
|
||||
@@ -1231,14 +1348,22 @@ async def _astream_to_console(
|
||||
dt = state.thinking_text.rstrip()
|
||||
if len(dt) > 500:
|
||||
dt = dt[:250] + "\n\u2026truncated\u2026\n" + dt[-250:]
|
||||
console.print(Panel(Text(dt, style="dim"), title="Thinking", border_style="blue"))
|
||||
console.print(
|
||||
Panel(Text(dt, style="dim"), title="Thinking", border_style="blue")
|
||||
)
|
||||
|
||||
# Summarization
|
||||
if state.summarization_text:
|
||||
st = state.summarization_text.rstrip()
|
||||
if len(st) > 500:
|
||||
st = st[:500] + " ..."
|
||||
console.print(Panel(Text(st, style="dim italic"), title="Context Summarized", border_style="#f59e0b"))
|
||||
console.print(
|
||||
Panel(
|
||||
Text(st, style="dim italic"),
|
||||
title="Context Summarized",
|
||||
border_style="#f59e0b",
|
||||
)
|
||||
)
|
||||
|
||||
# 1) Regular (non-task) tools — above Task List
|
||||
for i, tc in enumerate(state.tool_calls):
|
||||
|
||||
@@ -11,6 +11,7 @@ from typing import Any, Dict
|
||||
@dataclass
|
||||
class StreamEvent:
|
||||
"""Unified stream event."""
|
||||
|
||||
type: str
|
||||
data: Dict[str, Any]
|
||||
|
||||
@@ -21,7 +22,9 @@ class StreamEventEmitter:
|
||||
@staticmethod
|
||||
def thinking(content: str, thinking_id: int = 0) -> StreamEvent:
|
||||
"""Thinking content event."""
|
||||
return StreamEvent("thinking", {"type": "thinking", "content": content, "id": thinking_id})
|
||||
return StreamEvent(
|
||||
"thinking", {"type": "thinking", "content": content, "id": thinking_id}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def text(content: str) -> StreamEvent:
|
||||
@@ -31,52 +34,67 @@ class StreamEventEmitter:
|
||||
@staticmethod
|
||||
def tool_call(name: str, args: Dict[str, Any], tool_id: str = "") -> StreamEvent:
|
||||
"""Tool call event."""
|
||||
return StreamEvent("tool_call", {"type": "tool_call", "name": name, "args": args, "id": tool_id})
|
||||
return StreamEvent(
|
||||
"tool_call",
|
||||
{"type": "tool_call", "name": name, "args": args, "id": tool_id},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def tool_result(name: str, content: str, success: bool = True) -> StreamEvent:
|
||||
"""Tool result event."""
|
||||
return StreamEvent("tool_result", {
|
||||
"type": "tool_result",
|
||||
"name": name,
|
||||
"content": content,
|
||||
"success": success,
|
||||
})
|
||||
return StreamEvent(
|
||||
"tool_result",
|
||||
{
|
||||
"type": "tool_result",
|
||||
"name": name,
|
||||
"content": content,
|
||||
"success": success,
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def subagent_start(name: str, description: str) -> StreamEvent:
|
||||
"""Sub-agent delegation started."""
|
||||
return StreamEvent("subagent_start", {
|
||||
"type": "subagent_start",
|
||||
"name": name,
|
||||
"description": description,
|
||||
})
|
||||
return StreamEvent(
|
||||
"subagent_start",
|
||||
{
|
||||
"type": "subagent_start",
|
||||
"name": name,
|
||||
"description": description,
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def subagent_tool_call(
|
||||
subagent: str, name: str, args: Dict[str, Any], tool_id: str = ""
|
||||
) -> StreamEvent:
|
||||
"""Tool call from inside a sub-agent."""
|
||||
return StreamEvent("subagent_tool_call", {
|
||||
"type": "subagent_tool_call",
|
||||
"subagent": subagent,
|
||||
"name": name,
|
||||
"args": args,
|
||||
"id": tool_id,
|
||||
})
|
||||
return StreamEvent(
|
||||
"subagent_tool_call",
|
||||
{
|
||||
"type": "subagent_tool_call",
|
||||
"subagent": subagent,
|
||||
"name": name,
|
||||
"args": args,
|
||||
"id": tool_id,
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def subagent_tool_result(
|
||||
subagent: str, name: str, content: str, success: bool = True
|
||||
) -> StreamEvent:
|
||||
"""Tool result from inside a sub-agent."""
|
||||
return StreamEvent("subagent_tool_result", {
|
||||
"type": "subagent_tool_result",
|
||||
"subagent": subagent,
|
||||
"name": name,
|
||||
"content": content,
|
||||
"success": success,
|
||||
})
|
||||
return StreamEvent(
|
||||
"subagent_tool_result",
|
||||
{
|
||||
"type": "subagent_tool_result",
|
||||
"subagent": subagent,
|
||||
"name": name,
|
||||
"content": content,
|
||||
"success": success,
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def subagent_end(name: str) -> StreamEvent:
|
||||
@@ -86,45 +104,62 @@ class StreamEventEmitter:
|
||||
@staticmethod
|
||||
def done(response: str = "") -> StreamEvent:
|
||||
"""Done event."""
|
||||
return StreamEvent("done", {"type": "done", "content": response, "response": response})
|
||||
return StreamEvent(
|
||||
"done", {"type": "done", "content": response, "response": response}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def usage_stats(input_tokens: int, output_tokens: int) -> StreamEvent:
|
||||
"""Token usage statistics event."""
|
||||
return StreamEvent("usage_stats", {
|
||||
"type": "usage_stats",
|
||||
"input_tokens": input_tokens,
|
||||
"output_tokens": output_tokens,
|
||||
})
|
||||
return StreamEvent(
|
||||
"usage_stats",
|
||||
{
|
||||
"type": "usage_stats",
|
||||
"input_tokens": input_tokens,
|
||||
"output_tokens": output_tokens,
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def interrupt(
|
||||
interrupt_id: str, action_requests: list, review_configs: list | None = None,
|
||||
interrupt_id: str,
|
||||
action_requests: list,
|
||||
review_configs: list | None = None,
|
||||
) -> StreamEvent:
|
||||
"""Human-in-the-loop interrupt event."""
|
||||
return StreamEvent("interrupt", {
|
||||
"type": "interrupt",
|
||||
"interrupt_id": interrupt_id,
|
||||
"action_requests": action_requests,
|
||||
"review_configs": review_configs or [],
|
||||
})
|
||||
return StreamEvent(
|
||||
"interrupt",
|
||||
{
|
||||
"type": "interrupt",
|
||||
"interrupt_id": interrupt_id,
|
||||
"action_requests": action_requests,
|
||||
"review_configs": review_configs or [],
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def ask_user_interrupt(
|
||||
interrupt_id: str, questions: list, tool_call_id: str = "",
|
||||
interrupt_id: str,
|
||||
questions: list,
|
||||
tool_call_id: str = "",
|
||||
) -> StreamEvent:
|
||||
"""Agent-initiated ask_user interrupt event."""
|
||||
return StreamEvent("ask_user", {
|
||||
"type": "ask_user",
|
||||
"interrupt_id": interrupt_id,
|
||||
"questions": questions,
|
||||
"tool_call_id": tool_call_id,
|
||||
})
|
||||
return StreamEvent(
|
||||
"ask_user",
|
||||
{
|
||||
"type": "ask_user",
|
||||
"interrupt_id": interrupt_id,
|
||||
"questions": questions,
|
||||
"tool_call_id": tool_call_id,
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def summarization(content: str) -> StreamEvent:
|
||||
"""Context summarization event."""
|
||||
return StreamEvent("summarization", {"type": "summarization", "content": content})
|
||||
return StreamEvent(
|
||||
"summarization", {"type": "summarization", "content": content}
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def error(message: str) -> StreamEvent:
|
||||
|
||||
+106
-33
@@ -17,7 +17,14 @@ from .tracker import ToolCallTracker
|
||||
from .utils import DisplayLimits, is_success
|
||||
|
||||
# Image media types returned by DeepAgents read_file
|
||||
_IMAGE_MEDIA_TYPES = {"image/png", "image/jpeg", "image/gif", "image/webp", "image/bmp", "image/svg+xml"}
|
||||
_IMAGE_MEDIA_TYPES = {
|
||||
"image/png",
|
||||
"image/jpeg",
|
||||
"image/gif",
|
||||
"image/webp",
|
||||
"image/bmp",
|
||||
"image/svg+xml",
|
||||
}
|
||||
|
||||
|
||||
def _extract_tool_content(msg) -> tuple[str, bool]:
|
||||
@@ -119,10 +126,10 @@ async def stream_agent_events(
|
||||
full_response = ""
|
||||
|
||||
# Track sub-agent names
|
||||
_key_to_name: dict[str, str] = {} # subagent_key -> display name (cache)
|
||||
_announced_names: list[str] = [] # ordered queue of announced task names
|
||||
_assigned_names: set[str] = set() # names already assigned to a namespace
|
||||
_announced_task_ids: list[str] = [] # ordered task tool_call_ids
|
||||
_key_to_name: dict[str, str] = {} # subagent_key -> display name (cache)
|
||||
_announced_names: list[str] = [] # ordered queue of announced task names
|
||||
_assigned_names: set[str] = set() # names already assigned to a namespace
|
||||
_announced_task_ids: list[str] = [] # ordered task tool_call_ids
|
||||
_task_id_to_name: dict[str, str] = {} # tool_call_id -> sub-agent name
|
||||
_subagent_trackers: dict[str, ToolCallTracker] = {} # namespace_key -> tracker
|
||||
|
||||
@@ -210,7 +217,13 @@ async def stream_agent_events(
|
||||
if meta_task_id:
|
||||
return f"task:{meta_task_id}"
|
||||
if metadata:
|
||||
for key in ("parent_run_id", "root_run_id", "run_id", "graph_id", "node_id"):
|
||||
for key in (
|
||||
"parent_run_id",
|
||||
"root_run_id",
|
||||
"run_id",
|
||||
"graph_id",
|
||||
"node_id",
|
||||
):
|
||||
val = metadata.get(key)
|
||||
if val:
|
||||
return f"{key}:{val}"
|
||||
@@ -239,8 +252,12 @@ async def stream_agent_events(
|
||||
lc_name = lc_name.strip()
|
||||
# Filter out generic/framework names
|
||||
if lc_name and lc_name not in (
|
||||
"sub-agent", "agent", "tools", "EvoScientist",
|
||||
"LangGraph", "",
|
||||
"sub-agent",
|
||||
"agent",
|
||||
"tools",
|
||||
"EvoScientist",
|
||||
"LangGraph",
|
||||
"",
|
||||
):
|
||||
_key_to_name[key] = lc_name
|
||||
return lc_name
|
||||
@@ -287,6 +304,7 @@ async def stream_agent_events(
|
||||
content_blocks: list[dict[str, Any]] = []
|
||||
if message:
|
||||
content_blocks.append({"type": "text", "text": message})
|
||||
|
||||
def _read_file_b64(path: str) -> str:
|
||||
with open(path, "rb") as fh:
|
||||
return base64.b64encode(fh.read()).decode("ascii")
|
||||
@@ -294,22 +312,30 @@ async def stream_agent_events(
|
||||
file_refs: list[str] = []
|
||||
for path in media:
|
||||
ext = os.path.splitext(path)[1].lower()
|
||||
is_image = ext in _IMAGE_EXTS and await asyncio.to_thread(os.path.isfile, path)
|
||||
is_image = ext in _IMAGE_EXTS and await asyncio.to_thread(
|
||||
os.path.isfile, path
|
||||
)
|
||||
if is_image:
|
||||
fsize = await asyncio.to_thread(os.path.getsize, path)
|
||||
if fsize <= _MAX_INLINE_SIZE:
|
||||
mime = mimetypes.guess_type(path)[0] or "image/png"
|
||||
b64 = await asyncio.to_thread(_read_file_b64, path)
|
||||
content_blocks.append({"type": "image_url", "image_url": {
|
||||
"url": f"data:{mime};base64,{b64}",
|
||||
}})
|
||||
content_blocks.append(
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": f"data:{mime};base64,{b64}",
|
||||
},
|
||||
}
|
||||
)
|
||||
else:
|
||||
file_refs.append(path)
|
||||
else:
|
||||
file_refs.append(path)
|
||||
if file_refs:
|
||||
ref_text = "\n".join(
|
||||
f"[attached file: {os.path.basename(p)}] path: {p}" for p in file_refs
|
||||
f"[attached file: {os.path.basename(p)}] path: {p}"
|
||||
for p in file_refs
|
||||
)
|
||||
content_blocks.append({"type": "text", "text": ref_text})
|
||||
if content_blocks:
|
||||
@@ -378,9 +404,15 @@ async def stream_agent_events(
|
||||
if isinstance(interrupt_value, dict)
|
||||
else getattr(interrupt_value, "tool_call_id", "")
|
||||
)
|
||||
ns_parts = interrupt_obj.get("ns", [""]) if isinstance(interrupt_obj, dict) else getattr(interrupt_obj, "ns", [""])
|
||||
ns_parts = (
|
||||
interrupt_obj.get("ns", [""])
|
||||
if isinstance(interrupt_obj, dict)
|
||||
else getattr(interrupt_obj, "ns", [""])
|
||||
)
|
||||
interrupt_id = str(ns_parts[0]) if ns_parts else "default"
|
||||
yield emitter.ask_user_interrupt(interrupt_id, questions, tc_id).data
|
||||
yield emitter.ask_user_interrupt(
|
||||
interrupt_id, questions, tc_id
|
||||
).data
|
||||
continue
|
||||
|
||||
# Standard HITL approval interrupt
|
||||
@@ -388,12 +420,20 @@ async def stream_agent_events(
|
||||
action_reqs = interrupt_value.get("action_requests", [])
|
||||
review_cfgs = interrupt_value.get("review_configs", [])
|
||||
else:
|
||||
action_reqs = getattr(interrupt_value, "action_requests", [])
|
||||
action_reqs = getattr(
|
||||
interrupt_value, "action_requests", []
|
||||
)
|
||||
review_cfgs = getattr(interrupt_value, "review_configs", [])
|
||||
if action_reqs:
|
||||
ns_parts = interrupt_obj.get("ns", [""]) if isinstance(interrupt_obj, dict) else getattr(interrupt_obj, "ns", [""])
|
||||
ns_parts = (
|
||||
interrupt_obj.get("ns", [""])
|
||||
if isinstance(interrupt_obj, dict)
|
||||
else getattr(interrupt_obj, "ns", [""])
|
||||
)
|
||||
interrupt_id = str(ns_parts[0]) if ns_parts else "default"
|
||||
yield emitter.interrupt(interrupt_id, action_reqs, review_cfgs).data
|
||||
yield emitter.interrupt(
|
||||
interrupt_id, action_reqs, review_cfgs
|
||||
).data
|
||||
continue
|
||||
if mode_str != "messages":
|
||||
continue
|
||||
@@ -410,7 +450,10 @@ async def stream_agent_events(
|
||||
# Accumulate summarization middleware chunks and emit text incrementally.
|
||||
# The summarization LLM streams AIMessageChunks; content may be a
|
||||
# plain string or a list of content blocks (provider-dependent).
|
||||
if isinstance(metadata, dict) and metadata.get("lc_source") == "summarization":
|
||||
if (
|
||||
isinstance(metadata, dict)
|
||||
and metadata.get("lc_source") == "summarization"
|
||||
):
|
||||
if not _summarization_in_progress:
|
||||
_summarization_in_progress = True
|
||||
chunk_text = _extract_summarization_text(msg)
|
||||
@@ -422,14 +465,24 @@ async def stream_agent_events(
|
||||
subagent_tracker = None
|
||||
if subagent:
|
||||
tracker_key = _get_subagent_key(namespace, metadata) or str(namespace)
|
||||
subagent_tracker = _subagent_trackers.setdefault(tracker_key, ToolCallTracker())
|
||||
subagent_tracker = _subagent_trackers.setdefault(
|
||||
tracker_key, ToolCallTracker()
|
||||
)
|
||||
|
||||
# Extract token usage from main-agent AIMessages
|
||||
if isinstance(msg, (AIMessageChunk, AIMessage)) and not subagent:
|
||||
usage = getattr(msg, "usage_metadata", None)
|
||||
if usage:
|
||||
inp = usage.get("input_tokens", 0) if isinstance(usage, dict) else getattr(usage, "input_tokens", 0)
|
||||
out = usage.get("output_tokens", 0) if isinstance(usage, dict) else getattr(usage, "output_tokens", 0)
|
||||
inp = (
|
||||
usage.get("input_tokens", 0)
|
||||
if isinstance(usage, dict)
|
||||
else getattr(usage, "input_tokens", 0)
|
||||
)
|
||||
out = (
|
||||
usage.get("output_tokens", 0)
|
||||
if isinstance(usage, dict)
|
||||
else getattr(usage, "output_tokens", 0)
|
||||
)
|
||||
if inp or out:
|
||||
yield emitter.usage_stats(inp, out).data
|
||||
|
||||
@@ -440,7 +493,10 @@ async def stream_agent_events(
|
||||
for ev in _process_chunk_content(msg, emitter, subagent_tracker):
|
||||
if ev.type == "tool_call":
|
||||
yield emitter.subagent_tool_call(
|
||||
subagent, ev.data["name"], ev.data["args"], ev.data.get("id", "")
|
||||
subagent,
|
||||
ev.data["name"],
|
||||
ev.data["args"],
|
||||
ev.data.get("id", ""),
|
||||
).data
|
||||
# Skip text/thinking from sub-agents (too noisy)
|
||||
|
||||
@@ -453,7 +509,10 @@ async def stream_agent_events(
|
||||
if not name and not tool_id:
|
||||
continue
|
||||
yield emitter.subagent_tool_call(
|
||||
subagent, name, args if isinstance(args, dict) else {}, tool_id
|
||||
subagent,
|
||||
name,
|
||||
args if isinstance(args, dict) else {},
|
||||
tool_id,
|
||||
).data
|
||||
else:
|
||||
# Main agent content
|
||||
@@ -463,15 +522,21 @@ async def stream_agent_events(
|
||||
yield ev.data
|
||||
|
||||
if hasattr(msg, "tool_calls") and msg.tool_calls:
|
||||
for ev in _process_tool_calls(msg.tool_calls, emitter, main_tracker):
|
||||
for ev in _process_tool_calls(
|
||||
msg.tool_calls, emitter, main_tracker
|
||||
):
|
||||
yield ev.data
|
||||
# Detect task tool calls -> announce sub-agent
|
||||
tc_data = ev.data
|
||||
if tc_data.get("name") == "task":
|
||||
started_name = _register_task_tool_call(tc_data)
|
||||
if started_name:
|
||||
desc = str(tc_data.get("args", {}).get("description", "")).strip()
|
||||
yield emitter.subagent_start(started_name, desc).data
|
||||
desc = str(
|
||||
tc_data.get("args", {}).get("description", "")
|
||||
).strip()
|
||||
yield emitter.subagent_start(
|
||||
started_name, desc
|
||||
).data
|
||||
|
||||
# Process ToolMessage (tool execution result)
|
||||
elif hasattr(msg, "type") and msg.type == "tool":
|
||||
@@ -487,9 +552,11 @@ async def stream_agent_events(
|
||||
).data
|
||||
name = getattr(msg, "name", "unknown")
|
||||
raw_content, _is_img = _extract_tool_content(msg)
|
||||
content = raw_content[:DisplayLimits.TOOL_RESULT_MAX]
|
||||
content = raw_content[: DisplayLimits.TOOL_RESULT_MAX]
|
||||
success = is_success(content)
|
||||
yield emitter.subagent_tool_result(subagent, name, content, success).data
|
||||
yield emitter.subagent_tool_result(
|
||||
subagent, name, content, success
|
||||
).data
|
||||
else:
|
||||
for ev in _process_tool_result(msg, emitter, main_tracker):
|
||||
yield ev.data
|
||||
@@ -497,7 +564,9 @@ async def stream_agent_events(
|
||||
if ev.type == "tool_call" and ev.data.get("name") == "task":
|
||||
started_name = _register_task_tool_call(ev.data)
|
||||
if started_name:
|
||||
desc = str(ev.data.get("args", {}).get("description", "")).strip()
|
||||
desc = str(
|
||||
ev.data.get("args", {}).get("description", "")
|
||||
).strip()
|
||||
yield emitter.subagent_start(started_name, desc).data
|
||||
# Check if this is a task result -> sub-agent ended
|
||||
name = getattr(msg, "name", "")
|
||||
@@ -514,7 +583,9 @@ async def stream_agent_events(
|
||||
yield emitter.done(full_response).data
|
||||
|
||||
|
||||
def _process_chunk_content(chunk, emitter: StreamEventEmitter, tracker: ToolCallTracker):
|
||||
def _process_chunk_content(
|
||||
chunk, emitter: StreamEventEmitter, tracker: ToolCallTracker
|
||||
):
|
||||
"""Process content blocks from an AI message chunk."""
|
||||
content = chunk.content
|
||||
|
||||
@@ -587,7 +658,9 @@ def _process_chunk_content(chunk, emitter: StreamEventEmitter, tracker: ToolCall
|
||||
tracker.append_json_delta(partial_args, block.get("index", 0))
|
||||
|
||||
|
||||
def _process_tool_calls(tool_calls: list, emitter: StreamEventEmitter, tracker: ToolCallTracker):
|
||||
def _process_tool_calls(
|
||||
tool_calls: list, emitter: StreamEventEmitter, tracker: ToolCallTracker
|
||||
):
|
||||
"""Process tool_calls from chunk.tool_calls attribute."""
|
||||
for tc in tool_calls:
|
||||
tool_id = tc.get("id", "")
|
||||
@@ -612,7 +685,7 @@ def _process_tool_result(chunk, emitter: StreamEventEmitter, tracker: ToolCallTr
|
||||
|
||||
name = getattr(chunk, "name", "unknown")
|
||||
raw_content, _is_img = _extract_tool_content(chunk)
|
||||
content = raw_content[:DisplayLimits.TOOL_RESULT_MAX]
|
||||
content = raw_content[: DisplayLimits.TOOL_RESULT_MAX]
|
||||
if len(raw_content) > DisplayLimits.TOOL_RESULT_MAX:
|
||||
content += "\n... (truncated)"
|
||||
|
||||
|
||||
@@ -20,6 +20,7 @@ from .utils import SUCCESS_PREFIX, FAILURE_PREFIX, is_success as _is_success, tr
|
||||
|
||||
class ContentType(Enum):
|
||||
"""Content type categories."""
|
||||
|
||||
SUCCESS = "success"
|
||||
ERROR = "error"
|
||||
JSON = "json"
|
||||
@@ -30,6 +31,7 @@ class ContentType(Enum):
|
||||
@dataclass
|
||||
class FormattedResult:
|
||||
"""Formatted result container."""
|
||||
|
||||
content_type: ContentType
|
||||
elements: List[Any] # Rich renderable elements
|
||||
success: bool = True
|
||||
@@ -85,7 +87,9 @@ class ToolResultFormatter:
|
||||
formatter = formatter_map.get(content_type, self._format_text)
|
||||
elements = formatter(name, content, max_length)
|
||||
|
||||
return FormattedResult(content_type=content_type, elements=elements, success=success)
|
||||
return FormattedResult(
|
||||
content_type=content_type, elements=elements, success=success
|
||||
)
|
||||
|
||||
def _extract_body(self, content: str) -> str:
|
||||
"""Extract body after status prefix."""
|
||||
@@ -96,8 +100,9 @@ class ToolResultFormatter:
|
||||
content = content.strip()
|
||||
if not content:
|
||||
return False
|
||||
if (content.startswith('{') and content.endswith('}')) or \
|
||||
(content.startswith('[') and content.endswith(']')):
|
||||
if (content.startswith("{") and content.endswith("}")) or (
|
||||
content.startswith("[") and content.endswith("]")
|
||||
):
|
||||
try:
|
||||
json.loads(content)
|
||||
return True
|
||||
@@ -108,33 +113,37 @@ class ToolResultFormatter:
|
||||
def _is_error(self, content: str) -> bool:
|
||||
head = "\n".join(content.splitlines()[:3])
|
||||
error_patterns = [
|
||||
'Traceback (most recent call last)',
|
||||
'Exception:',
|
||||
'Error:',
|
||||
'Error invoking tool',
|
||||
'Failed ',
|
||||
"Traceback (most recent call last)",
|
||||
"Exception:",
|
||||
"Error:",
|
||||
"Error invoking tool",
|
||||
"Failed ",
|
||||
]
|
||||
return any(pattern in head for pattern in error_patterns)
|
||||
|
||||
def _is_markdown(self, content: str) -> bool:
|
||||
md_patterns = ['```', '**', '##', '- **']
|
||||
return content.startswith('#') or any(p in content for p in md_patterns)
|
||||
md_patterns = ["```", "**", "##", "- **"]
|
||||
return content.startswith("#") or any(p in content for p in md_patterns)
|
||||
|
||||
def _format_success(self, name: str, content: str, max_length: int) -> List[Any]:
|
||||
display = truncate(content, max_length)
|
||||
return [Panel(
|
||||
Text(display, style="green"),
|
||||
title=f"{escape(name)} OK",
|
||||
border_style="green",
|
||||
)]
|
||||
return [
|
||||
Panel(
|
||||
Text(display, style="green"),
|
||||
title=f"{escape(name)} OK",
|
||||
border_style="green",
|
||||
)
|
||||
]
|
||||
|
||||
def _format_error(self, name: str, content: str, max_length: int) -> List[Any]:
|
||||
display = truncate(content, max_length)
|
||||
return [Panel(
|
||||
Text(display, style="red"),
|
||||
title=f"{escape(name)} FAILED",
|
||||
border_style="red",
|
||||
)]
|
||||
return [
|
||||
Panel(
|
||||
Text(display, style="red"),
|
||||
title=f"{escape(name)} FAILED",
|
||||
border_style="red",
|
||||
)
|
||||
]
|
||||
|
||||
def _format_json(self, name: str, content: str, max_length: int) -> List[Any]:
|
||||
json_content = content
|
||||
@@ -154,11 +163,13 @@ class ToolResultFormatter:
|
||||
|
||||
def _format_markdown(self, name: str, content: str, max_length: int) -> List[Any]:
|
||||
display = truncate(content, max_length)
|
||||
return [Panel(
|
||||
Markdown(display),
|
||||
title=escape(name),
|
||||
border_style="cyan dim",
|
||||
)]
|
||||
return [
|
||||
Panel(
|
||||
Markdown(display),
|
||||
title=escape(name),
|
||||
border_style="cyan dim",
|
||||
)
|
||||
]
|
||||
|
||||
def _format_text(self, name: str, content: str, max_length: int) -> List[Any]:
|
||||
display = truncate(content, max_length)
|
||||
|
||||
@@ -107,13 +107,16 @@ class StreamState:
|
||||
def get_response_markdown(self):
|
||||
"""Return cached Markdown object, only re-parsing when text changes."""
|
||||
from rich.markdown import Markdown # type: ignore[import-untyped]
|
||||
|
||||
text = (self.response_text or "").strip()
|
||||
if text != self._cached_md_text:
|
||||
self._cached_md_text = text
|
||||
self._cached_md = Markdown(text) if text else None
|
||||
return self._cached_md
|
||||
|
||||
def _get_or_create_subagent(self, name: str, description: str = "") -> SubAgentState:
|
||||
def _get_or_create_subagent(
|
||||
self, name: str, description: str = ""
|
||||
) -> SubAgentState:
|
||||
if name not in self._subagent_map:
|
||||
# Case 1: real name arrives, "sub-agent" entry exists -> rename it
|
||||
if name != "sub-agent" and "sub-agent" in self._subagent_map:
|
||||
@@ -127,7 +130,8 @@ class StreamState:
|
||||
# exists with no tool calls -> merge into it
|
||||
if name == "sub-agent":
|
||||
active_named = [
|
||||
sa for sa in self.subagents
|
||||
sa
|
||||
for sa in self.subagents
|
||||
if sa.is_active and sa.name != "sub-agent"
|
||||
]
|
||||
if len(active_named) == 1 and not active_named[0].tool_calls:
|
||||
@@ -151,8 +155,7 @@ class StreamState:
|
||||
if name != "sub-agent":
|
||||
return name
|
||||
active_named = [
|
||||
sa.name for sa in self.subagents
|
||||
if sa.is_active and sa.name != "sub-agent"
|
||||
sa.name for sa in self.subagents if sa.is_active and sa.name != "sub-agent"
|
||||
]
|
||||
if len(active_named) == 1:
|
||||
return active_named[0]
|
||||
@@ -214,10 +217,12 @@ class StreamState:
|
||||
if result_name not in _INTERNAL_TOOLS:
|
||||
self.is_processing = True
|
||||
result_content = event.get("content", "")
|
||||
self.tool_results.append({
|
||||
"name": result_name,
|
||||
"content": result_content,
|
||||
})
|
||||
self.tool_results.append(
|
||||
{
|
||||
"name": result_name,
|
||||
"content": result_content,
|
||||
}
|
||||
)
|
||||
# Update todo list from write_todos / read_todos results (fallback)
|
||||
if result_name in ("write_todos", "read_todos"):
|
||||
parsed = _parse_todo_items(result_content)
|
||||
@@ -344,7 +349,7 @@ def _parse_todo_items(content: str) -> list[dict] | None:
|
||||
if bracket_start != -1:
|
||||
bracket_end = content.rfind("]")
|
||||
if bracket_end > bracket_start:
|
||||
embedded = content[bracket_start:bracket_end + 1]
|
||||
embedded = content[bracket_start : bracket_end + 1]
|
||||
result = _try_parse(embedded)
|
||||
if result:
|
||||
return result
|
||||
@@ -356,7 +361,7 @@ def _parse_todo_items(content: str) -> list[dict] | None:
|
||||
start = line.find("[")
|
||||
end = line.rfind("]")
|
||||
if end > start:
|
||||
result = _try_parse(line[start:end + 1])
|
||||
result = _try_parse(line[start : end + 1])
|
||||
if result:
|
||||
return result
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ from typing import Dict, Optional
|
||||
@dataclass
|
||||
class ToolCallInfo:
|
||||
"""Tool call information."""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
args: Dict = field(default_factory=dict)
|
||||
|
||||
@@ -18,19 +18,17 @@ FAILURE_PREFIX = "[FAILED]"
|
||||
# === Tool status indicators ===
|
||||
class ToolStatus(str, Enum):
|
||||
"""Tool execution status indicators."""
|
||||
RUNNING = "\u25cf" # Running - yellow
|
||||
SUCCESS = "\u25cf" # Success - green
|
||||
ERROR = "\u25cf" # Failed - red
|
||||
PENDING = "\u25cb" # Pending - gray
|
||||
|
||||
RUNNING = "\u25cf" # Running - yellow
|
||||
SUCCESS = "\u25cf" # Success - green
|
||||
ERROR = "\u25cf" # Failed - red
|
||||
PENDING = "\u25cb" # Pending - gray
|
||||
|
||||
|
||||
def get_status_symbol(status: ToolStatus) -> str:
|
||||
"""Get status symbol with ASCII fallback for terminals without Unicode."""
|
||||
try:
|
||||
supports_unicode = (
|
||||
sys.stdout.encoding
|
||||
and 'utf' in sys.stdout.encoding.lower()
|
||||
)
|
||||
supports_unicode = sys.stdout.encoding and "utf" in sys.stdout.encoding.lower()
|
||||
except Exception:
|
||||
supports_unicode = False
|
||||
|
||||
@@ -49,6 +47,7 @@ def get_status_symbol(status: ToolStatus) -> str:
|
||||
# === Display limit constants ===
|
||||
class DisplayLimits:
|
||||
"""Display length limits."""
|
||||
|
||||
THINKING_STREAM = 1000
|
||||
THINKING_FINAL = 2000
|
||||
ARGS_INLINE = 100
|
||||
@@ -78,11 +77,11 @@ def is_success(content: str) -> bool:
|
||||
# not buried deep inside file content or command output.
|
||||
head = "\n".join(content.splitlines()[:3])
|
||||
error_patterns = [
|
||||
'Traceback (most recent call last)',
|
||||
'Exception:',
|
||||
'Error:',
|
||||
'Error invoking tool',
|
||||
'Failed ',
|
||||
"Traceback (most recent call last)",
|
||||
"Exception:",
|
||||
"Error:",
|
||||
"Error invoking tool",
|
||||
"Failed ",
|
||||
]
|
||||
return not any(pattern in head for pattern in error_patterns)
|
||||
|
||||
@@ -96,6 +95,7 @@ def truncate(content: str, max_length: int, suffix: str = "\n... (truncated)") -
|
||||
|
||||
# === Compact formatting for deepagents tools ===
|
||||
|
||||
|
||||
def _shorten_path(path: str, max_len: int = 40) -> str:
|
||||
"""Shorten a file path for display."""
|
||||
if len(path) <= max_len:
|
||||
|
||||
@@ -44,7 +44,12 @@ def skill_manager(
|
||||
Returns:
|
||||
Result message
|
||||
"""
|
||||
from .skills_manager import install_skill, list_skills, uninstall_skill, get_skill_info
|
||||
from .skills_manager import (
|
||||
install_skill,
|
||||
list_skills,
|
||||
uninstall_skill,
|
||||
get_skill_info,
|
||||
)
|
||||
|
||||
if action == "install":
|
||||
if not source:
|
||||
@@ -117,4 +122,6 @@ def skill_manager(
|
||||
)
|
||||
|
||||
else:
|
||||
return f"Unknown action: {action}. Use 'install', 'list', 'uninstall', or 'info'."
|
||||
return (
|
||||
f"Unknown action: {action}. Use 'install', 'list', 'uninstall', or 'info'."
|
||||
)
|
||||
|
||||
@@ -140,7 +140,9 @@ def _clone_repo(repo: str, ref: str | None, dest: str) -> None:
|
||||
cmd += [clone_url, dest]
|
||||
|
||||
try:
|
||||
result = subprocess.run(cmd, capture_output=True, text=True, timeout=_CLONE_TIMEOUT)
|
||||
result = subprocess.run(
|
||||
cmd, capture_output=True, text=True, timeout=_CLONE_TIMEOUT
|
||||
)
|
||||
except subprocess.TimeoutExpired:
|
||||
raise RuntimeError(f"git clone timed out after {_CLONE_TIMEOUT}s for {repo}")
|
||||
if result.returncode != 0:
|
||||
@@ -187,7 +189,8 @@ def _scan_skill_dirs(root: Path) -> list[Path]:
|
||||
else:
|
||||
# Non-skill directory — scan its children (level 2)
|
||||
found.extend(
|
||||
gc for gc in sorted(child.iterdir())
|
||||
gc
|
||||
for gc in sorted(child.iterdir())
|
||||
if gc.is_dir() and _validate_skill_dir(gc)
|
||||
)
|
||||
return found
|
||||
@@ -281,9 +284,7 @@ def _install_from_local(source: str, dest_dir: str) -> dict:
|
||||
return _install_single_local(source_path, dest_dir)
|
||||
|
||||
|
||||
def _install_single_local(
|
||||
source_path: Path, dest_dir: str, *, ignore_fn=None
|
||||
) -> dict:
|
||||
def _install_single_local(source_path: Path, dest_dir: str, *, ignore_fn=None) -> dict:
|
||||
"""Install one skill directory into *dest_dir*."""
|
||||
skill_info = _parse_skill_md(source_path / "SKILL.md")
|
||||
skill_name = _sanitize_name(skill_info["name"])
|
||||
@@ -379,14 +380,18 @@ def _install_from_github(source: str, dest_dir: str) -> dict:
|
||||
if resolved:
|
||||
skill_source = resolved
|
||||
else:
|
||||
return {"success": False, "error": f"No SKILL.md found at '{path}' (also searched subdirectories) in: {source}"}
|
||||
return {
|
||||
"success": False,
|
||||
"error": f"No SKILL.md found at '{path}' (also searched subdirectories) in: {source}",
|
||||
}
|
||||
else:
|
||||
return {"success": False, "error": f"No SKILL.md found in: {source}"}
|
||||
return {
|
||||
"success": False,
|
||||
"error": f"No SKILL.md found in: {source}",
|
||||
}
|
||||
|
||||
# Single skill — install it
|
||||
result = _install_single_local(
|
||||
skill_source, dest_dir, ignore_fn=ignore_git
|
||||
)
|
||||
result = _install_single_local(skill_source, dest_dir, ignore_fn=ignore_git)
|
||||
if result.get("success"):
|
||||
result["source"] = source
|
||||
return result
|
||||
|
||||
@@ -78,7 +78,6 @@ def format_messages(messages):
|
||||
console.print(Panel(content, title=f"📝 {msg_type}", border_style="white"))
|
||||
|
||||
|
||||
|
||||
def show_prompt(prompt_text: str, title: str = "Prompt", border_style: str = "blue"):
|
||||
"""Display a prompt with rich formatting and XML tag highlighting.
|
||||
|
||||
@@ -160,7 +159,9 @@ def load_subagents(
|
||||
if "system_prompt_ref" in spec:
|
||||
ref = spec["system_prompt_ref"]
|
||||
if ref not in prompt_refs:
|
||||
raise ValueError(f"Unknown system_prompt_ref '{ref}' for subagent '{name}'")
|
||||
raise ValueError(
|
||||
f"Unknown system_prompt_ref '{ref}' for subagent '{name}'"
|
||||
)
|
||||
subagent["system_prompt"] = prompt_refs[ref]
|
||||
else:
|
||||
subagent["system_prompt"] = spec.get("system_prompt", "")
|
||||
|
||||
+31
-6
@@ -1,6 +1,5 @@
|
||||
"""Shared fixtures for EvoScientist tests."""
|
||||
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
@@ -44,11 +43,37 @@ def sample_events():
|
||||
return [
|
||||
{"type": "thinking", "content": "Let me think..."},
|
||||
{"type": "text", "content": "Here is the answer."},
|
||||
{"type": "tool_call", "id": "tc_001", "name": "execute", "args": {"command": "ls"}},
|
||||
{"type": "tool_result", "name": "execute", "content": "[OK] done", "success": True},
|
||||
{"type": "subagent_start", "name": "research-agent", "description": "Find papers"},
|
||||
{"type": "subagent_tool_call", "subagent": "research-agent", "name": "tavily_search", "args": {"query": "test"}, "id": "tc_sa_001"},
|
||||
{"type": "subagent_tool_result", "subagent": "research-agent", "name": "tavily_search", "content": "Results...", "success": True},
|
||||
{
|
||||
"type": "tool_call",
|
||||
"id": "tc_001",
|
||||
"name": "execute",
|
||||
"args": {"command": "ls"},
|
||||
},
|
||||
{
|
||||
"type": "tool_result",
|
||||
"name": "execute",
|
||||
"content": "[OK] done",
|
||||
"success": True,
|
||||
},
|
||||
{
|
||||
"type": "subagent_start",
|
||||
"name": "research-agent",
|
||||
"description": "Find papers",
|
||||
},
|
||||
{
|
||||
"type": "subagent_tool_call",
|
||||
"subagent": "research-agent",
|
||||
"name": "tavily_search",
|
||||
"args": {"query": "test"},
|
||||
"id": "tc_sa_001",
|
||||
},
|
||||
{
|
||||
"type": "subagent_tool_result",
|
||||
"subagent": "research-agent",
|
||||
"name": "tavily_search",
|
||||
"content": "Results...",
|
||||
"success": True,
|
||||
},
|
||||
{"type": "subagent_end", "name": "research-agent"},
|
||||
{"type": "done", "response": "Here is the answer."},
|
||||
]
|
||||
|
||||
@@ -4,7 +4,10 @@ import pytest
|
||||
|
||||
from EvoScientist.channels.base import ChannelError, OutboundMessage
|
||||
from EvoScientist.channels.email.channel import EmailChannel, EmailConfig
|
||||
from EvoScientist.channels.imessage.channel_rpc import IMessageChannelRpc, IMessageConfig
|
||||
from EvoScientist.channels.imessage.channel_rpc import (
|
||||
IMessageChannelRpc,
|
||||
IMessageConfig,
|
||||
)
|
||||
from EvoScientist.channels.qq.channel import QQChannel, QQConfig
|
||||
from EvoScientist.channels.signal.channel import SignalChannel, SignalConfig
|
||||
|
||||
@@ -14,7 +17,9 @@ from tests.conftest import run_async as _run
|
||||
class TestEmailChannelSmoke:
|
||||
def test_start_raises_without_required_imap_settings(self):
|
||||
channel = EmailChannel(EmailConfig())
|
||||
with pytest.raises(ChannelError, match="imap_host and imap_username are required"):
|
||||
with pytest.raises(
|
||||
ChannelError, match="imap_host and imap_username are required"
|
||||
):
|
||||
_run(channel.start())
|
||||
|
||||
def test_send_returns_false_when_smtp_not_ready(self):
|
||||
|
||||
@@ -143,9 +143,7 @@ class TestValidateQuestions:
|
||||
from EvoScientist.middleware.ask_user import _validate_questions
|
||||
|
||||
with pytest.raises(ValueError, match="non-empty 'choices' list"):
|
||||
_validate_questions(
|
||||
[{"question": "Q?", "type": "multiple_choice"}]
|
||||
)
|
||||
_validate_questions([{"question": "Q?", "type": "multiple_choice"}])
|
||||
|
||||
def test_text_with_choices_raises(self):
|
||||
from EvoScientist.middleware.ask_user import _validate_questions
|
||||
|
||||
+31
-11
@@ -1,6 +1,5 @@
|
||||
"""Tests for EvoScientist/backends.py — validate_command, path conversion, resolve_path."""
|
||||
|
||||
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
@@ -13,6 +12,7 @@ from EvoScientist.backends import (
|
||||
|
||||
# === validate_command ===
|
||||
|
||||
|
||||
class TestValidateCommand:
|
||||
def test_safe_ls(self):
|
||||
assert validate_command("ls -la") is None
|
||||
@@ -62,6 +62,7 @@ class TestValidateCommand:
|
||||
|
||||
# === convert_virtual_paths_in_command ===
|
||||
|
||||
|
||||
class TestConvertVirtualPaths:
|
||||
def test_absolute_to_relative(self):
|
||||
result = convert_virtual_paths_in_command("python /main.py")
|
||||
@@ -144,13 +145,15 @@ class TestConvertVirtualPaths:
|
||||
def test_system_path_without_workspace_unchanged(self):
|
||||
"""System paths not referencing workspace fall through to normal ./"""
|
||||
result = convert_virtual_paths_in_command(
|
||||
"cat /tmp/somefile", workspace_name="workspace",
|
||||
"cat /tmp/somefile",
|
||||
workspace_name="workspace",
|
||||
)
|
||||
assert result == "cat ./tmp/somefile"
|
||||
|
||||
|
||||
# === CustomSandboxBackend._resolve_path ===
|
||||
|
||||
|
||||
class TestResolvePath:
|
||||
def test_strip_workspace_prefix(self, tmp_workspace):
|
||||
backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True)
|
||||
@@ -200,6 +203,7 @@ class TestResolvePath:
|
||||
|
||||
# === CustomSandboxBackend.id ===
|
||||
|
||||
|
||||
class TestSandboxId:
|
||||
def test_sandbox_has_id(self, tmp_workspace):
|
||||
backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True)
|
||||
@@ -218,12 +222,13 @@ class TestSandboxId:
|
||||
|
||||
def test_sandbox_id_hex_suffix(self, tmp_workspace):
|
||||
backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True)
|
||||
suffix = backend.id[len("evosci-"):]
|
||||
suffix = backend.id[len("evosci-") :]
|
||||
assert re.fullmatch(r"[0-9a-f]{8}", suffix)
|
||||
|
||||
|
||||
# === execute() literal cwd sanitization ===
|
||||
|
||||
|
||||
class TestExecuteCwdSanitization:
|
||||
def test_literal_workspace_path_replaced(self, tmp_workspace):
|
||||
"""execute() should replace literal workspace root path with ./"""
|
||||
@@ -238,11 +243,13 @@ class TestExecuteCwdSanitization:
|
||||
|
||||
# === execute() output truncation ===
|
||||
|
||||
|
||||
class TestExecuteTruncation:
|
||||
def test_execute_truncates_large_output(self, tmp_workspace):
|
||||
backend = CustomSandboxBackend(
|
||||
root_dir=tmp_workspace,
|
||||
virtual_mode=True, max_output_bytes=100,
|
||||
virtual_mode=True,
|
||||
max_output_bytes=100,
|
||||
)
|
||||
# Generate output larger than 100 bytes
|
||||
resp = backend.execute("python3 -c \"print('A' * 200)\"")
|
||||
@@ -255,7 +262,8 @@ class TestExecuteTruncation:
|
||||
def test_execute_no_truncation_small_output(self, tmp_workspace):
|
||||
backend = CustomSandboxBackend(
|
||||
root_dir=tmp_workspace,
|
||||
virtual_mode=True, max_output_bytes=100_000,
|
||||
virtual_mode=True,
|
||||
max_output_bytes=100_000,
|
||||
)
|
||||
resp = backend.execute("echo hello")
|
||||
assert resp.truncated is False
|
||||
@@ -264,25 +272,31 @@ class TestExecuteTruncation:
|
||||
|
||||
# === execute() stderr attribution ===
|
||||
|
||||
|
||||
class TestExecuteStderr:
|
||||
def test_execute_stderr_attribution(self, tmp_workspace):
|
||||
backend = CustomSandboxBackend(
|
||||
root_dir=tmp_workspace, virtual_mode=True,
|
||||
root_dir=tmp_workspace,
|
||||
virtual_mode=True,
|
||||
)
|
||||
resp = backend.execute(
|
||||
"python3 -c \"import sys; sys.stderr.write('warning\\n')\""
|
||||
)
|
||||
resp = backend.execute("python3 -c \"import sys; sys.stderr.write('warning\\n')\"")
|
||||
assert "[stderr] warning" in resp.output
|
||||
|
||||
def test_execute_nonzero_exit_code_in_output(self, tmp_workspace):
|
||||
backend = CustomSandboxBackend(
|
||||
root_dir=tmp_workspace, virtual_mode=True,
|
||||
root_dir=tmp_workspace,
|
||||
virtual_mode=True,
|
||||
)
|
||||
resp = backend.execute("python3 -c \"raise SystemExit(42)\"")
|
||||
resp = backend.execute('python3 -c "raise SystemExit(42)"')
|
||||
assert resp.exit_code == 42
|
||||
assert "Exit code: 42" in resp.output
|
||||
|
||||
def test_execute_mixed_stdout_stderr(self, tmp_workspace):
|
||||
backend = CustomSandboxBackend(
|
||||
root_dir=tmp_workspace, virtual_mode=True,
|
||||
root_dir=tmp_workspace,
|
||||
virtual_mode=True,
|
||||
)
|
||||
resp = backend.execute(
|
||||
"python3 -c \"import sys; print('out'); sys.stderr.write('err\\n')\""
|
||||
@@ -292,7 +306,8 @@ class TestExecuteStderr:
|
||||
|
||||
def test_execute_success_no_exit_code(self, tmp_workspace):
|
||||
backend = CustomSandboxBackend(
|
||||
root_dir=tmp_workspace, virtual_mode=True,
|
||||
root_dir=tmp_workspace,
|
||||
virtual_mode=True,
|
||||
)
|
||||
resp = backend.execute("echo ok")
|
||||
assert resp.exit_code == 0
|
||||
@@ -301,6 +316,7 @@ class TestExecuteStderr:
|
||||
|
||||
# === execute() timeout kwarg ===
|
||||
|
||||
|
||||
class TestExecuteTimeout:
|
||||
def test_execute_accepts_timeout_kwarg(self, tmp_workspace):
|
||||
backend = CustomSandboxBackend(root_dir=tmp_workspace, virtual_mode=True)
|
||||
@@ -315,12 +331,14 @@ class TestExecuteTimeout:
|
||||
|
||||
def test_execute_accepts_timeout_introspection(self):
|
||||
from deepagents.backends.protocol import execute_accepts_timeout
|
||||
|
||||
execute_accepts_timeout.cache_clear()
|
||||
assert execute_accepts_timeout(CustomSandboxBackend) is True
|
||||
|
||||
|
||||
# === '..' traversal false-positive fix ===
|
||||
|
||||
|
||||
class TestTraversalFalsePositiveFix:
|
||||
def test_dotdot_in_filename_allowed(self):
|
||||
assert validate_command("echo foo..bar.txt") is None
|
||||
@@ -337,6 +355,7 @@ class TestTraversalFalsePositiveFix:
|
||||
|
||||
# === Pipeline command validation ===
|
||||
|
||||
|
||||
class TestPipelineCommandValidation:
|
||||
def test_pipe_blocked_command(self):
|
||||
"""sudo after pipe should be caught."""
|
||||
@@ -371,6 +390,7 @@ class TestPipelineCommandValidation:
|
||||
|
||||
# === execute() timeout recovery guidance ===
|
||||
|
||||
|
||||
class TestExecuteTimeoutRecovery:
|
||||
def test_timeout_includes_recovery_guidance(self, tmp_workspace):
|
||||
backend = CustomSandboxBackend(root_dir=tmp_workspace, timeout=1)
|
||||
|
||||
@@ -73,6 +73,7 @@ class TestBusInboundConsumer:
|
||||
_message_queue,
|
||||
_set_channel_response,
|
||||
)
|
||||
|
||||
_drain_queue(_message_queue)
|
||||
|
||||
async def _test():
|
||||
@@ -81,16 +82,16 @@ class TestBusInboundConsumer:
|
||||
ch = FakeChannel()
|
||||
manager.register(ch)
|
||||
|
||||
consumer = asyncio.create_task(
|
||||
_bus_inbound_consumer(bus, manager)
|
||||
)
|
||||
consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager))
|
||||
|
||||
await bus.publish_inbound(InboundMessage(
|
||||
channel="fake",
|
||||
sender_id="user1",
|
||||
chat_id="chat1",
|
||||
content="hello agent",
|
||||
))
|
||||
await bus.publish_inbound(
|
||||
InboundMessage(
|
||||
channel="fake",
|
||||
sender_id="user1",
|
||||
chat_id="chat1",
|
||||
content="hello agent",
|
||||
)
|
||||
)
|
||||
|
||||
# Wait for consumer to enqueue the message
|
||||
for _ in range(20):
|
||||
@@ -107,7 +108,8 @@ class TestBusInboundConsumer:
|
||||
_set_channel_response(msg.msg_id, "Reply to: hello agent")
|
||||
|
||||
outbound = await asyncio.wait_for(
|
||||
bus.consume_outbound(), timeout=2.0,
|
||||
bus.consume_outbound(),
|
||||
timeout=2.0,
|
||||
)
|
||||
assert outbound.channel == "fake"
|
||||
assert outbound.chat_id == "chat1"
|
||||
@@ -128,6 +130,7 @@ class TestBusInboundConsumer:
|
||||
_message_queue,
|
||||
_set_channel_response,
|
||||
)
|
||||
|
||||
_drain_queue(_message_queue)
|
||||
|
||||
async def _test():
|
||||
@@ -136,16 +139,16 @@ class TestBusInboundConsumer:
|
||||
ch = FakeChannel()
|
||||
manager.register(ch)
|
||||
|
||||
consumer = asyncio.create_task(
|
||||
_bus_inbound_consumer(bus, manager)
|
||||
)
|
||||
consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager))
|
||||
|
||||
await bus.publish_inbound(InboundMessage(
|
||||
channel="fake",
|
||||
sender_id="user1",
|
||||
chat_id="chat1",
|
||||
content="test",
|
||||
))
|
||||
await bus.publish_inbound(
|
||||
InboundMessage(
|
||||
channel="fake",
|
||||
sender_id="user1",
|
||||
chat_id="chat1",
|
||||
content="test",
|
||||
)
|
||||
)
|
||||
|
||||
for _ in range(20):
|
||||
if not _message_queue.empty():
|
||||
@@ -157,7 +160,8 @@ class TestBusInboundConsumer:
|
||||
_set_channel_response(msg.msg_id, "")
|
||||
|
||||
outbound = await asyncio.wait_for(
|
||||
bus.consume_outbound(), timeout=2.0,
|
||||
bus.consume_outbound(),
|
||||
timeout=2.0,
|
||||
)
|
||||
assert outbound.content == "No response"
|
||||
|
||||
@@ -176,6 +180,7 @@ class TestBusInboundConsumer:
|
||||
_message_queue,
|
||||
_set_channel_response,
|
||||
)
|
||||
|
||||
_drain_queue(_message_queue)
|
||||
|
||||
async def _test():
|
||||
@@ -184,16 +189,16 @@ class TestBusInboundConsumer:
|
||||
ch = FakeChannel()
|
||||
manager.register(ch)
|
||||
|
||||
consumer = asyncio.create_task(
|
||||
_bus_inbound_consumer(bus, manager)
|
||||
)
|
||||
consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager))
|
||||
|
||||
await bus.publish_inbound(InboundMessage(
|
||||
channel="fake",
|
||||
sender_id="u1",
|
||||
chat_id="c1",
|
||||
content="test",
|
||||
))
|
||||
await bus.publish_inbound(
|
||||
InboundMessage(
|
||||
channel="fake",
|
||||
sender_id="u1",
|
||||
chat_id="c1",
|
||||
content="test",
|
||||
)
|
||||
)
|
||||
|
||||
for _ in range(20):
|
||||
if not _message_queue.empty():
|
||||
@@ -223,6 +228,7 @@ class TestBusInboundConsumer:
|
||||
_message_queue,
|
||||
_set_channel_response,
|
||||
)
|
||||
|
||||
_drain_queue(_message_queue)
|
||||
|
||||
async def _test():
|
||||
@@ -231,18 +237,18 @@ class TestBusInboundConsumer:
|
||||
ch = FakeChannel()
|
||||
manager.register(ch)
|
||||
|
||||
consumer = asyncio.create_task(
|
||||
_bus_inbound_consumer(bus, manager)
|
||||
)
|
||||
consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager))
|
||||
|
||||
await bus.publish_inbound(InboundMessage(
|
||||
channel="fake",
|
||||
sender_id="user1",
|
||||
chat_id="chat1",
|
||||
content="with metadata",
|
||||
metadata={"key": "value"},
|
||||
message_id="msg-123",
|
||||
))
|
||||
await bus.publish_inbound(
|
||||
InboundMessage(
|
||||
channel="fake",
|
||||
sender_id="user1",
|
||||
chat_id="chat1",
|
||||
content="with metadata",
|
||||
metadata={"key": "value"},
|
||||
message_id="msg-123",
|
||||
)
|
||||
)
|
||||
|
||||
for _ in range(20):
|
||||
if not _message_queue.empty():
|
||||
@@ -259,7 +265,8 @@ class TestBusInboundConsumer:
|
||||
_set_channel_response(msg.msg_id, "done")
|
||||
|
||||
outbound = await asyncio.wait_for(
|
||||
bus.consume_outbound(), timeout=2.0,
|
||||
bus.consume_outbound(),
|
||||
timeout=2.0,
|
||||
)
|
||||
assert outbound.reply_to == "msg-123"
|
||||
|
||||
|
||||
@@ -78,6 +78,7 @@ class TestIsCcproxyRunning:
|
||||
@patch("httpx.get")
|
||||
def test_not_running(self, mock_get):
|
||||
import httpx
|
||||
|
||||
mock_get.side_effect = httpx.ConnectError("Connection refused")
|
||||
assert is_ccproxy_running(8000) is False
|
||||
|
||||
@@ -215,7 +216,10 @@ class TestMaybeStartCcproxy:
|
||||
with pytest.raises(RuntimeError, match="not found"):
|
||||
maybe_start_ccproxy(config)
|
||||
|
||||
@patch("EvoScientist.ccproxy_manager.check_ccproxy_auth", return_value=(False, "expired"))
|
||||
@patch(
|
||||
"EvoScientist.ccproxy_manager.check_ccproxy_auth",
|
||||
return_value=(False, "expired"),
|
||||
)
|
||||
@patch("EvoScientist.ccproxy_manager.is_ccproxy_available", return_value=True)
|
||||
def test_oauth_mode_raises_no_auth(self, mock_avail, mock_auth):
|
||||
config = MagicMock()
|
||||
|
||||
+304
-131
File diff suppressed because it is too large
Load Diff
@@ -59,7 +59,9 @@ def test_cmd_channel_running_path_passes_send_thinking(monkeypatch):
|
||||
config = SimpleNamespace(channel_enabled="telegram")
|
||||
|
||||
monkeypatch.setattr(config_mod, "load_config", lambda: config)
|
||||
monkeypatch.setattr(channel_cli, "_channels_is_running", lambda _channel_type=None: True)
|
||||
monkeypatch.setattr(
|
||||
channel_cli, "_channels_is_running", lambda _channel_type=None: True
|
||||
)
|
||||
monkeypatch.setattr(channel_cli, "_channels_running_list", lambda: [])
|
||||
monkeypatch.setattr(channel_cli, "_print_channel_panel", lambda _rows: None)
|
||||
monkeypatch.setattr(
|
||||
@@ -93,7 +95,9 @@ def test_cmd_channel_start_path_passes_send_thinking(monkeypatch):
|
||||
config = SimpleNamespace(channel_enabled="telegram")
|
||||
|
||||
monkeypatch.setattr(config_mod, "load_config", lambda: config)
|
||||
monkeypatch.setattr(channel_cli, "_channels_is_running", lambda _channel_type=None: False)
|
||||
monkeypatch.setattr(
|
||||
channel_cli, "_channels_is_running", lambda _channel_type=None: False
|
||||
)
|
||||
monkeypatch.setattr(channel_cli, "_print_channel_panel", lambda _rows: None)
|
||||
monkeypatch.setattr(
|
||||
channel_cli,
|
||||
|
||||
@@ -62,7 +62,9 @@ def _run_serve_once(
|
||||
monkeypatch.setattr(commands, "set_workspace_root", _fake_set_workspace_root)
|
||||
monkeypatch.setattr(commands, "ensure_dirs", _fake_ensure_dirs)
|
||||
monkeypatch.setattr(commands, "_load_agent", _fake_load_agent)
|
||||
monkeypatch.setattr(commands, "_start_channels_bus_mode", _fake_start_channels_bus_mode)
|
||||
monkeypatch.setattr(
|
||||
commands, "_start_channels_bus_mode", _fake_start_channels_bus_mode
|
||||
)
|
||||
monkeypatch.setattr(commands, "_channels_stop", _fake_channels_stop)
|
||||
monkeypatch.setattr(commands, "_message_queue", _InterruptQueue())
|
||||
|
||||
@@ -79,7 +81,9 @@ def _run_serve_once(
|
||||
return order, captured
|
||||
|
||||
|
||||
def test_serve_workdir_has_highest_priority_and_sets_root_before_ensure(monkeypatch, tmp_path):
|
||||
def test_serve_workdir_has_highest_priority_and_sets_root_before_ensure(
|
||||
monkeypatch, tmp_path
|
||||
):
|
||||
cfg_ws = tmp_path / "cfg_ws"
|
||||
cli_ws = tmp_path / "cli_ws"
|
||||
config = _make_config(default_workdir=str(cfg_ws), channel_send_thinking=True)
|
||||
|
||||
+118
-33
@@ -72,11 +72,25 @@ class TestCompactCutoffZero:
|
||||
mock_middleware_cls = MagicMock(return_value=mock_middleware_inst)
|
||||
|
||||
with (
|
||||
patch("EvoScientist.EvoScientist._ensure_chat_model", return_value=MagicMock()),
|
||||
patch("EvoScientist.EvoScientist._get_default_backend", return_value=MagicMock()),
|
||||
patch("deepagents.middleware.summarization.SummarizationMiddleware", mock_middleware_cls),
|
||||
patch("deepagents.middleware.summarization.compute_summarization_defaults", return_value={"keep": ("messages", 6)}),
|
||||
patch("langchain_core.messages.utils.count_tokens_approximately", return_value=500),
|
||||
patch(
|
||||
"EvoScientist.EvoScientist._ensure_chat_model", return_value=MagicMock()
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.EvoScientist._get_default_backend",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"deepagents.middleware.summarization.SummarizationMiddleware",
|
||||
mock_middleware_cls,
|
||||
),
|
||||
patch(
|
||||
"deepagents.middleware.summarization.compute_summarization_defaults",
|
||||
return_value={"keep": ("messages", 6)},
|
||||
),
|
||||
patch(
|
||||
"langchain_core.messages.utils.count_tokens_approximately",
|
||||
return_value=500,
|
||||
),
|
||||
):
|
||||
result = _run(compact_conversation(agent=agent, thread_id="tid-1"))
|
||||
|
||||
@@ -93,7 +107,9 @@ class TestCompactNegligibleSavings:
|
||||
|
||||
agent = MagicMock()
|
||||
msgs = [MagicMock() for _ in range(15)]
|
||||
snapshot = SimpleNamespace(values={"messages": msgs, "_summarization_event": None})
|
||||
snapshot = SimpleNamespace(
|
||||
values={"messages": msgs, "_summarization_event": None}
|
||||
)
|
||||
agent.aget_state = AsyncMock(return_value=snapshot)
|
||||
|
||||
mock_middleware_inst = MagicMock()
|
||||
@@ -108,11 +124,25 @@ class TestCompactNegligibleSavings:
|
||||
token_values = iter([200, 22000])
|
||||
|
||||
with (
|
||||
patch("EvoScientist.EvoScientist._ensure_chat_model", return_value=MagicMock()),
|
||||
patch("EvoScientist.EvoScientist._get_default_backend", return_value=MagicMock()),
|
||||
patch("deepagents.middleware.summarization.SummarizationMiddleware", mock_middleware_cls),
|
||||
patch("deepagents.middleware.summarization.compute_summarization_defaults", return_value={"keep": ("messages", 6)}),
|
||||
patch("langchain_core.messages.utils.count_tokens_approximately", side_effect=lambda x: next(token_values)),
|
||||
patch(
|
||||
"EvoScientist.EvoScientist._ensure_chat_model", return_value=MagicMock()
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.EvoScientist._get_default_backend",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"deepagents.middleware.summarization.SummarizationMiddleware",
|
||||
mock_middleware_cls,
|
||||
),
|
||||
patch(
|
||||
"deepagents.middleware.summarization.compute_summarization_defaults",
|
||||
return_value={"keep": ("messages", 6)},
|
||||
),
|
||||
patch(
|
||||
"langchain_core.messages.utils.count_tokens_approximately",
|
||||
side_effect=lambda x: next(token_values),
|
||||
),
|
||||
):
|
||||
result = _run(compact_conversation(agent=agent, thread_id="tid-1"))
|
||||
|
||||
@@ -128,7 +158,9 @@ class TestCompactNegligibleSavings:
|
||||
|
||||
agent = MagicMock()
|
||||
msgs = [MagicMock() for _ in range(10)]
|
||||
snapshot = SimpleNamespace(values={"messages": msgs, "_summarization_event": None})
|
||||
snapshot = SimpleNamespace(
|
||||
values={"messages": msgs, "_summarization_event": None}
|
||||
)
|
||||
agent.aget_state = AsyncMock(return_value=snapshot)
|
||||
agent.aupdate_state = AsyncMock()
|
||||
|
||||
@@ -149,11 +181,25 @@ class TestCompactNegligibleSavings:
|
||||
token_values = iter([5000, 15000, 500])
|
||||
|
||||
with (
|
||||
patch("EvoScientist.EvoScientist._ensure_chat_model", return_value=MagicMock()),
|
||||
patch("EvoScientist.EvoScientist._get_default_backend", return_value=MagicMock()),
|
||||
patch("deepagents.middleware.summarization.SummarizationMiddleware", mock_middleware_cls),
|
||||
patch("deepagents.middleware.summarization.compute_summarization_defaults", return_value={"keep": ("messages", 6)}),
|
||||
patch("langchain_core.messages.utils.count_tokens_approximately", side_effect=lambda x: next(token_values)),
|
||||
patch(
|
||||
"EvoScientist.EvoScientist._ensure_chat_model", return_value=MagicMock()
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.EvoScientist._get_default_backend",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"deepagents.middleware.summarization.SummarizationMiddleware",
|
||||
mock_middleware_cls,
|
||||
),
|
||||
patch(
|
||||
"deepagents.middleware.summarization.compute_summarization_defaults",
|
||||
return_value={"keep": ("messages", 6)},
|
||||
),
|
||||
patch(
|
||||
"langchain_core.messages.utils.count_tokens_approximately",
|
||||
side_effect=lambda x: next(token_values),
|
||||
),
|
||||
):
|
||||
result = _run(compact_conversation(agent=agent, thread_id="tid-1"))
|
||||
|
||||
@@ -170,7 +216,9 @@ class TestCompactSuccess:
|
||||
|
||||
agent = MagicMock()
|
||||
msgs = [MagicMock() for _ in range(20)]
|
||||
snapshot = SimpleNamespace(values={"messages": msgs, "_summarization_event": None})
|
||||
snapshot = SimpleNamespace(
|
||||
values={"messages": msgs, "_summarization_event": None}
|
||||
)
|
||||
agent.aget_state = AsyncMock(return_value=snapshot)
|
||||
agent.aupdate_state = AsyncMock()
|
||||
|
||||
@@ -183,7 +231,9 @@ class TestCompactSuccess:
|
||||
mock_middleware_inst._determine_cutoff_index.return_value = 15
|
||||
mock_middleware_inst._partition_messages.return_value = (to_summarize, to_keep)
|
||||
mock_middleware_inst._acreate_summary = AsyncMock(return_value="Summary text")
|
||||
mock_middleware_inst._aoffload_to_backend = AsyncMock(return_value="/conversation_history/tid.md")
|
||||
mock_middleware_inst._aoffload_to_backend = AsyncMock(
|
||||
return_value="/conversation_history/tid.md"
|
||||
)
|
||||
mock_middleware_inst._build_new_messages_with_path.return_value = [summary_msg]
|
||||
mock_middleware_inst._compute_state_cutoff.return_value = 15
|
||||
|
||||
@@ -193,11 +243,25 @@ class TestCompactSuccess:
|
||||
token_values = iter([5000, 1000, 200])
|
||||
|
||||
with (
|
||||
patch("EvoScientist.EvoScientist._ensure_chat_model", return_value=MagicMock()),
|
||||
patch("EvoScientist.EvoScientist._get_default_backend", return_value=MagicMock()),
|
||||
patch("deepagents.middleware.summarization.SummarizationMiddleware", mock_middleware_cls),
|
||||
patch("deepagents.middleware.summarization.compute_summarization_defaults", return_value={"keep": ("messages", 6)}),
|
||||
patch("langchain_core.messages.utils.count_tokens_approximately", side_effect=lambda x: next(token_values)),
|
||||
patch(
|
||||
"EvoScientist.EvoScientist._ensure_chat_model", return_value=MagicMock()
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.EvoScientist._get_default_backend",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"deepagents.middleware.summarization.SummarizationMiddleware",
|
||||
mock_middleware_cls,
|
||||
),
|
||||
patch(
|
||||
"deepagents.middleware.summarization.compute_summarization_defaults",
|
||||
return_value={"keep": ("messages", 6)},
|
||||
),
|
||||
patch(
|
||||
"langchain_core.messages.utils.count_tokens_approximately",
|
||||
side_effect=lambda x: next(token_values),
|
||||
),
|
||||
):
|
||||
result = _run(compact_conversation(agent=agent, thread_id="tid-1"))
|
||||
|
||||
@@ -222,7 +286,9 @@ class TestCompactSuccess:
|
||||
|
||||
agent = MagicMock()
|
||||
msgs = [MagicMock() for _ in range(10)]
|
||||
snapshot = SimpleNamespace(values={"messages": msgs, "_summarization_event": None})
|
||||
snapshot = SimpleNamespace(
|
||||
values={"messages": msgs, "_summarization_event": None}
|
||||
)
|
||||
agent.aget_state = AsyncMock(return_value=snapshot)
|
||||
agent.aupdate_state = AsyncMock()
|
||||
|
||||
@@ -233,18 +299,34 @@ class TestCompactSuccess:
|
||||
mock_middleware_inst._determine_cutoff_index.return_value = 7
|
||||
mock_middleware_inst._partition_messages.return_value = (msgs[:7], msgs[7:])
|
||||
mock_middleware_inst._acreate_summary = AsyncMock(return_value="Summary")
|
||||
mock_middleware_inst._aoffload_to_backend = AsyncMock(side_effect=RuntimeError("write failed"))
|
||||
mock_middleware_inst._aoffload_to_backend = AsyncMock(
|
||||
side_effect=RuntimeError("write failed")
|
||||
)
|
||||
mock_middleware_inst._build_new_messages_with_path.return_value = [summary_msg]
|
||||
mock_middleware_inst._compute_state_cutoff.return_value = 7
|
||||
|
||||
mock_middleware_cls = MagicMock(return_value=mock_middleware_inst)
|
||||
|
||||
with (
|
||||
patch("EvoScientist.EvoScientist._ensure_chat_model", return_value=MagicMock()),
|
||||
patch("EvoScientist.EvoScientist._get_default_backend", return_value=MagicMock()),
|
||||
patch("deepagents.middleware.summarization.SummarizationMiddleware", mock_middleware_cls),
|
||||
patch("deepagents.middleware.summarization.compute_summarization_defaults", return_value={"keep": ("messages", 6)}),
|
||||
patch("langchain_core.messages.utils.count_tokens_approximately", return_value=1000),
|
||||
patch(
|
||||
"EvoScientist.EvoScientist._ensure_chat_model", return_value=MagicMock()
|
||||
),
|
||||
patch(
|
||||
"EvoScientist.EvoScientist._get_default_backend",
|
||||
return_value=MagicMock(),
|
||||
),
|
||||
patch(
|
||||
"deepagents.middleware.summarization.SummarizationMiddleware",
|
||||
mock_middleware_cls,
|
||||
),
|
||||
patch(
|
||||
"deepagents.middleware.summarization.compute_summarization_defaults",
|
||||
return_value={"keep": ("messages", 6)},
|
||||
),
|
||||
patch(
|
||||
"langchain_core.messages.utils.count_tokens_approximately",
|
||||
return_value=1000,
|
||||
),
|
||||
):
|
||||
result = _run(compact_conversation(agent=agent, thread_id="tid-1"))
|
||||
|
||||
@@ -271,7 +353,9 @@ class TestRenderCompactResult:
|
||||
def test_render_noop_no_tokens(self):
|
||||
from EvoScientist.cli.commands import CompactResult, render_compact_result
|
||||
|
||||
result = CompactResult("noop", "Nothing to compact — no messages in conversation.")
|
||||
result = CompactResult(
|
||||
"noop", "Nothing to compact — no messages in conversation."
|
||||
)
|
||||
text = render_compact_result(result)
|
||||
assert "Nothing to compact" in text.plain
|
||||
|
||||
@@ -286,7 +370,8 @@ class TestRenderCompactResult:
|
||||
from EvoScientist.cli.commands import CompactResult, render_compact_result
|
||||
|
||||
result = CompactResult(
|
||||
"ok", "Compacted",
|
||||
"ok",
|
||||
"Compacted",
|
||||
messages_compacted=15,
|
||||
messages_kept=5,
|
||||
tokens_before=6000,
|
||||
|
||||
@@ -188,11 +188,14 @@ class TestLoadSaveReset:
|
||||
config_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
with open(config_path, "w") as f:
|
||||
yaml.safe_dump({
|
||||
"provider": "openai",
|
||||
"unknown_field": "should_be_ignored",
|
||||
"another_bad": 123,
|
||||
}, f)
|
||||
yaml.safe_dump(
|
||||
{
|
||||
"provider": "openai",
|
||||
"unknown_field": "should_be_ignored",
|
||||
"another_bad": 123,
|
||||
},
|
||||
f,
|
||||
)
|
||||
|
||||
config = load_config()
|
||||
assert config.provider == "openai"
|
||||
|
||||
@@ -17,6 +17,7 @@ from EvoScientist.stream.diff_format import (
|
||||
# _escape_markup
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestEscapeMarkup:
|
||||
def test_escapes_brackets(self):
|
||||
assert _escape_markup("[bold]text[/bold]") == r"\[bold\]text\[/bold\]"
|
||||
@@ -35,6 +36,7 @@ class TestEscapeMarkup:
|
||||
# _detect_unicode_support
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestDetectUnicodeSupport:
|
||||
def test_utf8_encoding(self):
|
||||
with mock.patch("sys.stdout") as mock_stdout:
|
||||
@@ -60,6 +62,7 @@ class TestDetectUnicodeSupport:
|
||||
# format_diff_rich
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFormatDiffRich:
|
||||
def test_empty_diff_returns_dim_message(self):
|
||||
result = format_diff_rich("")
|
||||
@@ -143,6 +146,7 @@ class TestFormatDiffRich:
|
||||
# build_edit_diff
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestBuildEditDiff:
|
||||
def test_returns_none_when_equal(self):
|
||||
assert build_edit_diff("/foo.py", "same", "same") is None
|
||||
@@ -195,6 +199,7 @@ class TestBuildEditDiff:
|
||||
# Integration with format_tool_result_compact
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestFormatToolResultCompactEditFile:
|
||||
def test_edit_file_with_tool_args_shows_diff(self):
|
||||
from EvoScientist.stream.display import format_tool_result_compact
|
||||
|
||||
@@ -80,6 +80,7 @@ class TestDingTalkChannel:
|
||||
|
||||
def test_capabilities(self):
|
||||
from EvoScientist.channels.capabilities import DINGTALK
|
||||
|
||||
config = DingTalkConfig()
|
||||
channel = DingTalkChannel(config)
|
||||
assert channel.capabilities is DINGTALK
|
||||
@@ -268,6 +269,7 @@ class TestDingTalkSendChunk:
|
||||
class TestDingTalkChannelRegistration:
|
||||
def test_dingtalk_registered(self):
|
||||
from EvoScientist.channels.channel_manager import available_channels
|
||||
|
||||
channels = available_channels()
|
||||
assert "dingtalk" in channels
|
||||
|
||||
@@ -275,16 +277,19 @@ class TestDingTalkChannelRegistration:
|
||||
class TestDingTalkProbe:
|
||||
def test_missing_credentials(self):
|
||||
from EvoScientist.channels.dingtalk.probe import validate_dingtalk
|
||||
|
||||
ok, msg = _run(validate_dingtalk("", ""))
|
||||
assert ok is False
|
||||
assert "required" in msg
|
||||
|
||||
def test_missing_client_id(self):
|
||||
from EvoScientist.channels.dingtalk.probe import validate_dingtalk
|
||||
|
||||
ok, msg = _run(validate_dingtalk("", "secret"))
|
||||
assert ok is False
|
||||
|
||||
def test_missing_client_secret(self):
|
||||
from EvoScientist.channels.dingtalk.probe import validate_dingtalk
|
||||
|
||||
ok, msg = _run(validate_dingtalk("id", ""))
|
||||
assert ok is False
|
||||
|
||||
@@ -111,9 +111,11 @@ class TestMultipleStreamingCalls:
|
||||
pass
|
||||
|
||||
# Patch the stream_agent_events function
|
||||
with patch('EvoScientist.stream.display.stream_agent_events', side_effect=mock_stream):
|
||||
with patch(
|
||||
"EvoScientist.stream.display.stream_agent_events", side_effect=mock_stream
|
||||
):
|
||||
# Patch Live to avoid terminal output during tests
|
||||
with patch('EvoScientist.stream.display.Live'):
|
||||
with patch("EvoScientist.stream.display.Live"):
|
||||
# First call
|
||||
_run_streaming(
|
||||
agent=mock_agent,
|
||||
|
||||
@@ -88,6 +88,7 @@ class TestFeishuChannel:
|
||||
|
||||
def test_capabilities(self):
|
||||
from EvoScientist.channels.capabilities import FEISHU
|
||||
|
||||
config = FeishuConfig()
|
||||
channel = FeishuChannel(config)
|
||||
assert channel.capabilities is FEISHU
|
||||
@@ -105,7 +106,10 @@ class TestFeishuChannel:
|
||||
"zh_cn": {
|
||||
"title": "Test Title",
|
||||
"content": [
|
||||
[{"tag": "text", "text": "Hello "}, {"tag": "a", "text": "world", "href": "http://example.com"}],
|
||||
[
|
||||
{"tag": "text", "text": "Hello "},
|
||||
{"tag": "a", "text": "world", "href": "http://example.com"},
|
||||
],
|
||||
[{"tag": "text", "text": "Second line"}],
|
||||
],
|
||||
}
|
||||
@@ -407,8 +411,7 @@ class TestFeishuMarkdownConversion:
|
||||
def test_bold_text(self):
|
||||
elements = _parse_inline_text("**bold text**")
|
||||
assert any(
|
||||
e.get("style") == ["bold"] and e["text"] == "bold text"
|
||||
for e in elements
|
||||
e.get("style") == ["bold"] and e["text"] == "bold text" for e in elements
|
||||
)
|
||||
|
||||
def test_inline_code(self):
|
||||
@@ -420,10 +423,7 @@ class TestFeishuMarkdownConversion:
|
||||
|
||||
def test_link(self):
|
||||
elements = _parse_inline_text("[click](http://example.com)")
|
||||
assert any(
|
||||
e.get("tag") == "a" and e["text"] == "click"
|
||||
for e in elements
|
||||
)
|
||||
assert any(e.get("tag") == "a" and e["text"] == "click" for e in elements)
|
||||
|
||||
def test_strikethrough(self):
|
||||
elements = _parse_inline_text("~~deleted~~")
|
||||
@@ -442,8 +442,7 @@ class TestFeishuMarkdownConversion:
|
||||
def test_heading(self):
|
||||
elements = _parse_inline_elements("## My Heading")
|
||||
assert any(
|
||||
e.get("style") == ["bold"] and e["text"] == "My Heading"
|
||||
for e in elements
|
||||
e.get("style") == ["bold"] and e["text"] == "My Heading" for e in elements
|
||||
)
|
||||
|
||||
def test_blockquote(self):
|
||||
@@ -479,6 +478,7 @@ class TestFeishuMarkdownConversion:
|
||||
class TestFeishuChannelRegistration:
|
||||
def test_feishu_registered(self):
|
||||
from EvoScientist.channels.channel_manager import available_channels
|
||||
|
||||
channels = available_channels()
|
||||
assert "feishu" in channels
|
||||
|
||||
@@ -486,12 +486,14 @@ class TestFeishuChannelRegistration:
|
||||
class TestFeishuProbe:
|
||||
def test_missing_app_id(self):
|
||||
from EvoScientist.channels.feishu.probe import validate_feishu_credentials
|
||||
|
||||
ok, msg = _run(validate_feishu_credentials("", "secret"))
|
||||
assert ok is False
|
||||
assert "app_id" in msg
|
||||
|
||||
def test_missing_app_secret(self):
|
||||
from EvoScientist.channels.feishu.probe import validate_feishu_credentials
|
||||
|
||||
ok, msg = _run(validate_feishu_credentials("id", ""))
|
||||
assert ok is False
|
||||
assert "app_secret" in msg
|
||||
|
||||
+171
-65
@@ -11,6 +11,7 @@ from EvoScientist.stream.state import StreamState
|
||||
# StreamEventEmitter.interrupt()
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestInterruptEmitter:
|
||||
def test_interrupt_event_structure(self):
|
||||
ev = StreamEventEmitter.interrupt(
|
||||
@@ -43,6 +44,7 @@ class TestInterruptEmitter:
|
||||
# StreamState.handle_event("interrupt")
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestStreamStateInterrupt:
|
||||
def test_handle_interrupt_sets_pending(self):
|
||||
state = StreamState()
|
||||
@@ -64,23 +66,27 @@ class TestStreamStateInterrupt:
|
||||
def test_interrupt_does_not_affect_other_state(self):
|
||||
state = StreamState()
|
||||
state.handle_event({"type": "text", "content": "hello"})
|
||||
state.handle_event({
|
||||
"type": "interrupt",
|
||||
"interrupt_id": "main",
|
||||
"action_requests": [{"name": "execute"}],
|
||||
"review_configs": [],
|
||||
})
|
||||
state.handle_event(
|
||||
{
|
||||
"type": "interrupt",
|
||||
"interrupt_id": "main",
|
||||
"action_requests": [{"name": "execute"}],
|
||||
"review_configs": [],
|
||||
}
|
||||
)
|
||||
assert state.response_text == "hello"
|
||||
assert state.pending_interrupt is not None
|
||||
|
||||
def test_done_after_interrupt_preserves_pending(self):
|
||||
state = StreamState()
|
||||
state.handle_event({
|
||||
"type": "interrupt",
|
||||
"interrupt_id": "main",
|
||||
"action_requests": [{"name": "execute"}],
|
||||
"review_configs": [],
|
||||
})
|
||||
state.handle_event(
|
||||
{
|
||||
"type": "interrupt",
|
||||
"interrupt_id": "main",
|
||||
"action_requests": [{"name": "execute"}],
|
||||
"review_configs": [],
|
||||
}
|
||||
)
|
||||
state.handle_event({"type": "done", "response": ""})
|
||||
# pending_interrupt should still be set
|
||||
assert state.pending_interrupt is not None
|
||||
@@ -90,30 +96,37 @@ class TestStreamStateInterrupt:
|
||||
# _matches_shell_allow_list
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestMatchesShellAllowList:
|
||||
def test_matches_prefix(self):
|
||||
from EvoScientist.stream.display import _matches_shell_allow_list
|
||||
|
||||
assert _matches_shell_allow_list("ls -la", ["ls", "cat"]) is True
|
||||
assert _matches_shell_allow_list("cat file.txt", ["ls", "cat"]) is True
|
||||
|
||||
def test_no_match(self):
|
||||
from EvoScientist.stream.display import _matches_shell_allow_list
|
||||
|
||||
assert _matches_shell_allow_list("rm -rf /", ["ls", "cat"]) is False
|
||||
|
||||
def test_empty_allow_list(self):
|
||||
from EvoScientist.stream.display import _matches_shell_allow_list
|
||||
|
||||
assert _matches_shell_allow_list("ls", []) is False
|
||||
|
||||
def test_whitespace_handling(self):
|
||||
from EvoScientist.stream.display import _matches_shell_allow_list
|
||||
|
||||
assert _matches_shell_allow_list(" ls -la", ["ls"]) is True
|
||||
|
||||
def test_exact_match(self):
|
||||
from EvoScientist.stream.display import _matches_shell_allow_list
|
||||
|
||||
assert _matches_shell_allow_list("python", ["python"]) is True
|
||||
|
||||
def test_partial_word_match(self):
|
||||
from EvoScientist.stream.display import _matches_shell_allow_list
|
||||
|
||||
# "ls" prefix matches "lsof" — this is by design (prefix matching)
|
||||
assert _matches_shell_allow_list("lsof", ["ls"]) is True
|
||||
|
||||
@@ -122,20 +135,27 @@ class TestMatchesShellAllowList:
|
||||
# _resolve_hitl_approval
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestResolveHitlApproval:
|
||||
def test_empty_requests_auto_approves(self):
|
||||
from EvoScientist.stream.display import _resolve_hitl_approval
|
||||
|
||||
result = _resolve_hitl_approval({"action_requests": []})
|
||||
assert result == [{"type": "approve"}]
|
||||
|
||||
def test_session_auto_approve(self):
|
||||
import EvoScientist.stream.display as disp
|
||||
|
||||
original = disp._session_auto_approve
|
||||
try:
|
||||
disp._session_auto_approve = True
|
||||
result = disp._resolve_hitl_approval({
|
||||
"action_requests": [{"name": "execute", "args": {"command": "rm -rf /"}}],
|
||||
})
|
||||
result = disp._resolve_hitl_approval(
|
||||
{
|
||||
"action_requests": [
|
||||
{"name": "execute", "args": {"command": "rm -rf /"}}
|
||||
],
|
||||
}
|
||||
)
|
||||
assert result == [{"type": "approve"}]
|
||||
finally:
|
||||
disp._session_auto_approve = original
|
||||
@@ -143,16 +163,23 @@ class TestResolveHitlApproval:
|
||||
def test_config_auto_approve(self):
|
||||
from EvoScientist.stream.display import _resolve_hitl_approval
|
||||
import EvoScientist.stream.display as disp
|
||||
|
||||
original = disp._session_auto_approve
|
||||
try:
|
||||
disp._session_auto_approve = False
|
||||
mock_cfg = MagicMock()
|
||||
mock_cfg.auto_approve = True
|
||||
mock_cfg.shell_allow_list = ""
|
||||
with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg):
|
||||
result = _resolve_hitl_approval({
|
||||
"action_requests": [{"name": "execute", "args": {"command": "rm"}}],
|
||||
})
|
||||
with patch(
|
||||
"EvoScientist.config.settings.load_config", return_value=mock_cfg
|
||||
):
|
||||
result = _resolve_hitl_approval(
|
||||
{
|
||||
"action_requests": [
|
||||
{"name": "execute", "args": {"command": "rm"}}
|
||||
],
|
||||
}
|
||||
)
|
||||
assert result == [{"type": "approve"}]
|
||||
finally:
|
||||
disp._session_auto_approve = original
|
||||
@@ -160,16 +187,23 @@ class TestResolveHitlApproval:
|
||||
def test_non_execute_tool_auto_approves(self):
|
||||
from EvoScientist.stream.display import _resolve_hitl_approval
|
||||
import EvoScientist.stream.display as disp
|
||||
|
||||
original = disp._session_auto_approve
|
||||
try:
|
||||
disp._session_auto_approve = False
|
||||
mock_cfg = MagicMock()
|
||||
mock_cfg.auto_approve = False
|
||||
mock_cfg.shell_allow_list = ""
|
||||
with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg):
|
||||
result = _resolve_hitl_approval({
|
||||
"action_requests": [{"name": "write_file", "args": {"path": "/out.txt"}}],
|
||||
})
|
||||
with patch(
|
||||
"EvoScientist.config.settings.load_config", return_value=mock_cfg
|
||||
):
|
||||
result = _resolve_hitl_approval(
|
||||
{
|
||||
"action_requests": [
|
||||
{"name": "write_file", "args": {"path": "/out.txt"}}
|
||||
],
|
||||
}
|
||||
)
|
||||
assert result == [{"type": "approve"}]
|
||||
finally:
|
||||
disp._session_auto_approve = original
|
||||
@@ -177,16 +211,23 @@ class TestResolveHitlApproval:
|
||||
def test_execute_with_matching_allow_list(self):
|
||||
from EvoScientist.stream.display import _resolve_hitl_approval
|
||||
import EvoScientist.stream.display as disp
|
||||
|
||||
original = disp._session_auto_approve
|
||||
try:
|
||||
disp._session_auto_approve = False
|
||||
mock_cfg = MagicMock()
|
||||
mock_cfg.auto_approve = False
|
||||
mock_cfg.shell_allow_list = "ls,cat,python"
|
||||
with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg):
|
||||
result = _resolve_hitl_approval({
|
||||
"action_requests": [{"name": "execute", "args": {"command": "ls -la"}}],
|
||||
})
|
||||
with patch(
|
||||
"EvoScientist.config.settings.load_config", return_value=mock_cfg
|
||||
):
|
||||
result = _resolve_hitl_approval(
|
||||
{
|
||||
"action_requests": [
|
||||
{"name": "execute", "args": {"command": "ls -la"}}
|
||||
],
|
||||
}
|
||||
)
|
||||
assert result == [{"type": "approve"}]
|
||||
finally:
|
||||
disp._session_auto_approve = original
|
||||
@@ -194,18 +235,27 @@ class TestResolveHitlApproval:
|
||||
def test_execute_not_in_allow_list_prompts(self):
|
||||
from EvoScientist.stream.display import _resolve_hitl_approval
|
||||
import EvoScientist.stream.display as disp
|
||||
|
||||
original = disp._session_auto_approve
|
||||
try:
|
||||
disp._session_auto_approve = False
|
||||
mock_cfg = MagicMock()
|
||||
mock_cfg.auto_approve = False
|
||||
mock_cfg.shell_allow_list = "ls,cat"
|
||||
with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg):
|
||||
with patch("EvoScientist.stream.display._prompt_hitl_approval") as mock_prompt:
|
||||
with patch(
|
||||
"EvoScientist.config.settings.load_config", return_value=mock_cfg
|
||||
):
|
||||
with patch(
|
||||
"EvoScientist.stream.display._prompt_hitl_approval"
|
||||
) as mock_prompt:
|
||||
mock_prompt.return_value = [{"type": "approve"}]
|
||||
result = _resolve_hitl_approval({
|
||||
"action_requests": [{"name": "execute", "args": {"command": "rm -rf /"}}],
|
||||
})
|
||||
result = _resolve_hitl_approval(
|
||||
{
|
||||
"action_requests": [
|
||||
{"name": "execute", "args": {"command": "rm -rf /"}}
|
||||
],
|
||||
}
|
||||
)
|
||||
assert result == [{"type": "approve"}]
|
||||
mock_prompt.assert_called_once()
|
||||
finally:
|
||||
@@ -216,24 +266,29 @@ class TestResolveHitlApproval:
|
||||
# Config fields
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestHitlConfig:
|
||||
def test_auto_approve_default(self):
|
||||
from EvoScientist.config.settings import EvoScientistConfig
|
||||
|
||||
cfg = EvoScientistConfig()
|
||||
assert cfg.auto_approve is False
|
||||
|
||||
def test_shell_allow_list_default(self):
|
||||
from EvoScientist.config.settings import EvoScientistConfig
|
||||
|
||||
cfg = EvoScientistConfig()
|
||||
assert cfg.shell_allow_list == ""
|
||||
|
||||
def test_auto_approve_set(self):
|
||||
from EvoScientist.config.settings import EvoScientistConfig
|
||||
|
||||
cfg = EvoScientistConfig(auto_approve=True)
|
||||
assert cfg.auto_approve is True
|
||||
|
||||
def test_shell_allow_list_set(self):
|
||||
from EvoScientist.config.settings import EvoScientistConfig
|
||||
|
||||
cfg = EvoScientistConfig(shell_allow_list="ls,cat,python")
|
||||
assert cfg.shell_allow_list == "ls,cat,python"
|
||||
|
||||
@@ -242,6 +297,7 @@ class TestHitlConfig:
|
||||
# Interrupt event parsing in stream_agent_events
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestInterruptEventParsing:
|
||||
def _run_async(self, coro):
|
||||
"""Run async code with a fresh event loop."""
|
||||
@@ -261,18 +317,23 @@ class TestInterruptEventParsing:
|
||||
ai_chunk = AIMessageChunk(content="thinking...", id="msg1")
|
||||
|
||||
interrupt_data = {
|
||||
"__interrupt__": [{
|
||||
"value": {
|
||||
"action_requests": [
|
||||
{"name": "execute", "args": {"command": "ls"}, "id": "tc1"}
|
||||
],
|
||||
"review_configs": [
|
||||
{"action_name": "execute", "allowed_decisions": ["approve", "reject"]}
|
||||
],
|
||||
},
|
||||
"ns": ["main"],
|
||||
"resumable": True,
|
||||
}]
|
||||
"__interrupt__": [
|
||||
{
|
||||
"value": {
|
||||
"action_requests": [
|
||||
{"name": "execute", "args": {"command": "ls"}, "id": "tc1"}
|
||||
],
|
||||
"review_configs": [
|
||||
{
|
||||
"action_name": "execute",
|
||||
"allowed_decisions": ["approve", "reject"],
|
||||
}
|
||||
],
|
||||
},
|
||||
"ns": ["main"],
|
||||
"resumable": True,
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
chunks = [
|
||||
@@ -336,33 +397,41 @@ class TestInterruptEventParsing:
|
||||
# Channel consumer HITL helpers
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestConsumerHitlHelpers:
|
||||
def test_parse_approval_approve(self):
|
||||
from EvoScientist.channels.consumer import _parse_approval_reply
|
||||
|
||||
for text in ("1", "y", "yes", "approve", "ok", " 1 ", " Y "):
|
||||
assert _parse_approval_reply(text) == "approve", f"Failed for: {text!r}"
|
||||
|
||||
def test_parse_approval_reject(self):
|
||||
from EvoScientist.channels.consumer import _parse_approval_reply
|
||||
|
||||
for text in ("2", "n", "no", "reject"):
|
||||
assert _parse_approval_reply(text) == "reject", f"Failed for: {text!r}"
|
||||
|
||||
def test_parse_approval_auto(self):
|
||||
from EvoScientist.channels.consumer import _parse_approval_reply
|
||||
|
||||
for text in ("3", "a", "auto", "approve all"):
|
||||
assert _parse_approval_reply(text) == "auto", f"Failed for: {text!r}"
|
||||
|
||||
def test_parse_approval_unrecognized(self):
|
||||
from EvoScientist.channels.consumer import _parse_approval_reply
|
||||
|
||||
assert _parse_approval_reply("hello world") is None
|
||||
assert _parse_approval_reply("") is None
|
||||
assert _parse_approval_reply("maybe") is None
|
||||
|
||||
def test_format_approval_prompt(self):
|
||||
from EvoScientist.channels.consumer import _format_approval_prompt
|
||||
prompt = _format_approval_prompt([
|
||||
{"name": "execute", "args": {"command": "ls -la"}},
|
||||
])
|
||||
|
||||
prompt = _format_approval_prompt(
|
||||
[
|
||||
{"name": "execute", "args": {"command": "ls -la"}},
|
||||
]
|
||||
)
|
||||
assert "Approval Required" in prompt
|
||||
assert "execute" in prompt
|
||||
assert "ls -la" in prompt
|
||||
@@ -371,53 +440,67 @@ class TestConsumerHitlHelpers:
|
||||
|
||||
def test_format_approval_prompt_multiple(self):
|
||||
from EvoScientist.channels.consumer import _format_approval_prompt
|
||||
prompt = _format_approval_prompt([
|
||||
{"name": "execute", "args": {"command": "ls"}},
|
||||
{"name": "write_file", "args": {"path": "/out.txt"}},
|
||||
])
|
||||
|
||||
prompt = _format_approval_prompt(
|
||||
[
|
||||
{"name": "execute", "args": {"command": "ls"}},
|
||||
{"name": "write_file", "args": {"path": "/out.txt"}},
|
||||
]
|
||||
)
|
||||
assert "1. execute: ls" in prompt
|
||||
assert "2. write_file: /out.txt" in prompt
|
||||
|
||||
def test_should_auto_approve_non_execute(self):
|
||||
from EvoScientist.channels.consumer import _should_auto_approve
|
||||
|
||||
assert _should_auto_approve([{"name": "write_file", "args": {}}]) is True
|
||||
|
||||
def test_should_auto_approve_empty(self):
|
||||
from EvoScientist.channels.consumer import _should_auto_approve
|
||||
|
||||
assert _should_auto_approve([]) is True
|
||||
|
||||
def test_should_auto_approve_execute_no_allowlist(self):
|
||||
from EvoScientist.channels.consumer import _should_auto_approve
|
||||
|
||||
# With default config (auto_approve=False, shell_allow_list=""),
|
||||
# execute should NOT auto-approve
|
||||
mock_cfg = MagicMock()
|
||||
mock_cfg.auto_approve = False
|
||||
mock_cfg.shell_allow_list = ""
|
||||
with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg):
|
||||
result = _should_auto_approve([
|
||||
{"name": "execute", "args": {"command": "rm -rf /"}},
|
||||
])
|
||||
result = _should_auto_approve(
|
||||
[
|
||||
{"name": "execute", "args": {"command": "rm -rf /"}},
|
||||
]
|
||||
)
|
||||
assert result is False
|
||||
|
||||
def test_should_auto_approve_config_true(self):
|
||||
from EvoScientist.channels.consumer import _should_auto_approve
|
||||
|
||||
mock_cfg = MagicMock()
|
||||
mock_cfg.auto_approve = True
|
||||
with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg):
|
||||
result = _should_auto_approve([
|
||||
{"name": "execute", "args": {"command": "rm -rf /"}},
|
||||
])
|
||||
result = _should_auto_approve(
|
||||
[
|
||||
{"name": "execute", "args": {"command": "rm -rf /"}},
|
||||
]
|
||||
)
|
||||
assert result is True
|
||||
|
||||
def test_should_auto_approve_allowlist_match(self):
|
||||
from EvoScientist.channels.consumer import _should_auto_approve
|
||||
|
||||
mock_cfg = MagicMock()
|
||||
mock_cfg.auto_approve = False
|
||||
mock_cfg.shell_allow_list = "ls,python"
|
||||
with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg):
|
||||
result = _should_auto_approve([
|
||||
{"name": "execute", "args": {"command": "ls -la"}},
|
||||
])
|
||||
result = _should_auto_approve(
|
||||
[
|
||||
{"name": "execute", "args": {"command": "ls -la"}},
|
||||
]
|
||||
)
|
||||
assert result is True
|
||||
|
||||
|
||||
@@ -425,6 +508,7 @@ class TestConsumerHitlHelpers:
|
||||
# Channel HITL intercept mechanism (channel.py)
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestChannelHitlIntercept:
|
||||
def test_register_and_set_hitl_reply(self):
|
||||
from EvoScientist.cli.channel import (
|
||||
@@ -432,6 +516,7 @@ class TestChannelHitlIntercept:
|
||||
_try_set_hitl_reply,
|
||||
_pop_hitl_reply,
|
||||
)
|
||||
|
||||
event = _register_hitl_wait("telegram", "chat123")
|
||||
assert not event.is_set()
|
||||
|
||||
@@ -445,11 +530,13 @@ class TestChannelHitlIntercept:
|
||||
|
||||
def test_try_set_hitl_reply_no_pending(self):
|
||||
from EvoScientist.cli.channel import _try_set_hitl_reply
|
||||
|
||||
# No pending HITL — should not intercept
|
||||
assert _try_set_hitl_reply("discord", "no_pending", "y") is False
|
||||
|
||||
def test_pop_hitl_reply_no_pending(self):
|
||||
from EvoScientist.cli.channel import _pop_hitl_reply
|
||||
|
||||
assert _pop_hitl_reply("discord", "no_pending") is None
|
||||
|
||||
def test_hitl_reply_timeout(self):
|
||||
@@ -457,6 +544,7 @@ class TestChannelHitlIntercept:
|
||||
_register_hitl_wait,
|
||||
_pop_hitl_reply,
|
||||
)
|
||||
|
||||
event = _register_hitl_wait("telegram", "timeout_chat")
|
||||
# Don't set reply — simulate timeout
|
||||
replied = event.wait(timeout=0.01)
|
||||
@@ -470,10 +558,12 @@ class TestChannelHitlIntercept:
|
||||
# _resolve_hitl_approval with custom prompt_fn
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestResolveHitlApprovalWithPromptFn:
|
||||
def test_prompt_fn_called_for_execute(self):
|
||||
from EvoScientist.stream.display import _resolve_hitl_approval
|
||||
import EvoScientist.stream.display as disp
|
||||
|
||||
original = disp._session_auto_approve
|
||||
try:
|
||||
disp._session_auto_approve = False
|
||||
@@ -482,9 +572,15 @@ class TestResolveHitlApprovalWithPromptFn:
|
||||
mock_cfg.shell_allow_list = ""
|
||||
custom_decisions = [{"type": "approve"}]
|
||||
mock_fn = MagicMock(return_value=custom_decisions)
|
||||
with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg):
|
||||
with patch(
|
||||
"EvoScientist.config.settings.load_config", return_value=mock_cfg
|
||||
):
|
||||
result = _resolve_hitl_approval(
|
||||
{"action_requests": [{"name": "execute", "args": {"command": "rm -rf /"}}]},
|
||||
{
|
||||
"action_requests": [
|
||||
{"name": "execute", "args": {"command": "rm -rf /"}}
|
||||
]
|
||||
},
|
||||
prompt_fn=mock_fn,
|
||||
)
|
||||
assert result == custom_decisions
|
||||
@@ -495,15 +591,22 @@ class TestResolveHitlApprovalWithPromptFn:
|
||||
def test_prompt_fn_not_called_for_auto_approve(self):
|
||||
from EvoScientist.stream.display import _resolve_hitl_approval
|
||||
import EvoScientist.stream.display as disp
|
||||
|
||||
original = disp._session_auto_approve
|
||||
try:
|
||||
disp._session_auto_approve = False
|
||||
mock_cfg = MagicMock()
|
||||
mock_cfg.auto_approve = True
|
||||
mock_fn = MagicMock()
|
||||
with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg):
|
||||
with patch(
|
||||
"EvoScientist.config.settings.load_config", return_value=mock_cfg
|
||||
):
|
||||
result = _resolve_hitl_approval(
|
||||
{"action_requests": [{"name": "execute", "args": {"command": "rm"}}]},
|
||||
{
|
||||
"action_requests": [
|
||||
{"name": "execute", "args": {"command": "rm"}}
|
||||
]
|
||||
},
|
||||
prompt_fn=mock_fn,
|
||||
)
|
||||
assert result == [{"type": "approve"}]
|
||||
@@ -514,6 +617,7 @@ class TestResolveHitlApprovalWithPromptFn:
|
||||
def test_prompt_fn_not_called_for_non_execute(self):
|
||||
from EvoScientist.stream.display import _resolve_hitl_approval
|
||||
import EvoScientist.stream.display as disp
|
||||
|
||||
original = disp._session_auto_approve
|
||||
try:
|
||||
disp._session_auto_approve = False
|
||||
@@ -521,7 +625,9 @@ class TestResolveHitlApprovalWithPromptFn:
|
||||
mock_cfg.auto_approve = False
|
||||
mock_cfg.shell_allow_list = ""
|
||||
mock_fn = MagicMock()
|
||||
with patch("EvoScientist.config.settings.load_config", return_value=mock_cfg):
|
||||
with patch(
|
||||
"EvoScientist.config.settings.load_config", return_value=mock_cfg
|
||||
):
|
||||
result = _resolve_hitl_approval(
|
||||
{"action_requests": [{"name": "write_file", "args": {}}]},
|
||||
prompt_fn=mock_fn,
|
||||
|
||||
+19
-7
@@ -37,13 +37,26 @@ class TestModelsRegistry:
|
||||
|
||||
def test_entries_are_valid_tuples(self):
|
||||
"""Test that _MODEL_ENTRIES contains valid (name, model_id, provider) tuples."""
|
||||
valid_providers = {"anthropic", "openai", "google-genai", "nvidia", "siliconflow", "openrouter", "zhipu", "zhipu-code"}
|
||||
valid_providers = {
|
||||
"anthropic",
|
||||
"openai",
|
||||
"google-genai",
|
||||
"nvidia",
|
||||
"siliconflow",
|
||||
"openrouter",
|
||||
"zhipu",
|
||||
"zhipu-code",
|
||||
"custom-openai",
|
||||
"custom-anthropic",
|
||||
}
|
||||
for entry in _MODEL_ENTRIES:
|
||||
assert len(entry) == 3, f"Entry {entry} doesn't have 3 elements"
|
||||
name, model_id, provider = entry
|
||||
assert isinstance(name, str)
|
||||
assert isinstance(model_id, str)
|
||||
assert provider in valid_providers, f"Unknown provider '{provider}' for '{name}'"
|
||||
assert provider in valid_providers, (
|
||||
f"Unknown provider '{provider}' for '{name}'"
|
||||
)
|
||||
|
||||
def test_get_models_for_provider(self):
|
||||
"""Test that get_models_for_provider returns correct models."""
|
||||
@@ -150,7 +163,7 @@ class TestGetChatModel:
|
||||
get_chat_model("claude-opus-4-5")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["model"] == "claude-opus-4-5-20251101"
|
||||
assert call_kwargs["model"] == "claude-opus-4-5"
|
||||
assert call_kwargs["model_provider"] == "anthropic"
|
||||
|
||||
@patch("EvoScientist.llm.models.init_chat_model")
|
||||
@@ -382,10 +395,10 @@ class TestThirdPartyRouting:
|
||||
def test_custom_routes_through_openai(self, mock_init, monkeypatch):
|
||||
"""Custom provider should route through OpenAI with env-configured base_url."""
|
||||
mock_init.return_value = "mock_model"
|
||||
monkeypatch.setenv("CUSTOM_BASE_URL", "https://my-llm.example.com/v1")
|
||||
monkeypatch.setenv("CUSTOM_API_KEY", "custom-key-789")
|
||||
monkeypatch.setenv("CUSTOM_OPENAI_BASE_URL", "https://my-llm.example.com/v1")
|
||||
monkeypatch.setenv("CUSTOM_OPENAI_API_KEY", "custom-key-789")
|
||||
|
||||
get_chat_model("my-custom-model", provider="custom")
|
||||
get_chat_model("my-custom-model", provider="custom-openai")
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["model_provider"] == "openai"
|
||||
@@ -530,4 +543,3 @@ class TestAutoConfig:
|
||||
|
||||
call_kwargs = mock_init.call_args[1]
|
||||
assert call_kwargs["include_thoughts"] is True
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user