From 8a0ab179362d2eec7c79e9046c1654ccfb91605f Mon Sep 17 00:00:00 2001 From: m4 Date: Mon, 20 Jul 2026 20:15:38 +0800 Subject: [PATCH] chore: baseline WIP before unified model configuration implementation Pre-existing uncommitted work (runtime snapshots, message budget middleware) preserved as baseline. --- EvoScientist/EvoScientist.py | 14 +- EvoScientist/config/provider_profiles.py | 244 ++++++++++- EvoScientist/langgraph_dev/http.py | 45 ++ EvoScientist/llm/models.py | 178 +++++++- EvoScientist/llm/runtime_snapshots.py | 401 ++++++++++++++++++ EvoScientist/middleware/__init__.py | 8 + EvoScientist/middleware/configurable_model.py | 55 ++- EvoScientist/middleware/message_budget.py | 327 ++++++++++++++ tests/test_langgraph_dev_http.py | 44 ++ tests/test_llm.py | 23 + tests/test_message_budget.py | 79 ++++ tests/test_provider_profiles.py | 48 ++- tests/test_runtime_snapshots.py | 152 +++++++ 13 files changed, 1590 insertions(+), 28 deletions(-) create mode 100644 EvoScientist/llm/runtime_snapshots.py create mode 100644 EvoScientist/middleware/message_budget.py create mode 100644 tests/test_message_budget.py create mode 100644 tests/test_runtime_snapshots.py diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index eb5a717..3a621a5 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -683,6 +683,7 @@ def _get_default_middleware( workspace_dir: str | Path | None = None, cfg=None, chat_model=None, + backend=None, memory_source_agent: str = "EvoScientist", ): """Build the default middleware list. @@ -713,6 +714,7 @@ def _get_default_middleware( create_context_editing_middleware, create_memory_lifecycle_middleware, create_memory_middleware, + create_message_budget_middleware, create_runtime_context_middleware, create_scheduler_middleware, create_tool_selector_middleware, @@ -724,6 +726,13 @@ def _get_default_middleware( 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 + # explicit model/configuration (notably tests and subagent assembly). + # Production graph factories always pass their real composite backend. + from deepagents.backends import StateBackend + + backend = StateBackend() memory_dir = str(_paths_mod.MEMORIES_DIR) source_type = ( MemorySourceType.SUBAGENT if for_async_subagent else MemorySourceType.TURN @@ -770,6 +779,7 @@ def _get_default_middleware( tool_selector_model = get_chat_model(model=aux_model, provider=aux_provider) mw = [ ConfigurableModelMiddleware(), + create_message_budget_middleware(model, backend), create_context_editing_middleware(model), ModelFallbackMiddleware(), ContextOverflowMapperMiddleware(), @@ -857,7 +867,7 @@ def _get_default_agent(): cfg = _ensure_config() be = _get_default_backend() - mw = _get_default_middleware() + mw = _get_default_middleware(backend=be) # HITL on main agent only (mirrors create_cli_agent). Use middleware, # not interrupt_on= kwarg — the kwarg propagates to every subagent and @@ -1022,7 +1032,7 @@ def create_cli_agent( # CLI agent never drifts from the default chain. Anything CLI-specific # (e.g. ``HumanInTheLoopMiddleware``) is appended below. mw: list[AgentMiddleware] = _get_default_middleware( - workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model + workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model, backend=be ) # HITL on main agent only — passing `interrupt_on=` to create_deep_agent diff --git a/EvoScientist/config/provider_profiles.py b/EvoScientist/config/provider_profiles.py index c8a4eb6..687c86c 100644 --- a/EvoScientist/config/provider_profiles.py +++ b/EvoScientist/config/provider_profiles.py @@ -22,7 +22,7 @@ import yaml from .settings import get_config_dir -PROVIDER_PROFILES_VERSION = 2 +PROVIDER_PROFILES_VERSION = 3 SUPPORTED_PROVIDER_ADAPTERS = ( "openai", "anthropic", @@ -79,12 +79,41 @@ class ProviderProfileError(ValueError): """Raised when a provider profile document or selection is invalid.""" +@dataclass(frozen=True) +class ProviderRuntime: + timeout_seconds: int = 120 + max_retries: int = 2 + default_temperature: float | None = None + default_top_p: float | None = None + default_reasoning_effort: str = "auto" + + +@dataclass(frozen=True) +class ModelRuntime: + limit_mode: str | None = None + context_window_tokens: int | None = None + max_input_tokens: int | None = None + max_output_tokens: int = 4096 + min_effective_input_tokens: int = 4096 + limits_status: str = "needs_confirmation" + limits_source: str = "user" + temperature: float | None = None + top_p: float | None = None + reasoning_effort: str = "auto" + capabilities: tuple[tuple[str, bool | str], ...] = ( + ("tools", "auto"), + ("vision", "auto"), + ("structured_output", "auto"), + ) + + @dataclass(frozen=True) class ProviderModel: id: str name: str model_id: str enabled: bool = True + runtime: ModelRuntime = ModelRuntime() @dataclass(frozen=True) @@ -97,6 +126,7 @@ class ProviderProfile: enabled: bool models: tuple[ProviderModel, ...] auth_mode: str = "api_key" + runtime: ProviderRuntime = ProviderRuntime() @dataclass(frozen=True) @@ -151,6 +181,150 @@ def _validate_id(value: str, *, context: str) -> str: return value +def _optional_number( + raw: dict[str, Any], key: str, *, context: str, minimum: float, maximum: float +) -> float | None: + value = raw.get(key) + if value is None: + return None + if isinstance(value, bool) or not isinstance(value, (int, float)): + raise ProviderProfileError(f"{context}.{key} must be a number or null.") + numeric = float(value) + if not minimum <= numeric <= maximum: + raise ProviderProfileError( + f"{context}.{key} must be between {minimum:g} and {maximum:g}." + ) + return numeric + + +def _optional_positive_int(raw: dict[str, Any], key: str, *, context: str) -> int | None: + value = raw.get(key) + if value is None: + return None + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + raise ProviderProfileError(f"{context}.{key} must be a positive integer or null.") + return value + + +def _parse_provider_runtime(raw: Any, *, context: str) -> ProviderRuntime: + if raw is None: + return ProviderRuntime() + if not isinstance(raw, dict): + raise ProviderProfileError(f"{context}.runtime must be an object.") + timeout_seconds = raw.get("timeout_seconds", 120) + max_retries = raw.get("max_retries", 2) + if isinstance(timeout_seconds, bool) or not isinstance(timeout_seconds, int): + raise ProviderProfileError(f"{context}.runtime.timeout_seconds must be an integer.") + if not 10 <= timeout_seconds <= 600: + raise ProviderProfileError( + f"{context}.runtime.timeout_seconds must be between 10 and 600." + ) + if isinstance(max_retries, bool) or not isinstance(max_retries, int): + raise ProviderProfileError(f"{context}.runtime.max_retries must be an integer.") + if not 0 <= max_retries <= 5: + raise ProviderProfileError( + f"{context}.runtime.max_retries must be between 0 and 5." + ) + effort = raw.get("default_reasoning_effort", "auto") + if effort not in {"auto", "low", "medium", "high"}: + raise ProviderProfileError( + f"{context}.runtime.default_reasoning_effort is invalid." + ) + return ProviderRuntime( + timeout_seconds=timeout_seconds, + max_retries=max_retries, + default_temperature=_optional_number( + raw, "default_temperature", context=f"{context}.runtime", minimum=0, maximum=2 + ), + default_top_p=_optional_number( + raw, "default_top_p", context=f"{context}.runtime", minimum=0.000001, maximum=1 + ), + default_reasoning_effort=effort, + ) + + +def resolve_model_input_limit(runtime: ModelRuntime) -> int: + """Return the prompt limit declared by a confirmed model runtime.""" + if runtime.limits_status != "confirmed" or runtime.limit_mode is None: + raise ProviderProfileError("Model limits require confirmation before use.") + if runtime.limit_mode == "combined": + if runtime.context_window_tokens is None: + raise ProviderProfileError("context_window_tokens is required for combined limits.") + return runtime.context_window_tokens - runtime.max_output_tokens + if runtime.limit_mode == "input_only": + if runtime.max_input_tokens is None: + raise ProviderProfileError("max_input_tokens is required for input_only limits.") + return runtime.max_input_tokens + raise ProviderProfileError("limit_mode must be combined or input_only.") + + +def _parse_model_runtime(raw: Any, *, context: str) -> ModelRuntime: + if raw is None: + return ModelRuntime() + if not isinstance(raw, dict): + raise ProviderProfileError(f"{context}.runtime must be an object.") + status = raw.get("limits_status", "needs_confirmation") + if status not in {"confirmed", "needs_confirmation"}: + raise ProviderProfileError(f"{context}.runtime.limits_status is invalid.") + source = raw.get("limits_source", "user") + if source not in {"catalog", "provider", "user"}: + raise ProviderProfileError(f"{context}.runtime.limits_source is invalid.") + limit_mode = raw.get("limit_mode") + if limit_mode is not None and limit_mode not in {"combined", "input_only"}: + raise ProviderProfileError(f"{context}.runtime.limit_mode is invalid.") + max_output = raw.get("max_output_tokens", 4096) + min_effective = raw.get("min_effective_input_tokens", 4096) + if isinstance(max_output, bool) or not isinstance(max_output, int) or max_output <= 0: + raise ProviderProfileError(f"{context}.runtime.max_output_tokens must be positive.") + if isinstance(min_effective, bool) or not isinstance(min_effective, int) or min_effective < 1024: + raise ProviderProfileError( + f"{context}.runtime.min_effective_input_tokens must be at least 1024." + ) + effort = raw.get("reasoning_effort", "auto") + if effort not in {"auto", "low", "medium", "high"}: + raise ProviderProfileError(f"{context}.runtime.reasoning_effort is invalid.") + capabilities_raw = raw.get("capabilities", {}) + if not isinstance(capabilities_raw, dict): + raise ProviderProfileError(f"{context}.runtime.capabilities must be an object.") + capabilities: list[tuple[str, bool | str]] = [] + for key in ("tools", "vision", "structured_output"): + value = capabilities_raw.get(key, "auto") + if value not in {True, False, "auto"}: + raise ProviderProfileError( + f"{context}.runtime.capabilities.{key} must be true, false, or auto." + ) + capabilities.append((key, value)) + runtime = ModelRuntime( + limit_mode=limit_mode, + context_window_tokens=_optional_positive_int( + raw, "context_window_tokens", context=f"{context}.runtime" + ), + max_input_tokens=_optional_positive_int( + raw, "max_input_tokens", context=f"{context}.runtime" + ), + max_output_tokens=max_output, + min_effective_input_tokens=min_effective, + limits_status=status, + limits_source=source, + temperature=_optional_number( + raw, "temperature", context=f"{context}.runtime", minimum=0, maximum=2 + ), + top_p=_optional_number( + raw, "top_p", context=f"{context}.runtime", minimum=0.000001, maximum=1 + ), + reasoning_effort=effort, + capabilities=tuple(capabilities), + ) + if status == "confirmed": + limit = resolve_model_input_limit(runtime) + safety = max(2048, int(limit * 0.10 + 0.999999)) + if limit - safety < min_effective: + raise ProviderProfileError( + f"{context}.runtime leaves less than min_effective_input_tokens after safety reserve." + ) + return runtime + + def _parse_model(raw: Any, *, provider_id: str, index: int) -> ProviderModel: context = f"providers[{provider_id}].models[{index}]" if not isinstance(raw, dict): @@ -164,6 +338,7 @@ def _parse_model(raw: Any, *, provider_id: str, index: int) -> ProviderModel: name=_required_string(raw, "name", context=context, max_length=120), model_id=_required_string(raw, "model_id", context=context, max_length=300), enabled=bool(raw.get("enabled", True)), + runtime=_parse_model_runtime(raw.get("runtime"), context=context), ) @@ -233,6 +408,20 @@ def _parse_profile(raw: Any, *, index: int, builtin: bool = False) -> ProviderPr raise ProviderProfileError( f"Provider {provider_id!r} contains duplicate model IDs." ) + portable_reasoning_adapters = { + "openai", + "openai-compatible", + "grok", + "antigravity", + "openrouter", + } + if adapter not in portable_reasoning_adapters: + for model in models: + if model.runtime.reasoning_effort != "auto": + raise ProviderProfileError( + f"{context}.models[{model.id}].runtime.reasoning_effort is " + f"not supported by adapter {adapter!r}." + ) return ProviderProfile( id=provider_id, @@ -243,6 +432,7 @@ def _parse_profile(raw: Any, *, index: int, builtin: bool = False) -> ProviderPr enabled=bool(raw.get("enabled", True)), models=models, auth_mode=auth_mode, + runtime=_parse_provider_runtime(raw.get("runtime"), context=context), ) @@ -251,9 +441,11 @@ def _parse_document(raw: Any) -> ProviderProfiles: return ProviderProfiles() if not isinstance(raw, dict): raise ProviderProfileError("Provider profile document must be an object.") - version = raw.get("version", 1) - if version not in {1, PROVIDER_PROFILES_VERSION}: - raise ProviderProfileError(f"Unsupported provider profile version {version!r}.") + version = raw.get("version") + if version != PROVIDER_PROFILES_VERSION: + raise ProviderProfileError( + "PROVIDER_PROFILE_RESET_REQUIRED: only version 3 provider profiles are supported." + ) builtins_raw = raw.get("builtins", []) if not isinstance(builtins_raw, list): raise ProviderProfileError("builtins must be a list.") @@ -290,12 +482,32 @@ def _profile_to_dict(profile: ProviderProfile) -> dict[str, Any]: "api_key": profile.api_key, "auth_mode": profile.auth_mode, "enabled": profile.enabled, + "runtime": { + "timeout_seconds": profile.runtime.timeout_seconds, + "max_retries": profile.runtime.max_retries, + "default_temperature": profile.runtime.default_temperature, + "default_top_p": profile.runtime.default_top_p, + "default_reasoning_effort": profile.runtime.default_reasoning_effort, + }, "models": [ { "id": model.id, "name": model.name, "model_id": model.model_id, "enabled": model.enabled, + "runtime": { + "limit_mode": model.runtime.limit_mode, + "context_window_tokens": model.runtime.context_window_tokens, + "max_input_tokens": model.runtime.max_input_tokens, + "max_output_tokens": model.runtime.max_output_tokens, + "min_effective_input_tokens": model.runtime.min_effective_input_tokens, + "limits_status": model.runtime.limits_status, + "limits_source": model.runtime.limits_source, + "temperature": model.runtime.temperature, + "top_p": model.runtime.top_p, + "reasoning_effort": model.runtime.reasoning_effort, + "capabilities": dict(model.runtime.capabilities), + }, } for model in profile.models ], @@ -471,12 +683,32 @@ def _provider_profiles_public_payload(document: ProviderProfiles) -> dict[str, A "enabled": profile.enabled, "api_key_configured": bool(profile.api_key), "api_key_hint": _api_key_hint(profile.api_key), + "runtime": { + "timeout_seconds": profile.runtime.timeout_seconds, + "max_retries": profile.runtime.max_retries, + "default_temperature": profile.runtime.default_temperature, + "default_top_p": profile.runtime.default_top_p, + "default_reasoning_effort": profile.runtime.default_reasoning_effort, + }, "models": [ { "id": model.id, "name": model.name, "model_id": model.model_id, "enabled": model.enabled, + "runtime": { + "limit_mode": model.runtime.limit_mode, + "context_window_tokens": model.runtime.context_window_tokens, + "max_input_tokens": model.runtime.max_input_tokens, + "max_output_tokens": model.runtime.max_output_tokens, + "min_effective_input_tokens": model.runtime.min_effective_input_tokens, + "limits_status": model.runtime.limits_status, + "limits_source": model.runtime.limits_source, + "temperature": model.runtime.temperature, + "top_p": model.runtime.top_p, + "reasoning_effort": model.runtime.reasoning_effort, + "capabilities": dict(model.runtime.capabilities), + }, } for model in profile.models ], @@ -565,6 +797,10 @@ def resolve_provider_model( raise ProviderProfileError( f"Model {model_name!r} in provider {provider_id!r} is disabled." ) + if model.runtime.limits_status != "confirmed": + raise ProviderProfileError( + f"RUNTIME_LIMITS_CONFIRMATION_REQUIRED: model {model_name!r} needs confirmed limits." + ) return profile, model raise ProviderProfileError( f"Model {model_name!r} is not configured for provider {provider_id!r}." diff --git a/EvoScientist/langgraph_dev/http.py b/EvoScientist/langgraph_dev/http.py index 2662425..b48b655 100644 --- a/EvoScientist/langgraph_dev/http.py +++ b/EvoScientist/langgraph_dev/http.py @@ -76,6 +76,7 @@ from EvoScientist.llm.provider_operations import ( discover_provider_models, test_provider_model, ) +from EvoScientist.llm.runtime_snapshots import create_run_runtime_snapshot from EvoScientist.sessions import ( MAIN_THREAD_FILTER_PARAMS, MAIN_THREAD_FILTER_SQL, @@ -655,6 +656,45 @@ async def provider_profiles_endpoint(request: Request) -> JSONResponse: ) +async def runtime_snapshot_endpoint(request: Request) -> JSONResponse: + """Create an opaque server-side snapshot for one forthcoming graph run.""" + auth_error = _provider_admin_error(request) + if auth_error is not None: + return auth_error + + try: + payload = await request.json() + if not isinstance(payload, dict): + raise ProviderProfileError("Request body must be an object.") + snapshot_id = payload.get("snapshot_id") + model = payload.get("model") + provider = payload.get("provider") + if not isinstance(snapshot_id, str): + raise ProviderProfileError("snapshot_id is required.") + if not isinstance(model, str) or not isinstance(provider, str): + raise ProviderProfileError("model and provider are required.") + snapshot = await asyncio.to_thread( + create_run_runtime_snapshot, + snapshot_id, + model=model, + provider=provider, + ) + return JSONResponse( + {"snapshot": snapshot.public_payload() if snapshot is not None else None} + ) + except (json.JSONDecodeError, UnicodeDecodeError): + return JSONResponse( + {"error": "Request body must be valid JSON."}, status_code=400 + ) + except ProviderProfileError as exc: + return JSONResponse({"error": str(exc)}, status_code=400) + except Exception: + _logger.exception("Runtime snapshot request failed") + return JSONResponse( + {"error": "Runtime snapshot request failed."}, status_code=500 + ) + + async def llm_config_endpoint(request: Request) -> JSONResponse: """Read or patch the LLM-related fields persisted in config.yaml.""" auth_error = _provider_admin_error(request) @@ -1636,6 +1676,11 @@ app = Starlette( provider_profiles_endpoint, methods=["GET", "PUT"], ), + Route( + "/api/runtime-snapshots", + runtime_snapshot_endpoint, + methods=["POST"], + ), Route( "/api/config", llm_config_endpoint, diff --git a/EvoScientist/llm/models.py b/EvoScientist/llm/models.py index f40c79c..2f038fb 100644 --- a/EvoScientist/llm/models.py +++ b/EvoScientist/llm/models.py @@ -15,6 +15,7 @@ import json import os import re import warnings +from dataclasses import dataclass from typing import Any from langchain.chat_models import init_chat_model @@ -29,6 +30,7 @@ from ..config.provider_profiles import ( get_provider_profiles_path, list_configured_builtin_model_entries, list_configured_model_entries, + resolve_model_input_limit, resolve_provider_model, ) from .context_window import apply_known_context_window @@ -53,6 +55,20 @@ _DEEPSEEK_BASE_URL = "https://api.deepseek.com" _MOONSHOT_BASE_URL = "https://api.moonshot.cn/v1" _KIMI_CODING_BASE_URL = "https://api.kimi.com/coding/" +# ``glm`` is a common shorthand in existing configuration files. The +# supported provider identifier is ``zhipu`` because it describes the API +# vendor rather than a single model family. Normalize it at the boundary so +# legacy configuration reaches the same authenticated OpenAI-compatible route. +_PROVIDER_ALIASES: dict[str, str] = { + "glm": "zhipu", +} + + +def normalize_provider_id(provider: str) -> str: + """Return the canonical provider ID accepted by the model registry.""" + normalized = provider.strip().lower() + return _PROVIDER_ALIASES.get(normalized, normalized) + # Providers routed through the OpenAI provider with a custom base_url. # Maps provider name → (base_url or None, env var for API key). _OPENAI_ROUTED_PROVIDERS: dict[str, tuple[str | None, str]] = { @@ -84,6 +100,74 @@ _THINKING_CAPABLE_PROVIDERS: set[str] = {"minimax"} _TRUTHY_ENV_VALUES = {"1", "true", "yes", "on"} _FALSEY_ENV_VALUES = {"0", "false", "no", "off"} + +@dataclass(frozen=True) +class ResolvedRuntimeOptions: + """Provider-neutral parameters resolved once for a configured model.""" + + profile_revision: str | None + model_id: str + adapter_id: str + limit_mode: str + resolved_input_limit: int + max_output_tokens: int + min_effective_input_tokens: int + timeout_seconds: int + max_retries: int + temperature: float | None + top_p: float | None + reasoning_effort: str + limits_source: str + capabilities: tuple[tuple[str, bool | str], ...] + + +def resolve_runtime_options( + profile: ProviderProfile, + model: ProviderModel, + call_overrides: dict[str, Any] | None = None, +) -> ResolvedRuntimeOptions: + """Resolve the bounded runtime options for one configured model. + + Only sampling values and a lower output budget may be overridden by an + internal caller. Connection and capacity fields always come from the + persisted profile. + """ + overrides = call_overrides or {} + runtime = model.runtime + input_limit = resolve_model_input_limit(runtime) + output_limit = runtime.max_output_tokens + override_output = overrides.get("max_tokens") + if isinstance(override_output, int) and not isinstance(override_output, bool): + if override_output <= 0 or override_output > output_limit: + raise ValueError("max_tokens override must be positive and no greater than the model limit.") + output_limit = override_output + temperature = overrides.get("temperature", runtime.temperature) + if temperature is None: + temperature = profile.runtime.default_temperature + top_p = overrides.get("top_p", runtime.top_p) + if top_p is None: + top_p = profile.runtime.default_top_p + return ResolvedRuntimeOptions( + profile_revision=get_provider_profile_revision(profile.id), + model_id=model.model_id, + adapter_id=profile.adapter, + limit_mode=runtime.limit_mode or "input_only", + resolved_input_limit=input_limit, + max_output_tokens=output_limit, + min_effective_input_tokens=runtime.min_effective_input_tokens, + timeout_seconds=profile.runtime.timeout_seconds, + max_retries=profile.runtime.max_retries, + temperature=temperature, + top_p=top_p, + reasoning_effort=( + runtime.reasoning_effort + if runtime.reasoning_effort != "auto" + else profile.runtime.default_reasoning_effort + ), + limits_source=runtime.limits_source, + capabilities=runtime.capabilities, + ) + # Model registry: list of (short_name, model_id, provider) # Allows same short_name across different providers. _MODEL_ENTRIES: list[tuple[str, str, str]] = [ @@ -261,6 +345,11 @@ _STATIC_PROVIDER_IDS = {provider for _, _, provider in _MODEL_ENTRIES} | {"ollam _MODEL_CATALOG_ID_PATTERN = re.compile(r"^[a-z0-9][a-z0-9._-]{0,63}$") +def is_static_provider_id(provider: str) -> bool: + """Return whether a provider is configured outside custom profiles.""" + return normalize_provider_id(provider) in _STATIC_PROVIDER_IDS + + def normalize_builtin_model_catalog(raw: Any) -> list[dict[str, Any]] | None: """Validate and normalize the optional config.yaml built-in model catalog.""" if raw is None: @@ -341,7 +430,13 @@ def list_builtin_model_catalog_entries( def _resolve_builtin_catalog_model(provider: str, model: str) -> str | None: from ..config.settings import load_config - profile = get_builtin_provider_profile(provider) + try: + profile = get_builtin_provider_profile(provider) + except ProviderProfileError: + # Custom profiles use an explicit v3 reset path during development. + # Built-in model selection must still work from config.yaml while that + # reset has not happened, because it never consumes the old profile. + profile = None if profile is not None: if not profile.enabled: return None @@ -366,7 +461,10 @@ def get_model_runtime_revision(provider: str | None) -> str | None: if not provider: return None if provider in _STATIC_PROVIDER_IDS: - managed_revision = get_provider_profile_revision(provider) + try: + managed_revision = get_provider_profile_revision(provider) + except ProviderProfileError: + managed_revision = None if managed_revision is not None: return managed_revision from ..config.settings import load_config @@ -398,7 +496,10 @@ def get_models_for_provider(provider: str) -> list[tuple[str, str]]: List of (short_name, model_id) tuples for the provider. """ if provider in _STATIC_PROVIDER_IDS: - managed = get_builtin_provider_profile(provider) + try: + managed = get_builtin_provider_profile(provider) + except ProviderProfileError: + managed = None if managed is not None: if not managed.enabled: return [] @@ -553,9 +654,12 @@ def get_chat_model( >>> model = get_chat_model("claude-3-opus-20240229", provider="anthropic") # Full ID """ model = model or DEFAULT_MODEL + if provider is not None: + provider = normalize_provider_id(provider) dynamic_profile = profile_override dynamic_model = model_override + resolved_runtime: ResolvedRuntimeOptions | None = None if dynamic_profile is not None: provider = dynamic_profile.id if dynamic_model is None: @@ -575,10 +679,27 @@ def get_chat_model( try: dynamic_resolution = resolve_provider_model(provider, model) except ProviderProfileError as exc: - raise ValueError(str(exc)) from exc + # An opaque provider override remains a valid LangChain escape + # hatch. A stale v1/v2 registry cannot be allowed to block it; + # configured custom profiles are rejected earlier by the snapshot + # endpoint and never silently selected from that document. + if str(exc).startswith("PROVIDER_PROFILE_RESET_REQUIRED"): + dynamic_resolution = None + else: + raise ValueError(str(exc)) from exc if dynamic_resolution is not None: dynamic_profile, dynamic_model = dynamic_resolution + if dynamic_profile is not None and dynamic_model is not None: + resolved_runtime = resolve_runtime_options(dynamic_profile, dynamic_model, kwargs) + kwargs.setdefault("timeout", resolved_runtime.timeout_seconds) + kwargs.setdefault("max_retries", resolved_runtime.max_retries) + kwargs.setdefault("max_tokens", resolved_runtime.max_output_tokens) + if resolved_runtime.temperature is not None: + kwargs.setdefault("temperature", resolved_runtime.temperature) + if resolved_runtime.top_p is not None: + kwargs.setdefault("top_p", resolved_runtime.top_p) + # Look up short name in the configured catalog, then the legacy registry. model_id = dynamic_model.model_id if dynamic_model is not None else None if model_id is None and provider in _STATIC_PROVIDER_IDS: @@ -616,11 +737,13 @@ def get_chat_model( ) _is_openai_proxy = False _original_provider: str | None = None - managed_builtin_profile = ( - get_builtin_provider_profile(provider) - if provider in _STATIC_PROVIDER_IDS - else None - ) + if provider in _STATIC_PROVIDER_IDS: + try: + managed_builtin_profile = get_builtin_provider_profile(provider) + except ProviderProfileError: + managed_builtin_profile = None + else: + managed_builtin_profile = None if managed_builtin_profile is not None: if not managed_builtin_profile.enabled: raise ValueError(f"Provider {provider!r} is disabled.") @@ -679,6 +802,23 @@ def get_chat_model( kwargs.setdefault("reasoning", {"effort": effort, "summary": "auto"}) _patch_openrouter_strip_responses_reasoning() + if resolved_runtime is not None and resolved_runtime.reasoning_effort != "auto": + if dynamic_profile.adapter in { + "openai", + "openai-compatible", + "grok", + "antigravity", + "openrouter", + }: + kwargs.setdefault( + "reasoning", {"effort": resolved_runtime.reasoning_effort} + ) + else: + raise ValueError( + f"REASONING_EFFORT_UNSUPPORTED: adapter {dynamic_profile.adapter!r} " + "does not support a portable reasoning effort setting." + ) + elif provider == "anthropic": base_url = os.environ.get("ANTHROPIC_BASE_URL", "") if base_url: @@ -729,8 +869,10 @@ def get_chat_model( if base_url: kwargs.setdefault("base_url", base_url) api_key = os.environ.get(api_key_env, "") - if api_key: - kwargs.setdefault("api_key", api_key) + # ChatOpenAI otherwise falls back to OPENAI_API_KEY when ``api_key`` + # is omitted. That is unsafe for routed providers: a missing Zhipu + # credential must never send an unrelated OpenAI credential to Zhipu. + kwargs.setdefault("api_key", api_key or None) # SiliconFlow: disable thinking — LangChain drops reasoning_content # from history, causing error 20015 on multi-turn requests. if provider == "siliconflow": @@ -864,6 +1006,20 @@ def get_chat_model( _patch_ccproxy_system_to_developer(chat_model) apply_known_context_window(chat_model) + if resolved_runtime is not None: + profile_data = getattr(chat_model, "profile", None) + base_profile = profile_data if isinstance(profile_data, dict) else {} + try: + chat_model.profile = { + **base_profile, + "max_input_tokens": resolved_runtime.resolved_input_limit, + "min_effective_input_tokens": ( + resolved_runtime.min_effective_input_tokens + ), + "runtime_capabilities": dict(resolved_runtime.capabilities), + } + except Exception: + pass return chat_model diff --git a/EvoScientist/llm/runtime_snapshots.py b/EvoScientist/llm/runtime_snapshots.py new file mode 100644 index 0000000..3dcfee2 --- /dev/null +++ b/EvoScientist/llm/runtime_snapshots.py @@ -0,0 +1,401 @@ +"""Server-side, per-run snapshots for custom provider model configuration. + +The browser and LangGraph run configuration carry only a random snapshot ID. +The SQLite record stores non-secret connection metadata and runtime limits. +Credentials stay in a process-local cache for the lifetime of a development +deployment, so a restarted deployment fails an old run explicitly instead of +silently using a newly edited credential. +""" + +from __future__ import annotations + +import json +import sqlite3 +import threading +import time +from dataclasses import dataclass, replace +from pathlib import Path +from typing import Any + +from ..config.provider_profiles import ( + ModelRuntime, + ProviderModel, + ProviderProfile, + ProviderProfileError, + ProviderRuntime, + get_provider_profile_revision, + resolve_model_input_limit, + resolve_provider_model, +) +from ..config.settings import get_config_dir +from .models import ( + ResolvedRuntimeOptions, + is_static_provider_id, + normalize_provider_id, + resolve_runtime_options, +) + +_SNAPSHOT_TTL_SECONDS = 7 * 24 * 60 * 60 +_MAX_SNAPSHOT_ID_LENGTH = 128 +_DB_LOCK = threading.RLock() +_MODEL_CACHE_LOCK = threading.RLock() +_MODEL_CACHE: dict[str, Any] = {} +_SNAPSHOT_SECRETS: dict[str, str] = {} + + +def _validate_static_provider_credentials(provider: str) -> None: + """Reject known static selections that lack their provider-specific key. + + Runtime snapshots are created by the WebUI before a run is queued. This is + the latest point at which we can return a useful client error without + making the graph fail after the user has already submitted a message. + """ + if provider not in {"zhipu", "zhipu-code"}: + return + + from ..config.settings import get_effective_config + from .provider_operations import build_builtin_provider_profile + + config = get_effective_config() + try: + profile = build_builtin_provider_profile(config, provider) + except ProviderProfileError as exc: + # A stale custom registry must not hide configuration supplied through + # config.yaml or the environment for a built-in provider. + if not str(exc).startswith("PROVIDER_PROFILE_RESET_REQUIRED"): + raise + profile = build_builtin_provider_profile( + config, provider, load_saved=False + ) + if not profile.api_key.strip(): + raise ProviderProfileError( + "ZHIPU_API_KEY_NOT_CONFIGURED: GLM requires a valid Zhipu API key. " + "Configure ZHIPU_API_KEY before starting a conversation." + ) + + +@dataclass(frozen=True) +class RunRuntimeSnapshot: + """An immutable server-side configuration captured before one run starts.""" + + snapshot_id: str + created_at: int + expires_at: int + profile_revision: str | None + profile: ProviderProfile + model: ProviderModel + options: ResolvedRuntimeOptions + + def public_payload(self) -> dict[str, Any]: + """Return diagnostic fields that are safe to attach to a run.""" + return { + "snapshot_id": self.snapshot_id, + "provider": self.profile.id, + "model": self.model.id, + "profile_revision": self.profile_revision, + "created_at": self.created_at, + "expires_at": self.expires_at, + "runtime": { + "max_input_tokens": self.options.resolved_input_limit, + "max_output_tokens": self.options.max_output_tokens, + "min_effective_input_tokens": self.options.min_effective_input_tokens, + "timeout_seconds": self.options.timeout_seconds, + "max_retries": self.options.max_retries, + "temperature": self.options.temperature, + "top_p": self.options.top_p, + "reasoning_effort": self.options.reasoning_effort, + "limits_source": self.options.limits_source, + "capabilities": dict(self.options.capabilities), + }, + } + + +def _database_path() -> Path: + return get_config_dir() / "run-runtime-snapshots.sqlite3" + + +def _connect() -> sqlite3.Connection: + path = _database_path() + path.parent.mkdir(parents=True, exist_ok=True) + try: + path.parent.chmod(0o700) + except OSError: + pass + connection = sqlite3.connect(path) + connection.execute( + """ + CREATE TABLE IF NOT EXISTS run_runtime_snapshots ( + snapshot_id TEXT PRIMARY KEY, + expires_at INTEGER NOT NULL, + payload_json TEXT NOT NULL + ) + """ + ) + try: + path.chmod(0o600) + except OSError: + pass + return connection + + +def _validate_snapshot_id(snapshot_id: str) -> str: + normalized = snapshot_id.strip() + if not normalized or len(normalized) > _MAX_SNAPSHOT_ID_LENGTH: + raise ProviderProfileError("runtime snapshot ID is invalid.") + if any(ord(character) < 33 or ord(character) > 126 for character in normalized): + raise ProviderProfileError("runtime snapshot ID is invalid.") + return normalized + + +def _provider_runtime_payload(runtime: ProviderRuntime) -> dict[str, Any]: + return { + "timeout_seconds": runtime.timeout_seconds, + "max_retries": runtime.max_retries, + "default_temperature": runtime.default_temperature, + "default_top_p": runtime.default_top_p, + "default_reasoning_effort": runtime.default_reasoning_effort, + } + + +def _model_runtime_payload(runtime: ModelRuntime) -> dict[str, Any]: + return { + "limit_mode": runtime.limit_mode, + "context_window_tokens": runtime.context_window_tokens, + "max_input_tokens": runtime.max_input_tokens, + "max_output_tokens": runtime.max_output_tokens, + "min_effective_input_tokens": runtime.min_effective_input_tokens, + "limits_status": runtime.limits_status, + "limits_source": runtime.limits_source, + "temperature": runtime.temperature, + "top_p": runtime.top_p, + "reasoning_effort": runtime.reasoning_effort, + "capabilities": dict(runtime.capabilities), + } + + +def _profile_payload(profile: ProviderProfile, model: ProviderModel) -> dict[str, Any]: + return { + "id": profile.id, + "name": profile.name, + "adapter": profile.adapter, + "base_url": profile.base_url, + # Credentials stay in the process-local secret cache. The persisted + # snapshot remains safe to inspect through diagnostics and never + # duplicates a Provider Profile API key. + "api_key": "", + "auth_mode": profile.auth_mode, + "enabled": profile.enabled, + "runtime": _provider_runtime_payload(profile.runtime), + "models": [ + { + "id": model.id, + "name": model.name, + "model_id": model.model_id, + "enabled": model.enabled, + "runtime": _model_runtime_payload(model.runtime), + } + ], + } + + +def _parse_snapshot_payload( + payload: dict[str, Any], *, require_secret: bool = True +) -> RunRuntimeSnapshot: + try: + profile_raw = payload["profile"] + if not isinstance(profile_raw, dict): + raise ProviderProfileError("runtime snapshot profile is invalid.") + runtime_raw = payload["runtime"] + if not isinstance(runtime_raw, dict): + raise ProviderProfileError("runtime snapshot options are invalid.") + + # The source is a private snapshot, not a browser draft. Reuse the + # validated profile parser without looking up the current registry. + from ..config.provider_profiles import _parse_profile + + profile = _parse_profile(profile_raw, index=0) + model = profile.models[0] + resolved_input_limit = int(runtime_raw["resolved_input_limit"]) + if resolved_input_limit != resolve_model_input_limit(model.runtime): + raise ProviderProfileError("runtime snapshot limits are invalid.") + capabilities_raw = dict(runtime_raw["capabilities"]) + options = ResolvedRuntimeOptions( + profile_revision=( + payload["profile_revision"] + if isinstance(payload.get("profile_revision"), str) + else None + ), + model_id=str(runtime_raw["model_id"]), + adapter_id=str(runtime_raw["adapter_id"]), + limit_mode=str(runtime_raw["limit_mode"]), + resolved_input_limit=resolved_input_limit, + max_output_tokens=int(runtime_raw["max_output_tokens"]), + min_effective_input_tokens=int( + runtime_raw["min_effective_input_tokens"] + ), + timeout_seconds=int(runtime_raw["timeout_seconds"]), + max_retries=int(runtime_raw["max_retries"]), + temperature=runtime_raw.get("temperature"), + top_p=runtime_raw.get("top_p"), + reasoning_effort=str(runtime_raw["reasoning_effort"]), + limits_source=str(runtime_raw["limits_source"]), + capabilities=tuple( + (key, capabilities_raw[key]) + for key in ("tools", "vision", "structured_output") + if capabilities_raw.get(key) in {True, False, "auto"} + ), + ) + snapshot_id = _validate_snapshot_id(str(payload["snapshot_id"])) + secret = _SNAPSHOT_SECRETS.get(snapshot_id) + if require_secret and snapshot_id not in _SNAPSHOT_SECRETS: + raise ProviderProfileError( + "RUN_RUNTIME_SNAPSHOT_CREDENTIAL_UNAVAILABLE: the deployment was " + "restarted after this run began. Start the message again." + ) + return RunRuntimeSnapshot( + snapshot_id=snapshot_id, + created_at=int(payload["created_at"]), + expires_at=int(payload["expires_at"]), + profile_revision=options.profile_revision, + profile=replace(profile, api_key=secret or ""), + model=model, + options=options, + ) + except (KeyError, TypeError, ValueError) as exc: + raise ProviderProfileError("runtime snapshot is invalid.") from exc + + +def _snapshot_payload(snapshot: RunRuntimeSnapshot) -> dict[str, Any]: + return { + "snapshot_id": snapshot.snapshot_id, + "created_at": snapshot.created_at, + "expires_at": snapshot.expires_at, + "profile_revision": snapshot.profile_revision, + "profile": _profile_payload(snapshot.profile, snapshot.model), + "runtime": { + "model_id": snapshot.options.model_id, + "adapter_id": snapshot.options.adapter_id, + "limit_mode": snapshot.options.limit_mode, + "resolved_input_limit": snapshot.options.resolved_input_limit, + "max_output_tokens": snapshot.options.max_output_tokens, + "min_effective_input_tokens": snapshot.options.min_effective_input_tokens, + "timeout_seconds": snapshot.options.timeout_seconds, + "max_retries": snapshot.options.max_retries, + "temperature": snapshot.options.temperature, + "top_p": snapshot.options.top_p, + "reasoning_effort": snapshot.options.reasoning_effort, + "limits_source": snapshot.options.limits_source, + "capabilities": dict(snapshot.options.capabilities), + }, + } + + +def create_run_runtime_snapshot( + snapshot_id: str, + *, + model: str, + provider: str, +) -> RunRuntimeSnapshot | None: + """Freeze one custom provider model configuration for a forthcoming run. + + Static built-in providers remain managed by their existing deployment + configuration and return ``None``. A missing custom model is an error so + a model selection never silently falls back to a different model. + """ + snapshot_id = _validate_snapshot_id(snapshot_id) + model = model.strip() + provider = normalize_provider_id(provider) + if not model or not provider: + raise ProviderProfileError("model and provider are required for a runtime snapshot.") + + # Built-in providers are resolved from the deployment configuration and + # have no browser-managed secret or runtime profile to freeze. Avoid + # loading custom profiles at all, so a stale development registry cannot + # block a normal built-in run. + if is_static_provider_id(provider): + _validate_static_provider_credentials(provider) + return None + + existing = get_run_runtime_snapshot(snapshot_id) + if existing is not None: + if existing.model.id != model or existing.profile.id != provider: + raise ProviderProfileError( + "runtime snapshot ID is already bound to a different model selection." + ) + return existing + + resolved = resolve_provider_model(provider, model) + if resolved is None: + return None + profile, provider_model = resolved + options = resolve_runtime_options(profile, provider_model) + now = int(time.time()) + snapshot = RunRuntimeSnapshot( + snapshot_id=snapshot_id, + created_at=now, + expires_at=now + _SNAPSHOT_TTL_SECONDS, + profile_revision=get_provider_profile_revision(provider), + profile=profile, + model=provider_model, + options=options, + ) + encoded = json.dumps( + _snapshot_payload(snapshot), sort_keys=True, separators=(",", ":") + ) + with _DB_LOCK, _connect() as connection: + connection.execute( + "DELETE FROM run_runtime_snapshots WHERE expires_at <= ?", (now,) + ) + connection.execute( + "INSERT OR REPLACE INTO run_runtime_snapshots(snapshot_id, expires_at, payload_json) " + "VALUES (?, ?, ?)", + (snapshot.snapshot_id, snapshot.expires_at, encoded), + ) + connection.commit() + _SNAPSHOT_SECRETS[snapshot.snapshot_id] = profile.api_key + return snapshot + + +def get_run_runtime_snapshot(snapshot_id: str) -> RunRuntimeSnapshot | None: + """Load a still-valid runtime snapshot by its opaque ID.""" + snapshot_id = _validate_snapshot_id(snapshot_id) + now = int(time.time()) + with _DB_LOCK, _connect() as connection: + connection.execute( + "DELETE FROM run_runtime_snapshots WHERE expires_at <= ?", (now,) + ) + row = connection.execute( + "SELECT expires_at, payload_json FROM run_runtime_snapshots WHERE snapshot_id = ?", + (snapshot_id,), + ).fetchone() + connection.commit() + if row is None or int(row[0]) <= now: + _SNAPSHOT_SECRETS.pop(snapshot_id, None) + return None + try: + payload = json.loads(str(row[1])) + except json.JSONDecodeError as exc: + raise ProviderProfileError("runtime snapshot is unreadable.") from exc + if not isinstance(payload, dict): + raise ProviderProfileError("runtime snapshot is unreadable.") + return _parse_snapshot_payload(payload) + + +def get_snapshot_chat_model(snapshot_id: str) -> Any: + """Build (or reuse) the chat model bound to a frozen runtime snapshot.""" + snapshot = get_run_runtime_snapshot(snapshot_id) + if snapshot is None: + raise ProviderProfileError( + "RUN_RUNTIME_SNAPSHOT_UNAVAILABLE: the run configuration snapshot " + "expired or is unavailable. Start the message again." + ) + with _MODEL_CACHE_LOCK: + cached = _MODEL_CACHE.get(snapshot.snapshot_id) + if cached is not None: + return cached + from .models import get_profile_chat_model + + model = get_profile_chat_model(snapshot.profile, snapshot.model) + with _MODEL_CACHE_LOCK: + _MODEL_CACHE[snapshot.snapshot_id] = model + return model diff --git a/EvoScientist/middleware/__init__.py b/EvoScientist/middleware/__init__.py index cd28de6..c1975bd 100644 --- a/EvoScientist/middleware/__init__.py +++ b/EvoScientist/middleware/__init__.py @@ -27,6 +27,11 @@ from .memory_lifecycle import ( create_memory_lifecycle_middleware, default_memory_scheduler, ) +from .message_budget import ( + MessageReservePolicy, + 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 ( @@ -46,16 +51,19 @@ __all__ = [ "ContextOverflowMapperMiddleware", "EvoMemoryLifecycleMiddleware", "EvoMemoryMiddleware", + "MessageReservePolicy", "ModelFallbackMiddleware", "Question", "RuntimeContextMiddleware", "SchedulerMiddleware", "ToolErrorHandlerMiddleware", "compute_context_editing_trigger", + "count_message_text_tokens", "create_code_interpreter_middleware", "create_context_editing_middleware", "create_memory_lifecycle_middleware", "create_memory_middleware", + "create_message_budget_middleware", "create_runtime_context_middleware", "create_scheduler_middleware", "create_tool_selector_middleware", diff --git a/EvoScientist/middleware/configurable_model.py b/EvoScientist/middleware/configurable_model.py index 656f1bb..534216a 100644 --- a/EvoScientist/middleware/configurable_model.py +++ b/EvoScientist/middleware/configurable_model.py @@ -74,6 +74,23 @@ def _read_model_override() -> tuple[str | None, str | 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. @@ -128,8 +145,32 @@ class ConfigurableModelMiddleware(AgentMiddleware): provider, ) - def _resolve(self, model: str, provider: str | None) -> Any: + 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 + + 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) @@ -156,9 +197,12 @@ class ConfigurableModelMiddleware(AgentMiddleware): model_name, provider = _read_model_override() if model_name is None: return handler(request) + runtime_snapshot_id = _read_runtime_snapshot_id() try: - new_model = self._resolve(model_name, provider) + 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", @@ -178,6 +222,7 @@ class ConfigurableModelMiddleware(AgentMiddleware): model_name, provider = _read_model_override() if model_name 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 @@ -185,8 +230,12 @@ class ConfigurableModelMiddleware(AgentMiddleware): # ``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) + 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", diff --git a/EvoScientist/middleware/message_budget.py b/EvoScientist/middleware/message_budget.py new file mode 100644 index 0000000..257c404 --- /dev/null +++ b/EvoScientist/middleware/message_budget.py @@ -0,0 +1,327 @@ +"""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. +""" + +from __future__ import annotations + +import threading +from collections.abc import Iterable, Mapping +from contextvars import ContextVar +from dataclasses import dataclass +from typing import Any + +from langchain_core.messages import AnyMessage, SystemMessage + + +@dataclass(frozen=True) +class MessageReservePolicy: + """Static non-message reserves used without full-request token counting.""" + + system_tokens: int = 4_096 + memory_tokens: int = 8_192 + tool_tokens: int = 8_192 + attachment_tokens: int = 8_192 + safety_fraction: float = 0.10 + minimum_safety_tokens: int = 2_048 + soft_fraction: float = 0.70 + keep_fraction: float = 0.35 + + def fixed_tokens( + self, *, has_tools: bool = False, has_attachments: bool = False + ) -> int: + """Return only the fixed reserves applicable to this model request.""" + return ( + self.system_tokens + + self.memory_tokens + + (self.tool_tokens if has_tools else 0) + + (self.attachment_tokens if has_attachments else 0) + ) + + def hard_budget( + self, + input_limit: int, + *, + has_tools: bool = False, + has_attachments: bool = False, + ) -> int: + safety = max( + self.minimum_safety_tokens, int(input_limit * self.safety_fraction) + ) + return max( + 1_024, + input_limit + - self.fixed_tokens( + has_tools=has_tools, has_attachments=has_attachments + ) + - safety, + ) + + def soft_budget(self, input_limit: int, **mode: bool) -> int: + return max( + 1_024, int(self.hard_budget(input_limit, **mode) * self.soft_fraction) + ) + + def keep_budget(self, input_limit: int, **mode: bool) -> int: + return max( + 1_024, int(self.hard_budget(input_limit, **mode) * self.keep_fraction) + ) + + +@dataclass(frozen=True) +class MessageBudget: + """The message budget selected for one model invocation.""" + + input_limit: int + hard_tokens: int + soft_tokens: int + keep_tokens: int + min_effective_input_tokens: int + has_tools: bool + has_attachments: bool + + +class ContextBudgetUnsatisfiableError(ValueError): + """Raised before a call when static reserves leave no safe message budget.""" + + +_ACTIVE_BUDGET: ContextVar[MessageBudget | None] = ContextVar( + "evoscientist_message_budget", default=None +) + + +def _content_text(content: Any) -> str: + """Extract text without serializing image/file blocks or tool schemas.""" + if isinstance(content, str): + return content + if not isinstance(content, list): + return "" + parts: list[str] = [] + for block in content: + if isinstance(block, str): + parts.append(block) + continue + if not isinstance(block, Mapping): + continue + block_type = block.get("type") + if block_type in {"image", "image_url", "file", "document", "audio", "video"}: + continue + text = block.get("text") + if isinstance(text, str): + parts.append(text) + return "\n".join(parts) + + +def count_message_text_tokens(messages: Iterable[AnyMessage]) -> int: + """Conservative, provider-neutral token estimate for compactable text only.""" + characters = sum(len(_content_text(message.content)) for message in messages) + # A four-character estimate deliberately errs slightly high for Chinese and + # mixed code while remaining cheap enough to run before every model call. + return max(0, (characters + 3) // 4) + + +def _has_attachments(messages: Iterable[AnyMessage]) -> bool: + attachment_types = {"image", "image_url", "file", "document", "audio", "video"} + for message in messages: + if not isinstance(message.content, list): + continue + for block in message.content: + if isinstance(block, Mapping) and block.get("type") in attachment_types: + return True + return False + + +def _current_configurable() -> Mapping[str, Any]: + try: + from langgraph.config import get_config + + config = get_config() + except Exception: + return {} + if not isinstance(config, Mapping): + return {} + configurable = config.get("configurable") + return configurable if isinstance(configurable, Mapping) else {} + + +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.""" + from deepagents.middleware.summarization import SummarizationMiddleware + + class _RuntimeMessageBudgetMiddleware(SummarizationMiddleware): + name = "message_budget" + + def __init__(self) -> None: + self._fallback_model = model + self._model_cache: dict[tuple[str, str | None], Any] = {} + self._model_cache_lock = threading.RLock() + self._policy = policy or MessageReservePolicy() + # Triggering and cutoff are overridden below. The base class is + # still used for safe AI/tool-pair handling, offloading, and + # persisted summarization events. + super().__init__( + model=model, + backend=backend, + trigger=("tokens", 1_000_000_000), + keep=("messages", 8), + token_counter=count_message_text_tokens, + trim_tokens_to_summarize=4_000, + truncate_args_settings={ + "trigger": ("tokens", 2_048), + "keep": ("messages", 8), + "max_length": 2_000, + "truncation_text": "...(tool arguments compacted)", + }, + ) + + @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: + 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) + if cached is not None: + return cached + from ..llm.models import get_chat_model + + resolved = get_chat_model(model=model_name, provider=provider_name) + with self._model_cache_lock: + self._model_cache[key] = resolved + return resolved + + def _input_limit(self) -> int: + profile = getattr(self.model, "profile", None) + 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): + 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", [])) + mode = { + "has_tools": 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: + raise ContextBudgetUnsatisfiableError( + "CONTEXT_BUDGET_UNSATISFIABLE: configured input limit " + f"{budget.input_limit:,} leaves only {budget.hard_tokens:,} " + "tokens after fixed reserves; increase the model window, reduce " + "the output budget, or disable tools/attachments." + ) + return budget + + 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, + ) + + def wrap_model_call(self, request: Any, handler: Any) -> Any: + token = _ACTIVE_BUDGET.set(self._budget_for_request(request)) + try: + return super().wrap_model_call(request, handler) + finally: + _ACTIVE_BUDGET.reset(token) + + async def awrap_model_call(self, request: Any, handler: Any) -> Any: + token = _ACTIVE_BUDGET.set(self._budget_for_request(request)) + try: + return await super().awrap_model_call(request, handler) + finally: + _ACTIVE_BUDGET.reset(token) + + def _get_profile_limits(self) -> int | None: + return self._active_budget().hard_tokens + + def _count_tokens( + self, + messages: list[AnyMessage], + system_message: SystemMessage | None, + tools: list[Any] | None, + ) -> int: + del system_message, tools + return count_message_text_tokens(messages) + + def _should_summarize( + self, messages: list[AnyMessage], total_tokens: int + ) -> bool: + del messages + return total_tokens >= self._active_budget().soft_tokens + + def _determine_cutoff_index(self, messages: list[AnyMessage]) -> int: + keep_tokens = self._active_budget().keep_tokens + retained = 0 + cutoff = len(messages) + for index in range(len(messages) - 1, -1, -1): + message_tokens = count_message_text_tokens([messages[index]]) + if retained + message_tokens > keep_tokens: + # Keep at least the newest message even when it alone + # exceeds the retention budget. The safe-cutoff helper + # below expands that boundary when it is a tool result. + cutoff = index if retained == 0 else index + 1 + break + retained += message_tokens + cutoff = index + if cutoff <= 0: + return 0 + return self._lc_helper._find_safe_cutoff_point(messages, cutoff) + + return _RuntimeMessageBudgetMiddleware() + + +def create_message_budget_middleware( + model: Any, backend: Any, *, policy: MessageReservePolicy | None = None +): + """Construct automatic compaction middleware for a graph backend.""" + return MessageBudgetMiddleware.create(model, backend, policy=policy) diff --git a/tests/test_langgraph_dev_http.py b/tests/test_langgraph_dev_http.py index 8089a67..4c3a6c8 100644 --- a/tests/test_langgraph_dev_http.py +++ b/tests/test_langgraph_dev_http.py @@ -124,12 +124,21 @@ def test_provider_profiles_api_round_trip_redacts_secret(tmp_path, monkeypatch): "base_url": "https://llm.example.test/v1", "api_key": "provider-secret", "enabled": True, + "runtime": {"timeout_seconds": 90, "max_retries": 1}, "models": [ { "id": "lab-model", "name": "Lab Model", "model_id": "vendor/model", "enabled": True, + "runtime": { + "limit_mode": "combined", + "context_window_tokens": 32_768, + "max_output_tokens": 4_096, + "min_effective_input_tokens": 4_096, + "limits_status": "confirmed", + "limits_source": "provider", + }, } ], } @@ -147,6 +156,41 @@ def test_provider_profiles_api_round_trip_redacts_secret(tmp_path, monkeypatch): assert get_response.json()["providers"][0]["models"][0]["id"] == "lab-model" assert "provider-secret" not in get_response.text + snapshot_response = client.post( + "/api/runtime-snapshots", + headers=headers, + json={ + "snapshot_id": "run-request-a", + "model": "lab-model", + "provider": "lab-openai", + }, + ) + assert snapshot_response.status_code == 200 + snapshot = snapshot_response.json()["snapshot"] + assert snapshot["runtime"]["max_input_tokens"] == 28_672 + assert "provider-secret" not in snapshot_response.text + + +def test_runtime_snapshot_reports_missing_zhipu_key_before_creating_a_run( + monkeypatch, +): + monkeypatch.setenv("EVOSCIENTIST_PROVIDER_ADMIN_TOKEN", "admin-secret") + monkeypatch.delenv("ZHIPU_API_KEY", raising=False) + monkeypatch.setenv("OPENAI_API_KEY", "unrelated-openai-key") + + response = client.post( + "/api/runtime-snapshots", + headers={"X-EvoScientist-Admin-Token": "admin-secret"}, + json={ + "snapshot_id": "run-request-zhipu", + "model": "glm-5.2", + "provider": "glm", + }, + ) + + assert response.status_code == 400 + assert response.json()["error"].startswith("ZHIPU_API_KEY_NOT_CONFIGURED") + def test_llm_config_api_requires_admin_token_header(monkeypatch): monkeypatch.delenv("EVOSCIENTIST_PROVIDER_ADMIN_TOKEN", raising=False) diff --git a/tests/test_llm.py b/tests/test_llm.py index 23168d0..25048fc 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -218,6 +218,29 @@ class TestGetChatModel: call_kwargs = mock_init.call_args[1] assert call_kwargs["model_provider"] == "custom_provider" + @patch("EvoScientist.llm.models.init_chat_model") + def test_glm_provider_alias_routes_to_zhipu(self, mock_init): + """The historic ``glm`` shorthand uses the Zhipu compatible endpoint.""" + mock_init.return_value = "mock_model" + + get_chat_model("glm-5.2", provider="glm") + + call_kwargs = mock_init.call_args[1] + assert call_kwargs["model"] == "glm-5.2" + assert call_kwargs["model_provider"] == "openai" + assert call_kwargs["base_url"] == "https://open.bigmodel.cn/api/paas/v4" + + @patch("EvoScientist.llm.models.init_chat_model") + def test_zhipu_does_not_fall_back_to_openai_api_key(self, mock_init, monkeypatch): + """A missing Zhipu key must not be sent from OPENAI_API_KEY instead.""" + monkeypatch.setenv("OPENAI_API_KEY", "unrelated-openai-key") + monkeypatch.delenv("ZHIPU_API_KEY", raising=False) + mock_init.return_value = "mock_model" + + get_chat_model("glm-5.2", provider="zhipu") + + assert mock_init.call_args.kwargs["api_key"] is None + @patch("EvoScientist.llm.models.init_chat_model") def test_passes_kwargs(self, mock_init): """Test that additional kwargs are passed through.""" diff --git a/tests/test_message_budget.py b/tests/test_message_budget.py new file mode 100644 index 0000000..4ee324c --- /dev/null +++ b/tests/test_message_budget.py @@ -0,0 +1,79 @@ +"""Tests for message-only context budgeting.""" + +from __future__ import annotations + +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +from langchain_core.messages import AIMessage, HumanMessage, ToolMessage + +from EvoScientist.middleware.message_budget import ( + ContextBudgetUnsatisfiableError, + MessageReservePolicy, + count_message_text_tokens, + create_message_budget_middleware, +) + + +def test_text_counter_excludes_attachment_payloads_and_counts_tool_results(): + messages = [ + HumanMessage( + content=[ + {"type": "text", "text": "abcd"}, + {"type": "image", "base64": "x" * 100_000}, + {"type": "file", "data": "y" * 100_000}, + ] + ), + ToolMessage(content="wxyz", tool_call_id="tool-1"), + ] + + 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 + + +def test_budget_middleware_uses_message_threshold_and_safe_tool_cutoff(): + model = MagicMock() + model.profile = {"max_input_tokens": 32_768} + middleware = create_message_budget_middleware(model, MagicMock()) + messages = [ + HumanMessage(content="a" * 30_000), + 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) + 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_provider_profiles.py b/tests/test_provider_profiles.py index aef9370..cafeaae 100644 --- a/tests/test_provider_profiles.py +++ b/tests/test_provider_profiles.py @@ -30,6 +30,7 @@ def provider_config_dir(tmp_path, monkeypatch): def _document(api_key: str = "secret-key") -> dict: return { + "version": 3, "providers": [ { "id": "lab-openai", @@ -38,18 +39,41 @@ def _document(api_key: str = "secret-key") -> dict: "base_url": "https://llm.example.test/v1/", "api_key": api_key, "enabled": True, + "runtime": { + "timeout_seconds": 90, + "max_retries": 1, + "default_temperature": 0.3, + "default_top_p": None, + "default_reasoning_effort": "auto", + }, "models": [ { "id": "research-model", "name": "Research Model", "model_id": "vendor/research-1", "enabled": True, + "runtime": { + "limit_mode": "combined", + "context_window_tokens": 32768, + "max_output_tokens": 4096, + "min_effective_input_tokens": 4096, + "limits_status": "confirmed", + "limits_source": "user", + }, }, { "id": "disabled-model", "name": "Disabled Model", "model_id": "vendor/disabled", "enabled": False, + "runtime": { + "limit_mode": "combined", + "context_window_tokens": 32768, + "max_output_tokens": 4096, + "min_effective_input_tokens": 4096, + "limits_status": "confirmed", + "limits_source": "user", + }, }, ], } @@ -59,6 +83,7 @@ def _document(api_key: str = "secret-key") -> dict: def _builtin_document(api_key: str = "builtin-secret") -> dict: return { + "version": 3, "builtins": [ { "id": "openai", @@ -68,12 +93,21 @@ def _builtin_document(api_key: str = "builtin-secret") -> dict: "api_key": api_key, "auth_mode": "api_key", "enabled": True, + "runtime": {}, "models": [ { "id": "chat-main", "name": "Chat Main", "model_id": "gpt-upstream", "enabled": True, + "runtime": { + "limit_mode": "combined", + "context_window_tokens": 32768, + "max_output_tokens": 4096, + "min_effective_input_tokens": 4096, + "limits_status": "confirmed", + "limits_source": "user", + }, } ], } @@ -99,7 +133,7 @@ def test_replace_round_trip_and_redacts_api_key(provider_config_dir): assert path.stat().st_mode & 0o777 == 0o600 -def test_v1_document_loads_and_saves_as_v2(provider_config_dir): +def test_v1_document_requires_reset(provider_config_dir): path = get_provider_profiles_path() path.parent.mkdir(parents=True) path.write_text( @@ -114,13 +148,8 @@ def test_v1_document_loads_and_saves_as_v2(provider_config_dir): encoding="utf-8", ) - loaded = load_provider_profiles() - assert loaded.version == 2 - assert loaded.builtins == () - assert loaded.providers[0].id == "lab-openai" - - replace_provider_profiles({"providers": _document()["providers"]}) - assert path.read_text(encoding="utf-8").startswith("version: 2\nbuiltins:") + with pytest.raises(ProviderProfileError, match="PROVIDER_PROFILE_RESET_REQUIRED"): + load_provider_profiles() def test_builtin_profiles_round_trip_and_redact_secret(provider_config_dir): @@ -300,6 +329,9 @@ def test_dynamic_profile_routes_through_selected_adapter(provider_config_dir): assert kwargs["model_provider"] == "openai" assert kwargs["base_url"] == "https://llm.example.test/v1" assert kwargs["api_key"] == "secret-key" + assert kwargs["timeout"] == 90 + assert kwargs["max_retries"] == 1 + assert kwargs["max_tokens"] == 4096 assert kwargs["default_headers"]["User-Agent"] == "codex_cli_rs/0.0.0" diff --git a/tests/test_runtime_snapshots.py b/tests/test_runtime_snapshots.py new file mode 100644 index 0000000..469bf15 --- /dev/null +++ b/tests/test_runtime_snapshots.py @@ -0,0 +1,152 @@ +"""Tests for server-side run configuration snapshots.""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +from EvoScientist.config.provider_profiles import ( + ProviderProfileError, + replace_provider_profiles, +) +from EvoScientist.llm.runtime_snapshots import ( + create_run_runtime_snapshot, + get_run_runtime_snapshot, + get_snapshot_chat_model, +) + + +@pytest.fixture +def provider_config_dir(tmp_path, monkeypatch): + monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path)) + return tmp_path / "evoscientist" + + +def _document(api_key: str = "first-secret", base_url: str = "https://one.test/v1"): + return { + "version": 3, + "providers": [ + { + "id": "lab-openai", + "name": "Lab OpenAI", + "adapter": "openai-compatible", + "base_url": base_url, + "api_key": api_key, + "enabled": True, + "runtime": {"timeout_seconds": 90, "max_retries": 1}, + "models": [ + { + "id": "research-model", + "name": "Research Model", + "model_id": "vendor/research-1", + "enabled": True, + "runtime": { + "limit_mode": "combined", + "context_window_tokens": 32_768, + "max_output_tokens": 4_096, + "min_effective_input_tokens": 4_096, + "limits_status": "confirmed", + "limits_source": "provider", + }, + } + ], + } + ], + } + + +def test_snapshot_keeps_connection_and_runtime_after_profile_changes(provider_config_dir): + replace_provider_profiles(_document()) + + created = create_run_runtime_snapshot( + "run-request-a", model="research-model", provider="lab-openai" + ) + assert created is not None + assert created.options.resolved_input_limit == 28_672 + assert created.profile.api_key == "first-secret" + assert "first-secret" not in str(created.public_payload()) + assert "first-secret" not in ( + provider_config_dir / "run-runtime-snapshots.sqlite3" + ).read_bytes().decode("utf-8", errors="ignore") + + replace_provider_profiles( + _document("second-secret", "https://two.test/v1") + ) + loaded = get_run_runtime_snapshot("run-request-a") + + assert loaded is not None + assert loaded.profile.api_key == "first-secret" + assert loaded.profile.base_url == "https://one.test/v1" + assert loaded.options.timeout_seconds == 90 + + +def test_snapshot_id_is_idempotent_and_cannot_change_model(provider_config_dir): + replace_provider_profiles(_document()) + + first = create_run_runtime_snapshot( + "run-request-a", model="research-model", provider="lab-openai" + ) + repeated = create_run_runtime_snapshot( + "run-request-a", model="research-model", provider="lab-openai" + ) + + assert repeated == first + with pytest.raises(ValueError, match="already bound"): + create_run_runtime_snapshot( + "run-request-a", model="another-model", provider="lab-openai" + ) + + +def test_static_provider_does_not_need_a_custom_snapshot(provider_config_dir): + assert ( + create_run_runtime_snapshot( + "run-request-a", model="gpt-5", provider="openai" + ) + is None + ) + + +def test_glm_alias_is_static_even_when_custom_registry_is_stale( + provider_config_dir, monkeypatch +): + monkeypatch.setenv("ZHIPU_API_KEY", "test-zhipu-key") + provider_config_dir.mkdir(parents=True) + (provider_config_dir / "providers.yaml").write_text( + "version: 2\nproviders: []\n", encoding="utf-8" + ) + + assert ( + create_run_runtime_snapshot( + "run-request-glm", model="glm-5.2", provider="glm" + ) + is None + ) + + +def test_zhipu_snapshot_requires_its_own_api_key(provider_config_dir, monkeypatch): + monkeypatch.delenv("ZHIPU_API_KEY", raising=False) + monkeypatch.setenv("OPENAI_API_KEY", "unrelated-openai-key") + + with pytest.raises(ProviderProfileError, match="ZHIPU_API_KEY_NOT_CONFIGURED"): + create_run_runtime_snapshot( + "run-request-zhipu", model="glm-5.2", provider="zhipu" + ) + + +def test_snapshot_model_builder_uses_frozen_profile(provider_config_dir): + replace_provider_profiles(_document()) + create_run_runtime_snapshot( + "run-request-a", model="research-model", provider="lab-openai" + ) + model = MagicMock() + + with patch( + "EvoScientist.llm.models.get_profile_chat_model", return_value=model + ) as factory: + resolved = get_snapshot_chat_model("run-request-a") + + assert resolved is model + profile, provider_model = factory.call_args.args + assert profile.api_key == "first-secret" + assert provider_model.model_id == "vendor/research-1"