Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| e0acc6155e |
@@ -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 |
@@ -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 |
@@ -48,3 +48,4 @@ conversation_history/
|
||||
*meals/
|
||||
botpy.log
|
||||
large_tool_results/
|
||||
runs/
|
||||
|
||||
+25
-136
@@ -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,
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"],
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
@@ -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
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
*,
|
||||
|
||||
@@ -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(),
|
||||
)
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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. "
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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()
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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>
|
||||
|
||||
|
||||
@@ -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.
|
||||
@@ -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
@@ -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
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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"}]
|
||||
|
||||
@@ -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
@@ -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())
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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."
|
||||
|
||||
+710
-623
File diff suppressed because it is too large
Load Diff
+58
-27
@@ -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
@@ -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()
|
||||
|
||||
@@ -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,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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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=[],
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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"
|
||||
@@ -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
@@ -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))
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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
@@ -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(
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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, []))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
|
||||
@@ -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)."""
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")]
|
||||
@@ -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
Reference in New Issue
Block a user