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:
X-iZhang
2026-03-15 21:12:51 +00:00
parent f27521d62b
commit c5a4d559a2
116 changed files with 4592 additions and 2236 deletions
+4 -2
View File
@@ -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)
+3
View File
@@ -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
View File
@@ -1,4 +1,5 @@
"""Enable `python -m EvoScientist` execution."""
from EvoScientist.cli import main
main()
+46 -26
View File
@@ -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:
+6 -1
View File
@@ -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
View File
@@ -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)
+6 -5
View File
@@ -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."""
+24 -24
View File
@@ -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"),
)
+36 -26
View File
@@ -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
+7 -6
View File
@@ -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:
+128 -79
View File
@@ -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]
+8 -6
View File
@@ -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)
+144 -62
View File
@@ -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"![image]({file_path})" + (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"![image]({file_path})"
+ (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 -6
View File
@@ -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)
+32 -16
View File
@@ -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,
)
)
+3 -1
View File
@@ -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:
+21 -19
View File
@@ -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)
+107 -33
View File
@@ -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", "")
+10 -2
View File
@@ -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():
+12 -10
View File
@@ -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)
+147 -100
View File
@@ -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:
+8
View File
@@ -98,10 +98,12 @@ def convert_markdown(
return text
# ═════════════════════════════════════════════════════════════════════
# Shared helpers
# ═════════════════════════════════════════════════════════════════════
def _escape_html(text: str) -> str:
return text.replace("&", "&amp;").replace("<", "&lt;").replace(">", "&gt;")
@@ -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.
+21 -15
View File
@@ -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.
+5 -2
View File
@@ -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
+16 -9
View File
@@ -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)
+54 -12
View File
@@ -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
+28 -10
View File
@@ -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.
+22 -4
View File
@@ -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]
+7 -5
View File
@@ -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)
+49 -26
View File
@@ -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()
+9 -7
View File
@@ -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)
+9 -7
View File
@@ -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)
+82 -34
View File
@@ -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 ────────────────────────────────────────────────
+5 -1
View File
@@ -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()}"
+9 -7
View File
@@ -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)
+41 -30
View File
@@ -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]}"
+16 -7
View File
@@ -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 -5
View File
@@ -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)
+63 -39
View File
@@ -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
+3 -1
View File
@@ -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:
+102 -65
View File
@@ -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)
+16 -10
View File
@@ -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")
+3 -1
View File
@@ -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()
+3 -1
View File
@@ -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
+12 -3
View File
@@ -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
View File
@@ -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()]
+6 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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")
+15 -5
View File
@@ -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", []))
File diff suppressed because it is too large Load Diff
+6 -3
View File
@@ -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),
+18 -4
View File
@@ -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")
+6 -2
View File
@@ -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}"
+19 -5
View File
@@ -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)."""
+6 -2
View File
@@ -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)))
+5 -1
View File
@@ -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:
+10 -10
View File
@@ -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),
)
)
+8 -2
View File
@@ -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", ""))
+1
View File
@@ -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}")
+54 -23
View File
@@ -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)
+20 -8
View File
@@ -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
View File
@@ -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:
+1 -2
View File
@@ -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:
+72 -18
View File
@@ -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
View File
@@ -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")
+1
View File
@@ -330,6 +330,7 @@ Finding one with context [1]. Another insight [2].
# Combined exports
# =============================================================================
def get_system_prompt() -> str:
"""Generate the complete system prompt.
+4
View File
@@ -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,
+21 -17
View File
@@ -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
View File
@@ -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'>&gt; Your answer:</style></b> ")).strip()
raw = pt_prompt(
HTML(" <b><style fg='#42a5f5'>&gt; 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'>&gt; Answer:</style></b> ")).strip()
raw = pt_prompt(
HTML(" <b><style fg='#42a5f5'>&gt; 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):
+83 -48
View File
@@ -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
View File
@@ -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)"
+36 -25
View File
@@ -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)
+15 -10
View File
@@ -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
+1
View File
@@ -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)
+13 -13
View File
@@ -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:
+9 -2
View File
@@ -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'."
)
+15 -10
View File
@@ -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
+3 -2
View File
@@ -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
View File
@@ -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."},
]
+7 -2
View File
@@ -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):
+1 -3
View File
@@ -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
View File
@@ -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)
+48 -41
View File
@@ -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"
+5 -1
View File
@@ -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()
File diff suppressed because it is too large Load Diff
+6 -2
View File
@@ -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,
+6 -2
View File
@@ -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
View File
@@ -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,
+8 -5
View File
@@ -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"
+5
View File
@@ -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
+5
View File
@@ -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
+4 -2
View File
@@ -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,
+11 -9
View File
@@ -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
View File
@@ -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
View File
@@ -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