diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index 3a621a5..9210437 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -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( diff --git a/EvoScientist/cli/tui_interactive.py b/EvoScientist/cli/tui_interactive.py index cdfac47..49767ab 100644 --- a/EvoScientist/cli/tui_interactive.py +++ b/EvoScientist/cli/tui_interactive.py @@ -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 diff --git a/EvoScientist/commands/implementation/__init__.py b/EvoScientist/commands/implementation/__init__.py index 2f748e7..998bf51 100644 --- a/EvoScientist/commands/implementation/__init__.py +++ b/EvoScientist/commands/implementation/__init__.py @@ -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", diff --git a/EvoScientist/commands/implementation/model_fallback.py b/EvoScientist/commands/implementation/model_fallback.py deleted file mode 100644 index 7e21cda..0000000 --- a/EvoScientist/commands/implementation/model_fallback.py +++ /dev/null @@ -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} ", - 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 [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 " - "(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 [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()) diff --git a/EvoScientist/llm/patches.py b/EvoScientist/llm/patches.py index a0839fc..08a3982 100644 --- a/EvoScientist/llm/patches.py +++ b/EvoScientist/llm/patches.py @@ -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": , "model_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 diff --git a/EvoScientist/middleware/__init__.py b/EvoScientist/middleware/__init__.py index c1975bd..550baf3 100644 --- a/EvoScientist/middleware/__init__.py +++ b/EvoScientist/middleware/__init__.py @@ -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", ] diff --git a/EvoScientist/middleware/configurable_model.py b/EvoScientist/middleware/configurable_model.py index 534216a..dd638b2 100644 --- a/EvoScientist/middleware/configurable_model.py +++ b/EvoScientist/middleware/configurable_model.py @@ -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)) diff --git a/EvoScientist/middleware/message_budget.py b/EvoScientist/middleware/message_budget.py index 257c404..dc71319 100644 --- a/EvoScientist/middleware/message_budget.py +++ b/EvoScientist/middleware/message_budget.py @@ -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 + ) diff --git a/EvoScientist/middleware/model_fallback.py b/EvoScientist/middleware/model_fallback.py deleted file mode 100644 index 0dee6ff..0000000 --- a/EvoScientist/middleware/model_fallback.py +++ /dev/null @@ -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) diff --git a/EvoScientist/model_registry/resolver.py b/EvoScientist/model_registry/resolver.py index 2c6be6c..de7d995 100644 --- a/EvoScientist/model_registry/resolver.py +++ b/EvoScientist/model_registry/resolver.py @@ -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. diff --git a/EvoScientist/model_registry/runtime.py b/EvoScientist/model_registry/runtime.py new file mode 100644 index 0000000..1933d96 --- /dev/null +++ b/EvoScientist/model_registry/runtime.py @@ -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 diff --git a/EvoScientist/subagents/_factory.py b/EvoScientist/subagents/_factory.py index 21aeaca..1d0ff54 100644 --- a/EvoScientist/subagents/_factory.py +++ b/EvoScientist/subagents/_factory.py @@ -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( diff --git a/tests/conftest.py b/tests/conftest.py index a6458d7..8e08194 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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.""" diff --git a/tests/registry_fixtures.py b/tests/registry_fixtures.py new file mode 100644 index 0000000..01524a9 --- /dev/null +++ b/tests/registry_fixtures.py @@ -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 diff --git a/tests/test_ask_user.py b/tests/test_ask_user.py index 7f8c494..f1f9787 100644 --- a/tests/test_ask_user.py +++ b/tests/test_ask_user.py @@ -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 # --------------------------------------------------------------------------- diff --git a/tests/test_async_subagent_factory.py b/tests/test_async_subagent_factory.py index f8e0c3c..665a0c2 100644 --- a/tests/test_async_subagent_factory.py +++ b/tests/test_async_subagent_factory.py @@ -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}) diff --git a/tests/test_auxiliary_model.py b/tests/test_auxiliary_model.py index ccb3c5b..d10aa74 100644 --- a/tests/test_auxiliary_model.py +++ b/tests/test_auxiliary_model.py @@ -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 diff --git a/tests/test_cli_completion.py b/tests/test_cli_completion.py index cc6e67f..2750835 100644 --- a/tests/test_cli_completion.py +++ b/tests/test_cli_completion.py @@ -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() diff --git a/tests/test_command_completions.py b/tests/test_command_completions.py index 58193de..f19baf5 100644 --- a/tests/test_command_completions.py +++ b/tests/test_command_completions.py @@ -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: diff --git a/tests/test_command_manager.py b/tests/test_command_manager.py index 80f9ccc..49de965 100644 --- a/tests/test_command_manager.py +++ b/tests/test_command_manager.py @@ -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 diff --git a/tests/test_configurable_model_middleware.py b/tests/test_configurable_model_middleware.py index 6314186..b0d89a8 100644 --- a/tests/test_configurable_model_middleware.py +++ b/tests/test_configurable_model_middleware.py @@ -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) diff --git a/tests/test_message_budget.py b/tests/test_message_budget.py index 4ee324c..84a7da0 100644 --- a/tests/test_message_budget.py +++ b/tests/test_message_budget.py @@ -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) diff --git a/tests/test_model_fallback.py b/tests/test_model_fallback.py deleted file mode 100644 index fa9e346..0000000 --- a/tests/test_model_fallback.py +++ /dev/null @@ -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) diff --git a/tests/test_model_passthrough_patch.py b/tests/test_model_passthrough_patch.py index cafc36c..ee9c74d 100644 --- a/tests/test_model_passthrough_patch.py +++ b/tests/test_model_passthrough_patch.py @@ -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": {}} diff --git a/tests/test_snapshot_runtime.py b/tests/test_snapshot_runtime.py new file mode 100644 index 0000000..140807f --- /dev/null +++ b/tests/test_snapshot_runtime.py @@ -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) diff --git a/tests/test_usage_tracking.py b/tests/test_usage_tracking.py index 006b248..7990233 100644 --- a/tests/test_usage_tracking.py +++ b/tests/test_usage_tracking.py @@ -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",