feat(runtime)!: switch middleware and agent factories to snapshot-driven models

Replace config.yaml-driven model selection with registry snapshot resolution
across the runtime chain:

- ConfigurableModelMiddleware reads configurable["runtime_snapshot_id"] only;
  model/model_provider overrides are rejected with MODEL_CONFIG_OUTSIDE_SNAPSHOT
- MessageBudgetMiddleware derives budgets from snapshot reserves
  (system/tools/attachments) and re-resolves the summarizer per snapshot
- Agent factory and subagent factory resolve models via SnapshotRuntime
  (auxiliary/tool_selector/scheduler -> defaults.auxiliary ?? defaults.primary)
- Remove ModelFallbackMiddleware, /model-fallback command, and fallback chain
- Add model_registry/runtime.py SnapshotRuntime glue layer

Legacy config.yaml LLM fields, /model command, and llm/models.py remain for
Task 7. Report: .superpowers/sdd/briefs/task-6-report.md
This commit is contained in:
m4
2026-07-21 12:54:35 +08:00
parent b2660fc38c
commit 57176b359a
26 changed files with 2054 additions and 2014 deletions
+56 -49
View File
@@ -63,11 +63,12 @@ _chat_model = None
_chat_model_key: tuple[str | None, str | None] | None = None
# Auxiliary model for background/helper LLM calls (memory workers + main-agent
# tool selector). Cached separately from the main model; falls back to the main
# instance when the auxiliary_* config fields are empty (see
# tool selector). Cached separately from the main model; resolved through the
# registry ``auxiliary`` role mapping and falls back to the main instance when
# no distinct auxiliary default is configured (see
# _ensure_auxiliary_chat_model).
_auxiliary_chat_model = None
_auxiliary_chat_model_key: tuple[str | None, str | None] | None = None
_auxiliary_chat_model_key: tuple[str, str, int] | None = None
# Cache MCP tools by the effective config signature to avoid reconnecting
# to MCP servers on every `/new` when config is unchanged.
@@ -166,25 +167,38 @@ def _ensure_chat_model():
def _ensure_auxiliary_chat_model():
"""Return the auxiliary chat model for background/helper LLM calls.
Resolves ``(cfg.auxiliary_model or cfg.model, cfg.auxiliary_provider or
cfg.provider)``. When the auxiliary fields are empty — or resolve to the same
``(model, provider)`` pair as the main model — returns the main
``_ensure_chat_model()`` instance directly, so no second client is built.
Otherwise it is cached separately under its own key. Onboard sets the
provider alongside the model, so the ``or cfg.provider`` fallback only
matters for a model set without an explicit auxiliary provider.
Resolves the ``auxiliary`` role through the model registry role mapping
(design doc 6.1): ``registry.defaults.auxiliary ?? defaults.primary``,
constructed via ``build_chat_model`` from the resolved configuration.
The legacy ``cfg.auxiliary_model``/``cfg.auxiliary_provider`` free
strings are no longer consulted.
When the registry is still in bootstrap (the CLI local snapshot entry
point is wired separately) or no auxiliary default is configured — or it
matches the primary default — returns the main ``_ensure_chat_model()``
instance directly, so no second client is built. Otherwise the model is
cached under its ``(provider_id, model_key, registry_revision)`` key.
"""
global _auxiliary_chat_model, _auxiliary_chat_model_key
from .llm import get_chat_model
from .model_registry.errors import MODEL_REGISTRY_NOT_READY, ModelRegistryError
from .model_registry.runtime import get_snapshot_runtime
cfg = _ensure_config()
aux_model = cfg.auxiliary_model or cfg.model
aux_provider = cfg.auxiliary_provider or cfg.provider
if (aux_model, aux_provider) == (cfg.model, cfg.provider):
runtime = get_snapshot_runtime()
try:
primary_ref, auxiliary_ref, revision = runtime.registry_defaults()
except ModelRegistryError as exc:
if exc.code == MODEL_REGISTRY_NOT_READY:
return _ensure_chat_model()
raise
if auxiliary_ref is None or auxiliary_ref == primary_ref:
return _ensure_chat_model()
key = (aux_model, aux_provider)
key: tuple[str, str, int] = (
auxiliary_ref.provider_id,
auxiliary_ref.model_key,
revision,
)
if _auxiliary_chat_model is None or _auxiliary_chat_model_key != key:
_auxiliary_chat_model = get_chat_model(model=aux_model, provider=aux_provider)
_auxiliary_chat_model = runtime.build_default_role_model("auxiliary")
_auxiliary_chat_model_key = key
return _auxiliary_chat_model
@@ -201,8 +215,9 @@ def set_chat_model(model: str, provider: str | None = None):
"""
from .llm import get_chat_model
# Invalidate the auxiliary cache too: when auxiliary_* is empty it mirrors
# the main model, so a /model switch must let it re-resolve to the new main.
# Invalidate the auxiliary cache too: when the registry has no distinct
# auxiliary default the auxiliary model mirrors the main one, so a /model
# switch must let it re-resolve to the new main.
global _auxiliary_chat_model, _auxiliary_chat_model_key
_auxiliary_chat_model = None
_auxiliary_chat_model_key = None
@@ -341,10 +356,7 @@ def _inject_subagent_middleware(
ToolErrorHandlerMiddleware(),
ContextOverflowMapperMiddleware(),
]
if (
memory_controls.memory_enabled
and cfg.workspace_isolation != "required"
):
if memory_controls.memory_enabled and cfg.workspace_isolation != "required":
middleware.append(memory_middleware)
if (
memory_controls.worker_needed(MemoryObservationTarget.SUBAGENT_WORKER)
@@ -469,9 +481,10 @@ def _maybe_swap_async_subagents(
middleware.append(AsyncWatcherMiddleware(agent_specs))
# Forward the CLI's live (model, provider) into deepagents'
# start/update_async_task tool calls so the deployed graph can
# re-resolve its chat model per run via ConfigurableModelMiddleware.
# Wrap deepagents' start/update_async_task tool calls so workspace scope
# and usage correlation metadata reach the deployed graph's runs. Model
# configuration is NOT forwarded: the deployed graph resolves its model
# per run from ``runtime_snapshot_id`` via ConfigurableModelMiddleware.
# Idempotent — safe to call on every CLI startup.
if agent_specs:
from .llm.patches import _patch_deepagents_model_passthrough
@@ -659,7 +672,9 @@ def _get_default_backend():
if cfg.workspace_isolation == "legacy":
return _get_legacy_backend()
if cfg.workspace_isolation == "required" and cfg.dangerous_mode:
raise RuntimeError("dangerous_mode is incompatible with required workspace isolation")
raise RuntimeError(
"dangerous_mode is incompatible with required workspace isolation"
)
if cfg.workspace_isolation == "required":
verify_required_cutover(_paths_mod.WORKSPACE_ROOT)
verify_required_executor()
@@ -670,9 +685,7 @@ def _get_default_backend():
# The CLI and its stripped async-subagent service keep their configured
# shared workspace. The WebUI deployment must receive a scope from the
# trusted WebUI/API boundary instead.
allow_unscoped_legacy=os.environ.get(
"EVOSCIENTIST_DEPLOY_MODE", ""
).lower()
allow_unscoped_legacy=os.environ.get("EVOSCIENTIST_DEPLOY_MODE", "").lower()
!= "full",
)
@@ -685,6 +698,7 @@ def _get_default_middleware(
chat_model=None,
backend=None,
memory_source_agent: str = "EvoScientist",
snapshot_role: str = "primary",
):
"""Build the default middleware list.
@@ -704,11 +718,14 @@ def _get_default_middleware(
(avoids writing module globals on the pure path).
memory_source_agent: Attribution name for profile/observation writes.
Async sub-agent factories pass their deployed agent name here.
snapshot_role: The model role ``ConfigurableModelMiddleware``
resolves from the run snapshot (``primary`` for the main agent
and working sub-agents, ``auxiliary`` for unattended helper
agents such as the scheduler).
"""
from .middleware import (
ConfigurableModelMiddleware,
ContextOverflowMapperMiddleware,
ModelFallbackMiddleware,
ToolErrorHandlerMiddleware,
create_code_interpreter_middleware,
create_context_editing_middleware,
@@ -719,12 +736,9 @@ def _get_default_middleware(
create_scheduler_middleware,
create_tool_selector_middleware,
default_memory_scheduler,
load_fallback_chain,
)
cfg = cfg if cfg is not None else _ensure_config()
if cfg.model_fallbacks:
load_fallback_chain(cfg.model_fallbacks)
model = chat_model if chat_model is not None else _ensure_chat_model()
if backend is None:
# Preserve the factory's pure path for callers that provide an
@@ -744,10 +758,8 @@ def _get_default_middleware(
if for_async_subagent
else MemoryObservationTarget.TURN_WORKER
)
# ``ConfigurableModelMiddleware`` is placed first so it wraps
# ``ModelFallbackMiddleware``: a configurable.model override sets the
# PRIMARY model only, leaving the fallback chain free to try its own
# alternatives instead of re-overriding every retry to the same model.
# ``ConfigurableModelMiddleware`` sits first so the snapshot-driven model
# override applies before any other middleware inspects the request.
memory_middleware = create_memory_middleware(
memory_dir,
workspace_dir=workspace_dir,
@@ -761,7 +773,7 @@ def _get_default_middleware(
memory_scheduler=memory_scheduler,
)
# Main-agent tool selection may use the auxiliary model; async sub-agents
# keep the main model (they do real work, not a one-off helper call).
# keep their own model (they do real work, not a one-off helper call).
# context_editing stays on the main model — its model only sizes the
# context-window trigger for the main agent's own history.
if for_async_subagent:
@@ -769,19 +781,14 @@ def _get_default_middleware(
elif chat_model is None:
tool_selector_model = _ensure_auxiliary_chat_model()
else:
aux_model = cfg.auxiliary_model or cfg.model
aux_provider = cfg.auxiliary_provider or cfg.provider
if (aux_model, aux_provider) == (cfg.model, cfg.provider):
tool_selector_model = model
else:
from .llm import get_chat_model
tool_selector_model = get_chat_model(model=aux_model, provider=aux_provider)
# Pure path (explicit model + config): no free-string auxiliary
# resolution — the threaded model stands in until the local snapshot
# entry point is wired (Task 7).
tool_selector_model = model
mw = [
ConfigurableModelMiddleware(),
ConfigurableModelMiddleware(role=snapshot_role),
create_message_budget_middleware(model, backend),
create_context_editing_middleware(model),
ModelFallbackMiddleware(),
ContextOverflowMapperMiddleware(),
ToolErrorHandlerMiddleware(),
*create_tool_selector_middleware(
-9
View File
@@ -794,12 +794,6 @@ def run_textual_interactive(
yield Static("", id="status")
def on_mount(self) -> None:
# Register fallback middleware UI callback so messages appear
# as SystemMessage widgets in the chat container.
from ..middleware.model_fallback import set_ui_emit
set_ui_emit(lambda text, style: self._append_system(text, style))
self._render_welcome()
self._render_status()
self.set_interval(1.0, self._render_status)
@@ -2970,9 +2964,6 @@ def run_textual_interactive(
def _do_exit(self) -> None:
"""Clean up channels, unregister callbacks, and exit."""
from ..middleware.model_fallback import set_ui_emit
set_ui_emit(None)
if self._channel_timer is not None:
self._channel_timer.stop()
self._channel_timer = None
@@ -6,7 +6,6 @@ from . import (
general,
mcp,
model,
model_fallback,
schedule,
session,
skills,
@@ -18,7 +17,6 @@ __all__ = [
"general",
"mcp",
"model",
"model_fallback",
"schedule",
"session",
"skills",
@@ -1,304 +0,0 @@
"""Slash command for managing the model fallback chain.
Provides ``/model-fallback`` (alias ``/fallback``) with subcommands to
add, remove, list, clear, save, and display help for fallback models.
"""
from __future__ import annotations
from typing import ClassVar
from ..base import Argument, Command, CommandContext, SubCommand
from ..manager import manager
class ModelFallbackCommand(Command):
"""Manage the model fallback chain."""
name = "/model-fallback"
alias: ClassVar[list[str]] = ["/fallback"]
description = "Manage fallback models (add/remove/list/clear)"
category = "Model"
arguments: ClassVar[list[Argument]] = [
Argument(
name="action",
type=str,
description="add|remove|list|clear|save|help",
required=False,
),
]
subcommands: ClassVar[list[SubCommand]] = [
SubCommand("list", "Display the current fallback chain"),
SubCommand("add", "Append a model to the fallback chain"),
SubCommand("remove", "Remove a model by position"),
SubCommand("clear", "Remove all fallback entries"),
SubCommand("save", "Persist the chain to config"),
SubCommand("help", "Show subcommand reference"),
]
async def execute(self, ctx: CommandContext, args: list[str]) -> None:
from ...llm.models import MODELS
from ...middleware.model_fallback import (
add_fallback,
clear_fallbacks,
get_fallback_chain,
remove_fallback_at,
serialize_fallback_chain,
)
save = "--save" in args
args = [a for a in args if a != "--save"]
if not args:
await self._show_list(ctx, get_fallback_chain())
return
action = args[0].lower()
if action == "list":
await self._show_list(ctx, get_fallback_chain())
elif action == "add":
if len(args) >= 2:
model_name = args[1]
provider = args[2] if len(args) > 2 else None
if provider is None:
if model_name in MODELS:
_, provider = MODELS[model_name]
else:
ctx.ui.append_system(
f"Unknown model '{model_name}'. Specify provider explicitly: "
f"/model-fallback add {model_name} <provider>",
style="red",
)
return
else:
picked = await self._pick_model(ctx)
if picked is None:
return
model_name, provider = picked
if add_fallback(model_name, provider):
ctx.ui.append_system(
f"Added {model_name} ({provider}) to fallback chain", style="green"
)
else:
ctx.ui.append_system(
f"{model_name} ({provider}) is already in the fallback chain",
style="yellow",
)
return
if save:
self._save_to_config(serialize_fallback_chain())
elif action == "remove":
chain = get_fallback_chain()
if not chain:
ctx.ui.append_system("Fallback chain is empty", style="yellow")
return
if len(args) >= 2:
arg = args[1]
try:
idx = int(arg) - 1
except ValueError:
ctx.ui.append_system(
f"Expected a position number (1-{len(chain)}), got '{arg}'. "
"Use /model-fallback list to see positions.",
style="red",
)
return
removed = remove_fallback_at(idx)
if removed is None:
ctx.ui.append_system(
f"Invalid position {arg}. "
f"Use a number between 1 and {len(chain)}.",
style="red",
)
return
model_name, provider = removed
else:
picked = await self._pick_fallback_to_remove(ctx, chain)
if picked is None:
return
model_name, provider = picked
live_chain = get_fallback_chain()
try:
idx = live_chain.index((model_name, provider))
except ValueError:
ctx.ui.append_system(
f"{model_name} ({provider}) is no longer in the fallback chain",
style="yellow",
)
return
remove_fallback_at(idx)
ctx.ui.append_system(
f"Removed {model_name} ({provider}) from fallback chain",
style="green",
)
if save:
self._save_to_config(serialize_fallback_chain())
elif action == "clear":
clear_fallbacks()
ctx.ui.append_system("Cleared all fallback models", style="green")
if save:
self._save_to_config("")
elif action == "save":
self._save_to_config(serialize_fallback_chain())
ctx.ui.append_system("Fallback chain saved to config", style="green")
elif action == "help":
self._show_help(ctx)
else:
self._show_help(ctx)
async def _pick_model(self, ctx: CommandContext) -> tuple[str, str] | None:
"""Open the interactive model picker to select a fallback model.
Falls back to a usage hint when the UI does not support interactive
widgets (CLI mode without a model argument).
Args:
ctx: Current command context.
Returns:
``(model_name, provider)`` tuple, or ``None`` if cancelled.
"""
if not ctx.ui.supports_interactive:
ctx.ui.append_system(
"Usage: /model-fallback add <model> [provider]", style="yellow"
)
return None
from ...EvoScientist import _ensure_config
from ...llm.models import list_model_picker_entries
cfg = _ensure_config()
entries = await list_model_picker_entries(
getattr(cfg, "ollama_base_url", None),
include_custom_ollama=True,
)
result = await ctx.ui.wait_for_model_pick(
entries,
current_model=cfg.model,
current_provider=cfg.provider,
)
if result is None:
return None
name, provider = result
if provider == "ollama" and name in (
"Custom Ollama model...",
"__custom_ollama__",
):
return None
return name, provider
async def _pick_fallback_to_remove(
self, ctx: CommandContext, chain: list[tuple[str, str]]
) -> tuple[str, str] | None:
"""Open the model picker populated with the current fallback chain.
Falls back to a usage hint in CLI mode.
Args:
ctx: Current command context.
chain: The current fallback chain to choose from.
Returns:
``(model_name, provider)`` tuple, or ``None`` if cancelled.
"""
if not ctx.ui.supports_interactive:
ctx.ui.append_system(
"Usage: /model-fallback remove <position> "
"(use /model-fallback list to see positions)",
style="yellow",
)
return None
entries = [(m, m, p) for m, p in chain]
result = await ctx.ui.wait_for_model_pick(
entries, current_model=None, current_provider=None
)
if result is None:
return None
return result
def _show_help(self, ctx: CommandContext) -> None:
"""Render the subcommand reference table.
Args:
ctx: Current command context.
"""
from rich.text import Text
text = Text("/model-fallback subcommands:\n", style="bold")
for cmd, desc in (
(
"add [model] [provider]",
"Add a fallback model (opens picker if omitted)",
),
(
"remove [position]",
"Remove a fallback by position (opens picker in TUI)",
),
("list", "Show the current fallback chain"),
("clear", "Remove all fallback models"),
("save", "Save current fallback chain to config file"),
("help", "Show this help message"),
):
text.append(f" {cmd:<26}", style="cyan")
text.append(f"{desc}\n", style="dim")
text.append(
"\nAdd --save to add/remove/clear to persist the change immediately.\n",
style="dim",
)
ctx.ui.mount_renderable(text)
async def _show_list(
self, ctx: CommandContext, chain: list[tuple[str, str]]
) -> None:
"""Display the current fallback chain as a numbered list.
Args:
ctx: Current command context.
chain: The fallback chain to display.
"""
if not chain:
ctx.ui.append_system("No fallback models configured", style="dim")
ctx.ui.append_system(
"Use /model-fallback add <model> [provider] to add one",
style="dim",
)
return
from rich.text import Text
text = Text("Fallback chain:\n", style="bold")
for idx, (model, provider) in enumerate(chain, 1):
text.append(f" {idx}. ", style="dim")
text.append(model, style="cyan")
text.append(f" ({provider})\n", style="dim")
ctx.ui.mount_renderable(text)
def _save_to_config(self, value: str) -> None:
"""Persist the fallback chain string to the config file.
Args:
value: Serialized chain (``"model:provider,..."``).
"""
from ...config.settings import set_config_value
set_config_value("model_fallbacks", value)
manager.register(ModelFallbackCommand())
+25 -50
View File
@@ -873,26 +873,23 @@ def _patch_deepseek_reasoning_passback(model: Any) -> None:
# ---------------------------------------------------------------------------
# Patch: forward CLI's live (model, model_provider) into deepagents'
# start_async_task / update_async_task tool calls so the deployed graph
# (running in a separate ``langgraph dev`` subprocess) re-resolves the
# chat model per run.
# Patch: forward workspace scope and usage correlation context into
# deepagents' start_async_task / update_async_task tool calls so the deployed
# graph (running in a separate ``langgraph dev`` subprocess) inherits the
# parent run's scoping and accounting metadata.
#
# Without this, async sub-agents stay on the model their graph was compiled
# with at langgraph dev boot — `/model` switches in the CLI never reach
# them because they live in another process.
# Model configuration is deliberately NOT forwarded: the deployed graph
# resolves its chat model per run from ``configurable.runtime_snapshot_id``
# via ``ConfigurableModelMiddleware`` (design doc 8.2/8.3). Injecting
# ``model``/``model_provider`` here would be rejected with
# ``MODEL_CONFIG_OUTSIDE_SNAPSHOT``.
#
# Mechanism: wrap deepagents' ``_build_start_tool`` and ``_build_update_tool``
# factories. Each wrapped factory calls the original with a proxied client
# cache that intercepts ``runs.create(...)`` calls only and injects
# ``config={"configurable": {"model": <cfg.model>, "model_provider": <cfg.provider>}}``.
# All other client methods (``threads.create``, ``runs.get``, ``runs.cancel``,
# ``runs.join_stream``) pass through unchanged. The deployed graph picks up
# ``configurable.model`` via ``ConfigurableModelMiddleware``.
#
# Reads ``_ensure_config()`` at tool-call time (not patch time) so a
# ``/model`` switch in the CLI is reflected on the very next async tool
# call without an agent rebuild.
# cache that intercepts ``runs.create(...)`` calls only and merges the
# inherited scope/usage context into ``config``/``metadata``. All other
# client methods (``threads.create``, ``runs.get``, ``runs.cancel``,
# ``runs.join_stream``) pass through unchanged.
#
# Upstream PR opportunity: passing ``config`` through ``client.runs.create``
# is generic functionality; worth contributing back to ``langchain-ai/deepagents``
@@ -901,38 +898,11 @@ def _patch_deepseek_reasoning_passback(model: Any) -> None:
_model_passthrough_patched = False
def _read_cfg_configurable() -> dict[str, str]:
"""Read live ``(model, provider)`` from EvoScientist config.
Returns a dict suitable for inserting under
``RunnableConfig.configurable``. Empty dict on any failure (so the
patch degrades to a no-op rather than breaking async tool calls).
"""
try:
from EvoScientist.EvoScientist import _ensure_config
cfg = _ensure_config()
except Exception:
return {}
out: dict[str, str] = {}
model = getattr(cfg, "model", None)
provider = getattr(cfg, "provider", None)
if isinstance(model, str) and model:
out["model"] = model
if isinstance(provider, str) and provider:
out["model_provider"] = provider
return out
def _merge_runs_config_kwargs(kwargs: dict) -> dict:
"""Merge model and usage correlation context into ``runs.create``.
"""Merge scope and usage correlation context into ``runs.create``.
Preserves any caller-supplied ``config.configurable`` keys. EvoScientist's
keys take precedence on conflict (callers shouldn't be passing model
overrides — the CLI is the source of truth).
Preserves any caller-supplied ``config.configurable`` keys.
"""
overrides = _read_cfg_configurable()
usage_enabled = os.getenv("EVOSCIENTIST_USAGE_TRACKING", "").strip().lower() in {
"1",
"true",
@@ -985,7 +955,9 @@ def _merge_runs_config_kwargs(kwargs: dict) -> dict:
str(inherited_scope["workspace_scope_id"]),
owner_type="derived_run",
resource_id=target_thread_id,
parent_owner_id=str(inherited_scope["workspace_scope_owner_id"]),
parent_owner_id=str(
inherited_scope["workspace_scope_owner_id"]
),
state="active",
)
except ScopeConflictError:
@@ -1001,7 +973,10 @@ def _merge_runs_config_kwargs(kwargs: dict) -> dict:
# The required-mode backend factory rejects missing/invalid
# ownership at execution time. Optional mode keeps legacy
# integrations working when they cannot register a child.
if os.getenv("EVOSCIENTIST_WORKSPACE_ISOLATION", "optional").lower() == "required":
if (
os.getenv("EVOSCIENTIST_WORKSPACE_ISOLATION", "optional").lower()
== "required"
):
raise
existing = kwargs.get("config")
@@ -1010,9 +985,9 @@ def _merge_runs_config_kwargs(kwargs: dict) -> dict:
existing_configurable = existing.get("configurable")
if not isinstance(existing_configurable, dict):
existing_configurable = {}
merged_configurable = {**existing_configurable, **overrides, **inherited_scope}
merged_configurable = {**existing_configurable, **inherited_scope}
kwargs = dict(kwargs)
if overrides or "config" in kwargs:
if merged_configurable or "config" in kwargs:
kwargs["config"] = {**existing, "configurable": merged_configurable}
if not usage_enabled:
@@ -1121,7 +1096,7 @@ class _ClientCacheProxy:
def _patch_deepagents_model_passthrough() -> None:
"""Wrap deepagents' async-launch tool factories to inject CLI model.
"""Wrap deepagents' async-launch tool factories to inherit run context.
Idempotent: re-invocation is a no-op once the patch is active. Safe to
call from ``_maybe_swap_async_subagents`` on every CLI startup; both
-3
View File
@@ -32,7 +32,6 @@ from .message_budget import (
count_message_text_tokens,
create_message_budget_middleware,
)
from .model_fallback import ModelFallbackMiddleware, load_fallback_chain
from .runtime_context import RuntimeContextMiddleware, create_runtime_context_middleware
from .scheduler import (
SchedulerMiddleware,
@@ -52,7 +51,6 @@ __all__ = [
"EvoMemoryLifecycleMiddleware",
"EvoMemoryMiddleware",
"MessageReservePolicy",
"ModelFallbackMiddleware",
"Question",
"RuntimeContextMiddleware",
"SchedulerMiddleware",
@@ -69,5 +67,4 @@ __all__ = [
"create_tool_selector_middleware",
"default_memory_scheduler",
"disable_thinking",
"load_fallback_chain",
]
+158 -169
View File
@@ -1,23 +1,22 @@
"""Middleware that resolves the chat model from RunnableConfig.configurable per call.
"""Middleware that resolves the per-call chat model from a run snapshot.
The deployed async sub-agents run in a separate ``langgraph dev`` subprocess
and have their model frozen into the compiled graph at subprocess boot time
(see ``EvoScientist/subagents/_factory.py``). When the user runs ``/model``
in the CLI, only the CLI process's model state changes — the subprocess
graph still uses the boot-time model.
Design doc 8.3: the middleware no longer reads ``config.yaml``, global
aliases, or ``model``/``model_provider`` overrides. The only model input a
run may carry is ``configurable["runtime_snapshot_id"]``; the middleware
loads the frozen snapshot (deployment/thread binding verified), resolves
the middleware's role through the section 6.1 role mapping, resolves the
credential against the frozen ``credential_revision``, and constructs the
model via ``build_chat_model`` with both safe HTTP clients.
This middleware closes that gap by reading ``model`` / ``model_provider``
from ``RunnableConfig.configurable`` on every model call. The CLI's patched
``start_async_task`` / ``update_async_task`` (see ``llm/patches.py``) injects
those fields into ``client.runs.create(config=...)``; the deployed graph
hits this middleware and re-resolves the chat model fresh.
``configurable`` carrying ``model``, ``model_provider``, or any other
out-of-snapshot model parameter is rejected with
``MODEL_CONFIG_OUTSIDE_SNAPSHOT`` (422 semantics, section 8.2) — run
creation must never mix snapshot and non-snapshot model configuration.
When ``configurable.model`` is absent, the middleware is a pass-through —
safe to install on the CLI's in-process agent too.
The middleware mirrors the pattern used by ``ModelFallbackMiddleware``:
``request.override(model=new_model)`` does not break tool binding, because
the downstream model-invocation node re-binds tools per request.
When no ``runtime_snapshot_id`` is present the middleware is a pass-through:
local entry points (CLI/channels/scheduler) are wired to create local
snapshots separately, and their compile-time model binding applies until
then.
**Reading the config**: ``Runtime`` (per its own docstring) does NOT include
``config``. The official path to reach ``RunnableConfig`` from inside any
@@ -32,8 +31,8 @@ from __future__ import annotations
import asyncio
import logging
import threading
from collections.abc import Awaitable, Callable
from typing import Any
from collections.abc import Awaitable, Callable, Mapping
from typing import Any, get_args
from langchain.agents.middleware.types import (
AgentMiddleware,
@@ -41,177 +40,184 @@ from langchain.agents.middleware.types import (
ModelResponse,
)
from ..model_registry.errors import (
MODEL_CONFIG_OUTSIDE_SNAPSHOT,
ModelRegistryError,
)
from ..model_registry.runtime import SnapshotRuntime, get_snapshot_runtime
from ..model_registry.schemas import ModelRole
from ..model_registry.snapshots import RuntimeSnapshot
logger = logging.getLogger(__name__)
_MODEL_ROLES = get_args(ModelRole)
def _read_model_override() -> tuple[str | None, str | None]:
"""Pull ``(model, model_provider)`` from the active ``RunnableConfig``.
# The only model-related keys a run's ``configurable`` may never carry: the
# snapshot is the single model configuration entry point (section 8.2).
_OUTSIDE_SNAPSHOT_MODEL_KEYS = ("model", "model_provider")
Reads via ``langgraph.config.get_config()`` (the documented entry point
for accessing the per-run ``RunnableConfig`` from inside any runnable
context — middleware, node, tool). Returns ``(None, None)`` when the
config has no ``configurable.model`` override or when called outside a
runnable context.
"""
def _current_configurable() -> Mapping[str, Any]:
"""Return the active run's ``configurable`` mapping (empty outside runs)."""
try:
from langgraph.config import get_config
cfg = get_config()
config = get_config()
except Exception:
# Outside a runnable context (most common in tests) or
# langgraph not importable — nothing to override.
return None, None
if not isinstance(cfg, dict):
return None, None
configurable = cfg.get("configurable") or {}
if not isinstance(configurable, dict):
return None, None
model = configurable.get("model")
provider = configurable.get("model_provider")
# Outside a runnable context (most common in tests) or langgraph not
# importable — treat as "no per-run configuration".
return {}
if not isinstance(config, Mapping):
return {}
configurable = config.get("configurable")
return configurable if isinstance(configurable, Mapping) else {}
def check_no_outside_snapshot_model_config(configurable: Mapping[str, Any]) -> None:
"""Reject any model configuration carried outside the run snapshot."""
for key in _OUTSIDE_SNAPSHOT_MODEL_KEYS:
if configurable.get(key) is not None:
raise ModelRegistryError(
MODEL_CONFIG_OUTSIDE_SNAPSHOT,
"Model configuration must come from the run snapshot only; "
f"configurable carries {key!r}.",
details=[{"path": key, "code": MODEL_CONFIG_OUTSIDE_SNAPSHOT}],
)
def read_snapshot_binding(
configurable: Mapping[str, Any],
) -> tuple[str, str | None, str | None] | None:
"""Extract ``(snapshot_id, deployment_id, thread_id)`` from configurable.
Returns ``None`` when the run carries no ``runtime_snapshot_id``. The
deployment ID falls back to ``None`` (caller substitutes the platform's
local deployment ID); a missing thread ID fails closed later because it
can never match the snapshot's binding.
"""
snapshot_id = configurable.get("runtime_snapshot_id")
if not isinstance(snapshot_id, str) or not snapshot_id:
return None
deployment_id = configurable.get("workspace_deployment_id")
thread_id = configurable.get("thread_id")
return (
model if isinstance(model, str) and model else None,
provider if isinstance(provider, str) and provider else None,
snapshot_id,
deployment_id if isinstance(deployment_id, str) and deployment_id else None,
thread_id if isinstance(thread_id, str) else None,
)
def _read_runtime_snapshot_id() -> str | None:
"""Return the opaque server-side runtime snapshot ID for this run."""
try:
from langgraph.config import get_config
cfg = get_config()
except Exception:
return None
if not isinstance(cfg, dict):
return None
configurable = cfg.get("configurable") or {}
if not isinstance(configurable, dict):
return None
snapshot_id = configurable.get("runtime_snapshot_id")
return snapshot_id if isinstance(snapshot_id, str) and snapshot_id else None
class ConfigurableModelMiddleware(AgentMiddleware):
"""Re-resolve the chat model from RunnableConfig.configurable on every call.
"""Re-resolve the chat model from the run snapshot on every call.
Reads ``model`` and ``model_provider`` from the active ``RunnableConfig``
via ``langgraph.config.get_config()`` — the documented entry point for
accessing per-run config from any runnable context (middleware, node, tool).
When the override is present, calls
``EvoScientist.llm.get_chat_model(model=..., provider=...)`` and replaces
``request.model`` via ``request.override``. When absent, the middleware
passes through unchanged.
``role`` selects which frozen configuration of the snapshot feeds this
agent's model calls: ``primary`` for the main agent and working
sub-agents, ``auxiliary`` for unattended helper agents (e.g. the
scheduler). The section 6.1 mapping (``auxiliary``/``summary``/
``tool_selector`` → ``snapshot.auxiliary ?? snapshot.primary``) is
applied by the snapshot layer.
Note: ``Runtime`` (per its own docstring) does NOT include ``config`` as a
field — an earlier version of this middleware tried to read
``request.runtime.config`` and silently no-op'd because that attribute does
not exist. Stick with ``get_config()``.
A per-instance cache keyed by ``(model, provider)`` avoids rebuilding
identical models within a turn. The cache is a plain dict guarded by a
``threading.Lock`` because middleware instances are shared across
A per-instance cache keyed by snapshot ID avoids rebuilding the model on
every call within a run; snapshots are immutable once created, so the
cached instance stays valid for the run's lifetime. The cache is guarded
by a ``threading.Lock`` because middleware instances are shared across
concurrent requests in long-lived deployments (e.g. ``langgraph dev``
workers).
"""
name = "configurable_model"
def __init__(self) -> None:
def __init__(
self, role: ModelRole = "primary", runtime: SnapshotRuntime | None = None
) -> None:
super().__init__()
self._cache: dict[tuple[str, str | None, str | None], Any] = {}
if role not in _MODEL_ROLES:
raise ValueError(
f"Unknown model role {role!r}; expected one of {list(_MODEL_ROLES)}."
)
self._role = role
self._runtime = runtime
self._cache: dict[str, Any] = {}
self._lock = threading.Lock()
# Track the last (model, provider) pair we INFO-logged so we only
# surface a banner on transition. Without this, every LLM call in a
# long async run would emit an identical INFO line.
self._last_logged_key: tuple[str, str | None] | None = None
# Track the last snapshot we INFO-logged so we only surface a banner
# on transition. Without this, every LLM call in a long run would
# emit an identical INFO line.
self._last_logged_snapshot_id: str | None = None
def _log_override(self, model_name: str, provider: str | None) -> None:
"""INFO on transition; DEBUG on subsequent calls with same key."""
key = (model_name, provider)
def _snapshot_runtime(self) -> SnapshotRuntime:
return self._runtime if self._runtime is not None else get_snapshot_runtime()
def _log_override(self, snapshot: RuntimeSnapshot) -> None:
"""INFO on transition; DEBUG on subsequent calls of the same snapshot."""
from ..model_registry.snapshots import config_for_role
config = config_for_role(snapshot, self._role)
with self._lock:
transitioned = key != self._last_logged_key
transitioned = snapshot.snapshot_id != self._last_logged_snapshot_id
if transitioned:
self._last_logged_key = key
self._last_logged_snapshot_id = snapshot.snapshot_id
message_args = (
self._role,
config.model_ref.provider_id,
config.model_ref.model_key,
snapshot.snapshot_id,
)
if transitioned:
logger.info(
"ConfigurableModelMiddleware: overriding model to %s (%s)",
model_name,
provider,
"ConfigurableModelMiddleware: role %s bound to %s/%s (snapshot %s)",
*message_args,
)
else:
logger.debug(
"ConfigurableModelMiddleware: reusing override model=%s provider=%s",
model_name,
provider,
"ConfigurableModelMiddleware: role %s reusing %s/%s (snapshot %s)",
*message_args,
)
def _resolve(
self, model: str, provider: str | None, runtime_snapshot_id: str | None
) -> Any:
"""Return a cached or freshly-built chat model for ``(model, provider)``."""
if runtime_snapshot_id is not None:
from ..llm.models import get_profile_chat_model
from ..llm.runtime_snapshots import get_run_runtime_snapshot
def _load_snapshot(self, snapshot_id: str) -> RuntimeSnapshot:
"""Load the snapshot, verifying its deployment/thread binding."""
configurable = _current_configurable()
binding = read_snapshot_binding(configurable)
assert binding is not None # guarded by the caller
_, deployment_id, thread_id = binding
runtime = self._snapshot_runtime()
return runtime.get_snapshot(
snapshot_id,
deployment_id=deployment_id or runtime.local_deployment_id,
thread_id=thread_id or "",
)
snapshot = get_run_runtime_snapshot(runtime_snapshot_id)
if snapshot is None:
raise ValueError(
"RUN_RUNTIME_SNAPSHOT_UNAVAILABLE: the run configuration snapshot "
"expired or is unavailable. Start the message again."
)
if snapshot.model.id != model or snapshot.profile.id != provider:
raise ValueError("RUN_RUNTIME_SNAPSHOT_MISMATCH: run configuration is invalid.")
key = ("snapshot", runtime_snapshot_id, snapshot.profile_revision)
with self._lock:
cached = self._cache.get(key)
if cached is not None:
return cached
new_model = get_profile_chat_model(snapshot.profile, snapshot.model)
with self._lock:
self._cache[key] = new_model
return new_model
from ..llm.models import get_model_runtime_revision
revision = get_model_runtime_revision(provider)
key = (model, provider, revision)
def _resolve(self, snapshot_id: str) -> Any:
"""Return a cached or freshly-built chat model for the snapshot."""
with self._lock:
cached = self._cache.get(key)
cached = self._cache.get(snapshot_id)
if cached is not None:
return cached
# Build outside the lock (network/SDK init can be slow); two
# concurrent first-time misses for the same key may build twice but
# the second result simply overwrites the first — both are equivalent.
from ..llm import get_chat_model
new_model = get_chat_model(model=model, provider=provider)
snapshot = self._load_snapshot(snapshot_id)
# Build outside the lock (SDK init can be slow); two concurrent
# first-time misses for the same snapshot may build twice but the
# second result simply overwrites the first — both are equivalent.
new_model = self._snapshot_runtime().build_role_model(snapshot, self._role)
with self._lock:
self._cache[key] = new_model
self._cache[snapshot_id] = new_model
self._log_override(snapshot)
return new_model
def _snapshot_id_for_request(self) -> str | None:
"""Return the run's snapshot ID after the outside-config check."""
configurable = _current_configurable()
check_no_outside_snapshot_model_config(configurable)
binding = read_snapshot_binding(configurable)
return binding[0] if binding is not None else None
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelResponse:
model_name, provider = _read_model_override()
if model_name is None:
snapshot_id = self._snapshot_id_for_request()
if snapshot_id is None:
return handler(request)
runtime_snapshot_id = _read_runtime_snapshot_id()
try:
new_model = self._resolve(model_name, provider, runtime_snapshot_id)
except Exception:
if runtime_snapshot_id is not None:
raise
logger.warning(
"ConfigurableModelMiddleware failed to resolve model=%r "
"provider=%r; falling back to compile-time model",
model_name,
provider,
exc_info=True,
)
return handler(request)
self._log_override(model_name, provider)
new_model = self._resolve(snapshot_id)
return handler(request.override(model=new_model))
async def awrap_model_call(
@@ -219,30 +225,13 @@ class ConfigurableModelMiddleware(AgentMiddleware):
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse:
model_name, provider = _read_model_override()
if model_name is None:
snapshot_id = self._snapshot_id_for_request()
if snapshot_id is None:
return await handler(request)
runtime_snapshot_id = _read_runtime_snapshot_id()
try:
# Offload first-call SDK init off the event loop. ``_resolve`` calls
# ``get_chat_model`` on a cache miss, which can spend hundreds of ms
# building HTTP clients. Doing this synchronously inside an
# ``async def`` would block every other coroutine on the same
# langgraph dev event loop. Cache hits are still fast (a dict
# lookup); the thread-pool overhead is irrelevant once warm.
new_model = await asyncio.to_thread(
self._resolve, model_name, provider, runtime_snapshot_id
)
except Exception:
if runtime_snapshot_id is not None:
raise
logger.warning(
"ConfigurableModelMiddleware failed to resolve model=%r "
"provider=%r; falling back to compile-time model",
model_name,
provider,
exc_info=True,
)
return await handler(request)
self._log_override(model_name, provider)
# Offload first-call SDK init off the event loop: ``_resolve`` reads
# SQLite and can spend hundreds of ms building HTTP clients on a
# cache miss, which would block every other coroutine on the same
# langgraph dev event loop. Cache hits are still fast; the
# thread-pool overhead is irrelevant once warm.
new_model = await asyncio.to_thread(self._resolve, snapshot_id)
return await handler(request.override(model=new_model))
+175 -84
View File
@@ -1,10 +1,42 @@
"""Message-only context budgeting and automatic conversation compaction.
The middleware deliberately does not attempt to count the whole provider
request. System instructions, memories, tools, and attachments instead consume
fixed conservative reserves. The measured value is only textual conversation
messages and tool-result text, which is the part compaction can actually
reduce.
request. System instructions, memories, tools, and attachments instead
consume the fixed reserves frozen into the run snapshot. The measured value
is only textual conversation messages and tool-result text, which is the
part compaction can actually reduce.
Snapshot mode (design doc 6.5, 8.3): when the run carries a
``runtime_snapshot_id``, the input limit and the three fixed reserves come
from the frozen ``ResolvedModelConfig.budget`` — never from the legacy
``llm/runtime_snapshots.py`` payloads, and never from an implicit 32K
default. The per-call message budget is recomputed on every invocation
from those frozen reserves with the current ``has_tools``/``has_attachments``
mode; it is never carried over from a previous call:
message_budget = resolved_input_limit
- fixed_system_reserve_tokens
- (has_tools ? fixed_tools_reserve_tokens : 0)
- (has_attachments ? fixed_attachments_reserve_tokens : 0)
``has_tools`` follows the conservative rule: it is decided once at agent
construction (the agent has tool capability and a configured toolset), not
by counting the tools bound to a single call, so it does not depend on
tool-selector middleware ordering. ``has_attachments`` is true only when
the current messages carry file/image/other attachment blocks. The frozen
base ``budget.message_budget`` (no tools, no attachments) is never reused
directly.
Because ``count_message_text_tokens`` is a conservative character estimate,
the effective hard budget is ``message_budget x 0.90`` and the soft trigger
scales proportionally (``x 0.70``), replacing the removed ``safety_reserve``
(section 6.5). Budget satisfiability against ``min_effective_input_tokens``
is guaranteed at snapshot creation (run-creation stage, all four
tool/attachment modes), so the middleware does not re-check it per call.
Interim local mode: runs without a snapshot ID (CLI local path before the
local snapshot entry point is wired) keep the legacy behavior — limits read
from the compile-time model's ``profile`` with the default reserve policy.
"""
from __future__ import annotations
@@ -13,14 +45,26 @@ import threading
from collections.abc import Iterable, Mapping
from contextvars import ContextVar
from dataclasses import dataclass
from typing import Any
from typing import TYPE_CHECKING, Any
from langchain_core.messages import AnyMessage, SystemMessage
from .configurable_model import _current_configurable, read_snapshot_binding
if TYPE_CHECKING:
from ..model_registry.runtime import SnapshotRuntime
from ..model_registry.snapshots import RuntimeSnapshot
# Character-estimation budget scaling (section 6.5). The conservative char
# counter replaces the removed ``safety_reserve`` deduction.
_ESTIMATE_HARD_FRACTION = 0.90
_SOFT_FRACTION = 0.70
_KEEP_FRACTION = 0.35
@dataclass(frozen=True)
class MessageReservePolicy:
"""Static non-message reserves used without full-request token counting."""
"""Static non-message reserves for the interim no-snapshot local path."""
system_tokens: int = 4_096
memory_tokens: int = 8_192
@@ -55,9 +99,7 @@ class MessageReservePolicy:
return max(
1_024,
input_limit
- self.fixed_tokens(
has_tools=has_tools, has_attachments=has_attachments
)
- self.fixed_tokens(has_tools=has_tools, has_attachments=has_attachments)
- safety,
)
@@ -80,7 +122,6 @@ class MessageBudget:
hard_tokens: int
soft_tokens: int
keep_tokens: int
min_effective_input_tokens: int
has_tools: bool
has_attachments: bool
@@ -135,25 +176,56 @@ def _has_attachments(messages: Iterable[AnyMessage]) -> bool:
return False
def _current_configurable() -> Mapping[str, Any]:
try:
from langgraph.config import get_config
def _snapshot_message_budget(
snapshot: RuntimeSnapshot, *, has_tools: bool, has_attachments: bool
) -> MessageBudget:
"""Recompute the section 6.5 per-call budget from the frozen reserves."""
from ..model_registry.snapshots import config_for_role
config = get_config()
except Exception:
return {}
if not isinstance(config, Mapping):
return {}
configurable = config.get("configurable")
return configurable if isinstance(configurable, Mapping) else {}
budget = config_for_role(snapshot, "primary").budget
reserves = budget.fixed_reserves
message_budget = (
budget.resolved_input_limit
- reserves.fixed_system_reserve_tokens
- (reserves.fixed_tools_reserve_tokens if has_tools else 0)
- (reserves.fixed_attachments_reserve_tokens if has_attachments else 0)
)
hard = max(1, int(message_budget * _ESTIMATE_HARD_FRACTION))
return MessageBudget(
input_limit=budget.resolved_input_limit,
hard_tokens=hard,
soft_tokens=max(1, int(hard * _SOFT_FRACTION)),
keep_tokens=max(1, int(hard * _KEEP_FRACTION)),
has_tools=has_tools,
has_attachments=has_attachments,
)
class MessageBudgetMiddleware:
"""Factory namespace kept separate from DeepAgents' concrete middleware."""
@staticmethod
def create(model: Any, backend: Any, *, policy: MessageReservePolicy | None = None):
"""Create the runtime-aware DeepAgents summarization middleware."""
def create(
model: Any,
backend: Any,
*,
policy: MessageReservePolicy | None = None,
has_tools: bool = True,
runtime: SnapshotRuntime | None = None,
):
"""Create the runtime-aware DeepAgents summarization middleware.
Args:
model: Compile-time fallback model, used for the summarizer and
for budget limits when the run carries no snapshot.
backend: Graph backend forwarded to the summarization middleware.
policy: Reserve policy for the interim no-snapshot local path.
has_tools: Conservative section 6.5 tool-mode flag: true when the
agent has tool capability and a configured toolset. Decided
once here, never from a single request's bound tool count.
runtime: Snapshot runtime override; defaults to the shared
process runtime (tests inject an isolated one).
"""
from deepagents.middleware.summarization import SummarizationMiddleware
class _RuntimeMessageBudgetMiddleware(SummarizationMiddleware):
@@ -161,9 +233,11 @@ class MessageBudgetMiddleware:
def __init__(self) -> None:
self._fallback_model = model
self._model_cache: dict[tuple[str, str | None], Any] = {}
self._snapshot_model_cache: dict[str, Any] = {}
self._model_cache_lock = threading.RLock()
self._policy = policy or MessageReservePolicy()
self._has_tools = has_tools
self._runtime = runtime
# Triggering and cutoff are overridden below. The base class is
# still used for safe AI/tool-pair handling, offloading, and
# persisted summarization events.
@@ -182,90 +256,100 @@ class MessageBudgetMiddleware:
},
)
def _snapshot_runtime(self) -> SnapshotRuntime:
if self._runtime is not None:
return self._runtime
from ..model_registry.runtime import get_snapshot_runtime
return get_snapshot_runtime()
def _snapshot(self) -> RuntimeSnapshot | None:
"""Load the run's snapshot, verifying its binding."""
binding = read_snapshot_binding(_current_configurable())
if binding is None:
return None
snapshot_id, deployment_id, thread_id = binding
runtime = self._snapshot_runtime()
return runtime.get_snapshot(
snapshot_id,
deployment_id=deployment_id or runtime.local_deployment_id,
thread_id=thread_id or "",
)
@property
def model(self) -> Any: # type: ignore[override]
configurable = _current_configurable()
snapshot_id = configurable.get("runtime_snapshot_id")
if isinstance(snapshot_id, str) and snapshot_id:
from ..llm.runtime_snapshots import get_snapshot_chat_model
return get_snapshot_chat_model(snapshot_id)
model_name = configurable.get("model")
provider = configurable.get("model_provider")
if not isinstance(model_name, str) or not model_name:
"""Summarizer model: the snapshot's ``summary`` role mapping."""
snapshot = self._snapshot()
if snapshot is None:
return self._fallback_model
provider_name = provider if isinstance(provider, str) and provider else None
key = (model_name, provider_name)
with self._model_cache_lock:
cached = self._model_cache.get(key)
cached = self._snapshot_model_cache.get(snapshot.snapshot_id)
if cached is not None:
return cached
from ..llm.models import get_chat_model
resolved = get_chat_model(model=model_name, provider=provider_name)
resolved = self._snapshot_runtime().build_role_model(
snapshot, "summary"
)
with self._model_cache_lock:
self._model_cache[key] = resolved
self._snapshot_model_cache[snapshot.snapshot_id] = resolved
return resolved
def _input_limit(self) -> int:
profile = getattr(self.model, "profile", None)
def _interim_budget(self, *, has_attachments: bool) -> MessageBudget:
"""Legacy local path: limits from the compile-time profile."""
profile = getattr(self._fallback_model, "profile", None)
input_limit = 32_768
minimum_effective = 1_024
if isinstance(profile, Mapping):
value = profile.get("max_input_tokens")
if isinstance(value, int) and not isinstance(value, bool) and value > 0:
return value
# Existing static providers that lack profile data use a
# conservative lower fallback. Custom profiles always supply
# `max_input_tokens` through the frozen runtime options.
return 32_768
def _minimum_effective_input(self) -> int:
profile = getattr(self.model, "profile", None)
if isinstance(profile, Mapping):
if (
isinstance(value, int)
and not isinstance(value, bool)
and value > 0
):
input_limit = value
value = profile.get("min_effective_input_tokens")
if isinstance(value, int) and not isinstance(value, bool) and value > 0:
return value
return 1_024
def _budget_for_request(self, request: Any) -> MessageBudget:
input_limit = self._input_limit()
has_tools = bool(getattr(request, "tools", None))
has_attachments = _has_attachments(getattr(request, "messages", []))
if (
isinstance(value, int)
and not isinstance(value, bool)
and value > 0
):
minimum_effective = value
mode = {
"has_tools": has_tools,
"has_tools": self._has_tools,
"has_attachments": has_attachments,
}
budget = MessageBudget(
input_limit=input_limit,
hard_tokens=self._policy.hard_budget(input_limit, **mode),
soft_tokens=self._policy.soft_budget(input_limit, **mode),
keep_tokens=self._policy.keep_budget(input_limit, **mode),
min_effective_input_tokens=self._minimum_effective_input(),
has_tools=has_tools,
has_attachments=has_attachments,
)
if budget.hard_tokens < budget.min_effective_input_tokens:
hard = self._policy.hard_budget(input_limit, **mode)
if hard < minimum_effective:
raise ContextBudgetUnsatisfiableError(
"CONTEXT_BUDGET_UNSATISFIABLE: configured input limit "
f"{budget.input_limit:,} leaves only {budget.hard_tokens:,} "
f"{input_limit:,} leaves only {hard:,} "
"tokens after fixed reserves; increase the model window, reduce "
"the output budget, or disable tools/attachments."
)
return budget
return MessageBudget(
input_limit=input_limit,
hard_tokens=hard,
soft_tokens=self._policy.soft_budget(input_limit, **mode),
keep_tokens=self._policy.keep_budget(input_limit, **mode),
has_tools=self._has_tools,
has_attachments=has_attachments,
)
def _budget_for_request(self, request: Any) -> MessageBudget:
has_attachments = _has_attachments(getattr(request, "messages", []))
snapshot = self._snapshot()
if snapshot is not None:
return _snapshot_message_budget(
snapshot,
has_tools=self._has_tools,
has_attachments=has_attachments,
)
return self._interim_budget(has_attachments=has_attachments)
def _active_budget(self) -> MessageBudget:
active = _ACTIVE_BUDGET.get()
if active is not None:
return active
input_limit = self._input_limit()
return MessageBudget(
input_limit=input_limit,
hard_tokens=self._policy.hard_budget(input_limit),
soft_tokens=self._policy.soft_budget(input_limit),
keep_tokens=self._policy.keep_budget(input_limit),
min_effective_input_tokens=self._minimum_effective_input(),
has_tools=False,
has_attachments=False,
)
return self._interim_budget(has_attachments=False)
def wrap_model_call(self, request: Any, handler: Any) -> Any:
token = _ACTIVE_BUDGET.set(self._budget_for_request(request))
@@ -321,7 +405,14 @@ class MessageBudgetMiddleware:
def create_message_budget_middleware(
model: Any, backend: Any, *, policy: MessageReservePolicy | None = None
model: Any,
backend: Any,
*,
policy: MessageReservePolicy | None = None,
has_tools: bool = True,
runtime: SnapshotRuntime | None = None,
):
"""Construct automatic compaction middleware for a graph backend."""
return MessageBudgetMiddleware.create(model, backend, policy=policy)
return MessageBudgetMiddleware.create(
model, backend, policy=policy, has_tools=has_tools, runtime=runtime
)
-378
View File
@@ -1,378 +0,0 @@
"""Middleware that implements model fallback on LLM call failures.
Uses LangChain's AgentMiddleware to intercept model calls. When the primary
model raises an exception, the middleware walks the configured fallback chain,
trying each alternative model in order. Every fallback attempt and its
outcome is surfaced to the user via the registered UI callback.
Errors that indicate a client-side bug (malformed request / HTTP 400) or a
context-length breach are not eligible for fallback and are re-raised
immediately so the correct handler (user or ContextOverflowMapperMiddleware)
can deal with them.
"""
from __future__ import annotations
import logging
import threading
from collections.abc import Awaitable, Callable
from langchain.agents.middleware.types import (
AgentMiddleware,
ModelRequest,
ModelResponse,
)
logger = logging.getLogger(__name__)
_ui_emit_fn: Callable[[str, str], None] | None = None
"""UI callback registered by the CLI/TUI entrypoint. ``None`` until set."""
_fallback_chain_lock = threading.Lock()
_fallback_chain: list[tuple[str, str]] = []
"""Ordered list of ``(model_name, provider)`` fallback entries."""
_CONTEXT_LIMIT_PATTERNS: list[str] = [
"context_length_exceeded",
"context length exceeded",
"too many tokens",
"maximum context length",
"output too large",
"context_window_exceeded",
"string_too_long",
"max_tokens_exceeded",
]
"""Substrings that identify a context-length error in provider messages."""
_MALFORMED_REQUEST_PATTERNS: list[str] = [
"invalid_request_error",
"invalid request",
"malformed",
]
"""Substrings that identify a malformed request (client-side bug)."""
_AUTH_ERROR_PATTERNS: list[str] = [
"invalid_api_key",
"authentication",
"permission",
]
"""Substrings that identify auth/permission errors.
These are intentionally *not* treated as non-fallbackable because a different
provider in the chain may have valid credentials."""
def set_ui_emit(fn: Callable[[str, str], None] | None) -> None:
"""Register (or clear) the UI callback for fallback status messages.
Args:
fn: Callable with signature ``fn(text, style)`` where *style* is a
Rich style string (``"yellow"``, ``"red"``, ``"green"``).
Pass ``None`` to unregister.
"""
global _ui_emit_fn
_ui_emit_fn = fn
def _emit(text: str, style: str = "yellow") -> None:
"""Surface a fallback status message to the user.
Dispatches to the registered UI callback when available (TUI mode),
otherwise falls back to the shared Rich console on stdout (CLI mode).
Args:
text: Plain-text message to display.
style: Rich style string applied to the message.
"""
if _ui_emit_fn is not None:
try:
_ui_emit_fn(text, style)
return
except Exception:
pass
from ..stream.console import console
console.print(text, style=style)
def get_fallback_chain() -> list[tuple[str, str]]:
"""Return a snapshot of the current fallback chain.
Returns:
List of ``(model_name, provider)`` tuples in priority order.
"""
with _fallback_chain_lock:
return list(_fallback_chain)
def set_fallback_chain(chain: list[tuple[str, str]]) -> None:
"""Replace the entire fallback chain.
Args:
chain: New list of ``(model_name, provider)`` tuples.
"""
global _fallback_chain
with _fallback_chain_lock:
_fallback_chain = list(chain)
def add_fallback(model: str, provider: str) -> bool:
"""Append a model to the end of the fallback chain.
Args:
model: Short model name (e.g. ``"gpt-5.5"``).
provider: Provider identifier (e.g. ``"openai"``).
Returns:
``True`` if added, ``False`` if the entry was already present.
"""
entry = (model, provider)
with _fallback_chain_lock:
if entry in _fallback_chain:
return False
_fallback_chain.append(entry)
return True
def remove_fallback(model: str) -> bool:
"""Remove all entries matching *model* regardless of provider.
Args:
model: Short model name to remove.
Returns:
``True`` if at least one entry was removed.
"""
global _fallback_chain
with _fallback_chain_lock:
before = len(_fallback_chain)
_fallback_chain = [(m, p) for m, p in _fallback_chain if m != model]
return len(_fallback_chain) < before
def remove_fallback_at(index: int) -> tuple[str, str] | None:
"""Remove the entry at a 0-based index.
Args:
index: Position in the chain (0-based).
Returns:
The removed ``(model, provider)`` tuple, or ``None`` if out of range.
"""
with _fallback_chain_lock:
if 0 <= index < len(_fallback_chain):
return _fallback_chain.pop(index)
return None
def clear_fallbacks() -> None:
"""Remove every entry from the fallback chain."""
global _fallback_chain
with _fallback_chain_lock:
_fallback_chain = []
def serialize_fallback_chain() -> str:
"""Serialize the chain to a config-friendly string.
Returns:
Comma-separated ``"model:provider,model:provider"`` string.
"""
with _fallback_chain_lock:
return ",".join(f"{m}:{p}" for m, p in _fallback_chain)
def load_fallback_chain(raw: str) -> None:
"""Populate the chain from a serialized config string.
Args:
raw: Comma-separated ``"model:provider"`` pairs. Empty or
whitespace-only segments are silently skipped.
"""
global _fallback_chain
chain: list[tuple[str, str]] = []
for part in raw.split(","):
part = part.strip()
if not part:
continue
if ":" in part:
model, provider = part.rsplit(":", 1)
chain.append((model.strip(), provider.strip()))
with _fallback_chain_lock:
_fallback_chain = chain
def _is_non_fallbackable(exc: Exception) -> str | None:
"""Determine whether an exception should bypass the fallback chain.
Args:
exc: The exception raised by a model call.
Returns:
A human-readable reason string if the error must *not* trigger
fallback, or ``None`` if fallback should proceed.
"""
from langchain_core.exceptions import ContextOverflowError
if isinstance(exc, ContextOverflowError):
return "context length exceeded"
err_msg = str(exc).lower()
is_400 = "400" in err_msg or "bad request" in err_msg
if is_400 and any(p in err_msg for p in _CONTEXT_LIMIT_PATTERNS):
return "context length exceeded"
if is_400 and any(p in err_msg for p in _MALFORMED_REQUEST_PATTERNS):
return "malformed request (client-side error)"
return None
async def _try_fallbacks(
request: ModelRequest,
invoke: Callable[[ModelRequest], Awaitable[ModelResponse]],
primary_exc: Exception,
) -> ModelResponse:
"""Walk the fallback chain, trying each model until one succeeds.
Shared implementation for both sync and async middleware entry points.
The *invoke* callable is an async function that calls the handler with
a given request — the sync path wraps the synchronous handler in a
trivial coroutine so both paths converge here.
Args:
request: The original model request.
invoke: Async callable that invokes the handler on a request.
primary_exc: The exception raised by the primary model.
Returns:
The ``ModelResponse`` from the first successful fallback.
Raises:
Exception: Re-raises the last exception if all fallbacks fail.
"""
from ..llm.models import get_chat_model
_emit(
f"Primary model failed: {type(primary_exc).__name__}: {primary_exc}",
style="yellow",
)
logger.warning(
"Primary model failed: %s: %s", type(primary_exc).__name__, primary_exc
)
last_exc = primary_exc
for model_name, provider in get_fallback_chain():
_emit(
f" -> Falling back to {model_name} ({provider}) "
f"due to: {type(last_exc).__name__}: {last_exc}",
style="yellow",
)
try:
fallback_model = get_chat_model(model=model_name, provider=provider)
fb_request = request.override(model=fallback_model)
result = await invoke(fb_request)
_emit(
f" Fallback to {model_name} ({provider}) succeeded",
style="green",
)
logger.info("Fallback to %s (%s) succeeded", model_name, provider)
return result
except Exception as fb_exc:
reason = _is_non_fallbackable(fb_exc)
if reason is not None:
_emit(
f" {model_name} hit non-fallbackable error ({reason}) "
f"-- aborting fallback chain",
style="red",
)
raise
last_exc = fb_exc
_emit(
f" x {model_name} also failed: {type(fb_exc).__name__}: {fb_exc}",
style="red",
)
logger.warning(
"Fallback %s (provider=%s) failed: %s: %s",
model_name,
provider,
type(fb_exc).__name__,
fb_exc,
)
_emit(" All fallbacks exhausted -- re-raising last error", style="red")
raise last_exc
def _guard_and_fallback(
primary_exc: Exception,
request: ModelRequest,
invoke: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> Awaitable[ModelResponse]:
"""Check non-fallbackable conditions, then delegate to ``_try_fallbacks``.
Args:
primary_exc: The exception raised by the primary model.
request: The original model request.
invoke: Async callable that invokes the handler on a request.
Returns:
Coroutine that resolves to the fallback ``ModelResponse``.
Raises:
Exception: Re-raises immediately for non-fallbackable errors.
"""
reason = _is_non_fallbackable(primary_exc)
if reason is not None:
_emit(
f"Model error ({reason}) -- not eligible for fallback, re-raising",
style="red",
)
raise primary_exc
return _try_fallbacks(request, invoke, primary_exc)
class ModelFallbackMiddleware(AgentMiddleware):
"""LangChain AgentMiddleware that retries failed model calls on fallbacks.
On each invocation the middleware reads the module-level
``_fallback_chain`` so that ``/model-fallback add`` takes effect
immediately without rebuilding the agent.
Attributes:
name: Middleware identifier used by the framework.
"""
name = "model_fallback"
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelResponse:
if not _fallback_chain:
return handler(request)
try:
return handler(request)
except Exception as exc:
async def _sync_invoke(r: ModelRequest) -> ModelResponse:
return handler(r)
import asyncio
return asyncio.run(_guard_and_fallback(exc, request, _sync_invoke))
async def awrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse:
if not _fallback_chain:
return await handler(request)
try:
return await handler(request)
except Exception as exc:
return await _guard_and_fallback(exc, request, handler)
+5 -3
View File
@@ -2,9 +2,11 @@
``resolve(ModelRef, role)`` validates the provider, model, credentials,
capabilities, limits, and input budget, then freezes a complete
``ResolvedModelConfig`` (section 6.4). ``resolve_for_test`` relaxes only the
"model must be enabled" user-visibility check (section 9.4); every other
validation still runs. ``compute_availability`` implements the section 4.3
``ResolvedModelConfig`` (section 6.4). ``resolve_for_test`` relaxes the two
run-time visibility gates — the model need not be enabled and no passing
verification record is required (section 9.4) — because the provider test is
exactly what produces that verification; every other validation still runs.
``compute_availability`` implements the section 4.3
six-state judgement order and is the only availability computation — callers
must not derive state from ``enabled`` flags or test timestamps on their own.
+246
View File
@@ -0,0 +1,246 @@
"""Runtime-link glue for snapshot-driven model resolution (design doc 8.3).
``SnapshotRuntime`` bundles the store, resolver, and snapshot service into
the single accessor the running agent link (middleware, agent factories,
auxiliary roles) uses. It provides two resolution paths:
- ``build_role_model(snapshot, role)`` — per-call construction from a frozen
run snapshot (section 6.1 role mapping, section 5.2 per-call credential
resolution). This is the only path the per-call middleware uses.
- ``build_default_role_model(role)`` — build-time construction from the
active registry's ``defaults`` for auxiliary roles that bind a model when
a graph is compiled (tool selector, scheduler, memory workers). Auxiliary,
summary, and tool-selector roles map to ``defaults.auxiliary ??
defaults.primary``; no free model/provider strings are consulted anywhere.
Every model is constructed through ``build_chat_model`` with both safe HTTP
clients (section 6.4). Ollama adapters additionally carry ``max_retries``
on the safe transports because langchain-ollama exposes no client-level
retry option. Secrets are resolved per call and never logged.
"""
from __future__ import annotations
import threading
from pathlib import Path
import yaml
from langchain_core.language_models.chat_models import BaseChatModel
from EvoScientist.config.settings import get_config_path
from .endpoint_policy import EndpointPolicy
from .errors import MODEL_REGISTRY_NOT_READY, ModelRegistryError
from .factory import build_chat_model
from .platform import DEFAULT_LOCAL_DEPLOYMENT_ID, PlatformConfigError
from .resolver import ModelRegistryResolver
from .safe_transport import build_safe_async_http_client, build_safe_http_client
from .schemas import DevelopmentEndpoint, ModelRef, ModelRole, ResolvedModelConfig
from .snapshots import RuntimeSnapshot, SnapshotService, config_for_role
from .store import ModelRuntimeStore
class SnapshotRuntime:
"""Store/resolver/snapshot bundle shared by the whole runtime link."""
def __init__(
self,
store: ModelRuntimeStore,
*,
endpoint_policy: EndpointPolicy | None = None,
local_deployment_id: str = DEFAULT_LOCAL_DEPLOYMENT_ID,
) -> None:
self._store = store
self._resolver = ModelRegistryResolver(store)
self._snapshots = SnapshotService(store, self._resolver)
self._endpoint_policy = (
endpoint_policy if endpoint_policy is not None else EndpointPolicy(())
)
self._local_deployment_id = local_deployment_id
@property
def store(self) -> ModelRuntimeStore:
return self._store
@property
def resolver(self) -> ModelRegistryResolver:
return self._resolver
@property
def snapshots(self) -> SnapshotService:
return self._snapshots
@property
def local_deployment_id(self) -> str:
return self._local_deployment_id
# --- snapshot-driven (per-call) resolution ------------------------------
def get_snapshot(
self, snapshot_id: str, *, deployment_id: str, thread_id: str
) -> RuntimeSnapshot:
"""Read a snapshot after verifying its deployment/thread binding."""
return self._snapshots.get(
snapshot_id, deployment_id=deployment_id, thread_id=thread_id
)
def build_role_model(
self, snapshot: RuntimeSnapshot, role: ModelRole
) -> BaseChatModel:
"""Build the chat model for one role of a frozen run snapshot."""
config = config_for_role(snapshot, role)
credential = self._snapshots.resolve_snapshot_credential(snapshot, role)
return self._build(config, credential)
# --- registry-default (build-time) resolution ----------------------------
def registry_defaults(self) -> tuple[ModelRef, ModelRef | None, int]:
"""Return ``(primary, auxiliary, revision)`` for an active registry."""
registry = self._store.load_registry()
if registry.state != "active" or registry.defaults.primary is None:
raise ModelRegistryError(
MODEL_REGISTRY_NOT_READY,
"The model registry is in bootstrap; configure and enable a "
"primary model first.",
)
return registry.defaults.primary, registry.defaults.auxiliary, registry.revision
def resolve_role_config(self, role: ModelRole = "primary") -> ResolvedModelConfig:
"""Resolve a role against the registry defaults (section 6.1).
``primary`` resolves ``defaults.primary``; every auxiliary role
(``auxiliary``, ``summary``, ``tool_selector``) resolves
``defaults.auxiliary ?? defaults.primary``.
"""
primary, auxiliary, _ = self.registry_defaults()
if role == "primary" or auxiliary is None:
return self._resolver.resolve(primary, "primary")
return self._resolver.resolve(auxiliary, "auxiliary")
def build_default_role_model(self, role: ModelRole = "primary") -> BaseChatModel:
"""Build the chat model for a role from the registry defaults."""
config = self.resolve_role_config(role)
return self._build(config, self._resolve_credential(config))
# --- internals -------------------------------------------------------------
def _resolve_credential(self, config: ResolvedModelConfig) -> str:
auth_ref = config.auth_ref
if auth_ref.mode == "none":
return ""
assert auth_ref.credential_id is not None # AuthSpec validation
assert auth_ref.credential_revision is not None
return self._store.resolve_credential(
auth_ref.credential_id, auth_ref.credential_revision
)
def _build(self, config: ResolvedModelConfig, credential: str) -> BaseChatModel:
# langchain-ollama exposes no client-level retry option, so the safe
# transports carry the frozen retry budget for ollama adapters; other
# adapters receive retries through their own client options.
retries = (
config.client_options.max_retries if config.adapter_id == "ollama" else 0
)
timeout = config.client_options.timeout_seconds
return build_chat_model(
config,
http_client=build_safe_http_client(
self._endpoint_policy, timeout=timeout, retries=retries
),
http_async_client=build_safe_async_http_client(
self._endpoint_policy, timeout=timeout, retries=retries
),
credential=credential,
)
# --- process-shared default runtime -------------------------------------------
_default_runtime: SnapshotRuntime | None = None
_default_runtime_lock = threading.Lock()
def _read_local_platform_fields() -> tuple[
str, Path | None, tuple[DevelopmentEndpoint, ...]
]:
"""Read the runtime-relevant platform fields without the BFF auth gates.
Unlike ``load_platform_security_config`` (which rightly refuses to serve
the Config API without a BFF token and delegation keys), the runtime
link only needs ``local_deployment_id``, ``model_runtime_db``, and
``development_endpoints`` — all safe to read with the documented
defaults when ``config.yaml`` is absent.
"""
path = get_config_path()
if not path.exists():
return DEFAULT_LOCAL_DEPLOYMENT_ID, None, ()
try:
with open(path, encoding="utf-8") as handle:
data = yaml.safe_load(handle) or {}
except yaml.YAMLError as exc:
raise PlatformConfigError(f"config.yaml is not valid YAML: {exc}.") from exc
if not isinstance(data, dict):
raise PlatformConfigError("config.yaml must contain a mapping at the top.")
deployment_id = data.get("local_deployment_id")
if deployment_id is not None and (
not isinstance(deployment_id, str) or not deployment_id.strip()
):
raise PlatformConfigError(
"config.yaml field 'local_deployment_id' must be a non-empty string."
)
database = data.get("model_runtime_db")
if database is not None and not isinstance(database, str):
raise PlatformConfigError(
"config.yaml field 'model_runtime_db' must be a string."
)
raw_endpoints = data.get("development_endpoints")
endpoints: tuple[DevelopmentEndpoint, ...] = ()
if raw_endpoints is not None:
if not isinstance(raw_endpoints, list):
raise PlatformConfigError(
"config.yaml field 'development_endpoints' must be a list."
)
try:
endpoints = tuple(
DevelopmentEndpoint.model_validate(entry) for entry in raw_endpoints
)
except ValueError as exc:
raise PlatformConfigError(
f"Invalid development_endpoints entry: {exc}."
) from exc
return (
deployment_id.strip() if deployment_id else DEFAULT_LOCAL_DEPLOYMENT_ID,
Path(database).expanduser() if database else None,
endpoints,
)
def _build_default_runtime() -> SnapshotRuntime:
deployment_id, database_path, endpoints = _read_local_platform_fields()
store = (
ModelRuntimeStore(database_path=database_path)
if database_path is not None
else ModelRuntimeStore()
)
return SnapshotRuntime(
store,
endpoint_policy=EndpointPolicy(endpoints),
local_deployment_id=deployment_id,
)
def get_snapshot_runtime() -> SnapshotRuntime:
"""Return the process-shared runtime, building it on first access."""
global _default_runtime
if _default_runtime is None:
with _default_runtime_lock:
if _default_runtime is None:
_default_runtime = _build_default_runtime()
return _default_runtime
def set_snapshot_runtime_for_tests(runtime: SnapshotRuntime | None) -> None:
"""Install (or clear) the shared runtime so tests stay hermetic."""
global _default_runtime
_default_runtime = runtime
+21 -12
View File
@@ -7,10 +7,13 @@ construction utility, not a deployment concern. Any deployment surface
servers) can call ``build_async_subagent_graph(name)`` to materialize the
runnable graph.
Reuses the main EvoScientist agent's chat model, backend, and middleware so
the deployed sub-agent has full capability parity with its in-process
synchronous counterpart: same workspace files, same ``/skills/`` and
``/memories/`` routes, same error-handling and context-overflow middleware.
Reuses the main EvoScientist agent's backend and middleware so the deployed
sub-agent has full capability parity with its in-process synchronous
counterpart: same workspace files, same ``/skills/`` and ``/memories/``
routes, same error-handling and context-overflow middleware. The chat model
is resolved from the active model registry's defaults (``auxiliary`` for the
scheduler, ``primary`` otherwise) and re-resolved per run from the run
snapshot by ``ConfigurableModelMiddleware``.
"""
from __future__ import annotations
@@ -40,13 +43,12 @@ def build_async_subagent_graph(name: str) -> Any:
from EvoScientist.config import apply_config_to_env, get_effective_config
from EvoScientist.EvoScientist import (
SUBAGENTS_CONFIG,
_ensure_auxiliary_chat_model,
_ensure_chat_model,
_ensure_general_purpose_subagent,
_get_default_backend,
_get_default_middleware,
_inject_subagent_middleware,
)
from EvoScientist.model_registry.runtime import get_snapshot_runtime
from EvoScientist.tools import skill_manager, tavily_search, think_tool
from EvoScientist.utils import load_subagents
@@ -99,18 +101,25 @@ def build_async_subagent_graph(name: str) -> Any:
#
# Memory middleware is included so async sub-agents get the same profile
# context and `/memories/profile/...` file guidance as the main agent.
#
# The compile-time model binding is resolved from the active registry's
# defaults via build_chat_model — never from config.yaml free strings.
# Per-run calls are re-resolved from the run snapshot by
# ConfigurableModelMiddleware; the scheduler (an unattended timer task)
# binds the cheaper auxiliary role, working sub-agents bind primary.
runtime = get_snapshot_runtime()
snapshot_role = "auxiliary" if name == "scheduler" else "primary"
model = runtime.build_default_role_model(snapshot_role)
subagents = []
_ensure_general_purpose_subagent(subagents)
_inject_subagent_middleware(subagents)
_inject_subagent_middleware(subagents, chat_model=model)
middleware = _get_default_middleware(
for_async_subagent=True,
memory_source_agent=name,
)
# Scheduler is an unattended timer task → use the cheaper auxiliary model.
model = (
_ensure_auxiliary_chat_model() if name == "scheduler" else _ensure_chat_model()
chat_model=model,
snapshot_role=snapshot_role,
)
return create_deep_agent(
+23
View File
@@ -23,6 +23,29 @@ def _reset_tool_selection_state():
reset_tool_selection_state_for_tests()
@pytest.fixture(autouse=True)
def isolated_snapshot_runtime(tmp_path):
"""Isolate the shared snapshot runtime around every test.
``model_registry.runtime.get_snapshot_runtime()`` otherwise builds a
store at the real ``~/.config/evoscientist`` — tests must never create
or read that database. Each test gets a fresh runtime backed by a
tmp-path store in ``bootstrap`` state; tests that need an active
registry populate ``runtime.store`` themselves (see
``tests/registry_fixtures.py``).
"""
from EvoScientist.model_registry.runtime import (
SnapshotRuntime,
set_snapshot_runtime_for_tests,
)
from EvoScientist.model_registry.store import ModelRuntimeStore
runtime = SnapshotRuntime(ModelRuntimeStore(config_dir=tmp_path / "model-runtime"))
set_snapshot_runtime_for_tests(runtime)
yield runtime
set_snapshot_runtime_for_tests(None)
@pytest.fixture
def sample_tool_call():
"""A minimal tool call dict."""
+186
View File
@@ -0,0 +1,186 @@
"""Shared builders for model-registry integration tests.
Centralizes the "active registry + verified models + snapshot" fixture
graph so middleware and runtime tests don't each re-derive the RegistryV4
payloads. Mirrors the fixtures in ``tests/test_snapshots.py``; that module
keeps its own copies to stay self-contained (Task 4 deliverable).
"""
from __future__ import annotations
from EvoScientist.model_registry.adapters import find_adapter_spec
from EvoScientist.model_registry.hashing import configuration_hash
from EvoScientist.model_registry.resolver import ModelRegistryResolver
from EvoScientist.model_registry.schemas import (
CredentialWrite,
ModelRef,
RegistryV4,
)
from EvoScientist.model_registry.snapshots import (
SnapshotCreateRequest,
SnapshotService,
)
from EvoScientist.model_registry.store import ModelRuntimeStore
ZHIPU_REF = ModelRef(provider_id="zhipu-glm", model_key="glm-5.2")
OLLAMA_REF = ModelRef(provider_id="local-ollama", model_key="qwen3")
ZHIPU_SECRET = "sk-live-9876abcd"
DEFAULT_MODEL_RUNTIME = {
"limit_mode": "combined",
"context_window_tokens": 1048576,
"max_input_tokens": None,
"max_output_tokens": 32768,
"min_effective_input_tokens": 8192,
"fixed_system_reserve_tokens": 4096,
"fixed_tools_reserve_tokens": 8192,
"fixed_attachments_reserve_tokens": 4096,
"limits_status": "confirmed",
"limits_source": "provider",
"temperature": None,
"top_p": None,
"reasoning_effort": "auto",
"declared_capabilities": {
"tools": True,
"vision": False,
"structured_output": True,
},
}
def model_runtime_payload(**overrides) -> dict:
payload = dict(DEFAULT_MODEL_RUNTIME)
payload.update(overrides)
return payload
def zhipu_provider_payload(**model_overrides) -> dict:
return {
"id": "zhipu-glm",
"name": "Zhipu GLM",
"adapter": "openai-compatible",
"base_url": "https://open.bigmodel.cn/api/paas/v4",
"auth": {"mode": "api_key", "credential_id": "zhipu-primary"},
"enabled": True,
"runtime": {
"timeout_seconds": 120,
"max_retries": 2,
"default_temperature": 0.7,
"default_top_p": 0.95,
"default_reasoning_effort": "auto",
},
"models": [
{
"key": "glm-5.2",
"name": "GLM-5.2",
"upstream_model_id": "glm-5.2",
"enabled": True,
"runtime": model_runtime_payload(**model_overrides),
}
],
}
def ollama_provider_payload(**model_overrides) -> dict:
return {
"id": "local-ollama",
"name": "Local Ollama",
"adapter": "ollama",
"base_url": "http://localhost:11434",
"auth": {"mode": "none", "credential_id": None},
"enabled": True,
"runtime": {"timeout_seconds": 120, "max_retries": 2},
"models": [
{
"key": "qwen3",
"name": "Qwen3",
"upstream_model_id": "qwen3",
"enabled": True,
"runtime": model_runtime_payload(**model_overrides),
}
],
}
def registry_payload(*, auxiliary_default: bool = True, **model_overrides) -> dict:
return {
"version": 4,
"revision": 1,
"state": "bootstrap",
"defaults": {
"primary": {"provider_id": "zhipu-glm", "model_key": "glm-5.2"},
"auxiliary": (
{"provider_id": "local-ollama", "model_key": "qwen3"}
if auxiliary_default
else None
),
},
"providers": [
zhipu_provider_payload(**model_overrides),
ollama_provider_payload(**model_overrides),
],
}
def _verify_model(store, registry, provider_id, model_key):
provider = registry.find_provider(provider_id)
model = provider.find_model(model_key)
spec = find_adapter_spec(provider.adapter, model.upstream_model_id)
store.record_model_verification(
provider_id=provider_id,
model_key=model_key,
configuration_hash=configuration_hash(provider, model),
credential_revision=1 if provider.auth.credential_id else 0,
adapter_spec_revision=spec.spec_revision,
result="passed",
verified_capabilities={
"tools": True,
"vision": False,
"structured_output": True,
},
)
def activate_store(store, *, auxiliary_default: bool = True, **model_overrides):
"""Persist an active registry with verified models into *store*."""
registry = store.save_registry(
expected_revision=1,
registry=RegistryV4.model_validate(
registry_payload(auxiliary_default=auxiliary_default, **model_overrides)
),
credential_writes=[
CredentialWrite(credential_id="zhipu-primary", secret_value=ZHIPU_SECRET)
],
)
assert registry.state == "active"
_verify_model(store, registry, "zhipu-glm", "glm-5.2")
_verify_model(store, registry, "local-ollama", "qwen3")
return store
def make_active_store(config_dir, *, auxiliary_default: bool = True, **model_overrides):
"""Persist an active registry with verified models into a fresh store."""
return activate_store(
ModelRuntimeStore(config_dir=config_dir),
auxiliary_default=auxiliary_default,
**model_overrides,
)
def make_snapshot_request(**overrides) -> SnapshotCreateRequest:
payload = {
"run_request_id": "req-1",
"thread_id": "thread-1",
"deployment_id": "local",
"model_selection_revision": 0,
"primary": None,
"auxiliary": None,
}
payload.update(overrides)
return SnapshotCreateRequest.model_validate(payload)
def make_snapshot(store, **overrides):
"""Create a snapshot against the store's registry defaults."""
service = SnapshotService(store, ModelRegistryResolver(store))
return service.create(make_snapshot_request(**overrides)).snapshot
+4 -1
View File
@@ -618,7 +618,10 @@ def test_for_async_subagent_omits_ask_user_middleware(
# Other middleware must remain — only ask_user is filtered.
assert "ConfigurableModelMiddleware" in async_names
assert "ContextEditingMiddleware" in async_names
assert "ModelFallbackMiddleware" in async_names
# The fallback chain was removed (design doc 8.3): it must appear in
# neither the default nor the async sub-agent middleware stack.
assert "ModelFallbackMiddleware" not in async_names
assert "ModelFallbackMiddleware" not in default_names
# ---------------------------------------------------------------------------
+68 -3
View File
@@ -42,7 +42,7 @@ def _assert_subagent_memory_middleware(subagent: dict, *, source_agent: str) ->
@patch("EvoScientist.EvoScientist._load_mcp_tools_cached", return_value={})
@patch("EvoScientist.EvoScientist._get_default_middleware", return_value=[])
@patch("EvoScientist.EvoScientist._get_default_backend")
@patch("EvoScientist.EvoScientist._ensure_chat_model")
@patch("EvoScientist.model_registry.runtime.get_snapshot_runtime")
@patch("EvoScientist.utils.load_subagents")
@patch("EvoScientist.config.apply_config_to_env")
@patch("EvoScientist.config.get_effective_config")
@@ -50,7 +50,7 @@ def test_factory_requests_async_safe_middleware(
mock_get_cfg,
mock_apply_env,
mock_load_subs,
mock_chat,
mock_get_runtime,
mock_backend,
mock_get_mw,
mock_mcp,
@@ -71,6 +71,9 @@ def test_factory_requests_async_safe_middleware(
cfg.memory_observation_writer = MemoryObservationWriter.ALL
cfg.memory_workers_enabled = True
mock_get_cfg.return_value = cfg
# The compile-time model binding comes from the registry-driven runtime.
model = MagicMock(name="registry_model")
mock_get_runtime.return_value.build_default_role_model.return_value = model
# Factory looks up the requested name in the loaded subagent specs;
# any matching name is fine.
mock_load_subs.return_value = [
@@ -89,11 +92,19 @@ def test_factory_requests_async_safe_middleware(
build_async_subagent_graph("writing-agent")
# The contract: factory MUST pass async-safe mode and the source agent name.
# The model is registry-resolved for the working (primary) role…
mock_get_runtime.return_value.build_default_role_model.assert_called_once_with(
"primary"
)
# …and the contract: factory MUST pass async-safe mode, the source agent
# name, the registry model, and the snapshot role.
mock_get_mw.assert_called_once_with(
for_async_subagent=True,
memory_source_agent="writing-agent",
chat_model=model,
snapshot_role="primary",
)
assert mock_create.call_args.kwargs["model"] is model
subagents = mock_create.call_args.kwargs["subagents"]
assert subagents[0]["name"] == "general-purpose"
_assert_subagent_memory_middleware(
@@ -102,6 +113,60 @@ def test_factory_requests_async_safe_middleware(
)
@patch("deepagents.create_deep_agent")
@patch("EvoScientist.EvoScientist._load_mcp_tools_cached", return_value={})
@patch("EvoScientist.EvoScientist._get_default_middleware", return_value=[])
@patch("EvoScientist.EvoScientist._get_default_backend")
@patch("EvoScientist.model_registry.runtime.get_snapshot_runtime")
@patch("EvoScientist.utils.load_subagents")
@patch("EvoScientist.config.apply_config_to_env")
@patch("EvoScientist.config.get_effective_config")
def test_scheduler_binds_auxiliary_role(
mock_get_cfg,
mock_apply_env,
mock_load_subs,
mock_get_runtime,
mock_backend,
mock_get_mw,
mock_mcp,
mock_create,
):
"""The scheduler is an unattended timer task → auxiliary role mapping."""
cfg = MagicMock()
cfg.recursion_limit = 1_000_000
cfg.memory_profile_enabled = False
cfg.memory_observations_enabled = False
cfg.memory_observation_writer = MemoryObservationWriter.ALL
cfg.memory_workers_enabled = False
mock_get_cfg.return_value = cfg
model = MagicMock(name="auxiliary_model")
mock_get_runtime.return_value.build_default_role_model.return_value = model
mock_load_subs.return_value = [
{
"name": "scheduler",
"system_prompt": "",
"tools": [],
"skills": None,
}
]
mock_create.return_value.with_config.return_value = MagicMock()
from EvoScientist.subagents._factory import build_async_subagent_graph
build_async_subagent_graph("scheduler")
mock_get_runtime.return_value.build_default_role_model.assert_called_once_with(
"auxiliary"
)
mock_get_mw.assert_called_once_with(
for_async_subagent=True,
memory_source_agent="scheduler",
chat_model=model,
snapshot_role="auxiliary",
)
assert mock_create.call_args.kwargs["model"] is model
@patch("EvoScientist.EvoScientist._ensure_chat_model")
def test_inject_subagent_adds_memory_middleware(mock_model, tmp_path):
mock_model.return_value = MagicMock(profile={"max_input_tokens": 200_000})
+89 -75
View File
@@ -1,17 +1,20 @@
"""Tests for the auxiliary-model resolver and its middleware scoping.
Covers ``EvoScientist.EvoScientist._ensure_auxiliary_chat_model`` (fallback to
the main model when unset) and the wiring in ``_get_default_middleware`` that
routes the main agent's tool selector to the auxiliary model while keeping
context editing — and async sub-agents — on the main model.
Covers ``EvoScientist.EvoScientist._ensure_auxiliary_chat_model`` — now
resolved through the model registry role mapping (design doc 6.1:
``defaults.auxiliary ?? defaults.primary``) instead of the legacy
``cfg.auxiliary_model``/``auxiliary_provider`` free strings — and the wiring
in ``_get_default_middleware`` that routes the main agent's tool selector to
the auxiliary model while keeping context editing — and async sub-agents —
on the main model.
"""
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import pytest
import EvoScientist.EvoScientist as E
from tests.registry_fixtures import activate_store
@pytest.fixture(autouse=True)
@@ -23,64 +26,66 @@ def _reset_model_caches(monkeypatch):
monkeypatch.setattr(E, "_auxiliary_chat_model_key", None, raising=False)
def _cfg(**over):
base = {
"model": "main-m",
"provider": "main-p",
"auxiliary_model": "",
"auxiliary_provider": "",
}
base.update(over)
return SimpleNamespace(**base)
class TestAuxiliaryResolver:
def test_empty_returns_main_instance(self, monkeypatch):
def test_bootstrap_registry_returns_main_instance(
self, monkeypatch, isolated_snapshot_runtime
):
"""Pre-cutover CLI (bootstrap registry): auxiliary mirrors main."""
main = object()
monkeypatch.setattr(E, "_ensure_config", lambda config=None: _cfg())
monkeypatch.setattr(E, "_ensure_chat_model", lambda: main)
assert E._ensure_auxiliary_chat_model() is main
def test_aux_equal_to_main_reuses_main_instance(self, monkeypatch):
def test_no_auxiliary_default_returns_main_instance(
self, monkeypatch, isolated_snapshot_runtime
):
main = object()
monkeypatch.setattr(
E,
"_ensure_config",
lambda config=None: _cfg(
auxiliary_model="main-m", auxiliary_provider="main-p"
),
)
activate_store(isolated_snapshot_runtime.store, auxiliary_default=False)
monkeypatch.setattr(E, "_ensure_chat_model", lambda: main)
assert E._ensure_auxiliary_chat_model() is main
def test_set_builds_auxiliary(self, monkeypatch):
def test_auxiliary_equal_to_primary_reuses_main_instance(
self, monkeypatch, isolated_snapshot_runtime
):
main = object()
store = activate_store(isolated_snapshot_runtime.store)
registry = store.load_registry()
registry.defaults.auxiliary = registry.defaults.primary
store.save_registry(expected_revision=registry.revision, registry=registry)
monkeypatch.setattr(E, "_ensure_chat_model", lambda: main)
assert E._ensure_auxiliary_chat_model() is main
def test_auxiliary_default_builds_via_role_mapping(
self, monkeypatch, isolated_snapshot_runtime
):
fake = object()
get_chat_model = MagicMock(return_value=fake)
build = MagicMock(return_value=fake)
activate_store(isolated_snapshot_runtime.store)
monkeypatch.setattr(
E,
"_ensure_config",
lambda config=None: _cfg(
auxiliary_model="aux-m", auxiliary_provider="aux-p"
),
isolated_snapshot_runtime, "build_default_role_model", build
)
monkeypatch.setattr(
E, "_ensure_chat_model", lambda: pytest.fail("must not build main")
)
monkeypatch.setattr("EvoScientist.llm.get_chat_model", get_chat_model)
assert E._ensure_auxiliary_chat_model() is fake
get_chat_model.assert_called_once_with(model="aux-m", provider="aux-p")
build.assert_called_once_with("auxiliary")
def test_empty_provider_falls_back_to_main_provider(self, monkeypatch):
get_chat_model = MagicMock(return_value=object())
def test_auxiliary_cache_reused_within_same_registry_revision(
self, monkeypatch, isolated_snapshot_runtime
):
build = MagicMock(side_effect=[object(), object()])
activate_store(isolated_snapshot_runtime.store)
monkeypatch.setattr(
E,
"_ensure_config",
lambda config=None: _cfg(auxiliary_model="aux-m", auxiliary_provider=""),
isolated_snapshot_runtime, "build_default_role_model", build
)
monkeypatch.setattr("EvoScientist.llm.get_chat_model", get_chat_model)
E._ensure_auxiliary_chat_model()
get_chat_model.assert_called_once_with(model="aux-m", provider="main-p")
first = E._ensure_auxiliary_chat_model()
assert E._ensure_auxiliary_chat_model() is first
assert build.call_count == 1
def test_set_chat_model_resets_aux_cache(self, monkeypatch):
monkeypatch.setattr(E, "_auxiliary_chat_model", object(), raising=False)
monkeypatch.setattr(E, "_auxiliary_chat_model_key", ("x", "y"), raising=False)
monkeypatch.setattr(
E, "_auxiliary_chat_model_key", ("x", "y", 1), raising=False
)
monkeypatch.setattr(
"EvoScientist.llm.get_chat_model", MagicMock(return_value=object())
)
@@ -94,9 +99,6 @@ def _mock_cfg():
cfg.enable_ask_user = False
cfg.auto_mode = False
cfg.auto_approve = False
cfg.model_fallbacks = None
cfg.auxiliary_model = ""
cfg.auxiliary_provider = ""
return cfg
@@ -158,13 +160,11 @@ class TestAuxiliaryMiddlewareScope:
assert cap["tool_selector"] is main_model
assert cap["context_editing"] is main_model
def test_pure_path_tool_selector_uses_threaded_main_when_aux_empty(self):
def test_pure_path_tool_selector_uses_threaded_model(self):
"""The pure path never resolves auxiliary free strings (Task 7 wires
the local snapshot entry); the threaded model stands in."""
cap, fake_ts, fake_ce = self._capture()
cfg = _mock_cfg()
cfg.model = "new-main"
cfg.provider = "new-provider"
cfg.auxiliary_model = ""
cfg.auxiliary_provider = ""
main_model = object()
with (
@@ -185,33 +185,47 @@ class TestAuxiliaryMiddlewareScope:
assert cap["tool_selector"] is main_model
assert cap["context_editing"] is main_model
def test_pure_path_tool_selector_builds_aux_from_threaded_config(self):
cap, fake_ts, fake_ce = self._capture()
def test_snapshot_role_forwarded_to_configurable_model_middleware(self):
cfg = _mock_cfg()
cfg.model = "new-main"
cfg.provider = "new-provider"
cfg.auxiliary_model = "new-aux"
cfg.auxiliary_provider = "aux-provider"
main_model, aux_model = object(), object()
main_model = object()
with (
patch.object(E, "_ensure_config", side_effect=AssertionError),
patch.object(E, "_ensure_chat_model", side_effect=AssertionError),
patch.object(E, "_ensure_auxiliary_chat_model", side_effect=AssertionError),
patch(
"EvoScientist.llm.get_chat_model", return_value=aux_model
) as get_model,
patch.object(E, "_ensure_config", return_value=cfg),
patch.object(E, "_ensure_chat_model", return_value=main_model),
patch.object(E, "_ensure_auxiliary_chat_model", return_value=main_model),
patch(
"EvoScientist.middleware.create_tool_selector_middleware",
side_effect=fake_ts,
),
patch(
"EvoScientist.middleware.create_context_editing_middleware",
side_effect=fake_ce,
return_value=[MagicMock()],
),
):
E._get_default_middleware(cfg=cfg, chat_model=main_model)
mw = E._get_default_middleware(snapshot_role="auxiliary")
get_model.assert_called_once_with(model="new-aux", provider="aux-provider")
assert cap["tool_selector"] is aux_model
assert cap["context_editing"] is main_model
configurable = next(
m for m in mw if type(m).__name__ == "ConfigurableModelMiddleware"
)
assert configurable._role == "auxiliary"
def test_memory_agent_factory_uses_auxiliary_role(monkeypatch):
"""Memory workers bind the auxiliary role model at graph build."""
sentinel = object()
monkeypatch.setattr(E, "_ensure_auxiliary_chat_model", lambda: sentinel)
captured: dict[str, object] = {}
def fake_create_deep_agent(**kwargs):
captured.update(kwargs)
return MagicMock()
monkeypatch.setattr("deepagents.create_deep_agent", fake_create_deep_agent)
from EvoScientist.memory.agents._factory import build_memory_agent_graph
build_memory_agent_graph(
name="worker",
system_prompt="",
memory_dir="/tmp/m",
workspace_dir="/tmp/w",
tools=[],
middleware=[],
backend=MagicMock(),
)
assert captured["model"] is sentinel
+25 -8
View File
@@ -31,17 +31,34 @@ class TestSlashCommandCompleter:
def test_exact_command_hides_even_when_prefix_of_another(self):
"""Regression for #293: ``/model`` must hide on exact match so Enter
submits it, even though ``/model-fallback`` shares the prefix. Before
submits it, even when another command shares the prefix. Before
the fix the popup stayed visible (two prefix matches) and the TUI's
Enter handler completed the text instead of executing the command.
The original prefix sibling (``/model-fallback``) was removed with
the fallback chain; a temporary stub command recreates the scenario.
"""
completer = SlashCommandCompleter()
# Sanity: both commands share the prefix, so a partial prefix lists both.
partial = {c.text for c in completer.get_completions(_doc("/mode"), None)}
assert {"/model", "/model-fallback"} <= partial
# Exact ``/model`` with no trailing space → hide.
completions = list(completer.get_completions(_doc("/model"), None))
assert completions == []
from EvoScientist.commands.base import Command
from EvoScientist.commands.manager import manager
class _StubCommand(Command):
name = "/model-extra"
description = "stub"
async def execute(self, ctx, args):
return None
manager.register(_StubCommand())
try:
completer = SlashCommandCompleter()
# Sanity: both commands share the prefix, so a partial prefix lists both.
partial = {c.text for c in completer.get_completions(_doc("/mode"), None)}
assert {"/model", "/model-extra"} <= partial
# Exact ``/model`` with no trailing space → hide.
completions = list(completer.get_completions(_doc("/model"), None))
assert completions == []
finally:
manager._commands.pop("/model-extra", None)
def test_non_slash_returns_empty(self):
completer = SlashCommandCompleter()
+10 -10
View File
@@ -45,22 +45,22 @@ class TestTopLevelCompletions:
class TestAliasVisibility:
def test_alias_prefix_matches(self):
r = compute_completions("/fa", 3)
r = compute_completions("/skills-re", 10)
names = [c.text for c in r.candidates]
assert "/model-fallback" in names
assert "/autoskills" in names
def test_alias_exact_hides_leaf(self):
r = compute_completions("/quit", 5)
assert r.kind == "empty"
def test_alias_exact_hides_without_space(self):
r = compute_completions("/fallback", 9)
r = compute_completions("/q", 2)
assert r.kind == "empty"
def test_alias_space_shows_subcommands(self):
r = compute_completions("/fallback ", 10)
r = compute_completions("/skills-review ", 15)
names = [c.text for c in r.candidates]
assert "add" in names
assert "approve" in names
assert "list" in names
@@ -79,7 +79,7 @@ class TestSubcommandCompletions:
assert "list" not in names
def test_leaf_subcommand_stops(self):
r = compute_completions("/model-fallback help ", 21)
r = compute_completions("/autoskills help ", 17)
assert r.kind == "empty"
def test_channel_all_types(self):
@@ -89,11 +89,11 @@ class TestSubcommandCompletions:
assert "telegram" in names
assert "discord" in names
def test_model_fallback_subcommands(self):
r = compute_completions("/model-fallback ", 16)
def test_autoskills_subcommands(self):
r = compute_completions("/autoskills ", 12)
names = [c.text for c in r.candidates]
assert "add" in names
assert "clear" in names
assert "approve" in names
assert "reject" in names
class TestDynamicCompletions:
+5 -17
View File
@@ -34,18 +34,6 @@ class TestCommandManagerSubcommands:
names = {name for name, _desc in scs}
assert names == {"list", "config", "add", "edit", "remove", "install"}
def test_model_fallback_has_subcommands(self):
"""``/model-fallback`` must expose its subcommands."""
manager = CommandManager()
from EvoScientist.commands.implementation.model_fallback import (
ModelFallbackCommand,
)
manager.register(ModelFallbackCommand())
scs = manager.list_subcommands("/model-fallback")
names = {name for name, _desc in scs}
assert {"list", "add", "remove", "clear", "save", "help"} <= names
def test_channel_has_subcommands(self):
"""``/channel`` must expose status, stop, and channel type subcommands."""
manager = CommandManager()
@@ -74,11 +62,11 @@ class TestCommandManagerSubcommands:
def test_subcommand_via_alias(self):
"""Registry by alias should still expose subcommands."""
manager = CommandManager()
from EvoScientist.commands.implementation.model_fallback import (
ModelFallbackCommand,
from EvoScientist.commands.implementation.autoskills import (
AutoSkillsCommand,
)
manager.register(ModelFallbackCommand())
scs = manager.list_subcommands("/fallback")
manager.register(AutoSkillsCommand())
scs = manager.list_subcommands("/skills-review")
names = {name for name, _desc in scs}
assert {"list", "add", "remove"} <= names
assert {"list", "approve", "reject"} <= names
+284 -205
View File
@@ -1,9 +1,10 @@
"""Tests for ``EvoScientist.middleware.configurable_model``.
"""Tests for ``EvoScientist.middleware.configurable_model`` (design doc 8.3).
Verifies that the middleware reads ``model`` / ``model_provider`` from
the active ``RunnableConfig.configurable`` (via ``langgraph.config.get_config``)
and overrides ``request.model`` accordingly, without breaking the no-override
pass-through path or the per-instance cache.
Verifies the snapshot-only contract: the middleware resolves the frozen
``ResolvedModelConfig`` for its role from ``configurable["runtime_snapshot_id"]``
and builds the model via ``build_chat_model``; ``model``/``model_provider``
in ``configurable`` are rejected with ``MODEL_CONFIG_OUTSIDE_SNAPSHOT``; a
run without a snapshot ID passes through unchanged.
"""
from __future__ import annotations
@@ -11,9 +12,26 @@ from __future__ import annotations
from contextlib import contextmanager
from unittest.mock import MagicMock, patch
import pytest
from langchain_ollama import ChatOllama
from langchain_openai import ChatOpenAI
from EvoScientist.middleware.configurable_model import (
ConfigurableModelMiddleware,
_read_model_override,
check_no_outside_snapshot_model_config,
read_snapshot_binding,
)
from EvoScientist.model_registry.errors import (
MODEL_CONFIG_OUTSIDE_SNAPSHOT,
SNAPSHOT_EXPIRED,
SNAPSHOT_NOT_FOUND,
ModelRegistryError,
)
from EvoScientist.model_registry.runtime import SnapshotRuntime
from tests.registry_fixtures import (
ZHIPU_SECRET,
make_active_store,
make_snapshot,
)
@@ -69,61 +87,104 @@ def _make_request():
return req
# =============================================================================
# 1. _read_model_override — input parsing
# =============================================================================
@pytest.fixture
def store(tmp_path):
return make_active_store(tmp_path / "model-runtime")
class TestReadModelOverride:
"""Verify the helper that pulls (model, provider) from active config."""
@pytest.fixture
def runtime(store):
return SnapshotRuntime(store)
def test_returns_override_when_both_present(self):
with _patched_config({"model": "gpt-5", "model_provider": "openai"}):
assert _read_model_override() == ("gpt-5", "openai")
def test_provider_optional(self):
with _patched_config({"model": "claude-haiku-4-5"}):
assert _read_model_override() == ("claude-haiku-4-5", None)
def test_no_configurable_key(self):
with _patched_config({}):
assert _read_model_override() == (None, None)
def test_outside_runnable_context(self):
"""``get_config`` raises outside a runnable — middleware must no-op."""
with _patched_config(None):
assert _read_model_override() == (None, None)
def test_empty_string_treated_as_absent(self):
with _patched_config({"model": "", "model_provider": ""}):
assert _read_model_override() == (None, None)
def test_non_string_ignored(self):
with _patched_config({"model": 42, "model_provider": object()}):
assert _read_model_override() == (None, None)
def test_non_dict_configurable_safe(self):
with _patched_config({"configurable": "not-a-dict"}):
# Inner ``configurable`` is the wrong type — patched_config
# wraps it again so we end up with {"configurable": {"configurable": "..."}}
# which has no model/model_provider keys → no override.
assert _read_model_override() == (None, None)
def test_non_dict_config_safe(self):
with _patched_config("garbage"):
assert _read_model_override() == (None, None)
def _configurable_for(snapshot, **overrides):
configurable = {
"runtime_snapshot_id": snapshot.snapshot_id,
"thread_id": snapshot.thread_id,
}
configurable.update(overrides)
return configurable
# =============================================================================
# 2. ConfigurableModelMiddleware — pass-through behavior
# 1. Outside-snapshot model configuration is rejected (section 8.2)
# =============================================================================
class TestOutsideSnapshotConfig:
@pytest.mark.parametrize("key", ["model", "model_provider"])
def test_model_keys_rejected(self, key):
with pytest.raises(ModelRegistryError) as excinfo:
check_no_outside_snapshot_model_config({key: "gpt-5"})
assert excinfo.value.code == MODEL_CONFIG_OUTSIDE_SNAPSHOT
assert excinfo.value.http_status == 422
assert excinfo.value.details[0].path == key
def test_rejected_alongside_snapshot_id(self, store, runtime):
mw = ConfigurableModelMiddleware(runtime=runtime)
snapshot = make_snapshot(store)
req = _make_request()
with (
_patched_config(_configurable_for(snapshot, model="gpt-5")),
pytest.raises(ModelRegistryError) as excinfo,
):
mw.wrap_model_call(req, MagicMock())
assert excinfo.value.code == MODEL_CONFIG_OUTSIDE_SNAPSHOT
def test_rejected_without_snapshot_id(self, runtime):
"""The check applies even on the pass-through path (验收 13)."""
mw = ConfigurableModelMiddleware(runtime=runtime)
req = _make_request()
handler = MagicMock()
with (
_patched_config({"model_provider": "openai"}),
pytest.raises(ModelRegistryError) as excinfo,
):
mw.wrap_model_call(req, handler)
assert excinfo.value.code == MODEL_CONFIG_OUTSIDE_SNAPSHOT
handler.assert_not_called()
def test_none_values_treated_as_absent(self):
check_no_outside_snapshot_model_config({"model": None, "thread_id": "t"})
# =============================================================================
# 2. Snapshot binding parsing
# =============================================================================
class TestReadSnapshotBinding:
def test_full_binding(self):
binding = read_snapshot_binding(
{
"runtime_snapshot_id": "snap-1",
"workspace_deployment_id": "deploy-1",
"thread_id": "thread-1",
}
)
assert binding == ("snap-1", "deploy-1", "thread-1")
def test_missing_snapshot_id_returns_none(self):
assert read_snapshot_binding({}) is None
assert read_snapshot_binding({"runtime_snapshot_id": ""}) is None
assert read_snapshot_binding({"runtime_snapshot_id": 42}) is None
def test_missing_deployment_and_thread_fall_back(self):
assert read_snapshot_binding({"runtime_snapshot_id": "snap-1"}) == (
"snap-1",
None,
None,
)
# =============================================================================
# 3. Pass-through behavior (no snapshot in configurable)
# =============================================================================
class TestPassThrough:
"""When no override is present, the middleware must not touch the request."""
def test_sync_no_override_passes_request_unchanged(self):
mw = ConfigurableModelMiddleware()
def test_sync_no_snapshot_passes_request_unchanged(self, runtime):
mw = ConfigurableModelMiddleware(runtime=runtime)
req = _make_request()
sentinel = object()
handler = MagicMock(return_value=sentinel)
@@ -133,8 +194,8 @@ class TestPassThrough:
handler.assert_called_once_with(req)
req.override.assert_not_called()
async def test_async_no_override_passes_request_unchanged(self):
mw = ConfigurableModelMiddleware()
async def test_async_no_snapshot_passes_request_unchanged(self, runtime):
mw = ConfigurableModelMiddleware(runtime=runtime)
req = _make_request()
async def handler(r):
@@ -146,8 +207,8 @@ class TestPassThrough:
assert result == "ok"
req.override.assert_not_called()
def test_outside_runnable_context_passes_through(self):
mw = ConfigurableModelMiddleware()
def test_outside_runnable_context_passes_through(self, runtime):
mw = ConfigurableModelMiddleware(runtime=runtime)
req = _make_request()
handler = MagicMock(return_value="ok")
with _patched_config(None):
@@ -156,190 +217,212 @@ class TestPassThrough:
# =============================================================================
# 3. ConfigurableModelMiddleware — override behavior
# 4. Snapshot-driven model construction (full chain)
# =============================================================================
class TestModelOverride:
"""When override present, middleware resolves model and overrides request."""
def test_sync_override_calls_get_chat_model_and_overrides(self):
mw = ConfigurableModelMiddleware()
class TestSnapshotDrivenConstruction:
def test_sync_builds_snapshot_primary_model(self, store, runtime):
mw = ConfigurableModelMiddleware(runtime=runtime)
snapshot = make_snapshot(store)
req = _make_request()
new_model = MagicMock(name="resolved_chat_model")
handler = MagicMock(return_value="response")
with (
_patched_config({"model": "gpt-5", "model_provider": "openai"}),
patch(
"EvoScientist.llm.get_chat_model", return_value=new_model
) as mock_get,
):
with _patched_config(_configurable_for(snapshot)):
mw.wrap_model_call(req, handler)
mock_get.assert_called_once_with(model="gpt-5", provider="openai")
req.override.assert_called_once_with(model=new_model)
# Handler must receive the OVERRIDDEN request, not the original.
called_with = handler.call_args[0][0]
assert called_with is not req
assert called_with.model is new_model
model = called_with.model
assert isinstance(model, ChatOpenAI)
assert model.model_name == "glm-5.2"
assert model.openai_api_base == "https://open.bigmodel.cn/api/paas/v4"
assert model.openai_api_key.get_secret_value() == ZHIPU_SECRET
assert model.max_retries == 2
async def test_async_override_path_parity(self):
mw = ConfigurableModelMiddleware()
async def test_async_path_parity(self, store, runtime):
mw = ConfigurableModelMiddleware(runtime=runtime)
snapshot = make_snapshot(store)
req = _make_request()
new_model = MagicMock()
async def handler(r):
assert r.model is new_model
return "ok"
return r.model
with _patched_config(_configurable_for(snapshot)):
model = await mw.awrap_model_call(req, handler)
assert isinstance(model, ChatOpenAI)
assert model.model_name == "glm-5.2"
def test_auxiliary_role_maps_to_snapshot_auxiliary(self, store, runtime):
mw = ConfigurableModelMiddleware(role="auxiliary", runtime=runtime)
snapshot = make_snapshot(store)
req = _make_request()
handler = MagicMock(return_value="ok")
with _patched_config(_configurable_for(snapshot)):
mw.wrap_model_call(req, handler)
model = handler.call_args[0][0].model
assert isinstance(model, ChatOllama)
assert model.model == "qwen3"
def test_summary_role_falls_back_to_snapshot_primary(self, tmp_path):
store = make_active_store(tmp_path / "db", auxiliary_default=False)
mw = ConfigurableModelMiddleware(role="summary", runtime=SnapshotRuntime(store))
snapshot = make_snapshot(store)
req = _make_request()
handler = MagicMock(return_value="ok")
with _patched_config(_configurable_for(snapshot)):
mw.wrap_model_call(req, handler)
model = handler.call_args[0][0].model
assert isinstance(model, ChatOpenAI)
assert model.model_name == "glm-5.2"
def test_deployment_id_from_configurable_is_verified(self, store, runtime):
mw = ConfigurableModelMiddleware(runtime=runtime)
snapshot = make_snapshot(store)
req = _make_request()
handler = MagicMock(return_value="ok")
with _patched_config(
_configurable_for(snapshot, workspace_deployment_id="local")
):
mw.wrap_model_call(req, handler)
handler.assert_called_once()
def test_wrong_thread_binding_fails_closed(self, store, runtime):
mw = ConfigurableModelMiddleware(runtime=runtime)
snapshot = make_snapshot(store)
req = _make_request()
with (
_patched_config(_configurable_for(snapshot, thread_id="other-thread")),
pytest.raises(ModelRegistryError) as excinfo,
):
mw.wrap_model_call(req, MagicMock())
assert excinfo.value.code == SNAPSHOT_NOT_FOUND
def test_wrong_deployment_binding_fails_closed(self, store, runtime):
mw = ConfigurableModelMiddleware(runtime=runtime)
snapshot = make_snapshot(store)
req = _make_request()
with (
_patched_config(
{"model": "claude-opus-4-8", "model_provider": "anthropic"}
_configurable_for(snapshot, workspace_deployment_id="other-deploy")
),
patch(
"EvoScientist.llm.get_chat_model", return_value=new_model
) as mock_get,
pytest.raises(ModelRegistryError) as excinfo,
):
result = await mw.awrap_model_call(req, handler)
mw.wrap_model_call(req, MagicMock())
assert excinfo.value.code == SNAPSHOT_NOT_FOUND
assert result == "ok"
mock_get.assert_called_once_with(model="claude-opus-4-8", provider="anthropic")
def test_provider_omitted_passed_as_none(self):
mw = ConfigurableModelMiddleware()
def test_expired_snapshot_fails_loudly(self, store, runtime):
mw = ConfigurableModelMiddleware(runtime=runtime)
snapshot = make_snapshot(store)
store.abort_run_snapshot(snapshot.snapshot_id)
req = _make_request()
new_model = MagicMock()
handler = MagicMock(return_value="ok")
with (
_patched_config({"model": "gpt-5"}),
patch(
"EvoScientist.llm.get_chat_model", return_value=new_model
) as mock_get,
_patched_config(_configurable_for(snapshot)),
pytest.raises(ModelRegistryError) as excinfo,
):
mw.wrap_model_call(req, handler)
mock_get.assert_called_once_with(model="gpt-5", provider=None)
mw.wrap_model_call(req, MagicMock())
assert excinfo.value.code == SNAPSHOT_EXPIRED
# =============================================================================
# 4. ConfigurableModelMiddleware — caching
# 5. Frozen credential revision resolution
# =============================================================================
class TestCredentialRevision:
def test_snapshot_keeps_its_frozen_credential_revision(self, store, runtime):
"""Rotating the credential after freezing must not affect the run."""
mw = ConfigurableModelMiddleware(runtime=runtime)
snapshot = make_snapshot(store)
store.write_credential_version("zhipu-primary", "sk-rotated-0000")
req = _make_request()
handler = MagicMock(return_value="ok")
with _patched_config(_configurable_for(snapshot)):
mw.wrap_model_call(req, handler)
model = handler.call_args[0][0].model
assert model.openai_api_key.get_secret_value() == ZHIPU_SECRET
# =============================================================================
# 6. Per-instance caching
# =============================================================================
class TestCache:
"""Two consecutive calls with same (model, provider) should hit cache."""
def test_cache_hit_avoids_second_get_chat_model(self):
mw = ConfigurableModelMiddleware()
req1 = _make_request()
req2 = _make_request()
new_model = MagicMock()
handler = MagicMock(return_value="ok")
with (
_patched_config({"model": "gpt-5", "model_provider": "openai"}),
patch(
"EvoScientist.llm.get_chat_model", return_value=new_model
) as mock_get,
):
mw.wrap_model_call(req1, handler)
mw.wrap_model_call(req2, handler)
# First call resolves via factory, second hits the cache.
assert mock_get.call_count == 1
def test_cache_miss_on_different_provider(self):
mw = ConfigurableModelMiddleware()
req = _make_request()
handler = MagicMock(return_value="ok")
with patch(
"EvoScientist.llm.get_chat_model", side_effect=[MagicMock(), MagicMock()]
) as mock_get:
with _patched_config(
{"model": "claude-sonnet-4-6", "model_provider": "anthropic"}
):
mw.wrap_model_call(req, handler)
with _patched_config(
{"model": "claude-sonnet-4-6", "model_provider": "custom-anthropic"}
):
mw.wrap_model_call(req, handler)
assert mock_get.call_count == 2
def test_independent_instances_have_independent_caches(self):
mw_a = ConfigurableModelMiddleware()
mw_b = ConfigurableModelMiddleware()
def test_second_call_hits_cache(self, store, runtime):
mw = ConfigurableModelMiddleware(runtime=runtime)
snapshot = make_snapshot(store)
req = _make_request()
handler = MagicMock(return_value="ok")
with (
_patched_config({"model": "gpt-5", "model_provider": "openai"}),
patch(
"EvoScientist.llm.get_chat_model",
side_effect=[MagicMock(), MagicMock()],
) as mock_get,
):
mw_a.wrap_model_call(req, handler)
mw_b.wrap_model_call(req, handler)
# Different instances must each resolve once.
assert mock_get.call_count == 2
# =============================================================================
# 5. Resilience — get_chat_model raising
# =============================================================================
class TestResolveFailure:
"""If get_chat_model raises, middleware must fall back to original model."""
def test_sync_falls_back_when_resolve_raises(self):
mw = ConfigurableModelMiddleware()
req = _make_request()
handler = MagicMock(return_value="ok")
with (
_patched_config({"model": "doesnotexist", "model_provider": "openai"}),
patch(
"EvoScientist.llm.get_chat_model",
side_effect=ValueError("unknown model"),
),
_patched_config(_configurable_for(snapshot)),
patch.object(
SnapshotRuntime, "build_role_model", wraps=runtime.build_role_model
) as build,
):
mw.wrap_model_call(req, handler)
mw.wrap_model_call(req, handler)
# Override never happened — handler called with original request.
handler.assert_called_once_with(req)
req.override.assert_not_called()
assert build.call_count == 1
async def test_async_falls_back_when_resolve_raises(self):
mw = ConfigurableModelMiddleware()
def test_independent_instances_have_independent_caches(self, store, runtime):
snapshot = make_snapshot(store)
req = _make_request()
called = []
async def handler(r):
called.append(r)
return "ok"
handler = MagicMock(return_value="ok")
with (
_patched_config({"model": "doesnotexist", "model_provider": "openai"}),
patch(
"EvoScientist.llm.get_chat_model",
side_effect=ValueError("unknown model"),
),
_patched_config(_configurable_for(snapshot)),
patch.object(
SnapshotRuntime, "build_role_model", wraps=runtime.build_role_model
) as build,
):
result = await mw.awrap_model_call(req, handler)
ConfigurableModelMiddleware(runtime=runtime).wrap_model_call(req, handler)
ConfigurableModelMiddleware(runtime=runtime).wrap_model_call(req, handler)
assert result == "ok"
assert called == [req]
assert build.call_count == 2
# =============================================================================
# 6. Integration — real langgraph contextvar (no get_config mock)
# 7. Constructor validation
# =============================================================================
class TestRoleValidation:
def test_unknown_role_rejected(self):
with pytest.raises(ValueError, match="Unknown model role"):
ConfigurableModelMiddleware(role="bogus")
# =============================================================================
# 8. Fallback removal startup assertion (design doc 8.3)
# =============================================================================
class TestFallbackChainRemoved:
def test_fallback_middleware_gone_from_package(self):
import importlib
import EvoScientist.middleware as mw
assert not hasattr(mw, "ModelFallbackMiddleware")
assert not hasattr(mw, "load_fallback_chain")
with pytest.raises(ModuleNotFoundError):
importlib.import_module("EvoScientist.middleware.model_fallback")
# =============================================================================
# 9. Integration — real langgraph contextvar (no get_config mock)
# =============================================================================
@@ -352,37 +435,33 @@ class TestRunnableContextVarIntegration:
different module, or the contextvar mechanism changes).
"""
def test_real_contextvar_drives_override(self):
"""Without mocking get_config, set the contextvar and verify override."""
def test_real_contextvar_drives_snapshot_resolution(self, store, runtime):
from langchain_core.runnables.config import var_child_runnable_config
mw = ConfigurableModelMiddleware()
mw = ConfigurableModelMiddleware(runtime=runtime)
snapshot = make_snapshot(store)
req = _make_request()
new_model = MagicMock(name="resolved")
handler = MagicMock(return_value="ok")
token = var_child_runnable_config.set(
{"configurable": {"model": "gpt-5.5", "model_provider": "openai"}}
{"configurable": _configurable_for(snapshot)}
)
try:
with patch(
"EvoScientist.llm.get_chat_model", return_value=new_model
) as mock_get:
mw.wrap_model_call(req, handler)
mw.wrap_model_call(req, handler)
finally:
var_child_runnable_config.reset(token)
mock_get.assert_called_once_with(model="gpt-5.5", provider="openai")
req.override.assert_called_once_with(model=new_model)
model = handler.call_args[0][0].model
assert isinstance(model, ChatOpenAI)
assert model.model_name == "glm-5.2"
def test_real_contextvar_unset_passes_through(self):
"""When no contextvar is set, get_config() raises → no override."""
def test_real_contextvar_unset_passes_through(self, runtime):
from langchain_core.runnables.config import var_child_runnable_config
# Defensive: ensure no leftover contextvar from another test.
token = var_child_runnable_config.set(None)
try:
mw = ConfigurableModelMiddleware()
mw = ConfigurableModelMiddleware(runtime=runtime)
req = _make_request()
handler = MagicMock(return_value="ok")
mw.wrap_model_call(req, handler)
+278 -32
View File
@@ -1,19 +1,78 @@
"""Tests for message-only context budgeting."""
"""Tests for message-only context budgeting (design doc 6.5, 8.3).
Covers both budget modes:
- Snapshot mode: input limit and fixed reserves come from the frozen
``ResolvedModelConfig.budget``; the per-call budget is recomputed on every
invocation with the current has_tools/has_attachments mode; the
summarizer model resolves through the snapshot's ``summary`` role.
- Interim local mode (no snapshot ID): limits from the compile-time model
profile with the default reserve policy.
"""
from __future__ import annotations
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import MagicMock
from unittest.mock import MagicMock, patch
import pytest
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
from langchain_ollama import ChatOllama
from EvoScientist.middleware.message_budget import (
ContextBudgetUnsatisfiableError,
MessageReservePolicy,
_snapshot_message_budget,
count_message_text_tokens,
create_message_budget_middleware,
)
from EvoScientist.model_registry.runtime import SnapshotRuntime
from tests.registry_fixtures import make_active_store, make_snapshot
@contextmanager
def _patched_config(configurable: dict | None):
"""Patch ``langgraph.config.get_config`` to expose ``configurable``."""
import langgraph.config as _lg_cfg
if configurable is None:
with patch.object(
_lg_cfg,
"get_config",
side_effect=RuntimeError("Called get_config outside of a runnable context"),
):
yield
else:
with patch.object(
_lg_cfg,
"get_config",
return_value={"configurable": configurable},
):
yield
@pytest.fixture
def store(tmp_path):
return make_active_store(tmp_path / "model-runtime")
@pytest.fixture
def runtime(store):
return SnapshotRuntime(store)
def _request(messages, tools=()):
return SimpleNamespace(messages=messages, tools=list(tools))
def _configurable_for(snapshot, **overrides):
configurable = {
"runtime_snapshot_id": snapshot.snapshot_id,
"thread_id": snapshot.thread_id,
}
configurable.update(overrides)
return configurable
def test_text_counter_excludes_attachment_payloads_and_counts_tool_results():
@@ -31,49 +90,236 @@ def test_text_counter_excludes_attachment_payloads_and_counts_tool_results():
assert count_message_text_tokens(messages) == 2
def test_reserve_policy_uses_fixed_overhead_without_request_token_counting():
policy = MessageReservePolicy()
assert policy.hard_budget(32_768) == 17_204
assert policy.soft_budget(32_768) == 12_042
assert policy.keep_budget(32_768) == 6_021
assert policy.hard_budget(32_768, has_tools=True) == 9_012
# =============================================================================
# Snapshot-driven budget recomputation (section 6.5)
# =============================================================================
def test_budget_middleware_uses_message_threshold_and_safe_tool_cutoff():
class TestSnapshotMessageBudget:
"""The frozen fixture registry resolves to:
resolved_input_limit = 1048576 - 32768 = 1015808
reserves: system 4096, tools 8192, attachments 4096
"""
def test_base_mode_deducts_only_system_reserve(self, store):
snapshot = make_snapshot(store)
budget = _snapshot_message_budget(
snapshot, has_tools=False, has_attachments=False
)
message_budget = 1015808 - 4096
assert budget.input_limit == 1015808
assert budget.hard_tokens == int(message_budget * 0.90)
assert budget.soft_tokens == int(budget.hard_tokens * 0.70)
assert budget.keep_tokens == int(budget.hard_tokens * 0.35)
def test_tools_and_attachments_deduct_their_reserves(self, store):
snapshot = make_snapshot(store)
tools_only = _snapshot_message_budget(
snapshot, has_tools=True, has_attachments=False
)
full = _snapshot_message_budget(snapshot, has_tools=True, has_attachments=True)
assert tools_only.hard_tokens == int((1015808 - 4096 - 8192) * 0.90)
assert full.hard_tokens == int((1015808 - 4096 - 8192 - 4096) * 0.90)
class TestSnapshotModeMiddleware:
def test_budget_uses_frozen_reserves_not_model_profile(self, store, runtime):
"""A misleading compile-time profile must not affect snapshot mode."""
model = MagicMock()
model.profile = {"max_input_tokens": 32_768}
middleware = create_message_budget_middleware(
model, MagicMock(), runtime=runtime
)
snapshot = make_snapshot(store)
with _patched_config(_configurable_for(snapshot)):
budget = middleware._budget_for_request(_request([HumanMessage("hi")]))
assert budget.input_limit == 1015808
assert budget.has_tools is True # conservative construction-time flag
assert budget.has_attachments is False
assert budget.hard_tokens == int((1015808 - 4096 - 8192) * 0.90)
def test_recomputed_per_call_with_current_attachments(self, store, runtime):
middleware = create_message_budget_middleware(
MagicMock(), MagicMock(), runtime=runtime
)
snapshot = make_snapshot(store)
plain = _request([HumanMessage(content="plain")])
attached = _request(
[HumanMessage(content=[{"type": "image", "url": "https://example.test/a"}])]
)
with _patched_config(_configurable_for(snapshot)):
first = middleware._budget_for_request(plain)
second = middleware._budget_for_request(attached)
third = middleware._budget_for_request(plain)
assert first.has_attachments is False
assert second.has_attachments is True
assert second.hard_tokens == int((1015808 - 4096 - 8192 - 4096) * 0.90)
# No caching across calls: the third call recomputes from scratch.
assert third == first
def test_has_tools_is_conservative_not_request_derived(self, store, runtime):
"""Tool mode comes from construction, not the request's bound tools."""
middleware = create_message_budget_middleware(
MagicMock(), MagicMock(), has_tools=True, runtime=runtime
)
snapshot = make_snapshot(store)
with _patched_config(_configurable_for(snapshot)):
budget = middleware._budget_for_request(_request([HumanMessage("hi")]))
# The request binds no tools at all, yet the tools reserve applies.
assert budget.has_tools is True
assert budget.hard_tokens == int((1015808 - 4096 - 8192) * 0.90)
def test_has_tools_false_when_agent_has_no_toolset(self, store, runtime):
middleware = create_message_budget_middleware(
MagicMock(), MagicMock(), has_tools=False, runtime=runtime
)
snapshot = make_snapshot(store)
with _patched_config(_configurable_for(snapshot)):
budget = middleware._budget_for_request(
_request([HumanMessage("hi")], tools=[{"name": "read_file"}])
)
assert budget.has_tools is False
assert budget.hard_tokens == int((1015808 - 4096) * 0.90)
def test_snapshot_binding_is_verified(self, store, runtime):
from EvoScientist.model_registry.errors import (
SNAPSHOT_NOT_FOUND,
ModelRegistryError,
)
middleware = create_message_budget_middleware(
MagicMock(), MagicMock(), runtime=runtime
)
snapshot = make_snapshot(store)
with (
_patched_config(_configurable_for(snapshot, thread_id="other")),
pytest.raises(ModelRegistryError) as excinfo,
):
middleware._budget_for_request(_request([HumanMessage("hi")]))
assert excinfo.value.code == SNAPSHOT_NOT_FOUND
def test_summarizer_model_uses_summary_role(self, store, runtime):
middleware = create_message_budget_middleware(
MagicMock(), MagicMock(), runtime=runtime
)
snapshot = make_snapshot(store)
with _patched_config(_configurable_for(snapshot)):
model = middleware.model
# summary → snapshot.auxiliary ?? snapshot.primary (section 6.1)
assert isinstance(model, ChatOllama)
assert model.model == "qwen3"
def test_summarizer_model_cached_per_snapshot(self, store, runtime):
middleware = create_message_budget_middleware(
MagicMock(), MagicMock(), runtime=runtime
)
snapshot = make_snapshot(store)
with _patched_config(_configurable_for(snapshot)):
assert middleware.model is middleware.model
# =============================================================================
# Interim local mode (no snapshot ID in configurable)
# =============================================================================
class TestInterimMode:
def test_budget_from_compile_time_profile(self, runtime):
model = MagicMock()
model.profile = {"max_input_tokens": 32_768}
middleware = create_message_budget_middleware(
model, MagicMock(), has_tools=False, runtime=runtime
)
with _patched_config({}):
budget = middleware._budget_for_request(_request([HumanMessage("hi")]))
policy = MessageReservePolicy()
assert budget.hard_tokens == policy.hard_budget(32_768, has_tools=False)
def test_rejects_when_reserves_exceed_model_limit(self, runtime):
model = MagicMock()
model.profile = {
"max_input_tokens": 28_672,
"min_effective_input_tokens": 4_096,
}
middleware = create_message_budget_middleware(
model, MagicMock(), runtime=runtime
)
request = _request(
[HumanMessage(content=[{"type": "image", "url": "https://x.test/a"}])]
)
with (
_patched_config({}),
pytest.raises(
ContextBudgetUnsatisfiableError, match="CONTEXT_BUDGET_UNSATISFIABLE"
),
):
middleware._budget_for_request(request)
def test_summarizer_falls_back_to_compile_time_model(self, runtime):
fallback = MagicMock()
middleware = create_message_budget_middleware(
fallback, MagicMock(), runtime=runtime
)
with _patched_config({}):
assert middleware.model is fallback
def test_outside_runnable_context_uses_interim_mode(self, runtime):
fallback = MagicMock()
middleware = create_message_budget_middleware(
fallback, MagicMock(), has_tools=False, runtime=runtime
)
with _patched_config(None):
budget = middleware._budget_for_request(_request([HumanMessage("hi")]))
assert budget.input_limit == 32_768
# =============================================================================
# Cutoff behavior (budget-driven summarization trigger)
# =============================================================================
def test_budget_middleware_uses_message_threshold_and_safe_tool_cutoff(runtime):
model = MagicMock()
model.profile = {"max_input_tokens": 32_768}
middleware = create_message_budget_middleware(model, MagicMock())
middleware = create_message_budget_middleware(
model, MagicMock(), has_tools=False, runtime=runtime
)
messages = [
HumanMessage(content="a" * 30_000),
AIMessage(content="", tool_calls=[{"name": "read_file", "args": {}, "id": "tool-1"}]),
AIMessage(
content="", tool_calls=[{"name": "read_file", "args": {}, "id": "tool-1"}]
),
ToolMessage(content="b" * 30_000, tool_call_id="tool-1"),
HumanMessage(content="c" * 30_000),
]
total = count_message_text_tokens(messages)
assert middleware._should_summarize(messages, total) is True
cutoff = middleware._determine_cutoff_index(messages)
from EvoScientist.middleware.message_budget import _ACTIVE_BUDGET
with _patched_config({}):
# Simulate the per-call budget activation that wrap_model_call sets.
token = _ACTIVE_BUDGET.set(middleware._budget_for_request(_request(messages)))
try:
assert middleware._should_summarize(messages, total) is True
cutoff = middleware._determine_cutoff_index(messages)
finally:
_ACTIVE_BUDGET.reset(token)
assert cutoff in {1, 3}
# A cutoff never leaves the tool response without its matching AI tool call.
if cutoff == 1:
assert isinstance(messages[cutoff], AIMessage)
def test_budget_rejects_tools_and_attachments_when_reserves_exceed_model_limit():
model = MagicMock()
model.profile = {
"max_input_tokens": 28_672,
"min_effective_input_tokens": 4_096,
}
middleware = create_message_budget_middleware(model, MagicMock())
request = SimpleNamespace(
messages=[
HumanMessage(content=[{"type": "image", "url": "https://example.test/a"}])
],
tools=[{"name": "read_file"}],
)
with pytest.raises(ContextBudgetUnsatisfiableError, match="CONTEXT_BUDGET_UNSATISFIABLE"):
middleware._budget_for_request(request)
-322
View File
@@ -1,322 +0,0 @@
"""Tests for the model fallback middleware.
Covers error classification (_is_non_fallbackable) and the end-to-end
fallback chain behaviour via _try_fallbacks / _guard_and_fallback.
"""
from __future__ import annotations
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from langchain_core.exceptions import ContextOverflowError
from langchain_core.messages import AIMessage, HumanMessage
from EvoScientist.middleware.model_fallback import (
_guard_and_fallback,
_is_non_fallbackable,
_try_fallbacks,
add_fallback,
clear_fallbacks,
set_ui_emit,
)
# ── Helpers ──────────────────────────────────────────────────────
def _fake_request():
"""Build a minimal ModelRequest stub with an .override() method."""
req = MagicMock()
req.override = MagicMock(side_effect=lambda **kw: req)
req.messages = [HumanMessage(content="hi")]
return req
AI_RESPONSE = AIMessage(content="ok")
@pytest.fixture(autouse=True)
def _clean_chain():
"""Ensure a clean fallback chain and no UI callback for every test."""
clear_fallbacks()
set_ui_emit(None)
yield
clear_fallbacks()
set_ui_emit(None)
# ═════════════════════════════════════════════════════════════════
# 1. _is_non_fallbackable — error classification
# ═════════════════════════════════════════════════════════════════
class TestIsNonFallbackable:
"""Verify which errors block fallback and which allow it."""
# ── Context-length errors: must NOT fallback ────────────────
def test_context_overflow_error_instance(self):
exc = ContextOverflowError("too long")
assert _is_non_fallbackable(exc) == "context length exceeded"
@pytest.mark.parametrize(
"msg",
[
"Error 400: context_length_exceeded",
"Bad Request: context length exceeded in prompt",
"400 too many tokens for this model",
"Bad Request: maximum context length is 128k",
"Error 400: output too large",
"400 Bad Request: context_window_exceeded",
"400: string_too_long",
"Bad Request: max_tokens_exceeded",
],
)
def test_context_limit_400_patterns(self, msg):
assert _is_non_fallbackable(Exception(msg)) == "context length exceeded"
# ── Malformed request errors: must NOT fallback ─────────────
@pytest.mark.parametrize(
"msg",
[
"Error 400: invalid_request_error",
"400 Bad Request: invalid request body",
"400: malformed JSON in request",
],
)
def test_malformed_request_400_patterns(self, msg):
assert (
_is_non_fallbackable(Exception(msg))
== "malformed request (client-side error)"
)
# ── Auth errors: SHOULD fallback (different provider may work) ──
@pytest.mark.parametrize(
"msg",
[
"400 Bad Request: invalid_api_key",
"400: authentication failed",
"400 Bad Request: permission denied",
],
)
def test_auth_errors_are_fallbackable(self, msg):
assert _is_non_fallbackable(Exception(msg)) is None
# ── Server / transient errors: SHOULD fallback ──────────────
@pytest.mark.parametrize(
"msg",
[
"Error 500: internal server error",
"429 Too Many Requests: rate limit exceeded",
"503 Service Unavailable",
"Connection timed out",
"HTTPSConnectionPool: Read timed out",
"502 Bad Gateway",
"overloaded_error: the server is temporarily overloaded",
],
)
def test_server_errors_are_fallbackable(self, msg):
assert _is_non_fallbackable(Exception(msg)) is None
# ── Edge: 400 without a known pattern → fallbackable ────────
def test_400_unknown_pattern_is_fallbackable(self):
assert _is_non_fallbackable(Exception("400: unknown_field 'foo'")) is None
# ── Edge: pattern present but no 400 → fallbackable ─────────
def test_context_pattern_without_400_is_fallbackable(self):
exc = Exception("context_length_exceeded (warning only)")
assert _is_non_fallbackable(exc) is None
def test_malformed_pattern_without_400_is_fallbackable(self):
exc = Exception("invalid_request_error logged for debugging")
assert _is_non_fallbackable(exc) is None
# ═════════════════════════════════════════════════════════════════
# 2. _try_fallbacks — chain walk behaviour
# ═════════════════════════════════════════════════════════════════
class TestTryFallbacks:
"""End-to-end tests for the fallback chain traversal."""
async def test_first_fallback_succeeds(self):
"""When the first fallback model works, return its response."""
add_fallback("fb-model", "fb-provider")
req = _fake_request()
invoke = AsyncMock(return_value=AI_RESPONSE)
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = await _try_fallbacks(req, invoke, Exception("503 boom"))
assert result is AI_RESPONSE
invoke.assert_awaited_once()
mock_gcm.assert_called_once_with(model="fb-model", provider="fb-provider")
async def test_skips_failing_fallback_tries_next(self):
"""When the first fallback fails, try the second."""
add_fallback("fb-bad", "prov-a")
add_fallback("fb-good", "prov-b")
req = _fake_request()
call_count = 0
async def _invoke(r):
nonlocal call_count
call_count += 1
if call_count == 1:
raise Exception("429 rate limited")
return AI_RESPONSE
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = await _try_fallbacks(req, _invoke, Exception("503 boom"))
assert result is AI_RESPONSE
assert call_count == 2
async def test_all_fallbacks_exhausted_raises_last(self):
"""When every fallback fails, re-raise the last exception."""
add_fallback("fb-a", "prov-a")
add_fallback("fb-b", "prov-b")
req = _fake_request()
last_error = Exception("429 from fb-b")
call_count = 0
async def _invoke(r):
nonlocal call_count
call_count += 1
if call_count == 1:
raise Exception("500 from fb-a")
raise last_error
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
with pytest.raises(Exception, match="429 from fb-b") as exc_info:
await _try_fallbacks(req, _invoke, Exception("503 primary"))
assert exc_info.value is last_error
async def test_non_fallbackable_in_chain_aborts_immediately(self):
"""A non-fallbackable error from a fallback model aborts the chain."""
add_fallback("fb-a", "prov-a")
add_fallback("fb-b", "prov-b") # should never be reached
req = _fake_request()
async def _invoke(r):
raise Exception("400: context_length_exceeded")
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
with pytest.raises(Exception, match="context_length_exceeded"):
await _try_fallbacks(req, _invoke, Exception("503 primary"))
# get_chat_model should only have been called once (for fb-a),
# fb-b should never be reached.
assert mock_gcm.call_count == 1
# ═════════════════════════════════════════════════════════════════
# 3. _guard_and_fallback — pre-check before chain walk
# ═════════════════════════════════════════════════════════════════
class TestGuardAndFallback:
"""Verify that non-fallbackable errors are re-raised before trying the chain."""
async def test_context_overflow_raises_immediately(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock()
with pytest.raises(ContextOverflowError):
await _guard_and_fallback(ContextOverflowError("overflow"), req, invoke)
invoke.assert_not_awaited()
async def test_malformed_400_raises_immediately(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock()
with pytest.raises(Exception, match="invalid_request_error"):
await _guard_and_fallback(
Exception("400: invalid_request_error"), req, invoke
)
invoke.assert_not_awaited()
async def test_server_error_proceeds_to_fallback(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock(return_value=AI_RESPONSE)
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = await _guard_and_fallback(Exception("503 overloaded"), req, invoke)
assert result is AI_RESPONSE
invoke.assert_awaited_once()
async def test_auth_error_proceeds_to_fallback(self):
"""Auth errors should try the fallback chain (different provider)."""
add_fallback("fb", "other-prov")
req = _fake_request()
invoke = AsyncMock(return_value=AI_RESPONSE)
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
result = await _guard_and_fallback(
Exception("400 Bad Request: invalid_api_key"), req, invoke
)
assert result is AI_RESPONSE
invoke.assert_awaited_once()
# ═════════════════════════════════════════════════════════════════
# 4. UI emit callback
# ═════════════════════════════════════════════════════════════════
class TestUiEmit:
"""Verify that fallback events are surfaced via the registered callback."""
async def test_emit_captures_messages(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock(return_value=AI_RESPONSE)
messages: list[tuple[str, str]] = []
set_ui_emit(lambda text, style: messages.append((text, style)))
with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm:
mock_gcm.return_value = MagicMock()
await _try_fallbacks(req, invoke, Exception("503 down"))
texts = [t for t, _ in messages]
assert any("Primary model failed" in t for t in texts)
assert any("Falling back to fb (prov)" in t for t in texts)
assert any("succeeded" in t for t in texts)
async def test_emit_shows_non_fallbackable_rejection(self):
add_fallback("fb", "prov")
req = _fake_request()
invoke = AsyncMock()
messages: list[tuple[str, str]] = []
set_ui_emit(lambda text, style: messages.append((text, style)))
with pytest.raises(ContextOverflowError):
await _guard_and_fallback(ContextOverflowError("overflow"), req, invoke)
texts = [t for t, _ in messages]
assert any("not eligible for fallback" in t for t in texts)
+185 -277
View File
@@ -1,17 +1,21 @@
"""Tests for the deepagents model-passthrough patch.
"""Tests for the deepagents async-task runs.create patch.
Verifies that ``_patch_deepagents_model_passthrough`` wraps
``_build_start_tool`` / ``_build_update_tool`` so that ``client.runs.create``
calls inside the launched async-task tools carry
``config={"configurable": {"model": ..., "model_provider": ...}}``,
without affecting other client methods (``threads.create``, ``runs.get``,
``runs.cancel``).
calls inside the launched async-task tools inherit workspace scope and usage
correlation context — and that NO model configuration is injected: model
resolution is snapshot-driven (``runtime_snapshot_id``), and injecting
``model``/``model_provider`` would be rejected with
``MODEL_CONFIG_OUTSIDE_SNAPSHOT`` (design doc 8.2/8.3).
Other client methods (``threads.create``, ``runs.get``, ``runs.cancel``)
must pass through unmodified.
"""
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import AsyncMock, MagicMock
import pytest
@@ -22,11 +26,6 @@ from EvoScientist.llm import patches as patches_mod
# =============================================================================
def _stub_cfg(model: str = "claude-sonnet-4-6", provider: str = "anthropic"):
"""Build a stand-in for ``_ensure_config()`` return value."""
return SimpleNamespace(model=model, provider=provider)
def _make_client_cache(
*,
create_run_id: str = "run-001",
@@ -72,6 +71,41 @@ def _runtime_stub():
return SimpleNamespace(tool_call_id="tc-001", state={})
def _agent_map(name: str = "writing-agent") -> dict:
return {
name: {
"name": name,
"description": "Draft paper",
"graph_id": name,
"url": "http://localhost:6174",
}
}
_SCOPE = {
"workspace_scope_id": "11111111-1111-1111-1111-111111111111",
"workspace_scope_owner_id": "22222222-2222-2222-2222-222222222222",
"workspace_scope_revision": 3,
"workspace_deployment_id": "local",
"thread_id": "thread-001",
}
def _set_runnable_config(configurable: dict, metadata: dict | None = None):
"""Populate the real langgraph contextvar the patch reads."""
from langchain_core.runnables.config import var_child_runnable_config
return var_child_runnable_config.set(
{"configurable": configurable, "metadata": metadata or {}}
)
def _reset_runnable_config(token) -> None:
from langchain_core.runnables.config import var_child_runnable_config
var_child_runnable_config.reset(token)
# =============================================================================
# 1. Idempotence
# =============================================================================
@@ -106,14 +140,14 @@ class TestIdempotence:
# =============================================================================
# 2. start_async_task injects config into runs.create
# 2. No model configuration is injected (design doc 8.2)
# =============================================================================
class TestStartAsyncTaskInjection:
"""Sync and async start_async_task must inject configurable.model."""
class TestNoModelInjection:
"""runs.create must never carry model/model_provider overrides."""
def test_sync_start_injects_config(self, restore_model_passthrough_patch):
def test_sync_start_injects_no_model_keys(self, restore_model_passthrough_patch):
try:
from deepagents.middleware import async_subagents as ds_mod
except ImportError:
@@ -122,36 +156,24 @@ class TestStartAsyncTaskInjection:
patches_mod._patch_deepagents_model_passthrough()
cache, runs_sync, _ = _make_client_cache()
agent_map = {
"writing-agent": {
"name": "writing-agent",
"description": "Draft paper",
"graph_id": "writing-agent",
"url": "http://localhost:6174",
}
}
tool = ds_mod._build_start_tool(agent_map, cache, "desc")
with patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=_stub_cfg(model="gpt-5", provider="openai"),
):
tool.func(
description="hello",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
tool = ds_mod._build_start_tool(_agent_map(), cache, "desc")
tool.func(
description="hello",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
runs_sync.create.assert_called_once()
kwargs = runs_sync.create.call_args.kwargs
assert kwargs["thread_id"] == "thread-001"
assert kwargs["assistant_id"] == "writing-agent"
assert kwargs["config"] == {
"configurable": {"model": "gpt-5", "model_provider": "openai"}
}
configurable = (kwargs.get("config") or {}).get("configurable") or {}
assert "model" not in configurable
assert "model_provider" not in configurable
async def test_async_start_injects_config(self, restore_model_passthrough_patch):
async def test_async_start_injects_no_model_keys(
self, restore_model_passthrough_patch
):
try:
from deepagents.middleware import async_subagents as ds_mod
except ImportError:
@@ -160,153 +182,21 @@ class TestStartAsyncTaskInjection:
patches_mod._patch_deepagents_model_passthrough()
cache, _, runs_async = _make_client_cache()
agent_map = {
"writing-agent": {
"name": "writing-agent",
"description": "Draft paper",
"graph_id": "writing-agent",
"url": "http://localhost:6174",
}
}
tool = ds_mod._build_start_tool(agent_map, cache, "desc")
with patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=_stub_cfg(model="claude-haiku-4-5", provider="anthropic"),
):
await tool.coroutine(
description="hi",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
runs_async.create.assert_awaited_once()
kwargs = runs_async.create.call_args.kwargs
assert kwargs["config"] == {
"configurable": {
"model": "claude-haiku-4-5",
"model_provider": "anthropic",
}
}
# =============================================================================
# 3. Live config read at tool-call time (post-/model switch behavior)
# =============================================================================
class TestLiveConfigRead:
"""The patch must read cfg fresh on every tool call, not at patch time."""
def test_two_calls_reflect_separate_cfg(self, restore_model_passthrough_patch):
try:
from deepagents.middleware import async_subagents as ds_mod
except ImportError:
pytest.skip("deepagents not available")
patches_mod._patch_deepagents_model_passthrough()
cache, runs_sync, _ = _make_client_cache()
agent_map = {
"writing-agent": {
"name": "writing-agent",
"description": "x",
"graph_id": "writing-agent",
"url": "http://localhost:6174",
}
}
tool = ds_mod._build_start_tool(agent_map, cache, "desc")
with patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=_stub_cfg(model="model-a", provider="anthropic"),
):
tool.func(
description="t1",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
first_kwargs = runs_sync.create.call_args.kwargs
with patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=_stub_cfg(model="model-b", provider="openai"),
):
tool.func(
description="t2",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
second_kwargs = runs_sync.create.call_args.kwargs
assert first_kwargs["config"]["configurable"]["model"] == "model-a"
assert second_kwargs["config"]["configurable"]["model"] == "model-b"
# =============================================================================
# 4. update_async_task also injects config
# =============================================================================
class TestUpdateAsyncTaskInjection:
"""update_async_task must inject config too — not just start."""
def _tracked_task(self, agent_name: str = "writing-agent") -> dict:
return {
"task_id": "thread-001",
"agent_name": agent_name,
"thread_id": "thread-001",
"run_id": "old-run",
"status": "running",
"created_at": "2026-05-07T00:00:00Z",
"last_checked_at": "2026-05-07T00:00:00Z",
"last_updated_at": "2026-05-07T00:00:00Z",
}
async def test_async_update_injects_config(self, restore_model_passthrough_patch):
"""The async coroutine path must inject config too."""
try:
from deepagents.middleware import async_subagents as ds_mod
except ImportError:
pytest.skip("deepagents not available")
patches_mod._patch_deepagents_model_passthrough()
cache, _, runs_async = _make_client_cache()
agent_map = {
"writing-agent": {
"name": "writing-agent",
"description": "x",
"graph_id": "writing-agent",
"url": "http://localhost:6174",
}
}
runtime = SimpleNamespace(
tool_call_id="tc-002",
state={"async_tasks": {"thread-001": self._tracked_task()}},
tool = ds_mod._build_start_tool(_agent_map(), cache, "desc")
await tool.coroutine(
description="hi",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
tool = ds_mod._build_update_tool(agent_map, cache)
with patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=_stub_cfg(model="gpt-5", provider="openai"),
):
await tool.coroutine(
task_id="thread-001",
message="follow up async",
runtime=runtime,
)
runs_async.create.assert_awaited_once()
kwargs = runs_async.create.call_args.kwargs
assert kwargs["config"] == {
"configurable": {"model": "gpt-5", "model_provider": "openai"}
}
assert kwargs.get("multitask_strategy") == "interrupt"
configurable = (runs_async.create.call_args.kwargs.get("config") or {}).get(
"configurable"
) or {}
assert "model" not in configurable
assert "model_provider" not in configurable
def test_sync_update_injects_config(self, restore_model_passthrough_patch):
def test_sync_update_injects_no_model_keys(self, restore_model_passthrough_patch):
try:
from deepagents.middleware import async_subagents as ds_mod
except ImportError:
@@ -315,15 +205,6 @@ class TestUpdateAsyncTaskInjection:
patches_mod._patch_deepagents_model_passthrough()
cache, runs_sync, _ = _make_client_cache()
agent_map = {
"writing-agent": {
"name": "writing-agent",
"description": "x",
"graph_id": "writing-agent",
"url": "http://localhost:6174",
}
}
# update_async_task reads tracked task from runtime.state
tracked_task = {
"task_id": "thread-001",
"agent_name": "writing-agent",
@@ -339,26 +220,94 @@ class TestUpdateAsyncTaskInjection:
state={"async_tasks": {"thread-001": tracked_task}},
)
tool = ds_mod._build_update_tool(agent_map, cache)
with patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=_stub_cfg(model="gpt-5", provider="openai"),
):
tool.func(
task_id="thread-001",
message="follow up",
runtime=runtime,
)
tool = ds_mod._build_update_tool(_agent_map(), cache)
tool.func(
task_id="thread-001",
message="follow up",
runtime=runtime,
)
runs_sync.create.assert_called_once()
kwargs = runs_sync.create.call_args.kwargs
assert kwargs["config"]["configurable"]["model"] == "gpt-5"
assert kwargs["config"]["configurable"]["model_provider"] == "openai"
configurable = (kwargs.get("config") or {}).get("configurable") or {}
assert "model" not in configurable
assert "model_provider" not in configurable
# update preserves multitask_strategy="interrupt" — verify no regression
assert kwargs.get("multitask_strategy") == "interrupt"
# =============================================================================
# 3. Workspace scope inheritance is preserved
# =============================================================================
class TestScopeInheritance:
def test_scope_keys_forwarded(self, restore_model_passthrough_patch):
try:
from deepagents.middleware import async_subagents as ds_mod
except ImportError:
pytest.skip("deepagents not available")
patches_mod._patch_deepagents_model_passthrough()
cache, runs_sync, _ = _make_client_cache()
tool = ds_mod._build_start_tool(_agent_map(), cache, "desc")
token = _set_runnable_config(dict(_SCOPE))
try:
tool.func(
description="hello",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
finally:
_reset_runnable_config(token)
configurable = runs_sync.create.call_args.kwargs["config"]["configurable"]
for key in (
"workspace_scope_id",
"workspace_scope_owner_id",
"workspace_scope_revision",
"workspace_deployment_id",
):
assert configurable[key] == _SCOPE[key]
# =============================================================================
# 4. Usage correlation metadata is preserved
# =============================================================================
class TestUsageMetadata:
def test_usage_scope_tagged(self, restore_model_passthrough_patch, monkeypatch):
try:
from deepagents.middleware import async_subagents as ds_mod
except ImportError:
pytest.skip("deepagents not available")
monkeypatch.setenv("EVOSCIENTIST_USAGE_TRACKING", "1")
patches_mod._patch_deepagents_model_passthrough()
cache, runs_sync, _ = _make_client_cache()
tool = ds_mod._build_start_tool(_agent_map(), cache, "desc")
token = _set_runnable_config(
{"thread_id": "thread-001"}, metadata={"turn_id": "turn-9"}
)
try:
tool.func(
description="hello",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
finally:
_reset_runnable_config(token)
metadata = runs_sync.create.call_args.kwargs["metadata"]
assert metadata["usage_scope"] == "async_subagent"
assert metadata["turn_id"] == "turn-9"
# =============================================================================
# 5. Other client methods are unaffected
# =============================================================================
@@ -382,39 +331,24 @@ class TestNonInterceptedMethods:
cache, _runs_sync, _ = _make_client_cache()
threads_create = cache.get_sync.return_value.threads.create
agent_map = {
"writing-agent": {
"name": "writing-agent",
"description": "x",
"graph_id": "writing-agent",
"url": "http://localhost:6174",
}
}
tool = ds_mod._build_start_tool(agent_map, cache, "desc")
with patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=_stub_cfg(),
):
tool.func(
description="hi",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
tool = ds_mod._build_start_tool(_agent_map(), cache, "desc")
tool.func(
description="hi",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
# threads.create() is called with no kwargs (deepagents pattern).
threads_create.assert_called_once_with()
# =============================================================================
# 6. Empty cfg → no config kwarg added
# 6. No context to inherit → no config kwarg added
# =============================================================================
class TestEmptyCfg:
"""If neither model nor provider is set, don't inject anything."""
def test_empty_cfg_no_config_kwarg(self, restore_model_passthrough_patch):
class TestNoContextNoConfig:
def test_no_scope_no_config_kwarg(self, restore_model_passthrough_patch):
try:
from deepagents.middleware import async_subagents as ds_mod
except ImportError:
@@ -423,29 +357,16 @@ class TestEmptyCfg:
patches_mod._patch_deepagents_model_passthrough()
cache, runs_sync, _ = _make_client_cache()
agent_map = {
"writing-agent": {
"name": "writing-agent",
"description": "x",
"graph_id": "writing-agent",
"url": "http://localhost:6174",
}
}
tool = ds_mod._build_start_tool(agent_map, cache, "desc")
with patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=SimpleNamespace(model=None, provider=None),
):
tool.func(
description="hi",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
tool = ds_mod._build_start_tool(_agent_map(), cache, "desc")
tool.func(
description="hi",
subagent_type="writing-agent",
runtime=_runtime_stub(),
)
runs_sync.create.assert_called_once()
kwargs = runs_sync.create.call_args.kwargs
# No config kwarg should be added when there's nothing to override.
# No config kwarg should be added when there's nothing to inherit.
assert "config" not in kwargs
@@ -455,29 +376,22 @@ class TestEmptyCfg:
class TestPreserveExistingConfig:
"""If a caller already supplied config.configurable.X, our merge keeps it."""
"""If a caller already supplied config.configurable.X, the merge keeps it."""
def test_existing_configurable_preserved(self):
"""Direct unit test of the merge helper (integration covered above)."""
with patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=_stub_cfg(model="gpt-5", provider="openai"),
):
merged = patches_mod._merge_runs_config_kwargs(
{
"thread_id": "t1",
"config": {
"configurable": {"thread_id": "outer-t", "extra": 42},
"tags": ["debug"],
},
}
)
merged = patches_mod._merge_runs_config_kwargs(
{
"thread_id": "t1",
"config": {
"configurable": {"thread_id": "outer-t", "extra": 42},
"tags": ["debug"],
},
}
)
assert merged["thread_id"] == "t1"
assert merged["config"]["tags"] == ["debug"]
assert merged["config"]["configurable"]["thread_id"] == "outer-t"
assert merged["config"]["configurable"]["extra"] == 42
assert merged["config"]["configurable"]["model"] == "gpt-5"
assert merged["config"]["configurable"]["model_provider"] == "openai"
def test_non_dict_config_replaced(self):
"""Non-dict ``config`` (e.g. a Pydantic RunnableConfig) is replaced.
@@ -494,16 +408,10 @@ class TestPreserveExistingConfig:
class _Sentinel:
"""Stand-in for any non-dict config-shaped object."""
with patch(
"EvoScientist.EvoScientist._ensure_config",
return_value=_stub_cfg(model="gpt-5", provider="openai"),
):
merged = patches_mod._merge_runs_config_kwargs(
{"thread_id": "t1", "config": _Sentinel()}
)
merged = patches_mod._merge_runs_config_kwargs(
{"thread_id": "t1", "config": _Sentinel()}
)
assert merged["thread_id"] == "t1"
# Non-dict input was replaced with a fresh dict carrying only our
# injected keys.
assert merged["config"] == {
"configurable": {"model": "gpt-5", "model_provider": "openai"}
}
# Non-dict input was replaced with a fresh dict; nothing is injected
# when there is no runnable context to inherit from.
assert merged["config"] == {"configurable": {}}
+211
View File
@@ -0,0 +1,211 @@
"""Tests for ``EvoScientist.model_registry.runtime`` (design doc 8.3).
Covers the runtime-link glue: registry-default role resolution (auxiliary ??
primary), snapshot role construction with per-call credential resolution,
safe-client wiring (ollama carries retries on the transports), the shared
default-runtime accessor, and the tolerant local platform field reader.
"""
from __future__ import annotations
import pytest
from langchain_ollama import ChatOllama
from langchain_openai import ChatOpenAI
from EvoScientist.model_registry.endpoint_policy import EndpointPolicy
from EvoScientist.model_registry.errors import (
MODEL_REGISTRY_NOT_READY,
ModelRegistryError,
)
from EvoScientist.model_registry.platform import (
DEFAULT_LOCAL_DEPLOYMENT_ID,
PlatformConfigError,
)
from EvoScientist.model_registry.runtime import (
SnapshotRuntime,
_read_local_platform_fields,
get_snapshot_runtime,
set_snapshot_runtime_for_tests,
)
from EvoScientist.model_registry.schemas import DevelopmentEndpoint
from EvoScientist.model_registry.store import ModelRuntimeStore
from tests.registry_fixtures import (
OLLAMA_REF,
ZHIPU_REF,
ZHIPU_SECRET,
make_active_store,
make_snapshot,
)
@pytest.fixture
def store(tmp_path):
return make_active_store(tmp_path / "model-runtime")
@pytest.fixture
def runtime(store):
return SnapshotRuntime(store)
@pytest.fixture(autouse=True)
def _reset_default_runtime():
set_snapshot_runtime_for_tests(None)
yield
set_snapshot_runtime_for_tests(None)
class TestRegistryDefaults:
def test_bootstrap_registry_is_not_ready(self, tmp_path):
runtime = SnapshotRuntime(ModelRuntimeStore(config_dir=tmp_path / "db"))
with pytest.raises(ModelRegistryError) as excinfo:
runtime.registry_defaults()
assert excinfo.value.code == MODEL_REGISTRY_NOT_READY
assert excinfo.value.http_status == 422
def test_active_registry_returns_defaults(self, runtime):
primary, auxiliary, revision = runtime.registry_defaults()
assert primary == ZHIPU_REF
assert auxiliary == OLLAMA_REF
assert revision == 2
class TestResolveRoleConfig:
def test_primary_role_resolves_primary_default(self, runtime):
config = runtime.resolve_role_config("primary")
assert config.model_ref == ZHIPU_REF
assert config.role == "primary"
def test_auxiliary_role_resolves_auxiliary_default(self, runtime):
for role in ("auxiliary", "summary", "tool_selector"):
config = runtime.resolve_role_config(role)
assert config.model_ref == OLLAMA_REF, role
assert config.role == "auxiliary"
def test_auxiliary_falls_back_to_primary(self, tmp_path):
store = make_active_store(tmp_path / "db", auxiliary_default=False)
runtime = SnapshotRuntime(store)
config = runtime.resolve_role_config("auxiliary")
assert config.model_ref == ZHIPU_REF
assert config.role == "primary"
class TestBuildDefaultRoleModel:
def test_primary_builds_chat_openai_with_frozen_options(self, runtime):
model = runtime.build_default_role_model("primary")
assert isinstance(model, ChatOpenAI)
assert model.model_name == "glm-5.2"
assert model.openai_api_base == "https://open.bigmodel.cn/api/paas/v4"
assert model.max_retries == 2
assert model.temperature == 0.7
assert model.top_p == 0.95
def test_auxiliary_builds_ollama_without_credential(self, runtime):
model = runtime.build_default_role_model("auxiliary")
assert isinstance(model, ChatOllama)
assert model.model == "qwen3"
def test_openai_compatible_never_reads_env_api_key(self, runtime, monkeypatch):
"""The credential comes from the frozen auth_ref, never OPENAI_API_KEY."""
monkeypatch.setenv("OPENAI_API_KEY", "sk-env-must-not-leak")
model = runtime.build_default_role_model("primary")
assert model.openai_api_key.get_secret_value() == ZHIPU_SECRET
class TestBuildRoleModelFromSnapshot:
def test_primary_role_uses_snapshot_primary(self, store, runtime):
snapshot = make_snapshot(store)
model = runtime.build_role_model(snapshot, "primary")
assert isinstance(model, ChatOpenAI)
assert model.model_name == "glm-5.2"
assert model.openai_api_key.get_secret_value() == ZHIPU_SECRET
def test_auxiliary_roles_use_snapshot_auxiliary(self, store, runtime):
snapshot = make_snapshot(store)
for role in ("auxiliary", "summary", "tool_selector"):
model = runtime.build_role_model(snapshot, role)
assert isinstance(model, ChatOllama), role
assert model.model == "qwen3"
def test_auxiliary_roles_fall_back_to_snapshot_primary(self, tmp_path):
store = make_active_store(tmp_path / "db", auxiliary_default=False)
runtime = SnapshotRuntime(store)
snapshot = make_snapshot(store)
model = runtime.build_role_model(snapshot, "summary")
assert isinstance(model, ChatOpenAI)
assert model.model_name == "glm-5.2"
class TestSharedDefaultRuntime:
def test_set_and_get_runtime(self, runtime):
set_snapshot_runtime_for_tests(runtime)
assert get_snapshot_runtime() is runtime
def test_builds_from_missing_config_yaml(self, tmp_path, monkeypatch):
monkeypatch.setattr(
"EvoScientist.model_registry.runtime.get_config_path",
lambda: tmp_path / "config.yaml",
)
monkeypatch.setattr(
"EvoScientist.model_registry.store.DEFAULT_CONFIG_DIR",
tmp_path / "config-dir",
)
runtime = get_snapshot_runtime()
assert runtime.local_deployment_id == DEFAULT_LOCAL_DEPLOYMENT_ID
assert get_snapshot_runtime() is runtime
class TestLocalPlatformFields:
def _read(self, tmp_path, monkeypatch, text: str | None):
config_path = tmp_path / "config.yaml"
if text is not None:
config_path.write_text(text, encoding="utf-8")
monkeypatch.setattr(
"EvoScientist.model_registry.runtime.get_config_path", lambda: config_path
)
return _read_local_platform_fields()
def test_missing_config_uses_defaults(self, tmp_path, monkeypatch):
deployment_id, db_path, endpoints = self._read(tmp_path, monkeypatch, None)
assert deployment_id == DEFAULT_LOCAL_DEPLOYMENT_ID
assert db_path is None
assert endpoints == ()
def test_reads_runtime_fields(self, tmp_path, monkeypatch):
deployment_id, db_path, endpoints = self._read(
tmp_path,
monkeypatch,
"local_deployment_id: dev-deploy\n"
"model_runtime_db: /tmp/mr.sqlite3\n"
"development_endpoints:\n"
" - {id: ollama, url: 'http://localhost:11434', label: Local}\n",
)
assert deployment_id == "dev-deploy"
assert db_path == __import__("pathlib").Path("/tmp/mr.sqlite3")
assert endpoints == (
DevelopmentEndpoint(
id="ollama", url="http://localhost:11434", label="Local"
),
)
def test_invalid_yaml_raises(self, tmp_path, monkeypatch):
with pytest.raises(PlatformConfigError):
self._read(tmp_path, monkeypatch, "local_deployment_id: [unclosed")
def test_invalid_field_type_raises(self, tmp_path, monkeypatch):
with pytest.raises(PlatformConfigError):
self._read(tmp_path, monkeypatch, "local_deployment_id: 42")
class TestEndpointPolicyWiring:
def test_development_endpoints_allow_loopback_ollama(self, store):
policy = EndpointPolicy(
(
DevelopmentEndpoint(
id="ollama", url="http://localhost:11434", label="Local"
),
)
)
runtime = SnapshotRuntime(store, endpoint_policy=policy)
model = runtime.build_default_role_model("auxiliary")
assert isinstance(model, ChatOllama)
-1
View File
@@ -151,7 +151,6 @@ def test_async_run_inherits_usage_context_only_when_tracking_enabled(
from EvoScientist.llm import patches
monkeypatch.setenv("EVOSCIENTIST_USAGE_TRACKING", "true")
monkeypatch.setattr(patches, "_read_cfg_configurable", lambda: {})
active = {
"metadata": {
"turn_id": "turn-a",