1 Commits

Author SHA1 Message Date
m4 e0acc6155e feat: improve WebUI run recovery
Build / build (push) Has been cancelled
Docker / build (push) Has been cancelled
Lint / ruff (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.11) (push) Has been cancelled
Test / pytest (ubuntu-latest, 3.12) (push) Has been cancelled
Test / pytest (windows-latest, 3.11) (push) Has been cancelled
Test / pytest (windows-latest, 3.12) (push) Has been cancelled
2026-07-10 17:35:44 +08:00
119 changed files with 4902 additions and 10284 deletions
+1 -1
View File
@@ -5,5 +5,5 @@
<rect x="54" y="5" width="62" height="24" rx="6" fill="#1565c0"/>
<text x="85" y="22" text-anchor="middle"
font-family="Inter, -apple-system, system-ui, sans-serif"
font-size="13" font-weight="700" fill="#ffffff">v0.2.2</text>
font-size="13" font-weight="700" fill="#ffffff">v0.2.1</text>
</svg>

Before

Width:  |  Height:  |  Size: 555 B

After

Width:  |  Height:  |  Size: 555 B

+1 -1
View File
@@ -5,5 +5,5 @@
<rect x="54" y="5" width="62" height="24" rx="6" fill="#2563eb"/>
<text x="85" y="22" text-anchor="middle"
font-family="Inter, -apple-system, system-ui, sans-serif"
font-size="13" font-weight="700" fill="#ffffff">v0.2.2</text>
font-size="13" font-weight="700" fill="#ffffff">v0.2.1</text>
</svg>

Before

Width:  |  Height:  |  Size: 555 B

After

Width:  |  Height:  |  Size: 555 B

Binary file not shown.

Before

Width:  |  Height:  |  Size: 287 KiB

After

Width:  |  Height:  |  Size: 286 KiB

+1
View File
@@ -48,3 +48,4 @@ conversation_history/
*meals/
botpy.log
large_tool_results/
runs/
+25 -136
View File
@@ -19,7 +19,6 @@ Usage:
import json
import logging
import os
from collections.abc import Sequence
from pathlib import Path
from typing import TYPE_CHECKING
@@ -305,12 +304,8 @@ def _inject_subagent_middleware(
path doesn't fall back to the global-writing ``_ensure_chat_model()``.
"""
from .middleware import (
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
ContextOverflowMapperMiddleware,
ErrorNormalizationMiddleware,
RepetitiveToolCallGuardMiddleware,
ToolErrorHandlerMiddleware,
ToolProtocolGuardMiddleware,
create_context_editing_middleware,
create_memory_lifecycle_middleware,
create_memory_middleware,
@@ -319,16 +314,6 @@ def _inject_subagent_middleware(
)
cfg = cfg if cfg is not None else _ensure_config()
repetitive_tool_call_threshold = getattr(
cfg,
"repetitive_tool_call_threshold",
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
)
if not isinstance(repetitive_tool_call_threshold, int):
repetitive_tool_call_threshold = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD
max_consecutive_tool_errors = getattr(cfg, "max_consecutive_tool_errors", 3)
if not isinstance(max_consecutive_tool_errors, int):
max_consecutive_tool_errors = 3
memory_controls = MemoryControls.from_config(cfg)
memory_dir = str(_paths_mod.MEMORIES_DIR)
memory_scheduler = default_memory_scheduler()
@@ -348,16 +333,6 @@ def _inject_subagent_middleware(
memory_scheduler=memory_scheduler,
)
middleware = [
# Outermost — catches provider-SDK exceptions from the
# model call (including inner middlewares) and normalizes
# them into a non-dataclass envelope wrapper before
# anything downstream sees them.
ErrorNormalizationMiddleware(),
RepetitiveToolCallGuardMiddleware(
threshold=repetitive_tool_call_threshold,
max_consecutive_errors=max_consecutive_tool_errors,
),
ToolProtocolGuardMiddleware(),
# Subagents share the main agent's model: use the threaded
# ``chat_model`` on the pure path, else defer to the factory's
# ``_ensure_chat_model()`` fallback (when ``chat_model=None``).
@@ -666,14 +641,9 @@ def _get_default_middleware(
*,
for_async_subagent: bool = False,
workspace_dir: str | Path | None = None,
memory_dir: str | Path | None = None,
cfg=None,
chat_model=None,
memory_source_agent: str = "EvoScientist",
tool_selector_threshold: int | None = None,
memory_max_inline_profile_chars: int | None = None,
enable_background_execution: bool = True,
enable_legacy_model_fallback: bool = True,
):
"""Build the default middleware list.
@@ -695,14 +665,10 @@ def _get_default_middleware(
Async sub-agent factories pass their deployed agent name here.
"""
from .middleware import (
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
ConfigurableModelMiddleware,
ContextOverflowMapperMiddleware,
ErrorNormalizationMiddleware,
ModelFallbackMiddleware,
RepetitiveToolCallGuardMiddleware,
ToolErrorHandlerMiddleware,
ToolProtocolGuardMiddleware,
create_code_interpreter_middleware,
create_context_editing_middleware,
create_memory_lifecycle_middleware,
@@ -715,20 +681,10 @@ def _get_default_middleware(
)
cfg = cfg if cfg is not None else _ensure_config()
repetitive_tool_call_threshold = getattr(
cfg,
"repetitive_tool_call_threshold",
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
)
if not isinstance(repetitive_tool_call_threshold, int):
repetitive_tool_call_threshold = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD
max_consecutive_tool_errors = getattr(cfg, "max_consecutive_tool_errors", 3)
if not isinstance(max_consecutive_tool_errors, int):
max_consecutive_tool_errors = 3
if cfg.model_fallbacks:
load_fallback_chain(cfg.model_fallbacks)
model = chat_model if chat_model is not None else _ensure_chat_model()
memory_dir = str(memory_dir or _paths_mod.MEMORIES_DIR)
memory_dir = str(_paths_mod.MEMORIES_DIR)
source_type = (
MemorySourceType.SUBAGENT if for_async_subagent else MemorySourceType.TURN
)
@@ -743,20 +699,18 @@ def _get_default_middleware(
# ``ModelFallbackMiddleware``: a configurable.model override sets the
# PRIMARY model only, leaving the fallback chain free to try its own
# alternatives instead of re-overriding every retry to the same model.
memory_kwargs = {
"workspace_dir": workspace_dir,
"source_type": source_type,
"source_agent": memory_source_agent,
"enable_profile_memory": memory_controls.profile_enabled,
"enable_observation_memory": memory_controls.observations_enabled,
"enable_observation_tool": memory_controls.observation_tool_enabled(
memory_middleware = create_memory_middleware(
memory_dir,
workspace_dir=workspace_dir,
source_type=source_type,
source_agent=memory_source_agent,
enable_profile_memory=memory_controls.profile_enabled,
enable_observation_memory=memory_controls.observations_enabled,
enable_observation_tool=memory_controls.observation_tool_enabled(
MemoryObservationTarget.AGENT
),
"memory_scheduler": memory_scheduler,
}
if memory_max_inline_profile_chars is not None:
memory_kwargs["max_inline_profile_chars"] = memory_max_inline_profile_chars
memory_middleware = create_memory_middleware(memory_dir, **memory_kwargs)
memory_scheduler=memory_scheduler,
)
# Main-agent tool selection may use the auxiliary model; async sub-agents
# keep the main model (they do real work, not a one-off helper call).
# context_editing stays on the main model — its model only sizes the
@@ -774,32 +728,16 @@ def _get_default_middleware(
from .llm import get_chat_model
tool_selector_model = get_chat_model(model=aux_model, provider=aux_provider)
selector_middlewares = create_tool_selector_middleware(
**(
{"threshold": tool_selector_threshold}
if tool_selector_threshold is not None
else {}
),
model=tool_selector_model,
track_stream_selection=not for_async_subagent,
)
mw = [
# Outermost — catches provider-SDK exceptions from the model
# call (including exceptions surfaced through inner
# middlewares) and normalizes them into a non-dataclass
# envelope wrapper before anything downstream sees them.
ErrorNormalizationMiddleware(),
ConfigurableModelMiddleware(),
create_context_editing_middleware(model),
*([ModelFallbackMiddleware()] if enable_legacy_model_fallback else []),
RepetitiveToolCallGuardMiddleware(
threshold=repetitive_tool_call_threshold,
max_consecutive_errors=max_consecutive_tool_errors,
),
ModelFallbackMiddleware(),
ContextOverflowMapperMiddleware(),
ToolErrorHandlerMiddleware(),
*selector_middlewares,
ToolProtocolGuardMiddleware(),
*create_tool_selector_middleware(
model=tool_selector_model,
track_stream_selection=not for_async_subagent,
),
# Interpreter prompt must land before runtime/memory context, so this
# middleware sits ahead of runtime_context in the stack.
create_code_interpreter_middleware(
@@ -832,7 +770,7 @@ def _get_default_middleware(
# Background-process tools (run_in_background / check_process / stop_process /
# list_processes) — main agent only. Async sub-agents run on langgraph-dev and
# must not spawn local OS processes.
if not for_async_subagent and enable_background_execution:
if not for_async_subagent:
from .middleware.background import BackgroundExecutionMiddleware
mw.append(BackgroundExecutionMiddleware())
@@ -930,14 +868,6 @@ def create_cli_agent(
chat_model=None,
*,
on_mcp_progress=None,
workspace_backend=None,
memory_dir: str | Path | None = None,
tool_selector_threshold: int | None = None,
memory_max_inline_profile_chars: int | None = None,
enable_subagents: bool = True,
enable_background_execution: bool = True,
main_agent_outer_middlewares: Sequence[AgentMiddleware] | None = None,
main_agent_route_middleware: AgentMiddleware | None = None,
) -> "CompiledStateGraph":
"""Create agent with checkpointer for CLI multi-turn support.
@@ -964,22 +894,6 @@ def create_cli_agent(
chat_model: Optional pre-built chat model. Only triggers the pure
path when ``config`` is also explicit; otherwise it is ignored in
favor of the ``_ensure_chat_model()`` fallback.
workspace_backend: Optional host-provided backend for the workspace
route. The default remains ``CustomSandboxBackend``.
memory_dir: Optional memory root used by both the backend route and
memory middleware.
tool_selector_threshold: Optional adaptive tool-selection threshold.
memory_max_inline_profile_chars: Optional memory profile injection cap.
enable_subagents: Whether configured subagents are available to the agent.
enable_background_execution: Whether local background-process tools are
installed. Embedding hosts should disable this when process execution
is provided by an external backend.
main_agent_outer_middlewares: Optional host-owned middleware installed
only on the top-level agent, outside EvoScientist's default chain.
main_agent_route_middleware: Optional host-owned route middleware placed
after ConfigurableModelMiddleware and before tool selection. When
provided, EvoScientist's legacy model fallback is disabled for the
top-level agent so the host is the only fallback authority.
"""
import os as _os
@@ -1021,21 +935,19 @@ def create_cli_agent(
workspace_dir = str(_paths.WORKSPACE_ROOT)
# Read paths dynamically so runtime set_workspace_root() changes are picked up
_mem_dir = str(memory_dir or _paths.MEMORIES_DIR)
_mem_dir = str(_paths.MEMORIES_DIR)
_usr_skills_dir = str(_paths.USER_SKILLS_DIR)
_global_skills_dir = str(_paths.GLOBAL_SKILLS_DIR)
# Always construct fresh backends from current paths (avoids stale
# module-level backend when workspace root changed at runtime).
set_active_workspace(workspace_dir)
ws_backend = workspace_backend
if ws_backend is None:
ws_backend = CustomSandboxBackend(
root_dir=workspace_dir,
virtual_mode=True,
timeout=cfg.sandbox_execute_timeout,
dangerous=cfg.dangerous_mode,
)
ws_backend = CustomSandboxBackend(
root_dir=workspace_dir,
virtual_mode=True,
timeout=cfg.sandbox_execute_timeout,
dangerous=cfg.dangerous_mode,
)
sk_backend = MergedSkillsBackend(
primary_dir=_usr_skills_dir,
global_dir=_global_skills_dir,
@@ -1057,29 +969,8 @@ def create_cli_agent(
# CLI agent never drifts from the default chain. Anything CLI-specific
# (e.g. ``HumanInTheLoopMiddleware``) is appended below.
mw: list[AgentMiddleware] = _get_default_middleware(
workspace_dir=workspace_dir,
memory_dir=_mem_dir,
cfg=cfg,
chat_model=chat_model,
tool_selector_threshold=tool_selector_threshold,
memory_max_inline_profile_chars=memory_max_inline_profile_chars,
enable_background_execution=enable_background_execution,
enable_legacy_model_fallback=main_agent_route_middleware is None,
workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model
)
if main_agent_route_middleware is not None:
configurable_index = next(
(
index
for index, middleware in enumerate(mw)
if getattr(middleware, "name", "") == "configurable_model"
),
None,
)
if configurable_index is None:
raise RuntimeError("ConfigurableModelMiddleware route slot is unavailable")
mw.insert(configurable_index + 1, main_agent_route_middleware)
if main_agent_outer_middlewares:
mw = [*main_agent_outer_middlewares, *mw]
# HITL on main agent only — passing `interrupt_on=` to create_deep_agent
# would propagate it to every subagent, breaking parallel execute calls
@@ -1104,8 +995,6 @@ def create_cli_agent(
chat_model=chat_model,
workspace_dir=workspace_dir,
)
if not enable_subagents:
kwargs = {**kwargs, "subagents": []}
return create_deep_agent(
**kwargs,
-2
View File
@@ -9,8 +9,6 @@ from __future__ import annotations
from importlib import import_module
__version__ = "0.2.2"
_EXPORTS: dict[str, tuple[str, str]] = {
# Agent graph (lazy to avoid expensive initialization at import time)
"EvoScientist_agent": (".EvoScientist", "EvoScientist_agent"),
+6 -57
View File
@@ -20,9 +20,6 @@ from EvoScientist.config import EvoScientistConfig
logger = logging.getLogger(__name__)
_CCPROXY_AUTH_TIMEOUT_SECONDS = 30
_CCPROXY_HEALTH_TIMEOUT_SECONDS = 180
# =============================================================================
# Availability & auth checks
@@ -130,11 +127,7 @@ def check_ccproxy_auth(provider: str = "claude_api") -> tuple[bool, str]:
[exe, "auth", "status", provider],
capture_output=True,
text=True,
# ccproxy's CLI initializes its full plugin system on every
# invocation — a cold start takes ~10s on Apple Silicon, so a
# 10s timeout made OAuth startup fail intermittently with
# "Auth check timed out".
timeout=_CCPROXY_AUTH_TIMEOUT_SECONDS,
timeout=10,
)
import re as _re
@@ -183,33 +176,6 @@ def is_ccproxy_running(port: int) -> bool:
return False
def write_ccproxy_config() -> str:
"""Write the ccproxy config file EvoScientist passes to ``serve --config``.
Disables ccproxy's default Codex model mappings, which rewrite any
``gpt-*``/``o1-*``/``o3-*``/``claude-*`` model to ``gpt-5.3-codex``
before forwarding — silently overriding the model the user configured
(and failing outright on accounts where ``gpt-5.3-codex`` is not
served). With no mappings, the requested model reaches the Codex
backend unmodified.
Returns:
Absolute path to the generated config file.
"""
from EvoScientist.config import get_config_dir
path = get_config_dir() / "ccproxy.toml"
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(
"# Generated by EvoScientist (ccproxy_manager) — do not edit;\n"
"# regenerated on every ccproxy start.\n"
"[plugins.codex]\n"
"model_mappings = []\n",
encoding="utf-8",
)
return str(path)
def start_ccproxy(port: int) -> subprocess.Popen:
"""Start ccproxy serve as a background process.
@@ -220,32 +186,18 @@ def start_ccproxy(port: int) -> subprocess.Popen:
The Popen handle for the ccproxy process.
Raises:
RuntimeError: If ccproxy fails to become healthy within
``_CCPROXY_HEALTH_TIMEOUT_SECONDS``.
RuntimeError: If ccproxy fails to become healthy within 30 seconds.
FileNotFoundError: If ccproxy binary is not found.
"""
exe = _ccproxy_exe() or "ccproxy"
cmd = [exe, "serve", "--port", str(port)]
try:
cmd += ["--config", write_ccproxy_config()]
except (OSError, UnicodeError) as exc:
logger.warning(
"Could not write ccproxy config (%s); starting with defaults — "
"Codex model mappings will rewrite gpt-* models to gpt-5.3-codex",
exc,
)
logger.warning(
"Starting ccproxy on port %d; first startup may take up to %d seconds",
port,
_CCPROXY_HEALTH_TIMEOUT_SECONDS,
)
proc = subprocess.Popen(
cmd,
[exe, "serve", "--port", str(port)],
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
)
deadline = time.monotonic() + _CCPROXY_HEALTH_TIMEOUT_SECONDS
# Wait for health (ccproxy can take up to ~11s on first start)
deadline = time.monotonic() + 30
while time.monotonic() < deadline:
if proc.poll() is not None:
raise RuntimeError(
@@ -261,10 +213,7 @@ def start_ccproxy(port: int) -> subprocess.Popen:
proc.wait(timeout=3)
except subprocess.TimeoutExpired:
proc.kill()
raise RuntimeError(
"ccproxy did not become healthy within "
f"{_CCPROXY_HEALTH_TIMEOUT_SECONDS} seconds"
)
raise RuntimeError("ccproxy did not become healthy within 30 seconds")
def stop_ccproxy(proc: subprocess.Popen | None) -> None:
+5 -10
View File
@@ -75,13 +75,11 @@ class DedupCache:
max_size: int = _DEDUP_MAX,
trim_to: int = _DEDUP_TRIM,
ttl_seconds: float = _DEDUP_TTL,
clock: Callable[[], float] | None = None,
) -> None:
self._seen: OrderedDict[str, float] = OrderedDict()
self._max = max_size
self._trim = trim_to
self._ttl = ttl_seconds
self._clock = clock or time.monotonic
# ── public API ──────────────────────────────────────────────────
@@ -95,16 +93,15 @@ class DedupCache:
if not msg_id:
return False
now = self._clock()
self._prune(now)
self._prune()
if msg_id in self._seen:
# LRU: refresh position and timestamp
self._seen.move_to_end(msg_id)
self._seen[msg_id] = now
self._seen[msg_id] = time.monotonic()
return True
self._seen[msg_id] = now
self._seen[msg_id] = time.monotonic()
if len(self._seen) > self._max:
while len(self._seen) > self._trim:
self._seen.popitem(last=False)
@@ -121,9 +118,9 @@ class DedupCache:
# ── internal ────────────────────────────────────────────────────
def _prune(self, now: float | None = None) -> None:
def _prune(self) -> None:
"""Remove entries older than *ttl_seconds*."""
cutoff = (self._clock() if now is None else now) - self._ttl
cutoff = time.monotonic() - self._ttl
# OrderedDict is insertion-ordered; oldest entries are first.
while self._seen:
_key, ts = next(iter(self._seen.items()))
@@ -431,13 +428,11 @@ class DedupMiddleware(InboundMiddleware):
max_size: int = 1000,
trim_to: int = 500,
ttl_seconds: float = 3600.0,
clock: Callable[[], float] | None = None,
) -> None:
self._cache = DedupCache(
max_size=max_size,
trim_to=trim_to,
ttl_seconds=ttl_seconds,
clock=clock,
)
async def process_inbound(
+109 -66
View File
@@ -371,31 +371,14 @@ def _step_minimax_region(config: EvoScientistConfig) -> str:
return _MINIMAX_REGIONS[region]
def _step_oauth_auth_mode(
config: EvoScientistConfig,
*,
provider_label: str,
ccproxy_provider: str,
config_attr: str,
prompt_login_label: str,
oauth_choice_label: str | None = None,
status_label: str | None = None,
question_label: str | None = None,
) -> str:
"""Select API-key vs ccproxy OAuth authentication for a provider.
def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str:
"""Step 2a: Select Anthropic authentication mode (API key vs OAuth).
Args:
config: Current configuration.
provider_label: Provider display name for direct API-key access.
ccproxy_provider: ccproxy auth provider name.
config_attr: Config attribute storing this provider's auth mode.
prompt_login_label: Label used in "Log in to ..." prompts.
oauth_choice_label: Optional display label for the OAuth choice.
status_label: Optional display label for status messages.
question_label: Optional prompt label override.
Returns:
Selected auth mode: "api_key" or "oauth".
Selected auth mode: "api_key", "oauth", or "auto".
"""
from ...ccproxy_manager import check_ccproxy_auth, is_ccproxy_available
@@ -403,14 +386,10 @@ def _step_oauth_auth_mode(
from .prompter import BACK_SENTINEL, GoBack, install_navigation_keys
oauth_label = oauth_choice_label or f"{prompt_login_label} OAuth"
auth_status_label = status_label or oauth_label
auth_question_label = question_label or f"{provider_label} authentication mode"
choices = [
Choice(title=f"API Key (direct {provider_label} access)", value="api_key"),
Choice(title="API Key (direct Anthropic access)", value="api_key"),
Choice(
title=f"{oauth_label} (via ccproxy — no API key needed)"
title="Claude Code OAuth (via ccproxy — no API key needed)"
+ (
""
if ccproxy_available
@@ -422,12 +401,12 @@ def _step_oauth_auth_mode(
Choice(title="← Back (re-pick provider)", value=BACK_SENTINEL),
]
current = getattr(config, config_attr)
current = config.anthropic_auth_mode
if current not in ("api_key", "oauth"):
current = "api_key"
question = questionary.select(
f"{auth_question_label} [Esc/← to go back]:",
"Authentication mode [Esc/← to go back]:",
choices=choices,
default=current,
style=WIZARD_STYLE,
@@ -469,9 +448,11 @@ def _step_oauth_auth_mode(
if auth_mode == "oauth":
_prompt_ccproxy_port(config)
authed, msg = check_ccproxy_auth(ccproxy_provider)
# If OAuth selected, check auth status and offer login
if auth_mode in ("oauth", "auto"):
authed, msg = check_ccproxy_auth()
if authed:
console.print(f" [green]✓ {auth_status_label}: {msg}[/green]")
console.print(f" [green]✓ OAuth: {msg}[/green]")
relogin = questionary.confirm(
"Re-authenticate to refresh credentials?",
default=False,
@@ -481,13 +462,11 @@ def _step_oauth_auth_mode(
if relogin is None:
raise KeyboardInterrupt()
if relogin:
_run_ccproxy_login(ccproxy_provider, auth_status_label)
_run_ccproxy_login("claude_api", "OAuth")
else:
console.print(
f" [yellow]{auth_status_label} not authenticated: {msg}[/yellow]"
)
console.print(f" [yellow]OAuth not authenticated: {msg}[/yellow]")
login = questionary.confirm(
f"Log in to {prompt_login_label} now?",
"Log in to Claude now?",
default=True,
style=CONFIRM_STYLE,
qmark=QMARK,
@@ -495,32 +474,11 @@ def _step_oauth_auth_mode(
if login is None:
raise KeyboardInterrupt()
if login:
_run_ccproxy_login(ccproxy_provider, auth_status_label)
_run_ccproxy_login("claude_api", "OAuth")
return auth_mode
def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str:
"""Step 2a: Select Anthropic authentication mode (API key vs OAuth).
Args:
config: Current configuration.
Returns:
Selected auth mode: "api_key" or "oauth".
"""
return _step_oauth_auth_mode(
config,
provider_label="Anthropic",
ccproxy_provider="claude_api",
config_attr="anthropic_auth_mode",
prompt_login_label="Claude",
oauth_choice_label="Claude Code OAuth",
status_label="OAuth",
question_label="Authentication mode",
)
def _step_openai_auth_mode(config: EvoScientistConfig) -> str:
"""Step 2b: Select OpenAI authentication mode (API key vs Codex OAuth).
@@ -530,16 +488,101 @@ def _step_openai_auth_mode(config: EvoScientistConfig) -> str:
Returns:
Selected auth mode: "api_key" or "oauth".
"""
return _step_oauth_auth_mode(
config,
provider_label="OpenAI",
ccproxy_provider="codex",
config_attr="openai_auth_mode",
prompt_login_label="Codex",
oauth_choice_label="Codex OAuth",
status_label="Codex OAuth",
question_label="OpenAI authentication mode",
from ...ccproxy_manager import check_ccproxy_auth, is_ccproxy_available
ccproxy_available = is_ccproxy_available()
from .prompter import BACK_SENTINEL, GoBack, install_navigation_keys
choices = [
Choice(title="API Key (direct OpenAI access)", value="api_key"),
Choice(
title="Codex OAuth (via ccproxy — no API key needed)"
+ (
""
if ccproxy_available
else " [requires: pip install evoscientist[oauth]]"
),
value="oauth",
),
questionary.Separator(),
Choice(title="← Back (re-pick provider)", value=BACK_SENTINEL),
]
current = config.openai_auth_mode
if current not in ("api_key", "oauth"):
current = "api_key"
question = questionary.select(
"OpenAI authentication mode [Esc/← to go back]:",
choices=choices,
default=current,
style=WIZARD_STYLE,
qmark=QMARK,
use_indicator=True,
)
install_navigation_keys(question, with_back=True)
auth_mode = question.ask()
if auth_mode is None:
raise KeyboardInterrupt()
if auth_mode == BACK_SENTINEL:
raise GoBack()
if auth_mode == "oauth" and not ccproxy_available:
console.print(" [yellow]✗ ccproxy not installed[/yellow]")
console.print()
install = questionary.confirm(
'Install ccproxy now? (pip install "evoscientist[oauth]")',
default=True,
style=WIZARD_STYLE,
qmark=f" {QMARK}",
).ask()
if install is None:
raise KeyboardInterrupt()
if install:
console.print()
if _install_ccproxy():
console.print(" [green]✓ ccproxy installed successfully.[/green]")
else:
console.print(" [yellow]Falling back to API key mode.[/yellow]")
return "api_key"
else:
console.print(
' [dim]Skipped. Install manually: pip install "evoscientist[oauth]"[/dim]'
)
return "api_key"
# If OAuth selected, prompt for port and check auth status
if auth_mode == "oauth":
_prompt_ccproxy_port(config)
authed, msg = check_ccproxy_auth("codex")
if authed:
console.print(f" [green]✓ Codex OAuth: {msg}[/green]")
relogin = questionary.confirm(
"Re-authenticate to refresh credentials?",
default=False,
style=CONFIRM_STYLE,
qmark=QMARK,
).ask()
if relogin is None:
raise KeyboardInterrupt()
if relogin:
_run_ccproxy_login("codex", "Codex OAuth")
else:
console.print(f" [yellow]Codex OAuth not authenticated: {msg}[/yellow]")
login = questionary.confirm(
"Log in to Codex now?",
default=True,
style=CONFIRM_STYLE,
qmark=QMARK,
).ask()
if login is None:
raise KeyboardInterrupt()
if login:
_run_ccproxy_login("codex", "Codex OAuth")
return auth_mode
def _step_provider_api_key(
+181 -258
View File
@@ -129,12 +129,6 @@ _PROVIDER_KEY_ATTR = {
"custom-anthropic": "custom_anthropic_api_key",
}
_MINIMAX_GLOBAL_BASE_URL = "https://api.minimax.io/anthropic"
_CUSTOM_PROVIDER_BASE_URL = {
"custom-openai": ("custom_openai_base_url", "CUSTOM_OPENAI_BASE_URL"),
"custom-anthropic": ("custom_anthropic_base_url", "CUSTOM_ANTHROPIC_BASE_URL"),
}
def _autosave(config: EvoScientistConfig) -> None:
"""Persist current config to disk between phases.
@@ -148,201 +142,6 @@ def _autosave(config: EvoScientistConfig) -> None:
pass
def _configure_provider_base_url(
config: EvoScientistConfig,
provider: str,
*,
strict: bool,
) -> list[str]:
"""Configure provider-specific base URL/region and return Ollama models."""
if provider in _CUSTOM_PROVIDER_BASE_URL:
attr_name, env_name = _CUSTOM_PROVIDER_BASE_URL[provider]
current_base_url = getattr(config, attr_name) or os.environ.get(env_name, "")
if strict:
if not current_base_url:
raise RuntimeError(
f"--non-interactive: {provider} provider needs a base URL. "
f"Set the {env_name} env var or run without --non-interactive."
)
setattr(config, attr_name, current_base_url)
else:
setattr(
config,
attr_name,
_step_base_url(config, current_value=current_base_url),
)
elif provider == "minimax":
if strict:
config.minimax_base_url = (
config.minimax_base_url or _MINIMAX_GLOBAL_BASE_URL
)
else:
config.minimax_base_url = _step_minimax_region(config)
elif provider == "ollama":
if strict:
config.ollama_base_url = (
config.ollama_base_url
or os.environ.get("OLLAMA_BASE_URL", "")
or "http://localhost:11434"
)
else:
ollama_url, ollama_detected_models = _step_ollama_base_url(config)
config.ollama_base_url = ollama_url
return ollama_detected_models
return []
def _configure_provider_auth_mode(
config: EvoScientistConfig,
provider: str,
*,
strict: bool,
) -> None:
"""Configure Anthropic/OpenAI auth mode for the selected provider."""
if provider == "anthropic":
if strict:
config.anthropic_auth_mode = "api_key"
else:
config.anthropic_auth_mode = _step_anthropic_auth_mode(config)
elif provider == "openai":
if strict:
config.openai_auth_mode = "api_key"
else:
config.openai_auth_mode = _step_openai_auth_mode(config)
def _active_llm_providers(config: EvoScientistConfig) -> set[str]:
"""Return providers currently selected by the main and auxiliary models."""
providers = {config.provider}
if config.auxiliary_provider:
providers.add(config.auxiliary_provider)
return providers
def _reconcile_oauth_modes(config: EvoScientistConfig) -> None:
"""Clear OAuth flags for providers no selected model uses."""
active_providers = _active_llm_providers(config)
if "anthropic" not in active_providers:
config.anthropic_auth_mode = "api_key"
if "openai" not in active_providers:
config.openai_auth_mode = "api_key"
def _provider_uses_oauth(config: EvoScientistConfig, provider: str) -> bool:
return (provider == "anthropic" and config.anthropic_auth_mode == "oauth") or (
provider == "openai" and config.openai_auth_mode == "oauth"
)
def _apply_preset_provider_api_key(
config: EvoScientistConfig,
provider: str,
preset_api_key: str,
*,
skip_validation: bool,
) -> None:
"""Validate and store a CLI-supplied provider API key."""
if not skip_validation:
from .helpers import _provider_key_info
_info = _provider_key_info(config, provider)
validate_fn = _info[2] if _info else None
if validate_fn is not None:
console.print(" [dim]Validating preset API key...[/dim]", end="")
valid, msg = validate_fn(preset_api_key)
if valid:
console.print(f"\r [green]✓ {msg}[/green] ")
else:
console.print(f"\r [red]✗ {msg}[/red] ")
raise RuntimeError(
f"--api-key rejected by {provider} validator: {msg}. "
"Pass --skip-validation to override."
)
key_attr = _PROVIDER_KEY_ATTR.get(provider, "openai_api_key")
setattr(config, key_attr, preset_api_key)
console.print(
f" [green]✓ API key: ***{preset_api_key[-4:]}[/green] [dim](--api-key)[/dim]"
)
def _configure_provider_api_key(
config: EvoScientistConfig,
provider: str,
*,
skip_validation: bool,
preset_api_key: str | None = None,
require_api_key=None,
) -> None:
"""Configure provider API key unless the provider does not need one."""
if provider == "ollama" or _provider_uses_oauth(config, provider):
return
key_attr = _PROVIDER_KEY_ATTR.get(provider, "openai_api_key")
if preset_api_key is not None:
_apply_preset_provider_api_key(
config,
provider,
preset_api_key,
skip_validation=skip_validation,
)
return
if require_api_key is not None:
require_api_key()
new_key = _step_provider_api_key(config, provider, skip_validation)
if new_key is not None:
setattr(config, key_attr, new_key)
elif not getattr(config, key_attr):
_print_step_skipped("API Key", "not set")
def _provider_connection_configured(config: EvoScientistConfig, provider: str) -> bool:
"""Return True when provider-level setup can be safely reused."""
if provider == "ollama":
return bool(config.ollama_base_url)
if provider == "custom-openai" and not config.custom_openai_base_url:
return False
if provider == "custom-anthropic" and not config.custom_anthropic_base_url:
return False
if provider == "minimax" and not config.minimax_base_url:
return False
if _provider_uses_oauth(config, provider):
return True
key_attr = _PROVIDER_KEY_ATTR.get(provider, "openai_api_key")
return bool(getattr(config, key_attr))
def _configure_provider_connection(
config: EvoScientistConfig,
provider: str,
*,
strict: bool,
skip_validation: bool,
preset_api_key: str | None = None,
require_api_key=None,
) -> list[str]:
"""Configure provider base URL/region, auth mode, and API key."""
ollama_detected_models = _configure_provider_base_url(
config,
provider,
strict=strict,
)
_configure_provider_auth_mode(
config,
provider,
strict=strict,
)
_configure_provider_api_key(
config,
provider,
skip_validation=skip_validation,
preset_api_key=preset_api_key,
require_api_key=require_api_key,
)
return ollama_detected_models
# Sections offered in Keep/Modify/Reset → which step labels they enable.
_SECTION_LABELS: list[tuple[str, str]] = [
("ui", "UI backend"),
@@ -680,17 +479,102 @@ def run_onboard(
provider = _step_provider(config)
config.provider = provider
try:
ollama_detected_models = _configure_provider_connection(
config,
provider,
strict=strict,
skip_validation=skip_validation,
preset_api_key=_preset("api_key"),
require_api_key=lambda provider=provider: _require(
"api_key", f"{provider} API key"
),
# Step 2a: Base URL (custom-openai, custom-anthropic,
# minimax, ollama). In strict non-interactive mode we
# never call the interactive _step_base_url /
# _step_minimax_region / _step_ollama_base_url helpers —
# fall back to the existing config value or the
# CUSTOM_*_BASE_URL / OLLAMA_BASE_URL env var instead.
# If neither is set for a provider that needs it, raise
# so the user sees the same "missing required answer"
# error as for other required prompts.
if provider == "custom-openai":
current_base_url = (
config.custom_openai_base_url
or os.environ.get("CUSTOM_OPENAI_BASE_URL", "")
)
if strict:
if not current_base_url:
raise RuntimeError(
"--non-interactive: custom-openai provider "
"needs a base URL. Set the "
"CUSTOM_OPENAI_BASE_URL env var or run "
"without --non-interactive."
)
config.custom_openai_base_url = current_base_url
else:
config.custom_openai_base_url = _step_base_url(
config, current_value=current_base_url
)
elif provider == "custom-anthropic":
current_base_url = (
config.custom_anthropic_base_url
or os.environ.get("CUSTOM_ANTHROPIC_BASE_URL", "")
)
if strict:
if not current_base_url:
raise RuntimeError(
"--non-interactive: custom-anthropic "
"provider needs a base URL. Set the "
"CUSTOM_ANTHROPIC_BASE_URL env var or run "
"without --non-interactive."
)
config.custom_anthropic_base_url = current_base_url
else:
config.custom_anthropic_base_url = _step_base_url(
config, current_value=current_base_url
)
elif provider == "minimax":
if strict:
# MiniMax has 2 region URLs; default to whatever
# is already in config, else the Global endpoint.
config.minimax_base_url = (
config.minimax_base_url
or "https://api.minimax.io/anthropic"
)
else:
config.minimax_base_url = _step_minimax_region(config)
elif provider == "ollama":
if strict:
# Ollama: existing config value > env var >
# localhost default. Skip the live connection
# validation under strict — model discovery
# happens at runtime anyway.
config.ollama_base_url = (
config.ollama_base_url
or os.environ.get("OLLAMA_BASE_URL", "")
or "http://localhost:11434"
)
# ollama_detected_models stays [] — model picker
# will fall back to free-text or the preset.
else:
ollama_url, ollama_detected_models = _step_ollama_base_url(
config
)
config.ollama_base_url = ollama_url
# Step 2b: Auth mode (Anthropic or OpenAI — API key vs OAuth).
# In strict non-interactive mode we assume "api_key".
# The prompt offers a `← Back` choice that raises GoBack so
# the user can re-pick the provider without exiting the wizard.
try:
if provider == "anthropic":
if strict:
config.anthropic_auth_mode = "api_key"
else:
config.anthropic_auth_mode = _step_anthropic_auth_mode(
config
)
elif provider == "openai":
if strict:
config.openai_auth_mode = "api_key"
else:
config.openai_auth_mode = _step_openai_auth_mode(config)
else:
# Non-Anthropic/OpenAI provider: reset OAuth modes to
# avoid stale oauth config triggering ccproxy at startup.
config.anthropic_auth_mode = "api_key"
config.openai_auth_mode = "api_key"
except GoBack:
# User picked "← Back" — restore config to its state at the
# top of this iteration (drops any base_url / region /
@@ -710,9 +594,60 @@ def run_onboard(
ollama_detected_models = []
console.print(" [dim]↩ Returning to provider selection.[/dim]")
continue
break # Provider setup succeeded — exit sub-loop
break # auth_mode succeeded — exit sub-loop
_reconcile_oauth_modes(config)
# Step 2c: Provider API Key (skip for Ollama and pure OAuth)
_skip_api_key = (
provider == "ollama"
or (
provider == "anthropic"
and config.anthropic_auth_mode == "oauth"
)
or (provider == "openai" and config.openai_auth_mode == "oauth")
)
if not _skip_api_key:
key_attr = _PROVIDER_KEY_ATTR.get(provider, "openai_api_key")
preset_api_key = _preset("api_key")
if preset_api_key is not None:
# Validate the preset key against the same validator
# the interactive path uses, unless --skip-validation
# was passed. Interactive flow shows a "Save anyway?"
# confirm on failure; the non-interactive path has no
# way to ask, so a failed validation is fatal.
if not skip_validation:
from .helpers import _provider_key_info
_info = _provider_key_info(config, provider)
validate_fn = _info[2] if _info else None
if validate_fn is not None:
console.print(
" [dim]Validating preset API key...[/dim]",
end="",
)
valid, msg = validate_fn(preset_api_key)
if valid:
console.print(f"\r [green]✓ {msg}[/green] ")
else:
console.print(f"\r [red]✗ {msg}[/red] ")
raise RuntimeError(
f"--api-key rejected by {provider} "
f"validator: {msg}. Pass "
"--skip-validation to override."
)
setattr(config, key_attr, preset_api_key)
console.print(
f" [green]✓ API key: ***{preset_api_key[-4:]}[/green]"
" [dim](--api-key)[/dim]"
)
else:
_require("api_key", f"{provider} API key")
new_key = _step_provider_api_key(
config, provider, skip_validation
)
if new_key is not None:
setattr(config, key_attr, new_key)
elif not getattr(config, key_attr):
_print_step_skipped("API Key", "not set")
_autosave(config)
else:
# Provider section skipped — keep prior provider value to drive
@@ -745,55 +680,44 @@ def run_onboard(
"kept current" if config.auxiliary_model else "not set",
)
elif _step_auxiliary_enable(config):
from .prompter import GoBack
aux_ollama_detected_models: list[str] = []
while True:
loop_snapshot = copy.deepcopy(config)
aux_provider = _step_provider(
# Assemble: pick provider -> base URL (custom) -> key -> model,
# mirroring the main flow's order. Keys/base URLs are stored
# per provider, so when the auxiliary provider matches the main
# one they're already set and the user just keeps them (Enter).
# Ollama needs no key. Re-runs default to the saved auxiliary
# provider/model rather than the main ones.
aux_provider = _step_provider(
config,
label="co-pilot",
default_value=config.auxiliary_provider,
)
config.auxiliary_provider = aux_provider
if aux_provider == "custom-openai":
config.custom_openai_base_url = _step_base_url(
config,
label="co-pilot",
default_value=config.auxiliary_provider,
current_value=config.custom_openai_base_url
or os.environ.get("CUSTOM_OPENAI_BASE_URL", ""),
)
config.auxiliary_provider = aux_provider
if (
aux_provider == config.provider
and _provider_connection_configured(config, aux_provider)
):
if aux_provider == "ollama":
aux_ollama_detected_models = ollama_detected_models
_print_step_skipped(
"Co-pilot credentials",
"reusing main provider settings",
)
else:
try:
aux_ollama_detected_models = (
_configure_provider_connection(
config,
aux_provider,
strict=False,
skip_validation=skip_validation,
)
)
except GoBack:
for field_name in vars(loop_snapshot):
setattr(
config,
field_name,
getattr(loop_snapshot, field_name),
)
aux_ollama_detected_models = []
console.print(
" [dim]↩ Returning to co-pilot provider "
"selection.[/dim]"
)
continue
break
elif aux_provider == "custom-anthropic":
config.custom_anthropic_base_url = _step_base_url(
config,
current_value=config.custom_anthropic_base_url
or os.environ.get("CUSTOM_ANTHROPIC_BASE_URL", ""),
)
elif aux_provider == "minimax":
config.minimax_base_url = _step_minimax_region(config)
if aux_provider != "ollama":
aux_key_attr = _PROVIDER_KEY_ATTR.get(
aux_provider, "openai_api_key"
)
new_aux_key = _step_provider_api_key(
config, aux_provider, skip_validation
)
if new_aux_key is not None:
setattr(config, aux_key_attr, new_aux_key)
config.auxiliary_model = _step_model(
config,
aux_provider,
ollama_detected_models=aux_ollama_detected_models,
label="co-pilot",
default_value=config.auxiliary_model,
)
@@ -801,7 +725,6 @@ def run_onboard(
# Skip: single driver — clear any prior auxiliary config.
config.auxiliary_provider = ""
config.auxiliary_model = ""
_reconcile_oauth_modes(config)
_autosave(config)
if "tavily" in sections_to_run:
+3 -87
View File
@@ -106,23 +106,11 @@ def _normalize_hhmm(value: Any) -> str | None:
def get_config_dir() -> Path:
"""Get the configuration directory path.
Priority:
1. EVOSCIENTIST_CONFIG_DIR
2. EVOSCIENTIST_HOME/config
3. XDG_CONFIG_HOME/evoscientist
4. ~/.config/evoscientist
Uses XDG_CONFIG_HOME if set, otherwise ~/.config/evoscientist/
"""
configured = os.environ.get("EVOSCIENTIST_CONFIG_DIR")
if configured:
return Path(configured).expanduser().resolve()
home = os.environ.get("EVOSCIENTIST_HOME")
if home:
return Path(home).expanduser().resolve() / "config"
xdg_config = os.environ.get("XDG_CONFIG_HOME")
if xdg_config:
return Path(xdg_config).expanduser() / "evoscientist"
return Path(xdg_config) / "evoscientist"
return Path.home() / ".config" / "evoscientist"
@@ -135,16 +123,6 @@ def get_config_path() -> Path:
# Configuration dataclass
# =============================================================================
# OpenRouter app-attribution defaults (issue #339). Single source of truth: the
# EvoScientistConfig fields below default to these, and llm/models.py imports
# them for its env-fallback, so the values never drift across the two layers.
OPENROUTER_DEFAULT_HTTP_REFERER = "https://github.com/EvoScientist/EvoScientist"
OPENROUTER_DEFAULT_APP_TITLE = "EvoScientist"
# OpenRouter honors only the first 2 categories per request (server-side limit)
# and silently ignores the rest, so keep the two most relevant ones. Chosen per
# maintainer review — creative-writing is a less competitive marketplace group.
OPENROUTER_DEFAULT_APP_CATEGORIES = "creative-writing,personal-agent"
@dataclass
class EvoScientistConfig:
@@ -262,13 +240,6 @@ class EvoScientistConfig:
# Lower (e.g., 5000) if you want a tighter safety net against runaway loops.
recursion_limit: int = 1_000_000
# Number of consecutive model rounds with the same structured tool name and
# arguments that activates provider-facing loop repair. Set 0 to disable.
repetitive_tool_call_threshold: int = 2
# Number of consecutive deterministic tool errors allowed before the next
# model call is blocked. Transient provider/network errors are not counted.
max_consecutive_tool_errors: int = 3
# Memory Settings
# Profile memory injects and maintains `/memories/profile/...` files.
memory_profile_enabled: bool = True
@@ -307,21 +278,10 @@ class EvoScientistConfig:
# a deploy-style langgraph server instead of the in-terminal CLI/TUI.
ui_backend: Literal["cli", "tui", "webui"] = "tui"
log_level: str = "warning"
# Empty means use the provider/model default. A non-empty value is an
# explicit user override exported as EVOSCIENTIST_REASONING_EFFORT.
reasoning_effort: str = ""
reasoning_effort: str = "high"
# Anthropic prompt caching for OpenRouter anthropic/* models. Opt out if
# cache-write costs outweigh the benefit for a workflow.
openrouter_anthropic_prompt_cache: bool = True
# OpenRouter app attribution (issue #339). Sent only for the openrouter
# provider; identifies EvoScientist in OpenRouter's app rankings/analytics.
# Override (e.g. a private fork) via these fields or their env vars.
# Defaults live in the module constants above (also imported by llm/models.py).
openrouter_http_referer: str = OPENROUTER_DEFAULT_HTTP_REFERER
openrouter_app_title: str = OPENROUTER_DEFAULT_APP_TITLE
# Comma-separated; split into a list before being passed to
# langchain-openrouter (its app_categories kwarg expects list[str]).
openrouter_app_categories: str = OPENROUTER_DEFAULT_APP_CATEGORIES
# Channel Settings
channel_enabled: str = "" # "imessage" | "telegram" | "discord" | "slack" | "wechat" | "dingtalk" | "feishu" | "email" | "qq" | "signal" | "" (comma-separated for multiple)
@@ -480,14 +440,6 @@ class EvoScientistConfig:
stt_compute_type: str = "int8" # "int8" | "float16" | "float32"
def __post_init__(self) -> None:
for field_name in (
"repetitive_tool_call_threshold",
"max_consecutive_tool_errors",
):
value = getattr(self, field_name)
if not isinstance(value, int) or isinstance(value, bool) or value < 0:
raise ValueError(f"{field_name} must be a non-negative integer")
# A non-positive or non-int sandbox_execute_timeout (e.g. a hand-edited
# config file value — load_config does not coerce file values — or a
# 0/negative env value) would raise inside CustomSandboxBackend.__init__
@@ -596,10 +548,6 @@ def save_config(config: EvoScientistConfig) -> None:
"""
config_path = get_config_path()
config_path.parent.mkdir(parents=True, exist_ok=True)
try:
config_path.parent.chmod(0o700)
except OSError:
pass
data = _config_to_dict(config)
@@ -612,10 +560,6 @@ def save_config(config: EvoScientistConfig) -> None:
sort_keys=False,
allow_unicode=True,
)
try:
config_path.chmod(0o600)
except OSError:
pass
def reset_config() -> None:
@@ -750,11 +694,6 @@ def set_config_value(key: str, value: Any) -> bool:
if key == "sandbox_execute_timeout" and value <= 0:
return False
if key in {
"repetitive_tool_call_threshold",
"max_consecutive_tool_errors",
} and (isinstance(value, bool) or value < 0):
return False
if key == "memory_skill_synthesis_time":
value = _normalize_hhmm(value)
if value is None:
@@ -814,9 +753,6 @@ _ENV_MAPPINGS = {
"openrouter_anthropic_prompt_cache": (
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE"
),
"openrouter_http_referer": "EVOSCIENTIST_OPENROUTER_HTTP_REFERER",
"openrouter_app_title": "EVOSCIENTIST_OPENROUTER_APP_TITLE",
"openrouter_app_categories": "EVOSCIENTIST_OPENROUTER_APP_CATEGORIES",
"dangerous_mode": "EVOSCIENTIST_DANGEROUS_MODE",
"channel_debug_tracing": "EVOSCIENTIST_CHANNEL_DEBUG_TRACING",
"ccproxy_port": "EVOSCIENTIST_CCPROXY_PORT",
@@ -833,10 +769,6 @@ _ENV_MAPPINGS = {
"langgraph_dev_file_persistence": "EVOSCIENTIST_LANGGRAPH_DEV_FILE_PERSISTENCE",
"langgraph_dev_jobs_per_worker": "EVOSCIENTIST_LANGGRAPH_DEV_JOBS_PER_WORKER",
"recursion_limit": "EVOSCIENTIST_RECURSION_LIMIT",
"repetitive_tool_call_threshold": (
"EVOSCIENTIST_REPETITIVE_TOOL_CALL_THRESHOLD"
),
"max_consecutive_tool_errors": "EVOSCIENTIST_MAX_CONSECUTIVE_TOOL_ERRORS",
"memory_profile_enabled": "EVOSCIENTIST_MEMORY_PROFILE_ENABLED",
"memory_observations_enabled": "EVOSCIENTIST_MEMORY_OBSERVATIONS_ENABLED",
"memory_observation_writer": "EVOSCIENTIST_MEMORY_OBSERVATION_WRITER",
@@ -952,22 +884,6 @@ def apply_config_to_env(config: EvoScientistConfig) -> None:
os.environ["TAVILY_API_KEY"] = config.tavily_api_key
if config.reasoning_effort and not os.environ.get("EVOSCIENTIST_REASONING_EFFORT"):
os.environ["EVOSCIENTIST_REASONING_EFFORT"] = config.reasoning_effort
if config.openrouter_http_referer and not os.environ.get(
"EVOSCIENTIST_OPENROUTER_HTTP_REFERER"
):
os.environ["EVOSCIENTIST_OPENROUTER_HTTP_REFERER"] = (
config.openrouter_http_referer
)
if config.openrouter_app_title and not os.environ.get(
"EVOSCIENTIST_OPENROUTER_APP_TITLE"
):
os.environ["EVOSCIENTIST_OPENROUTER_APP_TITLE"] = config.openrouter_app_title
if config.openrouter_app_categories and not os.environ.get(
"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES"
):
os.environ["EVOSCIENTIST_OPENROUTER_APP_CATEGORIES"] = (
config.openrouter_app_categories
)
if not config.openrouter_anthropic_prompt_cache and not os.environ.get(
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE"
):
+9 -2
View File
@@ -41,9 +41,15 @@ from ..stream.console import console
# Front-end npm package + spec. ``@latest`` → always the newest published UI.
_WEBUI_PACKAGE = "@evoscientist/webui@latest"
_WEBUI_PACKAGE_ENV = "EVOSCIENTIST_WEBUI_PACKAGE"
_DEFAULT_WEBUI_PORT = 4716
def _resolve_webui_package() -> str:
"""Return the npm package spec to launch for the WebUI front-end."""
return os.getenv(_WEBUI_PACKAGE_ENV) or _WEBUI_PACKAGE
def run_webui(config: Any, workspace_dir: str | None = None) -> None:
"""Start the deploy-style backend + the WebUI front-end, then block.
@@ -203,6 +209,7 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
"PORT": str(webui_port),
}
)
webui_package = _resolve_webui_package()
console.print(
Panel(
Text.from_markup(
@@ -211,7 +218,7 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
f"[bold]WebUI:[/bold] http://localhost:{webui_port} "
f"[dim](opens in your browser)[/dim]\n"
f"[bold]Logs:[/bold] {_shorten(str(RUNTIME.log_file))}\n\n"
f"[dim]Fetching {_WEBUI_PACKAGE} via npx (first run may take a "
f"[dim]Fetching {webui_package} via npx (first run may take a "
f"moment)… Press Ctrl+C to stop.[/dim]"
),
title="[bold green]✓ EvoScientist WebUI[/bold green]",
@@ -228,7 +235,7 @@ def run_webui(config: Any, workspace_dir: str | None = None) -> None:
popen_kwargs["creationflags"] = subprocess.CREATE_NEW_PROCESS_GROUP
try:
webui_proc = subprocess.Popen(
[npx, "--yes", _WEBUI_PACKAGE, "--port", str(webui_port)],
[npx, "--yes", webui_package, "--port", str(webui_port)],
**popen_kwargs,
)
except Exception as exc:
+282 -1
View File
@@ -23,7 +23,17 @@ memory.
from __future__ import annotations
import asyncio
import json
import logging
import os
import sqlite3
from pathlib import Path
from typing import Any
import aiosqlite
import httpx
from langgraph.checkpoint.serde.jsonplus import JsonPlusSerializer
from langgraph.checkpoint.sqlite.aio import AsyncSqliteSaver
from starlette.applications import Starlette
from starlette.requests import Request
from starlette.responses import JSONResponse
@@ -31,6 +41,15 @@ from starlette.routing import Route
from EvoScientist.config import get_effective_config
from EvoScientist.llm.models import list_model_picker_entries
from EvoScientist.sessions import (
MAIN_THREAD_FILTER_PARAMS,
MAIN_THREAD_FILTER_SQL,
_load_checkpoint_messages,
_table_exists,
_to_short_path,
)
_logger = logging.getLogger(__name__)
async def get_models(_request: Request) -> JSONResponse:
@@ -43,7 +62,7 @@ async def get_models(_request: Request) -> JSONResponse:
``discover_ollama_models()`` call, same 1.5-s timeout, same
fail-soft semantics (the probe returns ``[]`` on any error, never
raises). The TUI's "Custom Ollama model…" sentinel is intentionally
omitted — that's a widget-specific input affordance, not part of
omitted: that's a widget-specific input affordance, not part of
the registry surface.
``default`` reflects the deployment's currently-configured fallback
@@ -75,8 +94,270 @@ async def get_models(_request: Request) -> JSONResponse:
)
def _message_type(message: Any) -> str | None:
if isinstance(message, dict):
role = message.get("role")
if role == "assistant":
return "ai"
if role == "user":
return "human"
value = message.get("type")
return str(value) if value is not None else None
value = getattr(message, "type", None)
return str(value) if value is not None else None
def _message_content(message: Any) -> Any:
if isinstance(message, dict):
return message.get("content")
return getattr(message, "content", None)
def _extract_text_content(content: Any) -> str:
if isinstance(content, str):
return content
if not isinstance(content, list):
return ""
parts: list[str] = []
for block in content:
if isinstance(block, str):
parts.append(block)
continue
if not isinstance(block, dict):
continue
block_type = block.get("type")
if block_type not in {"text", "output_text"}:
continue
text = block.get("text") or block.get("content")
if isinstance(text, str):
parts.append(text)
return "\n\n".join(part for part in parts if part)
def _is_tool_selection_payload(raw: str) -> bool:
try:
parsed = json.loads(raw)
except json.JSONDecodeError:
return False
return (
isinstance(parsed, dict)
and set(parsed) == {"tools"}
and isinstance(parsed["tools"], list)
and all(isinstance(tool, str) for tool in parsed["tools"])
)
def _split_json_objects(raw: str) -> list[str] | None:
objects: list[str] = []
depth = 0
start = -1
in_string = False
escaping = False
for i, char in enumerate(raw):
if in_string:
if escaping:
escaping = False
elif char == "\\":
escaping = True
elif char == '"':
in_string = False
continue
if char == '"':
if depth == 0:
return None
in_string = True
continue
if char == "{":
if depth == 0:
start = i
depth += 1
continue
if char == "}":
depth -= 1
if depth < 0 or start < 0:
return None
if depth == 0:
objects.append(raw[start : i + 1])
start = -1
continue
if depth == 0 and not char.isspace():
return None
if depth != 0 or in_string or not objects:
return None
return objects
def _is_tool_selection_text(text: str) -> bool:
stripped = text.strip()
if not stripped or '"tools"' not in stripped:
return False
if _is_tool_selection_payload(stripped):
return True
objects = _split_json_objects(stripped)
return objects is not None and all(_is_tool_selection_payload(obj) for obj in objects)
def _extract_final_answer(messages: list[Any]) -> str:
"""Return displayable text from the latest AI message in *messages*."""
for message in reversed(messages):
if _message_type(message) != "ai":
continue
content = _extract_text_content(_message_content(message)).strip()
if content and not _is_tool_selection_text(content):
return content
return ""
def _sessions_db_path_for_http() -> Path:
data_dir = os.getenv("EVOSCIENTIST_DATA_DIR")
base = Path(data_dir).expanduser() if data_dir else Path.home() / ".evoscientist"
return Path(_to_short_path(str(base))) / "sessions.db"
async def _get_thread_metadata_for_http(thread_id: str) -> dict | None:
try:
async with aiosqlite.connect(
str(_sessions_db_path_for_http()), timeout=30.0
) as conn:
if not await _table_exists(conn, "checkpoints"):
return None
query = f"""
SELECT json_extract(metadata, '$.workspace_dir') as workspace_dir,
json_extract(metadata, '$.model') as model,
json_extract(metadata, '$.updated_at') as updated_at
FROM checkpoints
WHERE thread_id = ?
AND {MAIN_THREAD_FILTER_SQL}
ORDER BY checkpoint_id DESC
LIMIT 1
"""
async with conn.execute(
query, (thread_id, *MAIN_THREAD_FILTER_PARAMS)
) as cur:
row = await cur.fetchone()
except (OSError, sqlite3.Error):
return None
if not row:
return None
return {
"workspace_dir": row[0],
"model": row[1],
"updated_at": row[2],
}
async def _get_thread_messages_for_http(thread_id: str) -> list:
try:
async with aiosqlite.connect(
str(_sessions_db_path_for_http()), timeout=30.0
) as conn:
if not await _table_exists(conn, "checkpoints"):
return []
check = f"""
SELECT 1 FROM checkpoints
WHERE thread_id = ? AND {MAIN_THREAD_FILTER_SQL}
LIMIT 1
"""
async with conn.execute(
check, (thread_id, *MAIN_THREAD_FILTER_PARAMS)
) as cur:
if not await cur.fetchone():
return []
serde = JsonPlusSerializer()
saver = AsyncSqliteSaver(conn, serde=serde)
return await _load_checkpoint_messages(saver, thread_id)
except (OSError, sqlite3.Error):
return []
async def _read_thread_runtime_state(request: Request, thread_id: str) -> dict[str, Any]:
"""Read thread status from the co-hosted langgraph-api endpoints."""
base_url = f"{request.url.scheme}://{request.url.netloc}"
timeout = httpx.Timeout(2.0, connect=0.5)
async with httpx.AsyncClient(base_url=base_url, timeout=timeout) as client:
thread_resp = await client.get(f"/threads/{thread_id}")
if thread_resp.status_code == 404:
return {"found": False}
thread_resp.raise_for_status()
thread = thread_resp.json()
state: dict[str, Any] = {}
state_resp = await client.get(f"/threads/{thread_id}/state")
if state_resp.status_code == 404:
return {"found": False}
if state_resp.status_code < 400:
state = state_resp.json()
next_nodes = state.get("next")
is_terminal_checkpoint = isinstance(next_nodes, (list, tuple)) and not next_nodes
status = thread.get("status")
complete = status == "idle" or is_terminal_checkpoint
completed_at = None
if complete:
completed_at = (
thread.get("state_updated_at")
or thread.get("updated_at")
or (state.get("metadata") or {}).get("updated_at")
)
return {
"found": True,
"complete": complete,
"completed_at": completed_at,
}
async def get_final_answer(request: Request) -> JSONResponse:
"""Return the latest checkpointed assistant answer for a WebUI thread.
This is a recovery surface for the WebUI stream consumer: when browser-side
SSE is interrupted but the langgraph run continues server-side, the final
answer is already persisted in ``sessions.db``. The route centralizes the
non-trivial "latest AIMessage text only" extraction so the browser does not
render reasoning blocks, tool calls, or tool-selection JSON fragments.
"""
thread_id = request.path_params["thread_id"]
metadata = await _get_thread_metadata_for_http(thread_id)
if metadata is None:
return JSONResponse({"error": "thread not found"}, status_code=404)
messages = await _get_thread_messages_for_http(thread_id)
content = _extract_final_answer(messages)
runtime: dict[str, Any] = {"found": True, "complete": False, "completed_at": None}
try:
runtime = await _read_thread_runtime_state(request, thread_id)
except Exception as exc:
_logger.debug(
"Could not read langgraph runtime state for thread %s: %s",
thread_id,
exc,
)
if runtime.get("found") is False:
return JSONResponse({"error": "thread not found"}, status_code=404)
completed_at = runtime.get("completed_at")
if completed_at is None and runtime.get("complete"):
completed_at = metadata.get("updated_at")
return JSONResponse(
{
"content": content,
"completed_at": completed_at,
"complete": bool(runtime.get("complete")),
}
)
app = Starlette(
routes=[
Route("/api/models", get_models, methods=["GET"]),
Route(
"/api/threads/{thread_id}/final-answer",
get_final_answer,
methods=["GET"],
),
]
)
+1 -237
View File
@@ -5,244 +5,8 @@ in ``EvoScientist/EvoScientist.py`` so it doesn't construct on plain
``import EvoScientist``. ``langgraph dev`` 's symbol resolver inspects
module attributes directly and doesn't trigger ``__getattr__``, so we
re-export here to make it visible.
Before re-export we upgrade the compiled graph's class in place to
``_EvoFilteredGraph``, which strips ``PrivateStateAttr``-marked fields
(currently just ``_quickjs_snapshot_payload``) from ``get_state`` /
``get_state_history`` responses. Upstream ``langchain_quickjs`` annotates
the field ``PrivateStateAttr = OmitFromSchema(input=True, output=True)``,
but LangGraph's ``_prepare_state_snapshot`` doesn't honor that on
checkpoint reads — every ``getState`` materializes the delta chain back
into a full ~1.4 MB blob, which the WebUI then downloads. The subclass
closes the gap without touching the middleware's write path, preserving
cross-turn REPL persistence as ``langchain-ai/deepagents#3064`` shipped it.
"""
from langgraph.graph.state import CompiledStateGraph
from langgraph.types import PregelTask, StateSnapshot
from EvoScientist.EvoScientist import EvoScientist_agent as _agent
_PRIVATE_STATE_FIELDS = frozenset({"_quickjs_snapshot_payload"})
# Sanity check on the LangGraph internals ``_strip_private`` scrubs. If any
# of these attributes disappear or get renamed in a future upstream bump,
# the assertion fires at import time — the deployment refuses to start,
# instead of silently degrading (the filter would ``.get()`` its way to a
# no-op and the private-field payload would come back on the wire without
# anyone noticing until a user reports slow thread switches again).
#
# Doesn't cover every internal we depend on — ``metadata["writes"]`` /
# ``metadata["counters_since_delta_snapshot"]`` dict keys aren't a canary
# target because ``dict.get`` already tolerates their absence. What we
# canary here is the ``NamedTuple`` field set: renames there would be the
# highest-impact silent regression.
_EXPECTED_SNAPSHOT_FIELDS = frozenset({"values", "metadata", "tasks"})
_EXPECTED_TASK_FIELDS = frozenset({"result", "state"})
_missing_snap = _EXPECTED_SNAPSHOT_FIELDS - set(StateSnapshot._fields)
_missing_task = _EXPECTED_TASK_FIELDS - set(PregelTask._fields)
if _missing_snap or _missing_task:
raise RuntimeError(
"LangGraph state shape drifted from the version _strip_private was "
f"written against. Missing StateSnapshot fields: {_missing_snap or set()}. "
f"Missing PregelTask fields: {_missing_task or set()}. Review "
"_strip_private and re-verify against the current upstream shape "
"before removing this assertion."
)
def _strip_private(snap):
"""Strip ``PrivateStateAttr``-marked fields from a ``StateSnapshot``.
Empirically verified against a live history response for a thread with
a single touched turn: the private field leaks on four surfaces — three
trivial, one heavy:
* ``snap.values`` — the materialized channel state exposed as the main
payload. For DeltaChannels this is the delta chain replayed into full
bytes (~1.4 MB for the quickjs snapshot). ``get_state`` and every
history entry.
* ``snap.metadata['writes']`` — ``{node_name: {channel: value}}`` map of
the raw writes that produced each checkpoint. On the ``after_agent``
step that first snapshots the REPL, ``value`` is the encoded write
record ``("snap", full_bytes)`` ≈ 1.4 MB.
* ``snap.tasks[*].result`` — the return dict of each completed
``PregelTask``. ``after_agent`` returns
``{"_quickjs_snapshot_payload": ("snap", bytes)}``; this dict becomes
the task's ``result`` field, which the API surfaces verbatim under
``tasks[*].result`` (``langgraph_api.state:106``). This is the
dominant leak: 1.7 MB in the last history entry of any thread whose
most-recent-in-window checkpoint had a snapshot anchor.
* ``snap.metadata['counters_since_delta_snapshot']`` — DeltaChannel's
snapshot cadence bookkeeping ``{channel: [count, superstep]}``. Tiny
(~20 B) but exposes the private field name; strip for cleanliness.
* ``snap.tasks[*].state`` (nested ``StateSnapshot``) — populated when the
caller passes ``subgraphs=True``. Repeats all of the above surfaces
for each subgraph task, so recurse into it. Not exercised by the
current WebUI (which doesn't pass ``subgraphs=True`` on REST reads),
but SDK / curl / gRPC callers can.
"""
if snap is None:
return snap
values = {k: v for k, v in snap.values.items() if k not in _PRIVATE_STATE_FIELDS}
metadata = snap.metadata
if metadata:
new_metadata = metadata
if new_metadata.get("writes"):
scrubbed_writes = {
node: {
k: v for k, v in ch_writes.items() if k not in _PRIVATE_STATE_FIELDS
}
for node, ch_writes in new_metadata["writes"].items()
}
new_metadata = {**new_metadata, "writes": scrubbed_writes}
if new_metadata.get("counters_since_delta_snapshot"):
scrubbed_counters = {
k: v
for k, v in new_metadata["counters_since_delta_snapshot"].items()
if k not in _PRIVATE_STATE_FIELDS
}
new_metadata = {
**new_metadata,
"counters_since_delta_snapshot": scrubbed_counters,
}
metadata = new_metadata
tasks = snap.tasks
if tasks:
new_tasks = []
changed = False
for t in tasks:
replace_kwargs: dict = {}
result = getattr(t, "result", None)
if isinstance(result, dict) and any(
k in result for k in _PRIVATE_STATE_FIELDS
):
replace_kwargs["result"] = {
k: v for k, v in result.items() if k not in _PRIVATE_STATE_FIELDS
}
# ``t.state`` is a ``RunnableConfig | StateSnapshot | None`` per
# ``PregelTask``'s typing. When ``subgraphs=True`` on the caller,
# this holds the subgraph's fully-materialized ``StateSnapshot`` —
# which repeats the same four leak surfaces (``values``,
# ``metadata.writes``, ``metadata.counters_since_delta_snapshot``,
# ``tasks[*].result/state``). Recurse so the whole tree is clean.
nested_state = getattr(t, "state", None)
if isinstance(nested_state, StateSnapshot):
scrubbed_state = _strip_private(nested_state)
if scrubbed_state is not nested_state:
replace_kwargs["state"] = scrubbed_state
if replace_kwargs:
new_tasks.append(t._replace(**replace_kwargs))
changed = True
else:
new_tasks.append(t)
if changed:
tasks = tuple(new_tasks)
return snap._replace(values=values, metadata=metadata, tasks=tasks)
class _EvoFilteredGraph(CompiledStateGraph):
"""Filters ``PrivateStateAttr``-marked state fields from checkpoint reads.
``Pregel.copy`` uses ``self.__class__(**attrs)`` so this subclass
survives the ``graph_obj.copy(update=...)`` call in
``langgraph_api.graph.get_graph`` that binds the checkpointer / store
before yielding to endpoint handlers.
**Known gap — streaming paths.** The overrides only cover ``get_state``
/ ``get_state_history``. On this compiled graph,
``self.output_channels`` correctly excludes ``_quickjs_snapshot_payload``
(respects ``OmitFromSchema(output=True)``), but
``self.stream_channels_asis`` includes it alongside other private
fields (``jump_to``, ``_summarization_event``) — the two lists are
built by ``langgraph.graph.state``'s graph builder and only the first
checks the output schema. So a client streaming with
``stream_mode="values"`` or ``stream_mode="events"`` (which fall back
to ``stream_channels_asis`` when ``output_keys`` is ``None``) can pull
the anchor blob in per-run event data. Empirically the WebUI's
``stream_mode=["updates"]`` path is clean, so this is transient per-run
rather than the persistent per-getState download this PR targets.
Filter here first; extend into the stream layer if a client relying on
``values`` / ``events`` reports it.
"""
async def aget_state(self, config, *, subgraphs=False):
return _strip_private(await super().aget_state(config, subgraphs=subgraphs))
def get_state(self, config, *, subgraphs=False):
return _strip_private(super().get_state(config, subgraphs=subgraphs))
async def aget_state_history(self, config, **kw):
async for snap in super().aget_state_history(config, **kw):
yield _strip_private(snap)
def get_state_history(self, config, **kw):
for snap in super().get_state_history(config, **kw):
yield _strip_private(snap)
# In-place ``__class__`` swap: the subclass adds only methods (no new
# instance attributes) so the memory layout is identical and the swap is
# safe. Constructing a fresh ``_EvoFilteredGraph`` via ``.copy()`` would
# require reproducing the deep-agent build pipeline; the swap avoids that.
_agent.__class__ = _EvoFilteredGraph
EvoScientist_agent = _agent
def _apply_filter_to_all_registered_graphs() -> None:
"""Extend the class swap to every graph registered in ``langgraph.json``.
``EvoScientist.py:_build_middleware_stack`` installs
``create_code_interpreter_middleware`` unconditionally — it's not gated
on the ``for_async_subagent`` flag — so every subagent (sync ``task``
dispatch and async ``start_async_task``) carries the QuickJS REPL and
can produce ``_quickjs_snapshot_payload`` writes on its own checkpoint
namespace.
Async subagents get their own ``thread_id`` and their ``/threads/{id}/state``
endpoint is served by their own compiled graph. Without swapping the
class on those graphs, the filter we applied to ``EvoScientist_agent``
doesn't reach that endpoint and any real code_interpreter touch inside
a subagent leaks the anchor snapshot verbatim.
Reads the graph registry straight from ``langgraph.json`` so a new
subagent added to the config picks up the swap automatically — no
hardcoded list to keep in sync.
Idempotent (skips graphs already swapped) and safe on graphs that don't
use the middleware — ``_strip_private`` returns snapshots unchanged when
the private field is absent. Best-effort: if the config is unreadable
or an entry can't be resolved, the deployment still starts — only the
unresolvable subagents remain unfiltered.
"""
import json
from importlib import import_module
from pathlib import Path
config_path = Path(__file__).parent / "langgraph.json"
try:
config = json.loads(config_path.read_text())
except (OSError, json.JSONDecodeError):
return
for path in config.get("graphs", {}).values():
# Format: "module.dotted.path:attr_name"
if ":" not in path:
continue
module_path, attr = path.rsplit(":", 1)
try:
module = import_module(module_path)
except ImportError:
continue
graph = getattr(module, attr, None)
if isinstance(graph, CompiledStateGraph) and not isinstance(
graph, _EvoFilteredGraph
):
graph.__class__ = _EvoFilteredGraph
_apply_filter_to_all_registered_graphs()
from EvoScientist.EvoScientist import EvoScientist_agent
__all__ = ["EvoScientist_agent"]
+5 -11
View File
@@ -306,19 +306,14 @@ def is_async_subagents_available() -> bool:
def _langgraph_exe() -> str | None:
"""Return the path to the langgraph CLI binary, or None if not found."""
import sys
executable_dir = os.path.dirname(sys.executable)
candidate_names = (
["langgraph.exe", "langgraph"] if os.name == "nt" else ["langgraph"]
)
for candidate_name in candidate_names:
candidate = os.path.join(executable_dir, candidate_name)
if os.path.isfile(candidate) and os.access(candidate, os.X_OK):
return candidate
found = shutil.which("langgraph")
if found:
return found
import sys as _sys
candidate = os.path.join(os.path.dirname(_sys.executable), "langgraph")
if os.path.isfile(candidate) and os.access(candidate, os.X_OK):
return candidate
return None
@@ -714,7 +709,6 @@ def start_langgraph_dev(
sub_env["EVOSCIENTIST_DEPLOY_MODE"] = "full" if deploy_mode else "stripped"
try:
logger.info("Starting langgraph dev with CLI: %s", exe)
proc = subprocess.Popen(
[
exe,
+2 -6
View File
@@ -20,9 +20,9 @@ _KNOWN_MODEL_CONTEXT_WINDOWS: dict[str, int] = {
# Qwen 3.7 closed-source tiers — Max flagship and Plus (1M).
"qwen3.7-max": 1_000_000,
"qwen3.7-plus": 1_000_000,
# xAI Grok — per-model windows (build-0.1: 256K, 4.5: 500K).
# xAI Grok — per-model windows (build-0.1: 256K, 4.3: 1M).
"grok-build-0.1": 256_000,
"grok-4.5": 500_000,
"grok-4.3": 1_000_000,
# Claude Haiku 4.5 — exception to the ``claude-`` family (200K, not 1M).
"claude-haiku-4-5": 200_000,
# MiniMax M3 — 1M context (M2.x variants stay at provider default ~204K).
@@ -32,8 +32,6 @@ _KNOWN_MODEL_CONTEXT_WINDOWS: dict[str, int] = {
# Zhipu GLM-5.2 — 1M context, an exception to the ``glm-5`` family (203K).
# Matches OpenRouter ``z-ai/glm-5.2`` via split('/')[-1].
"glm-5.2": 1_000_000,
# Tencent Hunyuan HY3 — 262K context (OpenRouter ``tencent/hy3``).
"hy3": 262_000,
}
# Family-level fallbacks: tried only after exact-name lookup misses.
@@ -42,8 +40,6 @@ _KNOWN_MODEL_CONTEXT_WINDOWS: dict[str, int] = {
_KNOWN_MODEL_FAMILIES: list[tuple[str, int]] = [
# All Claude — 1M via the ``context-1m-2025-08-07`` beta header.
("claude-", 1_000_000),
# OpenAI GPT-5.6 family — sol, terra, luna variants
("gpt-5.6", 1_050_000),
# OpenAI GPT-5.5 family — base, pro, future variants
("gpt-5.5", 1_050_000),
# Google Gemini 3.x family — flash, flash-lite, pro (1.05M). Excludes 2.5.
-386
View File
@@ -1,386 +0,0 @@
"""Provider-error surface for langgraph SSE frames.
Provides :class:`ProviderStreamError` — a normalized, non-dataclass
exception raised by ``ErrorNormalizationMiddleware`` in place of the
provider SDK exception that a chat model call raised. Non-dataclass on
purpose: since orjson 3.0, dataclass instances are serialized natively
via their field enumeration, skipping the ``default=`` hook that
would otherwise build our SSE envelope. Some provider SDKs (openrouter
today) decorate their exceptions with ``@dataclass``, so their errors
emerge on the wire as raw dataclass fields — no envelope, no way for
the WebUI to distinguish quota / auth / rate-limit. Wrapping them in
a plain ``Exception`` subclass here keeps orjson on the ``default=``
path, which then calls :meth:`ProviderStreamError.model_dump`
(upstream ``langgraph_api.serde.default`` checks that hook before its
``BaseException`` branch) — no serde monkey-patch needed.
Also lives here: the pure-function helpers the middleware uses to
build the envelope (provider tag from ``ModelRequest.model``, SDK
field extractors, env-driven API-key redaction). They stay next to
:class:`ProviderStreamError` because the middleware is their only
consumer.
"""
from __future__ import annotations
import os
import re
from typing import Any
# ---------------------------------------------------------------------------
# ProviderStreamError
# ---------------------------------------------------------------------------
class AgentControlError(Exception):
"""Host-defined terminal control error that must bypass model fallback."""
non_fallbackable = True
def __init__(
self,
code: str,
message: str,
*,
status_code: int = 403,
retryable: bool = False,
) -> None:
super().__init__(message)
self.code = code
self.message = message
self.status_code = status_code
self.retryable = retryable
def model_dump(self) -> dict[str, Any]:
return {
"error": type(self).__name__,
"code": self.code,
"message": self.message,
"status_code": self.status_code,
"retryable": self.retryable,
}
class ModelToolProtocolError(AgentControlError):
"""A completed model response contained an invalid tool-call protocol."""
def __init__(
self,
reason: str,
*,
provider: str | None = None,
model: str | None = None,
route_key: str | None = None,
config_generation: int | None = None,
api_mode: str | None = None,
endpoint: str | None = None,
tool_call_transport: str | None = None,
call_id: str | None = None,
call_diagnostic: dict[str, Any] | None = None,
) -> None:
super().__init__(
"MODEL_TOOL_PROTOCOL_INVALID",
"The model returned an invalid structured tool call.",
status_code=502,
retryable=False,
)
self.reason = reason
self.provider = provider
self.model = model
self.route_key = route_key
self.config_generation = config_generation
self.api_mode = api_mode
self.endpoint = endpoint
self.tool_call_transport = tool_call_transport
self.call_id = call_id
# Internal-only, redacted structure for server logs. Deliberately omitted
# from model_dump() so it never becomes part of the public SSE contract.
self.call_diagnostic = dict(call_diagnostic or {})
self.fallbackable = True
self.recoverable = True
def model_dump(self) -> dict[str, Any]:
payload = super().model_dump()
payload.update(
{
"reason": self.reason,
"fallbackable": self.fallbackable,
"recoverable": self.recoverable,
}
)
for key in (
"provider",
"model",
"route_key",
"config_generation",
"api_mode",
"endpoint",
"tool_call_transport",
"call_id",
):
value = getattr(self, key)
if value is not None:
payload[key] = value
return payload
class ProviderStreamError(Exception):
"""Envelope-shaped wrapper for a provider SDK exception raised
inside a chat model call.
Attributes mirror the SSE envelope one-for-one:
- ``provider`` — concrete provider tag (``openai`` / ``anthropic``
/ ``deepseek`` / ``openrouter`` / ``openai_compat`` / …)
- ``class_qualname`` — fully qualified name of the underlying
exception's class (e.g. ``openrouter.errors.…``)
- ``message`` — API-key-redacted ``str(exc)``
- ``status_code`` — HTTP status if the SDK exposed one
- ``code`` — provider error code (``insufficient_quota``, …)
- ``err_type`` — provider error type label (openai's ``.type``)
- ``request_id`` — SDK-provided correlation id
The underlying exception is available via ``__cause__`` (set by
``raise ProviderStreamError(...) from exc`` in the middleware).
"""
def __init__(
self,
provider: str,
class_qualname: str,
message: str,
*,
status_code: int | None = None,
code: str | None = None,
err_type: str | None = None,
request_id: str | None = None,
) -> None:
super().__init__(message)
self.provider = provider
self.class_qualname = class_qualname
self.message = message
self.status_code = status_code
self.code = code
self.err_type = err_type
self.request_id = request_id
def as_envelope(self) -> dict[str, Any]:
"""Return the SSE envelope dict — the shape the WebUI consumes."""
payload: dict[str, Any] = {
"error": self.class_qualname.rsplit(".", 1)[-1],
"class": self.class_qualname,
"message": self.message,
"provider": self.provider,
}
if self.status_code is not None:
payload["status_code"] = self.status_code
if self.code is not None:
payload["code"] = self.code
if self.err_type is not None:
payload["type"] = self.err_type
if self.request_id:
payload["request_id"] = self.request_id
return payload
def model_dump(self) -> dict[str, Any]:
"""Serialization hook consumed by ``langgraph_api.serde.default``.
Upstream's dispatch checks ``hasattr(obj, 'model_dump')`` BEFORE
the ``isinstance(obj, BaseException)`` branch, so exposing this
method lets upstream emit our envelope with no monkey-patch on
its ``default`` callable. The name matches Pydantic's
convention deliberately — it's the hook upstream is looking
for.
"""
return self.as_envelope()
# ---------------------------------------------------------------------------
# API-key redaction — env-driven, prefix-only
# ---------------------------------------------------------------------------
#
# Redaction is built from credentials actually deployed via env vars,
# not from generic key shapes. Rationale: (a) zero false positives —
# we only scrub strings we know are secrets, (b) defense-in-depth —
# the compiled regex holds only the first 8 chars of each key, so a
# leak of the regex object itself (traceback locals, process dump)
# can't expose the secret. Suffix-greedy match consumes the rest of
# the key shape at runtime. The table is rebuilt on every
# ``_redact_api_keys`` call so credentials loaded after import
# (typically ``load_dotenv`` in a main entry point) still get
# scrubbed. ``re.compile`` caches by source string internally, so an
# unchanged env costs a dict lookup.
_API_KEY_ENV_SUFFIXES = ("_API_KEY", "_TOKEN", "_SECRET")
_API_KEY_MIN_LEN = 12
_API_KEY_PREFIX_LEN = 8
def _build_env_key_redaction_re() -> re.Pattern[str] | None:
prefixes: list[str] = []
for k, v in os.environ.items():
if not k.endswith(_API_KEY_ENV_SUFFIXES):
continue
if not isinstance(v, str) or len(v) < _API_KEY_MIN_LEN:
continue
prefixes.append(re.escape(v[:_API_KEY_PREFIX_LEN]))
if not prefixes:
return None
alternation = "|".join(f"{p}[A-Za-z0-9_+/=.-]*" for p in prefixes)
return re.compile(alternation)
def _redact_api_keys(message: str) -> str:
"""Replace any deployed key prefix in *message* with ``<redacted>``.
Defensive; provider error messages occasionally echo the
authorization header back. Rebuilt per call so credentials loaded
after import (typical ``load_dotenv`` pattern) are still redacted.
"""
pattern = _build_env_key_redaction_re()
if pattern is None:
return message
return pattern.sub("<redacted>", message)
# ---------------------------------------------------------------------------
# Provider inference from ModelRequest.model
# ---------------------------------------------------------------------------
#
# Host → concrete provider. Hand-maintained snapshot mirroring the
# routed-provider tables in ``llm/models.py``
# (``_OPENAI_ROUTED_PROVIDERS`` + ``_ANTHROPIC_ROUTED_PROVIDERS``).
# Kept here rather than imported from ``models.py`` to keep the
# import surface of ``errors.py`` minimal — importing ``models.py``
# would pull in every langchain chat-model client at first
# middleware access. Consumed by ``_lookup_host_or_compat``; unknown
# hosts fall back to ``<module>_compat`` so the WebUI knows
# "openai/anthropic SDK, but not native" instead of getting a
# misleading concrete tag. Update when a new routed provider is
# added to ``models.py``.
#
# Related sibling: ``_PROVIDER_EXC_MODULE_PREFIXES`` in
# ``middleware/error_normalization.py`` — the exception-side
# provider allow-list. Adding a whole new provider SDK (not just a
# new base_url routed through an existing one) means updating that
# list too.
_HOST_TO_PROVIDER: dict[str, str] = {
"api.openai.com": "openai",
"api.anthropic.com": "anthropic",
"api.deepseek.com": "deepseek",
"api.moonshot.cn": "moonshot",
"api.siliconflow.cn": "siliconflow",
"open.bigmodel.cn": "zhipu", # zhipu + zhipu-code share this host
"ark.cn-beijing.volces.com": "volcengine",
"dashscope.aliyuncs.com": "dashscope",
"coding.dashscope.aliyuncs.com": "dashscope",
"api.minimaxi.com": "minimax",
"api.kimi.com": "kimi", # kimi-coding shares this host
"openrouter.ai": "openrouter",
}
def _provider_from_model(model: Any) -> str | None:
"""Derive the concrete provider tag from a chat model instance.
Class-based dispatch for unambiguous providers (``ChatOpenRouter``,
``ChatGoogleGenerativeAI``); ``openai_api_base`` /
``anthropic_api_url`` looked up in ``_HOST_TO_PROVIDER`` for
openai/anthropic-shape clients (native + routed). Returns ``None``
when the model isn't from a recognized provider SDK — the caller
(``ErrorNormalizationMiddleware``) then passes the exception
through unchanged.
"""
cls_module = type(model).__module__ or ""
if cls_module.startswith("langchain_openrouter"):
return "openrouter"
if cls_module.startswith("langchain_google_genai"):
return "google_genai"
if cls_module.startswith("langchain_openai"):
return _lookup_host_or_compat(
getattr(model, "openai_api_base", None), module_tag="openai"
)
if cls_module.startswith("langchain_anthropic"):
return _lookup_host_or_compat(
getattr(model, "anthropic_api_url", None), module_tag="anthropic"
)
return None
def _lookup_host_or_compat(base_url: str | None, module_tag: str) -> str:
"""Extract host from *base_url* and look up in ``_HOST_TO_PROVIDER``.
Falls back to *module_tag* when no ``base_url`` is set (native SDK
default endpoint) or ``<module_tag>_compat`` for an unrecognized
host — the honest "openai SDK shape but unknown upstream" tag.
"""
if not base_url:
return module_tag
try:
from urllib.parse import urlparse
host = urlparse(base_url).hostname
except Exception:
host = None
if not host:
return module_tag
return _HOST_TO_PROVIDER.get(host.lower(), f"{module_tag}_compat")
# ---------------------------------------------------------------------------
# SDK-field extractors — populate the envelope's optional fields
# ---------------------------------------------------------------------------
def _extract_status_code(exc: BaseException) -> int | None:
"""Best-effort HTTP status code from a provider SDK exception.
Order matters: openai/anthropic store it on ``.status_code``;
httpx-wrappers expose it via ``.response.status_code``;
``google.genai.errors.APIError`` (unusually) stores it as an
integer ``.code`` — type-disambiguated from openai/anthropic's
string ``.code`` (provider error code, surfaced separately).
"""
status_code = getattr(exc, "status_code", None)
if isinstance(status_code, int):
return status_code
response = getattr(exc, "response", None)
if response is not None:
rsc = getattr(response, "status_code", None)
if isinstance(rsc, int):
return rsc
code = getattr(exc, "code", None)
if isinstance(code, int):
return code
return None
def _extract_provider_code(exc: BaseException) -> str | None:
"""Provider error code (e.g. ``insufficient_quota``,
``invalid_api_key``). Distinct from HTTP status; higher signal for
a WebUI toast than the integer alone.
"""
code = getattr(exc, "code", None)
if isinstance(code, str) and code:
return code
return None
def _extract_error_type(exc: BaseException) -> str | None:
"""Provider error type label.
- openai exposes this as ``.type`` (``rate_limit_error`` etc.)
- ``google.genai.errors.APIError`` stores a string label at
``.status`` (``"NOT_FOUND"``, ``"RESOURCE_EXHAUSTED"``, …) — a
good fit for the same field.
``.type`` takes precedence when both are set.
"""
err_type = getattr(exc, "type", None)
if isinstance(err_type, str) and err_type:
return err_type
status = getattr(exc, "status", None)
if isinstance(status, str) and status:
return status
return None
+38 -245
View File
@@ -10,19 +10,11 @@ endpoints) and convenient short names for common models.
from __future__ import annotations
import os
import re
import subprocess
import warnings
from functools import lru_cache
from typing import Any
from langchain.chat_models import init_chat_model
from ..config.settings import (
OPENROUTER_DEFAULT_APP_CATEGORIES,
OPENROUTER_DEFAULT_APP_TITLE,
OPENROUTER_DEFAULT_HTTP_REFERER,
)
from .context_window import apply_known_context_window
from .patches import (
_is_ccproxy_codex,
@@ -45,50 +37,6 @@ _DEEPSEEK_BASE_URL = "https://api.deepseek.com"
_MOONSHOT_BASE_URL = "https://api.moonshot.cn/v1"
_KIMI_CODING_BASE_URL = "https://api.kimi.com/coding/"
# Minimum Codex CLI version advertised when no explicit override is set. Newer
# installed versions are advertised automatically.
_CODEX_CLIENT_VERSION_FALLBACK = "0.144.1"
@lru_cache(maxsize=1)
def _installed_codex_client_version() -> str:
"""Return the installed Codex CLI version, or an empty string."""
try:
result = subprocess.run(
["codex", "--version"],
capture_output=True,
text=True,
timeout=2,
check=False,
)
except (OSError, subprocess.TimeoutExpired):
return ""
if result.returncode != 0:
return ""
match = re.search(r"\b(\d+\.\d+\.\d+)\b", result.stdout + result.stderr)
return match.group(1) if match else ""
def _resolve_codex_client_version() -> str:
"""Resolve an explicit override or the newer of installed and minimum versions."""
override = os.environ.get("EVOSCIENTIST_CODEX_CLIENT_VERSION", "").strip()
if override:
return override
installed = _installed_codex_client_version()
if installed and tuple(map(int, installed.split("."))) >= tuple(
map(int, _CODEX_CLIENT_VERSION_FALLBACK.split("."))
):
return installed
return _CODEX_CLIENT_VERSION_FALLBACK
def _resolve_reasoning_effort(default: str) -> str:
"""Return the configured reasoning effort or a provider-specific default."""
return os.environ.get("EVOSCIENTIST_REASONING_EFFORT", "").strip() or default
# Providers routed through the OpenAI provider with a custom base_url.
# Maps provider name → (base_url or None, env var for API key).
_OPENAI_ROUTED_PROVIDERS: dict[str, tuple[str | None, str]] = {
@@ -120,19 +68,6 @@ _THINKING_CAPABLE_PROVIDERS: set[str] = {"minimax"}
_TRUTHY_ENV_VALUES = {"1", "true", "yes", "on"}
_FALSEY_ENV_VALUES = {"0", "false", "no", "off"}
# OpenRouter app attribution (issue #339). Default values are the single source
# of truth in config/settings.py (imported above); langchain-openrouter maps
# app_url → HTTP-Referer, app_title → X-Title, app_categories →
# X-OpenRouter-Categories. OpenRouter honors at most this many categories per
# request (server-side limit) and silently ignores the rest, so the sent list is
# capped to this many below. https://openrouter.ai/docs/app-attribution
_OPENROUTER_MAX_CATEGORIES_PER_REQUEST = 2
# Legacy/provider-specific options that are not accepted by the installed
# LangChain chat model constructors. Leaving them at the top level makes
# LangChain move them into model_kwargs and can later leak them into SDK calls.
_UNSUPPORTED_CHAT_MODEL_KWARGS = frozenset({"sanitize_openai_sdk_headers"})
# Model registry: list of (short_name, model_id, provider)
# Allows same short_name across different providers.
_MODEL_ENTRIES: list[tuple[str, str, str]] = [
@@ -154,9 +89,6 @@ _MODEL_ENTRIES: list[tuple[str, str, str]] = [
("claude-sonnet-4-6", "claude-sonnet-4-6", "anthropic"),
("claude-haiku-4-5", "claude-haiku-4-5", "anthropic"),
# OpenAI
("gpt-5.6-sol", "gpt-5.6-sol", "openai"),
("gpt-5.6-terra", "gpt-5.6-terra", "openai"),
("gpt-5.6-luna", "gpt-5.6-luna", "openai"),
("gpt-5.5-pro", "gpt-5.5-pro", "openai"),
("gpt-5.5", "gpt-5.5", "openai"),
("gpt-5.4", "gpt-5.4", "openai"),
@@ -213,9 +145,6 @@ _MODEL_ENTRIES: list[tuple[str, str, str]] = [
("claude-opus-4.8-fast", "anthropic/claude-opus-4.8-fast", "openrouter"),
("claude-sonnet-5", "anthropic/claude-sonnet-5", "openrouter"),
("claude-sonnet-4.6", "anthropic/claude-sonnet-4.6", "openrouter"),
("gpt-5.6-sol", "openai/gpt-5.6-sol", "openrouter"),
("gpt-5.6-terra", "openai/gpt-5.6-terra", "openrouter"),
("gpt-5.6-luna", "openai/gpt-5.6-luna", "openrouter"),
("gpt-5.5-pro", "openai/gpt-5.5-pro", "openrouter"),
("gpt-5.5", "openai/gpt-5.5", "openrouter"),
("gpt-5.4", "openai/gpt-5.4", "openrouter"),
@@ -230,8 +159,7 @@ _MODEL_ENTRIES: list[tuple[str, str, str]] = [
("mimo-v2.5-pro", "xiaomi/mimo-v2.5-pro", "openrouter"),
("mimo-v2.5", "xiaomi/mimo-v2.5", "openrouter"),
("grok-build-0.1", "x-ai/grok-build-0.1", "openrouter"),
("grok-4.5", "x-ai/grok-4.5", "openrouter"),
("hy3", "tencent/hy3", "openrouter"),
("grok-4.3", "x-ai/grok-4.3", "openrouter"),
("qwen3.7-max", "qwen/qwen3.7-max", "openrouter"),
("qwen3.7-plus", "qwen/qwen3.7-plus", "openrouter"),
("qwen3.6-flash", "qwen/qwen3.6-flash", "openrouter"),
@@ -329,15 +257,6 @@ def _env_flag_disabled(name: str) -> bool:
return value is not None and value.strip().lower() in _FALSEY_ENV_VALUES
def _drop_unsupported_chat_model_kwargs(kwargs: dict[str, Any]) -> None:
for key in _UNSUPPORTED_CHAT_MODEL_KWARGS:
kwargs.pop(key, None)
model_kwargs = kwargs.get("model_kwargs")
if isinstance(model_kwargs, dict):
for key in _UNSUPPORTED_CHAT_MODEL_KWARGS:
model_kwargs.pop(key, None)
def _supports_openrouter_anthropic_prompt_cache(provider: str, model_id: str) -> bool:
"""Return whether EvoScientist should declare OpenRouter Claude caching."""
return provider == "openrouter" and model_id.startswith(
@@ -395,16 +314,8 @@ def _apply_auto_config(
Mutates *kwargs* in place. Only sets keys that the caller hasn't already
provided, so explicit user settings are never overridden.
"""
disable_reasoning = bool(kwargs.pop("_disable_reasoning", False))
disable_thinking = bool(kwargs.pop("_disable_thinking", False))
if disable_reasoning:
kwargs.pop("reasoning", None)
kwargs.pop("include_thoughts", None)
if disable_thinking:
kwargs.pop("thinking", None)
# Anthropic: extended thinking
if provider == "anthropic" and not disable_thinking and "thinking" not in kwargs:
if provider == "anthropic" and "thinking" not in kwargs:
_supports_thinking = original_provider in _THINKING_CAPABLE_PROVIDERS
# Detect local proxy (e.g. ccproxy): thinking blocks in conversation
# history cause 422 errors because the proxy doesn't accept 'thinking'
@@ -423,31 +334,24 @@ def _apply_auto_config(
kwargs["thinking"] = {"type": "enabled", "budget_tokens": 10000}
# OpenAI (native, not third-party routed): reasoning
if (
provider == "openai"
and not is_third_party
and not disable_reasoning
and "reasoning" not in kwargs
):
_default_effort = (
"xhigh"
if (
"5.4" in model_id
or "5.5" in model_id
or "5.6" in model_id
or "codex" in model_id
if provider == "openai" and not is_third_party and "reasoning" not in kwargs:
if _is_ccproxy_codex():
# ccproxy uses Chat Completions which doesn't support reasoning.
pass
else:
_eff = (
"xhigh"
if ("5.4" in model_id or "5.5" in model_id or "codex" in model_id)
else "high"
)
else "high"
)
_eff = _resolve_reasoning_effort(_default_effort)
kwargs["reasoning"] = {"effort": _eff, "summary": "auto"}
kwargs["reasoning"] = {"effort": _eff, "summary": "auto"}
# Google GenAI: surface thinking traces
if provider == "google-genai" and not disable_reasoning:
if provider == "google-genai":
kwargs.setdefault("include_thoughts", True)
# Ollama: separate reasoning content from response for thinking models
if provider == "ollama" and not disable_reasoning and "reasoning" not in kwargs:
if provider == "ollama" and "reasoning" not in kwargs:
kwargs["reasoning"] = True
@@ -474,46 +378,7 @@ def get_chat_model(
>>> model = get_chat_model("gpt-4o") # OpenAI model
>>> model = get_chat_model("claude-3-opus-20240229", provider="anthropic") # Full ID
"""
skip_runtime_resolver = bool(kwargs.pop("_skip_runtime_model_resolver", False))
runtime_provider_name: str | None = None
runtime_supports_reasoning: bool | None = None
runtime_resolved = None
if not skip_runtime_resolver:
from EvoScientist.runtime_integrations import resolve_runtime_model
runtime_resolved = resolve_runtime_model(model, provider)
if runtime_resolved is not None:
resolved_params = dict(getattr(runtime_resolved, "params", {}) or {})
extra_body = resolved_params.pop("_extra_body", None)
default_headers = resolved_params.pop("_default_headers", None)
if extra_body:
resolved_params["extra_body"] = extra_body
if default_headers:
resolved_params["default_headers"] = default_headers
resolved_params.update(kwargs)
kwargs = resolved_params
resolved_api_key = str(getattr(runtime_resolved, "api_key", "") or "")
resolved_base_url = str(getattr(runtime_resolved, "base_url", "") or "")
if resolved_api_key:
kwargs.setdefault("api_key", resolved_api_key)
if resolved_base_url:
kwargs.setdefault("base_url", resolved_base_url.rstrip("/"))
runtime_provider_name = str(
getattr(runtime_resolved, "provider_name", "") or ""
)
runtime_supports_reasoning = bool(
getattr(runtime_resolved, "supports_reasoning", False)
)
if not runtime_supports_reasoning:
kwargs.setdefault("_disable_reasoning", True)
kwargs.setdefault("_disable_thinking", True)
model = str(runtime_resolved.model_id)
provider = str(runtime_resolved.protocol)
else:
model = model or DEFAULT_MODEL
model = model or DEFAULT_MODEL
# Look up short name in registry (provider-aware)
model_id = None
@@ -548,35 +413,22 @@ def get_chat_model(
_is_third_party = (
provider in _OPENAI_ROUTED_PROVIDERS or provider in _ANTHROPIC_ROUTED_PROVIDERS
)
if runtime_provider_name and runtime_provider_name != provider:
_is_third_party = True
if (
runtime_resolved is not None
and provider == "openai"
and resolved_base_url
and "api.openai.com" not in resolved_base_url.lower()
):
_is_third_party = True
_is_openai_proxy = False
_original_provider: str | None = (
runtime_provider_name if runtime_provider_name != provider else None
)
_original_provider: str | None = None
if provider == "anthropic":
base_url = os.environ.get("ANTHROPIC_BASE_URL", "")
if base_url:
kwargs.setdefault("base_url", base_url)
kwargs["base_url"] = base_url
api_key = os.environ.get("ANTHROPIC_API_KEY", "")
if api_key:
kwargs.setdefault("api_key", api_key)
kwargs["api_key"] = api_key
# Native OpenAI base_url override (e.g. ccproxy Codex at localhost:8000/codex/v1)
elif provider == "openai":
base_url = os.environ.get("OPENAI_BASE_URL", "")
if base_url:
kwargs.setdefault("base_url", base_url)
_is_openai_proxy = _is_ccproxy_codex(
kwargs.get("base_url"), kwargs.get("api_key")
)
kwargs["base_url"] = base_url
_is_openai_proxy = _is_ccproxy_codex()
if _is_openai_proxy:
# Use Responses API for ccproxy: bypasses the format chain
# converter (Chat→Responses→Chat) which returns 502 on
@@ -589,23 +441,9 @@ def get_chat_model(
# for Chat Completions tool_call duplication — not an issue
# with the Responses API SSE format.)
kwargs.pop("streaming", None) # remove if set elsewhere
# ccproxy forwards client headers upstream and only
# gap-fills its own, so the Codex backend sees this
# client's identity. Without Codex-CLI-shaped headers it
# rejects current models ("The '<model>' model requires
# a newer version of Codex").
_codex_ver = _resolve_codex_client_version()
_headers = kwargs.get("default_headers") or {}
kwargs["default_headers"] = _headers
_headers.setdefault("originator", "codex_cli_rs")
_headers.setdefault("version", _codex_ver)
_headers.setdefault(
"User-Agent",
f"codex_cli_rs/{_headers['version']} (EvoScientist)",
)
api_key = os.environ.get("OPENAI_API_KEY", "")
if api_key:
kwargs.setdefault("api_key", api_key)
kwargs["api_key"] = api_key
# OpenAI-routed providers → route through OpenAI provider with base_url
elif provider in _OPENAI_ROUTED_PROVIDERS:
@@ -623,10 +461,10 @@ def get_chat_model(
else:
base_url = base_url_default
if base_url:
kwargs.setdefault("base_url", base_url)
kwargs["base_url"] = base_url
api_key = os.environ.get(api_key_env, "")
if api_key:
kwargs.setdefault("api_key", api_key)
kwargs["api_key"] = api_key
# SiliconFlow: disable thinking — LangChain drops reasoning_content
# from history, causing error 20015 on multi-turn requests.
if provider == "siliconflow":
@@ -636,6 +474,15 @@ def get_chat_model(
# Even native thinking models like kimi-k2-thinking operate in non-thinking mode.
if provider == "moonshot":
kwargs.setdefault("extra_body", {})["thinking"] = {"type": "disabled"}
# custom-openai: some Codex-orientated reseller gateways (e.g.
# hunnuapi.top) blacklist the openai SDK's default
# "OpenAI/Python" User-Agent at the app layer and return
# 403 "Your request was blocked." to any langchain-openai client.
# Spoof a Codex-CLI UA so requests go through.
if provider == "custom-openai":
kwargs.setdefault("default_headers", {}).setdefault(
"User-Agent", "codex_cli_rs/0.0.0"
)
provider = "openai"
# OpenRouter → native ChatOpenRouter via init_chat_model.
@@ -643,61 +490,15 @@ def get_chat_model(
_is_third_party = True
api_key = os.environ.get("OPENROUTER_API_KEY", "")
if api_key:
kwargs.setdefault("api_key", api_key)
kwargs["api_key"] = api_key
# Reasoning via `effort` + `summary: "auto"` so a readable reasoning
# summary is returned for display. OpenAI-Responses also emits encrypted
# reasoning items (`rs_*` id) that can't be replayed on multi-turn
# passback (OpenRouter's `/responses` beta is stateless, store=false —
# "Item with id 'rs_...' not found"); the patch strips them on passback,
# so enabling `summary` is safe. See langchain-ai/langchain#37777.
effort = _resolve_reasoning_effort("high")
effort = os.environ.get("EVOSCIENTIST_REASONING_EFFORT", "").strip() or "high"
kwargs.setdefault("reasoning", {"effort": effort, "summary": "auto"})
# App attribution (issue #339): identify EvoScientist to OpenRouter so
# usage is credited to the project (app rankings, model app tabs,
# analytics) rather than langchain-openrouter's LangChain-branded
# defaults. setdefault so an explicit caller kwarg wins; values are
# configurable via EVOSCIENTIST_OPENROUTER_* env (fed from the config
# file by apply_config_to_env). Applied only here, so no other provider
# ever receives these kwargs.
kwargs.setdefault(
"app_url",
os.environ.get("EVOSCIENTIST_OPENROUTER_HTTP_REFERER", "").strip()
or OPENROUTER_DEFAULT_HTTP_REFERER,
)
kwargs.setdefault(
"app_title",
os.environ.get("EVOSCIENTIST_OPENROUTER_APP_TITLE", "").strip()
or OPENROUTER_DEFAULT_APP_TITLE,
)
# app_categories must be a list[str] (langchain-openrouter joins it into
# the X-OpenRouter-Categories header); split the comma-separated config
# value and drop blanks so a stray comma/space can't emit an empty one.
_app_categories_raw = (
os.environ.get("EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "").strip()
or OPENROUTER_DEFAULT_APP_CATEGORIES
)
_app_categories = [
c.strip() for c in _app_categories_raw.split(",") if c.strip()
]
# Cap to the per-request limit and warn, so a misconfigured extra is
# dropped predictably here (and surfaced to the user) rather than being
# silently truncated server-side.
_limit = _OPENROUTER_MAX_CATEGORIES_PER_REQUEST
if len(_app_categories) > _limit:
warnings.warn(
f"OpenRouter accepts at most {_limit} app categories per "
f"request, so only the first {_limit} are sent: "
f"{_app_categories[:_limit]}. Ignoring the rest: "
f"{_app_categories[_limit:]}. Set "
f"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES (or the "
f"openrouter_app_categories config) to at most {_limit} "
f"categories to silence this warning.",
UserWarning,
stacklevel=2,
)
_app_categories = _app_categories[:_limit]
if _app_categories:
kwargs.setdefault("app_categories", _app_categories)
_patch_openrouter_strip_responses_reasoning()
# Anthropic-routed providers → route through Anthropic provider with base_url
@@ -718,10 +519,10 @@ def get_chat_model(
else:
base_url = base_url_default
if base_url:
kwargs.setdefault("base_url", base_url)
kwargs["base_url"] = base_url
api_key = os.environ.get(api_key_env, "")
if api_key:
kwargs.setdefault("api_key", api_key)
kwargs["api_key"] = api_key
# Kimi Coding Plan requires claude-code User-Agent header
if provider == "kimi-coding":
kwargs.setdefault("default_headers", {})["User-Agent"] = "claude-code/0.1.0"
@@ -730,9 +531,8 @@ def get_chat_model(
elif provider == "ollama":
base_url = os.environ.get("OLLAMA_BASE_URL", "")
if base_url:
kwargs.setdefault("base_url", base_url)
kwargs["base_url"] = base_url
_drop_unsupported_chat_model_kwargs(kwargs)
_apply_auto_config(provider, model_id, _is_third_party, kwargs, _original_provider)
_apply_openrouter_anthropic_prompt_cache(provider, model_id, kwargs)
@@ -749,14 +549,7 @@ def get_chat_model(
elif _responses_api_setting == "true":
kwargs["use_responses_api"] = True
anthropic_auth_token = None
if provider == "anthropic" and kwargs.get("api_key"):
anthropic_auth_token = os.environ.pop("ANTHROPIC_AUTH_TOKEN", None)
try:
chat_model = init_chat_model(model=model_id, model_provider=provider, **kwargs)
finally:
if anthropic_auth_token is not None:
os.environ["ANTHROPIC_AUTH_TOKEN"] = anthropic_auth_token
chat_model = init_chat_model(model=model_id, model_provider=provider, **kwargs)
# Flatten list content to strings for strict OpenAI-compatible providers
# (DeepSeek, SiliconFlow, OpenRouter, custom-openai, etc.) and
+3 -367
View File
@@ -25,7 +25,6 @@ Utilities:
from __future__ import annotations
import hashlib
import os
from typing import Any
@@ -179,20 +178,15 @@ _patch_ccproxy_codex_compat()
# ---------------------------------------------------------------------------
# Utility: detect ccproxy's Codex adapter (as opposed to generic localhost).
# ---------------------------------------------------------------------------
def _is_ccproxy_codex(
base_url: str | None = None,
api_key: str | None = None,
) -> bool:
def _is_ccproxy_codex() -> bool:
"""Return True if the OpenAI endpoint is ccproxy's Codex adapter.
Checks for the ccproxy-specific markers set by ``setup_codex_env()``
in ``ccproxy_manager.py``: the sentinel API key and the ``/codex/v1``
path. Plain localhost endpoints (vLLM, Ollama, etc.) are not affected.
"""
if base_url is None:
base_url = os.environ.get("OPENAI_BASE_URL", "")
if api_key is None:
api_key = os.environ.get("OPENAI_API_KEY", "")
base_url = os.environ.get("OPENAI_BASE_URL", "")
api_key = os.environ.get("OPENAI_API_KEY", "")
return (
("127.0.0.1" in base_url or "localhost" in base_url)
and api_key == "ccproxy-oauth"
@@ -273,299 +267,6 @@ def _flatten_message_content(content: Any) -> str | list[Any] | Any:
return "\n\n".join(parts) if parts else ""
def _stable_tool_call_id(message: Any, message_index: int, call_index: int) -> str:
seed = ":".join(
(
str(getattr(message, "id", "") or "message"),
str(message_index),
str(call_index),
)
)
return "call_" + hashlib.sha256(seed.encode("utf-8")).hexdigest()[:24]
def _tool_message_match_index(
tool_messages: list[Any],
used_indexes: set[int],
*,
call_id: str,
call_name: str,
) -> int | None:
"""Find the best unused result for one assistant tool call."""
def _matches(index: int, *, require_id: bool, require_name: bool) -> bool:
if index in used_indexes:
return False
message = tool_messages[index]
result_id = str(getattr(message, "tool_call_id", "") or "")
result_name = str(getattr(message, "name", "") or "")
if require_id and result_id != call_id:
return False
if not require_id and result_id:
return False
return not require_name or not result_name or result_name == call_name
if call_id:
for require_name in (True, False):
for index in range(len(tool_messages)):
if _matches(index, require_id=True, require_name=require_name):
return index
for require_name in (True, False):
for index in range(len(tool_messages)):
if _matches(index, require_id=False, require_name=require_name):
return index
return None
# A result-side identifier is more authoritative than a generated fallback.
for require_name in (True, False):
for index, message in enumerate(tool_messages):
if index in used_indexes:
continue
result_id = str(getattr(message, "tool_call_id", "") or "")
result_name = str(getattr(message, "name", "") or "")
if result_id and (
not require_name or not result_name or result_name == call_name
):
return index
for require_name in (True, False):
for index in range(len(tool_messages)):
if _matches(index, require_id=False, require_name=require_name):
return index
return None
def _copy_ai_message_with_tool_pairs(
message: Any,
message_index: int,
tool_messages: list[Any],
) -> tuple[Any | None, list[Any]]:
"""Return a replay-safe assistant message and its matched tool results."""
import copy
copied = copy.copy(message)
additional_kwargs = dict(getattr(message, "additional_kwargs", None) or {})
# Parsed tool_calls are canonical. Raw copies can otherwise re-introduce an
# invalid call after invalid_tool_calls has been cleared.
additional_kwargs.pop("tool_calls", None)
copied.additional_kwargs = additional_kwargs
copied.invalid_tool_calls = []
original_calls = list(getattr(message, "tool_calls", None) or [])
used_results: set[int] = set()
matched_calls: list[dict[str, Any]] = []
matched_result_indexes: list[int] = []
original_to_matched_call: dict[int, tuple[str, str]] = {}
for call_index, original_call in enumerate(original_calls):
call = dict(original_call)
call_id = str(call.get("id") or "")
call_name = str(call.get("name") or "").strip()
# A missing name is structurally unreplayable. Never infer it from
# arguments or retain its paired ToolMessage in provider history.
if not call_name:
continue
call["name"] = call_name
result_index = _tool_message_match_index(
tool_messages,
used_results,
call_id=call_id,
call_name=call_name,
)
# A historical client-side function call is only replayable together
# with its result. Incomplete calls are discarded instead of asking the
# provider to continue a broken tool turn.
if result_index is None:
continue
if not call_id:
result_id = str(
getattr(tool_messages[result_index], "tool_call_id", "") or ""
)
call_id = result_id or _stable_tool_call_id(
message, message_index, call_index
)
call["id"] = call_id
matched_calls.append(call)
matched_result_indexes.append(result_index)
original_to_matched_call[call_index] = (call_id, call_name)
used_results.add(result_index)
copied.tool_calls = matched_calls
if isinstance(copied.content, list):
original_call_index = 0
blocks: list[Any] = []
for original_block in copied.content:
if not isinstance(original_block, dict):
blocks.append(original_block)
continue
block = dict(original_block)
if block.get("type") in {"tool_call", "function_call"}:
matched_call = original_to_matched_call.get(original_call_index)
original_call_index += 1
if matched_call is None:
continue
call_id, call_name = matched_call
# LangChain content blocks use id; the Responses converter later
# maps it to call_id.
block["id"] = call_id
block["name"] = call_name
if isinstance(block.get("function"), dict):
block["function"] = {**block["function"], "name": call_name}
blocks.append(block)
copied.content = blocks
matched_results: list[Any] = []
result_to_call_id = {
result_index: matched_calls[index]["id"]
for index, result_index in enumerate(matched_result_indexes)
}
for result_index, result in enumerate(tool_messages):
call_id = result_to_call_id.get(result_index)
if call_id is None:
continue
copied_result = copy.copy(result)
copied_result.tool_call_id = call_id
matched_results.append(copied_result)
had_tool_protocol = bool(original_calls) or bool(
getattr(message, "invalid_tool_calls", None)
)
if not matched_calls and had_tool_protocol:
replayable_content = _flatten_message_content(copied.content)
if not replayable_content:
return None, matched_results
return copied, matched_results
def _sanitize_openai_tool_history(messages: list[Any]) -> list[Any]:
"""Copy history while retaining only complete, replayable tool turns."""
normalized: list[Any] = []
index = 0
while index < len(messages):
message = messages[index]
message_type = getattr(message, "type", None)
if message_type == "tool":
# A tool result without its immediately preceding assistant call is
# invalid for both Chat Completions and Responses APIs.
index += 1
continue
if message_type != "ai":
normalized.append(message)
index += 1
continue
next_index = index + 1
tool_messages: list[Any] = []
while (
next_index < len(messages)
and getattr(messages[next_index], "type", None) == "tool"
):
tool_messages.append(messages[next_index])
next_index += 1
copied, matched_results = _copy_ai_message_with_tool_pairs(
message,
index,
tool_messages,
)
if copied is not None:
normalized.append(copied)
normalized.extend(matched_results)
index = next_index
return normalized
def _ensure_openai_tool_call_ids(messages: list[Any]) -> list[Any]:
"""Backward-compatible alias for replay-safe tool history normalization."""
return _sanitize_openai_tool_history(messages)
def _has_assistant_tool_protocol(messages: list[Any]) -> bool:
"""Return whether history contains assistant-side tool protocol state."""
for message in messages:
if getattr(message, "type", None) != "ai":
continue
if getattr(message, "tool_calls", None) or getattr(
message, "invalid_tool_calls", None
):
return True
additional_kwargs = getattr(message, "additional_kwargs", None) or {}
if additional_kwargs.get("tool_calls"):
return True
content = getattr(message, "content", None)
if isinstance(content, list) and any(
isinstance(block, dict)
and block.get("type") in {"tool_call", "function_call"}
for block in content
):
return True
return False
def _validate_openai_tool_history(messages: list[Any]) -> None:
"""Raise when sanitized history still contains an invalid tool protocol."""
available_call_ids: set[str] = set()
for message in messages:
message_type = getattr(message, "type", None)
if message_type == "ai":
if getattr(message, "invalid_tool_calls", None):
raise ValueError("invalid_tool_calls must not be replayed")
response_call_ids: set[str] = set()
response_calls: dict[str, str] = {}
for call in getattr(message, "tool_calls", None) or []:
call_name = str(call.get("name") or "").strip()
if not call_name:
raise ValueError("assistant tool call is missing a name")
call_id = str(call.get("id") or "").strip()
if not call_id:
raise ValueError("assistant tool call is missing an id")
if call_id in response_call_ids or call_id in available_call_ids:
raise ValueError(
"assistant tool call id is duplicated while outstanding"
)
response_call_ids.add(call_id)
available_call_ids.add(call_id)
response_calls[call_id] = call_name
content = getattr(message, "content", None)
content_call_ids: set[str] = set()
if isinstance(content, list):
for block in content:
if not isinstance(block, dict) or block.get("type") not in {
"tool_call",
"function_call",
}:
continue
block_id = str(
block.get("id") or block.get("call_id") or ""
).strip()
block_name = block.get("name") or block.get("tool_name")
function = block.get("function")
if not block_name and isinstance(function, dict):
block_name = function.get("name")
block_name = str(block_name or "").strip()
if (
not block_id
or block_id in content_call_ids
or response_calls.get(block_id) != block_name
):
raise ValueError(
"assistant content block does not match parsed tool call"
)
content_call_ids.add(block_id)
elif message_type == "tool":
call_id = str(getattr(message, "tool_call_id", "") or "")
if not call_id or call_id not in available_call_ids:
raise ValueError("tool result does not match a prior tool call")
available_call_ids.remove(call_id)
if available_call_ids:
raise ValueError("assistant tool call is missing its tool result")
def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> list[Any]:
"""Flatten list content for OpenAI-compatible APIs, preserving media.
@@ -581,9 +282,6 @@ def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> li
from langchain_core.messages import HumanMessage
sanitize_tool_history = _has_assistant_tool_protocol(messages)
if sanitize_tool_history:
messages = _sanitize_openai_tool_history(messages)
out: list[Any] = []
pending_media: list[Any] = [] # media hoisted out of a run of tool messages
@@ -622,8 +320,6 @@ def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> li
msg.content = flat
out.append(msg)
_flush() # conversation may end with tool messages
if sanitize_tool_history:
_validate_openai_tool_history(out)
return out
@@ -1033,66 +729,6 @@ def _patch_openai_capture_reasoning_content() -> None:
_patch_openai_capture_reasoning_content()
# ---------------------------------------------------------------------------
# Patch (module-level): silence langgraph_api's OpenAPI schema-generation
# warnings for endpoints whose docstrings aren't valid YAML.
#
# Upstream ``langgraph_api.utils.SchemaGenerator.get_schema`` calls
# ``parse_docstring`` (inherited from Starlette's ``BaseSchemaGenerator``)
# on every registered endpoint. When the docstring is prose with stray
# ``:`` characters, ``yaml.safe_load`` raises and upstream logs the
# failure + full traceback at WARNING level. It then falls back to
# ``{"description": docstring}`` — the endpoint still ends up in the
# schema with its prose as the description, just without structured
# ``parameters``/``responses``/``tags`` fields.
#
# The fallback path is fine; the warning + traceback is just noise. And
# it's only triggered for our deploy because mounting any custom Starlette
# app (``EvoScientist/langgraph_dev/http.py``) makes upstream call
# ``update_openapi_spec`` at startup — which iterates EVERY route,
# including upstream's own endpoints whose prose docstrings predate the
# YAML convention.
#
# Fix: wrap ``parse_docstring`` itself and absorb ``yaml.YAMLError`` by
# returning the same fallback shape upstream's except branch produces.
# Non-YAML exceptions are deliberately left to propagate — upstream's
# ``get_schema`` already catches them and logs WARNING + traceback, so
# unexpected failures remain debuggable. Patching ``parse_docstring`` (a
# small, stable method) instead of ``get_schema`` (the larger loop body)
# minimizes our exposure to upstream churn.
# ---------------------------------------------------------------------------
_langgraph_schema_silenced_patched = False
def _patch_langgraph_schema_generator_silence_warnings() -> None:
global _langgraph_schema_silenced_patched
if _langgraph_schema_silenced_patched:
return
try:
import langgraph_api.utils as _lgapi_utils
import yaml
_SchemaGenerator = _lgapi_utils.SchemaGenerator
_orig_parse_docstring = _SchemaGenerator.parse_docstring
def _patched_parse_docstring(self: Any, func: Any) -> dict[str, Any]:
try:
return _orig_parse_docstring(self, func)
except yaml.YAMLError:
return {"description": getattr(func, "__doc__", None) or ""}
_SchemaGenerator.parse_docstring = _patched_parse_docstring
_langgraph_schema_silenced_patched = True
except Exception:
# Patches are loader-safe: never crash the import. Silent failure
# here just leaves the upstream warnings visible in deploy logs,
# which is a benign fallback.
pass
_patch_langgraph_schema_generator_silence_warnings()
# ---------------------------------------------------------------------------
# Patch (lazy, OpenRouter only): strip OpenAI-Responses encrypted reasoning
# items from outgoing assistant messages.
-302
View File
@@ -1,302 +0,0 @@
"""Shared logging configuration helpers."""
from __future__ import annotations
import logging
import os
import sys
from datetime import UTC, datetime
from pathlib import Path
from typing import Any, TextIO
DEFAULT_LOG_RETENTION_DAYS = 30
DEFAULT_LOG_FORMAT = "%(asctime)s [%(levelname)s] %(name)s: %(message)s"
DEFAULT_LOG_DATE_FORMAT = "%Y-%m-%d %H:%M:%S"
MANAGED_HANDLER_ATTR = "_evoscientist_managed_handler"
def resolve_log_level(level: int | str | None, default: int = logging.INFO) -> int:
"""Resolve a logging level from config or environment input."""
if isinstance(level, int):
return level
raw = str(level or "").strip()
if not raw:
return default
if raw.isdigit():
return int(raw)
normalized = raw.upper()
if normalized == "WARN":
normalized = "WARNING"
resolved = logging.getLevelNamesMapping().get(normalized)
return resolved if isinstance(resolved, int) else default
def _mark_managed(handler: logging.Handler, kind: str) -> logging.Handler:
setattr(handler, MANAGED_HANDLER_ATTR, kind)
return handler
def _managed_kind(handler: logging.Handler) -> str | None:
kind = getattr(handler, MANAGED_HANDLER_ATTR, None)
return kind if isinstance(kind, str) else None
def remove_managed_handlers(
logger: logging.Logger | None = None,
*,
kinds: set[str] | None = None,
) -> None:
"""Remove handlers installed by this module without touching external ones."""
target = logger or logging.getLogger()
for handler in target.handlers[:]:
kind = _managed_kind(handler)
if kind and (kinds is None or kind in kinds):
target.removeHandler(handler)
handler.close()
def _standard_formatter() -> logging.Formatter:
return logging.Formatter(DEFAULT_LOG_FORMAT, datefmt=DEFAULT_LOG_DATE_FORMAT)
class DailyLogFileHandler(logging.FileHandler):
"""File handler that writes the active log to a date-based filename."""
def __init__(
self,
log_dir: str | Path,
*,
prefix: str = "evoscientist",
retention_days: int = DEFAULT_LOG_RETENTION_DAYS,
encoding: str = "utf-8",
utc: bool = False,
) -> None:
self.log_dir = Path(log_dir).expanduser()
self.prefix = prefix
self.retention_days = max(1, retention_days)
self.utc = utc
self.log_dir.mkdir(parents=True, exist_ok=True)
super().__init__(self._dated_log_path(), encoding=encoding, delay=True)
@property
def active_log_path(self) -> Path:
"""Return the active log path for the current date."""
return self._dated_log_path()
def _dated_log_path(self) -> Path:
now = datetime.now(UTC if self.utc else None)
return self.log_dir / f"{self.prefix}-{now:%Y-%m-%d}.log"
def emit(self, record: logging.LogRecord) -> None:
try:
expected = str(self.active_log_path)
if self.baseFilename != expected:
if self.stream:
self.stream.close()
self.stream = None
self.baseFilename = expected
self._delete_expired_logs()
super().emit(record)
except OSError:
self.handleError(record)
def getFilesToDelete(self) -> list[str]:
candidates = sorted(self.log_dir.glob(f"{self.prefix}-????-??-??.log"))
if len(candidates) <= self.retention_days:
return []
return [str(path) for path in candidates[: -self.retention_days]]
def _delete_expired_logs(self) -> None:
for path in self.getFilesToDelete():
try:
os.remove(path)
except OSError:
pass
def default_log_dir() -> Path:
"""Return the default runtime log directory."""
env_dir = os.environ.get("EVOSCIENTIST_LOG_DIR")
if env_dir:
return Path(env_dir).expanduser()
from EvoScientist.paths import DATA_DIR
return DATA_DIR / "logs"
def configure_daily_file_logging(
logger: logging.Logger | None = None,
*,
log_dir: str | Path | None = None,
prefix: str = "evoscientist",
level: int | str = logging.INFO,
retention_days: int = DEFAULT_LOG_RETENTION_DAYS,
) -> DailyLogFileHandler:
"""Attach a daily file handler, replacing older matching handlers."""
target = logger or logging.getLogger()
resolved_level = resolve_log_level(level, default=logging.INFO)
retention_days = max(1, int(retention_days))
resolved_dir = Path(log_dir).expanduser() if log_dir else default_log_dir()
for handler in target.handlers[:]:
if (
isinstance(handler, DailyLogFileHandler)
and handler.prefix == prefix
and handler.log_dir == resolved_dir
):
target.removeHandler(handler)
handler.close()
handler = DailyLogFileHandler(
resolved_dir,
prefix=prefix,
retention_days=retention_days,
)
_mark_managed(handler, "file")
handler.setLevel(resolved_level)
handler.setFormatter(_standard_formatter())
target.addHandler(handler)
if target.level == logging.NOTSET or target.level > resolved_level:
target.setLevel(resolved_level)
return handler
def configure_console_logging(
logger: logging.Logger | None = None,
*,
level: int | str | None = logging.INFO,
stream: TextIO | None = None,
replace: bool = True,
) -> logging.StreamHandler:
"""Attach a standard console handler for non-interactive entry points."""
target = logger or logging.getLogger()
resolved_level = resolve_log_level(level, default=logging.INFO)
if replace:
remove_managed_handlers(target, kinds={"console", "rich"})
handler = logging.StreamHandler(stream or sys.stderr)
_mark_managed(handler, "console")
handler.setLevel(resolved_level)
handler.setFormatter(_standard_formatter())
target.addHandler(handler)
target.setLevel(resolved_level)
return handler
def configure_rich_console_logging(
logger: logging.Logger | None = None,
*,
level: int | str | None = logging.INFO,
console: Any = None,
replace: bool = True,
dim_warnings: bool = False,
show_time: bool | None = None,
show_path: bool | None = None,
show_level: bool | None = None,
) -> logging.Handler:
"""Attach a Rich console handler for interactive CLI output."""
from rich.logging import RichHandler
from rich.markup import escape
target = logger or logging.getLogger()
resolved_level = resolve_log_level(level, default=logging.INFO)
verbose = resolved_level <= logging.DEBUG
if replace:
remove_managed_handlers(target, kinds={"console", "rich"})
class DimWarningHandler(RichHandler):
def emit(self, record: logging.LogRecord) -> None:
if dim_warnings and record.levelno == logging.WARNING and console is not None:
console.print(
"[dim yellow]\u26a0\ufe0f Warning:[/dim yellow] "
f"[dim]{escape(record.getMessage())}[/dim]"
)
return
super().emit(record)
handler = DimWarningHandler(
console=console,
show_time=verbose if show_time is None else show_time,
show_path=verbose if show_path is None else show_path,
show_level=verbose if show_level is None else show_level,
)
_mark_managed(handler, "rich")
handler.setLevel(resolved_level)
target.addHandler(handler)
target.setLevel(resolved_level)
return handler
def configure_logging(
logger: logging.Logger | None = None,
*,
level: int | str | None = logging.INFO,
log_dir: str | Path | None = None,
retention_days: int = DEFAULT_LOG_RETENTION_DAYS,
prefix: str = "evoscientist",
console: bool = True,
file: bool = True,
replace_managed: bool = True,
) -> list[logging.Handler]:
"""Configure standard EvoScientist console and daily file logging."""
target = logger or logging.getLogger()
resolved_level = resolve_log_level(level, default=logging.INFO)
if replace_managed:
remove_managed_handlers(target, kinds={"console", "rich", "file"})
handlers: list[logging.Handler] = []
if console:
handlers.append(
configure_console_logging(target, level=resolved_level, replace=False)
)
if file:
handlers.append(
configure_daily_file_logging(
target,
log_dir=log_dir,
prefix=prefix,
level=resolved_level,
retention_days=retention_days,
)
)
target.setLevel(resolved_level)
return handlers
def configure_logging_from_settings(
logger: logging.Logger | None = None,
*,
default_level: int = logging.INFO,
prefix: str = "evoscientist",
console: bool = True,
file: bool = True,
) -> list[logging.Handler]:
"""Configure logging from EvoScientist settings and environment overrides."""
level: int | str | None = os.environ.get("EVOSCIENTIST_LOG_LEVEL")
log_dir: str | Path | None = os.environ.get("EVOSCIENTIST_LOG_DIR") or None
retention_days = int(
os.environ.get("EVOSCIENTIST_LOG_RETENTION_DAYS", DEFAULT_LOG_RETENTION_DAYS)
)
try:
from EvoScientist.config import get_effective_config
cfg = get_effective_config()
level = level or getattr(cfg, "log_level", None)
log_dir = log_dir or getattr(cfg, "log_dir", None) or None
retention_days = int(
getattr(cfg, "log_retention_days", DEFAULT_LOG_RETENTION_DAYS)
)
except Exception:
level = level or default_level
return configure_logging(
logger,
level=resolve_log_level(level, default=default_level),
log_dir=log_dir,
retention_days=retention_days,
prefix=prefix,
console=console,
file=file,
)
-2
View File
@@ -10,7 +10,6 @@ from .client import (
build_mcp_add_kwargs,
build_mcp_edit_fields,
edit_mcp_server,
get_mcp_server_errors,
load_mcp_config,
load_mcp_tools,
parse_mcp_add_args,
@@ -39,7 +38,6 @@ __all__ = [
"find_server_by_name",
"get_all_tags",
"get_installed_names",
"get_mcp_server_errors",
"install_mcp_server",
"install_mcp_servers",
"load_mcp_config",
+1 -16
View File
@@ -114,10 +114,6 @@ _URL_TRANSPORTS = {"http", "streamable_http", "sse", "websocket"}
# still parallelizing the common 3–7 server case to completion.
_MAX_CONCURRENT_CONNECTIONS = 8
# Last connection error per configured server. This is process-local runtime
# diagnostics for the Web/CLI status surfaces, not persisted configuration.
_MCP_SERVER_ERRORS: dict[str, str] = {}
# Env vars forwarded to stdio MCP subprocesses on top of the MCP SDK's
# minimal default set (HOME/PATH/USER/…). Without this, servers behind
# a proxy or with a custom CA bundle silently fail with long timeouts.
@@ -768,9 +764,6 @@ async def _load_tools(
if not connections:
return {}
for stale_name in set(_MCP_SERVER_ERRORS) - set(connections):
_MCP_SERVER_ERRORS.pop(stale_name, None)
client = MultiServerMCPClient(connections) # type: ignore[invalid-argument-type]
def _report(event: str, name: str, detail: str = "") -> None:
@@ -794,13 +787,10 @@ async def _load_tools(
_report("start", name)
try:
tools = await client.get_tools(server_name=name)
_MCP_SERVER_ERRORS.pop(name, None)
logger.info("MCP server %r: loaded %d tool(s)", name, len(tools))
_report("success", name, str(len(tools)))
return name, tools
except Exception as exc:
detail = str(exc) or type(exc).__name__
_MCP_SERVER_ERRORS[name] = detail
# When the caller wired up ``on_progress`` they own the
# user-facing display; downgrade the logger so we don't
# double-print.
@@ -808,7 +798,7 @@ async def _load_tools(
logger.warning("MCP server %r: failed to load tools: %s", name, exc)
else:
logger.debug("MCP server %r: failed to load tools: %s", name, exc)
_report("error", name, detail)
_report("error", name, str(exc))
return name, []
# ``return_exceptions=False`` is fine because ``_fetch`` already
@@ -817,11 +807,6 @@ async def _load_tools(
return dict(results)
def get_mcp_server_errors() -> dict[str, str]:
"""Return a snapshot of the most recent per-server connection errors."""
return dict(_MCP_SERVER_ERRORS)
async def aload_mcp_tools(
config: dict[str, Any] | None = None,
*,
+11 -17
View File
@@ -428,7 +428,6 @@ def _memory_worker_middleware(
enable_observation_memory: bool = True,
):
"""Build middleware for memory workers, excluding task execution tools."""
from ...middleware.error_normalization import ErrorNormalizationMiddleware
from ...middleware.memory import create_memory_middleware
memory_controls = MemoryControls(
@@ -440,23 +439,18 @@ def _memory_worker_middleware(
enable_observation_tool = memory_controls.observation_tool_enabled(
_memory_worker_observation_target(source_type)
)
return [
# Outermost — normalize provider-SDK exceptions from the
# auxiliary model call before any inner middleware sees them.
ErrorNormalizationMiddleware(),
*memory_agent_middleware(
create_memory_middleware(
str(memory_dir),
workspace_dir=workspace_dir,
source_type=source_type,
source_agent=_memory_worker_agent_name(source_type),
enable_profile_memory=enable_profile_memory,
enable_observation_memory=enable_observation_memory,
enable_observation_tool=enable_observation_tool,
),
excluded_tools=_MEMORY_WORKER_EXCLUDED_TOOLS,
return memory_agent_middleware(
create_memory_middleware(
str(memory_dir),
workspace_dir=workspace_dir,
source_type=source_type,
source_agent=_memory_worker_agent_name(source_type),
enable_profile_memory=enable_profile_memory,
enable_observation_memory=enable_observation_memory,
enable_observation_tool=enable_observation_tool,
),
]
excluded_tools=_MEMORY_WORKER_EXCLUDED_TOOLS,
)
def _build_memory_worker_agent(
@@ -71,8 +71,6 @@ def build_observation_linker_graph(
workspace_dir: str | Path | None = None,
) -> CompiledStateGraph:
"""Build the registered LangGraph observation linker."""
from ...middleware.error_normalization import ErrorNormalizationMiddleware
agent_paths = resolve_memory_agent_paths(
memory_dir=memory_dir,
workspace_dir=workspace_dir,
@@ -87,7 +85,5 @@ def build_observation_linker_graph(
tools=tools,
memory_dir=agent_paths.memory_dir,
workspace_dir=agent_paths.workspace_dir,
# Outermost — normalize provider-SDK exceptions from the
# auxiliary model call before any inner middleware sees them.
middleware=[ErrorNormalizationMiddleware(), *memory_agent_middleware()],
middleware=memory_agent_middleware(),
)
-14
View File
@@ -18,7 +18,6 @@ from .context_editing import (
create_context_editing_middleware,
)
from .context_overflow import ContextOverflowMapperMiddleware
from .error_normalization import ErrorNormalizationMiddleware
from .memory import (
EvoMemoryMiddleware,
create_memory_middleware,
@@ -29,42 +28,29 @@ from .memory_lifecycle import (
default_memory_scheduler,
)
from .model_fallback import ModelFallbackMiddleware, load_fallback_chain
from .repetitive_tool_guard import (
DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS,
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
RepetitiveToolCallGuardMiddleware,
collapse_repetitive_tool_rounds,
)
from .runtime_context import RuntimeContextMiddleware, create_runtime_context_middleware
from .scheduler import (
SchedulerMiddleware,
create_scheduler_middleware,
)
from .tool_error_handler import ToolErrorHandlerMiddleware
from .tool_protocol_guard import ToolProtocolGuardMiddleware
from .tool_selector import create_tool_selector_middleware
from .utils import disable_thinking
__all__ = [
"DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS",
"DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD",
"AskUserMiddleware",
"AskUserRequest",
"AskUserWidgetResult",
"Choice",
"ConfigurableModelMiddleware",
"ContextOverflowMapperMiddleware",
"ErrorNormalizationMiddleware",
"EvoMemoryLifecycleMiddleware",
"EvoMemoryMiddleware",
"ModelFallbackMiddleware",
"Question",
"RepetitiveToolCallGuardMiddleware",
"RuntimeContextMiddleware",
"SchedulerMiddleware",
"ToolErrorHandlerMiddleware",
"ToolProtocolGuardMiddleware",
"collapse_repetitive_tool_rounds",
"compute_context_editing_trigger",
"create_code_interpreter_middleware",
"create_context_editing_middleware",
+1 -15
View File
@@ -45,21 +45,7 @@ _MEMORY_FIRST_INTERPRETER_PROMPT = (
class EvoCodeInterpreterMiddleware(CodeInterpreterMiddleware):
"""Code interpreter middleware with EvoScientist's memory preflight hint.
``after_agent`` / ``aafter_agent`` are intentionally NOT overridden. An
earlier "conditional snapshot" gate that skipped ``after_agent`` on turns
where ``code_interpreter`` wasn't called saved ~50 ms/turn of
``create_snapshot()`` work, but also skipped the slot eviction upstream
performs in the same hook (``finally: self._registry.evict(thread_id)``
in ``langchain_quickjs.middleware.CodeInterpreterMiddleware.after_agent``).
``before_agent`` restores the REPL on every turn that follows a touched
one via ``self._registry.get(thread_id)`` (get-or-create), so skipping
eviction leaked one ``ThreadWorker`` + QuickJS Runtime per persistent
``thread_id`` that ever went touched → quiet. The regression test
``test_after_agent_evicts_slot_on_untouched_turn`` guards against
reintroducing the gate.
"""
"""Code interpreter middleware with EvoScientist's memory preflight hint."""
def _prepare_for_call(self, request: ModelRequest) -> str:
return super()._prepare_for_call(request) + _MEMORY_FIRST_INTERPRETER_PROMPT
@@ -1,240 +0,0 @@
"""ErrorNormalizationMiddleware — catch provider-SDK exceptions at the
model boundary and re-raise as a normalized non-dataclass wrapper.
Some provider SDKs (openrouter.errors.* today) decorate their exception
classes with ``@dataclass``. When langgraph_api emits an SSE error
frame via ``json_dumpb`` → ``orjson.dumps(obj, default=default,
option=OPT_SERIALIZE_DATACLASS)``, orjson's dataclass fast-path
enumerates the fields directly and skips the ``default=`` hook that
builds our envelope. The wire payload comes out as
``{"message": …, "status_code": …, "body": …, "headers": null,
"raw_response": null, "data": {…}}`` with no ``error`` / ``class`` /
``provider`` envelope and no way for the WebUI to distinguish quota /
auth / rate-limit / model-not-found.
This middleware sits at the model-call boundary. It catches
``BaseException`` from ``handler()``, and if ``request.model`` is a
recognized provider SDK client, wraps the exception in a
:class:`~EvoScientist.llm.errors.ProviderStreamError` (a plain
``Exception`` subclass, not a dataclass). The wrapper carries the SSE
envelope pre-baked on its instance attributes.
Contract: the wrap decision is based on the **model**, not the
exception, after platform and graph control signals have been excluded.
Provider SDK exceptions, httpx errors, langchain-wrapper failures, and
even builtins like ``RuntimeError`` get wrapped for a recognized model.
At the middleware boundary we can tell which provider was in use, but
not the exception's precise origin; a uniform envelope is more useful
to the WebUI than gambling on the exception class. If the model isn't
from a recognized provider, or the request carries no ``.model``, the
exception re-raises unchanged and upstream's whitelist / catch-all
behavior takes over.
"""
from __future__ import annotations
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING
from langchain.agents.middleware.types import (
AgentMiddleware,
ModelRequest,
ModelResponse,
)
if TYPE_CHECKING:
from ..llm.errors import ProviderStreamError
def _should_pass_through(exc: BaseException) -> bool:
"""True if *exc* is a LangGraph-level signal that must propagate
untouched — either a control-flow signal or a structural error
that isn't a provider failure.
Covers everything in ``langgraph.errors.*``:
- **Control flow** (breaking these would corrupt the interrupt /
resume protocol): ``GraphBubbleUp`` and its subclasses
``GraphInterrupt``, ``NodeInterrupt``, ``ParentCommand``,
``GraphDrained``.
- **Structural** (wrapping would mis-attribute a graph-level
issue as a provider failure): ``InvalidUpdateError``,
``EmptyInputError``, ``EmptyChannelError``, ``TaskNotFound``,
``GraphRecursionError``, ``NodeCancelledError``,
``NodeTimeoutError``.
Symmetric with upstream ``langgraph_api.serde.default``'s
whitelist, which also exposes these classes' ``str(exc)`` untouched
rather than swallowing them behind a provider envelope.
``KeyboardInterrupt``, ``SystemExit``, and ``asyncio.CancelledError``
are handled implicitly by catching ``Exception`` — they inherit
from ``BaseException``.
"""
return (type(exc).__module__ or "").startswith("langgraph.errors")
# Module prefixes for provider SDK exceptions. Consumed by
# ``_is_provider_error`` to decide whether an exception raised inside
# a model call should surface as a provider incident or gracefully
# degrade (used by ``_ConditionalToolSelectorMiddleware``).
#
# Related sibling: ``_HOST_TO_PROVIDER`` in ``llm/errors.py`` — the
# host-side allow-list. Adding a whole new provider SDK means updating
# both; adding a new routed provider (new base_url through an existing
# SDK) only touches ``_HOST_TO_PROVIDER``.
_PROVIDER_EXC_MODULE_PREFIXES: tuple[str, ...] = (
"openai",
"anthropic",
"google.genai",
"google.api_core",
"openrouter",
"langchain_openai",
"langchain_anthropic",
"langchain_google_genai",
"langchain_openrouter",
"httpx",
)
def _is_provider_error(exc: BaseException) -> bool:
"""True if *exc* looks like it originated inside a provider SDK
(openai, anthropic, google.genai, openrouter, httpx, or their
langchain wrappers), as opposed to a shape / config error (structured
output not supported, malformed schema, missing tool, …).
Used by callers that need to decide whether an exception from the
model call is worth surfacing to the user (provider errors) or
can be silently degraded around (shape errors). Cheap alternative
to inspecting ``status_code`` / ``request`` because some provider
errors — connection errors, timeouts — don't carry those attributes.
"""
module = type(exc).__module__ or ""
return any(module.startswith(p) for p in _PROVIDER_EXC_MODULE_PREFIXES)
def _normalize(request: ModelRequest, exc: BaseException) -> ProviderStreamError | None:
"""Return a :class:`ProviderStreamError` wrapping *exc* if the model
on *request* comes from a recognized provider SDK, or ``None`` if
the caller should re-raise *exc* unchanged.
Provider is read from ``request.model`` — the definitive config
the exception was raised under, not inferred from the exception
class / URL. Status / code / redaction still come from the raised
exception because those fields are populated by the SDK at raise
time.
Returns ``None`` (caller re-raises unchanged) for:
- Already-normalized wrappers (would double-attribute).
- LangGraph control-flow / structural errors — see
``_should_pass_through``. This gate lives here so every caller
of ``_normalize`` (not just the wrap sites of this middleware)
gets the protection automatically. Notably
``ModelFallbackMiddleware`` also calls ``_normalize`` at the
raise point of its fallback chain.
- ``ContextOverflowError`` — a cross-layer control signal that
deepagents' ``SummarizationMiddleware`` catches by type from
**outside** the user middleware stack to compress history and
retry. Wrapping it here would change the type and break that
self-healing fallback.
- ``AgentControlError`` — a platform-owned typed decision. Gateway route
fallback and canonical error mapping depend on its concrete type and
structured fields, so it must never become a provider incident.
- Models we don't recognize as a provider SDK.
"""
from langchain_core.exceptions import ContextOverflowError
from ..llm.errors import (
AgentControlError,
ProviderStreamError,
_extract_error_type,
_extract_provider_code,
_extract_status_code,
_provider_from_model,
_redact_api_keys,
)
# Already normalized (e.g. by ModelFallbackMiddleware wrapping against
# the actual failing model rather than the original request's model).
# Pass through — re-wrapping would double-attribute.
if isinstance(exc, ProviderStreamError):
return None
# Platform control errors are raised by inner middleware after the provider
# response has already been interpreted. Wrapping them would erase routing,
# retry and recovery semantics such as ModelToolProtocolError.fallbackable.
if isinstance(exc, AgentControlError):
return None
# LangGraph control-flow / structural signals must propagate
# untouched, regardless of which caller invoked us.
if _should_pass_through(exc):
return None
# SummarizationMiddleware sits outside our stack and catches this
# by exact type to trigger reactive history compression + retry.
if isinstance(exc, ContextOverflowError):
return None
provider = _provider_from_model(getattr(request, "model", None))
if provider is None:
return None
cls = type(exc)
mod = cls.__module__ or ""
class_qualname = f"{mod}.{cls.__qualname__}" if mod else cls.__qualname__
request_id_attr = getattr(exc, "request_id", None)
request_id = (
request_id_attr
if isinstance(request_id_attr, str) and request_id_attr
else None
)
return ProviderStreamError(
provider=provider,
class_qualname=class_qualname,
message=_redact_api_keys(str(exc)),
status_code=_extract_status_code(exc),
code=_extract_provider_code(exc),
err_type=_extract_error_type(exc),
request_id=request_id,
)
class ErrorNormalizationMiddleware(AgentMiddleware):
"""Wrap the model call in try/except and normalize provider SDK
exceptions into a non-dataclass envelope wrapper.
Place this middleware **outermost** in the chain (first in the
middleware list) so it catches exceptions raised by inner
middlewares as well as the model handler itself.
"""
name = "error_normalization"
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelResponse:
try:
return handler(request)
except Exception as exc:
normalized = _normalize(request, exc)
if normalized is None:
raise
raise normalized from exc
async def awrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse:
try:
return await handler(request)
except Exception as exc:
normalized = _normalize(request, exc)
if normalized is None:
raise
raise normalized from exc
+3 -34
View File
@@ -48,8 +48,6 @@ _MALFORMED_REQUEST_PATTERNS: list[str] = [
"invalid_request_error",
"invalid request",
"malformed",
"repetitive tool calls",
"identical name and arguments",
]
"""Substrings that identify a malformed request (client-side bug)."""
@@ -217,9 +215,6 @@ def _is_non_fallbackable(exc: Exception) -> str | None:
"""
from langchain_core.exceptions import ContextOverflowError
if getattr(exc, "non_fallbackable", False):
return f"platform control error: {getattr(exc, 'code', type(exc).__name__)}"
if isinstance(exc, ContextOverflowError):
return "context length exceeded"
@@ -268,15 +263,7 @@ async def _try_fallbacks(
"Primary model failed: %s: %s", type(primary_exc).__name__, primary_exc
)
# Track the request whose model actually raised ``last_exc`` so we
# can attribute the exception to the failing model, not the
# original ``request.model``. Without this, a fallback chain
# ``deepseek → moonshot`` where moonshot exhausts its quota would
# surface as ``provider: deepseek`` — the model the user never
# actually saw fail.
last_exc = primary_exc
last_failing_request = request
for model_name, provider in get_fallback_chain():
_emit(
f" -> Falling back to {model_name} ({provider}) "
@@ -301,9 +288,8 @@ async def _try_fallbacks(
f"-- aborting fallback chain",
style="red",
)
_raise_normalized(fb_request, fb_exc)
raise
last_exc = fb_exc
last_failing_request = fb_request
_emit(
f" x {model_name} also failed: {type(fb_exc).__name__}: {fb_exc}",
style="red",
@@ -317,24 +303,7 @@ async def _try_fallbacks(
)
_emit(" All fallbacks exhausted -- re-raising last error", style="red")
_raise_normalized(last_failing_request, last_exc)
def _raise_normalized(request: ModelRequest, exc: Exception) -> None:
"""Wrap *exc* in a ``ProviderStreamError`` attributed to
``request.model`` and raise, so the outer chain sees the failure
tagged with the model that actually raised.
Falls back to a plain ``raise`` when the model isn't from a
recognized provider (``_normalize`` returns None) — nothing useful
to add.
"""
from .error_normalization import _normalize
normalized = _normalize(request, exc)
if normalized is not None:
raise normalized from exc
raise exc
raise last_exc
def _guard_and_fallback(
@@ -361,7 +330,7 @@ def _guard_and_fallback(
f"Model error ({reason}) -- not eligible for fallback, re-raising",
style="red",
)
_raise_normalized(request, primary_exc)
raise primary_exc
return _try_fallbacks(request, invoke, primary_exc)
@@ -1,350 +0,0 @@
"""Detect deterministic tool loops and compact only provider-facing history."""
from __future__ import annotations
import json
import logging
import re
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from typing import Any
from langchain.agents.middleware.types import (
AgentMiddleware,
ModelRequest,
ModelResponse,
)
from ..llm.errors import AgentControlError
logger = logging.getLogger(__name__)
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD = 2
DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS = 3
_TRANSIENT_PATTERNS = (
"timeout",
"timed out",
"cancelled",
"canceled",
"connection",
"rate limit",
"too many requests",
"temporarily unavailable",
"service unavailable",
"overloaded",
"bad gateway",
"gateway timeout",
"http 500",
"http 502",
"http 503",
"http 504",
)
_DETERMINISTIC_PATTERNS: tuple[tuple[str, tuple[str, ...]], ...] = (
(
"INVALID_ARGUMENTS",
("invalid argument", "validation error", "schema", "bad input"),
),
("UNKNOWN_TOOL", ("not a valid tool", "unknown tool", "tool not found")),
("UNSUPPORTED", ("not supported", "unsupported", "not implemented")),
(
"POLICY_DENIED",
("permission denied", "forbidden", "policy denied", "not allowed"),
),
)
_SAFE_CODE_RE = re.compile(r"^[A-Za-z][A-Za-z0-9_.:-]{0,95}$")
_DETERMINISTIC_CODE_MARKERS = (
"INVALID",
"VALIDATION",
"SCHEMA",
"UNKNOWN_TOOL",
"NOT_FOUND",
"UNSUPPORTED",
"NOT_IMPLEMENTED",
"POLICY",
"PERMISSION",
"FORBIDDEN",
"DENIED",
)
@dataclass(frozen=True, slots=True)
class RepetitiveToolHistoryRepair:
messages: list[Any]
blocked_tool_names: frozenset[str]
removed_rounds: int
tail_repetitions: int = 0
tail_consecutive_errors: int = 0
@dataclass(frozen=True, slots=True)
class _ToolRound:
messages: tuple[Any, ...]
signature: tuple[tuple[str, str, str], ...]
tool_names: frozenset[str]
deterministic_error: bool
def _canonical_tool_args(value: Any) -> str:
if isinstance(value, str):
try:
value = json.loads(value)
except json.JSONDecodeError:
return value.strip()
try:
return json.dumps(
value,
ensure_ascii=False,
sort_keys=True,
separators=(",", ":"),
default=str,
)
except (TypeError, ValueError):
return repr(value)
def _deterministic_result_code(message: Any) -> str | None:
additional = getattr(message, "additional_kwargs", None)
additional = additional if isinstance(additional, Mapping) else {}
raw_code = additional.get("error_code") or additional.get("code")
status = str(getattr(message, "status", "") or "").lower()
content = str(getattr(message, "content", "") or "")
lowered = content.lower()
if any(pattern in lowered for pattern in _TRANSIENT_PATTERNS):
return None
if isinstance(raw_code, str) and _SAFE_CODE_RE.fullmatch(raw_code.strip()):
normalized = raw_code.strip().upper()
if any(
pattern.replace(" ", "_") in normalized for pattern in _TRANSIENT_PATTERNS
):
return None
if any(marker in normalized for marker in _DETERMINISTIC_CODE_MARKERS):
return normalized
return None
is_error = status == "error" or lowered.startswith("error:")
if not is_error:
return None
for code, patterns in _DETERMINISTIC_PATTERNS:
if any(pattern in lowered for pattern in patterns):
return code
return None
def _parse_tool_round(
messages: Sequence[Any], start: int
) -> tuple[_ToolRound, int] | None:
assistant = messages[start]
if getattr(assistant, "type", None) != "ai":
return None
raw_calls = list(getattr(assistant, "tool_calls", None) or [])
calls = [call for call in raw_calls if isinstance(call, Mapping)]
if not calls or len(calls) != len(raw_calls):
return None
end = start + 1
results: list[Any] = []
while end < len(messages) and getattr(messages[end], "type", None) == "tool":
results.append(messages[end])
end += 1
if not results:
return None
results_by_id = {
str(getattr(result, "tool_call_id", "") or "").strip(): result
for result in results
if str(getattr(result, "tool_call_id", "") or "").strip()
}
signature: list[tuple[str, str, str]] = []
tool_names: set[str] = set()
for index, call in enumerate(calls):
name = str(call.get("name") or "").strip()
call_id = str(call.get("id") or "").strip()
if not name or not call_id:
return None
result = results_by_id.get(call_id)
if result is None and index < len(results):
candidate = results[index]
if not str(getattr(candidate, "tool_call_id", "") or "").strip():
result = candidate
if result is None:
return None
result_code = _deterministic_result_code(result)
if result_code is None:
return _ToolRound(
messages=(assistant, *results),
signature=(),
tool_names=frozenset(),
deterministic_error=False,
), end
signature.append((name, _canonical_tool_args(call.get("args")), result_code))
tool_names.add(name)
return (
_ToolRound(
messages=(assistant, *results),
signature=tuple(signature),
tool_names=frozenset(tool_names),
deterministic_error=True,
),
end,
)
def collapse_repetitive_tool_rounds(
messages: Sequence[Any],
*,
threshold: int = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
) -> RepetitiveToolHistoryRepair:
"""Build a provider-only projection while preserving audit history.
Only the middle rounds of three-or-more identical deterministic error
groups are omitted. The first and last observations remain, and callers
must never persist this projection back to a checkpoint.
"""
if not isinstance(threshold, int) or isinstance(threshold, bool) or threshold < 0:
raise ValueError("repetitive tool call threshold must be non-negative")
original = list(messages)
segments: list[Any | _ToolRound] = []
index = 0
while index < len(original):
parsed = _parse_tool_round(original, index)
if parsed is None:
segments.append(original[index])
index += 1
continue
tool_round, index = parsed
segments.append(tool_round)
tail_repetitions = 0
tail_consecutive_errors = 0
if segments and isinstance(segments[-1], _ToolRound):
tail = segments[-1]
if tail.deterministic_error:
cursor = len(segments) - 1
while cursor >= 0 and isinstance(segments[cursor], _ToolRound):
current = segments[cursor]
if not current.deterministic_error:
break
tail_consecutive_errors += 1
cursor -= 1
cursor = len(segments) - 1
while cursor >= 0 and isinstance(segments[cursor], _ToolRound):
current = segments[cursor]
if (
not current.deterministic_error
or current.signature != tail.signature
):
break
tail_repetitions += 1
cursor -= 1
projected: list[Any] = []
removed_rounds = 0
index = 0
while index < len(segments):
segment = segments[index]
if not isinstance(segment, _ToolRound) or not segment.deterministic_error:
if isinstance(segment, _ToolRound):
projected.extend(segment.messages)
else:
projected.append(segment)
index += 1
continue
end = index + 1
while (
end < len(segments)
and isinstance(segments[end], _ToolRound)
and segments[end].deterministic_error
and segments[end].signature == segment.signature
):
end += 1
group = segments[index:end]
should_compact = threshold > 0 and len(group) >= threshold and len(group) > 2
if should_compact:
projected.extend(group[0].messages)
projected.extend(group[-1].messages)
removed_rounds += len(group) - 2
else:
for item in group:
projected.extend(item.messages)
index = end
blocked = (
segments[-1].tool_names
if tail_repetitions and isinstance(segments[-1], _ToolRound)
else frozenset()
)
return RepetitiveToolHistoryRepair(
messages=projected,
blocked_tool_names=blocked,
removed_rounds=removed_rounds,
tail_repetitions=tail_repetitions,
tail_consecutive_errors=tail_consecutive_errors,
)
class RepetitiveToolCallGuardMiddleware(AgentMiddleware):
"""Stop deterministic loops before another model request is made."""
name = "repetitive_tool_call_guard"
def __init__(
self,
*,
threshold: int = DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
max_consecutive_errors: int = DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS,
) -> None:
super().__init__()
for name, value in {
"threshold": threshold,
"max_consecutive_errors": max_consecutive_errors,
}.items():
if not isinstance(value, int) or isinstance(value, bool) or value < 0:
raise ValueError(f"{name} must be a non-negative integer")
self.threshold = threshold
self.max_consecutive_errors = max_consecutive_errors
def _prepare_request(self, request: ModelRequest) -> ModelRequest:
repair = collapse_repetitive_tool_rounds(
request.messages,
threshold=self.threshold,
)
if self.threshold and repair.tail_repetitions >= self.threshold:
raise AgentControlError(
"MODEL_TOOL_LOOP_DETECTED",
"A deterministic repeated tool-call loop was stopped.",
status_code=422,
retryable=False,
)
if (
self.max_consecutive_errors
and repair.tail_consecutive_errors >= self.max_consecutive_errors
):
raise AgentControlError(
"MODEL_TOOL_ERROR_LIMIT",
"Too many consecutive deterministic tool errors were stopped.",
status_code=422,
retryable=False,
)
if repair.removed_rounds:
logger.info(
"Compacted deterministic tool errors for provider projection: removed_rounds=%d",
repair.removed_rounds,
)
return request.override(messages=repair.messages)
return request
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelResponse:
return handler(self._prepare_request(request))
async def awrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse:
return await handler(self._prepare_request(request))
@@ -1,361 +0,0 @@
"""Validate completed model tool calls before they can reach ToolNode."""
from __future__ import annotations
import hashlib
import json
from collections.abc import Awaitable, Callable, Mapping, Sequence
from typing import Any
from langchain.agents.middleware.types import (
AgentMiddleware,
ExtendedModelResponse,
ModelRequest,
ModelResponse,
)
from langchain_core.messages import AIMessage
from langchain_core.tools import BaseTool
from ..llm.errors import ModelToolProtocolError, _provider_from_model
_TOOL_BLOCK_TYPES = frozenset({"tool_call", "function_call"})
_MAX_DIAGNOSTIC_KEYS = 16
_MAX_DIAGNOSTIC_KEY_CHARS = 64
def _tool_name(tool: BaseTool | Mapping[str, Any] | Any) -> str | None:
if isinstance(tool, BaseTool):
return tool.name.strip() or None
if isinstance(tool, Mapping):
value = tool.get("name")
if not value and isinstance(tool.get("function"), Mapping):
value = tool["function"].get("name")
if isinstance(value, str) and value.strip():
return value.strip()
return None
value = getattr(tool, "name", None)
return value.strip() if isinstance(value, str) and value.strip() else None
def _ai_messages(response: Any) -> list[AIMessage]:
"""Extract final AI messages from every LangChain middleware response shape."""
if isinstance(response, AIMessage):
return [response]
if isinstance(response, ExtendedModelResponse):
response = response.model_response
elif not isinstance(response, ModelResponse):
nested = getattr(response, "model_response", None)
if nested is not None:
response = nested
result = getattr(response, "result", None)
if not isinstance(result, Sequence) or isinstance(result, str | bytes):
return []
return [message for message in result if isinstance(message, AIMessage)]
def _block_identity(block: Mapping[str, Any]) -> tuple[str, str]:
call_id = str(block.get("id") or block.get("call_id") or "").strip()
name = block.get("name") or block.get("tool_name")
function = block.get("function")
if not name and isinstance(function, Mapping):
name = function.get("name")
return call_id, str(name or "").strip()
def _value_digest(value: Any) -> str:
try:
encoded = json.dumps(
value,
ensure_ascii=False,
sort_keys=True,
separators=(",", ":"),
default=lambda item: f"<{type(item).__name__}>",
).encode("utf-8")
except (TypeError, ValueError):
encoded = f"<{type(value).__name__}:unserializable>".encode()
return "sha256:" + hashlib.sha256(encoded).hexdigest()[:16]
def _argument_diagnostic(value: Any, *, present: bool) -> dict[str, Any]:
if not present:
return {"args_present": False, "args_type": "missing"}
if isinstance(value, Mapping):
keys = sorted(str(key)[:_MAX_DIAGNOSTIC_KEY_CHARS] for key in value)
return {
"args_present": True,
"args_type": "object",
"args_key_count": len(keys),
"args_keys": keys[:_MAX_DIAGNOSTIC_KEYS],
"args_keys_truncated": len(keys) > _MAX_DIAGNOSTIC_KEYS,
"args_digest": _value_digest(value),
}
if isinstance(value, Sequence) and not isinstance(value, str | bytes):
value_type = "array"
elif isinstance(value, str):
value_type = "string"
elif value is None:
value_type = "null"
else:
value_type = type(value).__name__
return {
"args_present": True,
"args_type": value_type,
"args_digest": _value_digest(value),
}
def _summarize_call(call: Any) -> dict[str, Any]:
if not isinstance(call, Mapping):
return {"call_type": type(call).__name__}
function = call.get("function")
function = function if isinstance(function, Mapping) else {}
call_id = str(call.get("id") or call.get("call_id") or "").strip()
name = call.get("name") or call.get("tool_name") or function.get("name")
name = str(name or "").strip()
if "args" in call:
args = call.get("args")
args_present = True
elif "arguments" in call:
args = call.get("arguments")
args_present = True
elif "arguments" in function:
args = function.get("arguments")
args_present = True
else:
args = None
args_present = False
summary = {
"call_type": "object",
"name": name or "<missing>",
"id_present": bool(call_id),
**_argument_diagnostic(args, present=args_present),
}
if call_id:
summary["id_fingerprint"] = _value_digest(call_id)
return summary
def _raw_openai_call(message: AIMessage, call_index: int) -> Any | None:
additional = getattr(message, "additional_kwargs", None)
additional = additional if isinstance(additional, Mapping) else {}
raw_calls = additional.get("tool_calls")
if (
isinstance(raw_calls, Sequence)
and not isinstance(raw_calls, str | bytes)
and call_index < len(raw_calls)
):
return raw_calls[call_index]
return None
def _call_diagnostic(
message: AIMessage,
call: Any,
*,
source: str,
call_index: int,
call_count: int,
) -> dict[str, Any]:
diagnostic = {
"source": source,
"call_index": call_index,
"call_count": call_count,
**_summarize_call(call),
}
raw_call = _raw_openai_call(message, call_index)
diagnostic["raw_openai_call_available"] = raw_call is not None
if raw_call is not None:
diagnostic["raw_openai_call"] = _summarize_call(raw_call)
return diagnostic
def _route_metadata(request: ModelRequest) -> dict[str, Any]:
model = request.model
metadata = getattr(model, "metadata", None)
metadata = metadata if isinstance(metadata, Mapping) else {}
provider = metadata.get("route_provider") or _provider_from_model(model)
model_id = metadata.get("route_model")
if not model_id:
model_id = (
getattr(model, "model_name", None)
or getattr(model, "model", None)
or getattr(model, "model_id", None)
)
generation = metadata.get("route_config_generation")
try:
config_generation = int(generation) if generation is not None else None
except (TypeError, ValueError):
config_generation = None
return {
"provider": str(provider) if provider else None,
"model": str(model_id) if model_id else None,
"route_key": str(metadata.get("route_key"))
if metadata.get("route_key")
else None,
"config_generation": config_generation,
"api_mode": str(metadata.get("route_api_mode"))
if metadata.get("route_api_mode")
else None,
"endpoint": str(metadata.get("route_endpoint"))
if metadata.get("route_endpoint")
else None,
"tool_call_transport": str(metadata.get("route_tool_call_transport"))
if metadata.get("route_tool_call_transport")
else None,
}
def _raise_protocol_error(
request: ModelRequest,
reason: str,
*,
call_id: str | None = None,
call_diagnostic: dict[str, Any] | None = None,
) -> None:
raise ModelToolProtocolError(
reason,
call_id=call_id or None,
call_diagnostic=call_diagnostic,
**_route_metadata(request),
)
def _validate_message(
message: AIMessage,
request: ModelRequest,
allowed_names: frozenset[str],
) -> None:
invalid_calls = list(getattr(message, "invalid_tool_calls", None) or [])
if invalid_calls:
invalid = invalid_calls[0]
call_id = str(invalid.get("id") or "") if isinstance(invalid, Mapping) else ""
_raise_protocol_error(
request,
"invalid_final_call",
call_id=call_id,
call_diagnostic=_call_diagnostic(
message,
invalid,
source="invalid_tool_calls",
call_index=0,
call_count=len(invalid_calls),
),
)
parsed_by_id: dict[str, str] = {}
parsed_calls = list(getattr(message, "tool_calls", None) or [])
for call_index, raw_call in enumerate(parsed_calls):
diagnostic = _call_diagnostic(
message,
raw_call,
source="parsed_tool_calls",
call_index=call_index,
call_count=len(parsed_calls),
)
if not isinstance(raw_call, Mapping):
_raise_protocol_error(
request, "invalid_final_call", call_diagnostic=diagnostic
)
call_id = str(raw_call.get("id") or raw_call.get("call_id") or "").strip()
name = str(raw_call.get("name") or "").strip()
if not name:
_raise_protocol_error(
request,
"missing_name",
call_id=call_id,
call_diagnostic=diagnostic,
)
if name not in allowed_names:
_raise_protocol_error(
request,
"unknown_name",
call_id=call_id,
call_diagnostic=diagnostic,
)
if not call_id:
_raise_protocol_error(request, "missing_id", call_diagnostic=diagnostic)
if call_id in parsed_by_id:
_raise_protocol_error(
request,
"duplicate_id",
call_id=call_id,
call_diagnostic=diagnostic,
)
args = raw_call.get("args")
if not isinstance(args, Mapping):
_raise_protocol_error(
request,
"invalid_args",
call_id=call_id,
call_diagnostic=diagnostic,
)
parsed_by_id[call_id] = name
content = getattr(message, "content", None)
if not isinstance(content, list):
return
seen_block_ids: set[str] = set()
tool_blocks = [
block
for block in content
if isinstance(block, Mapping) and block.get("type") in _TOOL_BLOCK_TYPES
]
for block_index, block in enumerate(tool_blocks):
diagnostic = _call_diagnostic(
message,
block,
source="content_blocks",
call_index=block_index,
call_count=len(tool_blocks),
)
call_id, name = _block_identity(block)
if not call_id:
_raise_protocol_error(request, "missing_id", call_diagnostic=diagnostic)
if call_id in seen_block_ids:
_raise_protocol_error(
request,
"duplicate_id",
call_id=call_id,
call_diagnostic=diagnostic,
)
seen_block_ids.add(call_id)
parsed_name = parsed_by_id.get(call_id)
if parsed_name is None or (name and name != parsed_name):
_raise_protocol_error(
request,
"inconsistent_block",
call_id=call_id,
call_diagnostic=diagnostic,
)
class ToolProtocolGuardMiddleware(AgentMiddleware):
"""Fail closed on malformed final tool calls using the actual request tools."""
name = "tool_protocol_guard"
@staticmethod
def _validate(response: Any, request: ModelRequest) -> None:
allowed_names = frozenset(
name for tool in request.tools if (name := _tool_name(tool)) is not None
)
for message in _ai_messages(response):
_validate_message(message, request, allowed_names)
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelResponse:
response = handler(request)
self._validate(response, request)
return response
async def awrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse:
response = await handler(request)
self._validate(response, request)
return response
+3 -31
View File
@@ -48,7 +48,6 @@ DEFAULT_ALWAYS_INCLUDE_TOOLS: frozenset[str] = frozenset(
"read_memory",
"record_observation",
"search_observations",
"write_todos",
}
)
@@ -133,21 +132,10 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
return self._build_selector(request).wrap_model_call(
request, _handler_after_selection
)
except Exception as exc:
except Exception:
if _handler_called:
raise # Error from downstream model — don't retry
from ..llm.errors import ProviderStreamError
from .error_normalization import _is_provider_error
if isinstance(exc, ProviderStreamError) or _is_provider_error(exc):
# Auth / quota / connection failures on the selector's
# own model. Falling back to "use all tools" would hit
# the same provider anyway (same client, likely same
# credentials). Surface it instead so the user sees
# the real cause.
raise
# Structured-output shape / config failure — gracefully
# degrade to using all tools.
# Selector itself failed (e.g., structured output not supported).
logger.debug("Tool selector failed, using all tools", exc_info=True)
if self._track_stream_selection:
_selector_active = False
@@ -183,16 +171,9 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
return await self._build_selector(request).awrap_model_call(
request, _handler_after_selection
)
except Exception as exc:
except Exception:
if _handler_called:
raise
from ..llm.errors import ProviderStreamError
from .error_normalization import _is_provider_error
if isinstance(exc, ProviderStreamError) or _is_provider_error(exc):
# See sync path — surface provider errors, degrade only
# on shape / config failures.
raise
logger.debug("Tool selector failed, using all tools", exc_info=True)
if self._track_stream_selection:
_selector_active = False
@@ -277,15 +258,6 @@ def create_tool_selector_middleware(
model = _ensure_chat_model()
safe_model = disable_thinking(model)
safe_model = safe_model.model_copy(
update={
"tags": [*(safe_model.tags or []), "metering:tool_selector"],
"metadata": {
**(safe_model.metadata or {}),
"metering_scope": "tool_selector",
},
}
)
system_prompt = (
"You are selecting tools for a scientific research agent. "
-63
View File
@@ -5,7 +5,6 @@ from __future__ import annotations
import logging
import os
import shutil
from collections.abc import Iterator
from datetime import datetime
from pathlib import Path
@@ -216,65 +215,3 @@ def resolve_virtual_path(virtual_path: str) -> Path:
"""Resolve a virtual workspace path (e.g. /image.png) to a real filesystem path."""
vpath = virtual_path if virtual_path.startswith("/") else "/" + virtual_path
return (_active_workspace / vpath.lstrip("/")).resolve()
def evoscientist_root() -> Path:
"""Return the application root used by Gateway-managed runtime data."""
env_root = os.environ.get("EVOSCIENTIST_HOME")
if env_root:
return Path(env_root).expanduser().resolve()
return DATA_DIR.expanduser().resolve()
_EVOSCIENTIST_DATA_ROOT: Path | None = None
def _data_root() -> Path:
"""Return the root directory for isolated Web user workspaces."""
global _EVOSCIENTIST_DATA_ROOT
if _EVOSCIENTIST_DATA_ROOT is not None:
return _EVOSCIENTIST_DATA_ROOT
env_root = os.environ.get("EVOSCIENTIST_DATA_ROOT")
if env_root:
root = Path(env_root).expanduser().resolve()
else:
root = evoscientist_root() / "data"
_EVOSCIENTIST_DATA_ROOT = root
return root
def user_data_dir(user_id: str) -> Path:
"""Return and create the isolated data directory for a Web user."""
path = _data_root() / user_id
path.mkdir(parents=True, exist_ok=True)
return path
def iter_user_data_dirs() -> Iterator[Path]:
"""Yield existing Web user directories without creating the data root."""
root = _data_root()
if not root.exists():
return
for path in root.iterdir():
if path.is_dir():
yield path
def thread_data_dir(user_id: str, thread_id: str) -> Path:
"""Return and create a user's isolated thread workspace."""
path = user_data_dir(user_id) / thread_id
path.mkdir(parents=True, exist_ok=True)
return path
def global_data_dir(user_id: str) -> Path:
"""Return and create a user's directory shared across all threads."""
path = user_data_dir(user_id) / "__global__"
path.mkdir(parents=True, exist_ok=True)
return path
def uploads_dir() -> Path:
"""Return the Gateway upload staging directory."""
return evoscientist_root() / "uploads"
-116
View File
@@ -1,116 +0,0 @@
"""Optional runtime services supplied by an application embedding EvoScientist.
The CLI package must not import a concrete web gateway. Applications such as
Ai4Sci-Web can register their database, storage, metering, and media services
at process startup through this module.
"""
from __future__ import annotations
from collections.abc import Awaitable, Callable
from dataclasses import dataclass, replace
from datetime import date
from pathlib import Path
from typing import Any
AsyncProvider = Callable[[], Awaitable[Any]]
AsyncFileHandler = Callable[[Path], Awaitable[Any]]
AsyncUsageRecorder = Callable[[str, str], Awaitable[Any]]
ModelResolver = Callable[[str | None, str | None], Any | None]
class RuntimeIntegrationUnavailable(RuntimeError):
"""Raised when an optional host-provided service is not configured."""
@dataclass(frozen=True)
class RuntimeIntegrations:
app_connection_provider: AsyncProvider | None = None
session_connection_provider: AsyncProvider | None = None
session_dsn_provider: Callable[[], str | None] | None = None
current_date_provider: Callable[[], date] | None = None
user_storage_root_provider: Callable[[str], Path] | None = None
knowledge_file_handler: AsyncFileHandler | None = None
usage_recorder: AsyncUsageRecorder | None = None
image_backend_factory: Callable[[], Any] | None = None
model_resolver: ModelResolver | None = None
_integrations = RuntimeIntegrations()
def configure_runtime_integrations(**services: Any) -> RuntimeIntegrations:
"""Register host-provided services and return the resulting configuration."""
global _integrations
_integrations = replace(_integrations, **services)
return _integrations
def reset_runtime_integrations() -> None:
"""Clear all host-provided services, primarily for tests."""
global _integrations
_integrations = RuntimeIntegrations()
def has_session_connection_provider() -> bool:
return _integrations.session_connection_provider is not None
def get_session_dsn() -> str | None:
provider = _integrations.session_dsn_provider
return provider() if provider is not None else None
async def get_session_connection() -> Any:
provider = _integrations.session_connection_provider
if provider is None:
raise RuntimeIntegrationUnavailable(
"No session connection provider is configured"
)
return await provider()
async def get_app_connection() -> Any:
provider = _integrations.app_connection_provider
if provider is None:
raise RuntimeIntegrationUnavailable(
"No application connection provider is configured"
)
return await provider()
def current_date() -> date:
provider = _integrations.current_date_provider
return provider() if provider is not None else date.today()
def resolve_user_storage_root(user_id: str) -> Path | None:
provider = _integrations.user_storage_root_provider
return provider(user_id) if provider is not None else None
def resolve_runtime_model(model: str | None, provider: str | None = None) -> Any | None:
"""Resolve a host-managed model configuration when one is registered."""
resolver = _integrations.model_resolver
return resolver(model, provider) if resolver is not None else None
async def handle_knowledge_file(path: Path) -> None:
handler = _integrations.knowledge_file_handler
if handler is not None:
await handler(path)
async def record_service_usage(service: str, action: str) -> None:
recorder = _integrations.usage_recorder
if recorder is not None:
await recorder(service, action)
def get_image_backend() -> Any:
factory = _integrations.image_backend_factory
if factory is None:
raise RuntimeIntegrationUnavailable(
"Image generation is unavailable in this runtime. Configure an image backend first."
)
return factory()
+2 -32
View File
@@ -7,15 +7,6 @@ All events contain a type and associated data dict.
from dataclasses import dataclass
from typing import Any
STREAM_PROTOCOL_CAPABILITIES = frozenset(
{
"task_snapshot_v1",
"complete_tool_call_v1",
"correlated_tool_call_id_v1",
"final_invalid_tool_call_v1",
}
)
@dataclass
class StreamEvent:
@@ -167,14 +158,6 @@ class StreamEventEmitter:
},
)
@staticmethod
def task_snapshot(source: str, items: list[dict[str, Any]]) -> StreamEvent:
"""Emit the complete root-agent task state without product-specific IDs."""
return StreamEvent(
"task_snapshot",
{"type": "task_snapshot", "source": source, "items": items},
)
@staticmethod
def interrupt(
interrupt_id: str,
@@ -230,19 +213,6 @@ class StreamEventEmitter:
)
@staticmethod
def error(
message: str,
*,
code: str | None = None,
recoverable: bool | None = None,
details: dict[str, Any] | None = None,
) -> StreamEvent:
def error(message: str) -> StreamEvent:
"""Error event."""
data: dict[str, Any] = {"type": "error", "message": message}
if code is not None:
data["code"] = code
if recoverable is not None:
data["recoverable"] = recoverable
if details is not None:
data["details"] = details
return StreamEvent("error", data)
return StreamEvent("error", {"type": "error", "message": message})
+45 -284
View File
@@ -8,15 +8,13 @@ import base64
import inspect
import mimetypes
import os
import warnings
from collections.abc import AsyncGenerator, AsyncIterator, Mapping
from dataclasses import dataclass
from typing import Any, TypeAlias
from langchain_core._api import LangChainBetaWarning
from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, ToolMessage
from langgraph.graph import END
from langgraph.types import Command, Interrupt, Overwrite
from langgraph.types import Command, Interrupt
from ..memory.worker_activity import clear_completed_memory_activity_counts
from .emitter import StreamEventEmitter
@@ -45,12 +43,6 @@ GraphRunInput: TypeAlias = str | Command
LangGraphStreamInput: TypeAlias = dict[str, list[dict[str, object]]] | Command
_ValueMessageKey: TypeAlias = tuple[str, ...]
warnings.filterwarnings(
"ignore",
message=r"The v3 streaming protocol on Pregel is experimental\.",
category=LangChainBetaWarning,
)
@dataclass(frozen=True, slots=True)
class _AssistantValueMessage:
@@ -95,10 +87,9 @@ async def _clear_interrupted_graph_state(
no output and leaves the messages channel unchanged. From the user's side the
conversation looks like it lost all history because the agent stops responding.
Recovery first removes malformed/incomplete tool protocol from the messages
channel, then ``aupdate_state(config, None, as_node=END)`` clears pending
tasks and writes a checkpoint whose ``next`` is the empty tuple. Completed
tool call/result pairs and all non-tool history are preserved.
The fix: ``aupdate_state(config, None, as_node=END)`` clears all pending tasks
and writes a checkpoint whose ``next`` is the empty tuple, without touching
any channel values (message history is preserved).
Critically, this only runs when the stuck state is *not* a legitimate
human-in-the-loop interrupt. The agent pauses via ``interrupt()`` /
@@ -115,9 +106,10 @@ async def _clear_interrupted_graph_state(
_log = logging.getLogger(__name__)
try:
snapshot = await agent.aget_state(config)
if not snapshot:
# Only act when the graph is genuinely stuck (non-empty next tuple)...
if not snapshot or not getattr(snapshot, "next", None):
return
# Never alter a real human-in-the-loop pause.
# ...and not parked at a real human-in-the-loop interrupt.
if _snapshot_has_pending_interrupt(snapshot):
_log.debug(
"Leaving interrupted graph state intact for thread %s: "
@@ -127,13 +119,6 @@ async def _clear_interrupted_graph_state(
)
return
await _repair_malformed_tool_history(agent, config, snapshot=snapshot)
# Only force END when the graph is genuinely stuck. Message repair also
# applies to failures that already left next empty.
if not getattr(snapshot, "next", None):
return
stuck_at = snapshot.next
await agent.aupdate_state(config, None, as_node=END)
_log.debug(
@@ -149,49 +134,6 @@ async def _clear_interrupted_graph_state(
)
async def _repair_malformed_tool_history(
agent: Any,
config: dict[str, Any],
*,
snapshot: Any | None = None,
) -> bool:
"""Rewrite a checkpoint's messages to a replay-safe tool history.
Only structurally invalid protocol is removed. Completed tool call/result
pairs, including repeated successes and repeated errors, are audit and
billing facts and must remain in persistent history.
"""
import logging
from ..llm.patches import _sanitize_openai_tool_history
_log = logging.getLogger(__name__)
if snapshot is None:
snapshot = await agent.aget_state(config)
if not snapshot or _snapshot_has_pending_interrupt(snapshot):
return False
values = getattr(snapshot, "values", None)
if not isinstance(values, Mapping):
return False
messages = values.get("messages")
if not isinstance(messages, list):
return False
repaired = _sanitize_openai_tool_history(messages)
if repaired == messages:
return False
await agent.aupdate_state(config, {"messages": Overwrite(repaired)})
_log.warning(
"Repaired structurally invalid tool history for thread %s: messages %d -> %d",
config.get("configurable", {}).get("thread_id", "?"),
len(messages),
len(repaired),
)
return True
@dataclass(frozen=True)
class _SubagentInfo:
path: tuple[str, ...]
@@ -267,12 +209,7 @@ class _V3EventProcessor:
tuple[tuple[str, ...], str], tuple[str, dict[str, Any]]
] = {}
self._emitted_tool_calls: set[tuple[tuple[str, ...], str]] = set()
self._pending_tool_calls: dict[
tuple[tuple[str, ...], str], tuple[str, dict[str, Any]]
] = {}
self._emitted_interrupts: set[str] = set()
self._pending_invalid_tool_calls: dict[str, tuple[str, str]] = {}
self._last_task_snapshot: tuple[tuple[str, str], ...] | None = None
self._selector = _ToolSelectionSuppressor(emitter)
@staticmethod
@@ -300,26 +237,13 @@ class _V3EventProcessor:
if method == "tools":
return self._process_tool_event(namespace, _event_data(event), subagent)
if method == "updates":
return self._process_update_event(
_event_data(event), namespace=namespace, source="update"
)
return self._process_update_event(_event_data(event))
if method == "values":
events: list[dict[str, Any]] = []
params = event.get("params") or {}
interrupts = params.get("interrupts") or ()
if interrupts:
events.extend(
self._process_update_event(
{"__interrupt__": interrupts},
namespace=namespace,
source="values",
)
)
events.extend(
self._process_update_event(
_event_data(event), namespace=namespace, source="values"
)
)
events.extend(self._process_update_event({"__interrupt__": interrupts}))
if self._process_value_message_snapshots and not namespace:
events.extend(self._process_value_messages(_event_data(event)))
return events
@@ -458,7 +382,6 @@ class _V3EventProcessor:
inp, out = _usage_counts(usage) if usage is not None else (0, 0)
if inp or out:
events.append(self.emitter.usage_stats(inp, out).data)
events.extend(self._flush_invalid_tool_calls())
return events
return []
@@ -484,12 +407,14 @@ class _V3EventProcessor:
if tool_call is None:
return events
tool_name, args, tool_call_id = tool_call
self._pending_invalid_tool_calls.pop(tool_call_id, None)
self._pending_tool_calls[
(self._tool_scope(namespace, subagent), tool_call_id)
] = (
tool_name,
args,
events.extend(
self._emit_tool_call_once(
namespace=namespace,
subagent=subagent,
name=tool_name,
args=args,
tool_call_id=tool_call_id,
)
)
return events
@@ -532,20 +457,6 @@ class _V3EventProcessor:
]
return [self.emitter.tool_call(name, args, tool_call_id).data]
def _pending_call_id(
self,
*,
scope: tuple[str, ...],
name: str,
args: dict[str, Any],
) -> str:
matches = [
call_id
for (candidate_scope, call_id), candidate in self._pending_tool_calls.items()
if candidate_scope == scope and candidate == (name, args)
]
return matches[0] if len(matches) == 1 else ""
def _process_whole_message(
self,
msg: AIMessage | AIMessageChunk,
@@ -553,16 +464,6 @@ class _V3EventProcessor:
namespace: tuple[str, ...],
) -> list[dict[str, Any]]:
events: list[dict[str, Any]] = []
for invalid in getattr(msg, "invalid_tool_calls", ()) or ():
invalid_map = _as_raw_map(invalid)
if invalid_map is None:
continue
call_id = str(
invalid_map.get("id") or invalid_map.get("tool_call_id") or ""
)
name = str(invalid_map.get("name") or invalid_map.get("tool_name") or "")
key = call_id or f"chunk_{len(self._pending_invalid_tool_calls)}"
self._pending_invalid_tool_calls[key] = (call_id, name)
additional = msg.additional_kwargs
reasoning = additional.get("reasoning_content")
emitted_reasoning = False
@@ -585,12 +486,14 @@ class _V3EventProcessor:
if tool_call is None:
continue
tool_name, args, tool_call_id = tool_call
self._pending_invalid_tool_calls.pop(tool_call_id, None)
self._pending_tool_calls[
(self._tool_scope(namespace, subagent), tool_call_id)
] = (
tool_name,
args,
events.extend(
self._emit_tool_call_once(
namespace=namespace,
subagent=subagent,
name=tool_name,
args=args,
tool_call_id=tool_call_id,
)
)
if subagent is None:
@@ -623,9 +526,6 @@ class _V3EventProcessor:
name,
args,
)
self._pending_tool_calls.pop(
(self._tool_scope(namespace, subagent), tool_call_id), None
)
events.extend(
self._emit_tool_call_once(
namespace=namespace,
@@ -672,7 +572,6 @@ class _V3EventProcessor:
content += "\n... (truncated)"
success = is_success(content)
lifecycle_key = (self._tool_scope(namespace, subagent), tool_call_id)
if subagent is not None:
events.append(
self.emitter.subagent_tool_result(
@@ -684,78 +583,22 @@ class _V3EventProcessor:
instance_id=subagent.instance_id,
).data
)
else:
events.append(
self.emitter.tool_result(
name, content, success, tool_call_id=tool_call_id
).data
)
self._emitted_tool_calls.discard(lifecycle_key)
self._pending_tool_calls.pop(lifecycle_key, None)
return events
events.append(
self.emitter.tool_result(
name, content, success, tool_call_id=tool_call_id
).data
)
return events
return []
@staticmethod
def _normalize_task_items(value: object) -> list[dict[str, str]] | None:
if not isinstance(value, list):
return None
aliases = {
"todo": "pending",
"pending": "pending",
"active": "in_progress",
"in-progress": "in_progress",
"in_progress": "in_progress",
"done": "completed",
"completed": "completed",
}
items: list[dict[str, str]] = []
for raw in value:
raw_map = _as_raw_map(raw)
if raw_map is None:
continue
content = str(raw_map.get("content") or raw_map.get("task") or "").strip()
if not content:
continue
status = aliases.get(str(raw_map.get("status") or "pending").lower())
if status is None:
continue
items.append({"content": content, "status": status})
return items
@classmethod
def _find_task_items(cls, data: object) -> list[dict[str, str]] | None:
data_map = _as_raw_map(data)
if data_map is None:
return None
if "todos" in data_map:
return cls._normalize_task_items(data_map["todos"])
for value in data_map.values():
nested = _as_raw_map(value)
if nested is not None and "todos" in nested:
return cls._normalize_task_items(nested["todos"])
return None
def _process_update_event(
self,
data: object,
*,
namespace: tuple[str, ...] = (),
source: str = "update",
) -> list[dict[str, Any]]:
def _process_update_event(self, data: object) -> list[dict[str, Any]]:
events: list[dict[str, Any]] = []
data_map = _as_raw_map(data)
if data_map is not None and "__interrupt__" in data_map:
events.extend(self._process_interrupts(data_map["__interrupt__"]))
if not namespace:
items = self._find_task_items(data)
if items is not None:
signature = tuple((item["content"], item["status"]) for item in items)
if signature != self._last_task_snapshot:
self._last_task_snapshot = signature
events.append(self.emitter.task_snapshot(source, items).data)
summarization_event = _find_summarization_event_payload(data)
if summarization_event and not self._summarization_in_progress:
signature = _summarization_event_signature(summarization_event)
@@ -771,10 +614,6 @@ class _V3EventProcessor:
events.extend(self._emit_summarization_text(summary_text))
return events
def _flush_invalid_tool_calls(self) -> list[dict[str, Any]]:
self._pending_invalid_tool_calls.clear()
return []
def _process_interrupts(self, interrupts: object) -> list[dict[str, Any]]:
events: list[dict[str, Any]] = []
if not isinstance(interrupts, list | tuple):
@@ -813,74 +652,26 @@ class _V3EventProcessor:
raw_questions = interrupt_map.get("questions")
questions = raw_questions if isinstance(raw_questions, list) else []
tc_id = str(interrupt_map.get("tool_call_id", ""))
events: list[dict[str, Any]] = []
candidate = self._pending_tool_calls.get(((), tc_id)) if tc_id else None
if candidate is not None:
events.extend(
self._emit_tool_call_once(
namespace=(),
subagent=None,
name=candidate[0],
args=candidate[1],
tool_call_id=tc_id,
)
)
events.extend(
self._dedupe_interrupt_event(
self.emitter.ask_user_interrupt(
interrupt_id,
questions,
tc_id,
).data
)
return self._dedupe_interrupt_event(
self.emitter.ask_user_interrupt(
interrupt_id,
questions,
tc_id,
).data
)
return events
raw_action_reqs = interrupt_map.get("action_requests")
action_reqs = raw_action_reqs if isinstance(raw_action_reqs, list) else []
raw_review_cfgs = interrupt_map.get("review_configs")
review_cfgs = raw_review_cfgs if isinstance(raw_review_cfgs, list) else None
if action_reqs:
events: list[dict[str, Any]] = []
for raw_request in action_reqs:
request_map = _as_raw_map(raw_request)
if request_map is None:
continue
call_id = str(
request_map.get("id") or request_map.get("tool_call_id") or ""
)
name = str(request_map.get("name") or request_map.get("tool_name") or "")
args_map = _as_raw_map(
request_map.get("args")
if "args" in request_map
else request_map.get("input")
)
if not call_id and name and args_map is not None:
call_id = self._pending_call_id(
scope=(),
name=name,
args=dict(args_map),
)
if call_id and name and args_map is not None:
events.extend(
self._emit_tool_call_once(
namespace=(),
subagent=None,
name=name,
args=dict(args_map),
tool_call_id=call_id,
)
)
events.extend(
self._dedupe_interrupt_event(
self.emitter.interrupt(
interrupt_id,
action_reqs,
review_cfgs,
).data
)
return self._dedupe_interrupt_event(
self.emitter.interrupt(
interrupt_id,
action_reqs,
review_cfgs,
).data
)
return events
return []
def _dedupe_interrupt_event(self, event: dict[str, Any]) -> list[dict[str, Any]]:
@@ -1006,8 +797,6 @@ async def stream_agent_events(
thread_id: str,
metadata: dict[str, Any] | None = None,
media: list[str] | None = None,
callbacks: list[Any] | None = None,
error_mode: str = "emit",
) -> AsyncGenerator[dict[str, Any], None]:
"""Stream events from a DeepAgents/LangGraph v3 run.
@@ -1023,9 +812,6 @@ async def stream_agent_events(
metadata: Optional metadata dict merged into the LangGraph config
(e.g. agent_name, updated_at for checkpoint persistence).
media: Optional list of local file paths for attachments.
callbacks: Optional Runnable callbacks propagated to all nested model calls.
error_mode: ``emit`` preserves the generic error event; ``raise`` lets an
embedding host produce the single terminal error envelope.
Yields:
Event dicts: thinking, text, tool_call, tool_result,
@@ -1035,8 +821,6 @@ async def stream_agent_events(
config: dict[str, Any] = {"configurable": {"thread_id": thread_id}}
if metadata:
config["metadata"] = metadata
if callbacks:
config["callbacks"] = callbacks
emitter = StreamEventEmitter()
existing_summarization_event: Mapping[str, object] | None = None
try:
@@ -1165,30 +949,7 @@ async def stream_agent_events(
yield item
except Exception as e:
_run_raised = True
if error_mode == "emit":
payload = e.model_dump() if hasattr(e, "model_dump") else {}
if not isinstance(payload, Mapping):
payload = {}
code = str(payload.get("code") or "") or None
details = {
key: payload[key]
for key in (
"reason",
"provider",
"model",
"route_key",
"config_generation",
"api_mode",
"call_id",
)
if payload.get(key) is not None
}
yield emitter.error(
str(payload.get("message") or e),
code=code,
recoverable=bool(payload.get("recoverable", True)) if code else None,
details=details or None,
).data
yield emitter.error(str(e)).data
raise
finally:
if stream is not None:
+34 -2
View File
@@ -137,6 +137,34 @@ Moving beyond traditional human-in-the-loop systems, EvoScientist adopts a human
> [!TIP]
> Looking for ready-to-use research skills? Check out [**EvoSkills**](https://github.com/EvoScientist/EvoSkills) — powered by [**EvoScientist**](https://github.com/EvoScientist/EvoScientist)'s engine and installable skills, the entire end-to-end research lifecycle is covered out of the box. [**EvoSkills**](https://github.com/EvoScientist/EvoSkills) are also compatible with other CLI coding agents.
## WebUI Streaming Fix Notes
The WebUI streaming fix is based on one rule: a browser-side SSE disconnect is
only a transport event, not proof that the LangGraph run has finished.
- Keep the composer action button as **Stop** while the backend thread still has
pending work. The UI should not switch back to **Send** merely because the
live SSE stream ended.
- Derive the visible busy state from both the live stream and backend thread
state. In the patched WebUI this is represented as `isRunActive`, backed by
`threads.getState(threadId)`.
- Restore **Send** only after the backend reports a terminal state: idle,
completed final answer, failed, or cancelled.
- Preserve human-in-the-loop controls during tool approvals. Approval buttons
remain actionable, while the bottom composer still shows **Stop** to avoid
duplicate submissions.
- Treat leaked `{"tools":[...]}` payloads as internal tool-selection control
messages, not assistant text, and filter them from the transcript.
- Recover dropped long-response tails through the checkpoint-backed
`/api/threads/{thread_id}/final-answer` endpoint instead of relying only on
browser stream state.
- Use a long enough final-answer polling window for real research runs; short
fallback windows can expire before the backend completes and make the browser
appear truncated until refresh.
- Validate with real browser flows: normal streaming, simulated SSE disconnect
and reload, pending tool approval, rejection/cancellation, and final transition
back to **Send**.
## 🔥 News
- **[03 Jun 2026]** 🥈 Ranked #2 overall — and 🥇 #1 among `GPT-5.4`-based agents — on [ResearchClawBench](https://github.com/InternScience/ResearchClawBench) (Agent Mode)! [**Leaderboard**](https://internscience.github.io/ResearchClawBench-Home/) 👈
- **[18 Apr 2026]** 🥇 Ranked #1 on [DeepResearch Bench](https://deepresearch-bench.github.io/) at submission time! [**Leaderboard**](https://huggingface.co/spaces/muset-ai/DeepResearch-Bench-Leaderboard) 👈
@@ -151,7 +179,7 @@ Moving beyond traditional human-in-the-loop systems, EvoScientist adopts a human
<details>
<summary>📦 Release Highlights — version changelog</summary>
- **[11 Jul 2026]** **[v0.2.2](https://github.com/EvoScientist/EvoScientist/releases/tag/v0.2.2)** — New models selectable in onboarding and `/model`: GPT-5.6 (sol, terra, luna) for OpenAI and OpenRouter, plus Grok 4.5 and Tencent Hunyuan HY3 on OpenRouter; tighter config-file permissions and a reworked onboarding OAuth flow for auxiliary models.
- **[07 Jul 2026]** **[v0.2.2](https://github.com/EvoScientist/EvoScientist/releases/tag/v0.2.2)** — WebUI streaming resilience hotfix: SSE disconnects no longer imply run completion; the composer stays on **Stop** until backend thread state is terminal; final-answer checkpoint recovery fills dropped response tails; tool-selection JSON payloads are filtered from live transcripts; `EVOSCIENTIST_WEBUI_PACKAGE` lets the launcher run a local patched WebUI package for validation before npm publication.
- **[05 Jul 2026]** **[v0.2.1](https://github.com/EvoScientist/EvoScientist/releases/tag/v0.2.1)** — AutoSkills: EvoMemory drafts reusable skills from its own observation clusters for you to review via `/autoskills`; a new `--output-format stream-json` for headless / SDK clients; richer slash-command completions; Windows UTF-8 config reads; a TUI welcome-banner fix; langchain-openrouter 0.2.5.
- **[26 Jun 2026]** **[v0.2.0](https://github.com/EvoScientist/EvoScientist/releases/tag/v0.2.0)** — Scheduled tasks: cron-style recurring runs via `/schedule` or natural language, run unattended with shell-access gating; self-linking memory that connects observations into a knowledge graph (complements / contradicts / supersedes); a read-only `GET /api/models` endpoint for the WebUI model picker.
- **[23 Jun 2026]** **[v0.1.9](https://github.com/EvoScientist/EvoScientist/releases/tag/v0.1.9)** — Hotfix for fresh installs: the first message crashed with `The subagent `task` tool cannot be exposed via `ptc`` after deepagents 0.6.11 / langchain-quickjs 0.3 reserved `task` as the REPL global. Removed `task` from the code-interpreter PTC allowlist (`task()` stays available as the REPL global; async dispatch stays in PTC) and pinned `deepagents[quickjs]~=0.6.11`.
@@ -181,6 +209,7 @@ Moving beyond traditional human-in-the-loop systems, EvoScientist adopts a human
- [📦 Installation](#-installation)
- [🔑 Configuration](#-configuration)
- [⚡ Quick Start](#-quick-start)
- [WebUI Streaming Fix Notes](#webui-streaming-fix-notes)
- [⏰ Scheduled Tasks](#-scheduled-tasks)
- [🍪 Examples & Recipes](#-examples--recipes)
- [🔌 MCP Integration](#-mcp-integration)
@@ -434,7 +463,10 @@ EvoSci deploy # standalone LangGraph server for external UIs
EvoSci -p "query" --output-format stream-json --auto-mode # JSONL event stream on stdout (for programmatic clients)
```
`--output-format stream-json` makes a single-shot (`-p`) run emit its native events as line-delimited JSON on stdout (one object per line), with all human output on stderr — the integration surface for headless clients (e.g. an agent runtime). See [docs/guides/stream-json.md](docs/guides/stream-json.md) for the event schema.
`--output-format stream-json` makes a single-shot (`-p`) run emit its native
events as line-delimited JSON on stdout (one object per line), with all human
output on stderr — the integration surface for headless clients (e.g. an agent
runtime). See [docs/stream-json.md](docs/stream-json.md) for the event schema.
</details>
-5
View File
@@ -11,11 +11,6 @@
|------------------------------------------------------------|---------------------------------------------------------------------------------|
| [macOS 24/7 Deployment](https://github.com/EvoScientist/EvoScientist/blob/main/docs/recipes/deployment-macos-24h.md#running-evoscientist-247-on-macos-telegram-bot--stt--ccproxy) | Run EvoScientist as an always-on service on macOS with OAuth + Telegram + STT |
| Guide | Description |
|------------------------------------------------------------|---------------------------------------------------------------------------------|
| [`stream-json` output protocol](https://github.com/EvoScientist/EvoScientist/blob/main/docs/guides/stream-json.md#stream-json-output-protocol) | Line-delimited JSON event stream (`--output-format stream-json`) for driving EvoScientist headlessly from SDK / programmatic clients |
## Contributing a Recipe
See the [Contributing Guide](../CONTRIBUTING.md) for general guidelines. When adding a new recipe:
@@ -0,0 +1,303 @@
# TUI as Primary Interface + Run Resilience
**Date:** 2026-07-06
**Status:** Design — pending implementation plan
**Scope:** EvoScientist CLI/TUI launch path, run-interruption recovery, WebUI opt-in hardening
## 1. Background & Root Cause
A user running the WebUI backend (`ui_backend = "webui"`) saw a long research run
(127 s) render truncated output in the browser, ending mid-token at `O-8282d7e`.
Investigation established:
- **The server-side run completed cleanly.** The thread checkpoint in
`~/.evoscientist/sessions.db` holds the full 6818-char final answer
(`status: idle`, `error: None`). The text `O-8282d7e` does not appear at the
end of any stored content — it is a mid-stream artifact.
- **The disconnect is browser-side.** `langgraph dev` sends an SSE heartbeat
every 5 s (`langgraph_api/sse.py:102`), `send_timeout` is `None` (never kills
slow sends), and the server supports resume via the `last-event-id` header
(`langgraph_api/api/runs.py:623`). The `runs/stream` endpoint is **POST**, so
the browser uses `fetch()` to read it; native `EventSource` auto-reconnect does
not apply, and the `@evoscientist/webui` front-end does not re-POST with
`last-event-id` to resume.
- **`stream_resumable=True`** lets the server-side run finish after the browser
drops, which is why "server succeeded" ≠ "browser received".
- A separate symptom from the same session — `Port 6174 cannot be bound` — was a
leftover `langgraph dev` process (PID 15288) whose `atexit` cleanup never ran
(parent killed hard).
**Conclusion:** the SSE disconnect is structural to the WebUI path. The TUI path
does not use SSE at all — `tui_interactive.py:297` calls `create_runtime_gateways()`
with no `backend`, so it defaults to `backend="local"` (`gateway/runtime.py`),
which runs the compiled graph **in-process** and is immune.
## 2. Goals
1. **Make the in-process TUI the primary interface.** Long runs no longer touch
SSE, so the class of failure the user hit cannot occur.
2. **Guarantee the user always sees the complete answer.** Even if a TUI run is
interrupted (Ctrl-C, crash, terminal close), the next invocation can surface
the latest persisted assistant message and optionally resume.
3. **Keep the WebUI available as opt-in** (`--web`), with the launch path
hardened against orphaned processes and port conflicts.
4. **Eliminate the port-6174 / orphaned-`langgraph dev` failure mode** for the
default (TUI) path: it must not start `langgraph dev` at all.
## 3. Non-Goals
- Patching `langgraph_api` upstream.
- Replacing the WebUI front-end's fetch-SSE transport. (`@evoscientist/webui`
lives in a separate npm repo with its own release cycle; a client-side resume
fix is noted as future work, not in scope here.)
- Changing the graph, middleware, or checkpointer semantics.
- Touching the `EvoSci deploy` standalone-server command.
## 4. Architecture
Three launch modes, selected by a single source of truth:
| Mode | `ui_backend` | Runs graph via | Starts `langgraph dev`? |
|---|---|---|---|
| **TUI (default)** | `tui` | `LocalGraphGateway` (in-process) | No |
| CLI | `cli` | `LocalGraphGateway` (in-process) | No |
| WebUI (opt-in) | `webui` | `langgraph dev` over HTTP/SSE | Yes |
A new `--web` flag on the top-level `evoscientist` command forces `webui` for
that invocation regardless of config; `--tui` forces `tui`. This lets users keep
`ui_backend = "webui"` in config but drop into the TUI without editing config,
and vice-versa.
### 4.1 Launch-mode resolution (single source of truth)
Today the default is inconsistent: `config/settings.py:279` defaults `ui_backend`
to `"tui"`, but `cli/tui_runtime.py:12` has `DEFAULT_UI_BACKEND = "cli"`. This
design collapses to one rule, evaluated in order:
1. `--web` flag → `webui`
2. `--tui` flag → `tui`
3. `config.ui_backend` (resolved via existing `resolve_ui_backend`)
4. `"tui"` (hardcoded final fallback)
`resolve_ui_backend` (`cli/tui_runtime.py:41`) stays the single normalizer; the
new flags feed into it rather than bypassing it.
### 4.2 Why TUI is immune (and what can still interrupt it)
In-process execution via `stream_agent_events` has no network hop. The only
"loss" modes are process-level:
- Ctrl-C mid-run
- Terminal close / SIGHUP
- Python crash / OOM
Because the checkpointer is `AsyncSqliteSaver` writing to
`~/.evoscientist/sessions.db` (`sessions.py:108` `get_db_path`), every completed
super-step is on disk before the next begins. An interrupted run therefore
leaves a recoverable checkpoint whose `next` tuple is non-empty (pending nodes)
and whose latest assistant message — if the model node finished — is the same
content the WebUI failed to deliver.
## 5. Components
### C1 — Launch-mode flags + TUI default
**Files:**
- `EvoScientist/cli/__init__.py` (Typer app entry) — add `--web` / `--tui`
options on the root command; thread the resolved `ui_backend` into the existing
dispatch.
- `EvoScientist/cli/tui_runtime.py:12` — remove the divergent
`DEFAULT_UI_BACKEND = "cli"`; route through `resolve_ui_backend` so
`settings.py`'s `"tui"` default wins.
- `EvoScientist/cli/commands.py` — wherever `ui_backend` is currently read, accept
the flag override.
**Behavior:** Running `evoscientist` with no flags and default config launches
the TUI and starts **no** `langgraph dev` subprocess. No port is bound; the
orphan-process class of bug disappears for the default path.
**Acceptance:** `evoscientist` (no args, default config) does not open a
listening socket on 6174 or any other port.
### C2 — TUI run resilience (Approach B)
Two cooperating pieces, both backed by `sessions.db`.
#### C2a — Interrupt-fallback in the streaming loop
Wrap the TUI streaming consumer so that on any interruption
(`KeyboardInterrupt`, `asyncio.CancelledError`, unexpected `Exception`) it reads
the thread's latest checkpoint and renders the latest assistant message before
exiting.
**Files:**
- `EvoScientist/stream/display.py` — `_run_streaming` and
`_astream_to_console` (the two streaming entry points used by
`cli/tui_backends.py`). Add a `finally`/except guard that, after an
interruption, calls a new `render_latest_assistant(thread_id)` helper.
- `EvoScientist/sessions.py` — add `get_latest_assistant_message(thread_id) ->
str | None`: read the most recent checkpoint for the thread, walk `messages`
backwards to the last `AIMessage` with non-empty content, return its text
(joined content blocks). Reuses existing `AsyncSqliteSaver` access already in
this module.
**Behavior on Ctrl-C mid-run:**
1. The guard catches the interrupt.
2. Prints a one-line notice: `Run interrupted — recovering latest output from
checkpoint…`.
3. Renders the latest persisted assistant message via the same Markdown renderer
the normal completion path uses.
4. If no assistant message is persisted yet (model node hadn't completed),
prints `Run interrupted before any output was saved. Resume with: evoscientist
--thread <id>` and exits.
**Non-goal:** this does not *continue* the run; it only guarantees visibility of
whatever completed. Continuation is C2b.
#### C2b — Startup detection of interrupted runs
Runs only when the TUI is launched to **resume a specific thread** (via the
existing thread-resume flag, i.e. the command `resume_hint.py` already prints at
exit). For a brand-new session there is nothing to detect and the prompt is
skipped. When resuming, check the thread's latest checkpoint: `next` is
non-empty **and** no run is currently active for that thread. If so, prompt:
```
You have an unfinished run in thread <id>:
[s] show what completed
[r] resume the run from its last checkpoint
[n] start fresh
```
**Files:**
- `EvoScientist/sessions.py` — add `detect_interrupted_run(thread_id) ->
InterruptedRunInfo | None` returning the checkpoint id, the pending node names
(`next`), and the timestamp. Implemented as a thin read over the existing
checkpointer's `aget_tuple`.
- `EvoScientist/cli/tui_interactive.py` — near the existing thread-load path
(around the `create_runtime_gateways()` call at line 297), invoke the detector
for the active thread and render the prompt above.
**Resume semantics:** selecting `r` re-invokes the graph with the same
`thread_id`. LangGraph resumes from the last checkpoint automatically (this is
the same mechanism `evoscientist --thread <id>` already relies on for session
continuity). No new graph code needed.
**Acceptance:** After killing a TUI mid-run with Ctrl-C, relaunching with
`--thread <id>` offers `[s]`/`[r]`/`[n]`; `[s]` prints the persisted final
answer; `[r]` continues execution from where it stopped.
### C3 — WebUI mode hardening (`--web`)
Applies only when `ui_backend == "webui"`. The goal is to make the existing
`deploy/webui.py` launch path robust and to mitigate (not eliminate) the SSE
limitation.
**Files:**
- `EvoScientist/deploy/webui.py` — three changes:
1. **Process-group + signal cleanup.** Today `atexit.register(stop_langgraph_dev,
...)` only fires on clean exit. Wrap the spawned `langgraph dev` (and the
`npx` child) in a process group (`start_new_session=True` on POSIX; on
Windows use a job object via the existing `_winloop.py` helpers) and
install `SIGINT`/`SIGTERM` handlers that tear the group down. This closes
the orphan-PID-15288 path for normal terminations. `SIGKILL` cannot run
handlers on either platform — document that hard kills may still orphan
the child, and the port-conflict UX in (2) is the user-facing recovery.
2. **Port-conflict UX.** The current "Port X is occupied by another process"
message is already correct; extend it to detect *whether the occupant is a
`langgraph dev`* (via the existing `/ok` health probe) and, if so, offer to
reuse it rather than bail. If the occupant is foreign, print the `lsof`
hint already present.
3. **Resume window.** Set `RESUMABLE_STREAM_TTL_SECONDS=600` in the env passed
to `start_langgraph_dev` so a browser that does reconnect can replay events
for runs up to 10 minutes long (the observed run was 127 s vs. the 120 s
default).
- `EvoScientist/deploy/server.py` — `start_langgraph_dev` already accepts env;
pass `RESUMABLE_STREAM_TTL_SECONDS` through.
**What C3 does NOT do:** it cannot make the browser auto-resume, because that
logic lives in `@evoscientist/webui`. It only widens the server-side window and
stops the orphan process. Documented in the WebUI exit message: *for long runs,
prefer the TUI (`evoscientist --tui`); the browser client does not resume a
dropped stream.*
### C4 — Documentation
- `README` / CLI `--help`: `--web` (browser), `--tui` (terminal, default,
recommended for long runs).
- A short "Why did my WebUI output stop?" note pointing at this design doc, so
future users reading the symptom can find the explanation.
## 6. Data Flow
### Default (TUI) path
```
evoscientist ─▶ resolve_ui_backend ─▶ "tui"
│
(no langgraph dev, no port bound)
▼
create_runtime_gateways(backend="local")
│
LocalGraphGateway.stream_events
│
stream_agent_events (in-process) ── checkpoint per super-step ──▶ sessions.db
│
on interrupt ──▶ get_latest_assistant_message(thread_id) ──▶ render
```
### WebUI (`--web`) path
unchanged from today (langgraph dev + `npx @evoscientist/webui`), plus C3's
process-group cleanup and `RESUMABLE_STREAM_TTL_SECONDS=600`.
## 7. Error Handling
| Failure | TUI default path | WebUI `--web` path |
|---|---|---|
| Ctrl-C mid-run | C2a renders latest persisted assistant message | Browser truncates; user runs `evoscientist --tui --thread <id>` → `[s]` to recover |
| Crash / SIGHUP | Checkpoint on disk; next `--thread <id>` shows `[s]`/`[r]` prompt (C2b) | Same recovery via TUI |
| Port 6174 bound | N/A — no port bound | C3 reuses if `langgraph dev`, else prints `lsof` hint |
| Orphaned `langgraph dev` | Cannot originate from TUI path | C3 process-group + signal handlers tear it down |
## 8. Testing
- **Unit** (`tests/test_sessions_*.py`): `get_latest_assistant_message` returns
the full final answer for a thread whose last run completed; returns `None`
for a thread interrupted before the model node.
- **Unit**: `detect_interrupted_run` returns pending `next` nodes for an
interrupted thread; `None` for an idle one.
- **Integration** (new `tests/test_tui_interrupt_recovery.py`): start a TUI run
against a graph whose model node sleeps, send `SIGINT` mid-run, assert the
process exits 0 after printing the recovered message.
- **Integration**: relaunch with `--thread <id>`, assert the `[s]/[r]/[n]`
prompt appears and `[r]` produces a completed run.
- **Launch** (`tests/test_cli_launch.py`, extend): `evoscientist` with default
config binds no port; `--web` binds 6174 and tears it down on `SIGTERM`.
- **WebUI** (manual / smoke): `--web`, kill parent with `SIGKILL`, assert no
orphan `langgraph dev` remains after the process-group change (best-effort —
`SIGKILL` cannot run handlers, but the process group lets the OS reap
children when the session ends; document this limit).
## 9. Rollout / Migration
- User's config currently has `ui_backend = "webui"`. After this change, that
config still selects WebUI (back-compatible). The user can run `evoscientist
--tui` immediately to get the resilient path, or `EvoSci config set ui_backend
tui` to make it permanent.
- No migration of `sessions.db` — schema is unchanged.
- No breaking change to `EvoSci deploy`.
## 10. Open Questions / Future Work
- **`@evoscientist/webui` client-side resume.** The front-end could re-POST
`/runs/stream` with `last-event-id` (or use the thread's SSE endpoint with
`since`) after a disconnect. This is the only way to fully fix the WebUI
symptom; it lives in the npm repo and is out of scope here. File a
cross-repo issue referencing this design.
- Whether `[r]` resume should replay the *interrupted* model call or only
continue from the *next* node. LangGraph's default (resume from pending
writes) is correct for most cases; revisit if users hit re-execution of
expensive tool calls.
- `cli` backend vs `tui` backend distinction — both are in-process; consider
deprecating the `cli`/`tui` split in a follow-up once the resilience work
lands, to reduce the three-way default confusion that caused the original
`DEFAULT_UI_BACKEND` divergence.
@@ -0,0 +1,196 @@
# WebUI SSE Truncation Recovery via Checkpoint Fallback (α)
**Date:** 2026-07-06
**Status:** Design — pending implementation plan
**Scope:** `@evoscientist/webui` front-end (separate npm repo) + one optional convenience route in this Python repo
## 1. Background & Root Cause
A WebUI user saw a 127 s research run render truncated output, ending mid-token at
`O-8282d7e`. Investigation established:
- The server-side run **completed cleanly**; the thread checkpoint in
`~/.evoscientist/sessions.db` holds the full 6818-char final answer
(`status: idle`, `error: None`).
- The disconnect is **browser-side**: `langgraph dev` sends an SSE heartbeat every
5 s, `send_timeout` is `None`, and the server supports resume via
`last-event-id` (`langgraph_api/api/runs.py:623`). But `/runs/stream` is
**POST**, so the browser reads it via `fetch()`; native `EventSource`
auto-reconnect does not apply, and `@evoscientist/webui` does not re-POST to
resume. When the connection drops mid-run, the front-end is left showing the
last partial token and never fetches the completed answer.
**Core insight:** the server already has the complete answer in the checkpoint.
The fix is to make the front-end fall back to that checkpoint when the stream
ends abnormally. This is small, surgical, and directly addresses the symptom.
## 2. Goal
**The user always sees the complete final answer in the WebUI, even if the SSE
stream drops mid-run.**
## 3. Non-Goals
- Seamless stream resume (tracking `last-event-id` / `seq` and replaying
buffered events). That is option β — better UX, ~3–5× the code, deferred.
- The TUI pivot / launch-mode decoupling (option γ). Separable; only relevant if
the port-conflict / orphan-process issues need addressing.
- Patching `langgraph_api` upstream.
## 4. Detection Logic — What Counts as "Truncated"
The langgraph SSE protocol emits terminal events when a run finishes
(success → `end`, failure → `error`; exact names to be confirmed against the
`@langchain/langgraph-sdk` version during implementation). The front-end's
streaming reader loop currently processes events until the `fetch()` ReadableStream
closes.
**Truncation** = the stream closed (reader returned) **without** a terminal event
having been received. Causes: browser tab throttled, network blip, proxy idle
timeout. On truncation, the in-progress assistant bubble is left showing a prefix
of the real answer.
## 5. Components
### F1 — Truncation detector (front-end)
In the streaming reader loop, track a `terminalSeen` flag. Set it when a
terminal event arrives. When the reader returns, if `!terminalSeen`, mark the
run `truncated` and trigger F2.
**File:** `@evoscientist/webui` — the run-stream consumer (the module that wraps
the langgraph-ts SDK's `runs.stream` and renders into the message list).
### F2 — Checkpoint fallback fetch (front-end)
On `truncated`, fetch the complete final assistant message and **replace** the
in-progress bubble's content with it (not append — the partial tokens are a
prefix of the same message).
Two implementation choices, pick one:
- **F2a (zero server change):** call the standard langgraph endpoint
`GET /threads/{thread_id}/state`, walk `values.messages` backwards to the last
`AIMessage`, join its content blocks. No new route, but the front-end parses
raw state and reimplements "find latest assistant text" logic.
- **F2b (recommended, uses S1 below):** call `GET /api/threads/{id}/final-answer`
→ `{content, completed_at, complete}`. Server owns the parsing; front-end
stays dumb.
Render a subtle affordance (e.g., a dim "⚠ stream dropped — recovered from
checkpoint" line above the message) so the user knows a disconnect happened and
the displayed text is authoritative, not stale streaming.
### F3 — Retry until the run settles (front-end)
The browser may drop while the run is **still executing** server-side. A single
fallback fetch at drop-time could return a not-yet-final message. So F2 retries
with backoff until either:
- `complete == true` (run finished — render and stop), or
- a max wait elapses (default 240 s — comfortably exceeds the observed 127 s
run even if the drop happens at t=0), then render whatever is latest and
stop retrying.
Backoff: poll at 1 s, 2 s, 4 s, then every 5 s up to the max. Each poll replaces
the bubble with the latest content, so if the run finishes mid-retry the user
sees it update to the final form.
### S1 — Convenience route `GET /api/threads/{id}/final-answer` (this repo, optional but recommended)
Enables F2b. Mounted on the existing custom ASGI app already wired via
`langgraph.json` → `EvoScientist.langgraph_dev.http:app` (currently only serves
`/api/models`).
**Contract:**
```
GET /api/threads/{thread_id}/final-answer
200 → { "content": str, # last AIMessage text, blocks joined
"completed_at": iso8601, # run end timestamp, or null
"complete": bool } # true iff run reached terminal state
404 → thread not found
```
**`complete` derivation** (run reached a terminal state):
`thread.status == "idle"` **or** latest checkpoint's `next` tuple is empty.
**File:** `EvoScientist/langgraph_dev/http.py` — add `get_final_answer` route
beside the existing `get_models`. Implementation reads the thread state via the
shared checkpointer (`AsyncSqliteSaver` at `~/.evoscientist/sessions.db`,
already used by `sessions.py`). If direct checkpointer access is awkward from
inside the mounted app, fall back to an internal `langgraph_sdk.get_client()`
call to `GET /threads/{id}/state` on localhost — the contract stays identical.
**Why prefer F2b+S1 over F2a:** the "find latest AIMessage + join content blocks
+ decide complete?" logic is non-trivial (content is a list of typed blocks,
including `reasoning` blocks that should be excluded from the rendered answer —
see the checkpoint: block[0] was `reasoning`, block[1] was `text`). Centralizing
it server-side keeps the front-end thin and makes the same logic reusable by the
TUI's future recovery path if γ is ever pursued.
## 6. Data Flow
```
browser fetch(/runs/stream, POST) ──SSE──▶ langgraph dev
│ │
│ (mid-run, connection drops) │ run continues to completion
▼ │
reader returns, no terminal event ──▶ truncated │
│ │
│ GET /api/threads/{id}/final-answer │
▼ ▼
http.py::get_final_answer ──read checkpoint──▶ sessions.db
│
▼
{content, complete} ──▶ if !complete, retry (F3); else render (F2)
```
## 7. Error Handling
| Case | Behavior |
|---|---|
| Stream ends normally (`end`/`error` received) | F1 sets `terminalSeen`; no fallback; existing render path unchanged |
| Stream drops, run already finished server-side | First fallback fetch returns `complete=true`; render immediately |
| Stream drops, run still executing | F3 retries with backoff; each retry renders latest; stops when `complete` or 120 s max |
| Thread unknown / deleted | 404; front-end shows "stream interrupted and the session could not be recovered" |
| Convenience route unreachable (older server without S1) | Front-end falls back to F2a (raw state endpoint) — degrade gracefully |
## 8. Testing
**Front-end (`@evoscientist/webui`):**
- Unit: feed the reader a synthetic event stream that closes without a terminal
event → assert `truncated == true`; feed one with `end` → `false`.
- Integration: mock `fetch` to drop mid-stream; assert the fallback fetch fires
and the bubble's content is replaced with the checkpoint value.
**Server (this repo, if S1):** new `tests/test_http_final_answer.py`
- Completed thread → 200, `complete=true`, content matches the last AIMessage
text (not the reasoning block).
- Mid-run thread (forced `busy` / non-empty `next`) → 200, `complete=false`.
- Unknown thread → 404.
- Content with mixed `reasoning` + `text` blocks → only `text` in `content`.
## 9. Cross-Repo Coordination & Rollout
- **Order:** land S1 in this Python repo first (route + tests, behind no flag —
it's a pure addition). Then F1/F2/F3 in `@evoscientist/webui`.
- **Versioning:** the WebUI launches via `npx @evoscientist/webui@latest`, so
the front-end fix reaches users on their next launch automatically once
published. No coordinated upgrade required on the Python side.
- **Backward compatibility:** if a user runs a new front-end against an older
Python server without S1, the front-end must detect the 404 and fall back to
F2a (raw state endpoint). If an old front-end runs against a new server, S1
is simply unused — no harm.
## 10. Open Questions / Future Work
- **Exact terminal-event names** for the SDK version in use (`end`/`error` vs.
`done`/`result`). Confirm against `@langchain/langgraph-sdk` during
implementation; the detector is parameterized on this.
- **True resume (β):** track `last-event-id` (or the V2 thread-stream `seq`) and
replay buffered events on reconnect — seamless, no visible "recovered" flash.
Larger front-end effort; defer until α is validated.
- **TUI pivot (γ):** separable work to make the in-process TUI the default and
eliminate the port-conflict / orphan-process issues. Independent spec if
pursued.
-4
View File
@@ -48,7 +48,6 @@ dependencies = [
[dependency-groups]
dev = [
"pytest>=8.0",
"pytest-asyncio>=1.0",
"pytest-cov>=5.0",
"pytest-timeout>=2.4",
"ruff>=0.5",
@@ -59,7 +58,6 @@ dev = [
[project.optional-dependencies]
dev = [
"pytest>=8.0",
"pytest-asyncio>=1.0",
"pytest-cov>=5.0",
"pytest-timeout>=2.4",
"ruff>=0.5",
@@ -119,8 +117,6 @@ EvoScientist = [
[tool.pytest.ini_options]
testpaths = ["tests"]
asyncio_mode = "auto"
asyncio_default_fixture_loop_scope = "function"
filterwarnings = [
"ignore::UserWarning:langchain_nvidia_ai_endpoints",
]
+26 -23
View File
@@ -1,10 +1,34 @@
"""Shared fixtures for EvoScientist tests."""
from pathlib import Path
import asyncio
import pytest
_NONEXISTENT_DOTENV = str(Path(__file__).with_name(".pytest-dotenv-does-not-exist"))
def run_async(coro):
"""Run an async coroutine safely, cancelling pending tasks before closing.
This prevents 'Event loop is closed' errors from asyncio.Queue cleanup
when tasks are still waiting on Queue.get() at teardown time.
"""
loop = asyncio.new_event_loop()
try:
return loop.run_until_complete(coro)
finally:
# Cancel all pending tasks so Queue getters don't raise on close
pending = asyncio.all_tasks(loop)
for task in pending:
task.cancel()
if pending:
loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True))
loop.run_until_complete(loop.shutdown_asyncgens())
loop.close()
@pytest.fixture(name="run_async")
def run_async_fixture():
"""Pytest fixture that exposes run_async as a callable for test functions."""
return run_async
@pytest.fixture(autouse=True)
@@ -168,24 +192,3 @@ def restore_model_passthrough_patch():
yield
finally:
_reset()
@pytest.fixture(autouse=True)
def _isolate_dotenv(monkeypatch):
"""Keep the developer's real .env out of the test environment.
``get_effective_config`` runs ``load_dotenv(find_dotenv(usecwd=True),
override=True)``, so any test that loads config injects the repo's
real .env into ``os.environ`` for the rest of the pytest process.
An empty-valued line like ``MINIMAX_BASE_URL=`` then makes
``os.environ.get(key, default)`` return "" instead of the default,
breaking unrelated tests later in the run (see issue #322).
Pointing ``find_dotenv`` at a fixed path that does not exist makes
``load_dotenv`` a no-op without creating a temporary directory for
every test.
"""
monkeypatch.setattr(
"EvoScientist.config.settings.find_dotenv",
lambda *args, **kwargs: _NONEXISTENT_DOTENV,
)
+15 -10
View File
@@ -9,6 +9,7 @@ from typing import Any
from unittest.mock import MagicMock
from EvoScientist.stream.events import stream_agent_events
from tests.conftest import run_async
async def async_iter(items: Iterable[Any]) -> AsyncIterator[Any]:
@@ -16,20 +17,24 @@ async def async_iter(items: Iterable[Any]) -> AsyncIterator[Any]:
yield item
async def collect_events(
def collect_events(
agent,
message: str = "hi",
thread_id: str = "t1",
):
"""Collect stream_agent_events output for tests."""
events = []
async for ev in stream_agent_events(
agent,
message,
thread_id,
):
events.append(ev)
return events
"""Collect stream_agent_events output for synchronous tests."""
async def _run():
events = []
async for ev in stream_agent_events(
agent,
message,
thread_id,
):
events.append(ev)
return events
return run_async(_run())
def protocol_event(
+19 -20
View File
@@ -10,17 +10,18 @@ from EvoScientist.channels.imessage.channel_rpc import (
)
from EvoScientist.channels.qq.channel import QQChannel, QQConfig
from EvoScientist.channels.signal.channel import SignalChannel, SignalConfig
from tests.conftest import run_async as _run
class TestEmailChannelSmoke:
async def test_start_raises_without_required_imap_settings(self):
def test_start_raises_without_required_imap_settings(self):
channel = EmailChannel(EmailConfig())
with pytest.raises(
ChannelError, match="imap_host and imap_username are required"
):
await channel.start()
_run(channel.start())
async def test_send_returns_false_when_smtp_not_ready(self):
def test_send_returns_false_when_smtp_not_ready(self):
channel = EmailChannel(EmailConfig())
msg = OutboundMessage(
channel="email",
@@ -28,16 +29,16 @@ class TestEmailChannelSmoke:
content="hello",
metadata={"chat_id": "user@example.com"},
)
assert await channel.send(msg) is False
assert _run(channel.send(msg)) is False
class TestSignalChannelSmoke:
async def test_start_raises_without_phone_number(self):
def test_start_raises_without_phone_number(self):
channel = SignalChannel(SignalConfig())
with pytest.raises(ChannelError, match="phone_number is required"):
await channel.start()
_run(channel.start())
async def test_send_returns_false_when_not_connected(self):
def test_send_returns_false_when_not_connected(self):
channel = SignalChannel(SignalConfig(phone_number="+123456789"))
msg = OutboundMessage(
channel="signal",
@@ -45,29 +46,27 @@ class TestSignalChannelSmoke:
content="hello",
metadata={"chat_id": "+123456789"},
)
assert await channel.send(msg) is False
assert _run(channel.send(msg)) is False
class TestQQChannelSmoke:
async def test_start_raises_when_sdk_missing(self, monkeypatch):
def test_start_raises_when_sdk_missing(self, monkeypatch):
from EvoScientist.channels.qq import channel as qq_module
monkeypatch.setattr(qq_module, "QQ_AVAILABLE", False)
channel = QQChannel(QQConfig(app_id="id", app_secret="secret"))
with pytest.raises(ChannelError, match="SDK not installed"):
await channel.start()
_run(channel.start())
async def test_start_raises_without_credentials_when_sdk_available(
self, monkeypatch
):
def test_start_raises_without_credentials_when_sdk_available(self, monkeypatch):
from EvoScientist.channels.qq import channel as qq_module
monkeypatch.setattr(qq_module, "QQ_AVAILABLE", True)
channel = QQChannel(QQConfig(app_id="", app_secret=""))
with pytest.raises(ChannelError, match="app_id and app_secret are required"):
await channel.start()
_run(channel.start())
async def test_send_returns_false_without_client(self):
def test_send_returns_false_without_client(self):
channel = QQChannel(QQConfig(app_id="id", app_secret="secret"))
msg = OutboundMessage(
channel="qq",
@@ -75,11 +74,11 @@ class TestQQChannelSmoke:
content="hello",
metadata={"chat_id": "openid"},
)
assert await channel.send(msg) is False
assert _run(channel.send(msg)) is False
class TestIMessageChannelSmoke:
async def test_start_wraps_rpc_bootstrap_error(self, monkeypatch):
def test_start_wraps_rpc_bootstrap_error(self, monkeypatch):
async def _broken_start(self):
raise RuntimeError("imsg not found")
@@ -88,9 +87,9 @@ class TestIMessageChannelSmoke:
monkeypatch.setattr(imessage_module.ImsgRpcClient, "start", _broken_start)
channel = IMessageChannelRpc(IMessageConfig())
with pytest.raises(ChannelError, match="Failed to start imsg"):
await channel.start()
_run(channel.start())
async def test_send_returns_false_without_rpc_client(self):
def test_send_returns_false_without_rpc_client(self):
channel = IMessageChannelRpc(IMessageConfig())
msg = OutboundMessage(
channel="imessage",
@@ -98,4 +97,4 @@ class TestIMessageChannelSmoke:
content="hello",
metadata={"chat_id": "+123456789"},
)
assert await channel.send(msg) is False
assert _run(channel.send(msg)) is False
-137
View File
@@ -1,137 +0,0 @@
def test_create_cli_agent_accepts_host_backend_and_memory_options(
monkeypatch, tmp_path
):
import EvoScientist.EvoScientist as agent_module
from EvoScientist.config.settings import EvoScientistConfig
calls = {}
workspace_backend = object()
chat_model = object()
class _CompositeBackend:
def __init__(self, *, default, routes):
calls["default_backend"] = default
calls["routes"] = routes
class _MemoryBackend:
def __init__(self, **kwargs):
calls["memory_backend_kwargs"] = kwargs
class _SkillsBackend:
def __init__(self, **kwargs):
calls["skills_backend_kwargs"] = kwargs
class _Agent:
def with_config(self, config):
calls["agent_config"] = config
return self
cfg = EvoScientistConfig(auto_approve=True, recursion_limit=321)
monkeypatch.setattr("deepagents.backends.CompositeBackend", _CompositeBackend)
monkeypatch.setattr("deepagents.create_deep_agent", lambda **kwargs: _Agent())
monkeypatch.setattr("EvoScientist.backends.MemoryFilesystemBackend", _MemoryBackend)
monkeypatch.setattr("EvoScientist.backends.MergedSkillsBackend", _SkillsBackend)
monkeypatch.setattr(agent_module, "set_active_workspace", lambda path: None)
monkeypatch.setattr(
agent_module,
"_get_default_middleware",
lambda **kwargs: calls.setdefault("middleware_kwargs", kwargs) or [],
)
monkeypatch.setattr(
agent_module,
"load_mcp_and_build_kwargs",
lambda *args, **kwargs: {"subagents": [{"name": "research"}]},
)
memory_dir = tmp_path / "memory"
result = agent_module.create_cli_agent(
workspace_dir=str(tmp_path / "workspace"),
checkpointer=object(),
config=cfg,
chat_model=chat_model,
workspace_backend=workspace_backend,
memory_dir=memory_dir,
tool_selector_threshold=8,
memory_max_inline_profile_chars=1000,
enable_subagents=False,
enable_background_execution=False,
)
assert isinstance(result, _Agent)
assert calls["default_backend"] is workspace_backend
assert calls["memory_backend_kwargs"] == {
"root_dir": str(memory_dir),
"virtual_mode": True,
}
assert calls["middleware_kwargs"]["memory_dir"] == str(memory_dir)
assert calls["middleware_kwargs"]["tool_selector_threshold"] == 8
assert calls["middleware_kwargs"]["memory_max_inline_profile_chars"] == 1000
assert calls["middleware_kwargs"]["enable_background_execution"] is False
assert calls["agent_config"] == {"recursion_limit": 321}
def test_create_cli_agent_installs_route_middleware_in_fixed_slot(monkeypatch, tmp_path):
import EvoScientist.EvoScientist as agent_module
from EvoScientist.config.settings import EvoScientistConfig
calls = {}
class _Middleware:
def __init__(self, name):
self.name = name
class _Backend:
def __init__(self, **_kwargs):
pass
class _CompositeBackend:
def __init__(self, **_kwargs):
pass
class _Agent:
def with_config(self, _config):
return self
default_chain = [
_Middleware("error_normalization"),
_Middleware("configurable_model"),
_Middleware("context_editing"),
_Middleware("tool_protocol_guard"),
]
route = _Middleware("gateway_route_fallback")
monkeypatch.setattr("deepagents.backends.CompositeBackend", _CompositeBackend)
monkeypatch.setattr("deepagents.create_deep_agent", lambda **_kwargs: _Agent())
monkeypatch.setattr("EvoScientist.backends.MemoryFilesystemBackend", _Backend)
monkeypatch.setattr("EvoScientist.backends.MergedSkillsBackend", _Backend)
monkeypatch.setattr(agent_module, "set_active_workspace", lambda _path: None)
def fake_default_middleware(**kwargs):
calls["middleware_kwargs"] = kwargs
return list(default_chain)
monkeypatch.setattr(agent_module, "_get_default_middleware", fake_default_middleware)
def fake_load(_backend, middleware, **_kwargs):
calls["middleware"] = middleware
return {"subagents": []}
monkeypatch.setattr(agent_module, "load_mcp_and_build_kwargs", fake_load)
agent_module.create_cli_agent(
workspace_dir=str(tmp_path),
checkpointer=object(),
config=EvoScientistConfig(auto_approve=True),
chat_model=object(),
workspace_backend=object(),
main_agent_route_middleware=route,
)
assert calls["middleware_kwargs"]["enable_legacy_model_fallback"] is False
assert [middleware.name for middleware in calls["middleware"][:5]] == [
"error_normalization",
"configurable_model",
"gateway_route_fallback",
"context_editing",
"tool_protocol_guard",
]
+146 -123
View File
@@ -3,7 +3,6 @@
from __future__ import annotations
import asyncio
import threading
import pytest
@@ -94,95 +93,88 @@ def _make_loader_fn(agent_value="AGENT", fail_with=None, capture=None):
return _loader
class _GatedThreadLoader:
"""Callable loader that blocks until tests explicitly release it."""
def __init__(self, agent_value="AGENT", progress_events=()):
self.agent_value = agent_value
self.progress_events = tuple(progress_events)
self.started = threading.Event()
self.release = threading.Event()
self.finished = threading.Event()
def __call__(self, *, on_mcp_progress=None):
self.started.set()
self.release.wait(timeout=1)
try:
if on_mcp_progress is not None:
for event in self.progress_events:
on_mcp_progress(*event)
return self.agent_value
finally:
self.finished.set()
async def _wait_for_event(event, timeout=1):
return await asyncio.to_thread(event.wait, timeout)
def _run(coro):
return asyncio.run(coro)
class TestBackgroundAgentLoaderStart:
async def test_start_creates_task_and_forwards_kwargs(self):
def test_start_creates_task_and_forwards_kwargs(self):
captured: dict = {}
loader = BackgroundAgentLoader(_make_loader_fn(capture=captured))
loader.start(workspace_dir="/ws", checkpointer="CK")
assert loader.task is not None
assert loader.is_pending
await loader.await_ready()
async def _go():
loader.start(workspace_dir="/ws", checkpointer="CK")
assert loader.task is not None
assert loader.is_pending
await loader.await_ready()
_run(_go())
assert captured["kwargs"][0] == {"workspace_dir": "/ws", "checkpointer": "CK"}
async def test_start_bumps_load_id(self):
def test_start_bumps_load_id(self):
loader = BackgroundAgentLoader(_make_loader_fn())
assert loader._load_id == 0
loader.start()
assert loader._load_id == 1
loader.start()
assert loader._load_id == 2
await loader.await_ready()
async def _go():
assert loader._load_id == 0
loader.start()
assert loader._load_id == 1
loader.start()
assert loader._load_id == 2
await loader.await_ready()
async def test_start_cancels_in_flight_prior_task(self):
blocking = _GatedThreadLoader("LATE")
_run(_go())
loader = BackgroundAgentLoader(blocking)
loader.start()
first_task = loader.task
assert first_task is not None
assert await _wait_for_event(blocking.started)
# Supersede immediately; asyncio.to_thread wrapper gets cancelled.
loader._loader_fn = _make_loader_fn("FRESH")
loader.start()
agent = await loader.await_ready()
assert agent == "FRESH"
blocking.release.set()
try:
await first_task
except asyncio.CancelledError:
pass
assert first_task.cancelled() or first_task.done()
def test_start_cancels_in_flight_prior_task(self):
import time
def _blocking(*, on_mcp_progress=None):
time.sleep(0.05)
return "LATE"
async def _go():
loader = BackgroundAgentLoader(_blocking)
loader.start()
first_task = loader.task
# Supersede immediately; asyncio.to_thread wrapper gets cancelled.
loader._loader_fn = _make_loader_fn("FRESH")
loader.start()
agent = await loader.await_ready()
assert agent == "FRESH"
# Let the first thread drain so its done callback (gated) fires.
await asyncio.sleep(0.1)
assert first_task.cancelled() or first_task.done()
_run(_go())
class TestBackgroundAgentLoaderCallbacks:
async def test_progress_hook_sees_events_in_order(self):
def test_progress_hook_sees_events_in_order(self):
events: list[tuple[str, str, str]] = []
loader = BackgroundAgentLoader(
_make_loader_fn(capture={}),
on_progress=lambda e, s, d: events.append((e, s, d)),
)
loader.start()
await loader.await_ready()
async def _go():
loader.start()
await loader.await_ready()
_run(_go())
assert events == [("start", "srv", ""), ("success", "srv", "1")]
async def test_stale_progress_events_are_dropped(self):
def test_stale_progress_events_are_dropped(self):
"""A progress event fired after a newer `start` must not reach the hook."""
slow_loader = _GatedThreadLoader(
"slow-agent", progress_events=[("success", "from-slow", "1")]
)
import time
seen: list[str] = []
# Loader 1 sleeps so its progress event fires AFTER load 2 starts.
def slow_loader(*, on_mcp_progress=None):
time.sleep(0.08)
if on_mcp_progress is not None:
on_mcp_progress("success", "from-slow", "1")
return "slow-agent"
def fast_loader(*, on_mcp_progress=None):
if on_mcp_progress is not None:
on_mcp_progress("success", "from-fast", "1")
@@ -192,32 +184,36 @@ class TestBackgroundAgentLoaderCallbacks:
slow_loader, on_progress=lambda e, s, d: seen.append(s)
)
loader.start()
assert await _wait_for_event(slow_loader.started)
# Loader 1 waits so its progress event fires AFTER load 2 starts.
loader._loader_fn = fast_loader
loader.start()
await loader.await_ready()
slow_loader.release.set()
assert await _wait_for_event(slow_loader.finished)
async def _go():
loader.start()
# Supersede before the slow thread's event fires.
await asyncio.sleep(0.01)
loader._loader_fn = fast_loader
loader.start()
await loader.await_ready()
# Let the superseded thread finish (its event is gated out).
await asyncio.sleep(0.1)
_run(_go())
assert "from-fast" in seen
assert "from-slow" not in seen
async def test_success_callback_fires_on_completion(self):
def test_success_callback_fires_on_completion(self):
got = []
loader = BackgroundAgentLoader(
_make_loader_fn("MY_AGENT"),
on_success=lambda a: got.append(a),
)
loader.start()
await loader.await_ready()
await asyncio.sleep(0) # let done-callback run
async def _go():
loader.start()
await loader.await_ready()
await asyncio.sleep(0) # let done-callback run
_run(_go())
assert got == ["MY_AGENT"]
async def test_failure_callback_fires_on_error(self):
def test_failure_callback_fires_on_error(self):
err = RuntimeError("load failed")
got_failures = []
got_successes = []
@@ -227,33 +223,40 @@ class TestBackgroundAgentLoaderCallbacks:
on_failure=lambda e: got_failures.append(e),
)
loader.start()
with pytest.raises(RuntimeError, match="load failed"):
await loader.await_ready()
await asyncio.sleep(0)
async def _go():
loader.start()
with pytest.raises(RuntimeError, match="load failed"):
await loader.await_ready()
await asyncio.sleep(0)
_run(_go())
assert got_failures == [err]
assert got_successes == []
class TestBackgroundAgentLoaderAwaitReady:
async def test_returns_cached_agent_without_reawaiting(self):
def test_returns_cached_agent_without_reawaiting(self):
captured: dict = {}
loader = BackgroundAgentLoader(_make_loader_fn("A", capture=captured))
loader.start()
assert await loader.await_ready() == "A"
assert await loader.await_ready() == "A"
async def _go():
loader.start()
assert await loader.await_ready() == "A"
assert await loader.await_ready() == "A"
_run(_go())
assert len(captured["kwargs"]) == 1
async def test_raises_if_started_not_called(self):
def test_raises_if_started_not_called(self):
loader = BackgroundAgentLoader(_make_loader_fn())
with pytest.raises(RuntimeError, match="before start"):
await loader.await_ready()
async def _go():
with pytest.raises(RuntimeError, match="before start"):
await loader.await_ready()
async def test_reraises_real_error_on_subsequent_awaits(self):
_run(_go())
def test_reraises_real_error_on_subsequent_awaits(self):
"""After a failure, ``await_ready`` must keep raising the real exception —
not the "before start()" sentinel — until ``start`` is called again."""
@@ -262,13 +265,16 @@ class TestBackgroundAgentLoaderAwaitReady:
loader = BackgroundAgentLoader(_fail)
loader.start()
with pytest.raises(RuntimeError, match="bad MCP config"):
await loader.await_ready()
with pytest.raises(RuntimeError, match="bad MCP config"):
await loader.await_ready()
async def _go():
loader.start()
with pytest.raises(RuntimeError, match="bad MCP config"):
await loader.await_ready()
with pytest.raises(RuntimeError, match="bad MCP config"):
await loader.await_ready()
async def test_needs_restart_flags_failed_load_for_retry(self):
_run(_go())
def test_needs_restart_flags_failed_load_for_retry(self):
calls = {"n": 0}
def flaky(*, on_mcp_progress=None):
@@ -279,14 +285,17 @@ class TestBackgroundAgentLoaderAwaitReady:
loader = BackgroundAgentLoader(flaky)
assert loader.needs_restart # never started
loader.start()
with pytest.raises(RuntimeError):
await loader.await_ready()
assert loader.needs_restart # failed, caller may retry
loader.start()
assert await loader.await_ready() == "SECOND"
assert not loader.needs_restart # success → no retry
async def _go():
assert loader.needs_restart # never started
loader.start()
with pytest.raises(RuntimeError):
await loader.await_ready()
assert loader.needs_restart # failed, caller may retry
loader.start()
assert await loader.await_ready() == "SECOND"
assert not loader.needs_restart # success → no retry
_run(_go())
class TestBackgroundAgentLoaderAdopt:
@@ -296,19 +305,26 @@ class TestBackgroundAgentLoaderAdopt:
assert loader.agent == "EXTERNAL"
assert not loader.is_pending
async def test_adopt_supersedes_in_flight_load(self):
def test_adopt_supersedes_in_flight_load(self):
"""A late background completion must not overwrite an adopted agent."""
slow_loader = _GatedThreadLoader("FROM_BACKGROUND")
import time
loader = BackgroundAgentLoader(slow_loader)
def _slow(*, on_mcp_progress=None):
time.sleep(0.08)
return "FROM_BACKGROUND"
loader.start()
assert await _wait_for_event(slow_loader.started)
loader.adopt("FROM_MODEL")
slow_loader.release.set()
assert await _wait_for_event(slow_loader.finished)
await asyncio.sleep(0)
assert loader.agent == "FROM_MODEL"
loader = BackgroundAgentLoader(_slow)
async def _go():
loader.start()
await asyncio.sleep(0.01)
loader.adopt("FROM_MODEL")
# Give the background thread time to finish and fire its
# done-callback; the generation token should make it a no-op.
await asyncio.sleep(0.1)
assert loader.agent == "FROM_MODEL"
_run(_go())
class TestBackgroundAgentLoaderIsPending:
@@ -316,22 +332,29 @@ class TestBackgroundAgentLoaderIsPending:
loader = BackgroundAgentLoader(_make_loader_fn())
assert not loader.is_pending
async def test_false_after_completion(self):
def test_false_after_completion(self):
loader = BackgroundAgentLoader(_make_loader_fn())
loader.start()
await loader.await_ready()
async def _go():
loader.start()
await loader.await_ready()
_run(_go())
assert not loader.is_pending
async def test_true_between_start_and_completion(self):
wait_loader = _GatedThreadLoader("ok")
def test_true_between_start_and_completion(self):
import time
loader = BackgroundAgentLoader(wait_loader)
def _wait_loader(*, on_mcp_progress=None):
time.sleep(0.05)
return "ok"
loader.start()
assert await _wait_for_event(wait_loader.started)
assert loader.is_pending
wait_loader.release.set()
await loader.await_ready()
assert not loader.is_pending
loader = BackgroundAgentLoader(_wait_loader)
async def _go():
loader.start()
assert loader.is_pending
await loader.await_ready()
assert not loader.is_pending
_run(_go())
+201 -129
View File
@@ -5,8 +5,6 @@ import queue
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import pytest
from EvoScientist.cli import async_notifier
from EvoScientist.cli.async_notifier import (
dedup_notifications,
@@ -30,6 +28,12 @@ def test_notification_dataclass_fields():
def test_notification_queue_is_module_level_fifo():
# Drain anything left over from other tests
while True:
try:
async_notifier._notification_queue.get_nowait()
except queue.Empty:
break
n1 = async_notifier.AsyncTaskNotification("a", "x", "success", "")
n2 = async_notifier.AsyncTaskNotification("b", "x", "success", "")
async_notifier._notification_queue.put(n1)
@@ -47,7 +51,7 @@ def _drain_queue(q):
return items
async def test_read_async_tasks_from_gateway_reads_state_values():
def test_read_async_tasks_from_gateway_reads_state_values(run_async):
gateway = FakeGraphGateway(
state_values={
"async_tasks": {
@@ -56,16 +60,18 @@ async def test_read_async_tasks_from_gateway_reads_state_values():
}
)
tasks = await async_notifier.read_async_tasks_from_gateway(
gateway,
GraphTarget(local_graph=MagicMock()),
"tid",
tasks = run_async(
async_notifier.read_async_tasks_from_gateway(
gateway,
GraphTarget(local_graph=MagicMock()),
"tid",
)
)
assert tasks == {"task-1": {"status": "success"}}
async def test_watcher_pushes_notification_on_stream_end():
def test_watcher_pushes_notification_on_stream_end(run_async):
# Stream yields one "values" chunk with the final state, then closes
final_state = {
"messages": [{"type": "ai", "content": "Quantum superposition is..."}]
@@ -81,7 +87,10 @@ async def test_watcher_pushes_notification_on_stream_end():
# runs.get is used to fetch terminal status when stream ends
client.runs.get = AsyncMock(return_value={"status": "success"})
await async_notifier.watch_run_and_notify(client, "thr-1", "run-1", "writing-agent")
_drain_all(async_notifier)
run_async(
async_notifier.watch_run_and_notify(client, "thr-1", "run-1", "writing-agent")
)
notifs = _drain_queue(async_notifier._notification_queue)
assert len(notifs) == 1
@@ -90,7 +99,7 @@ async def test_watcher_pushes_notification_on_stream_end():
assert notifs[0].status == "success"
async def test_watcher_pushes_error_status_on_stream_exception():
def test_watcher_pushes_error_status_on_stream_exception(run_async):
async def fake_stream(*a, **kw):
raise RuntimeError("network broken")
yield # unreachable; makes this an async generator
@@ -102,13 +111,14 @@ async def test_watcher_pushes_error_status_on_stream_exception():
return_value={"status": "error", "error": "network broken"}
)
await async_notifier.watch_run_and_notify(client, "thr-4", "run-4", "agentZ")
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thr-4", "run-4", "agentZ"))
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "error"
async def test_spawn_watcher_replaces_existing_for_same_thread():
def test_spawn_watcher_replaces_existing_for_same_thread(run_async):
"""A second spawn_watcher with the same thread_id cancels the old watcher
and registers the new one — supports update_async_task creating a new
run_id on the same thread_id."""
@@ -128,34 +138,43 @@ async def test_spawn_watcher_replaces_existing_for_same_thread():
client.runs.join_stream = fake_stream_long
client.runs.get = AsyncMock(return_value={"status": "success"})
# First spawn for thread X, run R1
t1 = async_notifier.spawn_watcher(client, "thr-X", "R1", "agent")
assert t1 is not None
assert async_notifier._watcher_by_thread["thr-X"] is t1
await asyncio.sleep(0.02) # let it start streaming
async def scenario():
# Clear all queues and the watcher registries
async_notifier._active_watchers.clear()
async_notifier._watcher_by_thread.clear()
_drain_all(async_notifier)
# Second spawn for SAME thread X, NEW run R2
t2 = async_notifier.spawn_watcher(client, "thr-X", "R2", "agent")
assert t2 is not None
assert t2 is not t1
assert async_notifier._watcher_by_thread["thr-X"] is t2
# First spawn for thread X, run R1
t1 = async_notifier.spawn_watcher(client, "thr-X", "R1", "agent")
assert t1 is not None
assert async_notifier._watcher_by_thread["thr-X"] is t1
await asyncio.sleep(0.02) # let it start streaming
# Old watcher should be cancelled
await asyncio.sleep(0.02)
assert t1.cancelled() or t1.done()
# Second spawn for SAME thread X, NEW run R2
t2 = async_notifier.spawn_watcher(client, "thr-X", "R2", "agent")
assert t2 is not None
assert t2 is not t1
assert async_notifier._watcher_by_thread["thr-X"] is t2
# Cleanup the new task too
t2.cancel()
try:
await t2
except asyncio.CancelledError:
pass
# Old watcher should be cancelled
await asyncio.sleep(0.02)
assert t1.cancelled() or t1.done()
# Cancelled watchers don't push notifications
assert _drain_one_queue_helper(async_notifier._notification_queue) == []
assert _drain_one_queue_helper(async_notifier._unrouted_queue) == []
for q in async_notifier._notifications_by_thread.values():
assert _drain_one_queue_helper(q) == []
# Cleanup the new task too
t2.cancel()
try:
await t2
except asyncio.CancelledError:
pass
# Cancelled watchers don't push notifications
assert _drain_one_queue_helper(async_notifier._notification_queue) == []
assert _drain_one_queue_helper(async_notifier._unrouted_queue) == []
if hasattr(async_notifier, "_notifications_by_thread"):
for q in async_notifier._notifications_by_thread.values():
assert _drain_one_queue_helper(q) == []
run_async(scenario())
# ============================================================================
@@ -305,6 +324,13 @@ def test_format_notification_lines_timeout_uses_warning_icon():
def test_drain_returns_all_pending_and_empties_queue():
"""drain_notifications pulls every pending notification and empties queue."""
# Clear the queue first
while True:
try:
async_notifier._notification_queue.get_nowait()
except queue.Empty:
break
# Add three notifications
for tid in ("a", "b", "c"):
async_notifier._notification_queue.put(
@@ -437,12 +463,17 @@ def test_dedup_preserves_order():
# ============================================================================
async def test_consume_notifications_calls_runner_with_batched_message():
def test_consume_notifications_calls_runner_with_batched_message(run_async):
"""When notifications arrive and agent is idle, consume_notifications fires
the supplied async runner once with the formatted batch message and notifs list."""
from EvoScientist.cli import async_notifier as an
# Set up two pending notifications, no dedup match
while True:
try:
an._notification_queue.get_nowait()
except queue.Empty:
break
an._notification_queue.put(an.AsyncTaskNotification("t1", "wA", "success", "", ""))
an._notification_queue.put(an.AsyncTaskNotification("t2", "wB", "success", "", ""))
@@ -455,15 +486,21 @@ async def test_consume_notifications_calls_runner_with_batched_message():
async def fake_state_reader() -> dict:
return {} # no dedup info
await an.consume_notifications(fake_runner, fake_state_reader)
run_async(an.consume_notifications(fake_runner, fake_state_reader))
assert "wA" in captured["text"]
assert "wB" in captured["text"]
assert len(captured["notifs"]) == 2
async def test_consume_notifications_no_op_when_queue_empty():
def test_consume_notifications_no_op_when_queue_empty(run_async):
from EvoScientist.cli import async_notifier as an
while True:
try:
an._notification_queue.get_nowait()
except queue.Empty:
break
called = False
async def fake_runner(text: str, notifs: list):
@@ -473,7 +510,7 @@ async def test_consume_notifications_no_op_when_queue_empty():
async def fake_state_reader():
return {}
await an.consume_notifications(fake_runner, fake_state_reader)
run_async(an.consume_notifications(fake_runner, fake_state_reader))
assert called is False
@@ -484,7 +521,7 @@ async def test_consume_notifications_no_op_when_queue_empty():
# ============================================================================
async def test_notification_consuming_flag_prevents_reentry():
def test_notification_consuming_flag_prevents_reentry(run_async):
"""The _notification_consuming guard prevents two overlapping consumers.
Verifies the flag contract used by _consume_notifications_tui:
@@ -499,6 +536,13 @@ async def test_notification_consuming_flag_prevents_reentry():
"""
from EvoScientist.cli import async_notifier as an
# Clear the queue
while True:
try:
an._notification_queue.get_nowait()
except queue.Empty:
break
state = {"inject_count": 0, "consuming": False}
async def counting_runner(text: str, notifs: list) -> None:
@@ -521,42 +565,45 @@ async def test_notification_consuming_flag_prevents_reentry():
n1 = an.AsyncTaskNotification("g1", "writing-agent", "success", "", "")
n2 = an.AsyncTaskNotification("g2", "data-agent", "success", "", "")
# Scenario 1: normal flow — flag cleared, second consumer runs fine.
await guarded_consume(n1)
assert state["inject_count"] == 1
assert state["consuming"] is False # finally ran
async def scenario():
# Scenario 1: normal flow — flag cleared, second consumer runs fine.
await guarded_consume(n1)
assert state["inject_count"] == 1
assert state["consuming"] is False # finally ran
state["inject_count"] = 0
await guarded_consume(n2)
assert state["inject_count"] == 1
assert state["consuming"] is False
state["inject_count"] = 0
await guarded_consume(n2)
assert state["inject_count"] == 1
assert state["consuming"] is False
# Scenario 2: flag pre-set (first consumer in-flight) → second bails.
state["inject_count"] = 0
state["consuming"] = True # simulate first consumer running
an._notification_queue.put(n1)
await guarded_consume(n1) # should be blocked immediately
assert state["inject_count"] == 0 # runner never called
state["consuming"] = False # cleanup
# Scenario 2: flag pre-set (first consumer in-flight) → second bails.
state["inject_count"] = 0
state["consuming"] = True # simulate first consumer running
an._notification_queue.put(n1)
await guarded_consume(n1) # should be blocked immediately
assert state["inject_count"] == 0 # runner never called
state["consuming"] = False # cleanup
# Scenario 3: exception in runner → flag still cleared by finally.
async def raising_runner(text: str, notifs: list) -> None:
raise RuntimeError("boom")
# Scenario 3: exception in runner → flag still cleared by finally.
async def raising_runner(text: str, notifs: list) -> None:
raise RuntimeError("boom")
async def guarded_consume_raising(notif):
if state["consuming"]:
return
state["consuming"] = True
try:
an._notification_queue.put(notif)
await an.consume_notifications(raising_runner, fake_state_reader)
except RuntimeError:
pass
finally:
state["consuming"] = False
async def guarded_consume_raising(notif):
if state["consuming"]:
return
state["consuming"] = True
try:
an._notification_queue.put(notif)
await an.consume_notifications(raising_runner, fake_state_reader)
except RuntimeError:
pass
finally:
state["consuming"] = False
await guarded_consume_raising(n2)
assert state["consuming"] is False # cleared despite exception
await guarded_consume_raising(n2)
assert state["consuming"] is False # cleared despite exception
run_async(scenario())
# ============================================================================
@@ -566,42 +613,33 @@ async def test_notification_consuming_flag_prevents_reentry():
def _drain_all(an_mod):
"""Drain every queue (per-thread + unrouted) so tests start clean."""
while True:
try:
an_mod._notification_queue.get_nowait()
except queue.Empty:
break
for q in list(an_mod._notifications_by_thread.values()):
if hasattr(an_mod, "_notification_queue"):
while True:
try:
q.get_nowait()
an_mod._notification_queue.get_nowait()
except queue.Empty:
break
if hasattr(an_mod, "_notifications_by_thread"):
for q in list(an_mod._notifications_by_thread.values()):
while True:
try:
q.get_nowait()
except queue.Empty:
break
if hasattr(an_mod, "_unrouted_queue"):
while True:
try:
an_mod._unrouted_queue.get_nowait()
except queue.Empty:
break
while True:
try:
an_mod._unrouted_queue.get_nowait()
except queue.Empty:
break
def _reset_notifier_state(an_mod):
_drain_all(an_mod)
an_mod._active_watchers.clear()
an_mod._watcher_by_thread.clear()
@pytest.fixture(autouse=True)
def _clean_async_notifier_state():
_reset_notifier_state(async_notifier)
yield
_reset_notifier_state(async_notifier)
async def test_consume_only_drains_matching_thread():
def test_consume_only_drains_matching_thread(run_async):
"""Notifications tagged with origin_cli_thread_id only drain when the
consumer is invoked with the matching current_thread_id."""
from EvoScientist.cli import async_notifier as an
_drain_all(an)
n_a = an.AsyncTaskNotification(
"tA", "writing-agent", "success", "", "", origin_cli_thread_id="threadA"
)
@@ -619,17 +657,21 @@ async def test_consume_only_drains_matching_thread():
async def state_reader() -> dict:
return {}
await an.consume_notifications(runner, state_reader, current_thread_id="threadA")
run_async(
an.consume_notifications(runner, state_reader, current_thread_id="threadA")
)
assert captured["runs"] == [["tA"]]
# B's notification should still be queued
assert an.has_pending_notifications("threadB")
_drain_all(an)
async def test_unrouted_notifications_drain_on_any_thread():
def test_unrouted_notifications_drain_on_any_thread(run_async):
"""Notifications without origin_cli_thread_id (legacy / direct put) drain
regardless of the current_thread_id arg."""
from EvoScientist.cli import async_notifier as an
_drain_all(an)
an._notification_queue.put(
an.AsyncTaskNotification("tU", "writing-agent", "success", "", "")
)
@@ -642,15 +684,19 @@ async def test_unrouted_notifications_drain_on_any_thread():
async def state_reader() -> dict:
return {}
await an.consume_notifications(runner, state_reader, current_thread_id="anything")
run_async(
an.consume_notifications(runner, state_reader, current_thread_id="anything")
)
assert [n.task_id for n in captured["notifs"]] == ["tU"]
_drain_all(an)
async def test_thread_switch_drains_pending():
def test_thread_switch_drains_pending(run_async):
"""Pending notifications for thread B are not delivered while consumer
asks for thread A; once consumer runs with thread B they drain."""
from EvoScientist.cli import async_notifier as an
_drain_all(an)
an._enqueue(
an.AsyncTaskNotification(
"tB", "writing-agent", "success", "", "", origin_cli_thread_id="threadB"
@@ -666,19 +712,25 @@ async def test_thread_switch_drains_pending():
return {}
# First consume in thread A → no drain, B's notif still queued
await an.consume_notifications(runner, state_reader, current_thread_id="threadA")
run_async(
an.consume_notifications(runner, state_reader, current_thread_id="threadA")
)
assert captured["runs"] == []
assert an.has_pending_notifications("threadB")
# Now switch to thread B → drains
await an.consume_notifications(runner, state_reader, current_thread_id="threadB")
run_async(
an.consume_notifications(runner, state_reader, current_thread_id="threadB")
)
assert captured["runs"] == [["tB"]]
_drain_all(an)
def test_has_pending_notifications_respects_routing():
"""has_pending_notifications returns true only for matching or unrouted."""
from EvoScientist.cli import async_notifier as an
_drain_all(an)
# Unrouted always counts
an._notification_queue.put(
an.AsyncTaskNotification("tU", "writing-agent", "success", "", "")
@@ -696,6 +748,7 @@ def test_has_pending_notifications_respects_routing():
assert an.has_pending_notifications("threadA") is True
assert an.has_pending_notifications("threadB") is False
assert an.has_pending_notifications() is False # no unrouted, no current_thread
_drain_all(an)
# ============================================================================
@@ -709,7 +762,7 @@ def test_has_pending_notifications_respects_routing():
# ============================================================================
async def test_watcher_reports_error_on_in_band_error_event():
def test_watcher_reports_error_on_in_band_error_event(run_async):
"""SSE error event in the stream → notification.status == 'error'."""
async def fake_stream(*a, **kw):
@@ -724,7 +777,8 @@ async def test_watcher_reports_error_on_in_band_error_event():
return_value={"status": "success"}
) # would mislead — should NOT be consulted
await async_notifier.watch_run_and_notify(client, "thrE", "rE", "agentE")
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrE", "rE", "agentE"))
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "error"
@@ -732,7 +786,7 @@ async def test_watcher_reports_error_on_in_band_error_event():
client.runs.get.assert_not_awaited()
async def test_watcher_clean_exit_with_runs_get_success_is_success():
def test_watcher_clean_exit_with_runs_get_success_is_success(run_async):
"""Clean stream exit + runs.get reports success → status=success."""
async def fake_stream(*a, **kw):
@@ -744,14 +798,15 @@ async def test_watcher_clean_exit_with_runs_get_success_is_success():
client.runs.join_stream = fake_stream
client.runs.get = AsyncMock(return_value={"status": "success"})
await async_notifier.watch_run_and_notify(client, "thrS", "rS", "agentS")
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrS", "rS", "agentS"))
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "success"
client.runs.get.assert_awaited_once()
async def test_watcher_clean_exit_with_runs_get_error_is_race_safe():
def test_watcher_clean_exit_with_runs_get_error_is_race_safe(run_async):
"""Clean stream exit + no in-band error event + runs.get returns 'error'
→ status=success (race-safe).
@@ -772,13 +827,14 @@ async def test_watcher_clean_exit_with_runs_get_error_is_race_safe():
client.runs.join_stream = fake_stream
client.runs.get = AsyncMock(return_value={"status": "error"})
await async_notifier.watch_run_and_notify(client, "thrS", "rS", "agentS")
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrS", "rS", "agentS"))
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "success"
async def test_watcher_clean_exit_with_runs_get_running_drops_notification():
def test_watcher_clean_exit_with_runs_get_running_drops_notification(run_async):
"""Reproduces the production bug: clean SSE close while run is still
actually running (HTTP keep-alive timeout under concurrency).
@@ -800,8 +856,11 @@ async def test_watcher_clean_exit_with_runs_get_running_drops_notification():
client.runs.join_stream = fake_stream
client.runs.get = AsyncMock(return_value={"status": "running"})
await async_notifier.watch_run_and_notify(
client, "thr-bug", "rB", "data-analysis-agent"
_drain_all(async_notifier)
run_async(
async_notifier.watch_run_and_notify(
client, "thr-bug", "rB", "data-analysis-agent"
)
)
# No notification should have been enqueued anywhere.
@@ -813,7 +872,7 @@ async def test_watcher_clean_exit_with_runs_get_running_drops_notification():
assert client.runs.get.await_count >= 1
async def test_watcher_unknown_status_treated_as_non_terminal():
def test_watcher_unknown_status_treated_as_non_terminal(run_async):
"""Future / unrecognized status values should trigger a re-join, not a
false-positive notification.
@@ -834,7 +893,8 @@ async def test_watcher_unknown_status_treated_as_non_terminal():
side_effect=[{"status": "queued"}, {"status": "success"}]
)
await async_notifier.watch_run_and_notify(client, "thrU", "rU", "agentU")
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrU", "rU", "agentU"))
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "success"
@@ -842,7 +902,7 @@ async def test_watcher_unknown_status_treated_as_non_terminal():
assert client.runs.get.await_count == 2
async def test_watcher_runs_get_persistent_failure_drops_notification(monkeypatch):
def test_watcher_runs_get_persistent_failure_drops_notification(run_async, monkeypatch):
"""If ``runs.get`` keeps raising, the watcher cannot verify terminal
state and MUST drop the notification rather than default to
``"success"`` — otherwise a transient server outage reintroduces the
@@ -861,20 +921,22 @@ async def test_watcher_runs_get_persistent_failure_drops_notification(monkeypatc
monkeypatch.setattr(async_notifier.asyncio, "sleep", _no_sleep)
await async_notifier.watch_run_and_notify(client, "thrG", "rG", "agentG")
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrG", "rG", "agentG"))
# No notification — watcher exhausted the reconnect budget. Check every
# queue routing could send to so a future routing change can't make this
# test silently false-pass.
assert _drain_one_queue_helper(async_notifier._unrouted_queue) == []
assert _drain_one_queue_helper(async_notifier._notification_queue) == []
for q in async_notifier._notifications_by_thread.values():
assert _drain_one_queue_helper(q) == []
if hasattr(async_notifier, "_notifications_by_thread"):
for q in async_notifier._notifications_by_thread.values():
assert _drain_one_queue_helper(q) == []
# 1 initial + _MAX_RECONNECT_ATTEMPTS retries = 11 calls total.
assert client.runs.get.await_count == async_notifier._MAX_RECONNECT_ATTEMPTS + 1
async def test_watcher_runs_get_transient_failure_recovers(monkeypatch):
def test_watcher_runs_get_transient_failure_recovers(run_async, monkeypatch):
"""A single ``runs.get`` failure followed by a successful response on
retry must produce a correct notification — verifies the bounded
retry path actually recovers from transient outages instead of just
@@ -895,14 +957,15 @@ async def test_watcher_runs_get_transient_failure_recovers(monkeypatch):
monkeypatch.setattr(async_notifier.asyncio, "sleep", _no_sleep)
await async_notifier.watch_run_and_notify(client, "thrT", "rT", "agentT")
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrT", "rT", "agentT"))
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "success"
assert client.runs.get.await_count == 2
async def test_watcher_re_joins_stream_until_terminal_status():
def test_watcher_re_joins_stream_until_terminal_status(run_async):
"""When runs.get returns 'running' on attempt N but a terminal status
on attempt N+1, the watcher re-joins, observes the terminal status,
and enqueues the notification correctly."""
@@ -917,7 +980,8 @@ async def test_watcher_re_joins_stream_until_terminal_status():
side_effect=[{"status": "running"}, {"status": "success"}]
)
await async_notifier.watch_run_and_notify(client, "thrR", "rR", "agentR")
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrR", "rR", "agentR"))
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "success"
@@ -931,12 +995,15 @@ async def test_watcher_re_joins_stream_until_terminal_status():
# ============================================================================
async def test_consume_notifications_propagates_inject_exception():
def test_consume_notifications_propagates_inject_exception(run_async):
"""If the run_message callback raises, consume_notifications propagates
the exception to the caller — pollers wrap it in try/except so the
poller task does not die."""
import pytest
from EvoScientist.cli import async_notifier as an
_drain_all(an)
an._notification_queue.put(
an.AsyncTaskNotification("tX", "writing-agent", "success", "", "")
)
@@ -948,10 +1015,11 @@ async def test_consume_notifications_propagates_inject_exception():
return {}
with pytest.raises(RuntimeError, match="kaboom"):
await an.consume_notifications(boom_runner, state_reader)
run_async(an.consume_notifications(boom_runner, state_reader))
_drain_all(an)
async def test_watcher_skips_notification_on_stream_fail_with_nonterminal_status():
def test_watcher_skips_notification_on_stream_fail_with_nonterminal_status(run_async):
"""When the SSE stream errors AND runs.get returns a non-terminal status
(e.g. ``pending`` because the run is still alive), the watcher must
NOT enqueue a notification — otherwise the user sees a confusing
@@ -967,13 +1035,15 @@ async def test_watcher_skips_notification_on_stream_fail_with_nonterminal_status
client.runs.join_stream = fake_stream
client.runs.get = AsyncMock(return_value={"status": "pending"})
await async_notifier.watch_run_and_notify(client, "thrP", "rP", "agentP")
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrP", "rP", "agentP"))
# No notification should have been enqueued in any queue.
assert _drain_one_queue_helper(async_notifier._unrouted_queue) == []
assert _drain_one_queue_helper(async_notifier._notification_queue) == []
for q in async_notifier._notifications_by_thread.values():
assert _drain_one_queue_helper(q) == []
if hasattr(async_notifier, "_notifications_by_thread"):
for q in async_notifier._notifications_by_thread.values():
assert _drain_one_queue_helper(q) == []
def _drain_one_queue_helper(q):
@@ -990,6 +1060,8 @@ def test_active_watchers_grace_filters_by_thread():
(otherwise consume_notifications grace period would block thread A by up
to 3s waiting for thread B's unrelated watchers to finish)."""
async_notifier._active_watchers.clear()
# Sentinel handles — only their identity matters here, not their type
handle_a = object()
handle_b = object()
+21 -22
View File
@@ -7,6 +7,7 @@ deepagents internals. It hooks into ``awrap_tool_call`` and only fires on
from __future__ import annotations
import asyncio
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
@@ -104,7 +105,7 @@ def _make_middleware():
return mw, fake_client
async def test_middleware_spawns_watcher_on_start_async_task():
def test_middleware_spawns_watcher_on_start_async_task():
"""A successful start_async_task tool call must spawn one watcher per task."""
from langgraph.types import Command
@@ -141,7 +142,7 @@ async def test_middleware_spawns_watcher_on_start_async_task():
)
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
result = await mw.awrap_tool_call(request, fake_handler)
result = asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert isinstance(result, Command)
assert spawn_calls == [
@@ -149,7 +150,7 @@ async def test_middleware_spawns_watcher_on_start_async_task():
]
async def test_middleware_spawns_watcher_on_update_async_task():
def test_middleware_spawns_watcher_on_update_async_task():
"""A successful update_async_task call must also spawn a (replacement) watcher."""
from langgraph.types import Command
@@ -182,7 +183,7 @@ async def test_middleware_spawns_watcher_on_update_async_task():
)
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
await mw.awrap_tool_call(request, fake_handler)
asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert len(spawn_calls) == 1
args, kwargs = spawn_calls[0]
@@ -194,7 +195,7 @@ async def test_middleware_spawns_watcher_on_update_async_task():
assert kwargs["origin_cli_thread_id"] == "cli-thread-A"
async def test_middleware_pre_cancels_old_watcher_on_update():
def test_middleware_pre_cancels_old_watcher_on_update():
"""update_async_task must cancel the existing watcher BEFORE invoking the handler.
Otherwise the new run interrupts the old run's stream, which closes
@@ -220,14 +221,14 @@ async def test_middleware_pre_cancels_old_watcher_on_update():
try:
with patch.object(async_notifier, "spawn_watcher"):
await mw.awrap_tool_call(request, fake_handler)
asyncio.run(mw.awrap_tool_call(request, fake_handler))
finally:
async_notifier._watcher_by_thread.pop("task-1", None)
assert cancel_observed_before_handler["value"] is True
async def test_middleware_passes_through_unrelated_tools():
def test_middleware_passes_through_unrelated_tools():
"""A non-launch tool call must not spawn any watcher and must return result unchanged."""
mw, _ = _make_middleware()
@@ -239,13 +240,13 @@ async def test_middleware_passes_through_unrelated_tools():
request = _build_request("ls", {"path": "/"}, thread_id="t")
with patch.object(async_notifier, "spawn_watcher") as mock_spawn:
result = await mw.awrap_tool_call(request, fake_handler)
result = asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert result is sentinel
assert mock_spawn.call_count == 0
async def test_middleware_handles_non_command_results_gracefully():
def test_middleware_handles_non_command_results_gracefully():
"""If the launch tool returns a string (validation error), no watcher is spawned."""
mw, _ = _make_middleware()
@@ -259,13 +260,13 @@ async def test_middleware_handles_non_command_results_gracefully():
)
with patch.object(async_notifier, "spawn_watcher") as mock_spawn:
result = await mw.awrap_tool_call(request, fake_handler)
result = asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert result == "Unknown async subagent type `bogus`"
assert mock_spawn.call_count == 0
async def test_middleware_origin_thread_id_is_none_when_runtime_config_missing():
def test_middleware_origin_thread_id_is_none_when_runtime_config_missing():
"""When runtime.config is empty, origin_cli_thread_id must be None (not crash)."""
from langgraph.types import Command
@@ -298,12 +299,12 @@ async def test_middleware_origin_thread_id_is_none_when_runtime_config_missing()
)
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
await mw.awrap_tool_call(request, fake_handler)
asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert captured.get("origin_cli_thread_id") is None
async def test_middleware_swallows_spawn_exceptions():
def test_middleware_swallows_spawn_exceptions():
"""spawn_watcher errors must not propagate up — middleware logs and continues."""
from langgraph.types import Command
@@ -335,7 +336,7 @@ async def test_middleware_swallows_spawn_exceptions():
with patch.object(async_notifier, "spawn_watcher", side_effect=boom):
# Should not raise.
result = await mw.awrap_tool_call(request, fake_handler)
result = asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert isinstance(result, Command)
@@ -355,9 +356,7 @@ async def test_middleware_swallows_spawn_exceptions():
),
],
)
async def test_middleware_picks_correct_prompt_field_per_tool(
tool_name, args, prompt_field
):
def test_middleware_picks_correct_prompt_field_per_tool(tool_name, args, prompt_field):
"""start_async_task uses 'description'; update_async_task uses 'message'."""
from langgraph.types import Command
@@ -386,12 +385,12 @@ async def test_middleware_picks_correct_prompt_field_per_tool(
request = _build_request(tool_name, args, thread_id="t")
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
await mw.awrap_tool_call(request, fake_handler)
asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert captured_prompt["value"] == prompt_field
async def test_middleware_prompt_field_is_tool_name_gated_not_fallback_chained():
def test_middleware_prompt_field_is_tool_name_gated_not_fallback_chained():
"""update_async_task with extra `description` arg must still use `message`.
Guards against the previous `args.get('description') or args.get('message')`
@@ -433,12 +432,12 @@ async def test_middleware_prompt_field_is_tool_name_gated_not_fallback_chained()
)
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
await mw.awrap_tool_call(request, fake_handler)
asyncio.run(mw.awrap_tool_call(request, fake_handler))
assert captured_prompt["value"] == "use this"
async def test_middleware_pre_cancel_swallows_unexpected_errors():
def test_middleware_pre_cancel_swallows_unexpected_errors():
"""A faulty old-watcher handle must not block the handler from running."""
from langgraph.types import Command
@@ -461,7 +460,7 @@ async def test_middleware_pre_cancel_swallows_unexpected_errors():
try:
with patch.object(async_notifier, "spawn_watcher"):
# Should not raise.
await mw.awrap_tool_call(request, fake_handler)
asyncio.run(mw.awrap_tool_call(request, fake_handler))
finally:
async_notifier._watcher_by_thread.pop("t1", None)
+7 -6
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
import asyncio
import json
from types import SimpleNamespace
@@ -916,16 +917,16 @@ class _AsyncFakeCrons:
return [{"cron_id": "cron-async"}]
async def test_alist_autoskill_schedules_uses_async_client_and_explicit_limit(
monkeypatch,
):
def test_alist_autoskill_schedules_uses_async_client_and_explicit_limit(monkeypatch):
crons = _AsyncFakeCrons()
client = SimpleNamespace(crons=crons)
monkeypatch.setattr("langgraph_sdk.get_client", lambda **_kwargs: client)
rows = await alist_autoskill_schedules(
EvoScientistConfig(),
limit=3,
rows = asyncio.run(
alist_autoskill_schedules(
EvoScientistConfig(),
limit=3,
)
)
assert rows == [{"cron_id": "cron-async"}]
+2 -2
View File
@@ -91,7 +91,7 @@ def test_stop_already_finished_is_graceful(tmp_path):
assert "already finished" in bg.stop(pid)
def test_exited_elapsed_is_frozen(tmp_path, monkeypatch):
def test_exited_elapsed_is_frozen(tmp_path):
"""Elapsed for an exited process freezes at its runtime, it must not keep growing."""
pid = bg.launch(_true_cmd(), str(tmp_path))
assert _wait_until(lambda: bg._PROCESSES[pid].finished_ts is not None)
@@ -99,7 +99,7 @@ def test_exited_elapsed_is_frozen(tmp_path, monkeypatch):
proc = bg._PROCESSES[pid]
assert proc.finished_ts is not None
first = bg._elapsed(proc)
monkeypatch.setattr(bg.time, "time", lambda: proc.finished_ts + 100.0)
time.sleep(1.1) # intentional: prove elapsed stays frozen, not ticking up
assert bg._elapsed(proc) == first
+403 -369
View File
@@ -12,6 +12,7 @@ import pytest
from EvoScientist.channels.bus.events import InboundMessage
from EvoScientist.channels.bus.message_bus import MessageBus
from EvoScientist.channels.channel_manager import ChannelManager
from tests.conftest import run_async as _run
from tests.fakes import QueueFakeChannel as FakeChannel
@@ -57,7 +58,7 @@ def clean_channel_state():
class TestBusInboundConsumer:
"""Test the _bus_inbound_consumer queue bridge."""
async def test_processes_inbound_and_publishes_outbound(self):
def test_processes_inbound_and_publishes_outbound(self):
"""InboundMessage -> queue -> response -> OutboundMessage flow."""
from EvoScientist.cli.channel import (
_bus_inbound_consumer,
@@ -67,51 +68,54 @@ class TestBusInboundConsumer:
_drain_queue(_message_queue)
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
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):
if not _message_queue.empty():
break
await asyncio.sleep(0.05)
# Wait for consumer to enqueue the message
for _ in range(20):
if not _message_queue.empty():
break
await asyncio.sleep(0.05)
msg = _message_queue.get_nowait()
assert msg.content == "hello agent"
assert msg.sender == "user1"
assert msg.channel_type == "fake"
msg = _message_queue.get_nowait()
assert msg.content == "hello agent"
assert msg.sender == "user1"
assert msg.channel_type == "fake"
# Simulate main-thread response
_set_channel_response(msg.msg_id, "Reply to: hello agent")
# Simulate main-thread response
_set_channel_response(msg.msg_id, "Reply to: hello agent")
outbound = await asyncio.wait_for(
bus.consume_outbound(),
timeout=2.0,
)
assert outbound.channel == "fake"
assert outbound.chat_id == "chat1"
assert "Reply to: hello agent" in outbound.content
outbound = await asyncio.wait_for(
bus.consume_outbound(),
timeout=2.0,
)
assert outbound.channel == "fake"
assert outbound.chat_id == "chat1"
assert "Reply to: hello agent" in outbound.content
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
async def test_no_response_fallback(self):
_run(_test())
def test_no_response_fallback(self):
"""Empty response is replaced with 'No response' fallback."""
from EvoScientist.cli.channel import (
_bus_inbound_consumer,
@@ -121,44 +125,47 @@ class TestBusInboundConsumer:
_drain_queue(_message_queue)
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
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():
break
await asyncio.sleep(0.05)
for _ in range(20):
if not _message_queue.empty():
break
await asyncio.sleep(0.05)
msg = _message_queue.get_nowait()
# Set empty response — falsy, so consumer falls back to "No response"
_set_channel_response(msg.msg_id, "")
msg = _message_queue.get_nowait()
# Set empty response — falsy, so consumer falls back to "No response"
_set_channel_response(msg.msg_id, "")
outbound = await asyncio.wait_for(
bus.consume_outbound(),
timeout=2.0,
)
assert outbound.content == "No response"
outbound = await asyncio.wait_for(
bus.consume_outbound(),
timeout=2.0,
)
assert outbound.content == "No response"
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
async def test_late_response_after_timeout_still_publishes(self, monkeypatch):
_run(_test())
def test_late_response_after_timeout_still_publishes(self, monkeypatch):
"""A response that arrives after the bridge timeout is still forwarded."""
from EvoScientist.cli import channel as channel_mod
from EvoScientist.cli.channel import (
@@ -172,53 +179,56 @@ class TestBusInboundConsumer:
_drain_queue(_message_queue)
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
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="slow request",
message_id="msg-123",
await bus.publish_inbound(
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="slow request",
message_id="msg-123",
)
)
)
for _ in range(20):
if not _message_queue.empty():
break
await asyncio.sleep(0.05)
for _ in range(20):
if not _message_queue.empty():
break
await asyncio.sleep(0.05)
msg = _message_queue.get_nowait()
msg = _message_queue.get_nowait()
notice = await asyncio.wait_for(
bus.consume_outbound(),
timeout=1.0,
)
assert "Still working on it" in notice.content
assert notice.reply_to == "msg-123"
notice = await asyncio.wait_for(
bus.consume_outbound(),
timeout=1.0,
)
assert "Still working on it" in notice.content
assert notice.reply_to == "msg-123"
_set_channel_response(msg.msg_id, "final answer")
_set_channel_response(msg.msg_id, "final answer")
outbound = await asyncio.wait_for(
bus.consume_outbound(),
timeout=1.0,
)
assert outbound.content == "final answer"
assert outbound.reply_to == "msg-123"
outbound = await asyncio.wait_for(
bus.consume_outbound(),
timeout=1.0,
)
assert outbound.content == "final answer"
assert outbound.reply_to == "msg-123"
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
async def test_late_timeout_keeps_active_request_cancellable(self, monkeypatch):
_run(_test())
def test_late_timeout_keeps_active_request_cancellable(self, monkeypatch):
"""Late timeout must not discard an active request's cancel scope."""
from EvoScientist.cli import channel as channel_mod
from EvoScientist.cli.channel import (
@@ -233,179 +243,191 @@ class TestBusInboundConsumer:
monkeypatch.setattr(channel_mod, "_RESPONSE_TIMEOUT", 0.05)
monkeypatch.setattr(channel_mod, "_LATE_RESPONSE_TIMEOUT", 0.05)
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
task = asyncio.create_task(
_handle_bus_message(
task = asyncio.create_task(
_handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="still running",
message_id="msg-active",
),
)
)
queued = None
for _ in range(20):
with _message_queue.mutex:
queued = _message_queue.queue[0] if _message_queue.queue else None
if queued is not None:
break
await asyncio.sleep(0.05)
assert queued is not None
assert _claim_channel_request(queued) is True
notice = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
assert "Still working on it" in notice.content
await task
assert _channel_request_state(queued.msg_id) == "active"
cancel_scope = _channel_message_cancel_scope(queued)
assert not display_mod.is_stream_cancel_requested(cancel_scope)
await _handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="still running",
message_id="msg-active",
content="/stop",
message_id="msg-stop-active",
),
)
)
queued = None
for _ in range(20):
with _message_queue.mutex:
queued = _message_queue.queue[0] if _message_queue.queue else None
if queued is not None:
break
await asyncio.sleep(0.05)
ack = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
assert ack.content == "Stopped."
assert ack.reply_to == "msg-stop-active"
assert display_mod.is_stream_cancel_requested(cancel_scope)
assert queued is not None
assert _claim_channel_request(queued) is True
_run(_test())
notice = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
assert "Still working on it" in notice.content
await task
assert _channel_request_state(queued.msg_id) == "active"
cancel_scope = _channel_message_cancel_scope(queued)
assert not display_mod.is_stream_cancel_requested(cancel_scope)
await _handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="/stop",
message_id="msg-stop-active",
),
)
ack = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
assert ack.content == "Stopped."
assert ack.reply_to == "msg-stop-active"
assert display_mod.is_stream_cancel_requested(cancel_scope)
async def test_cancelled_wait_cleans_pending_response(self):
def test_cancelled_wait_cleans_pending_response(self):
"""Cancelling a pending bus message should not leak its response slot."""
from EvoScientist.cli import channel as channel_mod
from EvoScientist.cli.channel import _handle_bus_message, _message_queue
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
task = asyncio.create_task(
_handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="cancel me",
),
task = asyncio.create_task(
_handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="cancel me",
),
)
)
)
for _ in range(20):
if not _message_queue.empty():
break
await asyncio.sleep(0.05)
for _ in range(20):
if not _message_queue.empty():
break
await asyncio.sleep(0.05)
queued = _message_queue.get_nowait()
with channel_mod._response_lock:
assert queued.msg_id in channel_mod._pending_responses
queued = _message_queue.get_nowait()
with channel_mod._response_lock:
assert queued.msg_id in channel_mod._pending_responses
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
with channel_mod._response_lock:
assert queued.msg_id not in channel_mod._pending_responses
with channel_mod._response_lock:
assert queued.msg_id not in channel_mod._pending_responses
async def test_consumer_shutdown_cleans_pending_response(self):
_run(_test())
def test_consumer_shutdown_cleans_pending_response(self):
"""Stopping the consumer should cancel late waits and clear state."""
from EvoScientist.cli import channel as channel_mod
from EvoScientist.cli.channel import _bus_inbound_consumer, _message_queue
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
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="slow shutdown",
await bus.publish_inbound(
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="slow shutdown",
)
)
)
for _ in range(20):
if not _message_queue.empty():
break
await asyncio.sleep(0.05)
for _ in range(20):
if not _message_queue.empty():
break
await asyncio.sleep(0.05)
queued = _message_queue.get_nowait()
with channel_mod._response_lock:
assert queued.msg_id in channel_mod._pending_responses
queued = _message_queue.get_nowait()
with channel_mod._response_lock:
assert queued.msg_id in channel_mod._pending_responses
consumer.cancel()
await consumer
consumer.cancel()
await consumer
with channel_mod._response_lock:
assert queued.msg_id not in channel_mod._pending_responses
with channel_mod._response_lock:
assert queued.msg_id not in channel_mod._pending_responses
async def test_stop_during_hitl_wait_releases_wait_and_acks(self):
_run(_test())
def test_stop_during_hitl_wait_releases_wait_and_acks(self):
"""`/stop` should wake pending HITL wait and publish immediate ack."""
from EvoScientist.cli import channel as channel_mod
from EvoScientist.cli.channel import _bus_inbound_consumer, _message_queue
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
hitl_event = channel_mod._register_hitl_wait("fake", "chat1")
consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager))
hitl_event = channel_mod._register_hitl_wait("fake", "chat1")
consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager))
await bus.publish_inbound(
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="/stop",
message_id="m-stop-1",
await bus.publish_inbound(
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="/stop",
message_id="m-stop-1",
)
)
)
for _ in range(20):
if hitl_event.is_set():
break
await asyncio.sleep(0.05)
assert hitl_event.is_set()
assert channel_mod._pop_hitl_reply("fake", "chat1") == "/stop"
for _ in range(20):
if hitl_event.is_set():
break
await asyncio.sleep(0.05)
assert hitl_event.is_set()
assert channel_mod._pop_hitl_reply("fake", "chat1") == "/stop"
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=2.0)
assert outbound.content == "Stopped."
assert outbound.reply_to == "m-stop-1"
assert _message_queue.empty()
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=2.0)
assert outbound.content == "Stopped."
assert outbound.reply_to == "m-stop-1"
assert _message_queue.empty()
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
async def test_stop_cancels_queued_request_before_main_thread_processes_it(self):
_run(_test())
def test_stop_cancels_queued_request_before_main_thread_processes_it(self):
"""`/stop` should cancel a queued request instead of only acking."""
from EvoScientist.cli import channel as channel_mod
from EvoScientist.cli.channel import (
@@ -414,70 +436,73 @@ class TestBusInboundConsumer:
_message_queue,
)
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
task = asyncio.create_task(
_handle_bus_message(
task = asyncio.create_task(
_handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="please work",
message_id="m-work-1",
),
)
)
queued = None
for _ in range(20):
with _message_queue.mutex:
queued = _message_queue.queue[0] if _message_queue.queue else None
if queued is not None:
break
await asyncio.sleep(0.05)
assert queued is not None
with channel_mod._response_lock:
assert queued.msg_id in channel_mod._pending_responses
await _handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="please work",
message_id="m-work-1",
content="/stop",
message_id="m-stop-2",
),
)
)
queued = None
for _ in range(20):
with _message_queue.mutex:
queued = _message_queue.queue[0] if _message_queue.queue else None
if queued is not None:
break
await asyncio.sleep(0.05)
with pytest.raises(asyncio.CancelledError):
await task
assert queued is not None
with channel_mod._response_lock:
assert queued.msg_id in channel_mod._pending_responses
skipped = _message_queue.get_nowait()
assert skipped.msg_id == queued.msg_id
assert _claim_or_complete_channel_request(skipped) is False
await _handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="/stop",
message_id="m-stop-2",
),
)
with channel_mod._response_lock:
assert queued.msg_id not in channel_mod._pending_responses
with channel_mod._channel_request_lock:
assert queued.msg_id not in channel_mod._channel_requests
assert queued.msg_id not in channel_mod._cancelled_channel_messages
assert "fake:chat1" not in channel_mod._session_requests
with pytest.raises(asyncio.CancelledError):
await task
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=2.0)
assert outbound.content == "Stopped."
assert outbound.reply_to == "m-stop-2"
with pytest.raises(asyncio.TimeoutError):
await asyncio.wait_for(bus.consume_outbound(), timeout=0.2)
skipped = _message_queue.get_nowait()
assert skipped.msg_id == queued.msg_id
assert _claim_or_complete_channel_request(skipped) is False
_run(_test())
with channel_mod._response_lock:
assert queued.msg_id not in channel_mod._pending_responses
with channel_mod._channel_request_lock:
assert queued.msg_id not in channel_mod._channel_requests
assert queued.msg_id not in channel_mod._cancelled_channel_messages
assert "fake:chat1" not in channel_mod._session_requests
outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=2.0)
assert outbound.content == "Stopped."
assert outbound.reply_to == "m-stop-2"
with pytest.raises(asyncio.TimeoutError):
await asyncio.wait_for(bus.consume_outbound(), timeout=0.2)
async def test_stop_leaves_resolved_response_available_for_delivery(self):
def test_stop_leaves_resolved_response_available_for_delivery(self):
"""`/stop` must not steal a response whose waiter already resolved."""
from EvoScientist.cli import channel as channel_mod
from EvoScientist.cli.channel import (
@@ -490,39 +515,42 @@ class TestBusInboundConsumer:
_set_channel_response,
)
msg = ChannelMessage(
msg_id="msg-resolved",
content="already answered",
sender="user1",
channel_type="fake",
metadata={},
channel_ref=None,
bus_ref=None,
chat_id="chat1",
message_id="m-resolved",
)
async def _test():
msg = ChannelMessage(
msg_id="msg-resolved",
content="already answered",
sender="user1",
channel_type="fake",
metadata={},
channel_ref=None,
bus_ref=None,
chat_id="chat1",
message_id="m-resolved",
)
waiter = _enqueue_channel_message(msg)
assert _claim_channel_request(msg) is True
waiter = _enqueue_channel_message(msg)
assert _claim_channel_request(msg) is True
_set_channel_response(msg.msg_id, "final answer")
assert await asyncio.wait_for(asyncio.shield(waiter), timeout=1.0) == (
"final answer"
)
_set_channel_response(msg.msg_id, "final answer")
assert await asyncio.wait_for(asyncio.shield(waiter), timeout=1.0) == (
"final answer"
)
cancelled_count, active_count = _cancel_channel_session("fake", "chat1")
assert cancelled_count == 0
assert active_count == 0
cancelled_count, active_count = _cancel_channel_session("fake", "chat1")
assert cancelled_count == 0
assert active_count == 0
with channel_mod._response_lock:
assert msg.msg_id in channel_mod._pending_responses
with channel_mod._channel_request_lock:
assert msg.msg_id not in channel_mod._cancelled_channel_messages
with channel_mod._response_lock:
assert msg.msg_id in channel_mod._pending_responses
with channel_mod._channel_request_lock:
assert msg.msg_id not in channel_mod._cancelled_channel_messages
assert _pop_channel_response(msg.msg_id) == "final answer"
_complete_channel_request(msg.msg_id)
assert _pop_channel_response(msg.msg_id) == "final answer"
_complete_channel_request(msg.msg_id)
async def test_message_counting(self):
_run(_test())
def test_message_counting(self):
"""Messages are counted via record_message."""
from EvoScientist.cli.channel import (
_bus_inbound_consumer,
@@ -532,42 +560,45 @@ class TestBusInboundConsumer:
_drain_queue(_message_queue)
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
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():
break
await asyncio.sleep(0.05)
for _ in range(20):
if not _message_queue.empty():
break
await asyncio.sleep(0.05)
msg = _message_queue.get_nowait()
_set_channel_response(msg.msg_id, "ok")
msg = _message_queue.get_nowait()
_set_channel_response(msg.msg_id, "ok")
await asyncio.wait_for(bus.consume_outbound(), timeout=2.0)
await asyncio.wait_for(bus.consume_outbound(), timeout=2.0)
assert manager._message_counts["fake"]["received"] == 1
assert manager._message_counts["fake"]["sent"] == 1
assert manager._message_counts["fake"]["received"] == 1
assert manager._message_counts["fake"]["sent"] == 1
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
async def test_channel_message_carries_metadata(self):
_run(_test())
def test_channel_message_carries_metadata(self):
"""ChannelMessage carries metadata, chat_id, and message_id."""
from EvoScientist.cli.channel import (
_bus_inbound_consumer,
@@ -577,46 +608,49 @@ class TestBusInboundConsumer:
_drain_queue(_message_queue)
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
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():
break
await asyncio.sleep(0.05)
for _ in range(20):
if not _message_queue.empty():
break
await asyncio.sleep(0.05)
msg = _message_queue.get_nowait()
assert msg.content == "with metadata"
assert msg.metadata == {"key": "value"}
assert msg.chat_id == "chat1"
assert msg.message_id == "msg-123"
assert msg.channel_ref is ch
msg = _message_queue.get_nowait()
assert msg.content == "with metadata"
assert msg.metadata == {"key": "value"}
assert msg.chat_id == "chat1"
assert msg.message_id == "msg-123"
assert msg.channel_ref is ch
_set_channel_response(msg.msg_id, "done")
_set_channel_response(msg.msg_id, "done")
outbound = await asyncio.wait_for(
bus.consume_outbound(),
timeout=2.0,
)
assert outbound.reply_to == "msg-123"
outbound = await asyncio.wait_for(
bus.consume_outbound(),
timeout=2.0,
)
assert outbound.reply_to == "msg-123"
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
_run(_test())
+2 -65
View File
@@ -6,8 +6,6 @@ from unittest.mock import MagicMock, patch
import pytest
from EvoScientist.ccproxy_manager import (
_CCPROXY_AUTH_TIMEOUT_SECONDS,
_CCPROXY_HEALTH_TIMEOUT_SECONDS,
check_ccproxy_auth,
ensure_ccproxy,
is_ccproxy_available,
@@ -17,7 +15,6 @@ from EvoScientist.ccproxy_manager import (
setup_codex_env,
start_ccproxy,
stop_ccproxy,
write_ccproxy_config,
)
# =============================================================================
@@ -55,8 +52,6 @@ class TestCheckCcproxyAuth:
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
assert cmd[1:] == ["auth", "status", "claude_api"]
# ccproxy CLI cold start takes ~10s; timeout must leave headroom
assert mock_run.call_args[1]["timeout"] == _CCPROXY_AUTH_TIMEOUT_SECONDS
@patch("subprocess.run")
def test_valid_auth_codex(self, mock_run):
@@ -128,10 +123,9 @@ class TestIsCcproxyRunning:
class TestStartCcproxy:
@patch("EvoScientist.ccproxy_manager.logger.warning")
@patch("EvoScientist.ccproxy_manager.is_ccproxy_running")
@patch("subprocess.Popen")
def test_success(self, mock_popen, mock_running, mock_warning):
def test_success(self, mock_popen, mock_running):
proc = MagicMock()
proc.poll.return_value = None
mock_popen.return_value = proc
@@ -140,11 +134,6 @@ class TestStartCcproxy:
result = start_ccproxy(8000)
assert result is proc
mock_warning.assert_called_once_with(
"Starting ccproxy on port %d; first startup may take up to %d seconds",
8000,
_CCPROXY_HEALTH_TIMEOUT_SECONDS,
)
@patch("EvoScientist.ccproxy_manager.is_ccproxy_running", return_value=False)
@patch("EvoScientist.ccproxy_manager.time")
@@ -154,11 +143,7 @@ class TestStartCcproxy:
proc.poll.return_value = None
mock_popen.return_value = proc
# Simulate time passing beyond deadline
mock_time.monotonic.side_effect = [
0,
0,
_CCPROXY_HEALTH_TIMEOUT_SECONDS + 1,
]
mock_time.monotonic.side_effect = [0, 0, 31]
mock_time.sleep = MagicMock()
with pytest.raises(RuntimeError, match="did not become healthy"):
@@ -169,54 +154,6 @@ class TestStartCcproxy:
with pytest.raises(FileNotFoundError):
start_ccproxy(8000)
@patch("EvoScientist.ccproxy_manager.is_ccproxy_running")
@patch("subprocess.Popen")
def test_passes_generated_config(self, mock_popen, mock_running, tmp_path):
proc = MagicMock()
proc.poll.return_value = None
mock_popen.return_value = proc
mock_running.side_effect = [True]
with patch("EvoScientist.config.get_config_dir", return_value=tmp_path):
start_ccproxy(8000)
cmd = mock_popen.call_args[0][0]
assert "--config" in cmd
assert cmd[cmd.index("--config") + 1] == str(tmp_path / "ccproxy.toml")
@patch("EvoScientist.ccproxy_manager.is_ccproxy_running")
@patch("EvoScientist.ccproxy_manager.write_ccproxy_config", side_effect=OSError)
@patch("subprocess.Popen")
def test_config_write_failure_starts_without_config(
self, mock_popen, mock_write, mock_running
):
proc = MagicMock()
proc.poll.return_value = None
mock_popen.return_value = proc
mock_running.side_effect = [True]
start_ccproxy(8000)
cmd = mock_popen.call_args[0][0]
assert "--config" not in cmd
# =============================================================================
# write_ccproxy_config
# =============================================================================
class TestWriteCcproxyConfig:
def test_writes_codex_mapping_override(self, tmp_path):
config_dir = tmp_path / "missing" / "config"
with patch("EvoScientist.config.get_config_dir", return_value=config_dir):
path = write_ccproxy_config()
assert path == str(config_dir / "ccproxy.toml")
content = (config_dir / "ccproxy.toml").read_text(encoding="utf-8")
assert "[plugins.codex]" in content
assert "model_mappings = []" in content
# =============================================================================
# ensure_ccproxy
+12 -10
View File
@@ -3,6 +3,8 @@
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from tests.conftest import run_async as _run
def _ctx():
from EvoScientist.commands.base import ChannelRuntime, CommandContext
@@ -53,7 +55,7 @@ class TestNeedsAgent:
class TestStartPath:
"""Start flow must propagate agent/thread_id globals."""
async def test_start_binds_channel_runtime(self):
def test_start_binds_channel_runtime(self):
from EvoScientist.commands.implementation.channel import ChannelCommand
ctx, _ui = _ctx()
@@ -75,11 +77,11 @@ class TestStartPath:
return_value=config,
),
):
await ChannelCommand().execute(ctx, ["telegram"])
_run(ChannelCommand().execute(ctx, ["telegram"]))
assert ctx.channel_runtime.agent is ctx.agent
assert ctx.channel_runtime.thread_id == "tid-42"
async def test_start_propagates_send_thinking(self):
def test_start_propagates_send_thinking(self):
"""send_thinking flag must reach _start_channels_bus_mode."""
from EvoScientist.commands.implementation.channel import ChannelCommand
@@ -109,14 +111,14 @@ class TestStartPath:
return_value=config,
),
):
await ChannelCommand().execute(ctx, ["telegram"])
_run(ChannelCommand().execute(ctx, ["telegram"]))
assert captured["agent"] is ctx.agent
assert captured["thread_id"] == "tid-42"
assert captured["send_thinking"] is False
class TestAddToRunningPath:
async def test_add_to_running_binds_channel_runtime(self):
def test_add_to_running_binds_channel_runtime(self):
from EvoScientist.commands.implementation.channel import ChannelCommand
ctx, _ui = _ctx()
@@ -138,11 +140,11 @@ class TestAddToRunningPath:
return_value=config,
),
):
await ChannelCommand().execute(ctx, ["discord"])
_run(ChannelCommand().execute(ctx, ["discord"]))
assert ctx.channel_runtime.agent is ctx.agent
assert ctx.channel_runtime.thread_id == "tid-42"
async def test_add_to_running_propagates_send_thinking(self):
def test_add_to_running_propagates_send_thinking(self):
"""Adding to a running bus must honor config.channel_send_thinking."""
from EvoScientist.commands.implementation.channel import ChannelCommand
@@ -171,13 +173,13 @@ class TestAddToRunningPath:
return_value=config,
),
):
await ChannelCommand().execute(ctx, ["discord"])
_run(ChannelCommand().execute(ctx, ["discord"]))
assert captured["channel_type"] == "discord"
assert captured["send_thinking"] is True
class TestStatusPath:
async def test_status_without_running_channels(self):
def test_status_without_running_channels(self):
from EvoScientist.commands.implementation.channel import ChannelCommand
ctx, ui = _ctx()
@@ -196,6 +198,6 @@ class TestStatusPath:
return_value=config,
),
):
await ChannelCommand().execute(ctx, ["status"])
_run(ChannelCommand().execute(ctx, ["status"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("No messaging channels" in m for m in msgs)
+9 -8
View File
@@ -8,6 +8,7 @@ import pytest
from EvoScientist.commands.channel_ui import ChannelCommandUI
from EvoScientist.gateway import ThreadStore
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway, FakeThreadStore
@@ -56,7 +57,7 @@ def _sent_text(bus_ref) -> str:
)
async def test_handle_session_resume_sends_history_back_to_channel_without_local_duplicate():
def test_handle_session_resume_sends_history_back_to_channel_without_local_duplicate():
callback = AsyncMock()
bus_ref = SimpleNamespace(publish_outbound=AsyncMock())
@@ -71,7 +72,7 @@ async def test_handle_session_resume_sends_history_back_to_channel_without_local
thread_store=thread_store,
)
await _run_resume(ui, "thread-42", "/workspace")
_run(_run_resume(ui, "thread-42", "/workspace"))
callback.assert_awaited_once_with("thread-42", "/workspace")
assert thread_store.calls == [("get_thread_messages", "thread-42")]
@@ -83,7 +84,7 @@ async def test_handle_session_resume_sends_history_back_to_channel_without_local
assert "EvoScientist: Here is the saved answer." in text
async def test_handle_session_resume_propagates_callback_abort_without_history():
def test_handle_session_resume_propagates_callback_abort_without_history():
callback = AsyncMock(side_effect=RuntimeError("workspace conflict"))
bus_ref = SimpleNamespace(publish_outbound=AsyncMock())
thread_store = FakeThreadStore()
@@ -94,7 +95,7 @@ async def test_handle_session_resume_propagates_callback_abort_without_history()
)
with pytest.raises(RuntimeError, match="workspace conflict"):
await _run_resume(ui, "thread-42", "/workspace")
_run(_run_resume(ui, "thread-42", "/workspace"))
callback.assert_awaited_once_with("thread-42", "/workspace")
assert thread_store.calls == []
@@ -102,7 +103,7 @@ async def test_handle_session_resume_propagates_callback_abort_without_history()
assert captured == []
async def test_handle_session_resume_reports_history_load_error():
def test_handle_session_resume_reports_history_load_error():
callback = AsyncMock()
bus_ref = SimpleNamespace(publish_outbound=AsyncMock())
ui, captured = _make_ui(
@@ -113,7 +114,7 @@ async def test_handle_session_resume_reports_history_load_error():
),
)
await _run_resume(ui, "thread-42", "/workspace")
_run(_run_resume(ui, "thread-42", "/workspace"))
callback.assert_awaited_once_with("thread-42", "/workspace")
assert captured == []
@@ -122,7 +123,7 @@ async def test_handle_session_resume_reports_history_load_error():
assert "history unavailable: db locked" in text
async def test_handle_session_resume_distinguishes_non_displayable_messages():
def test_handle_session_resume_distinguishes_non_displayable_messages():
bus_ref = SimpleNamespace(publish_outbound=AsyncMock())
ui, captured = _make_ui(
bus_ref=bus_ref,
@@ -131,7 +132,7 @@ async def test_handle_session_resume_distinguishes_non_displayable_messages():
),
)
await _run_resume(ui, "thread-42", "/workspace")
_run(_run_resume(ui, "thread-42", "/workspace"))
assert captured == [
"Resumed session: thread-42\nNo displayable messages in this session."
File diff suppressed because it is too large Load Diff
+58 -27
View File
@@ -11,6 +11,8 @@ from EvoScientist.channels.debug import (
emit_debug_event_if,
)
from .conftest import run_async
def test_debug_trace_enabled_from_bool():
assert debug_trace_enabled(True) is True
@@ -73,10 +75,10 @@ def _make_channel_context(*, debug_trace=True, name="test_channel"):
return {"channel": channel}
async def test_middleware_dedup_emits_structured_event(caplog):
def test_middleware_dedup_emits_structured_event(caplog):
from EvoScientist.channels.middleware import DedupMiddleware
with caplog.at_level(logging.DEBUG):
async def _run():
mw = DedupMiddleware()
ctx = _make_channel_context()
raw = _make_raw(message_id="dup1")
@@ -89,40 +91,49 @@ async def test_middleware_dedup_emits_structured_event(caplog):
caplog.clear()
result = await mw.process_inbound(raw, ctx)
assert result is None
with caplog.at_level(logging.DEBUG):
run_async(_run())
assert "middleware_dedup_drop" in caplog.text
assert "message_id=dup1" in caplog.text
async def test_middleware_allowlist_emits_structured_event(caplog):
def test_middleware_allowlist_emits_structured_event(caplog):
from EvoScientist.channels.middleware import AllowListMiddleware
with caplog.at_level(logging.DEBUG):
async def _run():
mw = AllowListMiddleware(allowed_senders={"allowed_user"})
ctx = _make_channel_context()
raw = _make_raw(sender_id="blocked_user")
result = await mw.process_inbound(raw, ctx)
assert result is None
with caplog.at_level(logging.DEBUG):
run_async(_run())
assert "middleware_allowlist_drop" in caplog.text
assert "reason=sender_not_allowed" in caplog.text
async def test_middleware_mention_gating_emits_structured_event(caplog):
def test_middleware_mention_gating_emits_structured_event(caplog):
from EvoScientist.channels.middleware import MentionGatingMiddleware
with caplog.at_level(logging.DEBUG):
async def _run():
mw = MentionGatingMiddleware(require_mention="group")
ctx = _make_channel_context()
raw = _make_raw(is_group=True, was_mentioned=False)
result = await mw.process_inbound(raw, ctx)
assert result is None
with caplog.at_level(logging.DEBUG):
run_async(_run())
assert "middleware_mention_drop" in caplog.text
assert "policy=group" in caplog.text
async def test_typing_manager_emits_trace_events(caplog):
def test_typing_manager_emits_trace_events(caplog):
from EvoScientist.channels.middleware import TypingManager
with caplog.at_level(logging.DEBUG):
async def _run():
send_action = AsyncMock(side_effect=RuntimeError("typing api down"))
mgr = TypingManager(
send_action,
@@ -133,14 +144,17 @@ async def test_typing_manager_emits_trace_events(caplog):
await mgr.start("chat1")
await asyncio.sleep(0.01)
await mgr.stop("chat1")
with caplog.at_level(logging.DEBUG):
run_async(_run())
assert "typing_error" in caplog.text
assert "chat_id=chat1" in caplog.text
async def test_ack_reaction_emits_error_traces(caplog):
def test_ack_reaction_emits_error_traces(caplog):
from EvoScientist.channels.middleware import AckReactionMiddleware
with caplog.at_level(logging.DEBUG):
async def _run():
send_fn = AsyncMock()
remove_fn = AsyncMock(side_effect=RuntimeError("remove failed"))
ack = AckReactionMiddleware(
@@ -164,13 +178,16 @@ async def test_ack_reaction_emits_error_traces(caplog):
send_fn.reset_mock()
send_fn.side_effect = RuntimeError("api down")
await ack.send_ack("chat2", "msg2")
with caplog.at_level(logging.DEBUG):
run_async(_run())
assert "ack_send_error" in caplog.text
assert "ack_remove_error" in caplog.text
assert "api down" in caplog.text
assert "remove failed" in caplog.text
async def test_inbound_raw_event_emitted(caplog):
def test_inbound_raw_event_emitted(caplog):
"""Integration-style: _enqueue_raw emits inbound_raw at the top."""
from EvoScientist.channels.base import Channel, RawIncoming
@@ -201,16 +218,19 @@ async def test_inbound_raw_event_emitted(caplog):
config.ack_scope = "off"
config.dedup_ttl = 3600
with caplog.at_level(logging.DEBUG):
async def _run():
with patch.object(Channel, "__abstractmethods__", set()):
ch = _TestChannel(config)
raw = RawIncoming(sender_id="u1", chat_id="c1", text="hi", message_id="m1")
await ch._enqueue_raw(raw)
with caplog.at_level(logging.DEBUG):
run_async(_run())
assert "inbound_raw" in caplog.text
assert "sender_id=u1" in caplog.text
async def test_format_fallback_emits_event(caplog):
def test_format_fallback_emits_event(caplog):
"""_send_with_format_fallback emits outbound_format_fallback on fallback."""
from EvoScientist.channels.base import Channel
@@ -248,10 +268,13 @@ async def test_format_fallback_emits_event(caplog):
if call_count == 1:
raise ValueError("parse error in formatted text")
with caplog.at_level(logging.DEBUG):
async def _run():
with patch.object(Channel, "__abstractmethods__", set()):
ch = _TestChannel(config)
await ch._send_with_format_fallback(_failing_send, "<b>hi</b>", "hi")
with caplog.at_level(logging.DEBUG):
run_async(_run())
assert "outbound_format_fallback" in caplog.text
assert call_count == 2
@@ -281,7 +304,7 @@ def test_trace_mixin_trace_event(caplog):
assert "key=val" in caplog.text
async def test_standalone_dispatcher_treats_false_send_as_error(caplog):
def test_standalone_dispatcher_treats_false_send_as_error(caplog):
from EvoScientist.channels.bus import MessageBus
from EvoScientist.channels.bus.events import OutboundMessage
from EvoScientist.channels.standalone import standalone_outbound_dispatcher
@@ -292,7 +315,7 @@ async def test_standalone_dispatcher_treats_false_send_as_error(caplog):
channel.send = AsyncMock(return_value=False)
bus = MessageBus()
with caplog.at_level(logging.DEBUG):
async def _run():
task = asyncio.create_task(standalone_outbound_dispatcher(bus, channel))
await bus.publish_outbound(
OutboundMessage(channel="test", chat_id="c1", content="hi")
@@ -303,11 +326,14 @@ async def test_standalone_dispatcher_treats_false_send_as_error(caplog):
await task
except asyncio.CancelledError:
pass
with caplog.at_level(logging.DEBUG):
run_async(_run())
assert "standalone_dispatch_error" in caplog.text
assert "send() returned False" in caplog.text
async def test_standalone_dispatcher_sends_media():
def test_standalone_dispatcher_sends_media():
from EvoScientist.channels.bus import MessageBus
from EvoScientist.channels.bus.events import OutboundMessage
from EvoScientist.channels.standalone import standalone_outbound_dispatcher
@@ -319,16 +345,21 @@ async def test_standalone_dispatcher_sends_media():
channel.send_media = AsyncMock(return_value=True)
bus = MessageBus()
task = asyncio.create_task(standalone_outbound_dispatcher(bus, channel))
await bus.publish_outbound(
OutboundMessage(channel="test", chat_id="c1", content="", media=["/tmp/a.png"])
)
await asyncio.sleep(0.05)
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
async def _run():
task = asyncio.create_task(standalone_outbound_dispatcher(bus, channel))
await bus.publish_outbound(
OutboundMessage(
channel="test", chat_id="c1", content="", media=["/tmp/a.png"]
)
)
await asyncio.sleep(0.05)
task.cancel()
try:
await task
except asyncio.CancelledError:
pass
run_async(_run())
channel.send_media.assert_awaited_once_with(
recipient="c1",
file_path="/tmp/a.png",
+186 -151
View File
@@ -14,6 +14,7 @@ from EvoScientist.cli.channel import (
from EvoScientist.cli.channel import (
dispatch_channel_slash_command as _dispatch_channel_slash_command,
)
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway, FakeThreadStore
@@ -42,31 +43,12 @@ def _make_msg(
)
async def test_non_slash_returns_false():
def test_non_slash_returns_false():
"""Plain text messages must fall through to the agent."""
msg = _make_msg(content="hello agent")
append = MagicMock()
handled = await dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
)
assert handled is False
append.assert_not_called()
async def test_unresolved_slash_returns_false():
"""Unknown slash commands must fall through (matches TUI behavior)."""
msg = _make_msg(content="/unknown-cmd")
append = MagicMock()
with patch(
"EvoScientist.commands.manager.manager.resolve",
return_value=None,
):
handled = await dispatch_channel_slash_command(
handled = _run(
dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
@@ -74,10 +56,33 @@ async def test_unresolved_slash_returns_false():
checkpointer=None,
append_system=append,
)
)
assert handled is False
append.assert_not_called()
def test_unresolved_slash_returns_false():
"""Unknown slash commands must fall through (matches TUI behavior)."""
msg = _make_msg(content="/unknown-cmd")
append = MagicMock()
with patch(
"EvoScientist.commands.manager.manager.resolve",
return_value=None,
):
handled = _run(
dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
)
)
assert handled is False
async def test_successful_slash_execution_sets_response_and_breadcrumb():
def test_successful_slash_execution_sets_response_and_breadcrumb():
"""Known slash command: cmd_manager.execute ran, helper returns True,
sends a confirmation to the channel user, and appends a local log line."""
msg = _make_msg()
@@ -95,13 +100,15 @@ async def test_successful_slash_execution_sets_response_and_breadcrumb():
) as mock_execute,
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
handled = await dispatch_channel_slash_command(
msg,
agent="fake-agent",
thread_id="t1",
workspace_dir="/tmp",
checkpointer=None,
append_system=append,
handled = _run(
dispatch_channel_slash_command(
msg,
agent="fake-agent",
thread_id="t1",
workspace_dir="/tmp",
checkpointer=None,
append_system=append,
)
)
assert handled is True
mock_execute.assert_awaited_once()
@@ -112,7 +119,7 @@ async def test_successful_slash_execution_sets_response_and_breadcrumb():
assert any("Executed command from" in t for t in breadcrumbs)
async def test_slash_dispatch_passes_graph_gateway_to_command_context():
def test_slash_dispatch_passes_graph_gateway_to_command_context():
msg = _make_msg()
fake_cmd = MagicMock()
fake_cmd.needs_agent.return_value = False
@@ -135,21 +142,23 @@ async def test_slash_dispatch_passes_graph_gateway_to_command_context():
),
patch("EvoScientist.cli.channel._set_channel_response"),
):
handled = await dispatch_channel_slash_command(
msg,
agent="fake-agent",
thread_id="t1",
workspace_dir="/tmp",
checkpointer=None,
append_system=append,
graph_gateway=graph_gateway,
handled = _run(
dispatch_channel_slash_command(
msg,
agent="fake-agent",
thread_id="t1",
workspace_dir="/tmp",
checkpointer=None,
append_system=append,
graph_gateway=graph_gateway,
)
)
assert handled is True
assert captured["graph_gateway"] is graph_gateway
async def test_needs_agent_awaits_loader_and_passes_result():
def test_needs_agent_awaits_loader_and_passes_result():
"""Commands with needs_agent=True must await the loader and the
resulting agent must flow through the CommandContext."""
msg = _make_msg()
@@ -173,14 +182,16 @@ async def test_needs_agent_awaits_loader_and_passes_result():
) as mock_execute,
patch("EvoScientist.cli.channel._set_channel_response"),
):
handled = await dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
await_agent_ready=_await_ready,
handled = _run(
dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
await_agent_ready=_await_ready,
)
)
assert handled is True
await_called.assert_called_once()
@@ -190,7 +201,7 @@ async def test_needs_agent_awaits_loader_and_passes_result():
assert ctx_arg.agent == "ready-agent"
async def test_await_agent_ready_failure_sets_error_response():
def test_await_agent_ready_failure_sets_error_response():
msg = _make_msg()
fake_cmd = MagicMock()
fake_cmd.needs_agent.return_value = True
@@ -206,14 +217,16 @@ async def test_await_agent_ready_failure_sets_error_response():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
handled = await dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
await_agent_ready=_await_ready,
handled = _run(
dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
await_agent_ready=_await_ready,
)
)
assert handled is True
mock_set_resp.assert_called_once()
@@ -222,7 +235,7 @@ async def test_await_agent_ready_failure_sets_error_response():
assert "agent blew up" in resp_text
async def test_cmd_manager_raises_returns_true_with_error():
def test_cmd_manager_raises_returns_true_with_error():
"""If cmd_manager.execute raises past its own try/except, the helper
must absorb it, return True, and report via _set_channel_response."""
msg = _make_msg()
@@ -240,13 +253,15 @@ async def test_cmd_manager_raises_returns_true_with_error():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
handled = await dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
handled = _run(
dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
)
)
assert handled is True
mock_set_resp.assert_called_once()
@@ -255,7 +270,7 @@ async def test_cmd_manager_raises_returns_true_with_error():
assert "boom" in resp_text
async def test_on_cmd_completed_awaited_with_ctx_original_agent_and_cmd():
def test_on_cmd_completed_awaited_with_ctx_original_agent_and_cmd():
"""After a successful slash execute, the on_cmd_completed hook must
be awaited with (ctx, original_agent, cmd) so Rich CLI can adopt an
``/model`` agent swap and refresh status for state-mutating commands."""
@@ -287,14 +302,16 @@ async def test_on_cmd_completed_awaited_with_ctx_original_agent_and_cmd():
),
patch("EvoScientist.cli.channel._set_channel_response"),
):
handled = await dispatch_channel_slash_command(
msg,
agent="original-agent",
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
on_cmd_completed=_on_completed,
handled = _run(
dispatch_channel_slash_command(
msg,
agent="original-agent",
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
on_cmd_completed=_on_completed,
)
)
assert handled is True
assert captured["ctx_agent"] == "swapped-agent"
@@ -302,7 +319,7 @@ async def test_on_cmd_completed_awaited_with_ctx_original_agent_and_cmd():
assert captured["cmd_name"] == "/model"
async def test_on_cmd_completed_receives_cmd_for_new_and_compact():
def test_on_cmd_completed_receives_cmd_for_new_and_compact():
"""``/new`` / ``/compact`` invoked via channel must flow the cmd into
the hook so the callback can still refresh status when the agent
didn't swap — mirrors REPL ``interactive.py:1027-1030``."""
@@ -326,19 +343,21 @@ async def test_on_cmd_completed_receives_cmd_for_new_and_compact():
),
patch("EvoScientist.cli.channel._set_channel_response"),
):
await dispatch_channel_slash_command(
_make_msg(content=cmd_name),
agent="same-agent",
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=_on_completed,
_run(
dispatch_channel_slash_command(
_make_msg(content=cmd_name),
agent="same-agent",
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=_on_completed,
)
)
assert captured["cmd_name"] == cmd_name, cmd_name
async def test_on_cmd_completed_skipped_on_fall_through_and_error():
def test_on_cmd_completed_skipped_on_fall_through_and_error():
"""The hook must NOT fire for unresolved slash, non-slash text, or
when cmd_manager.execute raised."""
fake_cmd = MagicMock()
@@ -350,14 +369,16 @@ async def test_on_cmd_completed_skipped_on_fall_through_and_error():
# Non-slash
with patch("EvoScientist.cli.channel._set_channel_response"):
await dispatch_channel_slash_command(
_make_msg(content="hi"),
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=_noop,
_run(
dispatch_channel_slash_command(
_make_msg(content="hi"),
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=_noop,
)
)
# Unresolved slash
with (
@@ -367,14 +388,16 @@ async def test_on_cmd_completed_skipped_on_fall_through_and_error():
),
patch("EvoScientist.cli.channel._set_channel_response"),
):
await dispatch_channel_slash_command(
_make_msg(content="/nope"),
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=_noop,
_run(
dispatch_channel_slash_command(
_make_msg(content="/nope"),
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=_noop,
)
)
# Execute raises
with (
@@ -388,20 +411,22 @@ async def test_on_cmd_completed_skipped_on_fall_through_and_error():
),
patch("EvoScientist.cli.channel._set_channel_response"),
):
await dispatch_channel_slash_command(
_make_msg(),
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=_noop,
_run(
dispatch_channel_slash_command(
_make_msg(),
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=_noop,
)
)
completed.assert_not_called()
async def test_command_error_skips_completion_hook_and_reports_error():
def test_command_error_skips_completion_hook_and_reports_error():
"""A command caught as failed by CommandManager must not look successful."""
msg = _make_msg(content="/resume abc")
fake_cmd = MagicMock()
@@ -425,14 +450,16 @@ async def test_command_error_skips_completion_hook_and_reports_error():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
handled = await dispatch_channel_slash_command(
msg,
agent=None,
thread_id="old-thread",
workspace_dir="/old-workspace",
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=completed,
handled = _run(
dispatch_channel_slash_command(
msg,
agent=None,
thread_id="old-thread",
workspace_dir="/old-workspace",
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=completed,
)
)
assert handled is True
@@ -440,7 +467,7 @@ async def test_command_error_skips_completion_hook_and_reports_error():
mock_set_resp.assert_called_once_with("msg-1", "Command error: workspace conflict")
async def test_empty_command_error_still_reports_error():
def test_empty_command_error_still_reports_error():
"""An empty string error is still a command failure sentinel."""
msg = _make_msg(content="/resume abc")
fake_cmd = MagicMock()
@@ -462,14 +489,16 @@ async def test_empty_command_error_still_reports_error():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
handled = await dispatch_channel_slash_command(
msg,
agent=None,
thread_id="old-thread",
workspace_dir="/old-workspace",
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=completed,
handled = _run(
dispatch_channel_slash_command(
msg,
agent=None,
thread_id="old-thread",
workspace_dir="/old-workspace",
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=completed,
)
)
assert handled is True
@@ -477,7 +506,7 @@ async def test_empty_command_error_still_reports_error():
mock_set_resp.assert_called_once_with("msg-1", "Command error: (no details)")
async def test_on_cmd_completed_exception_is_absorbed():
def test_on_cmd_completed_exception_is_absorbed():
"""A raising hook must NOT prevent the channel response from being set."""
msg = _make_msg()
fake_cmd = MagicMock()
@@ -497,21 +526,23 @@ async def test_on_cmd_completed_exception_is_absorbed():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
handled = await dispatch_channel_slash_command(
msg,
agent="orig",
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=_boom,
handled = _run(
dispatch_channel_slash_command(
msg,
agent="orig",
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=_boom,
)
)
assert handled is True
mock_set_resp.assert_called_once()
assert "Command executed" in mock_set_resp.call_args[0][1]
async def test_top_level_exception_is_absorbed():
def test_top_level_exception_is_absorbed():
"""Last-ditch safety net: if anything inside the dispatch pipeline
raises unexpectedly (lazy import failure, ChannelCommandUI ctor,
terminal I/O from append_system, ...), the helper must NOT
@@ -526,13 +557,15 @@ async def test_top_level_exception_is_absorbed():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
handled = await dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
handled = _run(
dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
)
)
assert handled is True
mock_set_resp.assert_called_once()
@@ -541,7 +574,7 @@ async def test_top_level_exception_is_absorbed():
assert "exploded during resolve" in resp_text
async def test_cmd_execute_returning_false_falls_through():
def test_cmd_execute_returning_false_falls_through():
"""When cmd_manager.execute returns False (empty/unparseable input),
the helper must return False so the caller falls through to the agent."""
msg = _make_msg(content="/")
@@ -559,13 +592,15 @@ async def test_cmd_execute_returning_false_falls_through():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
handled = await dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
handled = _run(
dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
)
)
assert handled is False
mock_set_resp.assert_not_called()
+7 -4
View File
@@ -1,5 +1,6 @@
"""Tests for CLI interactive UI backend dispatch."""
import asyncio
from types import SimpleNamespace
import pytest
@@ -100,7 +101,7 @@ def test_background_agent_server_starts_even_when_async_subagents_disabled(
assert calls == [(config, "/tmp/workspace")]
async def test_resume_workspace_sync_runs_even_when_async_subagents_disabled(
def test_resume_workspace_sync_runs_even_when_async_subagents_disabled(
monkeypatch,
):
import EvoScientist.cli.commands as cmds
@@ -116,9 +117,11 @@ async def test_resume_workspace_sync_runs_even_when_async_subagents_disabled(
)
config = SimpleNamespace(enable_async_subagents=False)
await cmds._sync_background_agent_server_workspace(
config,
workspace_dir="/tmp/resumed-workspace",
asyncio.run(
cmds._sync_background_agent_server_workspace(
config,
workspace_dir="/tmp/resumed-workspace",
)
)
assert calls == [(config, "/tmp/resumed-workspace")]
+1 -292
View File
@@ -1,5 +1,4 @@
"""Regression tests for the code_interpreter PTC allowlist and the
``EvoCodeInterpreterMiddleware`` subclass shape.
"""Regression tests for the code_interpreter PTC allowlist.
langchain-quickjs >=0.3 reserves the ``task`` sub-agent dispatch tool as the
top-level REPL global and raises ``ValueError`` if ``task`` appears in the
@@ -9,10 +8,7 @@ allowlist (``task()`` stays reachable as the REPL global, with responseSchema).
from __future__ import annotations
from unittest.mock import MagicMock
import pytest
from langchain_core.messages import AIMessage, HumanMessage
from EvoScientist.middleware.code_interpreter import (
_DEFAULT_PTC_ALLOWLIST,
@@ -49,290 +45,3 @@ def test_filter_tools_for_ptc_accepts_default_allowlist():
def test_create_code_interpreter_middleware_builds():
assert create_code_interpreter_middleware() is not None
def test_middleware_uses_thread_mode():
"""Upstream ``mode="thread"`` (the default) preserves cross-turn REPL
state as ``langchain-ai/deepagents#3064`` shipped it. The wire-cost
bloat that motivated the earlier ``mode="turn"`` regression guard is
fixed at the API serialization layer (``EvoFilteredGraph`` in
``EvoScientist/langgraph_dev/main_graph.py``), not by revoking the
persistence feature.
"""
mw = create_code_interpreter_middleware()
assert mw._mode == "thread"
def test_after_agent_evicts_slot_on_untouched_turn():
"""Regression guard against reintroducing a conditional-snapshot gate
that skips ``after_agent`` on untouched turns.
Upstream ``after_agent`` in ``langchain_quickjs/middleware.py`` performs
two things: snapshot the REPL AND evict the slot (``finally:
self._registry.evict(thread_id)``). ``before_agent`` restores the REPL
on any turn that follows a touched one via ``self._registry.get`` —
which is get-or-create. So if ``after_agent`` returns early without
evicting, one ``ThreadWorker`` + QuickJS Runtime leaks per persistent
``thread_id`` that ever went touched → quiet.
Fix: don't override ``after_agent`` / ``aafter_agent`` at all — inherit
upstream's unconditional snapshot+evict behavior. This test creates a
slot the way ``before_agent`` would, calls ``after_agent`` with an
untouched-state input, and asserts the slot was evicted.
"""
mw = create_code_interpreter_middleware()
tid = mw._fallback_thread_id
# Simulate the slot creation that ``before_agent`` performs when it sees
# a prior turn's snapshot payload in state.
mw._registry.get(tid)
assert len(mw._registry._slots) == 1
# Untouched-turn state: no ``code_interpreter`` tool call between the
# last ``HumanMessage`` and end. Under the earlier buggy gate this
# returned ``{}`` without evicting — leaking the slot created above.
untouched_state = {
"_quickjs_snapshot_payload": b"payload-from-prior-turn",
"messages": [
HumanMessage(content="thanks"),
AIMessage(content="you're welcome"),
],
}
mw.after_agent(untouched_state, runtime=None)
assert len(mw._registry._slots) == 0, (
"after_agent must evict the slot even on untouched turns, because "
"before_agent already restored a REPL that owns a ThreadWorker + "
"QuickJS Runtime. Skipping eviction leaks those resources."
)
def test_evo_filtered_graph_strips_private_snapshot_field():
"""The ``StateSnapshot`` returned by ``EvoScientist_agent.get_state`` must
not contain ``_quickjs_snapshot_payload`` in either ``values`` (the
materialized channel payload) or ``metadata['writes']`` (the raw write
records surfaced by ``get_state_history``).
"""
from EvoScientist.langgraph_dev.main_graph import _EvoFilteredGraph, _strip_private
snap = MagicMock()
snap.values = {
"messages": ["m1"],
"_quickjs_snapshot_payload": b"x" * 100,
"skills_metadata": [],
}
snap.metadata = {
"source": "loop",
"step": 42,
"writes": {
"CodeInterpreterMiddleware.after_agent": {
"_quickjs_snapshot_payload": ("snap", b"y" * 1_400_000),
"messages": [],
},
"model": {"messages": ["m1"]},
},
"parents": {},
}
_strip_private(snap)
snap._replace.assert_called_once()
kwargs = snap._replace.call_args.kwargs
assert "_quickjs_snapshot_payload" not in kwargs["values"]
assert "messages" in kwargs["values"]
assert "skills_metadata" in kwargs["values"]
scrubbed_writes = kwargs["metadata"]["writes"]
assert (
"_quickjs_snapshot_payload"
not in scrubbed_writes["CodeInterpreterMiddleware.after_agent"]
)
assert "messages" in scrubbed_writes["CodeInterpreterMiddleware.after_agent"]
assert scrubbed_writes["model"] == {"messages": ["m1"]}
# Non-writes metadata keys are preserved.
assert kwargs["metadata"]["source"] == "loop"
assert kwargs["metadata"]["step"] == 42
# Sanity: the class exists and inherits from CompiledStateGraph.
from langgraph.graph.state import CompiledStateGraph
assert issubclass(_EvoFilteredGraph, CompiledStateGraph)
def test_strip_private_handles_missing_metadata_writes():
"""``metadata['writes']`` can be missing or ``None`` on some snapshots
(e.g. initial state). The filter must not crash and must still strip
values.
"""
from EvoScientist.langgraph_dev.main_graph import _strip_private
snap = MagicMock()
snap.values = {"_quickjs_snapshot_payload": b"x", "messages": []}
snap.metadata = {"source": "input", "step": -1, "writes": None}
snap.tasks = ()
_strip_private(snap)
kwargs = snap._replace.call_args.kwargs
assert "_quickjs_snapshot_payload" not in kwargs["values"]
# writes was None, metadata passes through unchanged.
assert kwargs["metadata"]["writes"] is None
def test_strip_private_scrubs_task_result_snapshot_blob():
"""``tasks[*].result`` is where ``after_agent``'s return dict lands.
When the middleware snapshots, ``result`` carries
``{"_quickjs_snapshot_payload": ("snap", ~1.4 MB bytes)}``. Verified
on live history: this is the dominant per-response leak, larger than
``values`` and ``metadata.writes`` combined for anchor checkpoints.
"""
from EvoScientist.langgraph_dev.main_graph import _strip_private
class FakeTask:
def __init__(self, id_, result):
self.id = id_
self.name = "CodeInterpreterMiddleware.after_agent"
self.result = result
def _replace(self, **kwargs):
for k, v in kwargs.items():
setattr(self, k, v)
return self
leaking_task = FakeTask(
"t1", {"_quickjs_snapshot_payload": ("snap", b"z" * 1_400_000), "messages": []}
)
clean_task = FakeTask("t2", {"messages": ["hi"]})
snap = MagicMock()
snap.values = {}
snap.metadata = {"source": "loop", "step": 5}
snap.tasks = (leaking_task, clean_task)
_strip_private(snap)
kwargs = snap._replace.call_args.kwargs
tasks_after = kwargs["tasks"]
assert "_quickjs_snapshot_payload" not in tasks_after[0].result
assert "messages" in tasks_after[0].result
# Clean task is passed through untouched.
assert tasks_after[1] is clean_task
def test_agent_uses_filtered_graph_class():
"""The ``__class__`` swap in ``main_graph.py`` is the load-bearing wiring
that makes ``_strip_private`` reach the langgraph-api endpoints.
``_strip_private`` and ``_EvoFilteredGraph`` in isolation don't prove the
swap ran; every other test in this file passes even if someone drops the
swap line. This asserts the compiled agent is actually the filtered
subclass at module-load time, and that the subclass survives
``Pregel.copy(update=...)`` — the call langgraph-api makes in
``get_graph`` before yielding the graph to endpoint handlers.
"""
from EvoScientist.langgraph_dev.main_graph import (
EvoScientist_agent,
_EvoFilteredGraph,
)
assert isinstance(EvoScientist_agent, _EvoFilteredGraph)
assert isinstance(EvoScientist_agent.copy(update={}), _EvoFilteredGraph)
def test_all_registered_graphs_use_filtered_graph_class():
"""Every graph registered in ``langgraph.json`` (main + all subagents)
gets the ``__class__`` swap via ``_apply_filter_to_all_registered_graphs``.
Iterating the config directly matches the auto-detect refactor: adding
a new subagent to ``langgraph.json`` should not require a corresponding
test update.
Subagents get ``create_code_interpreter_middleware`` unconditionally
(``EvoScientist.py:_build_middleware_stack``), so they can touch the
QuickJS REPL and write ``_quickjs_snapshot_payload`` on their own
checkpoint namespace. Async subagents also get their own ``thread_id``
and their ``/threads/{id}/state`` endpoint runs on their own compiled
graph — without the swap on those graphs, our filter would miss that
endpoint entirely.
"""
import json
from importlib import import_module
from pathlib import Path
# Import triggers ``main_graph``'s swap loop.
from EvoScientist.langgraph_dev import main_graph
from EvoScientist.langgraph_dev.main_graph import _EvoFilteredGraph
config_path = Path(main_graph.__file__).parent / "langgraph.json"
config = json.loads(config_path.read_text())
for name, path in config["graphs"].items():
module_path, attr = path.rsplit(":", 1)
graph = getattr(import_module(module_path), attr)
assert isinstance(graph, _EvoFilteredGraph), (
f"graph {name!r} ({path}) did not receive the class swap"
)
def test_strip_private_recurses_into_nested_subgraph_state():
"""When ``subgraphs=True``, ``PregelTask.state`` holds a nested
``StateSnapshot`` for the subgraph. Its ``values`` (and its own nested
tasks) can carry ``_quickjs_snapshot_payload`` just like the parent.
Recursion covers the compound leak path CodeRabbit flagged.
"""
from langgraph.types import StateSnapshot
from EvoScientist.langgraph_dev.main_graph import _strip_private
nested_snap = StateSnapshot(
values={"_quickjs_snapshot_payload": b"n" * 1_400_000, "messages": []},
next=(),
config={},
metadata={"source": "loop", "step": 3},
created_at="2026-07-01T12:00:00Z",
parent_config=None,
tasks=(),
interrupts=(),
)
class FakeTask:
def __init__(self, state):
self.id = "sub-1"
self.name = "subgraph"
self.result = None
self.state = state
def _replace(self, **kwargs):
for k, v in kwargs.items():
setattr(self, k, v)
return self
task_with_nested = FakeTask(nested_snap)
task_with_config_state = FakeTask({"configurable": {"thread_id": "t"}})
snap = MagicMock()
snap.values = {}
snap.metadata = {"source": "loop", "step": 5}
snap.tasks = (task_with_nested, task_with_config_state)
_strip_private(snap)
kwargs = snap._replace.call_args.kwargs
tasks_after = kwargs["tasks"]
# Nested StateSnapshot got recursively scrubbed.
assert "_quickjs_snapshot_payload" not in tasks_after[0].state.values
assert "messages" in tasks_after[0].state.values
# A dict (RunnableConfig-shaped) state passes through unchanged — we only
# recurse into ``StateSnapshot`` instances.
assert tasks_after[1].state == {"configurable": {"thread_id": "t"}}
def test_strip_private_scrubs_delta_counters():
"""``metadata['counters_since_delta_snapshot']`` is a small
``{channel: [count, superstep]}`` bookkeeping map. Not a size problem,
but leaks the channel name — strip for consistency with the private
annotation.
"""
from EvoScientist.langgraph_dev.main_graph import _strip_private
snap = MagicMock()
snap.values = {}
snap.metadata = {
"source": "loop",
"step": 5,
"counters_since_delta_snapshot": {
"_quickjs_snapshot_payload": [1, 14],
"messages": [3, 14],
},
}
snap.tasks = ()
_strip_private(snap)
kwargs = snap._replace.call_args.kwargs
counters = kwargs["metadata"]["counters_since_delta_snapshot"]
assert "_quickjs_snapshot_payload" not in counters
assert "messages" in counters
+27 -24
View File
@@ -4,12 +4,13 @@ from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
from EvoScientist.gateway import GraphTarget
from tests.conftest import run_async as _run
from tests.fakes import FakeCommandUI, FakeGraphGateway
_TARGET = GraphTarget()
async def _compact(
def _compact(
graph_gateway: FakeGraphGateway,
*,
thread_id: str = "tid-1",
@@ -17,28 +18,30 @@ async def _compact(
):
from EvoScientist.cli.commands import compact_conversation
return await compact_conversation(
graph_gateway=graph_gateway,
thread_id=thread_id,
target=_TARGET,
input_tokens_hint=input_tokens_hint,
return _run(
compact_conversation(
graph_gateway=graph_gateway,
thread_id=thread_id,
target=_TARGET,
input_tokens_hint=input_tokens_hint,
)
)
class TestCompactGuards:
"""Guard conditions that return early without touching the middleware."""
async def test_empty_messages(self):
def test_empty_messages(self):
graph_gateway = FakeGraphGateway(state_values={"messages": []})
result = await _compact(graph_gateway)
result = _compact(graph_gateway)
assert result.status == "noop"
assert "no messages" in result.message
async def test_state_read_failure(self):
def test_state_read_failure(self):
graph_gateway = FakeGraphGateway(state_error=RuntimeError("DB gone"))
result = await _compact(graph_gateway)
result = _compact(graph_gateway)
assert result.status == "error"
assert "Failed to read state" in result.message
@@ -46,7 +49,7 @@ class TestCompactGuards:
class TestCompactCutoffZero:
"""When cutoff == 0, conversation is within retention budget."""
async def test_nothing_to_compact_short_conversation(self):
def test_nothing_to_compact_short_conversation(self):
msgs = [MagicMock() for _ in range(3)]
graph_gateway = FakeGraphGateway(state_values={"messages": msgs})
@@ -76,7 +79,7 @@ class TestCompactCutoffZero:
return_value=500,
),
):
result = await _compact(graph_gateway)
result = _compact(graph_gateway)
assert result.status == "noop"
assert "within the retention budget" in result.message
@@ -86,7 +89,7 @@ class TestCompactCutoffZero:
class TestCompactNegligibleSavings:
"""When cutoff > 0 but savings are too small to be worth it."""
async def test_skip_when_few_messages_and_low_tokens(self):
def test_skip_when_few_messages_and_low_tokens(self):
msgs = [MagicMock() for _ in range(15)]
graph_gateway = FakeGraphGateway(
state_values={"messages": msgs, "_summarization_event": None}
@@ -123,14 +126,14 @@ class TestCompactNegligibleSavings:
side_effect=lambda x: next(token_values),
),
):
result = await _compact(graph_gateway)
result = _compact(graph_gateway)
assert result.status == "noop"
assert "not worth" in result.message
# No LLM call should have been made
mock_middleware_inst._acreate_summary.assert_not_called()
async def test_still_compacts_when_few_messages_but_high_tokens(self):
def test_still_compacts_when_few_messages_but_high_tokens(self):
"""2 messages but they account for >2% of tokens — should compact."""
from langchain_core.messages import HumanMessage
@@ -175,7 +178,7 @@ class TestCompactNegligibleSavings:
side_effect=lambda x: next(token_values),
),
):
result = await _compact(graph_gateway)
result = _compact(graph_gateway)
assert result.status == "ok"
assert len(graph_gateway.updated_states) == 1
@@ -184,7 +187,7 @@ class TestCompactNegligibleSavings:
class TestCompactSuccess:
"""Normal compaction flow."""
async def test_manual_threshold_blocks_low_context_compaction(self):
def test_manual_threshold_blocks_low_context_compaction(self):
msgs = [MagicMock() for _ in range(20)]
graph_gateway = FakeGraphGateway(
state_values={"messages": msgs, "_summarization_event": None}
@@ -214,7 +217,7 @@ class TestCompactSuccess:
return_value=30_000,
),
):
result = await _compact(graph_gateway)
result = _compact(graph_gateway)
assert result.status == "noop"
assert "40%" in result.message
@@ -222,7 +225,7 @@ class TestCompactSuccess:
mock_middleware_inst._determine_cutoff_index.assert_not_called()
mock_middleware_inst._acreate_summary.assert_not_called()
async def test_successful_compaction(self):
def test_successful_compaction(self):
from langchain_core.messages import HumanMessage
msgs = [MagicMock() for _ in range(20)]
@@ -270,7 +273,7 @@ class TestCompactSuccess:
side_effect=lambda x: next(token_values),
),
):
result = await _compact(graph_gateway)
result = _compact(graph_gateway)
assert result.status == "ok"
assert result.messages_compacted == 15
@@ -288,7 +291,7 @@ class TestCompactSuccess:
assert "_summarization_event" in event_data
assert event_data["_summarization_event"]["cutoff_index"] == 15
async def test_offload_failure_non_fatal(self):
def test_offload_failure_non_fatal(self):
"""Offload failure should not prevent compaction."""
from langchain_core.messages import HumanMessage
@@ -332,7 +335,7 @@ class TestCompactSuccess:
return_value=1000,
),
):
result = await _compact(graph_gateway)
result = _compact(graph_gateway)
assert result.status == "ok"
assert len(graph_gateway.updated_states) == 1
@@ -374,7 +377,7 @@ class TestRenderCompactResult:
class TestCompactCommandUI:
"""TUI-specific compact progress indicator behavior."""
async def test_command_uses_tui_indicator_when_available(self):
def test_command_uses_tui_indicator_when_available(self):
from EvoScientist.cli.commands import CompactResult
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.session import CompactCommand
@@ -410,7 +413,7 @@ class TestCompactCommandUI:
return_value="summary-panel",
),
):
await CompactCommand().execute(ctx, [])
_run(CompactCommand().execute(ctx, []))
assert ui.started == 1
assert ui.stopped == 1
+6 -162
View File
@@ -52,9 +52,12 @@ def _restore_dangerous_env():
def temp_config_dir(tmp_path, monkeypatch):
"""Use a temporary directory for config during tests."""
config_dir = tmp_path / "evoscientist"
monkeypatch.delenv("EVOSCIENTIST_CONFIG_DIR", raising=False)
monkeypatch.delenv("EVOSCIENTIST_HOME", raising=False)
monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path))
# Prevent load_dotenv from loading the project's real .env file
monkeypatch.setattr(
"EvoScientist.config.settings.find_dotenv",
lambda *a, **k: str(tmp_path / ".env"),
)
# Also clear any API keys from environment
for key in [
"ANTHROPIC_API_KEY",
@@ -74,9 +77,6 @@ def temp_config_dir(tmp_path, monkeypatch):
"EVOSCIENTIST_AUXILIARY_MODEL",
"EVOSCIENTIST_AUXILIARY_PROVIDER",
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE",
"EVOSCIENTIST_OPENROUTER_HTTP_REFERER",
"EVOSCIENTIST_OPENROUTER_APP_TITLE",
"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES",
"EVOSCIENTIST_DANGEROUS_MODE",
]:
monkeypatch.delenv(key, raising=False)
@@ -104,9 +104,6 @@ def clean_env(monkeypatch):
"EVOSCIENTIST_AUXILIARY_MODEL",
"EVOSCIENTIST_AUXILIARY_PROVIDER",
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE",
"EVOSCIENTIST_OPENROUTER_HTTP_REFERER",
"EVOSCIENTIST_OPENROUTER_APP_TITLE",
"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES",
"EVOSCIENTIST_DANGEROUS_MODE",
]:
monkeypatch.delenv(key, raising=False)
@@ -132,13 +129,8 @@ class TestEvoScientistConfig:
assert config.show_thinking is True
assert config.ui_backend == "tui"
assert config.log_level == "warning"
assert config.reasoning_effort == ""
assert config.reasoning_effort == "high"
assert config.openrouter_anthropic_prompt_cache is True
assert config.openrouter_http_referer == (
"https://github.com/EvoScientist/EvoScientist"
)
assert config.openrouter_app_title == "EvoScientist"
assert config.openrouter_app_categories == "creative-writing,personal-agent"
assert config.memory_profile_enabled is True
assert config.memory_observations_enabled is True
assert config.memory_observation_writer == MemoryObservationWriter.ALL
@@ -153,8 +145,6 @@ class TestEvoScientistConfig:
assert config.channel_debug_tracing is False
assert config.imessage_enabled is False
assert config.imessage_allowed_senders == ""
assert config.repetitive_tool_call_threshold == 2
assert config.max_consecutive_tool_errors == 3
def test_auth_mode_default(self):
"""Test that anthropic_auth_mode defaults to api_key."""
@@ -202,18 +192,6 @@ class TestEvoScientistConfig:
assert config.dangerous_mode is True
assert config.auto_approve is True
@pytest.mark.parametrize(
"kwargs",
[
{"repetitive_tool_call_threshold": -1},
{"max_consecutive_tool_errors": -1},
{"max_consecutive_tool_errors": True},
],
)
def test_tool_guard_thresholds_must_be_non_negative_integers(self, kwargs):
with pytest.raises(ValueError, match="non-negative integer"):
EvoScientistConfig(**kwargs)
# =============================================================================
# Test config path functions
@@ -221,34 +199,14 @@ class TestEvoScientistConfig:
class TestConfigPaths:
def test_get_config_dir_with_explicit_override(self, monkeypatch, tmp_path):
"""An explicit config directory has the highest priority."""
config_dir = tmp_path / "gateway-config"
monkeypatch.setenv("EVOSCIENTIST_CONFIG_DIR", str(config_dir))
monkeypatch.setenv("EVOSCIENTIST_HOME", str(tmp_path / "runtime-home"))
assert get_config_dir() == config_dir.resolve()
def test_get_config_dir_with_evoscientist_home(self, monkeypatch, tmp_path):
"""Runtime home keeps configuration and data under one root."""
home = tmp_path / "runtime-home"
monkeypatch.delenv("EVOSCIENTIST_CONFIG_DIR", raising=False)
monkeypatch.setenv("EVOSCIENTIST_HOME", str(home))
assert get_config_dir() == home.resolve() / "config"
def test_get_config_dir_with_xdg(self, monkeypatch, tmp_path):
"""Test config dir uses XDG_CONFIG_HOME when set."""
monkeypatch.delenv("EVOSCIENTIST_CONFIG_DIR", raising=False)
monkeypatch.delenv("EVOSCIENTIST_HOME", raising=False)
monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path))
config_dir = get_config_dir()
assert config_dir == tmp_path / "evoscientist"
def test_get_config_dir_default(self, monkeypatch):
"""Test config dir defaults to ~/.config/evoscientist."""
monkeypatch.delenv("EVOSCIENTIST_CONFIG_DIR", raising=False)
monkeypatch.delenv("EVOSCIENTIST_HOME", raising=False)
monkeypatch.delenv("XDG_CONFIG_HOME", raising=False)
config_dir = get_config_dir()
assert config_dir == Path.home() / ".config" / "evoscientist"
@@ -284,23 +242,6 @@ class TestLoadSaveReset:
assert data["provider"] == "openai"
assert data["model"] == "gpt-4o"
def test_save_restricts_config_permissions(self, temp_config_dir, clean_env):
"""Config file permissions should not depend on the process umask."""
original_umask = os.umask(0)
try:
save_config(EvoScientistConfig(anthropic_api_key="test-key"))
finally:
os.umask(original_umask)
config_path = get_config_path()
if os.name == "nt":
assert config_path.exists()
# Windows reports pseudo-permission bits, so we don't test them here.
return
assert config_path.parent.stat().st_mode & 0o777 == 0o700
assert config_path.stat().st_mode & 0o777 == 0o600
def test_load_reads_saved_config(self, temp_config_dir, clean_env):
"""Test that load reads previously saved config."""
original = EvoScientistConfig(
@@ -714,22 +655,6 @@ class TestPriorityChain:
config = get_effective_config()
assert config.openrouter_anthropic_prompt_cache is False
def test_env_openrouter_app_attribution_override(
self, temp_config_dir, monkeypatch
):
"""OpenRouter app-attribution env vars should override file config."""
save_config(EvoScientistConfig())
monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_HTTP_REFERER", "https://acme.test")
monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_TITLE", "Acme")
monkeypatch.setenv(
"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "cli-agent,programming-app"
)
config = get_effective_config()
assert config.openrouter_http_referer == "https://acme.test"
assert config.openrouter_app_title == "Acme"
assert config.openrouter_app_categories == "cli-agent,programming-app"
def test_set_openrouter_anthropic_prompt_cache(self, temp_config_dir, clean_env):
"""Test OpenRouter Anthropic prompt cache can be set through config."""
save_config(EvoScientistConfig())
@@ -790,60 +715,6 @@ class TestApplyConfigToEnv:
"false"
)
def test_openrouter_app_attribution_applied_to_env(self, clean_env, monkeypatch):
"""Config app-attribution values are exported to env for models.py."""
for env in (
"EVOSCIENTIST_OPENROUTER_HTTP_REFERER",
"EVOSCIENTIST_OPENROUTER_APP_TITLE",
"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES",
):
monkeypatch.delenv(env, raising=False)
config = EvoScientistConfig(
openrouter_http_referer="https://acme.test",
openrouter_app_title="Acme",
openrouter_app_categories="cli-agent,programming-app",
)
apply_config_to_env(config)
assert os.environ.get("EVOSCIENTIST_OPENROUTER_HTTP_REFERER") == (
"https://acme.test"
)
assert os.environ.get("EVOSCIENTIST_OPENROUTER_APP_TITLE") == "Acme"
assert os.environ.get("EVOSCIENTIST_OPENROUTER_APP_CATEGORIES") == (
"cli-agent,programming-app"
)
def test_openrouter_app_attribution_env_not_overwritten(
self, clean_env, monkeypatch
):
"""apply_config_to_env must not clobber an already-set attribution env var."""
monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_TITLE", "existing-title")
config = EvoScientistConfig(openrouter_app_title="config-title")
apply_config_to_env(config)
assert os.environ.get("EVOSCIENTIST_OPENROUTER_APP_TITLE") == "existing-title"
def test_openrouter_app_attribution_empty_config_not_applied(
self, clean_env, monkeypatch
):
"""Empty-string attribution config must not create env vars."""
for env in (
"EVOSCIENTIST_OPENROUTER_HTTP_REFERER",
"EVOSCIENTIST_OPENROUTER_APP_TITLE",
"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES",
):
monkeypatch.delenv(env, raising=False)
config = EvoScientistConfig(
openrouter_http_referer="",
openrouter_app_title="",
openrouter_app_categories="",
)
apply_config_to_env(config)
assert os.environ.get("EVOSCIENTIST_OPENROUTER_HTTP_REFERER") is None
assert os.environ.get("EVOSCIENTIST_OPENROUTER_APP_TITLE") is None
assert os.environ.get("EVOSCIENTIST_OPENROUTER_APP_CATEGORIES") is None
def test_dangerous_mode_round_trips_to_env(self, clean_env, monkeypatch):
"""dangerous_mode set via CLI override must survive a fresh re-read.
@@ -961,30 +832,3 @@ def test_scheduler_config_defaults_and_env(monkeypatch):
assert eff2.memory_skill_synthesis_mode == MemorySkillSynthesisMode.AUTO
assert eff2.memory_skill_synthesis_cadence == MemorySkillSynthesisCadence.MONTHLY
assert eff2.memory_skill_synthesis_time == "04:30"
# =============================================================================
# Dotenv isolation (issue #322)
# =============================================================================
class TestDotenvIsolation:
def test_env_file_not_leaked_into_process_env(self, tmp_path, monkeypatch):
"""A .env in cwd must not leak into os.environ during tests.
Without the suite-wide ``_isolate_dotenv`` fixture,
``get_effective_config`` loads the developer's real .env with
``override=True``; an empty-valued line like ``MINIMAX_BASE_URL=``
then poisons ``os.environ.get(key, default)`` lookups for every
test that runs afterwards in the same process.
"""
monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path))
repro_dir = tmp_path / "repro"
repro_dir.mkdir()
(repro_dir / ".env").write_text("MINIMAX_BASE_URL=\n")
monkeypatch.chdir(repro_dir)
monkeypatch.delenv("MINIMAX_BASE_URL", raising=False)
get_effective_config()
assert "MINIMAX_BASE_URL" not in os.environ
+7 -6
View File
@@ -15,6 +15,7 @@ from EvoScientist.middleware.configurable_model import (
ConfigurableModelMiddleware,
_read_model_override,
)
from tests.conftest import run_async as _run
@contextmanager
@@ -133,7 +134,7 @@ class TestPassThrough:
handler.assert_called_once_with(req)
req.override.assert_not_called()
async def test_async_no_override_passes_request_unchanged(self):
def test_async_no_override_passes_request_unchanged(self):
mw = ConfigurableModelMiddleware()
req = _make_request()
@@ -142,7 +143,7 @@ class TestPassThrough:
return "ok"
with _patched_config({}):
result = await mw.awrap_model_call(req, handler)
result = _run(mw.awrap_model_call(req, handler))
assert result == "ok"
req.override.assert_not_called()
@@ -184,7 +185,7 @@ class TestModelOverride:
assert called_with is not req
assert called_with.model is new_model
async def test_async_override_path_parity(self):
def test_async_override_path_parity(self):
mw = ConfigurableModelMiddleware()
req = _make_request()
new_model = MagicMock()
@@ -201,7 +202,7 @@ class TestModelOverride:
"EvoScientist.llm.get_chat_model", return_value=new_model
) as mock_get,
):
result = await mw.awrap_model_call(req, handler)
result = _run(mw.awrap_model_call(req, handler))
assert result == "ok"
mock_get.assert_called_once_with(model="claude-opus-4-8", provider="anthropic")
@@ -315,7 +316,7 @@ class TestResolveFailure:
handler.assert_called_once_with(req)
req.override.assert_not_called()
async def test_async_falls_back_when_resolve_raises(self):
def test_async_falls_back_when_resolve_raises(self):
mw = ConfigurableModelMiddleware()
req = _make_request()
@@ -332,7 +333,7 @@ class TestResolveFailure:
side_effect=ValueError("unknown model"),
),
):
result = await mw.awrap_model_call(req, handler)
result = _run(mw.awrap_model_call(req, handler))
assert result == "ok"
assert called == [req]
@@ -66,6 +66,7 @@ def test_wrap_model_call_raises_context_overflow():
assert handler.call_count == 1
@pytest.mark.anyio
async def test_awrap_model_call_raises_context_overflow():
# Setup mocks
msgs = [HumanMessage(content=f"msg {i}") for i in range(10)]
@@ -90,6 +91,7 @@ async def test_awrap_model_call_raises_context_overflow():
assert handler.call_count == 1
@pytest.mark.anyio
async def test_awrap_model_call_passes_through_other_errors():
request = ModelRequest(
messages=[],
+6 -4
View File
@@ -2,9 +2,11 @@
from unittest.mock import MagicMock
from tests.conftest import run_async as _run
class TestCurrentCommand:
async def test_prints_thread_workspace_and_memory(self):
def test_prints_thread_workspace_and_memory(self):
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.general import CurrentCommand
@@ -15,14 +17,14 @@ class TestCurrentCommand:
ui=ui,
workspace_dir="/tmp/ws",
)
await CurrentCommand().execute(ctx, [])
_run(CurrentCommand().execute(ctx, []))
# Three append_system calls: Thread, Workspace, Memory dir.
calls = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Thread: abc123" in s for s in calls)
assert any("Workspace:" in s for s in calls)
assert any("Memory dir:" in s for s in calls)
async def test_skips_workspace_when_none(self):
def test_skips_workspace_when_none(self):
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.general import CurrentCommand
@@ -33,7 +35,7 @@ class TestCurrentCommand:
ui=ui,
workspace_dir=None,
)
await CurrentCommand().execute(ctx, [])
_run(CurrentCommand().execute(ctx, []))
calls = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Thread: abc123" in s for s in calls)
assert not any("Workspace:" in s for s in calls)
+15 -14
View File
@@ -2,6 +2,7 @@
from unittest.mock import AsyncMock, MagicMock
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway, FakeThreadStore
@@ -20,17 +21,17 @@ def _ctx(thread_id="current", thread_store=None):
class TestDeleteCommand:
async def test_refuses_to_delete_current(self):
def test_refuses_to_delete_current(self):
from EvoScientist.commands.implementation.session import DeleteCommand
thread_store = FakeThreadStore(resolved_thread_id="current", deleted=True)
ctx, ui = _ctx(thread_id="current", thread_store=thread_store)
await DeleteCommand().execute(ctx, ["current"])
_run(DeleteCommand().execute(ctx, ["current"]))
assert ("delete_thread", "current") not in thread_store.calls
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Cannot delete the current session" in m for m in msgs)
async def test_happy_path_success(self):
def test_happy_path_success(self):
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx(
@@ -40,45 +41,45 @@ class TestDeleteCommand:
deleted=True,
),
)
await DeleteCommand().execute(ctx, ["other-thread"])
_run(DeleteCommand().execute(ctx, ["other-thread"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Deleted session other-thread" in m for m in msgs)
async def test_not_found(self):
def test_not_found(self):
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx()
await DeleteCommand().execute(ctx, ["missing"])
_run(DeleteCommand().execute(ctx, ["missing"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("not found" in m for m in msgs)
async def test_ambiguous_prefix(self):
def test_ambiguous_prefix(self):
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx(thread_store=FakeThreadStore(matches=["abc-one", "abc-two"]))
await DeleteCommand().execute(ctx, ["abc"])
_run(DeleteCommand().execute(ctx, ["abc"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Ambiguous" in m for m in msgs)
async def test_prefix_resolves_to_unique_match(self):
def test_prefix_resolves_to_unique_match(self):
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx(
thread_store=FakeThreadStore(resolved_thread_id="abc-one", deleted=True)
)
await DeleteCommand().execute(ctx, ["abc"])
_run(DeleteCommand().execute(ctx, ["abc"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Deleted session abc-one" in m for m in msgs)
async def test_no_arg_empty_sessions_prints_notice(self):
def test_no_arg_empty_sessions_prints_notice(self):
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx()
await DeleteCommand().execute(ctx, [])
_run(DeleteCommand().execute(ctx, []))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("No sessions to delete" in m for m in msgs)
async def test_no_arg_calls_picker_returns_none(self):
def test_no_arg_calls_picker_returns_none(self):
"""When no arg and picker returns None, nothing is deleted."""
from EvoScientist.commands.implementation.session import DeleteCommand
@@ -95,5 +96,5 @@ class TestDeleteCommand:
]
store = FakeThreadStore(threads=threads)
ctx.graph_gateway = FakeGraphGateway(thread_store=store)
await DeleteCommand().execute(ctx, [])
_run(DeleteCommand().execute(ctx, []))
ui.wait_for_thread_pick.assert_awaited_once()
+33 -32
View File
@@ -7,6 +7,7 @@ import pytest
from EvoScientist.channels.base import ChannelError, OutboundMessage
from EvoScientist.channels.dingtalk.channel import DingTalkChannel, DingTalkConfig
from tests.conftest import run_async as _run
class TestDingTalkConfig:
@@ -40,30 +41,30 @@ class TestDingTalkChannel:
assert channel._running is False
assert channel.name == "dingtalk"
async def test_start_raises_without_credentials(self):
def test_start_raises_without_credentials(self):
config = DingTalkConfig(client_id="", client_secret="")
channel = DingTalkChannel(config)
with pytest.raises(ChannelError, match="client_id and client_secret"):
await channel.start()
_run(channel.start())
async def test_start_raises_without_client_id(self):
def test_start_raises_without_client_id(self):
config = DingTalkConfig(client_id="", client_secret="secret")
channel = DingTalkChannel(config)
with pytest.raises(ChannelError, match="client_id and client_secret"):
await channel.start()
_run(channel.start())
async def test_start_raises_without_client_secret(self):
def test_start_raises_without_client_secret(self):
config = DingTalkConfig(client_id="id", client_secret="")
channel = DingTalkChannel(config)
with pytest.raises(ChannelError, match="client_id and client_secret"):
await channel.start()
_run(channel.start())
async def test_stop_when_not_running(self):
def test_stop_when_not_running(self):
config = DingTalkConfig(client_id="test-id", client_secret="test-secret")
channel = DingTalkChannel(config)
await channel.stop()
_run(channel.stop())
async def test_send_returns_false_without_client(self):
def test_send_returns_false_without_client(self):
config = DingTalkConfig(client_id="test-id", client_secret="test-secret")
channel = DingTalkChannel(config)
msg = OutboundMessage(
@@ -72,7 +73,7 @@ class TestDingTalkChannel:
content="hello",
metadata={"chat_id": "user123"},
)
result = await channel.send(msg)
result = _run(channel.send(msg))
assert result is False
def test_capabilities(self):
@@ -129,20 +130,20 @@ class TestDingTalkWsMessageParsing:
channel._token_expires = 9999999999
return channel
async def test_system_ping_ack(self):
def test_system_ping_ack(self):
channel = self._make_channel()
data = {
"type": "SYSTEM",
"headers": {"topic": "ping", "messageId": "ping-1"},
"data": "pong-data",
}
await channel._on_ws_message(data)
_run(channel._on_ws_message(data))
channel._ws_session.send_str.assert_called_once()
sent = json.loads(channel._ws_session.send_str.call_args[0][0])
assert sent["code"] == 200
assert sent["data"] == "pong-data"
async def test_callback_text_message(self):
def test_callback_text_message(self):
channel = self._make_channel()
channel._enqueue_raw = AsyncMock()
@@ -157,14 +158,14 @@ class TestDingTalkWsMessageParsing:
"headers": {"messageId": "msg-1", "contentType": "application/json"},
"data": json.dumps(payload),
}
await channel._on_ws_message(data)
_run(channel._on_ws_message(data))
channel._enqueue_raw.assert_called_once()
raw = channel._enqueue_raw.call_args[0][0]
assert raw.text == "hello bot"
assert raw.sender_id == "staff123"
assert raw.is_group is False
async def test_callback_group_message_mention(self):
def test_callback_group_message_mention(self):
channel = self._make_channel()
channel._enqueue_raw = AsyncMock()
@@ -180,12 +181,12 @@ class TestDingTalkWsMessageParsing:
"headers": {"messageId": "msg-2"},
"data": json.dumps(payload),
}
await channel._on_ws_message(data)
_run(channel._on_ws_message(data))
raw = channel._enqueue_raw.call_args[0][0]
assert raw.is_group is True
assert raw.was_mentioned is True
async def test_callback_group_no_mention(self):
def test_callback_group_no_mention(self):
channel = self._make_channel()
channel._enqueue_raw = AsyncMock()
@@ -200,12 +201,12 @@ class TestDingTalkWsMessageParsing:
"headers": {"messageId": "msg-3"},
"data": json.dumps(payload),
}
await channel._on_ws_message(data)
_run(channel._on_ws_message(data))
raw = channel._enqueue_raw.call_args[0][0]
assert raw.is_group is True
assert raw.was_mentioned is False
async def test_ignores_non_callback(self):
def test_ignores_non_callback(self):
channel = self._make_channel()
channel._enqueue_raw = AsyncMock()
@@ -214,10 +215,10 @@ class TestDingTalkWsMessageParsing:
"headers": {"messageId": "msg-x"},
"data": "{}",
}
await channel._on_ws_message(data)
_run(channel._on_ws_message(data))
channel._enqueue_raw.assert_not_called()
async def test_ignores_empty_content(self):
def test_ignores_empty_content(self):
channel = self._make_channel()
channel._enqueue_raw = AsyncMock()
@@ -231,20 +232,20 @@ class TestDingTalkWsMessageParsing:
"headers": {"messageId": "msg-e"},
"data": json.dumps(payload),
}
await channel._on_ws_message(data)
_run(channel._on_ws_message(data))
channel._enqueue_raw.assert_not_called()
async def test_non_dict_data_ignored(self):
def test_non_dict_data_ignored(self):
channel = self._make_channel()
channel._enqueue_raw = AsyncMock()
await channel._on_ws_message("not a dict")
_run(channel._on_ws_message("not a dict"))
channel._enqueue_raw.assert_not_called()
class TestDingTalkSendChunk:
"""Test _send_chunk with mocked HTTP client."""
async def test_send_chunk_calls_api(self):
def test_send_chunk_calls_api(self):
config = DingTalkConfig(client_id="test-app", client_secret="test-secret")
channel = DingTalkChannel(config)
channel._access_token = "fake-token"
@@ -255,7 +256,7 @@ class TestDingTalkSendChunk:
channel._http_client = MagicMock()
channel._http_client.post = AsyncMock(return_value=mock_response)
await channel._send_chunk("user1", "formatted", "raw text", None, {})
_run(channel._send_chunk("user1", "formatted", "raw text", None, {}))
channel._http_client.post.assert_called_once()
call_args = channel._http_client.post.call_args
body = call_args.kwargs.get("json") or call_args[1].get("json")
@@ -272,21 +273,21 @@ class TestDingTalkChannelRegistration:
class TestDingTalkProbe:
async def test_missing_credentials(self):
def test_missing_credentials(self):
from EvoScientist.channels.dingtalk.probe import validate_dingtalk
ok, msg = await validate_dingtalk("", "")
ok, msg = _run(validate_dingtalk("", ""))
assert ok is False
assert "required" in msg
async def test_missing_client_id(self):
def test_missing_client_id(self):
from EvoScientist.channels.dingtalk.probe import validate_dingtalk
ok, _msg = await validate_dingtalk("", "secret")
ok, _msg = _run(validate_dingtalk("", "secret"))
assert ok is False
async def test_missing_client_secret(self):
def test_missing_client_secret(self):
from EvoScientist.channels.dingtalk.probe import validate_dingtalk
ok, _msg = await validate_dingtalk("id", "")
ok, _msg = _run(validate_dingtalk("id", ""))
assert ok is False
+7 -6
View File
@@ -4,6 +4,7 @@ import pytest
from EvoScientist.channels.base import ChannelError
from EvoScientist.channels.discord.channel import DiscordChannel, DiscordConfig
from tests.conftest import run_async as _run
class TestDiscordChannel:
@@ -13,18 +14,18 @@ class TestDiscordChannel:
assert channel.config is config
assert channel._running is False
async def test_start_raises_without_token_or_library(self):
def test_start_raises_without_token_or_library(self):
config = DiscordConfig(bot_token="")
channel = DiscordChannel(config)
with pytest.raises(ChannelError):
await channel.start()
_run(channel.start())
async def test_stop_when_not_running(self):
def test_stop_when_not_running(self):
config = DiscordConfig(bot_token="test")
channel = DiscordChannel(config)
await channel.stop()
_run(channel.stop())
async def test_send_returns_false_without_client(self):
def test_send_returns_false_without_client(self):
from EvoScientist.channels.base import OutboundMessage
config = DiscordConfig(bot_token="test")
@@ -35,5 +36,5 @@ class TestDiscordChannel:
content="hello",
metadata={"chat_id": "123"},
)
result = await channel.send(msg)
result = _run(channel.send(msg))
assert result is False
@@ -1,440 +0,0 @@
"""Tests for ErrorNormalizationMiddleware + ProviderStreamError.
Verifies that provider-SDK exceptions from a chat model call get
wrapped into a non-dataclass ``ProviderStreamError`` at the model
boundary, and that non-provider exceptions pass through unchanged.
The provider tag is derived from ``request.model`` (class + base_url),
not from the raised exception.
"""
from __future__ import annotations
import asyncio
import dataclasses
from types import SimpleNamespace
import pytest
from EvoScientist.llm.errors import (
AgentControlError,
ModelToolProtocolError,
ProviderStreamError,
)
from EvoScientist.middleware.error_normalization import (
ErrorNormalizationMiddleware,
_normalize,
)
# ---------------------------------------------------------------------------
# Test fixtures — fake chat model instances + requests
# ---------------------------------------------------------------------------
def _fake_model(module: str, cls_name: str, **attrs):
"""Build a fake chat model instance whose ``type(model).__module__``
matches *module*, carrying arbitrary attributes for ``base_url`` /
``openai_api_base`` / ``anthropic_api_url`` lookup.
"""
cls = type(cls_name, (), {"__module__": module})
inst = cls()
for k, v in attrs.items():
setattr(inst, k, v)
return inst
def _request(model):
"""Fake ``ModelRequest`` with just the ``.model`` attribute the
middleware reads.
"""
return SimpleNamespace(model=model)
def _openai_model(base_url: str | None = None):
return _fake_model(
"langchain_openai.chat_models.base",
"ChatOpenAI",
openai_api_base=base_url,
)
def _anthropic_model(base_url: str | None = None):
return _fake_model(
"langchain_anthropic.chat_models",
"ChatAnthropic",
anthropic_api_url=base_url,
)
def _openrouter_model():
return _fake_model("langchain_openrouter.chat_models", "ChatOpenRouter")
def _google_model():
return _fake_model("langchain_google_genai.chat_models", "ChatGoogleGenerativeAI")
def _make_exc(cls_name: str = "APIError", message: str = "boom", **attrs):
"""Build a plain-Exception subclass carrying arbitrary attributes
(``status_code``, ``code``, ``type``, ``request_id`` …).
"""
cls = type(cls_name, (Exception,), attrs)
return cls(message)
# ---------------------------------------------------------------------------
# _normalize — provider inference from ModelRequest.model
# ---------------------------------------------------------------------------
class TestNormalize:
def test_openai_native_model_tags_openai(self):
req = _request(_openai_model())
exc = _make_exc(message="rate limited", status_code=429)
wrapped = _normalize(req, exc)
assert isinstance(wrapped, ProviderStreamError)
assert wrapped.provider == "openai"
assert wrapped.status_code == 429
def test_openai_routed_deepseek_tagged_by_base_url(self):
req = _request(_openai_model(base_url="https://api.deepseek.com"))
wrapped = _normalize(req, _make_exc(message="quota exceeded"))
assert wrapped.provider == "deepseek"
def test_openai_routed_moonshot_tagged_by_base_url(self):
req = _request(_openai_model(base_url="https://api.moonshot.cn/v1"))
assert _normalize(req, _make_exc()).provider == "moonshot"
def test_unknown_openai_compat_host_tagged_openai_compat(self):
req = _request(_openai_model(base_url="https://internal.corp/v1"))
assert _normalize(req, _make_exc()).provider == "openai_compat"
def test_anthropic_native_model_tags_anthropic(self):
req = _request(_anthropic_model(base_url="https://api.anthropic.com"))
assert _normalize(req, _make_exc()).provider == "anthropic"
def test_anthropic_routed_minimax_tagged_by_base_url(self):
req = _request(_anthropic_model(base_url="https://api.minimaxi.com/anthropic"))
assert _normalize(req, _make_exc()).provider == "minimax"
def test_unknown_anthropic_compat_host_tagged_anthropic_compat(self):
req = _request(_anthropic_model(base_url="https://internal.corp/v1"))
assert _normalize(req, _make_exc()).provider == "anthropic_compat"
def test_openrouter_tagged_from_class_alone(self):
req = _request(_openrouter_model())
wrapped = _normalize(req, _make_exc(cls_name="UnauthorizedResponseError"))
assert wrapped.provider == "openrouter"
assert wrapped.class_qualname.endswith(".UnauthorizedResponseError")
def test_google_genai_tagged_from_class_alone(self):
req = _request(_google_model())
assert _normalize(req, _make_exc()).provider == "google_genai"
def test_unrecognized_model_class_returns_none(self):
req = _request(_fake_model("some.other.pkg", "SomeModel"))
assert _normalize(req, _make_exc()) is None
def test_missing_model_on_request_returns_none(self):
"""If the request has no ``.model`` at all (defensive)."""
assert _normalize(SimpleNamespace(), _make_exc()) is None
def test_already_normalized_exception_passes_through(self):
"""``ModelFallbackMiddleware`` wraps against the failing model
before re-raising. The outer chain's ``_normalize`` must NOT
double-wrap — otherwise attribution flips back to the original
request's model.
"""
req = _request(_openrouter_model())
pre_wrapped = ProviderStreamError(
provider="moonshot",
class_qualname="openai.RateLimitError",
message="quota exceeded",
)
assert _normalize(req, pre_wrapped) is None
@pytest.mark.parametrize(
"error",
[
AgentControlError("MODEL_TOOL_LOOP_DETECTED", "loop stopped"),
ModelToolProtocolError(
"missing_name",
provider="openai",
model="gpt-example",
route_key="route-1",
),
],
)
def test_platform_control_error_passes_through(self, error):
req = _request(_openai_model())
assert _normalize(req, error) is None
# ---------------------------------------------------------------------------
# _is_provider_error — used by tool selector to distinguish provider
# failures (surface) from shape / config failures (degrade)
# ---------------------------------------------------------------------------
class TestIsProviderError:
def test_openai_module_is_provider_error(self):
from EvoScientist.middleware.error_normalization import _is_provider_error
assert _is_provider_error(_make_exc(__module__="openai"))
def test_httpx_timeout_is_provider_error(self):
from EvoScientist.middleware.error_normalization import _is_provider_error
assert _is_provider_error(
_make_exc(cls_name="TimeoutException", __module__="httpx")
)
def test_langchain_wrapper_module_is_provider_error(self):
from EvoScientist.middleware.error_normalization import _is_provider_error
assert _is_provider_error(
_make_exc(
cls_name="BadRequestError",
__module__="langchain_openai.chat_models",
)
)
def test_pydantic_validation_is_not_provider_error(self):
"""Structured-output shape failures come from pydantic /
langchain, NOT from a provider SDK — the tool selector's
graceful-degrade path is right for these.
"""
from EvoScientist.middleware.error_normalization import _is_provider_error
assert not _is_provider_error(
_make_exc(cls_name="ValidationError", __module__="pydantic")
)
def test_builtin_is_not_provider_error(self):
from EvoScientist.middleware.error_normalization import _is_provider_error
assert not _is_provider_error(RuntimeError("x"))
# ---------------------------------------------------------------------------
# ProviderStreamError envelope
# ---------------------------------------------------------------------------
class TestProviderStreamErrorEnvelope:
def test_envelope_contains_required_fields(self):
err = ProviderStreamError(
provider="deepseek",
class_qualname="openai.RateLimitError",
message="quota exceeded",
status_code=429,
code="insufficient_quota",
)
env = err.as_envelope()
assert env["error"] == "RateLimitError"
assert env["class"] == "openai.RateLimitError"
assert env["message"] == "quota exceeded"
assert env["provider"] == "deepseek"
assert env["status_code"] == 429
assert env["code"] == "insufficient_quota"
def test_envelope_omits_absent_optional_fields(self):
err = ProviderStreamError(
provider="openrouter",
class_qualname="openrouter.errors.foo.UnauthorizedResponseError",
message="User not found.",
)
env = err.as_envelope()
assert "status_code" not in env
assert "code" not in env
assert "type" not in env
assert "request_id" not in env
def test_provider_stream_error_is_not_a_dataclass(self):
"""The whole point of the wrapper — must not be a dataclass so
orjson's OPT_SERIALIZE_DATACLASS fast-path doesn't fire.
"""
err = ProviderStreamError("x", "y.Z", "msg")
assert not dataclasses.is_dataclass(err)
assert not dataclasses.is_dataclass(type(err))
def test_model_dump_returns_envelope(self):
"""Upstream ``serde.default`` calls ``model_dump()`` before its
exception branch — the hook that lets us skip the serde patch.
"""
err = ProviderStreamError(
provider="openrouter",
class_qualname="openrouter.errors.foo.UnauthorizedResponseError",
message="User not found.",
status_code=401,
)
assert err.model_dump() == err.as_envelope()
# ---------------------------------------------------------------------------
# Middleware behavior
# ---------------------------------------------------------------------------
class TestMiddleware:
def _run_awrap(self, mw, request, handler):
async def _go():
return await mw.awrap_model_call(request=request, handler=handler)
return asyncio.run(_go())
def test_awrap_normalizes_provider_exception(self):
raised = _make_exc(cls_name="UnauthorizedResponseError", message="boom")
async def handler(_req):
raise raised
req = _request(_openrouter_model())
mw = ErrorNormalizationMiddleware()
with pytest.raises(ProviderStreamError) as excinfo:
self._run_awrap(mw, req, handler)
assert excinfo.value.provider == "openrouter"
assert excinfo.value.__cause__ is raised
def test_awrap_passes_through_non_provider_model_exception(self):
"""If the model isn't a recognized provider SDK, the exception
passes through unwrapped — same as any non-model exception.
"""
raised = _make_exc(message="boom")
async def handler(_req):
raise raised
req = _request(_fake_model("some.other.pkg", "SomeModel"))
mw = ErrorNormalizationMiddleware()
with pytest.raises(Exception, match="boom") as excinfo:
self._run_awrap(mw, req, handler)
assert excinfo.value is raised
def _langgraph_error_samples(self):
"""Instances covering both branches of ``_should_pass_through``:
control-flow (``GraphBubbleUp`` + subclasses) and structural
errors. Constructor signatures vary — some need positional
args — so build each explicitly.
"""
from langgraph.errors import (
EmptyInputError,
GraphBubbleUp,
GraphInterrupt,
InvalidUpdateError,
NodeTimeoutError,
TaskNotFound,
)
return [
GraphBubbleUp(),
GraphInterrupt(),
InvalidUpdateError("bad update"),
EmptyInputError("no input"),
TaskNotFound(),
NodeTimeoutError("node-x", 1.5, kind="run", run_timeout=1.0),
]
def test_awrap_passes_through_langgraph_errors(self):
"""Exceptions from ``langgraph.errors.*`` must propagate
untouched even when the model is a recognized provider —
they're either control-flow signals (interrupts, HITL) or
graph-level structural errors, neither is a provider incident.
"""
req = _request(_openrouter_model()) # recognized — would normally wrap
mw = ErrorNormalizationMiddleware()
for raised in self._langgraph_error_samples():
async def handler(_req, _r=raised):
raise _r
with pytest.raises(type(raised)) as excinfo:
self._run_awrap(mw, req, handler)
assert excinfo.value is raised, (
f"{type(raised).__name__} got wrapped instead of propagated"
)
def test_awrap_passes_through_context_overflow_error(self):
"""``ContextOverflowError`` is a cross-layer control signal:
deepagents' ``SummarizationMiddleware`` sits outside our stack
and catches it by type to compress history and retry. Wrapping
it here would change the type and break that self-healing
fallback — regressing to a user-visible ``ProviderStreamError``
on any long conversation.
"""
from langchain_core.exceptions import ContextOverflowError
raised = ContextOverflowError("context length exceeded")
async def handler(_req):
raise raised
req = _request(_openrouter_model()) # recognized — would normally wrap
mw = ErrorNormalizationMiddleware()
with pytest.raises(ContextOverflowError) as excinfo:
self._run_awrap(mw, req, handler)
assert excinfo.value is raised
def test_awrap_preserves_model_tool_protocol_error_identity(self):
raised = ModelToolProtocolError(
"missing_name",
provider="openai",
model="gpt-example",
route_key="route-1",
)
async def handler(_req):
raise raised
req = _request(_openai_model())
with pytest.raises(ModelToolProtocolError) as excinfo:
self._run_awrap(ErrorNormalizationMiddleware(), req, handler)
assert excinfo.value is raised
assert excinfo.value.code == "MODEL_TOOL_PROTOCOL_INVALID"
assert excinfo.value.fallbackable is True
def test_awrap_wraps_any_exception_from_recognized_model(self):
"""Any exception raised inside a call to a provider-recognized
model gets wrapped — including builtins like ``RuntimeError``.
Rationale: at the middleware boundary we can tell the model is
a provider, but not the exception's origin (SDK vs
langchain-wrapper vs httpx vs our code). Wrapping uniformly
gives the WebUI a consistent envelope; upstream's
``RuntimeError``-whitelist would emit ``{"error":
"RuntimeError", "message": str(exc)}`` which isn't more
useful.
"""
raised = RuntimeError("internal glitch")
async def handler(_req):
raise raised
req = _request(_openai_model())
mw = ErrorNormalizationMiddleware()
with pytest.raises(ProviderStreamError) as excinfo:
self._run_awrap(mw, req, handler)
assert excinfo.value.provider == "openai"
assert excinfo.value.__cause__ is raised
assert excinfo.value.class_qualname == "builtins.RuntimeError"
def test_sync_wrap_normalizes_provider_exception(self):
raised = _make_exc(message="boom")
def handler(_req):
raise raised
req = _request(_openrouter_model())
mw = ErrorNormalizationMiddleware()
with pytest.raises(ProviderStreamError) as excinfo:
mw.wrap_model_call(request=req, handler=handler)
assert excinfo.value.provider == "openrouter"
def test_success_path_returns_handler_result(self):
async def handler(_req):
return "ok"
req = _request(_openrouter_model())
mw = ErrorNormalizationMiddleware()
assert self._run_awrap(mw, req, handler) == "ok"
+10 -8
View File
@@ -2,6 +2,8 @@
from unittest.mock import AsyncMock, MagicMock, patch
from tests.conftest import run_async as _run
def _ctx(supports_interactive=True):
from EvoScientist.commands.base import CommandContext
@@ -29,7 +31,7 @@ _INDEX = [
class TestInstallSkills:
async def test_picker_cancel_no_install(self):
def test_picker_cancel_no_install(self):
from EvoScientist.commands.implementation.skills import InstallSkills
ctx, ui = _ctx()
@@ -43,12 +45,12 @@ class TestInstallSkills:
"EvoScientist.tools.skills_manager.install_skill",
) as install_mock,
):
await InstallSkills().execute(ctx, [])
_run(InstallSkills().execute(ctx, []))
install_mock.assert_not_called()
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Browse cancelled" in m for m in msgs)
async def test_picker_returns_selections_installs_each(self):
def test_picker_returns_selections_installs_each(self):
from EvoScientist.commands.implementation.skills import InstallSkills
ctx, ui = _ctx()
@@ -66,10 +68,10 @@ class TestInstallSkills:
return_value={"success": True, "name": "x"},
) as install_mock,
):
await InstallSkills().execute(ctx, [])
_run(InstallSkills().execute(ctx, []))
assert install_mock.call_count == 2
async def test_channel_auto_install_on_tag(self):
def test_channel_auto_install_on_tag(self):
"""Non-interactive UI + tag arg → auto-installs matching skills."""
from EvoScientist.commands.implementation.skills import InstallSkills
@@ -84,12 +86,12 @@ class TestInstallSkills:
return_value={"success": True, "name": "x"},
) as install_mock,
):
await InstallSkills().execute(ctx, ["core"])
_run(InstallSkills().execute(ctx, ["core"]))
# "core" matches research-ideation only → 1 install, no picker call
assert install_mock.call_count == 1
ui.wait_for_skill_browse.assert_not_called()
async def test_fetch_failure_prints_error(self):
def test_fetch_failure_prints_error(self):
from EvoScientist.commands.implementation.skills import InstallSkills
ctx, ui = _ctx()
@@ -97,6 +99,6 @@ class TestInstallSkills:
"EvoScientist.tools.skills_manager.fetch_remote_skill_index",
side_effect=RuntimeError("network fail"),
):
await InstallSkills().execute(ctx, [])
_run(InstallSkills().execute(ctx, []))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Failed to fetch" in m for m in msgs)
+4 -2
View File
@@ -2,9 +2,11 @@
from unittest.mock import MagicMock
from tests.conftest import run_async as _run
class TestExitCommand:
async def test_execute_calls_force_quit(self):
def test_execute_calls_force_quit(self):
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.session import ExitCommand
@@ -15,7 +17,7 @@ class TestExitCommand:
ui=ui,
)
cmd = ExitCommand()
await cmd.execute(ctx, [])
_run(cmd.execute(ctx, []))
ui.force_quit.assert_called_once()
def test_aliases_registered(self):
+39 -38
View File
@@ -14,6 +14,7 @@ from EvoScientist.channels.feishu.channel import (
_parse_inline_elements,
_parse_inline_text,
)
from tests.conftest import run_async as _run
class TestFeishuConfig:
@@ -56,24 +57,24 @@ class TestFeishuChannel:
assert channel._running is False
assert channel.name == "feishu"
async def test_start_raises_without_app_id(self):
def test_start_raises_without_app_id(self):
config = FeishuConfig(app_id="", app_secret="test-secret")
channel = FeishuChannel(config)
with pytest.raises(ChannelError, match="app_id"):
await channel.start()
_run(channel.start())
async def test_start_raises_without_app_secret(self):
def test_start_raises_without_app_secret(self):
config = FeishuConfig(app_id="test-id", app_secret="")
channel = FeishuChannel(config)
with pytest.raises(ChannelError, match="app_secret"):
await channel.start()
_run(channel.start())
async def test_stop_when_not_running(self):
def test_stop_when_not_running(self):
config = FeishuConfig(app_id="test-id", app_secret="test-secret")
channel = FeishuChannel(config)
await channel.stop()
_run(channel.stop())
async def test_send_returns_false_without_client(self):
def test_send_returns_false_without_client(self):
config = FeishuConfig(app_id="test-id", app_secret="test-secret")
channel = FeishuChannel(config)
msg = OutboundMessage(
@@ -82,7 +83,7 @@ class TestFeishuChannel:
content="hello",
metadata={"chat_id": "oc_test"},
)
result = await channel.send(msg)
result = _run(channel.send(msg))
assert result is False
def test_capabilities(self):
@@ -200,7 +201,7 @@ class TestFeishuWebhookEvent:
channel._enqueue_raw = AsyncMock()
return channel
async def test_text_message_v2(self):
def test_text_message_v2(self):
channel = self._make_channel()
event = {
"sender": {
@@ -216,7 +217,7 @@ class TestFeishuWebhookEvent:
"create_time": "1700000000000",
},
}
await channel._on_message(event)
_run(channel._on_message(event))
channel._enqueue_raw.assert_called_once()
raw = channel._enqueue_raw.call_args[0][0]
assert raw.text == "hello feishu"
@@ -224,7 +225,7 @@ class TestFeishuWebhookEvent:
assert raw.chat_id == "oc_chat1"
assert raw.is_group is False
async def test_group_message_with_mention(self):
def test_group_message_with_mention(self):
channel = self._make_channel()
event = {
"sender": {
@@ -241,13 +242,13 @@ class TestFeishuWebhookEvent:
"mentions": [{"key": "@_user_1", "id": {}}],
},
}
await channel._on_message(event)
_run(channel._on_message(event))
raw = channel._enqueue_raw.call_args[0][0]
assert raw.is_group is True
assert raw.was_mentioned is True
assert channel._mention_names == ["@_user_1"]
async def test_group_message_no_mention(self):
def test_group_message_no_mention(self):
channel = self._make_channel()
event = {
"sender": {
@@ -263,12 +264,12 @@ class TestFeishuWebhookEvent:
"create_time": "1700000000000",
},
}
await channel._on_message(event)
_run(channel._on_message(event))
raw = channel._enqueue_raw.call_args[0][0]
assert raw.is_group is True
assert raw.was_mentioned is False
async def test_skips_bot_messages(self):
def test_skips_bot_messages(self):
channel = self._make_channel()
event = {
"sender": {
@@ -282,10 +283,10 @@ class TestFeishuWebhookEvent:
"content": json.dumps({"text": "bot reply"}),
},
}
await channel._on_message(event)
_run(channel._on_message(event))
channel._enqueue_raw.assert_not_called()
async def test_post_message(self):
def test_post_message(self):
channel = self._make_channel()
post_content = {
"zh_cn": {
@@ -307,12 +308,12 @@ class TestFeishuWebhookEvent:
"create_time": "1700000000000",
},
}
await channel._on_message(event)
_run(channel._on_message(event))
raw = channel._enqueue_raw.call_args[0][0]
assert "Test" in raw.text
assert "Post body" in raw.text
async def test_unsupported_msg_type_annotation(self):
def test_unsupported_msg_type_annotation(self):
channel = self._make_channel()
event = {
"sender": {
@@ -328,7 +329,7 @@ class TestFeishuWebhookEvent:
"create_time": "1700000000000",
},
}
await channel._on_message(event)
_run(channel._on_message(event))
raw = channel._enqueue_raw.call_args[0][0]
assert "share_chat" in raw.text
@@ -336,7 +337,7 @@ class TestFeishuWebhookEvent:
class TestFeishuSendChunk:
"""Test _send_chunk with mocked HTTP client."""
async def test_send_chunk_post_format(self):
def test_send_chunk_post_format(self):
config = FeishuConfig(app_id="test-app", app_secret="test-secret")
channel = FeishuChannel(config)
channel._access_token = "fake-token"
@@ -347,14 +348,14 @@ class TestFeishuSendChunk:
channel._http_client = MagicMock()
channel._http_client.post = AsyncMock(return_value=mock_response)
await channel._send_chunk("oc_chat1", "formatted", "raw **text**", None, {})
_run(channel._send_chunk("oc_chat1", "formatted", "raw **text**", None, {}))
channel._http_client.post.assert_called()
# Should try post format first
call_args = channel._http_client.post.call_args
body = call_args.kwargs.get("json") or call_args[1].get("json")
assert body["receive_id"] == "oc_chat1"
async def test_send_chunk_with_reply(self):
def test_send_chunk_with_reply(self):
config = FeishuConfig(app_id="test-app", app_secret="test-secret")
channel = FeishuChannel(config)
channel._access_token = "fake-token"
@@ -365,7 +366,7 @@ class TestFeishuSendChunk:
channel._http_client = MagicMock()
channel._http_client.post = AsyncMock(return_value=mock_response)
await channel._send_chunk("oc_chat1", "reply", "reply text", "om_reply_id", {})
_run(channel._send_chunk("oc_chat1", "reply", "reply text", "om_reply_id", {}))
# Should call the reply API endpoint
first_call_url = channel._http_client.post.call_args_list[0][0][0]
assert "reply" in first_call_url
@@ -483,17 +484,17 @@ class TestFeishuChannelRegistration:
class TestFeishuProbe:
async def test_missing_app_id(self):
def test_missing_app_id(self):
from EvoScientist.channels.feishu.probe import validate_feishu_credentials
ok, msg = await validate_feishu_credentials("", "secret")
ok, msg = _run(validate_feishu_credentials("", "secret"))
assert ok is False
assert "app_id" in msg
async def test_missing_app_secret(self):
def test_missing_app_secret(self):
from EvoScientist.channels.feishu.probe import validate_feishu_credentials
ok, msg = await validate_feishu_credentials("id", "")
ok, msg = _run(validate_feishu_credentials("id", ""))
assert ok is False
assert "app_secret" in msg
@@ -509,7 +510,7 @@ class TestFeishuWebSocketMode:
)
assert config.subscription_mode == "websocket"
async def test_start_websocket_raises_without_lark_oapi(self):
def test_start_websocket_raises_without_lark_oapi(self):
config = FeishuConfig(
app_id="test-id",
app_secret="test-secret",
@@ -519,9 +520,9 @@ class TestFeishuWebSocketMode:
# Temporarily hide lark_oapi if it's installed
with patch.dict(sys.modules, {"lark_oapi": None}):
with pytest.raises(ChannelError, match="lark-oapi"):
await channel.start()
_run(channel.start())
async def test_start_webhook_mode_still_works(self):
def test_start_webhook_mode_still_works(self):
"""Ensure subscription_mode='webhook' still validates as before."""
config = FeishuConfig(
app_id="",
@@ -530,9 +531,9 @@ class TestFeishuWebSocketMode:
)
channel = FeishuChannel(config)
with pytest.raises(ChannelError, match="app_id"):
await channel.start()
_run(channel.start())
async def test_invalid_subscription_mode_raises(self):
def test_invalid_subscription_mode_raises(self):
config = FeishuConfig(
app_id="test-id",
app_secret="test-secret",
@@ -540,9 +541,9 @@ class TestFeishuWebSocketMode:
)
channel = FeishuChannel(config)
with pytest.raises(ChannelError, match="Invalid feishu_subscription_mode"):
await channel.start()
_run(channel.start())
async def test_on_lark_sdk_message_bridges_to_on_message(self):
def test_on_lark_sdk_message_bridges_to_on_message(self):
"""Test that _on_lark_sdk_message enqueues event dict via queue."""
import queue as queue_mod
@@ -593,14 +594,14 @@ class TestFeishuWebSocketMode:
)
# Verify the consumer processes it correctly
await channel._on_message(event_dict)
_run(channel._on_message(event_dict))
channel._enqueue_raw.assert_called_once()
raw = channel._enqueue_raw.call_args[0][0]
assert raw.text == "hello from websocket"
assert raw.sender_id == "ou_test_ws"
assert raw.is_group is False
async def test_cleanup_websocket_mode(self):
def test_cleanup_websocket_mode(self):
config = FeishuConfig(
app_id="test-id",
app_secret="test-secret",
@@ -616,7 +617,7 @@ class TestFeishuWebSocketMode:
channel._ws_consumer_task = None
channel._access_token = "fake-token"
await channel._cleanup()
_run(channel._cleanup())
mock_client.aclose.assert_called_once()
assert channel._http_client is None
+41 -31
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
import asyncio
from types import SimpleNamespace
from unittest.mock import MagicMock
@@ -178,7 +179,7 @@ def test_launch_background_run_deletes_thread_when_run_creation_fails(monkeypatc
fake_client.threads.delete.assert_called_once_with("thread-1")
async def test_async_launch_background_run_deletes_thread_when_run_creation_fails(
def test_async_launch_background_run_deletes_thread_when_run_creation_fails(
monkeypatch,
):
monkeypatch.setattr(
@@ -203,8 +204,11 @@ async def test_async_launch_background_run_deletes_thread_when_run_creation_fail
lambda **_kwargs: SimpleNamespace(threads=_Threads(), runs=_Runs()),
)
with pytest.raises(RuntimeError, match="run creation failed"):
await background_runs.alaunch_background_run(_request())
async def run() -> None:
with pytest.raises(RuntimeError, match="run creation failed"):
await background_runs.alaunch_background_run(_request())
asyncio.run(run())
assert deleted == ["thread-1"]
@@ -274,7 +278,7 @@ def test_sync_status_watcher_preserves_thread_on_poll_failure(
assert deleted == []
async def test_async_status_watcher_aborts_and_deletes_thread_on_error_status():
def test_async_status_watcher_aborts_and_deletes_thread_on_error_status():
finished: list[background_runs.BackgroundRun] = []
aborted: list[background_runs.BackgroundRun] = []
deleted: list[str] = []
@@ -287,26 +291,29 @@ async def test_async_status_watcher_aborts_and_deletes_thread_on_error_status():
async def delete(self, thread_id: str):
deleted.append(thread_id)
await background_runs.awatch_background_run(
SimpleNamespace(runs=_Runs(), threads=_Threads()),
thread_id="thread-1",
run_id="run-1",
name="test worker",
hooks=background_runs.BackgroundRunHooks(
on_finished=finished.append,
on_aborted=aborted.append,
),
watcher_config=background_runs.BackgroundRunWatcherConfig(
poll_interval_seconds=0,
),
)
async def run() -> None:
await background_runs.awatch_background_run(
SimpleNamespace(runs=_Runs(), threads=_Threads()),
thread_id="thread-1",
run_id="run-1",
name="test worker",
hooks=background_runs.BackgroundRunHooks(
on_finished=finished.append,
on_aborted=aborted.append,
),
watcher_config=background_runs.BackgroundRunWatcherConfig(
poll_interval_seconds=0,
),
)
asyncio.run(run())
assert finished == []
assert [run.run_id for run in aborted] == ["run-1"]
assert deleted == ["thread-1"]
async def test_async_status_watcher_preserves_run_url():
def test_async_status_watcher_preserves_run_url():
finished: list[background_runs.BackgroundRun] = []
class _Runs:
@@ -317,18 +324,21 @@ async def test_async_status_watcher_preserves_run_url():
async def delete(self, _thread_id: str):
return None
await background_runs.awatch_background_run(
SimpleNamespace(runs=_Runs(), threads=_Threads()),
url="http://worker.example",
thread_id="thread-1",
run_id="run-1",
name="test worker",
hooks=background_runs.BackgroundRunHooks(
on_finished=finished.append,
),
watcher_config=background_runs.BackgroundRunWatcherConfig(
poll_interval_seconds=0,
),
)
async def run() -> None:
await background_runs.awatch_background_run(
SimpleNamespace(runs=_Runs(), threads=_Threads()),
url="http://worker.example",
thread_id="thread-1",
run_id="run-1",
name="test worker",
hooks=background_runs.BackgroundRunHooks(
on_finished=finished.append,
),
watcher_config=background_runs.BackgroundRunWatcherConfig(
poll_interval_seconds=0,
),
)
asyncio.run(run())
assert [run.url for run in finished] == ["http://worker.example"]
+75 -72
View File
@@ -20,6 +20,7 @@ from EvoScientist.gateway import (
)
from EvoScientist.gateway.server import _THREAD_SEARCH_LIMIT
from EvoScientist.stream import display as display_mod
from tests.conftest import run_async
from tests.fakes import (
FakeGraphGateway,
FakeLangGraphClient,
@@ -29,7 +30,7 @@ from tests.fakes import (
)
async def test_local_gateway_streams_from_injected_streamer():
def test_local_gateway_streams_from_injected_streamer():
seen: dict[str, Any] = {}
async def _streamer(agent, message, thread_id, **kwargs):
@@ -59,7 +60,7 @@ async def test_local_gateway_streams_from_injected_streamer():
return [event async for event in gateway.stream_events(request)]
with patch("EvoScientist.stream.events.stream_agent_events", new=_streamer):
events = await _collect()
events = run_async(_collect())
assert events == [
{"type": "text", "content": "hi"},
@@ -74,7 +75,7 @@ async def test_local_gateway_streams_from_injected_streamer():
}
async def test_local_graph_gateway_delegates_thread_operations():
def test_local_graph_gateway_delegates_thread_operations():
thread_store = FakeThreadStore(
generated_thread_id="new12345",
threads=[{"thread_id": "abc12345"}],
@@ -101,7 +102,7 @@ async def test_local_graph_gateway_delegates_thread_operations():
"deleted": await gateway.delete_thread("abc12345"),
}
result = await _run()
result = run_async(_run())
assert result["created"] == "new12345"
assert result["threads"] == [{"thread_id": "abc12345"}]
@@ -131,14 +132,16 @@ async def test_local_graph_gateway_delegates_thread_operations():
]
async def test_local_graph_gateway_reads_state_values():
def test_local_graph_gateway_reads_state_values():
agent = MagicMock()
agent.aget_state = AsyncMock(
return_value=SimpleNamespace(values={"async_tasks": {"task-1": {}}})
)
gateway = LocalGraphGateway()
values = await gateway.get_state_values(GraphTarget(local_graph=agent), "abc12345")
values = run_async(
gateway.get_state_values(GraphTarget(local_graph=agent), "abc12345")
)
assert values == {"async_tasks": {"task-1": {}}}
agent.aget_state.assert_awaited_once_with(
@@ -146,15 +149,17 @@ async def test_local_graph_gateway_reads_state_values():
)
async def test_local_graph_gateway_updates_state_values():
def test_local_graph_gateway_updates_state_values():
agent = MagicMock()
agent.aupdate_state = AsyncMock()
gateway = LocalGraphGateway()
await gateway.update_state_values(
GraphTarget(local_graph=agent),
"abc12345",
{"_summarization_event": {"cutoff_index": 2}},
run_async(
gateway.update_state_values(
GraphTarget(local_graph=agent),
"abc12345",
{"_summarization_event": {"cutoff_index": 2}},
)
)
agent.aupdate_state.assert_awaited_once_with(
@@ -164,7 +169,7 @@ async def test_local_graph_gateway_updates_state_values():
)
async def test_local_stream_events_delegates_aclose_to_inner():
def test_local_stream_events_delegates_aclose_to_inner():
cleanup_ran = False
async def _streamer(_agent, _message, _thread_id, **_kwargs):
@@ -189,7 +194,7 @@ async def test_local_stream_events_delegates_aclose_to_inner():
assert cleanup_ran is True
with patch("EvoScientist.stream.events.stream_agent_events", new=_streamer):
await _run()
run_async(_run())
def test_run_streaming_can_consume_injected_gateway():
@@ -223,7 +228,7 @@ def test_run_streaming_can_consume_injected_gateway():
]
async def test_resume_command_consumes_context_gateway():
def test_resume_command_consumes_context_gateway():
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.session import ResumeCommand
@@ -241,7 +246,7 @@ async def test_resume_command_consumes_context_gateway():
graph_gateway=FakeGraphGateway(thread_store=thread_store),
)
await ResumeCommand().execute(ctx, ["abc"])
run_async(ResumeCommand().execute(ctx, ["abc"]))
assert ctx.thread_id == "abc12345"
assert ctx.workspace_dir == "/restored"
@@ -282,7 +287,7 @@ def test_cmd_run_passes_local_graph_gateway(monkeypatch):
assert seen["gateway"].thread_store is thread_store
async def test_langgraph_server_thread_store_delegates_to_sdk_threads():
def test_langgraph_server_thread_store_delegates_to_sdk_threads():
threads = FakeLangGraphThreadsClient(
threads=[
{
@@ -330,7 +335,7 @@ async def test_langgraph_server_thread_store_delegates_to_sdk_threads():
"deleted": await store.delete_thread("abc12345"),
}
result = await _run()
result = run_async(_run())
assert result["created"] == "server-thread"
assert len(threads.created) == 1
@@ -374,7 +379,7 @@ async def test_langgraph_server_thread_store_delegates_to_sdk_threads():
assert threads.deleted == ["abc12345"]
async def test_langgraph_server_thread_store_limit_zero_pages_all_threads():
def test_langgraph_server_thread_store_limit_zero_pages_all_threads():
rows = [
{
"thread_id": f"thread-{index}",
@@ -387,7 +392,7 @@ async def test_langgraph_server_thread_store_limit_zero_pages_all_threads():
client=FakeLangGraphClient(threads),
)
result = await store.list_threads(limit=0)
result = run_async(store.list_threads(limit=0))
assert [row["thread_id"] for row in result] == [
f"thread-{index}" for index in range(_THREAD_SEARCH_LIMIT + 1)
@@ -398,7 +403,7 @@ async def test_langgraph_server_thread_store_limit_zero_pages_all_threads():
]
async def test_langgraph_server_thread_store_positive_limit_uses_single_search():
def test_langgraph_server_thread_store_positive_limit_uses_single_search():
threads = FakeLangGraphThreadsClient(
threads=[
{
@@ -412,7 +417,7 @@ async def test_langgraph_server_thread_store_positive_limit_uses_single_search()
client=FakeLangGraphClient(threads),
)
result = await store.list_threads(limit=2)
result = run_async(store.list_threads(limit=2))
assert [row["thread_id"] for row in result] == ["thread-0", "thread-1"]
assert [(search["limit"], search["offset"]) for search in threads.searches] == [
@@ -420,7 +425,7 @@ async def test_langgraph_server_thread_store_positive_limit_uses_single_search()
]
async def test_langgraph_server_thread_store_prefix_resolution_skips_exact_lookup():
def test_langgraph_server_thread_store_prefix_resolution_skips_exact_lookup():
threads = FakeLangGraphThreadsClient(
threads=[
{
@@ -433,14 +438,14 @@ async def test_langgraph_server_thread_store_prefix_resolution_skips_exact_looku
client=FakeLangGraphClient(threads),
)
result = await store.resolve_thread_id_prefix("abc")
result = run_async(store.resolve_thread_id_prefix("abc"))
assert result == ("abc12345", [])
assert threads.gets == []
assert len(threads.searches) == 1
async def test_langgraph_server_thread_store_prefix_resolution_pages_all_threads():
def test_langgraph_server_thread_store_prefix_resolution_pages_all_threads():
rows = [
{
"thread_id": f"thread-{index}",
@@ -459,7 +464,7 @@ async def test_langgraph_server_thread_store_prefix_resolution_pages_all_threads
client=FakeLangGraphClient(threads),
)
result = await store.resolve_thread_id_prefix("older-thread")
result = run_async(store.resolve_thread_id_prefix("older-thread"))
assert result == ("older-thread-match", [])
assert [(search["limit"], search["offset"]) for search in threads.searches] == [
@@ -468,7 +473,7 @@ async def test_langgraph_server_thread_store_prefix_resolution_pages_all_threads
]
async def test_langgraph_server_thread_store_uuid_resolution_uses_exact_lookup():
def test_langgraph_server_thread_store_uuid_resolution_uses_exact_lookup():
thread_id = "019ed9e4-4253-7f62-b50f-f0470a4b3c9f"
threads = FakeLangGraphThreadsClient(
threads=[
@@ -482,14 +487,14 @@ async def test_langgraph_server_thread_store_uuid_resolution_uses_exact_lookup()
client=FakeLangGraphClient(threads),
)
result = await store.resolve_thread_id_prefix(thread_id)
result = run_async(store.resolve_thread_id_prefix(thread_id))
assert result == (thread_id, [])
assert threads.gets == [thread_id]
assert threads.searches == []
async def test_langgraph_server_thread_store_uuid_resolution_filters_graph_id():
def test_langgraph_server_thread_store_uuid_resolution_filters_graph_id():
thread_id = "019ed9e4-4253-7f62-b50f-f0470a4b3c9f"
threads = FakeLangGraphThreadsClient(
threads=[
@@ -503,7 +508,7 @@ async def test_langgraph_server_thread_store_uuid_resolution_filters_graph_id():
client=FakeLangGraphClient(threads),
)
result = await store.resolve_thread_id_prefix(thread_id)
result = run_async(store.resolve_thread_id_prefix(thread_id))
assert result == (None, [])
assert threads.gets == [thread_id]
@@ -512,7 +517,7 @@ async def test_langgraph_server_thread_store_uuid_resolution_filters_graph_id():
]
async def test_langgraph_server_thread_store_clones_thread_with_metadata():
def test_langgraph_server_thread_store_clones_thread_with_metadata():
clone_metadata = {
"clone_purpose": "memory_extraction",
"source_thread_id": "source-thread",
@@ -529,8 +534,8 @@ async def test_langgraph_server_thread_store_clones_thread_with_metadata():
client=FakeLangGraphClient(threads),
)
cloned_thread_id = await store.clone_thread(
"source-thread", metadata=clone_metadata
cloned_thread_id = run_async(
store.clone_thread("source-thread", metadata=clone_metadata)
)
assert cloned_thread_id == "source-thread-copy"
@@ -547,7 +552,7 @@ async def test_langgraph_server_thread_store_clones_thread_with_metadata():
}
async def test_langgraph_server_thread_store_rejects_copy_without_thread_id():
def test_langgraph_server_thread_store_rejects_copy_without_thread_id():
threads = FakeLangGraphThreadsClient(
threads=[{"thread_id": "source-thread", "metadata": {"graph_id": "agent"}}],
copy_response=None,
@@ -560,10 +565,10 @@ async def test_langgraph_server_thread_store_rejects_copy_without_thread_id():
await store.clone_thread("source-thread")
with pytest.raises(RuntimeError, match="did not return a cloned thread id"):
await _run()
run_async(_run())
async def test_langgraph_server_gateway_clones_thread():
def test_langgraph_server_gateway_clones_thread():
threads = FakeLangGraphThreadsClient(
threads=[{"thread_id": "source-thread", "metadata": {"graph_id": "agent"}}]
)
@@ -573,10 +578,12 @@ async def test_langgraph_server_gateway_clones_thread():
)
)
cloned_thread_id = await gateway.clone_thread(
"source-thread",
metadata={"clone_purpose": "manual"},
target=GraphTarget(graph_id="agent"),
cloned_thread_id = run_async(
gateway.clone_thread(
"source-thread",
metadata={"clone_purpose": "manual"},
target=GraphTarget(graph_id="agent"),
)
)
assert cloned_thread_id == "source-thread-copy"
@@ -585,12 +592,12 @@ async def test_langgraph_server_gateway_clones_thread():
]
async def test_local_graph_gateway_clone_thread_is_explicitly_unsupported():
def test_local_graph_gateway_clone_thread_is_explicitly_unsupported():
async def _run():
await LocalGraphGateway().clone_thread("source-thread")
with pytest.raises(NotImplementedError, match="does not support thread cloning"):
await _run()
run_async(_run())
def test_runtime_gateways_can_use_langgraph_server_backend():
@@ -609,7 +616,7 @@ def test_runtime_gateways_can_use_langgraph_server_backend():
assert gateway.thread_store is runtime_gateways.thread_store
async def test_langgraph_server_gateway_reads_state_values():
def test_langgraph_server_gateway_reads_state_values():
threads = FakeLangGraphThreadsClient(
threads=[{"thread_id": "abc12345", "metadata": {"graph_id": "EvoScientist"}}],
states={"abc12345": {"values": {"async_tasks": {"task-1": {}}}}},
@@ -620,12 +627,12 @@ async def test_langgraph_server_gateway_reads_state_values():
)
)
values = await gateway.get_state_values(GraphTarget(), "abc12345")
values = run_async(gateway.get_state_values(GraphTarget(), "abc12345"))
assert values == {"async_tasks": {"task-1": {}}}
async def test_langgraph_server_gateway_messages_apply_summarization_event():
def test_langgraph_server_gateway_messages_apply_summarization_event():
threads = FakeLangGraphThreadsClient(
threads=[{"thread_id": "abc12345", "metadata": {"graph_id": "EvoScientist"}}],
states={
@@ -651,7 +658,7 @@ async def test_langgraph_server_gateway_messages_apply_summarization_event():
)
)
messages = await gateway.get_thread_messages("abc12345")
messages = run_async(gateway.get_thread_messages("abc12345"))
assert len(messages) == 2
assert isinstance(messages[0], AIMessage)
@@ -660,7 +667,7 @@ async def test_langgraph_server_gateway_messages_apply_summarization_event():
assert messages[1].content == "third"
async def test_langgraph_server_gateway_updates_state_values():
def test_langgraph_server_gateway_updates_state_values():
threads = FakeLangGraphThreadsClient(
threads=[{"thread_id": "abc12345", "metadata": {"graph_id": "EvoScientist"}}],
)
@@ -670,10 +677,12 @@ async def test_langgraph_server_gateway_updates_state_values():
)
)
await gateway.update_state_values(
GraphTarget(),
"abc12345",
{"_summarization_event": {"cutoff_index": 2}},
run_async(
gateway.update_state_values(
GraphTarget(),
"abc12345",
{"_summarization_event": {"cutoff_index": 2}},
)
)
assert threads.state_updates == [
@@ -681,7 +690,7 @@ async def test_langgraph_server_gateway_updates_state_values():
]
async def test_langgraph_server_gateway_streams_root_protocol_events():
def test_langgraph_server_gateway_streams_root_protocol_events():
stream = FakeLangGraphThreadStream(
"abc12345",
events=[
@@ -728,7 +737,7 @@ async def test_langgraph_server_gateway_streams_root_protocol_events():
)
]
events = await _collect()
events = run_async(_collect())
assert len(threads.created) == 1
assert threads.created[0]["thread_id"] == "abc12345"
@@ -795,7 +804,7 @@ def _root_message_finish() -> dict[str, object]:
}
async def _collect_server_gateway_stream(
def _collect_server_gateway_stream(
events: list[dict[str, object]],
*,
state_messages: list[dict[str, object]] | None = None,
@@ -823,11 +832,11 @@ async def _collect_server_gateway_stream(
)
]
return await _collect()
return run_async(_collect())
async def test_langgraph_server_gateway_streams_value_message_snapshots():
events = await _collect_server_gateway_stream(
def test_langgraph_server_gateway_streams_value_message_snapshots():
events = _collect_server_gateway_stream(
[
_value_snapshot([_OLD_AI, _HUMAN]),
_value_snapshot([_OLD_AI, _HUMAN, _NEW_AI]),
@@ -841,8 +850,8 @@ async def test_langgraph_server_gateway_streams_value_message_snapshots():
]
async def test_langgraph_server_gateway_values_do_not_duplicate_message_stream():
events = await _collect_server_gateway_stream(
def test_langgraph_server_gateway_values_do_not_duplicate_message_stream():
events = _collect_server_gateway_stream(
[
_root_text_delta("new"),
_root_message_finish(),
@@ -857,8 +866,8 @@ async def test_langgraph_server_gateway_values_do_not_duplicate_message_stream()
]
async def test_langgraph_server_gateway_ignores_non_root_value_messages():
events = await _collect_server_gateway_stream(
def test_langgraph_server_gateway_ignores_non_root_value_messages():
events = _collect_server_gateway_stream(
[
_value_snapshot(
[{"type": "ai", "content": "subagent text", "id": "subagent-ai"}],
@@ -871,7 +880,7 @@ async def test_langgraph_server_gateway_ignores_non_root_value_messages():
assert events[-1] == {"type": "done", "content": "", "response": ""}
async def test_langgraph_server_gateway_emits_state_interrupt_before_done():
def test_langgraph_server_gateway_emits_state_interrupt_before_done():
stream = FakeLangGraphThreadStream(
"abc12345",
events=[],
@@ -921,15 +930,9 @@ async def test_langgraph_server_gateway_emits_state_interrupt_before_done():
)
]
events = await _collect()
events = run_async(_collect())
assert events == [
{
"type": "tool_call",
"name": "execute",
"args": {"command": "echo hello"},
"id": "tool-1",
},
{
"type": "interrupt",
"interrupt_id": "interrupt-1",
@@ -951,7 +954,7 @@ async def test_langgraph_server_gateway_emits_state_interrupt_before_done():
]
async def test_langgraph_server_gateway_streams_subagent_protocol_events():
def test_langgraph_server_gateway_streams_subagent_protocol_events():
stream = FakeLangGraphThreadStream(
"abc12345",
events=[
@@ -1000,7 +1003,7 @@ async def test_langgraph_server_gateway_streams_subagent_protocol_events():
)
]
events = await _collect()
events = run_async(_collect())
assert events == [
{
@@ -1025,7 +1028,7 @@ async def test_langgraph_server_gateway_streams_subagent_protocol_events():
]
async def test_langgraph_server_gateway_resumes_interrupt_with_thread_stream():
def test_langgraph_server_gateway_resumes_interrupt_with_thread_stream():
from langgraph.types import Command
stream = FakeLangGraphThreadStream(
@@ -1055,7 +1058,7 @@ async def test_langgraph_server_gateway_resumes_interrupt_with_thread_stream():
)
]
events = await _collect()
events = run_async(_collect())
assert stream.run.starts == []
assert stream.run.responses == [
+4 -4
View File
@@ -378,7 +378,7 @@ class TestHitlConfig:
class TestInterruptEventParsing:
async def test_interrupt_from_updates_mode(self):
def test_interrupt_from_updates_mode(self):
"""__interrupt__ in updates mode yields interrupt event."""
interrupt_data = {
"__interrupt__": [
@@ -405,7 +405,7 @@ class TestInterruptEventParsing:
protocol_event("updates", interrupt_data),
]
)
events = await collect_events(agent, message="test", thread_id="thread-1")
events = collect_events(agent, message="test", thread_id="thread-1")
types = [e["type"] for e in events]
assert "interrupt" in types
@@ -415,14 +415,14 @@ class TestInterruptEventParsing:
assert interrupt_ev["action_requests"][0]["name"] == "execute"
assert interrupt_ev["interrupt_id"] == "main"
async def test_updates_without_interrupt_skipped(self):
def test_updates_without_interrupt_skipped(self):
"""Regular updates mode data is skipped as before."""
agent = FakeV3Agent(
[
protocol_event("updates", {"some_node": {"key": "value"}}),
]
)
events = await collect_events(agent, message="test", thread_id="thread-1")
events = collect_events(agent, message="test", thread_id="thread-1")
types = [e["type"] for e in events]
assert "interrupt" not in types
-13
View File
@@ -1,13 +0,0 @@
from EvoScientist.llm.errors import AgentControlError
from EvoScientist.middleware.model_fallback import _is_non_fallbackable
def test_agent_control_error_is_non_fallbackable():
error = AgentControlError(
"INSUFFICIENT_BALANCE",
"balance unavailable",
status_code=403,
)
assert "platform control error" in (_is_non_fallbackable(error) or "")
assert error.model_dump()["code"] == "INSUFFICIENT_BALANCE"
+12 -10
View File
@@ -2,6 +2,8 @@
from unittest.mock import MagicMock, patch
from tests.conftest import run_async as _run
def _ctx():
from EvoScientist.commands.base import CommandContext
@@ -12,15 +14,15 @@ def _ctx():
class TestInstallSkill:
async def test_usage_message_when_no_args(self):
def test_usage_message_when_no_args(self):
from EvoScientist.commands.implementation.skills import InstallSkill
ctx, ui = _ctx()
await InstallSkill().execute(ctx, [])
_run(InstallSkill().execute(ctx, []))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Usage:" in m for m in msgs)
async def test_happy_path(self):
def test_happy_path(self):
from EvoScientist.commands.implementation.skills import InstallSkill
ctx, ui = _ctx()
@@ -33,21 +35,21 @@ class TestInstallSkill:
"path": "/tmp/demo",
},
):
await InstallSkill().execute(ctx, ["./some-path"])
_run(InstallSkill().execute(ctx, ["./some-path"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Installed: demo-skill" in m for m in msgs)
class TestUninstallSkill:
async def test_usage_message_when_no_args(self):
def test_usage_message_when_no_args(self):
from EvoScientist.commands.implementation.skills import UninstallSkill
ctx, ui = _ctx()
await UninstallSkill().execute(ctx, [])
_run(UninstallSkill().execute(ctx, []))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Usage:" in m for m in msgs)
async def test_uninstall_success(self):
def test_uninstall_success(self):
from EvoScientist.commands.implementation.skills import UninstallSkill
ctx, ui = _ctx()
@@ -55,11 +57,11 @@ class TestUninstallSkill:
"EvoScientist.tools.skills_manager.uninstall_skill",
return_value={"success": True},
):
await UninstallSkill().execute(ctx, ["demo-skill"])
_run(UninstallSkill().execute(ctx, ["demo-skill"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Uninstalled: demo-skill" in m for m in msgs)
async def test_uninstall_failure(self):
def test_uninstall_failure(self):
from EvoScientist.commands.implementation.skills import UninstallSkill
ctx, ui = _ctx()
@@ -67,6 +69,6 @@ class TestUninstallSkill:
"EvoScientist.tools.skills_manager.uninstall_skill",
return_value={"success": False, "error": "not found"},
):
await UninstallSkill().execute(ctx, ["missing"])
_run(UninstallSkill().execute(ctx, ["missing"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Failed: not found" in m for m in msgs)
+10 -10
View File
@@ -19,7 +19,7 @@ async def _agen(items):
yield item
async def test_writes_each_event_as_one_jsonl_line():
def test_writes_each_event_as_one_jsonl_line(run_async):
"""Each event dict is serialized to exactly one JSON line, in order."""
events = [
{"type": "thinking", "content": "hmm", "id": 0},
@@ -34,7 +34,7 @@ async def test_writes_each_event_as_one_jsonl_line():
]
out = io.StringIO()
await write_events_as_json(_agen(events), out)
run_async(write_events_as_json(_agen(events), out))
lines = out.getvalue().splitlines()
assert len(lines) == len(events)
@@ -43,7 +43,7 @@ async def test_writes_each_event_as_one_jsonl_line():
assert parsed[2]["args"] == {"path": "a.md"}
async def test_returns_final_response_from_done_event():
def test_returns_final_response_from_done_event(run_async):
"""The sink returns the response text carried by the terminal `done` event."""
events = [
{"type": "text", "content": "partial"},
@@ -51,12 +51,12 @@ async def test_returns_final_response_from_done_event():
]
out = io.StringIO()
result = await write_events_as_json(_agen(events), out)
result = run_async(write_events_as_json(_agen(events), out))
assert result == "the answer"
async def test_non_serializable_arg_does_not_crash_the_stream():
def test_non_serializable_arg_does_not_crash_the_stream(run_async):
"""A non-JSON-serializable value degrades to its str form instead of raising."""
class Weird:
@@ -72,7 +72,7 @@ async def test_non_serializable_arg_does_not_crash_the_stream():
]
out = io.StringIO()
await write_events_as_json(_agen(events), out)
run_async(write_events_as_json(_agen(events), out))
lines = out.getvalue().splitlines()
# Both lines must be valid JSON; the non-serializable value falls back to str.
@@ -80,7 +80,7 @@ async def test_non_serializable_arg_does_not_crash_the_stream():
assert first["args"]["obj"] == "WEIRD"
async def test_stream_json_sources_events_from_gateway():
def test_stream_json_sources_events_from_gateway(run_async):
"""stream_json pulls events from gateway.stream_events(request) and serializes
them — it does not reach past the gateway abstraction."""
seen: dict[str, object] = {}
@@ -98,7 +98,7 @@ async def test_stream_json_sources_events_from_gateway():
return _agen(events)
out = io.StringIO()
result = await stream_json(_FakeGateway(), object(), out=out)
result = run_async(stream_json(_FakeGateway(), object(), out=out))
assert result == "hi"
assert "request" in seen # the request was forwarded to the gateway
@@ -106,7 +106,7 @@ async def test_stream_json_sources_events_from_gateway():
assert types == ["text", "done"]
async def test_stream_json_propagates_gateway_errors():
def test_stream_json_propagates_gateway_errors(run_async):
"""An error from the gateway stream propagates out of stream_json so the CLI
dispatch can turn it into a clean exit."""
@@ -124,4 +124,4 @@ async def test_stream_json_propagates_gateway_errors():
out = io.StringIO()
with pytest.raises(RuntimeError, match="boom"):
await stream_json(_FakeGateway(), object(), out=out)
run_async(stream_json(_FakeGateway(), object(), out=out))
+141
View File
@@ -7,6 +7,7 @@ from __future__ import annotations
from unittest.mock import patch
from langchain_core.messages import AIMessage, HumanMessage
from starlette.testclient import TestClient
from EvoScientist.config import EvoScientistConfig
@@ -161,3 +162,143 @@ def test_ollama_discovery_skipped_when_base_url_absent():
{"name": n, "model_id": m, "provider": p}
for n, m, p in list_models_by_provider()
]
def test_final_answer_extracts_latest_ai_text_blocks():
async def fake_metadata(_thread_id):
return {"updated_at": "2026-07-06T14:14:53+00:00"}
async def fake_messages(_thread_id):
return [
HumanMessage(content="question"),
AIMessage(content="old answer"),
AIMessage(
content=[
{"type": "reasoning", "text": "internal"},
{"type": "text", "text": "Part A"},
{"type": "tool_use", "name": "search"},
{"type": "output_text", "text": "Part B"},
]
),
]
async def fake_runtime(_request, _thread_id):
return {
"found": True,
"complete": True,
"completed_at": "2026-07-06T14:15:00+00:00",
}
with (
patch(
"EvoScientist.langgraph_dev.http._get_thread_metadata_for_http",
new=fake_metadata,
),
patch(
"EvoScientist.langgraph_dev.http._get_thread_messages_for_http",
new=fake_messages,
),
patch(
"EvoScientist.langgraph_dev.http._read_thread_runtime_state",
new=fake_runtime,
),
):
resp = client.get("/api/threads/thread-1/final-answer")
assert resp.status_code == 200
assert resp.json() == {
"content": "Part A\n\nPart B",
"completed_at": "2026-07-06T14:15:00+00:00",
"complete": True,
}
def test_final_answer_skips_tool_selection_json_text():
async def fake_metadata(_thread_id):
return {"updated_at": "2026-07-06T14:14:53+00:00"}
async def fake_messages(_thread_id):
return [
HumanMessage(content="question"),
AIMessage(content="stable answer"),
AIMessage(
content=(
'{"tools":["search_papers","get_abstract"]}'
'{"tools":["web_search_exa"]}'
)
),
]
async def fake_runtime(_request, _thread_id):
return {
"found": True,
"complete": True,
"completed_at": "2026-07-06T14:15:00+00:00",
}
with (
patch(
"EvoScientist.langgraph_dev.http._get_thread_metadata_for_http",
new=fake_metadata,
),
patch(
"EvoScientist.langgraph_dev.http._get_thread_messages_for_http",
new=fake_messages,
),
patch(
"EvoScientist.langgraph_dev.http._read_thread_runtime_state",
new=fake_runtime,
),
):
resp = client.get("/api/threads/thread-1/final-answer")
assert resp.status_code == 200
assert resp.json()["content"] == "stable answer"
def test_final_answer_returns_404_for_unknown_thread():
async def fake_metadata(_thread_id):
return None
with patch(
"EvoScientist.langgraph_dev.http._get_thread_metadata_for_http",
new=fake_metadata,
):
resp = client.get("/api/threads/missing/final-answer")
assert resp.status_code == 404
assert resp.json() == {"error": "thread not found"}
def test_final_answer_does_not_mark_complete_when_runtime_state_fails():
async def fake_metadata(_thread_id):
return {"updated_at": "2026-07-06T14:14:53+00:00"}
async def fake_messages(_thread_id):
return [AIMessage(content="checkpoint answer")]
async def fake_runtime(_request, _thread_id):
raise RuntimeError("langgraph runtime unavailable")
with (
patch(
"EvoScientist.langgraph_dev.http._get_thread_metadata_for_http",
new=fake_metadata,
),
patch(
"EvoScientist.langgraph_dev.http._get_thread_messages_for_http",
new=fake_messages,
),
patch(
"EvoScientist.langgraph_dev.http._read_thread_runtime_state",
new=fake_runtime,
),
):
resp = client.get("/api/threads/thread-1/final-answer")
assert resp.status_code == 200
assert resp.json() == {
"content": "checkpoint answer",
"completed_at": None,
"complete": False,
}
-72
View File
@@ -8,7 +8,6 @@ to be available.
from __future__ import annotations
import dataclasses
import sys
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
@@ -33,77 +32,6 @@ def reset_module_state():
manager._LOG_OFFSET_AT_START = 0
# =============================================================================
# langgraph CLI resolution
# =============================================================================
class TestLanggraphCliResolution:
def _make_executable(self, path):
path.write_text("#!/bin/sh\n", encoding="utf-8")
path.chmod(0o755)
def test_prefers_current_python_environment_over_path(self, tmp_path, monkeypatch):
local_bin = tmp_path / "local" / "bin"
local_bin.mkdir(parents=True)
local_langgraph = local_bin / "langgraph"
self._make_executable(local_langgraph)
path_bin = tmp_path / "path" / "bin"
path_bin.mkdir(parents=True)
path_langgraph = path_bin / "langgraph"
self._make_executable(path_langgraph)
monkeypatch.setattr(sys, "executable", str(local_bin / "python"))
monkeypatch.setattr(
manager.shutil,
"which",
lambda command: str(path_langgraph) if command == "langgraph" else None,
)
assert manager._langgraph_exe() == str(local_langgraph)
def test_falls_back_to_path_when_environment_binary_missing(
self, tmp_path, monkeypatch
):
path_bin = tmp_path / "path" / "bin"
path_bin.mkdir(parents=True)
path_langgraph = path_bin / "langgraph"
self._make_executable(path_langgraph)
monkeypatch.setattr(
sys, "executable", str(tmp_path / "local" / "bin" / "python")
)
monkeypatch.setattr(
manager.shutil,
"which",
lambda command: str(path_langgraph) if command == "langgraph" else None,
)
assert manager._langgraph_exe() == str(path_langgraph)
def test_checks_windows_suffix_next_to_current_python(self, tmp_path, monkeypatch):
scripts_dir = tmp_path / "Scripts"
scripts_dir.mkdir()
local_langgraph = scripts_dir / "langgraph.exe"
self._make_executable(local_langgraph)
path_bin = tmp_path / "path" / "bin"
path_bin.mkdir(parents=True)
path_langgraph = path_bin / "langgraph.exe"
self._make_executable(path_langgraph)
monkeypatch.setattr(sys, "executable", str(scripts_dir / "python.exe"))
monkeypatch.setattr(manager.os, "name", "nt", raising=False)
monkeypatch.setattr(
manager.shutil,
"which",
lambda command: str(path_langgraph) if command == "langgraph" else None,
)
assert manager._langgraph_exe() == str(local_langgraph)
# =============================================================================
# is_langgraph_dev_running
# =============================================================================
@@ -1,152 +0,0 @@
"""Regression tests for the langgraph_api SchemaGenerator silencing patch.
Reproducer: mounting our ``/api/models`` custom Starlette app makes
langgraph_api call ``update_openapi_spec`` at startup, which iterates
EVERY route (ours + upstream's). Endpoints whose docstrings aren't
valid YAML hit a warning + traceback in the deploy log — purely noise,
since the existing fallback path already produces a usable schema
entry. The patch keeps the fallback shape but silences the log spam.
"""
from __future__ import annotations
import os
# ``langgraph_api.config`` reads several required env vars at import
# time via starlette's ``Config(...)`` helper. We don't actually use the
# DB or Redis here — any non-empty string keeps the loader happy.
os.environ.setdefault("DATABASE_URI", "sqlite:///:memory:")
os.environ.setdefault("REDIS_URI", "redis://localhost:6379")
# Importing patches.py applies the eager module-level monkey-patch.
import langgraph_api.utils as _lgapi_utils
import EvoScientist.llm.patches as _patches
# Re-invoke the patch after env vars are set. Required because earlier test
# modules (e.g. test_llm.py) import patches.py *without* DATABASE_URI/
# REDIS_URI, which makes ``langgraph_api.utils`` fail to import inside the
# patch's bare ``except``; the loader swallows it and the flag stays False
# forever (Python won't re-run module-level code on subsequent imports).
# The patch function is idempotent (early-return on the flag), so calling
# it here is a no-op when the patch already landed and a successful retry
# when the prior import failed.
_patches._patch_langgraph_schema_generator_silence_warnings()
class _FakeEndpoint:
"""Minimal Starlette-like endpoint info for the schema generator."""
def __init__(self, path: str, method: str, func):
self.path = path
self.http_method = method
self.func = func
class _DocstringFixture:
"""The kinds of docstrings the patched generator must handle."""
@staticmethod
def prose_with_colon():
"""Endpoint summary.
Query params:
id: The thing you want.
"""
@staticmethod
def valid_yaml():
"""
summary: A valid YAML docstring.
description: Stays structured.
"""
@staticmethod
def no_docstring():
pass
def _generator():
return _lgapi_utils.SchemaGenerator(
{"openapi": "3.1.0", "info": {"title": "test", "version": "0"}}
)
def test_prose_docstring_no_longer_logs_warning():
"""The patched ``parse_docstring`` must silence upstream's structlog
WARNING when ``yaml.safe_load`` fails on a prose docstring.
Inverts the patch first to prove the fixture actually trips
``yaml.safe_load`` — without this baseline assertion the test would
pass vacuously if the fixture stopped triggering the failure path
(e.g. if upstream changed how docstrings are pre-processed).
"""
from structlog.testing import capture_logs
gen = _generator()
endpoint = _FakeEndpoint("/x", "get", _DocstringFixture.prose_with_colon)
gen.get_endpoints = lambda _routes: [endpoint]
patched_parse = _lgapi_utils.SchemaGenerator.parse_docstring
# Phase 1: baseline. Drop the subclass override so MRO falls through
# to Starlette's BaseSchemaGenerator.parse_docstring, which is what
# production hits before our patch installs.
del _lgapi_utils.SchemaGenerator.parse_docstring
try:
with capture_logs() as baseline_records:
gen.get_schema([])
finally:
_lgapi_utils.SchemaGenerator.parse_docstring = patched_parse
baseline_warnings = [r for r in baseline_records if r.get("log_level") == "warning"]
assert any(
"Unable to parse docstring" in r.get("event", "") for r in baseline_warnings
), "fixture no longer trips parse_docstring — test would pass vacuously"
# Phase 2: with the patch reinstated, the same call must emit no
# warning records.
with capture_logs() as patched_records:
schema = gen.get_schema([])
assert [r for r in patched_records if r.get("log_level") == "warning"] == []
# Schema still has the fallback shape — fixture's prose becomes the
# description verbatim (with leading/trailing whitespace from the
# docstring preserved by upstream's fallback path).
entry = schema["paths"]["/x"]["get"]
assert "description" in entry
assert "Query params" in entry["description"]
def test_valid_yaml_docstring_keeps_structured_parse():
"""Endpoints with parseable YAML keep their structured metadata —
we only changed the failure branch, not the success path.
"""
gen = _generator()
endpoint = _FakeEndpoint("/y", "get", _DocstringFixture.valid_yaml)
gen.get_endpoints = lambda _routes: [endpoint]
schema = gen.get_schema([])
entry = schema["paths"]["/y"]["get"]
assert entry.get("summary") == "A valid YAML docstring."
assert entry.get("description") == "Stays structured."
def test_no_docstring_still_handled():
"""Endpoints with ``__doc__ = None`` must not raise — fallback uses
empty string for ``description``.
"""
gen = _generator()
endpoint = _FakeEndpoint("/z", "get", _DocstringFixture.no_docstring)
gen.get_endpoints = lambda _routes: [endpoint]
schema = gen.get_schema([])
entry = schema["paths"]["/z"]["get"]
# Either description="" (fallback path) or structured (if YAML parse
# of None happens to succeed somehow — implementation detail).
# The contract is just "no exception, entry exists".
assert isinstance(entry, dict)
def test_patch_flag_set():
from EvoScientist.llm.patches import _langgraph_schema_silenced_patched
assert _langgraph_schema_silenced_patched is True
+7 -606
View File
@@ -1,6 +1,5 @@
"""Tests for EvoScientist LLM module."""
from types import SimpleNamespace
from unittest.mock import patch
import pytest
@@ -161,68 +160,6 @@ class TestGetModelInfo:
class TestGetChatModel:
@patch("EvoScientist.llm.models.init_chat_model")
def test_uses_host_model_resolver(self, mock_init):
"""An embedding host can provide model routing without a core dependency."""
from EvoScientist.runtime_integrations import (
configure_runtime_integrations,
reset_runtime_integrations,
)
mock_init.return_value = "mock_model"
resolved = SimpleNamespace(
provider_name="relay-a",
model_id="model-a",
protocol="openai",
api_key="sk-host",
base_url="https://relay.example/v1/",
params={"max_tokens": 8192, "_default_headers": {"X-Relay": "a"}},
supports_reasoning=False,
)
configure_runtime_integrations(model_resolver=lambda model, provider: resolved)
try:
assert get_chat_model("alias-a") == "mock_model"
finally:
reset_runtime_integrations()
call_kwargs = mock_init.call_args.kwargs
assert call_kwargs["model"] == "model-a"
assert call_kwargs["model_provider"] == "openai"
assert call_kwargs["api_key"] == "sk-host"
assert call_kwargs["base_url"] == "https://relay.example/v1"
assert call_kwargs["max_tokens"] == 8192
assert call_kwargs["default_headers"] == {"X-Relay": "a"}
assert "reasoning" not in call_kwargs
@patch("EvoScientist.llm.models._patch_openai_compat_content")
@patch("EvoScientist.llm.models.init_chat_model")
def test_host_openai_provider_with_custom_base_uses_compat_patch(
self, mock_init, mock_compat
):
from EvoScientist.runtime_integrations import (
configure_runtime_integrations,
reset_runtime_integrations,
)
model_instance = object()
mock_init.return_value = model_instance
resolved = SimpleNamespace(
provider_name="openai",
model_id="gpt-5.5",
protocol="openai",
api_key="sk-host",
base_url="https://relay.example/v1",
params={},
supports_reasoning=True,
)
configure_runtime_integrations(model_resolver=lambda model, provider: resolved)
try:
get_chat_model("gpt-5.5", provider="openai")
finally:
reset_runtime_integrations()
mock_compat.assert_called_once_with(model_instance, hoist_tool_media=True)
@patch("EvoScientist.llm.models.init_chat_model")
def test_uses_default_model_when_none(self, mock_init):
"""Test that get_chat_model uses default model when model=None."""
@@ -292,39 +229,6 @@ class TestGetChatModel:
assert call_kwargs["temperature"] == 0.7
assert call_kwargs["max_tokens"] == 1000
@patch("EvoScientist.llm.models.init_chat_model")
def test_drops_unsupported_legacy_model_kwargs(self, mock_init):
mock_init.return_value = "mock_model"
get_chat_model(
"gpt-5-nano",
provider="openai",
sanitize_openai_sdk_headers=True,
model_kwargs={"sanitize_openai_sdk_headers": False, "custom": "value"},
)
call_kwargs = mock_init.call_args.kwargs
assert "sanitize_openai_sdk_headers" not in call_kwargs
assert call_kwargs["model_kwargs"] == {"custom": "value"}
@patch("EvoScientist.llm.models.init_chat_model")
def test_explicit_credentials_override_environment(self, mock_init, monkeypatch):
"""Host-provided credentials take precedence over process defaults."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENAI_API_KEY", "sk-environment")
monkeypatch.setenv("OPENAI_BASE_URL", "https://environment.example/v1")
get_chat_model(
"gpt-5-nano",
provider="openai",
api_key="sk-explicit",
base_url="https://explicit.example/v1",
)
call_kwargs = mock_init.call_args.kwargs
assert call_kwargs["api_key"] == "sk-explicit"
assert call_kwargs["base_url"] == "https://explicit.example/v1"
@patch("EvoScientist.llm.models.init_chat_model")
def test_infers_openai_from_gpt_prefix(self, mock_init):
"""Test that OpenAI is inferred from gpt- prefix."""
@@ -535,200 +439,6 @@ class TestThirdPartyRouting:
call_kwargs = mock_init.call_args[1]
assert call_kwargs["reasoning"] == {"effort": "medium", "summary": "auto"}
# --- OpenRouter app attribution (issue #339) ---
_APP_ATTR_ENV = (
"EVOSCIENTIST_OPENROUTER_HTTP_REFERER",
"EVOSCIENTIST_OPENROUTER_APP_TITLE",
"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES",
)
@patch("EvoScientist.llm.models.init_chat_model")
def test_openrouter_app_attribution_defaults(self, mock_init, monkeypatch):
"""OpenRouter init should carry EvoScientist's default app attribution."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
# Isolate from any leaked env overrides so we assert the built-in defaults.
for _env in self._APP_ATTR_ENV:
monkeypatch.delenv(_env, raising=False)
get_chat_model("x-ai/grok-4.3", provider="openrouter")
call_kwargs = mock_init.call_args[1]
assert call_kwargs["app_url"] == "https://github.com/EvoScientist/EvoScientist"
assert call_kwargs["app_title"] == "EvoScientist"
# Must be a list[str] (not the comma string) — langchain-openrouter joins it.
assert call_kwargs["app_categories"] == ["creative-writing", "personal-agent"]
@patch("EvoScientist.llm.models.init_chat_model")
def test_openrouter_app_attribution_from_env(self, mock_init, monkeypatch):
"""Env vars should override the default app attribution values."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_HTTP_REFERER", "https://acme.test")
monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_TITLE", "Acme")
# Include a space to prove each category is stripped.
monkeypatch.setenv(
"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "cli-agent, programming-app"
)
get_chat_model("x-ai/grok-4.3", provider="openrouter")
call_kwargs = mock_init.call_args[1]
assert call_kwargs["app_url"] == "https://acme.test"
assert call_kwargs["app_title"] == "Acme"
assert call_kwargs["app_categories"] == ["cli-agent", "programming-app"]
@patch("EvoScientist.llm.models.init_chat_model")
def test_openrouter_app_attribution_user_override_not_clobbered(
self, mock_init, monkeypatch
):
"""Caller-supplied attribution kwargs must beat both env and defaults."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
# Env is also set, to prove an explicit kwarg outranks the env override
# (not just the built-in default).
monkeypatch.setenv(
"EVOSCIENTIST_OPENROUTER_HTTP_REFERER", "https://env.example"
)
monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_TITLE", "EnvTitle")
monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "env-cat")
get_chat_model(
"x-ai/grok-4.3",
provider="openrouter",
app_url="https://mine.example",
app_title="MyApp",
app_categories=["only-this"],
)
call_kwargs = mock_init.call_args[1]
assert call_kwargs["app_url"] == "https://mine.example"
assert call_kwargs["app_title"] == "MyApp"
# An explicit list is preserved verbatim, not re-split.
assert call_kwargs["app_categories"] == ["only-this"]
@patch("EvoScientist.llm.models.init_chat_model")
def test_non_openrouter_providers_get_no_app_attribution(
self, mock_init, monkeypatch
):
"""Only the openrouter provider should receive app-attribution kwargs."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-real")
monkeypatch.setenv("OLLAMA_BASE_URL", "http://localhost:11434")
for model, provider in (
("claude-sonnet-4-6", "anthropic"),
("llama3.1:8b", "ollama"),
):
get_chat_model(model, provider=provider)
call_kwargs = mock_init.call_args[1]
assert "app_url" not in call_kwargs
assert "app_title" not in call_kwargs
assert "app_categories" not in call_kwargs
@patch("EvoScientist.llm.models.init_chat_model")
def test_openrouter_app_attribution_coexists_with_reasoning_and_cache(
self, mock_init, monkeypatch
):
"""Attribution must not disturb reasoning or Anthropic prompt caching."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False)
monkeypatch.delenv(
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE", raising=False
)
for _env in self._APP_ATTR_ENV:
monkeypatch.delenv(_env, raising=False)
get_chat_model("claude-sonnet-4.6", provider="openrouter")
call_kwargs = mock_init.call_args[1]
# Existing behavior intact.
assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"}
assert call_kwargs["model_kwargs"]["cache_control"] == {"type": "ephemeral"}
# Attribution added alongside.
assert call_kwargs["app_url"] == "https://github.com/EvoScientist/EvoScientist"
assert call_kwargs["app_title"] == "EvoScientist"
assert call_kwargs["app_categories"] == ["creative-writing", "personal-agent"]
@patch("EvoScientist.llm.models.init_chat_model")
def test_openrouter_app_categories_env_strips_blank_items(
self, mock_init, monkeypatch
):
"""A messy comma value (stray commas / spaces) yields a clean list."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", "a,, b ")
get_chat_model("x-ai/grok-4.3", provider="openrouter")
assert mock_init.call_args[1]["app_categories"] == ["a", "b"]
@patch("EvoScientist.llm.models.init_chat_model")
def test_openrouter_app_categories_capped_to_per_request_limit(
self, mock_init, monkeypatch
):
"""Over-configuring categories caps to the first N and warns the user."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
monkeypatch.setenv(
"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES",
"cli-agent,programming-app,personal-agent,writing-assistant",
)
with pytest.warns(UserWarning, match="at most 2 app categories"):
get_chat_model("x-ai/grok-4.3", provider="openrouter")
# OpenRouter honors at most 2 per request, so only the first 2 are sent.
assert mock_init.call_args[1]["app_categories"] == [
"cli-agent",
"programming-app",
]
@patch("EvoScientist.llm.models.init_chat_model")
def test_openrouter_app_categories_all_separators_omit_kwarg(
self, mock_init, monkeypatch
):
"""A categories value with no real items omits the kwarg entirely."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
monkeypatch.setenv("EVOSCIENTIST_OPENROUTER_APP_CATEGORIES", " , , ")
get_chat_model("x-ai/grok-4.3", provider="openrouter")
# No app_categories kwarg at all — not an empty list (which the library
# would reject / send as an empty header).
assert "app_categories" not in mock_init.call_args[1]
def test_openrouter_app_attribution_lands_on_real_model(self, monkeypatch):
"""Build a REAL ChatOpenRouter (no mock) and assert the attribution
values land on the instance rather than being silently dumped into
model_kwargs.
The mocked tests above assert on the kwargs handed to init_chat_model,
so they cannot catch a param-name typo or a langchain-openrouter version
that accepts these only as passthrough model params (which the library
does with a warning, not an error). This test is the guard for both.
"""
from langchain_openrouter import ChatOpenRouter
monkeypatch.setenv("OPENROUTER_API_KEY", "or-key")
for _env in self._APP_ATTR_ENV:
monkeypatch.delenv(_env, raising=False)
model = get_chat_model("x-ai/grok-4.3", provider="openrouter")
assert isinstance(model, ChatOpenRouter)
assert model.app_url == "https://github.com/EvoScientist/EvoScientist"
assert model.app_title == "EvoScientist"
assert model.app_categories == ["creative-writing", "personal-agent"]
# Not silently swallowed into model_kwargs (the passthrough failure mode).
model_kwargs = model.model_kwargs or {}
assert "app_url" not in model_kwargs
assert "app_title" not in model_kwargs
assert "app_categories" not in model_kwargs
@patch("EvoScientist.llm.models.init_chat_model")
def test_openrouter_anthropic_prompt_cache_enabled_by_default(
self, mock_init, monkeypatch
@@ -1247,153 +957,6 @@ class TestPatchOpenAICompatContent:
model._astream = AsyncMock()
return model
def test_missing_tool_call_ids_are_repaired_without_mutating_history(self):
from langchain_core.messages import AIMessage, ToolMessage
from EvoScientist.llm.patches import _ensure_openai_tool_call_ids
ai = AIMessage(
content=[{"type": "tool_call", "id": "", "name": "execute", "args": {}}],
tool_calls=[{"id": "", "name": "execute", "args": {}}],
)
tool = ToolMessage(content="ok", tool_call_id="")
normalized = _ensure_openai_tool_call_ids([ai, tool])
call_id = normalized[0].tool_calls[0]["id"]
assert call_id.startswith("call_")
assert normalized[0].content[0]["id"] == call_id
assert normalized[1].tool_call_id == call_id
assert ai.tool_calls[0]["id"] == ""
assert tool.tool_call_id == ""
def test_missing_parallel_tool_call_ids_are_stable_and_ordered(self):
from langchain_core.messages import AIMessage, ToolMessage
from EvoScientist.llm.patches import _ensure_openai_tool_call_ids
messages = [
AIMessage(
id="assistant-1",
content="",
tool_calls=[
{"id": "", "name": "read_file", "args": {}},
{"id": "", "name": "execute", "args": {}},
],
),
ToolMessage(content="file", tool_call_id=""),
ToolMessage(content="command", tool_call_id=""),
]
first = _ensure_openai_tool_call_ids(messages)
second = _ensure_openai_tool_call_ids(messages)
call_ids = [call["id"] for call in first[0].tool_calls]
assert call_ids == [call["id"] for call in second[0].tool_calls]
assert len(set(call_ids)) == 2
assert [message.tool_call_id for message in first[1:]] == call_ids
def test_content_tool_block_is_normalized_to_parsed_call(self):
from langchain_core.messages import AIMessage, ToolMessage
from EvoScientist.llm.patches import _ensure_openai_tool_call_ids
normalized = _ensure_openai_tool_call_ids(
[
AIMessage(
content=[
{
"type": "tool_call",
"id": "wrong-id",
"name": "wrong-name",
"args": {},
}
],
tool_calls=[{"id": "call-1", "name": "execute", "args": {}}],
),
ToolMessage(content="ok", tool_call_id="call-1"),
]
)
assert normalized[0].content[0]["id"] == "call-1"
assert normalized[0].content[0]["name"] == "execute"
def test_invalid_tool_call_is_not_replayed_to_responses_api(self):
from langchain_core.messages import AIMessage, HumanMessage
from langchain_openai.chat_models.base import _construct_responses_api_input
from EvoScientist.llm.patches import _sanitize_messages
invalid = AIMessage(
content=[
{"type": "reasoning", "reasoning": "partial"},
{
"type": "tool_call",
"id": None,
"name": "execute",
"args": '{"command":',
},
],
invalid_tool_calls=[
{
"type": "invalid_tool_call",
"id": None,
"name": "execute",
"args": '{"command":',
"error": "Failed to parse tool call arguments as JSON",
}
],
)
normalized = _sanitize_messages([invalid, HumanMessage(content="retry")])
payload = _construct_responses_api_input(normalized)
assert all(item.get("type") != "function_call" for item in payload)
assert [message.type for message in normalized] == ["human"]
def test_invalid_tool_call_preserves_replayable_assistant_text(self):
from langchain_core.messages import AIMessage
from EvoScientist.llm.patches import _sanitize_messages
invalid = AIMessage(
content="I could not finish the tool request.",
invalid_tool_calls=[
{
"type": "invalid_tool_call",
"id": None,
"name": "execute",
"args": "{",
"error": "bad json",
}
],
)
normalized = _sanitize_messages([invalid])
assert len(normalized) == 1
assert normalized[0].content == "I could not finish the tool request."
assert normalized[0].invalid_tool_calls == []
def test_orphan_tool_results_and_unanswered_calls_are_removed(self):
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
from EvoScientist.llm.patches import _sanitize_messages
messages = [
ToolMessage(content="orphan", tool_call_id="missing"),
AIMessage(
content="waiting",
tool_calls=[{"id": "call_unanswered", "name": "execute", "args": {}}],
),
HumanMessage(content="continue"),
]
normalized = _sanitize_messages(messages)
assert [message.type for message in normalized] == ["ai", "human"]
assert normalized[0].tool_calls == []
def test_generate_flattened(self):
from langchain_core.messages import HumanMessage
@@ -1409,6 +972,7 @@ class TestPatchOpenAICompatContent:
called_msgs = orig.call_args[0][0]
assert called_msgs[0].content == "hello"
@pytest.mark.anyio
async def test_agenerate_flattened(self):
from langchain_core.messages import HumanMessage
@@ -1439,6 +1003,7 @@ class TestPatchOpenAICompatContent:
called_msgs = orig.call_args[0][0]
assert called_msgs[0].content == "hello"
@pytest.mark.anyio
async def test_astream_flattened(self):
from langchain_core.messages import HumanMessage
@@ -1479,6 +1044,7 @@ class TestPatchOpenAICompatContent:
called_msgs = orig.call_args[0][0]
assert called_msgs[0].content == [{"type": "text", "text": "see"}, img]
@pytest.mark.anyio
async def test_agenerate_preserves_media(self):
from langchain_core.messages import HumanMessage
@@ -1511,6 +1077,7 @@ class TestPatchOpenAICompatContent:
called_msgs = orig.call_args[0][0]
assert called_msgs[0].content == [{"type": "text", "text": "see"}, img]
@pytest.mark.anyio
async def test_astream_preserves_media(self):
from langchain_core.messages import HumanMessage
@@ -1983,6 +1550,7 @@ class TestNoVisionFallback:
assert out == ["x", "y"]
assert len(calls) == 2
@pytest.mark.anyio
async def test_astream_falls_back(self):
from unittest.mock import MagicMock
@@ -2675,38 +2243,6 @@ class TestPatchOpenrouterStripResponsesReasoning:
class TestAutoConfig:
@pytest.fixture(autouse=True)
def _clear_reasoning_effort_env(self, monkeypatch):
"""Keep auto-config tests independent of the developer environment."""
monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False)
@patch("EvoScientist.llm.models.init_chat_model")
def test_internal_sentinels_disable_auto_reasoning(self, mock_init, monkeypatch):
"""Internal callers can disable reasoning without leaking sentinels."""
mock_init.return_value = "mock_model"
monkeypatch.delenv("ANTHROPIC_BASE_URL", raising=False)
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
for model, provider in (
("claude-sonnet-4-6", "anthropic"),
("gpt-5-nano", "openai"),
("gemini-2.5-flash", "google-genai"),
("llama3.1:8b", "ollama"),
):
mock_init.reset_mock()
get_chat_model(
model,
provider=provider,
_disable_reasoning=True,
_disable_thinking=True,
)
call_kwargs = mock_init.call_args.kwargs
assert "_disable_reasoning" not in call_kwargs
assert "_disable_thinking" not in call_kwargs
assert "reasoning" not in call_kwargs
assert "thinking" not in call_kwargs
assert "include_thoughts" not in call_kwargs
@patch("EvoScientist.llm.models.init_chat_model")
def test_anthropic_4_5_thinking(self, mock_init, monkeypatch):
"""Anthropic 4-5 models get enabled thinking with budget."""
@@ -2796,7 +2332,6 @@ class TestAutoConfig:
"""gpt-5.4+ and codex models get xhigh reasoning."""
mock_init.return_value = "mock_model"
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False)
get_chat_model("gpt-5.4", provider="openai")
assert mock_init.call_args[1]["reasoning"] == {
@@ -2816,26 +2351,6 @@ class TestAutoConfig:
"summary": "auto",
}
get_chat_model("gpt-5.6-sol", provider="openai")
assert mock_init.call_args[1]["reasoning"] == {
"effort": "xhigh",
"summary": "auto",
}
@patch("EvoScientist.llm.models.init_chat_model")
def test_openai_reasoning_effort_from_env(self, mock_init, monkeypatch):
"""Native OpenAI reasoning effort should be configurable via env var."""
mock_init.return_value = "mock_model"
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
monkeypatch.setenv("EVOSCIENTIST_REASONING_EFFORT", "high")
get_chat_model("gpt-5.5", provider="openai")
assert mock_init.call_args[1]["reasoning"] == {
"effort": "high",
"summary": "auto",
}
@patch("EvoScientist.llm.models.init_chat_model")
def test_openai_reasoning_high_fallback(self, mock_init, monkeypatch):
"""Other OpenAI models get high reasoning effort."""
@@ -2867,8 +2382,8 @@ class TestAutoConfig:
assert call_kwargs["model_provider"] == "openai"
assert call_kwargs["base_url"] == "http://127.0.0.1:8000/codex/v1"
assert call_kwargs["api_key"] == "ccproxy-oauth"
# ccproxy uses the Responses API, so reasoning configuration is valid.
assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"}
# Proxy mode: reasoning skipped (ccproxy untested)
assert "reasoning" not in call_kwargs
# Proxy mode: Responses API (bypasses format chain), streaming ON
assert call_kwargs["use_responses_api"] is True
assert "streaming" not in call_kwargs
@@ -2901,120 +2416,6 @@ class TestAutoConfig:
call_kwargs = mock_init.call_args[1]
assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"}
assert "use_responses_api" not in call_kwargs
assert "default_headers" not in call_kwargs
@patch(
"EvoScientist.llm.models._installed_codex_client_version",
return_value="0.144.1",
)
@patch("EvoScientist.llm.models.init_chat_model")
def test_openai_ccproxy_codex_client_headers(
self, mock_init, mock_installed_version, monkeypatch
):
"""ccproxy Codex mode sends Codex-CLI-shaped client headers."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENAI_BASE_URL", "http://127.0.0.1:8000/codex/v1")
monkeypatch.setenv("OPENAI_API_KEY", "ccproxy-oauth")
monkeypatch.delenv("EVOSCIENTIST_CODEX_CLIENT_VERSION", raising=False)
get_chat_model("gpt-5.5", provider="openai")
headers = mock_init.call_args[1]["default_headers"]
assert headers["originator"] == "codex_cli_rs"
assert headers["version"] == "0.144.1"
assert headers["User-Agent"].startswith("codex_cli_rs/0.144.1")
mock_installed_version.assert_called_once_with()
assert mock_init.call_args[1]["reasoning"]["effort"] == "xhigh"
@patch("EvoScientist.llm.models.init_chat_model")
def test_openai_ccproxy_codex_client_version_env(self, mock_init, monkeypatch):
"""EVOSCIENTIST_CODEX_CLIENT_VERSION overrides the pinned version."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENAI_BASE_URL", "http://127.0.0.1:8000/codex/v1")
monkeypatch.setenv("OPENAI_API_KEY", "ccproxy-oauth")
monkeypatch.setenv("EVOSCIENTIST_CODEX_CLIENT_VERSION", "9.9.9")
get_chat_model("gpt-5.5", provider="openai")
headers = mock_init.call_args[1]["default_headers"]
assert headers["version"] == "9.9.9"
assert headers["User-Agent"].startswith("codex_cli_rs/9.9.9")
@patch("EvoScientist.llm.models.subprocess.run")
def test_installed_codex_client_version(self, mock_run):
"""The advertised version follows the installed Codex CLI."""
from EvoScientist.llm.models import _installed_codex_client_version
mock_run.return_value.returncode = 0
mock_run.return_value.stdout = "codex-cli 0.144.1\n"
mock_run.return_value.stderr = ""
_installed_codex_client_version.cache_clear()
try:
assert _installed_codex_client_version() == "0.144.1"
assert _installed_codex_client_version() == "0.144.1"
finally:
_installed_codex_client_version.cache_clear()
mock_run.assert_called_once_with(
["codex", "--version"],
capture_output=True,
text=True,
timeout=2,
check=False,
)
@patch(
"EvoScientist.llm.models._installed_codex_client_version",
return_value="0.140.0",
)
def test_older_installed_codex_uses_fallback(
self, mock_installed_version, monkeypatch
):
"""An outdated installed CLI must not undercut the safe fallback."""
from EvoScientist.llm.models import (
_CODEX_CLIENT_VERSION_FALLBACK,
_resolve_codex_client_version,
)
monkeypatch.delenv("EVOSCIENTIST_CODEX_CLIENT_VERSION", raising=False)
assert _resolve_codex_client_version() == _CODEX_CLIENT_VERSION_FALLBACK
mock_installed_version.assert_called_once_with()
@patch("EvoScientist.llm.models.init_chat_model")
def test_openai_ccproxy_codex_headers_respect_caller(self, mock_init, monkeypatch):
"""Caller-supplied default_headers keys are not overridden."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENAI_BASE_URL", "http://127.0.0.1:8000/codex/v1")
monkeypatch.setenv("OPENAI_API_KEY", "ccproxy-oauth")
get_chat_model(
"gpt-5.5",
provider="openai",
default_headers={"originator": "codex_vscode", "version": "9.9.9"},
)
headers = mock_init.call_args[1]["default_headers"]
assert headers["originator"] == "codex_vscode"
assert headers["version"] == "9.9.9"
assert headers["User-Agent"].startswith("codex_cli_rs/9.9.9")
@patch("EvoScientist.llm.models.init_chat_model")
def test_openai_ccproxy_codex_none_headers(self, mock_init, monkeypatch):
"""An explicit default_headers=None is normalized before gap-filling."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENAI_BASE_URL", "http://127.0.0.1:8000/codex/v1")
monkeypatch.setenv("OPENAI_API_KEY", "ccproxy-oauth")
monkeypatch.setenv("EVOSCIENTIST_CODEX_CLIENT_VERSION", "9.9.9")
get_chat_model(
"gpt-5.5",
provider="openai",
default_headers=None,
)
headers = mock_init.call_args[1]["default_headers"]
assert headers["originator"] == "codex_cli_rs"
assert headers["version"] == "9.9.9"
@patch("EvoScientist.llm.models.init_chat_model")
def test_openai_ccproxy_key_but_wrong_path_not_ccproxy(
-129
View File
@@ -1,129 +0,0 @@
import logging
from datetime import UTC, datetime
from io import StringIO
from EvoScientist.logging_config import (
DailyLogFileHandler,
configure_console_logging,
configure_daily_file_logging,
configure_logging,
default_log_dir,
resolve_log_level,
)
def test_daily_log_file_handler_uses_dated_active_file(tmp_path):
handler = DailyLogFileHandler(tmp_path, prefix="evoscientist", retention_days=30)
logger = logging.getLogger("tests.daily_log_file_handler")
logger.handlers.clear()
logger.propagate = False
logger.setLevel(logging.INFO)
logger.addHandler(handler)
logger.info("hello")
handler.close()
today = datetime.now().strftime("%Y-%m-%d")
assert (tmp_path / f"evoscientist-{today}.log").read_text(encoding="utf-8").strip()
def test_daily_log_file_handler_keeps_latest_retention_days(tmp_path):
for day in range(1, 33):
(tmp_path / f"evoscientist-2026-01-{day:02d}.log").write_text(
"x", encoding="utf-8"
)
handler = DailyLogFileHandler(tmp_path, prefix="evoscientist", retention_days=30)
handler._delete_expired_logs()
handler.close()
remaining = sorted(path.name for path in tmp_path.glob("evoscientist-*.log"))
assert len(remaining) == 30
assert remaining[0] == "evoscientist-2026-01-03.log"
def test_configure_daily_file_logging_replaces_matching_handler(tmp_path):
logger = logging.getLogger("tests.configure_daily_file_logging")
logger.handlers.clear()
logger.propagate = False
first = configure_daily_file_logging(logger, log_dir=tmp_path)
second = configure_daily_file_logging(logger, log_dir=tmp_path)
try:
handlers = [h for h in logger.handlers if isinstance(h, DailyLogFileHandler)]
assert handlers == [second]
assert first.stream is None
finally:
for handler in logger.handlers[:]:
logger.removeHandler(handler)
handler.close()
def test_configure_logging_replaces_only_managed_handlers(tmp_path):
logger = logging.getLogger("tests.configure_logging")
logger.handlers.clear()
logger.propagate = False
external = logging.NullHandler()
logger.addHandler(external)
configure_logging(logger, log_dir=tmp_path, level="debug")
configure_logging(logger, log_dir=tmp_path, level="info")
try:
daily_handlers = [h for h in logger.handlers if isinstance(h, DailyLogFileHandler)]
stream_handlers = [
h
for h in logger.handlers
if isinstance(h, logging.StreamHandler)
and not isinstance(h, DailyLogFileHandler)
]
assert external in logger.handlers
assert len(daily_handlers) == 1
assert len(stream_handlers) == 1
assert logger.level == logging.INFO
finally:
for handler in logger.handlers[:]:
logger.removeHandler(handler)
handler.close()
def test_configure_console_logging_emits_to_stream():
logger = logging.getLogger("tests.configure_console_logging")
logger.handlers.clear()
logger.propagate = False
stream = StringIO()
configure_console_logging(logger, level="INFO", stream=stream)
try:
logger.info("hello")
assert "tests.configure_console_logging: hello" in stream.getvalue()
finally:
for handler in logger.handlers[:]:
logger.removeHandler(handler)
handler.close()
def test_resolve_log_level_accepts_alias_numeric_and_fallback():
assert resolve_log_level("warn") == logging.WARNING
assert resolve_log_level("10") == logging.DEBUG
assert resolve_log_level("", default=logging.ERROR) == logging.ERROR
assert resolve_log_level("not-a-level", default=logging.CRITICAL) == logging.CRITICAL
def test_daily_log_file_handler_supports_utc(tmp_path):
handler = DailyLogFileHandler(tmp_path, utc=True)
try:
today_utc = datetime.now(UTC).strftime("%Y-%m-%d")
assert handler.active_log_path.name == f"evoscientist-{today_utc}.log"
finally:
handler.close()
def test_default_log_dir_uses_current_data_dir(monkeypatch, tmp_path):
import EvoScientist.paths as paths
monkeypatch.delenv("EVOSCIENTIST_LOG_DIR", raising=False)
monkeypatch.setattr(paths, "DATA_DIR", tmp_path / "data")
assert default_log_dir() == tmp_path / "data" / "logs"
+18 -23
View File
@@ -1399,7 +1399,9 @@ class TestLoadToolsProgressCallback:
monkeypatch.setattr(lc_client, "MultiServerMCPClient", _FakeClient)
async def test_success_emits_start_then_success_with_tool_count(self, monkeypatch):
def test_success_emits_start_then_success_with_tool_count(self, monkeypatch):
import asyncio
from EvoScientist.mcp.client import _load_tools
events: list[tuple[str, str, str]] = []
@@ -1413,45 +1415,36 @@ class TestLoadToolsProgressCallback:
def record(event, name, detail):
events.append((event, name, detail))
await _load_tools(config, on_progress=record)
asyncio.run(_load_tools(config, on_progress=record))
assert events == [
("start", "srv", ""),
("success", "srv", "3"),
]
async def test_failure_emits_start_then_error_with_detail(self, monkeypatch):
from EvoScientist.mcp import client as mcp_client
def test_failure_emits_start_then_error_with_detail(self, monkeypatch):
import asyncio
from EvoScientist.mcp.client import _load_tools
events: list[tuple[str, str, str]] = []
self._patch_client(monkeypatch, {"srv": RuntimeError("boom")})
monkeypatch.setattr(mcp_client, "_MCP_SERVER_ERRORS", {})
config = {"srv": {"transport": "stdio", "command": "demo"}}
def record(event, name, detail):
events.append((event, name, detail))
await mcp_client._load_tools(config, on_progress=record)
asyncio.run(_load_tools(config, on_progress=record))
assert events == [
("start", "srv", ""),
("error", "srv", "boom"),
]
assert mcp_client.get_mcp_server_errors() == {"srv": "boom"}
async def test_success_clears_previous_server_error(self, monkeypatch):
from EvoScientist.mcp import client as mcp_client
def test_mixed_fleet_reports_each_server_independently(self, monkeypatch):
import asyncio
self._patch_client(monkeypatch, {"srv": ["tool1"]})
monkeypatch.setattr(mcp_client, "_MCP_SERVER_ERRORS", {"srv": "old error"})
config = {"srv": {"transport": "stdio", "command": "demo"}}
await mcp_client._load_tools(config)
assert mcp_client.get_mcp_server_errors() == {}
async def test_mixed_fleet_reports_each_server_independently(self, monkeypatch):
from EvoScientist.mcp.client import _load_tools
events: list[tuple[str, str, str]] = []
@@ -1471,7 +1464,7 @@ class TestLoadToolsProgressCallback:
def record(event, name, detail):
events.append((event, name, detail))
await _load_tools(config, on_progress=record)
asyncio.run(_load_tools(config, on_progress=record))
by_server = {}
for ev, name, detail in events:
@@ -1479,7 +1472,9 @@ class TestLoadToolsProgressCallback:
assert by_server["ok_srv"] == [("start", ""), ("success", "1")]
assert by_server["bad_srv"] == [("start", ""), ("error", "refused")]
async def test_callback_errors_do_not_break_the_load(self, monkeypatch):
def test_callback_errors_do_not_break_the_load(self, monkeypatch):
import asyncio
from EvoScientist.mcp.client import _load_tools
self._patch_client(monkeypatch, {"srv": ["tool1"]})
@@ -1489,10 +1484,10 @@ class TestLoadToolsProgressCallback:
def bad_callback(event, name, detail):
raise RuntimeError("callback bug")
result = await _load_tools(config, on_progress=bad_callback)
result = asyncio.run(_load_tools(config, on_progress=bad_callback))
assert result == {"srv": ["tool1"]}
async def test_semaphore_caps_concurrent_connections(self, monkeypatch):
def test_semaphore_caps_concurrent_connections(self, monkeypatch):
"""Many configured servers must not all spawn at once."""
import asyncio
@@ -1521,7 +1516,7 @@ class TestLoadToolsProgressCallback:
config = {
f"srv{i}": {"transport": "stdio", "command": "demo"} for i in range(10)
}
await mcp_client._load_tools(config)
asyncio.run(mcp_client._load_tools(config))
assert inflight["peak"] <= 3
assert inflight["peak"] > 1 # sanity: we *are* parallelizing
+18 -16
View File
@@ -2,6 +2,8 @@
from unittest.mock import MagicMock, patch
from tests.conftest import run_async as _run
def _ctx():
from EvoScientist.commands.base import CommandContext
@@ -12,16 +14,16 @@ def _ctx():
class TestMCPCommandDispatch:
async def test_no_args_lists(self):
def test_no_args_lists(self):
from EvoScientist.commands.implementation.mcp import MCPCommand
ctx, ui = _ctx()
with patch("EvoScientist.mcp.load_mcp_config", return_value={}):
await MCPCommand().execute(ctx, [])
_run(MCPCommand().execute(ctx, []))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("No MCP servers configured" in m for m in msgs)
async def test_list_subcommand(self):
def test_list_subcommand(self):
from EvoScientist.commands.implementation.mcp import MCPCommand
ctx, ui = _ctx()
@@ -29,10 +31,10 @@ class TestMCPCommandDispatch:
"srv1": {"transport": "stdio", "tools": ["foo"], "expose_to": ["main"]},
}
with patch("EvoScientist.mcp.load_mcp_config", return_value=cfg):
await MCPCommand().execute(ctx, ["list"])
_run(MCPCommand().execute(ctx, ["list"]))
ui.mount_renderable.assert_called_once()
async def test_add_subcommand_dispatches(self):
def test_add_subcommand_dispatches(self):
from EvoScientist.commands.implementation.mcp import MCPCommand
ctx, _ui = _ctx()
@@ -46,10 +48,10 @@ class TestMCPCommandDispatch:
return_value={"transport": "stdio"},
) as add_mock,
):
await MCPCommand().execute(ctx, ["add", "srv1", "python"])
_run(MCPCommand().execute(ctx, ["add", "srv1", "python"]))
add_mock.assert_called_once()
async def test_edit_subcommand_dispatches(self):
def test_edit_subcommand_dispatches(self):
from EvoScientist.commands.implementation.mcp import MCPCommand
ctx, _ui = _ctx()
@@ -62,28 +64,28 @@ class TestMCPCommandDispatch:
"EvoScientist.mcp.edit_mcp_server",
) as edit_mock,
):
await MCPCommand().execute(ctx, ["edit", "srv1", "--tools", "bar"])
_run(MCPCommand().execute(ctx, ["edit", "srv1", "--tools", "bar"]))
edit_mock.assert_called_once_with("srv1", tools=["bar"])
async def test_remove_subcommand_success(self):
def test_remove_subcommand_success(self):
from EvoScientist.commands.implementation.mcp import MCPCommand
ctx, ui = _ctx()
with patch("EvoScientist.mcp.remove_mcp_server", return_value=True):
await MCPCommand().execute(ctx, ["remove", "srv1"])
_run(MCPCommand().execute(ctx, ["remove", "srv1"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Removed MCP server: srv1" in m for m in msgs)
async def test_remove_subcommand_not_found(self):
def test_remove_subcommand_not_found(self):
from EvoScientist.commands.implementation.mcp import MCPCommand
ctx, ui = _ctx()
with patch("EvoScientist.mcp.remove_mcp_server", return_value=False):
await MCPCommand().execute(ctx, ["remove", "missing"])
_run(MCPCommand().execute(ctx, ["remove", "missing"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Server not found" in m for m in msgs)
async def test_install_delegates_to_install_mcp_command(self):
def test_install_delegates_to_install_mcp_command(self):
"""/mcp install should instantiate InstallMCPCommand and execute it."""
from EvoScientist.commands.implementation.mcp import MCPCommand
@@ -99,13 +101,13 @@ class TestMCPCommandDispatch:
instance.execute = fake_execute
klass.return_value = instance
await MCPCommand().execute(ctx, ["install", "foo"])
_run(MCPCommand().execute(ctx, ["install", "foo"]))
klass.assert_called_once()
async def test_unknown_subcommand_prints_help(self):
def test_unknown_subcommand_prints_help(self):
from EvoScientist.commands.implementation.mcp import MCPCommand
ctx, ui = _ctx()
await MCPCommand().execute(ctx, ["bogus"])
_run(MCPCommand().execute(ctx, ["bogus"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("MCP commands:" in m for m in msgs)
+34 -34
View File
@@ -5,6 +5,8 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from tests.conftest import run_async as _run
class TestExtractModelAndProvider:
"""Unit tests for the argument parser helper."""
@@ -78,7 +80,7 @@ class TestExtractModelAndProvider:
class TestModelCommandUnknownModel:
"""Verify error message for unknown models."""
async def test_unknown_model_shows_error(self):
def test_unknown_model_shows_error(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
@@ -93,7 +95,7 @@ class TestModelCommandUnknownModel:
"EvoScientist.EvoScientist._ensure_config",
return_value=cfg,
):
await cmd.execute(ctx, ["nonexistent-model-xyz"])
_run(cmd.execute(ctx, ["nonexistent-model-xyz"]))
ui.append_system.assert_called_once()
call_args = ui.append_system.call_args
@@ -104,7 +106,7 @@ class TestModelCommandUnknownModel:
class TestModelCommandPickerCancelled:
"""Verify no-op when the interactive picker is cancelled."""
async def test_picker_returns_none(self):
def test_picker_returns_none(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
@@ -120,7 +122,7 @@ class TestModelCommandPickerCancelled:
"EvoScientist.EvoScientist._ensure_config",
return_value=cfg,
):
await cmd.execute(ctx, [])
_run(cmd.execute(ctx, []))
# No model switch should have happened
ui.append_system.assert_not_called()
@@ -129,7 +131,7 @@ class TestModelCommandPickerCancelled:
class TestModelCommandSwitch:
"""Verify a successful model switch updates config and rebuilds agent."""
async def test_switch_known_model(self):
def test_switch_known_model(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
@@ -156,7 +158,7 @@ class TestModelCommandSwitch:
return_value=new_agent,
),
):
await cmd.execute(ctx, ["claude-opus-4-8"])
_run(cmd.execute(ctx, ["claude-opus-4-8"]))
# The switch is committed via set_active_config(temp_cfg), not by
# mutating the original cfg object in place.
@@ -174,7 +176,7 @@ class TestModelCommandSwitch:
assert "claude-opus-4-8" in msg
assert "anthropic" in msg
async def test_switch_with_save_flag(self):
def test_switch_with_save_flag(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
@@ -201,7 +203,7 @@ class TestModelCommandSwitch:
),
patch("EvoScientist.config.settings.set_config_value") as mock_save,
):
await cmd.execute(ctx, ["claude-opus-4-8", "--save"])
_run(cmd.execute(ctx, ["claude-opus-4-8", "--save"]))
# Config file should be updated
mock_save.assert_any_call("model", "claude-opus-4-8")
@@ -211,7 +213,7 @@ class TestModelCommandSwitch:
msg = ui.append_system.call_args[0][0]
assert "saved to config" in msg
async def test_switch_without_save_flag_does_not_persist(self):
def test_switch_without_save_flag_does_not_persist(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
@@ -238,7 +240,7 @@ class TestModelCommandSwitch:
),
patch("EvoScientist.config.settings.set_config_value") as mock_save,
):
await cmd.execute(ctx, ["claude-opus-4-8"])
_run(cmd.execute(ctx, ["claude-opus-4-8"]))
# Config file should NOT be updated
mock_save.assert_not_called()
@@ -251,7 +253,7 @@ class TestModelCommandSwitch:
class TestModelCommandFailure:
"""Verify error handling when chat-model construction raises."""
async def test_build_chat_model_error(self):
def test_build_chat_model_error(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
@@ -274,7 +276,7 @@ class TestModelCommandFailure:
side_effect=RuntimeError("API key missing"),
) as mock_build,
):
await cmd.execute(ctx, ["claude-opus-4-8"])
_run(cmd.execute(ctx, ["claude-opus-4-8"]))
mock_build.assert_called_once()
ui.append_system.assert_called_once()
@@ -444,7 +446,7 @@ class TestApplyModelIntegration:
pair so we can assert on identity.
"""
async def test_new_agent_is_bound_to_newly_selected_model(self, evo_module_state):
def test_new_agent_is_bound_to_newly_selected_model(self, evo_module_state):
from EvoScientist.commands.implementation.model import ModelCommand
from EvoScientist.config.settings import EvoScientistConfig
@@ -500,7 +502,7 @@ class TestApplyModelIntegration:
),
):
cmd = ModelCommand()
await cmd._apply_model(ctx, "minimax-m2.7", "openrouter")
_run(cmd._apply_model(ctx, "minimax-m2.7", "openrouter"))
# The agent produced by _apply_model must be bound to the
# NEWLY requested model, threaded in via chat_model=.
@@ -537,9 +539,7 @@ class TestApplyModelPreservesConfigByReference:
switch (the held object stops being the active ``_config`` after the first).
"""
async def test_held_config_reference_tracks_repeated_switches(
self, evo_module_state
):
def test_held_config_reference_tracks_repeated_switches(self, evo_module_state):
from EvoScientist.commands.implementation.model import ModelCommand
from EvoScientist.config.settings import EvoScientistConfig
@@ -592,7 +592,7 @@ class TestApplyModelPreservesConfigByReference:
("minimax-m2.7", "openrouter"),
("claude-sonnet-4-6", "anthropic"),
]:
await cmd._apply_model(ctx, model, provider)
_run(cmd._apply_model(ctx, model, provider))
# The held reference must reflect the LATEST switch on every
# iteration — not just the first — and stay the active config.
assert agent_holder["config"].model == model
@@ -610,7 +610,7 @@ class TestModelCommandLoadAgentFailure:
the ordering could silently regress (e.g. if ``_apply_model`` were
reordered to call ``set_chat_model`` first)."""
async def test_load_agent_error_is_transactional(self):
def test_load_agent_error_is_transactional(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
@@ -646,7 +646,7 @@ class TestModelCommandLoadAgentFailure:
# Pass ``--save`` to strengthen the assertion: if the ordering
# ever regresses, ``set_config_value`` would be called with
# stale data.
await cmd.execute(ctx, ["claude-opus-4-8", "--save"])
_run(cmd.execute(ctx, ["claude-opus-4-8", "--save"]))
# _load_agent was attempted (transactional first step).
mock_load.assert_called_once()
@@ -677,7 +677,7 @@ class TestApplyModelLoadAgentFailureTransactional:
downstream setters never run on failure.
"""
async def test_globals_unchanged_when_load_agent_raises(self, evo_module_state):
def test_globals_unchanged_when_load_agent_raises(self, evo_module_state):
from EvoScientist.commands.implementation.model import ModelCommand
from EvoScientist.config.settings import EvoScientistConfig
@@ -721,7 +721,7 @@ class TestApplyModelLoadAgentFailureTransactional:
),
):
cmd = ModelCommand()
await cmd._apply_model(ctx, "minimax-m2.7", "openrouter")
_run(cmd._apply_model(ctx, "minimax-m2.7", "openrouter"))
# All four globals are unchanged — nothing was committed.
assert mod._config is cfg
@@ -753,7 +753,7 @@ class TestModelCommandOllamaPicker:
ctx.ui = ui
return ctx, cfg, ui
async def test_picker_entries_include_detected_ollama_models(self):
def test_picker_entries_include_detected_ollama_models(self):
"""When Ollama is reachable, detected models appear in entries with
provider='ollama' and the Custom sentinel is appended."""
from EvoScientist.commands.implementation.model import ModelCommand
@@ -773,7 +773,7 @@ class TestModelCommandOllamaPicker:
side_effect=fake_discover,
),
):
await ModelCommand().execute(ctx, [])
_run(ModelCommand().execute(ctx, []))
entries = ui.wait_for_model_pick.call_args[0][0]
ollama_rows = [(n, mid, p) for (n, mid, p) in entries if p == "ollama"]
@@ -785,7 +785,7 @@ class TestModelCommandOllamaPicker:
"ollama",
) in ollama_rows
async def test_picker_entries_include_sentinel_when_discovery_empty(self):
def test_picker_entries_include_sentinel_when_discovery_empty(self):
"""Daemon unreachable / no models pulled — sentinel is the user's
escape hatch and must always be present."""
from EvoScientist.commands.implementation.model import ModelCommand
@@ -805,7 +805,7 @@ class TestModelCommandOllamaPicker:
side_effect=fake_discover,
),
):
await ModelCommand().execute(ctx, [])
_run(ModelCommand().execute(ctx, []))
entries = ui.wait_for_model_pick.call_args[0][0]
ollama_rows = [(n, mid, p) for (n, mid, p) in entries if p == "ollama"]
@@ -813,7 +813,7 @@ class TestModelCommandOllamaPicker:
("Custom Ollama model...", "__custom_ollama__", "ollama")
]
async def test_picker_skips_ollama_section_when_not_configured(self):
def test_picker_skips_ollama_section_when_not_configured(self):
"""ollama_base_url unset → no discovery call, no ollama entries,
no sentinel (issue non-goal: no implicit localhost detection)."""
from EvoScientist.commands.implementation.model import ModelCommand
@@ -832,13 +832,13 @@ class TestModelCommandOllamaPicker:
discovery,
),
):
await ModelCommand().execute(ctx, [])
_run(ModelCommand().execute(ctx, []))
discovery.assert_not_called()
entries = ui.wait_for_model_pick.call_args[0][0]
assert not any(p == "ollama" for (_, _, p) in entries)
async def test_picker_handles_cfg_without_ollama_base_url_attr(self):
def test_picker_handles_cfg_without_ollama_base_url_attr(self):
"""getattr(cfg, 'ollama_base_url', None) fallback: old configs
(or SimpleNamespace test fixtures) may not carry the attribute
at all. Must not raise AttributeError, must not probe."""
@@ -864,13 +864,13 @@ class TestModelCommandOllamaPicker:
discovery,
),
):
await ModelCommand().execute(ctx, [])
_run(ModelCommand().execute(ctx, []))
discovery.assert_not_called()
entries = ui.wait_for_model_pick.call_args[0][0]
assert not any(p == "ollama" for (_, _, p) in entries)
async def test_picker_sentinel_result_is_treated_as_cancel(self):
def test_picker_sentinel_result_is_treated_as_cancel(self):
"""Defense-in-depth: if the widget ever returns the sentinel name
itself (shouldn't happen — it should substitute the typed name),
dispatch treats it as a cancel and does NOT call _apply_model."""
@@ -893,12 +893,12 @@ class TestModelCommandOllamaPicker:
),
patch("EvoScientist.cli.agent._load_agent") as load_agent,
):
await ModelCommand().execute(ctx, [])
_run(ModelCommand().execute(ctx, []))
load_agent.assert_not_called()
assert cfg.model == "claude-sonnet-4-6" # unchanged
async def test_picker_applies_detected_ollama_model(self):
def test_picker_applies_detected_ollama_model(self):
"""User picks a live-detected Ollama model → _apply_model is invoked
with (name, "ollama") and the agent is rebuilt."""
from EvoScientist.commands.implementation.model import ModelCommand
@@ -928,7 +928,7 @@ class TestModelCommandOllamaPicker:
return_value=MagicMock(),
),
):
await ModelCommand().execute(ctx, [])
_run(ModelCommand().execute(ctx, []))
# Committed via set_active_config(temp_cfg); original cfg untouched.
set_cfg.assert_called_once()
+27 -135
View File
@@ -6,7 +6,6 @@ fallback chain behaviour via _try_fallbacks / _guard_and_fallback.
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@@ -21,6 +20,7 @@ from EvoScientist.middleware.model_fallback import (
clear_fallbacks,
set_ui_emit,
)
from tests.conftest import run_async as _run
# ── Helpers ──────────────────────────────────────────────────────
@@ -84,7 +84,6 @@ class TestIsNonFallbackable:
"Error 400: invalid_request_error",
"400 Bad Request: invalid request body",
"400: malformed JSON in request",
"<400> InvalidParameter: Repetitive tool calls detected in history",
],
)
def test_malformed_request_400_patterns(self, msg):
@@ -147,7 +146,7 @@ class TestIsNonFallbackable:
class TestTryFallbacks:
"""End-to-end tests for the fallback chain traversal."""
async def test_first_fallback_succeeds(self):
def test_first_fallback_succeeds(self):
"""When the first fallback model works, return its response."""
add_fallback("fb-model", "fb-provider")
req = _fake_request()
@@ -155,13 +154,13 @@ class TestTryFallbacks:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = await _try_fallbacks(req, invoke, Exception("503 boom"))
result = _run(_try_fallbacks(req, invoke, Exception("503 boom")))
assert result is AI_RESPONSE
invoke.assert_awaited_once()
mock_gcm.assert_called_once_with(model="fb-model", provider="fb-provider")
async def test_skips_failing_fallback_tries_next(self):
def test_skips_failing_fallback_tries_next(self):
"""When the first fallback fails, try the second."""
add_fallback("fb-bad", "prov-a")
add_fallback("fb-good", "prov-b")
@@ -178,12 +177,12 @@ class TestTryFallbacks:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = await _try_fallbacks(req, _invoke, Exception("503 boom"))
result = _run(_try_fallbacks(req, _invoke, Exception("503 boom")))
assert result is AI_RESPONSE
assert call_count == 2
async def test_all_fallbacks_exhausted_raises_last(self):
def test_all_fallbacks_exhausted_raises_last(self):
"""When every fallback fails, re-raise the last exception."""
add_fallback("fb-a", "prov-a")
add_fallback("fb-b", "prov-b")
@@ -203,11 +202,11 @@ class TestTryFallbacks:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
with pytest.raises(Exception, match="429 from fb-b") as exc_info:
await _try_fallbacks(req, _invoke, Exception("503 primary"))
_run(_try_fallbacks(req, _invoke, Exception("503 primary")))
assert exc_info.value is last_error
async def test_non_fallbackable_in_chain_aborts_immediately(self):
def test_non_fallbackable_in_chain_aborts_immediately(self):
"""A non-fallbackable error from a fallback model aborts the chain."""
add_fallback("fb-a", "prov-a")
add_fallback("fb-b", "prov-b") # should never be reached
@@ -219,95 +218,12 @@ class TestTryFallbacks:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
with pytest.raises(Exception, match="context_length_exceeded"):
await _try_fallbacks(req, _invoke, Exception("503 primary"))
_run(_try_fallbacks(req, _invoke, Exception("503 primary")))
# get_chat_model should only have been called once (for fb-a),
# fb-b should never be reached.
assert mock_gcm.call_count == 1
async def test_exhausted_fallbacks_attribute_to_last_failing_model(self):
"""Regression: when every fallback fails, the raised
``ProviderStreamError`` must be attributed to the model that
ACTUALLY failed last, not the original ``request.model``.
Prevents a ``deepseek → moonshot`` chain from surfacing as
``provider: deepseek`` after moonshot exhausts its quota.
"""
from EvoScientist.llm.errors import ProviderStreamError
add_fallback("moonshot-model", "moonshot")
# Original request's model is openai-shape. Fallback's model
# will be openai-shape with a moonshot base_url.
req = _fake_request()
# ChatOpenAI-shape model instance so ``_provider_from_model``
# returns a recognized provider.
def _make_openai_model(base_url=None):
cls = type(
"ChatOpenAI",
(),
{"__module__": "langchain_openai.chat_models.base"},
)
inst = cls()
inst.openai_api_base = base_url
return inst
req.model = _make_openai_model() # primary
fallback_model = _make_openai_model(base_url="https://api.moonshot.cn/v1")
# ``request.override(model=...)`` must return the request with the
# new model so ``_try_fallbacks`` tracks the failing model.
req.override = MagicMock(
side_effect=lambda **kw: SimpleNamespace(model=kw.get("model", req.model))
)
async def _invoke(_r):
raise Exception("429 quota exceeded")
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = fallback_model
with pytest.raises(ProviderStreamError) as exc_info:
await _try_fallbacks(req, _invoke, Exception("openai primary failed"))
# Attribution flipped to moonshot (the failing fallback), not
# openai (the original request's model).
assert exc_info.value.provider == "moonshot"
assert "quota exceeded" in exc_info.value.message
async def test_langgraph_error_at_fallback_raise_point_passes_through(self):
"""Regression: ``_raise_normalized`` calls ``_normalize``
directly, so its ``_should_pass_through`` gate must fire even
without the ``ErrorNormalizationMiddleware`` wrap sites' own
check. Prevents a ``langgraph.errors.*`` exception hitting the
fallback chain from being wrapped as a provider incident.
"""
from langgraph.errors import InvalidUpdateError
add_fallback("fb-a", "prov-a")
req = _fake_request()
# Use a recognized-provider model so ``_provider_from_model``
# wouldn't short-circuit — the guard has to come from
# ``_should_pass_through``, not the provider check.
cls = type(
"ChatOpenAI", (), {"__module__": "langchain_openai.chat_models.base"}
)
model = cls()
model.openai_api_base = None
req.model = model
req.override = MagicMock(
side_effect=lambda **kw: SimpleNamespace(model=kw.get("model", req.model))
)
raised = InvalidUpdateError("state mismatch")
async def _invoke(_r):
raise raised
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = model
with pytest.raises(InvalidUpdateError) as exc_info:
await _try_fallbacks(req, _invoke, Exception("primary failed"))
assert exc_info.value is raised
# ═════════════════════════════════════════════════════════════════
# 3. _guard_and_fallback — pre-check before chain walk
@@ -317,69 +233,43 @@ class TestTryFallbacks:
class TestGuardAndFallback:
"""Verify that non-fallbackable errors are re-raised before trying the chain."""
async def test_context_overflow_raises_immediately(self):
def test_context_overflow_raises_immediately(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock()
with pytest.raises(ContextOverflowError):
await _guard_and_fallback(ContextOverflowError("overflow"), req, invoke)
_run(_guard_and_fallback(ContextOverflowError("overflow"), req, invoke))
invoke.assert_not_awaited()
async def test_context_overflow_with_provider_model_passes_through_unwrapped(self):
"""Regression: a ``ContextOverflowError`` entering
``_guard_and_fallback`` under a recognized-provider model must
come out unwrapped. Otherwise ``_raise_normalized`` →
``_normalize`` would wrap it as a ``ProviderStreamError`` and
deepagents' ``SummarizationMiddleware`` (which sits outside
the user middleware stack and catches by exact type) would
stop compressing history and retrying.
"""
add_fallback("fb", "prov")
req = _fake_request()
# Recognized provider — without the gate in ``_normalize`` this
# would wrap. With the gate, the raw type propagates.
cls = type(
"ChatOpenAI", (), {"__module__": "langchain_openai.chat_models.base"}
)
model = cls()
model.openai_api_base = None
req.model = model
invoke = AsyncMock()
raised = ContextOverflowError("context length exceeded")
with pytest.raises(ContextOverflowError) as exc_info:
await _guard_and_fallback(raised, req, invoke)
assert exc_info.value is raised
invoke.assert_not_awaited()
async def test_malformed_400_raises_immediately(self):
def test_malformed_400_raises_immediately(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock()
with pytest.raises(Exception, match="invalid_request_error"):
await _guard_and_fallback(
Exception("400: invalid_request_error"), req, invoke
_run(
_guard_and_fallback(
Exception("400: invalid_request_error"), req, invoke
)
)
invoke.assert_not_awaited()
async def test_server_error_proceeds_to_fallback(self):
def test_server_error_proceeds_to_fallback(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock(return_value=AI_RESPONSE)
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = await _guard_and_fallback(Exception("503 overloaded"), req, invoke)
result = _run(_guard_and_fallback(Exception("503 overloaded"), req, invoke))
assert result is AI_RESPONSE
invoke.assert_awaited_once()
async def test_auth_error_proceeds_to_fallback(self):
def test_auth_error_proceeds_to_fallback(self):
"""Auth errors should try the fallback chain (different provider)."""
add_fallback("fb", "other-prov")
req = _fake_request()
@@ -387,8 +277,10 @@ class TestGuardAndFallback:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = await _guard_and_fallback(
Exception("400 Bad Request: invalid_api_key"), req, invoke
result = _run(
_guard_and_fallback(
Exception("400 Bad Request: invalid_api_key"), req, invoke
)
)
assert result is AI_RESPONSE
@@ -403,7 +295,7 @@ class TestGuardAndFallback:
class TestUiEmit:
"""Verify that fallback events are surfaced via the registered callback."""
async def test_emit_captures_messages(self):
def test_emit_captures_messages(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock(return_value=AI_RESPONSE)
@@ -413,14 +305,14 @@ class TestUiEmit:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
await _try_fallbacks(req, invoke, Exception("503 down"))
_run(_try_fallbacks(req, invoke, Exception("503 down")))
texts = [t for t, _ in messages]
assert any("Primary model failed" in t for t in texts)
assert any("Falling back to fb (prov)" in t for t in texts)
assert any("succeeded" in t for t in texts)
async def test_emit_shows_non_fallbackable_rejection(self):
def test_emit_shows_non_fallbackable_rejection(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock()
@@ -429,7 +321,7 @@ class TestUiEmit:
set_ui_emit(lambda text, style: messages.append((text, style)))
with pytest.raises(ContextOverflowError):
await _guard_and_fallback(ContextOverflowError("overflow"), req, invoke)
_run(_guard_and_fallback(ContextOverflowError("overflow"), req, invoke))
texts = [t for t, _ in messages]
assert any("not eligible for fallback" in t for t in texts)
+15 -10
View File
@@ -16,6 +16,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from EvoScientist.llm import patches as patches_mod
from tests.conftest import run_async as _run
# =============================================================================
# Helpers
@@ -151,7 +152,7 @@ class TestStartAsyncTaskInjection:
"configurable": {"model": "gpt-5", "model_provider": "openai"}
}
async def test_async_start_injects_config(self, restore_model_passthrough_patch):
def test_async_start_injects_config(self, restore_model_passthrough_patch):
try:
from deepagents.middleware import async_subagents as ds_mod
except ImportError:
@@ -175,10 +176,12 @@ class TestStartAsyncTaskInjection:
"EvoScientist.EvoScientist._ensure_config",
return_value=_stub_cfg(model="claude-haiku-4-5", provider="anthropic"),
):
await tool.coroutine(
description="hi",
subagent_type="writing-agent",
runtime=_runtime_stub(),
_run(
tool.coroutine(
description="hi",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
)
runs_async.create.assert_awaited_once()
@@ -264,7 +267,7 @@ class TestUpdateAsyncTaskInjection:
"last_updated_at": "2026-05-07T00:00:00Z",
}
async def test_async_update_injects_config(self, restore_model_passthrough_patch):
def test_async_update_injects_config(self, restore_model_passthrough_patch):
"""The async coroutine path must inject config too."""
try:
from deepagents.middleware import async_subagents as ds_mod
@@ -293,10 +296,12 @@ class TestUpdateAsyncTaskInjection:
"EvoScientist.EvoScientist._ensure_config",
return_value=_stub_cfg(model="gpt-5", provider="openai"),
):
await tool.coroutine(
task_id="thread-001",
message="follow up async",
runtime=runtime,
_run(
tool.coroutine(
task_id="thread-001",
message="follow up async",
runtime=runtime,
)
)
runs_async.create.assert_awaited_once()
+6 -4
View File
@@ -2,9 +2,11 @@
from unittest.mock import AsyncMock, MagicMock
from tests.conftest import run_async as _run
class TestNewCommand:
async def test_execute_calls_start_new_session(self):
def test_execute_calls_start_new_session(self):
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.session import NewCommand
@@ -16,7 +18,7 @@ class TestNewCommand:
ui=ui,
workspace_dir="/old/ws",
)
await NewCommand().execute(ctx, [])
_run(NewCommand().execute(ctx, []))
ui.start_new_session.assert_awaited_once()
def test_requires_agent_false(self):
@@ -24,7 +26,7 @@ class TestNewCommand:
assert NewCommand().requires_agent is False
async def test_no_agent_access(self):
def test_no_agent_access(self):
"""Command body must not touch ctx.agent (it's still loading)."""
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.session import NewCommand
@@ -33,4 +35,4 @@ class TestNewCommand:
ui.start_new_session = AsyncMock()
ctx = CommandContext(agent=None, thread_id="tid", ui=ui)
# No AttributeError even though ctx.agent is None
await NewCommand().execute(ctx, [])
_run(NewCommand().execute(ctx, []))
+65 -54
View File
@@ -1655,7 +1655,9 @@ def test_turn_compaction_uses_latest_user_turn_only():
]
async def test_lifecycle_schedules_turn_worker_without_awaiting(tmp_path, monkeypatch):
def test_lifecycle_schedules_turn_worker_without_awaiting(
tmp_path, monkeypatch, run_async
):
memory_dir = tmp_path / "memories"
workspace_dir = tmp_path / "workspace"
calls = []
@@ -1680,18 +1682,21 @@ async def test_lifecycle_schedules_turn_worker_without_awaiting(tmp_path, monkey
)
runtime = _runtime("thread-1")
state: AgentState[object] = {
"messages": [
HumanMessage("previous turn"),
AIMessage("previous answer"),
HumanMessage("hi"),
AIMessage("done"),
]
}
await middleware.aafter_agent(
state,
runtime,
)
async def run():
state: AgentState[object] = {
"messages": [
HumanMessage("previous turn"),
AIMessage("previous answer"),
HumanMessage("hi"),
AIMessage("done"),
]
}
await middleware.aafter_agent(
state,
runtime,
)
run_async(run())
assert len(calls) == 1
request, hooks = calls[0]
@@ -2177,9 +2182,10 @@ def test_observation_linker_does_not_launch_when_observations_disabled(
launch_call.assert_not_called()
async def test_async_observation_linker_does_not_launch_when_observations_disabled(
def test_async_observation_linker_does_not_launch_when_observations_disabled(
tmp_path,
monkeypatch,
run_async,
):
context = _linker_context(
memory_dir=tmp_path / "memories",
@@ -2194,7 +2200,7 @@ async def test_async_observation_linker_does_not_launch_when_observations_disabl
launch_call = MagicMock()
monkeypatch.setattr(memory_launch, "alaunch_background_run", launch_call)
run = await memory_launch.alaunch_observation_linker(context)
run = run_async(memory_launch.alaunch_observation_linker(context))
assert run is None
launch_call.assert_not_called()
@@ -2294,11 +2300,7 @@ def test_memory_worker_observation_writer_modes(
observation_writer=observation_writer,
)
# ErrorNormalizationMiddleware wraps outermost so provider-SDK
# exceptions from the auxiliary model call get normalized before
# the tool-error handler sees them.
assert type(middleware[0]).__name__ == "ErrorNormalizationMiddleware"
assert type(middleware[1]).__name__ == "ToolErrorHandlerMiddleware"
assert type(middleware[0]).__name__ == "ToolErrorHandlerMiddleware"
assert _memory_tool_names(middleware) == expected_tools
@@ -2333,8 +2335,8 @@ def test_sync_memory_worker_watcher_untracks_without_counting_on_poll_abort(
assert status.observations_recorded == 0
async def test_async_memory_worker_watcher_untracks_without_counting_on_poll_abort(
tmp_path, monkeypatch
def test_async_memory_worker_watcher_untracks_without_counting_on_poll_abort(
tmp_path, monkeypatch, run_async
):
memory_dir = tmp_path / "memories"
_mark_worker_started(memory_dir)
@@ -2346,12 +2348,14 @@ async def test_async_memory_worker_watcher_untracks_without_counting_on_poll_abo
async def get(self, **_kwargs):
raise RuntimeError("poll failed")
await background_runs.awatch_background_run(
SimpleNamespace(runs=_Runs()),
thread_id="worker-thread",
run_id="run-1",
hooks=memory_launch._memory_worker_launch_hooks(memory_dir),
watcher_config=_fast_watcher_config(max_poll_failures=1),
run_async(
background_runs.awatch_background_run(
SimpleNamespace(runs=_Runs()),
thread_id="worker-thread",
run_id="run-1",
hooks=memory_launch._memory_worker_launch_hooks(memory_dir),
watcher_config=_fast_watcher_config(max_poll_failures=1),
)
)
status = worker_activity.memory_worker_status()
assert status.is_running is False
@@ -2359,8 +2363,8 @@ async def test_async_memory_worker_watcher_untracks_without_counting_on_poll_abo
assert status.observations_recorded == 0
async def test_async_memory_worker_watcher_counts_completion_under_blockbuster(
tmp_path,
def test_async_memory_worker_watcher_counts_completion_under_blockbuster(
tmp_path, run_async
):
memory_dir = tmp_path / "memories"
_mark_worker_started(memory_dir)
@@ -2372,17 +2376,20 @@ async def test_async_memory_worker_watcher_counts_completion_under_blockbuster(
async def get(self, **_kwargs):
return {"status": "success"}
blocker = BlockBuster(scanned_modules=[memory_worker, worker_activity])
blocker.activate()
try:
await background_runs.awatch_background_run(
SimpleNamespace(runs=_Runs()),
thread_id="worker-thread",
run_id="run-1",
hooks=memory_launch._memory_worker_launch_hooks(memory_dir),
)
finally:
blocker.deactivate()
async def run():
blocker = BlockBuster(scanned_modules=[memory_worker, worker_activity])
blocker.activate()
try:
await background_runs.awatch_background_run(
SimpleNamespace(runs=_Runs()),
thread_id="worker-thread",
run_id="run-1",
hooks=memory_launch._memory_worker_launch_hooks(memory_dir),
)
finally:
blocker.deactivate()
run_async(run())
status = worker_activity.memory_worker_status()
assert status.is_running is False
assert status.profile_updates == 1
@@ -2520,7 +2527,7 @@ def test_memory_worker_marks_active_status(tmp_path, monkeypatch):
assert status.observations_recorded == 1
async def test_async_memory_worker_offloads_blocking_work(tmp_path, monkeypatch):
def test_async_memory_worker_offloads_blocking_work(tmp_path, monkeypatch, run_async):
monkeypatch.setattr(
background_runs, "default_background_run_url", lambda: "http://x"
)
@@ -2562,18 +2569,22 @@ async def test_async_memory_worker_offloads_blocking_work(tmp_path, monkeypatch)
spawned: list[background_runs.BackgroundRun] = []
event_loop_thread = threading.get_ident()
context = _memory_source_context(
memory_dir=tmp_path / "memories",
workspace_dir=tmp_path / "workspace",
trajectory=[{"role": "human", "content": "hi"}],
)
request = memory_launch.memory_worker_launch_request(context)
await background_runs.alaunch_background_run(
request,
hooks=memory_launch._memory_worker_launch_hooks(tmp_path / "memories"),
spawn_status_watcher=spawned.append,
)
async def run():
event_loop_thread = threading.get_ident()
context = _memory_source_context(
memory_dir=tmp_path / "memories",
workspace_dir=tmp_path / "workspace",
trajectory=[{"role": "human", "content": "hi"}],
)
request = memory_launch.memory_worker_launch_request(context)
await background_runs.alaunch_background_run(
request,
hooks=memory_launch._memory_worker_launch_hooks(tmp_path / "memories"),
spawn_status_watcher=spawned.append,
)
return event_loop_thread
event_loop_thread = run_async(run())
assert [name for name, _thread_id in call_threads] == ["health", "snapshot"]
assert all(thread_id != event_loop_thread for _name, thread_id in call_threads)
assert worker_activity.memory_worker_status().is_running is True
+21 -20
View File
@@ -17,6 +17,7 @@ from EvoScientist.llm.ollama_discovery import (
discover_ollama_models,
validate_ollama_connection,
)
from tests.conftest import run_async as _run
class TestValidateOllamaConnection:
@@ -70,17 +71,17 @@ class TestValidateOllamaConnection:
class TestDiscoverOllamaModels:
"""Async probe — contract: never raise, return list[str]."""
async def test_empty_base_url_returns_empty_without_http(self):
def test_empty_base_url_returns_empty_without_http(self):
# No HTTP call should be made for an empty base_url — verified by
# the fact that no mock is set up and the test completes.
names = await discover_ollama_models("")
names = _run(discover_ollama_models(""))
assert names == []
async def test_none_base_url_returns_empty(self):
names = await discover_ollama_models(None)
def test_none_base_url_returns_empty(self):
names = _run(discover_ollama_models(None))
assert names == []
async def test_200_returns_names(self):
def test_200_returns_names(self):
async def fake_get(self, url):
resp = MagicMock()
resp.status_code = 200
@@ -92,10 +93,10 @@ class TestDiscoverOllamaModels:
return resp
with patch.object(httpx.AsyncClient, "get", fake_get):
names = await discover_ollama_models("http://localhost:11434")
names = _run(discover_ollama_models("http://localhost:11434"))
assert names == ["llama3.3:latest", "qwen3:8b"]
async def test_strips_entries_without_name(self):
def test_strips_entries_without_name(self):
async def fake_get(self, url):
resp = MagicMock()
resp.status_code = 200
@@ -111,36 +112,36 @@ class TestDiscoverOllamaModels:
return resp
with patch.object(httpx.AsyncClient, "get", fake_get):
names = await discover_ollama_models("http://localhost:11434")
names = _run(discover_ollama_models("http://localhost:11434"))
assert names == ["llama3.3"]
async def test_timeout_returns_empty(self):
def test_timeout_returns_empty(self):
async def fake_get(self, url):
raise httpx.TimeoutException("timed out")
with patch.object(httpx.AsyncClient, "get", fake_get):
names = await discover_ollama_models("http://localhost:11434")
names = _run(discover_ollama_models("http://localhost:11434"))
assert names == []
async def test_connect_error_returns_empty(self):
def test_connect_error_returns_empty(self):
async def fake_get(self, url):
raise httpx.ConnectError("refused")
with patch.object(httpx.AsyncClient, "get", fake_get):
names = await discover_ollama_models("http://localhost:11434")
names = _run(discover_ollama_models("http://localhost:11434"))
assert names == []
async def test_non_200_returns_empty(self):
def test_non_200_returns_empty(self):
async def fake_get(self, url):
resp = MagicMock()
resp.status_code = 500
return resp
with patch.object(httpx.AsyncClient, "get", fake_get):
names = await discover_ollama_models("http://localhost:11434")
names = _run(discover_ollama_models("http://localhost:11434"))
assert names == []
async def test_malformed_json_returns_empty(self):
def test_malformed_json_returns_empty(self):
async def fake_get(self, url):
resp = MagicMock()
resp.status_code = 200
@@ -148,10 +149,10 @@ class TestDiscoverOllamaModels:
return resp
with patch.object(httpx.AsyncClient, "get", fake_get):
names = await discover_ollama_models("http://localhost:11434")
names = _run(discover_ollama_models("http://localhost:11434"))
assert names == []
async def test_missing_models_key_returns_empty(self):
def test_missing_models_key_returns_empty(self):
async def fake_get(self, url):
resp = MagicMock()
resp.status_code = 200
@@ -159,10 +160,10 @@ class TestDiscoverOllamaModels:
return resp
with patch.object(httpx.AsyncClient, "get", fake_get):
names = await discover_ollama_models("http://localhost:11434")
names = _run(discover_ollama_models("http://localhost:11434"))
assert names == []
async def test_trailing_slash_stripped_from_url(self):
def test_trailing_slash_stripped_from_url(self):
called = {}
async def fake_get(self, url):
@@ -173,7 +174,7 @@ class TestDiscoverOllamaModels:
return resp
with patch.object(httpx.AsyncClient, "get", fake_get):
await discover_ollama_models("http://localhost:11434/")
_run(discover_ollama_models("http://localhost:11434/"))
assert called["url"] == "http://localhost:11434/api/tags"
-309
View File
@@ -184,38 +184,6 @@ class TestSharedConstantsAlignment:
)
class TestOAuthModeReconcile:
def test_reconcile_preserves_auxiliary_openai_oauth(self):
from EvoScientist.config.onboard.wizard import _reconcile_oauth_modes
config = EvoScientistConfig(
provider="minimax",
auxiliary_provider="openai",
auxiliary_model="gpt-5.5",
openai_auth_mode="oauth",
anthropic_auth_mode="oauth",
)
_reconcile_oauth_modes(config)
assert config.openai_auth_mode == "oauth"
assert config.anthropic_auth_mode == "api_key"
def test_reconcile_preserves_auxiliary_provider_without_model(self):
from EvoScientist.config.onboard.wizard import _reconcile_oauth_modes
config = EvoScientistConfig(
provider="minimax",
auxiliary_provider="openai",
auxiliary_model="",
openai_auth_mode="oauth",
)
_reconcile_oauth_modes(config)
assert config.openai_auth_mode == "oauth"
# =============================================================================
# Test render_progress
# =============================================================================
@@ -418,109 +386,6 @@ class TestStepProvider:
_step_provider(config)
class TestStepOAuthAuthMode:
@pytest.mark.parametrize(
(
"step_name",
"config_attr",
"provider_label",
"oauth_choice_label",
"ccproxy_provider",
"status_label",
"question_label",
"login_prompt",
),
[
(
"_step_anthropic_auth_mode",
"anthropic_auth_mode",
"Anthropic",
"Claude Code OAuth",
"claude_api",
"OAuth",
"Authentication mode",
"Log in to Claude now?",
),
(
"_step_openai_auth_mode",
"openai_auth_mode",
"OpenAI",
"Codex OAuth",
"codex",
"Codex OAuth",
"OpenAI authentication mode",
"Log in to Codex now?",
),
],
)
def test_oauth_wrappers_use_provider_specific_ccproxy_flow(
self,
step_name,
config_attr,
provider_label,
oauth_choice_label,
ccproxy_provider,
status_label,
question_label,
login_prompt,
):
"""Anthropic/OpenAI wrappers share flow but keep provider-specific IDs."""
from EvoScientist.config.onboard import steps as onboard_steps
config = EvoScientistConfig(**{config_attr: "oauth"})
select_question = MagicMock()
select_question.ask.return_value = "oauth"
confirm_question = MagicMock()
confirm_question.ask.return_value = True
with (
patch(
"EvoScientist.ccproxy_manager.is_ccproxy_available", return_value=True
),
patch(
"EvoScientist.ccproxy_manager.check_ccproxy_auth",
return_value=(False, "not authenticated"),
) as mock_check_auth,
patch(
"EvoScientist.config.onboard.prompter.install_navigation_keys"
) as mock_nav,
patch(
"EvoScientist.config.onboard.steps.questionary.select",
return_value=select_question,
) as mock_select,
patch(
"EvoScientist.config.onboard.steps.questionary.confirm",
return_value=confirm_question,
) as mock_confirm,
patch(
"EvoScientist.config.onboard.steps._prompt_ccproxy_port"
) as mock_port,
patch("EvoScientist.config.onboard.steps._run_ccproxy_login") as mock_login,
):
result = getattr(onboard_steps, step_name)(config)
assert result == "oauth"
mock_nav.assert_called_once_with(select_question, with_back=True)
mock_port.assert_called_once_with(config)
mock_check_auth.assert_called_once_with(ccproxy_provider)
mock_login.assert_called_once_with(ccproxy_provider, status_label)
select_call = mock_select.call_args
assert select_call.args[0] == f"{question_label} [Esc/← to go back]:"
assert select_call.kwargs["default"] == "oauth"
choice_titles = [
choice.title
for choice in select_call.kwargs["choices"]
if getattr(choice, "value", None) in {"api_key", "oauth"}
]
assert choice_titles == [
f"API Key (direct {provider_label} access)",
f"{oauth_choice_label} (via ccproxy — no API key needed)",
]
mock_confirm.assert_called_once()
assert mock_confirm.call_args.args[0] == login_prompt
class TestStepModel:
def test_returns_selected_model(self):
"""Test that _step_model returns selected model."""
@@ -1401,7 +1266,6 @@ class TestRunOnboard:
"claude-sonnet-4-6", # Model
"assemble", # Auxiliary: Assemble
"openai", # Auxiliary provider (a different company)
"api_key", # Auxiliary OpenAI auth mode
"gpt-5.5", # Auxiliary model
"daemon", # Workspace mode
True, # Show thinking
@@ -1425,185 +1289,12 @@ class TestRunOnboard:
final_config = mock_save.call_args_list[-1].args[0]
assert final_config.auxiliary_provider == "openai"
assert final_config.auxiliary_model == "gpt-5.5"
assert final_config.openai_auth_mode == "api_key"
# The auxiliary provider's key is stored in its per-provider field.
assert final_config.openai_api_key == "sk-aux-openai"
# Main agent is untouched.
assert final_config.provider == "anthropic"
assert final_config.model == "claude-sonnet-4-6"
def test_auxiliary_same_provider_reuses_main_credentials(self):
"""Same-provider co-pilot should not imply separate credentials exist."""
from EvoScientist.config.onboard.wizard import run_onboard
mock_q = MagicMock()
with (
_patch_all_questionary(mock_q),
patch("EvoScientist.config.onboard.wizard.load_config") as mock_load,
patch("EvoScientist.config.onboard.wizard.save_config") as mock_save,
patch("EvoScientist.config.onboard.wizard.console"),
patch("EvoScientist.config.onboard.steps.console"),
):
mock_load.return_value = EvoScientistConfig(
provider="openai",
model="gpt-5.5",
openai_api_key="sk-main-openai",
)
mock_q.select.return_value.ask.side_effect = [
"assemble", # Auxiliary: Assemble
"openai", # Same provider as the main model
"gpt-5.5", # Auxiliary model
]
mock_q.confirm.return_value.ask.side_effect = [True] # Save config
result = run_onboard(
skip_validation=True, only_sections={"auxiliary_model"}
)
assert result is True
final_config = mock_save.call_args_list[-1].args[0]
assert final_config.auxiliary_provider == "openai"
assert final_config.auxiliary_model == "gpt-5.5"
assert final_config.openai_api_key == "sk-main-openai"
mock_q.password.assert_not_called()
assert mock_q.select.return_value.ask.call_count == 3
def test_auxiliary_same_provider_prompts_when_shared_key_missing(self):
"""Same-provider reuse should not hide a missing shared API key."""
from EvoScientist.config.onboard.wizard import run_onboard
mock_q = MagicMock()
with (
_patch_all_questionary(mock_q),
patch("EvoScientist.config.onboard.wizard.load_config") as mock_load,
patch("EvoScientist.config.onboard.wizard.save_config") as mock_save,
patch("EvoScientist.config.onboard.wizard.console"),
patch("EvoScientist.config.onboard.steps.console"),
patch("EvoScientist.config.onboard.helpers.console"),
patch(
"EvoScientist.ccproxy_manager.is_ccproxy_available", return_value=True
),
):
mock_load.return_value = EvoScientistConfig(
provider="openai",
model="gpt-5.5",
openai_auth_mode="api_key",
openai_api_key="",
)
mock_q.select.return_value.ask.side_effect = [
"assemble", # Auxiliary: Assemble
"openai", # Same provider as the main model
"api_key", # Shared OpenAI auth mode
"gpt-5.5", # Auxiliary model
]
mock_q.password.return_value.ask.side_effect = [
"sk-shared-openai",
]
mock_q.confirm.return_value.ask.side_effect = [True] # Save config
result = run_onboard(
skip_validation=True, only_sections={"auxiliary_model"}
)
assert result is True
final_config = mock_save.call_args_list[-1].args[0]
assert final_config.auxiliary_provider == "openai"
assert final_config.auxiliary_model == "gpt-5.5"
assert final_config.openai_auth_mode == "api_key"
assert final_config.openai_api_key == "sk-shared-openai"
mock_q.password.assert_called_once()
assert mock_q.select.return_value.ask.call_count == 4
def test_auxiliary_openai_oauth_skips_api_key(self):
"""Auxiliary OpenAI now uses the shared auth flow and skips keys on OAuth."""
from EvoScientist.config.onboard.wizard import run_onboard
mock_q = MagicMock()
with (
_patch_all_questionary(mock_q),
patch("EvoScientist.config.onboard.wizard.load_config") as mock_load,
patch("EvoScientist.config.onboard.wizard.save_config") as mock_save,
patch("EvoScientist.config.onboard.wizard.console"),
patch("EvoScientist.config.onboard.steps.console"),
patch("EvoScientist.config.onboard.helpers.console"),
patch(
"EvoScientist.ccproxy_manager.is_ccproxy_available", return_value=True
),
patch(
"EvoScientist.ccproxy_manager.check_ccproxy_auth",
return_value=(False, "not authenticated"),
) as mock_auth,
):
mock_load.return_value = EvoScientistConfig()
mock_q.select.return_value.ask.side_effect = [
"assemble", # Auxiliary: Assemble
"openai", # Auxiliary provider
"oauth", # OpenAI auth mode
"gpt-5.5", # Auxiliary model
]
mock_q.text.return_value.ask.side_effect = [
"", # ccproxy port (keep default)
]
mock_q.confirm.return_value.ask.side_effect = [
False, # Do not log in to Codex now
True, # Save config
]
result = run_onboard(
skip_validation=True, only_sections={"auxiliary_model"}
)
assert result is True
final_config = mock_save.call_args_list[-1].args[0]
assert final_config.auxiliary_provider == "openai"
assert final_config.auxiliary_model == "gpt-5.5"
assert final_config.openai_auth_mode == "oauth"
assert final_config.openai_api_key == ""
mock_q.password.assert_not_called()
mock_auth.assert_called_once_with("codex")
def test_auxiliary_reconfigure_clears_unused_openai_oauth(self):
"""Switching co-pilot away from OpenAI clears stale OpenAI OAuth mode."""
from EvoScientist.config.onboard.wizard import run_onboard
mock_q = MagicMock()
with (
_patch_all_questionary(mock_q),
patch("EvoScientist.config.onboard.wizard.load_config") as mock_load,
patch("EvoScientist.config.onboard.wizard.save_config") as mock_save,
patch("EvoScientist.config.onboard.wizard.console"),
patch("EvoScientist.config.onboard.steps.console"),
patch("EvoScientist.config.onboard.helpers.console"),
):
mock_load.return_value = EvoScientistConfig(
provider="anthropic",
model="claude-sonnet-4-6",
anthropic_auth_mode="oauth",
auxiliary_provider="openai",
auxiliary_model="gpt-5.5",
openai_auth_mode="oauth",
)
mock_q.select.return_value.ask.side_effect = [
"assemble", # Auxiliary: Assemble
"minimax", # Auxiliary provider no longer uses OpenAI
"global", # MiniMax region
"minimax-m2", # Auxiliary model
]
mock_q.password.return_value.ask.side_effect = [
"sk-minimax", # MiniMax API key
]
mock_q.confirm.return_value.ask.side_effect = [True] # Save config
result = run_onboard(
skip_validation=True, only_sections={"auxiliary_model"}
)
assert result is True
final_config = mock_save.call_args_list[-1].args[0]
assert final_config.auxiliary_provider == "minimax"
assert final_config.openai_auth_mode == "api_key"
assert final_config.anthropic_auth_mode == "oauth"
def test_auxiliary_custom_provider_collects_base_url(self):
"""Regression for the custom-provider fix: a custom auxiliary provider
collects its base URL (provider -> base URL -> key -> model order)."""
-59
View File
@@ -21,7 +21,6 @@ def _restore_paths():
"GLOBAL_MEMORIES_DIR": paths.GLOBAL_MEMORIES_DIR,
"USER_SKILLS_DIR": paths.USER_SKILLS_DIR,
"_active_workspace": paths._active_workspace,
"_EVOSCIENTIST_DATA_ROOT": paths._EVOSCIENTIST_DATA_ROOT,
}
yield
paths.WORKSPACE_ROOT = orig["WORKSPACE_ROOT"]
@@ -33,7 +32,6 @@ def _restore_paths():
paths.GLOBAL_MEMORIES_DIR = orig["GLOBAL_MEMORIES_DIR"]
paths.USER_SKILLS_DIR = orig["USER_SKILLS_DIR"]
paths._active_workspace = orig["_active_workspace"]
paths._EVOSCIENTIST_DATA_ROOT = orig["_EVOSCIENTIST_DATA_ROOT"]
class TestSetWorkspaceRoot:
@@ -142,63 +140,6 @@ class TestDataDir:
assert paths.GLOBAL_MEMORIES_DIR == paths.DATA_DIR / "memories"
class TestGatewayDataDirs:
def test_evoscientist_root_prefers_home_override(self, tmp_path, monkeypatch):
home = tmp_path / "runtime-home"
monkeypatch.setenv("EVOSCIENTIST_HOME", str(home))
assert paths.evoscientist_root() == home.resolve()
def test_evoscientist_root_falls_back_to_data_dir(self, tmp_path, monkeypatch):
data_dir = tmp_path / "app-data"
monkeypatch.delenv("EVOSCIENTIST_HOME", raising=False)
monkeypatch.setattr(paths, "DATA_DIR", data_dir)
assert paths.evoscientist_root() == data_dir.resolve()
def test_data_root_respects_environment_override(self, tmp_path, monkeypatch):
data_root = tmp_path / "web-data"
monkeypatch.setenv("EVOSCIENTIST_DATA_ROOT", str(data_root))
paths._EVOSCIENTIST_DATA_ROOT = None
assert paths._data_root() == data_root.resolve()
def test_user_thread_and_global_dirs_are_created(self, tmp_path, monkeypatch):
data_root = tmp_path / "web-data"
monkeypatch.setenv("EVOSCIENTIST_DATA_ROOT", str(data_root))
paths._EVOSCIENTIST_DATA_ROOT = None
user_dir = paths.user_data_dir("user-a")
thread_dir = paths.thread_data_dir("user-a", "thread-1")
shared_dir = paths.global_data_dir("user-a")
assert user_dir == data_root / "user-a"
assert thread_dir == user_dir / "thread-1"
assert shared_dir == user_dir / "__global__"
assert user_dir.is_dir()
assert thread_dir.is_dir()
assert shared_dir.is_dir()
def test_iter_user_data_dirs_yields_directories_only(self, tmp_path, monkeypatch):
data_root = tmp_path / "web-data"
monkeypatch.setenv("EVOSCIENTIST_DATA_ROOT", str(data_root))
paths._EVOSCIENTIST_DATA_ROOT = None
paths.user_data_dir("user-a")
paths.user_data_dir("user-b")
(data_root / "metadata.json").write_text("{}", encoding="utf-8")
assert {path.name for path in paths.iter_user_data_dirs()} == {
"user-a",
"user-b",
}
def test_uploads_dir_uses_evoscientist_root(self, tmp_path, monkeypatch):
home = tmp_path / "runtime-home"
monkeypatch.setenv("EVOSCIENTIST_HOME", str(home))
assert paths.uploads_dir() == home.resolve() / "uploads"
class TestLegacySessionsDbMigration:
"""Tests for migrate_legacy_sessions_db() — transitional helper.
+6 -4
View File
@@ -79,13 +79,14 @@ class TestPickSkillsInteractive:
class TestInstallSkillsHandlesEmpty:
"""InstallSkills.execute must distinguish None vs [] from the picker."""
async def test_empty_list_suppresses_cancel_message(self):
def test_empty_list_suppresses_cancel_message(self):
"""When picker returns [], user should NOT see "Browse cancelled"
(the picker already printed its own message)."""
from unittest.mock import AsyncMock
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.skills import InstallSkills
from tests.conftest import run_async as _run
ui = MagicMock()
ui.supports_interactive = True
@@ -96,17 +97,18 @@ class TestInstallSkillsHandlesEmpty:
"EvoScientist.tools.skills_manager.fetch_remote_skill_index",
return_value=_INDEX,
):
await InstallSkills().execute(ctx, [])
_run(InstallSkills().execute(ctx, []))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert not any("Browse cancelled" in m for m in msgs)
async def test_none_shows_cancel_message(self):
def test_none_shows_cancel_message(self):
"""When picker returns None (actual cancel), user sees the message."""
from unittest.mock import AsyncMock
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.skills import InstallSkills
from tests.conftest import run_async as _run
ui = MagicMock()
ui.supports_interactive = True
@@ -117,7 +119,7 @@ class TestInstallSkillsHandlesEmpty:
"EvoScientist.tools.skills_manager.fetch_remote_skill_index",
return_value=_INDEX,
):
await InstallSkills().execute(ctx, [])
_run(InstallSkills().execute(ctx, []))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Browse cancelled" in m for m in msgs)
+20 -14
View File
@@ -344,7 +344,9 @@ def test_profile_memory_uses_path_pointers_when_profiles_exceed_budget(
)
async def test_profile_memory_async_path_bootstraps_and_injects(tmp_path, monkeypatch):
def test_profile_memory_async_path_bootstraps_and_injects(
tmp_path, monkeypatch, run_async
):
memories = tmp_path / "memories"
workspace = tmp_path / "workspace"
workspace.mkdir()
@@ -354,7 +356,7 @@ async def test_profile_memory_async_path_bootstraps_and_injects(tmp_path, monkey
return request
middleware = memory_module.create_memory_middleware(str(memories))
await middleware.awrap_model_call(_request(), _handler)
run_async(middleware.awrap_model_call(_request(), _handler))
assert (memories / "profile" / "USER_PROFILE.md").exists()
@@ -397,8 +399,8 @@ def test_profile_memory_read_failure_uses_path_pointers_without_overwriting(
assert soul_path.read_bytes() == original_bytes
async def test_profile_memory_async_path_inlines_content_under_blockbuster(
tmp_path, monkeypatch
def test_profile_memory_async_path_inlines_content_under_blockbuster(
tmp_path, monkeypatch, run_async
):
memories = tmp_path / "memories"
workspace = tmp_path / "workspace"
@@ -423,13 +425,17 @@ async def test_profile_memory_async_path_inlines_content_under_blockbuster(
monkeypatch.setattr(middleware, "_read_profile_memory", tracked_read_profile_memory)
event_loop_thread = threading.get_ident()
blocker = BlockBuster(scanned_modules=memory_module)
blocker.activate()
try:
modified = await middleware.amodify_request(_request())
finally:
blocker.deactivate()
async def run():
event_loop_thread = threading.get_ident()
blocker = BlockBuster(scanned_modules=memory_module)
blocker.activate()
try:
modified = await middleware.amodify_request(_request())
finally:
blocker.deactivate()
return event_loop_thread, modified
event_loop_thread, modified = run_async(run())
assert call_threads
assert all(thread_id != event_loop_thread for thread_id in call_threads)
@@ -528,8 +534,8 @@ def test_profile_memory_uses_explicit_workspace_for_project_profile(
).exists()
async def test_profile_memory_resolves_project_id_once_per_middleware(
tmp_path, monkeypatch
def test_profile_memory_resolves_project_id_once_per_middleware(
tmp_path, monkeypatch, run_async
):
memories = tmp_path / "memories"
workspace = tmp_path / "workspace"
@@ -546,7 +552,7 @@ async def test_profile_memory_resolves_project_id_once_per_middleware(
str(memories), workspace_dir=str(workspace), max_inline_profile_chars=10
)
middleware.modify_request(_request())
await middleware.amodify_request(_request())
run_async(middleware.amodify_request(_request()))
assert calls == [workspace]
assert middleware.project_id == "P-cached-project"
+38 -35
View File
@@ -8,6 +8,7 @@ from EvoScientist.channels.qq.channel import (
QQConfig,
_build_qq_keyboard,
)
from tests.conftest import run_async as _run
class TestQQChannelSend:
@@ -21,7 +22,7 @@ class TestQQChannelSend:
channel._client.api.post_group_message = AsyncMock()
return channel
async def test_send_prefers_native_markdown_for_c2c(self):
def test_send_prefers_native_markdown_for_c2c(self):
channel = self._make_ready_channel()
msg = OutboundMessage(
channel="qq",
@@ -34,7 +35,7 @@ class TestQQChannelSend:
},
)
assert await channel.send(msg) is True
assert _run(channel.send(msg)) is True
channel._client.api.post_c2c_message.assert_awaited_once()
sent = channel._client.api.post_c2c_message.await_args.kwargs
@@ -45,7 +46,7 @@ class TestQQChannelSend:
assert sent["msg_seq"] == 1
assert "content" not in sent
async def test_send_falls_back_to_plain_text_when_markdown_send_fails(self):
def test_send_falls_back_to_plain_text_when_markdown_send_fails(self):
channel = self._make_ready_channel()
channel._trace_event = MagicMock(side_effect=RuntimeError("trace failed"))
channel._client.api.post_c2c_message = AsyncMock(
@@ -62,7 +63,7 @@ class TestQQChannelSend:
},
)
assert await channel.send(msg) is True
assert _run(channel.send(msg)) is True
assert channel._client.api.post_c2c_message.await_count == 2
first = channel._client.api.post_c2c_message.await_args_list[0].kwargs
@@ -79,7 +80,7 @@ class TestQQChannelSend:
# trigger "duplicate msg_seq".
assert second["msg_seq"] == 2
async def test_send_does_not_fallback_on_transport_error(self):
def test_send_does_not_fallback_on_transport_error(self):
channel = self._make_ready_channel()
async def _send_once(coro_factory, max_retries=3):
@@ -100,13 +101,13 @@ class TestQQChannelSend:
},
)
assert await channel.send(msg) is False
assert _run(channel.send(msg)) is False
channel._client.api.post_c2c_message.assert_awaited_once()
sent = channel._client.api.post_c2c_message.await_args.kwargs
assert sent["msg_type"] == 2
assert "content" not in sent
async def test_send_does_not_fallback_when_transport_error_mentions_markdown(self):
def test_send_does_not_fallback_when_transport_error_mentions_markdown(self):
"""A transport-layer error whose message incidentally contains the word
"markdown" must NOT be reclassified as a markdown compatibility failure,
otherwise genuine send failures get silently swallowed as plain-text."""
@@ -132,10 +133,10 @@ class TestQQChannelSend:
},
)
assert await channel.send(msg) is False
assert _run(channel.send(msg)) is False
channel._client.api.post_c2c_message.assert_awaited_once()
async def test_send_falls_back_on_qq_server_error_code(self):
def test_send_falls_back_on_qq_server_error_code(self):
"""QQ server-side markdown errors (e.g. 304014 template not configured)
should trigger plain-text fallback with a fresh msg_seq."""
channel = self._make_ready_channel()
@@ -158,7 +159,7 @@ class TestQQChannelSend:
},
)
assert await channel.send(msg) is True
assert _run(channel.send(msg)) is True
assert channel._client.api.post_c2c_message.await_count == 2
first = channel._client.api.post_c2c_message.await_args_list[0].kwargs
@@ -230,7 +231,7 @@ class TestQQSendWithButtons:
channel._client.api.post_group_message = AsyncMock()
return channel
async def test_c2c_send_attaches_keyboard(self):
def test_c2c_send_attaches_keyboard(self):
channel = self._make_channel()
msg = OutboundMessage(
channel="qq",
@@ -246,7 +247,7 @@ class TestQQSendWithButtons:
],
},
)
assert await channel.send(msg) is True
assert _run(channel.send(msg)) is True
sent = channel._client.api.post_c2c_message.await_args.kwargs
assert sent["msg_type"] == 2
@@ -255,7 +256,7 @@ class TestQQSendWithButtons:
assert rows[0]["buttons"][0]["action"]["data"] == "1"
assert rows[1]["buttons"][0]["action"]["data"] == "2"
async def test_group_send_does_not_attach_keyboard(self):
def test_group_send_does_not_attach_keyboard(self):
"""Group keyboards are out of scope — silently dropped."""
channel = self._make_channel()
msg = OutboundMessage(
@@ -269,11 +270,11 @@ class TestQQSendWithButtons:
"buttons": [{"text": "Approve", "value": "1"}],
},
)
assert await channel.send(msg) is True
assert _run(channel.send(msg)) is True
sent = channel._client.api.post_group_message.await_args.kwargs
assert "keyboard" not in sent
async def test_fallback_appends_button_hint_when_keyboard_present(self):
def test_fallback_appends_button_hint_when_keyboard_present(self):
"""If markdown send fails and we fall back to plain text, the
keyboard is lost — append a textual hint so the user still has
a way to reply (the values still pass `_parse_approval_reply`).
@@ -301,7 +302,7 @@ class TestQQSendWithButtons:
],
},
)
assert await channel.send(msg) is True
assert _run(channel.send(msg)) is True
plain_call = channel._client.api.post_c2c_message.await_args_list[1].kwargs
assert plain_call["msg_type"] == 0
@@ -310,7 +311,7 @@ class TestQQSendWithButtons:
assert "1=Approve" in plain_call["content"]
assert "2=Reject" in plain_call["content"]
async def test_fallback_hint_handles_non_string_button_value(self):
def test_fallback_hint_handles_non_string_button_value(self):
"""Regression: integer/None button values must not crash the
plain-text fallback (the keyboard builder already coerces them)."""
channel = self._make_channel()
@@ -334,7 +335,7 @@ class TestQQSendWithButtons:
],
},
)
assert await channel.send(msg) is True
assert _run(channel.send(msg)) is True
plain_call = channel._client.api.post_c2c_message.await_args_list[1].kwargs
assert "42=OK" in plain_call["content"]
assert "Cancel=Cancel" in plain_call["content"]
@@ -391,9 +392,9 @@ class TestQQInteractionCallback:
)
return interaction
async def test_click_publishes_to_bus_with_button_data(self):
def test_click_publishes_to_bus_with_button_data(self):
channel = self._make_channel()
await channel._on_interaction(self._make_interaction("1"))
_run(channel._on_interaction(self._make_interaction("1")))
channel._bus.publish_inbound.assert_awaited_once()
inbound = channel._bus.publish_inbound.await_args[0][0]
@@ -404,61 +405,63 @@ class TestQQInteractionCallback:
assert inbound.metadata["button_value"] == "1"
assert inbound.metadata["msg_type"] == "c2c"
async def test_click_acks_interaction(self):
def test_click_acks_interaction(self):
channel = self._make_channel()
await channel._on_interaction(self._make_interaction("1"))
_run(channel._on_interaction(self._make_interaction("1")))
channel._client.api.on_interaction_result.assert_awaited_once_with("intr_1", 0)
async def test_click_bypasses_debounce(self):
def test_click_bypasses_debounce(self):
"""Click never hits queue_message (debounce buffer)."""
channel = self._make_channel()
channel.queue_message = AsyncMock()
await channel._on_interaction(self._make_interaction("3"))
_run(channel._on_interaction(self._make_interaction("3")))
channel.queue_message.assert_not_called()
channel._bus.publish_inbound.assert_awaited_once()
async def test_group_interaction_ignored(self):
def test_group_interaction_ignored(self):
"""No user_openid → group/guild click → don't publish."""
channel = self._make_channel()
intr = self._make_interaction(user_openid="")
intr.group_openid = "group_xxx"
await channel._on_interaction(intr)
_run(channel._on_interaction(intr))
channel._bus.publish_inbound.assert_not_called()
# ACK still fires — it runs first, before the group-skip return.
channel._client.api.on_interaction_result.assert_awaited_once_with("intr_1", 0)
async def test_click_dropped_when_middleware_rejects(self):
def test_click_dropped_when_middleware_rejects(self):
channel = self._make_channel()
channel._build_inbound_async = AsyncMock(return_value=None)
await channel._on_interaction(self._make_interaction("1"))
_run(channel._on_interaction(self._make_interaction("1")))
channel._bus.publish_inbound.assert_not_called()
# ACK still fires (we don't want the user staring at a stuck button)
channel._client.api.on_interaction_result.assert_awaited_once()
async def test_empty_button_data_falls_back_to_button_id(self):
def test_empty_button_data_falls_back_to_button_id(self):
channel = self._make_channel()
await channel._on_interaction(
self._make_interaction(button_data="", button_id="btn_3")
_run(
channel._on_interaction(
self._make_interaction(button_data="", button_id="btn_3")
)
)
inbound = channel._bus.publish_inbound.await_args[0][0]
assert inbound.content == "btn_3"
async def test_ack_fires_even_when_handler_throws(self):
def test_ack_fires_even_when_handler_throws(self):
"""ACK must run before downstream processing so the QQ button UI
stays responsive even if middleware/bus crashes."""
channel = self._make_channel()
channel._build_inbound_async = AsyncMock(side_effect=RuntimeError("boom"))
# Should not raise — handler swallows downstream errors.
await channel._on_interaction(self._make_interaction("1"))
_run(channel._on_interaction(self._make_interaction("1")))
channel._client.api.on_interaction_result.assert_awaited_once_with("intr_1", 0)
async def test_button_value_metadata_is_string_coerced(self):
def test_button_value_metadata_is_string_coerced(self):
"""Regression: metadata['button_value'] must be a string (was raw)."""
channel = self._make_channel()
resolved = MagicMock(button_id="btn_0", button_data=42, message_id="msg_orig")
data = MagicMock(type=None, resolved=resolved)
intr = MagicMock(id="intr_1", user_openid="u_x", group_openid=None, data=data)
await channel._on_interaction(intr)
_run(channel._on_interaction(intr))
inbound = channel._bus.publish_inbound.await_args[0][0]
assert inbound.content == "42"
assert inbound.metadata["button_value"] == "42"
-177
View File
@@ -1,177 +0,0 @@
"""Deterministic tool-loop guard and provider projection tests."""
from dataclasses import dataclass, replace
from typing import Any
import pytest
from langchain.agents.middleware.types import ModelResponse
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
from EvoScientist.llm.errors import AgentControlError
from EvoScientist.middleware.repetitive_tool_guard import (
RepetitiveToolCallGuardMiddleware,
collapse_repetitive_tool_rounds,
)
def _round(
call_id: str,
*,
name: str = "execute",
command: str = "pwd",
content: str = "Error: invalid argument: command rejected by schema",
status: str = "error",
) -> list[Any]:
return [
AIMessage(
content="",
tool_calls=[{"id": call_id, "name": name, "args": {"command": command}}],
),
ToolMessage(
content=content,
tool_call_id=call_id,
name=name,
status=status,
),
]
@dataclass(frozen=True)
class _Request:
messages: list[Any]
tools: list[Any]
def override(self, **updates: Any):
return replace(self, **updates)
def test_provider_projection_keeps_first_and_last_deterministic_error_rounds():
messages = [HumanMessage(content="inspect")]
for index in range(4):
messages.extend(_round(f"call-{index}"))
messages.append(HumanMessage(content="continue"))
repair = collapse_repetitive_tool_rounds(messages, threshold=2)
assert repair.removed_rounds == 2
assert [m.type for m in repair.messages] == [
"human",
"ai",
"tool",
"ai",
"tool",
"human",
]
assert repair.messages[1].tool_calls[0]["id"] == "call-0"
assert repair.messages[3].tool_calls[0]["id"] == "call-3"
def test_successful_repeated_calls_are_never_projected_away():
messages = [
*_round("call-1", content="ok", status="success"),
*_round("call-2", content="ok", status="success"),
*_round("call-3", content="ok", status="success"),
]
repair = collapse_repetitive_tool_rounds(messages)
assert repair.messages == messages
assert repair.removed_rounds == 0
assert repair.tail_repetitions == 0
def test_transient_and_unknown_errors_do_not_count_as_semantic_loop():
transient = [
*_round("call-1", content="Error: connection timeout"),
*_round("call-2", content="Error: connection timeout"),
]
unknown = [
*_round("call-3", content="Error: something unusual"),
*_round("call-4", content="Error: something unusual"),
]
assert collapse_repetitive_tool_rounds(transient).tail_repetitions == 0
assert collapse_repetitive_tool_rounds(unknown).tail_consecutive_errors == 0
def test_generic_raw_execution_error_code_remains_unknown():
messages = _round("call-1", content="Error: something unusual")
messages[1].additional_kwargs["error_code"] = "TOOL_EXECUTION_FAILED"
repair = collapse_repetitive_tool_rounds(messages)
assert repair.tail_consecutive_errors == 0
def test_identical_tail_loop_stops_before_next_model_call():
request = _Request(
messages=[*_round("call-1"), *_round("call-2")],
tools=[{"name": "execute"}],
)
called = False
def handler(_request):
nonlocal called
called = True
return ModelResponse(result=[AIMessage(content="should not run")])
with pytest.raises(AgentControlError) as caught:
RepetitiveToolCallGuardMiddleware(threshold=2).wrap_model_call(request, handler)
assert caught.value.code == "MODEL_TOOL_LOOP_DETECTED"
assert called is False
def test_different_deterministic_errors_hit_consecutive_limit():
request = _Request(
messages=[
*_round("one", name="execute"),
*_round("two", name="read_file"),
*_round("three", name="search"),
],
tools=[],
)
with pytest.raises(AgentControlError) as caught:
RepetitiveToolCallGuardMiddleware(
threshold=0, max_consecutive_errors=3
).wrap_model_call(request, lambda _request: None)
assert caught.value.code == "MODEL_TOOL_ERROR_LIMIT"
def test_user_message_breaks_tail_loop_but_historical_projection_is_temporary():
original = [
*_round("call-1"),
*_round("call-2"),
*_round("call-3"),
HumanMessage(content="try a new approach"),
]
request = _Request(messages=original, tools=[])
captured = []
def handler(prepared):
captured.append(prepared)
return ModelResponse(result=[AIMessage(content="continued")])
RepetitiveToolCallGuardMiddleware().wrap_model_call(request, handler)
assert len(captured[0].messages) == 5
assert len(original) == 7
def test_zero_thresholds_disable_only_semantic_loop_guards():
request = _Request(messages=[*_round("one"), *_round("two")], tools=[])
captured = []
middleware = RepetitiveToolCallGuardMiddleware(
threshold=0, max_consecutive_errors=0
)
middleware.wrap_model_call(
request,
lambda prepared: (
captured.append(prepared) or ModelResponse(result=[AIMessage(content="ok")])
),
)
assert captured == [request]
@pytest.mark.parametrize("kwargs", [{"threshold": -1}, {"max_consecutive_errors": -1}])
def test_negative_threshold_is_rejected(kwargs):
with pytest.raises(ValueError, match="non-negative"):
RepetitiveToolCallGuardMiddleware(**kwargs)
+17 -16
View File
@@ -2,6 +2,7 @@
from unittest.mock import AsyncMock, MagicMock
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway, FakeThreadStore
@@ -23,7 +24,7 @@ def _ctx(thread_id="current", workspace_dir="/ws", thread_store=None):
class TestResumeCommand:
async def test_with_arg_resolves_and_calls_ui(self):
def test_with_arg_resolves_and_calls_ui(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx(
@@ -32,23 +33,23 @@ class TestResumeCommand:
metadata={"workspace_dir": "/restored"},
)
)
await ResumeCommand().execute(ctx, ["target-tid"])
_run(ResumeCommand().execute(ctx, ["target-tid"]))
ui.handle_session_resume.assert_awaited_once_with("target-tid", "/restored")
# ctx mutations
assert ctx.thread_id == "target-tid"
assert ctx.workspace_dir == "/restored"
async def test_no_arg_empty_threads_prints_message(self):
def test_no_arg_empty_threads_prints_message(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx()
await ResumeCommand().execute(ctx, [])
_run(ResumeCommand().execute(ctx, []))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("No sessions to resume" in m for m in msgs)
ui.wait_for_thread_pick.assert_not_called()
ui.handle_session_resume.assert_not_called()
async def test_no_arg_calls_picker(self):
def test_no_arg_calls_picker(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx()
@@ -59,11 +60,11 @@ class TestResumeCommand:
resolved_thread_id="picked-tid",
)
ctx.graph_gateway = FakeGraphGateway(thread_store=store)
await ResumeCommand().execute(ctx, [])
_run(ResumeCommand().execute(ctx, []))
ui.wait_for_thread_pick.assert_awaited_once()
ui.handle_session_resume.assert_awaited_once()
async def test_picker_cancel_returns(self):
def test_picker_cancel_returns(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx()
@@ -71,28 +72,28 @@ class TestResumeCommand:
threads = [{"thread_id": "t1", "preview": "", "message_count": 0}]
store = FakeThreadStore(threads=threads)
ctx.graph_gateway = FakeGraphGateway(thread_store=store)
await ResumeCommand().execute(ctx, [])
_run(ResumeCommand().execute(ctx, []))
ui.handle_session_resume.assert_not_called()
async def test_ambiguous_prefix(self):
def test_ambiguous_prefix(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx(thread_store=FakeThreadStore(matches=["abc-one", "abc-two"]))
await ResumeCommand().execute(ctx, ["abc"])
_run(ResumeCommand().execute(ctx, ["abc"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Ambiguous" in m for m in msgs)
ui.handle_session_resume.assert_not_called()
async def test_not_found(self):
def test_not_found(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx()
await ResumeCommand().execute(ctx, ["missing"])
_run(ResumeCommand().execute(ctx, ["missing"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("not found" in m for m in msgs)
ui.handle_session_resume.assert_not_called()
async def test_prefix_resolves_to_unique_match(self):
def test_prefix_resolves_to_unique_match(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx(
@@ -101,18 +102,18 @@ class TestResumeCommand:
metadata={"workspace_dir": "/ws1"},
)
)
await ResumeCommand().execute(ctx, ["abc"])
_run(ResumeCommand().execute(ctx, ["abc"]))
ui.handle_session_resume.assert_awaited_once_with("abc-one", "/ws1")
assert ctx.thread_id == "abc-one"
async def test_empty_workspace_metadata_preserves_ctx_workspace(self):
def test_empty_workspace_metadata_preserves_ctx_workspace(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx(
workspace_dir="/keep",
thread_store=FakeThreadStore(resolved_thread_id="tid", metadata={}),
)
await ResumeCommand().execute(ctx, ["tid"])
_run(ResumeCommand().execute(ctx, ["tid"]))
# ResumeCommand only overwrites ctx.workspace_dir if metadata has one
assert ctx.workspace_dir == "/keep"
# Callback still fires with the metadata value (empty string)
+56 -46
View File
@@ -5,6 +5,8 @@ from unittest.mock import MagicMock
from rich.console import Console
from rich.table import Table
from tests.conftest import run_async as _run
def _make_ui(**kwargs):
"""Build a RichCLICommandUI backed by a MagicMock console."""
@@ -38,9 +40,9 @@ class TestBasicIO:
ui.mount_renderable(table)
console.print.assert_called_once_with(table)
async def test_flush_is_async_noop(self):
def test_flush_is_async_noop(self):
ui, console = _make_ui()
await ui.flush()
_run(ui.flush())
# flush should not print anything
console.print.assert_not_called()
@@ -48,29 +50,33 @@ class TestBasicIO:
class TestWaitForModelPick:
"""CLI model picker fallback: print table + return None."""
async def test_returns_none(self):
def test_returns_none(self):
ui, _ = _make_ui()
entries = [
("claude-sonnet-4-6", "anthropic/claude-sonnet", "anthropic"),
("gpt-4o", "openai/gpt-4o", "openai"),
]
result = await ui.wait_for_model_pick(
entries,
current_model="claude-sonnet-4-6",
current_provider="anthropic",
result = _run(
ui.wait_for_model_pick(
entries,
current_model="claude-sonnet-4-6",
current_provider="anthropic",
)
)
assert result is None
async def test_prints_table_with_current_model_marker(self):
def test_prints_table_with_current_model_marker(self):
ui, console = _make_ui()
entries = [
("claude-sonnet-4-6", "anthropic/claude-sonnet", "anthropic"),
("gpt-4o", "openai/gpt-4o", "openai"),
]
await ui.wait_for_model_pick(
entries,
current_model="claude-sonnet-4-6",
current_provider="anthropic",
_run(
ui.wait_for_model_pick(
entries,
current_model="claude-sonnet-4-6",
current_provider="anthropic",
)
)
# First call renders the Table (Rich renderable), second prints usage.
assert console.print.call_count == 2
@@ -81,23 +87,27 @@ class TestWaitForModelPick:
assert "Usage: /model" in usage_arg
assert "--save" in usage_arg
async def test_no_current_model_no_marker(self):
def test_no_current_model_no_marker(self):
ui, console = _make_ui()
entries = [("claude-sonnet-4-6", "anthropic/claude-sonnet", "anthropic")]
await ui.wait_for_model_pick(
entries,
current_model=None,
current_provider=None,
_run(
ui.wait_for_model_pick(
entries,
current_model=None,
current_provider=None,
)
)
# Just asserts the coroutine runs without marker-branch issues.
assert console.print.call_count == 2
async def test_empty_entries_still_prints_header_and_usage(self):
def test_empty_entries_still_prints_header_and_usage(self):
ui, console = _make_ui()
result = await ui.wait_for_model_pick(
[],
current_model=None,
current_provider=None,
result = _run(
ui.wait_for_model_pick(
[],
current_model=None,
current_provider=None,
)
)
assert result is None
# Header table + usage hint should still be printed even with
@@ -203,7 +213,7 @@ class TestWaitForThreadPick:
},
]
async def test_returns_selected_thread_id(self, monkeypatch):
def test_returns_selected_thread_id(self, monkeypatch):
import EvoScientist.cli.rich_command_ui as mod
ui, _ = _make_ui()
@@ -216,7 +226,7 @@ class TestWaitForThreadPick:
return prompt
monkeypatch.setattr("questionary.select", fake_select)
result = await ui.wait_for_thread_pick(self._threads(), "abc123", "pick:")
result = _run(ui.wait_for_thread_pick(self._threads(), "abc123", "pick:"))
assert result == "abc123"
assert called["title"] == "pick:"
# _build_items prepends a workspace header — choices has headers +
@@ -225,14 +235,14 @@ class TestWaitForThreadPick:
# Table import removed; this test no longer depends on console output.
assert mod.RichCLICommandUI is not None # sanity
async def test_cancel_returns_none(self, monkeypatch):
def test_cancel_returns_none(self, monkeypatch):
ui, _ = _make_ui()
prompt = self._fake_prompt(None)
monkeypatch.setattr("questionary.select", lambda *a, **k: prompt)
result = await ui.wait_for_thread_pick(self._threads(), "abc123", "pick:")
result = _run(ui.wait_for_thread_pick(self._threads(), "abc123", "pick:"))
assert result is None
async def test_current_thread_marker_in_label(self, monkeypatch):
def test_current_thread_marker_in_label(self, monkeypatch):
ui, _ = _make_ui()
prompt = self._fake_prompt(None)
captured_choices: list = []
@@ -242,7 +252,7 @@ class TestWaitForThreadPick:
return prompt
monkeypatch.setattr("questionary.select", fake_select)
await ui.wait_for_thread_pick(self._threads(), "abc123", "pick:")
_run(ui.wait_for_thread_pick(self._threads(), "abc123", "pick:"))
# At least one Choice title contains "abc123 *" (current marker)
choice_titles = [getattr(c, "title", "") for c in captured_choices]
assert any("abc123 *" in t for t in choice_titles)
@@ -272,38 +282,38 @@ class TestCompactIndicator:
class TestPhaseBMigrated:
"""Session lifecycle callbacks (start/resume) filled in Phase B."""
async def test_start_new_session_fires_callback(self):
def test_start_new_session_fires_callback(self):
from unittest.mock import AsyncMock
cb = AsyncMock()
ui, _ = _make_ui(on_start_new_session=cb)
await ui.start_new_session()
_run(ui.start_new_session())
cb.assert_awaited_once()
async def test_start_new_session_without_callback_is_noop(self):
def test_start_new_session_without_callback_is_noop(self):
ui, console = _make_ui()
await ui.start_new_session()
_run(ui.start_new_session())
console.print.assert_not_called()
async def test_handle_session_resume_awaits_callback(self):
def test_handle_session_resume_awaits_callback(self):
from unittest.mock import AsyncMock
cb = AsyncMock()
ui, _ = _make_ui(on_handle_session_resume=cb)
await ui.handle_session_resume("tid-x", "/workspace")
_run(ui.handle_session_resume("tid-x", "/workspace"))
cb.assert_awaited_once_with("tid-x", "/workspace")
async def test_handle_session_resume_without_callback_is_noop(self):
def test_handle_session_resume_without_callback_is_noop(self):
ui, _ = _make_ui()
# Should not raise
await ui.handle_session_resume("tid-x")
_run(ui.handle_session_resume("tid-x"))
async def test_handle_session_resume_workspace_defaults_none(self):
def test_handle_session_resume_workspace_defaults_none(self):
from unittest.mock import AsyncMock
cb = AsyncMock()
ui, _ = _make_ui(on_handle_session_resume=cb)
await ui.handle_session_resume("tid-x")
_run(ui.handle_session_resume("tid-x"))
cb.assert_awaited_once_with("tid-x", None)
@@ -311,7 +321,7 @@ class TestPhaseCMigrated:
"""Skill/MCP browse pickers delegate to questionary helpers via
``asyncio.to_thread`` since questionary blocks the event loop."""
async def test_skill_browse_delegates_to_picker(self, monkeypatch):
def test_skill_browse_delegates_to_picker(self, monkeypatch):
from unittest.mock import MagicMock
picker = MagicMock(return_value=["skill-a", "skill-b"])
@@ -320,11 +330,11 @@ class TestPhaseCMigrated:
picker,
)
ui, _ = _make_ui()
result = await ui.wait_for_skill_browse([{"name": "a"}], {"installed"}, "core")
result = _run(ui.wait_for_skill_browse([{"name": "a"}], {"installed"}, "core"))
assert result == ["skill-a", "skill-b"]
picker.assert_called_once_with([{"name": "a"}], {"installed"}, "core")
async def test_skill_browse_cancel_returns_none(self, monkeypatch):
def test_skill_browse_cancel_returns_none(self, monkeypatch):
from unittest.mock import MagicMock
monkeypatch.setattr(
@@ -332,10 +342,10 @@ class TestPhaseCMigrated:
MagicMock(return_value=None),
)
ui, _ = _make_ui()
result = await ui.wait_for_skill_browse([], set(), "")
result = _run(ui.wait_for_skill_browse([], set(), ""))
assert result is None
async def test_mcp_browse_delegates_to_picker(self, monkeypatch):
def test_mcp_browse_delegates_to_picker(self, monkeypatch):
from unittest.mock import MagicMock
sentinel_entries = [MagicMock(name="entry1"), MagicMock(name="entry2")]
@@ -345,11 +355,11 @@ class TestPhaseCMigrated:
picker,
)
ui, _ = _make_ui()
result = await ui.wait_for_mcp_browse([MagicMock()], {"configured"}, "")
result = _run(ui.wait_for_mcp_browse([MagicMock()], {"configured"}, ""))
assert result is sentinel_entries
picker.assert_called_once()
async def test_mcp_browse_cancel_returns_none(self, monkeypatch):
def test_mcp_browse_cancel_returns_none(self, monkeypatch):
from unittest.mock import MagicMock
monkeypatch.setattr(
@@ -357,5 +367,5 @@ class TestPhaseCMigrated:
MagicMock(return_value=None),
)
ui, _ = _make_ui()
result = await ui.wait_for_mcp_browse([], set(), "")
result = _run(ui.wait_for_mcp_browse([], set(), ""))
assert result is None
-117
View File
@@ -1,117 +0,0 @@
from __future__ import annotations
import ast
from pathlib import Path
import pytest
from EvoScientist.runtime_integrations import (
RuntimeIntegrationUnavailable,
configure_runtime_integrations,
get_app_connection,
get_image_backend,
get_session_connection,
get_session_dsn,
handle_knowledge_file,
record_service_usage,
reset_runtime_integrations,
resolve_runtime_model,
)
@pytest.fixture(autouse=True)
def reset_integrations():
reset_runtime_integrations()
yield
reset_runtime_integrations()
def test_core_package_does_not_import_gateway():
package_root = Path(__file__).resolve().parents[1] / "EvoScientist"
violations = []
for source_file in package_root.rglob("*.py"):
tree = ast.parse(
source_file.read_text(encoding="utf-8"), filename=str(source_file)
)
for node in ast.walk(tree):
if isinstance(node, ast.Import):
names = [alias.name for alias in node.names]
elif isinstance(node, ast.ImportFrom):
if node.level:
continue
names = [node.module or ""]
else:
continue
if any(name == "gateway" or name.startswith("gateway.") for name in names):
violations.append(
f"{source_file.relative_to(package_root)}:{node.lineno}"
)
assert violations == []
@pytest.mark.anyio
async def test_optional_integrations_are_safe_without_web_runtime(tmp_path):
assert get_session_dsn() is None
await handle_knowledge_file(tmp_path / "result.md")
await record_service_usage("search", "query")
with pytest.raises(RuntimeIntegrationUnavailable):
await get_app_connection()
with pytest.raises(RuntimeIntegrationUnavailable):
await get_session_connection()
with pytest.raises(RuntimeIntegrationUnavailable):
get_image_backend()
@pytest.mark.anyio
async def test_host_can_register_runtime_integrations(tmp_path):
app_connection = object()
session_connection = object()
knowledge_paths = []
usage = []
image_backend = object()
async def provide_app_connection():
return app_connection
async def provide_session_connection():
return session_connection
async def handle_file(path):
knowledge_paths.append(path)
async def record_usage(service, action):
usage.append((service, action))
configure_runtime_integrations(
app_connection_provider=provide_app_connection,
session_connection_provider=provide_session_connection,
session_dsn_provider=lambda: "postgresql://example/session",
knowledge_file_handler=handle_file,
usage_recorder=record_usage,
image_backend_factory=lambda: image_backend,
)
path = tmp_path / "result.md"
await handle_knowledge_file(path)
await record_service_usage("mineru", "parse")
assert await get_app_connection() is app_connection
assert await get_session_connection() is session_connection
assert get_session_dsn() == "postgresql://example/session"
assert get_image_backend() is image_backend
assert knowledge_paths == [path]
assert usage == [("mineru", "parse")]
def test_host_can_register_model_resolver():
resolved = object()
calls = []
def resolve_model(model, provider):
calls.append((model, provider))
return resolved
configure_runtime_integrations(model_resolver=resolve_model)
assert resolve_runtime_model("model-a", "provider-a") is resolved
assert calls == [("model-a", "provider-a")]
+28 -24
View File
@@ -2,6 +2,8 @@
from unittest.mock import MagicMock, patch
from tests.conftest import run_async as _run
def _ctx():
from EvoScientist.commands.base import CommandContext
@@ -10,17 +12,17 @@ def _ctx():
return CommandContext(agent=None, thread_id="tid", ui=ui), ui
async def test_list_when_backend_down():
def test_list_when_backend_down():
from EvoScientist.commands.implementation.schedule import ScheduleCommand
ctx, ui = _ctx()
with patch("EvoScientist.cron.schedule.is_available", return_value=False):
await ScheduleCommand().execute(ctx, ["list"])
_run(ScheduleCommand().execute(ctx, ["list"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("unavailable" in m.lower() for m in msgs)
async def test_add_parses_five_field_cron_and_prompt():
def test_add_parses_five_field_cron_and_prompt():
from EvoScientist.commands.implementation.schedule import ScheduleCommand
ctx, _ui = _ctx()
@@ -31,15 +33,17 @@ async def test_add_parses_five_field_cron_and_prompt():
return_value={"cron_id": "c-9"},
) as mk,
):
await ScheduleCommand().execute(
ctx, ["add", "*/10", "*", "*", "*", "*", "search", "uk", "weather"]
_run(
ScheduleCommand().execute(
ctx, ["add", "*/10", "*", "*", "*", "*", "search", "uk", "weather"]
)
)
kw = mk.call_args.kwargs
assert kw["schedule"] == "*/10 * * * *"
assert kw["prompt"] == "search uk weather"
async def test_list_renders_table():
def test_list_renders_table():
from EvoScientist.commands.implementation.schedule import ScheduleCommand
ctx, ui = _ctx()
@@ -56,11 +60,11 @@ async def test_list_renders_table():
patch("EvoScientist.cron.schedule.is_available", return_value=True),
patch("EvoScientist.cron.schedule.list_schedules", return_value=rows),
):
await ScheduleCommand().execute(ctx, ["list"])
_run(ScheduleCommand().execute(ctx, ["list"]))
ui.mount_renderable.assert_called_once()
async def test_add_parses_quoted_cron_and_prompt():
def test_add_parses_quoted_cron_and_prompt():
from EvoScientist.commands.implementation.schedule import ScheduleCommand
ctx, _ui = _ctx()
@@ -71,15 +75,15 @@ async def test_add_parses_quoted_cron_and_prompt():
return_value={"cron_id": "c-9"},
) as mk,
):
await ScheduleCommand().execute(
ctx, ["add", "*/10 * * * *", "search uk weather"]
_run(
ScheduleCommand().execute(ctx, ["add", "*/10 * * * *", "search uk weather"])
)
kw = mk.call_args.kwargs
assert kw["schedule"] == "*/10 * * * *"
assert kw["prompt"] == "search uk weather"
async def test_run_with_matching_prefix_fires_matched_prompt():
def test_run_with_matching_prefix_fires_matched_prompt():
from EvoScientist.commands.implementation.schedule import ScheduleCommand
ctx, _ui = _ctx()
@@ -92,11 +96,11 @@ async def test_run_with_matching_prefix_fires_matched_prompt():
return_value={"run_id": "r-1"},
) as rn,
):
await ScheduleCommand().execute(ctx, ["run", "c-123"])
_run(ScheduleCommand().execute(ctx, ["run", "c-123"]))
rn.assert_called_once_with("do the thing")
async def test_run_with_no_match_reports():
def test_run_with_no_match_reports():
from EvoScientist.commands.implementation.schedule import ScheduleCommand
ctx, ui = _ctx()
@@ -105,13 +109,13 @@ async def test_run_with_no_match_reports():
patch("EvoScientist.cron.schedule.list_schedules", return_value=[]),
patch("EvoScientist.cron.schedule.run_now") as rn,
):
await ScheduleCommand().execute(ctx, ["run", "nope"])
_run(ScheduleCommand().execute(ctx, ["run", "nope"]))
rn.assert_not_called()
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("No schedule matching" in m for m in msgs)
async def test_pause_resume_set_enabled_with_resolved_id():
def test_pause_resume_set_enabled_with_resolved_id():
from EvoScientist.commands.implementation.schedule import ScheduleCommand
rows = [{"cron_id": "c-abcdef", "metadata": {"name": "t"}}]
@@ -122,7 +126,7 @@ async def test_pause_resume_set_enabled_with_resolved_id():
patch("EvoScientist.cron.schedule.list_schedules", return_value=rows),
patch("EvoScientist.cron.schedule.set_enabled") as se,
):
await ScheduleCommand().execute(ctx, [sub, "c-abc"])
_run(ScheduleCommand().execute(ctx, [sub, "c-abc"]))
se.assert_called_once_with("c-abcdef", expected)
@@ -131,7 +135,7 @@ async def test_pause_resume_set_enabled_with_resolved_id():
# ---------------------------------------------------------------------------
async def test_list_error_shows_red_message_no_exception():
def test_list_error_shows_red_message_no_exception():
"""B1: list_schedules raising after is_available() shows a red error, not a traceback."""
from EvoScientist.commands.implementation.schedule import ScheduleCommand
@@ -143,7 +147,7 @@ async def test_list_error_shows_red_message_no_exception():
side_effect=RuntimeError("backend gone"),
),
):
await ScheduleCommand().execute(ctx, ["list"])
_run(ScheduleCommand().execute(ctx, ["list"]))
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Error:" in m for m in msgs)
# Verify no exception escaped (test would have raised above otherwise)
@@ -154,7 +158,7 @@ async def test_list_error_shows_red_message_no_exception():
# ---------------------------------------------------------------------------
async def test_remove_ambiguous_prefix_aborts_without_deleting():
def test_remove_ambiguous_prefix_aborts_without_deleting():
"""B2: two crons sharing a prefix → ambiguity message, delete NOT called."""
from EvoScientist.commands.implementation.schedule import ScheduleCommand
@@ -168,7 +172,7 @@ async def test_remove_ambiguous_prefix_aborts_without_deleting():
patch("EvoScientist.cron.schedule.list_schedules", return_value=rows),
patch("EvoScientist.cron.schedule.delete_schedule") as mk,
):
await ScheduleCommand().execute(ctx, ["remove", "abc"])
_run(ScheduleCommand().execute(ctx, ["remove", "abc"]))
mk.assert_not_called()
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Multiple" in m for m in msgs)
@@ -179,7 +183,7 @@ async def test_remove_ambiguous_prefix_aborts_without_deleting():
# ---------------------------------------------------------------------------
async def test_remove_backend_error_shows_red_error_not_no_match():
def test_remove_backend_error_shows_red_error_not_no_match():
"""FIX 1: list_schedules() crashing in _resolve → red 'Error:' message, not 'No schedule matching'."""
from EvoScientist.commands.implementation.schedule import ScheduleCommand
@@ -192,7 +196,7 @@ async def test_remove_backend_error_shows_red_error_not_no_match():
),
patch("EvoScientist.cron.schedule.delete_schedule") as mk,
):
await ScheduleCommand().execute(ctx, ["remove", "abc"])
_run(ScheduleCommand().execute(ctx, ["remove", "abc"]))
mk.assert_not_called()
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Error:" in m for m in msgs), f"Expected red Error: message, got: {msgs}"
@@ -205,7 +209,7 @@ async def test_remove_backend_error_shows_red_error_not_no_match():
)
async def test_add_name_sanitized_from_nasty_prompt():
def test_add_name_sanitized_from_nasty_prompt():
"""B3: prompt with newline / slashes / special chars → clean kebab-case name."""
import re
@@ -221,7 +225,7 @@ async def test_add_name_sanitized_from_nasty_prompt():
return_value={"cron_id": "c-x"},
) as mk,
):
await ScheduleCommand().execute(ctx, ["add", "*/5 * * * *", nasty_prompt])
_run(ScheduleCommand().execute(ctx, ["add", "*/5 * * * *", nasty_prompt]))
name = mk.call_args.kwargs["name"]
# Must be non-empty, no spaces, no newlines, no slashes
assert name

Some files were not shown because too many files have changed in this diff Show More