Compare commits
10 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 3ce5614254 | |||
| 4fc74e7da7 | |||
| 753c745405 | |||
| 88ac9f5ba1 | |||
| 952e68efe3 | |||
| 2b28c46caf | |||
| da6ca38d53 | |||
| 19888c2db6 | |||
| f3e65a446f | |||
| f72f7b93d5 |
+132
-21
@@ -19,6 +19,7 @@ Usage:
|
|||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
|
from collections.abc import Sequence
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING
|
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()``.
|
path doesn't fall back to the global-writing ``_ensure_chat_model()``.
|
||||||
"""
|
"""
|
||||||
from .middleware import (
|
from .middleware import (
|
||||||
|
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
||||||
ContextOverflowMapperMiddleware,
|
ContextOverflowMapperMiddleware,
|
||||||
|
ErrorNormalizationMiddleware,
|
||||||
|
RepetitiveToolCallGuardMiddleware,
|
||||||
ToolErrorHandlerMiddleware,
|
ToolErrorHandlerMiddleware,
|
||||||
|
ToolProtocolGuardMiddleware,
|
||||||
create_context_editing_middleware,
|
create_context_editing_middleware,
|
||||||
create_memory_lifecycle_middleware,
|
create_memory_lifecycle_middleware,
|
||||||
create_memory_middleware,
|
create_memory_middleware,
|
||||||
@@ -314,6 +319,16 @@ def _inject_subagent_middleware(
|
|||||||
)
|
)
|
||||||
|
|
||||||
cfg = cfg if cfg is not None else _ensure_config()
|
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_controls = MemoryControls.from_config(cfg)
|
||||||
memory_dir = str(_paths_mod.MEMORIES_DIR)
|
memory_dir = str(_paths_mod.MEMORIES_DIR)
|
||||||
memory_scheduler = default_memory_scheduler()
|
memory_scheduler = default_memory_scheduler()
|
||||||
@@ -333,6 +348,16 @@ def _inject_subagent_middleware(
|
|||||||
memory_scheduler=memory_scheduler,
|
memory_scheduler=memory_scheduler,
|
||||||
)
|
)
|
||||||
middleware = [
|
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
|
# Subagents share the main agent's model: use the threaded
|
||||||
# ``chat_model`` on the pure path, else defer to the factory's
|
# ``chat_model`` on the pure path, else defer to the factory's
|
||||||
# ``_ensure_chat_model()`` fallback (when ``chat_model=None``).
|
# ``_ensure_chat_model()`` fallback (when ``chat_model=None``).
|
||||||
@@ -641,9 +666,14 @@ def _get_default_middleware(
|
|||||||
*,
|
*,
|
||||||
for_async_subagent: bool = False,
|
for_async_subagent: bool = False,
|
||||||
workspace_dir: str | Path | None = None,
|
workspace_dir: str | Path | None = None,
|
||||||
|
memory_dir: str | Path | None = None,
|
||||||
cfg=None,
|
cfg=None,
|
||||||
chat_model=None,
|
chat_model=None,
|
||||||
memory_source_agent: str = "EvoScientist",
|
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.
|
"""Build the default middleware list.
|
||||||
|
|
||||||
@@ -665,10 +695,14 @@ def _get_default_middleware(
|
|||||||
Async sub-agent factories pass their deployed agent name here.
|
Async sub-agent factories pass their deployed agent name here.
|
||||||
"""
|
"""
|
||||||
from .middleware import (
|
from .middleware import (
|
||||||
|
DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD,
|
||||||
ConfigurableModelMiddleware,
|
ConfigurableModelMiddleware,
|
||||||
ContextOverflowMapperMiddleware,
|
ContextOverflowMapperMiddleware,
|
||||||
|
ErrorNormalizationMiddleware,
|
||||||
ModelFallbackMiddleware,
|
ModelFallbackMiddleware,
|
||||||
|
RepetitiveToolCallGuardMiddleware,
|
||||||
ToolErrorHandlerMiddleware,
|
ToolErrorHandlerMiddleware,
|
||||||
|
ToolProtocolGuardMiddleware,
|
||||||
create_code_interpreter_middleware,
|
create_code_interpreter_middleware,
|
||||||
create_context_editing_middleware,
|
create_context_editing_middleware,
|
||||||
create_memory_lifecycle_middleware,
|
create_memory_lifecycle_middleware,
|
||||||
@@ -681,10 +715,20 @@ def _get_default_middleware(
|
|||||||
)
|
)
|
||||||
|
|
||||||
cfg = cfg if cfg is not None else _ensure_config()
|
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:
|
if cfg.model_fallbacks:
|
||||||
load_fallback_chain(cfg.model_fallbacks)
|
load_fallback_chain(cfg.model_fallbacks)
|
||||||
model = chat_model if chat_model is not None else _ensure_chat_model()
|
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 = (
|
source_type = (
|
||||||
MemorySourceType.SUBAGENT if for_async_subagent else MemorySourceType.TURN
|
MemorySourceType.SUBAGENT if for_async_subagent else MemorySourceType.TURN
|
||||||
)
|
)
|
||||||
@@ -699,18 +743,20 @@ def _get_default_middleware(
|
|||||||
# ``ModelFallbackMiddleware``: a configurable.model override sets the
|
# ``ModelFallbackMiddleware``: a configurable.model override sets the
|
||||||
# PRIMARY model only, leaving the fallback chain free to try its own
|
# PRIMARY model only, leaving the fallback chain free to try its own
|
||||||
# alternatives instead of re-overriding every retry to the same model.
|
# alternatives instead of re-overriding every retry to the same model.
|
||||||
memory_middleware = create_memory_middleware(
|
memory_kwargs = {
|
||||||
memory_dir,
|
"workspace_dir": workspace_dir,
|
||||||
workspace_dir=workspace_dir,
|
"source_type": source_type,
|
||||||
source_type=source_type,
|
"source_agent": memory_source_agent,
|
||||||
source_agent=memory_source_agent,
|
"enable_profile_memory": memory_controls.profile_enabled,
|
||||||
enable_profile_memory=memory_controls.profile_enabled,
|
"enable_observation_memory": memory_controls.observations_enabled,
|
||||||
enable_observation_memory=memory_controls.observations_enabled,
|
"enable_observation_tool": memory_controls.observation_tool_enabled(
|
||||||
enable_observation_tool=memory_controls.observation_tool_enabled(
|
|
||||||
MemoryObservationTarget.AGENT
|
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
|
# 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).
|
# 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
|
# 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
|
from .llm import get_chat_model
|
||||||
|
|
||||||
tool_selector_model = get_chat_model(model=aux_model, provider=aux_provider)
|
tool_selector_model = get_chat_model(model=aux_model, provider=aux_provider)
|
||||||
mw = [
|
selector_middlewares = create_tool_selector_middleware(
|
||||||
ConfigurableModelMiddleware(),
|
**(
|
||||||
create_context_editing_middleware(model),
|
{"threshold": tool_selector_threshold}
|
||||||
ModelFallbackMiddleware(),
|
if tool_selector_threshold is not None
|
||||||
ContextOverflowMapperMiddleware(),
|
else {}
|
||||||
ToolErrorHandlerMiddleware(),
|
),
|
||||||
*create_tool_selector_middleware(
|
|
||||||
model=tool_selector_model,
|
model=tool_selector_model,
|
||||||
track_stream_selection=not for_async_subagent,
|
track_stream_selection=not for_async_subagent,
|
||||||
|
)
|
||||||
|
mw = [
|
||||||
|
# Outermost — catches provider-SDK exceptions from the model
|
||||||
|
# call (including exceptions surfaced through inner
|
||||||
|
# middlewares) and normalizes them into a non-dataclass
|
||||||
|
# envelope wrapper before anything downstream sees them.
|
||||||
|
ErrorNormalizationMiddleware(),
|
||||||
|
ConfigurableModelMiddleware(),
|
||||||
|
create_context_editing_middleware(model),
|
||||||
|
*([ModelFallbackMiddleware()] if enable_legacy_model_fallback else []),
|
||||||
|
RepetitiveToolCallGuardMiddleware(
|
||||||
|
threshold=repetitive_tool_call_threshold,
|
||||||
|
max_consecutive_errors=max_consecutive_tool_errors,
|
||||||
),
|
),
|
||||||
|
ContextOverflowMapperMiddleware(),
|
||||||
|
ToolErrorHandlerMiddleware(),
|
||||||
|
*selector_middlewares,
|
||||||
|
ToolProtocolGuardMiddleware(),
|
||||||
# Interpreter prompt must land before runtime/memory context, so this
|
# Interpreter prompt must land before runtime/memory context, so this
|
||||||
# middleware sits ahead of runtime_context in the stack.
|
# middleware sits ahead of runtime_context in the stack.
|
||||||
create_code_interpreter_middleware(
|
create_code_interpreter_middleware(
|
||||||
@@ -770,7 +832,7 @@ def _get_default_middleware(
|
|||||||
# Background-process tools (run_in_background / check_process / stop_process /
|
# Background-process tools (run_in_background / check_process / stop_process /
|
||||||
# list_processes) — main agent only. Async sub-agents run on langgraph-dev and
|
# list_processes) — main agent only. Async sub-agents run on langgraph-dev and
|
||||||
# must not spawn local OS processes.
|
# 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
|
from .middleware.background import BackgroundExecutionMiddleware
|
||||||
|
|
||||||
mw.append(BackgroundExecutionMiddleware())
|
mw.append(BackgroundExecutionMiddleware())
|
||||||
@@ -868,6 +930,14 @@ def create_cli_agent(
|
|||||||
chat_model=None,
|
chat_model=None,
|
||||||
*,
|
*,
|
||||||
on_mcp_progress=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":
|
) -> "CompiledStateGraph":
|
||||||
"""Create agent with checkpointer for CLI multi-turn support.
|
"""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
|
chat_model: Optional pre-built chat model. Only triggers the pure
|
||||||
path when ``config`` is also explicit; otherwise it is ignored in
|
path when ``config`` is also explicit; otherwise it is ignored in
|
||||||
favor of the ``_ensure_chat_model()`` fallback.
|
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
|
import os as _os
|
||||||
|
|
||||||
@@ -935,13 +1021,15 @@ def create_cli_agent(
|
|||||||
workspace_dir = str(_paths.WORKSPACE_ROOT)
|
workspace_dir = str(_paths.WORKSPACE_ROOT)
|
||||||
|
|
||||||
# Read paths dynamically so runtime set_workspace_root() changes are picked up
|
# 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)
|
_usr_skills_dir = str(_paths.USER_SKILLS_DIR)
|
||||||
_global_skills_dir = str(_paths.GLOBAL_SKILLS_DIR)
|
_global_skills_dir = str(_paths.GLOBAL_SKILLS_DIR)
|
||||||
|
|
||||||
# Always construct fresh backends from current paths (avoids stale
|
# Always construct fresh backends from current paths (avoids stale
|
||||||
# module-level backend when workspace root changed at runtime).
|
# module-level backend when workspace root changed at runtime).
|
||||||
set_active_workspace(workspace_dir)
|
set_active_workspace(workspace_dir)
|
||||||
|
ws_backend = workspace_backend
|
||||||
|
if ws_backend is None:
|
||||||
ws_backend = CustomSandboxBackend(
|
ws_backend = CustomSandboxBackend(
|
||||||
root_dir=workspace_dir,
|
root_dir=workspace_dir,
|
||||||
virtual_mode=True,
|
virtual_mode=True,
|
||||||
@@ -969,8 +1057,29 @@ def create_cli_agent(
|
|||||||
# CLI agent never drifts from the default chain. Anything CLI-specific
|
# CLI agent never drifts from the default chain. Anything CLI-specific
|
||||||
# (e.g. ``HumanInTheLoopMiddleware``) is appended below.
|
# (e.g. ``HumanInTheLoopMiddleware``) is appended below.
|
||||||
mw: list[AgentMiddleware] = _get_default_middleware(
|
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
|
# HITL on main agent only — passing `interrupt_on=` to create_deep_agent
|
||||||
# would propagate it to every subagent, breaking parallel execute calls
|
# would propagate it to every subagent, breaking parallel execute calls
|
||||||
@@ -995,6 +1104,8 @@ def create_cli_agent(
|
|||||||
chat_model=chat_model,
|
chat_model=chat_model,
|
||||||
workspace_dir=workspace_dir,
|
workspace_dir=workspace_dir,
|
||||||
)
|
)
|
||||||
|
if not enable_subagents:
|
||||||
|
kwargs = {**kwargs, "subagents": []}
|
||||||
|
|
||||||
return create_deep_agent(
|
return create_deep_agent(
|
||||||
**kwargs,
|
**kwargs,
|
||||||
|
|||||||
@@ -9,6 +9,8 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from importlib import import_module
|
from importlib import import_module
|
||||||
|
|
||||||
|
__version__ = "0.2.2"
|
||||||
|
|
||||||
_EXPORTS: dict[str, tuple[str, str]] = {
|
_EXPORTS: dict[str, tuple[str, str]] = {
|
||||||
# Agent graph (lazy to avoid expensive initialization at import time)
|
# Agent graph (lazy to avoid expensive initialization at import time)
|
||||||
"EvoScientist_agent": (".EvoScientist", "EvoScientist_agent"),
|
"EvoScientist_agent": (".EvoScientist", "EvoScientist_agent"),
|
||||||
|
|||||||
@@ -20,6 +20,9 @@ from EvoScientist.config import EvoScientistConfig
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_CCPROXY_AUTH_TIMEOUT_SECONDS = 30
|
||||||
|
_CCPROXY_HEALTH_TIMEOUT_SECONDS = 180
|
||||||
|
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
# Availability & auth checks
|
# Availability & auth checks
|
||||||
@@ -127,7 +130,11 @@ def check_ccproxy_auth(provider: str = "claude_api") -> tuple[bool, str]:
|
|||||||
[exe, "auth", "status", provider],
|
[exe, "auth", "status", provider],
|
||||||
capture_output=True,
|
capture_output=True,
|
||||||
text=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
|
import re as _re
|
||||||
|
|
||||||
@@ -176,6 +183,33 @@ def is_ccproxy_running(port: int) -> bool:
|
|||||||
return False
|
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:
|
def start_ccproxy(port: int) -> subprocess.Popen:
|
||||||
"""Start ccproxy serve as a background process.
|
"""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.
|
The Popen handle for the ccproxy process.
|
||||||
|
|
||||||
Raises:
|
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.
|
FileNotFoundError: If ccproxy binary is not found.
|
||||||
"""
|
"""
|
||||||
exe = _ccproxy_exe() or "ccproxy"
|
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(
|
proc = subprocess.Popen(
|
||||||
[exe, "serve", "--port", str(port)],
|
cmd,
|
||||||
stdout=subprocess.DEVNULL,
|
stdout=subprocess.DEVNULL,
|
||||||
stderr=subprocess.DEVNULL,
|
stderr=subprocess.DEVNULL,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Wait for health (ccproxy can take up to ~11s on first start)
|
deadline = time.monotonic() + _CCPROXY_HEALTH_TIMEOUT_SECONDS
|
||||||
deadline = time.monotonic() + 30
|
|
||||||
while time.monotonic() < deadline:
|
while time.monotonic() < deadline:
|
||||||
if proc.poll() is not None:
|
if proc.poll() is not None:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
@@ -213,7 +261,10 @@ def start_ccproxy(port: int) -> subprocess.Popen:
|
|||||||
proc.wait(timeout=3)
|
proc.wait(timeout=3)
|
||||||
except subprocess.TimeoutExpired:
|
except subprocess.TimeoutExpired:
|
||||||
proc.kill()
|
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:
|
def stop_ccproxy(proc: subprocess.Popen | None) -> None:
|
||||||
|
|||||||
@@ -106,11 +106,23 @@ def _normalize_hhmm(value: Any) -> str | None:
|
|||||||
def get_config_dir() -> Path:
|
def get_config_dir() -> Path:
|
||||||
"""Get the configuration directory 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")
|
xdg_config = os.environ.get("XDG_CONFIG_HOME")
|
||||||
if xdg_config:
|
if xdg_config:
|
||||||
return Path(xdg_config) / "evoscientist"
|
return Path(xdg_config).expanduser() / "evoscientist"
|
||||||
return Path.home() / ".config" / "evoscientist"
|
return Path.home() / ".config" / "evoscientist"
|
||||||
|
|
||||||
|
|
||||||
@@ -123,6 +135,16 @@ def get_config_path() -> Path:
|
|||||||
# Configuration dataclass
|
# 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
|
@dataclass
|
||||||
class EvoScientistConfig:
|
class EvoScientistConfig:
|
||||||
@@ -240,6 +262,13 @@ class EvoScientistConfig:
|
|||||||
# Lower (e.g., 5000) if you want a tighter safety net against runaway loops.
|
# Lower (e.g., 5000) if you want a tighter safety net against runaway loops.
|
||||||
recursion_limit: int = 1_000_000
|
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
|
# Memory Settings
|
||||||
# Profile memory injects and maintains `/memories/profile/...` files.
|
# Profile memory injects and maintains `/memories/profile/...` files.
|
||||||
memory_profile_enabled: bool = True
|
memory_profile_enabled: bool = True
|
||||||
@@ -278,10 +307,21 @@ class EvoScientistConfig:
|
|||||||
# a deploy-style langgraph server instead of the in-terminal CLI/TUI.
|
# a deploy-style langgraph server instead of the in-terminal CLI/TUI.
|
||||||
ui_backend: Literal["cli", "tui", "webui"] = "tui"
|
ui_backend: Literal["cli", "tui", "webui"] = "tui"
|
||||||
log_level: str = "warning"
|
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
|
# Anthropic prompt caching for OpenRouter anthropic/* models. Opt out if
|
||||||
# cache-write costs outweigh the benefit for a workflow.
|
# cache-write costs outweigh the benefit for a workflow.
|
||||||
openrouter_anthropic_prompt_cache: bool = True
|
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 Settings
|
||||||
channel_enabled: str = "" # "imessage" | "telegram" | "discord" | "slack" | "wechat" | "dingtalk" | "feishu" | "email" | "qq" | "signal" | "" (comma-separated for multiple)
|
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"
|
stt_compute_type: str = "int8" # "int8" | "float16" | "float32"
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
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
|
# 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
|
# config file value — load_config does not coerce file values — or a
|
||||||
# 0/negative env value) would raise inside CustomSandboxBackend.__init__
|
# 0/negative env value) would raise inside CustomSandboxBackend.__init__
|
||||||
@@ -702,6 +750,11 @@ def set_config_value(key: str, value: Any) -> bool:
|
|||||||
|
|
||||||
if key == "sandbox_execute_timeout" and value <= 0:
|
if key == "sandbox_execute_timeout" and value <= 0:
|
||||||
return False
|
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":
|
if key == "memory_skill_synthesis_time":
|
||||||
value = _normalize_hhmm(value)
|
value = _normalize_hhmm(value)
|
||||||
if value is None:
|
if value is None:
|
||||||
@@ -761,6 +814,9 @@ _ENV_MAPPINGS = {
|
|||||||
"openrouter_anthropic_prompt_cache": (
|
"openrouter_anthropic_prompt_cache": (
|
||||||
"EVOSCIENTIST_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",
|
"dangerous_mode": "EVOSCIENTIST_DANGEROUS_MODE",
|
||||||
"channel_debug_tracing": "EVOSCIENTIST_CHANNEL_DEBUG_TRACING",
|
"channel_debug_tracing": "EVOSCIENTIST_CHANNEL_DEBUG_TRACING",
|
||||||
"ccproxy_port": "EVOSCIENTIST_CCPROXY_PORT",
|
"ccproxy_port": "EVOSCIENTIST_CCPROXY_PORT",
|
||||||
@@ -777,6 +833,10 @@ _ENV_MAPPINGS = {
|
|||||||
"langgraph_dev_file_persistence": "EVOSCIENTIST_LANGGRAPH_DEV_FILE_PERSISTENCE",
|
"langgraph_dev_file_persistence": "EVOSCIENTIST_LANGGRAPH_DEV_FILE_PERSISTENCE",
|
||||||
"langgraph_dev_jobs_per_worker": "EVOSCIENTIST_LANGGRAPH_DEV_JOBS_PER_WORKER",
|
"langgraph_dev_jobs_per_worker": "EVOSCIENTIST_LANGGRAPH_DEV_JOBS_PER_WORKER",
|
||||||
"recursion_limit": "EVOSCIENTIST_RECURSION_LIMIT",
|
"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_profile_enabled": "EVOSCIENTIST_MEMORY_PROFILE_ENABLED",
|
||||||
"memory_observations_enabled": "EVOSCIENTIST_MEMORY_OBSERVATIONS_ENABLED",
|
"memory_observations_enabled": "EVOSCIENTIST_MEMORY_OBSERVATIONS_ENABLED",
|
||||||
"memory_observation_writer": "EVOSCIENTIST_MEMORY_OBSERVATION_WRITER",
|
"memory_observation_writer": "EVOSCIENTIST_MEMORY_OBSERVATION_WRITER",
|
||||||
@@ -892,6 +952,22 @@ def apply_config_to_env(config: EvoScientistConfig) -> None:
|
|||||||
os.environ["TAVILY_API_KEY"] = config.tavily_api_key
|
os.environ["TAVILY_API_KEY"] = config.tavily_api_key
|
||||||
if config.reasoning_effort and not os.environ.get("EVOSCIENTIST_REASONING_EFFORT"):
|
if config.reasoning_effort and not os.environ.get("EVOSCIENTIST_REASONING_EFFORT"):
|
||||||
os.environ["EVOSCIENTIST_REASONING_EFFORT"] = config.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(
|
if not config.openrouter_anthropic_prompt_cache and not os.environ.get(
|
||||||
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE"
|
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE"
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -43,7 +43,7 @@ async def get_models(_request: Request) -> JSONResponse:
|
|||||||
``discover_ollama_models()`` call, same 1.5-s timeout, same
|
``discover_ollama_models()`` call, same 1.5-s timeout, same
|
||||||
fail-soft semantics (the probe returns ``[]`` on any error, never
|
fail-soft semantics (the probe returns ``[]`` on any error, never
|
||||||
raises). The TUI's "Custom Ollama model…" sentinel is intentionally
|
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.
|
the registry surface.
|
||||||
|
|
||||||
``default`` reflects the deployment's currently-configured fallback
|
``default`` reflects the deployment's currently-configured fallback
|
||||||
|
|||||||
@@ -5,8 +5,244 @@ in ``EvoScientist/EvoScientist.py`` so it doesn't construct on plain
|
|||||||
``import EvoScientist``. ``langgraph dev`` 's symbol resolver inspects
|
``import EvoScientist``. ``langgraph dev`` 's symbol resolver inspects
|
||||||
module attributes directly and doesn't trigger ``__getattr__``, so we
|
module attributes directly and doesn't trigger ``__getattr__``, so we
|
||||||
re-export here to make it visible.
|
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"]
|
__all__ = ["EvoScientist_agent"]
|
||||||
|
|||||||
@@ -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
|
||||||
+232
-23
@@ -10,11 +10,19 @@ endpoints) and convenient short names for common models.
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import os
|
import os
|
||||||
|
import re
|
||||||
|
import subprocess
|
||||||
import warnings
|
import warnings
|
||||||
|
from functools import lru_cache
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from langchain.chat_models import init_chat_model
|
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 .context_window import apply_known_context_window
|
||||||
from .patches import (
|
from .patches import (
|
||||||
_is_ccproxy_codex,
|
_is_ccproxy_codex,
|
||||||
@@ -37,6 +45,50 @@ _DEEPSEEK_BASE_URL = "https://api.deepseek.com"
|
|||||||
_MOONSHOT_BASE_URL = "https://api.moonshot.cn/v1"
|
_MOONSHOT_BASE_URL = "https://api.moonshot.cn/v1"
|
||||||
_KIMI_CODING_BASE_URL = "https://api.kimi.com/coding/"
|
_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.
|
# Providers routed through the OpenAI provider with a custom base_url.
|
||||||
# Maps provider name → (base_url or None, env var for API key).
|
# Maps provider name → (base_url or None, env var for API key).
|
||||||
_OPENAI_ROUTED_PROVIDERS: dict[str, tuple[str | None, str]] = {
|
_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"}
|
_TRUTHY_ENV_VALUES = {"1", "true", "yes", "on"}
|
||||||
_FALSEY_ENV_VALUES = {"0", "false", "no", "off"}
|
_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)
|
# Model registry: list of (short_name, model_id, provider)
|
||||||
# Allows same short_name across different providers.
|
# Allows same short_name across different providers.
|
||||||
_MODEL_ENTRIES: list[tuple[str, str, str]] = [
|
_MODEL_ENTRIES: list[tuple[str, str, str]] = [
|
||||||
@@ -264,6 +329,15 @@ def _env_flag_disabled(name: str) -> bool:
|
|||||||
return value is not None and value.strip().lower() in _FALSEY_ENV_VALUES
|
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:
|
def _supports_openrouter_anthropic_prompt_cache(provider: str, model_id: str) -> bool:
|
||||||
"""Return whether EvoScientist should declare OpenRouter Claude caching."""
|
"""Return whether EvoScientist should declare OpenRouter Claude caching."""
|
||||||
return provider == "openrouter" and model_id.startswith(
|
return provider == "openrouter" and model_id.startswith(
|
||||||
@@ -321,8 +395,16 @@ def _apply_auto_config(
|
|||||||
Mutates *kwargs* in place. Only sets keys that the caller hasn't already
|
Mutates *kwargs* in place. Only sets keys that the caller hasn't already
|
||||||
provided, so explicit user settings are never overridden.
|
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
|
# 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
|
_supports_thinking = original_provider in _THINKING_CAPABLE_PROVIDERS
|
||||||
# Detect local proxy (e.g. ccproxy): thinking blocks in conversation
|
# Detect local proxy (e.g. ccproxy): thinking blocks in conversation
|
||||||
# history cause 422 errors because the proxy doesn't accept 'thinking'
|
# history cause 422 errors because the proxy doesn't accept 'thinking'
|
||||||
@@ -341,24 +423,31 @@ def _apply_auto_config(
|
|||||||
kwargs["thinking"] = {"type": "enabled", "budget_tokens": 10000}
|
kwargs["thinking"] = {"type": "enabled", "budget_tokens": 10000}
|
||||||
|
|
||||||
# OpenAI (native, not third-party routed): reasoning
|
# OpenAI (native, not third-party routed): reasoning
|
||||||
if provider == "openai" and not is_third_party and "reasoning" not in kwargs:
|
if (
|
||||||
if _is_ccproxy_codex():
|
provider == "openai"
|
||||||
# ccproxy uses Chat Completions which doesn't support reasoning.
|
and not is_third_party
|
||||||
pass
|
and not disable_reasoning
|
||||||
else:
|
and "reasoning" not in kwargs
|
||||||
_eff = (
|
):
|
||||||
|
_default_effort = (
|
||||||
"xhigh"
|
"xhigh"
|
||||||
if ("5.4" in model_id or "5.5" in model_id or "codex" in model_id)
|
if (
|
||||||
|
"5.4" in model_id
|
||||||
|
or "5.5" in model_id
|
||||||
|
or "5.6" in model_id
|
||||||
|
or "codex" in model_id
|
||||||
|
)
|
||||||
else "high"
|
else "high"
|
||||||
)
|
)
|
||||||
|
_eff = _resolve_reasoning_effort(_default_effort)
|
||||||
kwargs["reasoning"] = {"effort": _eff, "summary": "auto"}
|
kwargs["reasoning"] = {"effort": _eff, "summary": "auto"}
|
||||||
|
|
||||||
# Google GenAI: surface thinking traces
|
# Google GenAI: surface thinking traces
|
||||||
if provider == "google-genai":
|
if provider == "google-genai" and not disable_reasoning:
|
||||||
kwargs.setdefault("include_thoughts", True)
|
kwargs.setdefault("include_thoughts", True)
|
||||||
|
|
||||||
# Ollama: separate reasoning content from response for thinking models
|
# 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
|
kwargs["reasoning"] = True
|
||||||
|
|
||||||
|
|
||||||
@@ -385,6 +474,45 @@ def get_chat_model(
|
|||||||
>>> model = get_chat_model("gpt-4o") # OpenAI model
|
>>> model = get_chat_model("gpt-4o") # OpenAI model
|
||||||
>>> model = get_chat_model("claude-3-opus-20240229", provider="anthropic") # Full ID
|
>>> model = get_chat_model("claude-3-opus-20240229", provider="anthropic") # Full ID
|
||||||
"""
|
"""
|
||||||
|
skip_runtime_resolver = bool(kwargs.pop("_skip_runtime_model_resolver", False))
|
||||||
|
runtime_provider_name: str | None = None
|
||||||
|
runtime_supports_reasoning: bool | None = None
|
||||||
|
runtime_resolved = None
|
||||||
|
if not skip_runtime_resolver:
|
||||||
|
from EvoScientist.runtime_integrations import resolve_runtime_model
|
||||||
|
|
||||||
|
runtime_resolved = resolve_runtime_model(model, provider)
|
||||||
|
|
||||||
|
if runtime_resolved is not None:
|
||||||
|
resolved_params = dict(getattr(runtime_resolved, "params", {}) or {})
|
||||||
|
extra_body = resolved_params.pop("_extra_body", None)
|
||||||
|
default_headers = resolved_params.pop("_default_headers", None)
|
||||||
|
if extra_body:
|
||||||
|
resolved_params["extra_body"] = extra_body
|
||||||
|
if default_headers:
|
||||||
|
resolved_params["default_headers"] = default_headers
|
||||||
|
resolved_params.update(kwargs)
|
||||||
|
kwargs = resolved_params
|
||||||
|
|
||||||
|
resolved_api_key = str(getattr(runtime_resolved, "api_key", "") or "")
|
||||||
|
resolved_base_url = str(getattr(runtime_resolved, "base_url", "") or "")
|
||||||
|
if resolved_api_key:
|
||||||
|
kwargs.setdefault("api_key", resolved_api_key)
|
||||||
|
if resolved_base_url:
|
||||||
|
kwargs.setdefault("base_url", resolved_base_url.rstrip("/"))
|
||||||
|
|
||||||
|
runtime_provider_name = str(
|
||||||
|
getattr(runtime_resolved, "provider_name", "") or ""
|
||||||
|
)
|
||||||
|
runtime_supports_reasoning = bool(
|
||||||
|
getattr(runtime_resolved, "supports_reasoning", False)
|
||||||
|
)
|
||||||
|
if not runtime_supports_reasoning:
|
||||||
|
kwargs.setdefault("_disable_reasoning", True)
|
||||||
|
kwargs.setdefault("_disable_thinking", True)
|
||||||
|
model = str(runtime_resolved.model_id)
|
||||||
|
provider = str(runtime_resolved.protocol)
|
||||||
|
else:
|
||||||
model = model or DEFAULT_MODEL
|
model = model or DEFAULT_MODEL
|
||||||
|
|
||||||
# Look up short name in registry (provider-aware)
|
# Look up short name in registry (provider-aware)
|
||||||
@@ -420,22 +548,35 @@ def get_chat_model(
|
|||||||
_is_third_party = (
|
_is_third_party = (
|
||||||
provider in _OPENAI_ROUTED_PROVIDERS or provider in _ANTHROPIC_ROUTED_PROVIDERS
|
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
|
_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":
|
if provider == "anthropic":
|
||||||
base_url = os.environ.get("ANTHROPIC_BASE_URL", "")
|
base_url = os.environ.get("ANTHROPIC_BASE_URL", "")
|
||||||
if base_url:
|
if base_url:
|
||||||
kwargs["base_url"] = base_url
|
kwargs.setdefault("base_url", base_url)
|
||||||
api_key = os.environ.get("ANTHROPIC_API_KEY", "")
|
api_key = os.environ.get("ANTHROPIC_API_KEY", "")
|
||||||
if 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)
|
# Native OpenAI base_url override (e.g. ccproxy Codex at localhost:8000/codex/v1)
|
||||||
elif provider == "openai":
|
elif provider == "openai":
|
||||||
base_url = os.environ.get("OPENAI_BASE_URL", "")
|
base_url = os.environ.get("OPENAI_BASE_URL", "")
|
||||||
if base_url:
|
if base_url:
|
||||||
kwargs["base_url"] = base_url
|
kwargs.setdefault("base_url", base_url)
|
||||||
_is_openai_proxy = _is_ccproxy_codex()
|
_is_openai_proxy = _is_ccproxy_codex(
|
||||||
|
kwargs.get("base_url"), kwargs.get("api_key")
|
||||||
|
)
|
||||||
if _is_openai_proxy:
|
if _is_openai_proxy:
|
||||||
# Use Responses API for ccproxy: bypasses the format chain
|
# Use Responses API for ccproxy: bypasses the format chain
|
||||||
# converter (Chat→Responses→Chat) which returns 502 on
|
# converter (Chat→Responses→Chat) which returns 502 on
|
||||||
@@ -448,9 +589,23 @@ def get_chat_model(
|
|||||||
# for Chat Completions tool_call duplication — not an issue
|
# for Chat Completions tool_call duplication — not an issue
|
||||||
# with the Responses API SSE format.)
|
# with the Responses API SSE format.)
|
||||||
kwargs.pop("streaming", None) # remove if set elsewhere
|
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", "")
|
api_key = os.environ.get("OPENAI_API_KEY", "")
|
||||||
if 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
|
# OpenAI-routed providers → route through OpenAI provider with base_url
|
||||||
elif provider in _OPENAI_ROUTED_PROVIDERS:
|
elif provider in _OPENAI_ROUTED_PROVIDERS:
|
||||||
@@ -468,10 +623,10 @@ def get_chat_model(
|
|||||||
else:
|
else:
|
||||||
base_url = base_url_default
|
base_url = base_url_default
|
||||||
if base_url:
|
if base_url:
|
||||||
kwargs["base_url"] = base_url
|
kwargs.setdefault("base_url", base_url)
|
||||||
api_key = os.environ.get(api_key_env, "")
|
api_key = os.environ.get(api_key_env, "")
|
||||||
if api_key:
|
if api_key:
|
||||||
kwargs["api_key"] = api_key
|
kwargs.setdefault("api_key", api_key)
|
||||||
# SiliconFlow: disable thinking — LangChain drops reasoning_content
|
# SiliconFlow: disable thinking — LangChain drops reasoning_content
|
||||||
# from history, causing error 20015 on multi-turn requests.
|
# from history, causing error 20015 on multi-turn requests.
|
||||||
if provider == "siliconflow":
|
if provider == "siliconflow":
|
||||||
@@ -488,15 +643,61 @@ def get_chat_model(
|
|||||||
_is_third_party = True
|
_is_third_party = True
|
||||||
api_key = os.environ.get("OPENROUTER_API_KEY", "")
|
api_key = os.environ.get("OPENROUTER_API_KEY", "")
|
||||||
if 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
|
# Reasoning via `effort` + `summary: "auto"` so a readable reasoning
|
||||||
# summary is returned for display. OpenAI-Responses also emits encrypted
|
# summary is returned for display. OpenAI-Responses also emits encrypted
|
||||||
# reasoning items (`rs_*` id) that can't be replayed on multi-turn
|
# reasoning items (`rs_*` id) that can't be replayed on multi-turn
|
||||||
# passback (OpenRouter's `/responses` beta is stateless, store=false —
|
# passback (OpenRouter's `/responses` beta is stateless, store=false —
|
||||||
# "Item with id 'rs_...' not found"); the patch strips them on passback,
|
# "Item with id 'rs_...' not found"); the patch strips them on passback,
|
||||||
# so enabling `summary` is safe. See langchain-ai/langchain#37777.
|
# 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"})
|
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()
|
_patch_openrouter_strip_responses_reasoning()
|
||||||
|
|
||||||
# Anthropic-routed providers → route through Anthropic provider with base_url
|
# Anthropic-routed providers → route through Anthropic provider with base_url
|
||||||
@@ -517,10 +718,10 @@ def get_chat_model(
|
|||||||
else:
|
else:
|
||||||
base_url = base_url_default
|
base_url = base_url_default
|
||||||
if base_url:
|
if base_url:
|
||||||
kwargs["base_url"] = base_url
|
kwargs.setdefault("base_url", base_url)
|
||||||
api_key = os.environ.get(api_key_env, "")
|
api_key = os.environ.get(api_key_env, "")
|
||||||
if api_key:
|
if api_key:
|
||||||
kwargs["api_key"] = api_key
|
kwargs.setdefault("api_key", api_key)
|
||||||
# Kimi Coding Plan requires claude-code User-Agent header
|
# Kimi Coding Plan requires claude-code User-Agent header
|
||||||
if provider == "kimi-coding":
|
if provider == "kimi-coding":
|
||||||
kwargs.setdefault("default_headers", {})["User-Agent"] = "claude-code/0.1.0"
|
kwargs.setdefault("default_headers", {})["User-Agent"] = "claude-code/0.1.0"
|
||||||
@@ -529,8 +730,9 @@ def get_chat_model(
|
|||||||
elif provider == "ollama":
|
elif provider == "ollama":
|
||||||
base_url = os.environ.get("OLLAMA_BASE_URL", "")
|
base_url = os.environ.get("OLLAMA_BASE_URL", "")
|
||||||
if 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_auto_config(provider, model_id, _is_third_party, kwargs, _original_provider)
|
||||||
_apply_openrouter_anthropic_prompt_cache(provider, model_id, kwargs)
|
_apply_openrouter_anthropic_prompt_cache(provider, model_id, kwargs)
|
||||||
|
|
||||||
@@ -547,7 +749,14 @@ def get_chat_model(
|
|||||||
elif _responses_api_setting == "true":
|
elif _responses_api_setting == "true":
|
||||||
kwargs["use_responses_api"] = True
|
kwargs["use_responses_api"] = True
|
||||||
|
|
||||||
|
anthropic_auth_token = None
|
||||||
|
if provider == "anthropic" and kwargs.get("api_key"):
|
||||||
|
anthropic_auth_token = os.environ.pop("ANTHROPIC_AUTH_TOKEN", None)
|
||||||
|
try:
|
||||||
chat_model = init_chat_model(model=model_id, model_provider=provider, **kwargs)
|
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
|
# Flatten list content to strings for strict OpenAI-compatible providers
|
||||||
# (DeepSeek, SiliconFlow, OpenRouter, custom-openai, etc.) and
|
# (DeepSeek, SiliconFlow, OpenRouter, custom-openai, etc.) and
|
||||||
|
|||||||
+365
-1
@@ -25,6 +25,7 @@ Utilities:
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
import os
|
import os
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
@@ -178,14 +179,19 @@ _patch_ccproxy_codex_compat()
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Utility: detect ccproxy's Codex adapter (as opposed to generic localhost).
|
# 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.
|
"""Return True if the OpenAI endpoint is ccproxy's Codex adapter.
|
||||||
|
|
||||||
Checks for the ccproxy-specific markers set by ``setup_codex_env()``
|
Checks for the ccproxy-specific markers set by ``setup_codex_env()``
|
||||||
in ``ccproxy_manager.py``: the sentinel API key and the ``/codex/v1``
|
in ``ccproxy_manager.py``: the sentinel API key and the ``/codex/v1``
|
||||||
path. Plain localhost endpoints (vLLM, Ollama, etc.) are not affected.
|
path. Plain localhost endpoints (vLLM, Ollama, etc.) are not affected.
|
||||||
"""
|
"""
|
||||||
|
if base_url is None:
|
||||||
base_url = os.environ.get("OPENAI_BASE_URL", "")
|
base_url = os.environ.get("OPENAI_BASE_URL", "")
|
||||||
|
if api_key is None:
|
||||||
api_key = os.environ.get("OPENAI_API_KEY", "")
|
api_key = os.environ.get("OPENAI_API_KEY", "")
|
||||||
return (
|
return (
|
||||||
("127.0.0.1" in base_url or "localhost" in base_url)
|
("127.0.0.1" in base_url or "localhost" in base_url)
|
||||||
@@ -267,6 +273,299 @@ def _flatten_message_content(content: Any) -> str | list[Any] | Any:
|
|||||||
return "\n\n".join(parts) if parts else ""
|
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]:
|
def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> list[Any]:
|
||||||
"""Flatten list content for OpenAI-compatible APIs, preserving media.
|
"""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
|
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] = []
|
out: list[Any] = []
|
||||||
pending_media: list[Any] = [] # media hoisted out of a run of tool messages
|
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
|
msg.content = flat
|
||||||
out.append(msg)
|
out.append(msg)
|
||||||
_flush() # conversation may end with tool messages
|
_flush() # conversation may end with tool messages
|
||||||
|
if sanitize_tool_history:
|
||||||
|
_validate_openai_tool_history(out)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
@@ -729,6 +1033,66 @@ def _patch_openai_capture_reasoning_content() -> None:
|
|||||||
_patch_openai_capture_reasoning_content()
|
_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
|
# Patch (lazy, OpenRouter only): strip OpenAI-Responses encrypted reasoning
|
||||||
# items from outgoing assistant messages.
|
# items from outgoing assistant messages.
|
||||||
|
|||||||
@@ -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,
|
||||||
|
)
|
||||||
@@ -10,6 +10,7 @@ from .client import (
|
|||||||
build_mcp_add_kwargs,
|
build_mcp_add_kwargs,
|
||||||
build_mcp_edit_fields,
|
build_mcp_edit_fields,
|
||||||
edit_mcp_server,
|
edit_mcp_server,
|
||||||
|
get_mcp_server_errors,
|
||||||
load_mcp_config,
|
load_mcp_config,
|
||||||
load_mcp_tools,
|
load_mcp_tools,
|
||||||
parse_mcp_add_args,
|
parse_mcp_add_args,
|
||||||
@@ -38,6 +39,7 @@ __all__ = [
|
|||||||
"find_server_by_name",
|
"find_server_by_name",
|
||||||
"get_all_tags",
|
"get_all_tags",
|
||||||
"get_installed_names",
|
"get_installed_names",
|
||||||
|
"get_mcp_server_errors",
|
||||||
"install_mcp_server",
|
"install_mcp_server",
|
||||||
"install_mcp_servers",
|
"install_mcp_servers",
|
||||||
"load_mcp_config",
|
"load_mcp_config",
|
||||||
|
|||||||
@@ -114,6 +114,10 @@ _URL_TRANSPORTS = {"http", "streamable_http", "sse", "websocket"}
|
|||||||
# still parallelizing the common 3–7 server case to completion.
|
# still parallelizing the common 3–7 server case to completion.
|
||||||
_MAX_CONCURRENT_CONNECTIONS = 8
|
_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
|
# Env vars forwarded to stdio MCP subprocesses on top of the MCP SDK's
|
||||||
# minimal default set (HOME/PATH/USER/…). Without this, servers behind
|
# minimal default set (HOME/PATH/USER/…). Without this, servers behind
|
||||||
# a proxy or with a custom CA bundle silently fail with long timeouts.
|
# a proxy or with a custom CA bundle silently fail with long timeouts.
|
||||||
@@ -764,6 +768,9 @@ async def _load_tools(
|
|||||||
if not connections:
|
if not connections:
|
||||||
return {}
|
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]
|
client = MultiServerMCPClient(connections) # type: ignore[invalid-argument-type]
|
||||||
|
|
||||||
def _report(event: str, name: str, detail: str = "") -> None:
|
def _report(event: str, name: str, detail: str = "") -> None:
|
||||||
@@ -787,10 +794,13 @@ async def _load_tools(
|
|||||||
_report("start", name)
|
_report("start", name)
|
||||||
try:
|
try:
|
||||||
tools = await client.get_tools(server_name=name)
|
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))
|
logger.info("MCP server %r: loaded %d tool(s)", name, len(tools))
|
||||||
_report("success", name, str(len(tools)))
|
_report("success", name, str(len(tools)))
|
||||||
return name, tools
|
return name, tools
|
||||||
except Exception as exc:
|
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
|
# When the caller wired up ``on_progress`` they own the
|
||||||
# user-facing display; downgrade the logger so we don't
|
# user-facing display; downgrade the logger so we don't
|
||||||
# double-print.
|
# double-print.
|
||||||
@@ -798,7 +808,7 @@ async def _load_tools(
|
|||||||
logger.warning("MCP server %r: failed to load tools: %s", name, exc)
|
logger.warning("MCP server %r: failed to load tools: %s", name, exc)
|
||||||
else:
|
else:
|
||||||
logger.debug("MCP server %r: failed to load tools: %s", name, exc)
|
logger.debug("MCP server %r: failed to load tools: %s", name, exc)
|
||||||
_report("error", name, str(exc))
|
_report("error", name, detail)
|
||||||
return name, []
|
return name, []
|
||||||
|
|
||||||
# ``return_exceptions=False`` is fine because ``_fetch`` already
|
# ``return_exceptions=False`` is fine because ``_fetch`` already
|
||||||
@@ -807,6 +817,11 @@ async def _load_tools(
|
|||||||
return dict(results)
|
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(
|
async def aload_mcp_tools(
|
||||||
config: dict[str, Any] | None = None,
|
config: dict[str, Any] | None = None,
|
||||||
*,
|
*,
|
||||||
|
|||||||
@@ -428,6 +428,7 @@ def _memory_worker_middleware(
|
|||||||
enable_observation_memory: bool = True,
|
enable_observation_memory: bool = True,
|
||||||
):
|
):
|
||||||
"""Build middleware for memory workers, excluding task execution tools."""
|
"""Build middleware for memory workers, excluding task execution tools."""
|
||||||
|
from ...middleware.error_normalization import ErrorNormalizationMiddleware
|
||||||
from ...middleware.memory import create_memory_middleware
|
from ...middleware.memory import create_memory_middleware
|
||||||
|
|
||||||
memory_controls = MemoryControls(
|
memory_controls = MemoryControls(
|
||||||
@@ -439,7 +440,11 @@ def _memory_worker_middleware(
|
|||||||
enable_observation_tool = memory_controls.observation_tool_enabled(
|
enable_observation_tool = memory_controls.observation_tool_enabled(
|
||||||
_memory_worker_observation_target(source_type)
|
_memory_worker_observation_target(source_type)
|
||||||
)
|
)
|
||||||
return memory_agent_middleware(
|
return [
|
||||||
|
# Outermost — normalize provider-SDK exceptions from the
|
||||||
|
# auxiliary model call before any inner middleware sees them.
|
||||||
|
ErrorNormalizationMiddleware(),
|
||||||
|
*memory_agent_middleware(
|
||||||
create_memory_middleware(
|
create_memory_middleware(
|
||||||
str(memory_dir),
|
str(memory_dir),
|
||||||
workspace_dir=workspace_dir,
|
workspace_dir=workspace_dir,
|
||||||
@@ -450,7 +455,8 @@ def _memory_worker_middleware(
|
|||||||
enable_observation_tool=enable_observation_tool,
|
enable_observation_tool=enable_observation_tool,
|
||||||
),
|
),
|
||||||
excluded_tools=_MEMORY_WORKER_EXCLUDED_TOOLS,
|
excluded_tools=_MEMORY_WORKER_EXCLUDED_TOOLS,
|
||||||
)
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def _build_memory_worker_agent(
|
def _build_memory_worker_agent(
|
||||||
|
|||||||
@@ -71,6 +71,8 @@ def build_observation_linker_graph(
|
|||||||
workspace_dir: str | Path | None = None,
|
workspace_dir: str | Path | None = None,
|
||||||
) -> CompiledStateGraph:
|
) -> CompiledStateGraph:
|
||||||
"""Build the registered LangGraph observation linker."""
|
"""Build the registered LangGraph observation linker."""
|
||||||
|
from ...middleware.error_normalization import ErrorNormalizationMiddleware
|
||||||
|
|
||||||
agent_paths = resolve_memory_agent_paths(
|
agent_paths = resolve_memory_agent_paths(
|
||||||
memory_dir=memory_dir,
|
memory_dir=memory_dir,
|
||||||
workspace_dir=workspace_dir,
|
workspace_dir=workspace_dir,
|
||||||
@@ -85,5 +87,7 @@ def build_observation_linker_graph(
|
|||||||
tools=tools,
|
tools=tools,
|
||||||
memory_dir=agent_paths.memory_dir,
|
memory_dir=agent_paths.memory_dir,
|
||||||
workspace_dir=agent_paths.workspace_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()],
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ from .context_editing import (
|
|||||||
create_context_editing_middleware,
|
create_context_editing_middleware,
|
||||||
)
|
)
|
||||||
from .context_overflow import ContextOverflowMapperMiddleware
|
from .context_overflow import ContextOverflowMapperMiddleware
|
||||||
|
from .error_normalization import ErrorNormalizationMiddleware
|
||||||
from .memory import (
|
from .memory import (
|
||||||
EvoMemoryMiddleware,
|
EvoMemoryMiddleware,
|
||||||
create_memory_middleware,
|
create_memory_middleware,
|
||||||
@@ -28,29 +29,42 @@ from .memory_lifecycle import (
|
|||||||
default_memory_scheduler,
|
default_memory_scheduler,
|
||||||
)
|
)
|
||||||
from .model_fallback import ModelFallbackMiddleware, load_fallback_chain
|
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 .runtime_context import RuntimeContextMiddleware, create_runtime_context_middleware
|
||||||
from .scheduler import (
|
from .scheduler import (
|
||||||
SchedulerMiddleware,
|
SchedulerMiddleware,
|
||||||
create_scheduler_middleware,
|
create_scheduler_middleware,
|
||||||
)
|
)
|
||||||
from .tool_error_handler import ToolErrorHandlerMiddleware
|
from .tool_error_handler import ToolErrorHandlerMiddleware
|
||||||
|
from .tool_protocol_guard import ToolProtocolGuardMiddleware
|
||||||
from .tool_selector import create_tool_selector_middleware
|
from .tool_selector import create_tool_selector_middleware
|
||||||
from .utils import disable_thinking
|
from .utils import disable_thinking
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
"DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS",
|
||||||
|
"DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD",
|
||||||
"AskUserMiddleware",
|
"AskUserMiddleware",
|
||||||
"AskUserRequest",
|
"AskUserRequest",
|
||||||
"AskUserWidgetResult",
|
"AskUserWidgetResult",
|
||||||
"Choice",
|
"Choice",
|
||||||
"ConfigurableModelMiddleware",
|
"ConfigurableModelMiddleware",
|
||||||
"ContextOverflowMapperMiddleware",
|
"ContextOverflowMapperMiddleware",
|
||||||
|
"ErrorNormalizationMiddleware",
|
||||||
"EvoMemoryLifecycleMiddleware",
|
"EvoMemoryLifecycleMiddleware",
|
||||||
"EvoMemoryMiddleware",
|
"EvoMemoryMiddleware",
|
||||||
"ModelFallbackMiddleware",
|
"ModelFallbackMiddleware",
|
||||||
"Question",
|
"Question",
|
||||||
|
"RepetitiveToolCallGuardMiddleware",
|
||||||
"RuntimeContextMiddleware",
|
"RuntimeContextMiddleware",
|
||||||
"SchedulerMiddleware",
|
"SchedulerMiddleware",
|
||||||
"ToolErrorHandlerMiddleware",
|
"ToolErrorHandlerMiddleware",
|
||||||
|
"ToolProtocolGuardMiddleware",
|
||||||
|
"collapse_repetitive_tool_rounds",
|
||||||
"compute_context_editing_trigger",
|
"compute_context_editing_trigger",
|
||||||
"create_code_interpreter_middleware",
|
"create_code_interpreter_middleware",
|
||||||
"create_context_editing_middleware",
|
"create_context_editing_middleware",
|
||||||
|
|||||||
@@ -45,7 +45,21 @@ _MEMORY_FIRST_INTERPRETER_PROMPT = (
|
|||||||
|
|
||||||
|
|
||||||
class EvoCodeInterpreterMiddleware(CodeInterpreterMiddleware):
|
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:
|
def _prepare_for_call(self, request: ModelRequest) -> str:
|
||||||
return super()._prepare_for_call(request) + _MEMORY_FIRST_INTERPRETER_PROMPT
|
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
|
||||||
@@ -48,6 +48,8 @@ _MALFORMED_REQUEST_PATTERNS: list[str] = [
|
|||||||
"invalid_request_error",
|
"invalid_request_error",
|
||||||
"invalid request",
|
"invalid request",
|
||||||
"malformed",
|
"malformed",
|
||||||
|
"repetitive tool calls",
|
||||||
|
"identical name and arguments",
|
||||||
]
|
]
|
||||||
"""Substrings that identify a malformed request (client-side bug)."""
|
"""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
|
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):
|
if isinstance(exc, ContextOverflowError):
|
||||||
return "context length exceeded"
|
return "context length exceeded"
|
||||||
|
|
||||||
@@ -263,7 +268,15 @@ async def _try_fallbacks(
|
|||||||
"Primary model failed: %s: %s", type(primary_exc).__name__, primary_exc
|
"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_exc = primary_exc
|
||||||
|
last_failing_request = request
|
||||||
|
|
||||||
for model_name, provider in get_fallback_chain():
|
for model_name, provider in get_fallback_chain():
|
||||||
_emit(
|
_emit(
|
||||||
f" -> Falling back to {model_name} ({provider}) "
|
f" -> Falling back to {model_name} ({provider}) "
|
||||||
@@ -288,8 +301,9 @@ async def _try_fallbacks(
|
|||||||
f"-- aborting fallback chain",
|
f"-- aborting fallback chain",
|
||||||
style="red",
|
style="red",
|
||||||
)
|
)
|
||||||
raise
|
_raise_normalized(fb_request, fb_exc)
|
||||||
last_exc = fb_exc
|
last_exc = fb_exc
|
||||||
|
last_failing_request = fb_request
|
||||||
_emit(
|
_emit(
|
||||||
f" x {model_name} also failed: {type(fb_exc).__name__}: {fb_exc}",
|
f" x {model_name} also failed: {type(fb_exc).__name__}: {fb_exc}",
|
||||||
style="red",
|
style="red",
|
||||||
@@ -303,7 +317,24 @@ async def _try_fallbacks(
|
|||||||
)
|
)
|
||||||
|
|
||||||
_emit(" All fallbacks exhausted -- re-raising last error", style="red")
|
_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(
|
def _guard_and_fallback(
|
||||||
@@ -330,7 +361,7 @@ def _guard_and_fallback(
|
|||||||
f"Model error ({reason}) -- not eligible for fallback, re-raising",
|
f"Model error ({reason}) -- not eligible for fallback, re-raising",
|
||||||
style="red",
|
style="red",
|
||||||
)
|
)
|
||||||
raise primary_exc
|
_raise_normalized(request, primary_exc)
|
||||||
return _try_fallbacks(request, invoke, 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
|
||||||
@@ -48,6 +48,7 @@ DEFAULT_ALWAYS_INCLUDE_TOOLS: frozenset[str] = frozenset(
|
|||||||
"read_memory",
|
"read_memory",
|
||||||
"record_observation",
|
"record_observation",
|
||||||
"search_observations",
|
"search_observations",
|
||||||
|
"write_todos",
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -132,10 +133,21 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
|
|||||||
return self._build_selector(request).wrap_model_call(
|
return self._build_selector(request).wrap_model_call(
|
||||||
request, _handler_after_selection
|
request, _handler_after_selection
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception as exc:
|
||||||
if _handler_called:
|
if _handler_called:
|
||||||
raise # Error from downstream model — don't retry
|
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)
|
logger.debug("Tool selector failed, using all tools", exc_info=True)
|
||||||
if self._track_stream_selection:
|
if self._track_stream_selection:
|
||||||
_selector_active = False
|
_selector_active = False
|
||||||
@@ -171,9 +183,16 @@ class _ConditionalToolSelectorMiddleware(AgentMiddleware):
|
|||||||
return await self._build_selector(request).awrap_model_call(
|
return await self._build_selector(request).awrap_model_call(
|
||||||
request, _handler_after_selection
|
request, _handler_after_selection
|
||||||
)
|
)
|
||||||
except Exception:
|
except Exception as exc:
|
||||||
if _handler_called:
|
if _handler_called:
|
||||||
raise
|
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)
|
logger.debug("Tool selector failed, using all tools", exc_info=True)
|
||||||
if self._track_stream_selection:
|
if self._track_stream_selection:
|
||||||
_selector_active = False
|
_selector_active = False
|
||||||
@@ -258,6 +277,15 @@ def create_tool_selector_middleware(
|
|||||||
|
|
||||||
model = _ensure_chat_model()
|
model = _ensure_chat_model()
|
||||||
safe_model = disable_thinking(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 = (
|
system_prompt = (
|
||||||
"You are selecting tools for a scientific research agent. "
|
"You are selecting tools for a scientific research agent. "
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ from __future__ import annotations
|
|||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import shutil
|
import shutil
|
||||||
|
from collections.abc import Iterator
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from pathlib import Path
|
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."""
|
"""Resolve a virtual workspace path (e.g. /image.png) to a real filesystem path."""
|
||||||
vpath = virtual_path if virtual_path.startswith("/") else "/" + virtual_path
|
vpath = virtual_path if virtual_path.startswith("/") else "/" + virtual_path
|
||||||
return (_active_workspace / vpath.lstrip("/")).resolve()
|
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"
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -7,6 +7,15 @@ All events contain a type and associated data dict.
|
|||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any
|
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
|
@dataclass
|
||||||
class StreamEvent:
|
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
|
@staticmethod
|
||||||
def interrupt(
|
def interrupt(
|
||||||
interrupt_id: str,
|
interrupt_id: str,
|
||||||
@@ -213,6 +230,19 @@ class StreamEventEmitter:
|
|||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@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."""
|
"""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)
|
||||||
|
|||||||
+269
-30
@@ -8,13 +8,15 @@ import base64
|
|||||||
import inspect
|
import inspect
|
||||||
import mimetypes
|
import mimetypes
|
||||||
import os
|
import os
|
||||||
|
import warnings
|
||||||
from collections.abc import AsyncGenerator, AsyncIterator, Mapping
|
from collections.abc import AsyncGenerator, AsyncIterator, Mapping
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, TypeAlias
|
from typing import Any, TypeAlias
|
||||||
|
|
||||||
|
from langchain_core._api import LangChainBetaWarning
|
||||||
from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, ToolMessage
|
from langchain_core.messages import AIMessage, AIMessageChunk, BaseMessage, ToolMessage
|
||||||
from langgraph.graph import END
|
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 ..memory.worker_activity import clear_completed_memory_activity_counts
|
||||||
from .emitter import StreamEventEmitter
|
from .emitter import StreamEventEmitter
|
||||||
@@ -43,6 +45,12 @@ GraphRunInput: TypeAlias = str | Command
|
|||||||
LangGraphStreamInput: TypeAlias = dict[str, list[dict[str, object]]] | Command
|
LangGraphStreamInput: TypeAlias = dict[str, list[dict[str, object]]] | Command
|
||||||
_ValueMessageKey: TypeAlias = tuple[str, ...]
|
_ValueMessageKey: TypeAlias = tuple[str, ...]
|
||||||
|
|
||||||
|
warnings.filterwarnings(
|
||||||
|
"ignore",
|
||||||
|
message=r"The v3 streaming protocol on Pregel is experimental\.",
|
||||||
|
category=LangChainBetaWarning,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True, slots=True)
|
@dataclass(frozen=True, slots=True)
|
||||||
class _AssistantValueMessage:
|
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
|
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.
|
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
|
Recovery first removes malformed/incomplete tool protocol from the messages
|
||||||
and writes a checkpoint whose ``next`` is the empty tuple, without touching
|
channel, then ``aupdate_state(config, None, as_node=END)`` clears pending
|
||||||
any channel values (message history is preserved).
|
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
|
Critically, this only runs when the stuck state is *not* a legitimate
|
||||||
human-in-the-loop interrupt. The agent pauses via ``interrupt()`` /
|
human-in-the-loop interrupt. The agent pauses via ``interrupt()`` /
|
||||||
@@ -106,10 +115,9 @@ async def _clear_interrupted_graph_state(
|
|||||||
_log = logging.getLogger(__name__)
|
_log = logging.getLogger(__name__)
|
||||||
try:
|
try:
|
||||||
snapshot = await agent.aget_state(config)
|
snapshot = await agent.aget_state(config)
|
||||||
# Only act when the graph is genuinely stuck (non-empty next tuple)...
|
if not snapshot:
|
||||||
if not snapshot or not getattr(snapshot, "next", None):
|
|
||||||
return
|
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):
|
if _snapshot_has_pending_interrupt(snapshot):
|
||||||
_log.debug(
|
_log.debug(
|
||||||
"Leaving interrupted graph state intact for thread %s: "
|
"Leaving interrupted graph state intact for thread %s: "
|
||||||
@@ -119,6 +127,13 @@ async def _clear_interrupted_graph_state(
|
|||||||
)
|
)
|
||||||
return
|
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
|
stuck_at = snapshot.next
|
||||||
await agent.aupdate_state(config, None, as_node=END)
|
await agent.aupdate_state(config, None, as_node=END)
|
||||||
_log.debug(
|
_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)
|
@dataclass(frozen=True)
|
||||||
class _SubagentInfo:
|
class _SubagentInfo:
|
||||||
path: tuple[str, ...]
|
path: tuple[str, ...]
|
||||||
@@ -209,7 +267,12 @@ class _V3EventProcessor:
|
|||||||
tuple[tuple[str, ...], str], tuple[str, dict[str, Any]]
|
tuple[tuple[str, ...], str], tuple[str, dict[str, Any]]
|
||||||
] = {}
|
] = {}
|
||||||
self._emitted_tool_calls: set[tuple[tuple[str, ...], str]] = set()
|
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._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)
|
self._selector = _ToolSelectionSuppressor(emitter)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -237,13 +300,26 @@ class _V3EventProcessor:
|
|||||||
if method == "tools":
|
if method == "tools":
|
||||||
return self._process_tool_event(namespace, _event_data(event), subagent)
|
return self._process_tool_event(namespace, _event_data(event), subagent)
|
||||||
if method == "updates":
|
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":
|
if method == "values":
|
||||||
events: list[dict[str, Any]] = []
|
events: list[dict[str, Any]] = []
|
||||||
params = event.get("params") or {}
|
params = event.get("params") or {}
|
||||||
interrupts = params.get("interrupts") or ()
|
interrupts = params.get("interrupts") or ()
|
||||||
if interrupts:
|
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:
|
if self._process_value_message_snapshots and not namespace:
|
||||||
events.extend(self._process_value_messages(_event_data(event)))
|
events.extend(self._process_value_messages(_event_data(event)))
|
||||||
return events
|
return events
|
||||||
@@ -382,6 +458,7 @@ class _V3EventProcessor:
|
|||||||
inp, out = _usage_counts(usage) if usage is not None else (0, 0)
|
inp, out = _usage_counts(usage) if usage is not None else (0, 0)
|
||||||
if inp or out:
|
if inp or out:
|
||||||
events.append(self.emitter.usage_stats(inp, out).data)
|
events.append(self.emitter.usage_stats(inp, out).data)
|
||||||
|
events.extend(self._flush_invalid_tool_calls())
|
||||||
return events
|
return events
|
||||||
return []
|
return []
|
||||||
|
|
||||||
@@ -407,14 +484,12 @@ class _V3EventProcessor:
|
|||||||
if tool_call is None:
|
if tool_call is None:
|
||||||
return events
|
return events
|
||||||
tool_name, args, tool_call_id = tool_call
|
tool_name, args, tool_call_id = tool_call
|
||||||
events.extend(
|
self._pending_invalid_tool_calls.pop(tool_call_id, None)
|
||||||
self._emit_tool_call_once(
|
self._pending_tool_calls[
|
||||||
namespace=namespace,
|
(self._tool_scope(namespace, subagent), tool_call_id)
|
||||||
subagent=subagent,
|
] = (
|
||||||
name=tool_name,
|
tool_name,
|
||||||
args=args,
|
args,
|
||||||
tool_call_id=tool_call_id,
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
return events
|
return events
|
||||||
|
|
||||||
@@ -457,6 +532,20 @@ class _V3EventProcessor:
|
|||||||
]
|
]
|
||||||
return [self.emitter.tool_call(name, args, tool_call_id).data]
|
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(
|
def _process_whole_message(
|
||||||
self,
|
self,
|
||||||
msg: AIMessage | AIMessageChunk,
|
msg: AIMessage | AIMessageChunk,
|
||||||
@@ -464,6 +553,16 @@ class _V3EventProcessor:
|
|||||||
namespace: tuple[str, ...],
|
namespace: tuple[str, ...],
|
||||||
) -> list[dict[str, Any]]:
|
) -> list[dict[str, Any]]:
|
||||||
events: 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
|
additional = msg.additional_kwargs
|
||||||
reasoning = additional.get("reasoning_content")
|
reasoning = additional.get("reasoning_content")
|
||||||
emitted_reasoning = False
|
emitted_reasoning = False
|
||||||
@@ -486,14 +585,12 @@ class _V3EventProcessor:
|
|||||||
if tool_call is None:
|
if tool_call is None:
|
||||||
continue
|
continue
|
||||||
tool_name, args, tool_call_id = tool_call
|
tool_name, args, tool_call_id = tool_call
|
||||||
events.extend(
|
self._pending_invalid_tool_calls.pop(tool_call_id, None)
|
||||||
self._emit_tool_call_once(
|
self._pending_tool_calls[
|
||||||
namespace=namespace,
|
(self._tool_scope(namespace, subagent), tool_call_id)
|
||||||
subagent=subagent,
|
] = (
|
||||||
name=tool_name,
|
tool_name,
|
||||||
args=args,
|
args,
|
||||||
tool_call_id=tool_call_id,
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if subagent is None:
|
if subagent is None:
|
||||||
@@ -526,6 +623,9 @@ class _V3EventProcessor:
|
|||||||
name,
|
name,
|
||||||
args,
|
args,
|
||||||
)
|
)
|
||||||
|
self._pending_tool_calls.pop(
|
||||||
|
(self._tool_scope(namespace, subagent), tool_call_id), None
|
||||||
|
)
|
||||||
events.extend(
|
events.extend(
|
||||||
self._emit_tool_call_once(
|
self._emit_tool_call_once(
|
||||||
namespace=namespace,
|
namespace=namespace,
|
||||||
@@ -572,6 +672,7 @@ class _V3EventProcessor:
|
|||||||
content += "\n... (truncated)"
|
content += "\n... (truncated)"
|
||||||
success = is_success(content)
|
success = is_success(content)
|
||||||
|
|
||||||
|
lifecycle_key = (self._tool_scope(namespace, subagent), tool_call_id)
|
||||||
if subagent is not None:
|
if subagent is not None:
|
||||||
events.append(
|
events.append(
|
||||||
self.emitter.subagent_tool_result(
|
self.emitter.subagent_tool_result(
|
||||||
@@ -583,22 +684,78 @@ class _V3EventProcessor:
|
|||||||
instance_id=subagent.instance_id,
|
instance_id=subagent.instance_id,
|
||||||
).data
|
).data
|
||||||
)
|
)
|
||||||
return events
|
else:
|
||||||
events.append(
|
events.append(
|
||||||
self.emitter.tool_result(
|
self.emitter.tool_result(
|
||||||
name, content, success, tool_call_id=tool_call_id
|
name, content, success, tool_call_id=tool_call_id
|
||||||
).data
|
).data
|
||||||
)
|
)
|
||||||
|
self._emitted_tool_calls.discard(lifecycle_key)
|
||||||
|
self._pending_tool_calls.pop(lifecycle_key, None)
|
||||||
return events
|
return events
|
||||||
|
|
||||||
return []
|
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]] = []
|
events: list[dict[str, Any]] = []
|
||||||
data_map = _as_raw_map(data)
|
data_map = _as_raw_map(data)
|
||||||
if data_map is not None and "__interrupt__" in data_map:
|
if data_map is not None and "__interrupt__" in data_map:
|
||||||
events.extend(self._process_interrupts(data_map["__interrupt__"]))
|
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)
|
summarization_event = _find_summarization_event_payload(data)
|
||||||
if summarization_event and not self._summarization_in_progress:
|
if summarization_event and not self._summarization_in_progress:
|
||||||
signature = _summarization_event_signature(summarization_event)
|
signature = _summarization_event_signature(summarization_event)
|
||||||
@@ -614,6 +771,10 @@ class _V3EventProcessor:
|
|||||||
events.extend(self._emit_summarization_text(summary_text))
|
events.extend(self._emit_summarization_text(summary_text))
|
||||||
return events
|
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]]:
|
def _process_interrupts(self, interrupts: object) -> list[dict[str, Any]]:
|
||||||
events: list[dict[str, Any]] = []
|
events: list[dict[str, Any]] = []
|
||||||
if not isinstance(interrupts, list | tuple):
|
if not isinstance(interrupts, list | tuple):
|
||||||
@@ -652,26 +813,74 @@ class _V3EventProcessor:
|
|||||||
raw_questions = interrupt_map.get("questions")
|
raw_questions = interrupt_map.get("questions")
|
||||||
questions = raw_questions if isinstance(raw_questions, list) else []
|
questions = raw_questions if isinstance(raw_questions, list) else []
|
||||||
tc_id = str(interrupt_map.get("tool_call_id", ""))
|
tc_id = str(interrupt_map.get("tool_call_id", ""))
|
||||||
return self._dedupe_interrupt_event(
|
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(
|
self.emitter.ask_user_interrupt(
|
||||||
interrupt_id,
|
interrupt_id,
|
||||||
questions,
|
questions,
|
||||||
tc_id,
|
tc_id,
|
||||||
).data
|
).data
|
||||||
)
|
)
|
||||||
|
)
|
||||||
|
return events
|
||||||
|
|
||||||
raw_action_reqs = interrupt_map.get("action_requests")
|
raw_action_reqs = interrupt_map.get("action_requests")
|
||||||
action_reqs = raw_action_reqs if isinstance(raw_action_reqs, list) else []
|
action_reqs = raw_action_reqs if isinstance(raw_action_reqs, list) else []
|
||||||
raw_review_cfgs = interrupt_map.get("review_configs")
|
raw_review_cfgs = interrupt_map.get("review_configs")
|
||||||
review_cfgs = raw_review_cfgs if isinstance(raw_review_cfgs, list) else None
|
review_cfgs = raw_review_cfgs if isinstance(raw_review_cfgs, list) else None
|
||||||
if action_reqs:
|
if action_reqs:
|
||||||
return self._dedupe_interrupt_event(
|
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(
|
self.emitter.interrupt(
|
||||||
interrupt_id,
|
interrupt_id,
|
||||||
action_reqs,
|
action_reqs,
|
||||||
review_cfgs,
|
review_cfgs,
|
||||||
).data
|
).data
|
||||||
)
|
)
|
||||||
|
)
|
||||||
|
return events
|
||||||
return []
|
return []
|
||||||
|
|
||||||
def _dedupe_interrupt_event(self, event: dict[str, Any]) -> list[dict[str, Any]]:
|
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,
|
thread_id: str,
|
||||||
metadata: dict[str, Any] | None = None,
|
metadata: dict[str, Any] | None = None,
|
||||||
media: list[str] | None = None,
|
media: list[str] | None = None,
|
||||||
|
callbacks: list[Any] | None = None,
|
||||||
|
error_mode: str = "emit",
|
||||||
) -> AsyncGenerator[dict[str, Any], None]:
|
) -> AsyncGenerator[dict[str, Any], None]:
|
||||||
"""Stream events from a DeepAgents/LangGraph v3 run.
|
"""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
|
metadata: Optional metadata dict merged into the LangGraph config
|
||||||
(e.g. agent_name, updated_at for checkpoint persistence).
|
(e.g. agent_name, updated_at for checkpoint persistence).
|
||||||
media: Optional list of local file paths for attachments.
|
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:
|
Yields:
|
||||||
Event dicts: thinking, text, tool_call, tool_result,
|
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}}
|
config: dict[str, Any] = {"configurable": {"thread_id": thread_id}}
|
||||||
if metadata:
|
if metadata:
|
||||||
config["metadata"] = metadata
|
config["metadata"] = metadata
|
||||||
|
if callbacks:
|
||||||
|
config["callbacks"] = callbacks
|
||||||
emitter = StreamEventEmitter()
|
emitter = StreamEventEmitter()
|
||||||
existing_summarization_event: Mapping[str, object] | None = None
|
existing_summarization_event: Mapping[str, object] | None = None
|
||||||
try:
|
try:
|
||||||
@@ -949,7 +1165,30 @@ async def stream_agent_events(
|
|||||||
yield item
|
yield item
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
_run_raised = True
|
_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
|
raise
|
||||||
finally:
|
finally:
|
||||||
if stream is not None:
|
if stream is not None:
|
||||||
|
|||||||
@@ -1,7 +1,11 @@
|
|||||||
"""Shared fixtures for EvoScientist tests."""
|
"""Shared fixtures for EvoScientist tests."""
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
_NONEXISTENT_DOTENV = str(Path(__file__).with_name(".pytest-dotenv-does-not-exist"))
|
||||||
|
|
||||||
|
|
||||||
@pytest.fixture(autouse=True)
|
@pytest.fixture(autouse=True)
|
||||||
def _reset_tool_selection_state():
|
def _reset_tool_selection_state():
|
||||||
@@ -164,3 +168,24 @@ def restore_model_passthrough_patch():
|
|||||||
yield
|
yield
|
||||||
finally:
|
finally:
|
||||||
_reset()
|
_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,
|
||||||
|
)
|
||||||
|
|||||||
@@ -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",
|
||||||
|
]
|
||||||
@@ -6,6 +6,8 @@ from unittest.mock import MagicMock, patch
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from EvoScientist.ccproxy_manager import (
|
from EvoScientist.ccproxy_manager import (
|
||||||
|
_CCPROXY_AUTH_TIMEOUT_SECONDS,
|
||||||
|
_CCPROXY_HEALTH_TIMEOUT_SECONDS,
|
||||||
check_ccproxy_auth,
|
check_ccproxy_auth,
|
||||||
ensure_ccproxy,
|
ensure_ccproxy,
|
||||||
is_ccproxy_available,
|
is_ccproxy_available,
|
||||||
@@ -15,6 +17,7 @@ from EvoScientist.ccproxy_manager import (
|
|||||||
setup_codex_env,
|
setup_codex_env,
|
||||||
start_ccproxy,
|
start_ccproxy,
|
||||||
stop_ccproxy,
|
stop_ccproxy,
|
||||||
|
write_ccproxy_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
# =============================================================================
|
# =============================================================================
|
||||||
@@ -52,6 +55,8 @@ class TestCheckCcproxyAuth:
|
|||||||
mock_run.assert_called_once()
|
mock_run.assert_called_once()
|
||||||
cmd = mock_run.call_args[0][0]
|
cmd = mock_run.call_args[0][0]
|
||||||
assert cmd[1:] == ["auth", "status", "claude_api"]
|
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")
|
@patch("subprocess.run")
|
||||||
def test_valid_auth_codex(self, mock_run):
|
def test_valid_auth_codex(self, mock_run):
|
||||||
@@ -123,9 +128,10 @@ class TestIsCcproxyRunning:
|
|||||||
|
|
||||||
|
|
||||||
class TestStartCcproxy:
|
class TestStartCcproxy:
|
||||||
|
@patch("EvoScientist.ccproxy_manager.logger.warning")
|
||||||
@patch("EvoScientist.ccproxy_manager.is_ccproxy_running")
|
@patch("EvoScientist.ccproxy_manager.is_ccproxy_running")
|
||||||
@patch("subprocess.Popen")
|
@patch("subprocess.Popen")
|
||||||
def test_success(self, mock_popen, mock_running):
|
def test_success(self, mock_popen, mock_running, mock_warning):
|
||||||
proc = MagicMock()
|
proc = MagicMock()
|
||||||
proc.poll.return_value = None
|
proc.poll.return_value = None
|
||||||
mock_popen.return_value = proc
|
mock_popen.return_value = proc
|
||||||
@@ -134,6 +140,11 @@ class TestStartCcproxy:
|
|||||||
|
|
||||||
result = start_ccproxy(8000)
|
result = start_ccproxy(8000)
|
||||||
assert result is proc
|
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.is_ccproxy_running", return_value=False)
|
||||||
@patch("EvoScientist.ccproxy_manager.time")
|
@patch("EvoScientist.ccproxy_manager.time")
|
||||||
@@ -143,7 +154,11 @@ class TestStartCcproxy:
|
|||||||
proc.poll.return_value = None
|
proc.poll.return_value = None
|
||||||
mock_popen.return_value = proc
|
mock_popen.return_value = proc
|
||||||
# Simulate time passing beyond deadline
|
# 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()
|
mock_time.sleep = MagicMock()
|
||||||
|
|
||||||
with pytest.raises(RuntimeError, match="did not become healthy"):
|
with pytest.raises(RuntimeError, match="did not become healthy"):
|
||||||
@@ -154,6 +169,54 @@ class TestStartCcproxy:
|
|||||||
with pytest.raises(FileNotFoundError):
|
with pytest.raises(FileNotFoundError):
|
||||||
start_ccproxy(8000)
|
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
|
# ensure_ccproxy
|
||||||
|
|||||||
@@ -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
|
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
|
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 __future__ import annotations
|
||||||
|
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
from langchain_core.messages import AIMessage, HumanMessage
|
||||||
|
|
||||||
from EvoScientist.middleware.code_interpreter import (
|
from EvoScientist.middleware.code_interpreter import (
|
||||||
_DEFAULT_PTC_ALLOWLIST,
|
_DEFAULT_PTC_ALLOWLIST,
|
||||||
@@ -45,3 +49,290 @@ def test_filter_tools_for_ptc_accepts_default_allowlist():
|
|||||||
|
|
||||||
def test_create_code_interpreter_middleware_builds():
|
def test_create_code_interpreter_middleware_builds():
|
||||||
assert create_code_interpreter_middleware() is not None
|
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
|
||||||
|
|||||||
+145
-6
@@ -52,12 +52,9 @@ def _restore_dangerous_env():
|
|||||||
def temp_config_dir(tmp_path, monkeypatch):
|
def temp_config_dir(tmp_path, monkeypatch):
|
||||||
"""Use a temporary directory for config during tests."""
|
"""Use a temporary directory for config during tests."""
|
||||||
config_dir = tmp_path / "evoscientist"
|
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))
|
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
|
# Also clear any API keys from environment
|
||||||
for key in [
|
for key in [
|
||||||
"ANTHROPIC_API_KEY",
|
"ANTHROPIC_API_KEY",
|
||||||
@@ -77,6 +74,9 @@ def temp_config_dir(tmp_path, monkeypatch):
|
|||||||
"EVOSCIENTIST_AUXILIARY_MODEL",
|
"EVOSCIENTIST_AUXILIARY_MODEL",
|
||||||
"EVOSCIENTIST_AUXILIARY_PROVIDER",
|
"EVOSCIENTIST_AUXILIARY_PROVIDER",
|
||||||
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE",
|
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE",
|
||||||
|
"EVOSCIENTIST_OPENROUTER_HTTP_REFERER",
|
||||||
|
"EVOSCIENTIST_OPENROUTER_APP_TITLE",
|
||||||
|
"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES",
|
||||||
"EVOSCIENTIST_DANGEROUS_MODE",
|
"EVOSCIENTIST_DANGEROUS_MODE",
|
||||||
]:
|
]:
|
||||||
monkeypatch.delenv(key, raising=False)
|
monkeypatch.delenv(key, raising=False)
|
||||||
@@ -104,6 +104,9 @@ def clean_env(monkeypatch):
|
|||||||
"EVOSCIENTIST_AUXILIARY_MODEL",
|
"EVOSCIENTIST_AUXILIARY_MODEL",
|
||||||
"EVOSCIENTIST_AUXILIARY_PROVIDER",
|
"EVOSCIENTIST_AUXILIARY_PROVIDER",
|
||||||
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE",
|
"EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE",
|
||||||
|
"EVOSCIENTIST_OPENROUTER_HTTP_REFERER",
|
||||||
|
"EVOSCIENTIST_OPENROUTER_APP_TITLE",
|
||||||
|
"EVOSCIENTIST_OPENROUTER_APP_CATEGORIES",
|
||||||
"EVOSCIENTIST_DANGEROUS_MODE",
|
"EVOSCIENTIST_DANGEROUS_MODE",
|
||||||
]:
|
]:
|
||||||
monkeypatch.delenv(key, raising=False)
|
monkeypatch.delenv(key, raising=False)
|
||||||
@@ -129,8 +132,13 @@ class TestEvoScientistConfig:
|
|||||||
assert config.show_thinking is True
|
assert config.show_thinking is True
|
||||||
assert config.ui_backend == "tui"
|
assert config.ui_backend == "tui"
|
||||||
assert config.log_level == "warning"
|
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_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_profile_enabled is True
|
||||||
assert config.memory_observations_enabled is True
|
assert config.memory_observations_enabled is True
|
||||||
assert config.memory_observation_writer == MemoryObservationWriter.ALL
|
assert config.memory_observation_writer == MemoryObservationWriter.ALL
|
||||||
@@ -145,6 +153,8 @@ class TestEvoScientistConfig:
|
|||||||
assert config.channel_debug_tracing is False
|
assert config.channel_debug_tracing is False
|
||||||
assert config.imessage_enabled is False
|
assert config.imessage_enabled is False
|
||||||
assert config.imessage_allowed_senders == ""
|
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):
|
def test_auth_mode_default(self):
|
||||||
"""Test that anthropic_auth_mode defaults to api_key."""
|
"""Test that anthropic_auth_mode defaults to api_key."""
|
||||||
@@ -192,6 +202,18 @@ class TestEvoScientistConfig:
|
|||||||
assert config.dangerous_mode is True
|
assert config.dangerous_mode is True
|
||||||
assert config.auto_approve 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
|
# Test config path functions
|
||||||
@@ -199,14 +221,34 @@ class TestEvoScientistConfig:
|
|||||||
|
|
||||||
|
|
||||||
class TestConfigPaths:
|
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):
|
def test_get_config_dir_with_xdg(self, monkeypatch, tmp_path):
|
||||||
"""Test config dir uses XDG_CONFIG_HOME when set."""
|
"""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))
|
monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path))
|
||||||
config_dir = get_config_dir()
|
config_dir = get_config_dir()
|
||||||
assert config_dir == tmp_path / "evoscientist"
|
assert config_dir == tmp_path / "evoscientist"
|
||||||
|
|
||||||
def test_get_config_dir_default(self, monkeypatch):
|
def test_get_config_dir_default(self, monkeypatch):
|
||||||
"""Test config dir defaults to ~/.config/evoscientist."""
|
"""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)
|
monkeypatch.delenv("XDG_CONFIG_HOME", raising=False)
|
||||||
config_dir = get_config_dir()
|
config_dir = get_config_dir()
|
||||||
assert config_dir == Path.home() / ".config" / "evoscientist"
|
assert config_dir == Path.home() / ".config" / "evoscientist"
|
||||||
@@ -672,6 +714,22 @@ class TestPriorityChain:
|
|||||||
config = get_effective_config()
|
config = get_effective_config()
|
||||||
assert config.openrouter_anthropic_prompt_cache is False
|
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):
|
def test_set_openrouter_anthropic_prompt_cache(self, temp_config_dir, clean_env):
|
||||||
"""Test OpenRouter Anthropic prompt cache can be set through config."""
|
"""Test OpenRouter Anthropic prompt cache can be set through config."""
|
||||||
save_config(EvoScientistConfig())
|
save_config(EvoScientistConfig())
|
||||||
@@ -732,6 +790,60 @@ class TestApplyConfigToEnv:
|
|||||||
"false"
|
"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):
|
def test_dangerous_mode_round_trips_to_env(self, clean_env, monkeypatch):
|
||||||
"""dangerous_mode set via CLI override must survive a fresh re-read.
|
"""dangerous_mode set via CLI override must survive a fresh re-read.
|
||||||
|
|
||||||
@@ -849,3 +961,30 @@ def test_scheduler_config_defaults_and_env(monkeypatch):
|
|||||||
assert eff2.memory_skill_synthesis_mode == MemorySkillSynthesisMode.AUTO
|
assert eff2.memory_skill_synthesis_mode == MemorySkillSynthesisMode.AUTO
|
||||||
assert eff2.memory_skill_synthesis_cadence == MemorySkillSynthesisCadence.MONTHLY
|
assert eff2.memory_skill_synthesis_cadence == MemorySkillSynthesisCadence.MONTHLY
|
||||||
assert eff2.memory_skill_synthesis_time == "04:30"
|
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
|
||||||
|
|||||||
@@ -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"
|
||||||
@@ -924,6 +924,12 @@ async def test_langgraph_server_gateway_emits_state_interrupt_before_done():
|
|||||||
events = await _collect()
|
events = await _collect()
|
||||||
|
|
||||||
assert events == [
|
assert events == [
|
||||||
|
{
|
||||||
|
"type": "tool_call",
|
||||||
|
"name": "execute",
|
||||||
|
"args": {"command": "echo hello"},
|
||||||
|
"id": "tool-1",
|
||||||
|
},
|
||||||
{
|
{
|
||||||
"type": "interrupt",
|
"type": "interrupt",
|
||||||
"interrupt_id": "interrupt-1",
|
"interrupt_id": "interrupt-1",
|
||||||
|
|||||||
@@ -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"
|
||||||
@@ -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
-2
@@ -1,5 +1,6 @@
|
|||||||
"""Tests for EvoScientist LLM module."""
|
"""Tests for EvoScientist LLM module."""
|
||||||
|
|
||||||
|
from types import SimpleNamespace
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -160,6 +161,68 @@ class TestGetModelInfo:
|
|||||||
|
|
||||||
|
|
||||||
class TestGetChatModel:
|
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")
|
@patch("EvoScientist.llm.models.init_chat_model")
|
||||||
def test_uses_default_model_when_none(self, mock_init):
|
def test_uses_default_model_when_none(self, mock_init):
|
||||||
"""Test that get_chat_model uses default model when model=None."""
|
"""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["temperature"] == 0.7
|
||||||
assert call_kwargs["max_tokens"] == 1000
|
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")
|
@patch("EvoScientist.llm.models.init_chat_model")
|
||||||
def test_infers_openai_from_gpt_prefix(self, mock_init):
|
def test_infers_openai_from_gpt_prefix(self, mock_init):
|
||||||
"""Test that OpenAI is inferred from gpt- prefix."""
|
"""Test that OpenAI is inferred from gpt- prefix."""
|
||||||
@@ -439,6 +535,200 @@ class TestThirdPartyRouting:
|
|||||||
call_kwargs = mock_init.call_args[1]
|
call_kwargs = mock_init.call_args[1]
|
||||||
assert call_kwargs["reasoning"] == {"effort": "medium", "summary": "auto"}
|
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")
|
@patch("EvoScientist.llm.models.init_chat_model")
|
||||||
def test_openrouter_anthropic_prompt_cache_enabled_by_default(
|
def test_openrouter_anthropic_prompt_cache_enabled_by_default(
|
||||||
self, mock_init, monkeypatch
|
self, mock_init, monkeypatch
|
||||||
@@ -957,6 +1247,153 @@ class TestPatchOpenAICompatContent:
|
|||||||
model._astream = AsyncMock()
|
model._astream = AsyncMock()
|
||||||
return model
|
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):
|
def test_generate_flattened(self):
|
||||||
from langchain_core.messages import HumanMessage
|
from langchain_core.messages import HumanMessage
|
||||||
|
|
||||||
@@ -2238,6 +2675,38 @@ class TestPatchOpenrouterStripResponsesReasoning:
|
|||||||
|
|
||||||
|
|
||||||
class TestAutoConfig:
|
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")
|
@patch("EvoScientist.llm.models.init_chat_model")
|
||||||
def test_anthropic_4_5_thinking(self, mock_init, monkeypatch):
|
def test_anthropic_4_5_thinking(self, mock_init, monkeypatch):
|
||||||
"""Anthropic 4-5 models get enabled thinking with budget."""
|
"""Anthropic 4-5 models get enabled thinking with budget."""
|
||||||
@@ -2327,6 +2796,7 @@ class TestAutoConfig:
|
|||||||
"""gpt-5.4+ and codex models get xhigh reasoning."""
|
"""gpt-5.4+ and codex models get xhigh reasoning."""
|
||||||
mock_init.return_value = "mock_model"
|
mock_init.return_value = "mock_model"
|
||||||
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
|
monkeypatch.delenv("OPENAI_BASE_URL", raising=False)
|
||||||
|
monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False)
|
||||||
|
|
||||||
get_chat_model("gpt-5.4", provider="openai")
|
get_chat_model("gpt-5.4", provider="openai")
|
||||||
assert mock_init.call_args[1]["reasoning"] == {
|
assert mock_init.call_args[1]["reasoning"] == {
|
||||||
@@ -2346,6 +2816,26 @@ class TestAutoConfig:
|
|||||||
"summary": "auto",
|
"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")
|
@patch("EvoScientist.llm.models.init_chat_model")
|
||||||
def test_openai_reasoning_high_fallback(self, mock_init, monkeypatch):
|
def test_openai_reasoning_high_fallback(self, mock_init, monkeypatch):
|
||||||
"""Other OpenAI models get high reasoning effort."""
|
"""Other OpenAI models get high reasoning effort."""
|
||||||
@@ -2377,8 +2867,8 @@ class TestAutoConfig:
|
|||||||
assert call_kwargs["model_provider"] == "openai"
|
assert call_kwargs["model_provider"] == "openai"
|
||||||
assert call_kwargs["base_url"] == "http://127.0.0.1:8000/codex/v1"
|
assert call_kwargs["base_url"] == "http://127.0.0.1:8000/codex/v1"
|
||||||
assert call_kwargs["api_key"] == "ccproxy-oauth"
|
assert call_kwargs["api_key"] == "ccproxy-oauth"
|
||||||
# Proxy mode: reasoning skipped (ccproxy untested)
|
# ccproxy uses the Responses API, so reasoning configuration is valid.
|
||||||
assert "reasoning" not in call_kwargs
|
assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"}
|
||||||
# Proxy mode: Responses API (bypasses format chain), streaming ON
|
# Proxy mode: Responses API (bypasses format chain), streaming ON
|
||||||
assert call_kwargs["use_responses_api"] is True
|
assert call_kwargs["use_responses_api"] is True
|
||||||
assert "streaming" not in call_kwargs
|
assert "streaming" not in call_kwargs
|
||||||
@@ -2411,6 +2901,120 @@ class TestAutoConfig:
|
|||||||
call_kwargs = mock_init.call_args[1]
|
call_kwargs = mock_init.call_args[1]
|
||||||
assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"}
|
assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"}
|
||||||
assert "use_responses_api" not in call_kwargs
|
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")
|
@patch("EvoScientist.llm.models.init_chat_model")
|
||||||
def test_openai_ccproxy_key_but_wrong_path_not_ccproxy(
|
def test_openai_ccproxy_key_but_wrong_path_not_ccproxy(
|
||||||
|
|||||||
@@ -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"
|
||||||
@@ -1421,22 +1421,35 @@ class TestLoadToolsProgressCallback:
|
|||||||
]
|
]
|
||||||
|
|
||||||
async def test_failure_emits_start_then_error_with_detail(self, monkeypatch):
|
async def test_failure_emits_start_then_error_with_detail(self, monkeypatch):
|
||||||
from EvoScientist.mcp.client import _load_tools
|
from EvoScientist.mcp import client as mcp_client
|
||||||
|
|
||||||
events: list[tuple[str, str, str]] = []
|
events: list[tuple[str, str, str]] = []
|
||||||
self._patch_client(monkeypatch, {"srv": RuntimeError("boom")})
|
self._patch_client(monkeypatch, {"srv": RuntimeError("boom")})
|
||||||
|
monkeypatch.setattr(mcp_client, "_MCP_SERVER_ERRORS", {})
|
||||||
|
|
||||||
config = {"srv": {"transport": "stdio", "command": "demo"}}
|
config = {"srv": {"transport": "stdio", "command": "demo"}}
|
||||||
|
|
||||||
def record(event, name, detail):
|
def record(event, name, detail):
|
||||||
events.append((event, name, detail))
|
events.append((event, name, detail))
|
||||||
|
|
||||||
await _load_tools(config, on_progress=record)
|
await mcp_client._load_tools(config, on_progress=record)
|
||||||
|
|
||||||
assert events == [
|
assert events == [
|
||||||
("start", "srv", ""),
|
("start", "srv", ""),
|
||||||
("error", "srv", "boom"),
|
("error", "srv", "boom"),
|
||||||
]
|
]
|
||||||
|
assert mcp_client.get_mcp_server_errors() == {"srv": "boom"}
|
||||||
|
|
||||||
|
async def test_success_clears_previous_server_error(self, monkeypatch):
|
||||||
|
from EvoScientist.mcp import client as mcp_client
|
||||||
|
|
||||||
|
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):
|
async def test_mixed_fleet_reports_each_server_independently(self, monkeypatch):
|
||||||
from EvoScientist.mcp.client import _load_tools
|
from EvoScientist.mcp.client import _load_tools
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ fallback chain behaviour via _try_fallbacks / _guard_and_fallback.
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -83,6 +84,7 @@ class TestIsNonFallbackable:
|
|||||||
"Error 400: invalid_request_error",
|
"Error 400: invalid_request_error",
|
||||||
"400 Bad Request: invalid request body",
|
"400 Bad Request: invalid request body",
|
||||||
"400: malformed JSON in request",
|
"400: malformed JSON in request",
|
||||||
|
"<400> InvalidParameter: Repetitive tool calls detected in history",
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
def test_malformed_request_400_patterns(self, msg):
|
def test_malformed_request_400_patterns(self, msg):
|
||||||
@@ -223,6 +225,89 @@ class TestTryFallbacks:
|
|||||||
# fb-b should never be reached.
|
# fb-b should never be reached.
|
||||||
assert mock_gcm.call_count == 1
|
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
|
# 3. _guard_and_fallback — pre-check before chain walk
|
||||||
@@ -242,6 +327,34 @@ class TestGuardAndFallback:
|
|||||||
|
|
||||||
invoke.assert_not_awaited()
|
invoke.assert_not_awaited()
|
||||||
|
|
||||||
|
async def test_context_overflow_with_provider_model_passes_through_unwrapped(self):
|
||||||
|
"""Regression: a ``ContextOverflowError`` entering
|
||||||
|
``_guard_and_fallback`` under a recognized-provider model must
|
||||||
|
come out unwrapped. Otherwise ``_raise_normalized`` →
|
||||||
|
``_normalize`` would wrap it as a ``ProviderStreamError`` and
|
||||||
|
deepagents' ``SummarizationMiddleware`` (which sits outside
|
||||||
|
the user middleware stack and catches by exact type) would
|
||||||
|
stop compressing history and retrying.
|
||||||
|
"""
|
||||||
|
add_fallback("fb", "prov")
|
||||||
|
req = _fake_request()
|
||||||
|
# Recognized provider — without the gate in ``_normalize`` this
|
||||||
|
# would wrap. With the gate, the raw type propagates.
|
||||||
|
cls = type(
|
||||||
|
"ChatOpenAI", (), {"__module__": "langchain_openai.chat_models.base"}
|
||||||
|
)
|
||||||
|
model = cls()
|
||||||
|
model.openai_api_base = None
|
||||||
|
req.model = model
|
||||||
|
invoke = AsyncMock()
|
||||||
|
|
||||||
|
raised = ContextOverflowError("context length exceeded")
|
||||||
|
with pytest.raises(ContextOverflowError) as exc_info:
|
||||||
|
await _guard_and_fallback(raised, req, invoke)
|
||||||
|
|
||||||
|
assert exc_info.value is raised
|
||||||
|
invoke.assert_not_awaited()
|
||||||
|
|
||||||
async def test_malformed_400_raises_immediately(self):
|
async def test_malformed_400_raises_immediately(self):
|
||||||
add_fallback("fb", "prov")
|
add_fallback("fb", "prov")
|
||||||
req = _fake_request()
|
req = _fake_request()
|
||||||
|
|||||||
@@ -2294,7 +2294,11 @@ def test_memory_worker_observation_writer_modes(
|
|||||||
observation_writer=observation_writer,
|
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
|
assert _memory_tool_names(middleware) == expected_tools
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ def _restore_paths():
|
|||||||
"GLOBAL_MEMORIES_DIR": paths.GLOBAL_MEMORIES_DIR,
|
"GLOBAL_MEMORIES_DIR": paths.GLOBAL_MEMORIES_DIR,
|
||||||
"USER_SKILLS_DIR": paths.USER_SKILLS_DIR,
|
"USER_SKILLS_DIR": paths.USER_SKILLS_DIR,
|
||||||
"_active_workspace": paths._active_workspace,
|
"_active_workspace": paths._active_workspace,
|
||||||
|
"_EVOSCIENTIST_DATA_ROOT": paths._EVOSCIENTIST_DATA_ROOT,
|
||||||
}
|
}
|
||||||
yield
|
yield
|
||||||
paths.WORKSPACE_ROOT = orig["WORKSPACE_ROOT"]
|
paths.WORKSPACE_ROOT = orig["WORKSPACE_ROOT"]
|
||||||
@@ -32,6 +33,7 @@ def _restore_paths():
|
|||||||
paths.GLOBAL_MEMORIES_DIR = orig["GLOBAL_MEMORIES_DIR"]
|
paths.GLOBAL_MEMORIES_DIR = orig["GLOBAL_MEMORIES_DIR"]
|
||||||
paths.USER_SKILLS_DIR = orig["USER_SKILLS_DIR"]
|
paths.USER_SKILLS_DIR = orig["USER_SKILLS_DIR"]
|
||||||
paths._active_workspace = orig["_active_workspace"]
|
paths._active_workspace = orig["_active_workspace"]
|
||||||
|
paths._EVOSCIENTIST_DATA_ROOT = orig["_EVOSCIENTIST_DATA_ROOT"]
|
||||||
|
|
||||||
|
|
||||||
class TestSetWorkspaceRoot:
|
class TestSetWorkspaceRoot:
|
||||||
@@ -140,6 +142,63 @@ class TestDataDir:
|
|||||||
assert paths.GLOBAL_MEMORIES_DIR == paths.DATA_DIR / "memories"
|
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:
|
class TestLegacySessionsDbMigration:
|
||||||
"""Tests for migrate_legacy_sessions_db() — transitional helper.
|
"""Tests for migrate_legacy_sessions_db() — transitional helper.
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -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")]
|
||||||
@@ -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,
|
||||||
|
}
|
||||||
@@ -11,6 +11,7 @@ from langgraph.checkpoint.memory import InMemorySaver
|
|||||||
from langgraph.types import Command, Interrupt
|
from langgraph.types import Command, Interrupt
|
||||||
|
|
||||||
from EvoScientist.middleware.ask_user import AskUserMiddleware
|
from EvoScientist.middleware.ask_user import AskUserMiddleware
|
||||||
|
from EvoScientist.stream.emitter import STREAM_PROTOCOL_CAPABILITIES
|
||||||
from EvoScientist.stream.events import stream_agent_events
|
from EvoScientist.stream.events import stream_agent_events
|
||||||
from EvoScientist.stream.summarization import (
|
from EvoScientist.stream.summarization import (
|
||||||
_extract_summary_message_text,
|
_extract_summary_message_text,
|
||||||
@@ -1121,6 +1122,74 @@ class TestUsageStatsExtraction:
|
|||||||
assert len(usage_events) == 0
|
assert len(usage_events) == 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestCanonicalSourceCapabilities:
|
||||||
|
async def test_root_update_emits_full_task_snapshot_and_empty_clear(self):
|
||||||
|
agent = FakeV3Agent(
|
||||||
|
[
|
||||||
|
protocol_event(
|
||||||
|
"updates",
|
||||||
|
{"model": {"todos": [{"content": "Inspect", "status": "active"}]}},
|
||||||
|
),
|
||||||
|
protocol_event("updates", {"model": {"todos": []}}),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
events = await collect_events(agent)
|
||||||
|
snapshots = [event for event in events if event.get("type") == "task_snapshot"]
|
||||||
|
assert snapshots == [
|
||||||
|
{
|
||||||
|
"type": "task_snapshot",
|
||||||
|
"source": "update",
|
||||||
|
"items": [{"content": "Inspect", "status": "in_progress"}],
|
||||||
|
},
|
||||||
|
{"type": "task_snapshot", "source": "update", "items": []},
|
||||||
|
]
|
||||||
|
|
||||||
|
async def test_subagent_todos_do_not_replace_root_snapshot(self):
|
||||||
|
agent = FakeV3Agent(
|
||||||
|
[
|
||||||
|
protocol_event(
|
||||||
|
"updates",
|
||||||
|
{"todos": [{"content": "Nested", "status": "pending"}]},
|
||||||
|
namespace=("subagent",),
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
events = await collect_events(agent)
|
||||||
|
assert not any(event.get("type") == "task_snapshot" for event in events)
|
||||||
|
|
||||||
|
async def test_invalid_tool_call_candidate_is_not_committed_by_stream_processor(self):
|
||||||
|
invalid = AIMessage(
|
||||||
|
content="",
|
||||||
|
invalid_tool_calls=[
|
||||||
|
{
|
||||||
|
"name": "write_todos",
|
||||||
|
"args": "{bad",
|
||||||
|
"id": "call-invalid",
|
||||||
|
"error": "invalid json",
|
||||||
|
"type": "invalid_tool_call",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
agent = FakeV3Agent(
|
||||||
|
[
|
||||||
|
protocol_event("messages", (invalid, {})),
|
||||||
|
message_finish(),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
events = await collect_events(agent)
|
||||||
|
assert not any(event.get("type") in {"tool_call", "error"} for event in events)
|
||||||
|
|
||||||
|
def test_stream_capabilities_are_explicit(self):
|
||||||
|
assert STREAM_PROTOCOL_CAPABILITIES == frozenset(
|
||||||
|
{
|
||||||
|
"task_snapshot_v1",
|
||||||
|
"complete_tool_call_v1",
|
||||||
|
"correlated_tool_call_id_v1",
|
||||||
|
"final_invalid_tool_call_v1",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestSummarizationHelpers:
|
class TestSummarizationHelpers:
|
||||||
"""Summarization extraction helpers."""
|
"""Summarization extraction helpers."""
|
||||||
|
|
||||||
|
|||||||
@@ -9,10 +9,12 @@ so they actually verify the two claims the recovery rests on:
|
|||||||
left intact, so a pending question is never silently discarded.
|
left intact, so a pending question is never silently discarded.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from typing import TypedDict
|
from typing import Annotated, TypedDict
|
||||||
|
|
||||||
|
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
||||||
from langgraph.checkpoint.memory import InMemorySaver
|
from langgraph.checkpoint.memory import InMemorySaver
|
||||||
from langgraph.graph import END, START, StateGraph
|
from langgraph.graph import END, START, StateGraph
|
||||||
|
from langgraph.graph.message import add_messages
|
||||||
from langgraph.types import interrupt
|
from langgraph.types import interrupt
|
||||||
|
|
||||||
from EvoScientist.stream.events import _clear_interrupted_graph_state
|
from EvoScientist.stream.events import _clear_interrupted_graph_state
|
||||||
@@ -22,6 +24,10 @@ class _S(TypedDict):
|
|||||||
x: int
|
x: int
|
||||||
|
|
||||||
|
|
||||||
|
class _MessageState(TypedDict):
|
||||||
|
messages: Annotated[list, add_messages]
|
||||||
|
|
||||||
|
|
||||||
def _crashing_app():
|
def _crashing_app():
|
||||||
# Node 'b' crashes once, then succeeds — so a post-recovery run can complete
|
# Node 'b' crashes once, then succeeds — so a post-recovery run can complete
|
||||||
# and prove the graph is genuinely unstuck (not replaying the dead step).
|
# and prove the graph is genuinely unstuck (not replaying the dead step).
|
||||||
@@ -57,6 +63,75 @@ def _interrupting_app():
|
|||||||
return g.compile(checkpointer=InMemorySaver())
|
return g.compile(checkpointer=InMemorySaver())
|
||||||
|
|
||||||
|
|
||||||
|
def _invalid_tool_call_app():
|
||||||
|
def write_invalid_call(state):
|
||||||
|
return {
|
||||||
|
"messages": [
|
||||||
|
AIMessage(
|
||||||
|
content="",
|
||||||
|
invalid_tool_calls=[
|
||||||
|
{
|
||||||
|
"type": "invalid_tool_call",
|
||||||
|
"id": None,
|
||||||
|
"name": "execute",
|
||||||
|
"args": '{"command":',
|
||||||
|
"error": "bad json",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
def crash(state):
|
||||||
|
raise RuntimeError("provider stream failed")
|
||||||
|
|
||||||
|
g = StateGraph(_MessageState)
|
||||||
|
g.add_node("write_invalid_call", write_invalid_call)
|
||||||
|
g.add_node("crash", crash)
|
||||||
|
g.add_edge(START, "write_invalid_call")
|
||||||
|
g.add_edge("write_invalid_call", "crash")
|
||||||
|
g.add_edge("crash", END)
|
||||||
|
return g.compile(checkpointer=InMemorySaver())
|
||||||
|
|
||||||
|
|
||||||
|
def _repetitive_tool_call_app():
|
||||||
|
messages = [HumanMessage(content="inspect")]
|
||||||
|
for call_id in ("call-1", "call-2"):
|
||||||
|
messages.extend(
|
||||||
|
[
|
||||||
|
AIMessage(
|
||||||
|
content="",
|
||||||
|
tool_calls=[
|
||||||
|
{
|
||||||
|
"id": call_id,
|
||||||
|
"name": "execute",
|
||||||
|
"args": {"command": "pwd"},
|
||||||
|
}
|
||||||
|
],
|
||||||
|
),
|
||||||
|
ToolMessage(
|
||||||
|
content="/workspace",
|
||||||
|
tool_call_id=call_id,
|
||||||
|
name="execute",
|
||||||
|
),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
def write_repetitive_history(state):
|
||||||
|
return {"messages": messages}
|
||||||
|
|
||||||
|
def crash(state):
|
||||||
|
raise RuntimeError("provider rejected repetitive tool history")
|
||||||
|
|
||||||
|
g = StateGraph(_MessageState)
|
||||||
|
g.add_node("write_repetitive_history", write_repetitive_history)
|
||||||
|
g.add_node("crash", crash)
|
||||||
|
g.add_edge(START, "write_repetitive_history")
|
||||||
|
g.add_edge("write_repetitive_history", "crash")
|
||||||
|
g.add_edge("crash", END)
|
||||||
|
return g.compile(checkpointer=InMemorySaver())
|
||||||
|
|
||||||
|
|
||||||
async def test_recovery_clears_stuck_state_after_crash():
|
async def test_recovery_clears_stuck_state_after_crash():
|
||||||
app = _crashing_app()
|
app = _crashing_app()
|
||||||
cfg = {"configurable": {"thread_id": "t1"}}
|
cfg = {"configurable": {"thread_id": "t1"}}
|
||||||
@@ -91,3 +166,52 @@ async def test_recovery_preserves_pending_hitl_interrupt():
|
|||||||
after = app.get_state(cfg)
|
after = app.get_state(cfg)
|
||||||
assert after.next == ("ask",) # interrupt left intact, still resumable
|
assert after.next == ("ask",) # interrupt left intact, still resumable
|
||||||
assert after.interrupts
|
assert after.interrupts
|
||||||
|
|
||||||
|
|
||||||
|
async def test_recovery_removes_invalid_tool_call_from_checkpoint():
|
||||||
|
app = _invalid_tool_call_app()
|
||||||
|
cfg = {"configurable": {"thread_id": "tool-history"}}
|
||||||
|
try:
|
||||||
|
await app.ainvoke(
|
||||||
|
{"messages": [HumanMessage(content="run the command")]},
|
||||||
|
cfg,
|
||||||
|
)
|
||||||
|
except RuntimeError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
before = await app.aget_state(cfg)
|
||||||
|
assert before.next == ("crash",)
|
||||||
|
assert any(
|
||||||
|
isinstance(message, AIMessage) and message.invalid_tool_calls
|
||||||
|
for message in before.values["messages"]
|
||||||
|
)
|
||||||
|
|
||||||
|
await _clear_interrupted_graph_state(app, cfg)
|
||||||
|
|
||||||
|
after = await app.aget_state(cfg)
|
||||||
|
assert after.next == ()
|
||||||
|
assert [message.type for message in after.values["messages"]] == ["human"]
|
||||||
|
|
||||||
|
|
||||||
|
async def test_recovery_preserves_complete_repetitive_tool_rounds_in_checkpoint():
|
||||||
|
app = _repetitive_tool_call_app()
|
||||||
|
cfg = {"configurable": {"thread_id": "tool-loop-history"}}
|
||||||
|
try:
|
||||||
|
await app.ainvoke({"messages": []}, cfg)
|
||||||
|
except RuntimeError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
before = await app.aget_state(cfg)
|
||||||
|
assert before.next == ("crash",)
|
||||||
|
assert len(before.values["messages"]) == 5
|
||||||
|
|
||||||
|
await _clear_interrupted_graph_state(app, cfg)
|
||||||
|
|
||||||
|
after = await app.aget_state(cfg)
|
||||||
|
messages = after.values["messages"]
|
||||||
|
assert after.next == ()
|
||||||
|
assert [message.type for message in messages] == ["human", "ai", "tool", "ai", "tool"]
|
||||||
|
assert [messages[1].tool_calls[0]["id"], messages[3].tool_calls[0]["id"]] == [
|
||||||
|
"call-1",
|
||||||
|
"call-2",
|
||||||
|
]
|
||||||
|
|||||||
@@ -0,0 +1,246 @@
|
|||||||
|
"""Final model tool protocol validation tests."""
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field, replace
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from langchain.agents.middleware.types import ExtendedModelResponse, ModelResponse
|
||||||
|
from langchain_core.messages import AIMessage
|
||||||
|
|
||||||
|
from EvoScientist.llm.errors import ModelToolProtocolError
|
||||||
|
from EvoScientist.middleware.tool_protocol_guard import ToolProtocolGuardMiddleware
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class _Request:
|
||||||
|
tools: list[Any]
|
||||||
|
model: Any = field(default_factory=lambda: SimpleNamespace(metadata={}))
|
||||||
|
|
||||||
|
def override(self, **updates: Any):
|
||||||
|
return replace(self, **updates)
|
||||||
|
|
||||||
|
|
||||||
|
def _response(*calls: dict[str, Any], content: Any = "") -> ModelResponse:
|
||||||
|
return ModelResponse(result=[AIMessage(content=content, tool_calls=list(calls))])
|
||||||
|
|
||||||
|
|
||||||
|
def _call(call_id: str = "call-1", name: str = "search", args: Any = None):
|
||||||
|
return {"id": call_id, "name": name, "args": {} if args is None else args}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("call", "reason"),
|
||||||
|
[
|
||||||
|
(_call(name=""), "missing_name"),
|
||||||
|
(_call(name=" "), "missing_name"),
|
||||||
|
(_call(name="missing"), "unknown_name"),
|
||||||
|
(_call(call_id=""), "missing_id"),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_invalid_final_tool_call_fails_closed(call, reason):
|
||||||
|
middleware = ToolProtocolGuardMiddleware()
|
||||||
|
request = _Request(tools=[{"name": "search"}])
|
||||||
|
|
||||||
|
with pytest.raises(ModelToolProtocolError) as caught:
|
||||||
|
middleware.wrap_model_call(request, lambda _request: _response(call))
|
||||||
|
|
||||||
|
assert caught.value.reason == reason
|
||||||
|
assert caught.value.retryable is False
|
||||||
|
assert caught.value.fallbackable is True
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_mapping_args_are_rejected_if_adapter_bypasses_message_validation():
|
||||||
|
message = AIMessage(content="", tool_calls=[_call()])
|
||||||
|
message.tool_calls[0]["args"] = "{}"
|
||||||
|
|
||||||
|
with pytest.raises(ModelToolProtocolError) as caught:
|
||||||
|
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||||
|
_Request(tools=[{"name": "search"}]),
|
||||||
|
lambda _request: ModelResponse(result=[message]),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert caught.value.reason == "invalid_args"
|
||||||
|
|
||||||
|
|
||||||
|
def test_duplicate_parallel_call_id_rejects_whole_response():
|
||||||
|
request = _Request(tools=[{"name": "search"}, {"name": "read_file"}])
|
||||||
|
response = _response(_call(name="search"), _call(name="read_file"))
|
||||||
|
|
||||||
|
with pytest.raises(
|
||||||
|
ModelToolProtocolError, match="invalid structured tool call"
|
||||||
|
) as caught:
|
||||||
|
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||||
|
request, lambda _request: response
|
||||||
|
)
|
||||||
|
|
||||||
|
assert caught.value.reason == "duplicate_id"
|
||||||
|
|
||||||
|
|
||||||
|
def test_one_invalid_parallel_call_rejects_atomically():
|
||||||
|
request = _Request(tools=[{"name": "search"}])
|
||||||
|
response = _response(_call(call_id="one"), _call(call_id="two", name="missing"))
|
||||||
|
|
||||||
|
with pytest.raises(ModelToolProtocolError) as caught:
|
||||||
|
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||||
|
request, lambda _request: response
|
||||||
|
)
|
||||||
|
|
||||||
|
assert caught.value.reason == "unknown_name"
|
||||||
|
|
||||||
|
|
||||||
|
def test_final_invalid_tool_calls_are_rejected():
|
||||||
|
message = AIMessage(
|
||||||
|
content="",
|
||||||
|
invalid_tool_calls=[
|
||||||
|
{"id": "bad", "name": "search", "args": "{", "error": "bad json"}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ModelToolProtocolError) as caught:
|
||||||
|
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||||
|
_Request(tools=[{"name": "search"}]),
|
||||||
|
lambda _request: ModelResponse(result=[message]),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert caught.value.reason == "invalid_final_call"
|
||||||
|
assert caught.value.call_id == "bad"
|
||||||
|
|
||||||
|
|
||||||
|
def test_content_block_must_match_parsed_call():
|
||||||
|
response = _response(
|
||||||
|
_call(),
|
||||||
|
content=[
|
||||||
|
{"type": "tool_call", "id": "call-1", "name": "read_file", "args": {}}
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ModelToolProtocolError) as caught:
|
||||||
|
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||||
|
_Request(tools=[{"name": "search"}, {"name": "read_file"}]),
|
||||||
|
lambda _request: response,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert caught.value.reason == "inconsistent_block"
|
||||||
|
|
||||||
|
|
||||||
|
def test_parsed_only_valid_call_and_extended_response_pass():
|
||||||
|
response = ExtendedModelResponse(model_response=_response(_call()))
|
||||||
|
result = ToolProtocolGuardMiddleware().wrap_model_call(
|
||||||
|
_Request(tools=[{"type": "function", "function": {"name": "search"}}]),
|
||||||
|
lambda _request: response,
|
||||||
|
)
|
||||||
|
assert result is response
|
||||||
|
|
||||||
|
|
||||||
|
async def test_async_direct_ai_message_shape_passes():
|
||||||
|
response = AIMessage(content="", tool_calls=[_call()])
|
||||||
|
|
||||||
|
async def handler(_request):
|
||||||
|
return response
|
||||||
|
|
||||||
|
result = await ToolProtocolGuardMiddleware().awrap_model_call(
|
||||||
|
_Request(tools=[{"name": "search"}]), handler
|
||||||
|
)
|
||||||
|
assert result is response
|
||||||
|
|
||||||
|
|
||||||
|
def test_error_carries_safe_route_metadata():
|
||||||
|
model = SimpleNamespace(
|
||||||
|
metadata={
|
||||||
|
"route_provider": "openai",
|
||||||
|
"route_model": "gpt-example",
|
||||||
|
"route_key": "route-safe",
|
||||||
|
"route_config_generation": 12,
|
||||||
|
"route_api_mode": "chat_completions",
|
||||||
|
"route_endpoint": "primary",
|
||||||
|
"route_tool_call_transport": "streaming",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ModelToolProtocolError) as caught:
|
||||||
|
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||||
|
_Request(tools=[{"name": "search"}], model=model),
|
||||||
|
lambda _request: _response(_call(name="")),
|
||||||
|
)
|
||||||
|
|
||||||
|
payload = caught.value.model_dump()
|
||||||
|
assert payload["route_key"] == "route-safe"
|
||||||
|
assert payload["config_generation"] == 12
|
||||||
|
assert payload["endpoint"] == "primary"
|
||||||
|
assert payload["tool_call_transport"] == "streaming"
|
||||||
|
assert "args" not in payload
|
||||||
|
|
||||||
|
|
||||||
|
def test_missing_id_carries_redacted_call_diagnostic_only_for_internal_logging():
|
||||||
|
call = _call(call_id="", name="search", args={"query": "private search text"})
|
||||||
|
|
||||||
|
with pytest.raises(ModelToolProtocolError) as caught:
|
||||||
|
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||||
|
_Request(tools=[{"name": "search"}]),
|
||||||
|
lambda _request: _response(call),
|
||||||
|
)
|
||||||
|
|
||||||
|
diagnostic = caught.value.call_diagnostic
|
||||||
|
assert diagnostic == {
|
||||||
|
"source": "parsed_tool_calls",
|
||||||
|
"call_index": 0,
|
||||||
|
"call_count": 1,
|
||||||
|
"call_type": "object",
|
||||||
|
"name": "search",
|
||||||
|
"id_present": False,
|
||||||
|
"args_present": True,
|
||||||
|
"args_type": "object",
|
||||||
|
"args_key_count": 1,
|
||||||
|
"args_keys": ["query"],
|
||||||
|
"args_keys_truncated": False,
|
||||||
|
"args_digest": diagnostic["args_digest"],
|
||||||
|
"raw_openai_call_available": False,
|
||||||
|
}
|
||||||
|
assert diagnostic["args_digest"].startswith("sha256:")
|
||||||
|
assert "private search text" not in str(diagnostic)
|
||||||
|
assert "call_diagnostic" not in caught.value.model_dump()
|
||||||
|
|
||||||
|
|
||||||
|
def test_diagnostic_compares_parsed_and_preserved_raw_openai_call_shapes():
|
||||||
|
parsed = _call(call_id="", name="search", args={"query": "secret"})
|
||||||
|
raw = {
|
||||||
|
"id": "provider-call-id",
|
||||||
|
"type": "function",
|
||||||
|
"function": {"name": "search", "arguments": '{"query":"secret"}'},
|
||||||
|
}
|
||||||
|
message = AIMessage(
|
||||||
|
content="",
|
||||||
|
tool_calls=[parsed],
|
||||||
|
additional_kwargs={"tool_calls": [raw]},
|
||||||
|
)
|
||||||
|
|
||||||
|
with pytest.raises(ModelToolProtocolError) as caught:
|
||||||
|
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||||
|
_Request(tools=[{"name": "search"}]),
|
||||||
|
lambda _request: ModelResponse(result=[message]),
|
||||||
|
)
|
||||||
|
|
||||||
|
diagnostic = caught.value.call_diagnostic
|
||||||
|
assert diagnostic["id_present"] is False
|
||||||
|
assert diagnostic["raw_openai_call_available"] is True
|
||||||
|
assert diagnostic["raw_openai_call"]["id_present"] is True
|
||||||
|
assert diagnostic["raw_openai_call"]["name"] == "search"
|
||||||
|
assert "provider-call-id" not in str(diagnostic)
|
||||||
|
assert "secret" not in str(diagnostic)
|
||||||
|
|
||||||
|
|
||||||
|
def test_diagnostic_failure_cannot_mask_the_protocol_error():
|
||||||
|
circular: dict[str, Any] = {}
|
||||||
|
circular["self"] = circular
|
||||||
|
message = AIMessage(content="", tool_calls=[_call(call_id="", args={})])
|
||||||
|
message.tool_calls[0]["args"] = circular
|
||||||
|
|
||||||
|
with pytest.raises(ModelToolProtocolError) as caught:
|
||||||
|
ToolProtocolGuardMiddleware().wrap_model_call(
|
||||||
|
_Request(tools=[{"name": "search"}]),
|
||||||
|
lambda _request: ModelResponse(result=[message]),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert caught.value.reason == "missing_id"
|
||||||
|
assert caught.value.call_diagnostic["args_digest"].startswith("sha256:")
|
||||||
@@ -944,7 +944,7 @@ wheels = [
|
|||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "evoscientist"
|
name = "evoscientist"
|
||||||
version = "0.2.1"
|
version = "0.2.2"
|
||||||
source = { editable = "." }
|
source = { editable = "." }
|
||||||
dependencies = [
|
dependencies = [
|
||||||
{ name = "deepagents", extra = ["quickjs"] },
|
{ name = "deepagents", extra = ["quickjs"] },
|
||||||
|
|||||||
Reference in New Issue
Block a user