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:
@@ -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(
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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,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."""
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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})
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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": {}}
|
||||
|
||||
@@ -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)
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user