18 Commits

Author SHA1 Message Date
m4 3ce5614254 fix: harden tool-call protocol and fallback handling
Test / pytest (ubuntu-latest, 3.11) (pull_request) Has been cancelled
Build / build (pull_request) Has been cancelled
Docker / build (pull_request) Has been cancelled
Lint / ruff (pull_request) Has been cancelled
Test / pytest (ubuntu-latest, 3.12) (pull_request) Has been cancelled
Test / pytest (windows-latest, 3.11) (pull_request) Has been cancelled
Test / pytest (windows-latest, 3.12) (pull_request) Has been cancelled
2026-07-19 12:05:56 +08:00
m4 4fc74e7da7 EvoScientist Ai4Sci 2026-07-14 22:07:14 +08:00
jfilipiuk 753c745405 fix: silence YAML-docstring noise from custom-app OpenAPI scan (#317) 2026-07-13 15:38:28 +01:00
jfilipiuk 88ac9f5ba1 fix: surface real exception class+message in SSE error events (#315)
* fix: surface real exception class+message in SSE error events

* fix: tighten SSE error patch scope and key redaction

* fix: redact base64-style secret suffixes fully

* style: remove notes/ reference from the dosctring

* fix: rebuild env cache on each error call

* fix: route BaseException through serde.default on SSE/webhook paths

* fix: distinguish routed providers by request URL host

* feat: normalize provider-SDK exceptions via ErrorNormalizationMiddleware

* refactor: drop json_dumpb dataclass-bypass wrappers, superseded by middleware

* fix: guard _extract_host against SDK properties that raise

* refactor: derive provider tag from ModelRequest.model, not the exception

* refactor: drop serde.default patch and exception-based inference; ProviderStreamError.model_dump handles the emit

* refactor: move envelope helpers from patches.py to errors.py

* feat: extend ErrorNormalizationMiddleware coverage to every model-call path

* chore: clean up review findings from middleware pivot

* fix: pass through all langgraph.errors

* fix: move langgraph.errors pass-through into _normalize

* fix: pass through ContextOverflowError in _normalize
2026-07-13 14:17:56 +01:00
jfilipiuk 952e68efe3 fix: scope quickjs snapshot to turn to keep checkpoints small (#316)
* fix: scope quickjs snapshot to turn to keep checkpoints small

* fix: strip _quickjs_snapshot_payload from state/history responses instead of dropping mode=thread

* fix: recurse strip into nested subgraph StateSnapshot

* fix: drop conditional-snapshot gate that leaked repl slots

* fix: assert LangGraph state-shape invariants at import time

* refactor: discover graphs to filter from langgraph.json

* test: assert copy() preserves subclass; iterate langgraph.json for subagent coverage

---------

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
2026-07-13 12:24:52 +00:00
Mani Saint-Victor 2b28c46caf fix(llm): make gpt-5.x usable through ccproxy Codex OAuth (#324)
* fix(llm): make gpt-5.x usable through ccproxy Codex OAuth

Two independent blockers made current OpenAI models fail when routed
through ccproxy's Codex OAuth endpoint:

1. ccproxy's default Codex model mappings rewrite any gpt-*/o1-*/o3-*/
   claude-* model to gpt-5.3-codex before forwarding, silently overriding
   the configured model and failing outright on accounts where
   gpt-5.3-codex is not served ("The 'gpt-5.3-codex' model is not
   supported when using Codex with a ChatGPT account").
   start_ccproxy() now generates a config with empty codex model
   mappings and passes it via 'ccproxy serve --config'.

2. ccproxy forwards the client's own User-Agent upstream and only
   gap-fills its Codex headers, so the backend gates current models on
   the client identity ("The '<model>' model requires a newer version
   of Codex"). get_chat_model() now sends Codex-CLI-shaped
   originator/version/User-Agent headers when the ccproxy Codex adapter
   is detected, overridable via EVOSCIENTIST_CODEX_CLIENT_VERSION.

Verified live: gpt-5.5 and gpt-5.4 complete successfully through
ccproxy Codex OAuth on a ChatGPT Plus account with both fixes; each
fails without them.

* fix(ccproxy): harden Codex client routing

* fix(llm): keep Codex client identity consistent

* docs: clarify Codex version floor

* style: ruff format models.py after merge

---------

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
Co-authored-by: X-iZhang <zacharyzhang2022@gmail.com>
2026-07-13 11:04:47 +00:00
Mani Saint-Victor da6ca38d53 fix(llm): respect reasoning_effort setting on native OpenAI path (#321)
* fix(llm): respect reasoning_effort setting on native OpenAI path

The native OpenAI provider path hardcoded reasoning effort to xhigh for
gpt-5.4/5.5/codex models, silently ignoring the user's reasoning_effort
config setting. The OpenRouter path already honors the
EVOSCIENTIST_REASONING_EFFORT env var that settings.py exports from that
setting; this applies the same lookup on the native path, falling back
to the previous defaults when unset.

Adds a regression test and isolates the existing xhigh test from the
env var.

* fix(llm): preserve model reasoning defaults

* fix(llm): preserve GPT-5.6 reasoning default

---------

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
2026-07-13 10:42:55 +00:00
Mani Saint-Victor 19888c2db6 fix(ccproxy): raise startup timeouts (auth check 10s→30s, serve health 30s→180s) (#328)
* fix(ccproxy): raise auth status check timeout to 30s

ccproxy's CLI initializes its full plugin system on every invocation;
a cold 'ccproxy auth status' takes ~10s wall time on Apple Silicon,
so the 10s subprocess timeout made OAuth startup fail intermittently
with 'Auth check timed out' even when credentials were valid.

* fix(ccproxy): raise serve health deadline to 120s

ccproxy boot includes plugin init plus Codex CLI detection; measured
~76s to first healthy response on an Apple Silicon Mac (ccproxy-api
0.2.9). The 30s deadline in start_ccproxy() killed the process before
it could come up, failing OAuth startup with 'ccproxy did not become
healthy within 30 seconds'.

* fix(ccproxy): widen serve health deadline to 180s

Full startup measured at ~111s on a second cold run (Apple Silicon,
ccproxy-api 0.2.9); 120s left too little headroom for boot variance.

* fix(ccproxy): centralize startup timeouts

---------

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
2026-07-13 10:36:05 +00:00
Mani Saint-Victor f3e65a446f fix(tests): isolate the developer's real .env from the test suite (#329)
get_effective_config() runs load_dotenv(find_dotenv(usecwd=True),
override=True), so any test that loads config injected the repo's real
.env into os.environ for the rest of the pytest process. An
empty-valued line like MINIMAX_BASE_URL= then made
os.environ.get(key, default) return '' instead of the default,
failing the MiniMax routing tests in full-suite runs while they
passed in isolation.

Generalizes the find_dotenv redirect that test_config.py's
temp_config_dir fixture already applied locally into a suite-wide
autouse fixture, pointing at a never-created path so tests writing
their own tmp_path/.env cannot collide with it. Adds a regression
test reproducing the leak.

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
2026-07-13 10:27:40 +00:00
Zixin Dong f72f7b93d5 feat(llm): add OpenRouter app attribution headers (#339) (#344)
* feat(llm): add OpenRouter app attribution headers (#339)

Attach EvoScientist app-attribution at the shared model-init layer so all
OpenRouter calls are credited to the project. langchain-openrouter maps
app_url/app_title/app_categories -> HTTP-Referer / X-Title /
X-OpenRouter-Categories. Applied only for the openrouter provider, via
setdefault so explicit caller kwargs win. Configurable through new
openrouter_http_referer / openrouter_app_title / openrouter_app_categories
settings and their EVOSCIENTIST_OPENROUTER_* env vars.

Closes #339

* refactor(llm): centralize OpenRouter attribution defaults + cap categories

Address PR #344 review:
- Define the app-attribution default constants once in config/settings.py
  (the config fields and llm/models.py both use them) instead of duplicating
  the literals across the two modules.
- Reduce the default categories to creative-writing,personal-agent and cap the
  sent list to OpenRouter's 2-per-request limit, warning when a configured list
  exceeds it, so extras are dropped predictably (and surfaced) here rather than
  being silently truncated server-side.

---------

Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
2026-07-13 11:08:59 +01:00
dinos 49770949da fix(langgraph): prefer executable in venv over path (#341) 2026-07-11 11:59:27 +00:00
X-iZhang 6f10406d5b chore: update version to v0.2.2
Docker / build (push) Has been cancelled
2026-07-11 01:16:03 +01:00
X-iZhang 9042068094 feat(models): add support for GPT-5.6 variants and update context windows for Grok models 2026-07-11 00:51:06 +01:00
dependabot[bot] 81ff0519dc chore(deps): bump soupsieve in the uv group across 1 directory (#347)
Bumps the uv group with 1 update in the / directory: [soupsieve](https://github.com/facelessuser/soupsieve).


Updates `soupsieve` from 2.8.3 to 2.8.4
- [Release notes](https://github.com/facelessuser/soupsieve/releases)
- [Commits](https://github.com/facelessuser/soupsieve/compare/2.8.3...2.8.4)

---
updated-dependencies:
- dependency-name: soupsieve
  dependency-version: 2.8.4
  dependency-type: indirect
  dependency-group: uv
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-07-10 22:47:18 +01:00
dinos 690b903f85 test: standardize async tests on pytest-asyncio auto mode (#338)
* chore: add pytest-asyncio in auto mode

* test: migrate channel and stream tests to native async

Convert run_async() wrapper tests to plain 'async def test_*' under
pytest-asyncio auto mode. collect_events() in stream_v3_fakes becomes a
coroutine awaited at every call site.

* test: migrate command and model/middleware tests to native async

Convert run_async() wrappers (import, alias, and fixture forms) to plain
'async def test_*'. Multi-call tests merge onto one loop as sequential
awaits; none asserted on loop identity.

* test: migrate TUI, notifier, gateway, and session tests to native async

TUI/notifier/gateway files convert run_async wrappers to plain async
tests. test_sessions.py's unittest.TestCase classes move to
unittest.IsolatedAsyncioTestCase (pytest-asyncio does not await async
methods on plain TestCase; converting blindly would have made ~70 tests
silently vacuous). Its setUpClass keeps a one-shot asyncio.run() since
IsolatedAsyncioTestCase has no async class-level hook. TestLoadingWidget
in test_tui_widgets.py drops its TestCase base for the same reason.

* test: replace direct asyncio.run() calls with native async tests

Convert tests that called asyncio.run() (directly or via a local _run
helper) to plain 'async def test_*'; delete the local helpers.

* test: drop undeclared anyio markers and delete run_async helper

The @pytest.mark.anyio tests relied on anyio being a transitive dep of
httpx; auto-mode pytest-asyncio collects them natively. run_async() and
its fixture are unreferenced after the migration, so remove them —
pytest-asyncio's per-test loop teardown covers the pending-task
cancellation the helper existed for (verified: full suite runs with no
'Event loop is closed' errors or destroyed-task warnings).
2026-07-08 18:37:48 +00:00
dinos d2452c54d5 Refactor onboarding OAuth flow for auxiliary models (#337)
* refactor(onboard): shared flow for ccproxy providers

* feat(onboard): support oauth configuration for auxiliary models

* fix(onboard): reuse main model auth for same-provider auxiliary

* fix(onboard): reconcile oauth providers
2026-07-08 18:28:44 +00:00
dinos a7b9e175c1 fix(config): set config.yaml permissions to 0x600 (#336) 2026-07-08 19:25:01 +01:00
dinos be3dd272c3 test: deflake timing-dependent tests (#335)
* test: deflake timing-dependent tests

Inject a clock into channel dedup tests, replace fixed async sleeps with
events/explicit flushes, and avoid wall-clock waits in background tests.

* coderabbit nit
2026-07-07 08:25:35 +01:00
113 changed files with 10279 additions and 3914 deletions
+1 -1
View File
@@ -5,5 +5,5 @@
<rect x="54" y="5" width="62" height="24" rx="6" fill="#1565c0"/>
<text x="85" y="22" text-anchor="middle"
font-family="Inter, -apple-system, system-ui, sans-serif"
font-size="13" font-weight="700" fill="#ffffff">v0.2.1</text>
font-size="13" font-weight="700" fill="#ffffff">v0.2.2</text>
</svg>

Before

Width:  |  Height:  |  Size: 555 B

After

Width:  |  Height:  |  Size: 555 B

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

Before

Width:  |  Height:  |  Size: 555 B

After

Width:  |  Height:  |  Size: 555 B

Binary file not shown.

Before

Width:  |  Height:  |  Size: 286 KiB

After

Width:  |  Height:  |  Size: 287 KiB

+136 -25
View File
@@ -19,6 +19,7 @@ Usage:
import json
import logging
import os
from collections.abc import Sequence
from pathlib import Path
from typing import TYPE_CHECKING
@@ -304,8 +305,12 @@ 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,
@@ -314,6 +319,16 @@ 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()
@@ -333,6 +348,16 @@ 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``).
@@ -641,9 +666,14 @@ 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.
@@ -665,10 +695,14 @@ 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,
@@ -681,10 +715,20 @@ 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(_paths_mod.MEMORIES_DIR)
memory_dir = str(memory_dir or _paths_mod.MEMORIES_DIR)
source_type = (
MemorySourceType.SUBAGENT if for_async_subagent else MemorySourceType.TURN
)
@@ -699,18 +743,20 @@ 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_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(
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(
MemoryObservationTarget.AGENT
),
memory_scheduler=memory_scheduler,
)
"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)
# 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
@@ -728,16 +774,32 @@ 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(),
*([ModelFallbackMiddleware()] if enable_legacy_model_fallback else []),
RepetitiveToolCallGuardMiddleware(
threshold=repetitive_tool_call_threshold,
max_consecutive_errors=max_consecutive_tool_errors,
),
ContextOverflowMapperMiddleware(),
ToolErrorHandlerMiddleware(),
*create_tool_selector_middleware(
model=tool_selector_model,
track_stream_selection=not for_async_subagent,
),
*selector_middlewares,
ToolProtocolGuardMiddleware(),
# Interpreter prompt must land before runtime/memory context, so this
# middleware sits ahead of runtime_context in the stack.
create_code_interpreter_middleware(
@@ -770,7 +832,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:
if not for_async_subagent and enable_background_execution:
from .middleware.background import BackgroundExecutionMiddleware
mw.append(BackgroundExecutionMiddleware())
@@ -868,6 +930,14 @@ 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.
@@ -894,6 +964,22 @@ 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
@@ -935,19 +1021,21 @@ 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(_paths.MEMORIES_DIR)
_mem_dir = str(memory_dir or _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 = CustomSandboxBackend(
root_dir=workspace_dir,
virtual_mode=True,
timeout=cfg.sandbox_execute_timeout,
dangerous=cfg.dangerous_mode,
)
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,
)
sk_backend = MergedSkillsBackend(
primary_dir=_usr_skills_dir,
global_dir=_global_skills_dir,
@@ -969,8 +1057,29 @@ 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, cfg=cfg, chat_model=chat_model
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,
)
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
@@ -995,6 +1104,8 @@ def create_cli_agent(
chat_model=chat_model,
workspace_dir=workspace_dir,
)
if not enable_subagents:
kwargs = {**kwargs, "subagents": []}
return create_deep_agent(
**kwargs,
+2
View File
@@ -9,6 +9,8 @@ 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"),
+57 -6
View File
@@ -20,6 +20,9 @@ from EvoScientist.config import EvoScientistConfig
logger = logging.getLogger(__name__)
_CCPROXY_AUTH_TIMEOUT_SECONDS = 30
_CCPROXY_HEALTH_TIMEOUT_SECONDS = 180
# =============================================================================
# Availability & auth checks
@@ -127,7 +130,11 @@ def check_ccproxy_auth(provider: str = "claude_api") -> tuple[bool, str]:
[exe, "auth", "status", provider],
capture_output=True,
text=True,
timeout=10,
# 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,
)
import re as _re
@@ -176,6 +183,33 @@ 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.
@@ -186,18 +220,32 @@ def start_ccproxy(port: int) -> subprocess.Popen:
The Popen handle for the ccproxy process.
Raises:
RuntimeError: If ccproxy fails to become healthy within 30 seconds.
RuntimeError: If ccproxy fails to become healthy within
``_CCPROXY_HEALTH_TIMEOUT_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(
[exe, "serve", "--port", str(port)],
cmd,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
)
# Wait for health (ccproxy can take up to ~11s on first start)
deadline = time.monotonic() + 30
deadline = time.monotonic() + _CCPROXY_HEALTH_TIMEOUT_SECONDS
while time.monotonic() < deadline:
if proc.poll() is not None:
raise RuntimeError(
@@ -213,7 +261,10 @@ def start_ccproxy(port: int) -> subprocess.Popen:
proc.wait(timeout=3)
except subprocess.TimeoutExpired:
proc.kill()
raise RuntimeError("ccproxy did not become healthy within 30 seconds")
raise RuntimeError(
"ccproxy did not become healthy within "
f"{_CCPROXY_HEALTH_TIMEOUT_SECONDS} seconds"
)
def stop_ccproxy(proc: subprocess.Popen | None) -> None:
+10 -5
View File
@@ -75,11 +75,13 @@ 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 ──────────────────────────────────────────────────
@@ -93,15 +95,16 @@ class DedupCache:
if not msg_id:
return False
self._prune()
now = self._clock()
self._prune(now)
if msg_id in self._seen:
# LRU: refresh position and timestamp
self._seen.move_to_end(msg_id)
self._seen[msg_id] = time.monotonic()
self._seen[msg_id] = now
return True
self._seen[msg_id] = time.monotonic()
self._seen[msg_id] = now
if len(self._seen) > self._max:
while len(self._seen) > self._trim:
self._seen.popitem(last=False)
@@ -118,9 +121,9 @@ class DedupCache:
# ── internal ────────────────────────────────────────────────────
def _prune(self) -> None:
def _prune(self, now: float | None = None) -> None:
"""Remove entries older than *ttl_seconds*."""
cutoff = time.monotonic() - self._ttl
cutoff = (self._clock() if now is None else now) - self._ttl
# OrderedDict is insertion-ordered; oldest entries are first.
while self._seen:
_key, ts = next(iter(self._seen.items()))
@@ -428,11 +431,13 @@ 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(
+66 -109
View File
@@ -371,14 +371,31 @@ def _step_minimax_region(config: EvoScientistConfig) -> str:
return _MINIMAX_REGIONS[region]
def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str:
"""Step 2a: Select Anthropic authentication mode (API key vs OAuth).
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.
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", "oauth", or "auto".
Selected auth mode: "api_key" or "oauth".
"""
from ...ccproxy_manager import check_ccproxy_auth, is_ccproxy_available
@@ -386,10 +403,14 @@ def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str:
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="API Key (direct Anthropic access)", value="api_key"),
Choice(title=f"API Key (direct {provider_label} access)", value="api_key"),
Choice(
title="Claude Code OAuth (via ccproxy — no API key needed)"
title=f"{oauth_label} (via ccproxy — no API key needed)"
+ (
""
if ccproxy_available
@@ -401,12 +422,12 @@ def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str:
Choice(title="← Back (re-pick provider)", value=BACK_SENTINEL),
]
current = config.anthropic_auth_mode
current = getattr(config, config_attr)
if current not in ("api_key", "oauth"):
current = "api_key"
question = questionary.select(
"Authentication mode [Esc/← to go back]:",
f"{auth_question_label} [Esc/← to go back]:",
choices=choices,
default=current,
style=WIZARD_STYLE,
@@ -448,11 +469,9 @@ def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str:
if auth_mode == "oauth":
_prompt_ccproxy_port(config)
# If OAuth selected, check auth status and offer login
if auth_mode in ("oauth", "auto"):
authed, msg = check_ccproxy_auth()
authed, msg = check_ccproxy_auth(ccproxy_provider)
if authed:
console.print(f" [green]✓ OAuth: {msg}[/green]")
console.print(f" [green]✓ {auth_status_label}: {msg}[/green]")
relogin = questionary.confirm(
"Re-authenticate to refresh credentials?",
default=False,
@@ -462,11 +481,13 @@ def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str:
if relogin is None:
raise KeyboardInterrupt()
if relogin:
_run_ccproxy_login("claude_api", "OAuth")
_run_ccproxy_login(ccproxy_provider, auth_status_label)
else:
console.print(f" [yellow]OAuth not authenticated: {msg}[/yellow]")
console.print(
f" [yellow]{auth_status_label} not authenticated: {msg}[/yellow]"
)
login = questionary.confirm(
"Log in to Claude now?",
f"Log in to {prompt_login_label} now?",
default=True,
style=CONFIRM_STYLE,
qmark=QMARK,
@@ -474,11 +495,32 @@ def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str:
if login is None:
raise KeyboardInterrupt()
if login:
_run_ccproxy_login("claude_api", "OAuth")
_run_ccproxy_login(ccproxy_provider, auth_status_label)
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).
@@ -488,101 +530,16 @@ def _step_openai_auth_mode(config: EvoScientistConfig) -> str:
Returns:
Selected auth mode: "api_key" or "oauth".
"""
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,
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",
)
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(
+258 -181
View File
@@ -129,6 +129,12 @@ _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.
@@ -142,6 +148,201 @@ 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"),
@@ -479,102 +680,17 @@ def run_onboard(
provider = _step_provider(config)
config.provider = provider
# 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"
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"
),
)
except GoBack:
# User picked "← Back" — restore config to its state at the
# top of this iteration (drops any base_url / region /
@@ -594,60 +710,9 @@ def run_onboard(
ollama_detected_models = []
console.print(" [dim]↩ Returning to provider selection.[/dim]")
continue
break # auth_mode succeeded — exit sub-loop
break # Provider setup succeeded — exit sub-loop
# 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")
_reconcile_oauth_modes(config)
_autosave(config)
else:
# Provider section skipped — keep prior provider value to drive
@@ -680,44 +745,55 @@ def run_onboard(
"kept current" if config.auxiliary_model else "not set",
)
elif _step_auxiliary_enable(config):
# 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(
from .prompter import GoBack
aux_ollama_detected_models: list[str] = []
while True:
loop_snapshot = copy.deepcopy(config)
aux_provider = _step_provider(
config,
current_value=config.custom_openai_base_url
or os.environ.get("CUSTOM_OPENAI_BASE_URL", ""),
label="co-pilot",
default_value=config.auxiliary_provider,
)
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_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
config.auxiliary_model = _step_model(
config,
aux_provider,
ollama_detected_models=aux_ollama_detected_models,
label="co-pilot",
default_value=config.auxiliary_model,
)
@@ -725,6 +801,7 @@ 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:
+87 -3
View File
@@ -106,11 +106,23 @@ def _normalize_hhmm(value: Any) -> str | None:
def get_config_dir() -> Path:
"""Get the configuration directory path.
Uses XDG_CONFIG_HOME if set, otherwise ~/.config/evoscientist/
Priority:
1. EVOSCIENTIST_CONFIG_DIR
2. EVOSCIENTIST_HOME/config
3. XDG_CONFIG_HOME/evoscientist
4. ~/.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) / "evoscientist"
return Path(xdg_config).expanduser() / "evoscientist"
return Path.home() / ".config" / "evoscientist"
@@ -123,6 +135,16 @@ 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:
@@ -240,6 +262,13 @@ 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
@@ -278,10 +307,21 @@ 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"
reasoning_effort: str = "high"
# Empty means use the provider/model default. A non-empty value is an
# explicit user override exported as EVOSCIENTIST_REASONING_EFFORT.
reasoning_effort: str = ""
# 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)
@@ -440,6 +480,14 @@ 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__
@@ -548,6 +596,10 @@ 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)
@@ -560,6 +612,10 @@ def save_config(config: EvoScientistConfig) -> None:
sort_keys=False,
allow_unicode=True,
)
try:
config_path.chmod(0o600)
except OSError:
pass
def reset_config() -> None:
@@ -694,6 +750,11 @@ 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:
@@ -753,6 +814,9 @@ _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",
@@ -769,6 +833,10 @@ _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",
@@ -884,6 +952,22 @@ 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"
):
+1 -1
View File
@@ -43,7 +43,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
+237 -1
View File
@@ -5,8 +5,244 @@ 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 EvoScientist.EvoScientist import EvoScientist_agent
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()
__all__ = ["EvoScientist_agent"]
+11 -5
View File
@@ -306,14 +306,19 @@ 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
@@ -709,6 +714,7 @@ 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,
+6 -2
View File
@@ -20,9 +20,9 @@ _KNOWN_MODEL_CONTEXT_WINDOWS: dict[str, int] = {
# Qwen 3.7 closed-source tiers — Max flagship and Plus (1M).
"qwen3.7-max": 1_000_000,
"qwen3.7-plus": 1_000_000,
# xAI Grok — per-model windows (build-0.1: 256K, 4.3: 1M).
# xAI Grok — per-model windows (build-0.1: 256K, 4.5: 500K).
"grok-build-0.1": 256_000,
"grok-4.3": 1_000_000,
"grok-4.5": 500_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,6 +32,8 @@ _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.
@@ -40,6 +42,8 @@ _KNOWN_MODEL_CONTEXT_WINDOWS: dict[str, int] = {
_KNOWN_MODEL_FAMILIES: list[tuple[str, int]] = [
# All Claude — 1M via the ``context-1m-2025-08-07`` beta header.
("claude-", 1_000_000),
# OpenAI GPT-5.6 family — sol, terra, luna variants
("gpt-5.6", 1_050_000),
# OpenAI GPT-5.5 family — base, pro, future variants
("gpt-5.5", 1_050_000),
# Google Gemini 3.x family — flash, flash-lite, pro (1.05M). Excludes 2.5.
+386
View File
@@ -0,0 +1,386 @@
"""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
+245 -29
View File
@@ -10,11 +10,19 @@ 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,
@@ -37,6 +45,50 @@ _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]] = {
@@ -68,6 +120,19 @@ _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]] = [
@@ -89,6 +154,9 @@ _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"),
@@ -145,6 +213,9 @@ _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"),
@@ -159,7 +230,8 @@ _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.3", "x-ai/grok-4.3", "openrouter"),
("grok-4.5", "x-ai/grok-4.5", "openrouter"),
("hy3", "tencent/hy3", "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"),
@@ -257,6 +329,15 @@ 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(
@@ -314,8 +395,16 @@ 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 "thinking" not in kwargs:
if provider == "anthropic" and not disable_thinking 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'
@@ -334,24 +423,31 @@ 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 "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"
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
)
kwargs["reasoning"] = {"effort": _eff, "summary": "auto"}
else "high"
)
_eff = _resolve_reasoning_effort(_default_effort)
kwargs["reasoning"] = {"effort": _eff, "summary": "auto"}
# Google GenAI: surface thinking traces
if provider == "google-genai":
if provider == "google-genai" and not disable_reasoning:
kwargs.setdefault("include_thoughts", True)
# Ollama: separate reasoning content from response for thinking models
if provider == "ollama" and "reasoning" not in kwargs:
if provider == "ollama" and not disable_reasoning and "reasoning" not in kwargs:
kwargs["reasoning"] = True
@@ -378,7 +474,46 @@ def get_chat_model(
>>> model = get_chat_model("gpt-4o") # OpenAI model
>>> model = get_chat_model("claude-3-opus-20240229", provider="anthropic") # Full ID
"""
model = model or DEFAULT_MODEL
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
# Look up short name in registry (provider-aware)
model_id = None
@@ -413,22 +548,35 @@ 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 = None
_original_provider: str | None = (
runtime_provider_name if runtime_provider_name != provider else None
)
if provider == "anthropic":
base_url = os.environ.get("ANTHROPIC_BASE_URL", "")
if base_url:
kwargs["base_url"] = base_url
kwargs.setdefault("base_url", base_url)
api_key = os.environ.get("ANTHROPIC_API_KEY", "")
if api_key:
kwargs["api_key"] = api_key
kwargs.setdefault("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["base_url"] = base_url
_is_openai_proxy = _is_ccproxy_codex()
kwargs.setdefault("base_url", base_url)
_is_openai_proxy = _is_ccproxy_codex(
kwargs.get("base_url"), kwargs.get("api_key")
)
if _is_openai_proxy:
# Use Responses API for ccproxy: bypasses the format chain
# converter (Chat→Responses→Chat) which returns 502 on
@@ -441,9 +589,23 @@ 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["api_key"] = api_key
kwargs.setdefault("api_key", api_key)
# OpenAI-routed providers → route through OpenAI provider with base_url
elif provider in _OPENAI_ROUTED_PROVIDERS:
@@ -461,10 +623,10 @@ def get_chat_model(
else:
base_url = base_url_default
if base_url:
kwargs["base_url"] = base_url
kwargs.setdefault("base_url", base_url)
api_key = os.environ.get(api_key_env, "")
if api_key:
kwargs["api_key"] = api_key
kwargs.setdefault("api_key", api_key)
# SiliconFlow: disable thinking — LangChain drops reasoning_content
# from history, causing error 20015 on multi-turn requests.
if provider == "siliconflow":
@@ -481,15 +643,61 @@ def get_chat_model(
_is_third_party = True
api_key = os.environ.get("OPENROUTER_API_KEY", "")
if api_key:
kwargs["api_key"] = api_key
kwargs.setdefault("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 = os.environ.get("EVOSCIENTIST_REASONING_EFFORT", "").strip() or "high"
effort = _resolve_reasoning_effort("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
@@ -510,10 +718,10 @@ def get_chat_model(
else:
base_url = base_url_default
if base_url:
kwargs["base_url"] = base_url
kwargs.setdefault("base_url", base_url)
api_key = os.environ.get(api_key_env, "")
if api_key:
kwargs["api_key"] = api_key
kwargs.setdefault("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"
@@ -522,8 +730,9 @@ def get_chat_model(
elif provider == "ollama":
base_url = os.environ.get("OLLAMA_BASE_URL", "")
if base_url:
kwargs["base_url"] = base_url
kwargs.setdefault("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)
@@ -540,7 +749,14 @@ def get_chat_model(
elif _responses_api_setting == "true":
kwargs["use_responses_api"] = True
chat_model = init_chat_model(model=model_id, model_provider=provider, **kwargs)
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
# Flatten list content to strings for strict OpenAI-compatible providers
# (DeepSeek, SiliconFlow, OpenRouter, custom-openai, etc.) and
+367 -3
View File
@@ -25,6 +25,7 @@ Utilities:
from __future__ import annotations
import hashlib
import os
from typing import Any
@@ -178,15 +179,20 @@ _patch_ccproxy_codex_compat()
# ---------------------------------------------------------------------------
# Utility: detect ccproxy's Codex adapter (as opposed to generic localhost).
# ---------------------------------------------------------------------------
def _is_ccproxy_codex() -> bool:
def _is_ccproxy_codex(
base_url: str | None = None,
api_key: str | None = None,
) -> 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.
"""
base_url = os.environ.get("OPENAI_BASE_URL", "")
api_key = os.environ.get("OPENAI_API_KEY", "")
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", "")
return (
("127.0.0.1" in base_url or "localhost" in base_url)
and api_key == "ccproxy-oauth"
@@ -267,6 +273,299 @@ 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.
@@ -282,6 +581,9 @@ 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
@@ -320,6 +622,8 @@ 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
@@ -729,6 +1033,66 @@ def _patch_openai_capture_reasoning_content() -> None:
_patch_openai_capture_reasoning_content()
# ---------------------------------------------------------------------------
# Patch (module-level): silence langgraph_api's OpenAPI schema-generation
# warnings for endpoints whose docstrings aren't valid YAML.
#
# Upstream ``langgraph_api.utils.SchemaGenerator.get_schema`` calls
# ``parse_docstring`` (inherited from Starlette's ``BaseSchemaGenerator``)
# on every registered endpoint. When the docstring is prose with stray
# ``:`` characters, ``yaml.safe_load`` raises and upstream logs the
# failure + full traceback at WARNING level. It then falls back to
# ``{"description": docstring}`` — the endpoint still ends up in the
# schema with its prose as the description, just without structured
# ``parameters``/``responses``/``tags`` fields.
#
# The fallback path is fine; the warning + traceback is just noise. And
# it's only triggered for our deploy because mounting any custom Starlette
# app (``EvoScientist/langgraph_dev/http.py``) makes upstream call
# ``update_openapi_spec`` at startup — which iterates EVERY route,
# including upstream's own endpoints whose prose docstrings predate the
# YAML convention.
#
# Fix: wrap ``parse_docstring`` itself and absorb ``yaml.YAMLError`` by
# returning the same fallback shape upstream's except branch produces.
# Non-YAML exceptions are deliberately left to propagate — upstream's
# ``get_schema`` already catches them and logs WARNING + traceback, so
# unexpected failures remain debuggable. Patching ``parse_docstring`` (a
# small, stable method) instead of ``get_schema`` (the larger loop body)
# minimizes our exposure to upstream churn.
# ---------------------------------------------------------------------------
_langgraph_schema_silenced_patched = False
def _patch_langgraph_schema_generator_silence_warnings() -> None:
global _langgraph_schema_silenced_patched
if _langgraph_schema_silenced_patched:
return
try:
import langgraph_api.utils as _lgapi_utils
import yaml
_SchemaGenerator = _lgapi_utils.SchemaGenerator
_orig_parse_docstring = _SchemaGenerator.parse_docstring
def _patched_parse_docstring(self: Any, func: Any) -> dict[str, Any]:
try:
return _orig_parse_docstring(self, func)
except yaml.YAMLError:
return {"description": getattr(func, "__doc__", None) or ""}
_SchemaGenerator.parse_docstring = _patched_parse_docstring
_langgraph_schema_silenced_patched = True
except Exception:
# Patches are loader-safe: never crash the import. Silent failure
# here just leaves the upstream warnings visible in deploy logs,
# which is a benign fallback.
pass
_patch_langgraph_schema_generator_silence_warnings()
# ---------------------------------------------------------------------------
# Patch (lazy, OpenRouter only): strip OpenAI-Responses encrypted reasoning
# items from outgoing assistant messages.
+302
View File
@@ -0,0 +1,302 @@
"""Shared logging configuration helpers."""
from __future__ import annotations
import logging
import os
import sys
from datetime import UTC, datetime
from pathlib import Path
from typing import Any, TextIO
DEFAULT_LOG_RETENTION_DAYS = 30
DEFAULT_LOG_FORMAT = "%(asctime)s [%(levelname)s] %(name)s: %(message)s"
DEFAULT_LOG_DATE_FORMAT = "%Y-%m-%d %H:%M:%S"
MANAGED_HANDLER_ATTR = "_evoscientist_managed_handler"
def resolve_log_level(level: int | str | None, default: int = logging.INFO) -> int:
"""Resolve a logging level from config or environment input."""
if isinstance(level, int):
return level
raw = str(level or "").strip()
if not raw:
return default
if raw.isdigit():
return int(raw)
normalized = raw.upper()
if normalized == "WARN":
normalized = "WARNING"
resolved = logging.getLevelNamesMapping().get(normalized)
return resolved if isinstance(resolved, int) else default
def _mark_managed(handler: logging.Handler, kind: str) -> logging.Handler:
setattr(handler, MANAGED_HANDLER_ATTR, kind)
return handler
def _managed_kind(handler: logging.Handler) -> str | None:
kind = getattr(handler, MANAGED_HANDLER_ATTR, None)
return kind if isinstance(kind, str) else None
def remove_managed_handlers(
logger: logging.Logger | None = None,
*,
kinds: set[str] | None = None,
) -> None:
"""Remove handlers installed by this module without touching external ones."""
target = logger or logging.getLogger()
for handler in target.handlers[:]:
kind = _managed_kind(handler)
if kind and (kinds is None or kind in kinds):
target.removeHandler(handler)
handler.close()
def _standard_formatter() -> logging.Formatter:
return logging.Formatter(DEFAULT_LOG_FORMAT, datefmt=DEFAULT_LOG_DATE_FORMAT)
class DailyLogFileHandler(logging.FileHandler):
"""File handler that writes the active log to a date-based filename."""
def __init__(
self,
log_dir: str | Path,
*,
prefix: str = "evoscientist",
retention_days: int = DEFAULT_LOG_RETENTION_DAYS,
encoding: str = "utf-8",
utc: bool = False,
) -> None:
self.log_dir = Path(log_dir).expanduser()
self.prefix = prefix
self.retention_days = max(1, retention_days)
self.utc = utc
self.log_dir.mkdir(parents=True, exist_ok=True)
super().__init__(self._dated_log_path(), encoding=encoding, delay=True)
@property
def active_log_path(self) -> Path:
"""Return the active log path for the current date."""
return self._dated_log_path()
def _dated_log_path(self) -> Path:
now = datetime.now(UTC if self.utc else None)
return self.log_dir / f"{self.prefix}-{now:%Y-%m-%d}.log"
def emit(self, record: logging.LogRecord) -> None:
try:
expected = str(self.active_log_path)
if self.baseFilename != expected:
if self.stream:
self.stream.close()
self.stream = None
self.baseFilename = expected
self._delete_expired_logs()
super().emit(record)
except OSError:
self.handleError(record)
def getFilesToDelete(self) -> list[str]:
candidates = sorted(self.log_dir.glob(f"{self.prefix}-????-??-??.log"))
if len(candidates) <= self.retention_days:
return []
return [str(path) for path in candidates[: -self.retention_days]]
def _delete_expired_logs(self) -> None:
for path in self.getFilesToDelete():
try:
os.remove(path)
except OSError:
pass
def default_log_dir() -> Path:
"""Return the default runtime log directory."""
env_dir = os.environ.get("EVOSCIENTIST_LOG_DIR")
if env_dir:
return Path(env_dir).expanduser()
from EvoScientist.paths import DATA_DIR
return DATA_DIR / "logs"
def configure_daily_file_logging(
logger: logging.Logger | None = None,
*,
log_dir: str | Path | None = None,
prefix: str = "evoscientist",
level: int | str = logging.INFO,
retention_days: int = DEFAULT_LOG_RETENTION_DAYS,
) -> DailyLogFileHandler:
"""Attach a daily file handler, replacing older matching handlers."""
target = logger or logging.getLogger()
resolved_level = resolve_log_level(level, default=logging.INFO)
retention_days = max(1, int(retention_days))
resolved_dir = Path(log_dir).expanduser() if log_dir else default_log_dir()
for handler in target.handlers[:]:
if (
isinstance(handler, DailyLogFileHandler)
and handler.prefix == prefix
and handler.log_dir == resolved_dir
):
target.removeHandler(handler)
handler.close()
handler = DailyLogFileHandler(
resolved_dir,
prefix=prefix,
retention_days=retention_days,
)
_mark_managed(handler, "file")
handler.setLevel(resolved_level)
handler.setFormatter(_standard_formatter())
target.addHandler(handler)
if target.level == logging.NOTSET or target.level > resolved_level:
target.setLevel(resolved_level)
return handler
def configure_console_logging(
logger: logging.Logger | None = None,
*,
level: int | str | None = logging.INFO,
stream: TextIO | None = None,
replace: bool = True,
) -> logging.StreamHandler:
"""Attach a standard console handler for non-interactive entry points."""
target = logger or logging.getLogger()
resolved_level = resolve_log_level(level, default=logging.INFO)
if replace:
remove_managed_handlers(target, kinds={"console", "rich"})
handler = logging.StreamHandler(stream or sys.stderr)
_mark_managed(handler, "console")
handler.setLevel(resolved_level)
handler.setFormatter(_standard_formatter())
target.addHandler(handler)
target.setLevel(resolved_level)
return handler
def configure_rich_console_logging(
logger: logging.Logger | None = None,
*,
level: int | str | None = logging.INFO,
console: Any = None,
replace: bool = True,
dim_warnings: bool = False,
show_time: bool | None = None,
show_path: bool | None = None,
show_level: bool | None = None,
) -> logging.Handler:
"""Attach a Rich console handler for interactive CLI output."""
from rich.logging import RichHandler
from rich.markup import escape
target = logger or logging.getLogger()
resolved_level = resolve_log_level(level, default=logging.INFO)
verbose = resolved_level <= logging.DEBUG
if replace:
remove_managed_handlers(target, kinds={"console", "rich"})
class DimWarningHandler(RichHandler):
def emit(self, record: logging.LogRecord) -> None:
if dim_warnings and record.levelno == logging.WARNING and console is not None:
console.print(
"[dim yellow]\u26a0\ufe0f Warning:[/dim yellow] "
f"[dim]{escape(record.getMessage())}[/dim]"
)
return
super().emit(record)
handler = DimWarningHandler(
console=console,
show_time=verbose if show_time is None else show_time,
show_path=verbose if show_path is None else show_path,
show_level=verbose if show_level is None else show_level,
)
_mark_managed(handler, "rich")
handler.setLevel(resolved_level)
target.addHandler(handler)
target.setLevel(resolved_level)
return handler
def configure_logging(
logger: logging.Logger | None = None,
*,
level: int | str | None = logging.INFO,
log_dir: str | Path | None = None,
retention_days: int = DEFAULT_LOG_RETENTION_DAYS,
prefix: str = "evoscientist",
console: bool = True,
file: bool = True,
replace_managed: bool = True,
) -> list[logging.Handler]:
"""Configure standard EvoScientist console and daily file logging."""
target = logger or logging.getLogger()
resolved_level = resolve_log_level(level, default=logging.INFO)
if replace_managed:
remove_managed_handlers(target, kinds={"console", "rich", "file"})
handlers: list[logging.Handler] = []
if console:
handlers.append(
configure_console_logging(target, level=resolved_level, replace=False)
)
if file:
handlers.append(
configure_daily_file_logging(
target,
log_dir=log_dir,
prefix=prefix,
level=resolved_level,
retention_days=retention_days,
)
)
target.setLevel(resolved_level)
return handlers
def configure_logging_from_settings(
logger: logging.Logger | None = None,
*,
default_level: int = logging.INFO,
prefix: str = "evoscientist",
console: bool = True,
file: bool = True,
) -> list[logging.Handler]:
"""Configure logging from EvoScientist settings and environment overrides."""
level: int | str | None = os.environ.get("EVOSCIENTIST_LOG_LEVEL")
log_dir: str | Path | None = os.environ.get("EVOSCIENTIST_LOG_DIR") or None
retention_days = int(
os.environ.get("EVOSCIENTIST_LOG_RETENTION_DAYS", DEFAULT_LOG_RETENTION_DAYS)
)
try:
from EvoScientist.config import get_effective_config
cfg = get_effective_config()
level = level or getattr(cfg, "log_level", None)
log_dir = log_dir or getattr(cfg, "log_dir", None) or None
retention_days = int(
getattr(cfg, "log_retention_days", DEFAULT_LOG_RETENTION_DAYS)
)
except Exception:
level = level or default_level
return configure_logging(
logger,
level=resolve_log_level(level, default=default_level),
log_dir=log_dir,
retention_days=retention_days,
prefix=prefix,
console=console,
file=file,
)
+2
View File
@@ -10,6 +10,7 @@ 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,
@@ -38,6 +39,7 @@ __all__ = [
"find_server_by_name",
"get_all_tags",
"get_installed_names",
"get_mcp_server_errors",
"install_mcp_server",
"install_mcp_servers",
"load_mcp_config",
+16 -1
View File
@@ -114,6 +114,10 @@ _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.
@@ -764,6 +768,9 @@ 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:
@@ -787,10 +794,13 @@ 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.
@@ -798,7 +808,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, str(exc))
_report("error", name, detail)
return name, []
# ``return_exceptions=False`` is fine because ``_fetch`` already
@@ -807,6 +817,11 @@ 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,
*,
+17 -11
View File
@@ -428,6 +428,7 @@ 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(
@@ -439,18 +440,23 @@ def _memory_worker_middleware(
enable_observation_tool = memory_controls.observation_tool_enabled(
_memory_worker_observation_target(source_type)
)
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,
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,
),
excluded_tools=_MEMORY_WORKER_EXCLUDED_TOOLS,
)
]
def _build_memory_worker_agent(
@@ -71,6 +71,8 @@ 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,
@@ -85,5 +87,7 @@ def build_observation_linker_graph(
tools=tools,
memory_dir=agent_paths.memory_dir,
workspace_dir=agent_paths.workspace_dir,
middleware=memory_agent_middleware(),
# Outermost — normalize provider-SDK exceptions from the
# auxiliary model call before any inner middleware sees them.
middleware=[ErrorNormalizationMiddleware(), *memory_agent_middleware()],
)
+14
View File
@@ -18,6 +18,7 @@ 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,
@@ -28,29 +29,42 @@ 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",
+15 -1
View File
@@ -45,7 +45,21 @@ _MEMORY_FIRST_INTERPRETER_PROMPT = (
class EvoCodeInterpreterMiddleware(CodeInterpreterMiddleware):
"""Code interpreter middleware with EvoScientist's memory preflight hint."""
"""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.
"""
def _prepare_for_call(self, request: ModelRequest) -> str:
return super()._prepare_for_call(request) + _MEMORY_FIRST_INTERPRETER_PROMPT
@@ -0,0 +1,240 @@
"""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
+34 -3
View File
@@ -48,6 +48,8 @@ _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)."""
@@ -215,6 +217,9 @@ 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"
@@ -263,7 +268,15 @@ 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}) "
@@ -288,8 +301,9 @@ async def _try_fallbacks(
f"-- aborting fallback chain",
style="red",
)
raise
_raise_normalized(fb_request, fb_exc)
last_exc = fb_exc
last_failing_request = fb_request
_emit(
f" x {model_name} also failed: {type(fb_exc).__name__}: {fb_exc}",
style="red",
@@ -303,7 +317,24 @@ async def _try_fallbacks(
)
_emit(" All fallbacks exhausted -- re-raising last error", style="red")
raise last_exc
_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
def _guard_and_fallback(
@@ -330,7 +361,7 @@ def _guard_and_fallback(
f"Model error ({reason}) -- not eligible for fallback, re-raising",
style="red",
)
raise primary_exc
_raise_normalized(request, primary_exc)
return _try_fallbacks(request, invoke, primary_exc)
@@ -0,0 +1,350 @@
"""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))
@@ -0,0 +1,361 @@
"""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
+31 -3
View File
@@ -48,6 +48,7 @@ DEFAULT_ALWAYS_INCLUDE_TOOLS: frozenset[str] = frozenset(
"read_memory",
"record_observation",
"search_observations",
"write_todos",
}
)
@@ -132,10 +133,21 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
return self._build_selector(request).wrap_model_call(
request, _handler_after_selection
)
except Exception:
except Exception as exc:
if _handler_called:
raise # Error from downstream model — don't retry
# Selector itself failed (e.g., structured output not supported).
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.
logger.debug("Tool selector failed, using all tools", exc_info=True)
if self._track_stream_selection:
_selector_active = False
@@ -171,9 +183,16 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
return await self._build_selector(request).awrap_model_call(
request, _handler_after_selection
)
except Exception:
except Exception as exc:
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
@@ -258,6 +277,15 @@ def create_tool_selector_middleware(
model = _ensure_chat_model()
safe_model = disable_thinking(model)
safe_model = safe_model.model_copy(
update={
"tags": [*(safe_model.tags or []), "metering:tool_selector"],
"metadata": {
**(safe_model.metadata or {}),
"metering_scope": "tool_selector",
},
}
)
system_prompt = (
"You are selecting tools for a scientific research agent. "
+63
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import logging
import os
import shutil
from collections.abc import Iterator
from datetime import datetime
from pathlib import Path
@@ -215,3 +216,65 @@ def resolve_virtual_path(virtual_path: str) -> Path:
"""Resolve a virtual workspace path (e.g. /image.png) to a real filesystem path."""
vpath = virtual_path if virtual_path.startswith("/") else "/" + virtual_path
return (_active_workspace / vpath.lstrip("/")).resolve()
def evoscientist_root() -> Path:
"""Return the application root used by Gateway-managed runtime data."""
env_root = os.environ.get("EVOSCIENTIST_HOME")
if env_root:
return Path(env_root).expanduser().resolve()
return DATA_DIR.expanduser().resolve()
_EVOSCIENTIST_DATA_ROOT: Path | None = None
def _data_root() -> Path:
"""Return the root directory for isolated Web user workspaces."""
global _EVOSCIENTIST_DATA_ROOT
if _EVOSCIENTIST_DATA_ROOT is not None:
return _EVOSCIENTIST_DATA_ROOT
env_root = os.environ.get("EVOSCIENTIST_DATA_ROOT")
if env_root:
root = Path(env_root).expanduser().resolve()
else:
root = evoscientist_root() / "data"
_EVOSCIENTIST_DATA_ROOT = root
return root
def user_data_dir(user_id: str) -> Path:
"""Return and create the isolated data directory for a Web user."""
path = _data_root() / user_id
path.mkdir(parents=True, exist_ok=True)
return path
def iter_user_data_dirs() -> Iterator[Path]:
"""Yield existing Web user directories without creating the data root."""
root = _data_root()
if not root.exists():
return
for path in root.iterdir():
if path.is_dir():
yield path
def thread_data_dir(user_id: str, thread_id: str) -> Path:
"""Return and create a user's isolated thread workspace."""
path = user_data_dir(user_id) / thread_id
path.mkdir(parents=True, exist_ok=True)
return path
def global_data_dir(user_id: str) -> Path:
"""Return and create a user's directory shared across all threads."""
path = user_data_dir(user_id) / "__global__"
path.mkdir(parents=True, exist_ok=True)
return path
def uploads_dir() -> Path:
"""Return the Gateway upload staging directory."""
return evoscientist_root() / "uploads"
+116
View File
@@ -0,0 +1,116 @@
"""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()
+32 -2
View File
@@ -7,6 +7,15 @@ 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:
@@ -158,6 +167,14 @@ 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,
@@ -213,6 +230,19 @@ class StreamEventEmitter:
)
@staticmethod
def error(message: str) -> StreamEvent:
def error(
message: str,
*,
code: str | None = None,
recoverable: bool | None = None,
details: dict[str, Any] | None = None,
) -> StreamEvent:
"""Error event."""
return StreamEvent("error", {"type": "error", "message": message})
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)
+284 -45
View File
@@ -8,13 +8,15 @@ 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
from langgraph.types import Command, Interrupt, Overwrite
from ..memory.worker_activity import clear_completed_memory_activity_counts
from .emitter import StreamEventEmitter
@@ -43,6 +45,12 @@ 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:
@@ -87,9 +95,10 @@ 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.
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).
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.
Critically, this only runs when the stuck state is *not* a legitimate
human-in-the-loop interrupt. The agent pauses via ``interrupt()`` /
@@ -106,10 +115,9 @@ async def _clear_interrupted_graph_state(
_log = logging.getLogger(__name__)
try:
snapshot = await agent.aget_state(config)
# Only act when the graph is genuinely stuck (non-empty next tuple)...
if not snapshot or not getattr(snapshot, "next", None):
if not snapshot:
return
# ...and not parked at a real human-in-the-loop interrupt.
# Never alter a real human-in-the-loop pause.
if _snapshot_has_pending_interrupt(snapshot):
_log.debug(
"Leaving interrupted graph state intact for thread %s: "
@@ -119,6 +127,13 @@ 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(
@@ -134,6 +149,49 @@ 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, ...]
@@ -209,7 +267,12 @@ 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
@@ -237,13 +300,26 @@ 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))
return self._process_update_event(
_event_data(event), namespace=namespace, source="update"
)
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}))
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"
)
)
if self._process_value_message_snapshots and not namespace:
events.extend(self._process_value_messages(_event_data(event)))
return events
@@ -382,6 +458,7 @@ 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 []
@@ -407,14 +484,12 @@ class _V3EventProcessor:
if tool_call is None:
return events
tool_name, args, tool_call_id = tool_call
events.extend(
self._emit_tool_call_once(
namespace=namespace,
subagent=subagent,
name=tool_name,
args=args,
tool_call_id=tool_call_id,
)
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,
)
return events
@@ -457,6 +532,20 @@ 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,
@@ -464,6 +553,16 @@ 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
@@ -486,14 +585,12 @@ class _V3EventProcessor:
if tool_call is None:
continue
tool_name, args, tool_call_id = tool_call
events.extend(
self._emit_tool_call_once(
namespace=namespace,
subagent=subagent,
name=tool_name,
args=args,
tool_call_id=tool_call_id,
)
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,
)
if subagent is None:
@@ -526,6 +623,9 @@ 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,
@@ -572,6 +672,7 @@ 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(
@@ -583,22 +684,78 @@ class _V3EventProcessor:
instance_id=subagent.instance_id,
).data
)
return events
events.append(
self.emitter.tool_result(
name, content, success, tool_call_id=tool_call_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
return []
def _process_update_event(self, data: object) -> list[dict[str, Any]]:
@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]]:
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)
@@ -614,6 +771,10 @@ 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):
@@ -652,26 +813,74 @@ 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", ""))
return self._dedupe_interrupt_event(
self.emitter.ask_user_interrupt(
interrupt_id,
questions,
tc_id,
).data
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 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:
return self._dedupe_interrupt_event(
self.emitter.interrupt(
interrupt_id,
action_reqs,
review_cfgs,
).data
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 events
return []
def _dedupe_interrupt_event(self, event: dict[str, Any]) -> list[dict[str, Any]]:
@@ -797,6 +1006,8 @@ 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.
@@ -812,6 +1023,9 @@ 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,
@@ -821,6 +1035,8 @@ 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:
@@ -949,7 +1165,30 @@ async def stream_agent_events(
yield item
except Exception as e:
_run_raised = True
yield emitter.error(str(e)).data
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
raise
finally:
if stream is not None:
+3 -5
View File
@@ -10,7 +10,7 @@
<a href="https://pypi.org/project/EvoScientist/"><picture>
<source media="(prefers-color-scheme: light)" srcset="https://raw.githubusercontent.com/EvoScientist/EvoScientist/main/.github/assets/badge-pypi-light.svg">
<source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/EvoScientist/EvoScientist/main/.github/assets/badge-pypi-dark.svg">
<img alt="PyPI v0.2.1" src="https://raw.githubusercontent.com/EvoScientist/EvoScientist/main/.github/assets/badge-pypi-light.svg" height="28">
<img alt="PyPI v0.2.2" src="https://raw.githubusercontent.com/EvoScientist/EvoScientist/main/.github/assets/badge-pypi-light.svg" height="28">
</picture></a><a href="https://EvoScientist.github.io/"><picture>
<source media="(prefers-color-scheme: light)" srcset="https://raw.githubusercontent.com/EvoScientist/EvoScientist/main/.github/assets/badge-website-light.svg">
<source media="(prefers-color-scheme: dark)" srcset="https://raw.githubusercontent.com/EvoScientist/EvoScientist/main/.github/assets/badge-website-dark.svg">
@@ -151,6 +151,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.
- **[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`.
@@ -433,10 +434,7 @@ 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/stream-json.md](docs/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/guides/stream-json.md](docs/guides/stream-json.md) for the event schema.
</details>
+5
View File
@@ -11,6 +11,11 @@
|------------------------------------------------------------|---------------------------------------------------------------------------------|
| [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:
+5 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "EvoScientist"
version = "0.2.1"
version = "0.2.2"
description = "EvoScientist: Towards Self-Evolving AI Scientists for End-to-End Scientific Discovery"
readme = "README.md"
requires-python = ">=3.11"
@@ -48,6 +48,7 @@ dependencies = [
[dependency-groups]
dev = [
"pytest>=8.0",
"pytest-asyncio>=1.0",
"pytest-cov>=5.0",
"pytest-timeout>=2.4",
"ruff>=0.5",
@@ -58,6 +59,7 @@ dev = [
[project.optional-dependencies]
dev = [
"pytest>=8.0",
"pytest-asyncio>=1.0",
"pytest-cov>=5.0",
"pytest-timeout>=2.4",
"ruff>=0.5",
@@ -117,6 +119,8 @@ EvoScientist = [
[tool.pytest.ini_options]
testpaths = ["tests"]
asyncio_mode = "auto"
asyncio_default_fixture_loop_scope = "function"
filterwarnings = [
"ignore::UserWarning:langchain_nvidia_ai_endpoints",
]
+23 -26
View File
@@ -1,34 +1,10 @@
"""Shared fixtures for EvoScientist tests."""
import asyncio
from pathlib import Path
import pytest
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
_NONEXISTENT_DOTENV = str(Path(__file__).with_name(".pytest-dotenv-does-not-exist"))
@pytest.fixture(autouse=True)
@@ -192,3 +168,24 @@ 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,
)
+10 -15
View File
@@ -9,7 +9,6 @@ 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]:
@@ -17,24 +16,20 @@ async def async_iter(items: Iterable[Any]) -> AsyncIterator[Any]:
yield item
def collect_events(
async def collect_events(
agent,
message: str = "hi",
thread_id: str = "t1",
):
"""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())
"""Collect stream_agent_events output for tests."""
events = []
async for ev in stream_agent_events(
agent,
message,
thread_id,
):
events.append(ev)
return events
def protocol_event(
+20 -19
View File
@@ -10,18 +10,17 @@ 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:
def test_start_raises_without_required_imap_settings(self):
async def test_start_raises_without_required_imap_settings(self):
channel = EmailChannel(EmailConfig())
with pytest.raises(
ChannelError, match="imap_host and imap_username are required"
):
_run(channel.start())
await channel.start()
def test_send_returns_false_when_smtp_not_ready(self):
async def test_send_returns_false_when_smtp_not_ready(self):
channel = EmailChannel(EmailConfig())
msg = OutboundMessage(
channel="email",
@@ -29,16 +28,16 @@ class TestEmailChannelSmoke:
content="hello",
metadata={"chat_id": "user@example.com"},
)
assert _run(channel.send(msg)) is False
assert await channel.send(msg) is False
class TestSignalChannelSmoke:
def test_start_raises_without_phone_number(self):
async def test_start_raises_without_phone_number(self):
channel = SignalChannel(SignalConfig())
with pytest.raises(ChannelError, match="phone_number is required"):
_run(channel.start())
await channel.start()
def test_send_returns_false_when_not_connected(self):
async def test_send_returns_false_when_not_connected(self):
channel = SignalChannel(SignalConfig(phone_number="+123456789"))
msg = OutboundMessage(
channel="signal",
@@ -46,27 +45,29 @@ class TestSignalChannelSmoke:
content="hello",
metadata={"chat_id": "+123456789"},
)
assert _run(channel.send(msg)) is False
assert await channel.send(msg) is False
class TestQQChannelSmoke:
def test_start_raises_when_sdk_missing(self, monkeypatch):
async 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"):
_run(channel.start())
await channel.start()
def test_start_raises_without_credentials_when_sdk_available(self, monkeypatch):
async 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"):
_run(channel.start())
await channel.start()
def test_send_returns_false_without_client(self):
async def test_send_returns_false_without_client(self):
channel = QQChannel(QQConfig(app_id="id", app_secret="secret"))
msg = OutboundMessage(
channel="qq",
@@ -74,11 +75,11 @@ class TestQQChannelSmoke:
content="hello",
metadata={"chat_id": "openid"},
)
assert _run(channel.send(msg)) is False
assert await channel.send(msg) is False
class TestIMessageChannelSmoke:
def test_start_wraps_rpc_bootstrap_error(self, monkeypatch):
async def test_start_wraps_rpc_bootstrap_error(self, monkeypatch):
async def _broken_start(self):
raise RuntimeError("imsg not found")
@@ -87,9 +88,9 @@ class TestIMessageChannelSmoke:
monkeypatch.setattr(imessage_module.ImsgRpcClient, "start", _broken_start)
channel = IMessageChannelRpc(IMessageConfig())
with pytest.raises(ChannelError, match="Failed to start imsg"):
_run(channel.start())
await channel.start()
def test_send_returns_false_without_rpc_client(self):
async def test_send_returns_false_without_rpc_client(self):
channel = IMessageChannelRpc(IMessageConfig())
msg = OutboundMessage(
channel="imessage",
@@ -97,4 +98,4 @@ class TestIMessageChannelSmoke:
content="hello",
metadata={"chat_id": "+123456789"},
)
assert _run(channel.send(msg)) is False
assert await channel.send(msg) is False
+137
View File
@@ -0,0 +1,137 @@
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",
]
+123 -146
View File
@@ -3,6 +3,7 @@
from __future__ import annotations
import asyncio
import threading
import pytest
@@ -93,88 +94,95 @@ def _make_loader_fn(agent_value="AGENT", fail_with=None, capture=None):
return _loader
def _run(coro):
return asyncio.run(coro)
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)
class TestBackgroundAgentLoaderStart:
def test_start_creates_task_and_forwards_kwargs(self):
async def test_start_creates_task_and_forwards_kwargs(self):
captured: dict = {}
loader = BackgroundAgentLoader(_make_loader_fn(capture=captured))
async def _go():
loader.start(workspace_dir="/ws", checkpointer="CK")
assert loader.task is not None
assert loader.is_pending
await loader.await_ready()
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"}
def test_start_bumps_load_id(self):
async def test_start_bumps_load_id(self):
loader = BackgroundAgentLoader(_make_loader_fn())
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()
assert loader._load_id == 0
loader.start()
assert loader._load_id == 1
loader.start()
assert loader._load_id == 2
await loader.await_ready()
_run(_go())
async def test_start_cancels_in_flight_prior_task(self):
blocking = _GatedThreadLoader("LATE")
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())
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()
class TestBackgroundAgentLoaderCallbacks:
def test_progress_hook_sees_events_in_order(self):
async 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)),
)
async def _go():
loader.start()
await loader.await_ready()
loader.start()
await loader.await_ready()
_run(_go())
assert events == [("start", "srv", ""), ("success", "srv", "1")]
def test_stale_progress_events_are_dropped(self):
async def test_stale_progress_events_are_dropped(self):
"""A progress event fired after a newer `start` must not reach the hook."""
import time
slow_loader = _GatedThreadLoader(
"slow-agent", progress_events=[("success", "from-slow", "1")]
)
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")
@@ -184,36 +192,32 @@ class TestBackgroundAgentLoaderCallbacks:
slow_loader, on_progress=lambda e, s, d: seen.append(s)
)
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)
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)
_run(_go())
assert "from-fast" in seen
assert "from-slow" not in seen
def test_success_callback_fires_on_completion(self):
async def test_success_callback_fires_on_completion(self):
got = []
loader = BackgroundAgentLoader(
_make_loader_fn("MY_AGENT"),
on_success=lambda a: got.append(a),
)
async def _go():
loader.start()
await loader.await_ready()
await asyncio.sleep(0) # let done-callback run
loader.start()
await loader.await_ready()
await asyncio.sleep(0) # let done-callback run
_run(_go())
assert got == ["MY_AGENT"]
def test_failure_callback_fires_on_error(self):
async def test_failure_callback_fires_on_error(self):
err = RuntimeError("load failed")
got_failures = []
got_successes = []
@@ -223,40 +227,33 @@ class TestBackgroundAgentLoaderCallbacks:
on_failure=lambda e: got_failures.append(e),
)
async def _go():
loader.start()
with pytest.raises(RuntimeError, match="load failed"):
await loader.await_ready()
await asyncio.sleep(0)
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:
def test_returns_cached_agent_without_reawaiting(self):
async def test_returns_cached_agent_without_reawaiting(self):
captured: dict = {}
loader = BackgroundAgentLoader(_make_loader_fn("A", capture=captured))
async def _go():
loader.start()
assert await loader.await_ready() == "A"
assert await loader.await_ready() == "A"
loader.start()
assert await loader.await_ready() == "A"
assert await loader.await_ready() == "A"
_run(_go())
assert len(captured["kwargs"]) == 1
def test_raises_if_started_not_called(self):
async def test_raises_if_started_not_called(self):
loader = BackgroundAgentLoader(_make_loader_fn())
async def _go():
with pytest.raises(RuntimeError, match="before start"):
await loader.await_ready()
with pytest.raises(RuntimeError, match="before start"):
await loader.await_ready()
_run(_go())
def test_reraises_real_error_on_subsequent_awaits(self):
async 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."""
@@ -265,16 +262,13 @@ class TestBackgroundAgentLoaderAwaitReady:
loader = BackgroundAgentLoader(_fail)
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()
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()
_run(_go())
def test_needs_restart_flags_failed_load_for_retry(self):
async def test_needs_restart_flags_failed_load_for_retry(self):
calls = {"n": 0}
def flaky(*, on_mcp_progress=None):
@@ -285,17 +279,14 @@ class TestBackgroundAgentLoaderAwaitReady:
loader = BackgroundAgentLoader(flaky)
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())
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
class TestBackgroundAgentLoaderAdopt:
@@ -305,26 +296,19 @@ class TestBackgroundAgentLoaderAdopt:
assert loader.agent == "EXTERNAL"
assert not loader.is_pending
def test_adopt_supersedes_in_flight_load(self):
async def test_adopt_supersedes_in_flight_load(self):
"""A late background completion must not overwrite an adopted agent."""
import time
slow_loader = _GatedThreadLoader("FROM_BACKGROUND")
def _slow(*, on_mcp_progress=None):
time.sleep(0.08)
return "FROM_BACKGROUND"
loader = BackgroundAgentLoader(slow_loader)
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())
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"
class TestBackgroundAgentLoaderIsPending:
@@ -332,29 +316,22 @@ class TestBackgroundAgentLoaderIsPending:
loader = BackgroundAgentLoader(_make_loader_fn())
assert not loader.is_pending
def test_false_after_completion(self):
async def test_false_after_completion(self):
loader = BackgroundAgentLoader(_make_loader_fn())
async def _go():
loader.start()
await loader.await_ready()
loader.start()
await loader.await_ready()
_run(_go())
assert not loader.is_pending
def test_true_between_start_and_completion(self):
import time
async def test_true_between_start_and_completion(self):
wait_loader = _GatedThreadLoader("ok")
def _wait_loader(*, on_mcp_progress=None):
time.sleep(0.05)
return "ok"
loader = BackgroundAgentLoader(wait_loader)
loader = BackgroundAgentLoader(_wait_loader)
async def _go():
loader.start()
assert loader.is_pending
await loader.await_ready()
assert not loader.is_pending
_run(_go())
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
+129 -201
View File
@@ -5,6 +5,8 @@ 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,
@@ -28,12 +30,6 @@ 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)
@@ -51,7 +47,7 @@ def _drain_queue(q):
return items
def test_read_async_tasks_from_gateway_reads_state_values(run_async):
async def test_read_async_tasks_from_gateway_reads_state_values():
gateway = FakeGraphGateway(
state_values={
"async_tasks": {
@@ -60,18 +56,16 @@ def test_read_async_tasks_from_gateway_reads_state_values(run_async):
}
)
tasks = run_async(
async_notifier.read_async_tasks_from_gateway(
gateway,
GraphTarget(local_graph=MagicMock()),
"tid",
)
tasks = await async_notifier.read_async_tasks_from_gateway(
gateway,
GraphTarget(local_graph=MagicMock()),
"tid",
)
assert tasks == {"task-1": {"status": "success"}}
def test_watcher_pushes_notification_on_stream_end(run_async):
async def test_watcher_pushes_notification_on_stream_end():
# Stream yields one "values" chunk with the final state, then closes
final_state = {
"messages": [{"type": "ai", "content": "Quantum superposition is..."}]
@@ -87,10 +81,7 @@ def test_watcher_pushes_notification_on_stream_end(run_async):
# runs.get is used to fetch terminal status when stream ends
client.runs.get = AsyncMock(return_value={"status": "success"})
_drain_all(async_notifier)
run_async(
async_notifier.watch_run_and_notify(client, "thr-1", "run-1", "writing-agent")
)
await async_notifier.watch_run_and_notify(client, "thr-1", "run-1", "writing-agent")
notifs = _drain_queue(async_notifier._notification_queue)
assert len(notifs) == 1
@@ -99,7 +90,7 @@ def test_watcher_pushes_notification_on_stream_end(run_async):
assert notifs[0].status == "success"
def test_watcher_pushes_error_status_on_stream_exception(run_async):
async def test_watcher_pushes_error_status_on_stream_exception():
async def fake_stream(*a, **kw):
raise RuntimeError("network broken")
yield # unreachable; makes this an async generator
@@ -111,14 +102,13 @@ def test_watcher_pushes_error_status_on_stream_exception(run_async):
return_value={"status": "error", "error": "network broken"}
)
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thr-4", "run-4", "agentZ"))
await async_notifier.watch_run_and_notify(client, "thr-4", "run-4", "agentZ")
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "error"
def test_spawn_watcher_replaces_existing_for_same_thread(run_async):
async def test_spawn_watcher_replaces_existing_for_same_thread():
"""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."""
@@ -138,43 +128,34 @@ def test_spawn_watcher_replaces_existing_for_same_thread(run_async):
client.runs.join_stream = fake_stream_long
client.runs.get = AsyncMock(return_value={"status": "success"})
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)
# 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
# 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
# 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
# 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
# Old watcher should be cancelled
await asyncio.sleep(0.02)
assert t1.cancelled() or t1.done()
# Old watcher should be cancelled
await asyncio.sleep(0.02)
assert t1.cancelled() or t1.done()
# Cleanup the new task too
t2.cancel()
try:
await t2
except asyncio.CancelledError:
pass
# 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())
# 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) == []
# ============================================================================
@@ -324,13 +305,6 @@ 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(
@@ -463,17 +437,12 @@ def test_dedup_preserves_order():
# ============================================================================
def test_consume_notifications_calls_runner_with_batched_message(run_async):
async def test_consume_notifications_calls_runner_with_batched_message():
"""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", "", ""))
@@ -486,21 +455,15 @@ def test_consume_notifications_calls_runner_with_batched_message(run_async):
async def fake_state_reader() -> dict:
return {} # no dedup info
run_async(an.consume_notifications(fake_runner, fake_state_reader))
await an.consume_notifications(fake_runner, fake_state_reader)
assert "wA" in captured["text"]
assert "wB" in captured["text"]
assert len(captured["notifs"]) == 2
def test_consume_notifications_no_op_when_queue_empty(run_async):
async def test_consume_notifications_no_op_when_queue_empty():
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):
@@ -510,7 +473,7 @@ def test_consume_notifications_no_op_when_queue_empty(run_async):
async def fake_state_reader():
return {}
run_async(an.consume_notifications(fake_runner, fake_state_reader))
await an.consume_notifications(fake_runner, fake_state_reader)
assert called is False
@@ -521,7 +484,7 @@ def test_consume_notifications_no_op_when_queue_empty(run_async):
# ============================================================================
def test_notification_consuming_flag_prevents_reentry(run_async):
async def test_notification_consuming_flag_prevents_reentry():
"""The _notification_consuming guard prevents two overlapping consumers.
Verifies the flag contract used by _consume_notifications_tui:
@@ -536,13 +499,6 @@ def test_notification_consuming_flag_prevents_reentry(run_async):
"""
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:
@@ -565,45 +521,42 @@ def test_notification_consuming_flag_prevents_reentry(run_async):
n1 = an.AsyncTaskNotification("g1", "writing-agent", "success", "", "")
n2 = an.AsyncTaskNotification("g2", "data-agent", "success", "", "")
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
# 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
run_async(scenario())
await guarded_consume_raising(n2)
assert state["consuming"] is False # cleared despite exception
# ============================================================================
@@ -613,33 +566,42 @@ def test_notification_consuming_flag_prevents_reentry(run_async):
def _drain_all(an_mod):
"""Drain every queue (per-thread + unrouted) so tests start clean."""
if hasattr(an_mod, "_notification_queue"):
while True:
try:
an_mod._notification_queue.get_nowait()
except queue.Empty:
break
for q in list(an_mod._notifications_by_thread.values()):
while True:
try:
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()
q.get_nowait()
except queue.Empty:
break
while True:
try:
an_mod._unrouted_queue.get_nowait()
except queue.Empty:
break
def test_consume_only_drains_matching_thread(run_async):
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():
"""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"
)
@@ -657,21 +619,17 @@ def test_consume_only_drains_matching_thread(run_async):
async def state_reader() -> dict:
return {}
run_async(
an.consume_notifications(runner, state_reader, current_thread_id="threadA")
)
await 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)
def test_unrouted_notifications_drain_on_any_thread(run_async):
async def test_unrouted_notifications_drain_on_any_thread():
"""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", "", "")
)
@@ -684,19 +642,15 @@ def test_unrouted_notifications_drain_on_any_thread(run_async):
async def state_reader() -> dict:
return {}
run_async(
an.consume_notifications(runner, state_reader, current_thread_id="anything")
)
await an.consume_notifications(runner, state_reader, current_thread_id="anything")
assert [n.task_id for n in captured["notifs"]] == ["tU"]
_drain_all(an)
def test_thread_switch_drains_pending(run_async):
async def test_thread_switch_drains_pending():
"""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"
@@ -712,25 +666,19 @@ def test_thread_switch_drains_pending(run_async):
return {}
# First consume in thread A → no drain, B's notif still queued
run_async(
an.consume_notifications(runner, state_reader, current_thread_id="threadA")
)
await 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
run_async(
an.consume_notifications(runner, state_reader, current_thread_id="threadB")
)
await 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", "", "")
@@ -748,7 +696,6 @@ 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)
# ============================================================================
@@ -762,7 +709,7 @@ def test_has_pending_notifications_respects_routing():
# ============================================================================
def test_watcher_reports_error_on_in_band_error_event(run_async):
async def test_watcher_reports_error_on_in_band_error_event():
"""SSE error event in the stream → notification.status == 'error'."""
async def fake_stream(*a, **kw):
@@ -777,8 +724,7 @@ def test_watcher_reports_error_on_in_band_error_event(run_async):
return_value={"status": "success"}
) # would mislead — should NOT be consulted
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrE", "rE", "agentE"))
await async_notifier.watch_run_and_notify(client, "thrE", "rE", "agentE")
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "error"
@@ -786,7 +732,7 @@ def test_watcher_reports_error_on_in_band_error_event(run_async):
client.runs.get.assert_not_awaited()
def test_watcher_clean_exit_with_runs_get_success_is_success(run_async):
async def test_watcher_clean_exit_with_runs_get_success_is_success():
"""Clean stream exit + runs.get reports success → status=success."""
async def fake_stream(*a, **kw):
@@ -798,15 +744,14 @@ def test_watcher_clean_exit_with_runs_get_success_is_success(run_async):
client.runs.join_stream = fake_stream
client.runs.get = AsyncMock(return_value={"status": "success"})
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrS", "rS", "agentS"))
await 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()
def test_watcher_clean_exit_with_runs_get_error_is_race_safe(run_async):
async def test_watcher_clean_exit_with_runs_get_error_is_race_safe():
"""Clean stream exit + no in-band error event + runs.get returns 'error'
→ status=success (race-safe).
@@ -827,14 +772,13 @@ def test_watcher_clean_exit_with_runs_get_error_is_race_safe(run_async):
client.runs.join_stream = fake_stream
client.runs.get = AsyncMock(return_value={"status": "error"})
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrS", "rS", "agentS"))
await async_notifier.watch_run_and_notify(client, "thrS", "rS", "agentS")
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "success"
def test_watcher_clean_exit_with_runs_get_running_drops_notification(run_async):
async def test_watcher_clean_exit_with_runs_get_running_drops_notification():
"""Reproduces the production bug: clean SSE close while run is still
actually running (HTTP keep-alive timeout under concurrency).
@@ -856,11 +800,8 @@ def test_watcher_clean_exit_with_runs_get_running_drops_notification(run_async):
client.runs.join_stream = fake_stream
client.runs.get = AsyncMock(return_value={"status": "running"})
_drain_all(async_notifier)
run_async(
async_notifier.watch_run_and_notify(
client, "thr-bug", "rB", "data-analysis-agent"
)
await async_notifier.watch_run_and_notify(
client, "thr-bug", "rB", "data-analysis-agent"
)
# No notification should have been enqueued anywhere.
@@ -872,7 +813,7 @@ def test_watcher_clean_exit_with_runs_get_running_drops_notification(run_async):
assert client.runs.get.await_count >= 1
def test_watcher_unknown_status_treated_as_non_terminal(run_async):
async def test_watcher_unknown_status_treated_as_non_terminal():
"""Future / unrecognized status values should trigger a re-join, not a
false-positive notification.
@@ -893,8 +834,7 @@ def test_watcher_unknown_status_treated_as_non_terminal(run_async):
side_effect=[{"status": "queued"}, {"status": "success"}]
)
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrU", "rU", "agentU"))
await async_notifier.watch_run_and_notify(client, "thrU", "rU", "agentU")
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "success"
@@ -902,7 +842,7 @@ def test_watcher_unknown_status_treated_as_non_terminal(run_async):
assert client.runs.get.await_count == 2
def test_watcher_runs_get_persistent_failure_drops_notification(run_async, monkeypatch):
async def test_watcher_runs_get_persistent_failure_drops_notification(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
@@ -921,22 +861,20 @@ def test_watcher_runs_get_persistent_failure_drops_notification(run_async, monke
monkeypatch.setattr(async_notifier.asyncio, "sleep", _no_sleep)
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrG", "rG", "agentG"))
await 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) == []
if hasattr(async_notifier, "_notifications_by_thread"):
for q in async_notifier._notifications_by_thread.values():
assert _drain_one_queue_helper(q) == []
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
def test_watcher_runs_get_transient_failure_recovers(run_async, monkeypatch):
async def test_watcher_runs_get_transient_failure_recovers(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
@@ -957,15 +895,14 @@ def test_watcher_runs_get_transient_failure_recovers(run_async, monkeypatch):
monkeypatch.setattr(async_notifier.asyncio, "sleep", _no_sleep)
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrT", "rT", "agentT"))
await 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
def test_watcher_re_joins_stream_until_terminal_status(run_async):
async def test_watcher_re_joins_stream_until_terminal_status():
"""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."""
@@ -980,8 +917,7 @@ def test_watcher_re_joins_stream_until_terminal_status(run_async):
side_effect=[{"status": "running"}, {"status": "success"}]
)
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrR", "rR", "agentR"))
await async_notifier.watch_run_and_notify(client, "thrR", "rR", "agentR")
notif = async_notifier._notification_queue.get_nowait()
assert notif.status == "success"
@@ -995,15 +931,12 @@ def test_watcher_re_joins_stream_until_terminal_status(run_async):
# ============================================================================
def test_consume_notifications_propagates_inject_exception(run_async):
async def test_consume_notifications_propagates_inject_exception():
"""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", "", "")
)
@@ -1015,11 +948,10 @@ def test_consume_notifications_propagates_inject_exception(run_async):
return {}
with pytest.raises(RuntimeError, match="kaboom"):
run_async(an.consume_notifications(boom_runner, state_reader))
_drain_all(an)
await an.consume_notifications(boom_runner, state_reader)
def test_watcher_skips_notification_on_stream_fail_with_nonterminal_status(run_async):
async def test_watcher_skips_notification_on_stream_fail_with_nonterminal_status():
"""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
@@ -1035,15 +967,13 @@ def test_watcher_skips_notification_on_stream_fail_with_nonterminal_status(run_a
client.runs.join_stream = fake_stream
client.runs.get = AsyncMock(return_value={"status": "pending"})
_drain_all(async_notifier)
run_async(async_notifier.watch_run_and_notify(client, "thrP", "rP", "agentP"))
await 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) == []
if hasattr(async_notifier, "_notifications_by_thread"):
for q in async_notifier._notifications_by_thread.values():
assert _drain_one_queue_helper(q) == []
for q in async_notifier._notifications_by_thread.values():
assert _drain_one_queue_helper(q) == []
def _drain_one_queue_helper(q):
@@ -1060,8 +990,6 @@ 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()
+22 -21
View File
@@ -7,7 +7,6 @@ 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
@@ -105,7 +104,7 @@ def _make_middleware():
return mw, fake_client
def test_middleware_spawns_watcher_on_start_async_task():
async 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
@@ -142,7 +141,7 @@ def test_middleware_spawns_watcher_on_start_async_task():
)
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
result = asyncio.run(mw.awrap_tool_call(request, fake_handler))
result = await mw.awrap_tool_call(request, fake_handler)
assert isinstance(result, Command)
assert spawn_calls == [
@@ -150,7 +149,7 @@ def test_middleware_spawns_watcher_on_start_async_task():
]
def test_middleware_spawns_watcher_on_update_async_task():
async 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
@@ -183,7 +182,7 @@ def test_middleware_spawns_watcher_on_update_async_task():
)
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
asyncio.run(mw.awrap_tool_call(request, fake_handler))
await mw.awrap_tool_call(request, fake_handler)
assert len(spawn_calls) == 1
args, kwargs = spawn_calls[0]
@@ -195,7 +194,7 @@ def test_middleware_spawns_watcher_on_update_async_task():
assert kwargs["origin_cli_thread_id"] == "cli-thread-A"
def test_middleware_pre_cancels_old_watcher_on_update():
async 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
@@ -221,14 +220,14 @@ def test_middleware_pre_cancels_old_watcher_on_update():
try:
with patch.object(async_notifier, "spawn_watcher"):
asyncio.run(mw.awrap_tool_call(request, fake_handler))
await 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
def test_middleware_passes_through_unrelated_tools():
async 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()
@@ -240,13 +239,13 @@ 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 = asyncio.run(mw.awrap_tool_call(request, fake_handler))
result = await mw.awrap_tool_call(request, fake_handler)
assert result is sentinel
assert mock_spawn.call_count == 0
def test_middleware_handles_non_command_results_gracefully():
async def test_middleware_handles_non_command_results_gracefully():
"""If the launch tool returns a string (validation error), no watcher is spawned."""
mw, _ = _make_middleware()
@@ -260,13 +259,13 @@ def test_middleware_handles_non_command_results_gracefully():
)
with patch.object(async_notifier, "spawn_watcher") as mock_spawn:
result = asyncio.run(mw.awrap_tool_call(request, fake_handler))
result = await mw.awrap_tool_call(request, fake_handler)
assert result == "Unknown async subagent type `bogus`"
assert mock_spawn.call_count == 0
def test_middleware_origin_thread_id_is_none_when_runtime_config_missing():
async 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
@@ -299,12 +298,12 @@ def test_middleware_origin_thread_id_is_none_when_runtime_config_missing():
)
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
asyncio.run(mw.awrap_tool_call(request, fake_handler))
await mw.awrap_tool_call(request, fake_handler)
assert captured.get("origin_cli_thread_id") is None
def test_middleware_swallows_spawn_exceptions():
async def test_middleware_swallows_spawn_exceptions():
"""spawn_watcher errors must not propagate up — middleware logs and continues."""
from langgraph.types import Command
@@ -336,7 +335,7 @@ def test_middleware_swallows_spawn_exceptions():
with patch.object(async_notifier, "spawn_watcher", side_effect=boom):
# Should not raise.
result = asyncio.run(mw.awrap_tool_call(request, fake_handler))
result = await mw.awrap_tool_call(request, fake_handler)
assert isinstance(result, Command)
@@ -356,7 +355,9 @@ def test_middleware_swallows_spawn_exceptions():
),
],
)
def test_middleware_picks_correct_prompt_field_per_tool(tool_name, args, prompt_field):
async 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
@@ -385,12 +386,12 @@ def test_middleware_picks_correct_prompt_field_per_tool(tool_name, args, prompt_
request = _build_request(tool_name, args, thread_id="t")
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
asyncio.run(mw.awrap_tool_call(request, fake_handler))
await mw.awrap_tool_call(request, fake_handler)
assert captured_prompt["value"] == prompt_field
def test_middleware_prompt_field_is_tool_name_gated_not_fallback_chained():
async 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')`
@@ -432,12 +433,12 @@ def test_middleware_prompt_field_is_tool_name_gated_not_fallback_chained():
)
with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn):
asyncio.run(mw.awrap_tool_call(request, fake_handler))
await mw.awrap_tool_call(request, fake_handler)
assert captured_prompt["value"] == "use this"
def test_middleware_pre_cancel_swallows_unexpected_errors():
async 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
@@ -460,7 +461,7 @@ def test_middleware_pre_cancel_swallows_unexpected_errors():
try:
with patch.object(async_notifier, "spawn_watcher"):
# Should not raise.
asyncio.run(mw.awrap_tool_call(request, fake_handler))
await mw.awrap_tool_call(request, fake_handler)
finally:
async_notifier._watcher_by_thread.pop("t1", None)
+6 -7
View File
@@ -1,6 +1,5 @@
from __future__ import annotations
import asyncio
import json
from types import SimpleNamespace
@@ -917,16 +916,16 @@ class _AsyncFakeCrons:
return [{"cron_id": "cron-async"}]
def test_alist_autoskill_schedules_uses_async_client_and_explicit_limit(monkeypatch):
async 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 = asyncio.run(
alist_autoskill_schedules(
EvoScientistConfig(),
limit=3,
)
rows = await alist_autoskill_schedules(
EvoScientistConfig(),
limit=3,
)
assert rows == [{"cron_id": "cron-async"}]
+2 -2
View File
@@ -91,7 +91,7 @@ def test_stop_already_finished_is_graceful(tmp_path):
assert "already finished" in bg.stop(pid)
def test_exited_elapsed_is_frozen(tmp_path):
def test_exited_elapsed_is_frozen(tmp_path, monkeypatch):
"""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):
proc = bg._PROCESSES[pid]
assert proc.finished_ts is not None
first = bg._elapsed(proc)
time.sleep(1.1) # intentional: prove elapsed stays frozen, not ticking up
monkeypatch.setattr(bg.time, "time", lambda: proc.finished_ts + 100.0)
assert bg._elapsed(proc) == first
+369 -403
View File
@@ -12,7 +12,6 @@ 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
@@ -58,7 +57,7 @@ def clean_channel_state():
class TestBusInboundConsumer:
"""Test the _bus_inbound_consumer queue bridge."""
def test_processes_inbound_and_publishes_outbound(self):
async def test_processes_inbound_and_publishes_outbound(self):
"""InboundMessage -> queue -> response -> OutboundMessage flow."""
from EvoScientist.cli.channel import (
_bus_inbound_consumer,
@@ -68,54 +67,51 @@ class TestBusInboundConsumer:
_drain_queue(_message_queue)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
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
_run(_test())
def test_no_response_fallback(self):
async def test_no_response_fallback(self):
"""Empty response is replaced with 'No response' fallback."""
from EvoScientist.cli.channel import (
_bus_inbound_consumer,
@@ -125,47 +121,44 @@ class TestBusInboundConsumer:
_drain_queue(_message_queue)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
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
_run(_test())
def test_late_response_after_timeout_still_publishes(self, monkeypatch):
async 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 (
@@ -179,56 +172,53 @@ class TestBusInboundConsumer:
_drain_queue(_message_queue)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
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
_run(_test())
def test_late_timeout_keeps_active_request_cancellable(self, monkeypatch):
async 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 (
@@ -243,191 +233,179 @@ class TestBusInboundConsumer:
monkeypatch.setattr(channel_mod, "_RESPONSE_TIMEOUT", 0.05)
monkeypatch.setattr(channel_mod, "_LATE_RESPONSE_TIMEOUT", 0.05)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
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="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(
task = asyncio.create_task(
_handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="/stop",
message_id="msg-stop-active",
content="still running",
message_id="msg-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)
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)
_run(_test())
assert queued is not None
assert _claim_channel_request(queued) is True
def test_cancelled_wait_cleans_pending_response(self):
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):
"""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
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
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
_run(_test())
def test_consumer_shutdown_cleans_pending_response(self):
async 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
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
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
_run(_test())
def test_stop_during_hitl_wait_releases_wait_and_acks(self):
async 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
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
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
_run(_test())
def test_stop_cancels_queued_request_before_main_thread_processes_it(self):
async 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 (
@@ -436,73 +414,70 @@ class TestBusInboundConsumer:
_message_queue,
)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
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="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(
task = asyncio.create_task(
_handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="/stop",
message_id="m-stop-2",
content="please work",
message_id="m-work-1",
),
)
)
with pytest.raises(asyncio.CancelledError):
await task
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)
skipped = _message_queue.get_nowait()
assert skipped.msg_id == queued.msg_id
assert _claim_or_complete_channel_request(skipped) is False
assert queued is not None
with channel_mod._response_lock:
assert queued.msg_id in channel_mod._pending_responses
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
await _handle_bus_message(
bus,
manager,
InboundMessage(
channel="fake",
sender_id="user1",
chat_id="chat1",
content="/stop",
message_id="m-stop-2",
),
)
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)
with pytest.raises(asyncio.CancelledError):
await task
_run(_test())
skipped = _message_queue.get_nowait()
assert skipped.msg_id == queued.msg_id
assert _claim_or_complete_channel_request(skipped) is False
def test_stop_leaves_resolved_response_available_for_delivery(self):
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):
"""`/stop` must not steal a response whose waiter already resolved."""
from EvoScientist.cli import channel as channel_mod
from EvoScientist.cli.channel import (
@@ -515,42 +490,39 @@ class TestBusInboundConsumer:
_set_channel_response,
)
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",
)
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)
_run(_test())
def test_message_counting(self):
async def test_message_counting(self):
"""Messages are counted via record_message."""
from EvoScientist.cli.channel import (
_bus_inbound_consumer,
@@ -560,45 +532,42 @@ class TestBusInboundConsumer:
_drain_queue(_message_queue)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
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
_run(_test())
def test_channel_message_carries_metadata(self):
async def test_channel_message_carries_metadata(self):
"""ChannelMessage carries metadata, chat_id, and message_id."""
from EvoScientist.cli.channel import (
_bus_inbound_consumer,
@@ -608,49 +577,46 @@ class TestBusInboundConsumer:
_drain_queue(_message_queue)
async def _test():
bus = MessageBus()
manager = ChannelManager(bus)
ch = FakeChannel()
manager.register(ch)
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
_run(_test())
consumer.cancel()
try:
await consumer
except asyncio.CancelledError:
pass
+65 -2
View File
@@ -6,6 +6,8 @@ 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,
@@ -15,6 +17,7 @@ from EvoScientist.ccproxy_manager import (
setup_codex_env,
start_ccproxy,
stop_ccproxy,
write_ccproxy_config,
)
# =============================================================================
@@ -52,6 +55,8 @@ 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):
@@ -123,9 +128,10 @@ 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):
def test_success(self, mock_popen, mock_running, mock_warning):
proc = MagicMock()
proc.poll.return_value = None
mock_popen.return_value = proc
@@ -134,6 +140,11 @@ 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")
@@ -143,7 +154,11 @@ class TestStartCcproxy:
proc.poll.return_value = None
mock_popen.return_value = proc
# Simulate time passing beyond deadline
mock_time.monotonic.side_effect = [0, 0, 31]
mock_time.monotonic.side_effect = [
0,
0,
_CCPROXY_HEALTH_TIMEOUT_SECONDS + 1,
]
mock_time.sleep = MagicMock()
with pytest.raises(RuntimeError, match="did not become healthy"):
@@ -154,6 +169,54 @@ 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
+10 -12
View File
@@ -3,8 +3,6 @@
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
@@ -55,7 +53,7 @@ class TestNeedsAgent:
class TestStartPath:
"""Start flow must propagate agent/thread_id globals."""
def test_start_binds_channel_runtime(self):
async def test_start_binds_channel_runtime(self):
from EvoScientist.commands.implementation.channel import ChannelCommand
ctx, _ui = _ctx()
@@ -77,11 +75,11 @@ class TestStartPath:
return_value=config,
),
):
_run(ChannelCommand().execute(ctx, ["telegram"]))
await ChannelCommand().execute(ctx, ["telegram"])
assert ctx.channel_runtime.agent is ctx.agent
assert ctx.channel_runtime.thread_id == "tid-42"
def test_start_propagates_send_thinking(self):
async def test_start_propagates_send_thinking(self):
"""send_thinking flag must reach _start_channels_bus_mode."""
from EvoScientist.commands.implementation.channel import ChannelCommand
@@ -111,14 +109,14 @@ class TestStartPath:
return_value=config,
),
):
_run(ChannelCommand().execute(ctx, ["telegram"]))
await ChannelCommand().execute(ctx, ["telegram"])
assert captured["agent"] is ctx.agent
assert captured["thread_id"] == "tid-42"
assert captured["send_thinking"] is False
class TestAddToRunningPath:
def test_add_to_running_binds_channel_runtime(self):
async def test_add_to_running_binds_channel_runtime(self):
from EvoScientist.commands.implementation.channel import ChannelCommand
ctx, _ui = _ctx()
@@ -140,11 +138,11 @@ class TestAddToRunningPath:
return_value=config,
),
):
_run(ChannelCommand().execute(ctx, ["discord"]))
await ChannelCommand().execute(ctx, ["discord"])
assert ctx.channel_runtime.agent is ctx.agent
assert ctx.channel_runtime.thread_id == "tid-42"
def test_add_to_running_propagates_send_thinking(self):
async 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
@@ -173,13 +171,13 @@ class TestAddToRunningPath:
return_value=config,
),
):
_run(ChannelCommand().execute(ctx, ["discord"]))
await ChannelCommand().execute(ctx, ["discord"])
assert captured["channel_type"] == "discord"
assert captured["send_thinking"] is True
class TestStatusPath:
def test_status_without_running_channels(self):
async def test_status_without_running_channels(self):
from EvoScientist.commands.implementation.channel import ChannelCommand
ctx, ui = _ctx()
@@ -198,6 +196,6 @@ class TestStatusPath:
return_value=config,
),
):
_run(ChannelCommand().execute(ctx, ["status"]))
await 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 -9
View File
@@ -8,7 +8,6 @@ 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
@@ -57,7 +56,7 @@ def _sent_text(bus_ref) -> str:
)
def test_handle_session_resume_sends_history_back_to_channel_without_local_duplicate():
async def test_handle_session_resume_sends_history_back_to_channel_without_local_duplicate():
callback = AsyncMock()
bus_ref = SimpleNamespace(publish_outbound=AsyncMock())
@@ -72,7 +71,7 @@ def test_handle_session_resume_sends_history_back_to_channel_without_local_dupli
thread_store=thread_store,
)
_run(_run_resume(ui, "thread-42", "/workspace"))
await _run_resume(ui, "thread-42", "/workspace")
callback.assert_awaited_once_with("thread-42", "/workspace")
assert thread_store.calls == [("get_thread_messages", "thread-42")]
@@ -84,7 +83,7 @@ def test_handle_session_resume_sends_history_back_to_channel_without_local_dupli
assert "EvoScientist: Here is the saved answer." in text
def test_handle_session_resume_propagates_callback_abort_without_history():
async 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()
@@ -95,7 +94,7 @@ def test_handle_session_resume_propagates_callback_abort_without_history():
)
with pytest.raises(RuntimeError, match="workspace conflict"):
_run(_run_resume(ui, "thread-42", "/workspace"))
await _run_resume(ui, "thread-42", "/workspace")
callback.assert_awaited_once_with("thread-42", "/workspace")
assert thread_store.calls == []
@@ -103,7 +102,7 @@ def test_handle_session_resume_propagates_callback_abort_without_history():
assert captured == []
def test_handle_session_resume_reports_history_load_error():
async def test_handle_session_resume_reports_history_load_error():
callback = AsyncMock()
bus_ref = SimpleNamespace(publish_outbound=AsyncMock())
ui, captured = _make_ui(
@@ -114,7 +113,7 @@ def test_handle_session_resume_reports_history_load_error():
),
)
_run(_run_resume(ui, "thread-42", "/workspace"))
await _run_resume(ui, "thread-42", "/workspace")
callback.assert_awaited_once_with("thread-42", "/workspace")
assert captured == []
@@ -123,7 +122,7 @@ def test_handle_session_resume_reports_history_load_error():
assert "history unavailable: db locked" in text
def test_handle_session_resume_distinguishes_non_displayable_messages():
async def test_handle_session_resume_distinguishes_non_displayable_messages():
bus_ref = SimpleNamespace(publish_outbound=AsyncMock())
ui, captured = _make_ui(
bus_ref=bus_ref,
@@ -132,7 +131,7 @@ def test_handle_session_resume_distinguishes_non_displayable_messages():
),
)
_run(_run_resume(ui, "thread-42", "/workspace"))
await _run_resume(ui, "thread-42", "/workspace")
assert captured == [
"Resumed session: thread-42\nNo displayable messages in this session."
File diff suppressed because it is too large Load Diff
+27 -58
View File
@@ -11,8 +11,6 @@ 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
@@ -75,10 +73,10 @@ def _make_channel_context(*, debug_trace=True, name="test_channel"):
return {"channel": channel}
def test_middleware_dedup_emits_structured_event(caplog):
async def test_middleware_dedup_emits_structured_event(caplog):
from EvoScientist.channels.middleware import DedupMiddleware
async def _run():
with caplog.at_level(logging.DEBUG):
mw = DedupMiddleware()
ctx = _make_channel_context()
raw = _make_raw(message_id="dup1")
@@ -91,49 +89,40 @@ 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
def test_middleware_allowlist_emits_structured_event(caplog):
async def test_middleware_allowlist_emits_structured_event(caplog):
from EvoScientist.channels.middleware import AllowListMiddleware
async def _run():
with caplog.at_level(logging.DEBUG):
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
def test_middleware_mention_gating_emits_structured_event(caplog):
async def test_middleware_mention_gating_emits_structured_event(caplog):
from EvoScientist.channels.middleware import MentionGatingMiddleware
async def _run():
with caplog.at_level(logging.DEBUG):
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
def test_typing_manager_emits_trace_events(caplog):
async def test_typing_manager_emits_trace_events(caplog):
from EvoScientist.channels.middleware import TypingManager
async def _run():
with caplog.at_level(logging.DEBUG):
send_action = AsyncMock(side_effect=RuntimeError("typing api down"))
mgr = TypingManager(
send_action,
@@ -144,17 +133,14 @@ 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
def test_ack_reaction_emits_error_traces(caplog):
async def test_ack_reaction_emits_error_traces(caplog):
from EvoScientist.channels.middleware import AckReactionMiddleware
async def _run():
with caplog.at_level(logging.DEBUG):
send_fn = AsyncMock()
remove_fn = AsyncMock(side_effect=RuntimeError("remove failed"))
ack = AckReactionMiddleware(
@@ -178,16 +164,13 @@ 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
def test_inbound_raw_event_emitted(caplog):
async def test_inbound_raw_event_emitted(caplog):
"""Integration-style: _enqueue_raw emits inbound_raw at the top."""
from EvoScientist.channels.base import Channel, RawIncoming
@@ -218,19 +201,16 @@ def test_inbound_raw_event_emitted(caplog):
config.ack_scope = "off"
config.dedup_ttl = 3600
async def _run():
with caplog.at_level(logging.DEBUG):
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
def test_format_fallback_emits_event(caplog):
async def test_format_fallback_emits_event(caplog):
"""_send_with_format_fallback emits outbound_format_fallback on fallback."""
from EvoScientist.channels.base import Channel
@@ -268,13 +248,10 @@ def test_format_fallback_emits_event(caplog):
if call_count == 1:
raise ValueError("parse error in formatted text")
async def _run():
with caplog.at_level(logging.DEBUG):
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
@@ -304,7 +281,7 @@ def test_trace_mixin_trace_event(caplog):
assert "key=val" in caplog.text
def test_standalone_dispatcher_treats_false_send_as_error(caplog):
async 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
@@ -315,7 +292,7 @@ def test_standalone_dispatcher_treats_false_send_as_error(caplog):
channel.send = AsyncMock(return_value=False)
bus = MessageBus()
async def _run():
with caplog.at_level(logging.DEBUG):
task = asyncio.create_task(standalone_outbound_dispatcher(bus, channel))
await bus.publish_outbound(
OutboundMessage(channel="test", chat_id="c1", content="hi")
@@ -326,14 +303,11 @@ 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
def test_standalone_dispatcher_sends_media():
async 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
@@ -345,21 +319,16 @@ def test_standalone_dispatcher_sends_media():
channel.send_media = AsyncMock(return_value=True)
bus = MessageBus()
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())
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
channel.send_media.assert_awaited_once_with(
recipient="c1",
file_path="/tmp/a.png",
+145 -180
View File
@@ -14,7 +14,6 @@ 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
@@ -43,25 +42,23 @@ def _make_msg(
)
def test_non_slash_returns_false():
async def test_non_slash_returns_false():
"""Plain text messages must fall through to the agent."""
msg = _make_msg(content="hello agent")
append = MagicMock()
handled = _run(
dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
)
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()
def test_unresolved_slash_returns_false():
async def test_unresolved_slash_returns_false():
"""Unknown slash commands must fall through (matches TUI behavior)."""
msg = _make_msg(content="/unknown-cmd")
append = MagicMock()
@@ -69,20 +66,18 @@ def test_unresolved_slash_returns_false():
"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,
)
handled = await dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
)
assert handled is False
def test_successful_slash_execution_sets_response_and_breadcrumb():
async 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()
@@ -100,15 +95,13 @@ def test_successful_slash_execution_sets_response_and_breadcrumb():
) as mock_execute,
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
handled = _run(
dispatch_channel_slash_command(
msg,
agent="fake-agent",
thread_id="t1",
workspace_dir="/tmp",
checkpointer=None,
append_system=append,
)
handled = await 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()
@@ -119,7 +112,7 @@ def test_successful_slash_execution_sets_response_and_breadcrumb():
assert any("Executed command from" in t for t in breadcrumbs)
def test_slash_dispatch_passes_graph_gateway_to_command_context():
async def test_slash_dispatch_passes_graph_gateway_to_command_context():
msg = _make_msg()
fake_cmd = MagicMock()
fake_cmd.needs_agent.return_value = False
@@ -142,23 +135,21 @@ def test_slash_dispatch_passes_graph_gateway_to_command_context():
),
patch("EvoScientist.cli.channel._set_channel_response"),
):
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,
)
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,
)
assert handled is True
assert captured["graph_gateway"] is graph_gateway
def test_needs_agent_awaits_loader_and_passes_result():
async 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()
@@ -182,16 +173,14 @@ def test_needs_agent_awaits_loader_and_passes_result():
) as mock_execute,
patch("EvoScientist.cli.channel._set_channel_response"),
):
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,
)
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,
)
assert handled is True
await_called.assert_called_once()
@@ -201,7 +190,7 @@ def test_needs_agent_awaits_loader_and_passes_result():
assert ctx_arg.agent == "ready-agent"
def test_await_agent_ready_failure_sets_error_response():
async def test_await_agent_ready_failure_sets_error_response():
msg = _make_msg()
fake_cmd = MagicMock()
fake_cmd.needs_agent.return_value = True
@@ -217,16 +206,14 @@ def test_await_agent_ready_failure_sets_error_response():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
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,
)
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,
)
assert handled is True
mock_set_resp.assert_called_once()
@@ -235,7 +222,7 @@ def test_await_agent_ready_failure_sets_error_response():
assert "agent blew up" in resp_text
def test_cmd_manager_raises_returns_true_with_error():
async 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()
@@ -253,15 +240,13 @@ def test_cmd_manager_raises_returns_true_with_error():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
handled = _run(
dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
)
handled = await 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()
@@ -270,7 +255,7 @@ def test_cmd_manager_raises_returns_true_with_error():
assert "boom" in resp_text
def test_on_cmd_completed_awaited_with_ctx_original_agent_and_cmd():
async 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."""
@@ -302,16 +287,14 @@ def test_on_cmd_completed_awaited_with_ctx_original_agent_and_cmd():
),
patch("EvoScientist.cli.channel._set_channel_response"),
):
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,
)
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,
)
assert handled is True
assert captured["ctx_agent"] == "swapped-agent"
@@ -319,7 +302,7 @@ def test_on_cmd_completed_awaited_with_ctx_original_agent_and_cmd():
assert captured["cmd_name"] == "/model"
def test_on_cmd_completed_receives_cmd_for_new_and_compact():
async 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``."""
@@ -343,21 +326,19 @@ def test_on_cmd_completed_receives_cmd_for_new_and_compact():
),
patch("EvoScientist.cli.channel._set_channel_response"),
):
_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,
)
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,
)
assert captured["cmd_name"] == cmd_name, cmd_name
def test_on_cmd_completed_skipped_on_fall_through_and_error():
async 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()
@@ -369,16 +350,14 @@ def test_on_cmd_completed_skipped_on_fall_through_and_error():
# Non-slash
with patch("EvoScientist.cli.channel._set_channel_response"):
_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,
)
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,
)
# Unresolved slash
with (
@@ -388,16 +367,14 @@ def test_on_cmd_completed_skipped_on_fall_through_and_error():
),
patch("EvoScientist.cli.channel._set_channel_response"),
):
_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,
)
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,
)
# Execute raises
with (
@@ -411,22 +388,20 @@ def test_on_cmd_completed_skipped_on_fall_through_and_error():
),
patch("EvoScientist.cli.channel._set_channel_response"),
):
_run(
dispatch_channel_slash_command(
_make_msg(),
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=_noop,
)
await 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()
def test_command_error_skips_completion_hook_and_reports_error():
async 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()
@@ -450,16 +425,14 @@ def test_command_error_skips_completion_hook_and_reports_error():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
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,
)
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,
)
assert handled is True
@@ -467,7 +440,7 @@ def test_command_error_skips_completion_hook_and_reports_error():
mock_set_resp.assert_called_once_with("msg-1", "Command error: workspace conflict")
def test_empty_command_error_still_reports_error():
async 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()
@@ -489,16 +462,14 @@ def test_empty_command_error_still_reports_error():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
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,
)
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,
)
assert handled is True
@@ -506,7 +477,7 @@ def test_empty_command_error_still_reports_error():
mock_set_resp.assert_called_once_with("msg-1", "Command error: (no details)")
def test_on_cmd_completed_exception_is_absorbed():
async 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()
@@ -526,23 +497,21 @@ def test_on_cmd_completed_exception_is_absorbed():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
handled = _run(
dispatch_channel_slash_command(
msg,
agent="orig",
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
on_cmd_completed=_boom,
)
handled = await 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]
def test_top_level_exception_is_absorbed():
async 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
@@ -557,15 +526,13 @@ def test_top_level_exception_is_absorbed():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
handled = _run(
dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=MagicMock(),
)
handled = await 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()
@@ -574,7 +541,7 @@ def test_top_level_exception_is_absorbed():
assert "exploded during resolve" in resp_text
def test_cmd_execute_returning_false_falls_through():
async 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="/")
@@ -592,15 +559,13 @@ def test_cmd_execute_returning_false_falls_through():
),
patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp,
):
handled = _run(
dispatch_channel_slash_command(
msg,
agent=None,
thread_id="t1",
workspace_dir=None,
checkpointer=None,
append_system=append,
)
handled = await 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()
+4 -7
View File
@@ -1,6 +1,5 @@
"""Tests for CLI interactive UI backend dispatch."""
import asyncio
from types import SimpleNamespace
import pytest
@@ -101,7 +100,7 @@ def test_background_agent_server_starts_even_when_async_subagents_disabled(
assert calls == [(config, "/tmp/workspace")]
def test_resume_workspace_sync_runs_even_when_async_subagents_disabled(
async def test_resume_workspace_sync_runs_even_when_async_subagents_disabled(
monkeypatch,
):
import EvoScientist.cli.commands as cmds
@@ -117,11 +116,9 @@ def test_resume_workspace_sync_runs_even_when_async_subagents_disabled(
)
config = SimpleNamespace(enable_async_subagents=False)
asyncio.run(
cmds._sync_background_agent_server_workspace(
config,
workspace_dir="/tmp/resumed-workspace",
)
await cmds._sync_background_agent_server_workspace(
config,
workspace_dir="/tmp/resumed-workspace",
)
assert calls == [(config, "/tmp/resumed-workspace")]
+292 -1
View File
@@ -1,4 +1,5 @@
"""Regression tests for the code_interpreter PTC allowlist.
"""Regression tests for the code_interpreter PTC allowlist and the
``EvoCodeInterpreterMiddleware`` subclass shape.
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
@@ -8,7 +9,10 @@ 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,
@@ -45,3 +49,290 @@ 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
+24 -27
View File
@@ -4,13 +4,12 @@ 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()
def _compact(
async def _compact(
graph_gateway: FakeGraphGateway,
*,
thread_id: str = "tid-1",
@@ -18,30 +17,28 @@ def _compact(
):
from EvoScientist.cli.commands import compact_conversation
return _run(
compact_conversation(
graph_gateway=graph_gateway,
thread_id=thread_id,
target=_TARGET,
input_tokens_hint=input_tokens_hint,
)
return await 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."""
def test_empty_messages(self):
async def test_empty_messages(self):
graph_gateway = FakeGraphGateway(state_values={"messages": []})
result = _compact(graph_gateway)
result = await _compact(graph_gateway)
assert result.status == "noop"
assert "no messages" in result.message
def test_state_read_failure(self):
async def test_state_read_failure(self):
graph_gateway = FakeGraphGateway(state_error=RuntimeError("DB gone"))
result = _compact(graph_gateway)
result = await _compact(graph_gateway)
assert result.status == "error"
assert "Failed to read state" in result.message
@@ -49,7 +46,7 @@ class TestCompactGuards:
class TestCompactCutoffZero:
"""When cutoff == 0, conversation is within retention budget."""
def test_nothing_to_compact_short_conversation(self):
async def test_nothing_to_compact_short_conversation(self):
msgs = [MagicMock() for _ in range(3)]
graph_gateway = FakeGraphGateway(state_values={"messages": msgs})
@@ -79,7 +76,7 @@ class TestCompactCutoffZero:
return_value=500,
),
):
result = _compact(graph_gateway)
result = await _compact(graph_gateway)
assert result.status == "noop"
assert "within the retention budget" in result.message
@@ -89,7 +86,7 @@ class TestCompactCutoffZero:
class TestCompactNegligibleSavings:
"""When cutoff > 0 but savings are too small to be worth it."""
def test_skip_when_few_messages_and_low_tokens(self):
async 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}
@@ -126,14 +123,14 @@ class TestCompactNegligibleSavings:
side_effect=lambda x: next(token_values),
),
):
result = _compact(graph_gateway)
result = await _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()
def test_still_compacts_when_few_messages_but_high_tokens(self):
async 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
@@ -178,7 +175,7 @@ class TestCompactNegligibleSavings:
side_effect=lambda x: next(token_values),
),
):
result = _compact(graph_gateway)
result = await _compact(graph_gateway)
assert result.status == "ok"
assert len(graph_gateway.updated_states) == 1
@@ -187,7 +184,7 @@ class TestCompactNegligibleSavings:
class TestCompactSuccess:
"""Normal compaction flow."""
def test_manual_threshold_blocks_low_context_compaction(self):
async 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}
@@ -217,7 +214,7 @@ class TestCompactSuccess:
return_value=30_000,
),
):
result = _compact(graph_gateway)
result = await _compact(graph_gateway)
assert result.status == "noop"
assert "40%" in result.message
@@ -225,7 +222,7 @@ class TestCompactSuccess:
mock_middleware_inst._determine_cutoff_index.assert_not_called()
mock_middleware_inst._acreate_summary.assert_not_called()
def test_successful_compaction(self):
async def test_successful_compaction(self):
from langchain_core.messages import HumanMessage
msgs = [MagicMock() for _ in range(20)]
@@ -273,7 +270,7 @@ class TestCompactSuccess:
side_effect=lambda x: next(token_values),
),
):
result = _compact(graph_gateway)
result = await _compact(graph_gateway)
assert result.status == "ok"
assert result.messages_compacted == 15
@@ -291,7 +288,7 @@ class TestCompactSuccess:
assert "_summarization_event" in event_data
assert event_data["_summarization_event"]["cutoff_index"] == 15
def test_offload_failure_non_fatal(self):
async def test_offload_failure_non_fatal(self):
"""Offload failure should not prevent compaction."""
from langchain_core.messages import HumanMessage
@@ -335,7 +332,7 @@ class TestCompactSuccess:
return_value=1000,
),
):
result = _compact(graph_gateway)
result = await _compact(graph_gateway)
assert result.status == "ok"
assert len(graph_gateway.updated_states) == 1
@@ -377,7 +374,7 @@ class TestRenderCompactResult:
class TestCompactCommandUI:
"""TUI-specific compact progress indicator behavior."""
def test_command_uses_tui_indicator_when_available(self):
async 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
@@ -413,7 +410,7 @@ class TestCompactCommandUI:
return_value="summary-panel",
),
):
_run(CompactCommand().execute(ctx, []))
await CompactCommand().execute(ctx, [])
assert ui.started == 1
assert ui.stopped == 1
+162 -6
View File
@@ -52,12 +52,9 @@ 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",
@@ -77,6 +74,9 @@ 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,6 +104,9 @@ 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)
@@ -129,8 +132,13 @@ class TestEvoScientistConfig:
assert config.show_thinking is True
assert config.ui_backend == "tui"
assert config.log_level == "warning"
assert config.reasoning_effort == "high"
assert config.reasoning_effort == ""
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
@@ -145,6 +153,8 @@ 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."""
@@ -192,6 +202,18 @@ 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
@@ -199,14 +221,34 @@ 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"
@@ -242,6 +284,23 @@ 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(
@@ -655,6 +714,22 @@ 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())
@@ -715,6 +790,60 @@ 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.
@@ -832,3 +961,30 @@ 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
+6 -7
View File
@@ -15,7 +15,6 @@ from EvoScientist.middleware.configurable_model import (
ConfigurableModelMiddleware,
_read_model_override,
)
from tests.conftest import run_async as _run
@contextmanager
@@ -134,7 +133,7 @@ class TestPassThrough:
handler.assert_called_once_with(req)
req.override.assert_not_called()
def test_async_no_override_passes_request_unchanged(self):
async def test_async_no_override_passes_request_unchanged(self):
mw = ConfigurableModelMiddleware()
req = _make_request()
@@ -143,7 +142,7 @@ class TestPassThrough:
return "ok"
with _patched_config({}):
result = _run(mw.awrap_model_call(req, handler))
result = await mw.awrap_model_call(req, handler)
assert result == "ok"
req.override.assert_not_called()
@@ -185,7 +184,7 @@ class TestModelOverride:
assert called_with is not req
assert called_with.model is new_model
def test_async_override_path_parity(self):
async def test_async_override_path_parity(self):
mw = ConfigurableModelMiddleware()
req = _make_request()
new_model = MagicMock()
@@ -202,7 +201,7 @@ class TestModelOverride:
"EvoScientist.llm.get_chat_model", return_value=new_model
) as mock_get,
):
result = _run(mw.awrap_model_call(req, handler))
result = await mw.awrap_model_call(req, handler)
assert result == "ok"
mock_get.assert_called_once_with(model="claude-opus-4-8", provider="anthropic")
@@ -316,7 +315,7 @@ class TestResolveFailure:
handler.assert_called_once_with(req)
req.override.assert_not_called()
def test_async_falls_back_when_resolve_raises(self):
async def test_async_falls_back_when_resolve_raises(self):
mw = ConfigurableModelMiddleware()
req = _make_request()
@@ -333,7 +332,7 @@ class TestResolveFailure:
side_effect=ValueError("unknown model"),
),
):
result = _run(mw.awrap_model_call(req, handler))
result = await mw.awrap_model_call(req, handler)
assert result == "ok"
assert called == [req]
@@ -66,7 +66,6 @@ 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)]
@@ -91,7 +90,6 @@ 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=[],
+4 -6
View File
@@ -2,11 +2,9 @@
from unittest.mock import MagicMock
from tests.conftest import run_async as _run
class TestCurrentCommand:
def test_prints_thread_workspace_and_memory(self):
async def test_prints_thread_workspace_and_memory(self):
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.general import CurrentCommand
@@ -17,14 +15,14 @@ class TestCurrentCommand:
ui=ui,
workspace_dir="/tmp/ws",
)
_run(CurrentCommand().execute(ctx, []))
await 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)
def test_skips_workspace_when_none(self):
async def test_skips_workspace_when_none(self):
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.general import CurrentCommand
@@ -35,7 +33,7 @@ class TestCurrentCommand:
ui=ui,
workspace_dir=None,
)
_run(CurrentCommand().execute(ctx, []))
await 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)
+14 -15
View File
@@ -2,7 +2,6 @@
from unittest.mock import AsyncMock, MagicMock
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway, FakeThreadStore
@@ -21,17 +20,17 @@ def _ctx(thread_id="current", thread_store=None):
class TestDeleteCommand:
def test_refuses_to_delete_current(self):
async 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)
_run(DeleteCommand().execute(ctx, ["current"]))
await 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)
def test_happy_path_success(self):
async def test_happy_path_success(self):
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx(
@@ -41,45 +40,45 @@ class TestDeleteCommand:
deleted=True,
),
)
_run(DeleteCommand().execute(ctx, ["other-thread"]))
await 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)
def test_not_found(self):
async def test_not_found(self):
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx()
_run(DeleteCommand().execute(ctx, ["missing"]))
await 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)
def test_ambiguous_prefix(self):
async def test_ambiguous_prefix(self):
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx(thread_store=FakeThreadStore(matches=["abc-one", "abc-two"]))
_run(DeleteCommand().execute(ctx, ["abc"]))
await 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)
def test_prefix_resolves_to_unique_match(self):
async 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)
)
_run(DeleteCommand().execute(ctx, ["abc"]))
await 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)
def test_no_arg_empty_sessions_prints_notice(self):
async def test_no_arg_empty_sessions_prints_notice(self):
from EvoScientist.commands.implementation.session import DeleteCommand
ctx, ui = _ctx()
_run(DeleteCommand().execute(ctx, []))
await 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)
def test_no_arg_calls_picker_returns_none(self):
async 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
@@ -96,5 +95,5 @@ class TestDeleteCommand:
]
store = FakeThreadStore(threads=threads)
ctx.graph_gateway = FakeGraphGateway(thread_store=store)
_run(DeleteCommand().execute(ctx, []))
await DeleteCommand().execute(ctx, [])
ui.wait_for_thread_pick.assert_awaited_once()
+32 -33
View File
@@ -7,7 +7,6 @@ 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:
@@ -41,30 +40,30 @@ class TestDingTalkChannel:
assert channel._running is False
assert channel.name == "dingtalk"
def test_start_raises_without_credentials(self):
async 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"):
_run(channel.start())
await channel.start()
def test_start_raises_without_client_id(self):
async 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"):
_run(channel.start())
await channel.start()
def test_start_raises_without_client_secret(self):
async 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"):
_run(channel.start())
await channel.start()
def test_stop_when_not_running(self):
async def test_stop_when_not_running(self):
config = DingTalkConfig(client_id="test-id", client_secret="test-secret")
channel = DingTalkChannel(config)
_run(channel.stop())
await channel.stop()
def test_send_returns_false_without_client(self):
async def test_send_returns_false_without_client(self):
config = DingTalkConfig(client_id="test-id", client_secret="test-secret")
channel = DingTalkChannel(config)
msg = OutboundMessage(
@@ -73,7 +72,7 @@ class TestDingTalkChannel:
content="hello",
metadata={"chat_id": "user123"},
)
result = _run(channel.send(msg))
result = await channel.send(msg)
assert result is False
def test_capabilities(self):
@@ -130,20 +129,20 @@ class TestDingTalkWsMessageParsing:
channel._token_expires = 9999999999
return channel
def test_system_ping_ack(self):
async def test_system_ping_ack(self):
channel = self._make_channel()
data = {
"type": "SYSTEM",
"headers": {"topic": "ping", "messageId": "ping-1"},
"data": "pong-data",
}
_run(channel._on_ws_message(data))
await 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"
def test_callback_text_message(self):
async def test_callback_text_message(self):
channel = self._make_channel()
channel._enqueue_raw = AsyncMock()
@@ -158,14 +157,14 @@ class TestDingTalkWsMessageParsing:
"headers": {"messageId": "msg-1", "contentType": "application/json"},
"data": json.dumps(payload),
}
_run(channel._on_ws_message(data))
await 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
def test_callback_group_message_mention(self):
async def test_callback_group_message_mention(self):
channel = self._make_channel()
channel._enqueue_raw = AsyncMock()
@@ -181,12 +180,12 @@ class TestDingTalkWsMessageParsing:
"headers": {"messageId": "msg-2"},
"data": json.dumps(payload),
}
_run(channel._on_ws_message(data))
await 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
def test_callback_group_no_mention(self):
async def test_callback_group_no_mention(self):
channel = self._make_channel()
channel._enqueue_raw = AsyncMock()
@@ -201,12 +200,12 @@ class TestDingTalkWsMessageParsing:
"headers": {"messageId": "msg-3"},
"data": json.dumps(payload),
}
_run(channel._on_ws_message(data))
await 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
def test_ignores_non_callback(self):
async def test_ignores_non_callback(self):
channel = self._make_channel()
channel._enqueue_raw = AsyncMock()
@@ -215,10 +214,10 @@ class TestDingTalkWsMessageParsing:
"headers": {"messageId": "msg-x"},
"data": "{}",
}
_run(channel._on_ws_message(data))
await channel._on_ws_message(data)
channel._enqueue_raw.assert_not_called()
def test_ignores_empty_content(self):
async def test_ignores_empty_content(self):
channel = self._make_channel()
channel._enqueue_raw = AsyncMock()
@@ -232,20 +231,20 @@ class TestDingTalkWsMessageParsing:
"headers": {"messageId": "msg-e"},
"data": json.dumps(payload),
}
_run(channel._on_ws_message(data))
await channel._on_ws_message(data)
channel._enqueue_raw.assert_not_called()
def test_non_dict_data_ignored(self):
async def test_non_dict_data_ignored(self):
channel = self._make_channel()
channel._enqueue_raw = AsyncMock()
_run(channel._on_ws_message("not a dict"))
await channel._on_ws_message("not a dict")
channel._enqueue_raw.assert_not_called()
class TestDingTalkSendChunk:
"""Test _send_chunk with mocked HTTP client."""
def test_send_chunk_calls_api(self):
async 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"
@@ -256,7 +255,7 @@ class TestDingTalkSendChunk:
channel._http_client = MagicMock()
channel._http_client.post = AsyncMock(return_value=mock_response)
_run(channel._send_chunk("user1", "formatted", "raw text", None, {}))
await 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")
@@ -273,21 +272,21 @@ class TestDingTalkChannelRegistration:
class TestDingTalkProbe:
def test_missing_credentials(self):
async def test_missing_credentials(self):
from EvoScientist.channels.dingtalk.probe import validate_dingtalk
ok, msg = _run(validate_dingtalk("", ""))
ok, msg = await validate_dingtalk("", "")
assert ok is False
assert "required" in msg
def test_missing_client_id(self):
async def test_missing_client_id(self):
from EvoScientist.channels.dingtalk.probe import validate_dingtalk
ok, _msg = _run(validate_dingtalk("", "secret"))
ok, _msg = await validate_dingtalk("", "secret")
assert ok is False
def test_missing_client_secret(self):
async def test_missing_client_secret(self):
from EvoScientist.channels.dingtalk.probe import validate_dingtalk
ok, _msg = _run(validate_dingtalk("id", ""))
ok, _msg = await validate_dingtalk("id", "")
assert ok is False
+6 -7
View File
@@ -4,7 +4,6 @@ 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:
@@ -14,18 +13,18 @@ class TestDiscordChannel:
assert channel.config is config
assert channel._running is False
def test_start_raises_without_token_or_library(self):
async def test_start_raises_without_token_or_library(self):
config = DiscordConfig(bot_token="")
channel = DiscordChannel(config)
with pytest.raises(ChannelError):
_run(channel.start())
await channel.start()
def test_stop_when_not_running(self):
async def test_stop_when_not_running(self):
config = DiscordConfig(bot_token="test")
channel = DiscordChannel(config)
_run(channel.stop())
await channel.stop()
def test_send_returns_false_without_client(self):
async def test_send_returns_false_without_client(self):
from EvoScientist.channels.base import OutboundMessage
config = DiscordConfig(bot_token="test")
@@ -36,5 +35,5 @@ class TestDiscordChannel:
content="hello",
metadata={"chat_id": "123"},
)
result = _run(channel.send(msg))
result = await channel.send(msg)
assert result is False
@@ -0,0 +1,440 @@
"""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"
+8 -10
View File
@@ -2,8 +2,6 @@
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
@@ -31,7 +29,7 @@ _INDEX = [
class TestInstallSkills:
def test_picker_cancel_no_install(self):
async def test_picker_cancel_no_install(self):
from EvoScientist.commands.implementation.skills import InstallSkills
ctx, ui = _ctx()
@@ -45,12 +43,12 @@ class TestInstallSkills:
"EvoScientist.tools.skills_manager.install_skill",
) as install_mock,
):
_run(InstallSkills().execute(ctx, []))
await 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)
def test_picker_returns_selections_installs_each(self):
async def test_picker_returns_selections_installs_each(self):
from EvoScientist.commands.implementation.skills import InstallSkills
ctx, ui = _ctx()
@@ -68,10 +66,10 @@ class TestInstallSkills:
return_value={"success": True, "name": "x"},
) as install_mock,
):
_run(InstallSkills().execute(ctx, []))
await InstallSkills().execute(ctx, [])
assert install_mock.call_count == 2
def test_channel_auto_install_on_tag(self):
async def test_channel_auto_install_on_tag(self):
"""Non-interactive UI + tag arg → auto-installs matching skills."""
from EvoScientist.commands.implementation.skills import InstallSkills
@@ -86,12 +84,12 @@ class TestInstallSkills:
return_value={"success": True, "name": "x"},
) as install_mock,
):
_run(InstallSkills().execute(ctx, ["core"]))
await 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()
def test_fetch_failure_prints_error(self):
async def test_fetch_failure_prints_error(self):
from EvoScientist.commands.implementation.skills import InstallSkills
ctx, ui = _ctx()
@@ -99,6 +97,6 @@ class TestInstallSkills:
"EvoScientist.tools.skills_manager.fetch_remote_skill_index",
side_effect=RuntimeError("network fail"),
):
_run(InstallSkills().execute(ctx, []))
await 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 -4
View File
@@ -2,11 +2,9 @@
from unittest.mock import MagicMock
from tests.conftest import run_async as _run
class TestExitCommand:
def test_execute_calls_force_quit(self):
async def test_execute_calls_force_quit(self):
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.session import ExitCommand
@@ -17,7 +15,7 @@ class TestExitCommand:
ui=ui,
)
cmd = ExitCommand()
_run(cmd.execute(ctx, []))
await cmd.execute(ctx, [])
ui.force_quit.assert_called_once()
def test_aliases_registered(self):
+38 -39
View File
@@ -14,7 +14,6 @@ from EvoScientist.channels.feishu.channel import (
_parse_inline_elements,
_parse_inline_text,
)
from tests.conftest import run_async as _run
class TestFeishuConfig:
@@ -57,24 +56,24 @@ class TestFeishuChannel:
assert channel._running is False
assert channel.name == "feishu"
def test_start_raises_without_app_id(self):
async 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"):
_run(channel.start())
await channel.start()
def test_start_raises_without_app_secret(self):
async 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"):
_run(channel.start())
await channel.start()
def test_stop_when_not_running(self):
async def test_stop_when_not_running(self):
config = FeishuConfig(app_id="test-id", app_secret="test-secret")
channel = FeishuChannel(config)
_run(channel.stop())
await channel.stop()
def test_send_returns_false_without_client(self):
async def test_send_returns_false_without_client(self):
config = FeishuConfig(app_id="test-id", app_secret="test-secret")
channel = FeishuChannel(config)
msg = OutboundMessage(
@@ -83,7 +82,7 @@ class TestFeishuChannel:
content="hello",
metadata={"chat_id": "oc_test"},
)
result = _run(channel.send(msg))
result = await channel.send(msg)
assert result is False
def test_capabilities(self):
@@ -201,7 +200,7 @@ class TestFeishuWebhookEvent:
channel._enqueue_raw = AsyncMock()
return channel
def test_text_message_v2(self):
async def test_text_message_v2(self):
channel = self._make_channel()
event = {
"sender": {
@@ -217,7 +216,7 @@ class TestFeishuWebhookEvent:
"create_time": "1700000000000",
},
}
_run(channel._on_message(event))
await channel._on_message(event)
channel._enqueue_raw.assert_called_once()
raw = channel._enqueue_raw.call_args[0][0]
assert raw.text == "hello feishu"
@@ -225,7 +224,7 @@ class TestFeishuWebhookEvent:
assert raw.chat_id == "oc_chat1"
assert raw.is_group is False
def test_group_message_with_mention(self):
async def test_group_message_with_mention(self):
channel = self._make_channel()
event = {
"sender": {
@@ -242,13 +241,13 @@ class TestFeishuWebhookEvent:
"mentions": [{"key": "@_user_1", "id": {}}],
},
}
_run(channel._on_message(event))
await 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"]
def test_group_message_no_mention(self):
async def test_group_message_no_mention(self):
channel = self._make_channel()
event = {
"sender": {
@@ -264,12 +263,12 @@ class TestFeishuWebhookEvent:
"create_time": "1700000000000",
},
}
_run(channel._on_message(event))
await channel._on_message(event)
raw = channel._enqueue_raw.call_args[0][0]
assert raw.is_group is True
assert raw.was_mentioned is False
def test_skips_bot_messages(self):
async def test_skips_bot_messages(self):
channel = self._make_channel()
event = {
"sender": {
@@ -283,10 +282,10 @@ class TestFeishuWebhookEvent:
"content": json.dumps({"text": "bot reply"}),
},
}
_run(channel._on_message(event))
await channel._on_message(event)
channel._enqueue_raw.assert_not_called()
def test_post_message(self):
async def test_post_message(self):
channel = self._make_channel()
post_content = {
"zh_cn": {
@@ -308,12 +307,12 @@ class TestFeishuWebhookEvent:
"create_time": "1700000000000",
},
}
_run(channel._on_message(event))
await channel._on_message(event)
raw = channel._enqueue_raw.call_args[0][0]
assert "Test" in raw.text
assert "Post body" in raw.text
def test_unsupported_msg_type_annotation(self):
async def test_unsupported_msg_type_annotation(self):
channel = self._make_channel()
event = {
"sender": {
@@ -329,7 +328,7 @@ class TestFeishuWebhookEvent:
"create_time": "1700000000000",
},
}
_run(channel._on_message(event))
await channel._on_message(event)
raw = channel._enqueue_raw.call_args[0][0]
assert "share_chat" in raw.text
@@ -337,7 +336,7 @@ class TestFeishuWebhookEvent:
class TestFeishuSendChunk:
"""Test _send_chunk with mocked HTTP client."""
def test_send_chunk_post_format(self):
async 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"
@@ -348,14 +347,14 @@ class TestFeishuSendChunk:
channel._http_client = MagicMock()
channel._http_client.post = AsyncMock(return_value=mock_response)
_run(channel._send_chunk("oc_chat1", "formatted", "raw **text**", None, {}))
await 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"
def test_send_chunk_with_reply(self):
async 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"
@@ -366,7 +365,7 @@ class TestFeishuSendChunk:
channel._http_client = MagicMock()
channel._http_client.post = AsyncMock(return_value=mock_response)
_run(channel._send_chunk("oc_chat1", "reply", "reply text", "om_reply_id", {}))
await 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
@@ -484,17 +483,17 @@ class TestFeishuChannelRegistration:
class TestFeishuProbe:
def test_missing_app_id(self):
async def test_missing_app_id(self):
from EvoScientist.channels.feishu.probe import validate_feishu_credentials
ok, msg = _run(validate_feishu_credentials("", "secret"))
ok, msg = await validate_feishu_credentials("", "secret")
assert ok is False
assert "app_id" in msg
def test_missing_app_secret(self):
async def test_missing_app_secret(self):
from EvoScientist.channels.feishu.probe import validate_feishu_credentials
ok, msg = _run(validate_feishu_credentials("id", ""))
ok, msg = await validate_feishu_credentials("id", "")
assert ok is False
assert "app_secret" in msg
@@ -510,7 +509,7 @@ class TestFeishuWebSocketMode:
)
assert config.subscription_mode == "websocket"
def test_start_websocket_raises_without_lark_oapi(self):
async def test_start_websocket_raises_without_lark_oapi(self):
config = FeishuConfig(
app_id="test-id",
app_secret="test-secret",
@@ -520,9 +519,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"):
_run(channel.start())
await channel.start()
def test_start_webhook_mode_still_works(self):
async def test_start_webhook_mode_still_works(self):
"""Ensure subscription_mode='webhook' still validates as before."""
config = FeishuConfig(
app_id="",
@@ -531,9 +530,9 @@ class TestFeishuWebSocketMode:
)
channel = FeishuChannel(config)
with pytest.raises(ChannelError, match="app_id"):
_run(channel.start())
await channel.start()
def test_invalid_subscription_mode_raises(self):
async def test_invalid_subscription_mode_raises(self):
config = FeishuConfig(
app_id="test-id",
app_secret="test-secret",
@@ -541,9 +540,9 @@ class TestFeishuWebSocketMode:
)
channel = FeishuChannel(config)
with pytest.raises(ChannelError, match="Invalid feishu_subscription_mode"):
_run(channel.start())
await channel.start()
def test_on_lark_sdk_message_bridges_to_on_message(self):
async 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
@@ -594,14 +593,14 @@ class TestFeishuWebSocketMode:
)
# Verify the consumer processes it correctly
_run(channel._on_message(event_dict))
await 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
def test_cleanup_websocket_mode(self):
async def test_cleanup_websocket_mode(self):
config = FeishuConfig(
app_id="test-id",
app_secret="test-secret",
@@ -617,7 +616,7 @@ class TestFeishuWebSocketMode:
channel._ws_consumer_task = None
channel._access_token = "fake-token"
_run(channel._cleanup())
await channel._cleanup()
mock_client.aclose.assert_called_once()
assert channel._http_client is None
+31 -41
View File
@@ -1,6 +1,5 @@
from __future__ import annotations
import asyncio
from types import SimpleNamespace
from unittest.mock import MagicMock
@@ -179,7 +178,7 @@ def test_launch_background_run_deletes_thread_when_run_creation_fails(monkeypatc
fake_client.threads.delete.assert_called_once_with("thread-1")
def test_async_launch_background_run_deletes_thread_when_run_creation_fails(
async def test_async_launch_background_run_deletes_thread_when_run_creation_fails(
monkeypatch,
):
monkeypatch.setattr(
@@ -204,11 +203,8 @@ def test_async_launch_background_run_deletes_thread_when_run_creation_fails(
lambda **_kwargs: SimpleNamespace(threads=_Threads(), runs=_Runs()),
)
async def run() -> None:
with pytest.raises(RuntimeError, match="run creation failed"):
await background_runs.alaunch_background_run(_request())
asyncio.run(run())
with pytest.raises(RuntimeError, match="run creation failed"):
await background_runs.alaunch_background_run(_request())
assert deleted == ["thread-1"]
@@ -278,7 +274,7 @@ def test_sync_status_watcher_preserves_thread_on_poll_failure(
assert deleted == []
def test_async_status_watcher_aborts_and_deletes_thread_on_error_status():
async 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] = []
@@ -291,29 +287,26 @@ def test_async_status_watcher_aborts_and_deletes_thread_on_error_status():
async def delete(self, thread_id: str):
deleted.append(thread_id)
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())
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,
),
)
assert finished == []
assert [run.run_id for run in aborted] == ["run-1"]
assert deleted == ["thread-1"]
def test_async_status_watcher_preserves_run_url():
async def test_async_status_watcher_preserves_run_url():
finished: list[background_runs.BackgroundRun] = []
class _Runs:
@@ -324,21 +317,18 @@ def test_async_status_watcher_preserves_run_url():
async def delete(self, _thread_id: str):
return None
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())
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,
),
)
assert [run.url for run in finished] == ["http://worker.example"]
+72 -75
View File
@@ -20,7 +20,6 @@ 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,
@@ -30,7 +29,7 @@ from tests.fakes import (
)
def test_local_gateway_streams_from_injected_streamer():
async def test_local_gateway_streams_from_injected_streamer():
seen: dict[str, Any] = {}
async def _streamer(agent, message, thread_id, **kwargs):
@@ -60,7 +59,7 @@ 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 = run_async(_collect())
events = await _collect()
assert events == [
{"type": "text", "content": "hi"},
@@ -75,7 +74,7 @@ def test_local_gateway_streams_from_injected_streamer():
}
def test_local_graph_gateway_delegates_thread_operations():
async def test_local_graph_gateway_delegates_thread_operations():
thread_store = FakeThreadStore(
generated_thread_id="new12345",
threads=[{"thread_id": "abc12345"}],
@@ -102,7 +101,7 @@ def test_local_graph_gateway_delegates_thread_operations():
"deleted": await gateway.delete_thread("abc12345"),
}
result = run_async(_run())
result = await _run()
assert result["created"] == "new12345"
assert result["threads"] == [{"thread_id": "abc12345"}]
@@ -132,16 +131,14 @@ def test_local_graph_gateway_delegates_thread_operations():
]
def test_local_graph_gateway_reads_state_values():
async 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 = run_async(
gateway.get_state_values(GraphTarget(local_graph=agent), "abc12345")
)
values = await gateway.get_state_values(GraphTarget(local_graph=agent), "abc12345")
assert values == {"async_tasks": {"task-1": {}}}
agent.aget_state.assert_awaited_once_with(
@@ -149,17 +146,15 @@ def test_local_graph_gateway_reads_state_values():
)
def test_local_graph_gateway_updates_state_values():
async def test_local_graph_gateway_updates_state_values():
agent = MagicMock()
agent.aupdate_state = AsyncMock()
gateway = LocalGraphGateway()
run_async(
gateway.update_state_values(
GraphTarget(local_graph=agent),
"abc12345",
{"_summarization_event": {"cutoff_index": 2}},
)
await gateway.update_state_values(
GraphTarget(local_graph=agent),
"abc12345",
{"_summarization_event": {"cutoff_index": 2}},
)
agent.aupdate_state.assert_awaited_once_with(
@@ -169,7 +164,7 @@ def test_local_graph_gateway_updates_state_values():
)
def test_local_stream_events_delegates_aclose_to_inner():
async def test_local_stream_events_delegates_aclose_to_inner():
cleanup_ran = False
async def _streamer(_agent, _message, _thread_id, **_kwargs):
@@ -194,7 +189,7 @@ def test_local_stream_events_delegates_aclose_to_inner():
assert cleanup_ran is True
with patch("EvoScientist.stream.events.stream_agent_events", new=_streamer):
run_async(_run())
await _run()
def test_run_streaming_can_consume_injected_gateway():
@@ -228,7 +223,7 @@ def test_run_streaming_can_consume_injected_gateway():
]
def test_resume_command_consumes_context_gateway():
async def test_resume_command_consumes_context_gateway():
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.session import ResumeCommand
@@ -246,7 +241,7 @@ def test_resume_command_consumes_context_gateway():
graph_gateway=FakeGraphGateway(thread_store=thread_store),
)
run_async(ResumeCommand().execute(ctx, ["abc"]))
await ResumeCommand().execute(ctx, ["abc"])
assert ctx.thread_id == "abc12345"
assert ctx.workspace_dir == "/restored"
@@ -287,7 +282,7 @@ def test_cmd_run_passes_local_graph_gateway(monkeypatch):
assert seen["gateway"].thread_store is thread_store
def test_langgraph_server_thread_store_delegates_to_sdk_threads():
async def test_langgraph_server_thread_store_delegates_to_sdk_threads():
threads = FakeLangGraphThreadsClient(
threads=[
{
@@ -335,7 +330,7 @@ def test_langgraph_server_thread_store_delegates_to_sdk_threads():
"deleted": await store.delete_thread("abc12345"),
}
result = run_async(_run())
result = await _run()
assert result["created"] == "server-thread"
assert len(threads.created) == 1
@@ -379,7 +374,7 @@ def test_langgraph_server_thread_store_delegates_to_sdk_threads():
assert threads.deleted == ["abc12345"]
def test_langgraph_server_thread_store_limit_zero_pages_all_threads():
async def test_langgraph_server_thread_store_limit_zero_pages_all_threads():
rows = [
{
"thread_id": f"thread-{index}",
@@ -392,7 +387,7 @@ def test_langgraph_server_thread_store_limit_zero_pages_all_threads():
client=FakeLangGraphClient(threads),
)
result = run_async(store.list_threads(limit=0))
result = await store.list_threads(limit=0)
assert [row["thread_id"] for row in result] == [
f"thread-{index}" for index in range(_THREAD_SEARCH_LIMIT + 1)
@@ -403,7 +398,7 @@ def test_langgraph_server_thread_store_limit_zero_pages_all_threads():
]
def test_langgraph_server_thread_store_positive_limit_uses_single_search():
async def test_langgraph_server_thread_store_positive_limit_uses_single_search():
threads = FakeLangGraphThreadsClient(
threads=[
{
@@ -417,7 +412,7 @@ def test_langgraph_server_thread_store_positive_limit_uses_single_search():
client=FakeLangGraphClient(threads),
)
result = run_async(store.list_threads(limit=2))
result = await 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] == [
@@ -425,7 +420,7 @@ def test_langgraph_server_thread_store_positive_limit_uses_single_search():
]
def test_langgraph_server_thread_store_prefix_resolution_skips_exact_lookup():
async def test_langgraph_server_thread_store_prefix_resolution_skips_exact_lookup():
threads = FakeLangGraphThreadsClient(
threads=[
{
@@ -438,14 +433,14 @@ def test_langgraph_server_thread_store_prefix_resolution_skips_exact_lookup():
client=FakeLangGraphClient(threads),
)
result = run_async(store.resolve_thread_id_prefix("abc"))
result = await store.resolve_thread_id_prefix("abc")
assert result == ("abc12345", [])
assert threads.gets == []
assert len(threads.searches) == 1
def test_langgraph_server_thread_store_prefix_resolution_pages_all_threads():
async def test_langgraph_server_thread_store_prefix_resolution_pages_all_threads():
rows = [
{
"thread_id": f"thread-{index}",
@@ -464,7 +459,7 @@ def test_langgraph_server_thread_store_prefix_resolution_pages_all_threads():
client=FakeLangGraphClient(threads),
)
result = run_async(store.resolve_thread_id_prefix("older-thread"))
result = await store.resolve_thread_id_prefix("older-thread")
assert result == ("older-thread-match", [])
assert [(search["limit"], search["offset"]) for search in threads.searches] == [
@@ -473,7 +468,7 @@ def test_langgraph_server_thread_store_prefix_resolution_pages_all_threads():
]
def test_langgraph_server_thread_store_uuid_resolution_uses_exact_lookup():
async def test_langgraph_server_thread_store_uuid_resolution_uses_exact_lookup():
thread_id = "019ed9e4-4253-7f62-b50f-f0470a4b3c9f"
threads = FakeLangGraphThreadsClient(
threads=[
@@ -487,14 +482,14 @@ def test_langgraph_server_thread_store_uuid_resolution_uses_exact_lookup():
client=FakeLangGraphClient(threads),
)
result = run_async(store.resolve_thread_id_prefix(thread_id))
result = await store.resolve_thread_id_prefix(thread_id)
assert result == (thread_id, [])
assert threads.gets == [thread_id]
assert threads.searches == []
def test_langgraph_server_thread_store_uuid_resolution_filters_graph_id():
async def test_langgraph_server_thread_store_uuid_resolution_filters_graph_id():
thread_id = "019ed9e4-4253-7f62-b50f-f0470a4b3c9f"
threads = FakeLangGraphThreadsClient(
threads=[
@@ -508,7 +503,7 @@ def test_langgraph_server_thread_store_uuid_resolution_filters_graph_id():
client=FakeLangGraphClient(threads),
)
result = run_async(store.resolve_thread_id_prefix(thread_id))
result = await store.resolve_thread_id_prefix(thread_id)
assert result == (None, [])
assert threads.gets == [thread_id]
@@ -517,7 +512,7 @@ def test_langgraph_server_thread_store_uuid_resolution_filters_graph_id():
]
def test_langgraph_server_thread_store_clones_thread_with_metadata():
async def test_langgraph_server_thread_store_clones_thread_with_metadata():
clone_metadata = {
"clone_purpose": "memory_extraction",
"source_thread_id": "source-thread",
@@ -534,8 +529,8 @@ def test_langgraph_server_thread_store_clones_thread_with_metadata():
client=FakeLangGraphClient(threads),
)
cloned_thread_id = run_async(
store.clone_thread("source-thread", metadata=clone_metadata)
cloned_thread_id = await store.clone_thread(
"source-thread", metadata=clone_metadata
)
assert cloned_thread_id == "source-thread-copy"
@@ -552,7 +547,7 @@ def test_langgraph_server_thread_store_clones_thread_with_metadata():
}
def test_langgraph_server_thread_store_rejects_copy_without_thread_id():
async 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,
@@ -565,10 +560,10 @@ 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"):
run_async(_run())
await _run()
def test_langgraph_server_gateway_clones_thread():
async def test_langgraph_server_gateway_clones_thread():
threads = FakeLangGraphThreadsClient(
threads=[{"thread_id": "source-thread", "metadata": {"graph_id": "agent"}}]
)
@@ -578,12 +573,10 @@ def test_langgraph_server_gateway_clones_thread():
)
)
cloned_thread_id = run_async(
gateway.clone_thread(
"source-thread",
metadata={"clone_purpose": "manual"},
target=GraphTarget(graph_id="agent"),
)
cloned_thread_id = await gateway.clone_thread(
"source-thread",
metadata={"clone_purpose": "manual"},
target=GraphTarget(graph_id="agent"),
)
assert cloned_thread_id == "source-thread-copy"
@@ -592,12 +585,12 @@ def test_langgraph_server_gateway_clones_thread():
]
def test_local_graph_gateway_clone_thread_is_explicitly_unsupported():
async 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"):
run_async(_run())
await _run()
def test_runtime_gateways_can_use_langgraph_server_backend():
@@ -616,7 +609,7 @@ def test_runtime_gateways_can_use_langgraph_server_backend():
assert gateway.thread_store is runtime_gateways.thread_store
def test_langgraph_server_gateway_reads_state_values():
async 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": {}}}}},
@@ -627,12 +620,12 @@ def test_langgraph_server_gateway_reads_state_values():
)
)
values = run_async(gateway.get_state_values(GraphTarget(), "abc12345"))
values = await gateway.get_state_values(GraphTarget(), "abc12345")
assert values == {"async_tasks": {"task-1": {}}}
def test_langgraph_server_gateway_messages_apply_summarization_event():
async def test_langgraph_server_gateway_messages_apply_summarization_event():
threads = FakeLangGraphThreadsClient(
threads=[{"thread_id": "abc12345", "metadata": {"graph_id": "EvoScientist"}}],
states={
@@ -658,7 +651,7 @@ def test_langgraph_server_gateway_messages_apply_summarization_event():
)
)
messages = run_async(gateway.get_thread_messages("abc12345"))
messages = await gateway.get_thread_messages("abc12345")
assert len(messages) == 2
assert isinstance(messages[0], AIMessage)
@@ -667,7 +660,7 @@ def test_langgraph_server_gateway_messages_apply_summarization_event():
assert messages[1].content == "third"
def test_langgraph_server_gateway_updates_state_values():
async def test_langgraph_server_gateway_updates_state_values():
threads = FakeLangGraphThreadsClient(
threads=[{"thread_id": "abc12345", "metadata": {"graph_id": "EvoScientist"}}],
)
@@ -677,12 +670,10 @@ def test_langgraph_server_gateway_updates_state_values():
)
)
run_async(
gateway.update_state_values(
GraphTarget(),
"abc12345",
{"_summarization_event": {"cutoff_index": 2}},
)
await gateway.update_state_values(
GraphTarget(),
"abc12345",
{"_summarization_event": {"cutoff_index": 2}},
)
assert threads.state_updates == [
@@ -690,7 +681,7 @@ def test_langgraph_server_gateway_updates_state_values():
]
def test_langgraph_server_gateway_streams_root_protocol_events():
async def test_langgraph_server_gateway_streams_root_protocol_events():
stream = FakeLangGraphThreadStream(
"abc12345",
events=[
@@ -737,7 +728,7 @@ def test_langgraph_server_gateway_streams_root_protocol_events():
)
]
events = run_async(_collect())
events = await _collect()
assert len(threads.created) == 1
assert threads.created[0]["thread_id"] == "abc12345"
@@ -804,7 +795,7 @@ def _root_message_finish() -> dict[str, object]:
}
def _collect_server_gateway_stream(
async def _collect_server_gateway_stream(
events: list[dict[str, object]],
*,
state_messages: list[dict[str, object]] | None = None,
@@ -832,11 +823,11 @@ def _collect_server_gateway_stream(
)
]
return run_async(_collect())
return await _collect()
def test_langgraph_server_gateway_streams_value_message_snapshots():
events = _collect_server_gateway_stream(
async def test_langgraph_server_gateway_streams_value_message_snapshots():
events = await _collect_server_gateway_stream(
[
_value_snapshot([_OLD_AI, _HUMAN]),
_value_snapshot([_OLD_AI, _HUMAN, _NEW_AI]),
@@ -850,8 +841,8 @@ def test_langgraph_server_gateway_streams_value_message_snapshots():
]
def test_langgraph_server_gateway_values_do_not_duplicate_message_stream():
events = _collect_server_gateway_stream(
async def test_langgraph_server_gateway_values_do_not_duplicate_message_stream():
events = await _collect_server_gateway_stream(
[
_root_text_delta("new"),
_root_message_finish(),
@@ -866,8 +857,8 @@ def test_langgraph_server_gateway_values_do_not_duplicate_message_stream():
]
def test_langgraph_server_gateway_ignores_non_root_value_messages():
events = _collect_server_gateway_stream(
async def test_langgraph_server_gateway_ignores_non_root_value_messages():
events = await _collect_server_gateway_stream(
[
_value_snapshot(
[{"type": "ai", "content": "subagent text", "id": "subagent-ai"}],
@@ -880,7 +871,7 @@ def test_langgraph_server_gateway_ignores_non_root_value_messages():
assert events[-1] == {"type": "done", "content": "", "response": ""}
def test_langgraph_server_gateway_emits_state_interrupt_before_done():
async def test_langgraph_server_gateway_emits_state_interrupt_before_done():
stream = FakeLangGraphThreadStream(
"abc12345",
events=[],
@@ -930,9 +921,15 @@ def test_langgraph_server_gateway_emits_state_interrupt_before_done():
)
]
events = run_async(_collect())
events = await _collect()
assert events == [
{
"type": "tool_call",
"name": "execute",
"args": {"command": "echo hello"},
"id": "tool-1",
},
{
"type": "interrupt",
"interrupt_id": "interrupt-1",
@@ -954,7 +951,7 @@ def test_langgraph_server_gateway_emits_state_interrupt_before_done():
]
def test_langgraph_server_gateway_streams_subagent_protocol_events():
async def test_langgraph_server_gateway_streams_subagent_protocol_events():
stream = FakeLangGraphThreadStream(
"abc12345",
events=[
@@ -1003,7 +1000,7 @@ def test_langgraph_server_gateway_streams_subagent_protocol_events():
)
]
events = run_async(_collect())
events = await _collect()
assert events == [
{
@@ -1028,7 +1025,7 @@ def test_langgraph_server_gateway_streams_subagent_protocol_events():
]
def test_langgraph_server_gateway_resumes_interrupt_with_thread_stream():
async def test_langgraph_server_gateway_resumes_interrupt_with_thread_stream():
from langgraph.types import Command
stream = FakeLangGraphThreadStream(
@@ -1058,7 +1055,7 @@ def test_langgraph_server_gateway_resumes_interrupt_with_thread_stream():
)
]
events = run_async(_collect())
events = await _collect()
assert stream.run.starts == []
assert stream.run.responses == [
+4 -4
View File
@@ -378,7 +378,7 @@ class TestHitlConfig:
class TestInterruptEventParsing:
def test_interrupt_from_updates_mode(self):
async 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 = collect_events(agent, message="test", thread_id="thread-1")
events = await 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"
def test_updates_without_interrupt_skipped(self):
async 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 = collect_events(agent, message="test", thread_id="thread-1")
events = await collect_events(agent, message="test", thread_id="thread-1")
types = [e["type"] for e in events]
assert "interrupt" not in types
+13
View File
@@ -0,0 +1,13 @@
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"
+10 -12
View File
@@ -2,8 +2,6 @@
from unittest.mock import MagicMock, patch
from tests.conftest import run_async as _run
def _ctx():
from EvoScientist.commands.base import CommandContext
@@ -14,15 +12,15 @@ def _ctx():
class TestInstallSkill:
def test_usage_message_when_no_args(self):
async def test_usage_message_when_no_args(self):
from EvoScientist.commands.implementation.skills import InstallSkill
ctx, ui = _ctx()
_run(InstallSkill().execute(ctx, []))
await 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)
def test_happy_path(self):
async def test_happy_path(self):
from EvoScientist.commands.implementation.skills import InstallSkill
ctx, ui = _ctx()
@@ -35,21 +33,21 @@ class TestInstallSkill:
"path": "/tmp/demo",
},
):
_run(InstallSkill().execute(ctx, ["./some-path"]))
await 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:
def test_usage_message_when_no_args(self):
async def test_usage_message_when_no_args(self):
from EvoScientist.commands.implementation.skills import UninstallSkill
ctx, ui = _ctx()
_run(UninstallSkill().execute(ctx, []))
await 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)
def test_uninstall_success(self):
async def test_uninstall_success(self):
from EvoScientist.commands.implementation.skills import UninstallSkill
ctx, ui = _ctx()
@@ -57,11 +55,11 @@ class TestUninstallSkill:
"EvoScientist.tools.skills_manager.uninstall_skill",
return_value={"success": True},
):
_run(UninstallSkill().execute(ctx, ["demo-skill"]))
await 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)
def test_uninstall_failure(self):
async def test_uninstall_failure(self):
from EvoScientist.commands.implementation.skills import UninstallSkill
ctx, ui = _ctx()
@@ -69,6 +67,6 @@ class TestUninstallSkill:
"EvoScientist.tools.skills_manager.uninstall_skill",
return_value={"success": False, "error": "not found"},
):
_run(UninstallSkill().execute(ctx, ["missing"]))
await UninstallSkill().execute(ctx, ["missing"])
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("Failed: not found" in m for m in msgs)
+10 -10
View File
@@ -19,7 +19,7 @@ async def _agen(items):
yield item
def test_writes_each_event_as_one_jsonl_line(run_async):
async def test_writes_each_event_as_one_jsonl_line():
"""Each event dict is serialized to exactly one JSON line, in order."""
events = [
{"type": "thinking", "content": "hmm", "id": 0},
@@ -34,7 +34,7 @@ def test_writes_each_event_as_one_jsonl_line(run_async):
]
out = io.StringIO()
run_async(write_events_as_json(_agen(events), out))
await write_events_as_json(_agen(events), out)
lines = out.getvalue().splitlines()
assert len(lines) == len(events)
@@ -43,7 +43,7 @@ def test_writes_each_event_as_one_jsonl_line(run_async):
assert parsed[2]["args"] == {"path": "a.md"}
def test_returns_final_response_from_done_event(run_async):
async def test_returns_final_response_from_done_event():
"""The sink returns the response text carried by the terminal `done` event."""
events = [
{"type": "text", "content": "partial"},
@@ -51,12 +51,12 @@ def test_returns_final_response_from_done_event(run_async):
]
out = io.StringIO()
result = run_async(write_events_as_json(_agen(events), out))
result = await write_events_as_json(_agen(events), out)
assert result == "the answer"
def test_non_serializable_arg_does_not_crash_the_stream(run_async):
async def test_non_serializable_arg_does_not_crash_the_stream():
"""A non-JSON-serializable value degrades to its str form instead of raising."""
class Weird:
@@ -72,7 +72,7 @@ def test_non_serializable_arg_does_not_crash_the_stream(run_async):
]
out = io.StringIO()
run_async(write_events_as_json(_agen(events), out))
await 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 @@ def test_non_serializable_arg_does_not_crash_the_stream(run_async):
assert first["args"]["obj"] == "WEIRD"
def test_stream_json_sources_events_from_gateway(run_async):
async def test_stream_json_sources_events_from_gateway():
"""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 @@ def test_stream_json_sources_events_from_gateway(run_async):
return _agen(events)
out = io.StringIO()
result = run_async(stream_json(_FakeGateway(), object(), out=out))
result = await stream_json(_FakeGateway(), object(), out=out)
assert result == "hi"
assert "request" in seen # the request was forwarded to the gateway
@@ -106,7 +106,7 @@ def test_stream_json_sources_events_from_gateway(run_async):
assert types == ["text", "done"]
def test_stream_json_propagates_gateway_errors(run_async):
async def test_stream_json_propagates_gateway_errors():
"""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 @@ def test_stream_json_propagates_gateway_errors(run_async):
out = io.StringIO()
with pytest.raises(RuntimeError, match="boom"):
run_async(stream_json(_FakeGateway(), object(), out=out))
await stream_json(_FakeGateway(), object(), out=out)
+72
View File
@@ -8,6 +8,7 @@ to be available.
from __future__ import annotations
import dataclasses
import sys
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
@@ -32,6 +33,77 @@ 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
# =============================================================================
@@ -0,0 +1,152 @@
"""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
+606 -7
View File
@@ -1,5 +1,6 @@
"""Tests for EvoScientist LLM module."""
from types import SimpleNamespace
from unittest.mock import patch
import pytest
@@ -160,6 +161,68 @@ 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."""
@@ -229,6 +292,39 @@ 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."""
@@ -439,6 +535,200 @@ 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
@@ -957,6 +1247,153 @@ 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
@@ -972,7 +1409,6 @@ 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
@@ -1003,7 +1439,6 @@ 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
@@ -1044,7 +1479,6 @@ 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
@@ -1077,7 +1511,6 @@ 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
@@ -1550,7 +1983,6 @@ 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
@@ -2243,6 +2675,38 @@ 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."""
@@ -2332,6 +2796,7 @@ 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"] == {
@@ -2351,6 +2816,26 @@ 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."""
@@ -2382,8 +2867,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"
# Proxy mode: reasoning skipped (ccproxy untested)
assert "reasoning" not in call_kwargs
# ccproxy uses the Responses API, so reasoning configuration is valid.
assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"}
# Proxy mode: Responses API (bypasses format chain), streaming ON
assert call_kwargs["use_responses_api"] is True
assert "streaming" not in call_kwargs
@@ -2416,6 +2901,120 @@ class TestAutoConfig:
call_kwargs = mock_init.call_args[1]
assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"}
assert "use_responses_api" not in call_kwargs
assert "default_headers" not in call_kwargs
@patch(
"EvoScientist.llm.models._installed_codex_client_version",
return_value="0.144.1",
)
@patch("EvoScientist.llm.models.init_chat_model")
def test_openai_ccproxy_codex_client_headers(
self, mock_init, mock_installed_version, monkeypatch
):
"""ccproxy Codex mode sends Codex-CLI-shaped client headers."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENAI_BASE_URL", "http://127.0.0.1:8000/codex/v1")
monkeypatch.setenv("OPENAI_API_KEY", "ccproxy-oauth")
monkeypatch.delenv("EVOSCIENTIST_CODEX_CLIENT_VERSION", raising=False)
get_chat_model("gpt-5.5", provider="openai")
headers = mock_init.call_args[1]["default_headers"]
assert headers["originator"] == "codex_cli_rs"
assert headers["version"] == "0.144.1"
assert headers["User-Agent"].startswith("codex_cli_rs/0.144.1")
mock_installed_version.assert_called_once_with()
assert mock_init.call_args[1]["reasoning"]["effort"] == "xhigh"
@patch("EvoScientist.llm.models.init_chat_model")
def test_openai_ccproxy_codex_client_version_env(self, mock_init, monkeypatch):
"""EVOSCIENTIST_CODEX_CLIENT_VERSION overrides the pinned version."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENAI_BASE_URL", "http://127.0.0.1:8000/codex/v1")
monkeypatch.setenv("OPENAI_API_KEY", "ccproxy-oauth")
monkeypatch.setenv("EVOSCIENTIST_CODEX_CLIENT_VERSION", "9.9.9")
get_chat_model("gpt-5.5", provider="openai")
headers = mock_init.call_args[1]["default_headers"]
assert headers["version"] == "9.9.9"
assert headers["User-Agent"].startswith("codex_cli_rs/9.9.9")
@patch("EvoScientist.llm.models.subprocess.run")
def test_installed_codex_client_version(self, mock_run):
"""The advertised version follows the installed Codex CLI."""
from EvoScientist.llm.models import _installed_codex_client_version
mock_run.return_value.returncode = 0
mock_run.return_value.stdout = "codex-cli 0.144.1\n"
mock_run.return_value.stderr = ""
_installed_codex_client_version.cache_clear()
try:
assert _installed_codex_client_version() == "0.144.1"
assert _installed_codex_client_version() == "0.144.1"
finally:
_installed_codex_client_version.cache_clear()
mock_run.assert_called_once_with(
["codex", "--version"],
capture_output=True,
text=True,
timeout=2,
check=False,
)
@patch(
"EvoScientist.llm.models._installed_codex_client_version",
return_value="0.140.0",
)
def test_older_installed_codex_uses_fallback(
self, mock_installed_version, monkeypatch
):
"""An outdated installed CLI must not undercut the safe fallback."""
from EvoScientist.llm.models import (
_CODEX_CLIENT_VERSION_FALLBACK,
_resolve_codex_client_version,
)
monkeypatch.delenv("EVOSCIENTIST_CODEX_CLIENT_VERSION", raising=False)
assert _resolve_codex_client_version() == _CODEX_CLIENT_VERSION_FALLBACK
mock_installed_version.assert_called_once_with()
@patch("EvoScientist.llm.models.init_chat_model")
def test_openai_ccproxy_codex_headers_respect_caller(self, mock_init, monkeypatch):
"""Caller-supplied default_headers keys are not overridden."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENAI_BASE_URL", "http://127.0.0.1:8000/codex/v1")
monkeypatch.setenv("OPENAI_API_KEY", "ccproxy-oauth")
get_chat_model(
"gpt-5.5",
provider="openai",
default_headers={"originator": "codex_vscode", "version": "9.9.9"},
)
headers = mock_init.call_args[1]["default_headers"]
assert headers["originator"] == "codex_vscode"
assert headers["version"] == "9.9.9"
assert headers["User-Agent"].startswith("codex_cli_rs/9.9.9")
@patch("EvoScientist.llm.models.init_chat_model")
def test_openai_ccproxy_codex_none_headers(self, mock_init, monkeypatch):
"""An explicit default_headers=None is normalized before gap-filling."""
mock_init.return_value = "mock_model"
monkeypatch.setenv("OPENAI_BASE_URL", "http://127.0.0.1:8000/codex/v1")
monkeypatch.setenv("OPENAI_API_KEY", "ccproxy-oauth")
monkeypatch.setenv("EVOSCIENTIST_CODEX_CLIENT_VERSION", "9.9.9")
get_chat_model(
"gpt-5.5",
provider="openai",
default_headers=None,
)
headers = mock_init.call_args[1]["default_headers"]
assert headers["originator"] == "codex_cli_rs"
assert headers["version"] == "9.9.9"
@patch("EvoScientist.llm.models.init_chat_model")
def test_openai_ccproxy_key_but_wrong_path_not_ccproxy(
+129
View File
@@ -0,0 +1,129 @@
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"
+23 -18
View File
@@ -1399,9 +1399,7 @@ class TestLoadToolsProgressCallback:
monkeypatch.setattr(lc_client, "MultiServerMCPClient", _FakeClient)
def test_success_emits_start_then_success_with_tool_count(self, monkeypatch):
import asyncio
async def test_success_emits_start_then_success_with_tool_count(self, monkeypatch):
from EvoScientist.mcp.client import _load_tools
events: list[tuple[str, str, str]] = []
@@ -1415,36 +1413,45 @@ class TestLoadToolsProgressCallback:
def record(event, name, detail):
events.append((event, name, detail))
asyncio.run(_load_tools(config, on_progress=record))
await _load_tools(config, on_progress=record)
assert events == [
("start", "srv", ""),
("success", "srv", "3"),
]
def test_failure_emits_start_then_error_with_detail(self, monkeypatch):
import asyncio
from EvoScientist.mcp.client import _load_tools
async def test_failure_emits_start_then_error_with_detail(self, monkeypatch):
from EvoScientist.mcp import client as mcp_client
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))
asyncio.run(_load_tools(config, on_progress=record))
await mcp_client._load_tools(config, on_progress=record)
assert events == [
("start", "srv", ""),
("error", "srv", "boom"),
]
assert mcp_client.get_mcp_server_errors() == {"srv": "boom"}
def test_mixed_fleet_reports_each_server_independently(self, monkeypatch):
import asyncio
async def test_success_clears_previous_server_error(self, monkeypatch):
from EvoScientist.mcp import client as mcp_client
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]] = []
@@ -1464,7 +1471,7 @@ class TestLoadToolsProgressCallback:
def record(event, name, detail):
events.append((event, name, detail))
asyncio.run(_load_tools(config, on_progress=record))
await _load_tools(config, on_progress=record)
by_server = {}
for ev, name, detail in events:
@@ -1472,9 +1479,7 @@ class TestLoadToolsProgressCallback:
assert by_server["ok_srv"] == [("start", ""), ("success", "1")]
assert by_server["bad_srv"] == [("start", ""), ("error", "refused")]
def test_callback_errors_do_not_break_the_load(self, monkeypatch):
import asyncio
async def test_callback_errors_do_not_break_the_load(self, monkeypatch):
from EvoScientist.mcp.client import _load_tools
self._patch_client(monkeypatch, {"srv": ["tool1"]})
@@ -1484,10 +1489,10 @@ class TestLoadToolsProgressCallback:
def bad_callback(event, name, detail):
raise RuntimeError("callback bug")
result = asyncio.run(_load_tools(config, on_progress=bad_callback))
result = await _load_tools(config, on_progress=bad_callback)
assert result == {"srv": ["tool1"]}
def test_semaphore_caps_concurrent_connections(self, monkeypatch):
async def test_semaphore_caps_concurrent_connections(self, monkeypatch):
"""Many configured servers must not all spawn at once."""
import asyncio
@@ -1516,7 +1521,7 @@ class TestLoadToolsProgressCallback:
config = {
f"srv{i}": {"transport": "stdio", "command": "demo"} for i in range(10)
}
asyncio.run(mcp_client._load_tools(config))
await mcp_client._load_tools(config)
assert inflight["peak"] <= 3
assert inflight["peak"] > 1 # sanity: we *are* parallelizing
+16 -18
View File
@@ -2,8 +2,6 @@
from unittest.mock import MagicMock, patch
from tests.conftest import run_async as _run
def _ctx():
from EvoScientist.commands.base import CommandContext
@@ -14,16 +12,16 @@ def _ctx():
class TestMCPCommandDispatch:
def test_no_args_lists(self):
async 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={}):
_run(MCPCommand().execute(ctx, []))
await 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)
def test_list_subcommand(self):
async def test_list_subcommand(self):
from EvoScientist.commands.implementation.mcp import MCPCommand
ctx, ui = _ctx()
@@ -31,10 +29,10 @@ class TestMCPCommandDispatch:
"srv1": {"transport": "stdio", "tools": ["foo"], "expose_to": ["main"]},
}
with patch("EvoScientist.mcp.load_mcp_config", return_value=cfg):
_run(MCPCommand().execute(ctx, ["list"]))
await MCPCommand().execute(ctx, ["list"])
ui.mount_renderable.assert_called_once()
def test_add_subcommand_dispatches(self):
async def test_add_subcommand_dispatches(self):
from EvoScientist.commands.implementation.mcp import MCPCommand
ctx, _ui = _ctx()
@@ -48,10 +46,10 @@ class TestMCPCommandDispatch:
return_value={"transport": "stdio"},
) as add_mock,
):
_run(MCPCommand().execute(ctx, ["add", "srv1", "python"]))
await MCPCommand().execute(ctx, ["add", "srv1", "python"])
add_mock.assert_called_once()
def test_edit_subcommand_dispatches(self):
async def test_edit_subcommand_dispatches(self):
from EvoScientist.commands.implementation.mcp import MCPCommand
ctx, _ui = _ctx()
@@ -64,28 +62,28 @@ class TestMCPCommandDispatch:
"EvoScientist.mcp.edit_mcp_server",
) as edit_mock,
):
_run(MCPCommand().execute(ctx, ["edit", "srv1", "--tools", "bar"]))
await MCPCommand().execute(ctx, ["edit", "srv1", "--tools", "bar"])
edit_mock.assert_called_once_with("srv1", tools=["bar"])
def test_remove_subcommand_success(self):
async 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):
_run(MCPCommand().execute(ctx, ["remove", "srv1"]))
await 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)
def test_remove_subcommand_not_found(self):
async 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):
_run(MCPCommand().execute(ctx, ["remove", "missing"]))
await 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)
def test_install_delegates_to_install_mcp_command(self):
async def test_install_delegates_to_install_mcp_command(self):
"""/mcp install should instantiate InstallMCPCommand and execute it."""
from EvoScientist.commands.implementation.mcp import MCPCommand
@@ -101,13 +99,13 @@ class TestMCPCommandDispatch:
instance.execute = fake_execute
klass.return_value = instance
_run(MCPCommand().execute(ctx, ["install", "foo"]))
await MCPCommand().execute(ctx, ["install", "foo"])
klass.assert_called_once()
def test_unknown_subcommand_prints_help(self):
async def test_unknown_subcommand_prints_help(self):
from EvoScientist.commands.implementation.mcp import MCPCommand
ctx, ui = _ctx()
_run(MCPCommand().execute(ctx, ["bogus"]))
await MCPCommand().execute(ctx, ["bogus"])
msgs = [c.args[0] for c in ui.append_system.call_args_list]
assert any("MCP commands:" in m for m in msgs)
+34 -34
View File
@@ -5,8 +5,6 @@ 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."""
@@ -80,7 +78,7 @@ class TestExtractModelAndProvider:
class TestModelCommandUnknownModel:
"""Verify error message for unknown models."""
def test_unknown_model_shows_error(self):
async def test_unknown_model_shows_error(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
@@ -95,7 +93,7 @@ class TestModelCommandUnknownModel:
"EvoScientist.EvoScientist._ensure_config",
return_value=cfg,
):
_run(cmd.execute(ctx, ["nonexistent-model-xyz"]))
await cmd.execute(ctx, ["nonexistent-model-xyz"])
ui.append_system.assert_called_once()
call_args = ui.append_system.call_args
@@ -106,7 +104,7 @@ class TestModelCommandUnknownModel:
class TestModelCommandPickerCancelled:
"""Verify no-op when the interactive picker is cancelled."""
def test_picker_returns_none(self):
async def test_picker_returns_none(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
@@ -122,7 +120,7 @@ class TestModelCommandPickerCancelled:
"EvoScientist.EvoScientist._ensure_config",
return_value=cfg,
):
_run(cmd.execute(ctx, []))
await cmd.execute(ctx, [])
# No model switch should have happened
ui.append_system.assert_not_called()
@@ -131,7 +129,7 @@ class TestModelCommandPickerCancelled:
class TestModelCommandSwitch:
"""Verify a successful model switch updates config and rebuilds agent."""
def test_switch_known_model(self):
async def test_switch_known_model(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
@@ -158,7 +156,7 @@ class TestModelCommandSwitch:
return_value=new_agent,
),
):
_run(cmd.execute(ctx, ["claude-opus-4-8"]))
await 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.
@@ -176,7 +174,7 @@ class TestModelCommandSwitch:
assert "claude-opus-4-8" in msg
assert "anthropic" in msg
def test_switch_with_save_flag(self):
async def test_switch_with_save_flag(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
@@ -203,7 +201,7 @@ class TestModelCommandSwitch:
),
patch("EvoScientist.config.settings.set_config_value") as mock_save,
):
_run(cmd.execute(ctx, ["claude-opus-4-8", "--save"]))
await cmd.execute(ctx, ["claude-opus-4-8", "--save"])
# Config file should be updated
mock_save.assert_any_call("model", "claude-opus-4-8")
@@ -213,7 +211,7 @@ class TestModelCommandSwitch:
msg = ui.append_system.call_args[0][0]
assert "saved to config" in msg
def test_switch_without_save_flag_does_not_persist(self):
async def test_switch_without_save_flag_does_not_persist(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
@@ -240,7 +238,7 @@ class TestModelCommandSwitch:
),
patch("EvoScientist.config.settings.set_config_value") as mock_save,
):
_run(cmd.execute(ctx, ["claude-opus-4-8"]))
await cmd.execute(ctx, ["claude-opus-4-8"])
# Config file should NOT be updated
mock_save.assert_not_called()
@@ -253,7 +251,7 @@ class TestModelCommandSwitch:
class TestModelCommandFailure:
"""Verify error handling when chat-model construction raises."""
def test_build_chat_model_error(self):
async def test_build_chat_model_error(self):
from EvoScientist.commands.implementation.model import ModelCommand
cmd = ModelCommand()
@@ -276,7 +274,7 @@ class TestModelCommandFailure:
side_effect=RuntimeError("API key missing"),
) as mock_build,
):
_run(cmd.execute(ctx, ["claude-opus-4-8"]))
await cmd.execute(ctx, ["claude-opus-4-8"])
mock_build.assert_called_once()
ui.append_system.assert_called_once()
@@ -446,7 +444,7 @@ class TestApplyModelIntegration:
pair so we can assert on identity.
"""
def test_new_agent_is_bound_to_newly_selected_model(self, evo_module_state):
async 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
@@ -502,7 +500,7 @@ class TestApplyModelIntegration:
),
):
cmd = ModelCommand()
_run(cmd._apply_model(ctx, "minimax-m2.7", "openrouter"))
await 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=.
@@ -539,7 +537,9 @@ class TestApplyModelPreservesConfigByReference:
switch (the held object stops being the active ``_config`` after the first).
"""
def test_held_config_reference_tracks_repeated_switches(self, evo_module_state):
async 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"),
]:
_run(cmd._apply_model(ctx, model, provider))
await 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)."""
def test_load_agent_error_is_transactional(self):
async 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.
_run(cmd.execute(ctx, ["claude-opus-4-8", "--save"]))
await 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.
"""
def test_globals_unchanged_when_load_agent_raises(self, evo_module_state):
async 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()
_run(cmd._apply_model(ctx, "minimax-m2.7", "openrouter"))
await 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
def test_picker_entries_include_detected_ollama_models(self):
async 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,
),
):
_run(ModelCommand().execute(ctx, []))
await 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
def test_picker_entries_include_sentinel_when_discovery_empty(self):
async 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,
),
):
_run(ModelCommand().execute(ctx, []))
await 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")
]
def test_picker_skips_ollama_section_when_not_configured(self):
async 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,
),
):
_run(ModelCommand().execute(ctx, []))
await 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)
def test_picker_handles_cfg_without_ollama_base_url_attr(self):
async 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,
),
):
_run(ModelCommand().execute(ctx, []))
await 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)
def test_picker_sentinel_result_is_treated_as_cancel(self):
async 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,
):
_run(ModelCommand().execute(ctx, []))
await ModelCommand().execute(ctx, [])
load_agent.assert_not_called()
assert cfg.model == "claude-sonnet-4-6" # unchanged
def test_picker_applies_detected_ollama_model(self):
async 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(),
),
):
_run(ModelCommand().execute(ctx, []))
await ModelCommand().execute(ctx, [])
# Committed via set_active_config(temp_cfg); original cfg untouched.
set_cfg.assert_called_once()
+135 -27
View File
@@ -6,6 +6,7 @@ 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
@@ -20,7 +21,6 @@ from EvoScientist.middleware.model_fallback import (
clear_fallbacks,
set_ui_emit,
)
from tests.conftest import run_async as _run
# ── Helpers ──────────────────────────────────────────────────────
@@ -84,6 +84,7 @@ 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):
@@ -146,7 +147,7 @@ class TestIsNonFallbackable:
class TestTryFallbacks:
"""End-to-end tests for the fallback chain traversal."""
def test_first_fallback_succeeds(self):
async def test_first_fallback_succeeds(self):
"""When the first fallback model works, return its response."""
add_fallback("fb-model", "fb-provider")
req = _fake_request()
@@ -154,13 +155,13 @@ class TestTryFallbacks:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = _run(_try_fallbacks(req, invoke, Exception("503 boom")))
result = await _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")
def test_skips_failing_fallback_tries_next(self):
async 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")
@@ -177,12 +178,12 @@ class TestTryFallbacks:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = _run(_try_fallbacks(req, _invoke, Exception("503 boom")))
result = await _try_fallbacks(req, _invoke, Exception("503 boom"))
assert result is AI_RESPONSE
assert call_count == 2
def test_all_fallbacks_exhausted_raises_last(self):
async 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")
@@ -202,11 +203,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:
_run(_try_fallbacks(req, _invoke, Exception("503 primary")))
await _try_fallbacks(req, _invoke, Exception("503 primary"))
assert exc_info.value is last_error
def test_non_fallbackable_in_chain_aborts_immediately(self):
async 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
@@ -218,12 +219,95 @@ 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"):
_run(_try_fallbacks(req, _invoke, Exception("503 primary")))
await _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
@@ -233,43 +317,69 @@ class TestTryFallbacks:
class TestGuardAndFallback:
"""Verify that non-fallbackable errors are re-raised before trying the chain."""
def test_context_overflow_raises_immediately(self):
async def test_context_overflow_raises_immediately(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock()
with pytest.raises(ContextOverflowError):
_run(_guard_and_fallback(ContextOverflowError("overflow"), req, invoke))
await _guard_and_fallback(ContextOverflowError("overflow"), req, invoke)
invoke.assert_not_awaited()
def test_malformed_400_raises_immediately(self):
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):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock()
with pytest.raises(Exception, match="invalid_request_error"):
_run(
_guard_and_fallback(
Exception("400: invalid_request_error"), req, invoke
)
await _guard_and_fallback(
Exception("400: invalid_request_error"), req, invoke
)
invoke.assert_not_awaited()
def test_server_error_proceeds_to_fallback(self):
async 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 = _run(_guard_and_fallback(Exception("503 overloaded"), req, invoke))
result = await _guard_and_fallback(Exception("503 overloaded"), req, invoke)
assert result is AI_RESPONSE
invoke.assert_awaited_once()
def test_auth_error_proceeds_to_fallback(self):
async 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()
@@ -277,10 +387,8 @@ class TestGuardAndFallback:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = _run(
_guard_and_fallback(
Exception("400 Bad Request: invalid_api_key"), req, invoke
)
result = await _guard_and_fallback(
Exception("400 Bad Request: invalid_api_key"), req, invoke
)
assert result is AI_RESPONSE
@@ -295,7 +403,7 @@ class TestGuardAndFallback:
class TestUiEmit:
"""Verify that fallback events are surfaced via the registered callback."""
def test_emit_captures_messages(self):
async def test_emit_captures_messages(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock(return_value=AI_RESPONSE)
@@ -305,14 +413,14 @@ class TestUiEmit:
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
_run(_try_fallbacks(req, invoke, Exception("503 down")))
await _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)
def test_emit_shows_non_fallbackable_rejection(self):
async def test_emit_shows_non_fallbackable_rejection(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock()
@@ -321,7 +429,7 @@ class TestUiEmit:
set_ui_emit(lambda text, style: messages.append((text, style)))
with pytest.raises(ContextOverflowError):
_run(_guard_and_fallback(ContextOverflowError("overflow"), req, invoke))
await _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)
+10 -15
View File
@@ -16,7 +16,6 @@ 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
@@ -152,7 +151,7 @@ class TestStartAsyncTaskInjection:
"configurable": {"model": "gpt-5", "model_provider": "openai"}
}
def test_async_start_injects_config(self, restore_model_passthrough_patch):
async def test_async_start_injects_config(self, restore_model_passthrough_patch):
try:
from deepagents.middleware import async_subagents as ds_mod
except ImportError:
@@ -176,12 +175,10 @@ class TestStartAsyncTaskInjection:
"EvoScientist.EvoScientist._ensure_config",
return_value=_stub_cfg(model="claude-haiku-4-5", provider="anthropic"),
):
_run(
tool.coroutine(
description="hi",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
await tool.coroutine(
description="hi",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
runs_async.create.assert_awaited_once()
@@ -267,7 +264,7 @@ class TestUpdateAsyncTaskInjection:
"last_updated_at": "2026-05-07T00:00:00Z",
}
def test_async_update_injects_config(self, restore_model_passthrough_patch):
async 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
@@ -296,12 +293,10 @@ class TestUpdateAsyncTaskInjection:
"EvoScientist.EvoScientist._ensure_config",
return_value=_stub_cfg(model="gpt-5", provider="openai"),
):
_run(
tool.coroutine(
task_id="thread-001",
message="follow up async",
runtime=runtime,
)
await tool.coroutine(
task_id="thread-001",
message="follow up async",
runtime=runtime,
)
runs_async.create.assert_awaited_once()
+4 -6
View File
@@ -2,11 +2,9 @@
from unittest.mock import AsyncMock, MagicMock
from tests.conftest import run_async as _run
class TestNewCommand:
def test_execute_calls_start_new_session(self):
async def test_execute_calls_start_new_session(self):
from EvoScientist.commands.base import CommandContext
from EvoScientist.commands.implementation.session import NewCommand
@@ -18,7 +16,7 @@ class TestNewCommand:
ui=ui,
workspace_dir="/old/ws",
)
_run(NewCommand().execute(ctx, []))
await NewCommand().execute(ctx, [])
ui.start_new_session.assert_awaited_once()
def test_requires_agent_false(self):
@@ -26,7 +24,7 @@ class TestNewCommand:
assert NewCommand().requires_agent is False
def test_no_agent_access(self):
async 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
@@ -35,4 +33,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
_run(NewCommand().execute(ctx, []))
await NewCommand().execute(ctx, [])
+54 -65
View File
@@ -1655,9 +1655,7 @@ def test_turn_compaction_uses_latest_user_turn_only():
]
def test_lifecycle_schedules_turn_worker_without_awaiting(
tmp_path, monkeypatch, run_async
):
async def test_lifecycle_schedules_turn_worker_without_awaiting(tmp_path, monkeypatch):
memory_dir = tmp_path / "memories"
workspace_dir = tmp_path / "workspace"
calls = []
@@ -1682,21 +1680,18 @@ def test_lifecycle_schedules_turn_worker_without_awaiting(
)
runtime = _runtime("thread-1")
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())
state: AgentState[object] = {
"messages": [
HumanMessage("previous turn"),
AIMessage("previous answer"),
HumanMessage("hi"),
AIMessage("done"),
]
}
await middleware.aafter_agent(
state,
runtime,
)
assert len(calls) == 1
request, hooks = calls[0]
@@ -2182,10 +2177,9 @@ def test_observation_linker_does_not_launch_when_observations_disabled(
launch_call.assert_not_called()
def test_async_observation_linker_does_not_launch_when_observations_disabled(
async def test_async_observation_linker_does_not_launch_when_observations_disabled(
tmp_path,
monkeypatch,
run_async,
):
context = _linker_context(
memory_dir=tmp_path / "memories",
@@ -2200,7 +2194,7 @@ def test_async_observation_linker_does_not_launch_when_observations_disabled(
launch_call = MagicMock()
monkeypatch.setattr(memory_launch, "alaunch_background_run", launch_call)
run = run_async(memory_launch.alaunch_observation_linker(context))
run = await memory_launch.alaunch_observation_linker(context)
assert run is None
launch_call.assert_not_called()
@@ -2300,7 +2294,11 @@ def test_memory_worker_observation_writer_modes(
observation_writer=observation_writer,
)
assert type(middleware[0]).__name__ == "ToolErrorHandlerMiddleware"
# 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 _memory_tool_names(middleware) == expected_tools
@@ -2335,8 +2333,8 @@ def test_sync_memory_worker_watcher_untracks_without_counting_on_poll_abort(
assert status.observations_recorded == 0
def test_async_memory_worker_watcher_untracks_without_counting_on_poll_abort(
tmp_path, monkeypatch, run_async
async def test_async_memory_worker_watcher_untracks_without_counting_on_poll_abort(
tmp_path, monkeypatch
):
memory_dir = tmp_path / "memories"
_mark_worker_started(memory_dir)
@@ -2348,14 +2346,12 @@ def test_async_memory_worker_watcher_untracks_without_counting_on_poll_abort(
async def get(self, **_kwargs):
raise RuntimeError("poll failed")
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),
)
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),
)
status = worker_activity.memory_worker_status()
assert status.is_running is False
@@ -2363,8 +2359,8 @@ def test_async_memory_worker_watcher_untracks_without_counting_on_poll_abort(
assert status.observations_recorded == 0
def test_async_memory_worker_watcher_counts_completion_under_blockbuster(
tmp_path, run_async
async def test_async_memory_worker_watcher_counts_completion_under_blockbuster(
tmp_path,
):
memory_dir = tmp_path / "memories"
_mark_worker_started(memory_dir)
@@ -2376,20 +2372,17 @@ def test_async_memory_worker_watcher_counts_completion_under_blockbuster(
async def get(self, **_kwargs):
return {"status": "success"}
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())
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()
status = worker_activity.memory_worker_status()
assert status.is_running is False
assert status.profile_updates == 1
@@ -2527,7 +2520,7 @@ def test_memory_worker_marks_active_status(tmp_path, monkeypatch):
assert status.observations_recorded == 1
def test_async_memory_worker_offloads_blocking_work(tmp_path, monkeypatch, run_async):
async def test_async_memory_worker_offloads_blocking_work(tmp_path, monkeypatch):
monkeypatch.setattr(
background_runs, "default_background_run_url", lambda: "http://x"
)
@@ -2569,22 +2562,18 @@ def test_async_memory_worker_offloads_blocking_work(tmp_path, monkeypatch, run_a
spawned: list[background_runs.BackgroundRun] = []
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())
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,
)
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
+20 -21
View File
@@ -17,7 +17,6 @@ from EvoScientist.llm.ollama_discovery import (
discover_ollama_models,
validate_ollama_connection,
)
from tests.conftest import run_async as _run
class TestValidateOllamaConnection:
@@ -71,17 +70,17 @@ class TestValidateOllamaConnection:
class TestDiscoverOllamaModels:
"""Async probe — contract: never raise, return list[str]."""
def test_empty_base_url_returns_empty_without_http(self):
async 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 = _run(discover_ollama_models(""))
names = await discover_ollama_models("")
assert names == []
def test_none_base_url_returns_empty(self):
names = _run(discover_ollama_models(None))
async def test_none_base_url_returns_empty(self):
names = await discover_ollama_models(None)
assert names == []
def test_200_returns_names(self):
async def test_200_returns_names(self):
async def fake_get(self, url):
resp = MagicMock()
resp.status_code = 200
@@ -93,10 +92,10 @@ class TestDiscoverOllamaModels:
return resp
with patch.object(httpx.AsyncClient, "get", fake_get):
names = _run(discover_ollama_models("http://localhost:11434"))
names = await discover_ollama_models("http://localhost:11434")
assert names == ["llama3.3:latest", "qwen3:8b"]
def test_strips_entries_without_name(self):
async def test_strips_entries_without_name(self):
async def fake_get(self, url):
resp = MagicMock()
resp.status_code = 200
@@ -112,36 +111,36 @@ class TestDiscoverOllamaModels:
return resp
with patch.object(httpx.AsyncClient, "get", fake_get):
names = _run(discover_ollama_models("http://localhost:11434"))
names = await discover_ollama_models("http://localhost:11434")
assert names == ["llama3.3"]
def test_timeout_returns_empty(self):
async 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 = _run(discover_ollama_models("http://localhost:11434"))
names = await discover_ollama_models("http://localhost:11434")
assert names == []
def test_connect_error_returns_empty(self):
async 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 = _run(discover_ollama_models("http://localhost:11434"))
names = await discover_ollama_models("http://localhost:11434")
assert names == []
def test_non_200_returns_empty(self):
async 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 = _run(discover_ollama_models("http://localhost:11434"))
names = await discover_ollama_models("http://localhost:11434")
assert names == []
def test_malformed_json_returns_empty(self):
async def test_malformed_json_returns_empty(self):
async def fake_get(self, url):
resp = MagicMock()
resp.status_code = 200
@@ -149,10 +148,10 @@ class TestDiscoverOllamaModels:
return resp
with patch.object(httpx.AsyncClient, "get", fake_get):
names = _run(discover_ollama_models("http://localhost:11434"))
names = await discover_ollama_models("http://localhost:11434")
assert names == []
def test_missing_models_key_returns_empty(self):
async def test_missing_models_key_returns_empty(self):
async def fake_get(self, url):
resp = MagicMock()
resp.status_code = 200
@@ -160,10 +159,10 @@ class TestDiscoverOllamaModels:
return resp
with patch.object(httpx.AsyncClient, "get", fake_get):
names = _run(discover_ollama_models("http://localhost:11434"))
names = await discover_ollama_models("http://localhost:11434")
assert names == []
def test_trailing_slash_stripped_from_url(self):
async def test_trailing_slash_stripped_from_url(self):
called = {}
async def fake_get(self, url):
@@ -174,7 +173,7 @@ class TestDiscoverOllamaModels:
return resp
with patch.object(httpx.AsyncClient, "get", fake_get):
_run(discover_ollama_models("http://localhost:11434/"))
await discover_ollama_models("http://localhost:11434/")
assert called["url"] == "http://localhost:11434/api/tags"
+309
View File
@@ -184,6 +184,38 @@ 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
# =============================================================================
@@ -386,6 +418,109 @@ 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."""
@@ -1266,6 +1401,7 @@ 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
@@ -1289,12 +1425,185 @@ class TestRunOnboard:
final_config = mock_save.call_args_list[-1].args[0]
assert final_config.auxiliary_provider == "openai"
assert final_config.auxiliary_model == "gpt-5.5"
assert final_config.openai_auth_mode == "api_key"
# The auxiliary provider's key is stored in its per-provider field.
assert final_config.openai_api_key == "sk-aux-openai"
# Main agent is untouched.
assert final_config.provider == "anthropic"
assert final_config.model == "claude-sonnet-4-6"
def test_auxiliary_same_provider_reuses_main_credentials(self):
"""Same-provider co-pilot should not imply separate credentials exist."""
from EvoScientist.config.onboard.wizard import run_onboard
mock_q = MagicMock()
with (
_patch_all_questionary(mock_q),
patch("EvoScientist.config.onboard.wizard.load_config") as mock_load,
patch("EvoScientist.config.onboard.wizard.save_config") as mock_save,
patch("EvoScientist.config.onboard.wizard.console"),
patch("EvoScientist.config.onboard.steps.console"),
):
mock_load.return_value = EvoScientistConfig(
provider="openai",
model="gpt-5.5",
openai_api_key="sk-main-openai",
)
mock_q.select.return_value.ask.side_effect = [
"assemble", # Auxiliary: Assemble
"openai", # Same provider as the main model
"gpt-5.5", # Auxiliary model
]
mock_q.confirm.return_value.ask.side_effect = [True] # Save config
result = run_onboard(
skip_validation=True, only_sections={"auxiliary_model"}
)
assert result is True
final_config = mock_save.call_args_list[-1].args[0]
assert final_config.auxiliary_provider == "openai"
assert final_config.auxiliary_model == "gpt-5.5"
assert final_config.openai_api_key == "sk-main-openai"
mock_q.password.assert_not_called()
assert mock_q.select.return_value.ask.call_count == 3
def test_auxiliary_same_provider_prompts_when_shared_key_missing(self):
"""Same-provider reuse should not hide a missing shared API key."""
from EvoScientist.config.onboard.wizard import run_onboard
mock_q = MagicMock()
with (
_patch_all_questionary(mock_q),
patch("EvoScientist.config.onboard.wizard.load_config") as mock_load,
patch("EvoScientist.config.onboard.wizard.save_config") as mock_save,
patch("EvoScientist.config.onboard.wizard.console"),
patch("EvoScientist.config.onboard.steps.console"),
patch("EvoScientist.config.onboard.helpers.console"),
patch(
"EvoScientist.ccproxy_manager.is_ccproxy_available", return_value=True
),
):
mock_load.return_value = EvoScientistConfig(
provider="openai",
model="gpt-5.5",
openai_auth_mode="api_key",
openai_api_key="",
)
mock_q.select.return_value.ask.side_effect = [
"assemble", # Auxiliary: Assemble
"openai", # Same provider as the main model
"api_key", # Shared OpenAI auth mode
"gpt-5.5", # Auxiliary model
]
mock_q.password.return_value.ask.side_effect = [
"sk-shared-openai",
]
mock_q.confirm.return_value.ask.side_effect = [True] # Save config
result = run_onboard(
skip_validation=True, only_sections={"auxiliary_model"}
)
assert result is True
final_config = mock_save.call_args_list[-1].args[0]
assert final_config.auxiliary_provider == "openai"
assert final_config.auxiliary_model == "gpt-5.5"
assert final_config.openai_auth_mode == "api_key"
assert final_config.openai_api_key == "sk-shared-openai"
mock_q.password.assert_called_once()
assert mock_q.select.return_value.ask.call_count == 4
def test_auxiliary_openai_oauth_skips_api_key(self):
"""Auxiliary OpenAI now uses the shared auth flow and skips keys on OAuth."""
from EvoScientist.config.onboard.wizard import run_onboard
mock_q = MagicMock()
with (
_patch_all_questionary(mock_q),
patch("EvoScientist.config.onboard.wizard.load_config") as mock_load,
patch("EvoScientist.config.onboard.wizard.save_config") as mock_save,
patch("EvoScientist.config.onboard.wizard.console"),
patch("EvoScientist.config.onboard.steps.console"),
patch("EvoScientist.config.onboard.helpers.console"),
patch(
"EvoScientist.ccproxy_manager.is_ccproxy_available", return_value=True
),
patch(
"EvoScientist.ccproxy_manager.check_ccproxy_auth",
return_value=(False, "not authenticated"),
) as mock_auth,
):
mock_load.return_value = EvoScientistConfig()
mock_q.select.return_value.ask.side_effect = [
"assemble", # Auxiliary: Assemble
"openai", # Auxiliary provider
"oauth", # OpenAI auth mode
"gpt-5.5", # Auxiliary model
]
mock_q.text.return_value.ask.side_effect = [
"", # ccproxy port (keep default)
]
mock_q.confirm.return_value.ask.side_effect = [
False, # Do not log in to Codex now
True, # Save config
]
result = run_onboard(
skip_validation=True, only_sections={"auxiliary_model"}
)
assert result is True
final_config = mock_save.call_args_list[-1].args[0]
assert final_config.auxiliary_provider == "openai"
assert final_config.auxiliary_model == "gpt-5.5"
assert final_config.openai_auth_mode == "oauth"
assert final_config.openai_api_key == ""
mock_q.password.assert_not_called()
mock_auth.assert_called_once_with("codex")
def test_auxiliary_reconfigure_clears_unused_openai_oauth(self):
"""Switching co-pilot away from OpenAI clears stale OpenAI OAuth mode."""
from EvoScientist.config.onboard.wizard import run_onboard
mock_q = MagicMock()
with (
_patch_all_questionary(mock_q),
patch("EvoScientist.config.onboard.wizard.load_config") as mock_load,
patch("EvoScientist.config.onboard.wizard.save_config") as mock_save,
patch("EvoScientist.config.onboard.wizard.console"),
patch("EvoScientist.config.onboard.steps.console"),
patch("EvoScientist.config.onboard.helpers.console"),
):
mock_load.return_value = EvoScientistConfig(
provider="anthropic",
model="claude-sonnet-4-6",
anthropic_auth_mode="oauth",
auxiliary_provider="openai",
auxiliary_model="gpt-5.5",
openai_auth_mode="oauth",
)
mock_q.select.return_value.ask.side_effect = [
"assemble", # Auxiliary: Assemble
"minimax", # Auxiliary provider no longer uses OpenAI
"global", # MiniMax region
"minimax-m2", # Auxiliary model
]
mock_q.password.return_value.ask.side_effect = [
"sk-minimax", # MiniMax API key
]
mock_q.confirm.return_value.ask.side_effect = [True] # Save config
result = run_onboard(
skip_validation=True, only_sections={"auxiliary_model"}
)
assert result is True
final_config = mock_save.call_args_list[-1].args[0]
assert final_config.auxiliary_provider == "minimax"
assert final_config.openai_auth_mode == "api_key"
assert final_config.anthropic_auth_mode == "oauth"
def test_auxiliary_custom_provider_collects_base_url(self):
"""Regression for the custom-provider fix: a custom auxiliary provider
collects its base URL (provider -> base URL -> key -> model order)."""
+59
View File
@@ -21,6 +21,7 @@ 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"]
@@ -32,6 +33,7 @@ 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:
@@ -140,6 +142,63 @@ 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.
+4 -6
View File
@@ -79,14 +79,13 @@ class TestPickSkillsInteractive:
class TestInstallSkillsHandlesEmpty:
"""InstallSkills.execute must distinguish None vs [] from the picker."""
def test_empty_list_suppresses_cancel_message(self):
async 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
@@ -97,18 +96,17 @@ class TestInstallSkillsHandlesEmpty:
"EvoScientist.tools.skills_manager.fetch_remote_skill_index",
return_value=_INDEX,
):
_run(InstallSkills().execute(ctx, []))
await 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)
def test_none_shows_cancel_message(self):
async 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
@@ -119,7 +117,7 @@ class TestInstallSkillsHandlesEmpty:
"EvoScientist.tools.skills_manager.fetch_remote_skill_index",
return_value=_INDEX,
):
_run(InstallSkills().execute(ctx, []))
await 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)
+14 -20
View File
@@ -344,9 +344,7 @@ def test_profile_memory_uses_path_pointers_when_profiles_exceed_budget(
)
def test_profile_memory_async_path_bootstraps_and_injects(
tmp_path, monkeypatch, run_async
):
async def test_profile_memory_async_path_bootstraps_and_injects(tmp_path, monkeypatch):
memories = tmp_path / "memories"
workspace = tmp_path / "workspace"
workspace.mkdir()
@@ -356,7 +354,7 @@ def test_profile_memory_async_path_bootstraps_and_injects(
return request
middleware = memory_module.create_memory_middleware(str(memories))
run_async(middleware.awrap_model_call(_request(), _handler))
await middleware.awrap_model_call(_request(), _handler)
assert (memories / "profile" / "USER_PROFILE.md").exists()
@@ -399,8 +397,8 @@ def test_profile_memory_read_failure_uses_path_pointers_without_overwriting(
assert soul_path.read_bytes() == original_bytes
def test_profile_memory_async_path_inlines_content_under_blockbuster(
tmp_path, monkeypatch, run_async
async def test_profile_memory_async_path_inlines_content_under_blockbuster(
tmp_path, monkeypatch
):
memories = tmp_path / "memories"
workspace = tmp_path / "workspace"
@@ -425,17 +423,13 @@ def test_profile_memory_async_path_inlines_content_under_blockbuster(
monkeypatch.setattr(middleware, "_read_profile_memory", tracked_read_profile_memory)
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())
event_loop_thread = threading.get_ident()
blocker = BlockBuster(scanned_modules=memory_module)
blocker.activate()
try:
modified = await middleware.amodify_request(_request())
finally:
blocker.deactivate()
assert call_threads
assert all(thread_id != event_loop_thread for thread_id in call_threads)
@@ -534,8 +528,8 @@ def test_profile_memory_uses_explicit_workspace_for_project_profile(
).exists()
def test_profile_memory_resolves_project_id_once_per_middleware(
tmp_path, monkeypatch, run_async
async def test_profile_memory_resolves_project_id_once_per_middleware(
tmp_path, monkeypatch
):
memories = tmp_path / "memories"
workspace = tmp_path / "workspace"
@@ -552,7 +546,7 @@ 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())
run_async(middleware.amodify_request(_request()))
await middleware.amodify_request(_request())
assert calls == [workspace]
assert middleware.project_id == "P-cached-project"
+35 -38
View File
@@ -8,7 +8,6 @@ from EvoScientist.channels.qq.channel import (
QQConfig,
_build_qq_keyboard,
)
from tests.conftest import run_async as _run
class TestQQChannelSend:
@@ -22,7 +21,7 @@ class TestQQChannelSend:
channel._client.api.post_group_message = AsyncMock()
return channel
def test_send_prefers_native_markdown_for_c2c(self):
async def test_send_prefers_native_markdown_for_c2c(self):
channel = self._make_ready_channel()
msg = OutboundMessage(
channel="qq",
@@ -35,7 +34,7 @@ class TestQQChannelSend:
},
)
assert _run(channel.send(msg)) is True
assert await channel.send(msg) is True
channel._client.api.post_c2c_message.assert_awaited_once()
sent = channel._client.api.post_c2c_message.await_args.kwargs
@@ -46,7 +45,7 @@ class TestQQChannelSend:
assert sent["msg_seq"] == 1
assert "content" not in sent
def test_send_falls_back_to_plain_text_when_markdown_send_fails(self):
async 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(
@@ -63,7 +62,7 @@ class TestQQChannelSend:
},
)
assert _run(channel.send(msg)) is True
assert await 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
@@ -80,7 +79,7 @@ class TestQQChannelSend:
# trigger "duplicate msg_seq".
assert second["msg_seq"] == 2
def test_send_does_not_fallback_on_transport_error(self):
async def test_send_does_not_fallback_on_transport_error(self):
channel = self._make_ready_channel()
async def _send_once(coro_factory, max_retries=3):
@@ -101,13 +100,13 @@ class TestQQChannelSend:
},
)
assert _run(channel.send(msg)) is False
assert await 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
def test_send_does_not_fallback_when_transport_error_mentions_markdown(self):
async 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."""
@@ -133,10 +132,10 @@ class TestQQChannelSend:
},
)
assert _run(channel.send(msg)) is False
assert await channel.send(msg) is False
channel._client.api.post_c2c_message.assert_awaited_once()
def test_send_falls_back_on_qq_server_error_code(self):
async 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()
@@ -159,7 +158,7 @@ class TestQQChannelSend:
},
)
assert _run(channel.send(msg)) is True
assert await 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
@@ -231,7 +230,7 @@ class TestQQSendWithButtons:
channel._client.api.post_group_message = AsyncMock()
return channel
def test_c2c_send_attaches_keyboard(self):
async def test_c2c_send_attaches_keyboard(self):
channel = self._make_channel()
msg = OutboundMessage(
channel="qq",
@@ -247,7 +246,7 @@ class TestQQSendWithButtons:
],
},
)
assert _run(channel.send(msg)) is True
assert await channel.send(msg) is True
sent = channel._client.api.post_c2c_message.await_args.kwargs
assert sent["msg_type"] == 2
@@ -256,7 +255,7 @@ class TestQQSendWithButtons:
assert rows[0]["buttons"][0]["action"]["data"] == "1"
assert rows[1]["buttons"][0]["action"]["data"] == "2"
def test_group_send_does_not_attach_keyboard(self):
async def test_group_send_does_not_attach_keyboard(self):
"""Group keyboards are out of scope — silently dropped."""
channel = self._make_channel()
msg = OutboundMessage(
@@ -270,11 +269,11 @@ class TestQQSendWithButtons:
"buttons": [{"text": "Approve", "value": "1"}],
},
)
assert _run(channel.send(msg)) is True
assert await channel.send(msg) is True
sent = channel._client.api.post_group_message.await_args.kwargs
assert "keyboard" not in sent
def test_fallback_appends_button_hint_when_keyboard_present(self):
async 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`).
@@ -302,7 +301,7 @@ class TestQQSendWithButtons:
],
},
)
assert _run(channel.send(msg)) is True
assert await channel.send(msg) is True
plain_call = channel._client.api.post_c2c_message.await_args_list[1].kwargs
assert plain_call["msg_type"] == 0
@@ -311,7 +310,7 @@ class TestQQSendWithButtons:
assert "1=Approve" in plain_call["content"]
assert "2=Reject" in plain_call["content"]
def test_fallback_hint_handles_non_string_button_value(self):
async 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()
@@ -335,7 +334,7 @@ class TestQQSendWithButtons:
],
},
)
assert _run(channel.send(msg)) is True
assert await 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"]
@@ -392,9 +391,9 @@ class TestQQInteractionCallback:
)
return interaction
def test_click_publishes_to_bus_with_button_data(self):
async def test_click_publishes_to_bus_with_button_data(self):
channel = self._make_channel()
_run(channel._on_interaction(self._make_interaction("1")))
await channel._on_interaction(self._make_interaction("1"))
channel._bus.publish_inbound.assert_awaited_once()
inbound = channel._bus.publish_inbound.await_args[0][0]
@@ -405,63 +404,61 @@ class TestQQInteractionCallback:
assert inbound.metadata["button_value"] == "1"
assert inbound.metadata["msg_type"] == "c2c"
def test_click_acks_interaction(self):
async def test_click_acks_interaction(self):
channel = self._make_channel()
_run(channel._on_interaction(self._make_interaction("1")))
await channel._on_interaction(self._make_interaction("1"))
channel._client.api.on_interaction_result.assert_awaited_once_with("intr_1", 0)
def test_click_bypasses_debounce(self):
async def test_click_bypasses_debounce(self):
"""Click never hits queue_message (debounce buffer)."""
channel = self._make_channel()
channel.queue_message = AsyncMock()
_run(channel._on_interaction(self._make_interaction("3")))
await channel._on_interaction(self._make_interaction("3"))
channel.queue_message.assert_not_called()
channel._bus.publish_inbound.assert_awaited_once()
def test_group_interaction_ignored(self):
async 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"
_run(channel._on_interaction(intr))
await 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)
def test_click_dropped_when_middleware_rejects(self):
async def test_click_dropped_when_middleware_rejects(self):
channel = self._make_channel()
channel._build_inbound_async = AsyncMock(return_value=None)
_run(channel._on_interaction(self._make_interaction("1")))
await 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()
def test_empty_button_data_falls_back_to_button_id(self):
async def test_empty_button_data_falls_back_to_button_id(self):
channel = self._make_channel()
_run(
channel._on_interaction(
self._make_interaction(button_data="", button_id="btn_3")
)
await 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"
def test_ack_fires_even_when_handler_throws(self):
async 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.
_run(channel._on_interaction(self._make_interaction("1")))
await channel._on_interaction(self._make_interaction("1"))
channel._client.api.on_interaction_result.assert_awaited_once_with("intr_1", 0)
def test_button_value_metadata_is_string_coerced(self):
async 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)
_run(channel._on_interaction(intr))
await channel._on_interaction(intr)
inbound = channel._bus.publish_inbound.await_args[0][0]
assert inbound.content == "42"
assert inbound.metadata["button_value"] == "42"
+177
View File
@@ -0,0 +1,177 @@
"""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)
+16 -17
View File
@@ -2,7 +2,6 @@
from unittest.mock import AsyncMock, MagicMock
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway, FakeThreadStore
@@ -24,7 +23,7 @@ def _ctx(thread_id="current", workspace_dir="/ws", thread_store=None):
class TestResumeCommand:
def test_with_arg_resolves_and_calls_ui(self):
async def test_with_arg_resolves_and_calls_ui(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx(
@@ -33,23 +32,23 @@ class TestResumeCommand:
metadata={"workspace_dir": "/restored"},
)
)
_run(ResumeCommand().execute(ctx, ["target-tid"]))
await 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"
def test_no_arg_empty_threads_prints_message(self):
async def test_no_arg_empty_threads_prints_message(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx()
_run(ResumeCommand().execute(ctx, []))
await 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()
def test_no_arg_calls_picker(self):
async def test_no_arg_calls_picker(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx()
@@ -60,11 +59,11 @@ class TestResumeCommand:
resolved_thread_id="picked-tid",
)
ctx.graph_gateway = FakeGraphGateway(thread_store=store)
_run(ResumeCommand().execute(ctx, []))
await ResumeCommand().execute(ctx, [])
ui.wait_for_thread_pick.assert_awaited_once()
ui.handle_session_resume.assert_awaited_once()
def test_picker_cancel_returns(self):
async def test_picker_cancel_returns(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx()
@@ -72,28 +71,28 @@ class TestResumeCommand:
threads = [{"thread_id": "t1", "preview": "", "message_count": 0}]
store = FakeThreadStore(threads=threads)
ctx.graph_gateway = FakeGraphGateway(thread_store=store)
_run(ResumeCommand().execute(ctx, []))
await ResumeCommand().execute(ctx, [])
ui.handle_session_resume.assert_not_called()
def test_ambiguous_prefix(self):
async def test_ambiguous_prefix(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx(thread_store=FakeThreadStore(matches=["abc-one", "abc-two"]))
_run(ResumeCommand().execute(ctx, ["abc"]))
await 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()
def test_not_found(self):
async def test_not_found(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx()
_run(ResumeCommand().execute(ctx, ["missing"]))
await 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()
def test_prefix_resolves_to_unique_match(self):
async def test_prefix_resolves_to_unique_match(self):
from EvoScientist.commands.implementation.session import ResumeCommand
ctx, ui = _ctx(
@@ -102,18 +101,18 @@ class TestResumeCommand:
metadata={"workspace_dir": "/ws1"},
)
)
_run(ResumeCommand().execute(ctx, ["abc"]))
await ResumeCommand().execute(ctx, ["abc"])
ui.handle_session_resume.assert_awaited_once_with("abc-one", "/ws1")
assert ctx.thread_id == "abc-one"
def test_empty_workspace_metadata_preserves_ctx_workspace(self):
async 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={}),
)
_run(ResumeCommand().execute(ctx, ["tid"]))
await 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)
+46 -56
View File
@@ -5,8 +5,6 @@ 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."""
@@ -40,9 +38,9 @@ class TestBasicIO:
ui.mount_renderable(table)
console.print.assert_called_once_with(table)
def test_flush_is_async_noop(self):
async def test_flush_is_async_noop(self):
ui, console = _make_ui()
_run(ui.flush())
await ui.flush()
# flush should not print anything
console.print.assert_not_called()
@@ -50,33 +48,29 @@ class TestBasicIO:
class TestWaitForModelPick:
"""CLI model picker fallback: print table + return None."""
def test_returns_none(self):
async def test_returns_none(self):
ui, _ = _make_ui()
entries = [
("claude-sonnet-4-6", "anthropic/claude-sonnet", "anthropic"),
("gpt-4o", "openai/gpt-4o", "openai"),
]
result = _run(
ui.wait_for_model_pick(
entries,
current_model="claude-sonnet-4-6",
current_provider="anthropic",
)
result = await ui.wait_for_model_pick(
entries,
current_model="claude-sonnet-4-6",
current_provider="anthropic",
)
assert result is None
def test_prints_table_with_current_model_marker(self):
async 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"),
]
_run(
ui.wait_for_model_pick(
entries,
current_model="claude-sonnet-4-6",
current_provider="anthropic",
)
await 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
@@ -87,27 +81,23 @@ class TestWaitForModelPick:
assert "Usage: /model" in usage_arg
assert "--save" in usage_arg
def test_no_current_model_no_marker(self):
async def test_no_current_model_no_marker(self):
ui, console = _make_ui()
entries = [("claude-sonnet-4-6", "anthropic/claude-sonnet", "anthropic")]
_run(
ui.wait_for_model_pick(
entries,
current_model=None,
current_provider=None,
)
await 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
def test_empty_entries_still_prints_header_and_usage(self):
async def test_empty_entries_still_prints_header_and_usage(self):
ui, console = _make_ui()
result = _run(
ui.wait_for_model_pick(
[],
current_model=None,
current_provider=None,
)
result = await 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
@@ -213,7 +203,7 @@ class TestWaitForThreadPick:
},
]
def test_returns_selected_thread_id(self, monkeypatch):
async def test_returns_selected_thread_id(self, monkeypatch):
import EvoScientist.cli.rich_command_ui as mod
ui, _ = _make_ui()
@@ -226,7 +216,7 @@ class TestWaitForThreadPick:
return prompt
monkeypatch.setattr("questionary.select", fake_select)
result = _run(ui.wait_for_thread_pick(self._threads(), "abc123", "pick:"))
result = await 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 +
@@ -235,14 +225,14 @@ class TestWaitForThreadPick:
# Table import removed; this test no longer depends on console output.
assert mod.RichCLICommandUI is not None # sanity
def test_cancel_returns_none(self, monkeypatch):
async def test_cancel_returns_none(self, monkeypatch):
ui, _ = _make_ui()
prompt = self._fake_prompt(None)
monkeypatch.setattr("questionary.select", lambda *a, **k: prompt)
result = _run(ui.wait_for_thread_pick(self._threads(), "abc123", "pick:"))
result = await ui.wait_for_thread_pick(self._threads(), "abc123", "pick:")
assert result is None
def test_current_thread_marker_in_label(self, monkeypatch):
async def test_current_thread_marker_in_label(self, monkeypatch):
ui, _ = _make_ui()
prompt = self._fake_prompt(None)
captured_choices: list = []
@@ -252,7 +242,7 @@ class TestWaitForThreadPick:
return prompt
monkeypatch.setattr("questionary.select", fake_select)
_run(ui.wait_for_thread_pick(self._threads(), "abc123", "pick:"))
await 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)
@@ -282,38 +272,38 @@ class TestCompactIndicator:
class TestPhaseBMigrated:
"""Session lifecycle callbacks (start/resume) filled in Phase B."""
def test_start_new_session_fires_callback(self):
async def test_start_new_session_fires_callback(self):
from unittest.mock import AsyncMock
cb = AsyncMock()
ui, _ = _make_ui(on_start_new_session=cb)
_run(ui.start_new_session())
await ui.start_new_session()
cb.assert_awaited_once()
def test_start_new_session_without_callback_is_noop(self):
async def test_start_new_session_without_callback_is_noop(self):
ui, console = _make_ui()
_run(ui.start_new_session())
await ui.start_new_session()
console.print.assert_not_called()
def test_handle_session_resume_awaits_callback(self):
async def test_handle_session_resume_awaits_callback(self):
from unittest.mock import AsyncMock
cb = AsyncMock()
ui, _ = _make_ui(on_handle_session_resume=cb)
_run(ui.handle_session_resume("tid-x", "/workspace"))
await ui.handle_session_resume("tid-x", "/workspace")
cb.assert_awaited_once_with("tid-x", "/workspace")
def test_handle_session_resume_without_callback_is_noop(self):
async def test_handle_session_resume_without_callback_is_noop(self):
ui, _ = _make_ui()
# Should not raise
_run(ui.handle_session_resume("tid-x"))
await ui.handle_session_resume("tid-x")
def test_handle_session_resume_workspace_defaults_none(self):
async def test_handle_session_resume_workspace_defaults_none(self):
from unittest.mock import AsyncMock
cb = AsyncMock()
ui, _ = _make_ui(on_handle_session_resume=cb)
_run(ui.handle_session_resume("tid-x"))
await ui.handle_session_resume("tid-x")
cb.assert_awaited_once_with("tid-x", None)
@@ -321,7 +311,7 @@ class TestPhaseCMigrated:
"""Skill/MCP browse pickers delegate to questionary helpers via
``asyncio.to_thread`` since questionary blocks the event loop."""
def test_skill_browse_delegates_to_picker(self, monkeypatch):
async def test_skill_browse_delegates_to_picker(self, monkeypatch):
from unittest.mock import MagicMock
picker = MagicMock(return_value=["skill-a", "skill-b"])
@@ -330,11 +320,11 @@ class TestPhaseCMigrated:
picker,
)
ui, _ = _make_ui()
result = _run(ui.wait_for_skill_browse([{"name": "a"}], {"installed"}, "core"))
result = await ui.wait_for_skill_browse([{"name": "a"}], {"installed"}, "core")
assert result == ["skill-a", "skill-b"]
picker.assert_called_once_with([{"name": "a"}], {"installed"}, "core")
def test_skill_browse_cancel_returns_none(self, monkeypatch):
async def test_skill_browse_cancel_returns_none(self, monkeypatch):
from unittest.mock import MagicMock
monkeypatch.setattr(
@@ -342,10 +332,10 @@ class TestPhaseCMigrated:
MagicMock(return_value=None),
)
ui, _ = _make_ui()
result = _run(ui.wait_for_skill_browse([], set(), ""))
result = await ui.wait_for_skill_browse([], set(), "")
assert result is None
def test_mcp_browse_delegates_to_picker(self, monkeypatch):
async def test_mcp_browse_delegates_to_picker(self, monkeypatch):
from unittest.mock import MagicMock
sentinel_entries = [MagicMock(name="entry1"), MagicMock(name="entry2")]
@@ -355,11 +345,11 @@ class TestPhaseCMigrated:
picker,
)
ui, _ = _make_ui()
result = _run(ui.wait_for_mcp_browse([MagicMock()], {"configured"}, ""))
result = await ui.wait_for_mcp_browse([MagicMock()], {"configured"}, "")
assert result is sentinel_entries
picker.assert_called_once()
def test_mcp_browse_cancel_returns_none(self, monkeypatch):
async def test_mcp_browse_cancel_returns_none(self, monkeypatch):
from unittest.mock import MagicMock
monkeypatch.setattr(
@@ -367,5 +357,5 @@ class TestPhaseCMigrated:
MagicMock(return_value=None),
)
ui, _ = _make_ui()
result = _run(ui.wait_for_mcp_browse([], set(), ""))
result = await ui.wait_for_mcp_browse([], set(), "")
assert result is None
+117
View File
@@ -0,0 +1,117 @@
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")]
+24 -28
View File
@@ -2,8 +2,6 @@
from unittest.mock import MagicMock, patch
from tests.conftest import run_async as _run
def _ctx():
from EvoScientist.commands.base import CommandContext
@@ -12,17 +10,17 @@ def _ctx():
return CommandContext(agent=None, thread_id="tid", ui=ui), ui
def test_list_when_backend_down():
async 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):
_run(ScheduleCommand().execute(ctx, ["list"]))
await 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)
def test_add_parses_five_field_cron_and_prompt():
async def test_add_parses_five_field_cron_and_prompt():
from EvoScientist.commands.implementation.schedule import ScheduleCommand
ctx, _ui = _ctx()
@@ -33,17 +31,15 @@ def test_add_parses_five_field_cron_and_prompt():
return_value={"cron_id": "c-9"},
) as mk,
):
_run(
ScheduleCommand().execute(
ctx, ["add", "*/10", "*", "*", "*", "*", "search", "uk", "weather"]
)
await ScheduleCommand().execute(
ctx, ["add", "*/10", "*", "*", "*", "*", "search", "uk", "weather"]
)
kw = mk.call_args.kwargs
assert kw["schedule"] == "*/10 * * * *"
assert kw["prompt"] == "search uk weather"
def test_list_renders_table():
async def test_list_renders_table():
from EvoScientist.commands.implementation.schedule import ScheduleCommand
ctx, ui = _ctx()
@@ -60,11 +56,11 @@ def test_list_renders_table():
patch("EvoScientist.cron.schedule.is_available", return_value=True),
patch("EvoScientist.cron.schedule.list_schedules", return_value=rows),
):
_run(ScheduleCommand().execute(ctx, ["list"]))
await ScheduleCommand().execute(ctx, ["list"])
ui.mount_renderable.assert_called_once()
def test_add_parses_quoted_cron_and_prompt():
async def test_add_parses_quoted_cron_and_prompt():
from EvoScientist.commands.implementation.schedule import ScheduleCommand
ctx, _ui = _ctx()
@@ -75,15 +71,15 @@ def test_add_parses_quoted_cron_and_prompt():
return_value={"cron_id": "c-9"},
) as mk,
):
_run(
ScheduleCommand().execute(ctx, ["add", "*/10 * * * *", "search uk weather"])
await ScheduleCommand().execute(
ctx, ["add", "*/10 * * * *", "search uk weather"]
)
kw = mk.call_args.kwargs
assert kw["schedule"] == "*/10 * * * *"
assert kw["prompt"] == "search uk weather"
def test_run_with_matching_prefix_fires_matched_prompt():
async def test_run_with_matching_prefix_fires_matched_prompt():
from EvoScientist.commands.implementation.schedule import ScheduleCommand
ctx, _ui = _ctx()
@@ -96,11 +92,11 @@ def test_run_with_matching_prefix_fires_matched_prompt():
return_value={"run_id": "r-1"},
) as rn,
):
_run(ScheduleCommand().execute(ctx, ["run", "c-123"]))
await ScheduleCommand().execute(ctx, ["run", "c-123"])
rn.assert_called_once_with("do the thing")
def test_run_with_no_match_reports():
async def test_run_with_no_match_reports():
from EvoScientist.commands.implementation.schedule import ScheduleCommand
ctx, ui = _ctx()
@@ -109,13 +105,13 @@ def test_run_with_no_match_reports():
patch("EvoScientist.cron.schedule.list_schedules", return_value=[]),
patch("EvoScientist.cron.schedule.run_now") as rn,
):
_run(ScheduleCommand().execute(ctx, ["run", "nope"]))
await 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)
def test_pause_resume_set_enabled_with_resolved_id():
async def test_pause_resume_set_enabled_with_resolved_id():
from EvoScientist.commands.implementation.schedule import ScheduleCommand
rows = [{"cron_id": "c-abcdef", "metadata": {"name": "t"}}]
@@ -126,7 +122,7 @@ 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,
):
_run(ScheduleCommand().execute(ctx, [sub, "c-abc"]))
await ScheduleCommand().execute(ctx, [sub, "c-abc"])
se.assert_called_once_with("c-abcdef", expected)
@@ -135,7 +131,7 @@ def test_pause_resume_set_enabled_with_resolved_id():
# ---------------------------------------------------------------------------
def test_list_error_shows_red_message_no_exception():
async 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
@@ -147,7 +143,7 @@ def test_list_error_shows_red_message_no_exception():
side_effect=RuntimeError("backend gone"),
),
):
_run(ScheduleCommand().execute(ctx, ["list"]))
await 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)
@@ -158,7 +154,7 @@ def test_list_error_shows_red_message_no_exception():
# ---------------------------------------------------------------------------
def test_remove_ambiguous_prefix_aborts_without_deleting():
async 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
@@ -172,7 +168,7 @@ 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,
):
_run(ScheduleCommand().execute(ctx, ["remove", "abc"]))
await 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)
@@ -183,7 +179,7 @@ def test_remove_ambiguous_prefix_aborts_without_deleting():
# ---------------------------------------------------------------------------
def test_remove_backend_error_shows_red_error_not_no_match():
async 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
@@ -196,7 +192,7 @@ def test_remove_backend_error_shows_red_error_not_no_match():
),
patch("EvoScientist.cron.schedule.delete_schedule") as mk,
):
_run(ScheduleCommand().execute(ctx, ["remove", "abc"]))
await 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}"
@@ -209,7 +205,7 @@ def test_remove_backend_error_shows_red_error_not_no_match():
)
def test_add_name_sanitized_from_nasty_prompt():
async def test_add_name_sanitized_from_nasty_prompt():
"""B3: prompt with newline / slashes / special chars → clean kebab-case name."""
import re
@@ -225,7 +221,7 @@ def test_add_name_sanitized_from_nasty_prompt():
return_value={"cron_id": "c-x"},
) as mk,
):
_run(ScheduleCommand().execute(ctx, ["add", "*/5 * * * *", nasty_prompt]))
await 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
+242
View File
@@ -0,0 +1,242 @@
"""Regression tests for the helpers ``ErrorNormalizationMiddleware``
uses to build the SSE error envelope.
- ``_redact_api_keys`` + ``_build_env_key_redaction_re`` — scrubs
deployed credentials that the SDK might echo back.
- ``_extract_status_code`` / ``_extract_provider_code`` /
``_extract_error_type`` — read SDK-specific fields off the raised
exception.
Middleware wire behavior + ``_provider_from_model`` live in
``test_error_normalization_middleware.py``. One end-to-end orjson test
at the bottom guards that a ``ProviderStreamError`` survives
langgraph_api's UNPATCHED ``serde.default`` under
``OPT_SERIALIZE_DATACLASS`` — the whole reason the wrapper exists.
"""
from __future__ import annotations
import os
import langgraph_api.serde as _serde_mod
from EvoScientist.llm.errors import (
_API_KEY_ENV_SUFFIXES,
_build_env_key_redaction_re,
_extract_error_type,
_extract_provider_code,
_extract_status_code,
_redact_api_keys,
)
# ---------------------------------------------------------------------------
# Redaction
# ---------------------------------------------------------------------------
def test_env_deployed_key_redacted_in_message(monkeypatch):
"""A credential exported via env var is scrubbed by
``_redact_api_keys``. The redaction table is rebuilt per call —
``monkeypatch.setenv`` alone is enough, no attribute reassignment.
"""
key = "sk-proj-aBcDeFgHiJkLmNoPqRsTuVwXyZ1234567890"
monkeypatch.setenv("OPENAI_API_KEY", key)
msg = (
f"Invalid API key: {key}. Get a new one at https://platform.openai.com/api-keys"
)
redacted = _redact_api_keys(msg)
assert key not in redacted
assert "<redacted>" in redacted
assert "Invalid API key" in redacted
assert "platform.openai.com" in redacted
def test_multiple_env_keys_redacted_independently(monkeypatch):
"""Each ``*_API_KEY`` / ``*_TOKEN`` / ``*_SECRET`` env var
contributes its own prefix to the alternation.
"""
k1 = "sk-or-aBcDeFg012345678901234"
k2 = "AIzaABCDEFGHIJ0123456789"
k3 = "ghp_p4t70k3n0123456789abcdef"
monkeypatch.setenv("OPENROUTER_API_KEY", k1)
monkeypatch.setenv("GOOGLE_API_KEY", k2)
monkeypatch.setenv("GITHUB_TOKEN", k3)
msg = _redact_api_keys(f"Failures: {k1}, {k2}, {k3}")
assert k1 not in msg
assert k2 not in msg
assert k3 not in msg
assert msg.count("<redacted>") == 3
def test_base64_suffix_secret_fully_redacted(monkeypatch):
"""A base64-style secret (``/`` ``+`` ``=``) must redact end-to-end,
not leak its tail past the first padding char.
"""
key = "AbCdEfGh/secret+tail=="
monkeypatch.setenv("SOME_SECRET", key)
msg = _redact_api_keys(f"auth failed with token={key} on retry")
assert "secret" not in msg
assert "tail" not in msg
assert "<redacted>" in msg
assert "auth failed" in msg
assert "on retry" in msg
def test_unknown_shape_not_redacted_without_env(monkeypatch):
"""Env-only redaction: a key-shaped string not deployed via env is
left alone. Tradeoff — we only scrub what we know is a secret.
"""
for k in list(os.environ):
if k.endswith(_API_KEY_ENV_SUFFIXES):
monkeypatch.delenv(k, raising=False)
msg = _redact_api_keys("Unknown key seen: sk-or-aBcDeFg012345678901234")
assert "sk-or-aBcDeFg012345678901234" in msg
assert "<redacted>" not in msg
def test_env_key_loaded_after_first_call_is_redacted(monkeypatch):
"""The pattern rebuilds every call so keys loaded after
``patches.py`` imports (typical ``load_dotenv`` sequence) are
still scrubbed on the next call.
"""
for k in list(os.environ):
if k.endswith(_API_KEY_ENV_SUFFIXES):
monkeypatch.delenv(k, raising=False)
key = "sk-proj-loaded_after_import_1234567890abcdef"
# Pass 1: env empty — key leaks.
assert key in _redact_api_keys(f"leak: {key}")
# Pass 2: after simulated load_dotenv.
monkeypatch.setenv("OPENAI_API_KEY", key)
redacted = _redact_api_keys(f"leak: {key}")
assert key not in redacted
assert "<redacted>" in redacted
def test_redaction_regex_holds_only_prefix(monkeypatch):
"""Defense-in-depth: the compiled regex must not embed the full key.
A process-memory leak (traceback locals, debugger) exposes at most
the first 8 chars — not the secret.
"""
key = "sk-proj-aBcDeFgHiJkLmNoPqRsTuVwXyZ1234567890_secret_suffix"
monkeypatch.setenv("OPENAI_API_KEY", key)
pattern = _build_env_key_redaction_re()
assert pattern is not None
assert key not in pattern.pattern
assert "aBcDeFgHiJkLmNoPqRs" not in pattern.pattern
# Sanity: still matches the full key at runtime via prefix + suffix
# greedy.
m = pattern.search(f"err: {key}")
assert m is not None
assert m.group(0) == key
# ---------------------------------------------------------------------------
# Field extractors
# ---------------------------------------------------------------------------
def _fake_exc(**attrs):
return type("APIError", (Exception,), attrs)("boom")
def test_status_code_read_from_direct_attribute():
"""openai / anthropic ``APIStatusError`` carries integer
``.status_code`` — the primary path.
"""
assert _extract_status_code(_fake_exc(status_code=429)) == 429
def test_status_code_read_via_response_attribute():
"""Wrappers that don't promote status to top level expose it via
``.response.status_code`` (httpx pattern).
"""
class FakeResponse:
status_code = 504
assert _extract_status_code(_fake_exc(response=FakeResponse())) == 504
def test_status_code_read_via_integer_code_attribute():
"""``google.genai.errors.APIError`` stores HTTP status as integer
``.code`` — type-disambiguated from openai/anthropic's string
``.code`` (provider error code).
"""
assert _extract_status_code(_fake_exc(code=400)) == 400
def test_provider_code_read_from_string_code_attribute():
"""Provider error code (``insufficient_quota`` etc.) is a string
``.code`` — higher signal than the integer HTTP status alone.
"""
assert (
_extract_provider_code(_fake_exc(code="insufficient_quota"))
== "insufficient_quota"
)
def test_provider_code_ignores_integer_code():
"""An integer ``.code`` is HTTP status (see above); must not bleed
into the provider-code path.
"""
assert _extract_provider_code(_fake_exc(code=429)) is None
def test_error_type_read_from_type_attribute():
"""openai exposes a ``.type`` label (``rate_limit_error``)."""
assert _extract_error_type(_fake_exc(type="rate_limit_error")) == "rate_limit_error"
def test_extractors_return_none_when_attributes_absent():
"""A bare exception with no SDK-shape attributes — every extractor
returns None so the envelope drops the optional fields.
"""
exc = _fake_exc()
assert _extract_status_code(exc) is None
assert _extract_provider_code(exc) is None
assert _extract_error_type(exc) is None
# ---------------------------------------------------------------------------
# End-to-end: ProviderStreamError survives orjson under
# OPT_SERIALIZE_DATACLASS via upstream's UNPATCHED serde.default.
# ---------------------------------------------------------------------------
def test_provider_stream_error_survives_orjson_dataclass_option():
"""Guard: ``ProviderStreamError`` — a plain Exception subclass with
a ``model_dump()`` hook — must emerge as the envelope on the wire
even under ``OPT_SERIALIZE_DATACLASS``, using ONLY upstream's
stock ``serde.default``. Proof that we no longer need to patch
the serde module.
"""
import orjson
from EvoScientist.llm.errors import ProviderStreamError
err = ProviderStreamError(
provider="openrouter",
class_qualname="openrouter.errors.foo.UnauthorizedResponseError",
message="User not found.",
status_code=401,
)
wire = orjson.dumps(
err,
default=_serde_mod.default, # upstream, unpatched
option=orjson.OPT_SERIALIZE_DATACLASS,
)
decoded = orjson.loads(wire)
assert decoded == {
"error": "UnauthorizedResponseError",
"class": "openrouter.errors.foo.UnauthorizedResponseError",
"message": "User not found.",
"provider": "openrouter",
"status_code": 401,
}
+35 -36
View File
@@ -27,7 +27,6 @@ from EvoScientist.cli.commands import (
from EvoScientist.commands.base import ChannelRuntime
from EvoScientist.config import EvoScientistConfig
from EvoScientist.gateway import RuntimeGateways, ThreadStore
from tests.conftest import run_async as _run
from tests.fakes import FakeGraphGateway, FakeThreadStore
@@ -71,7 +70,7 @@ def _runtime_state(
)
def test_hook_updates_runtime_state_on_agent_swap():
async def test_hook_updates_runtime_state_on_agent_swap():
"""``/model`` mutates ``ctx.agent`` to a new handle — the hook must
push that handle into the shared runtime state so the outer poll loop sees
it on the next message."""
@@ -86,12 +85,12 @@ def test_hook_updates_runtime_state_on_agent_swap():
cmd = MagicMock()
cmd.name = "/model"
_run(hook(ctx, original_agent, cmd))
await hook(ctx, original_agent, cmd)
assert state.agent is new_agent
def test_hook_syncs_channel_runtime():
async def test_hook_syncs_channel_runtime():
"""Other readers (the bus) look at ``ChannelRuntime.agent``; the
hook keeps the runtime in sync with the runtime state update."""
original_agent = _agent("original-agent")
@@ -109,13 +108,13 @@ def test_hook_syncs_channel_runtime():
cmd = MagicMock()
cmd.name = "/model"
_run(hook(ctx, original_agent, cmd))
await hook(ctx, original_agent, cmd)
assert runtime.agent is new_agent
assert runtime.thread_id == "t"
def test_hook_noop_when_agent_unchanged():
async def test_hook_noop_when_agent_unchanged():
"""Commands like ``/evoskills`` don't touch ``ctx.agent`` — the
runtime state must stay put."""
original_agent = _agent("original-agent")
@@ -128,12 +127,12 @@ def test_hook_noop_when_agent_unchanged():
cmd = MagicMock()
cmd.name = "/evoskills"
_run(hook(ctx, original_agent, cmd))
await hook(ctx, original_agent, cmd)
assert state.agent is original_agent
def test_hook_noop_when_ctx_agent_is_none():
async def test_hook_noop_when_ctx_agent_is_none():
"""Guard against commands that reset ``ctx.agent`` to ``None`` —
we never want to write ``None`` into runtime state."""
original_agent = _agent("original-agent")
@@ -146,12 +145,12 @@ def test_hook_noop_when_ctx_agent_is_none():
cmd = MagicMock()
cmd.name = "/whatever"
_run(hook(ctx, original_agent, cmd))
await hook(ctx, original_agent, cmd)
assert state.agent is original_agent
def test_hook_updates_thread_id_on_resume():
async def test_hook_updates_thread_id_on_resume():
"""``/resume`` mutates ``ctx.thread_id`` — the hook must push the
new id into runtime state so the outer poll loop runs subsequent
messages on the resumed thread."""
@@ -166,12 +165,12 @@ def test_hook_updates_thread_id_on_resume():
cmd = MagicMock()
cmd.name = "/resume"
_run(hook(ctx, agent, cmd))
await hook(ctx, agent, cmd)
assert state.thread_id == "new-tid"
def test_hook_updates_workspace_dir_on_resume():
async def test_hook_updates_workspace_dir_on_resume():
"""`/resume` can restore a different workspace; serve must reload for it."""
cfg = _config()
old_agent = _agent("old-agent")
@@ -201,7 +200,7 @@ def test_hook_updates_workspace_dir_on_resume():
return_value=reloaded_agent,
) as load_agent,
):
_run(hook(ctx, old_agent, cmd))
await hook(ctx, old_agent, cmd)
sync_server.assert_awaited_once_with(cfg, workspace_dir="/restored-ws")
load_agent.assert_called_once_with(workspace_dir="/restored-ws", config=cfg)
@@ -209,7 +208,7 @@ def test_hook_updates_workspace_dir_on_resume():
assert state.agent is reloaded_agent
def test_hook_syncs_channel_runtime_thread_id():
async def test_hook_syncs_channel_runtime_thread_id():
"""The bus reads ``ChannelRuntime.thread_id``; hook must sync it
alongside the runtime state update."""
agent = _agent("a")
@@ -224,12 +223,12 @@ def test_hook_syncs_channel_runtime_thread_id():
cmd = MagicMock()
cmd.name = "/resume"
_run(hook(ctx, agent, cmd))
await hook(ctx, agent, cmd)
assert runtime.thread_id == "new-tid"
def test_hook_noop_when_thread_id_unchanged():
async def test_hook_noop_when_thread_id_unchanged():
"""Most commands don't touch thread_id — runtime state stays put."""
agent = _agent("a")
state = _runtime_state(agent=agent, thread_id="same-tid")
@@ -241,12 +240,12 @@ def test_hook_noop_when_thread_id_unchanged():
cmd = MagicMock()
cmd.name = "/evoskills"
_run(hook(ctx, agent, cmd))
await hook(ctx, agent, cmd)
assert state.thread_id == "same-tid"
def test_hook_skips_resume_warning_when_thread_unchanged():
async def test_hook_skips_resume_warning_when_thread_unchanged():
"""Bare ``/resume`` with no argument prints usage but leaves
``ctx.thread_id`` unchanged — the in-memory-state warning must NOT
fire because no resume actually happened."""
@@ -261,13 +260,13 @@ def test_hook_skips_resume_warning_when_thread_unchanged():
cmd = MagicMock()
cmd.name = "/resume"
_run(hook(ctx, agent, cmd))
await hook(ctx, agent, cmd)
ctx.ui.append_system.assert_not_called()
ctx.ui.flush.assert_not_called()
def test_hook_emits_resume_warning_when_thread_changed():
async def test_hook_emits_resume_warning_when_thread_changed():
"""``/resume <tid>`` that actually changes thread_id must surface
the in-memory-state warning via ``ctx.ui``."""
agent = _agent("a")
@@ -283,7 +282,7 @@ def test_hook_emits_resume_warning_when_thread_changed():
cmd = MagicMock()
cmd.name = "/resume"
_run(hook(ctx, agent, cmd))
await hook(ctx, agent, cmd)
ctx.ui.append_system.assert_called_once()
warn_text, warn_kwargs = (
@@ -296,7 +295,7 @@ def test_hook_emits_resume_warning_when_thread_changed():
ctx.ui.flush.assert_awaited_once()
def test_start_new_session_cb_rotates_thread_id():
async def test_start_new_session_cb_rotates_thread_id():
"""``/new`` via channel calls this callback — must generate a new
thread id, push into runtime state, and sync the channel runtime."""
agent = _agent("a")
@@ -311,13 +310,13 @@ def test_start_new_session_cb_rotates_thread_id():
state,
runtime,
)
_run(cb())
await cb()
assert state.thread_id == "freshly-generated-tid"
assert runtime.thread_id == "freshly-generated-tid"
def test_start_new_session_cb_leaves_agent_alone():
async def test_start_new_session_cb_leaves_agent_alone():
"""``/new`` rotates thread only — agent handle must stay put
(serve's agent is a single pre-loaded instance, not per-thread)."""
agent = _agent("a")
@@ -328,12 +327,12 @@ def test_start_new_session_cb_leaves_agent_alone():
)
cb = _make_serve_start_new_session_cb(state)
_run(cb())
await cb()
assert state.agent is agent
def test_serve_resume_callback_syncs_reloads_and_adopts_workspace():
async def test_serve_resume_callback_syncs_reloads_and_adopts_workspace():
cfg = _config()
old_agent = _agent("old-agent")
reloaded_agent = _agent("reloaded-agent")
@@ -364,7 +363,7 @@ def test_serve_resume_callback_syncs_reloads_and_adopts_workspace():
side_effect=_load_agent,
) as load_agent,
):
_run(cb("new-tid", "/new-ws"))
await cb("new-tid", "/new-ws")
sync_server.assert_awaited_once_with(cfg, workspace_dir="/new-ws")
load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg)
@@ -376,7 +375,7 @@ def test_serve_resume_callback_syncs_reloads_and_adopts_workspace():
assert runtime.agent is reloaded_agent
def test_hook_emits_resume_warning_after_resume_callback_adopts_thread():
async def test_hook_emits_resume_warning_after_resume_callback_adopts_thread():
cfg = _config()
old_agent = _agent("old-agent")
reloaded_agent = _agent("reloaded-agent")
@@ -399,7 +398,7 @@ def test_hook_emits_resume_warning_after_resume_callback_adopts_thread():
return_value=reloaded_agent,
),
):
_run(cb("abc12345-resumed-tid", "/new-ws"))
await cb("abc12345-resumed-tid", "/new-ws")
hook = _make_serve_cmd_completed_hook(state, runtime, config=cfg)
ctx = MagicMock()
@@ -410,14 +409,14 @@ def test_hook_emits_resume_warning_after_resume_callback_adopts_thread():
cmd = MagicMock()
cmd.name = "/resume"
_run(hook(ctx, reloaded_agent, cmd))
await hook(ctx, reloaded_agent, cmd)
ctx.ui.append_system.assert_called_once()
assert "in-memory state" in ctx.ui.append_system.call_args.args[0]
ctx.ui.flush.assert_awaited_once()
def test_serve_resume_callback_preserves_state_when_sync_fails():
async def test_serve_resume_callback_preserves_state_when_sync_fails():
cfg = _config()
old_agent = _agent("old-agent")
loaded_but_not_adopted = _agent("loaded-but-not-adopted")
@@ -442,7 +441,7 @@ def test_serve_resume_callback_preserves_state_when_sync_fails():
patch("EvoScientist.cli.commands.set_active_workspace") as set_active,
pytest.raises(RuntimeError, match="workspace conflict"),
):
_run(cb("new-tid", "/new-ws"))
await cb("new-tid", "/new-ws")
load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg)
set_active.assert_called_once_with("/old-ws")
@@ -455,7 +454,7 @@ def test_serve_resume_callback_preserves_state_when_sync_fails():
assert runtime.thread_id == "old-tid"
def test_serve_resume_callback_load_failure_does_not_sync_or_adopt():
async def test_serve_resume_callback_load_failure_does_not_sync_or_adopt():
cfg = _config()
old_agent = _agent("old-agent")
state = _runtime_state(
@@ -479,7 +478,7 @@ def test_serve_resume_callback_load_failure_does_not_sync_or_adopt():
) as sync_server,
pytest.raises(RuntimeError, match="load failed"),
):
_run(cb("new-tid", "/new-ws"))
await cb("new-tid", "/new-ws")
load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg)
set_active.assert_called_once_with("/old-ws")
@@ -493,7 +492,7 @@ def test_serve_resume_callback_load_failure_does_not_sync_or_adopt():
assert runtime.thread_id == "old-tid"
def test_hook_handles_both_agent_and_thread_swap():
async def test_hook_handles_both_agent_and_thread_swap():
"""Edge case: a command that changes both (hypothetical). Both
updates must land in runtime state."""
old_agent = _agent("old-agent")
@@ -506,7 +505,7 @@ def test_hook_handles_both_agent_and_thread_swap():
ctx.thread_id = "new-tid"
cmd = MagicMock()
_run(hook(ctx, old_agent, cmd))
await hook(ctx, old_agent, cmd)
assert state.agent is new_agent
assert state.thread_id == "new-tid"
+222 -220
View File
File diff suppressed because it is too large Load Diff
+8 -9
View File
@@ -4,7 +4,6 @@ import pytest
from EvoScientist.channels.base import ChannelError
from EvoScientist.channels.slack.channel import SlackChannel, SlackConfig
from tests.conftest import run_async as _run
class TestSlackConfig:
@@ -38,24 +37,24 @@ class TestSlackChannel:
assert channel.config is config
assert channel._running is False
def test_start_raises_without_bot_token(self):
async def test_start_raises_without_bot_token(self):
config = SlackConfig(bot_token="", app_token="xapp-test")
channel = SlackChannel(config)
with pytest.raises(ChannelError, match="bot token"):
_run(channel.start())
await channel.start()
def test_start_raises_without_app_token(self):
async def test_start_raises_without_app_token(self):
config = SlackConfig(bot_token="xoxb-test", app_token="")
channel = SlackChannel(config)
with pytest.raises(ChannelError, match="app token"):
_run(channel.start())
await channel.start()
def test_stop_when_not_running(self):
async def test_stop_when_not_running(self):
config = SlackConfig(bot_token="xoxb-test", app_token="xapp-test")
channel = SlackChannel(config)
_run(channel.stop())
await channel.stop()
def test_send_returns_false_without_client(self):
async def test_send_returns_false_without_client(self):
from EvoScientist.channels.base import OutboundMessage
config = SlackConfig(bot_token="xoxb-test", app_token="xapp-test")
@@ -66,7 +65,7 @@ class TestSlackChannel:
content="hello",
metadata={"chat_id": "C123"},
)
result = _run(channel.send(msg))
result = await channel.send(msg)
assert result is False
+7 -12
View File
@@ -2,7 +2,6 @@
from __future__ import annotations
import asyncio
from datetime import datetime, timedelta
from typing import ClassVar
@@ -243,7 +242,7 @@ def test_build_status_text_uses_rich_styles():
assert text.spans
def test_build_session_status_snapshot_uses_fallback_window(monkeypatch):
async def test_build_session_status_snapshot_uses_fallback_window(monkeypatch):
class _FakeModel:
model_name: ClassVar[str] = "provider/demo-model"
profile: ClassVar[dict[str, object]] = {}
@@ -262,16 +261,12 @@ def test_build_session_status_snapshot_uses_fallback_window(monkeypatch):
_fake_count,
)
snapshot = asyncio.run(
build_session_status_snapshot(
"thread-1",
pending_user_text="pending",
graph_gateway=FakeGraphGateway(
thread_store=FakeThreadStore(
messages=[HumanMessage(content="existing")]
)
),
)
snapshot = await build_session_status_snapshot(
"thread-1",
pending_user_text="pending",
graph_gateway=FakeGraphGateway(
thread_store=FakeThreadStore(messages=[HumanMessage(content="existing")])
),
)
assert snapshot.model_full == "provider/demo-model"

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