From 5a581c78a23c48936f0d7db5b1952db88b5464a3 Mon Sep 17 00:00:00 2001 From: m4 Date: Fri, 14 Aug 2026 22:03:04 +0800 Subject: [PATCH] feat: add scoped model runtime configuration Introduce provider, model, and invocation contracts with encrypted configuration persistence. Add web runtime fencing, route fallback, recovery middleware, workspace scoping, and comprehensive tests. --- EvoScientist/EvoScientist.py | 189 +- EvoScientist/__init__.py | 1 + EvoScientist/backends.py | 13 +- EvoScientist/config/onboard/steps.py | 6 +- EvoScientist/config/settings.py | 39 +- EvoScientist/deploy/server.py | 2 +- EvoScientist/gateway/background_runs.py | 11 + EvoScientist/langgraph_dev/http.py | 346 ++ EvoScientist/langgraph_dev/langgraph.json | 2 +- EvoScientist/langgraph_dev/manager.py | 8 +- EvoScientist/langgraph_dev/sdk.py | 6 +- EvoScientist/llm/README.md | 50 + EvoScientist/llm/__init__.py | 55 +- EvoScientist/llm/adapter_registry.py | 1379 +++++++ EvoScientist/llm/config_admin.py | 2080 +++++++++++ EvoScientist/llm/configuration/__init__.py | 24 + EvoScientist/llm/configuration/model.py | 52 + EvoScientist/llm/configuration/provider.py | 61 + EvoScientist/llm/contracts.py | 870 +++++ EvoScientist/llm/crypto.py | 159 + EvoScientist/llm/errors.py | 87 +- EvoScientist/llm/gateway_proxy.py | 93 + EvoScientist/llm/gemini_interactions.py | 379 ++ EvoScientist/llm/invocation/__init__.py | 20 + EvoScientist/llm/invocation/contract.py | 133 + EvoScientist/llm/invocation/messages.py | 68 + EvoScientist/llm/model_config.py | 3309 +++++++++++++++++ EvoScientist/llm/model_config_v4.py | 2837 ++++++++++++++ EvoScientist/llm/models.py | 107 +- EvoScientist/llm/patches.py | 101 +- EvoScientist/llm/runtime.py | 3255 ++++++++++++++++ EvoScientist/llm/secret_store.py | 472 +++ EvoScientist/llm/user_options.py | 258 ++ EvoScientist/memory/agents/_factory.py | 9 +- EvoScientist/memory/agents/memory_worker.py | 4 + .../memory/agents/observation_linker.py | 9 +- EvoScientist/memory/launch.py | 84 +- EvoScientist/memory/observations/__init__.py | 3 +- EvoScientist/memory/observations/relations.py | 92 +- EvoScientist/memory/observations/store.py | 62 +- EvoScientist/memory/scheduler.py | 46 +- EvoScientist/memory/search.py | 14 +- EvoScientist/middleware/__init__.py | 18 + EvoScientist/middleware/configurable_model.py | 44 + EvoScientist/middleware/context_overflow.py | 2 + EvoScientist/middleware/disable_subagent.py | 77 + .../middleware/error_normalization.py | 18 +- EvoScientist/middleware/evo_route_fallback.py | 455 +++ EvoScientist/middleware/model_fallback.py | 34 +- EvoScientist/middleware/provider_context.py | 376 ++ .../middleware/recoverable_metering.py | 272 ++ EvoScientist/middleware/recoverable_tools.py | 205 + EvoScientist/middleware/skill_context.py | 161 + .../middleware/tool_call_normalizer.py | 464 +++ .../middleware/tool_protocol_guard.py | 86 +- EvoScientist/runtime_integrations.py | 8 - EvoScientist/scope_registry.py | 1279 +++++++ EvoScientist/sessions.py | 277 +- EvoScientist/stream/events.py | 14 +- EvoScientist/subagents/_factory.py | 11 +- EvoScientist/web_runtime.py | 170 + EvoScientist/workspace_scope.py | 541 +++ pyproject.toml | 1 + start-langgraph.sh | 30 + tests/__init__.py | 1 + tests/test_admin_control_v2.py | 162 + tests/test_agent_factory_extensions.py | 70 +- tests/test_async_subagent_factory.py | 8 +- tests/test_config.py | 18 +- tests/test_context_overflow_middleware.py | 10 + tests/test_error_normalization_middleware.py | 2 + tests/test_evo_route_fallback.py | 328 ++ tests/test_gateway_background_runs.py | 1 + tests/test_host_metering_extensions.py | 22 + tests/test_invocation_contract.py | 119 + tests/test_langgraph_dev_http.py | 54 + tests/test_llm.py | 236 +- tests/test_memory_agent_factory.py | 37 + tests/test_model_config_v3.py | 460 +++ tests/test_model_config_v4.py | 822 ++++ tests/test_model_fallback.py | 23 +- tests/test_model_secret_store.py | 48 + tests/test_observation_memory.py | 229 +- tests/test_provider_context_middleware.py | 176 + tests/test_provider_model_config_v3.py | 1037 ++++++ tests/test_recoverable_tools.py | 53 + tests/test_runtime_integrations.py | 15 - tests/test_sessions.py | 54 + tests/test_skill_context_middleware.py | 84 + tests/test_tool_protocol_guard.py | 191 +- tests/test_user_model_options.py | 103 + tests/test_user_options_runtime.py | 146 + tests/test_v3_contracts_and_fencing.py | 64 + tests/test_web_model_runtime.py | 1024 +++++ tests/test_web_tool_registry.py | 17 + tests/v3_fixtures.py | 161 + uv.lock | 11 + 97 files changed, 26670 insertions(+), 454 deletions(-) create mode 100644 EvoScientist/llm/README.md create mode 100644 EvoScientist/llm/adapter_registry.py create mode 100644 EvoScientist/llm/config_admin.py create mode 100644 EvoScientist/llm/configuration/__init__.py create mode 100644 EvoScientist/llm/configuration/model.py create mode 100644 EvoScientist/llm/configuration/provider.py create mode 100644 EvoScientist/llm/contracts.py create mode 100644 EvoScientist/llm/crypto.py create mode 100644 EvoScientist/llm/gateway_proxy.py create mode 100644 EvoScientist/llm/gemini_interactions.py create mode 100644 EvoScientist/llm/invocation/__init__.py create mode 100644 EvoScientist/llm/invocation/contract.py create mode 100644 EvoScientist/llm/invocation/messages.py create mode 100644 EvoScientist/llm/model_config.py create mode 100644 EvoScientist/llm/model_config_v4.py create mode 100644 EvoScientist/llm/runtime.py create mode 100644 EvoScientist/llm/secret_store.py create mode 100644 EvoScientist/llm/user_options.py create mode 100644 EvoScientist/middleware/disable_subagent.py create mode 100644 EvoScientist/middleware/evo_route_fallback.py create mode 100644 EvoScientist/middleware/provider_context.py create mode 100644 EvoScientist/middleware/recoverable_metering.py create mode 100644 EvoScientist/middleware/recoverable_tools.py create mode 100644 EvoScientist/middleware/skill_context.py create mode 100644 EvoScientist/middleware/tool_call_normalizer.py create mode 100644 EvoScientist/scope_registry.py create mode 100644 EvoScientist/web_runtime.py create mode 100644 EvoScientist/workspace_scope.py create mode 100755 start-langgraph.sh create mode 100644 tests/test_admin_control_v2.py create mode 100644 tests/test_evo_route_fallback.py create mode 100644 tests/test_invocation_contract.py create mode 100644 tests/test_memory_agent_factory.py create mode 100644 tests/test_model_config_v3.py create mode 100644 tests/test_model_config_v4.py create mode 100644 tests/test_model_secret_store.py create mode 100644 tests/test_provider_context_middleware.py create mode 100644 tests/test_provider_model_config_v3.py create mode 100644 tests/test_recoverable_tools.py create mode 100644 tests/test_skill_context_middleware.py create mode 100644 tests/test_user_model_options.py create mode 100644 tests/test_user_options_runtime.py create mode 100644 tests/test_v3_contracts_and_fencing.py create mode 100644 tests/test_web_model_runtime.py create mode 100644 tests/test_web_tool_registry.py create mode 100644 tests/v3_fixtures.py diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index 98356eb..7f8304e 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -308,6 +308,8 @@ def _inject_subagent_middleware( DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD, ContextOverflowMapperMiddleware, ErrorNormalizationMiddleware, + RecoverableMeteringMiddleware, + RecoverableToolEffectMiddleware, RepetitiveToolCallGuardMiddleware, ToolErrorHandlerMiddleware, ToolProtocolGuardMiddleware, @@ -353,6 +355,8 @@ def _inject_subagent_middleware( # them into a non-dataclass envelope wrapper before # anything downstream sees them. ErrorNormalizationMiddleware(), + RecoverableMeteringMiddleware(), + RecoverableToolEffectMiddleware(), RepetitiveToolCallGuardMiddleware( threshold=repetitive_tool_call_threshold, max_consecutive_errors=max_consecutive_tool_errors, @@ -399,6 +403,42 @@ def _ensure_general_purpose_subagent(subs: list[dict]) -> None: ) +def _apply_budgeted_skill_context(kwargs: dict, backend) -> dict: + """Replace DeepAgents' full-catalog skill prompts with bounded prompts.""" + + from .middleware import BudgetedSkillsMiddleware + + updated = dict(kwargs) + middleware = list(updated.get("middleware") or ()) + if not any(isinstance(item, BudgetedSkillsMiddleware) for item in middleware): + middleware.append( + BudgetedSkillsMiddleware( + backend=backend, + sources=list(DEFAULT_SKILL_SOURCES), + ) + ) + updated["middleware"] = middleware + updated["skills"] = None + + subagents = [] + for spec in updated.get("subagents") or (): + if not isinstance(spec, dict) or not spec.get("skills"): + subagents.append(spec) + continue + child = dict(spec) + raw_sources = child["skills"] + sources = [raw_sources] if isinstance(raw_sources, str) else list(raw_sources) + child["skills"] = None + child_middleware = list(child.get("middleware") or ()) + child_middleware.append( + BudgetedSkillsMiddleware(backend=backend, sources=sources) + ) + child["middleware"] = child_middleware + subagents.append(child) + updated["subagents"] = subagents + return updated + + def _maybe_swap_async_subagents( subs: list, middleware: list | None = None, *, cfg=None ) -> list: @@ -459,7 +499,9 @@ def _maybe_swap_async_subagents( from deepagents import AsyncSubAgent - port = int(getattr(cfg, "langgraph_dev_port", 6174)) + from .langgraph_dev.sdk import configured_langgraph_dev_url + + runtime_url = configured_langgraph_dev_url() out = [] agent_specs: dict[str, AsyncSubAgent] = {} # MCP tools routed to async sub-agents (via ``expose_to: `` in @@ -474,7 +516,7 @@ def _maybe_swap_async_subagents( name=name, description=async_specs[name], graph_id=name, - url=f"http://localhost:{port}", + url=runtime_url, ) agent_specs[name] = spec out.append(spec) @@ -619,8 +661,8 @@ def load_mcp_and_build_kwargs( # ============================================================================= -def _get_default_backend(): - """Build the default composite backend from current paths.""" +def _get_legacy_backend(): + """Build the deployment-root backend used outside Web full deploy.""" from deepagents.backends import CompositeBackend from .backends import ( @@ -662,6 +704,20 @@ def _get_default_backend(): ) +def _get_default_backend(): + """Use Origin's conversation-scoped backend for Web full deploy.""" + if os.environ.get("EVOSCIENTIST_DEPLOY_MODE", "").lower() != "full": + return _get_legacy_backend() + from .workspace_scope import create_workspace_backend_factory + + cfg = _ensure_config() + return create_workspace_backend_factory( + _get_legacy_backend, + dangerous=cfg.dangerous_mode, + allow_unscoped_legacy=False, + ) + + def _get_default_middleware( *, for_async_subagent: bool = False, @@ -674,6 +730,11 @@ def _get_default_middleware( memory_max_inline_profile_chars: int | None = None, enable_background_execution: bool = True, enable_legacy_model_fallback: bool = True, + tool_selector_model=None, + include_configurable_model: bool = True, + enable_scheduler: bool | None = None, + enable_memory_workers: bool | None = None, + install_subagent_guard: bool = False, ): """Build the default middleware list. @@ -698,8 +759,11 @@ def _get_default_middleware( DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD, ConfigurableModelMiddleware, ContextOverflowMapperMiddleware, + DisableSubagentToolMiddleware, ErrorNormalizationMiddleware, ModelFallbackMiddleware, + RecoverableMeteringMiddleware, + RecoverableToolEffectMiddleware, RepetitiveToolCallGuardMiddleware, ToolErrorHandlerMiddleware, ToolProtocolGuardMiddleware, @@ -761,26 +825,30 @@ def _get_default_middleware( # keep the main model (they do real work, not a one-off helper call). # context_editing stays on the main model — its model only sizes the # context-window trigger for the main agent's own history. - if for_async_subagent: - tool_selector_model = model + if tool_selector_model is not None: + resolved_selector_model = tool_selector_model + elif for_async_subagent: + resolved_selector_model = model elif chat_model is None: - tool_selector_model = _ensure_auxiliary_chat_model() + resolved_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 + resolved_selector_model = model else: from .llm import get_chat_model - tool_selector_model = get_chat_model(model=aux_model, provider=aux_provider) + resolved_selector_model = get_chat_model( + model=aux_model, provider=aux_provider + ) selector_middlewares = create_tool_selector_middleware( **( {"threshold": tool_selector_threshold} if tool_selector_threshold is not None else {} ), - model=tool_selector_model, + model=resolved_selector_model, track_stream_selection=not for_async_subagent, ) mw = [ @@ -789,7 +857,8 @@ def _get_default_middleware( # middlewares) and normalizes them into a non-dataclass # envelope wrapper before anything downstream sees them. ErrorNormalizationMiddleware(), - ConfigurableModelMiddleware(), + RecoverableMeteringMiddleware(), + RecoverableToolEffectMiddleware(), create_context_editing_middleware(model), *([ModelFallbackMiddleware()] if enable_legacy_model_fallback else []), RepetitiveToolCallGuardMiddleware( @@ -807,12 +876,18 @@ def _get_default_middleware( max_result_chars=cfg.code_interpreter_max_result_chars, ), ] - if cfg.enable_scheduler and not for_async_subagent: + if include_configurable_model: + mw.insert(1, ConfigurableModelMiddleware()) + if enable_scheduler is None: + enable_scheduler = bool(cfg.enable_scheduler) + if enable_scheduler and not for_async_subagent: mw.append(create_scheduler_middleware()) mw.append(create_runtime_context_middleware()) if memory_controls.memory_enabled: mw.append(memory_middleware) - if memory_controls.worker_needed(worker_target): + if enable_memory_workers is not False and memory_controls.worker_needed( + worker_target + ): mw.append( create_memory_lifecycle_middleware( memory_dir, @@ -837,6 +912,9 @@ def _get_default_middleware( mw.append(BackgroundExecutionMiddleware()) + if install_subagent_guard: + mw.append(DisableSubagentToolMiddleware()) + return mw @@ -898,6 +976,7 @@ def _get_default_agent(): mw, workspace_dir=str(_paths_mod.WORKSPACE_ROOT), ) + kwargs = _apply_budgeted_skill_context(kwargs, be) _EvoScientist_agent = create_deep_agent( **kwargs, @@ -938,6 +1017,8 @@ def create_cli_agent( enable_background_execution: bool = True, main_agent_outer_middlewares: Sequence[AgentMiddleware] | None = None, main_agent_route_middleware: AgentMiddleware | None = None, + execution_profile=None, + agent_model_set=None, ) -> "CompiledStateGraph": """Create agent with checkpointer for CLI multi-turn support. @@ -1004,6 +1085,24 @@ def create_cli_agent( cfg = _ensure_config(config) chat_model = None + profile = execution_profile + if agent_model_set is not None: + chat_model = agent_model_set.main_agent + if profile is not None: + import copy + + cfg = copy.copy(cfg) + cfg.enable_async_subagents = bool(profile.async_subagents) + cfg.enable_scheduler = bool(profile.scheduler) + cfg.memory_workers_enabled = bool(profile.memory_workers) + cfg.enable_ask_user = False + cfg.auto_mode = True + cfg.auto_approve = True + enable_subagents = bool(enable_subagents and profile.subagents) + enable_background_execution = bool( + enable_background_execution and profile.background_execution + ) + if checkpointer is None: from langgraph.checkpoint.memory import InMemorySaver @@ -1056,15 +1155,53 @@ def create_cli_agent( # Delegate middleware construction to the single source of truth so the # 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, - memory_dir=_mem_dir, - cfg=cfg, - chat_model=chat_model, - tool_selector_threshold=tool_selector_threshold, - memory_max_inline_profile_chars=memory_max_inline_profile_chars, - enable_background_execution=enable_background_execution, - enable_legacy_model_fallback=main_agent_route_middleware is None, + mw: list[AgentMiddleware] = list( + _get_default_middleware( + workspace_dir=workspace_dir, + memory_dir=_mem_dir, + cfg=cfg, + chat_model=chat_model, + tool_selector_threshold=tool_selector_threshold, + memory_max_inline_profile_chars=memory_max_inline_profile_chars, + enable_background_execution=enable_background_execution, + enable_legacy_model_fallback=( + main_agent_route_middleware is None + and not ( + profile is not None + and getattr(profile, "name", "") in {"web_v1", "web_v3"} + ) + ), + tool_selector_model=( + agent_model_set.tool_selector if agent_model_set is not None else None + ), + include_configurable_model=( + bool(profile.configurable_model_override) + if profile is not None + else True + ), + enable_scheduler=(bool(profile.scheduler) if profile is not None else None), + enable_memory_workers=( + bool(profile.memory_workers) if profile is not None else None + ), + install_subagent_guard=(profile is not None and not profile.subagents), + ) + ) + from .middleware import ProviderContextMediaMiddleware + + # Keep assistant-generated binary output out of both provider history and + # future checkpoints. The middleware persists media through the same + # workspace backend before replacing it with a content-addressed reference. + error_index = next( + ( + index + for index, middleware in enumerate(mw) + if getattr(middleware, "name", "") == "error_normalization" + ), + None, + ) + mw.insert( + (error_index + 1) if error_index is not None else 0, + ProviderContextMediaMiddleware(be), ) if main_agent_route_middleware is not None: configurable_index = next( @@ -1075,9 +1212,10 @@ def create_cli_agent( ), None, ) - if configurable_index is None: - raise RuntimeError("ConfigurableModelMiddleware route slot is unavailable") - mw.insert(configurable_index + 1, main_agent_route_middleware) + mw.insert( + (configurable_index + 1) if configurable_index is not None else 1, + main_agent_route_middleware, + ) if main_agent_outer_middlewares: mw = [*main_agent_outer_middlewares, *mw] @@ -1106,6 +1244,7 @@ def create_cli_agent( ) if not enable_subagents: kwargs = {**kwargs, "subagents": []} + kwargs = _apply_budgeted_skill_context(kwargs, be) return create_deep_agent( **kwargs, diff --git a/EvoScientist/__init__.py b/EvoScientist/__init__.py index a58e6aa..6eaebe2 100644 --- a/EvoScientist/__init__.py +++ b/EvoScientist/__init__.py @@ -30,6 +30,7 @@ _EXPORTS: dict[str, tuple[str, str]] = { "MODELS": (".llm", "MODELS"), "list_models": (".llm", "list_models"), "DEFAULT_MODEL": (".llm", "DEFAULT_MODEL"), + "EvoModelRuntime": (".llm.runtime", "EvoModelRuntime"), # Prompts "get_system_prompt": (".prompts", "get_system_prompt"), # Tools diff --git a/EvoScientist/backends.py b/EvoScientist/backends.py index 8e680b2..97b1919 100644 --- a/EvoScientist/backends.py +++ b/EvoScientist/backends.py @@ -20,6 +20,7 @@ from deepagents.backends.protocol import ( LsResult, WriteResult, ) +from filelock import FileLock from . import paths @@ -835,6 +836,15 @@ class MemoryFilesystemBackend(FilesystemBackend): "/memories/profile/... files. Use memory tools for observations." ) + def __init__( + self, + root_dir: str | Path | None = None, + virtual_mode: bool | None = None, + max_file_size_mb: int = 10, + ) -> None: + super().__init__(root_dir, virtual_mode, max_file_size_mb) + self._profile_write_lock = FileLock(str(self.cwd / ".profile-write.lock")) + @staticmethod def _is_profile_path(file_path: str) -> bool: normalized = posixpath.normpath("/" + file_path.strip().lstrip("/")) @@ -852,7 +862,8 @@ class MemoryFilesystemBackend(FilesystemBackend): ) -> EditResult: if not self._is_profile_path(file_path): return EditResult(error=self._RAW_EDIT_ERROR) - return super().edit(file_path, old_string, new_string, replace_all) + with self._profile_write_lock: + return super().edit(file_path, old_string, new_string, replace_all) def upload_files(self, files: list[tuple[str, bytes]]) -> list[FileUploadResponse]: return [ diff --git a/EvoScientist/config/onboard/steps.py b/EvoScientist/config/onboard/steps.py index 184e0bd..24987ab 100644 --- a/EvoScientist/config/onboard/steps.py +++ b/EvoScientist/config/onboard/steps.py @@ -90,11 +90,11 @@ def _step_langgraph_dev_port(config: EvoScientistConfig) -> int: """ if not getattr(config, "enable_async_subagents", True): # User has async disabled — port is irrelevant, no prompt. - return getattr(config, "langgraph_dev_port", 6174) + return getattr(config, "langgraph_dev_port", 3076) from ...langgraph_dev.manager import _is_port_occupied, is_langgraph_dev_running - current_port = getattr(config, "langgraph_dev_port", 6174) + current_port = getattr(config, "langgraph_dev_port", 3076) current_occupied = _is_port_occupied(current_port) if current_occupied and is_langgraph_dev_running(port=current_port): # Another EvoSci shell is already serving on this port — reuse, don't @@ -174,7 +174,7 @@ def _step_webui_port(config: EvoScientistConfig) -> int: from ...langgraph_dev.manager import _is_port_occupied current_port = getattr(config, "webui_port", 4716) - backend_port = getattr(config, "langgraph_dev_port", 6174) + backend_port = getattr(config, "langgraph_dev_port", 3076) occupied = _is_port_occupied(current_port) conflicts_backend = current_port == backend_port diff --git a/EvoScientist/config/settings.py b/EvoScientist/config/settings.py index 8188f2d..216a884 100644 --- a/EvoScientist/config/settings.py +++ b/EvoScientist/config/settings.py @@ -156,6 +156,7 @@ class EvoScientistConfig: nvidia_api_key: NVIDIA API key for NVIDIA models. google_api_key: Google API key for Gemini models. tavily_api_key: Tavily API key for web search. + semantic_scholar_api_key: Semantic Scholar API key for paper-navigator. provider: Default LLM provider ('anthropic', 'openai', 'google-genai', or 'nvidia'). model: Default model name (short name or full ID). auxiliary_provider: Provider for auxiliary_model (empty = use main provider). @@ -189,6 +190,7 @@ class EvoScientistConfig: custom_anthropic_base_url: str = "" ollama_base_url: str = "" tavily_api_key: str = "" + semantic_scholar_api_key: str = "" # LLM Settings provider: str = "anthropic" @@ -212,15 +214,12 @@ class EvoScientistConfig: # synchronous sub-agents (planner / research / code / debug). enable_async_subagents: bool = True - # Port for the auto-started langgraph dev subprocess. 6174 is Kaprekar's - # constant — a memorable EvoScientist-themed default that avoids collisions - # with common dev ports (3000/5000/8000/8080) and the langgraph CLI default - # 2024. Override if it conflicts with another local service. - langgraph_dev_port: int = 6174 + # Port for the auto-started langgraph dev subprocess. Keep this aligned with + # the Ai4Sci-Web Gateway's recoverable runtime URL. + langgraph_dev_port: int = 3076 # Port for the WebUI front-end (Next.js server from @evoscientist/webui), - # used only when ui_backend == "webui". 4716 is 6174 reversed — a memorable - # pairing with the langgraph dev port that it connects to. The backend keeps + # used only when ui_backend == "webui". The backend keeps # its own port (langgraph_dev_port); this is just the browser server. webui_port: int = 4716 @@ -256,11 +255,10 @@ class EvoScientistConfig: # (sessions.db), ContextEditingMiddleware (window management), and # EvoMemoryMiddleware (cross-turn memory). # - # 1,000,000 is "effectively unlimited" — typical research turns use - # 200-1000 steps; reaching 1M would cost ~$10K in tokens, by which point - # rate limits, context overflow, or API quota errors would trip first. - # Lower (e.g., 5000) if you want a tighter safety net against runaway loops. - recursion_limit: int = 1_000_000 + # Typical research turns use 200-1000 steps. 5,000 leaves room for + # legitimate long-running work while still stopping a runaway graph. + # This is a control-flow limit, not a model-call or cost budget. + recursion_limit: int = 5_000 # Number of consecutive model rounds with the same structured tool name and # arguments that activates provider-facing loop repair. Set 0 to disable. @@ -307,9 +305,6 @@ class EvoScientistConfig: # a deploy-style langgraph server instead of the in-terminal CLI/TUI. ui_backend: Literal["cli", "tui", "webui"] = "tui" log_level: str = "warning" - # Empty means use the provider/model default. A non-empty value is an - # explicit user override exported as EVOSCIENTIST_REASONING_EFFORT. - reasoning_effort: str = "" # Anthropic prompt caching for OpenRouter anthropic/* models. Opt out if # cache-write costs outweigh the benefit for a workflow. openrouter_anthropic_prompt_cache: bool = True @@ -466,9 +461,6 @@ class EvoScientistConfig: # DM access control policy dm_policy: str = "allowlist" - # OpenAI API mode - "" = auto, "true" = force Responses, "false" = force Completions - use_responses_api: str = "" - # ccproxy ccproxy_port: int = 8000 @@ -803,6 +795,7 @@ _ENV_MAPPINGS = { "custom_anthropic_base_url": "CUSTOM_ANTHROPIC_BASE_URL", "ollama_base_url": "OLLAMA_BASE_URL", "tavily_api_key": "TAVILY_API_KEY", + "semantic_scholar_api_key": "S2_API_KEY", "default_mode": "EVOSCIENTIST_DEFAULT_MODE", "default_workdir": "EVOSCIENTIST_WORKSPACE_DIR", "ui_backend": "EVOSCIENTIST_UI_BACKEND", @@ -810,7 +803,6 @@ _ENV_MAPPINGS = { "model_fallbacks": "EVOSCIENTIST_MODEL_FALLBACKS", "auxiliary_provider": "EVOSCIENTIST_AUXILIARY_PROVIDER", "auxiliary_model": "EVOSCIENTIST_AUXILIARY_MODEL", - "reasoning_effort": "EVOSCIENTIST_REASONING_EFFORT", "openrouter_anthropic_prompt_cache": ( "EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE" ), @@ -820,7 +812,6 @@ _ENV_MAPPINGS = { "dangerous_mode": "EVOSCIENTIST_DANGEROUS_MODE", "channel_debug_tracing": "EVOSCIENTIST_CHANNEL_DEBUG_TRACING", "ccproxy_port": "EVOSCIENTIST_CCPROXY_PORT", - "use_responses_api": "EVOSCIENTIST_USE_RESPONSES_API", "checkpoint_keep_per_thread": "EVOSCIENTIST_CHECKPOINT_KEEP_PER_THREAD", "enable_async_subagents": "EVOSCIENTIST_ENABLE_ASYNC_SUBAGENTS", "langgraph_dev_port": "EVOSCIENTIST_LANGGRAPH_DEV_PORT", @@ -950,8 +941,8 @@ def apply_config_to_env(config: EvoScientistConfig) -> None: os.environ["OLLAMA_BASE_URL"] = config.ollama_base_url if config.tavily_api_key and not os.environ.get("TAVILY_API_KEY"): os.environ["TAVILY_API_KEY"] = config.tavily_api_key - if config.reasoning_effort and not os.environ.get("EVOSCIENTIST_REASONING_EFFORT"): - os.environ["EVOSCIENTIST_REASONING_EFFORT"] = config.reasoning_effort + if config.semantic_scholar_api_key and not os.environ.get("S2_API_KEY"): + os.environ["S2_API_KEY"] = config.semantic_scholar_api_key if config.openrouter_http_referer and not os.environ.get( "EVOSCIENTIST_OPENROUTER_HTTP_REFERER" ): @@ -982,7 +973,3 @@ def apply_config_to_env(config: EvoScientistConfig) -> None: os.environ["EVOSCIENTIST_DANGEROUS_MODE"] = "true" else: os.environ.pop("EVOSCIENTIST_DANGEROUS_MODE", None) - if config.use_responses_api and not os.environ.get( - "EVOSCIENTIST_USE_RESPONSES_API" - ): - os.environ["EVOSCIENTIST_USE_RESPONSES_API"] = config.use_responses_api diff --git a/EvoScientist/deploy/server.py b/EvoScientist/deploy/server.py index 2af9a0b..cb7e52a 100644 --- a/EvoScientist/deploy/server.py +++ b/EvoScientist/deploy/server.py @@ -44,7 +44,7 @@ def deploy( port: int | None = typer.Option( None, "--port", - help="Port for langgraph dev (default: config.langgraph_dev_port = 6174)", + help="Port for langgraph dev (default: config.langgraph_dev_port = 3076)", ), tunnel: bool = typer.Option( False, diff --git a/EvoScientist/gateway/background_runs.py b/EvoScientist/gateway/background_runs.py index 5659dbc..6b74403 100644 --- a/EvoScientist/gateway/background_runs.py +++ b/EvoScientist/gateway/background_runs.py @@ -141,6 +141,7 @@ class BackgroundRun: run_id: str assistant_id: str metadata: Mapping[str, str] + configurable: Mapping[str, object] | None = None @dataclass(frozen=True) @@ -278,6 +279,7 @@ def _background_run_handle( run_id: str, payload: BackgroundRunPayload, ) -> BackgroundRun: + configurable = payload["config"].get("configurable") return BackgroundRun( name=request.name, url=url, @@ -286,6 +288,9 @@ def _background_run_handle( run_id=run_id, assistant_id=payload["assistant_id"], metadata=dict(payload["metadata"]), + configurable=( + dict(configurable) if isinstance(configurable, Mapping) else None + ), ) @@ -528,6 +533,7 @@ def spawn_background_run_status_thread( "graph_id": run.graph_id, "assistant_id": run.assistant_id, "metadata": run.metadata, + "configurable": run.configurable, "name": run.name, "headers": headers, "hooks": hooks, @@ -547,6 +553,7 @@ def watch_background_run_sync( graph_id: str = "", assistant_id: str = "", metadata: Mapping[str, str] | None = None, + configurable: Mapping[str, object] | None = None, name: str = "background run", headers: Mapping[str, str] | None = None, hooks: BackgroundRunHooks | None = None, @@ -565,6 +572,7 @@ def watch_background_run_sync( run_id=run_id, assistant_id=assistant_id, metadata=dict(metadata or {}), + configurable=dict(configurable or {}), ) failures = 0 confirmed_finished = False @@ -632,6 +640,7 @@ def spawn_background_run_status_task( graph_id=run.graph_id, assistant_id=run.assistant_id, metadata=run.metadata, + configurable=run.configurable, name=run.name, hooks=hooks, watcher_config=watcher_config, @@ -650,6 +659,7 @@ async def awatch_background_run( graph_id: str = "", assistant_id: str = "", metadata: Mapping[str, str] | None = None, + configurable: Mapping[str, object] | None = None, name: str = "background run", hooks: BackgroundRunHooks | None = None, watcher_config: BackgroundRunWatcherConfig | None = None, @@ -665,6 +675,7 @@ async def awatch_background_run( run_id=run_id, assistant_id=assistant_id, metadata=dict(metadata or {}), + configurable=dict(configurable or {}), ) failures = 0 confirmed_finished = False diff --git a/EvoScientist/langgraph_dev/http.py b/EvoScientist/langgraph_dev/http.py index 0d9c6f6..c7dcf55 100644 --- a/EvoScientist/langgraph_dev/http.py +++ b/EvoScientist/langgraph_dev/http.py @@ -23,6 +23,13 @@ memory. from __future__ import annotations import asyncio +import hashlib +import json +import os +import secrets +from pathlib import Path, PurePosixPath +from typing import Any +from uuid import UUID from starlette.applications import Starlette from starlette.requests import Request @@ -32,6 +39,23 @@ from starlette.routing import Route from EvoScientist.config import get_effective_config from EvoScientist.llm.models import list_model_picker_entries +_recoverable_run_lock = asyncio.Lock() + + +def _load_scope_service_token() -> str: + configured = ( + os.getenv("EVOSCIENTIST_BACKEND_SERVICE_TOKEN", "").strip() + or os.getenv("AI4SCI_EVO_RUNTIME_GRANT_SECRET", "").strip() + ) + if configured: + return configured + from EvoScientist.scope_registry import get_scope_service_token + + return get_scope_service_token() + + +_SCOPE_SERVICE_TOKEN = _load_scope_service_token() + async def get_models(_request: Request) -> JSONResponse: """Return the model registry as ``{entries, default}``. @@ -75,8 +99,330 @@ async def get_models(_request: Request) -> JSONResponse: ) +async def recoverable_run_capabilities(_request: Request) -> JSONResponse: + """Capabilities required by Ai4Sci's durable dispatch outbox.""" + + return JSONResponse( + { + "version": 1, + "deterministic_run_id": True, + "stream_resumable": True, + "durability_sync": True, + "multitask_enqueue": True, + "interrupt_resume": True, + "pending_interrupt_state": True, + "workspace_scope_v1": os.getenv("EVOSCIENTIST_DEPLOY_MODE", "").lower() == "full", + } + ) + + +def _scope_service_authorized(request: Request) -> JSONResponse | None: + if not _SCOPE_SERVICE_TOKEN: + return JSONResponse({"code": "WORKSPACE_SERVICE_UNAVAILABLE"}, status_code=503) + header = request.headers.get("authorization", "") + if not header.startswith("Bearer ") or not secrets.compare_digest( + header[7:], _SCOPE_SERVICE_TOKEN + ): + return JSONResponse({"code": "UNAUTHORIZED"}, status_code=401) + return None + + +def _scope_payload(record: Any) -> dict[str, Any]: + return { + "deployment_id": record.deployment_id, + "scope_id": record.scope_id, + "primary_thread_id": record.primary_thread_id, + "primary_owner_id": record.primary_owner_id, + "state": record.state, + "revision": record.revision, + } + + +def _run_payload(run: Any) -> dict[str, Any]: + return { + "run_request_id": run.run_request_id, + "turn_id": run.turn_id, + "interrupt_key": run.interrupt_key, + "request_hash": run.request_hash, + "run_owner_id": run.run_owner_id, + "run_id": run.run_id, + "state": run.state, + } + + +def _registry_call(method: str, *args: Any, **kwargs: Any) -> Any: + from EvoScientist.scope_registry import get_scope_registry + from EvoScientist.workspace_scope import current_deployment_id + + return getattr(get_scope_registry(), method)(current_deployment_id(), *args, **kwargs) + + +def _provision_scope(thread_id: str) -> Any: + from EvoScientist.workspace_scope import ( + current_deployment_id, + provision_conversation_scope, + ) + + return provision_conversation_scope(thread_id, deployment_id=current_deployment_id()) + + +async def provision_workspace_scope(request: Request) -> JSONResponse: + if denied := _scope_service_authorized(request): + return denied + try: + payload = await request.json() + except json.JSONDecodeError: + payload = None + if not isinstance(payload, dict) or not isinstance(payload.get("thread_id"), str): + return JSONResponse({"code": "INVALID_REQUEST"}, status_code=400) + try: + record = await asyncio.to_thread(_provision_scope, payload["thread_id"]) + except Exception as exc: + return JSONResponse({"code": "WORKSPACE_SCOPE_CONFLICT", "message": str(exc)}, status_code=409) + return JSONResponse(_scope_payload(record), status_code=201) + + +async def get_workspace_scope(request: Request) -> JSONResponse: + if denied := _scope_service_authorized(request): + return denied + try: + record = await asyncio.to_thread( + _registry_call, "get_by_thread", str(request.path_params["thread_id"]) + ) + except Exception as exc: + return JSONResponse({"code": "WORKSPACE_SCOPE_NOT_FOUND", "message": str(exc)}, status_code=404) + return JSONResponse(_scope_payload(record)) + + +async def reserve_workspace_run(request: Request) -> JSONResponse: + if denied := _scope_service_authorized(request): + return denied + try: + payload = await request.json() + except json.JSONDecodeError: + payload = None + if not isinstance(payload, dict): + return JSONResponse({"code": "INVALID_REQUEST"}, status_code=400) + try: + run = await asyncio.to_thread( + _registry_call, + "reserve_run", + str(request.path_params["scope_id"]), + str(payload["run_request_id"]), + str(payload["turn_id"]), + str(payload["request_hash"]), + interrupt_key=( + str(payload["interrupt_key"]) if payload.get("interrupt_key") else None + ), + ) + except Exception as exc: + code = ( + "INTERRUPT_ALREADY_RESOLVED" + if type(exc).__name__ == "ScopeInterruptResolvedError" + else "WORKSPACE_RUN_CONFLICT" + ) + return JSONResponse({"code": code, "message": str(exc)}, status_code=409) + return JSONResponse(_run_payload(run), status_code=201) + + +async def bind_workspace_run(request: Request) -> JSONResponse: + if denied := _scope_service_authorized(request): + return denied + try: + payload = await request.json() + except json.JSONDecodeError: + payload = None + if not isinstance(payload, dict) or not isinstance(payload.get("run_id"), str): + return JSONResponse({"code": "INVALID_REQUEST"}, status_code=400) + try: + run = await asyncio.to_thread( + _registry_call, + "bind_run", + str(request.path_params["scope_id"]), + str(request.path_params["run_request_id"]), + payload["run_id"], + ) + except Exception as exc: + return JSONResponse({"code": "WORKSPACE_RUN_CONFLICT", "message": str(exc)}, status_code=409) + return JSONResponse(_run_payload(run)) + + +def _materialize_target(scope_id: str, raw_path: str) -> Path: + from EvoScientist.workspace_scope import conversation_files_dir + + path = PurePosixPath(raw_path.replace("\\", "/")) + if path.is_absolute() or not path.parts or path.parts[0] != "uploads": + raise ValueError("only uploads/ paths are accepted") + if any(part in {"", ".", ".."} for part in path.parts): + raise ValueError("invalid upload path") + root = conversation_files_dir(scope_id).resolve() + target = root.joinpath(*path.parts) + target.parent.mkdir(parents=True, exist_ok=True) + try: + target.parent.resolve().relative_to(root) + except ValueError as exc: + raise ValueError("upload path escapes workspace scope") from exc + current = root + for part in path.parts[:-1]: + current = current / part + if current.is_symlink(): + raise ValueError("symlink parents are rejected") + if target.is_symlink(): + raise ValueError("symlink targets are rejected") + return target + + +async def materialize_workspace_file(request: Request) -> JSONResponse: + if denied := _scope_service_authorized(request): + return denied + scope_id = str(request.path_params["scope_id"]) + try: + await asyncio.to_thread(_registry_call, "get", scope_id) + target = await asyncio.to_thread( + _materialize_target, scope_id, str(request.path_params["path"]) + ) + except Exception as exc: + return JSONResponse({"code": "WORKSPACE_PATH_INVALID", "message": str(exc)}, status_code=400) + expected_hash = request.headers.get("x-content-sha256", "").lower() + expected_size = int(request.headers.get("content-length") or 0) + if expected_size > 100 * 1024 * 1024: + return JSONResponse({"code": "WORKSPACE_FILE_TOO_LARGE"}, status_code=413) + temporary = target.with_name(f".{target.name}.{secrets.token_hex(8)}.tmp") + digest = hashlib.sha256() + size = 0 + try: + with temporary.open("xb") as handle: + async for chunk in request.stream(): + size += len(chunk) + if size > 100 * 1024 * 1024: + raise ValueError("workspace file exceeds 100 MiB") + digest.update(chunk) + handle.write(chunk) + handle.flush() + os.fsync(handle.fileno()) + actual_hash = digest.hexdigest() + if expected_hash and not secrets.compare_digest(actual_hash, expected_hash): + raise ValueError("workspace file hash mismatch") + os.replace(temporary, target) + except Exception as exc: + temporary.unlink(missing_ok=True) + return JSONResponse({"code": "WORKSPACE_FILE_INVALID", "message": str(exc)}, status_code=409) + return JSONResponse({"virtual_path": str(request.path_params["path"]), "size": size, "sha256": actual_hash}) + + +async def create_recoverable_run(request: Request) -> JSONResponse: + """Create a LangGraph Run with a caller-owned deterministic UUID. + + LangGraph's public create endpoint always generates its own UUID. This + adapter performs lookup and insertion while holding the process-wide run + creation lock and passes the durable request UUID to ``create_valid_run``. + Retrying after a lost HTTP response therefore cannot create another Run. + """ + + if request.headers.get("x-auth-scheme") != "langsmith": + return JSONResponse({"code": "UNAUTHORIZED"}, status_code=401) + value = await request.json() + if not isinstance(value, dict): + return JSONResponse({"code": "INVALID_REQUEST"}, status_code=400) + try: + thread_id = str(UUID(str(value["thread_id"]))) + run_id = UUID(str(value["run_id"])) + run_request_id = str(UUID(str(value["run_request_id"]))) + request_hash = str(value["request_hash"]) + assistant_id = str(value["assistant_id"]) + operation = str(value.get("operation") or "start") + except (KeyError, TypeError, ValueError): + return JSONResponse({"code": "INVALID_REQUEST"}, status_code=400) + if ( + str(run_id) != run_request_id + or len(request_hash) != 64 + or operation not in {"start", "resume"} + ): + return JSONResponse({"code": "INVALID_IDEMPOTENCY_KEY"}, status_code=400) + command = value.get("command") + if operation == "resume": + if ( + value.get("input") is not None + or not isinstance(command, dict) + or set(command) != {"resume"} + ): + return JSONResponse({"code": "INVALID_RESUME_REQUEST"}, status_code=400) + elif command is not None: + return JSONResponse({"code": "INVALID_START_REQUEST"}, status_code=400) + + from langgraph_api.models.run import Runs, create_valid_run + from langgraph_api.utils import fetchone + from langgraph_runtime.database import connect + + payload = { + "assistant_id": assistant_id, + "input": value.get("input"), + "command": command, + "metadata": value.get("metadata") or {}, + "config": value.get("config") or {}, + "stream_mode": value.get("stream_mode") or ["messages", "updates", "tasks", "custom"], + "stream_resumable": True, + "durability": "sync", + "multitask_strategy": "enqueue", + "if_not_exists": "create", + } + payload["metadata"] = { + **payload["metadata"], + "run_request_id": run_request_id, + "request_hash": request_hash, + } + async with _recoverable_run_lock: + async with connect() as conn: + existing_iter = await Runs.get(conn, run_id, thread_id=UUID(thread_id)) + try: + existing = await fetchone(existing_iter) + except Exception as exc: + if getattr(exc, "status_code", None) != 404: + raise + existing = None + if existing is not None: + metadata = existing.get("metadata") or {} + if metadata.get("request_hash") != request_hash: + return JSONResponse( + { + "code": "RUN_REQUEST_CONFLICT", + "message": "run_request_id is bound to another request hash", + }, + status_code=409, + ) + return JSONResponse( + {"run_id": str(existing["run_id"]), "status": existing["status"], "created": False} + ) + created = await create_valid_run( + conn, + thread_id, + payload, + dict(request.headers), + run_id=run_id, + ) + return JSONResponse( + {"run_id": str(created["run_id"]), "status": created["status"], "created": True}, + status_code=201, + ) + + app = Starlette( routes=[ Route("/api/models", get_models, methods=["GET"]), + Route( + "/api/ai4sci/recoverable-runs/capabilities", + recoverable_run_capabilities, + methods=["GET"], + ), + Route( + "/api/ai4sci/recoverable-runs/create", + create_recoverable_run, + methods=["POST"], + ), + Route("/internal/workspace-scopes/provision", provision_workspace_scope, methods=["POST"]), + Route("/internal/workspace-scopes/by-thread/{thread_id}", get_workspace_scope, methods=["GET"]), + Route("/internal/workspace-scopes/{scope_id}/runs/reserve", reserve_workspace_run, methods=["POST"]), + Route("/internal/workspace-scopes/{scope_id}/runs/{run_request_id}", bind_workspace_run, methods=["PATCH"]), + Route("/internal/workspace-scopes/{scope_id}/files/{path:path}", materialize_workspace_file, methods=["PUT"]), ] ) diff --git a/EvoScientist/langgraph_dev/langgraph.json b/EvoScientist/langgraph_dev/langgraph.json index c834540..56f0e56 100644 --- a/EvoScientist/langgraph_dev/langgraph.json +++ b/EvoScientist/langgraph_dev/langgraph.json @@ -15,7 +15,7 @@ "path": "EvoScientist.sessions.create_checkpointer_for_langgraph_api" }, "config": { - "recursion_limit": 1000000 + "recursion_limit": 5000 }, "http": { "app": "EvoScientist.langgraph_dev.http:app" diff --git a/EvoScientist/langgraph_dev/manager.py b/EvoScientist/langgraph_dev/manager.py index 69e5068..30d2fd9 100644 --- a/EvoScientist/langgraph_dev/manager.py +++ b/EvoScientist/langgraph_dev/manager.py @@ -110,11 +110,11 @@ def needs_langgraph_dev(config: EvoScientistConfig) -> bool: _LOCK = threading.RLock() -# Default port (Kaprekar's constant — see config/settings.py for the rationale). +# Default port shared with the Ai4Sci-Web recoverable runtime. # Overridable per-call via ``start_langgraph_dev(port=...)`` / # ``ensure_langgraph_dev`` (which reads ``config.langgraph_dev_port``) and the # corresponding url= field on AsyncSubAgent specs. -_DEFAULT_PORT = 6174 +_DEFAULT_PORT = 3076 def _base_url(port: int = _DEFAULT_PORT) -> str: @@ -448,7 +448,7 @@ def _kill_owned_stale_process(port: int) -> bool: Why this matters: 1. ``net_connections`` may report any process bound to the port, - including user-run dev servers that legitimately took 6174. + including user-run dev servers that legitimately took 3076. SIGKILL'ing those is a data-loss event. 2. Even with PID-file ownership, the OS may have recycled the PID to an unrelated process between sessions (e.g., after a SIGKILL'd @@ -556,7 +556,7 @@ def start_langgraph_dev( Determines where deployed agents' filesystem operations land (``CustomSandboxBackend`` derives its workspace root from cwd via ``paths.WORKSPACE_ROOT``). Defaults to ``Path.cwd()``. - port: TCP port to bind. Defaults to 6174 (Kaprekar's constant). + port: TCP port to bind. Defaults to 3076. file_persistence: When True (default), langgraph dev writes its full ``.langgraph_api/`` cache so async-task / Store / scheduler state survives subprocess restarts. Set False to suppress periodic diff --git a/EvoScientist/langgraph_dev/sdk.py b/EvoScientist/langgraph_dev/sdk.py index 7eafd09..88978c7 100644 --- a/EvoScientist/langgraph_dev/sdk.py +++ b/EvoScientist/langgraph_dev/sdk.py @@ -2,14 +2,18 @@ from __future__ import annotations +import os from collections.abc import Mapping -DEFAULT_LANGGRAPH_DEV_PORT = 6174 +DEFAULT_LANGGRAPH_DEV_PORT = 3076 LANGGRAPH_DEV_AUTH_HEADERS = {"x-auth-scheme": "langsmith"} def langgraph_dev_url(config: object | None = None, *, port: int | None = None) -> str: """Return the local langgraph-dev base URL for a config or explicit port.""" + runtime_url = os.environ.get("LANGGRAPH_SERVER_URL", "").strip().rstrip("/") + if port is None and runtime_url: + return runtime_url selected_port = ( int(port) if port is not None diff --git a/EvoScientist/llm/README.md b/EvoScientist/llm/README.md new file mode 100644 index 0000000..7fb623f --- /dev/null +++ b/EvoScientist/llm/README.md @@ -0,0 +1,50 @@ +# Model Runtime Layout + +The model runtime has three configuration and execution boundaries. + +| Layer | Source | Owns | Must not own | +| --- | --- | --- | --- | +| Provider | `configuration/provider.py` | Adapter identity, credentials, endpoints, headers, connection defaults | Model capabilities, model token limits, derived tool transport | +| Model | `configuration/model.py` | Provider model ID, capabilities, limits, canonical parameters, access, billing | Credentials, base URL, SDK client options, derived tool transport | +| Invocation | `invocation/contract.py` | Immutable API mode, output parameter, tool transport, streaming flag, final SDK parameters | Admin persistence, credentials, routing decisions | + +Supporting modules have narrower responsibilities: + +- `model_config_v4.py` normalizes and persists the Provider + ModelProfile admin + contract, then projects it to the stable runtime schema. +- `model_config.py` parses and validates the runtime schema. It re-exports the + provider and model contracts for compatibility with existing integrations. +- `adapter_registry.py` declares provider/model-family support and converts + canonical model parameters into provider SDK parameters. +- `runtime.py` selects a frozen route, asks its adapter to compile parameters, + compiles an `InvocationPlan`, and constructs the provider client from that + plan only. + +The call chain is fixed: + +```text +V4 Provider + ModelProfile + -> normalize and validate + -> V3 runtime projection + -> select provider endpoint and model profile + -> merge canonical model parameters + -> provider adapter compilation + -> immutable InvocationPlan validation + -> provider SDK call +``` + +Important invariants: + +1. Environment variables may provide secrets, proxy settings, and timeouts; + they cannot select an API protocol or rewrite a compiled invocation. +2. `tool_call_transport` is not administrator configuration. It is derived as + `native` when `capabilities.tools=true`, otherwise `disabled`. +3. Exactly one provider output-limit parameter is allowed in a compiled plan: + `max_output_tokens`, `max_completion_tokens`, or `max_tokens`. +4. Provider-specific parameter names are selected by the adapter. Gateway, + frontend, and generic runtime code must not guess them from model names. +5. Runtime logs report the final non-secret plan and parameter names. They must + never include credentials, authorization headers, or raw secret values. +6. Provider input projection removes assistant history that has neither final + text nor a tool call. A newly completed empty response receives one bounded + same-route repair attempt, then fails as `MODEL_PROVIDER_RESPONSE_INVALID`. diff --git a/EvoScientist/llm/__init__.py b/EvoScientist/llm/__init__.py index 1704632..a993334 100644 --- a/EvoScientist/llm/__init__.py +++ b/EvoScientist/llm/__init__.py @@ -14,7 +14,20 @@ import lazy_loader as _lazy __getattr__, __dir__, __all__ = _lazy.attach( __name__, - submodules=["context_window", "models", "patches"], + submodules=[ + "context_window", + "models", + "patches", + "contracts", + "config_admin", + "configuration", + "crypto", + "invocation", + "model_config", + "runtime", + "adapter_registry", + "user_options", + ], submod_attrs={ "context_window": [ "DEFAULT_CONTEXT_WINDOW_FALLBACK", @@ -29,5 +42,45 @@ __getattr__, __dir__, __all__ = _lazy.attach( "get_models_for_provider", "list_models", ], + "contracts": [ + "AdmissionGrant", + "AgentExecutionProfile", + "AgentInputV3", + "AgentModelSet", + "EvoRuntimeError", + "EvoRuntimeEvent", + "HmacGrantAuthority", + "PreparedRunQuote", + "RoutePreparationGrant", + "WebHostContext", + ], + "model_config": [ + "EvoModelConfig", + "FileEvoModelConfigStore", + "SaveModelConfigCommand", + ], + "config_admin": ["EvoModelConfigAdminService"], + "configuration": [ + "EndpointConfig", + "ModelConfig", + "ProviderConfig", + "ResolvedSecret", + "SecretReference", + "SecretResolver", + ], + "invocation": [ + "InvocationPlan", + "compile_invocation_plan", + "derive_runtime_invocation", + "derive_tool_call_transport", + ], + "runtime": ["EvoModelRuntime"], + "adapter_registry": ["AdapterRegistry", "get_adapter_registry"], + "user_options": [ + "model_options_schema_hash", + "project_user_options_for_purpose", + "validate_parameter_constraints", + "validate_user_model_options", + ], }, ) diff --git a/EvoScientist/llm/adapter_registry.py b/EvoScientist/llm/adapter_registry.py new file mode 100644 index 0000000..895d58b --- /dev/null +++ b/EvoScientist/llm/adapter_registry.py @@ -0,0 +1,1379 @@ +"""Versioned provider adapter registry used by model configuration V3. + +The registry is the only place where provider protocol differences are +described. Configuration may select an exact registered revision, but cannot +invent authentication headers, request paths, or provider-specific fields. +""" + +from __future__ import annotations + +import asyncio +import base64 +import hashlib +import json +from collections.abc import Mapping +from dataclasses import asdict, dataclass, field +from typing import Any, Literal +from urllib.parse import quote + +from .contracts import EvoRuntimeError + +AdapterLifecycle = Literal["active", "deprecated", "blocked"] +_PROBE_RETRY_BASE_SECONDS = 0.25 +_PROBE_RETRY_MAX_SECONDS = 5.0 +_KIMI_CHAT_COMPLETION_TOKEN_MODELS = frozenset( + { + "k3", + "k3-256k", + "kimi-k3", + "kimi-for-coding", + "kimi-for-coding-highspeed", + } +) + + +@dataclass(frozen=True, slots=True) +class ErrorDisposition: + error_code: str + retryable: bool + health_scope: Literal["none", "provider_connection", "model_route"] + retry_after_ms: int | None = None + + +@dataclass(frozen=True, slots=True) +class NormalizedUsage: + input_tokens: int | None + output_tokens: int | None + cached_input_tokens: int | None + reasoning_tokens: int | None = None + total_tokens: int | None = None + provider_request_id_hash: str | None = None + finality: Literal["confirmed", "partial", "unconfirmed"] = "unconfirmed" + + def confirmed_projection(self) -> Mapping[str, int | str] | None: + if self.finality != "confirmed" or any( + value is None + for value in ( + self.input_tokens, + self.output_tokens, + self.cached_input_tokens, + ) + ): + return None + assert self.input_tokens is not None + assert self.output_tokens is not None + assert self.cached_input_tokens is not None + if ( + min(self.input_tokens, self.output_tokens, self.cached_input_tokens) < 0 + or self.cached_input_tokens > self.input_tokens + or ( + self.reasoning_tokens is not None + and ( + self.reasoning_tokens < 0 + or self.reasoning_tokens > self.output_tokens + ) + ) + ): + return None + result: dict[str, int | str] = { + "input_tokens": self.input_tokens, + "cached_input_tokens": self.cached_input_tokens, + "output_tokens": self.output_tokens, + } + if self.reasoning_tokens is not None: + result["reasoning_tokens"] = self.reasoning_tokens + if self.total_tokens is not None: + result["total_tokens"] = self.total_tokens + if self.provider_request_id_hash: + result["provider_request_id_hash"] = self.provider_request_id_hash + return result + + +@dataclass(frozen=True, slots=True) +class ParameterRule: + kind: Literal["boolean", "integer", "number", "string", "enum", "object"] + minimum: float | None = None + maximum: float | None = None + minimum_exclusive: bool = False + maximum_exclusive: bool = False + choices: tuple[str, ...] = () + + def validate(self, value: Any, path: str) -> None: + if self.kind == "boolean": + valid = isinstance(value, bool) + elif self.kind == "integer": + valid = isinstance(value, int) and not isinstance(value, bool) + elif self.kind == "number": + valid = isinstance(value, int | float) and not isinstance(value, bool) + elif self.kind in {"string", "enum"}: + valid = isinstance(value, str) + else: + valid = isinstance(value, Mapping) + if not valid: + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", + f"{path} has an invalid type", + details=({"path": path, "code": "PARAMETER_TYPE_INVALID"},), + ) + if self.kind == "enum" and value not in self.choices: + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", + f"{path} is not supported", + details=({"path": path, "code": "PARAMETER_VALUE_UNSUPPORTED"},), + ) + if self.kind in {"integer", "number"}: + number = float(value) + if self.minimum is not None and ( + number < self.minimum + or (self.minimum_exclusive and number == self.minimum) + ): + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", + f"{path} is below its minimum", + details=({"path": path, "code": "PARAMETER_OUT_OF_RANGE"},), + ) + if self.maximum is not None and ( + number > self.maximum + or (self.maximum_exclusive and number == self.maximum) + ): + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", + f"{path} exceeds its maximum", + details=({"path": path, "code": "PARAMETER_OUT_OF_RANGE"},), + ) + + +@dataclass(frozen=True, slots=True) +class ModelDescriptor: + context_tokens: int + max_output_tokens: int + capabilities: frozenset[str] + parameters: Mapping[str, ParameterRule] + + +@dataclass(frozen=True, slots=True) +class ModelDiscoveryDescriptor: + context_tokens: int | None + max_output_tokens: int | None + capabilities: frozenset[str] + reasoning_mode: Literal["none", "boolean", "effort"] = "none" + reasoning_efforts: tuple[str, ...] = () + default_reasoning_effort: str | None = None + source: str = "adapter_catalog" + + +@dataclass(frozen=True, slots=True) +class AdapterRegistration: + adapter_id: str + adapter_revision: str + display_name: str + lifecycle: AdapterLifecycle + supported_wire_protocols: tuple[str, ...] + supported_api_modes: tuple[str, ...] + recommended_api_mode: str + recommended_base_url: str + auth_profile: Literal["bearer", "anthropic_api_key", "google_api_key"] + runtime_provider: str + protocol_hard_limits: Mapping[str, tuple[int, int]] + parameter_schema: Mapping[str, ParameterRule] + legacy_parameter_schema: Mapping[str, ParameterRule] = field(default_factory=dict) + reasoning_mode: Literal["none", "boolean", "effort"] = "none" + exact_model_overrides: Mapping[str, ModelDescriptor] = field(default_factory=dict) + discovery_model_overrides: Mapping[str, ModelDiscoveryDescriptor] = field( + default_factory=dict + ) + discovery_capability: bool = False + server_tools_configurable: bool = False + replacement_revision: str | None = None + implementation_fingerprint: str = "" + + def metadata(self) -> Mapping[str, Any]: + return { + "adapter_id": self.adapter_id, + "adapter_revision": self.adapter_revision, + "implementation_fingerprint": self.implementation_fingerprint, + "lifecycle": self.lifecycle, + "replacement_revision": self.replacement_revision, + "display_name": self.display_name, + "wire_protocols": list(self.supported_wire_protocols), + "recommended_base_url": self.recommended_base_url, + "credential_kind": "api_key", + "api_modes": [ + {"id": value, "recommended": value == self.recommended_api_mode} + for value in self.supported_api_modes + ], + "server_tools_configurable": self.server_tools_configurable, + "reasoning_mode": self.reasoning_mode, + "parameter_schema": { + key: asdict(rule) for key, rule in sorted(self.parameter_schema.items()) + }, + } + + @property + def all_parameter_schema(self) -> Mapping[str, ParameterRule]: + return {**self.legacy_parameter_schema, **self.parameter_schema} + + def resolve_model_descriptor( + self, + provider_model_id: str, + api_mode: str, + *, + context_tokens: int | None, + max_output_tokens: int | None, + declared_capabilities: Mapping[str, bool], + ) -> ModelDescriptor: + if self.lifecycle == "blocked": + raise EvoRuntimeError("MODEL_ADAPTER_BLOCKED") + if api_mode not in self.supported_api_modes: + raise EvoRuntimeError("MODEL_API_MODE_UNSUPPORTED") + hard_context, hard_output = self.protocol_hard_limits[api_mode] + exact = self.exact_model_overrides.get(provider_model_id) + descriptor_context = exact.context_tokens if exact else hard_context + descriptor_output = exact.max_output_tokens if exact else hard_output + if context_tokens is None or max_output_tokens is None: + if exact is None: + raise EvoRuntimeError("TOKEN_BOUND_UNAVAILABLE") + context_tokens = context_tokens or descriptor_context + max_output_tokens = max_output_tokens or descriptor_output + if context_tokens > descriptor_context or max_output_tokens > descriptor_output: + raise EvoRuntimeError("MODEL_TOKEN_BOUND_EXCEEDED") + if max_output_tokens > context_tokens: + raise EvoRuntimeError("MODEL_TOKEN_BOUND_EXCEEDED") + # Exact descriptors are an advisory catalog: they supply known limits, + # defaults and UI hints, but do not overrule an administrator's product + # capability policy. A model catalog necessarily lags new model IDs and + # Provider releases; protocol validation happens when a real request is + # compiled instead. + supported = exact.capabilities if exact else frozenset({"text"}) + return ModelDescriptor( + context_tokens=context_tokens, + max_output_tokens=max_output_tokens, + capabilities=supported, + parameters={**self.parameter_schema, **(exact.parameters if exact else {})}, + ) + + def resolve_discovery_descriptor( + self, provider_model_id: str + ) -> ModelDiscoveryDescriptor | None: + exact = self.exact_model_overrides.get(provider_model_id) + if exact is not None: + supports_reasoning = ( + "thinking" in exact.capabilities and self.reasoning_mode != "none" + ) + return ModelDiscoveryDescriptor( + context_tokens=exact.context_tokens, + max_output_tokens=exact.max_output_tokens, + capabilities=exact.capabilities, + reasoning_mode=self.reasoning_mode if supports_reasoning else "none", + reasoning_efforts=("low", "medium", "high") + if supports_reasoning + else (), + default_reasoning_effort="medium" if supports_reasoning else None, + source="adapter_exact_model", + ) + return self.discovery_model_overrides.get(provider_model_id) + + def validate_parameters(self, values: Mapping[str, Any], *, path: str) -> None: + schema = self.all_parameter_schema + unknown = set(values) - set(schema) + if unknown: + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", + f"{path} contains unsupported parameters: {', '.join(sorted(unknown))}", + details=tuple( + {"path": f"{path}.{name}", "code": "PARAMETER_UNSUPPORTED"} + for name in sorted(unknown) + ), + ) + for name, value in values.items(): + schema[name].validate(value, f"{path}.{name}") + if "temperature" in values and "top_p" in values: + raise EvoRuntimeError( + "MODEL_PARAMETER_CONFLICT", + f"{path} sets temperature and top_p", + details=( + {"path": f"{path}.temperature", "code": "PARAMETER_CONFLICT"}, + {"path": f"{path}.top_p", "code": "PARAMETER_CONFLICT"}, + ), + ) + reasoning = values.get("reasoning") + if reasoning is None and "reasoning_effort" in values: + reasoning = ( + "off" + if values.get("reasoning_effort") == "disabled" + else values.get("reasoning_effort") + ) + if reasoning is None and "thinking_enabled" in values: + reasoning = "high" if values.get("thinking_enabled") else "off" + if self.adapter_id == "dashscope": + forced_nonthinking = bool(values.get("structured_output")) or ( + values.get("tool_choice") == "required" + ) + if forced_nonthinking and reasoning not in {None, "off"}: + raise EvoRuntimeError( + "MODEL_PARAMETER_CONFLICT", + f"{path} combines reasoning with a Provider non-thinking mode", + details=( + { + "path": f"{path}.reasoning", + "code": "PARAMETER_CONFLICT", + }, + ), + ) + if "reasoning_budget_tokens" in values and reasoning in {None, "off"}: + raise EvoRuntimeError( + "MODEL_PARAMETER_CONFLICT", + f"{path}.reasoning_budget_tokens requires reasoning", + details=( + { + "path": f"{path}.reasoning_budget_tokens", + "code": "PARAMETER_CONFLICT", + }, + ), + ) + + def compile_runtime_parameters( + self, + api_mode: str, + values: Mapping[str, Any], + output_token_limit: int, + *, + provider_model_id: str | None = None, + ) -> Mapping[str, Any]: + """Map canonical V3 parameters into installed LangChain constructors.""" + + result = dict(values) + result.pop("output_token_limit", None) + legacy_thinking = result.pop("thinking_enabled", None) + structured = result.pop("structured_output", None) + legacy_effort = result.pop("reasoning_effort", None) + reasoning = result.pop("reasoning", None) + reasoning_budget = result.pop("reasoning_budget_tokens", None) + if reasoning is None and legacy_effort is not None: + reasoning = "off" if legacy_effort == "disabled" else legacy_effort + if reasoning is None and legacy_thinking is not None: + reasoning = "high" if legacy_thinking else "off" + thinking = None if reasoning is None else reasoning != "off" + if self.adapter_id == "dashscope" and api_mode == "chat_completions": + forced_nonthinking = bool(structured) or result.get("tool_choice") == "required" + if forced_nonthinking and thinking: + raise EvoRuntimeError( + "MODEL_PARAMETER_CONFLICT", + "DashScope structured output and forced tools require thinking to be disabled", + ) + if forced_nonthinking: + thinking = False + extra_body = {"enable_thinking": thinking} if thinking is not None else {} + if thinking and reasoning_budget is not None: + if int(reasoning_budget) > output_token_limit: + raise EvoRuntimeError( + "MODEL_PARAMETER_CONFLICT", + "reasoning budget exceeds the effective output limit", + ) + extra_body["thinking_budget"] = int(reasoning_budget) + result.update({"max_completion_tokens": output_token_limit}) + if extra_body: + result["extra_body"] = extra_body + elif self.adapter_id == "anthropic": + result["max_tokens"] = output_token_limit + if thinking: + if output_token_limit <= 1_024: + raise EvoRuntimeError( + "MODEL_PARAMETER_CONFLICT", + "Anthropic thinking requires max_tokens greater than 1024", + ) + result["thinking"] = { + "type": "enabled", + "budget_tokens": int(reasoning_budget) + if reasoning_budget is not None + else min(8_192, max(1_024, output_token_limit // 2)), + } + elif api_mode == "responses": + result.update( + { + "max_output_tokens": output_token_limit, + "use_responses_api": True, + "store": False, + } + ) + if reasoning not in {None, "off"}: + result["reasoning"] = {"effort": reasoning} + elif self.adapter_id == "dashscope" and thinking: + result["reasoning"] = {"effort": "medium"} + elif self.adapter_id == "google-gemini": + result["max_output_tokens"] = output_token_limit + if api_mode == "interactions": + result["store"] = False + if thinking is not None: + result["thinking"] = thinking + else: + if self._uses_max_completion_tokens(provider_model_id, api_mode): + result["max_completion_tokens"] = output_token_limit + else: + result["max_tokens"] = output_token_limit + # V3 routes declare the API envelope explicitly. Do not leave + # LangChain to infer the Responses API from model naming or kwargs. + if self.runtime_provider == "openai": + result["use_responses_api"] = False + if reasoning not in {None, "off"}: + result["reasoning_effort"] = reasoning + if structured: + result["response_format"] = {"type": "json_object"} + return result + + def _uses_max_completion_tokens( + self, provider_model_id: str | None, api_mode: str + ) -> bool: + """Return whether an OpenAI Chat model rejects legacy ``max_tokens``.""" + + model_id = str(provider_model_id or "").strip().lower() + normalized = model_id.replace("-", "") + return ( + self.adapter_id == "openai" + and api_mode == "chat_completions" + and ( + normalized.startswith("gpt5") + or model_id in _KIMI_CHAT_COMPLETION_TOKEN_MODELS + ) + ) + + def classify_error(self, error: BaseException) -> ErrorDisposition: + """Classify typed SDK/HTTP failures without inspecting user-facing text.""" + + try: + import httpx + except ImportError: # pragma: no cover - httpx is a runtime dependency + httpx = None # type: ignore[assignment] + + status = getattr(error, "status_code", None) + if status is None: + response = getattr(error, "response", None) + status = getattr(response, "status_code", None) + retry_after_ms = None + headers = getattr(getattr(error, "response", None), "headers", None) + if headers: + try: + retry_after_ms = ( + int(float(headers.get("retry-after", 0)) * 1000) or None + ) + except (TypeError, ValueError): + retry_after_ms = None + if status in {401, 403}: + return ErrorDisposition( + "MODEL_AUTHENTICATION_FAILED", False, "provider_connection" + ) + if status == 404: + return ErrorDisposition("MODEL_NOT_FOUND", False, "model_route") + if status in {400, 422}: + return ErrorDisposition( + "MODEL_PROVIDER_REQUEST_REJECTED", False, "model_route" + ) + if status == 429: + return ErrorDisposition( + "MODEL_RATE_LIMITED", True, "provider_connection", retry_after_ms + ) + if isinstance(status, int) and status >= 500: + return ErrorDisposition( + "MODEL_PROVIDER_ERROR", True, "provider_connection", retry_after_ms + ) + if isinstance(error, (TimeoutError, ConnectionError)) or ( + httpx is not None and isinstance(error, httpx.TimeoutException) + ): + return ErrorDisposition("MODEL_TIMEOUT", True, "provider_connection") + if httpx is not None and isinstance(error, httpx.TransportError): + return ErrorDisposition("MODEL_PROVIDER_ERROR", True, "provider_connection") + if isinstance(error, json.JSONDecodeError): + return ErrorDisposition( + "MODEL_PROVIDER_RESPONSE_INVALID", True, "provider_connection" + ) + if isinstance(error, EvoRuntimeError): + return ErrorDisposition(error.code, False, "none") + return ErrorDisposition("MODEL_PROVIDER_ERROR", False, "model_route") + + async def discover_models( + self, + *, + base_url: str, + api_key: str, + timeout_seconds: int = 10, + ) -> tuple[Mapping[str, str], ...]: + if not self.discovery_capability: + raise EvoRuntimeError("MODEL_DISCOVERY_UNSUPPORTED") + import httpx + + root = base_url.rstrip("/") + if self.adapter_id == "anthropic": + url = root + "/v1/models" + headers = { + "x-api-key": api_key, + "anthropic-version": "2023-06-01", + } + elif self.adapter_id == "google-gemini": + url = root + "/v1beta/models" + headers = {"x-goog-api-key": api_key} + else: + url = root + "/models" + headers = {"Authorization": f"Bearer {api_key}"} + try: + async with httpx.AsyncClient( + timeout=httpx.Timeout(timeout_seconds), + follow_redirects=False, + trust_env=False, + ) as client: + response = await client.get(url, headers=headers) + response.raise_for_status() + payload = response.json() + except Exception as exc: + disposition = self.classify_error(exc) + raise EvoRuntimeError(disposition.error_code) from exc + items = payload.get("data") if isinstance(payload, Mapping) else None + if items is None and isinstance(payload, Mapping): + items = payload.get("models") + result: list[Mapping[str, str]] = [] + for item in items or []: + if not isinstance(item, Mapping): + continue + model_id = str(item.get("id") or item.get("name") or "").strip() + model_id = model_id.removeprefix("models/") + if not model_id or len(model_id) > 512: + continue + result.append( + { + "provider_model_id": model_id, + "display_name": str( + item.get("display_name") or item.get("displayName") or model_id + )[:512], + } + ) + if len(result) >= 500: + break + return tuple(result) + + def build_probe_request( + self, + *, + base_url: str, + api_key: str, + provider_model_id: str, + api_mode: str, + probe_kind: str, + ) -> tuple[str, Mapping[str, str], Mapping[str, Any]]: + """Compile a minimal real Provider request for one capability probe.""" + + if probe_kind in {"video", "documents"}: + raise EvoRuntimeError("MODEL_CAPABILITY_PROBE_UNSUPPORTED") + root = base_url.rstrip("/") + prompt = ( + 'Return exactly {"ok":true}.' + if probe_kind == "structured_output" + else "Call the probe_ok tool now." + if probe_kind == "tools" + else "Reply with the single word OK." + ) + png = ( + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR4nGNg" + "YAAAAAMAASsJTYQAAAAASUVORK5CYII=" + ) + if self.auth_profile == "anthropic_api_key": + headers = { + "x-api-key": api_key, + "anthropic-version": "2023-06-01", + "content-type": "application/json", + } + elif self.auth_profile == "google_api_key": + headers = {"x-goog-api-key": api_key, "content-type": "application/json"} + else: + headers = { + "Authorization": f"Bearer {api_key}", + "content-type": "application/json", + } + tool = { + "name": "probe_ok", + "description": "Return a probe acknowledgement.", + "parameters": {"type": "object", "properties": {}}, + } + if self.adapter_id == "anthropic": + content: Any = prompt + if probe_kind == "vision": + content = [ + { + "type": "image", + "source": { + "type": "base64", + "media_type": "image/png", + "data": png, + }, + }, + {"type": "text", "text": prompt}, + ] + body: dict[str, Any] = { + "model": provider_model_id, + "messages": [{"role": "user", "content": content}], + "max_tokens": 16, + } + if probe_kind == "tools": + body["tools"] = [{**tool, "input_schema": tool["parameters"]}] + body["tools"][0].pop("parameters") + body["tool_choice"] = {"type": "tool", "name": "probe_ok"} + elif probe_kind == "structured_output": + body["output_config"] = { + "format": { + "type": "json_schema", + "schema": { + "type": "object", + "properties": {"ok": {"type": "boolean"}}, + "required": ["ok"], + "additionalProperties": False, + }, + } + } + elif probe_kind == "reasoning": + body["max_tokens"] = 1_025 + body["thinking"] = {"type": "enabled", "budget_tokens": 1_024} + return root + "/v1/messages", headers, body + if self.adapter_id == "google-gemini": + if api_mode == "interactions": + interaction_content: list[dict[str, Any]] = [ + {"type": "text", "text": prompt} + ] + if probe_kind == "vision": + interaction_content.insert( + 0, + { + "type": "image", + "mime_type": "image/png", + "data": png, + }, + ) + interaction_body: dict[str, Any] = { + "model": provider_model_id, + "input": [{"role": "user", "content": interaction_content}], + "generation_config": {"max_output_tokens": 16}, + "store": False, + "stream": False, + } + if probe_kind == "tools": + interaction_body["tools"] = [{"type": "function", **tool}] + interaction_body["tool_choice"] = { + "type": "function", + "name": "probe_ok", + } + elif probe_kind == "structured_output": + interaction_body["generation_config"].update( + { + "response_mime_type": "application/json", + "response_schema": { + "type": "object", + "properties": {"ok": {"type": "boolean"}}, + }, + } + ) + elif probe_kind == "reasoning": + interaction_body["generation_config"].update( + {"thinking_level": "high", "thinking_summaries": "auto"} + ) + return root + "/v1beta/interactions", headers, interaction_body + parts: list[dict[str, Any]] = [{"text": prompt}] + if probe_kind == "vision": + parts.insert( + 0, + {"inline_data": {"mime_type": "image/png", "data": png}}, + ) + body = { + "contents": [{"role": "user", "parts": parts}], + "generationConfig": {"maxOutputTokens": 16}, + } + if probe_kind == "tools": + body["tools"] = [{"functionDeclarations": [tool]}] + body["toolConfig"] = { + "functionCallingConfig": { + "mode": "ANY", + "allowedFunctionNames": ["probe_ok"], + } + } + elif probe_kind == "structured_output": + body["generationConfig"].update( + { + "responseMimeType": "application/json", + "responseSchema": { + "type": "OBJECT", + "properties": {"ok": {"type": "BOOLEAN"}}, + }, + } + ) + elif probe_kind == "reasoning": + body["generationConfig"]["thinkingConfig"] = { + "thinkingBudget": 1_024 + } + return ( + root + + "/v1beta/models/" + + quote(provider_model_id, safe="") + + ":generateContent", + headers, + body, + ) + if api_mode == "responses": + input_value: Any = prompt + if probe_kind == "vision": + input_value = [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": prompt}, + { + "type": "input_image", + "image_url": f"data:image/png;base64,{png}", + }, + ], + } + ] + body = { + "model": provider_model_id, + "input": input_value, + "max_output_tokens": 16, + "store": False, + } + if probe_kind == "tools": + body["tools"] = [{"type": "function", **tool}] + body["tool_choice"] = {"type": "function", "name": "probe_ok"} + elif probe_kind == "structured_output": + body["text"] = { + "format": { + "type": "json_schema", + "name": "probe", + "schema": { + "type": "object", + "properties": {"ok": {"type": "boolean"}}, + "required": ["ok"], + "additionalProperties": False, + }, + "strict": True, + } + } + elif probe_kind == "reasoning": + body["reasoning"] = {"effort": "low"} + return root + "/responses", headers, body + content = prompt + if probe_kind == "vision": + content = [ + {"type": "text", "text": prompt}, + { + "type": "image_url", + "image_url": {"url": f"data:image/png;base64,{png}"}, + }, + ] + body = { + "model": provider_model_id, + "messages": [{"role": "user", "content": content}], + } + if self.adapter_id == "dashscope" or self._uses_max_completion_tokens( + provider_model_id, api_mode + ): + body["max_completion_tokens"] = 16 + if self.adapter_id == "dashscope": + body["enable_thinking"] = False + else: + body["max_tokens"] = 16 + if probe_kind == "tools": + body["tools"] = [{"type": "function", "function": tool}] + body["tool_choice"] = { + "type": "function", + "function": {"name": "probe_ok"}, + } + elif probe_kind == "structured_output": + body["response_format"] = {"type": "json_object"} + elif probe_kind == "reasoning": + if self.adapter_id == "dashscope": + body["enable_thinking"] = True + body["thinking_budget"] = 16 + body["stream"] = True + else: + body["reasoning_effort"] = "low" + return root + "/chat/completions", headers, body + + async def probe_model( + self, + *, + base_url: str, + api_key: str, + provider_model_id: str, + api_mode: str, + probe_kinds: tuple[str, ...], + timeout_seconds: int, + max_attempts: int = 2, + ) -> Mapping[str, str]: + import httpx + + attempts_limit = max(1, min(int(max_attempts), 3)) + results: dict[str, str] = {} + async with httpx.AsyncClient( + timeout=httpx.Timeout(timeout_seconds), + follow_redirects=False, + trust_env=False, + ) as client: + for probe_kind in probe_kinds: + for attempt in range(1, attempts_limit + 1): + try: + url, headers, body = self.build_probe_request( + base_url=base_url, + api_key=api_key, + provider_model_id=provider_model_id, + api_mode=api_mode, + probe_kind=probe_kind, + ) + async with client.stream( + "POST", url, headers=dict(headers), json=dict(body) + ) as response: + response.raise_for_status() + chunks: list[bytes] = [] + total = 0 + async for chunk in response.aiter_bytes(): + total += len(chunk) + if total > 1_048_576: + raise EvoRuntimeError( + "MODEL_PROVIDER_RESPONSE_INVALID" + ) + chunks.append(chunk) + payloads = self._decode_probe_payloads(b"".join(chunks)) + observed_revision = self._validate_probe_payloads( + probe_kind, + payloads, + ) + except Exception as exc: + disposition = self.classify_error(exc) + if disposition.retryable and attempt < attempts_limit: + server_delay = ( + disposition.retry_after_ms / 1_000 + if disposition.retry_after_ms is not None + else _PROBE_RETRY_BASE_SECONDS * (2 ** (attempt - 1)) + ) + await asyncio.sleep( + min( + _PROBE_RETRY_MAX_SECONDS, + max(0.0, server_delay), + ) + ) + continue + raise EvoRuntimeError( + disposition.error_code, + details=( + { + "path": f"probe.{probe_kind}", + "code": disposition.error_code, + "probe_kind": probe_kind, + "attempts": attempt, + "retryable": disposition.retryable, + }, + ), + ) from exc + results[probe_kind] = "supported" + if observed_revision: + results["resolved_model_revision"] = observed_revision + break + return results + + async def test_model_connection( + self, + *, + base_url: str, + api_key: str, + provider_model_id: str, + api_mode: str, + timeout_seconds: int, + ) -> Mapping[str, str]: + """Make one minimal text request for an administrator connection test. + + This intentionally does not infer or certify tools, reasoning, media, + or structured output. Those behaviours are exercised by real requests + and Adapter integration tests, never used as a model availability gate. + """ + + return await self.probe_model( + base_url=base_url, + api_key=api_key, + provider_model_id=provider_model_id, + api_mode=api_mode, + probe_kinds=("connectivity",), + timeout_seconds=timeout_seconds, + ) + + @staticmethod + def _decode_probe_payloads(raw: bytes) -> tuple[Mapping[str, Any], ...]: + if not raw or len(raw) > 1_048_576: + raise EvoRuntimeError("MODEL_PROVIDER_RESPONSE_INVALID") + try: + text = raw.decode("utf-8") + except UnicodeDecodeError as exc: + raise EvoRuntimeError("MODEL_PROVIDER_RESPONSE_INVALID") from exc + payloads: list[Mapping[str, Any]] = [] + stripped = text.strip() + try: + decoded = json.loads(stripped) + if isinstance(decoded, Mapping): + payloads.append(decoded) + except json.JSONDecodeError: + for line in text.splitlines(): + if not line.startswith("data:"): + continue + data = line[5:].strip() + if not data or data == "[DONE]": + continue + try: + decoded = json.loads(data) + except json.JSONDecodeError as exc: + raise EvoRuntimeError("MODEL_PROVIDER_RESPONSE_INVALID") from exc + if isinstance(decoded, Mapping): + payloads.append(decoded) + if not payloads: + raise EvoRuntimeError("MODEL_PROVIDER_RESPONSE_INVALID") + return tuple(payloads) + + @classmethod + def _validate_probe_payloads( + cls, + probe_kind: str, + payloads: tuple[Mapping[str, Any], ...], + ) -> str: + if any(cls._contains_probe_key(payload, "error") for payload in payloads): + raise EvoRuntimeError("MODEL_PROVIDER_RESPONSE_INVALID") + recognized = any( + any( + key in payload + for key in ( + "choices", + "content", + "candidates", + "output", + "outputs", + "response", + "id", + ) + ) + for payload in payloads + ) + if not recognized: + raise EvoRuntimeError("MODEL_PROVIDER_RESPONSE_INVALID") + if probe_kind == "tools" and not any( + cls._contains_probe_key(payload, key) + for payload in payloads + for key in ("tool_calls", "function_call", "functionCall", "tool_use") + ) and not any( + cls._contains_probe_value(payload, value) + for payload in payloads + for value in ("function_call", "tool_use") + ): + raise EvoRuntimeError("MODEL_CAPABILITY_PROBE_FAILED") + if probe_kind == "structured_output": + structured = False + for payload in payloads: + for value in cls._probe_text_values(payload): + candidate = ( + value.strip() + .removeprefix("```json") + .removesuffix("```") + .strip() + ) + try: + parsed = json.loads(candidate) + except json.JSONDecodeError: + continue + if isinstance(parsed, Mapping) and isinstance(parsed.get("ok"), bool): + structured = True + break + if not structured: + raise EvoRuntimeError("MODEL_CAPABILITY_PROBE_FAILED") + if probe_kind == "reasoning" and not any( + cls._contains_probe_key(payload, key) + for payload in payloads + for key in ( + "reasoning", + "reasoning_content", + "reasoning_tokens", + "thinking", + "thought", + "thoughts_token_count", + ) + ) and not any( + cls._contains_probe_value(payload, value) + for payload in payloads + for value in ("reasoning", "thinking") + ): + raise EvoRuntimeError("MODEL_CAPABILITY_PROBE_FAILED") + for payload in reversed(payloads): + for key in ( + "resolved_model_revision", + "modelVersion", + "model_version", + "model", + ): + value = payload.get(key) + if isinstance(value, str) and value.strip(): + return value.strip()[:512] + return "" + + @classmethod + def _contains_probe_key(cls, value: Any, name: str) -> bool: + if isinstance(value, Mapping): + if name in value and value[name] not in (None, "", [], {}): + return True + return any(cls._contains_probe_key(item, name) for item in value.values()) + if isinstance(value, list | tuple): + return any(cls._contains_probe_key(item, name) for item in value) + return False + + @classmethod + def _contains_probe_value(cls, value: Any, expected: str) -> bool: + if isinstance(value, Mapping): + return any( + (isinstance(item, str) and item == expected) + or cls._contains_probe_value(item, expected) + for item in value.values() + ) + if isinstance(value, list | tuple): + return any(cls._contains_probe_value(item, expected) for item in value) + return False + + @classmethod + def _probe_text_values(cls, value: Any) -> tuple[str, ...]: + result: list[str] = [] + if isinstance(value, Mapping): + for key, item in value.items(): + if key in {"text", "content", "output_text"} and isinstance(item, str): + result.append(item) + else: + result.extend(cls._probe_text_values(item)) + elif isinstance(value, list | tuple): + for item in value: + result.extend(cls._probe_text_values(item)) + return tuple(result) + + +class AdapterRegistry: + def __init__( + self, + registrations: tuple[AdapterRegistration, ...], + *, + policy_key_id: str, + manifest_public_key: str, + manifest_signature: str, + ) -> None: + self._registrations = { + (item.adapter_id, item.adapter_revision): item for item in registrations + } + if len(self._registrations) != len(registrations): + raise RuntimeError("duplicate adapter registration") + projection = [item.metadata() for item in registrations] + manifest = json.dumps( + projection, sort_keys=True, separators=(",", ":") + ).encode() + try: + from cryptography.hazmat.primitives.asymmetric.ed25519 import ( + Ed25519PublicKey, + ) + + if not manifest_signature: + raise ValueError("missing registry signature") + Ed25519PublicKey.from_public_bytes( + base64.b64decode(manifest_public_key, validate=True) + ).verify(base64.b64decode(manifest_signature, validate=True), manifest) + except Exception as exc: + raise RuntimeError( + "adapter registry manifest signature is invalid" + ) from exc + digest = hashlib.sha256(manifest).hexdigest() + self.registry_revision = f"sha256:{digest}" + self.policy_key_id = policy_key_id + self.manifest_signature = manifest_signature + + def get(self, adapter_id: str, adapter_revision: str) -> AdapterRegistration: + item = self._registrations.get((adapter_id, adapter_revision)) + if item is None: + raise EvoRuntimeError("MODEL_ADAPTER_UNAVAILABLE") + return item + + def metadata(self) -> Mapping[str, Any]: + return { + "registry_revision": self.registry_revision, + "registry_policy_key_id": self.policy_key_id, + "manifest_signature": self.manifest_signature, + "adapters": [ + item.metadata() + for item in sorted( + self._registrations.values(), + key=lambda value: (value.adapter_id, value.adapter_revision), + ) + ], + } + + +_LEGACY_REASONING = { + "reasoning_effort": ParameterRule( + "enum", choices=("disabled", "low", "medium", "high", "max") + ), + "thinking_enabled": ParameterRule("boolean"), +} + +_ANTHROPIC_PARAMETERS = { + "output_token_limit": ParameterRule("integer", minimum=1), + "temperature": ParameterRule( + "number", minimum=0, maximum=1, maximum_exclusive=False + ), + "top_p": ParameterRule("number", minimum=0, maximum=1, minimum_exclusive=True), + "top_k": ParameterRule("integer", minimum=0), + "reasoning": ParameterRule( + "enum", choices=("off", "low", "medium", "high") + ), + "reasoning_budget_tokens": ParameterRule("integer", minimum=1024), + "structured_output": ParameterRule("boolean"), + "tool_choice": ParameterRule("enum", choices=("auto", "none", "required")), +} + +_OPENAI_PARAMETERS = { + "output_token_limit": ParameterRule("integer", minimum=1), + "temperature": ParameterRule( + "number", minimum=0, maximum=2, maximum_exclusive=True + ), + "top_p": ParameterRule("number", minimum=0, maximum=1, minimum_exclusive=True), + "reasoning": ParameterRule( + "enum", choices=("off", "low", "medium", "high", "max") + ), + "structured_output": ParameterRule("boolean"), + "tool_choice": ParameterRule("enum", choices=("auto", "none", "required")), +} + +_GEMINI_PARAMETERS = { + "output_token_limit": ParameterRule("integer", minimum=1), + "temperature": ParameterRule( + "number", minimum=0, maximum=2, maximum_exclusive=True + ), + "top_p": ParameterRule("number", minimum=0, maximum=1, minimum_exclusive=True), + "top_k": ParameterRule("integer", minimum=1), + "reasoning": ParameterRule( + "enum", choices=("off", "low", "medium", "high") + ), + "structured_output": ParameterRule("boolean"), + "tool_choice": ParameterRule("enum", choices=("auto", "none", "required")), +} + +_XAI_PARAMETERS = { + "output_token_limit": ParameterRule("integer", minimum=1), + "temperature": ParameterRule( + "number", minimum=0, maximum=2, maximum_exclusive=True + ), + "top_p": ParameterRule("number", minimum=0, maximum=1, minimum_exclusive=True), + "reasoning": ParameterRule( + "enum", choices=("off", "low", "medium", "high") + ), + "structured_output": ParameterRule("boolean"), + "tool_choice": ParameterRule("enum", choices=("auto", "none", "required")), +} + +_DASHSCOPE_PARAMETERS = { + "output_token_limit": ParameterRule("integer", minimum=1, maximum=65_536), + "temperature": ParameterRule( + "number", minimum=0, maximum=2, maximum_exclusive=True + ), + "top_p": ParameterRule("number", minimum=0, maximum=1, minimum_exclusive=True), + "reasoning": ParameterRule( + "enum", choices=("off", "low", "medium", "high") + ), + "reasoning_budget_tokens": ParameterRule( + "integer", minimum=1, maximum=262_144 + ), + "structured_output": ParameterRule("boolean"), + "tool_choice": ParameterRule("enum", choices=("auto", "none", "required")), +} + +_GENERIC_OPENAI_PARAMETERS = { + "output_token_limit": ParameterRule("integer", minimum=1), + "temperature": ParameterRule( + "number", minimum=0, maximum=2, maximum_exclusive=True + ), + "top_p": ParameterRule("number", minimum=0, maximum=1, minimum_exclusive=True), + "structured_output": ParameterRule("boolean"), + "tool_choice": ParameterRule("enum", choices=("auto", "none", "required")), +} + + +def _fingerprint(adapter_id: str, revision: str) -> str: + return ( + "sha256:" + + hashlib.sha256( + f"{adapter_id}:{revision}:evoscientist-adapter-v1".encode() + ).hexdigest() + ) + + +def _registration( + adapter_id: str, + revision: str, + display_name: str, + wires: tuple[str, ...], + modes: tuple[str, ...], + recommended_mode: str, + base_url: str, + auth: Literal["bearer", "anthropic_api_key", "google_api_key"], + runtime_provider: str, + parameter_schema: Mapping[str, ParameterRule], + *, + reasoning_mode: Literal["none", "boolean", "effort"] = "none", + discovery: bool = True, + exact: Mapping[str, ModelDescriptor] | None = None, + discovery_models: Mapping[str, ModelDiscoveryDescriptor] | None = None, +) -> AdapterRegistration: + return AdapterRegistration( + adapter_id=adapter_id, + adapter_revision=revision, + display_name=display_name, + lifecycle="active", + supported_wire_protocols=wires, + supported_api_modes=modes, + recommended_api_mode=recommended_mode, + recommended_base_url=base_url, + auth_profile=auth, + runtime_provider=runtime_provider, + protocol_hard_limits=dict.fromkeys(modes, (2_000_000, 131_072)), + parameter_schema=parameter_schema, + legacy_parameter_schema=( + _LEGACY_REASONING if reasoning_mode != "none" else {} + ), + reasoning_mode=reasoning_mode, + exact_model_overrides=exact or {}, + discovery_model_overrides=discovery_models or {}, + discovery_capability=discovery, + implementation_fingerprint=_fingerprint(adapter_id, revision), + ) + + +_QWEN37 = ModelDescriptor( + context_tokens=1_000_000, + max_output_tokens=65_536, + # Advisory catalog defaults. Product policy still decides whether the + # application has implemented each input/output workflow. + capabilities=frozenset( + {"text", "vision", "tools", "thinking", "structured_output"} + ), + parameters=_DASHSCOPE_PARAMETERS, +) + +_KIMI_CODE_DISCOVERY_MODELS = { + "k3": ModelDiscoveryDescriptor( + context_tokens=1_048_576, + max_output_tokens=None, + capabilities=frozenset({"text", "thinking"}), + reasoning_mode="effort", + reasoning_efforts=("low", "high", "max"), + default_reasoning_effort="high", + source="kimi_code_official_catalog", + ), + "kimi-for-coding": ModelDiscoveryDescriptor( + context_tokens=262_144, + max_output_tokens=None, + capabilities=frozenset({"text", "thinking"}), + reasoning_mode="effort", + reasoning_efforts=("high",), + default_reasoning_effort="high", + source="kimi_code_official_catalog", + ), + "kimi-for-coding-highspeed": ModelDiscoveryDescriptor( + context_tokens=262_144, + max_output_tokens=None, + capabilities=frozenset({"text", "thinking"}), + reasoning_mode="effort", + reasoning_efforts=("high",), + default_reasoning_effort="high", + source="kimi_code_official_catalog", + ), +} + +BUILTIN_ADAPTER_REGISTRY = AdapterRegistry( + ( + _registration( + "anthropic", + "anthropic-v1", + "Anthropic", + ("anthropic_native",), + ("messages",), + "messages", + "https://api.anthropic.com", + "anthropic_api_key", + "anthropic", + _ANTHROPIC_PARAMETERS, + reasoning_mode="boolean", + ), + _registration( + "openai", + "openai-v1", + "OpenAI", + ("openai_native",), + ("responses", "chat_completions"), + "responses", + "https://api.openai.com/v1", + "bearer", + "openai", + _OPENAI_PARAMETERS, + reasoning_mode="effort", + discovery_models=_KIMI_CODE_DISCOVERY_MODELS, + ), + _registration( + "google-gemini", + "google-gemini-v1", + "Google Gemini Developer API", + ("gemini_native",), + ("interactions", "generate_content"), + "interactions", + "https://generativelanguage.googleapis.com", + "google_api_key", + "google_genai", + _GEMINI_PARAMETERS, + reasoning_mode="effort", + ), + _registration( + "xai", + "xai-v1", + "xAI Grok", + ("openai_compatible",), + ("responses", "chat_completions"), + "responses", + "https://api.x.ai/v1", + "bearer", + "openai", + _XAI_PARAMETERS, + reasoning_mode="effort", + ), + _registration( + "dashscope", + "dashscope-v1", + "Alibaba Cloud DashScope", + ("openai_compatible",), + ("chat_completions", "responses"), + "chat_completions", + "https://dashscope.aliyuncs.com/compatible-mode/v1", + "bearer", + "openai", + _DASHSCOPE_PARAMETERS, + reasoning_mode="boolean", + exact={"qwen3.7-plus": _QWEN37, "qwen3.7-plus-2026-05-26": _QWEN37}, + ), + _registration( + "generic-openai-compatible", + "generic-openai-compatible-v1", + "Generic OpenAI Compatible", + ("openai_compatible",), + ("chat_completions",), + "chat_completions", + "https://example.invalid/v1", + "bearer", + "openai", + _GENERIC_OPENAI_PARAMETERS, + ), + ), + policy_key_id="evoscientist-adapter-registry-release-v2", + manifest_public_key="9NVTpwijNh5L4+yykcr6uoJukNGgihp8V8DGSHSXunA=", + manifest_signature="7/3X6ptDOf5KLxjEepI2jDsHPIu2l/2XfHXzRbveQVIPbD9cLp82OIUM9QEPejKzzWcQqRnNekNyOUy2OrLWCA==", +) + + +def get_adapter_registry() -> AdapterRegistry: + return BUILTIN_ADAPTER_REGISTRY diff --git a/EvoScientist/llm/config_admin.py b/EvoScientist/llm/config_admin.py new file mode 100644 index 0000000..2425378 --- /dev/null +++ b/EvoScientist/llm/config_admin.py @@ -0,0 +1,2080 @@ +"""Authorized Validate -> Probe -> Commit control plane for model routes.""" + +from __future__ import annotations + +import inspect +import json +import math +import sqlite3 +import uuid +from collections.abc import Awaitable, Callable, Mapping +from dataclasses import asdict +from datetime import UTC, datetime +from typing import Any + +from .adapter_registry import get_adapter_registry +from .configuration import SecretResolver +from .contracts import ( + AdminProposalResult, + CancelProposalRequest, + CommitModelConfigRequest, + CommitModelConfigResult, + CommitProposalRequest, + ConcreteRouteProposal, + CreateProposalRequest, + DiscoverProviderModelsRequest, + DiscoverProviderModelsResult, + EvoRuntimeError, + GetModelConfigRequest, + GetModelConfigResult, + HmacGrantAuthority, + ProbeCandidateRouteRequest, + ProbeCandidateRouteResult, + ProbeProposalRequest, + RollbackConfigRequest, + RollbackConfigResult, + UpdateProposalRequest, + ValidateCandidateConfigRequest, + ValidateCandidateConfigResult, + ValidateProposalRequest, + now_ms, +) +from .crypto import HmacKeyRing, canonical_json_v1, hmac_id, sha256_id +from .model_config import ( + EvoModelConfig, + FileEvoModelConfigStore, + RouteRef, + adapter_revision, + convert_v2_to_v3_draft, + endpoint_fingerprint, + proposal_hash, + resolve_secret, + route_fingerprint, + route_semantics_hash, +) + +_PROPOSAL_TTL_MS = 15 * 60 * 1000 +_ROUTE_SEMANTICS_INFO = "ai4sci/route-semantics-hash/v3" +_ENDPOINT_FINGERPRINT_INFO = "ai4sci/endpoint-fingerprint/v3" +_FIXTURE_DIGEST = sha256_id( + { + "connectivity": "Respond with exactly OK.", + "tool_protocol": ( + "You must call probe_tool exactly once with value set to ok. " + "Do not answer in text." + ), + "version": 3, + } +) + +ProbeRunner = Callable[[EvoModelConfig, RouteRef, str], Awaitable[bool] | bool] +ProbeEventSink = Callable[[Mapping[str, Any]], Awaitable[str]] + + +class EvoModelConfigAdminService: + def __init__( + self, + store: FileEvoModelConfigStore, + *, + grant_authority: HmacGrantAuthority, + identity_key_ring: HmacKeyRing, + secret_resolver: SecretResolver | None = None, + probe_runner: ProbeRunner | None = None, + probe_event_sink: ProbeEventSink | None = None, + secret_store: Any | None = None, + ) -> None: + self.store = store + self.grant_authority = grant_authority + self.identity_key_ring = identity_key_ring + self.secret_resolver = secret_resolver + self.probe_runner = probe_runner or self._default_probe + self.probe_event_sink = probe_event_sink + self.secret_store = secret_store + self._init_schema() + self._recover_commit_operations() + + def _connect(self) -> sqlite3.Connection: + connection = sqlite3.connect(self.store.ops_path) + connection.row_factory = sqlite3.Row + connection.execute("PRAGMA foreign_keys=ON") + connection.execute("PRAGMA busy_timeout=5000") + connection.execute("PRAGMA synchronous=FULL") + connection.execute("PRAGMA journal_mode=WAL") + return connection + + def _init_schema(self) -> None: + with self._connect() as connection: + connection.executescript( + """ + CREATE TABLE IF NOT EXISTS config_candidates ( + proposal_hash TEXT PRIMARY KEY, + subject_id TEXT NOT NULL, + operation_id TEXT NOT NULL, + expected_revision INTEGER NOT NULL, + target_revision INTEGER NOT NULL, + config_identity_key_id TEXT NOT NULL, + canonical_payload TEXT NOT NULL, + routes_json TEXT NOT NULL, + expires_at INTEGER NOT NULL, + created_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS capability_evidence_ops ( + evidence_id TEXT PRIMARY KEY, + proposal_hash TEXT NOT NULL REFERENCES config_candidates(proposal_hash), + route_semantics_hash TEXT NOT NULL, + probe_kind TEXT NOT NULL, + status TEXT NOT NULL, + adapter_revision TEXT NOT NULL, + endpoint_fingerprint TEXT NOT NULL, + secret_fingerprints_json TEXT NOT NULL, + fixture_digest TEXT NOT NULL, + config_identity_key_id TEXT NOT NULL, + expires_at INTEGER NOT NULL, + created_at INTEGER NOT NULL, + UNIQUE (proposal_hash, route_semantics_hash, probe_kind) + ); + CREATE TABLE IF NOT EXISTS config_commit_audit ( + operation_id TEXT PRIMARY KEY, + subject_id TEXT NOT NULL, + expected_revision INTEGER NOT NULL, + actual_revision INTEGER NOT NULL, + proposal_hash TEXT NOT NULL, + evidence_ids_json TEXT NOT NULL, + redacted_diff_json TEXT NOT NULL, + committed_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS admin_operation_journal ( + subject_id TEXT NOT NULL, + action TEXT NOT NULL, + operation_id TEXT NOT NULL, + request_digest TEXT NOT NULL, + result_json TEXT NOT NULL, + committed_at INTEGER NOT NULL, + PRIMARY KEY (subject_id, action, operation_id) + ); + CREATE TABLE IF NOT EXISTS config_admin_schema ( + singleton INTEGER PRIMARY KEY CHECK (singleton = 1), + schema_version INTEGER NOT NULL, + upgraded_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS config_proposals ( + proposal_id TEXT PRIMARY KEY, + subject_id TEXT NOT NULL, + base_revision INTEGER NOT NULL, + target_revision INTEGER NOT NULL, + state TEXT NOT NULL CHECK (state IN ( + 'DRAFT','VALIDATED','PROBING','READY','COMMITTING', + 'COMMITTED','FAILED','CANCELLED','EXPIRED' + )), + state_version INTEGER NOT NULL, + canonical_draft TEXT NOT NULL, + draft_etag TEXT NOT NULL, + validated_digest TEXT, + adapter_registry_revision TEXT NOT NULL, + implementation_fingerprints_json TEXT NOT NULL, + failure_code TEXT, + expires_at INTEGER NOT NULL, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS config_proposal_transitions ( + proposal_id TEXT NOT NULL REFERENCES config_proposals(proposal_id), + state_version INTEGER NOT NULL, + from_state TEXT, + to_state TEXT NOT NULL, + operation_id TEXT NOT NULL, + request_digest TEXT NOT NULL, + created_at INTEGER NOT NULL, + PRIMARY KEY (proposal_id, state_version) + ); + CREATE TABLE IF NOT EXISTS config_proposal_evidence ( + evidence_id TEXT PRIMARY KEY, + proposal_id TEXT NOT NULL REFERENCES config_proposals(proposal_id), + validated_digest TEXT NOT NULL, + provider_ref TEXT NOT NULL, + model_ref TEXT NOT NULL, + route_semantics_hash TEXT NOT NULL, + probe_kind TEXT NOT NULL, + status TEXT NOT NULL CHECK (status IN ('supported','failed')), + adapter_id TEXT NOT NULL, + adapter_revision TEXT NOT NULL, + implementation_fingerprint TEXT NOT NULL, + base_url_fingerprint TEXT NOT NULL, + secret_version INTEGER NOT NULL, + fixture_digest TEXT NOT NULL, + evidence_payload_json TEXT NOT NULL, + expires_at INTEGER NOT NULL, + created_at INTEGER NOT NULL, + UNIQUE (proposal_id, validated_digest, route_semantics_hash, probe_kind) + ); + CREATE TABLE IF NOT EXISTS config_commit_operations ( + operation_id TEXT PRIMARY KEY, + proposal_id TEXT NOT NULL REFERENCES config_proposals(proposal_id), + expected_active_revision INTEGER NOT NULL, + target_revision INTEGER NOT NULL, + stage TEXT NOT NULL CHECK (stage IN ( + 'INTENT_WRITTEN','CONFIG_COMMITTED','SECRETS_COORDINATED', + 'COMPLETED','FAILED' + )), + old_config_ref TEXT, + new_config_ref TEXT, + old_secret_refs_json TEXT NOT NULL, + new_secret_refs_json TEXT NOT NULL, + error_code TEXT, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL + ); + """ + ) + connection.execute( + """INSERT INTO config_admin_schema(singleton, schema_version, upgraded_at) + VALUES (1, 2, ?) ON CONFLICT(singleton) DO UPDATE SET + schema_version=MAX(schema_version, 2), upgraded_at=excluded.upgraded_at""", + (now_ms(),), + ) + columns = { + str(row["name"]) + for row in connection.execute( + "PRAGMA table_info(capability_evidence_ops)" + ).fetchall() + } + if "endpoint_fingerprint" not in columns: + connection.execute( + "ALTER TABLE capability_evidence_ops " + "ADD COLUMN endpoint_fingerprint TEXT NOT NULL DEFAULT ''" + ) + if "secret_fingerprints_json" not in columns: + connection.execute( + "ALTER TABLE capability_evidence_ops " + "ADD COLUMN secret_fingerprints_json TEXT NOT NULL DEFAULT '{}'" + ) + + def _recover_commit_operations(self) -> None: + with self._connect() as connection: + rows = connection.execute( + """SELECT * FROM config_commit_operations + WHERE stage NOT IN ('COMPLETED', 'FAILED') ORDER BY created_at""" + ).fetchall() + for row in rows: + try: + active_revision = self.store.current_revision() + except EvoRuntimeError as exc: + if exc.code != "LLM_ROUTE_CONFIGURATION_REQUIRED": + raise + active_revision = 0 + target_revision = int(row["target_revision"]) + expected_revision = int(row["expected_active_revision"]) + operation_id = str(row["operation_id"]) + if active_revision == target_revision: + self._coordinate_secret_refs( + operation_id, + json.loads(str(row["old_secret_refs_json"])), + json.loads(str(row["new_secret_refs_json"])), + ) + recovered_at = now_ms() + with self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + proposal = connection.execute( + "SELECT state, state_version FROM config_proposals WHERE proposal_id=?", + (str(row["proposal_id"]),), + ).fetchone() + if proposal is not None and str(proposal["state"]) == "COMMITTING": + version = int(proposal["state_version"]) + 1 + connection.execute( + """UPDATE config_proposals SET state='COMMITTED', + state_version=?, updated_at=? WHERE proposal_id=?""", + (version, recovered_at, str(row["proposal_id"])), + ) + connection.execute( + """INSERT OR IGNORE INTO config_proposal_transitions + (proposal_id, state_version, from_state, to_state, + operation_id, request_digest, created_at) + VALUES (?, ?, 'COMMITTING', 'COMMITTED', ?, ?, ?)""", + ( + str(row["proposal_id"]), + version, + operation_id, + sha256_id({"recovery": operation_id}), + recovered_at, + ), + ) + connection.execute( + """UPDATE config_commit_operations SET stage='COMPLETED', + updated_at=? WHERE operation_id=?""", + (recovered_at, operation_id), + ) + continue + failure_code = ( + "COMMIT_NOT_APPLIED" + if active_revision == expected_revision + else "CONFIG_REVISION_CONFLICT" + ) + with self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + connection.execute( + """UPDATE config_commit_operations SET stage='FAILED', error_code=?, + updated_at=? WHERE operation_id=?""", + (failure_code, now_ms(), operation_id), + ) + connection.execute( + """UPDATE config_proposals SET state=?, failure_code=?, + state_version=state_version+1, updated_at=? + WHERE proposal_id=? AND state='COMMITTING'""", + ( + "READY" if failure_code == "COMMIT_NOT_APPLIED" else "FAILED", + failure_code, + now_ms(), + str(row["proposal_id"]), + ), + ) + + def _coordinate_secret_refs( + self, operation_id: str, old_refs: list[str], new_refs: list[str] + ) -> None: + if self.secret_store is None: + return + new_set = set(new_refs) + for reference in new_refs: + secret_id, version_text = reference[9:].rsplit("#", 1) + self.secret_store.activate( + secret_id, int(version_text), operation_id=operation_id + ) + for reference in old_refs: + if reference in new_set: + continue + secret_id, version_text = reference[9:].rsplit("#", 1) + self.secret_store.retire( + secret_id, int(version_text), operation_id=operation_id + ) + + def get_config(self, request: GetModelConfigRequest) -> GetModelConfigResult: + self._authorize( + request.admin_grant, + action="model_config:read", + operation_id=request.operation_id, + request_payload={"operation_id": request.operation_id}, + ) + replay = self._replay_operation( + request.admin_grant, + "model_config:read", + {"operation_id": request.operation_id}, + ) + if replay is not None: + return GetModelConfigResult(**replay) + try: + config = self.store.load() + except EvoRuntimeError as exc: + if exc.code != "LLM_ROUTE_CONFIGURATION_REQUIRED": + raise + result = GetModelConfigResult( + request.operation_id, + 0, + _unconfigured_template(self.identity_key_ring.current.key_id), + {"default_alias": "", "aliases": []}, + ) + self._record_operation( + request.admin_grant, + "model_config:read", + {"operation_id": request.operation_id}, + asdict(result), + ) + return result + catalog = { + "default_alias": config.main_routes.default_alias, + "aliases": sorted(config.main_routes.selectable), + } + result = GetModelConfigResult( + request.operation_id, + config.config_revision, + _redact_config(config.raw), + catalog, + ) + self._record_operation( + request.admin_grant, + "model_config:read", + {"operation_id": request.operation_id}, + asdict(result), + ) + return result + + def adapter_metadata(self) -> Mapping[str, Any]: + return get_adapter_registry().metadata() + + async def discover_provider_models( + self, request: DiscoverProviderModelsRequest + ) -> DiscoverProviderModelsResult: + action = "model_config:provider:discover" + request_payload = self._control_payload(request) + self._authorize( + request.admin_grant, + action=action, + operation_id=request.operation_id, + request_payload=request_payload, + ) + replay = self._replay_operation(request.admin_grant, action, request_payload) + if replay is not None: + return DiscoverProviderModelsResult( + operation_id=str(replay["operation_id"]), + proposal_id=str(replay["proposal_id"]), + provider_id=str(replay["provider_id"]), + source=str(replay["source"]), + discovered_at=int(replay["discovered_at"]), + models=tuple(replay["models"]), + ) + with self._connect() as connection: + row = self._proposal_row(connection, request.proposal_id) + self._require_proposal_owner(row, request.admin_grant.subject_id) + if str(row["state"]) in {"COMMITTED", "CANCELLED", "EXPIRED"}: + raise EvoRuntimeError("PROPOSAL_STATE_CONFLICT") + draft = json.loads(str(row["canonical_draft"])) + config = EvoModelConfig.parse(draft, require_evidence=False) + provider = config.providers.get(request.provider_id) + if provider is None: + raise EvoRuntimeError("LLM_ROUTE_CONFIGURATION_REQUIRED") + endpoint = provider.endpoints[request.provider_id] + secret = resolve_secret(endpoint.auth, secret_resolver=self.secret_resolver) + registration = get_adapter_registry().get( + provider.adapter_id, provider.adapter_revision + ) + models = await registration.discover_models( + base_url=endpoint.base_url, + api_key=secret.value, + timeout_seconds=int( + provider.connection_defaults.get("connect_timeout_seconds", 10) + ), + ) + result = DiscoverProviderModelsResult( + operation_id=request.operation_id, + proposal_id=request.proposal_id, + provider_id=request.provider_id, + source="provider_api", + discovered_at=now_ms(), + models=models, + ) + self._record_operation( + request.admin_grant, action, request_payload, asdict(result) + ) + return result + + def create_proposal(self, request: CreateProposalRequest) -> AdminProposalResult: + request_payload = self._control_payload(request) + self._authorize( + request.admin_grant, + action="model_config:proposal:create", + operation_id=request.operation_id, + request_payload=request_payload, + ) + replay = self._replay_operation( + request.admin_grant, "model_config:proposal:create", request_payload + ) + if replay is not None: + return AdminProposalResult(**replay) + active_revision = self.store.current_revision() + if active_revision != request.expected_active_revision: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + try: + active = self.store.load() + if active.schema_version == 3: + draft = dict(active.raw) + migration_report = None + else: + draft, report = convert_v2_to_v3_draft( + active.raw, + target_revision=active_revision + 1, + config_identity_key_id=self.identity_key_ring.current.key_id, + ) + migration_report = asdict(report) + except EvoRuntimeError as exc: + if exc.code not in { + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "CONFIG_INTEGRITY_MISMATCH", + }: + raise + draft = _unconfigured_v3_template(self.identity_key_ring.current.key_id) + migration_report = None + target = active_revision + 1 + draft.update( + { + "schema_version": 3, + "config_revision": target, + "config_identity_key_id": self.identity_key_ring.current.key_id, + "capability_evidence": [], + } + ) + proposal_id = str(uuid.uuid4()) + created = now_ms() + expires_at = created + 24 * 60 * 60 * 1000 + canonical = canonical_json_v1(draft).decode() + etag = self._draft_etag(draft) + registry = get_adapter_registry() + fingerprints = { + item["adapter_id"] + ":" + item["adapter_revision"]: item[ + "implementation_fingerprint" + ] + for item in registry.metadata()["adapters"] + } + with self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + connection.execute( + """INSERT INTO config_proposals + (proposal_id, subject_id, base_revision, target_revision, state, + state_version, canonical_draft, draft_etag, validated_digest, + adapter_registry_revision, implementation_fingerprints_json, + expires_at, created_at, updated_at) + VALUES (?, ?, ?, ?, 'DRAFT', 1, ?, ?, NULL, ?, ?, ?, ?, ?)""", + ( + proposal_id, + request.admin_grant.subject_id, + active_revision, + target, + canonical, + etag, + registry.registry_revision, + canonical_json_v1(fingerprints).decode(), + expires_at, + created, + created, + ), + ) + self._insert_transition( + connection, + proposal_id=proposal_id, + state_version=1, + from_state=None, + to_state="DRAFT", + request=request, + request_payload=request_payload, + ) + result = AdminProposalResult( + request.operation_id, + active_revision, + proposal_id, + "DRAFT", + 1, + etag, + active_revision, + target, + expires_at, + draft_payload=_redact_config(draft), + migration_report=migration_report, + ) + self._record_operation( + request.admin_grant, + "model_config:proposal:create", + request_payload, + asdict(result), + ) + return result + + def update_proposal(self, request: UpdateProposalRequest) -> AdminProposalResult: + action = "model_config:proposal:update" + request_payload = self._control_payload(request) + self._authorize( + request.admin_grant, + action=action, + operation_id=request.operation_id, + request_payload=request_payload, + ) + replay = self._replay_operation(request.admin_grant, action, request_payload) + if replay is not None: + return AdminProposalResult(**replay) + with self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + row = self._proposal_row(connection, request.proposal_id) + self._require_proposal_owner(row, request.admin_grant.subject_id) + self._require_proposal_cas( + row, + request.expected_state_version, + request.expected_draft_etag, + allowed_states={"DRAFT", "VALIDATED", "READY", "FAILED"}, + ) + draft = dict(request.draft_payload) + draft["schema_version"] = 3 + draft["config_revision"] = int(row["target_revision"]) + draft["config_identity_key_id"] = self.identity_key_ring.current.key_id + draft["capability_evidence"] = [] + etag = self._draft_etag(draft) + version = int(row["state_version"]) + 1 + connection.execute( + """UPDATE config_proposals SET state='DRAFT', state_version=?, + canonical_draft=?, draft_etag=?, validated_digest=NULL, + failure_code=NULL, updated_at=? + WHERE proposal_id=? AND state_version=?""", + ( + version, + canonical_json_v1(draft).decode(), + etag, + now_ms(), + request.proposal_id, + request.expected_state_version, + ), + ) + connection.execute( + "DELETE FROM config_proposal_evidence WHERE proposal_id=?", + (request.proposal_id,), + ) + self._insert_transition( + connection, + proposal_id=request.proposal_id, + state_version=version, + from_state=str(row["state"]), + to_state="DRAFT", + request=request, + request_payload=request_payload, + ) + updated = self._proposal_row(connection, request.proposal_id) + result = self._proposal_result(updated, request.operation_id, draft=draft) + self._record_operation( + request.admin_grant, action, request_payload, asdict(result) + ) + return result + + def validate_proposal( + self, request: ValidateProposalRequest + ) -> AdminProposalResult: + action = "model_config:proposal:validate" + request_payload = self._control_payload(request) + self._authorize( + request.admin_grant, + action=action, + operation_id=request.operation_id, + request_payload=request_payload, + ) + replay = self._replay_operation(request.admin_grant, action, request_payload) + if replay is not None: + return AdminProposalResult(**replay) + with self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + row = self._proposal_row(connection, request.proposal_id) + self._require_proposal_owner(row, request.admin_grant.subject_id) + self._require_proposal_cas( + row, + request.expected_state_version, + request.expected_draft_etag, + allowed_states={"DRAFT", "FAILED"}, + ) + draft = json.loads(str(row["canonical_draft"])) + config = EvoModelConfig.parse(draft, require_evidence=False) + routes = tuple(asdict(item) for item in self._build_proposals(config)) + validated_digest = self._validated_digest( + draft, str(row["adapter_registry_revision"]), routes + ) + version = int(row["state_version"]) + 1 + connection.execute( + """UPDATE config_proposals SET state='VALIDATED', state_version=?, + validated_digest=?, failure_code=NULL, updated_at=? + WHERE proposal_id=? AND state_version=?""", + ( + version, + validated_digest, + now_ms(), + request.proposal_id, + request.expected_state_version, + ), + ) + connection.execute( + "DELETE FROM config_proposal_evidence WHERE proposal_id=?", + (request.proposal_id,), + ) + self._insert_transition( + connection, + proposal_id=request.proposal_id, + state_version=version, + from_state=str(row["state"]), + to_state="VALIDATED", + request=request, + request_payload=request_payload, + ) + updated = self._proposal_row(connection, request.proposal_id) + result = self._proposal_result( + updated, request.operation_id, routes=routes, draft=draft + ) + self._record_operation( + request.admin_grant, action, request_payload, asdict(result) + ) + return result + + async def probe_proposal( + self, request: ProbeProposalRequest + ) -> AdminProposalResult: + action = "model_config:proposal:probe" + request_payload = self._control_payload(request) + self._authorize( + request.admin_grant, + action=action, + operation_id=request.operation_id, + request_payload=request_payload, + ) + replay = self._replay_operation(request.admin_grant, action, request_payload) + if replay is not None: + return AdminProposalResult(**replay) + with self._connect() as connection: + row = self._proposal_row(connection, request.proposal_id) + self._require_proposal_owner(row, request.admin_grant.subject_id) + if str(row["state"]) not in {"VALIDATED", "PROBING", "READY"}: + raise EvoRuntimeError("PROPOSAL_STATE_CONFLICT") + if str(row["validated_digest"] or "") != request.validated_digest: + raise EvoRuntimeError("PROPOSAL_ETAG_CONFLICT") + draft = json.loads(str(row["canonical_draft"])) + config = EvoModelConfig.parse(draft, require_evidence=False) + routes = tuple(self._build_proposals(config)) + route_proposal = next( + ( + item + for item in routes + if item.route_semantics_hash == request.route_semantics_hash + ), + None, + ) + if ( + route_proposal is None + or request.probe_kind not in route_proposal.required_probe_kinds + ): + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + route = _find_route(config, route_proposal.route) + supported = await self._run_probe(config, route, request.probe_kind) + provider = config.providers[route.provider] + endpoint = provider.endpoints[route.endpoint] + evidence_id = str(uuid.uuid4()) + expires_at = now_ms() + _PROPOSAL_TTL_MS + evidence_payload = { + "provider_ref": route.provider, + "model_ref": route.model, + "probe_kind": request.probe_kind, + "status": "supported" if supported else "failed", + } + with self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + current = self._proposal_row(connection, request.proposal_id) + if str(current["validated_digest"] or "") != request.validated_digest: + raise EvoRuntimeError("PROPOSAL_ETAG_CONFLICT") + connection.execute( + """INSERT INTO config_proposal_evidence + (evidence_id, proposal_id, validated_digest, provider_ref, + model_ref, route_semantics_hash, probe_kind, status, adapter_id, + adapter_revision, implementation_fingerprint, base_url_fingerprint, + secret_version, fixture_digest, evidence_payload_json, expires_at, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(proposal_id, validated_digest, route_semantics_hash, probe_kind) + DO UPDATE SET evidence_id=excluded.evidence_id, status=excluded.status, + evidence_payload_json=excluded.evidence_payload_json, + expires_at=excluded.expires_at, created_at=excluded.created_at""", + ( + evidence_id, + request.proposal_id, + request.validated_digest, + route.provider, + route.model, + request.route_semantics_hash, + request.probe_kind, + "supported" if supported else "failed", + provider.adapter_id, + provider.adapter_revision, + provider.implementation_fingerprint, + route_proposal.endpoint_fingerprint, + endpoint.auth.revision, + _FIXTURE_DIGEST, + canonical_json_v1(evidence_payload).decode(), + expires_at, + now_ms(), + ), + ) + evidence_rows = connection.execute( + """SELECT * FROM config_proposal_evidence + WHERE proposal_id=? AND validated_digest=?""", + (request.proposal_id, request.validated_digest), + ).fetchall() + ready = self._proposal_evidence_complete(routes, evidence_rows) + next_state = "READY" if ready else "VALIDATED" + version = int(current["state_version"]) + 1 + connection.execute( + """UPDATE config_proposals SET state=?, state_version=?, updated_at=? + WHERE proposal_id=? AND state_version=?""", + ( + next_state, + version, + now_ms(), + request.proposal_id, + int(current["state_version"]), + ), + ) + self._insert_transition( + connection, + proposal_id=request.proposal_id, + state_version=version, + from_state=str(current["state"]), + to_state=next_state, + request=request, + request_payload=request_payload, + ) + updated = self._proposal_row(connection, request.proposal_id) + result = self._proposal_result( + updated, + request.operation_id, + routes=tuple(asdict(item) for item in routes), + evidence_ids=tuple(str(item["evidence_id"]) for item in evidence_rows), + ) + self._record_operation( + request.admin_grant, action, request_payload, asdict(result) + ) + return result + + def commit_proposal(self, request: CommitProposalRequest) -> AdminProposalResult: + action = "model_config:proposal:commit" + request_payload = self._control_payload(request) + self._authorize( + request.admin_grant, + action=action, + operation_id=request.operation_id, + request_payload=request_payload, + ) + replay = self._replay_operation(request.admin_grant, action, request_payload) + if replay is not None: + return AdminProposalResult(**replay) + with self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + row = self._proposal_row(connection, request.proposal_id) + self._require_proposal_owner(row, request.admin_grant.subject_id) + self._require_proposal_cas( + row, + request.expected_state_version, + request.expected_draft_etag, + allowed_states={"READY"}, + ) + if ( + int(row["base_revision"]) != request.expected_active_revision + or self.store.current_revision() != request.expected_active_revision + or str(row["validated_digest"] or "") != request.validated_digest + ): + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + draft = json.loads(str(row["canonical_draft"])) + config = EvoModelConfig.parse(draft, require_evidence=False) + routes = tuple(self._build_proposals(config)) + evidence_rows = connection.execute( + """SELECT * FROM config_proposal_evidence + WHERE proposal_id=? AND validated_digest=? AND evidence_id IN ({})""".format( + ",".join("?" for _ in request.evidence_ids) or "NULL" + ), + ( + request.proposal_id, + request.validated_digest, + *request.evidence_ids, + ), + ).fetchall() + if not self._proposal_evidence_complete(routes, evidence_rows): + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + version = int(row["state_version"]) + 1 + connection.execute( + """UPDATE config_proposals SET state='COMMITTING', state_version=?, + updated_at=? WHERE proposal_id=? AND state_version=?""", + ( + version, + now_ms(), + request.proposal_id, + request.expected_state_version, + ), + ) + old_refs = ( + self._config_secret_refs(self.store.load().raw) + if request.expected_active_revision + else [] + ) + new_refs = self._config_secret_refs(draft) + connection.execute( + """INSERT INTO config_commit_operations + (operation_id, proposal_id, expected_active_revision, target_revision, + stage, old_config_ref, new_config_ref, old_secret_refs_json, + new_secret_refs_json, created_at, updated_at) + VALUES (?, ?, ?, ?, 'INTENT_WRITTEN', ?, ?, ?, ?, ?, ?)""", + ( + request.operation_id, + request.proposal_id, + request.expected_active_revision, + int(row["target_revision"]), + f"revision:{request.expected_active_revision}", + f"revision:{int(row['target_revision'])}", + canonical_json_v1(old_refs).decode(), + canonical_json_v1(new_refs).decode(), + now_ms(), + now_ms(), + ), + ) + self._insert_transition( + connection, + proposal_id=request.proposal_id, + state_version=version, + from_state=str(row["state"]), + to_state="COMMITTING", + request=request, + request_payload=request_payload, + ) + final_evidence = self._v3_evidence_payload(config, routes, evidence_rows) + final_payload = dict(draft) + final_payload["capability_evidence"] = final_evidence + revision = self.store.commit_validated( + final_payload, + expected_revision=request.expected_active_revision, + operation_id=request.operation_id, + ).config_revision + with self._connect() as connection: + connection.execute( + """UPDATE config_commit_operations SET stage='CONFIG_COMMITTED', + updated_at=? WHERE operation_id=?""", + (now_ms(), request.operation_id), + ) + self._coordinate_secret_refs(request.operation_id, old_refs, new_refs) + with self._connect() as connection: + connection.execute( + """UPDATE config_commit_operations SET stage='SECRETS_COORDINATED', + updated_at=? WHERE operation_id=?""", + (now_ms(), request.operation_id), + ) + committed_at = now_ms() + with self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + current = self._proposal_row(connection, request.proposal_id) + version = int(current["state_version"]) + 1 + connection.execute( + """UPDATE config_proposals SET state='COMMITTED', state_version=?, + updated_at=? WHERE proposal_id=? AND state_version=?""", + ( + version, + committed_at, + request.proposal_id, + int(current["state_version"]), + ), + ) + connection.execute( + """UPDATE config_commit_operations SET stage='COMPLETED', updated_at=? + WHERE operation_id=?""", + (committed_at, request.operation_id), + ) + self._insert_transition( + connection, + proposal_id=request.proposal_id, + state_version=version, + from_state=str(current["state"]), + to_state="COMMITTED", + request=request, + request_payload=request_payload, + ) + updated = self._proposal_row(connection, request.proposal_id) + result = self._proposal_result( + updated, + request.operation_id, + evidence_ids=tuple(sorted(request.evidence_ids)), + committed_at=committed_at, + ) + if result.active_revision != revision: + result = AdminProposalResult( + **{**asdict(result), "active_revision": revision} + ) + self._record_operation( + request.admin_grant, action, request_payload, asdict(result) + ) + return result + + def cancel_proposal(self, request: CancelProposalRequest) -> AdminProposalResult: + action = "model_config:proposal:cancel" + request_payload = self._control_payload(request) + self._authorize( + request.admin_grant, + action=action, + operation_id=request.operation_id, + request_payload=request_payload, + ) + replay = self._replay_operation(request.admin_grant, action, request_payload) + if replay is not None: + return AdminProposalResult(**replay) + with self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + row = self._proposal_row(connection, request.proposal_id) + self._require_proposal_owner(row, request.admin_grant.subject_id) + self._require_proposal_cas( + row, + request.expected_state_version, + None, + allowed_states={"DRAFT", "VALIDATED", "READY", "FAILED"}, + ) + version = int(row["state_version"]) + 1 + connection.execute( + """UPDATE config_proposals SET state='CANCELLED', state_version=?, + updated_at=? WHERE proposal_id=? AND state_version=?""", + ( + version, + now_ms(), + request.proposal_id, + request.expected_state_version, + ), + ) + self._insert_transition( + connection, + proposal_id=request.proposal_id, + state_version=version, + from_state=str(row["state"]), + to_state="CANCELLED", + request=request, + request_payload=request_payload, + ) + updated = self._proposal_row(connection, request.proposal_id) + result = self._proposal_result(updated, request.operation_id) + self._record_operation( + request.admin_grant, action, request_payload, asdict(result) + ) + return result + + def rollback_config(self, request: RollbackConfigRequest) -> RollbackConfigResult: + action = "model_config:rollback" + request_payload = self._control_payload(request) + self._authorize( + request.admin_grant, + action=action, + operation_id=request.operation_id, + request_payload=request_payload, + ) + replay = self._replay_operation(request.admin_grant, action, request_payload) + if replay is not None: + return RollbackConfigResult(**replay) + if request.target_revision == request.expected_active_revision: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + # Resolve the archived target before switching the active pointer. This + # rejects blocked adapters, stale evidence, and revoked credentials. + target = self.store.load_revision(request.target_revision) + for selector_id in target.route_selectors: + route = target.concrete_routes(selector_id)[0] + self._secret_fingerprints(target, route) + revision = self.store.activate_revision( + request.target_revision, + expected_revision=request.expected_active_revision, + operation_id=request.operation_id, + ).config_revision + result = RollbackConfigResult( + operation_id=request.operation_id, + previous_active_revision=request.expected_active_revision, + active_revision=revision, + target_revision=request.target_revision, + rolled_back_at=now_ms(), + ) + self._record_operation( + request.admin_grant, action, request_payload, asdict(result) + ) + return result + + @staticmethod + def _control_payload(request: Any) -> Mapping[str, Any]: + payload = asdict(request) + payload.pop("admin_grant", None) + return payload + + def _draft_etag(self, payload: Mapping[str, Any]) -> str: + _, key = self.identity_key_ring.derive_current("ai4sci/admin-draft-etag/v2") + return hmac_id(key, payload) + + def _validated_digest( + self, + draft: Mapping[str, Any], + registry_revision: str, + routes: tuple[Mapping[str, Any], ...], + ) -> str: + _, key = self.identity_key_ring.derive_current( + "ai4sci/admin-validated-digest/v2" + ) + return hmac_id( + key, + { + "draft": draft, + "registry_revision": registry_revision, + "routes": routes, + }, + ) + + @staticmethod + def _proposal_row(connection: sqlite3.Connection, proposal_id: str) -> sqlite3.Row: + row = connection.execute( + "SELECT * FROM config_proposals WHERE proposal_id=?", (proposal_id,) + ).fetchone() + if row is None: + raise EvoRuntimeError("PROPOSAL_NOT_FOUND") + if int(row["expires_at"]) < now_ms() and str(row["state"]) not in { + "COMMITTED", + "CANCELLED", + "EXPIRED", + }: + connection.execute( + "UPDATE config_proposals SET state='EXPIRED', updated_at=? WHERE proposal_id=?", + (now_ms(), proposal_id), + ) + raise EvoRuntimeError("PROPOSAL_EXPIRED") + return row + + @staticmethod + def _require_proposal_owner(row: sqlite3.Row, subject_id: str) -> None: + if str(row["subject_id"]) != subject_id: + raise EvoRuntimeError("ADMIN_CONFIG_FORBIDDEN") + + @staticmethod + def _require_proposal_cas( + row: sqlite3.Row, + expected_state_version: int, + expected_draft_etag: str | None, + *, + allowed_states: set[str], + ) -> None: + if str(row["state"]) not in allowed_states: + raise EvoRuntimeError("PROPOSAL_STATE_CONFLICT") + if int(row["state_version"]) != expected_state_version: + raise EvoRuntimeError("PROPOSAL_STATE_CONFLICT") + if ( + expected_draft_etag is not None + and str(row["draft_etag"]) != expected_draft_etag + ): + raise EvoRuntimeError("PROPOSAL_ETAG_CONFLICT") + + @staticmethod + def _insert_transition( + connection: sqlite3.Connection, + *, + proposal_id: str, + state_version: int, + from_state: str | None, + to_state: str, + request: Any, + request_payload: Mapping[str, Any], + ) -> None: + connection.execute( + """INSERT INTO config_proposal_transitions + (proposal_id, state_version, from_state, to_state, operation_id, + request_digest, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)""", + ( + proposal_id, + state_version, + from_state, + to_state, + request.operation_id, + sha256_id(request_payload), + now_ms(), + ), + ) + + def _proposal_result( + self, + row: sqlite3.Row, + operation_id: str, + *, + routes: tuple[Mapping[str, Any], ...] = (), + evidence_ids: tuple[str, ...] = (), + committed_at: int | None = None, + draft: Mapping[str, Any] | None = None, + ) -> AdminProposalResult: + return AdminProposalResult( + operation_id=operation_id, + active_revision=self.store.current_revision(), + proposal_id=str(row["proposal_id"]), + state=str(row["state"]), + state_version=int(row["state_version"]), + draft_etag=str(row["draft_etag"]), + base_revision=int(row["base_revision"]), + target_revision=int(row["target_revision"]), + expires_at=int(row["expires_at"]), + validated_digest=( + str(row["validated_digest"]) if row["validated_digest"] else None + ), + routes=routes, + evidence_ids=evidence_ids, + committed_at=committed_at, + draft_payload=_redact_config(draft) if draft is not None else None, + ) + + @staticmethod + def _proposal_evidence_complete( + routes: tuple[ConcreteRouteProposal, ...], rows: list[sqlite3.Row] + ) -> bool: + supported = { + (str(row["route_semantics_hash"]), str(row["probe_kind"])) + for row in rows + if str(row["status"]) == "supported" and int(row["expires_at"]) >= now_ms() + } + return all( + (route.route_semantics_hash, kind) in supported + for route in routes + for kind in route.required_probe_kinds + ) + + @staticmethod + def _config_secret_refs(payload: Mapping[str, Any]) -> list[str]: + return sorted( + { + str(provider.get("connection", {}).get("credential_ref")) + for provider in payload.get("providers") or [] + if provider.get("connection", {}).get("credential_ref") + } + ) + + @staticmethod + def _v3_evidence_payload( + config: EvoModelConfig, + routes: tuple[ConcreteRouteProposal, ...], + rows: list[sqlite3.Row], + ) -> list[Mapping[str, Any]]: + rows_by_route: dict[str, dict[str, sqlite3.Row]] = {} + for row in rows: + rows_by_route.setdefault(str(row["route_semantics_hash"]), {})[ + str(row["probe_kind"]) + ] = row + result = [] + seen: set[tuple[str, str]] = set() + for proposal in routes: + route = _find_route(config, proposal.route) + target = (route.provider, route.model) + if target in seen: + continue + seen.add(target) + provider = config.providers[route.provider] + model = provider.models[route.model] + endpoint = provider.endpoints[route.endpoint] + evidence = rows_by_route[proposal.route_semantics_hash] + results = { + key: ( + "supported" + if key == "text" + else "not_declared" + if not declared + else str(evidence[key]["status"]) + ) + for key, declared in model.capabilities.items() + } + results["connectivity"] = str(evidence["connectivity"]["status"]) + result.append( + { + "provider_ref": route.provider, + "model_ref": route.model, + "adapter_id": provider.adapter_id, + "adapter_revision": provider.adapter_revision, + "implementation_fingerprint": provider.implementation_fingerprint, + "wire_protocol": provider.wire_protocol, + "provider_model_id": model.model_id, + "resolved_model_revision": model.resolved_model_revision, + "version_policy": model.version_policy, + "reproducible": model.reproducible, + "api_mode": route.api_mode, + "tool_call_transport": route.tool_call_transport, + "base_url_fingerprint": proposal.endpoint_fingerprint, + "secret_version": endpoint.auth.revision, + "route_semantics_hash": proposal.route_semantics_hash, + "fixture_digest": _FIXTURE_DIGEST, + "verified_at": datetime.now(UTC).isoformat(), + "evidence_expires_at": datetime.fromtimestamp( + min(int(item["expires_at"]) for item in evidence.values()) + / 1000, + UTC, + ).isoformat(), + "results": results, + } + ) + return result + + def validate_candidate( + self, request: ValidateCandidateConfigRequest + ) -> ValidateCandidateConfigResult: + if int(request.payload.get("schema_version", 0)) == 3: + raise EvoRuntimeError("ADMIN_CONTROL_UPGRADE_REQUIRED") + try: + if self.store.load().schema_version == 3: + raise EvoRuntimeError("V2_CONFIG_WRITE_DISABLED") + except EvoRuntimeError as exc: + if exc.code != "LLM_ROUTE_CONFIGURATION_REQUIRED": + raise + request_payload = { + "operation_id": request.operation_id, + "expected_revision": request.expected_revision, + "payload": request.payload, + } + self._authorize( + request.admin_grant, + action="model_config:validate", + operation_id=request.operation_id, + request_payload=request_payload, + ) + replay = self._replay_operation( + request.admin_grant, "model_config:validate", request_payload + ) + if replay is not None: + return ValidateCandidateConfigResult( + operation_id=str(replay["operation_id"]), + target_revision=int(replay["target_revision"]), + config_identity_key_id=str(replay["config_identity_key_id"]), + proposal_hash=str(replay["proposal_hash"]), + concrete_routes=tuple( + ConcreteRouteProposal( + route=item["route"], + route_semantics_hash=item["route_semantics_hash"], + endpoint_fingerprint=item["endpoint_fingerprint"], + required_probe_kinds=tuple(item["required_probe_kinds"]), + ) + for item in replay["concrete_routes"] + ), + expires_at=int(replay["expires_at"]), + ) + if "capability_evidence" in request.payload: + raise EvoRuntimeError("CANDIDATE_EVIDENCE_FORBIDDEN") + if self.store.current_revision() != request.expected_revision: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + target = request.expected_revision + 1 + candidate_payload = dict(request.payload) + candidate_payload["config_revision"] = target + candidate_payload["capability_evidence"] = [] + config = EvoModelConfig.parse(candidate_payload, require_evidence=False) + if config.config_identity_key_id != self.identity_key_ring.current.key_id: + raise EvoRuntimeError("CONFIG_IDENTITY_KEY_UNKNOWN") + proposals = self._build_proposals(config) + digest = proposal_hash( + request.payload, + target_revision=target, + identity_key=self.identity_key_ring.derive_current( + "ai4sci/proposal-hash/v3" + )[1], + ) + expires_at = now_ms() + _PROPOSAL_TTL_MS + with self._connect() as connection: + connection.execute( + """INSERT INTO config_candidates + (proposal_hash, subject_id, operation_id, expected_revision, + target_revision, config_identity_key_id, canonical_payload, + routes_json, expires_at, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(proposal_hash) DO NOTHING""", + ( + digest, + request.admin_grant.subject_id, + request.operation_id, + request.expected_revision, + target, + config.config_identity_key_id, + canonical_json_v1(request.payload).decode(), + canonical_json_v1([asdict(item) for item in proposals]).decode(), + expires_at, + now_ms(), + ), + ) + result = ValidateCandidateConfigResult( + request.operation_id, + target, + config.config_identity_key_id, + digest, + tuple(proposals), + expires_at, + ) + self._record_operation( + request.admin_grant, + "model_config:validate", + request_payload, + asdict(result), + ) + return result + + async def probe_candidate( + self, request: ProbeCandidateRouteRequest + ) -> ProbeCandidateRouteResult: + request_payload = { + "operation_id": request.operation_id, + "proposal_hash": request.proposal_hash, + "route_semantics_hash": request.route_semantics_hash, + "probe_kind": request.probe_kind, + } + self._authorize( + request.admin_grant, + action="model_config:probe", + operation_id=request.operation_id, + request_payload=request_payload, + ) + replay = self._replay_operation( + request.admin_grant, "model_config:probe", request_payload + ) + if replay is not None: + return ProbeCandidateRouteResult(**replay) + with self._connect() as connection: + candidate = connection.execute( + "SELECT * FROM config_candidates WHERE proposal_hash=?", + (request.proposal_hash,), + ).fetchone() + if candidate is None or int(candidate["expires_at"]) < now_ms(): + raise EvoRuntimeError("CANDIDATE_EXPIRED") + route_rows = json.loads(candidate["routes_json"]) + route_row = next( + ( + item + for item in route_rows + if item["route_semantics_hash"] == request.route_semantics_hash + and request.probe_kind in item["required_probe_kinds"] + ), + None, + ) + if route_row is None: + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + payload = json.loads(candidate["canonical_payload"]) + payload["config_revision"] = int(candidate["target_revision"]) + payload["capability_evidence"] = [] + config = EvoModelConfig.parse(payload, require_evidence=False) + route = _find_route(config, route_row["route"]) + adapter = adapter_revision( + config.providers[route.provider].protocol, route.api_mode + ) + endpoint_digest = endpoint_fingerprint( + config, + route, + self.identity_key_ring.derive_current(_ENDPOINT_FINGERPRINT_INFO)[1], + ) + if endpoint_digest != route_row["endpoint_fingerprint"]: + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + secret_fingerprints = self._secret_fingerprints(config, route) + secret_fingerprints_json = canonical_json_v1(secret_fingerprints).decode() + attempt_id = str( + uuid.uuid5( + uuid.UUID("afc5ad1e-a40c-4d6d-a505-2d870128982e"), + "\0".join( + ( + request.admin_grant.subject_id, + request.operation_id, + request.proposal_hash, + request.route_semantics_hash, + request.probe_kind, + ) + ), + ) + ) + model = config.route_model(route) + quote = model.quote + input_bound = len( + ( + "Respond with exactly OK." + if request.probe_kind == "connectivity" + else "You must call probe_tool exactly once with value set to ok. Do not answer in text." + ).encode("utf-8") + ) + reserve = math.ceil( + ( + input_bound * quote.input_microunits_per_million + + 128 * quote.output_microunits_per_million + ) + / quote.unit_scale + ) + route_key = self.identity_key_ring.derive_current( + "ai4sci/route-fingerprint/v3" + )[1] + event_base = { + "attempt_id": attempt_id, + "operation_id": request.operation_id, + "subject_id": request.admin_grant.subject_id, + "proposal_hash": request.proposal_hash, + "route_semantics_hash": request.route_semantics_hash, + "route_fingerprint": route_fingerprint(config, route, route_key), + "quote_id": quote.quote_id, + "probe_kind": request.probe_kind, + "billing_intent": "platform_cost", + "provider_reserved_microunits": reserve, + "adapter_revision": adapter, + "fixture_digest": _FIXTURE_DIGEST, + } + if self.probe_event_sink is not None: + ingress = await self.probe_event_sink( + {**event_base, "outcome": "started", "provider_request_started": True} + ) + if ingress.startswith("terminal:supported"): + supported = True + elif ingress != "committed": + raise EvoRuntimeError("EVENT_COMMIT_INDETERMINATE") + else: + try: + supported = await self._run_probe(config, route, request.probe_kind) + except Exception: + await self.probe_event_sink( + { + **event_base, + "outcome": "usage_unconfirmed", + "provider_request_started": True, + "probe_result": "failed", + } + ) + raise + await self.probe_event_sink( + { + **event_base, + "outcome": "usage_unconfirmed", + "provider_request_started": True, + "probe_result": "supported" if supported else "failed", + } + ) + else: + supported = await self._run_probe(config, route, request.probe_kind) + if not supported: + raise EvoRuntimeError("CAPABILITY_PROBE_FAILED") + evidence_id = str(uuid.uuid4()) + expires_at = int(candidate["expires_at"]) + with self._connect() as connection: + existing = connection.execute( + """SELECT evidence_id, adapter_revision, endpoint_fingerprint, + secret_fingerprints_json, fixture_digest, + config_identity_key_id, expires_at + FROM capability_evidence_ops + WHERE proposal_hash=? AND route_semantics_hash=? AND probe_kind=?""", + ( + request.proposal_hash, + request.route_semantics_hash, + request.probe_kind, + ), + ).fetchone() + if ( + existing is not None + and existing["adapter_revision"] == adapter + and existing["endpoint_fingerprint"] == endpoint_digest + and existing["secret_fingerprints_json"] == secret_fingerprints_json + and existing["fixture_digest"] == _FIXTURE_DIGEST + and existing["config_identity_key_id"] == config.config_identity_key_id + and int(existing["expires_at"]) >= now_ms() + ): + evidence_id = str(existing["evidence_id"]) + else: + if existing is not None: + connection.execute( + "DELETE FROM capability_evidence_ops WHERE evidence_id=?", + (existing["evidence_id"],), + ) + connection.execute( + """INSERT INTO capability_evidence_ops + (evidence_id, proposal_hash, route_semantics_hash, + probe_kind, status, adapter_revision, + endpoint_fingerprint, secret_fingerprints_json, + fixture_digest, config_identity_key_id, expires_at, created_at) + VALUES (?, ?, ?, ?, 'supported', ?, ?, ?, ?, ?, ?, ?)""", + ( + evidence_id, + request.proposal_hash, + request.route_semantics_hash, + request.probe_kind, + adapter, + endpoint_digest, + secret_fingerprints_json, + _FIXTURE_DIGEST, + config.config_identity_key_id, + expires_at, + now_ms(), + ), + ) + result = ProbeCandidateRouteResult( + request.operation_id, + evidence_id, + request.route_semantics_hash, + "supported", + "supported" if request.probe_kind == "tool_protocol" else "unknown", + adapter, + _FIXTURE_DIGEST, + expires_at, + ) + self._record_operation( + request.admin_grant, + "model_config:probe", + request_payload, + asdict(result), + ) + return result + + async def _run_probe( + self, config: EvoModelConfig, route: RouteRef, probe_kind: str + ) -> bool: + result = self.probe_runner(config, route, probe_kind) + return await result if inspect.isawaitable(result) else bool(result) + + def commit(self, request: CommitModelConfigRequest) -> CommitModelConfigResult: + if int(request.payload.get("schema_version", 0)) == 3: + raise EvoRuntimeError("ADMIN_CONTROL_UPGRADE_REQUIRED") + try: + if self.store.load().schema_version == 3: + raise EvoRuntimeError("V2_CONFIG_WRITE_DISABLED") + except EvoRuntimeError as exc: + if exc.code != "LLM_ROUTE_CONFIGURATION_REQUIRED": + raise + request_payload = { + "operation_id": request.operation_id, + "expected_revision": request.expected_revision, + "proposal_hash": request.proposal_hash, + "payload": request.payload, + "evidence_ids": sorted(set(request.evidence_ids)), + } + self._authorize( + request.admin_grant, + action="model_config:commit", + operation_id=request.operation_id, + request_payload=request_payload, + ) + replay = self._replay_operation( + request.admin_grant, "model_config:commit", request_payload + ) + if replay is not None: + return CommitModelConfigResult(**replay) + with self._connect() as connection: + candidate = connection.execute( + "SELECT * FROM config_candidates WHERE proposal_hash=?", + (request.proposal_hash,), + ).fetchone() + placeholders = ",".join("?" for _ in request.evidence_ids) + evidence = ( + connection.execute( + f"""SELECT * FROM capability_evidence_ops + WHERE evidence_id IN ({placeholders})""", + tuple(request.evidence_ids), + ).fetchall() + if request.evidence_ids + else [] + ) + if ( + candidate is None + or int(candidate["expires_at"]) < now_ms() + or int(candidate["expected_revision"]) != request.expected_revision + or canonical_json_v1(request.payload).decode() + != candidate["canonical_payload"] + ): + raise EvoRuntimeError("CANDIDATE_EXPIRED") + expected_proposal_hash = proposal_hash( + request.payload, + target_revision=request.expected_revision + 1, + identity_key=self.identity_key_ring.derive_current( + "ai4sci/proposal-hash/v3" + )[1], + ) + if expected_proposal_hash != request.proposal_hash: + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + route_rows = json.loads(candidate["routes_json"]) + candidate_payload = dict(request.payload) + candidate_payload["config_revision"] = request.expected_revision + 1 + candidate_payload["capability_evidence"] = [] + candidate_config = EvoModelConfig.parse( + candidate_payload, require_evidence=False + ) + recomputed_routes = [ + asdict(item) for item in self._build_proposals(candidate_config) + ] + if canonical_json_v1(recomputed_routes) != canonical_json_v1(route_rows): + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + if len(evidence) != len(set(request.evidence_ids)): + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + evidence_by_route: dict[str, dict[str, sqlite3.Row]] = {} + for item in evidence: + if ( + item["proposal_hash"] != request.proposal_hash + or int(item["expires_at"]) < now_ms() + or item["status"] != "supported" + or item["fixture_digest"] != _FIXTURE_DIGEST + or item["config_identity_key_id"] != candidate["config_identity_key_id"] + ): + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + evidence_by_route.setdefault(str(item["route_semantics_hash"]), {})[ + str(item["probe_kind"]) + ] = item + final_evidence = [] + for route_row in route_rows: + route_evidence = evidence_by_route.get( + route_row["route_semantics_hash"], {} + ) + if not set(route_row["required_probe_kinds"]) <= set(route_evidence): + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + route = _find_route(candidate_config, route_row["route"]) + expected_adapter = adapter_revision( + candidate_config.providers[route.provider].protocol, + route.api_mode, + ) + expected_secrets = canonical_json_v1( + self._secret_fingerprints(candidate_config, route) + ).decode() + if any( + item["adapter_revision"] != expected_adapter + or item["endpoint_fingerprint"] != route_row["endpoint_fingerprint"] + or item["secret_fingerprints_json"] != expected_secrets + for item in route_evidence.values() + ): + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + representative = next(iter(route_evidence.values())) + final_evidence.append( + { + "route": route_row["route"], + "connectivity": "supported", + "tool_capability": ( + "supported" + if "tool_protocol" in route_row["required_probe_kinds"] + else "unknown" + ), + "probe": { + "route_semantics_hash": route_row["route_semantics_hash"], + "endpoint_fingerprint": route_row["endpoint_fingerprint"], + "config_identity_key_id": candidate["config_identity_key_id"], + "adapter_revision": representative["adapter_revision"], + "fixture_digest": representative["fixture_digest"], + "verified_at": datetime.now(UTC).isoformat(), + }, + } + ) + final_payload = dict(request.payload) + final_payload["capability_evidence"] = final_evidence + revision = self.store.commit_validated( + final_payload, + expected_revision=request.expected_revision, + operation_id=request.operation_id, + ).config_revision + committed_at = now_ms() + redacted_diff = { + "target_revision": revision, + "providers": sorted((request.payload.get("providers") or {}).keys()), + "route_selectors": sorted( + (request.payload.get("route_selectors") or {}).keys() + ), + } + with self._connect() as connection: + connection.execute( + """INSERT INTO config_commit_audit + (operation_id, subject_id, expected_revision, actual_revision, + proposal_hash, evidence_ids_json, redacted_diff_json, committed_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(operation_id) DO NOTHING""", + ( + request.operation_id, + request.admin_grant.subject_id, + request.expected_revision, + revision, + request.proposal_hash, + json.dumps(sorted(request.evidence_ids)), + json.dumps(redacted_diff, sort_keys=True), + committed_at, + ), + ) + result = CommitModelConfigResult( + request.operation_id, + revision, + committed_at, + redacted_diff, + ) + self._record_operation( + request.admin_grant, + "model_config:commit", + request_payload, + asdict(result), + ) + return result + + def _build_proposals(self, config: EvoModelConfig) -> list[ConcreteRouteProposal]: + semantics_key = self.identity_key_ring.derive_current(_ROUTE_SEMANTICS_INFO)[1] + endpoint_key = self.identity_key_ring.derive_current( + _ENDPOINT_FINGERPRINT_INFO + )[1] + routes_by_key = { + route.key(): route + for selector_id in config.route_selectors + for route in config.concrete_routes(selector_id) + } + proposals: list[ConcreteRouteProposal] = [] + for route_key, required_kinds in sorted( + config.required_concrete_routes().items() + ): + route = routes_by_key[route_key] + self._secret_fingerprints(config, route) + proposals.append( + ConcreteRouteProposal( + route={ + "provider": route.provider, + "endpoint": route.endpoint, + "model": route.model, + "api_mode": route.api_mode, + "tool_call_transport": route.tool_call_transport, + }, + route_semantics_hash=route_semantics_hash( + config, route, semantics_key + ), + endpoint_fingerprint=endpoint_fingerprint( + config, route, endpoint_key + ), + required_probe_kinds=required_kinds, + ) + ) + return proposals + + def _secret_fingerprints( + self, config: EvoModelConfig, route: RouteRef + ) -> dict[str, str]: + endpoint = config.providers[route.provider].endpoints[route.endpoint] + references = (endpoint.auth, *endpoint.header_refs.values()) + return { + reference.ref: resolve_secret( + reference, secret_resolver=self.secret_resolver + ).runtime_fingerprint + for reference in references + } + + def _replay_operation( + self, + grant: Any, + action: str, + request_payload: Mapping[str, Any], + ) -> dict[str, Any] | None: + digest = sha256_id(request_payload) + with self._connect() as connection: + row = connection.execute( + """SELECT request_digest, result_json + FROM admin_operation_journal + WHERE subject_id=? AND action=? AND operation_id=?""", + (grant.subject_id, action, grant.operation_id), + ).fetchone() + if row is None: + return None + if str(row["request_digest"]) != digest: + raise EvoRuntimeError("IDEMPOTENCY_CONFLICT") + return dict(json.loads(row["result_json"])) + + def _record_operation( + self, + grant: Any, + action: str, + request_payload: Mapping[str, Any], + result: Mapping[str, Any], + ) -> None: + digest = sha256_id(request_payload) + result_json = canonical_json_v1(result).decode("utf-8") + with self._connect() as connection: + existing = connection.execute( + """SELECT request_digest, result_json + FROM admin_operation_journal + WHERE subject_id=? AND action=? AND operation_id=?""", + (grant.subject_id, action, grant.operation_id), + ).fetchone() + if existing is not None: + if ( + str(existing["request_digest"]) != digest + or str(existing["result_json"]) != result_json + ): + raise EvoRuntimeError("IDEMPOTENCY_CONFLICT") + return + connection.execute( + """INSERT INTO admin_operation_journal + (subject_id, action, operation_id, request_digest, + result_json, committed_at) + VALUES (?, ?, ?, ?, ?, ?)""", + ( + grant.subject_id, + action, + grant.operation_id, + digest, + result_json, + now_ms(), + ), + ) + + def _authorize( + self, + grant: Any, + *, + action: str, + operation_id: str, + request_payload: Mapping[str, Any], + ) -> None: + self.grant_authority.require_admin(grant) + if ( + grant.action != action + or grant.operation_id != operation_id + or grant.request_digest != sha256_id(request_payload) + ): + raise EvoRuntimeError("ADMIN_CONFIG_FORBIDDEN") + + async def _default_probe( + self, config: EvoModelConfig, route: RouteRef, probe_kind: str + ) -> bool: + from .models import get_chat_model + + provider = config.providers[route.provider] + endpoint = provider.endpoints[route.endpoint] + model = provider.models[route.model] + auth = resolve_secret(endpoint.auth, secret_resolver=self.secret_resolver) + if config.schema_version == 3: + registration = get_adapter_registry().get( + provider.adapter_id, provider.adapter_revision + ) + params = dict( + registration.compile_runtime_parameters( + route.api_mode, model.params, min(128, model.max_output_tokens) + ) + ) + params.update( + { + "api_key": auth.value, + "base_url": endpoint.base_url, + "max_retries": 0, + "streaming": False, + } + ) + runtime_provider = registration.runtime_provider + else: + params = { + **provider.params, + **endpoint.params, + **model.params, + "api_key": auth.value, + "base_url": endpoint.base_url, + "max_retries": 0, + "max_tokens": 128, + "streaming": False, + "disable_streaming": "tool_calling", + "use_responses_api": route.api_mode == "responses", + } + if model.reasoning_mode == "boolean": + params = _merge_params(params, model.reasoning_enabled_params) + runtime_provider = provider.protocol + headers = dict(endpoint.headers) + for name, reference in endpoint.header_refs.items(): + headers[name] = resolve_secret( + reference, secret_resolver=self.secret_resolver + ).value + if headers: + params["default_headers"] = headers + if ( + config.schema_version == 3 + and provider.adapter_id == "google-gemini" + and route.api_mode == "interactions" + ): + from .gemini_interactions import create_gemini_interactions_model + + client = create_gemini_interactions_model( + model=model.model_id, + provider="google_interactions", + **params, + ) + else: + client = get_chat_model( + model=model.model_id, + provider=runtime_provider, + **params, + ) + if probe_kind == "connectivity": + await client.ainvoke("Respond with exactly OK.") + return True + tool_schema = { + "type": "function", + "function": { + "name": "probe_tool", + "description": "Protocol probe only", + "parameters": { + "type": "object", + "properties": {"value": {"type": "string"}}, + "required": ["value"], + }, + }, + } + response = await client.bind_tools([tool_schema]).ainvoke( + "You must call probe_tool exactly once with value set to ok. " + "Do not answer in text." + ) + calls = getattr(response, "tool_calls", None) or [] + return bool( + calls + and calls[0].get("name") == "probe_tool" + and (calls[0].get("args") or {}).get("value") == "ok" + ) + + +def _merge_params( + base: Mapping[str, Any], override: Mapping[str, Any] +) -> dict[str, Any]: + merged = dict(base) + for key, value in override.items(): + if isinstance(value, Mapping) and isinstance(merged.get(key), Mapping): + merged[key] = _merge_params(merged[key], value) + else: + merged[key] = value + return merged + + +def _find_route(config: EvoModelConfig, projection: Mapping[str, str]) -> RouteRef: + matches = { + route.key(): route + for selector_id in config.route_selectors + for route in config.concrete_routes(selector_id) + if route.provider == projection.get("provider") + and route.endpoint == projection.get("endpoint") + and route.model == projection.get("model") + and route.api_mode == projection.get("api_mode") + and route.tool_call_transport == projection.get("tool_call_transport") + } + if len(matches) != 1: + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + return next(iter(matches.values())) + + +def _redact_config(value: Any, key: str = "") -> Any: + if isinstance(value, Mapping): + return { + str(item_key): _redact_config(item, str(item_key)) + for item_key, item in value.items() + } + if isinstance(value, list | tuple): + return [_redact_config(item, key) for item in value] + # Config stores references and revisions, never secret values. Keeping the + # reference makes the redacted document safely round-trippable in the admin + # editor without exposing the resolved credential. + return value + + +def _unconfigured_template(identity_key_id: str) -> dict[str, Any]: + return { + "schema_version": 2, + "config_revision": 0, + "config_identity_key_id": identity_key_id, + "runtime_defaults": {"max_retries": 0}, + "purpose_defaults": {}, + "providers": {}, + "endpoint_pools": {}, + "route_health": { + "failure_threshold": 3, + "cooldown_seconds": 30, + "half_open_max_inflight": 1, + "counted_error_codes": [], + "open_immediately_error_codes": [], + }, + "route_selectors": {}, + "purpose_routes": {}, + "purpose_call_limits": {}, + "web_runtime": {}, + "tool_protocol_fallbacks": [], + } + + +def _unconfigured_v3_template(identity_key_id: str) -> dict[str, Any]: + return { + "schema_version": 3, + "config_revision": 1, + "config_identity_key_id": identity_key_id, + "runtime_defaults": {}, + "providers": [], + "aliases": [], + "purpose_defaults": { + "main_agent": {}, + "tool_selector": {}, + "deepagents_summarizer": {}, + "title": {}, + }, + "purpose_routes": { + "main_agent": {"default_alias": ""}, + "tool_selector": "inherit_main", + "deepagents_summarizer": "inherit_main", + "title": {"default_alias": ""}, + }, + "purpose_call_limits": { + "main_agent": {"max_output_tokens": 8192, "max_attempts_per_run": 2}, + "tool_selector": {"max_output_tokens": 4096, "max_attempts_per_run": 2}, + "deepagents_summarizer": { + "max_output_tokens": 4096, + "max_attempts_per_run": 2, + }, + "title": {"max_output_tokens": 256, "max_attempts_per_run": 1}, + }, + "health_policy": { + "provider_connection": {}, + "model_route": {}, + }, + "web_runtime": {}, + "capability_evidence": [], + } diff --git a/EvoScientist/llm/configuration/__init__.py b/EvoScientist/llm/configuration/__init__.py new file mode 100644 index 0000000..d2651d2 --- /dev/null +++ b/EvoScientist/llm/configuration/__init__.py @@ -0,0 +1,24 @@ +"""Canonical provider and model configuration contracts. + +Parsing and persistence remain in :mod:`EvoScientist.llm.model_config` for +backward compatibility. New code should import the owned contracts from this +package so provider connection data and model capability data stay separate. +""" + +from .model import ModelConfig +from .provider import ( + EndpointConfig, + ProviderConfig, + ResolvedSecret, + SecretReference, + SecretResolver, +) + +__all__ = [ + "EndpointConfig", + "ModelConfig", + "ProviderConfig", + "ResolvedSecret", + "SecretReference", + "SecretResolver", +] diff --git a/EvoScientist/llm/configuration/model.py b/EvoScientist/llm/configuration/model.py new file mode 100644 index 0000000..2faf645 --- /dev/null +++ b/EvoScientist/llm/configuration/model.py @@ -0,0 +1,52 @@ +"""Model-owned configuration. + +This module contains model identity, capability, limit, access, billing, and +canonical parameter declarations. It deliberately contains no credentials, +endpoint URLs, SDK clients, or derived wire-transport fields. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass, field +from typing import Any + +from ..contracts import PricingQuote + + +@dataclass(frozen=True, slots=True) +class ModelConfig: + """Normalized configuration owned by one provider model profile.""" + + model_id: str + params: Mapping[str, Any] + supports_vision: bool + supports_reasoning: bool + allowed_reasoning_efforts: tuple[str, ...] + context_window: int + max_output_tokens: int + reasoning_mode: str + reasoning_enabled_params: Mapping[str, Any] + reasoning_disabled_params: Mapping[str, Any] + allowed_plans: tuple[str, ...] + allowed_roles: tuple[str, ...] + quote: PricingQuote + model_key: str = "" + display_name: str = "" + description: str = "" + tags: tuple[str, ...] = () + enabled: bool = True + version_policy: str = "rolling" + resolved_model_revision: str | None = None + reproducible: bool = False + capabilities: Mapping[str, bool] = field(default_factory=dict) + purpose_overrides: Mapping[str, Mapping[str, Any]] = field(default_factory=dict) + user_options: Mapping[str, Mapping[str, Any]] = field(default_factory=dict) + parameter_constraints: tuple[Mapping[str, Any], ...] = () + access: Mapping[str, Any] = field(default_factory=dict) + max_inflight_requests: int | None = None + descriptor_parameters: Mapping[str, Any] = field(default_factory=dict) + + @property + def billing_sku(self) -> str: + return self.quote.billing_sku diff --git a/EvoScientist/llm/configuration/provider.py b/EvoScientist/llm/configuration/provider.py new file mode 100644 index 0000000..1625378 --- /dev/null +++ b/EvoScientist/llm/configuration/provider.py @@ -0,0 +1,61 @@ +"""Provider-owned connection and adapter configuration. + +Provider configuration owns credentials, endpoints, headers, connection +defaults, and adapter identity. Model capabilities and model-level parameters +are represented by :class:`ModelConfig`, not duplicated here. +""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping +from dataclasses import dataclass, field +from typing import Any + +from .model import ModelConfig + + +@dataclass(frozen=True, slots=True) +class SecretReference: + ref: str + revision: int + + +@dataclass(frozen=True, slots=True) +class ResolvedSecret: + value: str + declared_revision: int + authoritative_version: str | None + runtime_fingerprint: str + + +SecretResolver = Callable[[SecretReference], ResolvedSecret] + + +@dataclass(frozen=True, slots=True) +class EndpointConfig: + """Provider connection endpoint; never contains model behavior.""" + + name: str + base_url: str + auth: SecretReference + headers: Mapping[str, str] + header_refs: Mapping[str, SecretReference] + params: Mapping[str, Any] + + +@dataclass(frozen=True, slots=True) +class ProviderConfig: + """Provider adapter and connection configuration with owned model profiles.""" + + key: str + protocol: str + params: Mapping[str, Any] + endpoints: Mapping[str, EndpointConfig] + models: Mapping[str, ModelConfig] + display_name: str = "" + adapter_id: str = "" + adapter_revision: str = "" + wire_protocol: str = "" + enabled: bool = True + connection_defaults: Mapping[str, int] = field(default_factory=dict) + implementation_fingerprint: str = "" diff --git a/EvoScientist/llm/contracts.py b/EvoScientist/llm/contracts.py new file mode 100644 index 0000000..7f55ed9 --- /dev/null +++ b/EvoScientist/llm/contracts.py @@ -0,0 +1,870 @@ +"""Public V3 contracts for the embedded EvoScientist Web model runtime. + +The host can provide identity, storage and durable event services through +these DTOs and protocols. Provider connection details never cross this +boundary. +""" + +from __future__ import annotations + +import time +import uuid +from collections.abc import AsyncIterator, Mapping, Sequence +from dataclasses import asdict, dataclass, field, fields +from typing import Any, Literal, Protocol, TypeVar + +from .crypto import ( + HmacKeyRing, + KeyMaterial, + canonical_json_v1, + hmac_id, + sign_contract, + verify_contract, +) + +CONTRACT_VERSION = 3 +ADMIN_CONTROL_VERSION = 2 +GATEWAY_ISSUER = "ai4sci-gateway" +EVO_ISSUER = "evoscientist-runtime" +GATEWAY_AUDIENCE = EVO_ISSUER +EVO_AUDIENCE = GATEWAY_ISSUER +MAX_CLOCK_SKEW_MS = 5_000 +MAX_CONTRACT_TTL_MS = 120_000 +MAX_ADMIN_TTL_MS = 60_000 + +_GATEWAY_GRANT_INFO = "ai4sci/gateway-to-evo-grant/v3" +_EVO_QUOTE_INFO = "ai4sci/evo-to-gateway-quote/v3" +_INPUT_DIGEST_INFO = "ai4sci/agent-input-digest/v3" +_PREPARED_SNAPSHOT_INFO = "ai4sci/prepared-snapshot-digest/v3" +_PREPARED_INPUT_INFO = "ai4sci/prepared-input-digest/v3" +_TOOL_REGISTRY_INFO = "ai4sci/tool-registry-snapshot/v3" +_ADMIN_CONFIG_INFO = "ai4sci/admin-config/v2" + + +class EvoRuntimeError(RuntimeError): + """A stable error whose code may be projected to the host.""" + + def __init__( + self, + code: str, + message: str | None = None, + *, + details: Sequence[Mapping[str, Any]] = (), + ) -> None: + super().__init__(message or code) + self.code = code + self.details = tuple(dict(item) for item in details) + + +def now_ms() -> int: + return time.time_ns() // 1_000_000 + + +def _unsigned(value: Any) -> dict[str, Any]: + payload = asdict(value) + payload.pop("signature", None) + return payload + + +@dataclass(frozen=True, slots=True) +class RoutePreparationGrant: + issuer: str + audience: str + grant_id: str + request_id: str + turn_id: str + thread_id: str + subject_id: str + requested_model_ref: str + plan: str + roles: tuple[str, ...] + requires_vision: bool + reasoning_effort: str + title_policy: Literal["disabled", "best_effort"] + gateway_input_digest: str + checkpoint_thread_id: str + checkpoint_snapshot_id: str + turn_fencing_token: int + issued_at: int + expires_at: int + key_id: str + signature: str + contract_version: int = CONTRACT_VERSION + + def unsigned_payload(self) -> dict[str, Any]: + return _unsigned(self) + + +@dataclass(frozen=True, slots=True) +class RouteIdentity: + config_revision: int + config_identity_key_id: str + purpose: str + route_selector_id: str + route_fingerprint: str + provider_id: str + endpoint_name: str + model_id: str + protocol: str + api_mode: str + tool_call_transport: str + route_semantics_hash: str + billing_sku: str + pricing_revision: str + quote_id: str + + @property + def route_key(self) -> str: + return ":".join( + ( + self.provider_id, + self.endpoint_name, + self.model_id, + self.api_mode, + self.tool_call_transport, + ) + ) + + +@dataclass(frozen=True, slots=True) +class PricingQuote: + billing_sku: str + pricing_revision: str + currency: str + unit_scale: int + input_microunits_per_million: int + cached_input_microunits_per_million: int + output_microunits_per_million: int + quote_id: str + multiplier: str = "1" + + @property + def cached_microunits_per_million(self) -> int: + return self.cached_input_microunits_per_million + + +@dataclass(frozen=True, slots=True) +class RouteCallBound: + route_identity: RouteIdentity + max_output_tokens: int + payload_input_hard_cap: int + billable_input_cap: int + protocol_margin_tokens: int + attempt_reserve_microunits: int + + +@dataclass(frozen=True, slots=True) +class PreparedRunQuote: + issuer: str + audience: str + preparation_id: str + request_id: str + turn_id: str + thread_id: str + subject_id: str + requested_model_ref: str + plan: str + roles: tuple[str, ...] + requires_vision: bool + reasoning_effort: str + title_policy: Literal["disabled", "best_effort"] + gateway_input_digest: str + prepared_snapshot_digest: str + prepared_input_digest: str + config_revision: int + catalog_revision: int + enabled_purposes: tuple[str, ...] + purpose_routes: Mapping[str, Mapping[str, Any]] + purpose_route_call_bounds: Mapping[str, tuple[RouteCallBound, ...]] + purpose_attempt_limits: Mapping[str, int] + total_max_attempts: int + pricing_quotes: Mapping[str, PricingQuote] + quote_ids: tuple[str, ...] + provider_run_reserve_microunits: int + checkpoint_snapshot_id: str + tool_registry_snapshot_id: str + turn_fencing_token: int + route_semantics_hashes: tuple[str, ...] + issued_at: int + expires_at: int + key_id: str + signature: str + contract_version: int = CONTRACT_VERSION + + def unsigned_payload(self) -> dict[str, Any]: + return _unsigned(self) + + +@dataclass(frozen=True, slots=True) +class AdmissionGrant: + issuer: str + audience: str + grant_id: str + preparation_id: str + request_id: str + turn_id: str + thread_id: str + subject_id: str + requested_model_ref: str + plan: str + roles: tuple[str, ...] + requires_vision: bool + reasoning_effort: str + title_policy: Literal["disabled", "best_effort"] + gateway_input_digest: str + prepared_snapshot_digest: str + prepared_input_digest: str + config_revision: int + catalog_revision: int + purpose_attempt_limits: Mapping[str, int] + total_max_attempts: int + checkpoint_snapshot_id: str + tool_registry_snapshot_id: str + turn_fencing_token: int + admission_snapshot_id: str + admission_id: str + hold_id: str + billing_fencing_token: int + provider_run_reserve_microunits: int + billing_policy_version: str + issued_at: int + expires_at: int + key_id: str + signature: str + contract_version: int = CONTRACT_VERSION + + def unsigned_payload(self) -> dict[str, Any]: + return _unsigned(self) + + +@dataclass(frozen=True, slots=True) +class VerifiedModelSubject: + issuer: str + audience: str + grant_id: str + subject_id: str + plan: str + roles: tuple[str, ...] + issued_at: int + expires_at: int + key_id: str + signature: str + contract_version: int = CONTRACT_VERSION + + def unsigned_payload(self) -> dict[str, Any]: + return _unsigned(self) + + +@dataclass(frozen=True, slots=True) +class AdminConfigGrant: + issuer: str + audience: str + grant_id: str + subject_id: str + action: str + operation_id: str + request_digest: str + issued_at: int + expires_at: int + key_id: str + signature: str + contract_version: int = CONTRACT_VERSION + + def unsigned_payload(self) -> dict[str, Any]: + return _unsigned(self) + + +@dataclass(frozen=True, slots=True) +class GetModelConfigRequest: + operation_id: str + admin_grant: AdminConfigGrant + + +@dataclass(frozen=True, slots=True) +class GetModelConfigResult: + operation_id: str + config_revision: int + redacted_config: Mapping[str, Any] + catalog_projection: Mapping[str, Any] + + +@dataclass(frozen=True, slots=True) +class ConcreteRouteProposal: + route: Mapping[str, str] + route_semantics_hash: str + endpoint_fingerprint: str + required_probe_kinds: tuple[str, ...] + + +@dataclass(frozen=True, slots=True) +class ValidateCandidateConfigRequest: + operation_id: str + expected_revision: int + payload: Mapping[str, Any] + admin_grant: AdminConfigGrant + + +@dataclass(frozen=True, slots=True) +class ValidateCandidateConfigResult: + operation_id: str + target_revision: int + config_identity_key_id: str + proposal_hash: str + concrete_routes: tuple[ConcreteRouteProposal, ...] + expires_at: int + + +@dataclass(frozen=True, slots=True) +class ProbeCandidateRouteRequest: + operation_id: str + proposal_hash: str + route_semantics_hash: str + probe_kind: Literal["connectivity", "tool_protocol"] + admin_grant: AdminConfigGrant + + +@dataclass(frozen=True, slots=True) +class ProbeCandidateRouteResult: + operation_id: str + evidence_id: str + route_semantics_hash: str + connectivity: str + tool_capability: str + adapter_revision: str + fixture_digest: str + expires_at: int + + +@dataclass(frozen=True, slots=True) +class CommitModelConfigRequest: + operation_id: str + expected_revision: int + proposal_hash: str + payload: Mapping[str, Any] + evidence_ids: tuple[str, ...] + admin_grant: AdminConfigGrant + + +@dataclass(frozen=True, slots=True) +class CommitModelConfigResult: + operation_id: str + config_revision: int + committed_at: int + redacted_diff: Mapping[str, Any] + + +@dataclass(frozen=True, slots=True) +class CreateProposalRequest: + operation_id: str + expected_active_revision: int + admin_grant: AdminConfigGrant + admin_control_version: int = ADMIN_CONTROL_VERSION + + +@dataclass(frozen=True, slots=True) +class UpdateProposalRequest: + operation_id: str + proposal_id: str + expected_state_version: int + expected_draft_etag: str + draft_payload: Mapping[str, Any] + admin_grant: AdminConfigGrant + admin_control_version: int = ADMIN_CONTROL_VERSION + + +@dataclass(frozen=True, slots=True) +class ValidateProposalRequest: + operation_id: str + proposal_id: str + expected_state_version: int + expected_draft_etag: str + admin_grant: AdminConfigGrant + admin_control_version: int = ADMIN_CONTROL_VERSION + + +@dataclass(frozen=True, slots=True) +class ProbeProposalRequest: + operation_id: str + proposal_id: str + validated_digest: str + route_semantics_hash: str + probe_kind: str + admin_grant: AdminConfigGrant + admin_control_version: int = ADMIN_CONTROL_VERSION + + +@dataclass(frozen=True, slots=True) +class CommitProposalRequest: + operation_id: str + proposal_id: str + expected_active_revision: int + expected_state_version: int + expected_draft_etag: str + validated_digest: str + evidence_ids: tuple[str, ...] + admin_grant: AdminConfigGrant + admin_control_version: int = ADMIN_CONTROL_VERSION + + +@dataclass(frozen=True, slots=True) +class CancelProposalRequest: + operation_id: str + proposal_id: str + expected_state_version: int + admin_grant: AdminConfigGrant + admin_control_version: int = ADMIN_CONTROL_VERSION + + +@dataclass(frozen=True, slots=True) +class RollbackConfigRequest: + operation_id: str + expected_active_revision: int + target_revision: int + admin_grant: AdminConfigGrant + admin_control_version: int = ADMIN_CONTROL_VERSION + + +@dataclass(frozen=True, slots=True) +class DiscoverProviderModelsRequest: + operation_id: str + proposal_id: str + provider_id: str + admin_grant: AdminConfigGrant + admin_control_version: int = ADMIN_CONTROL_VERSION + + +@dataclass(frozen=True, slots=True) +class DiscoverProviderModelsResult: + operation_id: str + proposal_id: str + provider_id: str + source: str + discovered_at: int + models: tuple[Mapping[str, str], ...] + admin_control_version: int = ADMIN_CONTROL_VERSION + + +@dataclass(frozen=True, slots=True) +class AdminProposalResult: + operation_id: str + active_revision: int + proposal_id: str + state: str + state_version: int + draft_etag: str + base_revision: int + target_revision: int + expires_at: int + validated_digest: str | None = None + routes: tuple[Mapping[str, Any], ...] = () + evidence_ids: tuple[str, ...] = () + committed_at: int | None = None + draft_payload: Mapping[str, Any] | None = None + migration_report: Mapping[str, Any] | None = None + admin_control_version: int = ADMIN_CONTROL_VERSION + + +@dataclass(frozen=True, slots=True) +class RollbackConfigResult: + operation_id: str + previous_active_revision: int + active_revision: int + target_revision: int + rolled_back_at: int + admin_control_version: int = ADMIN_CONTROL_VERSION + + +class AdmissionGrantVerifier(Protocol): + def require_preparation(self, grant: RoutePreparationGrant) -> None: ... + + def require_admission(self, grant: AdmissionGrant) -> None: ... + + def verify_subject(self, subject: VerifiedModelSubject) -> bool: ... + + +class AdminConfigGrantVerifier(Protocol): + def require_admin(self, grant: AdminConfigGrant) -> None: ... + + +_ContractT = TypeVar( + "_ContractT", + RoutePreparationGrant, + PreparedRunQuote, + AdmissionGrant, + VerifiedModelSubject, + AdminConfigGrant, +) + + +class HmacGrantAuthority: + """Purpose-separated V3 contract signer with current/previous key support.""" + + def __init__( + self, + secret: str | bytes, + key_id: str = "runtime-current", + *, + previous_secret: str | bytes | None = None, + previous_key_id: str | None = None, + ) -> None: + previous = None + if (previous_secret is None) != (previous_key_id is None): + raise ValueError("previous secret and key id must be configured together") + if previous_secret is not None and previous_key_id is not None: + previous = KeyMaterial.create(previous_key_id, previous_secret) + self.key_ring = HmacKeyRing(KeyMaterial.create(key_id, secret), previous) + + def sign_preparation(self, **kwargs: Any) -> RoutePreparationGrant: + return self._sign( + RoutePreparationGrant, _GATEWAY_GRANT_INFO, gateway=True, **kwargs + ) + + def sign_admission(self, **kwargs: Any) -> AdmissionGrant: + return self._sign(AdmissionGrant, _GATEWAY_GRANT_INFO, gateway=True, **kwargs) + + def sign_subject(self, **kwargs: Any) -> VerifiedModelSubject: + return self._sign( + VerifiedModelSubject, _GATEWAY_GRANT_INFO, gateway=True, **kwargs + ) + + def sign_admin(self, **kwargs: Any) -> AdminConfigGrant: + return self._sign(AdminConfigGrant, _ADMIN_CONFIG_INFO, gateway=True, **kwargs) + + def sign_quote(self, **kwargs: Any) -> PreparedRunQuote: + return self._sign(PreparedRunQuote, _EVO_QUOTE_INFO, gateway=False, **kwargs) + + def agent_input_digest(self, payload: Any) -> str: + _, key = self.key_ring.derive_current(_INPUT_DIGEST_INFO) + return hmac_id(key, payload) + + def prepared_input_digest(self, payload: Any) -> str: + _, key = self.key_ring.derive_current(_PREPARED_INPUT_INFO) + return hmac_id(key, payload) + + def prepared_snapshot_digest(self, payload: Any) -> str: + _, key = self.key_ring.derive_current(_PREPARED_SNAPSHOT_INFO) + return hmac_id(key, payload) + + def tool_registry_snapshot_id(self, payload: Any) -> str: + _, key = self.key_ring.derive_current(_TOOL_REGISTRY_INFO) + return hmac_id(key, payload) + + def require_preparation(self, grant: RoutePreparationGrant) -> None: + if not isinstance(grant, RoutePreparationGrant): + raise EvoRuntimeError("CONTRACT_SIGNATURE_INVALID") + self._require(grant, _GATEWAY_GRANT_INFO, gateway=True) + + def require_admission(self, grant: AdmissionGrant) -> None: + if not isinstance(grant, AdmissionGrant): + raise EvoRuntimeError("CONTRACT_SIGNATURE_INVALID") + self._require(grant, _GATEWAY_GRANT_INFO, gateway=True) + + def require_admin(self, grant: AdminConfigGrant) -> None: + if not isinstance(grant, AdminConfigGrant): + raise EvoRuntimeError("CONTRACT_SIGNATURE_INVALID") + self._require( + grant, _ADMIN_CONFIG_INFO, gateway=True, max_ttl_ms=MAX_ADMIN_TTL_MS + ) + + def require_quote(self, quote: PreparedRunQuote) -> None: + if not isinstance(quote, PreparedRunQuote): + raise EvoRuntimeError("CONTRACT_SIGNATURE_INVALID") + self._require(quote, _EVO_QUOTE_INFO, gateway=False) + + def verify_admission(self, grant: AdmissionGrant) -> bool: + try: + self.require_admission(grant) + except EvoRuntimeError: + return False + return True + + def verify_subject(self, subject: VerifiedModelSubject) -> bool: + try: + self._require(subject, _GATEWAY_GRANT_INFO, gateway=True) + except EvoRuntimeError: + return False + return True + + def verify_admin(self, grant: AdminConfigGrant) -> bool: + try: + self.require_admin(grant) + except EvoRuntimeError: + return False + return True + + def _sign( + self, + contract_class: type[_ContractT], + info: str, + *, + gateway: bool, + **kwargs: Any, + ) -> _ContractT: + issued_at = int(kwargs.pop("issued_at", now_ms())) + ttl_ms = int(kwargs.pop("ttl_ms", 60_000)) + key_id, key = self.key_ring.derive_current(info) + defaults = { + "contract_version": CONTRACT_VERSION, + "issuer": GATEWAY_ISSUER if gateway else EVO_ISSUER, + "audience": GATEWAY_AUDIENCE if gateway else EVO_AUDIENCE, + "issued_at": issued_at, + "expires_at": int(kwargs.pop("expires_at", issued_at + ttl_ms)), + "key_id": key_id, + } + field_names = {item.name for item in fields(contract_class)} + if "grant_id" in field_names: + defaults["grant_id"] = str(kwargs.pop("grant_id", uuid.uuid4())) + payload = {**defaults, **kwargs} + if "roles" in payload: + payload["roles"] = tuple(sorted({str(role) for role in payload["roles"]})) + unsigned = { + key_name: value + for key_name, value in payload.items() + if key_name != "signature" + } + signature = sign_contract(contract_class.__name__, unsigned, key) + return contract_class(signature=signature, **payload) + + def _require( + self, + contract: Any, + info: str, + *, + gateway: bool, + max_ttl_ms: int = MAX_CONTRACT_TTL_MS, + ) -> None: + if int(contract.contract_version) != CONTRACT_VERSION: + raise EvoRuntimeError("CONTRACT_VERSION_UNSUPPORTED") + expected_issuer = GATEWAY_ISSUER if gateway else EVO_ISSUER + expected_audience = GATEWAY_AUDIENCE if gateway else EVO_AUDIENCE + if contract.issuer != expected_issuer or contract.audience != expected_audience: + raise EvoRuntimeError("CONTRACT_SIGNATURE_INVALID") + try: + key = self.key_ring.derive(contract.key_id, info) + except KeyError as exc: + raise EvoRuntimeError("CONTRACT_KEY_UNKNOWN") from exc + if not verify_contract( + type(contract).__name__, + contract.unsigned_payload(), + contract.signature, + key, + ): + raise EvoRuntimeError("CONTRACT_SIGNATURE_INVALID") + current = now_ms() + issued_at = int(contract.issued_at) + expires_at = int(contract.expires_at) + if expires_at <= issued_at or expires_at - issued_at > max_ttl_ms: + raise EvoRuntimeError("CONTRACT_EXPIRED") + if ( + issued_at > current + MAX_CLOCK_SKEW_MS + or expires_at < current - MAX_CLOCK_SKEW_MS + ): + raise EvoRuntimeError("CONTRACT_EXPIRED") + + +@dataclass(frozen=True, slots=True) +class MediaDescriptor: + media_id: str + media_type: str + content_hash: str + version: str + byte_length: int + token_bound: int + locator: str + + +@dataclass(frozen=True, slots=True) +class AgentInputV3: + message: Any + checkpoint_thread_id: str + metadata: Mapping[str, Any] = field(default_factory=dict) + media: Sequence[MediaDescriptor | Mapping[str, Any] | str] = () + content_blocks: Sequence[Mapping[str, Any]] = () + + def projection(self) -> Mapping[str, Any]: + allowed_metadata = { + key: self.metadata[key] + for key in sorted(self.metadata) + if key + in { + "force_context_repair", + "public_thread_id", + "source", + "user_id", + "model_options", + } + } + media = [ + asdict(item) if isinstance(item, MediaDescriptor) else item + for item in self.media + ] + return { + "message": self.message, + "checkpoint_thread_id": self.checkpoint_thread_id, + "metadata": allowed_metadata, + "media": media, + "content_blocks": list(self.content_blocks), + } + + def canonical_bytes(self) -> bytes: + return canonical_json_v1(self.projection()) + + +AgentInput = AgentInputV3 + + +class RuntimeEventSink(Protocol): + async def commit(self, event: EvoRuntimeEvent) -> str: ... + + async def confirm(self, event_id: str, payload_digest: str) -> str: ... + + +@dataclass(frozen=True, slots=True) +class WebHostContext: + workspace_dir: str + memory_dir: str + workspace_backend: Any + checkpointer: Any + tool_selector_threshold: int | None = None + memory_max_inline_profile_chars: int | None = None + on_mcp_progress: Any = None + runtime_event_sink: RuntimeEventSink | None = None + tool_registry: Sequence[Any] = () + tool_registry_revision: str = "static" + tool_registry_provider: Any = None + + +@dataclass(frozen=True, slots=True) +class AgentExecutionProfile: + name: str + configurable_model_override: bool = True + subagents: bool = True + async_subagents: bool = True + memory_workers: bool = True + scheduler: bool = True + background_execution: bool = True + + @classmethod + def web_v3(cls) -> AgentExecutionProfile: + return cls( + name="web_v3", + configurable_model_override=False, + subagents=False, + async_subagents=False, + memory_workers=False, + scheduler=False, + background_execution=False, + ) + + @classmethod + def web_v1(cls) -> AgentExecutionProfile: + return cls.web_v3() + + +@dataclass(frozen=True, slots=True) +class AgentModelSet: + main_agent: Any + tool_selector: Any + deepagents_summarizer: Any + title: Any | None = None + main_fallbacks: tuple[Any, ...] = () + route_health: Any = None + capacity: Any = None + + +@dataclass(frozen=True, slots=True) +class ModelCatalogEntry: + alias: str + provider: str + supports_vision: bool + supports_reasoning: bool + allowed_reasoning_efforts: tuple[str, ...] + context_window: int + max_output_tokens: int + reasoning_mode: str + billing_sku: str + quote: PricingQuote + default_reasoning_effort: str = "" + health: Literal["closed", "open", "half_open"] = "closed" + display_name: str = "" + provider_display_name: str = "" + description: str = "" + version_policy: str = "rolling" + resolved_model_revision: str | None = None + reproducible: bool = False + capabilities: tuple[str, ...] = () + user_options: Mapping[str, Mapping[str, Any]] = field(default_factory=dict) + parameter_constraints: tuple[Mapping[str, Any], ...] = () + options_schema_hash: str = "" + + +@dataclass(frozen=True, slots=True) +class ModelCatalog: + catalog_revision: int + default_alias: str + entries: tuple[ModelCatalogEntry, ...] + + +@dataclass(frozen=True, slots=True) +class ModelAttemptEvent: + request_id: str + turn_id: str + run_id: str + preparation_id: str + admission_snapshot_id: str + admission_id: str + hold_id: str + prepared_snapshot_digest: str + prepared_input_digest: str + billing_fencing_token: int + turn_fencing_token: int + logical_call_id: str + attempt_id: str + purpose: Literal["main_agent", "tool_selector", "deepagents_summarizer", "title"] + attempt_index: int + identity: RouteIdentity + quote_id: str + billing_intent: Literal["user_charge", "platform_cost"] + outcome: Literal["rejected", "started", "succeeded", "failed", "usage_unconfirmed"] + provider_request_started: bool + provider_input_bound_tokens: int + provider_reserved_microunits: int + usage: Mapping[str, int | str | None] | None = None + usage_available: bool = False + error_code: str | None = None + health_state: Literal["closed", "open", "half_open"] = "closed" + fallback_from_attempt_id: str | None = None + timestamp: int = field(default_factory=now_ms) + + +@dataclass(frozen=True, slots=True) +class EvoRuntimeEvent: + event_id: str + runtime_instance_id: str + run_id: str + request_id: str + turn_id: str + sequence: int + kind: Literal["agent", "model_attempt", "title", "run"] + payload: Mapping[str, Any] + emitted_at: int = field(default_factory=now_ms) + schema_version: int = CONTRACT_VERSION + + +class EvoWebRun(Protocol): + @property + def run_id(self) -> str: ... + + async def stream( + self, after_sequence: int | None = None + ) -> AsyncIterator[EvoRuntimeEvent]: ... + + async def cancel(self, reason: str) -> str: ... + + +def event_payload(value: Any) -> dict[str, Any]: + if isinstance(value, ModelAttemptEvent): + return asdict(value) + if isinstance(value, Mapping): + return dict(value) + return {"value": str(value)} diff --git a/EvoScientist/llm/crypto.py b/EvoScientist/llm/crypto.py new file mode 100644 index 0000000..44b5086 --- /dev/null +++ b/EvoScientist/llm/crypto.py @@ -0,0 +1,159 @@ +"""Cryptographic primitives for the Evo Web runtime V3 contracts.""" + +from __future__ import annotations + +import hashlib +import hmac +import unicodedata +from collections.abc import Mapping, Sequence +from dataclasses import asdict, dataclass, is_dataclass +from typing import Any + +import rfc8785 + +_HKDF_SALT = b"ai4sci-evo-runtime-v3" +_MIN_SECRET_BYTES = 32 + + +class CanonicalJsonError(ValueError): + """Raised when a value cannot be represented by canonical JSON.""" + + +def _normalize(value: Any) -> Any: + if is_dataclass(value) and not isinstance(value, type): + return _normalize(asdict(value)) + if value is None or isinstance(value, bool | int | float): + return value + if isinstance(value, str): + return unicodedata.normalize("NFC", value) + if isinstance(value, Mapping): + normalized: dict[str, Any] = {} + for key, item in value.items(): + if not isinstance(key, str): + raise CanonicalJsonError("JSON object keys must be strings") + normalized_key = unicodedata.normalize("NFC", key) + if normalized_key in normalized: + raise CanonicalJsonError("duplicate JSON key after NFC normalization") + normalized[normalized_key] = _normalize(item) + return normalized + if isinstance(value, Sequence) and not isinstance( + value, bytes | bytearray | memoryview + ): + return [_normalize(item) for item in value] + raise CanonicalJsonError(f"unsupported canonical JSON type: {type(value).__name__}") + + +def canonical_json_v1(value: Any) -> bytes: + """Encode NFC-normalized data with RFC 8785 JSON canonicalization.""" + + try: + return rfc8785.dumps(_normalize(value)) + except (rfc8785.CanonicalizationError, rfc8785.FloatDomainError, TypeError) as exc: + raise CanonicalJsonError("value is not valid RFC 8785 JSON") from exc + + +def sha256_id(value: Any) -> str: + return f"sha256:{hashlib.sha256(canonical_json_v1(value)).hexdigest()}" + + +def hkdf_sha256(root_key: bytes, *, info: str, length: int = 32) -> bytes: + """RFC 5869 HKDF-SHA256 with the protocol's fixed salt.""" + + if length < 1 or length > 255 * hashlib.sha256().digest_size: + raise ValueError("invalid HKDF output length") + prk = hmac.new(_HKDF_SALT, root_key, hashlib.sha256).digest() + output = bytearray() + previous = b"" + counter = 1 + info_bytes = info.encode("utf-8") + while len(output) < length: + previous = hmac.new( + prk, + previous + info_bytes + bytes((counter,)), + hashlib.sha256, + ).digest() + output.extend(previous) + counter += 1 + return bytes(output[:length]) + + +def hmac_id(key: bytes, value: Any) -> str: + digest = hmac.new(key, canonical_json_v1(value), hashlib.sha256).hexdigest() + return f"hmac-sha256:{digest}" + + +@dataclass(frozen=True, slots=True) +class KeyMaterial: + key_id: str + secret: bytes + + @classmethod + def create(cls, key_id: str, secret: str | bytes) -> KeyMaterial: + normalized_id = str(key_id or "").strip() + encoded = secret.encode("utf-8") if isinstance(secret, str) else bytes(secret) + if not normalized_id: + raise ValueError("key id is required") + if len(encoded) < _MIN_SECRET_BYTES: + raise ValueError("signing secret must contain at least 32 bytes") + return cls(normalized_id, encoded) + + +class HmacKeyRing: + """Versioned root keys with purpose-separated HKDF children. + + ``previous`` remains as a compatibility property for callers that still + expose a two-key deployment contract. New persistence code uses + ``retained``/``key_ids`` so immutable revisions can outlive one rotation. + """ + + def __init__( + self, + current: KeyMaterial, + previous: KeyMaterial | None = None, + *, + retained: Sequence[KeyMaterial] = (), + ) -> None: + if previous is not None and previous.key_id == current.key_id: + raise ValueError("current and previous key ids must differ") + self.current = current + self.previous = previous + self._roots = {current.key_id: current.secret} + if previous is not None: + self._roots[previous.key_id] = previous.secret + for item in retained: + existing = self._roots.get(item.key_id) + if existing is not None and existing != item.secret: + raise ValueError("duplicate signing key id has different material") + self._roots[item.key_id] = item.secret + + @property + def key_ids(self) -> frozenset[str]: + return frozenset(self._roots) + + def contains(self, key_id: str) -> bool: + return key_id in self._roots + + def derive_current(self, info: str) -> tuple[str, bytes]: + return self.current.key_id, hkdf_sha256(self.current.secret, info=info) + + def derive(self, key_id: str, info: str) -> bytes: + try: + root = self._roots[key_id] + except KeyError as exc: + raise KeyError("unknown signing key") from exc + return hkdf_sha256(root, info=info) + + +def sign_contract(contract_type: str, payload: Mapping[str, Any], key: bytes) -> str: + message = contract_type.encode("utf-8") + b"\0" + canonical_json_v1(payload) + return hmac.new(key, message, hashlib.sha256).hexdigest() + + +def verify_contract( + contract_type: str, + payload: Mapping[str, Any], + signature: str, + key: bytes, +) -> bool: + expected = sign_contract(contract_type, payload, key) + return hmac.compare_digest(expected, str(signature)) diff --git a/EvoScientist/llm/errors.py b/EvoScientist/llm/errors.py index d953d79..04b7be5 100644 --- a/EvoScientist/llm/errors.py +++ b/EvoScientist/llm/errors.py @@ -60,10 +60,24 @@ class AgentControlError(Exception): "retryable": self.retryable, } + @classmethod + def model_construct(cls, **payload: Any) -> AgentControlError: + """Rebuild the allowlisted checkpoint form without trusting extra fields.""" + return cls( + str(payload.get("code") or "MODEL_REQUEST_REJECTED"), + str(payload.get("message") or "Model request rejected."), + status_code=int(payload.get("status_code") or 403), + retryable=bool(payload.get("retryable", False)), + ) + class ModelToolProtocolError(AgentControlError): """A completed model response contained an invalid tool-call protocol.""" + # Unlike authorization and admission control errors, a malformed model + # response is safe to retry before the agent executes any tool. + non_fallbackable = False + def __init__( self, reason: str, @@ -82,7 +96,7 @@ class ModelToolProtocolError(AgentControlError): "MODEL_TOOL_PROTOCOL_INVALID", "The model returned an invalid structured tool call.", status_code=502, - retryable=False, + retryable=True, ) self.reason = reason self.provider = provider @@ -123,6 +137,56 @@ class ModelToolProtocolError(AgentControlError): payload[key] = value return payload + @classmethod + def model_construct(cls, **payload: Any) -> ModelToolProtocolError: + """Rebuild only the public, redacted checkpoint projection.""" + + def optional_text(name: str) -> str | None: + value = payload.get(name) + return str(value) if value is not None else None + + generation = payload.get("config_generation") + return cls( + str(payload.get("reason") or "invalid_tool_protocol"), + provider=optional_text("provider"), + model=optional_text("model"), + route_key=optional_text("route_key"), + config_generation=(int(generation) if generation is not None else None), + api_mode=optional_text("api_mode"), + endpoint=optional_text("endpoint"), + tool_call_transport=optional_text("tool_call_transport"), + call_id=optional_text("call_id"), + ) + + +class ModelProviderResponseError(AgentControlError): + """A completed provider response had no final text or tool call.""" + + non_fallbackable = False + + def __init__(self, reason: str = "empty_assistant_response") -> None: + super().__init__( + "MODEL_PROVIDER_RESPONSE_INVALID", + "The model returned no final text or structured tool call.", + status_code=502, + retryable=True, + ) + self.reason = reason + self.fallbackable = True + self.recoverable = True + + def model_dump(self) -> dict[str, Any]: + return { + **super().model_dump(), + "reason": self.reason, + "fallbackable": self.fallbackable, + "recoverable": self.recoverable, + } + + @classmethod + def model_construct(cls, **payload: Any) -> ModelProviderResponseError: + return cls(str(payload.get("reason") or "empty_assistant_response")) + class ProviderStreamError(Exception): """Envelope-shaped wrapper for a provider SDK exception raised @@ -194,6 +258,27 @@ class ProviderStreamError(Exception): """ return self.as_envelope() + @classmethod + def model_construct(cls, **payload: Any) -> ProviderStreamError: + """Rebuild the redacted provider envelope stored in a checkpoint.""" + + def optional_text(name: str) -> str | None: + value = payload.get(name) + return str(value) if value is not None else None + + status = payload.get("status_code") + return cls( + provider=str(payload.get("provider") or "unknown"), + class_qualname=str( + payload.get("class") or payload.get("error") or "ProviderError" + ), + message=str(payload.get("message") or "Provider request failed."), + status_code=int(status) if status is not None else None, + code=optional_text("code"), + err_type=optional_text("type"), + request_id=optional_text("request_id"), + ) + # --------------------------------------------------------------------------- # API-key redaction — env-driven, prefix-only diff --git a/EvoScientist/llm/gateway_proxy.py b/EvoScientist/llm/gateway_proxy.py new file mode 100644 index 0000000..c0cf3ae --- /dev/null +++ b/EvoScientist/llm/gateway_proxy.py @@ -0,0 +1,93 @@ +"""Secretless ChatModel proxy for Ai4Sci Graph-native runs.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from typing import Any + +import httpx +from langchain_core.language_models.chat_models import BaseChatModel +from langchain_core.messages import BaseMessage, messages_from_dict, messages_to_dict +from langchain_core.outputs import ChatGeneration, ChatResult +from langchain_core.tools import BaseTool +from langchain_core.utils.function_calling import convert_to_openai_tool +from pydantic import Field + + +class GatewayProxyChatModel(BaseChatModel): + gateway_url: str + run_id: str + envelope_signature: str + provider_id: str = "" + model_id: str = "" + bound_tools: list[dict[str, Any]] = Field(default_factory=list) + bound_tool_choice: Any = None + + @property + def _llm_type(self) -> str: + return "ai4sci-gateway-proxy" + + @property + def _identifying_params(self) -> dict[str, Any]: + return {"provider_id": self.provider_id, "model_id": self.model_id} + + def bind_tools( + self, + tools: Sequence[dict[str, Any] | type | BaseTool], + *, + tool_choice: str | dict[str, Any] | bool | None = None, + **kwargs: Any, + ): + del kwargs + serialized = [convert_to_openai_tool(tool) for tool in tools] + return self.model_copy( + update={"bound_tools": serialized, "bound_tool_choice": tool_choice} + ) + + def _generate(self, *args: Any, **kwargs: Any) -> ChatResult: + del args, kwargs + raise RuntimeError("AI4SCI_RECOVERABLE_RUN_REQUIRES_ASYNC_MODEL_PATH") + + async def _agenerate( + self, + messages: list[BaseMessage], + stop: list[str] | None = None, + run_manager: Any = None, + **kwargs: Any, + ) -> ChatResult: + del stop, kwargs + attempt_id = str(getattr(run_manager, "run_id", None) or self.run_id) + async with httpx.AsyncClient(timeout=httpx.Timeout(660.0, connect=5.0)) as client: + response = await client.post( + f"{self.gateway_url.rstrip('/')}/api/internal/recoverable-runs/model/invoke", + json={ + "run_id": self.run_id, + "attempt_id": attempt_id, + "envelope_signature": self.envelope_signature, + "messages": messages_to_dict(messages), + "tools": self.bound_tools, + "tool_choice": self.bound_tool_choice, + }, + ) + response.raise_for_status() + value = response.json() + parsed = messages_from_dict([value["message"]]) + if len(parsed) != 1: + raise RuntimeError("AI4SCI_MODEL_PROXY_RESPONSE_INVALID") + return ChatResult(generations=[ChatGeneration(message=parsed[0])]) + + +def proxy_from_config( + value: Mapping[str, Any], *, provider_id: str = "", model_id: str = "" +) -> GatewayProxyChatModel: + required = { + name: str(value.get(name) or "") + for name in ("gateway_url", "run_id", "envelope_signature") + } + if not all(required.values()): + raise RuntimeError("AI4SCI_MODEL_PROXY_CONFIG_INVALID") + return GatewayProxyChatModel( + **required, + provider_id=provider_id, + model_id=model_id, + ) diff --git a/EvoScientist/llm/gemini_interactions.py b/EvoScientist/llm/gemini_interactions.py new file mode 100644 index 0000000..7ed85c1 --- /dev/null +++ b/EvoScientist/llm/gemini_interactions.py @@ -0,0 +1,379 @@ +"""LangChain chat-model bridge for the stateless Gemini Interactions API.""" + +from __future__ import annotations + +import json +from collections.abc import AsyncIterator, Mapping, Sequence +from typing import Any + +from langchain_core.language_models.chat_models import BaseChatModel +from langchain_core.messages import ( + AIMessage, + AIMessageChunk, + BaseMessage, + HumanMessage, + SystemMessage, + ToolMessage, +) +from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult +from langchain_core.tools import BaseTool +from langchain_core.utils.function_calling import convert_to_openai_tool +from pydantic import Field, SecretStr + + +class GeminiInteractionsChatModel(BaseChatModel): + """Minimal native bridge that preserves signed Provider content blocks.""" + + model_name: str + api_key: SecretStr + base_url: str = "https://generativelanguage.googleapis.com" + max_output_tokens: int = 8192 + temperature: float | None = None + top_p: float | None = None + thinking: bool | None = None + store: bool = False + bound_tools: tuple[dict[str, Any], ...] = Field(default_factory=tuple) + + @property + def _llm_type(self) -> str: + return "google-gemini-interactions" + + @property + def _identifying_params(self) -> dict[str, Any]: + return {"model_name": self.model_name, "api_mode": "interactions"} + + def bind_tools( + self, + tools: Sequence[dict[str, Any] | type | BaseTool | Any], + *, + tool_choice: str | None = None, + **kwargs: Any, + ) -> Any: + _ = tool_choice, kwargs + compiled = [] + for tool in tools: + value = convert_to_openai_tool(tool) + function = value.get("function", value) + compiled.append( + { + "type": "function", + "name": function["name"], + "description": function.get("description", ""), + "parameters": function.get( + "parameters", {"type": "object", "properties": {}} + ), + } + ) + return self.model_copy(update={"bound_tools": tuple(compiled)}) + + def _client(self) -> Any: + from google import genai + from google.genai import types + + return genai.Client( + api_key=self.api_key.get_secret_value(), + http_options=types.HttpOptions(base_url=self.base_url.rstrip("/")), + ) + + def _request(self, messages: Sequence[BaseMessage]) -> dict[str, Any]: + turns, system_instruction = _compile_messages(messages) + generation_config: dict[str, Any] = { + "max_output_tokens": self.max_output_tokens, + } + if self.temperature is not None: + generation_config["temperature"] = self.temperature + if self.top_p is not None: + generation_config["top_p"] = self.top_p + if self.thinking is not None: + generation_config["thinking_level"] = "high" if self.thinking else "minimal" + generation_config["thinking_summaries"] = ( + "auto" if self.thinking else "none" + ) + return { + "model": self.model_name, + "input": turns, + "system_instruction": system_instruction or "", + "generation_config": generation_config, + "tools": list(self.bound_tools), + "store": False, + "stream": False, + } + + def _generate( + self, + messages: list[BaseMessage], + stop: list[str] | None = None, + **kwargs: Any, + ) -> ChatResult: + request = self._request(messages) + if stop: + request["generation_config"]["stop_sequences"] = stop + response = self._client().interactions.create(**request) + return _chat_result(response) + + async def _agenerate( + self, + messages: list[BaseMessage], + stop: list[str] | None = None, + **kwargs: Any, + ) -> ChatResult: + request = self._request(messages) + if stop: + request["generation_config"]["stop_sequences"] = stop + response = await self._client().aio.interactions.create(**request) + return _chat_result(response) + + async def _astream( + self, + messages: list[BaseMessage], + stop: list[str] | None = None, + **kwargs: Any, + ) -> AsyncIterator[ChatGenerationChunk]: + request = self._request(messages) + request["stream"] = True + if stop: + request["generation_config"]["stop_sequences"] = stop + stream = await self._client().aio.interactions.create(**request) + blocks: dict[int, dict[str, Any]] = {} + async for event in stream: + payload = _dump(event) + event_type = payload.get("event_type") + if event_type == "content.start": + blocks[int(payload["index"])] = dict(payload.get("content") or {}) + continue + if event_type == "content.delta": + index = int(payload["index"]) + delta = dict(payload.get("delta") or {}) + block = blocks.setdefault(index, {}) + _merge_stream_delta(block, delta) + if delta.get("type") == "text" and delta.get("text"): + yield ChatGenerationChunk( + message=AIMessageChunk(content=str(delta["text"])) + ) + continue + if event_type == "content.stop": + block = blocks.get(int(payload["index"]), {}) + if block.get("type") == "function_call": + yield ChatGenerationChunk( + message=AIMessageChunk( + content="", + tool_call_chunks=[ + { + "id": str(block.get("id") or ""), + "name": str(block.get("name") or ""), + "args": json.dumps( + block.get("arguments") or {}, + separators=(",", ":"), + ), + "index": int(payload["index"]), + "type": "tool_call_chunk", + } + ], + ) + ) + continue + if event_type == "error": + raise RuntimeError("MODEL_PROVIDER_PROTOCOL_ERROR") + if event_type == "interaction.complete": + interaction = payload.get("interaction") or {} + ordered_blocks = [blocks[index] for index in sorted(blocks)] + yield ChatGenerationChunk( + message=AIMessageChunk( + content="", + additional_kwargs={ + "gemini_interaction_content": ordered_blocks + }, + usage_metadata=_usage_metadata( + interaction.get("usage"), + provider_request_id=interaction.get("id"), + ), + response_metadata={ + "model_name": str( + (interaction.get("model") or {}).get("id") or "" + ), + "finish_reason": str( + interaction.get("status") or "unknown" + ), + }, + ) + ) + + +def _message_content(message: BaseMessage) -> list[dict[str, Any]]: + if isinstance(message, AIMessage): + preserved = message.additional_kwargs.get("gemini_interaction_content") + if isinstance(preserved, list): + return [dict(item) for item in preserved] + content = message.content + if isinstance(content, str): + blocks: list[dict[str, Any]] = [{"type": "text", "text": content}] + elif isinstance(content, list): + blocks = [] + for item in content: + if isinstance(item, str): + blocks.append({"type": "text", "text": item}) + elif isinstance(item, dict) and item.get("type") in { + "text", + "thought", + "function_call", + "function_result", + "image", + "audio", + "video", + "document", + }: + blocks.append(dict(item)) + else: + raise ValueError("MODEL_CONTENT_BLOCK_UNSUPPORTED") + else: + blocks = [{"type": "text", "text": str(content)}] + if isinstance(message, AIMessage): + for call in message.tool_calls: + blocks.append( + { + "type": "function_call", + "id": str(call["id"]), + "name": str(call["name"]), + "arguments": dict(call.get("args") or {}), + } + ) + return blocks + + +def _compile_messages( + messages: Sequence[BaseMessage], +) -> tuple[list[dict[str, Any]], str]: + turns: list[dict[str, Any]] = [] + system_parts: list[str] = [] + for message in messages: + if isinstance(message, SystemMessage): + system_parts.append(str(message.content)) + continue + if isinstance(message, ToolMessage): + turns.append( + { + "role": "user", + "content": [ + { + "type": "function_result", + "call_id": str(message.tool_call_id), + "name": str(getattr(message, "name", "") or ""), + "result": message.content, + } + ], + } + ) + continue + role = "model" if isinstance(message, AIMessage) else "user" + if not isinstance(message, (AIMessage, HumanMessage)): + role = "user" + turns.append({"role": role, "content": _message_content(message)}) + return turns, "\n\n".join(system_parts) + + +def _chat_result(response: Any) -> ChatResult: + blocks = [ + item.model_dump(mode="json", by_alias=True, exclude_none=True) + if hasattr(item, "model_dump") + else dict(item) + for item in (getattr(response, "outputs", None) or []) + ] + tool_calls = [ + { + "id": str(item.get("id") or ""), + "name": str(item.get("name") or ""), + "args": dict(item.get("arguments") or {}), + "type": "tool_call", + } + for item in blocks + if item.get("type") == "function_call" + ] + usage_metadata = _usage_metadata( + _dump(getattr(response, "usage", None)), + provider_request_id=getattr(response, "id", None), + ) + message = AIMessage( + content=blocks, + tool_calls=tool_calls, + additional_kwargs={"gemini_interaction_content": blocks}, + usage_metadata=usage_metadata, + response_metadata={ + "model_name": str(getattr(getattr(response, "model", None), "id", "")), + "finish_reason": str(getattr(response, "status", "unknown")), + }, + ) + return ChatResult(generations=[ChatGeneration(message=message)]) + + +def _dump(value: Any) -> dict[str, Any]: + if value is None: + return {} + if hasattr(value, "model_dump"): + return value.model_dump(mode="json", by_alias=True, exclude_none=True) + if isinstance(value, Mapping): + return dict(value) + return {} + + +def _merge_stream_delta(block: dict[str, Any], delta: Mapping[str, Any]) -> None: + kind = str(delta.get("type") or "") + if kind == "text": + block["type"] = "text" + block["text"] = str(block.get("text") or "") + str(delta.get("text") or "") + elif kind == "thought_signature": + block.setdefault("type", "thought") + block["signature"] = delta.get("signature") + elif kind == "thought_summary": + block.setdefault("type", "thought") + content = delta.get("content") + if content is not None: + block.setdefault("summary", []).append(content) + elif kind == "text_annotation": + block.setdefault("annotations", []).extend(delta.get("annotations") or []) + else: + block.update(delta) + + +def _usage_metadata( + usage: Mapping[str, Any] | None, *, provider_request_id: Any +) -> dict[str, Any] | None: + if not usage: + return None + required = ( + usage.get("total_input_tokens"), + usage.get("total_cached_tokens"), + usage.get("total_output_tokens"), + ) + if any(value is None for value in required): + return None + result: dict[str, Any] = { + "input_tokens": int(required[0]), + "cached_input_tokens": int(required[1]), + "output_tokens": int(required[2]), + "total_tokens": int( + usage.get("total_tokens") + if usage.get("total_tokens") is not None + else int(required[0]) + int(required[2]) + ), + "usage_finality": "confirmed", + } + optional = { + "reasoning_tokens": usage.get("total_thought_tokens"), + "provider_request_id": provider_request_id, + } + result.update({key: value for key, value in optional.items() if value is not None}) + return result + + +def create_gemini_interactions_model(**kwargs: Any) -> GeminiInteractionsChatModel: + return GeminiInteractionsChatModel( + model_name=str(kwargs["model"]), + api_key=SecretStr(str(kwargs["api_key"])), + base_url=str( + kwargs.get("base_url") or "https://generativelanguage.googleapis.com" + ), + max_output_tokens=int(kwargs.get("max_output_tokens") or 8192), + temperature=kwargs.get("temperature"), + top_p=kwargs.get("top_p"), + thinking=kwargs.get("thinking"), + ) diff --git a/EvoScientist/llm/invocation/__init__.py b/EvoScientist/llm/invocation/__init__.py new file mode 100644 index 0000000..6c56b40 --- /dev/null +++ b/EvoScientist/llm/invocation/__init__.py @@ -0,0 +1,20 @@ +"""Compilation of normalized configuration into provider invocation plans.""" + +from .contract import ( + InvocationPlan, + ToolCallTransport, + compile_invocation_plan, + derive_runtime_invocation, + derive_tool_call_transport, +) +from .messages import assistant_message_has_output, project_provider_messages + +__all__ = [ + "InvocationPlan", + "ToolCallTransport", + "assistant_message_has_output", + "compile_invocation_plan", + "derive_runtime_invocation", + "derive_tool_call_transport", + "project_provider_messages", +] diff --git a/EvoScientist/llm/invocation/contract.py b/EvoScientist/llm/invocation/contract.py new file mode 100644 index 0000000..16a594e --- /dev/null +++ b/EvoScientist/llm/invocation/contract.py @@ -0,0 +1,133 @@ +"""Pure invocation contract shared by configuration and runtime. + +Provider configuration owns connection/authentication and adapter selection. +Model configuration owns capabilities, limits, and canonical parameters. This +module is the only layer that derives the effective wire invocation from those +inputs; the derived fields are runtime state, not administrator configuration. +""" + +from __future__ import annotations + +import hashlib +from collections.abc import Mapping +from dataclasses import dataclass +from types import MappingProxyType +from typing import Any, Literal + +from ..contracts import EvoRuntimeError +from ..crypto import canonical_json_v1 + +ToolCallTransport = Literal["native", "disabled"] + + +def derive_tool_call_transport( + capabilities: Mapping[str, Any], +) -> ToolCallTransport: + """Derive the wire transport from the authoritative tools capability.""" + + return "native" if bool(capabilities.get("tools", False)) else "disabled" + + +def derive_runtime_invocation( + api_mode: str, + capabilities: Mapping[str, Any], +) -> dict[str, str]: + """Build the internal invocation projection for a normalized model.""" + + return { + "api_mode": str(api_mode), + "tool_call_transport": derive_tool_call_transport(capabilities), + } + + +@dataclass(frozen=True, slots=True) +class InvocationPlan: + """Complete, immutable, non-secret wire contract for one model call.""" + + api_mode: str + output_token_parameter: str + output_token_limit: int + tool_call_transport: ToolCallTransport + reasoning_effort: str + streaming: bool + sdk_params: Mapping[str, Any] + plan_hash: str + + def model_kwargs(self) -> dict[str, Any]: + return dict(self.sdk_params) + + def projection(self) -> dict[str, Any]: + return { + "api_mode": self.api_mode, + "output_token_parameter": self.output_token_parameter, + "output_token_limit": self.output_token_limit, + "tool_call_transport": self.tool_call_transport, + "reasoning_effort": self.reasoning_effort, + "streaming": self.streaming, + "sdk_params": dict(self.sdk_params), + "plan_hash": self.plan_hash, + } + + +def compile_invocation_plan( + *, + api_mode: str, + declared_tool_call_transport: str, + supports_tools: bool, + purpose: str, + output_token_limit: int, + reasoning_effort: str, + runtime_provider: str, + sdk_params: Mapping[str, Any], +) -> InvocationPlan: + """Validate and freeze adapter output before constructing a provider SDK.""" + + params = dict(sdk_params) + streaming = purpose == "main_agent" + params["streaming"] = streaming + if runtime_provider == "openai" and streaming: + params["stream_usage"] = True + token_fields = tuple( + key + for key in ("max_output_tokens", "max_completion_tokens", "max_tokens") + if key in params + ) + if len(token_fields) != 1 or params[token_fields[0]] != output_token_limit: + raise EvoRuntimeError("MODEL_ADAPTER_COMPILE_FAILED") + output_token_parameter = token_fields[0] + if api_mode == "responses" and output_token_parameter != "max_output_tokens": + raise EvoRuntimeError("MODEL_ADAPTER_COMPILE_FAILED") + if api_mode == "chat_completions" and output_token_parameter == "max_output_tokens": + raise EvoRuntimeError("MODEL_ADAPTER_COMPILE_FAILED") + if runtime_provider == "openai" and api_mode in { + "responses", + "chat_completions", + }: + if params.get("use_responses_api") is not (api_mode == "responses"): + raise EvoRuntimeError("MODEL_ADAPTER_COMPILE_FAILED") + + tool_call_transport: ToolCallTransport = ( + "native" if supports_tools else "disabled" + ) + if declared_tool_call_transport != tool_call_transport: + raise EvoRuntimeError("MODEL_ADAPTER_COMPILE_FAILED") + + plan_payload = { + "api_mode": api_mode, + "output_token_parameter": output_token_parameter, + "output_token_limit": output_token_limit, + "tool_call_transport": tool_call_transport, + "reasoning_effort": reasoning_effort, + "streaming": streaming, + "sdk_params": params, + } + return InvocationPlan( + api_mode=api_mode, + output_token_parameter=output_token_parameter, + output_token_limit=output_token_limit, + tool_call_transport=tool_call_transport, + reasoning_effort=reasoning_effort, + streaming=streaming, + sdk_params=MappingProxyType(params), + plan_hash=hashlib.sha256(canonical_json_v1(plan_payload)).hexdigest(), + ) diff --git a/EvoScientist/llm/invocation/messages.py b/EvoScientist/llm/invocation/messages.py new file mode 100644 index 0000000..f0ae675 --- /dev/null +++ b/EvoScientist/llm/invocation/messages.py @@ -0,0 +1,68 @@ +"""Provider-facing message projection and response validation.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from typing import Any + +from langchain_core.messages import AIMessage + +_TOOL_BLOCK_TYPES = frozenset({"tool_call", "tool_use", "function_call"}) +_NON_FINAL_BLOCK_TYPES = frozenset({"reasoning", "thinking"}) + + +def assistant_message_has_output(message: AIMessage) -> bool: + """Return whether an assistant message has provider-visible output.""" + + if ( + getattr(message, "tool_calls", None) + or getattr(message, "invalid_tool_calls", None) + or _additional_tool_calls(message) + ): + return True + content = getattr(message, "content", None) + if isinstance(content, str): + return bool(content.strip()) + if not isinstance(content, Sequence) or isinstance(content, str | bytes): + return content is not None + for block in content: + if isinstance(block, str): + if block.strip(): + return True + continue + if not isinstance(block, Mapping): + return True + block_type = str(block.get("type") or "").strip() + if block_type in _TOOL_BLOCK_TYPES: + return True + if block_type in _NON_FINAL_BLOCK_TYPES: + continue + text = block.get("text") + if isinstance(text, str): + if text.strip(): + return True + continue + # Unknown non-reasoning blocks may carry multimodal or refusal output. + if block: + return True + return False + + +def project_provider_messages(messages: Sequence[Any]) -> tuple[list[Any], int]: + """Drop unusable assistant history without mutating checkpoint objects.""" + + projected: list[Any] = [] + dropped = 0 + for message in messages: + if isinstance(message, AIMessage) and not assistant_message_has_output(message): + dropped += 1 + continue + projected.append(message) + return projected, dropped + + +def _additional_tool_calls(message: AIMessage) -> bool: + additional = getattr(message, "additional_kwargs", None) + if not isinstance(additional, Mapping): + return False + return bool(additional.get("tool_calls") or additional.get("function_call")) diff --git a/EvoScientist/llm/model_config.py b/EvoScientist/llm/model_config.py new file mode 100644 index 0000000..0e066d8 --- /dev/null +++ b/EvoScientist/llm/model_config.py @@ -0,0 +1,3309 @@ +"""Strict Evo-owned V2/V3 model-route configuration.""" + +from __future__ import annotations + +import hashlib +import hmac +import json +import math +import os +import re +import secrets +import sqlite3 +import stat +import tempfile +import time +from collections.abc import Mapping, Sequence +from dataclasses import asdict, dataclass, field +from decimal import Decimal, InvalidOperation +from pathlib import Path +from typing import Any + +import yaml +from filelock import FileLock + +from ..config.settings import get_config_dir +from .adapter_registry import get_adapter_registry +from .configuration import ( + EndpointConfig, + ModelConfig, + ProviderConfig, + ResolvedSecret, + SecretReference, + SecretResolver, +) +from .contracts import ( + AdminConfigGrant, + AdminConfigGrantVerifier, + EvoRuntimeError, + PricingQuote, +) +from .crypto import canonical_json_v1, hmac_id, sha256_id +from .user_options import ( + project_user_options_for_purpose, + validate_parameter_constraints, +) + +_PURPOSES = ("main_agent", "tool_selector", "deepagents_summarizer", "title") +_REASONING_EFFORTS = frozenset({"disabled", "low", "medium", "high", "max"}) +_REASONING_MODES = frozenset({"effort", "boolean"}) +_BLOCKED_PARAM_KEYS = frozenset( + { + "api_key", + "base_url", + "max_tokens", + "max_output_tokens", + "max_retries", + "streaming", + "disable_streaming", + "use_responses_api", + "reasoning", + "token", + "secret", + "password", + "authorization", + } +) +_SENSITIVE_HEADERS = frozenset( + {"authorization", "proxy-authorization", "cookie", "set-cookie", "x-api-key"} +) +_HEADER_RE = re.compile(r"^[!#$%&'*+.^_`|~0-9A-Za-z-]+$") +_ENV_RE = re.compile(r"^[A-Z_][A-Z0-9_]*$") +_BIGINT_MAX = 2**63 - 1 +_PARAM_ALLOWLIST: Mapping[tuple[str, str], frozenset[str]] = { + ("custom-openai", "chat_completions"): frozenset( + { + "temperature", + "top_p", + "seed", + "frequency_penalty", + "presence_penalty", + "timeout", + "extra_body", + "output_token_limit", + } + ), + ("custom-openai", "responses"): frozenset( + {"temperature", "top_p", "seed", "timeout", "extra_body", "output_token_limit"} + ), + ("openai", "chat_completions"): frozenset( + { + "temperature", + "top_p", + "seed", + "frequency_penalty", + "presence_penalty", + "timeout", + "extra_body", + "output_token_limit", + } + ), + ("openai", "responses"): frozenset( + {"temperature", "top_p", "seed", "timeout", "extra_body", "output_token_limit"} + ), + ("anthropic", "messages"): frozenset( + {"temperature", "top_p", "top_k", "timeout", "output_token_limit"} + ), +} + + +class _UniqueKeyLoader(yaml.SafeLoader): + pass + + +def _construct_mapping( + loader: yaml.SafeLoader, node: yaml.MappingNode, deep: bool = False +) -> Any: + loader.flatten_mapping(node) + result: dict[Any, Any] = {} + for key_node, value_node in node.value: + key = loader.construct_object(key_node, deep=deep) + if key in result: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"duplicate YAML key: {key}" + ) + result[key] = loader.construct_object(value_node, deep=deep) + return result + + +_UniqueKeyLoader.add_constructor( + yaml.resolver.BaseResolver.DEFAULT_MAPPING_TAG, + _construct_mapping, +) + + +def load_yaml_unique(text: str) -> Any: + try: + return yaml.load(text, Loader=_UniqueKeyLoader) + except EvoRuntimeError: + raise + except yaml.YAMLError as exc: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "invalid model route YAML" + ) from exc + + +def _mapping(value: Any, name: str) -> Mapping[str, Any]: + if not isinstance(value, Mapping): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} must be a mapping" + ) + if not all(isinstance(key, str) for key in value): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} keys must be strings" + ) + return value + + +def _sequence(value: Any, name: str, *, allow_empty: bool = True) -> Sequence[Any]: + if not isinstance(value, list | tuple): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} must be a list" + ) + if not allow_empty and not value: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} must not be empty" + ) + return value + + +def _strict_keys(value: Mapping[str, Any], *, allowed: set[str], name: str) -> None: + unknown = set(value) - allowed + if unknown: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + f"{name} has unknown fields: {', '.join(sorted(unknown))}", + ) + + +def _text(value: Any, name: str) -> str: + if not isinstance(value, str) or not value.strip(): + raise EvoRuntimeError("LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} is required") + return value.strip() + + +def _integer( + value: Any, name: str, *, minimum: int = 0, maximum: int = _BIGINT_MAX +) -> int: + if isinstance(value, bool) or not isinstance(value, int): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} must be an integer" + ) + if value < minimum or value > maximum: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} is out of range" + ) + return value + + +def _positive_decimal(value: Any, name: str, *, default: str = "1") -> str: + if value is None: + value = default + if isinstance(value, bool): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} must be a positive number" + ) + try: + parsed = Decimal(str(value)) + except (InvalidOperation, TypeError, ValueError) as exc: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} must be a positive number" + ) from exc + if not parsed.is_finite() or parsed <= 0: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} must be greater than zero" + ) + return format(parsed.normalize(), "f") + + +def _bool(value: Any, name: str) -> bool: + if not isinstance(value, bool): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} must be boolean" + ) + return value + + +def _string_set( + value: Any, name: str, *, allowed: frozenset[str] | None = None +) -> tuple[str, ...]: + items = _sequence(value, name) + normalized = tuple(sorted({_text(item, name) for item in items})) + if allowed is not None and not set(normalized) <= allowed: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} contains unsupported values" + ) + return normalized + + +def _validate_params( + value: Any, *, name: str, allowed: frozenset[str] +) -> Mapping[str, Any]: + params = dict(_mapping(value or {}, name)) + unknown = set(params) - allowed + if unknown: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + f"{name} has unsupported parameters: {', '.join(sorted(unknown))}", + ) + + def visit(item: Any, path: str) -> None: + if item is None or isinstance(item, bool | int | float | str): + canonical_json_v1(item) + return + if isinstance(item, Mapping): + for key, child in item.items(): + if not isinstance(key, str): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{path} key must be text" + ) + if key.lower() in _BLOCKED_PARAM_KEYS: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{path}.{key} is forbidden" + ) + visit(child, f"{path}.{key}") + return + if isinstance(item, list | tuple): + for index, child in enumerate(item): + visit(child, f"{path}[{index}]") + return + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{path} is not JSON data" + ) + + visit(params, name) + return params + + +@dataclass(frozen=True, slots=True) +class AliasConfig: + alias: str + display_name: str + provider_ref: str + model_ref: str + enabled: bool + access: Mapping[str, Any] + defaults: Mapping[str, Any] + + +@dataclass(frozen=True, slots=True) +class PoolEndpoint: + name: str + weight: int + + +@dataclass(frozen=True, slots=True) +class EndpointPool: + pool_id: str + provider: str + strategy: LiteralStrategy + endpoints: tuple[PoolEndpoint, ...] + + +LiteralStrategy = str + + +@dataclass(frozen=True, slots=True) +class RouteSelector: + selector_id: str + provider: str + endpoint: str | None + endpoint_pool: str | None + model: str + api_mode: str + tool_call_transport: str + identity_selector_id: str = "" + alias: str = "" + alias_defaults: Mapping[str, Any] = field(default_factory=dict) + + +@dataclass(frozen=True, slots=True) +class RouteRef: + selector_id: str + provider: str + endpoint: str + model: str + api_mode: str + tool_call_transport: str + + def key(self) -> str: + return ":".join( + ( + self.provider, + self.endpoint, + self.model, + self.api_mode, + self.tool_call_transport, + ) + ) + + +@dataclass(frozen=True, slots=True) +class PurposeRoutes: + default_alias: str + selectable: Mapping[str, str] + + +@dataclass(frozen=True, slots=True) +class PurposeCallLimit: + max_attempts_per_run: int + + +def _effective_output_limit(model: ModelConfig, params: Mapping[str, Any]) -> int: + override = params.get("output_token_limit") + return model.max_output_tokens if override is None else int(override) + + +def _validate_model_output_limits(model: ModelConfig) -> None: + candidates = [model.params.get("output_token_limit")] + candidates.extend( + values.get("output_token_limit") + for values in model.purpose_overrides.values() + ) + for value in candidates: + if value is not None and int(value) > model.max_output_tokens: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "model output_token_limit exceeds model capability", + ) + + +@dataclass(frozen=True, slots=True) +class RouteHealthPolicy: + failure_threshold: int + cooldown_seconds: int + half_open_max_inflight: int + counted_error_codes: tuple[str, ...] + open_immediately_error_codes: tuple[str, ...] + + +@dataclass(frozen=True, slots=True) +class WebRuntimePolicy: + title_start_timeout_seconds: int + prepare_ttl_seconds: int + turn_lease_grace_seconds: int + active_run_timeout_seconds: int + max_run_journal_events: int + max_run_journal_bytes: int + max_prepared_runs_per_subject: int + max_prepared_runs_total: int + + +@dataclass(frozen=True, slots=True) +class CapabilityEvidence: + route: RouteRef + connectivity: str + tool_capability: str + route_semantics_hash: str + endpoint_fingerprint: str + config_identity_key_id: str + adapter_revision: str + fixture_digest: str + verified_at: str + results: Mapping[str, str] = field(default_factory=dict) + evidence_expires_at: str = "" + implementation_fingerprint: str = "" + resolved_model_revision: str | None = None + reproducible: bool = False + + +@dataclass(frozen=True, slots=True) +class EvoModelConfig: + config_revision: int + config_identity_key_id: str + runtime_defaults: Mapping[str, Any] + purpose_defaults: Mapping[str, Mapping[str, str]] + providers: Mapping[str, ProviderConfig] + endpoint_pools: Mapping[str, EndpointPool] + route_health: RouteHealthPolicy + route_selectors: Mapping[str, RouteSelector] + main_routes: PurposeRoutes + title_selector_id: str + purpose_call_limits: Mapping[str, PurposeCallLimit] + web_runtime: WebRuntimePolicy + capability_evidence: Mapping[str, CapabilityEvidence] + tool_protocol_fallbacks: Mapping[str, tuple[str, ...]] + raw: Mapping[str, Any] + schema_version: int = 2 + aliases: Mapping[str, AliasConfig] = field(default_factory=dict) + provider_health: RouteHealthPolicy | None = None + adapter_registry_revision: str = "" + purpose_selector_ids: Mapping[str, str] = field(default_factory=dict) + + @classmethod + def parse(cls, payload: Any, *, require_evidence: bool = True) -> EvoModelConfig: + # ``require_evidence`` is retained for callers of the legacy schema API. + # Capability evidence is historical audit data, not a publication or + # invocation prerequisite. + raw = _mapping(payload, "model_routes") + if raw.get("schema_version") == 3: + return _parse_v3_config(raw, require_evidence=require_evidence) + _strict_keys( + raw, + allowed={ + "schema_version", + "config_revision", + "config_identity_key_id", + "runtime_defaults", + "purpose_defaults", + "providers", + "endpoint_pools", + "route_health", + "route_selectors", + "purpose_routes", + "purpose_call_limits", + "web_runtime", + "capability_evidence", + "tool_protocol_fallbacks", + }, + name="model_routes", + ) + if raw.get("schema_version") != 2: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "schema_version must be 2" + ) + revision = _integer(raw.get("config_revision"), "config_revision", minimum=1) + identity_key_id = _text( + raw.get("config_identity_key_id"), "config_identity_key_id" + ) + runtime_defaults = _mapping(raw.get("runtime_defaults"), "runtime_defaults") + _strict_keys(runtime_defaults, allowed={"max_retries"}, name="runtime_defaults") + if runtime_defaults.get("max_retries") != 0: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "runtime max_retries must be 0" + ) + purpose_defaults = _parse_purpose_defaults(raw.get("purpose_defaults")) + providers = _parse_providers(raw.get("providers")) + pools = _parse_endpoint_pools(raw.get("endpoint_pools"), providers) + health = _parse_health(raw.get("route_health")) + selectors = _parse_selectors(raw.get("route_selectors"), providers, pools) + main_routes, title_selector_id = _parse_purpose_routes( + raw.get("purpose_routes"), selectors + ) + limits = _parse_call_limits(raw.get("purpose_call_limits")) + web_runtime = _parse_web_runtime(raw.get("web_runtime")) + fallbacks = _parse_fallbacks(raw.get("tool_protocol_fallbacks"), selectors) + config = cls( + config_revision=revision, + config_identity_key_id=identity_key_id, + runtime_defaults=dict(runtime_defaults), + purpose_defaults=purpose_defaults, + providers=providers, + endpoint_pools=pools, + route_health=health, + route_selectors=selectors, + main_routes=main_routes, + title_selector_id=title_selector_id, + purpose_call_limits=limits, + web_runtime=web_runtime, + capability_evidence={}, + tool_protocol_fallbacks=fallbacks, + raw=dict(raw), + ) + config._validate(require_evidence=require_evidence) + return config + + @property + def catalog_revision(self) -> int: + return self.config_revision + + @property + def title_start_timeout_seconds(self) -> int: + return self.web_runtime.title_start_timeout_seconds + + def resolve_main_selector(self, alias: str | None) -> RouteSelector: + selected_alias = str(alias or "").strip() or self.main_routes.default_alias + selector_id = self.main_routes.selectable.get(selected_alias) + if selector_id is None: + raise EvoRuntimeError("MODEL_ACCESS_DENIED") + return self.route_selectors[selector_id] + + def concrete_routes(self, selector_id: str) -> tuple[RouteRef, ...]: + selector = self.route_selectors[selector_id] + if selector.endpoint is not None: + endpoints = (selector.endpoint,) + else: + assert selector.endpoint_pool is not None + endpoints = tuple( + item.name + for item in self.endpoint_pools[selector.endpoint_pool].endpoints + ) + return tuple( + RouteRef( + selector_id=selector_id, + provider=selector.provider, + endpoint=endpoint, + model=selector.model, + api_mode=selector.api_mode, + tool_call_transport=selector.tool_call_transport, + ) + for endpoint in endpoints + ) + + def fallback_selectors(self, primary_selector_id: str) -> tuple[str, ...]: + return self.tool_protocol_fallbacks.get(primary_selector_id, ()) + + def route_model(self, route: RouteRef) -> ModelConfig: + return self.providers[route.provider].models[route.model] + + def required_concrete_routes(self) -> Mapping[str, tuple[str, ...]]: + if self.schema_version == 3: + required: dict[str, tuple[str, ...]] = {} + selector_ids = set(self.main_routes.selectable.values()) | { + self.title_selector_id + } + for selector_id in selector_ids: + route = self.concrete_routes(selector_id)[0] + model = self.route_model(route) + probes = ["connectivity"] + probes.extend( + key + for key, enabled in model.capabilities.items() + if enabled and key != "text" + ) + required[route.key()] = tuple(dict.fromkeys(probes)) + return required + main_selector_ids = set(self.main_routes.selectable.values()) + reachable = set(main_selector_ids) + for selector_id in tuple(main_selector_ids): + reachable.update(self.fallback_selectors(selector_id)) + result: dict[str, tuple[str, ...]] = {} + for selector_id in reachable: + for route in self.concrete_routes(selector_id): + result[route.key()] = ("connectivity", "tool_protocol") + for route in self.concrete_routes(self.title_selector_id): + result.setdefault(route.key(), ("connectivity",)) + return result + + def _validate(self, *, require_evidence: bool = True) -> None: + if self.schema_version == 3: + self._validate_v3(require_evidence=require_evidence) + return + main_limit = self.purpose_call_limits["main_agent"].max_attempts_per_run + for selector in self.route_selectors.values(): + model = self.providers[selector.provider].models[selector.model] + _validate_model_output_limits(model) + for primary, fallbacks in self.tool_protocol_fallbacks.items(): + chain = (primary, *fallbacks) + if len(chain) != len(set(chain)) or len(chain) > main_limit: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "invalid fallback chain" + ) + primary_models = [ + self.route_model(route) for route in self.concrete_routes(primary) + ] + base = primary_models[0] + for selector_id in fallbacks: + for candidate in self.concrete_routes(selector_id): + model = self.route_model(candidate) + if model.model_id != base.model_id or model.quote != base.quote: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "fallback billing differs", + ) + + def _validate_v3(self, *, require_evidence: bool) -> None: + if self.endpoint_pools or self.tool_protocol_fallbacks: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "V3 direct routes cannot contain pools or fallback chains", + ) + if not self.aliases or not self.main_routes.selectable: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "aliases are required" + ) + total_attempts = sum( + value.max_attempts_per_run for value in self.purpose_call_limits.values() + ) + if total_attempts > 16: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "purpose attempts exceed the run limit", + ) + for alias, selector_id in self.main_routes.selectable.items(): + selector = self.route_selectors[selector_id] + model = self.providers[selector.provider].models[selector.model] + registration = get_adapter_registry().get( + self.providers[selector.provider].adapter_id, + self.providers[selector.provider].adapter_revision, + ) + if not model.enabled or not self.providers[selector.provider].enabled: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + f"enabled alias {alias} has a disabled target", + ) + _validate_model_output_limits(model) + alias_config = self.aliases[alias] + applicable_purposes = {"main_agent"} + for purpose in ("tool_selector", "deepagents_summarizer"): + explicit_selector = self.purpose_selector_ids.get(purpose) + if explicit_selector in {None, selector_id}: + applicable_purposes.add(purpose) + if self.title_selector_id == selector_id: + applicable_purposes.add("title") + for purpose in applicable_purposes: + params = dict(self.purpose_defaults[purpose]) + params.update( + project_user_options_for_purpose( + values=model.params, + user_options=model.user_options, + purpose=purpose, + ) + ) + for name, option in model.user_options.items(): + if "default" in option and purpose in set( + option.get("applies_to") or ("main_agent",) + ): + params.setdefault(name, option["default"]) + params.update( + project_user_options_for_purpose( + values=alias_config.defaults, + user_options=model.user_options, + purpose=purpose, + ) + ) + params.update(model.purpose_overrides.get(purpose, {})) + registration.validate_parameters( + params, path=f"aliases.{alias}.merged.{purpose}" + ) + registration.compile_runtime_parameters( + selector.api_mode, + params, + _effective_output_limit(model, params), + provider_model_id=model.model_id, + ) + + +_V3_RUNTIME_DEFAULTS: Mapping[str, tuple[int, int, int]] = { + "sdk_max_retries": (0, 0, 0), + "connect_timeout_seconds": (10, 1, 60), + "first_event_timeout_seconds": (60, 1, 300), + "stream_idle_timeout_seconds": (60, 1, 300), + "attempt_timeout_seconds": (600, 1, 3600), + "max_sse_event_bytes": (1_048_576, 4_096, 4_194_304), + "max_content_block_bytes": (4_194_304, 4_096, 16_777_216), + "max_output_bytes": (16_777_216, 65_536, 67_108_864), + "max_opaque_state_bytes": (8_388_608, 65_536, 33_554_432), + "stream_buffer_max_events": (256, 1, 1_024), + "stream_buffer_max_bytes": (2_097_152, 65_536, 8_388_608), + "max_tool_schema_bytes": (262_144, 4_096, 1_048_576), + "max_tool_arguments_bytes": (1_048_576, 4_096, 4_194_304), + "max_tool_schema_depth": (16, 1, 32), + "max_tool_argument_depth": (32, 1, 64), +} +_V3_CAPABILITIES = ( + "text", + "vision", + "video", + "documents", + "tools", + "structured_output", + "thinking", +) +_V3_ACCESS_KEYS = {"visibility", "roles", "groups", "users"} + + +def _parse_v3_access(value: Any, name: str) -> Mapping[str, Any]: + raw = _mapping(value, name) + _strict_keys(raw, allowed=_V3_ACCESS_KEYS, name=name) + visibility = _text(raw.get("visibility"), f"{name}.visibility") + if visibility not in {"authenticated", "role_based", "private"}: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name}.visibility is invalid" + ) + result = { + "visibility": visibility, + "roles": _string_set(raw.get("roles") or [], f"{name}.roles"), + "groups": _string_set(raw.get("groups") or [], f"{name}.groups"), + "users": _string_set(raw.get("users") or [], f"{name}.users"), + } + if visibility == "authenticated" and any( + result[key] for key in ("roles", "groups", "users") + ): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + f"{name} authenticated access must be unscoped", + ) + if visibility == "role_based" and not (result["roles"] or result["groups"]): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} role access is empty" + ) + if visibility == "private" and not result["users"]: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name} private access is empty" + ) + return result + + +def _normalize_v3_user_option( + value: Mapping[str, Any], rule: Any, path: str +) -> Mapping[str, Any]: + option = dict(value) + numeric = rule.kind in {"integer", "number"} + bound_names = ( + "minimum", + "minimum_exclusive", + "maximum", + "maximum_exclusive", + ) + if any(name in option for name in bound_names) and not numeric: + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", f"{path} bounds require a numeric parameter" + ) + if "minimum" in option and "minimum_exclusive" in option: + raise EvoRuntimeError("MODEL_PARAMETER_INVALID", f"{path} has two lower bounds") + if "maximum" in option and "maximum_exclusive" in option: + raise EvoRuntimeError("MODEL_PARAMETER_INVALID", f"{path} has two upper bounds") + for name in bound_names: + if name not in option: + continue + bound = option[name] + if ( + isinstance(bound, bool) + or not isinstance(bound, int | float) + or not math.isfinite(float(bound)) + ): + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", f"{path}.{name} must be finite" + ) + if numeric and rule.minimum is not None: + default_name = "minimum_exclusive" if rule.minimum_exclusive else "minimum" + option.setdefault(default_name, rule.minimum) + if numeric and rule.maximum is not None: + default_name = "maximum_exclusive" if rule.maximum_exclusive else "maximum" + option.setdefault(default_name, rule.maximum) + + lower_name = next( + (name for name in ("minimum", "minimum_exclusive") if name in option), None + ) + upper_name = next( + (name for name in ("maximum", "maximum_exclusive") if name in option), None + ) + if lower_name and rule.minimum is not None: + lower = float(option[lower_name]) + if lower < rule.minimum or ( + lower == rule.minimum and rule.minimum_exclusive and lower_name == "minimum" + ): + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", f"{path} loosens the adapter minimum" + ) + if upper_name and rule.maximum is not None: + upper = float(option[upper_name]) + if upper > rule.maximum or ( + upper == rule.maximum and rule.maximum_exclusive and upper_name == "maximum" + ): + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", f"{path} loosens the adapter maximum" + ) + if lower_name and upper_name: + lower = float(option[lower_name]) + upper = float(option[upper_name]) + if lower > upper or ( + lower == upper + and (lower_name == "minimum_exclusive" or upper_name == "maximum_exclusive") + ): + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", f"{path} has an empty range" + ) + + if "choices" in option: + choices = list( + _sequence(option["choices"], f"{path}.choices", allow_empty=False) + ) + for index, choice in enumerate(choices): + rule.validate(choice, f"{path}.choices[{index}]") + if len({canonical_json_v1(choice) for choice in choices}) != len(choices): + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", f"{path}.choices contains duplicates" + ) + option["choices"] = choices + elif rule.kind == "enum": + option["choices"] = list(rule.choices) + + if "default" in option: + default = option["default"] + rule.validate(default, f"{path}.default") + if lower_name and ( + default < option[lower_name] + or (lower_name == "minimum_exclusive" and default == option[lower_name]) + ): + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", f"{path}.default is below its range" + ) + if upper_name and ( + default > option[upper_name] + or (upper_name == "maximum_exclusive" and default == option[upper_name]) + ): + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", f"{path}.default exceeds its range" + ) + if option.get("choices") and default not in option["choices"]: + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", f"{path}.default is not an allowed choice" + ) + return option + + +def _parse_v3_billing(value: Any, name: str, *, require_evidence: bool) -> PricingQuote: + raw = _mapping(value, name) + _strict_keys( + raw, + allowed={ + "sku", + "pricing_revision", + "currency", + "unit_scale", + "input_microunits_per_million", + "output_microunits_per_million", + "cached_microunits_per_million", + "multiplier", + }, + name=name, + ) + payload = { + "billing_sku": _text(raw.get("sku"), f"{name}.sku"), + "pricing_revision": _text( + raw.get("pricing_revision"), f"{name}.pricing_revision" + ), + "currency": _text(raw.get("currency"), f"{name}.currency"), + "unit_scale": _integer(raw.get("unit_scale"), f"{name}.unit_scale", minimum=1), + "input_microunits_per_million": _integer( + raw.get("input_microunits_per_million"), f"{name}.input", minimum=0 + ), + "output_microunits_per_million": _integer( + raw.get("output_microunits_per_million"), f"{name}.output", minimum=0 + ), + "cached_input_microunits_per_million": _integer( + raw.get("cached_microunits_per_million"), f"{name}.cached", minimum=0 + ), + "multiplier": _positive_decimal(raw.get("multiplier"), f"{name}.multiplier"), + } + if payload["currency"] != "CNY" or payload["unit_scale"] != 1_000_000: + raise EvoRuntimeError("PRICING_DIMENSION_UNSUPPORTED") + if require_evidence and payload["pricing_revision"].startswith("draft-"): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "active pricing is not approved" + ) + return PricingQuote(**payload, quote_id=sha256_id(payload)) + + +def _parse_v3_runtime_defaults(value: Any) -> Mapping[str, int]: + raw = _mapping(value or {}, "runtime_defaults") + _strict_keys(raw, allowed=set(_V3_RUNTIME_DEFAULTS), name="runtime_defaults") + result = { + key: _integer( + raw.get(key, default), + f"runtime_defaults.{key}", + minimum=minimum, + maximum=maximum, + ) + for key, (default, minimum, maximum) in _V3_RUNTIME_DEFAULTS.items() + } + if result["sdk_max_retries"] != 0: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "sdk_max_retries must be 0" + ) + if result["attempt_timeout_seconds"] < max( + result["connect_timeout_seconds"], + result["first_event_timeout_seconds"], + result["stream_idle_timeout_seconds"], + ): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "attempt timeout is too small" + ) + return result + + +def _parse_v3_provider_defaults( + value: Any, runtime: Mapping[str, int] +) -> Mapping[str, int]: + raw = _mapping(value or {}, "provider.defaults") + allowed = { + "connect_timeout_seconds", + "first_event_timeout_seconds", + "stream_idle_timeout_seconds", + "attempt_timeout_seconds", + "max_inflight_requests", + "queue_timeout_seconds", + } + _strict_keys(raw, allowed=allowed, name="provider.defaults") + result: dict[str, int] = {} + for key in allowed - {"max_inflight_requests", "queue_timeout_seconds"}: + result[key] = _integer( + raw.get(key, runtime[key]), + f"provider.defaults.{key}", + minimum=1, + maximum=_V3_RUNTIME_DEFAULTS[key][2], + ) + result["max_inflight_requests"] = _integer( + raw.get("max_inflight_requests", 16), + "provider.defaults.max_inflight_requests", + minimum=1, + maximum=256, + ) + result["queue_timeout_seconds"] = _integer( + raw.get("queue_timeout_seconds", 5), + "provider.defaults.queue_timeout_seconds", + minimum=1, + maximum=60, + ) + if result["attempt_timeout_seconds"] < max( + result["connect_timeout_seconds"], + result["first_event_timeout_seconds"], + result["stream_idle_timeout_seconds"], + ): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "provider attempt timeout is too small" + ) + return result + + +def _parse_v3_health(value: Any, name: str) -> RouteHealthPolicy: + raw = _mapping(value or {}, name) + _strict_keys( + raw, + allowed={ + "failure_threshold", + "cooldown_seconds", + "half_open_max_inflight", + "counted_error_codes", + "open_immediately_error_codes", + }, + name=name, + ) + counted = _string_set( + raw.get("counted_error_codes") or [], f"{name}.counted_error_codes" + ) + immediate = _string_set( + raw.get("open_immediately_error_codes") or [], + f"{name}.open_immediately_error_codes", + ) + return RouteHealthPolicy( + _integer( + raw.get("failure_threshold", 3), + f"{name}.failure_threshold", + minimum=1, + maximum=20, + ), + _integer( + raw.get("cooldown_seconds", 30), + f"{name}.cooldown_seconds", + minimum=1, + maximum=3600, + ), + _integer( + raw.get("half_open_max_inflight", 1), + f"{name}.half_open_max_inflight", + minimum=1, + maximum=16, + ), + counted, + immediate, + ) + + +def _parse_v3_web_runtime(value: Any) -> WebRuntimePolicy: + raw = _mapping(value or {}, "web_runtime") + specs = { + "title_start_timeout_seconds": (30, 1, 300), + "prepare_ttl_seconds": (30, 5, 300), + "turn_lease_grace_seconds": (30, 1, 300), + "active_run_timeout_seconds": (1800, 60, 7200), + "max_run_journal_events": (10000, 100, 100000), + "max_run_journal_bytes": (16777216, 1048576, 268435456), + "max_prepared_runs_per_subject": (4, 1, 32), + "max_prepared_runs_total": (128, 1, 4096), + } + _strict_keys(raw, allowed=set(specs), name="web_runtime") + parsed = { + key: _integer( + raw.get(key, default), f"web_runtime.{key}", minimum=low, maximum=high + ) + for key, (default, low, high) in specs.items() + } + if parsed["max_prepared_runs_total"] < parsed["max_prepared_runs_per_subject"]: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "prepared run limits are inconsistent" + ) + return WebRuntimePolicy(**parsed) + + +def _parse_v3_config( + raw: Mapping[str, Any], *, require_evidence: bool +) -> EvoModelConfig: + _strict_keys( + raw, + allowed={ + "schema_version", + "config_revision", + "config_identity_key_id", + "runtime_defaults", + "providers", + "aliases", + "purpose_defaults", + "purpose_routes", + "purpose_call_limits", + "health_policy", + "web_runtime", + "capability_evidence", + }, + name="model_routes", + ) + revision = _integer(raw.get("config_revision"), "config_revision", minimum=1) + identity_key_id = _text(raw.get("config_identity_key_id"), "config_identity_key_id") + runtime_defaults = _parse_v3_runtime_defaults(raw.get("runtime_defaults")) + registry = get_adapter_registry() + providers_raw = _sequence(raw.get("providers"), "providers", allow_empty=False) + providers: dict[str, ProviderConfig] = {} + for provider_index, provider_value in enumerate(providers_raw): + name = f"providers[{provider_index}]" + item = _mapping(provider_value, name) + _strict_keys( + item, + allowed={ + "provider_id", + "display_name", + "adapter_id", + "adapter_revision", + "wire_protocol", + "enabled", + "connection", + "defaults", + "models", + }, + name=name, + ) + provider_id = _text(item.get("provider_id"), f"{name}.provider_id") + if provider_id in providers: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "duplicate provider_id" + ) + adapter_id = _text(item.get("adapter_id"), f"{name}.adapter_id") + exact_revision = _text(item.get("adapter_revision"), f"{name}.adapter_revision") + if exact_revision in {"latest", "current"}: + raise EvoRuntimeError("MODEL_ADAPTER_UNAVAILABLE") + registration = registry.get(adapter_id, exact_revision) + if registration.lifecycle == "blocked": + raise EvoRuntimeError("MODEL_ADAPTER_BLOCKED") + wire_protocol = _text(item.get("wire_protocol"), f"{name}.wire_protocol") + if wire_protocol not in registration.supported_wire_protocols: + raise EvoRuntimeError("MODEL_WIRE_PROTOCOL_UNSUPPORTED") + connection = _mapping(item.get("connection"), f"{name}.connection") + _strict_keys( + connection, + allowed={"base_url", "credential_ref"}, + name=f"{name}.connection", + ) + base_url = _normalize_base_url( + connection.get("base_url"), f"{name}.connection.base_url" + ) + legacy_ref = connection.get("credential_ref") + if legacy_ref is None: + credential_ref = f"provider://{provider_id}" + secret_version = 0 + else: + credential_ref = _text(legacy_ref, f"{name}.connection.credential_ref") + expected_prefix = f"secret://model-providers/{provider_id}#" + if not credential_ref.startswith(expected_prefix): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "credential_ref must be scoped to its provider", + ) + try: + secret_version = int(credential_ref.rsplit("#", 1)[1]) + except ValueError as exc: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "credential_ref version is invalid", + ) from exc + if secret_version < 1: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "credential_ref version is invalid", + ) + provider_defaults = _parse_v3_provider_defaults( + item.get("defaults"), runtime_defaults + ) + models: dict[str, ModelConfig] = {} + for model_index, model_value in enumerate( + _sequence(item.get("models"), f"{name}.models", allow_empty=False) + ): + model_name = f"{name}.models[{model_index}]" + model_raw = _mapping(model_value, model_name) + _strict_keys( + model_raw, + allowed={ + "model_key", + "provider_model_id", + "version_policy", + "resolved_model_revision", + "display_name", + "description", + "enabled", + "tags", + "invocation", + "capabilities", + "limits", + "parameters", + "access", + "billing", + }, + name=model_name, + ) + model_key = _text(model_raw.get("model_key"), f"{model_name}.model_key") + if model_key in models: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "duplicate model_key" + ) + provider_model_id = _text( + model_raw.get("provider_model_id"), f"{model_name}.provider_model_id" + ) + policy = _text( + model_raw.get("version_policy"), f"{model_name}.version_policy" + ) + if policy not in {"pinned", "rolling"}: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "invalid version_policy" + ) + resolved = model_raw.get("resolved_model_revision") + if resolved is not None: + resolved = _text(resolved, f"{model_name}.resolved_model_revision") + if policy == "pinned" and resolved is None: + resolved = provider_model_id + capabilities_raw = _mapping( + model_raw.get("capabilities"), f"{model_name}.capabilities" + ) + _strict_keys( + capabilities_raw, + allowed=set(_V3_CAPABILITIES), + name=f"{model_name}.capabilities", + ) + capabilities = { + key: _bool( + capabilities_raw.get(key, key == "text"), + f"{model_name}.capabilities.{key}", + ) + for key in _V3_CAPABILITIES + } + if not capabilities["text"]: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "text capability is required" + ) + invocation = _mapping( + model_raw.get("invocation"), f"{model_name}.invocation" + ) + _strict_keys( + invocation, + allowed={"api_mode", "tool_call_transport"}, + name=f"{model_name}.invocation", + ) + api_mode = _text( + invocation.get("api_mode"), f"{model_name}.invocation.api_mode" + ) + transport = _text( + invocation.get("tool_call_transport"), + f"{model_name}.invocation.tool_call_transport", + ) + if transport not in {"native", "prompt", "disabled"} or ( + capabilities["tools"] and transport != "native" + ): + raise EvoRuntimeError("MODEL_TOOL_TRANSPORT_UNSUPPORTED") + if not capabilities["tools"]: + # V4 save validation rejects this mismatch for new revisions. + # Keep old signed projections runnable while treating their + # stale native/prompt declaration as the only safe transport. + transport = "disabled" + limits = _mapping(model_raw.get("limits"), f"{model_name}.limits") + _strict_keys( + limits, + allowed={ + "context_tokens", + "max_output_tokens", + "max_inflight_requests", + }, + name=f"{model_name}.limits", + ) + context_tokens = limits.get("context_tokens") + max_output_tokens = limits.get("max_output_tokens") + if context_tokens is not None: + context_tokens = _integer( + context_tokens, f"{model_name}.limits.context_tokens", minimum=1 + ) + if max_output_tokens is not None: + max_output_tokens = _integer( + max_output_tokens, + f"{model_name}.limits.max_output_tokens", + minimum=1, + ) + descriptor = registration.resolve_model_descriptor( + provider_model_id, + api_mode, + context_tokens=context_tokens, + max_output_tokens=max_output_tokens, + declared_capabilities=capabilities, + ) + max_inflight = limits.get("max_inflight_requests") + if max_inflight is not None: + max_inflight = _integer( + max_inflight, + f"{model_name}.limits.max_inflight_requests", + minimum=1, + maximum=provider_defaults["max_inflight_requests"], + ) + parameters = _mapping( + model_raw.get("parameters") or {}, f"{model_name}.parameters" + ) + _strict_keys( + parameters, + allowed={ + "defaults", + "purpose_overrides", + "user_options", + "constraints", + "reasoning_policy", + }, + name=f"{model_name}.parameters", + ) + defaults = dict( + _mapping( + parameters.get("defaults") or {}, + f"{model_name}.parameters.defaults", + ) + ) + registration.validate_parameters( + defaults, path=f"{model_name}.parameters.defaults" + ) + overrides_raw = _mapping( + parameters.get("purpose_overrides") or {}, + f"{model_name}.parameters.purpose_overrides", + ) + if not set(overrides_raw) <= set(_PURPOSES): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "unknown purpose override" + ) + overrides = {} + for purpose, values in overrides_raw.items(): + parsed_values = dict( + _mapping( + values, f"{model_name}.parameters.purpose_overrides.{purpose}" + ) + ) + registration.validate_parameters( + parsed_values, + path=f"{model_name}.parameters.purpose_overrides.{purpose}", + ) + overrides[purpose] = parsed_values + user_options_raw = _mapping( + parameters.get("user_options") or {}, + f"{model_name}.parameters.user_options", + ) + unknown_options = set(user_options_raw) - set( + registration.all_parameter_schema + ) + if unknown_options: + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", + "user_options contains an unsupported parameter", + ) + user_options = {} + for key, value in user_options_raw.items(): + option_path = f"{model_name}.parameters.user_options.{key}" + option = dict(_mapping(value, option_path)) + _strict_keys( + option, + allowed={ + "default", + "applies_to", + "minimum", + "maximum", + "minimum_exclusive", + "maximum_exclusive", + "choices", + }, + name=option_path, + ) + rule = registration.all_parameter_schema[key] + option = dict(_normalize_v3_user_option(option, rule, option_path)) + applies_to = _string_set( + option.get("applies_to") or ["main_agent"], + f"{option_path}.applies_to", + allowed=frozenset(_PURPOSES), + ) + user_options[key] = { + **option, + "type": rule.kind, + "applies_to": applies_to, + } + constraints = tuple( + dict(_mapping(value, f"{model_name}.parameters.constraints")) + for value in _sequence( + parameters.get("constraints") or [], + f"{model_name}.parameters.constraints", + ) + ) + reasoning_policy = _mapping( + parameters.get("reasoning_policy") or {}, + f"{model_name}.parameters.reasoning_policy", + ) + _strict_keys( + reasoning_policy, + allowed={"mode", "allowed_efforts", "default_effort"}, + name=f"{model_name}.parameters.reasoning_policy", + ) + access = _parse_v3_access(model_raw.get("access"), f"{model_name}.access") + quote = _parse_v3_billing( + model_raw.get("billing"), + f"{model_name}.billing", + require_evidence=require_evidence, + ) + # A declared model capability is usable only when the selected + # Adapter has a concrete wire-level reasoning implementation. + supports_reasoning = ( + capabilities["thinking"] and registration.reasoning_mode != "none" + ) + if reasoning_policy and not supports_reasoning: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + f"{model_name}.parameters.reasoning_policy requires reasoning capability", + ) + reasoning_mode = registration.reasoning_mode + allowed_reasoning_efforts: tuple[str, ...] = () + default_reasoning_effort: str | None = None + if supports_reasoning: + reasoning_mode = str( + reasoning_policy.get("mode") or registration.reasoning_mode + ) + if reasoning_mode not in {"boolean", "effort"}: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + f"{model_name}.parameters.reasoning_policy.mode is invalid", + ) + allowed_reasoning_efforts = _string_set( + reasoning_policy.get("allowed_efforts") + or ["low", "medium", "high"], + f"{model_name}.parameters.reasoning_policy.allowed_efforts", + allowed=_REASONING_EFFORTS - {"disabled"}, + ) + default_reasoning_effort = str( + reasoning_policy.get("default_effort") or "high" + ) + if default_reasoning_effort not in allowed_reasoning_efforts: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + f"{model_name}.parameters.reasoning_policy.default_effort is not allowed", + ) + registration.validate_parameters( + {"reasoning": default_reasoning_effort}, + path=f"{model_name}.parameters.reasoning_policy", + ) + constraints = validate_parameter_constraints( + constraints, + allowed_names=set(user_options) + | ({"reasoning"} if supports_reasoning else set()), + ) + models[model_key] = ModelConfig( + provider_model_id, + defaults, + capabilities["vision"], + supports_reasoning, + allowed_reasoning_efforts, + descriptor.context_tokens, + descriptor.max_output_tokens, + reasoning_mode, + {"reasoning": default_reasoning_effort} + if default_reasoning_effort + else {}, + {"reasoning": "off"}, + (), + tuple(access["roles"]), + quote, + model_key=model_key, + display_name=_text( + model_raw.get("display_name"), f"{model_name}.display_name" + ), + description=str(model_raw.get("description") or ""), + tags=_string_set(model_raw.get("tags") or [], f"{model_name}.tags"), + enabled=_bool(model_raw.get("enabled"), f"{model_name}.enabled"), + version_policy=policy, + resolved_model_revision=resolved, + reproducible=policy == "pinned", + capabilities=capabilities, + purpose_overrides=overrides, + user_options=user_options, + parameter_constraints=constraints, + access=access, + max_inflight_requests=max_inflight, + descriptor_parameters=descriptor.parameters, + ) + endpoint = EndpointConfig( + provider_id, + base_url, + SecretReference(credential_ref, secret_version), + {}, + {}, + {}, + ) + providers[provider_id] = ProviderConfig( + provider_id, + adapter_id, + {}, + {provider_id: endpoint}, + models, + display_name=_text(item.get("display_name"), f"{name}.display_name"), + adapter_id=adapter_id, + adapter_revision=exact_revision, + wire_protocol=wire_protocol, + enabled=_bool(item.get("enabled"), f"{name}.enabled"), + connection_defaults=provider_defaults, + implementation_fingerprint=registration.implementation_fingerprint, + ) + aliases_raw = _sequence(raw.get("aliases"), "aliases", allow_empty=False) + aliases: dict[str, AliasConfig] = {} + selectors: dict[str, RouteSelector] = {} + selectable: dict[str, str] = {} + for index, alias_value in enumerate(aliases_raw): + name = f"aliases[{index}]" + item = _mapping(alias_value, name) + _strict_keys( + item, + allowed={ + "alias", + "display_name", + "provider_ref", + "model_ref", + "enabled", + "access", + "defaults", + }, + name=name, + ) + alias = _text(item.get("alias"), f"{name}.alias") + if alias in aliases: + raise EvoRuntimeError("LLM_ROUTE_CONFIGURATION_REQUIRED", "duplicate alias") + provider_ref = _text(item.get("provider_ref"), f"{name}.provider_ref") + model_ref = _text(item.get("model_ref"), f"{name}.model_ref") + if ( + provider_ref not in providers + or model_ref not in providers[provider_ref].models + ): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "alias target does not exist" + ) + access = _parse_v3_access(item.get("access"), f"{name}.access") + defaults = dict(_mapping(item.get("defaults") or {}, f"{name}.defaults")) + registration = registry.get( + providers[provider_ref].adapter_id, providers[provider_ref].adapter_revision + ) + registration.validate_parameters(defaults, path=f"{name}.defaults") + model = providers[provider_ref].models[model_ref] + if not set(defaults) <= set(model.user_options): + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", "alias defaults must use model user_options" + ) + revision_key = hashlib.sha256( + str(model.resolved_model_revision or model.model_id).encode() + ).hexdigest()[:12] + identity_selector = f"direct:{provider_ref}:{model_ref}:{revision_key}:{_find_model_api_mode(raw, provider_ref, model_ref)}:{_find_model_transport(raw, provider_ref, model_ref)}" + internal_selector = f"{identity_selector}:alias:{alias}" + selector = RouteSelector( + internal_selector, + provider_ref, + provider_ref, + None, + model_ref, + _find_model_api_mode(raw, provider_ref, model_ref), + _find_model_transport(raw, provider_ref, model_ref), + identity_selector_id=identity_selector, + alias=alias, + alias_defaults=defaults, + ) + selectors[internal_selector] = selector + enabled = _bool(item.get("enabled"), f"{name}.enabled") + aliases[alias] = AliasConfig( + alias, + _text(item.get("display_name"), f"{name}.display_name"), + provider_ref, + model_ref, + enabled, + access, + defaults, + ) + if enabled: + selectable[alias] = internal_selector + purpose_defaults_raw = _mapping( + raw.get("purpose_defaults") or {}, "purpose_defaults" + ) + if set(purpose_defaults_raw) != set(_PURPOSES): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "purpose_defaults must define all purposes", + ) + purpose_defaults = { + key: dict(_mapping(value, f"purpose_defaults.{key}")) + for key, value in purpose_defaults_raw.items() + } + purpose_routes = _mapping(raw.get("purpose_routes"), "purpose_routes") + if set(purpose_routes) != set(_PURPOSES): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "purpose_routes must define all purposes", + ) + main_route = _mapping(purpose_routes["main_agent"], "purpose_routes.main_agent") + _strict_keys( + main_route, allowed={"default_alias"}, name="purpose_routes.main_agent" + ) + default_alias = _text( + main_route.get("default_alias"), "purpose_routes.main_agent.default_alias" + ) + if default_alias not in selectable: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "default alias is unavailable" + ) + purpose_selector_ids: dict[str, str] = {} + for purpose in ("tool_selector", "deepagents_summarizer"): + route = purpose_routes[purpose] + if route != "inherit_main" and not ( + isinstance(route, Mapping) + and set(route) == {"default_alias"} + and route["default_alias"] in selectable + ): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{purpose} route is invalid" + ) + if isinstance(route, Mapping): + purpose_selector_ids[purpose] = selectable[str(route["default_alias"])] + title_route = _mapping(purpose_routes["title"], "purpose_routes.title") + _strict_keys(title_route, allowed={"default_alias"}, name="purpose_routes.title") + title_alias = _text( + title_route.get("default_alias"), "purpose_routes.title.default_alias" + ) + if title_alias not in selectable: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "title alias is unavailable" + ) + limits_raw = _mapping(raw.get("purpose_call_limits"), "purpose_call_limits") + if set(limits_raw) != set(_PURPOSES): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "purpose_call_limits must define all purposes", + ) + purpose_limits = {} + for purpose, value in limits_raw.items(): + item = _mapping(value, f"purpose_call_limits.{purpose}") + _strict_keys( + item, + allowed={"max_output_tokens", "max_attempts_per_run"}, + name=f"purpose_call_limits.{purpose}", + ) + purpose_limits[purpose] = PurposeCallLimit( + _integer( + item.get("max_attempts_per_run"), + f"purpose_call_limits.{purpose}.max_attempts_per_run", + minimum=1, + maximum=8, + ), + ) + health = _mapping(raw.get("health_policy") or {}, "health_policy") + _strict_keys( + health, allowed={"provider_connection", "model_route"}, name="health_policy" + ) + provider_health = _parse_v3_health( + health.get("provider_connection"), "health_policy.provider_connection" + ) + model_health = _parse_v3_health( + health.get("model_route"), "health_policy.model_route" + ) + config = EvoModelConfig( + revision, + identity_key_id, + runtime_defaults, + purpose_defaults, + providers, + {}, + model_health, + selectors, + PurposeRoutes(default_alias, selectable), + selectable[title_alias], + purpose_limits, + _parse_v3_web_runtime(raw.get("web_runtime")), + {}, + {}, + dict(raw), + schema_version=3, + aliases=aliases, + provider_health=provider_health, + adapter_registry_revision=registry.registry_revision, + purpose_selector_ids=purpose_selector_ids, + ) + config._validate(require_evidence=require_evidence) + return config + + +def _find_v3_model( + raw: Mapping[str, Any], provider_id: str, model_key: str +) -> Mapping[str, Any]: + for provider in raw.get("providers") or []: + if provider.get("provider_id") == provider_id: + for model in provider.get("models") or []: + if model.get("model_key") == model_key: + return model + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "model target does not exist" + ) + + +def _find_model_api_mode( + raw: Mapping[str, Any], provider_id: str, model_key: str +) -> str: + return _text( + _mapping( + _find_v3_model(raw, provider_id, model_key).get("invocation"), "invocation" + ).get("api_mode"), + "api_mode", + ) + + +def _find_model_transport( + raw: Mapping[str, Any], provider_id: str, model_key: str +) -> str: + return _text( + _mapping( + _find_v3_model(raw, provider_id, model_key).get("invocation"), "invocation" + ).get("tool_call_transport"), + "tool_call_transport", + ) + + +def _parse_v3_evidence( + value: Any, config: EvoModelConfig +) -> Mapping[str, CapabilityEvidence]: + result: dict[str, CapabilityEvidence] = {} + for index, evidence_value in enumerate(_sequence(value, "capability_evidence")): + name = f"capability_evidence[{index}]" + raw = _mapping(evidence_value, name) + allowed = { + "provider_ref", + "model_ref", + "adapter_id", + "adapter_revision", + "implementation_fingerprint", + "wire_protocol", + "provider_model_id", + "resolved_model_revision", + "version_policy", + "reproducible", + "api_mode", + "tool_call_transport", + "base_url_fingerprint", + "secret_version", + "route_semantics_hash", + "fixture_digest", + "verified_at", + "evidence_expires_at", + "results", + } + _strict_keys(raw, allowed=allowed, name=name) + provider_ref = _text(raw.get("provider_ref"), f"{name}.provider_ref") + model_ref = _text(raw.get("model_ref"), f"{name}.model_ref") + matches = [ + selector + for selector in config.route_selectors.values() + if selector.provider == provider_ref and selector.model == model_ref + ] + if not matches: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "evidence target does not exist" + ) + route = config.concrete_routes(matches[0].selector_id)[0] + results_raw = _mapping(raw.get("results"), f"{name}.results") + valid_statuses = {"supported", "failed", "not_verified", "not_declared"} + results = {} + for key, status in results_raw.items(): + status_text = _text(status, f"{name}.results.{key}") + if status_text not in valid_statuses: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "invalid evidence status" + ) + results[key] = status_text + model = config.providers[provider_ref].models[model_ref] + evidence = CapabilityEvidence( + route, + results.get("connectivity", "failed"), + results.get("tools", "not_declared"), + _text(raw.get("route_semantics_hash"), f"{name}.route_semantics_hash"), + _text(raw.get("base_url_fingerprint"), f"{name}.base_url_fingerprint"), + config.config_identity_key_id, + _text(raw.get("adapter_revision"), f"{name}.adapter_revision"), + _text(raw.get("fixture_digest"), f"{name}.fixture_digest"), + _text(raw.get("verified_at"), f"{name}.verified_at"), + results=results, + evidence_expires_at=_text( + raw.get("evidence_expires_at"), f"{name}.evidence_expires_at" + ), + implementation_fingerprint=_text( + raw.get("implementation_fingerprint"), + f"{name}.implementation_fingerprint", + ), + resolved_model_revision=str(raw.get("resolved_model_revision") or "") + or None, + reproducible=_bool(raw.get("reproducible"), f"{name}.reproducible"), + ) + if ( + raw.get("provider_model_id") != model.model_id + or raw.get("adapter_id") != config.providers[provider_ref].adapter_id + ): + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + result[route.key()] = evidence + return result + + +def _parse_secret(value: Any, name: str) -> SecretReference: + raw = _mapping(value, name) + _strict_keys(raw, allowed={"ref", "revision"}, name=name) + ref = _text(raw.get("ref"), f"{name}.ref") + revision = _integer(raw.get("revision"), f"{name}.revision", minimum=1) + if ref.startswith("env://"): + if not _ENV_RE.fullmatch(ref[6:]): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name}.ref is invalid" + ) + elif ref.startswith("secret://"): + if "#" not in ref[9:] or not ref.rsplit("#", 1)[1]: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name}.ref needs a version" + ) + else: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", f"{name}.ref scheme is unsupported" + ) + return SecretReference(ref, revision) + + +def _normalize_base_url(value: Any, name: str) -> str: + return str(value or "").strip().rstrip("/") + + +def _parse_providers(value: Any) -> Mapping[str, ProviderConfig]: + raw = _mapping(value, "providers") + if not raw: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "providers must not be empty" + ) + result: dict[str, ProviderConfig] = {} + for provider_key, provider_value in raw.items(): + key = _text(provider_key, "provider key") + provider = _mapping(provider_value, f"providers.{key}") + _strict_keys( + provider, + allowed={"protocol", "params", "endpoints", "models"}, + name=f"providers.{key}", + ) + protocol = _text(provider.get("protocol"), f"providers.{key}.protocol") + endpoints: dict[str, EndpointConfig] = {} + for index, value_item in enumerate( + _sequence(provider.get("endpoints"), "endpoints", allow_empty=False) + ): + item = _mapping(value_item, f"providers.{key}.endpoints[{index}]") + _strict_keys( + item, + allowed={ + "name", + "base_url", + "auth", + "headers", + "header_refs", + "params", + }, + name=f"providers.{key}.endpoints[{index}]", + ) + endpoint_name = _text(item.get("name"), "endpoint.name") + if endpoint_name in endpoints: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "duplicate endpoint" + ) + headers = dict(_mapping(item.get("headers") or {}, "endpoint.headers")) + header_refs_raw = _mapping( + item.get("header_refs") or {}, "endpoint.header_refs" + ) + parsed_headers: dict[str, str] = {} + for header, header_value in headers.items(): + if ( + not _HEADER_RE.fullmatch(header) + or header.lower() in _SENSITIVE_HEADERS + ): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "invalid static header" + ) + text = _text(header_value, f"headers.{header}") + if any(char in text for char in "\r\n\0"): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "invalid header value" + ) + parsed_headers[header] = text + parsed_header_refs: dict[str, SecretReference] = {} + for header, reference in header_refs_raw.items(): + if not _HEADER_RE.fullmatch(header): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "invalid secret header" + ) + if header.lower() in {name.lower() for name in parsed_headers}: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "duplicate header" + ) + parsed_header_refs[header] = _parse_secret( + reference, f"header_refs.{header}" + ) + endpoints[endpoint_name] = EndpointConfig( + endpoint_name, + _normalize_base_url(item.get("base_url"), "endpoint.base_url"), + _parse_secret(item.get("auth"), "endpoint.auth"), + parsed_headers, + parsed_header_refs, + dict(_mapping(item.get("params") or {}, "endpoint.params")), + ) + models: dict[str, ModelConfig] = {} + for index, value_item in enumerate( + _sequence(provider.get("models"), "models", allow_empty=False) + ): + item = _mapping(value_item, f"providers.{key}.models[{index}]") + _strict_keys( + item, + allowed={ + "id", + "params", + "supports_vision", + "supports_reasoning", + "allowed_reasoning_efforts", + "context_window", + "max_output_tokens", + "reasoning", + "access", + "billing", + }, + name=f"providers.{key}.models[{index}]", + ) + model_id = _text(item.get("id"), "model.id") + if model_id in models: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "duplicate model id" + ) + access = _mapping(item.get("access"), "model.access") + _strict_keys( + access, allowed={"allowed_plans", "allowed_roles"}, name="model.access" + ) + billing = _mapping(item.get("billing"), "model.billing") + _strict_keys( + billing, + allowed={ + "sku", + "pricing_revision", + "currency", + "unit_scale", + "input_microunits_per_million", + "output_microunits_per_million", + "cached_microunits_per_million", + "multiplier", + }, + name="model.billing", + ) + quote_payload = { + "billing_sku": _text(billing.get("sku"), "billing.sku"), + "pricing_revision": _text( + billing.get("pricing_revision"), "billing.pricing_revision" + ), + "currency": _text(billing.get("currency"), "billing.currency"), + "unit_scale": _integer( + billing.get("unit_scale"), "billing.unit_scale", minimum=1 + ), + "input_microunits_per_million": _integer( + billing.get("input_microunits_per_million"), "billing.input" + ), + "cached_input_microunits_per_million": _integer( + billing.get("cached_microunits_per_million"), "billing.cached" + ), + "output_microunits_per_million": _integer( + billing.get("output_microunits_per_million"), "billing.output" + ), + "multiplier": _positive_decimal( + billing.get("multiplier"), "billing.multiplier" + ), + } + if ( + quote_payload["currency"] != "CNY" + or quote_payload["unit_scale"] != 1_000_000 + ): + raise EvoRuntimeError("PRICING_DIMENSION_UNSUPPORTED") + quote = PricingQuote(**quote_payload, quote_id=sha256_id(quote_payload)) + supports_reasoning = _bool( + item.get("supports_reasoning"), "model.supports_reasoning" + ) + efforts = _string_set( + item.get("allowed_reasoning_efforts"), + "model.allowed_reasoning_efforts", + allowed=_REASONING_EFFORTS - {"disabled"}, + ) + if bool(efforts) != supports_reasoning: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "reasoning capability is inconsistent", + ) + reasoning = _parse_model_reasoning( + item.get("reasoning"), supports_reasoning=supports_reasoning + ) + context_window = _integer( + item.get("context_window"), "model.context_window", minimum=1 + ) + max_output_tokens = _integer( + item.get("max_output_tokens", context_window), + "model.max_output_tokens", + minimum=1, + maximum=context_window, + ) + models[model_id] = ModelConfig( + model_id, + dict(_mapping(item.get("params") or {}, "model.params")), + _bool(item.get("supports_vision"), "model.supports_vision"), + supports_reasoning, + efforts, + context_window, + max_output_tokens, + reasoning["mode"], + reasoning["enabled_params"], + reasoning["disabled_params"], + _string_set(access.get("allowed_plans"), "allowed_plans"), + _string_set(access.get("allowed_roles"), "allowed_roles"), + quote, + ) + result[key] = ProviderConfig( + key, + protocol, + dict(_mapping(provider.get("params") or {}, "provider.params")), + endpoints, + models, + ) + return result + + +def _parse_model_reasoning( + value: Any, *, supports_reasoning: bool +) -> Mapping[str, Any]: + """Parse provider-specific reasoning controls without exposing them to users.""" + + if value is None: + return { + "mode": "effort" if supports_reasoning else "boolean", + "enabled_params": {}, + "disabled_params": {}, + } + raw = _mapping(value, "model.reasoning") + _strict_keys( + raw, + allowed={"mode", "enabled_params", "disabled_params"}, + name="model.reasoning", + ) + mode = _text(raw.get("mode"), "model.reasoning.mode") + if mode not in _REASONING_MODES: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "invalid reasoning mode" + ) + if not supports_reasoning: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "non-reasoning model cannot define reasoning controls", + ) + return { + "mode": mode, + "enabled_params": dict( + _mapping(raw.get("enabled_params") or {}, "model.reasoning.enabled_params") + ), + "disabled_params": dict( + _mapping( + raw.get("disabled_params") or {}, "model.reasoning.disabled_params" + ) + ), + } + + +def _parse_endpoint_pools( + value: Any, providers: Mapping[str, ProviderConfig] +) -> Mapping[str, EndpointPool]: + raw = _mapping(value, "endpoint_pools") + result: dict[str, EndpointPool] = {} + for pool_key, pool_value in raw.items(): + pool_id = _text(pool_key, "pool id") + pool = _mapping(pool_value, f"endpoint_pools.{pool_id}") + _strict_keys( + pool, + allowed={"provider", "strategy", "endpoints"}, + name=f"endpoint_pools.{pool_id}", + ) + provider = _text(pool.get("provider"), "pool.provider") + if provider not in providers: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "unknown pool provider" + ) + strategy = _text(pool.get("strategy"), "pool.strategy") + if strategy != "smooth_weighted_round_robin": + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "unsupported pool strategy" + ) + members: list[PoolEndpoint] = [] + seen: set[str] = set() + for item_value in _sequence( + pool.get("endpoints"), "pool.endpoints", allow_empty=False + ): + item = _mapping(item_value, "pool endpoint") + _strict_keys(item, allowed={"name", "weight"}, name="pool endpoint") + name = _text(item.get("name"), "pool endpoint.name") + if name in seen or name not in providers[provider].endpoints: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "invalid pool endpoint" + ) + seen.add(name) + members.append( + PoolEndpoint( + name, + _integer(item.get("weight"), "pool endpoint.weight", minimum=1), + ) + ) + result[pool_id] = EndpointPool(pool_id, provider, strategy, tuple(members)) + return result + + +def _parse_selectors( + value: Any, + providers: Mapping[str, ProviderConfig], + pools: Mapping[str, EndpointPool], +) -> Mapping[str, RouteSelector]: + raw = _mapping(value, "route_selectors") + result: dict[str, RouteSelector] = {} + for selector_key, selector_value in raw.items(): + selector_id = _text(selector_key, "selector id") + item = _mapping(selector_value, f"route_selectors.{selector_id}") + _strict_keys( + item, + allowed={ + "provider", + "endpoint", + "endpoint_pool", + "model", + "api_mode", + "tool_call_transport", + }, + name=f"route_selectors.{selector_id}", + ) + provider_id = _text(item.get("provider"), "selector.provider") + provider = providers.get(provider_id) + if provider is None: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "unknown selector provider" + ) + endpoint = item.get("endpoint") + endpoint_pool = item.get("endpoint_pool") + if (endpoint is None) == (endpoint_pool is None): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "selector requires exactly one endpoint source", + ) + endpoint_name = ( + _text(endpoint, "selector.endpoint") if endpoint is not None else None + ) + pool_id = ( + _text(endpoint_pool, "selector.endpoint_pool") + if endpoint_pool is not None + else None + ) + if endpoint_name is not None and endpoint_name not in provider.endpoints: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "unknown selector endpoint" + ) + if pool_id is not None and ( + pool_id not in pools or pools[pool_id].provider != provider_id + ): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "selector pool provider mismatch" + ) + model_id = _text(item.get("model"), "selector.model") + if model_id not in provider.models: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "unknown selector model" + ) + api_mode = _text(item.get("api_mode"), "selector.api_mode") + allowed = _PARAM_ALLOWLIST.get((provider.protocol, api_mode)) + if allowed is None: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "provider/api mode has no adapter" + ) + transport = _text( + item.get("tool_call_transport"), "selector.tool_call_transport" + ) + if transport != "non_streaming": + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "Web tool routes must be non_streaming", + ) + _validate_params( + provider.params, name=f"providers.{provider_id}.params", allowed=allowed + ) + for concrete_endpoint in provider.endpoints.values(): + _validate_params( + concrete_endpoint.params, name="endpoint.params", allowed=allowed + ) + _validate_params( + provider.models[model_id].params, name="model.params", allowed=allowed + ) + result[selector_id] = RouteSelector( + selector_id, + provider_id, + endpoint_name, + pool_id, + model_id, + api_mode, + transport, + ) + if not result: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "route_selectors must not be empty" + ) + return result + + +def _parse_purpose_defaults(value: Any) -> Mapping[str, Mapping[str, str]]: + raw = _mapping(value, "purpose_defaults") + if set(raw) != set(_PURPOSES): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + "purpose_defaults must define all purposes", + ) + result = {} + for purpose in _PURPOSES: + item = _mapping(raw[purpose], f"purpose_defaults.{purpose}") + _strict_keys( + item, allowed={"reasoning_effort"}, name=f"purpose_defaults.{purpose}" + ) + effort = _text(item.get("reasoning_effort"), "reasoning_effort") + if effort not in _REASONING_EFFORTS: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "invalid reasoning effort" + ) + result[purpose] = {"reasoning_effort": effort} + return result + + +def _parse_purpose_routes( + value: Any, selectors: Mapping[str, RouteSelector] +) -> tuple[PurposeRoutes, str]: + raw = _mapping(value, "purpose_routes") + _strict_keys(raw, allowed={"main_agent", "title"}, name="purpose_routes") + main = _mapping(raw.get("main_agent"), "purpose_routes.main_agent") + _strict_keys( + main, allowed={"default_alias", "selectable"}, name="purpose_routes.main_agent" + ) + selectable_raw = _mapping(main.get("selectable"), "main_agent.selectable") + selectable = { + _text(alias, "model alias"): _text(selector_id, "selector id") + for alias, selector_id in selectable_raw.items() + } + if not selectable or any( + selector not in selectors for selector in selectable.values() + ): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "invalid selectable route" + ) + default_alias = _text(main.get("default_alias"), "main_agent.default_alias") + if default_alias not in selectable: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "default alias is not selectable" + ) + title = _mapping(raw.get("title"), "purpose_routes.title") + _strict_keys(title, allowed={"default"}, name="purpose_routes.title") + title_selector = _text(title.get("default"), "title.default") + if title_selector not in selectors: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "unknown title selector" + ) + return PurposeRoutes(default_alias, selectable), title_selector + + +def _parse_call_limits(value: Any) -> Mapping[str, PurposeCallLimit]: + raw = _mapping(value, "purpose_call_limits") + if set(raw) != set(_PURPOSES): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "call limits must define all purposes" + ) + result = {} + for purpose in _PURPOSES: + item = _mapping(raw[purpose], f"purpose_call_limits.{purpose}") + _strict_keys( + item, + allowed={"max_output_tokens", "max_attempts_per_run"}, + name=f"purpose_call_limits.{purpose}", + ) + result[purpose] = PurposeCallLimit( + _integer( + item.get("max_attempts_per_run"), "max_attempts_per_run", minimum=1 + ), + ) + return result + + +def _parse_health(value: Any) -> RouteHealthPolicy: + raw = _mapping(value, "route_health") + _strict_keys( + raw, + allowed={ + "failure_threshold", + "cooldown_seconds", + "half_open_max_inflight", + "counted_error_codes", + "open_immediately_error_codes", + }, + name="route_health", + ) + return RouteHealthPolicy( + _integer(raw.get("failure_threshold"), "failure_threshold", minimum=1), + _integer(raw.get("cooldown_seconds"), "cooldown_seconds", minimum=1), + _integer( + raw.get("half_open_max_inflight"), "half_open_max_inflight", minimum=1 + ), + _string_set(raw.get("counted_error_codes"), "counted_error_codes"), + _string_set( + raw.get("open_immediately_error_codes"), "open_immediately_error_codes" + ), + ) + + +def _parse_web_runtime(value: Any) -> WebRuntimePolicy: + raw = _mapping(value, "web_runtime") + names = { + "title_start_timeout_seconds", + "prepare_ttl_seconds", + "turn_lease_grace_seconds", + "active_run_timeout_seconds", + "max_run_journal_events", + "max_run_journal_bytes", + "max_prepared_runs_per_subject", + "max_prepared_runs_total", + } + _strict_keys(raw, allowed=names, name="web_runtime") + values = { + name: _integer(raw.get(name), f"web_runtime.{name}", minimum=1) + for name in names + } + if ( + values["prepare_ttl_seconds"] > 120 + or values["title_start_timeout_seconds"] > 300 + ): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "runtime TTL is out of range" + ) + if values["max_prepared_runs_per_subject"] > values["max_prepared_runs_total"]: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "prepared run limits are inconsistent" + ) + return WebRuntimePolicy(**values) + + +def _parse_fallbacks( + value: Any, selectors: Mapping[str, RouteSelector] +) -> Mapping[str, tuple[str, ...]]: + result: dict[str, tuple[str, ...]] = {} + for item_value in _sequence(value, "tool_protocol_fallbacks"): + item = _mapping(item_value, "tool_protocol_fallback") + _strict_keys( + item, allowed={"primary", "fallbacks"}, name="tool_protocol_fallback" + ) + primary = _text(item.get("primary"), "fallback.primary") + fallbacks = tuple( + _text(entry, "fallback selector") + for entry in _sequence(item.get("fallbacks"), "fallbacks") + ) + if ( + primary in result + or primary not in selectors + or any(entry not in selectors for entry in fallbacks) + ): + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "invalid fallback selector" + ) + result[primary] = fallbacks + return result + + +def _route_from_evidence(value: Any, config: EvoModelConfig) -> RouteRef: + raw = _mapping(value, "evidence.route") + _strict_keys( + raw, + allowed={"provider", "endpoint", "model", "api_mode", "tool_call_transport"}, + name="evidence.route", + ) + matches = [ + route + for selector_id in config.route_selectors + for route in config.concrete_routes(selector_id) + if route.provider == raw.get("provider") + and route.endpoint == raw.get("endpoint") + and route.model == raw.get("model") + and route.api_mode == raw.get("api_mode") + and route.tool_call_transport == raw.get("tool_call_transport") + ] + if len(matches) != 1: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "evidence route is ambiguous or unknown" + ) + return matches[0] + + +def _parse_evidence( + value: Any, config: EvoModelConfig +) -> Mapping[str, CapabilityEvidence]: + result: dict[str, CapabilityEvidence] = {} + for item_value in _sequence(value, "capability_evidence"): + item = _mapping(item_value, "capability_evidence item") + _strict_keys( + item, + allowed={"route", "connectivity", "tool_capability", "probe"}, + name="capability_evidence item", + ) + route = _route_from_evidence(item.get("route"), config) + probe = _mapping(item.get("probe"), "capability_evidence.probe") + _strict_keys( + probe, + allowed={ + "route_semantics_hash", + "endpoint_fingerprint", + "config_identity_key_id", + "adapter_revision", + "fixture_digest", + "verified_at", + }, + name="capability_evidence.probe", + ) + evidence = CapabilityEvidence( + route, + _text(item.get("connectivity"), "connectivity"), + _text(item.get("tool_capability"), "tool_capability"), + _text(probe.get("route_semantics_hash"), "route_semantics_hash"), + _text(probe.get("endpoint_fingerprint"), "endpoint_fingerprint"), + _text(probe.get("config_identity_key_id"), "config_identity_key_id"), + _text(probe.get("adapter_revision"), "adapter_revision"), + _text(probe.get("fixture_digest"), "fixture_digest"), + _text(probe.get("verified_at"), "verified_at"), + ) + if route.key() in result: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "duplicate capability evidence" + ) + result[route.key()] = evidence + return result + + +def route_semantics_payload( + config: EvoModelConfig, route: RouteRef +) -> Mapping[str, Any]: + provider = config.providers[route.provider] + endpoint = provider.endpoints[route.endpoint] + model = provider.models[route.model] + if config.schema_version == 3: + return { + "schema_version": 3, + "config_revision": config.config_revision, + "provider_id": provider.key, + "adapter_id": provider.adapter_id, + "adapter_revision": provider.adapter_revision, + "implementation_fingerprint": provider.implementation_fingerprint, + "wire_protocol": provider.wire_protocol, + "base_url": endpoint.base_url, + "credential_ref": endpoint.auth.ref, + "model_key": route.model, + "provider_model_id": model.model_id, + "version_policy": model.version_policy, + "resolved_model_revision": model.resolved_model_revision, + "api_mode": route.api_mode, + "tool_call_transport": route.tool_call_transport, + # Capability evidence is reusable by every alias that targets this + # provider/model. Alias, purpose and user parameters belong to the + # invocation fingerprint produced by the runtime. + "params": dict(model.params), + } + allowed = _PARAM_ALLOWLIST[(provider.protocol, route.api_mode)] + return { + "provider": provider.key, + "protocol": provider.protocol, + "endpoint": endpoint.name, + "base_url": endpoint.base_url, + "auth": asdict(endpoint.auth), + "headers": dict(endpoint.headers), + "header_refs": { + key: asdict(value) for key, value in endpoint.header_refs.items() + }, + "model": model.model_id, + "model_capabilities": { + "context_window": model.context_window, + "max_output_tokens": model.max_output_tokens, + "supports_vision": model.supports_vision, + "supports_reasoning": model.supports_reasoning, + "allowed_reasoning_efforts": model.allowed_reasoning_efforts, + "reasoning_mode": model.reasoning_mode, + "reasoning_enabled_params": model.reasoning_enabled_params, + "reasoning_disabled_params": model.reasoning_disabled_params, + }, + "params": { + "provider": _validate_params( + provider.params, name="provider.params", allowed=allowed + ), + "endpoint": _validate_params( + endpoint.params, name="endpoint.params", allowed=allowed + ), + "model": _validate_params( + model.params, name="model.params", allowed=allowed + ), + }, + "api_mode": route.api_mode, + "tool_call_transport": route.tool_call_transport, + "adapter_revision": adapter_revision(provider.protocol, route.api_mode), + } + + +def adapter_revision(protocol: str, api_mode: str) -> str: + return ( + f"ai4sci-provider-adapter-v3:{protocol}:{api_mode}:" + "bounds=model-cap-v1:reasoning=declarative-v1:" + "media=descriptor-v1:margin=table-v1" + ) + + +def route_semantics_hash( + config: EvoModelConfig, route: RouteRef, identity_key: bytes +) -> str: + return hmac_id(identity_key, route_semantics_payload(config, route)) + + +def route_fingerprint( + config: EvoModelConfig, route: RouteRef, identity_key: bytes +) -> str: + return hmac_id( + identity_key, {"route": route_semantics_payload(config, route), "kind": "route"} + ) + + +def invocation_fingerprint( + config: EvoModelConfig, + route: RouteRef, + purpose: str, + final_params: Mapping[str, Any], + identity_key: bytes, +) -> str: + """Bind an invocation identity to alias, purpose, and final parameters.""" + + return hmac_id( + identity_key, + { + "route": route_semantics_payload(config, route), + "kind": "invocation", + "alias": route.selector_id, + "purpose": purpose, + "final_params": dict(final_params), + }, + ) + + +def endpoint_fingerprint( + config: EvoModelConfig, route: RouteRef, identity_key: bytes +) -> str: + endpoint = config.providers[route.provider].endpoints[route.endpoint] + if config.schema_version == 3: + provider = config.providers[route.provider] + return hmac_id( + identity_key, + { + "provider_id": route.provider, + "base_url": endpoint.base_url, + "credential_ref": endpoint.auth.ref, + "adapter_revision": provider.adapter_revision, + }, + ) + return hmac_id( + identity_key, {"provider": route.provider, "endpoint": asdict(endpoint)} + ) + + +_PROCESS_SECRET_FINGERPRINT_KEY = secrets.token_bytes(32) + + +def resolve_secret( + reference: SecretReference, *, secret_resolver: SecretResolver | None = None +) -> ResolvedSecret: + if secret_resolver is not None: + resolved = secret_resolver(reference) + elif reference.ref.startswith("env://"): + value = os.environ.get(reference.ref[6:], "") + if not value: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + fingerprint = hmac.new( + _PROCESS_SECRET_FINGERPRINT_KEY, value.encode(), hashlib.sha256 + ).hexdigest() + resolved = ResolvedSecret(value, reference.revision, None, fingerprint) + else: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + if resolved.declared_revision != reference.revision: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + if reference.ref.startswith("secret://"): + expected = reference.ref.rsplit("#", 1)[1] + if resolved.authoritative_version != expected: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + if any(char in resolved.value for char in "\r\n\0"): + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + return resolved + + +@dataclass(frozen=True, slots=True) +class ConfigRevision: + config_revision: int + + +@dataclass(frozen=True, slots=True) +class SaveModelConfigCommand: + expected_revision: int + payload: Mapping[str, Any] + admin_grant: AdminConfigGrant + operation_id: str = "" + + +class FileEvoModelConfigStore: + """Integrity-checked V2/V3 store with a durable committed-head journal.""" + + def __init__( + self, + path: Path | None = None, + *, + admin_verifier: AdminConfigGrantVerifier | None = None, + ops_path: Path | None = None, + ) -> None: + self.path = path or (get_config_dir() / "model_routes.yaml") + self._lock = FileLock(str(self.path) + ".lock") + self._admin_verifier = admin_verifier + self.ops_path = ops_path or self.path.with_name("model_config_ops.sqlite") + self._recovery_conflict = False + self._init_ops() + self._recover_preparing_operations() + + def _connect(self) -> sqlite3.Connection: + connection = sqlite3.connect(self.ops_path) + connection.row_factory = sqlite3.Row + connection.execute("PRAGMA foreign_keys=ON") + connection.execute("PRAGMA busy_timeout=5000") + connection.execute("PRAGMA synchronous=FULL") + connection.execute("PRAGMA journal_mode=WAL") + return connection + + def _init_ops(self) -> None: + self.ops_path.parent.mkdir(parents=True, exist_ok=True) + with self._connect() as connection: + connection.executescript( + """ + CREATE TABLE IF NOT EXISTS committed_head ( + singleton INTEGER PRIMARY KEY CHECK (singleton = 1), + config_revision INTEGER NOT NULL, + canonical_payload_hash TEXT NOT NULL, + config_identity_key_id TEXT NOT NULL, + commit_operation_id TEXT NOT NULL + ); + CREATE TABLE IF NOT EXISTS config_operations ( + subject_id TEXT NOT NULL, + action TEXT NOT NULL, + operation_id TEXT NOT NULL, + request_digest TEXT NOT NULL, + result_digest TEXT, + status TEXT NOT NULL, + expected_revision INTEGER, + target_revision INTEGER, + proposal_hash TEXT, + payload_hash TEXT, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL, + PRIMARY KEY (subject_id, action, operation_id) + ); + """ + ) + columns = { + str(row[1]) + for row in connection.execute("PRAGMA table_info(config_operations)") + } + if "canonical_payload" not in columns: + connection.execute( + "ALTER TABLE config_operations ADD COLUMN canonical_payload TEXT" + ) + try: + os.chmod(self.ops_path, stat.S_IRUSR | stat.S_IWUSR) + except OSError: + pass + + def load(self) -> EvoModelConfig: + with self._lock: + return self._load_locked() + + def _load_locked(self) -> EvoModelConfig: + if self._recovery_conflict: + raise EvoRuntimeError("CONFIG_INTEGRITY_MISMATCH") + with self._connect() as connection: + head = connection.execute( + "SELECT * FROM committed_head WHERE singleton = 1" + ).fetchone() + if not self.path.exists(): + if head is None: + raise EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", "model routes are unconfigured" + ) + raise EvoRuntimeError("CONFIG_INTEGRITY_MISMATCH") + if head is None: + raise EvoRuntimeError("CONFIG_INTEGRITY_MISMATCH") + config = EvoModelConfig.parse( + load_yaml_unique(self.path.read_text(encoding="utf-8")) + ) + payload_hash = sha256_id(config.raw) + if ( + config.config_revision != int(head["config_revision"]) + or config.config_identity_key_id != str(head["config_identity_key_id"]) + or payload_hash != str(head["canonical_payload_hash"]) + ): + raise EvoRuntimeError("CONFIG_INTEGRITY_MISMATCH") + return config + + def bootstrap_for_development( + self, payload: Mapping[str, Any], *, operation_id: str = "bootstrap" + ) -> ConfigRevision: + """Explicitly initialize an empty development store; never overwrites a head.""" + + with self._lock: + with self._connect() as connection: + if connection.execute( + "SELECT 1 FROM committed_head WHERE singleton = 1" + ).fetchone(): + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + if self.path.exists(): + raise EvoRuntimeError("CONFIG_INTEGRITY_MISMATCH") + candidate = dict(payload) + candidate["config_revision"] = 1 + config = EvoModelConfig.parse(candidate) + self._write_committed(config.raw, operation_id=operation_id) + return ConfigRevision(1) + + def save(self, command: SaveModelConfigCommand) -> ConfigRevision: + if self._admin_verifier is None: + raise EvoRuntimeError("ADMIN_CONFIG_FORBIDDEN") + self._admin_verifier.require_admin(command.admin_grant) + if command.admin_grant.action != "model_config:commit": + raise EvoRuntimeError("ADMIN_CONFIG_FORBIDDEN") + operation_id = command.operation_id or command.admin_grant.operation_id + if operation_id != command.admin_grant.operation_id: + raise EvoRuntimeError("IDEMPOTENCY_CONFLICT") + with self._lock: + current_revision = 0 + with self._connect() as connection: + head = connection.execute( + "SELECT * FROM committed_head WHERE singleton = 1" + ).fetchone() + if head is not None: + current_revision = int(head["config_revision"]) + if current_revision != command.expected_revision: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + candidate = dict(command.payload) + candidate["config_revision"] = current_revision + 1 + config = EvoModelConfig.parse(candidate) + request_digest = sha256_id( + { + "expected_revision": command.expected_revision, + "payload": command.payload, + } + ) + if request_digest != command.admin_grant.request_digest: + raise EvoRuntimeError("CONTRACT_SIGNATURE_INVALID") + self._write_committed(config.raw, operation_id=operation_id) + return ConfigRevision(config.config_revision) + + def current_revision(self) -> int: + with self._connect() as connection: + row = connection.execute( + "SELECT config_revision FROM committed_head WHERE singleton=1" + ).fetchone() + return int(row[0]) if row is not None else 0 + + def load_revision(self, revision: int) -> EvoModelConfig: + with self._connect() as connection: + row = connection.execute( + """SELECT canonical_payload FROM config_operations + WHERE action='model_config:commit' AND status='COMMITTED' + AND target_revision=? AND canonical_payload IS NOT NULL + ORDER BY updated_at DESC LIMIT 1""", + (revision,), + ).fetchone() + if row is None: + raise EvoRuntimeError("CONFIG_REVISION_NOT_FOUND") + return EvoModelConfig.parse(json.loads(str(row["canonical_payload"]))) + + def commit_validated( + self, + payload: Mapping[str, Any], + *, + expected_revision: int, + operation_id: str, + ) -> ConfigRevision: + """Commit a payload already authorized and evidenced by config admin.""" + + with self._lock: + target = expected_revision + 1 + candidate = dict(payload) + candidate["config_revision"] = target + config = EvoModelConfig.parse(candidate) + payload_hash = sha256_id(config.raw) + replay = self._committed_operation_revision( + operation_id, payload_hash=payload_hash, target_revision=target + ) + if replay is not None: + return ConfigRevision(replay) + current = self.current_revision() + if current != expected_revision: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + self._write_committed(config.raw, operation_id=operation_id) + return ConfigRevision(config.config_revision) + + def activate_revision( + self, + target_revision: int, + *, + expected_revision: int, + operation_id: str, + ) -> ConfigRevision: + """Atomically point the active head at an immutable committed revision.""" + + with self._lock: + if self.current_revision() != expected_revision: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + with self._connect() as connection: + replay = connection.execute( + """SELECT status, target_revision FROM config_operations + WHERE subject_id='admin' AND action='model_config:rollback' + AND operation_id=?""", + (operation_id,), + ).fetchone() + if replay is not None: + if int(replay["target_revision"] or 0) != target_revision: + raise EvoRuntimeError("IDEMPOTENCY_CONFLICT") + if str(replay["status"]) == "COMMITTED": + return ConfigRevision(target_revision) + source = connection.execute( + """SELECT canonical_payload FROM config_operations + WHERE action='model_config:commit' AND status='COMMITTED' + AND target_revision=? AND canonical_payload IS NOT NULL + ORDER BY updated_at DESC LIMIT 1""", + (target_revision,), + ).fetchone() + if source is None: + raise EvoRuntimeError("CONFIG_REVISION_NOT_FOUND") + payload = json.loads(str(source["canonical_payload"])) + config = EvoModelConfig.parse(payload) + now = time.time_ns() // 1_000_000 + payload_hash = sha256_id(config.raw) + with self._connect() as connection: + connection.execute( + """INSERT INTO config_operations + (subject_id, action, operation_id, request_digest, result_digest, + status, expected_revision, target_revision, payload_hash, + canonical_payload, created_at, updated_at) + VALUES ('admin', 'model_config:rollback', ?, ?, ?, 'PREPARING', + ?, ?, ?, ?, ?, ?)""", + ( + operation_id, + sha256_id( + { + "expected_revision": expected_revision, + "target_revision": target_revision, + } + ), + sha256_id({"config_revision": target_revision}), + expected_revision, + target_revision, + payload_hash, + canonical_json_v1(config.raw).decode(), + now, + now, + ), + ) + self._atomic_replace_payload(config.raw) + with self._connect() as connection: + connection.execute( + """INSERT INTO committed_head + (singleton, config_revision, canonical_payload_hash, + config_identity_key_id, commit_operation_id) + VALUES (1, ?, ?, ?, ?) + ON CONFLICT(singleton) DO UPDATE SET + config_revision=excluded.config_revision, + canonical_payload_hash=excluded.canonical_payload_hash, + config_identity_key_id=excluded.config_identity_key_id, + commit_operation_id=excluded.commit_operation_id""", + ( + target_revision, + payload_hash, + config.config_identity_key_id, + operation_id, + ), + ) + connection.execute( + """UPDATE config_operations SET status='COMMITTED', updated_at=? + WHERE subject_id='admin' AND action='model_config:rollback' + AND operation_id=?""", + (now, operation_id), + ) + return ConfigRevision(target_revision) + + def _committed_operation_revision( + self, operation_id: str, *, payload_hash: str, target_revision: int + ) -> int | None: + subject_id = "development" if operation_id == "bootstrap" else "admin" + with self._connect() as connection: + row = connection.execute( + """SELECT status, payload_hash, target_revision + FROM config_operations + WHERE subject_id=? AND action='model_config:commit' AND operation_id=?""", + (subject_id, operation_id), + ).fetchone() + if row is None: + return None + if ( + str(row["payload_hash"]) != payload_hash + or int(row["target_revision"] or 0) != target_revision + ): + raise EvoRuntimeError("IDEMPOTENCY_CONFLICT") + if row["status"] == "COMMITTED": + return target_revision + if row["status"] == "CONFLICT": + raise EvoRuntimeError("CONFIG_INTEGRITY_MISMATCH") + return None + + def _write_committed( + self, payload: Mapping[str, Any], *, operation_id: str + ) -> None: + payload_hash = sha256_id(payload) + revision = int(payload["config_revision"]) + identity_key_id = str(payload["config_identity_key_id"]) + now = time.time_ns() // 1_000_000 + subject_id = "development" if operation_id == "bootstrap" else "admin" + self.path.parent.mkdir(parents=True, exist_ok=True) + canonical_payload = canonical_json_v1(payload).decode("utf-8") + with self._connect() as connection: + existing = connection.execute( + """SELECT status, payload_hash, target_revision + FROM config_operations + WHERE subject_id=? AND action='model_config:commit' AND operation_id=?""", + (subject_id, operation_id), + ).fetchone() + if existing is not None: + if ( + str(existing["payload_hash"]) != payload_hash + or int(existing["target_revision"] or 0) != revision + ): + raise EvoRuntimeError("IDEMPOTENCY_CONFLICT") + if existing["status"] == "COMMITTED": + return + if existing["status"] == "CONFLICT": + raise EvoRuntimeError("CONFIG_INTEGRITY_MISMATCH") + else: + connection.execute( + """INSERT INTO config_operations + (subject_id, action, operation_id, request_digest, status, expected_revision, + target_revision, payload_hash, canonical_payload, created_at, updated_at) + VALUES (?, 'model_config:commit', ?, ?, 'PREPARING', ?, ?, ?, ?, ?, ?)""", + ( + subject_id, + operation_id, + payload_hash, + revision - 1, + revision, + payload_hash, + canonical_payload, + now, + now, + ), + ) + self._atomic_replace_payload(payload) + with self._connect() as connection: + connection.execute( + """INSERT INTO committed_head + (singleton, config_revision, canonical_payload_hash, config_identity_key_id, commit_operation_id) + VALUES (1, ?, ?, ?, ?) + ON CONFLICT(singleton) DO UPDATE SET + config_revision=excluded.config_revision, + canonical_payload_hash=excluded.canonical_payload_hash, + config_identity_key_id=excluded.config_identity_key_id, + commit_operation_id=excluded.commit_operation_id""", + (revision, payload_hash, identity_key_id, operation_id), + ) + connection.execute( + """UPDATE config_operations + SET status='COMMITTED', result_digest=?, updated_at=? + WHERE subject_id=? AND action='model_config:commit' AND operation_id=?""", + ( + sha256_id({"config_revision": revision}), + now, + subject_id, + operation_id, + ), + ) + + def _atomic_replace_payload(self, payload: Mapping[str, Any]) -> None: + rendered = yaml.safe_dump(dict(payload), allow_unicode=True, sort_keys=False) + fd, temporary_name = tempfile.mkstemp( + prefix=".model_routes.", suffix=".tmp", dir=self.path.parent + ) + temporary = Path(temporary_name) + try: + with os.fdopen(fd, "w", encoding="utf-8") as handle: + handle.write(rendered) + handle.flush() + os.fsync(handle.fileno()) + os.chmod(temporary, stat.S_IRUSR | stat.S_IWUSR) + os.replace(temporary, self.path) + directory_fd = os.open(self.path.parent, os.O_RDONLY) + try: + os.fsync(directory_fd) + finally: + os.close(directory_fd) + finally: + temporary.unlink(missing_ok=True) + + def _recover_preparing_operations(self) -> None: + """Complete or safely replay interrupted file/head commits.""" + + with self._lock: + with self._connect() as connection: + rows = connection.execute( + """SELECT * FROM config_operations + WHERE action IN ('model_config:commit', 'model_config:rollback') + AND status='PREPARING' + ORDER BY created_at""" + ).fetchall() + for row in rows: + try: + payload_text = str(row["canonical_payload"] or "") + if not payload_text: + raise ValueError("missing recovery payload") + payload = json.loads(payload_text) + target = int(row["target_revision"]) + payload_hash = str(row["payload_hash"]) + if int(payload.get("config_revision", 0)) != target: + raise ValueError("recovery revision mismatch") + file_matches_target = False + if self.path.exists(): + current_payload = load_yaml_unique( + self.path.read_text(encoding="utf-8") + ) + file_matches_target = sha256_id(current_payload) == payload_hash + if not file_matches_target: + head_revision = self.current_revision() + if head_revision != int(row["expected_revision"] or 0): + raise ValueError("recovery head moved") + self._atomic_replace_payload(payload) + now = time.time_ns() // 1_000_000 + with self._connect() as connection: + connection.execute( + """INSERT INTO committed_head + (singleton, config_revision, canonical_payload_hash, + config_identity_key_id, commit_operation_id) + VALUES (1, ?, ?, ?, ?) + ON CONFLICT(singleton) DO UPDATE SET + config_revision=excluded.config_revision, + canonical_payload_hash=excluded.canonical_payload_hash, + config_identity_key_id=excluded.config_identity_key_id, + commit_operation_id=excluded.commit_operation_id""", + ( + target, + payload_hash, + str(payload["config_identity_key_id"]), + str(row["operation_id"]), + ), + ) + connection.execute( + """UPDATE config_operations + SET status='COMMITTED', result_digest=?, updated_at=? + WHERE subject_id=? AND action=? AND operation_id=?""", + ( + sha256_id({"config_revision": target}), + now, + row["subject_id"], + row["action"], + row["operation_id"], + ), + ) + except Exception: + self._recovery_conflict = True + with self._connect() as connection: + connection.execute( + """UPDATE config_operations SET status='CONFLICT', updated_at=? + WHERE subject_id=? AND action=? AND operation_id=?""", + ( + time.time_ns() // 1_000_000, + row["subject_id"], + row["action"], + row["operation_id"], + ), + ) + + +def proposal_payload( + payload: Mapping[str, Any], *, target_revision: int +) -> Mapping[str, Any]: + candidate = dict(payload) + candidate["config_revision"] = target_revision + candidate.pop("capability_evidence", None) + return candidate + + +def proposal_hash( + payload: Mapping[str, Any], *, target_revision: int, identity_key: bytes +) -> str: + return hmac_id( + identity_key, proposal_payload(payload, target_revision=target_revision) + ) + + +@dataclass(frozen=True, slots=True) +class V2MigrationReport: + converted_providers: tuple[str, ...] + converted_models: tuple[str, ...] + converted_aliases: tuple[str, ...] + blocking_issues: tuple[Mapping[str, str], ...] + warnings: tuple[Mapping[str, str], ...] + required_secret_mappings: tuple[Mapping[str, str], ...] + + +def convert_v2_to_v3_draft( + payload: Mapping[str, Any], + *, + target_revision: int, + config_identity_key_id: str, +) -> tuple[Mapping[str, Any], V2MigrationReport]: + """Convert unambiguous V2 routes to a V3 Draft without guessing endpoints.""" + + source = EvoModelConfig.parse(payload, require_evidence=False) + if source.schema_version != 2: + raise EvoRuntimeError("LLM_ROUTE_CONFIGURATION_REQUIRED", "source must be V2") + adapter_map = { + "openai": ("openai", "openai-v1", "openai_native"), + "anthropic": ("anthropic", "anthropic-v1", "anthropic_native"), + "custom-openai": ( + "generic-openai-compatible", + "generic-openai-compatible-v1", + "openai_compatible", + ), + } + providers: list[Mapping[str, Any]] = [] + aliases: list[Mapping[str, Any]] = [] + converted_models: list[str] = [] + blocking: list[Mapping[str, str]] = [] + warnings: list[Mapping[str, str]] = [] + secret_mappings: list[Mapping[str, str]] = [] + model_keys: dict[tuple[str, str], str] = {} + for provider_id, provider in source.providers.items(): + referenced = { + route.endpoint + for selector_id in source.route_selectors + for route in source.concrete_routes(selector_id) + if route.provider == provider_id + } + if len(referenced) != 1: + blocking.append( + { + "code": "MULTIPLE_ENDPOINTS_REQUIRE_SPLIT", + "provider_id": provider_id, + "detail": ",".join(sorted(referenced)), + } + ) + continue + mapped = adapter_map.get(provider.protocol) + if mapped is None: + blocking.append( + { + "code": "ADAPTER_MAPPING_REQUIRED", + "provider_id": provider_id, + "detail": provider.protocol, + } + ) + continue + adapter_id, revision, wire = mapped + if provider.protocol == "custom-openai": + warnings.append( + { + "code": "GENERIC_ADAPTER_REQUIRES_CONFIRMATION", + "provider_id": provider_id, + "detail": "Base URL was not used to infer a vendor adapter", + } + ) + endpoint = provider.endpoints[next(iter(referenced))] + credential_ref = ( + f"secret://model-providers/{provider_id}#{endpoint.auth.revision}" + ) + secret_mappings.append( + { + "provider_id": provider_id, + "old_ref": endpoint.auth.ref, + "new_ref": credential_ref, + } + ) + models: list[Mapping[str, Any]] = [] + for index, (legacy_key, model) in enumerate(provider.models.items()): + model_key = re.sub(r"[^A-Za-z0-9._-]+", "-", legacy_key).strip("-") + model_key = model_key or f"model-{index + 1}" + if any(item["model_key"] == model_key for item in models): + model_key = f"{model_key}-{index + 1}" + model_keys[(provider_id, legacy_key)] = model_key + converted_models.append(f"{provider_id}/{model_key}") + selector = next( + ( + item + for item in source.route_selectors.values() + if item.provider == provider_id and item.model == legacy_key + ), + None, + ) + models.append( + { + "model_key": model_key, + "provider_model_id": model.model_id, + "version_policy": "rolling", + "resolved_model_revision": None, + "display_name": model.model_id, + "description": "Migrated from model routes V2", + "enabled": True, + "tags": ["migrated-v2"], + "invocation": { + "api_mode": selector.api_mode + if selector + else "chat_completions", + "tool_call_transport": "native", + }, + "capabilities": { + "text": True, + "vision": model.supports_vision, + "video": False, + "documents": False, + "tools": True, + "structured_output": False, + "thinking": model.supports_reasoning, + }, + "limits": { + "context_tokens": model.context_window, + "max_output_tokens": model.max_output_tokens, + }, + "parameters": { + "defaults": dict(model.params), + "purpose_overrides": {}, + "user_options": {}, + "constraints": [], + }, + "access": { + "visibility": "role_based" + if model.allowed_roles + else "authenticated", + "roles": list(model.allowed_roles), + "groups": [], + "users": [], + }, + "billing": { + "sku": model.quote.billing_sku, + "pricing_revision": model.quote.pricing_revision, + "currency": model.quote.currency, + "unit_scale": model.quote.unit_scale, + "input_microunits_per_million": model.quote.input_microunits_per_million, + "output_microunits_per_million": model.quote.output_microunits_per_million, + "cached_microunits_per_million": model.quote.cached_input_microunits_per_million, + "multiplier": model.quote.multiplier, + }, + } + ) + providers.append( + { + "provider_id": provider_id, + "display_name": provider_id, + "adapter_id": adapter_id, + "adapter_revision": revision, + "wire_protocol": wire, + "enabled": True, + "connection": { + "base_url": endpoint.base_url, + "credential_ref": credential_ref, + }, + "defaults": {}, + "models": models, + } + ) + for alias, selector_id in source.main_routes.selectable.items(): + selector = source.route_selectors[selector_id] + key = model_keys.get((selector.provider, selector.model)) + if key is None: + continue + aliases.append( + { + "alias": alias, + "display_name": alias, + "provider_ref": selector.provider, + "model_ref": key, + "enabled": True, + "access": { + "visibility": "authenticated", + "roles": [], + "groups": [], + "users": [], + }, + "defaults": {}, + } + ) + default_alias = ( + source.main_routes.default_alias + if any(item["alias"] == source.main_routes.default_alias for item in aliases) + else (str(aliases[0]["alias"]) if aliases else "") + ) + title_route = source.concrete_routes(source.title_selector_id)[0] + title_alias = next( + ( + str(item["alias"]) + for item in aliases + if item["provider_ref"] == title_route.provider + and item["model_ref"] + == model_keys.get((title_route.provider, title_route.model)) + ), + default_alias, + ) + draft = { + "schema_version": 3, + "config_revision": target_revision, + "config_identity_key_id": config_identity_key_id, + "runtime_defaults": {}, + "providers": providers, + "aliases": aliases, + "purpose_defaults": {purpose: {} for purpose in _PURPOSES}, + "purpose_routes": { + "main_agent": {"default_alias": default_alias}, + "tool_selector": "inherit_main", + "deepagents_summarizer": "inherit_main", + "title": {"default_alias": title_alias}, + }, + "purpose_call_limits": { + purpose: asdict(limit) + for purpose, limit in source.purpose_call_limits.items() + }, + "health_policy": { + "provider_connection": asdict(source.route_health), + "model_route": asdict(source.route_health), + }, + "web_runtime": asdict(source.web_runtime), + "capability_evidence": [], + } + return draft, V2MigrationReport( + converted_providers=tuple(str(item["provider_id"]) for item in providers), + converted_models=tuple(converted_models), + converted_aliases=tuple(str(item["alias"]) for item in aliases), + blocking_issues=tuple(blocking), + warnings=tuple(warnings), + required_secret_mappings=tuple(secret_mappings), + ) diff --git a/EvoScientist/llm/model_config_v4.py b/EvoScientist/llm/model_config_v4.py new file mode 100644 index 0000000..db2e0ed --- /dev/null +++ b/EvoScientist/llm/model_config_v4.py @@ -0,0 +1,2837 @@ +"""Model configuration V4 normalization, projection, and unified persistence. + +The admin contract is Provider + ModelProfile. Runtime remains on the stable +V3 wire/config contract and consumes a deterministic projection stored beside +encrypted credentials, evidence, operations, and audit records in one SQLite +database. +""" + +from __future__ import annotations + +import base64 +import hashlib +import hmac +import json +import os +import sqlite3 +import uuid +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from datetime import UTC, datetime, timedelta +from pathlib import Path +from typing import Any + +from cryptography.fernet import Fernet, InvalidToken +from filelock import FileLock + +from ..config.settings import get_config_dir +from .adapter_registry import get_adapter_registry +from .configuration import ResolvedSecret, SecretReference +from .contracts import EvoRuntimeError +from .crypto import HmacKeyRing, canonical_json_v1, hmac_id, sha256_id +from .invocation import derive_runtime_invocation +from .model_config import ( + EvoModelConfig, + endpoint_fingerprint, + route_semantics_hash, + route_semantics_payload, +) + +_REVISION_HMAC_INFO = "ai4sci/model-config-revision/v4" +_EVIDENCE_HMAC_INFO = "ai4sci/model-capability-evidence/v4" +_CREDENTIAL_ETAG_INFO = "ai4sci/model-credential-etag/v4" +_CREDENTIAL_FINGERPRINT_INFO = "ai4sci/model-credential-fingerprint/v4" + + +OPERATIONAL_POLICY_V1: dict[str, Any] = { + "revision": "operational-policy-v1", + "runtime_defaults": { + "sdk_max_retries": 0, + "connect_timeout_seconds": 10, + "first_event_timeout_seconds": 60, + "stream_idle_timeout_seconds": 60, + "attempt_timeout_seconds": 600, + "max_sse_event_bytes": 1_048_576, + "max_content_block_bytes": 4_194_304, + "max_output_bytes": 16_777_216, + "max_opaque_state_bytes": 8_388_608, + "stream_buffer_max_events": 256, + "stream_buffer_max_bytes": 2_097_152, + "max_tool_schema_bytes": 262_144, + "max_tool_arguments_bytes": 1_048_576, + "max_tool_schema_depth": 16, + "max_tool_argument_depth": 32, + }, + "runtime_validation_bounds": { + "connect_timeout_seconds": [1, 60], + "first_event_timeout_seconds": [1, 300], + "stream_idle_timeout_seconds": [1, 300], + "attempt_timeout_seconds": [1, 3600], + "max_sse_event_bytes": [4096, 4_194_304], + "max_content_block_bytes": [4096, 16_777_216], + "max_output_bytes": [65_536, 67_108_864], + "max_opaque_state_bytes": [65_536, 33_554_432], + "stream_buffer_max_events": [1, 1024], + "stream_buffer_max_bytes": [65_536, 8_388_608], + "max_tool_schema_bytes": [4096, 1_048_576], + "max_tool_arguments_bytes": [4096, 4_194_304], + "max_tool_schema_depth": [1, 32], + "max_tool_argument_depth": [1, 64], + }, + "provider_health": { + "failure_threshold": 3, + "cooldown_seconds": 30, + "half_open_max_inflight": 1, + "counted_error_codes": [ + "MODEL_RATE_LIMITED", + "MODEL_PROVIDER_ERROR", + "MODEL_TIMEOUT", + ], + "open_immediately_error_codes": ["MODEL_AUTHENTICATION_FAILED"], + }, + "model_health": { + "failure_threshold": 3, + "cooldown_seconds": 30, + "half_open_max_inflight": 1, + "counted_error_codes": ["MODEL_PROVIDER_ERROR"], + "open_immediately_error_codes": ["MODEL_NOT_FOUND"], + }, + "web_runtime": { + "title_start_timeout_seconds": 30, + "prepare_ttl_seconds": 30, + "turn_lease_grace_seconds": 30, + "active_run_timeout_seconds": 1800, + "max_run_journal_events": 10_000, + "max_run_journal_bytes": 16_777_216, + "max_prepared_runs_per_subject": 4, + "max_prepared_runs_total": 128, + }, + "admin_operations": { + "probe_timeout_seconds": 60, + "max_probe_concurrency_per_provider": 4, + "max_probe_concurrency_total": 8, + "max_enabled_profiles_per_save": 200, + "save_deadline_seconds": 180, + "evidence_ttl_seconds": 86_400, + "evidence_retention_days": 180, + }, +} + + +@dataclass(frozen=True, slots=True) +class ActiveGeneration: + revision: int + security_epoch: int + evidence_epoch: int + adapter_policy_revision: str + + +@dataclass(frozen=True, slots=True) +class ProviderCredentialMetadata: + provider_id: str + configured: bool + masked_value: str + updated_at: str + credential_etag: str + + +@dataclass(frozen=True, slots=True) +class CredentialPlan: + bindings: Mapping[str, int] + fingerprints: Mapping[str, str] + plaintext: Mapping[str, str] + new_versions: frozenset[str] + base_security_epoch: int + + def resolve(self, reference: SecretReference) -> ResolvedSecret: + if not reference.ref.startswith("secret://model-providers/"): + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + provider_and_version = reference.ref[len("secret://model-providers/") :] + try: + provider_id, raw_version = provider_and_version.rsplit("#", 1) + version = int(raw_version) + except (TypeError, ValueError) as exc: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") from exc + if self.bindings.get(provider_id) != version or reference.revision != version: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + try: + value = self.plaintext[provider_id] + fingerprint = self.fingerprints[provider_id] + except KeyError as exc: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") from exc + return ResolvedSecret(value, version, str(version), fingerprint) + + +@dataclass(frozen=True, slots=True) +class UnifiedSaveResult: + operation_id: str + config_revision: int + changed: bool + revalidated: bool + status: str = "SUCCEEDED" + + +def _utc_now() -> str: + return datetime.now(UTC).isoformat(timespec="microseconds") + + +def _uuid7() -> str: + """Generate a sortable UUIDv7 without requiring Python 3.14's uuid.uuid7.""" + + timestamp_ms = int(datetime.now(UTC).timestamp() * 1000) & ((1 << 48) - 1) + random_bits = int.from_bytes(os.urandom(10), "big") + value = timestamp_ms << 80 + value |= 0x7 << 76 + value |= ((random_bits >> 68) & 0xFFF) << 64 + value |= 0b10 << 62 + value |= random_bits & ((1 << 62) - 1) + return str(uuid.UUID(int=value)) + + +def _configuration_error( + path: str, + message: str, + *, + detail_code: str = "MODEL_CONFIG_VALIDATION_FAILED", +) -> EvoRuntimeError: + return EvoRuntimeError( + "LLM_ROUTE_CONFIGURATION_REQUIRED", + message, + details=({"path": path, "code": detail_code},), + ) + + +def _mapping(value: Any, path: str) -> dict[str, Any]: + if not isinstance(value, Mapping) or not all(isinstance(key, str) for key in value): + raise _configuration_error( + path, + f"{path} must be an object", + detail_code="CONFIG_OBJECT_REQUIRED", + ) + return dict(value) + + +def _sequence(value: Any, path: str) -> list[Any]: + if not isinstance(value, list | tuple): + raise _configuration_error( + path, + f"{path} must be an array", + detail_code="CONFIG_ARRAY_REQUIRED", + ) + return list(value) + + +def _strict(value: Mapping[str, Any], allowed: set[str], path: str) -> None: + unknown = set(value) - allowed + if unknown: + raise _configuration_error( + path, + f"{path} has unknown fields: {', '.join(sorted(unknown))}", + detail_code="CONFIG_UNKNOWN_FIELDS", + ) + + +def _warn( + warnings: list[dict[str, Any]] | None, path: str, code: str, message: str +) -> None: + if warnings is not None: + warnings.append({"path": path, "code": code, "message": message}) + + +def _filter_fields( + value: dict[str, Any], + allowed: set[str], + path: str, + *, + lenient: bool, + warnings: list[dict[str, Any]] | None, +) -> None: + unknown = set(value) - allowed + if not unknown: + return + if not lenient: + raise _configuration_error( + path, + f"{path} has unknown fields: {', '.join(sorted(unknown))}", + detail_code="CONFIG_UNKNOWN_FIELDS", + ) + for key in sorted(unknown): + value.pop(key, None) + _warn( + warnings, path, "CONFIG_UNKNOWN_FIELDS", + f"已忽略未知字段: {', '.join(sorted(unknown))}", + ) + + +def _bounded_integer( + value: Any, + path: str, + *, + minimum: int, + maximum: int | None, + lenient: bool, + warnings: list[dict[str, Any]] | None, +) -> int: + if isinstance(value, bool) or not isinstance(value, int): + raise _configuration_error( + path, f"{path} must be integer", detail_code="CONFIG_INTEGER_REQUIRED" + ) + if value >= minimum and (maximum is None or value <= maximum): + return value + if not lenient: + raise _configuration_error( + path, f"{path} is out of range", detail_code="CONFIG_VALUE_OUT_OF_RANGE" + ) + clamped = max(minimum, value) if maximum is None else max(minimum, min(value, maximum)) + _warn( + warnings, path, "CONFIG_VALUE_OUT_OF_RANGE", + f"{path} 超出范围,已调整为 {clamped}", + ) + return clamped + + +def _uuid(value: Any, path: str) -> str: + try: + parsed = uuid.UUID(str(value)) + except (ValueError, TypeError, AttributeError) as exc: + raise _configuration_error( + path, + f"{path} must be UUID", + detail_code="CONFIG_UUID_INVALID", + ) from exc + return str(parsed) + + +def _text(value: Any, path: str, *, allow_empty: bool = False) -> str: + if not isinstance(value, str): + raise _configuration_error( + path, + f"{path} must be text", + detail_code="CONFIG_TEXT_REQUIRED", + ) + normalized = value.strip() + if (not allow_empty and not normalized) or any(char in normalized for char in "\r\n\0"): + raise _configuration_error( + path, + f"{path} is invalid", + detail_code="CONFIG_VALUE_INVALID", + ) + return normalized + + +def _integer(value: Any, path: str, *, minimum: int, maximum: int | None = None) -> int: + if isinstance(value, bool) or not isinstance(value, int): + raise _configuration_error( + path, + f"{path} must be integer", + detail_code="CONFIG_INTEGER_REQUIRED", + ) + if value < minimum or (maximum is not None and value > maximum): + raise _configuration_error( + path, + f"{path} is out of range", + detail_code="CONFIG_VALUE_OUT_OF_RANGE", + ) + return value + + +def _base_url(value: Any, path: str, *, enabled: bool) -> str: + return str(value or "").strip().rstrip("/") + + +def _normalize_connection( + value: Any, + path: str, + *, + enabled: bool, + lenient: bool = False, + warnings: list[dict[str, Any]] | None = None, +) -> dict[str, Any]: + raw = _mapping(value or {}, path) + allowed = { + "base_url", + "connect_timeout_seconds", + "first_event_timeout_seconds", + "stream_idle_timeout_seconds", + "attempt_timeout_seconds", + "max_inflight_requests", + "queue_timeout_seconds", + } + _filter_fields(raw, allowed, path, lenient=lenient, warnings=warnings) + defaults = OPERATIONAL_POLICY_V1["runtime_defaults"] + bounds = OPERATIONAL_POLICY_V1["runtime_validation_bounds"] + result: dict[str, Any] = { + "base_url": _base_url(raw.get("base_url", ""), f"{path}.base_url", enabled=enabled), + "connect_timeout_seconds": _bounded_integer( + raw.get("connect_timeout_seconds", defaults["connect_timeout_seconds"]), + f"{path}.connect_timeout_seconds", + minimum=1, + maximum=bounds["connect_timeout_seconds"][1], + lenient=lenient, + warnings=warnings, + ), + "first_event_timeout_seconds": _bounded_integer( + raw.get("first_event_timeout_seconds", defaults["first_event_timeout_seconds"]), + f"{path}.first_event_timeout_seconds", + minimum=1, + maximum=bounds["first_event_timeout_seconds"][1], + lenient=lenient, + warnings=warnings, + ), + "stream_idle_timeout_seconds": _bounded_integer( + raw.get("stream_idle_timeout_seconds", defaults["stream_idle_timeout_seconds"]), + f"{path}.stream_idle_timeout_seconds", + minimum=1, + maximum=bounds["stream_idle_timeout_seconds"][1], + lenient=lenient, + warnings=warnings, + ), + "attempt_timeout_seconds": _bounded_integer( + raw.get("attempt_timeout_seconds", defaults["attempt_timeout_seconds"]), + f"{path}.attempt_timeout_seconds", + minimum=1, + maximum=bounds["attempt_timeout_seconds"][1], + lenient=lenient, + warnings=warnings, + ), + "max_inflight_requests": _bounded_integer( + raw.get("max_inflight_requests", 16), + f"{path}.max_inflight_requests", + minimum=1, + maximum=256, + lenient=lenient, + warnings=warnings, + ), + "queue_timeout_seconds": _bounded_integer( + raw.get("queue_timeout_seconds", 5), + f"{path}.queue_timeout_seconds", + minimum=1, + maximum=60, + lenient=lenient, + warnings=warnings, + ), + } + minimum_attempt = max( + result["connect_timeout_seconds"], + result["first_event_timeout_seconds"], + result["stream_idle_timeout_seconds"], + ) + if result["attempt_timeout_seconds"] < minimum_attempt: + if not lenient: + raise _configuration_error( + path, + f"{path} timeouts conflict", + detail_code="CONFIG_VALUE_CONFLICT", + ) + result["attempt_timeout_seconds"] = minimum_attempt + _warn( + warnings, f"{path}.attempt_timeout_seconds", "CONFIG_VALUE_CONFLICT", + f"超时组合冲突,attempt_timeout_seconds 已提升为 {minimum_attempt}", + ) + return result + + +def normalize_v4_config( + value: Any, + *, + lenient: bool = False, + warnings: list[dict[str, Any]] | None = None, +) -> dict[str, Any]: + """Normalize an external V4 admin payload without secret material. + + Strict mode (default) rejects anything suspicious. Lenient mode keeps + hard guarantees (structure, identifiers, runtime constructibility) but + downgrades everything else to warnings, auto-correcting where possible. + """ + + raw = _mapping(value, "model_config") + raw.pop("config_revision", None) + allowed = { + "schema_version", + "default_model_profile_id", + "purpose_defaults", + "purpose_policy", + "purpose_call_limits", + "providers", + } + _filter_fields(raw, allowed, "model_config", lenient=lenient, warnings=warnings) + if raw.get("schema_version") != 4: + raise _configuration_error( + "schema_version", + "schema_version must be 4", + detail_code="CONFIG_VALUE_INVALID", + ) + registry = get_adapter_registry() + providers: list[dict[str, Any]] = [] + provider_ids: set[str] = set() + profile_ids: set[str] = set() + enabled_profiles: set[str] = set() + for provider_index, provider_value in enumerate(_sequence(raw.get("providers") or [], "providers")): + path = f"providers[{provider_index}]" + provider_raw = _mapping(provider_value, path) + _filter_fields( + provider_raw, + { + "provider_id", + "display_name", + "adapter_id", + "adapter_revision", + "enabled", + "connection", + "models", + }, + path, + lenient=lenient, + warnings=warnings, + ) + provider_id = _uuid(provider_raw.get("provider_id"), f"{path}.provider_id") + if provider_id in provider_ids: + raise _configuration_error( + f"{path}.provider_id", + "duplicate provider_id", + detail_code="CONFIG_DUPLICATE_ID", + ) + provider_ids.add(provider_id) + enabled = bool(provider_raw.get("enabled", True)) + adapter_id = _text(provider_raw.get("adapter_id"), f"{path}.adapter_id") + adapter_revision = _text( + provider_raw.get("adapter_revision"), f"{path}.adapter_revision" + ) + registration = registry.get(adapter_id, adapter_revision) + if registration.lifecycle == "blocked": + raise EvoRuntimeError("MODEL_ADAPTER_BLOCKED") + connection = _normalize_connection( + provider_raw.get("connection"), + f"{path}.connection", + enabled=enabled, + lenient=lenient, + warnings=warnings, + ) + models: list[dict[str, Any]] = [] + for model_index, model_value in enumerate( + _sequence(provider_raw.get("models") or [], f"{path}.models") + ): + model_path = f"{path}.models[{model_index}]" + model_raw = _mapping(model_value, model_path) + _filter_fields( + model_raw, + { + "model_profile_id", + "provider_model_id", + "display_name", + "description", + "enabled", + "version_policy", + "resolved_model_revision", + "invocation", + "capabilities", + "limits", + "parameters", + "access", + "billing", + "tags", + }, + model_path, + lenient=lenient, + warnings=warnings, + ) + profile_id = _uuid( + model_raw.get("model_profile_id"), f"{model_path}.model_profile_id" + ) + if profile_id in profile_ids: + raise _configuration_error( + f"{model_path}.model_profile_id", + "duplicate model_profile_id", + detail_code="CONFIG_DUPLICATE_ID", + ) + profile_ids.add(profile_id) + model_enabled = bool(model_raw.get("enabled", True)) + provider_model_id = _text( + model_raw.get("provider_model_id", ""), + f"{model_path}.provider_model_id", + allow_empty=not (enabled and model_enabled), + ) + invocation = _mapping(model_raw.get("invocation") or {}, f"{model_path}.invocation") + # tool_call_transport was a V4 admin input before the tool capability + # contract became authoritative. Keep accepting it while old signed + # revisions are migrated, but never persist or trust its value. + _filter_fields( + invocation, {"api_mode", "tool_call_transport"}, + f"{model_path}.invocation", lenient=lenient, warnings=warnings, + ) + api_mode = _text( + invocation.get("api_mode", registration.recommended_api_mode), + f"{model_path}.invocation.api_mode", + ) + capabilities_raw = _mapping( + model_raw.get("capabilities") or {}, f"{model_path}.capabilities" + ) + capability_names = { + "text", + "vision", + "video", + "documents", + "tools", + "structured_output", + "reasoning", + } + _filter_fields( + capabilities_raw, capability_names, + f"{model_path}.capabilities", lenient=lenient, warnings=warnings, + ) + capabilities = { + name: bool(capabilities_raw.get(name, name == "text")) + for name in capability_names + } + if not capabilities["text"]: + if not lenient: + raise _configuration_error( + f"{model_path}.capabilities.text", + "text capability is required", + detail_code="CONFIG_REQUIRED", + ) + capabilities["text"] = True + _warn( + warnings, f"{model_path}.capabilities.text", "CONFIG_REQUIRED", + "text 能力为必需,已强制开启", + ) + limits_raw = _mapping(model_raw.get("limits") or {}, f"{model_path}.limits") + _filter_fields( + limits_raw, + {"context_tokens", "max_output_tokens", "max_inflight_requests"}, + f"{model_path}.limits", + lenient=lenient, + warnings=warnings, + ) + limits: dict[str, Any] = {} + for name in ("context_tokens", "max_output_tokens", "max_inflight_requests"): + item = limits_raw.get(name) + limits[name] = None if item is None else _bounded_integer( + item, f"{model_path}.limits.{name}", minimum=1, maximum=None, + lenient=lenient, warnings=warnings, + ) + if enabled and model_enabled and ( + limits["context_tokens"] is None or limits["max_output_tokens"] is None + ): + try: + descriptor = registration.resolve_model_descriptor( + provider_model_id, + api_mode, + context_tokens=limits["context_tokens"], + max_output_tokens=limits["max_output_tokens"], + declared_capabilities={ + "text": capabilities["text"], + "vision": capabilities["vision"], + "video": capabilities["video"], + "documents": capabilities["documents"], + "tools": capabilities["tools"], + "structured_output": capabilities["structured_output"], + "thinking": capabilities["reasoning"], + }, + ) + except EvoRuntimeError: + if not lenient: + raise + _warn( + warnings, f"{model_path}.limits", "MODEL_DESCRIPTOR_INCOMPLETE", + "无法解析模型描述符,token 上限保留为空", + ) + else: + limits["context_tokens"] = descriptor.context_tokens + limits["max_output_tokens"] = descriptor.max_output_tokens + parameters = _mapping(model_raw.get("parameters") or {}, f"{model_path}.parameters") + _filter_fields( + parameters, + { + "defaults", + "purpose_overrides", + "user_options", + "constraints", + "reasoning_policy", + }, + f"{model_path}.parameters", + lenient=lenient, + warnings=warnings, + ) + defaults = _mapping( + parameters.get("defaults") or {}, + f"{model_path}.parameters.defaults", + ) + capability_output_limit = limits["max_output_tokens"] + + def _normalize_output_limit( + mapping: dict, + path: str, + *, + capability_output_limit: int | None = capability_output_limit, + ) -> None: + value = mapping.get("output_token_limit") + if value is None: + return + value = _bounded_integer( + value, f"{path}.output_token_limit", minimum=1, maximum=None, + lenient=lenient, warnings=warnings, + ) + if capability_output_limit is not None and value > capability_output_limit: + if not lenient: + raise _configuration_error( + f"{path}.output_token_limit", + "output_token_limit exceeds model capability", + detail_code="CONFIG_LIMIT_EXCEEDED", + ) + _warn( + warnings, f"{path}.output_token_limit", "CONFIG_LIMIT_EXCEEDED", + "output_token_limit 超过模型能力上限,已收敛到能力上限", + ) + value = capability_output_limit + mapping["output_token_limit"] = value + + _normalize_output_limit(defaults, f"{model_path}.parameters.defaults") + overrides_raw = _mapping( + parameters.get("purpose_overrides") or {}, + f"{model_path}.parameters.purpose_overrides", + ) + purpose_overrides = {} + for purpose, values in overrides_raw.items(): + if purpose not in { + "main_agent", + "tool_selector", + "deepagents_summarizer", + "title", + }: + if not lenient: + raise _configuration_error( + f"{model_path}.parameters.purpose_overrides.{purpose}", + "unknown purpose override", + detail_code="CONFIG_REFERENCE_INVALID", + ) + _warn( + warnings, + f"{model_path}.parameters.purpose_overrides.{purpose}", + "CONFIG_REFERENCE_INVALID", + f"未知用途 {purpose},已按原样保留", + ) + normalized_values = _mapping( + values, + f"{model_path}.parameters.purpose_overrides.{purpose}", + ) + _normalize_output_limit( + normalized_values, + f"{model_path}.parameters.purpose_overrides.{purpose}", + ) + purpose_overrides[purpose] = normalized_values + user_options = _mapping( + parameters.get("user_options") or {}, + f"{model_path}.parameters.user_options", + ) + user_options.pop("output_token_limit", None) + parameters["defaults"] = defaults + parameters["purpose_overrides"] = purpose_overrides + parameters["user_options"] = user_options + parameters.setdefault("constraints", []) + parameters.setdefault("reasoning_policy", {}) + access = _mapping( + model_raw.get("access") + or {"visibility": "authenticated", "roles": [], "groups": [], "users": []}, + f"{model_path}.access", + ) + billing_value = model_raw.get("billing") + if enabled and model_enabled and billing_value is None: + raise _configuration_error( + f"{model_path}.billing", + "enabled model requires an explicit billing configuration", + detail_code="CONFIG_PRICING_REQUIRED", + ) + billing = _mapping( + billing_value + if billing_value is not None + else { + "sku": f"internal/{profile_id}", + "pricing_revision": "internal-unmetered-v1", + "currency": "CNY", + "unit_scale": 1_000_000, + "input_microunits_per_million": 0, + "output_microunits_per_million": 0, + "cached_microunits_per_million": 0, + "multiplier": 1, + }, + f"{model_path}.billing", + ) + billing.setdefault("multiplier", 1) + policy = _text( + model_raw.get("version_policy", "rolling"), + f"{model_path}.version_policy", + ) + if policy not in {"pinned", "rolling"}: + raise _configuration_error( + f"{model_path}.version_policy", + "invalid version policy", + detail_code="CONFIG_VALUE_INVALID", + ) + resolved = model_raw.get("resolved_model_revision") + if resolved is not None: + resolved = _text(resolved, f"{model_path}.resolved_model_revision") + models.append( + { + "model_profile_id": profile_id, + "provider_model_id": provider_model_id, + "display_name": _text( + model_raw.get("display_name", provider_model_id or "Model"), + f"{model_path}.display_name", + ), + "description": str(model_raw.get("description") or ""), + "enabled": model_enabled, + "version_policy": policy, + "resolved_model_revision": resolved, + "invocation": { + "api_mode": api_mode, + }, + "capabilities": capabilities, + "limits": limits, + "parameters": parameters, + "access": access, + "billing": billing, + "tags": list(model_raw.get("tags") or []), + } + ) + if enabled and model_enabled: + enabled_profiles.add(profile_id) + if enabled and not any(model["enabled"] for model in models): + raise _configuration_error( + f"{path}.models", + "enabled provider needs a model", + detail_code="CONFIG_ENABLED_MODEL_REQUIRED", + ) + providers.append( + { + "provider_id": provider_id, + "display_name": _text( + provider_raw.get("display_name", provider_id), f"{path}.display_name" + ), + "adapter_id": adapter_id, + "adapter_revision": adapter_revision, + "enabled": enabled, + "connection": connection, + "models": models, + } + ) + if not enabled_profiles: + raise _configuration_error( + "providers", + "an enabled model is required", + detail_code="CONFIG_ENABLED_MODEL_REQUIRED", + ) + if len(enabled_profiles) > int( + OPERATIONAL_POLICY_V1["admin_operations"]["max_enabled_profiles_per_save"] + ): + if not lenient: + raise _configuration_error( + "providers", + "too many enabled models", + detail_code="CONFIG_LIMIT_EXCEEDED", + ) + _warn( + warnings, "providers", "CONFIG_LIMIT_EXCEEDED", + "启用模型数量超过 200,已按原样保留", + ) + default_profile = _uuid( + raw.get("default_model_profile_id"), "default_model_profile_id" + ) + if default_profile not in enabled_profiles: + raise _configuration_error( + "default_model_profile_id", + "default model is unavailable", + detail_code="CONFIG_REFERENCE_INVALID", + ) + purposes = ("main_agent", "tool_selector", "deepagents_summarizer", "title") + purpose_defaults = _mapping(raw.get("purpose_defaults") or {}, "purpose_defaults") + purpose_defaults = { + purpose: _mapping(purpose_defaults.get(purpose) or {}, f"purpose_defaults.{purpose}") + for purpose in purposes + } + purpose_limits_raw = _mapping(raw.get("purpose_call_limits") or {}, "purpose_call_limits") + default_attempts = { + # A tool-using main agent needs an initial decision, at least one + # post-tool continuation, and a final response. Two calls is not a + # viable default for ordinary Web workflows. + "main_agent": 4, + "tool_selector": 2, + "deepagents_summarizer": 2, + "title": 1, + } + purpose_limits = {} + for purpose, attempt_default in default_attempts.items(): + item = _mapping(purpose_limits_raw.get(purpose) or {}, f"purpose_call_limits.{purpose}") + purpose_limits[purpose] = { + "max_attempts_per_run": _bounded_integer( + item.get("max_attempts_per_run", attempt_default), + f"purpose_call_limits.{purpose}.max_attempts_per_run", + minimum=1, + maximum=8, + lenient=lenient, + warnings=warnings, + ), + } + if sum(item["max_attempts_per_run"] for item in purpose_limits.values()) > 16: + if not lenient: + raise _configuration_error( + "purpose_call_limits", + "attempt limit is too large", + detail_code="CONFIG_LIMIT_EXCEEDED", + ) + _warn( + warnings, "purpose_call_limits", "CONFIG_LIMIT_EXCEEDED", + "各用途 max_attempts_per_run 总和超过 16,已按原样保留", + ) + return { + "schema_version": 4, + "default_model_profile_id": default_profile, + "purpose_defaults": purpose_defaults, + "purpose_policy": { + "tool_selector": "inherit_thread_model", + "deepagents_summarizer": "inherit_thread_model", + "title": "use_default_model", + }, + "purpose_call_limits": purpose_limits, + "providers": providers, + } + + +def model_profile_identity_map(payload: Mapping[str, Any]) -> dict[str, tuple[str, ...]]: + result: dict[str, tuple[str, ...]] = {} + for provider in payload.get("providers") or []: + for model in provider.get("models") or []: + invocation = model["invocation"] + result[str(model["model_profile_id"])] = ( + str(provider["provider_id"]), + str(model["provider_model_id"]), + str(provider["adapter_id"]), + str(invocation["api_mode"]), + str(model["version_policy"]), + ) + return result + + +def project_v4_to_v3( + payload: Mapping[str, Any], + *, + revision: int, + identity_key_id: str, + bindings: Mapping[str, int], +) -> dict[str, Any]: + registry = get_adapter_registry() + providers = [] + aliases = [] + for provider in payload["providers"]: + if not provider["enabled"]: + continue + provider_id = provider["provider_id"] + try: + secret_version = int(bindings[provider_id]) + except (KeyError, TypeError, ValueError) as exc: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") from exc + registration = registry.get(provider["adapter_id"], provider["adapter_revision"]) + connection = provider["connection"] + models = [] + for profile in provider["models"]: + if not profile["enabled"]: + continue + caps = dict(profile["capabilities"]) + caps["thinking"] = bool(caps.pop("reasoning", False)) + invocation = derive_runtime_invocation( + profile["invocation"]["api_mode"], caps + ) + models.append( + { + "model_key": profile["model_profile_id"], + "provider_model_id": profile["provider_model_id"], + "version_policy": profile["version_policy"], + "resolved_model_revision": profile["resolved_model_revision"], + "display_name": profile["display_name"], + "description": profile["description"], + "enabled": True, + "tags": profile["tags"], + "invocation": invocation, + "capabilities": caps, + "limits": profile["limits"], + "parameters": profile["parameters"], + "access": profile["access"], + "billing": profile["billing"], + } + ) + aliases.append( + { + "alias": profile["model_profile_id"], + "display_name": profile["display_name"], + "provider_ref": provider_id, + "model_ref": profile["model_profile_id"], + "enabled": True, + "access": profile["access"], + "defaults": {}, + } + ) + providers.append( + { + "provider_id": provider_id, + "display_name": provider["display_name"], + "adapter_id": provider["adapter_id"], + "adapter_revision": provider["adapter_revision"], + "wire_protocol": registration.supported_wire_protocols[0], + "enabled": True, + "connection": { + "base_url": connection["base_url"], + "credential_ref": ( + f"secret://model-providers/{provider_id}#{secret_version}" + ), + }, + "defaults": { + key: connection[key] + for key in ( + "connect_timeout_seconds", + "first_event_timeout_seconds", + "stream_idle_timeout_seconds", + "attempt_timeout_seconds", + "max_inflight_requests", + "queue_timeout_seconds", + ) + }, + "models": models, + } + ) + default_profile = payload["default_model_profile_id"] + projection = { + "schema_version": 3, + "config_revision": revision, + "config_identity_key_id": identity_key_id, + "runtime_defaults": dict(OPERATIONAL_POLICY_V1["runtime_defaults"]), + "providers": providers, + "aliases": aliases, + "purpose_defaults": dict(payload["purpose_defaults"]), + "purpose_routes": { + "main_agent": {"default_alias": default_profile}, + "tool_selector": "inherit_main", + "deepagents_summarizer": "inherit_main", + "title": {"default_alias": default_profile}, + }, + "purpose_call_limits": dict(payload["purpose_call_limits"]), + "health_policy": { + "provider_connection": dict(OPERATIONAL_POLICY_V1["provider_health"]), + "model_route": dict(OPERATIONAL_POLICY_V1["model_health"]), + }, + "web_runtime": dict(OPERATIONAL_POLICY_V1["web_runtime"]), + "capability_evidence": [], + } + EvoModelConfig.parse(projection, require_evidence=False) + return projection + + +def build_supported_v3_evidence( + projection: Mapping[str, Any], + *, + identity_key_ring: HmacKeyRing, + verified_profiles: Mapping[str, frozenset[str]] | None = None, + probe_results: Mapping[str, Mapping[str, Mapping[str, str]]] | None = None, + evidence_windows: Mapping[ + str, Mapping[str, tuple[str, str]] + ] | None = None, +) -> list[dict[str, Any]]: + """Build signed-input V3 evidence after controlled Provider probes succeed. + + ``verified_profiles`` maps Provider ID to model-profile IDs observed by the + probe runner. The function does not perform or assume network success. + """ + + config = EvoModelConfig.parse(projection, require_evidence=False) + semantics_key = identity_key_ring.derive( + config.config_identity_key_id, "ai4sci/route-semantics-hash/v3" + ) + endpoint_key = identity_key_ring.derive( + config.config_identity_key_id, "ai4sci/endpoint-fingerprint/v3" + ) + verified_at = _utc_now() + expires_at = ( + datetime.now(UTC) + + timedelta( + seconds=int(OPERATIONAL_POLICY_V1["admin_operations"]["evidence_ttl_seconds"]) + ) + ).isoformat(timespec="microseconds") + evidence: list[dict[str, Any]] = [] + routes = { + route.key(): route + for selector_id in config.route_selectors + for route in config.concrete_routes(selector_id) + } + for route in routes.values(): + provider = config.providers[route.provider] + model = provider.models[route.model] + route_verified_at, route_expires_at = ( + evidence_windows or {} + ).get(route.provider, {}).get(route.model, (verified_at, expires_at)) + if probe_results is not None: + try: + observed = dict(probe_results[route.provider][route.model]) + except KeyError as exc: + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") from exc + results = {"connectivity": observed.get("connectivity", "failed")} + for name, enabled in model.capabilities.items(): + if name == "text": + continue + probe_name = "reasoning" if name == "thinking" else name + results[name] = ( + observed.get(probe_name, "not_verified") + if enabled + else "not_declared" + ) + else: + if route.model not in (verified_profiles or {}).get( + route.provider, frozenset() + ): + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + results = {"connectivity": "supported"} + results.update( + { + name: "supported" if enabled else "not_declared" + for name, enabled in model.capabilities.items() + if name != "text" + } + ) + if results["connectivity"] != "supported" or any( + results.get(name) != "supported" + for name, enabled in model.capabilities.items() + if enabled and name != "text" + ): + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + evidence.append( + { + "evidence_id": _uuid7(), + "provider_ref": route.provider, + "model_ref": route.model, + "adapter_id": provider.adapter_id, + "adapter_revision": provider.adapter_revision, + "implementation_fingerprint": provider.implementation_fingerprint, + "wire_protocol": provider.wire_protocol, + "provider_model_id": model.model_id, + "resolved_model_revision": model.resolved_model_revision, + "version_policy": model.version_policy, + "reproducible": model.reproducible, + "api_mode": route.api_mode, + "tool_call_transport": route.tool_call_transport, + "base_url_fingerprint": endpoint_fingerprint(config, route, endpoint_key), + "secret_version": provider.endpoints[route.endpoint].auth.revision, + "route_semantics_hash": route_semantics_hash(config, route, semantics_key), + "fixture_digest": "provider-model-config-v4", + "verified_at": route_verified_at, + "evidence_expires_at": route_expires_at, + "probe_kind": "connectivity", + "outcome": "supported", + "failure_code": None, + "results": results, + } + ) + return evidence + + +class UnifiedModelConfigStore: + """Single SQLite authority for V4 admin state and V3 runtime projection.""" + + def __init__( + self, + path: Path | None = None, + *, + identity_key_ring: HmacKeyRing, + encryption_keys: Mapping[str, str | bytes] | None = None, + current_encryption_key_id: str | None = None, + ) -> None: + self.path = path or (get_config_dir() / "model_config.sqlite") + self.path.parent.mkdir(parents=True, exist_ok=True) + self._lock = FileLock(str(self.path) + ".lock") + self.identity_key_ring = identity_key_ring + materials = dict(encryption_keys or self._environment_encryption_keys()) + self.current_encryption_key_id = str( + current_encryption_key_id + or os.environ.get("AI4SCI_EVO_MODEL_SECRET_KEY_ID") + or next(iter(materials), "model-secret-v1") + ) + if self.current_encryption_key_id not in materials: + raise RuntimeError("current model secret encryption key is unavailable") + self._fernets = { + key_id: Fernet( + base64.urlsafe_b64encode( + hashlib.sha256( + material.encode("utf-8") if isinstance(material, str) else bytes(material) + ).digest() + ) + ) + for key_id, material in materials.items() + } + self._init_schema() + self.recover_expired_operations() + + @staticmethod + def _environment_encryption_keys() -> dict[str, str]: + raw = os.environ.get("AI4SCI_EVO_MODEL_SECRET_KEYS_JSON", "").strip() + if raw: + try: + values = json.loads(raw) + except json.JSONDecodeError as exc: + raise RuntimeError("AI4SCI_EVO_MODEL_SECRET_KEYS_JSON is invalid") from exc + if not isinstance(values, Mapping): + raise RuntimeError("AI4SCI_EVO_MODEL_SECRET_KEYS_JSON must be an object") + result = {str(key): str(value) for key, value in values.items()} + else: + material = os.environ.get("AI4SCI_EVO_MODEL_SECRET_MASTER_KEY", "") or os.environ.get( + "AI4SCI_EVO_CONFIG_IDENTITY_SECRET", "" + ) + result = { + os.environ.get("AI4SCI_EVO_MODEL_SECRET_KEY_ID", "model-secret-v1"): material + } + if not result or any(len(value.encode("utf-8")) < 32 for value in result.values()): + raise RuntimeError("model secret encryption keys must contain at least 32 bytes") + return result + + def _connect(self) -> sqlite3.Connection: + connection = sqlite3.connect(self.path, timeout=5) + connection.row_factory = sqlite3.Row + connection.execute("PRAGMA foreign_keys=ON") + connection.execute("PRAGMA busy_timeout=5000") + connection.execute("PRAGMA synchronous=FULL") + connection.execute("PRAGMA journal_mode=WAL") + return connection + + def _init_schema(self) -> None: + with self._lock, self._connect() as connection: + connection.executescript( + """ + CREATE TABLE IF NOT EXISTS model_config_revisions ( + revision INTEGER PRIMARY KEY, + rollback_of_revision INTEGER REFERENCES model_config_revisions(revision), + schema_version INTEGER NOT NULL CHECK (schema_version = 4), + payload_json TEXT NOT NULL, + payload_hash TEXT NOT NULL, + runtime_projection_json TEXT NOT NULL, + runtime_projection_hash TEXT NOT NULL, + operational_policy_revision TEXT NOT NULL, + operational_policy_json TEXT NOT NULL, + operational_policy_hash TEXT NOT NULL, + identity_key_id TEXT NOT NULL, + revision_hmac TEXT NOT NULL, + adapter_registry_revision TEXT NOT NULL, + created_by TEXT NOT NULL, + created_at TEXT NOT NULL, + operation_id TEXT NOT NULL UNIQUE + ); + CREATE TABLE IF NOT EXISTS active_model_config ( + singleton INTEGER PRIMARY KEY CHECK (singleton = 1), + revision INTEGER NOT NULL REFERENCES model_config_revisions(revision), + security_epoch INTEGER NOT NULL DEFAULT 0, + evidence_epoch INTEGER NOT NULL DEFAULT 0, + updated_at TEXT NOT NULL + ); + CREATE TABLE IF NOT EXISTS provider_secret_versions ( + provider_id TEXT NOT NULL, + version INTEGER NOT NULL, + cipher_version TEXT NOT NULL, + encryption_key_id TEXT NOT NULL, + ciphertext BLOB NOT NULL, + credential_fingerprint TEXT NOT NULL, + masked_value TEXT NOT NULL, + status TEXT NOT NULL CHECK (status IN ('active','retired','revoked')), + created_by TEXT NOT NULL, + created_at TEXT NOT NULL, + revoked_at TEXT, + revoke_reason TEXT, + PRIMARY KEY (provider_id, version) + ); + CREATE UNIQUE INDEX IF NOT EXISTS uq_provider_active_secret + ON provider_secret_versions(provider_id) WHERE status = 'active'; + CREATE TABLE IF NOT EXISTS config_provider_secret_bindings ( + config_revision INTEGER NOT NULL REFERENCES model_config_revisions(revision), + provider_id TEXT NOT NULL, + secret_version INTEGER NOT NULL, + PRIMARY KEY (config_revision, provider_id), + FOREIGN KEY (provider_id, secret_version) + REFERENCES provider_secret_versions(provider_id, version) + ); + CREATE TABLE IF NOT EXISTS capability_evidence ( + evidence_id TEXT PRIMARY KEY, + config_revision INTEGER NOT NULL REFERENCES model_config_revisions(revision), + provider_id TEXT NOT NULL, + model_profile_id TEXT NOT NULL, + route_semantics_hash TEXT NOT NULL, + probe_kind TEXT NOT NULL, + outcome TEXT NOT NULL CHECK (outcome IN ('supported','failed')), + failure_code TEXT, + adapter_revision TEXT NOT NULL, + implementation_fingerprint TEXT NOT NULL, + credential_version INTEGER NOT NULL, + credential_fingerprint TEXT NOT NULL, + results_json TEXT NOT NULL, + verified_at TEXT NOT NULL, + expires_at TEXT NOT NULL, + identity_key_id TEXT NOT NULL, + evidence_hmac TEXT NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_capability_evidence_lookup + ON capability_evidence(config_revision, model_profile_id, + route_semantics_hash, probe_kind, verified_at DESC, evidence_id DESC); + CREATE TABLE IF NOT EXISTS model_config_operations ( + operation_id TEXT PRIMARY KEY, + request_hash TEXT NOT NULL, + expected_revision INTEGER NOT NULL, + result_revision INTEGER, + status TEXT NOT NULL CHECK (status IN ( + 'RECEIVED','VALIDATING','PROBING','COMMITTING', + 'SUCCEEDED','FAILED','FAILED_RETRYABLE' + )), + lease_owner TEXT, + lease_expires_at TEXT, + stage TEXT, + progress_json TEXT, + error_code TEXT, + error_details_json TEXT, + result_json TEXT, + actor TEXT NOT NULL, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL + ); + CREATE TABLE IF NOT EXISTS model_config_audit ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + operation_id TEXT NOT NULL, + actor TEXT NOT NULL, + action TEXT NOT NULL, + old_revision INTEGER, + new_revision INTEGER, + summary_json TEXT NOT NULL, + created_at TEXT NOT NULL + ); + CREATE TABLE IF NOT EXISTS model_runtime_observations ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + config_revision INTEGER NOT NULL, + provider_id TEXT NOT NULL, + model_profile_id TEXT NOT NULL, + provider_model_id TEXT NOT NULL, + api_mode TEXT NOT NULL, + purpose TEXT NOT NULL, + outcome TEXT NOT NULL, + error_code TEXT, + strategy_json TEXT NOT NULL, + observed_at TEXT NOT NULL + ); + CREATE INDEX IF NOT EXISTS idx_model_runtime_observations_profile + ON model_runtime_observations( + config_revision, model_profile_id, id DESC + ); + CREATE TABLE IF NOT EXISTS legacy_model_alias_mappings ( + legacy_alias TEXT PRIMARY KEY, + model_profile_id TEXT, + source_config_revision INTEGER NOT NULL, + migration_status TEXT NOT NULL CHECK (migration_status IN ('mapped','stale')), + detail_code TEXT + ); + """ + ) + active_columns = { + str(row[1]) + for row in connection.execute("PRAGMA table_info(active_model_config)") + } + if "evidence_epoch" not in active_columns: + connection.execute( + "ALTER TABLE active_model_config " + "ADD COLUMN evidence_epoch INTEGER NOT NULL DEFAULT 0" + ) + evidence_columns = { + str(row[1]) + for row in connection.execute("PRAGMA table_info(capability_evidence)") + } + if "credential_version" not in evidence_columns: + connection.execute( + "ALTER TABLE capability_evidence " + "ADD COLUMN credential_version INTEGER NOT NULL DEFAULT 0" + ) + connection.execute( + """UPDATE capability_evidence + SET credential_version=COALESCE(( + SELECT binding.secret_version + FROM config_provider_secret_bindings AS binding + WHERE binding.config_revision=capability_evidence.config_revision + AND binding.provider_id=capability_evidence.provider_id + ), 0)""" + ) + operation_columns = { + str(row[1]) + for row in connection.execute("PRAGMA table_info(model_config_operations)") + } + for name in ("stage", "progress_json", "error_details_json"): + if name not in operation_columns: + connection.execute( + f"ALTER TABLE model_config_operations ADD COLUMN {name} TEXT" + ) + connection.execute( + """CREATE INDEX IF NOT EXISTS idx_capability_evidence_semantics + ON capability_evidence( + provider_id, model_profile_id, route_semantics_hash, + credential_version, adapter_revision, probe_kind, + verified_at DESC, evidence_id DESC + )""" + ) + try: + os.chmod(self.path, 0o600) + except OSError: + pass + + def current_generation(self) -> ActiveGeneration: + with self._connect() as connection: + row = connection.execute( + """SELECT revision, security_epoch, evidence_epoch + FROM active_model_config WHERE singleton=1""" + ).fetchone() + return ActiveGeneration( + int(row["revision"]) if row else 0, + int(row["security_epoch"]) if row else 0, + int(row["evidence_epoch"]) if row else 0, + get_adapter_registry().registry_revision, + ) + + def current_revision(self) -> int: + return self.current_generation().revision + + def provider_available_for_active(self, provider_id: str) -> bool: + generation = self.current_generation() + if generation.revision < 1: + return False + with self._connect() as connection: + row = connection.execute( + """SELECT secret.status, secret.encryption_key_id + FROM config_provider_secret_bindings AS binding + JOIN provider_secret_versions AS secret + ON secret.provider_id=binding.provider_id + AND secret.version=binding.secret_version + WHERE binding.config_revision=? AND binding.provider_id=?""", + (generation.revision, provider_id), + ).fetchone() + return bool( + row is not None + and str(row["status"]) != "revoked" + and str(row["encryption_key_id"]) in self._fernets + ) + + def _projection_route_facts( + self, projection: Mapping[str, Any] + ) -> dict[tuple[str, str], dict[str, Any]]: + config = EvoModelConfig.parse(projection, require_evidence=False) + semantics_key = self.identity_key_ring.derive( + config.config_identity_key_id, "ai4sci/route-semantics-hash/v3" + ) + facts: dict[tuple[str, str], dict[str, Any]] = {} + for selector_id in config.route_selectors: + for route in config.concrete_routes(selector_id): + provider = config.providers[route.provider] + model = provider.models[route.model] + reusable_payload = dict(route_semantics_payload(config, route)) + reusable_payload.pop("config_revision", None) + facts[(route.provider, route.model)] = { + "route_semantics_hash": route_semantics_hash( + config, route, semantics_key + ), + "reusable_semantics_hash": hmac_id( + semantics_key, reusable_payload + ), + "adapter_revision": provider.adapter_revision, + "implementation_fingerprint": provider.implementation_fingerprint, + "required_capabilities": tuple( + name + for name, enabled in model.capabilities.items() + if enabled and name != "text" + ), + } + return facts + + @staticmethod + def _evidence_runtime_results(row: sqlite3.Row) -> Mapping[str, Any]: + try: + runtime_evidence = json.loads(str(row["results_json"])) + except (TypeError, json.JSONDecodeError) as exc: + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") from exc + results = runtime_evidence.get("results") + if not isinstance(results, Mapping): + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + return results + + @staticmethod + def _evidence_not_expired(row: sqlite3.Row, now: datetime) -> bool: + try: + return datetime.fromisoformat(str(row["expires_at"])) > now + except ValueError: + return False + + def reusable_probe_results( + self, + projection: Mapping[str, Any], + *, + fingerprints: Mapping[str, str], + versions: Mapping[str, int], + ) -> tuple[ + dict[str, dict[str, dict[str, str]]], + dict[str, dict[str, tuple[str, str]]], + frozenset[tuple[str, str]], + ]: + """Return valid observed probe results and profiles requiring a probe.""" + + facts = self._projection_route_facts(projection) + with self._connect() as connection: + rows = connection.execute( + """SELECT evidence.*, revision.runtime_projection_json + AS source_projection_json + FROM capability_evidence AS evidence + JOIN model_config_revisions AS revision + ON revision.revision=evidence.config_revision + WHERE evidence.outcome='supported' + ORDER BY evidence.verified_at DESC, evidence.evidence_id DESC""" + ).fetchall() + by_profile: dict[tuple[str, str], list[sqlite3.Row]] = {} + for row in rows: + by_profile.setdefault( + (str(row["provider_id"]), str(row["model_profile_id"])), [] + ).append(row) + now = datetime.now(UTC) + observed: dict[str, dict[str, dict[str, str]]] = {} + windows: dict[str, dict[str, tuple[str, str]]] = {} + missing: set[tuple[str, str]] = set() + source_facts: dict[int, dict[tuple[str, str], dict[str, Any]]] = {} + for key, fact in facts.items(): + provider_id, profile_id = key + matched: sqlite3.Row | None = None + for row in by_profile.get(key, []): + try: + self._verify_evidence_row(row) + except EvoRuntimeError: + continue + source_revision = int(row["config_revision"]) + if source_revision not in source_facts: + try: + source_projection = json.loads( + str(row["source_projection_json"]) + ) + source_facts[source_revision] = self._projection_route_facts( + source_projection + ) + except (TypeError, json.JSONDecodeError, EvoRuntimeError): + source_facts[source_revision] = {} + source_fact = source_facts[source_revision].get(key) + if ( + source_fact is None + or str(row["route_semantics_hash"]) + != source_fact["route_semantics_hash"] + or source_fact["reusable_semantics_hash"] + != fact["reusable_semantics_hash"] + or str(row["adapter_revision"]) != fact["adapter_revision"] + or str(row["implementation_fingerprint"]) + != fact["implementation_fingerprint"] + or int(row["credential_version"]) + != int(versions.get(provider_id, 0)) + or str(row["credential_fingerprint"]) + != str(fingerprints.get(provider_id, "")) + or not self._evidence_not_expired(row, now) + ): + continue + results = self._evidence_runtime_results(row) + if results.get("connectivity") != "supported" or any( + results.get(name) != "supported" + for name in fact["required_capabilities"] + ): + continue + matched = row + break + if matched is None: + missing.add(key) + continue + runtime_results = self._evidence_runtime_results(matched) + profile_results = {"connectivity": "supported"} + for name in fact["required_capabilities"]: + probe_name = "reasoning" if name == "thinking" else name + profile_results[probe_name] = str(runtime_results[name]) + observed.setdefault(provider_id, {})[profile_id] = profile_results + windows.setdefault(provider_id, {})[profile_id] = ( + str(matched["verified_at"]), + str(matched["expires_at"]), + ) + return observed, windows, frozenset(missing) + + def get_model_availability(self) -> dict[str, dict[str, Any]]: + """Compute browser/catalog availability from persisted config facts.""" + + generation = self.current_generation() + if generation.revision < 1: + return {} + with self._connect() as connection: + revision_row = connection.execute( + """SELECT payload_json, runtime_projection_json + FROM model_config_revisions WHERE revision=?""", + (generation.revision,), + ).fetchone() + binding_rows = connection.execute( + """SELECT binding.provider_id, secret.status, secret.encryption_key_id + FROM config_provider_secret_bindings AS binding + JOIN provider_secret_versions AS secret + ON secret.provider_id=binding.provider_id + AND secret.version=binding.secret_version + WHERE binding.config_revision=?""", + (generation.revision,), + ).fetchall() + if revision_row is None: + return {} + payload = json.loads(str(revision_row["payload_json"])) + projection = json.loads(str(revision_row["runtime_projection_json"])) + facts = self._projection_route_facts(projection) + bindings = {str(row["provider_id"]): row for row in binding_rows} + result: dict[str, dict[str, Any]] = {} + for provider in payload.get("providers") or []: + provider_id = str(provider["provider_id"]) + binding = bindings.get(provider_id) + for model in provider.get("models") or []: + profile_id = str(model["model_profile_id"]) + enabled = bool(provider.get("enabled", True)) and bool( + model.get("enabled", True) + ) + selectable = ( + enabled + and binding is not None + and str(binding["status"]) != "revoked" + and str(binding["encryption_key_id"]) in self._fernets + and facts.get((provider_id, profile_id)) is not None + ) + result[profile_id] = { + "model_profile_id": profile_id, + "enabled": enabled, + "selectable": selectable, + } + return result + + def recover_expired_operations(self) -> int: + """Make abandoned publish operations explicitly retryable after a crash.""" + + now = _utc_now() + with self._lock, self._connect() as connection: + cursor = connection.execute( + """UPDATE model_config_operations + SET status='FAILED_RETRYABLE', lease_owner=NULL, + lease_expires_at=NULL, error_code='OPERATION_LEASE_EXPIRED', + stage='RECOVERY', + error_details_json='[{"code":"OPERATION_LEASE_EXPIRED","path":"operation"}]', + updated_at=? + WHERE status IN ('RECEIVED','VALIDATING','PROBING','COMMITTING') + AND lease_expires_at IS NOT NULL AND lease_expires_at < ?""", + (now, now), + ) + return max(int(cursor.rowcount), 0) + + def _credential_fingerprint(self, provider_id: str, value: str) -> str: + _, key = self.identity_key_ring.derive_current(_CREDENTIAL_FINGERPRINT_INFO) + return hmac_id(key, {"provider_id": provider_id, "credential": value}) + + @staticmethod + def _mask(value: str) -> str: + return "*" * len(value) if len(value) <= 8 else f"{value[:4]}...{value[-4:]}" + + def _decrypt_row(self, row: sqlite3.Row) -> str: + try: + fernet = self._fernets[str(row["encryption_key_id"])] + return fernet.decrypt(bytes(row["ciphertext"])).decode("utf-8") + except (KeyError, InvalidToken, UnicodeDecodeError) as exc: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") from exc + + def plan_credentials( + self, + payload: Mapping[str, Any], + credential_updates: Mapping[str, str], + ) -> CredentialPlan: + generation = self.current_generation() + bindings: dict[str, int] = {} + fingerprints: dict[str, str] = {} + plaintext: dict[str, str] = {} + new_versions: set[str] = set() + with self._connect() as connection: + for provider in payload["providers"]: + if not provider["enabled"]: + continue + provider_id = str(provider["provider_id"]) + supplied = str(credential_updates.get(provider_id) or "") + if supplied and any(char in supplied for char in "\r\n\0"): + raise EvoRuntimeError("LLM_SECRET_INVALID") + current = connection.execute( + """SELECT * FROM provider_secret_versions + WHERE provider_id=? AND status='active'""", + (provider_id,), + ).fetchone() + if supplied: + current_value = self._decrypt_row(current) if current is not None else "" + if current is not None and hmac.compare_digest(current_value, supplied): + version = int(current["version"]) + fingerprint = str(current["credential_fingerprint"]) + else: + fingerprint = self._credential_fingerprint(provider_id, supplied) + maximum = connection.execute( + "SELECT COALESCE(MAX(version),0) AS version FROM provider_secret_versions WHERE provider_id=?", + (provider_id,), + ).fetchone() + version = int(maximum["version"]) + 1 + new_versions.add(provider_id) + value = supplied + elif current is not None: + version = int(current["version"]) + fingerprint = str(current["credential_fingerprint"]) + value = self._decrypt_row(current) + else: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + bindings[provider_id] = version + fingerprints[provider_id] = fingerprint + plaintext[provider_id] = value + return CredentialPlan( + bindings, + fingerprints, + plaintext, + frozenset(new_versions), + generation.security_epoch, + ) + + def preview_revision( + self, + payload: Mapping[str, Any], + bindings: Mapping[str, int], + *, + expected_revision: int, + revision_floor: int = 0, + ) -> tuple[bool, int]: + with self._connect() as connection: + active = connection.execute( + "SELECT revision FROM active_model_config WHERE singleton=1" + ).fetchone() + current = int(active["revision"]) if active else 0 + if current != expected_revision: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + row = connection.execute( + "SELECT payload_hash FROM model_config_revisions WHERE revision=?", + (current,), + ).fetchone() + current_bindings = { + str(item["provider_id"]): int(item["secret_version"]) + for item in connection.execute( + """SELECT provider_id, secret_version + FROM config_provider_secret_bindings WHERE config_revision=?""", + (current,), + ).fetchall() + } + changed = ( + row is None + or str(row["payload_hash"]) != sha256_id(payload) + or current_bindings != dict(bindings) + ) + return changed, ( + max(current + 1, int(revision_floor) + 1) if changed else current + ) + + def begin_operation( + self, + operation_id: str, + *, + request_hash: str, + expected_revision: int, + actor: str, + lease_owner: str, + lease_seconds: int = 120, + ) -> Mapping[str, Any] | None: + now = datetime.now(UTC) + expires = (now + timedelta(seconds=lease_seconds)).isoformat(timespec="microseconds") + with self._lock, self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + row = connection.execute( + "SELECT * FROM model_config_operations WHERE operation_id=?", + (operation_id,), + ).fetchone() + if row is not None: + if str(row["request_hash"]) != request_hash: + raise EvoRuntimeError("IDEMPOTENCY_CONFLICT") + lease_expired = bool( + row["lease_expires_at"] + and str(row["lease_expires_at"]) + < now.isoformat(timespec="microseconds") + ) + if str(row["status"]) == "FAILED_RETRYABLE" or ( + lease_expired + and str(row["status"]) + in {"RECEIVED", "VALIDATING", "PROBING", "COMMITTING"} + ): + connection.execute( + """UPDATE model_config_operations SET status='RECEIVED', + lease_owner=?, lease_expires_at=?, stage='RECEIVED', + error_code=NULL, error_details_json=NULL, progress_json=NULL, + updated_at=? WHERE operation_id=?""", + ( + lease_owner, + expires, + now.isoformat(timespec="microseconds"), + operation_id, + ), + ) + return None + return dict(row) + connection.execute( + """INSERT INTO model_config_operations + (operation_id, request_hash, expected_revision, status, + lease_owner, lease_expires_at, stage, actor, created_at, updated_at) + VALUES (?, ?, ?, 'RECEIVED', ?, ?, 'RECEIVED', ?, ?, ?)""", + ( + operation_id, + request_hash, + expected_revision, + lease_owner, + expires, + actor[:256], + now.isoformat(timespec="microseconds"), + now.isoformat(timespec="microseconds"), + ), + ) + return None + + def set_operation_status( + self, + operation_id: str, + status: str, + *, + error_code: str | None = None, + stage: str | None = None, + progress: Mapping[str, Any] | None = None, + error_details: Sequence[Mapping[str, Any]] | None = None, + result: Mapping[str, Any] | None = None, + ) -> None: + assignments = ["status=?", "stage=?", "error_code=?", "updated_at=?"] + values: list[Any] = [status, stage or status, error_code, _utc_now()] + if progress is not None: + assignments.append("progress_json=?") + values.append(canonical_json_v1(progress).decode()) + if error_details is not None: + assignments.append("error_details_json=?") + values.append(canonical_json_v1(error_details).decode()) + if result is not None: + assignments.append("result_json=?") + values.append(canonical_json_v1(result).decode()) + values.append(operation_id) + with self._connect() as connection: + connection.execute( + f"UPDATE model_config_operations SET {', '.join(assignments)} " + "WHERE operation_id=?", + values, + ) + + def get_operation(self, operation_id: str) -> Mapping[str, Any] | None: + with self._connect() as connection: + row = connection.execute( + "SELECT * FROM model_config_operations WHERE operation_id=?", + (operation_id,), + ).fetchone() + if row is None: + return None + result = dict(row) + for source, target in ( + ("progress_json", "progress"), + ("error_details_json", "error_details"), + ("result_json", "result"), + ): + if result.get(source): + result[target] = json.loads(str(result[source])) + result.pop(source, None) + return result + + def _revision_signature_payload( + self, + *, + revision: int, + payload_hash: str, + projection_hash: str, + policy_hash: str, + adapter_registry_revision: str, + bindings: Mapping[str, int], + ) -> dict[str, Any]: + return { + "revision": revision, + "payload_hash": payload_hash, + "runtime_projection_hash": projection_hash, + "operational_policy_revision": OPERATIONAL_POLICY_V1["revision"], + "operational_policy_hash": policy_hash, + "adapter_registry_revision": adapter_registry_revision, + "secret_bindings": dict(sorted(bindings.items())), + } + + def commit( + self, + payload: Mapping[str, Any], + *, + expected_revision: int, + operation_id: str, + actor: str, + request_hash: str, + credential_plan: CredentialPlan, + evidence: Sequence[Mapping[str, Any]], + rollback_of_revision: int | None = None, + revision_floor: int = 0, + legacy_alias_mappings: Mapping[str, str | None] | None = None, + ) -> UnifiedSaveResult: + # Kept in the public method signature for callers upgrading from the + # probe workflow. Evidence is no longer persisted or used to decide + # whether a configuration may be published or invoked. + del evidence + payload_json = canonical_json_v1(payload).decode() + payload_hash = sha256_id(payload) + policy_json = canonical_json_v1(OPERATIONAL_POLICY_V1).decode() + policy_hash = sha256_id(OPERATIONAL_POLICY_V1) + identity_key_id, revision_key = self.identity_key_ring.derive_current( + _REVISION_HMAC_INFO + ) + adapter_revision = get_adapter_registry().registry_revision + with self._lock, self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + operation = connection.execute( + "SELECT * FROM model_config_operations WHERE operation_id=?", + (operation_id,), + ).fetchone() + if operation is None or str(operation["request_hash"]) != request_hash: + raise EvoRuntimeError("IDEMPOTENCY_CONFLICT") + if str(operation["status"]) == "SUCCEEDED": + result = json.loads(str(operation["result_json"] or "{}")) + return UnifiedSaveResult(**result) + active = connection.execute( + """SELECT revision, security_epoch, evidence_epoch + FROM active_model_config WHERE singleton=1""" + ).fetchone() + current_revision = int(active["revision"]) if active else 0 + security_epoch = int(active["security_epoch"]) if active else 0 + evidence_epoch = int(active["evidence_epoch"]) if active else 0 + if current_revision != expected_revision: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + if security_epoch != credential_plan.base_security_epoch: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + for provider_id, version in credential_plan.bindings.items(): + current = connection.execute( + """SELECT * FROM provider_secret_versions + WHERE provider_id=? AND status='active'""", + (provider_id,), + ).fetchone() + if provider_id in credential_plan.new_versions: + maximum = connection.execute( + "SELECT COALESCE(MAX(version),0) AS version FROM provider_secret_versions WHERE provider_id=?", + (provider_id,), + ).fetchone() + if int(maximum["version"]) + 1 != version: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + connection.execute( + "UPDATE provider_secret_versions SET status='retired' WHERE provider_id=? AND status='active'", + (provider_id,), + ) + value = credential_plan.plaintext[provider_id] + connection.execute( + """INSERT INTO provider_secret_versions + (provider_id, version, cipher_version, encryption_key_id, + ciphertext, credential_fingerprint, masked_value, status, + created_by, created_at) + VALUES (?, ?, 'fernet-v1', ?, ?, ?, ?, 'active', ?, ?)""", + ( + provider_id, + version, + self.current_encryption_key_id, + self._fernets[self.current_encryption_key_id].encrypt( + value.encode("utf-8") + ), + credential_plan.fingerprints[provider_id], + self._mask(value), + actor[:256], + _utc_now(), + ), + ) + elif current is None or int(current["version"]) != version: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + existing_bindings = { + str(row["provider_id"]): int(row["secret_version"]) + for row in connection.execute( + "SELECT provider_id, secret_version FROM config_provider_secret_bindings WHERE config_revision=?", + (current_revision,), + ).fetchall() + } + changed = payload_hash != ( + str( + ( + connection.execute( + "SELECT payload_hash FROM model_config_revisions WHERE revision=?", + (current_revision,), + ).fetchone() + or {"payload_hash": ""} + )["payload_hash"] + ) + ) or dict(credential_plan.bindings) != existing_bindings + target_revision = ( + max(current_revision + 1, int(revision_floor) + 1) + if changed + else current_revision + ) + projection = project_v4_to_v3( + payload, + revision=target_revision, + identity_key_id=identity_key_id, + bindings=credential_plan.bindings, + ) + projection_json = canonical_json_v1(projection).decode() + projection_hash = sha256_id(projection) + if changed: + connection.execute( + "UPDATE provider_secret_versions SET status='retired' WHERE status='active'" + ) + for provider_id, version in credential_plan.bindings.items(): + cursor = connection.execute( + """UPDATE provider_secret_versions SET status='active' + WHERE provider_id=? AND version=? AND status='retired'""", + (provider_id, version), + ) + if cursor.rowcount != 1: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + signature_payload = self._revision_signature_payload( + revision=target_revision, + payload_hash=payload_hash, + projection_hash=projection_hash, + policy_hash=policy_hash, + adapter_registry_revision=adapter_revision, + bindings=credential_plan.bindings, + ) + connection.execute( + """INSERT INTO model_config_revisions + (revision, rollback_of_revision, schema_version, payload_json, + payload_hash, runtime_projection_json, runtime_projection_hash, + operational_policy_revision, operational_policy_json, + operational_policy_hash, identity_key_id, revision_hmac, + adapter_registry_revision, created_by, created_at, operation_id) + VALUES (?, ?, 4, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", + ( + target_revision, + rollback_of_revision, + payload_json, + payload_hash, + projection_json, + projection_hash, + OPERATIONAL_POLICY_V1["revision"], + policy_json, + policy_hash, + identity_key_id, + hmac_id(revision_key, signature_payload), + adapter_revision, + actor[:256], + _utc_now(), + operation_id, + ), + ) + connection.executemany( + """INSERT INTO config_provider_secret_bindings + (config_revision, provider_id, secret_version) VALUES (?, ?, ?)""", + [ + (target_revision, provider_id, version) + for provider_id, version in credential_plan.bindings.items() + ], + ) + connection.execute( + """INSERT INTO active_model_config + (singleton, revision, security_epoch, evidence_epoch, updated_at) + VALUES (1, ?, ?, ?, ?) + ON CONFLICT(singleton) DO UPDATE SET revision=excluded.revision, + security_epoch=excluded.security_epoch, + evidence_epoch=excluded.evidence_epoch, + updated_at=excluded.updated_at""", + (target_revision, security_epoch, evidence_epoch, _utc_now()), + ) + for legacy_alias, model_profile_id in dict( + legacy_alias_mappings or {} + ).items(): + connection.execute( + """INSERT INTO legacy_model_alias_mappings + (legacy_alias, model_profile_id, source_config_revision, + migration_status, detail_code) + VALUES (?, ?, ?, ?, ?) + ON CONFLICT(legacy_alias) DO UPDATE SET + model_profile_id=excluded.model_profile_id, + source_config_revision=excluded.source_config_revision, + migration_status=excluded.migration_status, + detail_code=excluded.detail_code""", + ( + str(legacy_alias), + model_profile_id, + int(revision_floor), + "mapped" if model_profile_id else "stale", + None if model_profile_id else "LEGACY_ALIAS_UNRESOLVED", + ), + ) + result = UnifiedSaveResult( + operation_id, + target_revision, + changed, + False, + ) + result_json = canonical_json_v1( + { + "operation_id": result.operation_id, + "config_revision": result.config_revision, + "changed": result.changed, + "revalidated": result.revalidated, + "status": result.status, + } + ).decode() + connection.execute( + """UPDATE model_config_operations SET status='SUCCEEDED', + stage='SUCCEEDED', error_code=NULL, error_details_json=NULL, + result_revision=?, result_json=?, lease_owner=NULL, + lease_expires_at=NULL, updated_at=? WHERE operation_id=?""", + (target_revision, result_json, _utc_now(), operation_id), + ) + connection.execute( + """INSERT INTO model_config_audit + (operation_id, actor, action, old_revision, new_revision, + summary_json, created_at) VALUES (?, ?, ?, ?, ?, ?, ?)""", + ( + operation_id, + actor[:256], + "CONFIG_PUBLISH" if changed else "CONFIG_REVALIDATE", + current_revision, + target_revision, + canonical_json_v1( + { + "changed": changed, + "credential_changed": bool(credential_plan.new_versions), + "evidence_count": 0, + } + ).decode(), + _utc_now(), + ), + ) + return result + + def _insert_evidence( + self, + connection: sqlite3.Connection, + *, + config_revision: int, + evidence: Sequence[Mapping[str, Any]], + fingerprints: Mapping[str, str], + versions: Mapping[str, int], + ) -> None: + identity_key_id, evidence_key = self.identity_key_ring.derive_current( + _EVIDENCE_HMAC_INFO + ) + for item in evidence: + provider_id = str(item["provider_ref"]) + model_profile_id = str(item["model_ref"]) + evidence_id = str(item.get("evidence_id") or _uuid7()) + verified_at = str(item["verified_at"]) + expires_at = str(item["evidence_expires_at"]) + runtime_evidence = { + key: value + for key, value in item.items() + if key + in { + "provider_ref", + "model_ref", + "adapter_id", + "adapter_revision", + "implementation_fingerprint", + "wire_protocol", + "provider_model_id", + "resolved_model_revision", + "version_policy", + "reproducible", + "api_mode", + "tool_call_transport", + "base_url_fingerprint", + "secret_version", + "route_semantics_hash", + "fixture_digest", + "verified_at", + "evidence_expires_at", + "results", + } + } + record = { + "evidence_id": evidence_id, + "config_revision": config_revision, + "provider_id": provider_id, + "model_profile_id": model_profile_id, + "route_semantics_hash": str(item["route_semantics_hash"]), + "probe_kind": str(item.get("probe_kind") or "connectivity"), + "outcome": str(item.get("outcome") or "supported"), + "failure_code": item.get("failure_code"), + "adapter_revision": str(item["adapter_revision"]), + "implementation_fingerprint": str(item["implementation_fingerprint"]), + "credential_fingerprint": str(fingerprints[provider_id]), + "results_json": canonical_json_v1(runtime_evidence).decode(), + "verified_at": verified_at, + "expires_at": expires_at, + "identity_key_id": identity_key_id, + } + connection.execute( + """INSERT INTO capability_evidence + (evidence_id, config_revision, provider_id, model_profile_id, + route_semantics_hash, probe_kind, outcome, failure_code, + adapter_revision, implementation_fingerprint, + credential_version, credential_fingerprint, results_json, + verified_at, expires_at, identity_key_id, evidence_hmac) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", + ( + evidence_id, + config_revision, + provider_id, + model_profile_id, + record["route_semantics_hash"], + record["probe_kind"], + record["outcome"], + record["failure_code"], + record["adapter_revision"], + record["implementation_fingerprint"], + int(versions[provider_id]), + record["credential_fingerprint"], + record["results_json"], + verified_at, + expires_at, + identity_key_id, + hmac_id(evidence_key, record), + ), + ) + + def _load_revision_row(self, revision: int) -> tuple[sqlite3.Row, list[sqlite3.Row]]: + with self._connect() as connection: + row = connection.execute( + "SELECT * FROM model_config_revisions WHERE revision=?", (revision,) + ).fetchone() + bindings = connection.execute( + """SELECT binding.provider_id, binding.secret_version, + secret.credential_fingerprint + FROM config_provider_secret_bindings AS binding + JOIN provider_secret_versions AS secret + ON secret.provider_id=binding.provider_id + AND secret.version=binding.secret_version + WHERE binding.config_revision=?""", + (revision,), + ).fetchall() + if row is None: + raise EvoRuntimeError("CONFIG_REVISION_NOT_FOUND") + try: + policy = json.loads(str(row["operational_policy_json"])) + except (TypeError, json.JSONDecodeError) as exc: + raise EvoRuntimeError("LLM_ROUTE_CONFIGURATION_REQUIRED") from exc + if ( + str(row["operational_policy_revision"]) != OPERATIONAL_POLICY_V1["revision"] + or policy != OPERATIONAL_POLICY_V1 + or sha256_id(policy) != str(row["operational_policy_hash"]) + ): + raise EvoRuntimeError("LLM_ROUTE_CONFIGURATION_REQUIRED") + signature_payload = self._revision_signature_payload( + revision=revision, + payload_hash=str(row["payload_hash"]), + projection_hash=str(row["runtime_projection_hash"]), + policy_hash=str(row["operational_policy_hash"]), + adapter_registry_revision=str(row["adapter_registry_revision"]), + bindings={str(item["provider_id"]): int(item["secret_version"]) for item in bindings}, + ) + try: + key = self.identity_key_ring.derive(str(row["identity_key_id"]), _REVISION_HMAC_INFO) + except KeyError as exc: + raise EvoRuntimeError("CONFIG_IDENTITY_KEY_UNKNOWN") from exc + if not hmac.compare_digest( + hmac_id(key, signature_payload), str(row["revision_hmac"]) + ): + raise EvoRuntimeError("LLM_ROUTE_CONFIGURATION_REQUIRED") + if sha256_id(json.loads(str(row["payload_json"]))) != str(row["payload_hash"]): + raise EvoRuntimeError("LLM_ROUTE_CONFIGURATION_REQUIRED") + projection = json.loads(str(row["runtime_projection_json"])) + if sha256_id(projection) != str(row["runtime_projection_hash"]): + raise EvoRuntimeError("LLM_ROUTE_CONFIGURATION_REQUIRED") + return row, [] + + def _verify_evidence_row(self, row: sqlite3.Row) -> None: + record = { + "evidence_id": str(row["evidence_id"]), + "config_revision": int(row["config_revision"]), + "provider_id": str(row["provider_id"]), + "model_profile_id": str(row["model_profile_id"]), + "route_semantics_hash": str(row["route_semantics_hash"]), + "probe_kind": str(row["probe_kind"]), + "outcome": str(row["outcome"]), + "failure_code": row["failure_code"], + "adapter_revision": str(row["adapter_revision"]), + "implementation_fingerprint": str(row["implementation_fingerprint"]), + "credential_fingerprint": str(row["credential_fingerprint"]), + "results_json": str(row["results_json"]), + "verified_at": str(row["verified_at"]), + "expires_at": str(row["expires_at"]), + "identity_key_id": str(row["identity_key_id"]), + } + try: + key = self.identity_key_ring.derive( + record["identity_key_id"], _EVIDENCE_HMAC_INFO + ) + except KeyError as exc: + raise EvoRuntimeError("CONFIG_IDENTITY_KEY_UNKNOWN") from exc + if not hmac.compare_digest( + hmac_id(key, record), str(row["evidence_hmac"]) + ): + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + try: + runtime_evidence = json.loads(record["results_json"]) + except json.JSONDecodeError as exc: + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") from exc + if ( + str(runtime_evidence.get("provider_ref") or "") != record["provider_id"] + or str(runtime_evidence.get("model_ref") or "") + != record["model_profile_id"] + or str(runtime_evidence.get("route_semantics_hash") or "") + != record["route_semantics_hash"] + or str(runtime_evidence.get("adapter_revision") or "") + != record["adapter_revision"] + or str(runtime_evidence.get("implementation_fingerprint") or "") + != record["implementation_fingerprint"] + or int(runtime_evidence.get("secret_version") or 0) + != int(row["credential_version"]) + ): + raise EvoRuntimeError("CAPABILITY_EVIDENCE_STALE") + + def load_revision(self, revision: int) -> EvoModelConfig: + row, _ = self._load_revision_row(revision) + projection = json.loads(str(row["runtime_projection_json"])) + # Historical evidence is deliberately excluded from the runtime + # projection. It is not configuration state and cannot gate calls. + projection["capability_evidence"] = [] + return EvoModelConfig.parse(projection, require_evidence=False) + + def get_runtime_projection(self, revision: int) -> dict[str, Any]: + row, _ = self._load_revision_row(revision) + projection = json.loads(str(row["runtime_projection_json"])) + projection["capability_evidence"] = [] + EvoModelConfig.parse(projection, require_evidence=False) + return projection + + def load(self) -> EvoModelConfig: + revision = self.current_revision() + if revision < 1: + raise EvoRuntimeError("LLM_ROUTE_CONFIGURATION_REQUIRED") + return self.load_revision(revision) + + def preflight(self) -> None: + """Verify the active snapshot, key registry, and ciphertext.""" + + revision = self.current_revision() + if revision < 1: + return + self.load_revision(revision) + with self._connect() as connection: + rows = connection.execute( + """SELECT secret.* + FROM config_provider_secret_bindings AS binding + JOIN provider_secret_versions AS secret + ON secret.provider_id=binding.provider_id + AND secret.version=binding.secret_version + WHERE binding.config_revision=?""", + (revision,), + ).fetchall() + for row in rows: + if str(row["status"]) != "revoked": + self._decrypt_row(row) + + def record_runtime_observation( + self, + *, + config_revision: int, + provider_id: str, + model_profile_id: str, + provider_model_id: str, + api_mode: str, + purpose: str, + outcome: str, + error_code: str | None, + strategy: Mapping[str, Any], + ) -> None: + """Persist a small, non-authoritative record of a real model call. + + This is deliberately best-effort telemetry. Its contents are never read + by publication, model selection, health admission, or request routing. + """ + + with self._lock, self._connect() as connection: + connection.execute( + """INSERT INTO model_runtime_observations + (config_revision, provider_id, model_profile_id, + provider_model_id, api_mode, purpose, outcome, error_code, + strategy_json, observed_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", + ( + int(config_revision), + str(provider_id)[:256], + str(model_profile_id)[:256], + str(provider_model_id)[:512], + str(api_mode)[:128], + str(purpose)[:128], + str(outcome)[:128], + str(error_code)[:128] if error_code else None, + canonical_json_v1(dict(strategy)).decode(), + _utc_now(), + ), + ) + + def runtime_observation_summary(self) -> list[dict[str, Any]]: + """Return the most recent outcome and aggregate counts per active model.""" + + revision = self.current_revision() + if revision < 1: + return [] + with self._connect() as connection: + rows = connection.execute( + """SELECT latest.model_profile_id, latest.provider_id, + latest.provider_model_id, latest.api_mode, + latest.purpose, latest.outcome AS latest_outcome, + latest.error_code AS latest_error_code, + latest.strategy_json, latest.observed_at AS latest_at, + stats.total_calls, stats.successful_calls + FROM model_runtime_observations AS latest + JOIN ( + SELECT model_profile_id, MAX(id) AS latest_id, + COUNT(*) AS total_calls, + SUM(CASE WHEN outcome = 'succeeded' THEN 1 ELSE 0 END) + AS successful_calls + FROM model_runtime_observations + WHERE config_revision=? + GROUP BY model_profile_id + ) AS stats ON stats.latest_id = latest.id + ORDER BY latest.id DESC""", + (revision,), + ).fetchall() + result: list[dict[str, Any]] = [] + for row in rows: + item = dict(row) + try: + item["strategy"] = json.loads(str(item.pop("strategy_json"))) + except (TypeError, json.JSONDecodeError): + item["strategy"] = {} + result.append(item) + return result + + def get_admin_config(self) -> tuple[int, dict[str, Any], dict[str, ProviderCredentialMetadata]]: + generation = self.current_generation() + if generation.revision < 1: + return 0, {}, {} + row, _ = self._load_revision_row(generation.revision) + payload = json.loads(str(row["payload_json"])) + metadata: dict[str, ProviderCredentialMetadata] = {} + with self._connect() as connection: + for provider in payload["providers"]: + provider_id = str(provider["provider_id"]) + secret = connection.execute( + """SELECT * FROM provider_secret_versions + WHERE provider_id=? AND status='active'""", + (provider_id,), + ).fetchone() + if secret is None: + metadata[provider_id] = ProviderCredentialMetadata( + provider_id, False, "", "", "" + ) + else: + metadata[provider_id] = ProviderCredentialMetadata( + provider_id, + True, + str(secret["masked_value"]), + str(secret["created_at"]), + self.credential_etag( + provider_id, + int(secret["version"]), + str(secret["credential_fingerprint"]), + generation.security_epoch, + ), + ) + return generation.revision, payload, metadata + + def credential_etag( + self, + provider_id: str, + version: int, + fingerprint: str, + security_epoch: int, + ) -> str: + _, key = self.identity_key_ring.derive_current(_CREDENTIAL_ETAG_INFO) + return hmac_id( + key, + { + "provider_id": provider_id, + "secret_version": version, + "credential_fingerprint": fingerprint, + "security_epoch": security_epoch, + }, + ) + + def resolve_legacy_alias(self, legacy_alias: str) -> str | None: + with self._connect() as connection: + row = connection.execute( + """SELECT model_profile_id, migration_status + FROM legacy_model_alias_mappings WHERE legacy_alias=?""", + (str(legacy_alias),), + ).fetchone() + if row is None or str(row["migration_status"]) != "mapped": + return None + return str(row["model_profile_id"] or "") or None + + def list_revisions(self, *, limit: int = 100) -> list[dict[str, Any]]: + bounded = max(1, min(int(limit), 500)) + active = self.current_revision() + with self._connect() as connection: + rows = connection.execute( + """SELECT revision, rollback_of_revision, payload_hash, + runtime_projection_hash, operational_policy_revision, + adapter_registry_revision, created_by, created_at, + operation_id + FROM model_config_revisions + ORDER BY revision DESC LIMIT ?""", + (bounded,), + ).fetchall() + return [ + { + **dict(row), + "active": int(row["revision"]) == active, + } + for row in rows + ] + + def rollback( + self, + *, + target_revision: int, + expected_revision: int, + operation_id: str, + actor: str, + request_hash: str, + ) -> UnifiedSaveResult: + """Publish a historical snapshot as a new monotonic revision.""" + + target_row, _ = self._load_revision_row(target_revision) + target_payload = json.loads(str(target_row["payload_json"])) + with self._connect() as connection: + binding_rows = connection.execute( + """SELECT binding.provider_id, binding.secret_version, + secret.status, secret.credential_fingerprint + FROM config_provider_secret_bindings AS binding + JOIN provider_secret_versions AS secret + ON secret.provider_id=binding.provider_id + AND secret.version=binding.secret_version + WHERE binding.config_revision=?""", + (target_revision,), + ).fetchall() + if any(str(item["status"]) == "revoked" for item in binding_rows): + raise EvoRuntimeError("MODEL_CREDENTIAL_REVOKED") + bindings = { + str(item["provider_id"]): int(item["secret_version"]) + for item in binding_rows + } + identity_key_id, revision_key = self.identity_key_ring.derive_current( + _REVISION_HMAC_INFO + ) + policy_json = canonical_json_v1(OPERATIONAL_POLICY_V1).decode() + policy_hash = sha256_id(OPERATIONAL_POLICY_V1) + adapter_revision = get_adapter_registry().registry_revision + with self._lock, self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + operation = connection.execute( + "SELECT * FROM model_config_operations WHERE operation_id=?", + (operation_id,), + ).fetchone() + if operation is None or str(operation["request_hash"]) != request_hash: + raise EvoRuntimeError("IDEMPOTENCY_CONFLICT") + if str(operation["status"]) == "SUCCEEDED": + return UnifiedSaveResult( + **json.loads(str(operation["result_json"] or "{}")) + ) + active = connection.execute( + """SELECT revision, security_epoch, evidence_epoch + FROM active_model_config WHERE singleton=1""" + ).fetchone() + current = int(active["revision"]) if active else 0 + security_epoch = int(active["security_epoch"]) if active else 0 + evidence_epoch = int(active["evidence_epoch"]) if active else 0 + if current != expected_revision: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + new_revision = current + 1 + connection.execute( + "UPDATE provider_secret_versions SET status='retired' WHERE status='active'" + ) + for provider_id, version in bindings.items(): + cursor = connection.execute( + """UPDATE provider_secret_versions SET status='active' + WHERE provider_id=? AND version=? AND status='retired'""", + (provider_id, version), + ) + if cursor.rowcount != 1: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + projection = project_v4_to_v3( + target_payload, + revision=new_revision, + identity_key_id=identity_key_id, + bindings=bindings, + ) + payload_hash = sha256_id(target_payload) + projection_hash = sha256_id(projection) + signature = self._revision_signature_payload( + revision=new_revision, + payload_hash=payload_hash, + projection_hash=projection_hash, + policy_hash=policy_hash, + adapter_registry_revision=adapter_revision, + bindings=bindings, + ) + created_at = _utc_now() + connection.execute( + """INSERT INTO model_config_revisions + (revision, rollback_of_revision, schema_version, payload_json, + payload_hash, runtime_projection_json, runtime_projection_hash, + operational_policy_revision, operational_policy_json, + operational_policy_hash, identity_key_id, revision_hmac, + adapter_registry_revision, created_by, created_at, operation_id) + VALUES (?, ?, 4, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", + ( + new_revision, + target_revision, + canonical_json_v1(target_payload).decode(), + payload_hash, + canonical_json_v1(projection).decode(), + projection_hash, + OPERATIONAL_POLICY_V1["revision"], + policy_json, + policy_hash, + identity_key_id, + hmac_id(revision_key, signature), + adapter_revision, + actor[:256], + created_at, + operation_id, + ), + ) + connection.executemany( + """INSERT INTO config_provider_secret_bindings + (config_revision, provider_id, secret_version) VALUES (?, ?, ?)""", + [ + (new_revision, provider_id, version) + for provider_id, version in bindings.items() + ], + ) + connection.execute( + """UPDATE active_model_config SET revision=?, security_epoch=?, + evidence_epoch=?, updated_at=? WHERE singleton=1""", + (new_revision, security_epoch, evidence_epoch, created_at), + ) + result = UnifiedSaveResult(operation_id, new_revision, True, False) + result_json = canonical_json_v1( + { + "operation_id": operation_id, + "config_revision": new_revision, + "changed": True, + "revalidated": False, + "status": "SUCCEEDED", + } + ).decode() + connection.execute( + """UPDATE model_config_operations SET status='SUCCEEDED', + stage='SUCCEEDED', error_code=NULL, error_details_json=NULL, + result_revision=?, result_json=?, lease_owner=NULL, + lease_expires_at=NULL, updated_at=? WHERE operation_id=?""", + (new_revision, result_json, created_at, operation_id), + ) + connection.execute( + """INSERT INTO model_config_audit + (operation_id, actor, action, old_revision, new_revision, + summary_json, created_at) + VALUES (?, ?, 'CONFIG_ROLLBACK', ?, ?, ?, ?)""", + ( + operation_id, + actor[:256], + current, + new_revision, + canonical_json_v1({"target_revision": target_revision}).decode(), + created_at, + ), + ) + return result + + def reencrypt_secrets(self, *, actor: str, batch_size: int = 100) -> int: + """Re-encrypt retained ciphertext with the current encryption key.""" + + bounded = max(1, min(int(batch_size), 1000)) + audit_id = str(uuid.uuid4()) + with self._lock, self._connect() as connection: + rows = connection.execute( + """SELECT * FROM provider_secret_versions + WHERE encryption_key_id<>? ORDER BY provider_id, version LIMIT ?""", + (self.current_encryption_key_id, bounded), + ).fetchall() + connection.execute("BEGIN IMMEDIATE") + for row in rows: + plaintext = self._decrypt_row(row) + connection.execute( + """UPDATE provider_secret_versions + SET ciphertext=?, encryption_key_id=?, cipher_version='fernet-v1' + WHERE provider_id=? AND version=?""", + ( + self._fernets[self.current_encryption_key_id].encrypt( + plaintext.encode("utf-8") + ), + self.current_encryption_key_id, + str(row["provider_id"]), + int(row["version"]), + ), + ) + if rows: + connection.execute( + """INSERT INTO model_config_audit + (operation_id, actor, action, old_revision, new_revision, + summary_json, created_at) + VALUES (?, ?, 'SECRET_REENCRYPT', NULL, NULL, ?, ?)""", + ( + audit_id, + actor[:256], + canonical_json_v1( + { + "count": len(rows), + "encryption_key_id": self.current_encryption_key_id, + } + ).decode(), + _utc_now(), + ), + ) + return len(rows) + + def revoke_credential( + self, + provider_id: str, + *, + expected_revision: int, + credential_etag: str, + reason: str, + actor: str, + operation_id: str, + request_hash: str, + ) -> Mapping[str, Any]: + clean_reason = str(reason or "").strip() + if not clean_reason: + raise EvoRuntimeError("LLM_SECRET_INVALID") + with self._lock, self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + operation = connection.execute( + "SELECT * FROM model_config_operations WHERE operation_id=?", + (operation_id,), + ).fetchone() + if operation is None or str(operation["request_hash"]) != request_hash: + raise EvoRuntimeError("IDEMPOTENCY_CONFLICT") + if str(operation["status"]) == "SUCCEEDED": + return dict(json.loads(str(operation["result_json"] or "{}"))) + active = connection.execute( + "SELECT revision, security_epoch FROM active_model_config WHERE singleton=1" + ).fetchone() + if active is None or int(active["revision"]) != expected_revision: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + secret = connection.execute( + """SELECT * FROM provider_secret_versions + WHERE provider_id=? AND status='active'""", + (provider_id,), + ).fetchone() + if secret is None: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + expected_etag = self.credential_etag( + provider_id, + int(secret["version"]), + str(secret["credential_fingerprint"]), + int(active["security_epoch"]), + ) + if expected_etag != credential_etag: + raise EvoRuntimeError("CONFIG_REVISION_CONFLICT") + next_epoch = int(active["security_epoch"]) + 1 + now = _utc_now() + connection.execute( + """UPDATE provider_secret_versions SET status='revoked', + revoked_at=?, revoke_reason=? + WHERE provider_id=? AND version=? AND status='active'""", + ( + now, + clean_reason[:1024], + provider_id, + int(secret["version"]), + ), + ) + connection.execute( + """UPDATE active_model_config SET security_epoch=?, updated_at=? + WHERE singleton=1""", + (next_epoch, now), + ) + result = { + "operation_id": operation_id, + "provider_id": provider_id, + "status": "revoked", + "security_epoch": next_epoch, + } + result_json = canonical_json_v1(result).decode() + connection.execute( + """UPDATE model_config_operations SET status='SUCCEEDED', + stage='SUCCEEDED', error_code=NULL, error_details_json=NULL, + result_revision=?, result_json=?, lease_owner=NULL, + lease_expires_at=NULL, updated_at=? WHERE operation_id=?""", + (expected_revision, result_json, now, operation_id), + ) + connection.execute( + """INSERT INTO model_config_audit + (operation_id, actor, action, old_revision, new_revision, + summary_json, created_at) VALUES (?, ?, 'SECRET_REVOKE', ?, ?, ?, ?)""", + ( + operation_id, + actor[:256], + expected_revision, + expected_revision, + canonical_json_v1( + {"provider_id": provider_id, "security_epoch": next_epoch} + ).decode(), + now, + ), + ) + return result + + def resolve(self, reference: SecretReference) -> ResolvedSecret: + ref = str(reference.ref) + if ref.startswith("provider://"): + provider_id = ref[len("provider://") :] + with self._connect() as connection: + row = connection.execute( + """SELECT * FROM provider_secret_versions + WHERE provider_id=? AND status='active'""", + (provider_id,), + ).fetchone() + elif ref.startswith("secret://model-providers/") and "#" in ref: + provider_and_version = ref[len("secret://model-providers/") :] + try: + provider_id, raw_version = provider_and_version.rsplit("#", 1) + version = int(raw_version) + except ValueError as exc: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") from exc + if reference.revision != version: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + with self._connect() as connection: + row = connection.execute( + """SELECT * FROM provider_secret_versions + WHERE provider_id=? AND version=?""", + (provider_id, version), + ).fetchone() + else: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + if row is None: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + if str(row["status"]) == "revoked": + raise EvoRuntimeError("MODEL_CREDENTIAL_REVOKED") + value = self._decrypt_row(row) + return ResolvedSecret( + value, + int(row["version"]), + str(row["version"]), + str(row["credential_fingerprint"]), + ) diff --git a/EvoScientist/llm/models.py b/EvoScientist/llm/models.py index f9d6f27..84a33b6 100644 --- a/EvoScientist/llm/models.py +++ b/EvoScientist/llm/models.py @@ -84,11 +84,6 @@ def _resolve_codex_client_version() -> str: return _CODEX_CLIENT_VERSION_FALLBACK -def _resolve_reasoning_effort(default: str) -> str: - """Return the configured reasoning effort or a provider-specific default.""" - return os.environ.get("EVOSCIENTIST_REASONING_EFFORT", "").strip() or default - - # Providers routed through the OpenAI provider with a custom base_url. # Maps provider name → (base_url or None, env var for API key). _OPENAI_ROUTED_PROVIDERS: dict[str, tuple[str | None, str]] = { @@ -439,8 +434,10 @@ def _apply_auto_config( ) else "high" ) - _eff = _resolve_reasoning_effort(_default_effort) - kwargs["reasoning"] = {"effort": _eff, "summary": "auto"} + # An explicit API envelope belongs to the compiled invocation plan. + # Do not add a legacy Responses-style reasoning object to a Chat plan. + if "use_responses_api" not in kwargs: + kwargs["reasoning"] = {"effort": _default_effort, "summary": "auto"} # Google GenAI: surface thinking traces if provider == "google-genai" and not disable_reasoning: @@ -474,46 +471,7 @@ def get_chat_model( >>> model = get_chat_model("gpt-4o") # OpenAI model >>> model = get_chat_model("claude-3-opus-20240229", provider="anthropic") # Full ID """ - skip_runtime_resolver = bool(kwargs.pop("_skip_runtime_model_resolver", False)) - runtime_provider_name: str | None = None - runtime_supports_reasoning: bool | None = None - runtime_resolved = None - if not skip_runtime_resolver: - from EvoScientist.runtime_integrations import resolve_runtime_model - - runtime_resolved = resolve_runtime_model(model, provider) - - if runtime_resolved is not None: - resolved_params = dict(getattr(runtime_resolved, "params", {}) or {}) - extra_body = resolved_params.pop("_extra_body", None) - default_headers = resolved_params.pop("_default_headers", None) - if extra_body: - resolved_params["extra_body"] = extra_body - if default_headers: - resolved_params["default_headers"] = default_headers - resolved_params.update(kwargs) - kwargs = resolved_params - - resolved_api_key = str(getattr(runtime_resolved, "api_key", "") or "") - resolved_base_url = str(getattr(runtime_resolved, "base_url", "") or "") - if resolved_api_key: - kwargs.setdefault("api_key", resolved_api_key) - if resolved_base_url: - kwargs.setdefault("base_url", resolved_base_url.rstrip("/")) - - runtime_provider_name = str( - getattr(runtime_resolved, "provider_name", "") or "" - ) - runtime_supports_reasoning = bool( - getattr(runtime_resolved, "supports_reasoning", False) - ) - if not runtime_supports_reasoning: - kwargs.setdefault("_disable_reasoning", True) - kwargs.setdefault("_disable_thinking", True) - model = str(runtime_resolved.model_id) - provider = str(runtime_resolved.protocol) - else: - model = model or DEFAULT_MODEL + model = model or DEFAULT_MODEL # Look up short name in registry (provider-aware) model_id = None @@ -548,19 +506,15 @@ def get_chat_model( _is_third_party = ( provider in _OPENAI_ROUTED_PROVIDERS or provider in _ANTHROPIC_ROUTED_PROVIDERS ) - if runtime_provider_name and runtime_provider_name != provider: - _is_third_party = True + explicit_base_url = str(kwargs.get("base_url") or "") if ( - runtime_resolved is not None - and provider == "openai" - and resolved_base_url - and "api.openai.com" not in resolved_base_url.lower() + provider == "openai" + and explicit_base_url + and "api.openai.com" not in explicit_base_url.lower() ): _is_third_party = True _is_openai_proxy = False - _original_provider: str | None = ( - runtime_provider_name if runtime_provider_name != provider else None - ) + _original_provider: str | None = None if provider == "anthropic": base_url = os.environ.get("ANTHROPIC_BASE_URL", "") if base_url: @@ -578,17 +532,6 @@ def get_chat_model( kwargs.get("base_url"), kwargs.get("api_key") ) if _is_openai_proxy: - # Use Responses API for ccproxy: bypasses the format chain - # converter (Chat→Responses→Chat) which returns 502 on - # complex responses. System messages are converted to - # developer role by _patch_ccproxy_system_to_developer(). - kwargs.setdefault("use_responses_api", True) - # Streaming must stay ON for Responses API: ccproxy's - # StreamingBufferService loses output when assembling - # non-streaming responses. (The old streaming=False was - # for Chat Completions tool_call duplication — not an issue - # with the Responses API SSE format.) - kwargs.pop("streaming", None) # remove if set elsewhere # ccproxy forwards client headers upstream and only # gap-fills its own, so the Codex backend sees this # client's identity. Without Codex-CLI-shaped headers it @@ -650,8 +593,7 @@ def get_chat_model( # passback (OpenRouter's `/responses` beta is stateless, store=false — # "Item with id 'rs_...' not found"); the patch strips them on passback, # so enabling `summary` is safe. See langchain-ai/langchain#37777. - effort = _resolve_reasoning_effort("high") - kwargs.setdefault("reasoning", {"effort": effort, "summary": "auto"}) + kwargs.setdefault("reasoning", {"effort": "high", "summary": "auto"}) # App attribution (issue #339): identify EvoScientist to OpenRouter so # usage is credited to the project (app rankings, model app tabs, # analytics) rather than langchain-openrouter's LangChain-branded @@ -736,19 +678,6 @@ def get_chat_model( _apply_auto_config(provider, model_id, _is_third_party, kwargs, _original_provider) _apply_openrouter_anthropic_prompt_cache(provider, model_id, kwargs) - # User-level override for the OpenAI Responses API vs Chat Completions. - # When "false", force Chat Completions and drop reasoning (which triggers - # the Responses API path in langchain-openai). Only applies to OpenAI. - if provider == "openai": - _responses_api_setting = ( - os.environ.get("EVOSCIENTIST_USE_RESPONSES_API", "").strip().lower() - ) - if _responses_api_setting == "false": - kwargs["use_responses_api"] = False - kwargs.pop("reasoning", None) - elif _responses_api_setting == "true": - kwargs["use_responses_api"] = True - anthropic_auth_token = None if provider == "anthropic" and kwargs.get("api_key"): anthropic_auth_token = os.environ.pop("ANTHROPIC_AUTH_TOKEN", None) @@ -769,7 +698,19 @@ def get_chat_model( # Anthropic-routed providers accept media in tool results natively; # only OpenAI-compatible providers need tool-media hoisting. _hoist = _original_provider not in _ANTHROPIC_ROUTED_PROVIDERS - _patch_openai_compat_content(chat_model, hoist_tool_media=_hoist) + _patch_openai_compat_content( + chat_model, + hoist_tool_media=_hoist, + # Generic OpenAI-compatible proxies must not receive hidden + # reasoning traces emitted by a different provider. DeepSeek has + # its own explicit passback patch below, so preserve that path. + drop_reasoning_metadata=( + _is_third_party + and provider == "openai" + and _original_provider is None + and not _is_openai_proxy + ), + ) # DeepSeek thinking mode requires reasoning_content passback in multi-turn # + tool_use scenarios. diff --git a/EvoScientist/llm/patches.py b/EvoScientist/llm/patches.py index ca3b46e..2b1db19 100644 --- a/EvoScientist/llm/patches.py +++ b/EvoScientist/llm/patches.py @@ -11,6 +11,8 @@ Patches: - _patch_openai_capture_reasoning_content: capture provider reasoning_content into AIMessage.additional_kwargs (module-level, applied at import) + - _patch_openai_empty_sse_keepalive: ignore blank SSE keepalive events + emitted by some OpenAI-compatible Responses endpoints - _patch_deepseek_reasoning_passback: re-inject reasoning_content into outgoing DeepSeek assistant messages for thinking-mode multi-turn / tool_use scenarios @@ -70,6 +72,65 @@ def _patch_anthropic_proxy_compat() -> None: _patch_anthropic_proxy_compat() +# --------------------------------------------------------------------------- +# Patch (module-level): some OpenAI-compatible endpoints emit SSE keepalive +# frames with an empty ``data`` field. OpenAI SDK 2.x unconditionally passes +# every event to ``json.loads``, so an otherwise harmless keepalive becomes a +# JSONDecodeError and aborts the stream. Filter only blank data frames before +# the SDK's parser; JSON events, error events, and [DONE] are unchanged. +# --------------------------------------------------------------------------- +_openai_empty_sse_keepalive_patched = False + + +def _is_blank_sse_keepalive(event: Any) -> bool: + """Return whether an SSE event has no JSON payload to parse.""" + + data = getattr(event, "data", None) + return data is None or (isinstance(data, str) and not data.strip()) + + +def _patch_openai_empty_sse_keepalive() -> None: + global _openai_empty_sse_keepalive_patched + if _openai_empty_sse_keepalive_patched: + return + try: + import functools + + from openai._streaming import AsyncStream as _AsyncStream + from openai._streaming import Stream as _Stream + + original_async = _AsyncStream._iter_events + if not getattr(original_async, "_evoscientist_skips_blank_sse", False): + + @functools.wraps(original_async) + async def _filtered_async_events(self: Any) -> Any: + async for event in original_async(self): + if not _is_blank_sse_keepalive(event): + yield event + + _filtered_async_events._evoscientist_skips_blank_sse = True # type: ignore[attr-defined] + _AsyncStream._iter_events = _filtered_async_events + + original_sync = _Stream._iter_events + if not getattr(original_sync, "_evoscientist_skips_blank_sse", False): + + @functools.wraps(original_sync) + def _filtered_sync_events(self: Any) -> Any: + for event in original_sync(self): + if not _is_blank_sse_keepalive(event): + yield event + + _filtered_sync_events._evoscientist_skips_blank_sse = True # type: ignore[attr-defined] + _Stream._iter_events = _filtered_sync_events + _openai_empty_sse_keepalive_patched = True + except Exception: + # The patch is only needed when the optional OpenAI SDK is available. + pass + + +_patch_openai_empty_sse_keepalive() + + # --------------------------------------------------------------------------- # Patch: ccproxy-api 0.2.7 Codex compatibility. # @@ -207,6 +268,9 @@ def _is_ccproxy_codex( # preserved, not flattened away. # --------------------------------------------------------------------------- _SKIP_CONTENT_TYPES = frozenset({"thinking", "reasoning", "reasoning_content"}) +_NONPORTABLE_REASONING_METADATA = frozenset( + {"reasoning_content", "reasoning_details"} +) # Media block types preserved when flattening (positive allowlist; # thinking/reasoning still dropped). Images + files (PDF/documents): both @@ -407,6 +471,8 @@ def _copy_ai_message_with_tool_pairs( # LangChain content blocks use id; the Responses converter later # maps it to call_id. block["id"] = call_id + if "call_id" in block: + block["call_id"] = call_id block["name"] = call_name if isinstance(block.get("function"), dict): block["function"] = {**block["function"], "name": call_name} @@ -540,7 +606,7 @@ def _validate_openai_tool_history(messages: list[Any]) -> None: }: continue block_id = str( - block.get("id") or block.get("call_id") or "" + block.get("call_id") or block.get("id") or "" ).strip() block_name = block.get("name") or block.get("tool_name") function = block.get("function") @@ -566,7 +632,12 @@ def _validate_openai_tool_history(messages: list[Any]) -> None: raise ValueError("assistant tool call is missing its tool result") -def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> list[Any]: +def _sanitize_messages( + messages: list[Any], + hoist_tool_media: bool = True, + *, + drop_reasoning_metadata: bool = False, +) -> list[Any]: """Flatten list content for OpenAI-compatible APIs, preserving media. Text/reasoning content is flattened to a string; image blocks are @@ -593,6 +664,15 @@ def _sanitize_messages(messages: list[Any], hoist_tool_media: bool = True) -> li pending_media.clear() for msg in messages: + if drop_reasoning_metadata: + additional_kwargs = getattr(msg, "additional_kwargs", None) or {} + if set(additional_kwargs) & _NONPORTABLE_REASONING_METADATA: + msg = copy.copy(msg) + msg.additional_kwargs = { + key: value + for key, value in additional_kwargs.items() + if key not in _NONPORTABLE_REASONING_METADATA + } is_tool = getattr(msg, "type", None) == "tool" if not is_tool: _flush() # emit hoisted media before any non-tool message @@ -714,7 +794,12 @@ def _strip_media_types(messages: list[Any], types: set[str]) -> list[Any]: return out -def _patch_openai_compat_content(model: Any, hoist_tool_media: bool = True) -> None: +def _patch_openai_compat_content( + model: Any, + hoist_tool_media: bool = True, + *, + drop_reasoning_metadata: bool = False, +) -> None: """Flatten list content to strings before OpenAI-compatible API calls. Wraps ``_generate`` / ``_agenerate`` / ``_stream`` / ``_astream`` to prevent @@ -749,11 +834,17 @@ def _patch_openai_compat_content(model: Any, hoist_tool_media: bool = True) -> N def _prepare(messages: list[BaseMessage]) -> list[BaseMessage]: msgs = _strip_media_types(messages, blocked) if blocked else messages - return _sanitize_messages(msgs, hoist_tool_media) + return _sanitize_messages( + msgs, + hoist_tool_media, + drop_reasoning_metadata=drop_reasoning_metadata, + ) def _stripped(messages: list[BaseMessage], suspects: set[str]) -> list[BaseMessage]: return _sanitize_messages( - _strip_media_types(messages, blocked | suspects), hoist_tool_media + _strip_media_types(messages, blocked | suspects), + hoist_tool_media, + drop_reasoning_metadata=drop_reasoning_metadata, ) orig_generate = getattr(model, "_generate", None) diff --git a/EvoScientist/llm/runtime.py b/EvoScientist/llm/runtime.py new file mode 100644 index 0000000..e24f432 --- /dev/null +++ b/EvoScientist/llm/runtime.py @@ -0,0 +1,3255 @@ +"""Transactional V3 model runtime owned by EvoScientist.""" + +from __future__ import annotations + +import asyncio +import base64 +import binascii +import hashlib +import json +import logging +import math +import re +import time +import traceback +import uuid +from collections import defaultdict +from collections.abc import AsyncIterator, Callable, Mapping, Sequence +from dataclasses import asdict, dataclass, field, replace +from typing import Any + +from langchain_core.callbacks import AsyncCallbackHandler +from langchain_core.messages import BaseMessage, message_to_dict + +from .adapter_registry import AdapterRegistration, NormalizedUsage, get_adapter_registry +from .configuration import SecretResolver +from .contracts import ( + AdmissionGrant, + AdmissionGrantVerifier, + AgentInputV3, + AgentModelSet, + EvoRuntimeError, + EvoRuntimeEvent, + EvoWebRun, + HmacGrantAuthority, + ModelAttemptEvent, + ModelCatalog, + ModelCatalogEntry, + PreparedRunQuote, + PricingQuote, + RouteCallBound, + RouteIdentity, + RoutePreparationGrant, + VerifiedModelSubject, + WebHostContext, + event_payload, + now_ms, +) +from .crypto import HmacKeyRing, canonical_json_v1 +from .invocation import InvocationPlan, compile_invocation_plan +from .model_config import ( + EvoModelConfig, + FileEvoModelConfigStore, + RouteRef, + endpoint_fingerprint, + invocation_fingerprint, + resolve_secret, + route_fingerprint, + route_semantics_hash, +) +from .user_options import ( + model_options_schema_hash, + project_user_options_for_purpose, + validate_user_model_options, +) + +_BIGINT_MAX = 2**63 - 1 +logger = logging.getLogger(__name__) +_ROUTE_SEMANTICS_INFO = "ai4sci/route-semantics-hash/v3" +_ROUTE_FINGERPRINT_INFO = "ai4sci/route-fingerprint/v3" +_ENDPOINT_FINGERPRINT_INFO = "ai4sci/endpoint-fingerprint/v3" +_PROTOCOL_MARGIN_TOKENS: Mapping[tuple[str, str], int] = { + ("custom-openai", "chat_completions"): 64, + ("custom-openai", "responses"): 96, + # Generic OpenAI-compatible providers use the Chat Completions message + # envelope and therefore have the same conservative framing margin. + ("generic-openai-compatible", "chat_completions"): 64, + ("openai", "chat_completions"): 64, + ("openai", "responses"): 96, + ("anthropic", "messages"): 64, + ("dashscope", "chat_completions"): 64, + ("dashscope", "responses"): 96, + ("google-gemini", "interactions"): 96, + ("google-gemini", "generate_content"): 64, + ("xai", "responses"): 96, + ("xai", "chat_completions"): 64, +} + +_DEBUG_INVOCATION_PARAMETER_KEYS = frozenset( + { + "disable_streaming", + "extra_body", + "max_completion_tokens", + "max_output_tokens", + "max_tokens", + "reasoning", + "reasoning_effort", + "response_format", + "store", + "streaming", + "temperature", + "thinking", + "tool_choice", + "top_p", + "use_responses_api", + } +) +_DEBUG_NESTED_PARAMETER_KEYS = frozenset( + { + "budget_tokens", + "effort", + "enable_thinking", + "thinking_budget", + "type", + } +) +_AUTH_VALUE_PATTERN = re.compile( + r"(?i)(bearer\s+|(?:api[_-]?key|authorization|token|secret)\s*[:=]\s*)" + r"([^\s,;}'\"]+)" +) +_OPENAI_RESPONSE_EVENT_TYPE_PATTERN = re.compile( + r'"type"\s*:\s*"(response\.[A-Za-z0-9_.:-]{1,120})"' +) + + +def _enabled_purposes(title_policy: str) -> tuple[str, ...]: + purposes = ["main_agent", "tool_selector", "deepagents_summarizer"] + if title_policy == "best_effort": + purposes.append("title") + return tuple(purposes) + + +def _merge_params( + base: Mapping[str, Any], overlay: Mapping[str, Any] +) -> dict[str, Any]: + """Merge provider adapter parameters without discarding sibling JSON fields.""" + + result = dict(base) + for key, value in overlay.items(): + existing = result.get(key) + if isinstance(existing, Mapping) and isinstance(value, Mapping): + result[key] = _merge_params(existing, value) + else: + result[key] = value + return result + + +@dataclass(frozen=True, slots=True) +class ResolvedRoute: + ref: RouteRef + identity: RouteIdentity + quote: PricingQuote + context_window: int + max_output_tokens: int + reasoning_mode: str + reasoning_enabled_params: Mapping[str, Any] + reasoning_disabled_params: Mapping[str, Any] + base_url: str + params: Mapping[str, Any] + api_key: str + default_headers: Mapping[str, str] + secret_fingerprints: Mapping[str, str] + adapter: AdapterRegistration | None = None + provider_max_inflight: int = 16 + model_max_inflight: int = 16 + queue_timeout_seconds: int = 5 + provider_capacity_key: str = "" + model_capacity_key: str = "" + attempt_timeout_seconds: int = 600 + runtime_provider: str = "" + invocation_plan: InvocationPlan | None = None + supports_tools: bool = False + + +@dataclass(frozen=True, slots=True) +class ModelRuntimeSnapshot: + config_revision: int + catalog_revision: int + preparation_id: str + prepared_snapshot_digest: str + prepared_input_digest: str + purpose_attempt_limits: Mapping[str, int] + purpose_routes: Mapping[str, tuple[ResolvedRoute, ...]] + purpose_route_call_bounds: Mapping[str, tuple[RouteCallBound, ...]] + title_start_timeout_seconds: int + active_run_timeout_seconds: int + max_run_journal_events: int + max_run_journal_bytes: int + + @property + def main_routes(self) -> tuple[ResolvedRoute, ...]: + return self.purpose_routes["main_agent"] + + @property + def title_route(self) -> ResolvedRoute | None: + routes = self.purpose_routes.get("title") + return routes[0] if routes else None + + +@dataclass(slots=True) +class _PreparedHandle: + grant: RoutePreparationGrant + quote: PreparedRunQuote + input: AgentInputV3 + host: WebHostContext + snapshot: ModelRuntimeSnapshot + tool_registry_payload: Mapping[str, Any] + secret_fingerprints: Mapping[str, str] + request_digest: str + state: str = "PREPARED" + run: _EvoWebRun | None = None + start_grant_digest: bytes | None = None + lock: asyncio.Lock = field(default_factory=asyncio.Lock) + + +@dataclass(frozen=True, slots=True) +class _PreparedTombstone: + request_digest: str + state: str + retained_until: int + + +class RouteHealthBook: + """Single-worker circuit breaker with a bounded half-open lease.""" + + def __init__(self) -> None: + self._revision = -1 + self._policy: Any = None + self._failures: dict[str, int] = {} + self._states: dict[str, str] = {} + self._opened_at: dict[str, float] = {} + self._half_open_inflight: dict[str, int] = defaultdict(int) + + def configure(self, config: EvoModelConfig, policy: Any | None = None) -> None: + if self._revision != config.config_revision: + self._revision = config.config_revision + self._policy = policy or config.route_health + self._failures.clear() + self._states.clear() + self._opened_at.clear() + self._half_open_inflight.clear() + + def status(self, route_key: str) -> str: + if self._states.get(route_key) == "open": + assert self._policy is not None + if ( + time.monotonic() - self._opened_at.get(route_key, 0) + >= self._policy.cooldown_seconds + ): + return "half_open" + return self._states.get(route_key, "closed") + + def is_open(self, route_key: str) -> bool: + return self.status(route_key) == "open" + + def acquire_half_open(self, route_key: str) -> bool: + if self.status(route_key) != "half_open": + return True + assert self._policy is not None + if self._half_open_inflight[route_key] >= self._policy.half_open_max_inflight: + return False + self._half_open_inflight[route_key] += 1 + return True + + def release_half_open(self, route_key: str) -> None: + if self._half_open_inflight.get(route_key, 0) <= 1: + self._half_open_inflight.pop(route_key, None) + else: + self._half_open_inflight[route_key] -= 1 + + def record_success(self, route_key: str) -> None: + self._failures.pop(route_key, None) + self._states[route_key] = "closed" + self._opened_at.pop(route_key, None) + self._half_open_inflight.pop(route_key, None) + + def record_failure(self, route_key: str, error_code: str) -> None: + assert self._policy is not None + self._half_open_inflight.pop(route_key, None) + if error_code in self._policy.open_immediately_error_codes: + self._open(route_key) + return + if error_code not in self._policy.counted_error_codes: + return + failures = self._failures.get(route_key, 0) + 1 + self._failures[route_key] = failures + if failures >= self._policy.failure_threshold: + self._open(route_key) + + def _open(self, route_key: str) -> None: + self._states[route_key] = "open" + self._opened_at[route_key] = time.monotonic() + + +class _SmoothWeightedRoundRobin: + def __init__(self) -> None: + self._revision = -1 + self._current: dict[tuple[str, str], int] = {} + self._locks: dict[str, asyncio.Lock] = {} + + def configure(self, revision: int) -> None: + if revision != self._revision: + self._revision = revision + self._current.clear() + self._locks.clear() + + async def choose( + self, pool_id: str, weighted_names: Sequence[tuple[str, int]] + ) -> str: + lock = self._locks.setdefault(pool_id, asyncio.Lock()) + async with lock: + total = sum(weight for _, weight in weighted_names) + best_name = "" + best_value: int | None = None + for name, weight in weighted_names: + key = (pool_id, name) + current = self._current.get(key, 0) + weight + self._current[key] = current + if best_value is None or current > best_value: + best_name, best_value = name, current + self._current[(pool_id, best_name)] -= total + return best_name + + +class InvocationCapacityBook: + """Process-local Provider and Model concurrency gates for Web runtime calls.""" + + def __init__(self) -> None: + self._semaphores: dict[tuple[str, int], asyncio.BoundedSemaphore] = {} + + def _semaphore(self, key: str, limit: int) -> asyncio.BoundedSemaphore: + return self._semaphores.setdefault( + (key, limit), asyncio.BoundedSemaphore(limit) + ) + + async def acquire(self, model: Any) -> tuple[asyncio.BoundedSemaphore, ...]: + metadata = getattr(model, "metadata", None) or {} + provider_key = str(metadata.get("capacity_provider_key") or "") + model_key = str(metadata.get("capacity_model_key") or "") + if not provider_key or not model_key: + return () + provider = self._semaphore( + provider_key, int(metadata.get("provider_max_inflight") or 16) + ) + model_gate = self._semaphore( + model_key, int(metadata.get("model_max_inflight") or 16) + ) + acquired: list[asyncio.BoundedSemaphore] = [] + try: + async with asyncio.timeout( + int(metadata.get("capacity_queue_timeout_seconds") or 5) + ): + await provider.acquire() + acquired.append(provider) + await model_gate.acquire() + acquired.append(model_gate) + except TimeoutError as exc: + for gate in reversed(acquired): + gate.release() + raise EvoRuntimeError("MODEL_CAPACITY_EXHAUSTED") from exc + return tuple(acquired) + + @staticmethod + def release(lease: tuple[asyncio.BoundedSemaphore, ...]) -> None: + for gate in reversed(lease): + gate.release() + + +class EvoModelRuntime: + """Prepare, authorize and execute one frozen Evo Agent transaction.""" + + def __init__( + self, + store: FileEvoModelConfigStore, + *, + admission_verifier: AdmissionGrantVerifier, + quote_authority: HmacGrantAuthority | None = None, + identity_key_ring: HmacKeyRing | None = None, + secret_resolver: SecretResolver | None = None, + model_factory: Callable[..., Any] | None = None, + agent_factory: Callable[ + [ModelRuntimeSnapshot, WebHostContext, AgentModelSet], Any + ] + | None = None, + runtime_instance_id: str | None = None, + ) -> None: + self.store = store + self.admission_verifier = admission_verifier + if quote_authority is None and isinstance( + admission_verifier, HmacGrantAuthority + ): + quote_authority = admission_verifier + if quote_authority is None: + raise ValueError("quote authority is required") + if identity_key_ring is None: + raise ValueError("config identity key ring is required") + self.quote_authority = quote_authority + self.identity_key_ring = identity_key_ring + self.secret_resolver = secret_resolver + self.model_factory = model_factory or self._default_model_factory + self.agent_factory = agent_factory or self._default_agent_factory + self.runtime_instance_id = runtime_instance_id or str(uuid.uuid4()) + self._route_health = RouteHealthBook() + self._provider_health = RouteHealthBook() + self._pool = _SmoothWeightedRoundRobin() + self._capacity = InvocationCapacityBook() + self._prepared: dict[str, _PreparedHandle] = {} + self._prepared_requests: dict[tuple[str, str, str], str] = {} + self._prepared_tombstones: dict[tuple[str, str, str], _PreparedTombstone] = {} + self._started_grants: dict[str, tuple[bytes, _EvoWebRun]] = {} + self._registry_lock = asyncio.Lock() + + def _record_runtime_observation( + self, + route: ResolvedRoute, + *, + purpose: str, + outcome: str, + error_code: str | None, + ) -> None: + """Record call metadata when the configured store supports telemetry. + + Observability must never change the outcome of a request, including + when the telemetry database is unavailable. + """ + + recorder = getattr(self.store, "record_runtime_observation", None) + if not callable(recorder): + return + params = ( + route.invocation_plan.sdk_params + if route.invocation_plan is not None + else {} + ) + extra_body = params.get("extra_body") + if not isinstance(extra_body, Mapping): + extra_body = {} + response_format = params.get("response_format") + strategy = { + "enable_thinking": extra_body.get("enable_thinking"), + "structured_output": bool(response_format), + "tool_choice": params.get("tool_choice"), + } + try: + recorder( + config_revision=route.identity.config_revision, + provider_id=route.identity.provider_id, + model_profile_id=route.ref.model, + provider_model_id=route.identity.model_id, + api_mode=route.identity.api_mode, + purpose=purpose, + outcome=outcome, + error_code=error_code, + strategy=strategy, + ) + except Exception: + # Telemetry is intentionally non-blocking and non-authoritative. + return + + async def prepare_model_run( + self, + grant: RoutePreparationGrant, + agent_input: AgentInputV3, + host: WebHostContext, + ) -> PreparedRunQuote: + self.admission_verifier.require_preparation(grant) + self._validate_preparation_echo(grant, agent_input) + tool_payload = self._tool_registry_payload(host) + request_key = (grant.subject_id, grant.request_id, grant.turn_id) + request_digest = self._preparation_request_digest( + grant, agent_input, tool_payload + ) + async with self._registry_lock: + self._expire_prepared_locked() + replay = self._prepared_replay_locked(request_key, request_digest) + if replay is not None: + return replay + config = self.store.load() + self._route_health.configure(config) + self._provider_health.configure( + config, config.provider_health or config.route_health + ) + self._pool.configure(config.config_revision) + if not self.identity_key_ring.contains(config.config_identity_key_id): + raise EvoRuntimeError("CONFIG_IDENTITY_KEY_UNKNOWN") + main_selector = config.resolve_main_selector(grant.requested_model_ref) + main_model = config.providers[main_selector.provider].models[ + main_selector.model + ] + main_provider = config.providers[main_selector.provider] + current_options_schema_hash = model_options_schema_hash( + model_profile_id=grant.requested_model_ref, + user_options=main_model.user_options, + supports_reasoning=main_model.supports_reasoning, + reasoning_mode=main_model.reasoning_mode, + allowed_reasoning_efforts=main_model.allowed_reasoning_efforts, + parameter_constraints=main_model.parameter_constraints, + adapter_id=main_provider.adapter_id, + adapter_revision=main_provider.adapter_revision, + ) + supplied_options_schema_hash = str( + agent_input.metadata.get("model_options_schema_hash") or "" + ) + if ( + supplied_options_schema_hash + and supplied_options_schema_hash != current_options_schema_hash + ): + raise EvoRuntimeError("MODEL_OPTIONS_STALE") + supplied_model_options = dict(agent_input.metadata.get("model_options") or {}) + combined_user_options = dict(supplied_model_options) + if main_model.supports_reasoning: + combined_user_options["reasoning"] = ( + "off" + if grant.reasoning_effort == "disabled" + else grant.reasoning_effort + ) + validated_user_options = validate_user_model_options( + supplied=combined_user_options, + user_options=main_model.user_options, + supports_reasoning=main_model.supports_reasoning, + reasoning_mode=main_model.reasoning_mode, + allowed_reasoning_efforts=main_model.allowed_reasoning_efforts, + default_reasoning_effort=str( + main_model.reasoning_enabled_params.get("reasoning") or "" + ), + parameter_constraints=main_model.parameter_constraints, + ) + validated_user_options.pop("reasoning", None) + async with self._registry_lock: + per_subject = sum( + handle.state == "PREPARED" + and handle.grant.subject_id == grant.subject_id + for handle in self._prepared.values() + ) + active_total = sum( + handle.state == "PREPARED" for handle in self._prepared.values() + ) + if ( + per_subject >= config.web_runtime.max_prepared_runs_per_subject + or active_total >= config.web_runtime.max_prepared_runs_total + ): + raise EvoRuntimeError("PREPARATION_CAPACITY_EXCEEDED") + media_token_bound = 0 + for item in agent_input.media: + if hasattr(item, "token_bound"): + bound = int(item.token_bound) + elif isinstance(item, Mapping) and "token_bound" in item: + bound = int(item["token_bound"]) + else: + raise EvoRuntimeError("TOKEN_BOUND_UNAVAILABLE") + if bound < 0: + raise EvoRuntimeError("TOKEN_BOUND_UNAVAILABLE") + media_token_bound += bound + tool_snapshot_id = self.quote_authority.tool_registry_snapshot_id(tool_payload) + prepared_input_payload = { + "agent_input": agent_input.projection(), + "checkpoint_snapshot_id": grant.checkpoint_snapshot_id, + "tool_registry_snapshot_id": tool_snapshot_id, + } + prepared_input_digest = self.quote_authority.prepared_input_digest( + prepared_input_payload + ) + enabled_purposes = _enabled_purposes(grant.title_policy) + purpose_routes = await self._freeze_purpose_routes( + config, + grant, + validated_user_options, + ) + purpose_attempt_limits = { + purpose: config.purpose_call_limits[purpose].max_attempts_per_run + for purpose in enabled_purposes + } + bounds = self._build_call_bounds(purpose_routes) + purpose_routes = self._compile_purpose_routes( + purpose_routes, + bounds, + grant.reasoning_effort, + ) + initial_input_bound = ( + _payload_token_bound( + {"input": agent_input.projection(), "tool_registry": tool_payload} + ) + + media_token_bound + ) + if any( + initial_input_bound > bound.payload_input_hard_cap + for purpose_bounds in bounds.values() + for bound in purpose_bounds + ): + raise EvoRuntimeError("MODEL_CONTEXT_WINDOW_EXCEEDED") + reserve = self._run_reserve(purpose_attempt_limits, bounds) + total_attempts = sum(purpose_attempt_limits.values()) + public_routes = { + purpose: { + "primary": routes[0].identity, + "fallbacks": tuple(route.identity for route in routes[1:]), + } + for purpose, routes in purpose_routes.items() + } + public_invocation_plans = { + purpose: { + "primary": routes[0].invocation_plan.projection() + if routes[0].invocation_plan + else {}, + "fallbacks": tuple( + route.invocation_plan.projection() if route.invocation_plan else {} + for route in routes[1:] + ), + } + for purpose, routes in purpose_routes.items() + } + public_bounds = { + purpose: tuple(bounds[purpose]) for purpose in purpose_attempt_limits + } + quotes = { + route.quote.quote_id: route.quote + for routes in purpose_routes.values() + for route in routes + } + snapshot_payload = { + "request_id": grant.request_id, + "turn_id": grant.turn_id, + "thread_id": grant.thread_id, + "subject_id": grant.subject_id, + "requested_model_ref": grant.requested_model_ref, + "config_revision": config.config_revision, + "gateway_input_digest": grant.gateway_input_digest, + "prepared_input_digest": prepared_input_digest, + "checkpoint_snapshot_id": grant.checkpoint_snapshot_id, + "tool_registry_snapshot_id": tool_snapshot_id, + "purpose_routes": public_routes, + "purpose_invocation_plans": public_invocation_plans, + "purpose_route_call_bounds": public_bounds, + "purpose_attempt_limits": purpose_attempt_limits, + "provider_run_reserve_microunits": reserve, + "route_health": asdict(config.route_health), + } + prepared_snapshot_digest = self.quote_authority.prepared_snapshot_digest( + snapshot_payload + ) + preparation_id = str(uuid.uuid4()) + issued_at = now_ms() + quote = self.quote_authority.sign_quote( + preparation_id=preparation_id, + request_id=grant.request_id, + turn_id=grant.turn_id, + thread_id=grant.thread_id, + subject_id=grant.subject_id, + requested_model_ref=grant.requested_model_ref, + plan=grant.plan, + roles=grant.roles, + requires_vision=grant.requires_vision, + reasoning_effort=grant.reasoning_effort, + title_policy=grant.title_policy, + gateway_input_digest=grant.gateway_input_digest, + prepared_snapshot_digest=prepared_snapshot_digest, + prepared_input_digest=prepared_input_digest, + config_revision=config.config_revision, + catalog_revision=config.catalog_revision, + enabled_purposes=enabled_purposes, + purpose_routes=public_routes, + purpose_route_call_bounds=bounds, + purpose_attempt_limits=purpose_attempt_limits, + total_max_attempts=total_attempts, + pricing_quotes=quotes, + quote_ids=tuple(sorted(quotes)), + provider_run_reserve_microunits=reserve, + checkpoint_snapshot_id=grant.checkpoint_snapshot_id, + tool_registry_snapshot_id=tool_snapshot_id, + turn_fencing_token=grant.turn_fencing_token, + route_semantics_hashes=tuple( + sorted( + { + route.identity.route_semantics_hash + for routes in purpose_routes.values() + for route in routes + } + ) + ), + issued_at=issued_at, + expires_at=min( + grant.expires_at, + issued_at + config.web_runtime.prepare_ttl_seconds * 1000, + ), + ) + snapshot = ModelRuntimeSnapshot( + config_revision=config.config_revision, + catalog_revision=config.catalog_revision, + preparation_id=preparation_id, + prepared_snapshot_digest=prepared_snapshot_digest, + prepared_input_digest=prepared_input_digest, + purpose_attempt_limits=purpose_attempt_limits, + purpose_routes=purpose_routes, + purpose_route_call_bounds=public_bounds, + title_start_timeout_seconds=config.web_runtime.title_start_timeout_seconds, + active_run_timeout_seconds=config.web_runtime.active_run_timeout_seconds, + max_run_journal_events=config.web_runtime.max_run_journal_events, + max_run_journal_bytes=config.web_runtime.max_run_journal_bytes, + ) + fingerprints = { + secret_ref: fingerprint + for routes in purpose_routes.values() + for route in routes + for secret_ref, fingerprint in route.secret_fingerprints.items() + } + handle = _PreparedHandle( + grant=grant, + quote=quote, + input=agent_input, + host=host, + snapshot=snapshot, + tool_registry_payload=tool_payload, + secret_fingerprints=fingerprints, + request_digest=request_digest, + ) + async with self._registry_lock: + self._expire_prepared_locked() + replay = self._prepared_replay_locked(request_key, request_digest) + if replay is not None: + return replay + if grant.grant_id in ( + item.grant.grant_id for item in self._prepared.values() + ): + raise EvoRuntimeError("PREPARATION_CONFLICT") + per_subject = sum( + item.state == "PREPARED" and item.grant.subject_id == grant.subject_id + for item in self._prepared.values() + ) + active_total = sum( + item.state == "PREPARED" for item in self._prepared.values() + ) + if ( + per_subject >= config.web_runtime.max_prepared_runs_per_subject + or active_total >= config.web_runtime.max_prepared_runs_total + ): + raise EvoRuntimeError("PREPARATION_CAPACITY_EXCEEDED") + self._prepared[preparation_id] = handle + self._prepared_requests[request_key] = preparation_id + return quote + + async def start_web_run(self, admission: AdmissionGrant) -> EvoWebRun: + self.admission_verifier.require_admission(admission) + digest = canonical_json_v1(admission.unsigned_payload()) + async with self._registry_lock: + replay = self._started_grants.get(admission.grant_id) + if replay is not None: + if replay[0] != digest: + raise EvoRuntimeError("CONTRACT_REPLAYED") + return replay[1] + handle = self._prepared.get(admission.preparation_id) + if handle is None: + raise EvoRuntimeError("RUN_LOST") + async with handle.lock: + if handle.state == "STARTED" and handle.run is not None: + return handle.run + if handle.state != "PREPARED": + raise EvoRuntimeError("CONTRACT_REPLAYED") + if handle.quote.expires_at < now_ms(): + handle.state = "EXPIRED" + raise EvoRuntimeError("PREPARATION_STALE") + self._validate_admission_echo(handle.quote, admission) + self._validate_still_fresh(handle) + if handle.host.runtime_event_sink is None: + raise EvoRuntimeError("EVENT_INGRESS_UNAVAILABLE") + model_set = self._build_model_set( + handle.snapshot, admission.reasoning_effort + ) + agent = self.agent_factory(handle.snapshot, handle.host, model_set) + run = _EvoWebRun( + runtime=self, + admission=admission, + prepared=handle, + agent=agent, + model_set=model_set, + ) + handle.state = "STARTED" + handle.run = run + handle.start_grant_digest = digest + async with self._registry_lock: + self._started_grants[admission.grant_id] = (digest, run) + return run + + async def cancel_prepared_run(self, preparation_id: str, *, reason: str) -> bool: + """Idempotently invalidate a prepared handle that never reached Start.""" + + _ = reason + async with self._registry_lock: + handle = self._prepared.get(str(preparation_id)) + if handle is None: + return False + async with handle.lock: + if handle.state in {"CANCELLED", "EXPIRED"}: + return True + if handle.state != "PREPARED": + return False + handle.state = "CANCELLED" + async with self._registry_lock: + self._retire_prepared_locked(handle) + return True + + async def get_catalog(self, subject: VerifiedModelSubject) -> ModelCatalog: + if not self.admission_verifier.verify_subject(subject): + raise EvoRuntimeError("CATALOG_SUBJECT_INVALID") + config = self.store.load() + entries: list[ModelCatalogEntry] = [] + for alias, selector_id in config.main_routes.selectable.items(): + selector = config.route_selectors[selector_id] + model = config.providers[selector.provider].models[selector.model] + alias_config = config.aliases.get(alias) + if self._access_allowed( + model, subject.plan, subject.roles, subject.subject_id + ) and ( + alias_config is None + or self._access_policy_allowed( + alias_config.access, subject.roles, subject.subject_id + ) + ): + health_states = [ + self._route_health.status(route.key()) + for route in config.concrete_routes(selector_id) + ] + aggregate = ( + "closed" + if "closed" in health_states + else ("half_open" if "half_open" in health_states else "open") + ) + entries.append( + # Alias defaults are display defaults only. The runtime + # already compiles them into the frozen route. + ModelCatalogEntry( + alias=alias, + provider=selector.provider, + supports_vision=model.supports_vision, + supports_reasoning=model.supports_reasoning, + allowed_reasoning_efforts=model.allowed_reasoning_efforts, + context_window=model.context_window, + max_output_tokens=model.max_output_tokens, + reasoning_mode=model.reasoning_mode, + default_reasoning_effort=str( + model.reasoning_enabled_params.get("reasoning") or "" + ), + billing_sku=model.billing_sku, + quote=model.quote, + health=aggregate, + display_name=( + alias_config.display_name if alias_config else alias + ), + provider_display_name=config.providers[ + selector.provider + ].display_name, + description=model.description, + version_policy=model.version_policy, + resolved_model_revision=model.resolved_model_revision, + reproducible=model.reproducible, + capabilities=tuple( + key + for key, enabled in model.capabilities.items() + if enabled + ), + user_options={ + name: { + **dict(option), + **( + {"default": alias_config.defaults[name]} + if alias_config is not None + and name in alias_config.defaults + else {"default": model.params[name]} + if name in model.params + else {} + ), + } + for name, option in model.user_options.items() + }, + parameter_constraints=model.parameter_constraints, + options_schema_hash=model_options_schema_hash( + model_profile_id=alias, + user_options=model.user_options, + supports_reasoning=model.supports_reasoning, + reasoning_mode=model.reasoning_mode, + allowed_reasoning_efforts=model.allowed_reasoning_efforts, + parameter_constraints=model.parameter_constraints, + adapter_id=config.providers[selector.provider].adapter_id, + adapter_revision=config.providers[ + selector.provider + ].adapter_revision, + ), + ) + ) + allowed_aliases = {entry.alias for entry in entries} + default_alias = config.main_routes.default_alias + if default_alias not in allowed_aliases: + default_alias = min(allowed_aliases) if allowed_aliases else "" + return ModelCatalog(config.catalog_revision, default_alias, tuple(entries)) + + def turn_lease_ttl_seconds(self) -> int: + config = self.store.load() + return ( + config.web_runtime.prepare_ttl_seconds + + config.web_runtime.turn_lease_grace_seconds + ) + + async def _freeze_purpose_routes( + self, + config: EvoModelConfig, + grant: RoutePreparationGrant, + model_options: Mapping[str, Any] | None = None, + ) -> Mapping[str, tuple[ResolvedRoute, ...]]: + main_selector = config.resolve_main_selector(grant.requested_model_ref) + main_refs = [await self._select_concrete(config, main_selector.selector_id)] + for fallback_selector in config.fallback_selectors(main_selector.selector_id): + main_refs.append(await self._select_concrete(config, fallback_selector)) + main_model = config.route_model(main_refs[0]) + alias_config = config.aliases.get(main_selector.alias) + if not self._access_allowed( + main_model, grant.plan, grant.roles, grant.subject_id + ) or ( + alias_config is not None + and not self._access_policy_allowed( + alias_config.access, grant.roles, grant.subject_id + ) + ): + raise EvoRuntimeError("MODEL_ACCESS_DENIED") + if grant.requires_vision and not main_model.supports_vision: + raise EvoRuntimeError("MODEL_CAPABILITY_UNAVAILABLE") + if grant.reasoning_effort != "disabled": + if ( + not main_model.supports_reasoning + or grant.reasoning_effort not in main_model.allowed_reasoning_efforts + ): + raise EvoRuntimeError("MODEL_CAPABILITY_UNAVAILABLE") + result: dict[str, tuple[ResolvedRoute, ...]] = {} + result["main_agent"] = tuple( + self._resolve_route( + config, route, "main_agent", model_options=model_options or {} + ) + for route in main_refs + ) + for purpose in ("tool_selector", "deepagents_summarizer"): + selector_id = config.purpose_selector_ids.get(purpose) + refs = ( + [await self._select_concrete(config, selector_id)] + if selector_id is not None + else main_refs + ) + for ref in refs: + self._require_route_access(config, ref, grant) + result[purpose] = tuple( + self._resolve_route( + config, + route, + purpose, + model_options=model_options or {} if selector_id is None else {}, + ) + for route in refs + ) + if grant.title_policy == "best_effort": + title_ref = await self._select_concrete(config, config.title_selector_id) + self._require_route_access(config, title_ref, grant) + result["title"] = (self._resolve_route(config, title_ref, "title"),) + return result + + @classmethod + def _require_route_access( + cls, + config: EvoModelConfig, + route: RouteRef, + grant: RoutePreparationGrant, + ) -> None: + selector = config.route_selectors[route.selector_id] + model = config.route_model(route) + alias_config = config.aliases.get(selector.alias) + if not cls._access_allowed( + model, grant.plan, grant.roles, grant.subject_id + ) or ( + alias_config is not None + and not cls._access_policy_allowed( + alias_config.access, grant.roles, grant.subject_id + ) + ): + raise EvoRuntimeError("MODEL_ACCESS_DENIED") + + async def _select_concrete( + self, config: EvoModelConfig, selector_id: str + ) -> RouteRef: + selector = config.route_selectors[selector_id] + routes = list(config.concrete_routes(selector_id)) + closed = [ + route + for route in routes + if self._route_health.status(route.key()) == "closed" + ] + candidates = closed or [ + route + for route in routes + if self._route_health.status(route.key()) == "half_open" + ] + if not candidates: + raise EvoRuntimeError("MODEL_ROUTE_UNAVAILABLE") + if selector.endpoint_pool is None or len(candidates) == 1: + return candidates[0] + pool = config.endpoint_pools[selector.endpoint_pool] + candidate_names = {route.endpoint for route in candidates} + weighted = [ + (member.name, member.weight) + for member in pool.endpoints + if member.name in candidate_names + ] + chosen = await self._pool.choose(pool.pool_id, weighted) + return next(route for route in candidates if route.endpoint == chosen) + + def _resolve_route( + self, + config: EvoModelConfig, + route: RouteRef, + purpose: str, + *, + model_options: Mapping[str, Any] | None = None, + ) -> ResolvedRoute: + provider = config.providers[route.provider] + endpoint = provider.endpoints[route.endpoint] + model = provider.models[route.model] + semantics_key = self.identity_key_ring.derive( + config.config_identity_key_id, _ROUTE_SEMANTICS_INFO + ) + fingerprint_key = self.identity_key_ring.derive( + config.config_identity_key_id, _ROUTE_FINGERPRINT_INFO + ) + endpoint_key = self.identity_key_ring.derive( + config.config_identity_key_id, _ENDPOINT_FINGERPRINT_INFO + ) + semantics_hash = route_semantics_hash(config, route, semantics_key) + auth = resolve_secret(endpoint.auth, secret_resolver=self.secret_resolver) + secret_fingerprints = {endpoint.auth.ref: auth.runtime_fingerprint} + headers = dict(endpoint.headers) + for header, reference in endpoint.header_refs.items(): + resolved = resolve_secret(reference, secret_resolver=self.secret_resolver) + headers[header] = resolved.value + secret_fingerprints[reference.ref] = resolved.runtime_fingerprint + selector = config.route_selectors[route.selector_id] + if config.schema_version == 3: + params = dict(config.purpose_defaults.get(purpose, {})) + params.update( + project_user_options_for_purpose( + values=model.params, + user_options=model.user_options, + purpose=purpose, + ) + ) + for name, option in model.user_options.items(): + if "default" in option and purpose in set( + option.get("applies_to") or ("main_agent",) + ): + params.setdefault(name, option["default"]) + params.update( + project_user_options_for_purpose( + values=selector.alias_defaults, + user_options=model.user_options, + purpose=purpose, + ) + ) + params.update(model.purpose_overrides.get(purpose, {})) + registration = get_adapter_registry().get( + provider.adapter_id, provider.adapter_revision + ) + options = self._validated_user_options( + model, + registration, + project_user_options_for_purpose( + values=model_options or {}, + user_options=model.user_options, + purpose=purpose, + ), + purpose, + ) + params.update(options) + registration.validate_parameters(params, path=f"invocation.{purpose}") + if registration.lifecycle == "blocked": + raise EvoRuntimeError("MODEL_ADAPTER_BLOCKED") + invocation_identity = invocation_fingerprint( + config, route, purpose, params, fingerprint_key + ) + else: + params = { + **config.runtime_defaults, + **provider.params, + **endpoint.params, + **model.params, + } + registration = None + invocation_identity = route_fingerprint(config, route, fingerprint_key) + supports_tools = bool(model.capabilities.get("tools", False)) + effective_tool_transport = ( + route.tool_call_transport if supports_tools else "disabled" + ) + identity = RouteIdentity( + config_revision=config.config_revision, + config_identity_key_id=config.config_identity_key_id, + purpose=purpose, + route_selector_id=selector.identity_selector_id or route.selector_id, + route_fingerprint=invocation_identity, + provider_id=route.provider, + endpoint_name=route.endpoint, + model_id=model.model_id, + protocol=provider.adapter_id or provider.protocol, + api_mode=route.api_mode, + tool_call_transport=effective_tool_transport, + route_semantics_hash=semantics_hash, + billing_sku=model.billing_sku, + pricing_revision=model.quote.pricing_revision, + quote_id=model.quote.quote_id, + ) + return ResolvedRoute( + route, + identity, + model.quote, + model.context_window, + model.max_output_tokens, + model.reasoning_mode, + model.reasoning_enabled_params, + model.reasoning_disabled_params, + endpoint.base_url, + params, + auth.value, + headers, + secret_fingerprints, + registration, + int(provider.connection_defaults.get("max_inflight_requests", 16)), + int( + model.max_inflight_requests + or provider.connection_defaults.get("max_inflight_requests", 16) + ), + int(provider.connection_defaults.get("queue_timeout_seconds", 5)), + f"provider:{route.provider}:{endpoint_fingerprint(config, route, endpoint_key)}", + f"model:{route.key()}:{semantics_hash}", + int(provider.connection_defaults.get("attempt_timeout_seconds", 600)), + supports_tools=supports_tools, + ) + + @staticmethod + def _validated_user_options( + model: Any, + registration: AdapterRegistration, + supplied: Mapping[str, Any], + purpose: str, + ) -> Mapping[str, Any]: + if {"reasoning", "reasoning_effort"} & set(supplied): + raise EvoRuntimeError("AGENT_INPUT_MISMATCH") + validated = validate_user_model_options( + supplied=supplied, + user_options=model.user_options, + supports_reasoning=model.supports_reasoning, + reasoning_mode=model.reasoning_mode, + allowed_reasoning_efforts=model.allowed_reasoning_efforts, + default_reasoning_effort=str( + model.reasoning_enabled_params.get("reasoning") or "" + ), + parameter_constraints=model.parameter_constraints, + purpose=purpose, + allow_reasoning=False, + ) + for name, value in validated.items(): + registration.all_parameter_schema[name].validate( + value, f"model_options.{name}" + ) + return validated + + def _build_call_bounds( + self, + purpose_routes: Mapping[str, tuple[ResolvedRoute, ...]], + ) -> Mapping[str, tuple[RouteCallBound, ...]]: + result: dict[str, tuple[RouteCallBound, ...]] = {} + for purpose, routes in purpose_routes.items(): + bounds = [] + for route in routes: + override = route.params.get("output_token_limit") + output_limit = ( + route.max_output_tokens if override is None else int(override) + ) + if output_limit > route.max_output_tokens: + raise EvoRuntimeError("MODEL_OUTPUT_LIMIT_EXCEEDED") + margin = _PROTOCOL_MARGIN_TOKENS.get( + (route.identity.protocol, route.identity.api_mode) + ) + if margin is None: + raise EvoRuntimeError("TOKEN_BOUND_UNAVAILABLE") + billable_cap = route.context_window - output_limit + payload_cap = billable_cap - margin + if payload_cap <= 0: + raise EvoRuntimeError("MODEL_CONTEXT_WINDOW_EXCEEDED") + reserve_rate = max( + route.quote.input_microunits_per_million, + route.quote.cached_input_microunits_per_million, + ) + reserve = _checked_ceil_cost( + billable_cap, + reserve_rate, + output_limit, + route.quote.output_microunits_per_million, + route.quote.unit_scale, + ) + bounds.append( + RouteCallBound( + route.identity, + output_limit, + payload_cap, + billable_cap, + margin, + reserve, + ) + ) + result[purpose] = tuple(bounds) + return result + + @staticmethod + def _run_reserve( + attempt_limits: Mapping[str, int], + bounds: Mapping[str, tuple[RouteCallBound, ...]], + ) -> int: + """Do not turn model-call estimates into a runtime admission limit. + + Agent progress is bounded by LangGraph's per-run ``recursion_limit``. + The legacy purpose counts remain in the signed quote for wire + compatibility and reporting, but a normal model/tool loop must not be + stopped because it outgrows that estimate. + """ + del attempt_limits, bounds + return 0 + + def _build_model_set( + self, snapshot: ModelRuntimeSnapshot, user_reasoning_effort: str + ) -> AgentModelSet: + del user_reasoning_effort + + def models(purpose: str) -> tuple[Any, ...]: + return tuple( + self._build_model(route, purpose, snapshot) + for route in snapshot.purpose_routes[purpose] + ) + + main = models("main_agent") + # A fallback chain is a route choice, not a workflow-call budget. + # Do not duplicate the primary model from max_attempts_per_run: that + # field is legacy quote metadata and must not cap normal agent turns. + retry_models = main[1:] + return AgentModelSet( + main_agent=main[0], + tool_selector=models("tool_selector")[0], + deepagents_summarizer=models("deepagents_summarizer")[0], + title=models("title")[0] if "title" in snapshot.purpose_routes else None, + main_fallbacks=retry_models, + route_health=self._route_health, + capacity=self._capacity, + ) + + def _build_model( + self, + route: ResolvedRoute, + purpose: str, + snapshot: ModelRuntimeSnapshot, + ) -> Any: + bound = next( + item + for item in snapshot.purpose_route_call_bounds[purpose] + if item.route_identity == route.identity + ) + if not route.runtime_provider or route.invocation_plan is None: + raise EvoRuntimeError("MODEL_ADAPTER_COMPILE_FAILED") + model = self.model_factory( + model=route.identity.model_id, + provider=route.runtime_provider, + **self._model_factory_kwargs(route), + ) + return self._attach_route_metadata(model, route, purpose, bound) + + @staticmethod + def _model_factory_kwargs(route: ResolvedRoute) -> dict[str, Any]: + plan = route.invocation_plan + if plan is None: + raise EvoRuntimeError("MODEL_ADAPTER_COMPILE_FAILED") + params = plan.model_kwargs() + params.update( + { + "api_key": route.api_key, + "base_url": route.base_url, + "max_retries": 0, + "timeout": route.attempt_timeout_seconds, + } + ) + if route.default_headers: + params["default_headers"] = dict(route.default_headers) + return params + + def _compile_purpose_routes( + self, + purpose_routes: Mapping[str, tuple[ResolvedRoute, ...]], + bounds: Mapping[str, tuple[RouteCallBound, ...]], + user_reasoning_effort: str, + ) -> Mapping[str, tuple[ResolvedRoute, ...]]: + """Compile each frozen purpose route exactly once during Prepare.""" + + compiled: dict[str, tuple[ResolvedRoute, ...]] = {} + for purpose, routes in purpose_routes.items(): + by_identity = {item.route_identity: item for item in bounds[purpose]} + compiled[purpose] = tuple( + self._compile_route( + route, + purpose, + by_identity[route.identity], + user_reasoning_effort, + ) + for route in routes + ) + return compiled + + @staticmethod + def _compile_route( + route: ResolvedRoute, + purpose: str, + bound: RouteCallBound, + user_reasoning_effort: str, + ) -> ResolvedRoute: + params = dict(route.params) + effort = user_reasoning_effort if purpose == "main_agent" else "disabled" + if route.adapter is not None: + if effort != "disabled": + params["reasoning_effort"] = effort + else: + params["thinking_enabled"] = False + compiled = dict( + route.adapter.compile_runtime_parameters( + route.identity.api_mode, + params, + bound.max_output_tokens, + provider_model_id=route.identity.model_id, + ) + ) + runtime_provider = ( + "google_interactions" + if route.adapter.adapter_id == "google-gemini" + and route.identity.api_mode == "interactions" + else route.adapter.runtime_provider + ) + plan = compile_invocation_plan( + api_mode=route.identity.api_mode, + declared_tool_call_transport=route.identity.tool_call_transport, + supports_tools=route.supports_tools, + purpose=purpose, + output_token_limit=bound.max_output_tokens, + reasoning_effort=effort, + runtime_provider=runtime_provider, + sdk_params=compiled, + ) + return replace( + route, + runtime_provider=runtime_provider, + invocation_plan=plan, + ) + + params.update( + { + "max_tokens": bound.max_output_tokens, + "disable_streaming": "tool_calling", + "streaming": False + if purpose != "main_agent" + else params.get("streaming", False), + "use_responses_api": route.identity.api_mode == "responses", + } + ) + if route.default_headers: + params["default_headers"] = dict(route.default_headers) + if effort != "disabled" and route.reasoning_mode == "effort": + params["reasoning_effort"] = effort + elif route.reasoning_mode == "boolean": + params.pop("reasoning_effort", None) + params = _merge_params( + params, + route.reasoning_enabled_params + if effort != "disabled" + else route.reasoning_disabled_params, + ) + else: + params.pop("reasoning_effort", None) + plan = compile_invocation_plan( + api_mode=route.identity.api_mode, + declared_tool_call_transport=route.identity.tool_call_transport, + supports_tools=route.supports_tools, + purpose=purpose, + output_token_limit=bound.max_output_tokens, + reasoning_effort=effort, + runtime_provider=route.identity.protocol, + sdk_params=params, + ) + return replace( + route, + runtime_provider=route.identity.protocol, + invocation_plan=plan, + ) + + @staticmethod + def _attach_route_metadata( + model: Any, route: ResolvedRoute, purpose: str, bound: RouteCallBound + ) -> Any: + metadata = { + "route_key": route.identity.route_key, + "route_provider": route.identity.provider_id, + "route_model": route.identity.model_id, + "route_endpoint": route.identity.endpoint_name, + "route_api_mode": route.invocation_plan.api_mode + if route.invocation_plan + else route.identity.api_mode, + "route_tool_call_transport": route.invocation_plan.tool_call_transport + if route.invocation_plan + else route.identity.tool_call_transport, + "route_supports_tools": route.supports_tools, + "route_invocation_plan_hash": route.invocation_plan.plan_hash + if route.invocation_plan + else "", + "route_output_token_parameter": route.invocation_plan.output_token_parameter + if route.invocation_plan + else "", + "route_adapter_id": route.adapter.adapter_id if route.adapter else "", + "route_adapter_revision": route.adapter.adapter_revision + if route.adapter + else "", + "route_config_generation": route.identity.config_revision, + "runtime_purpose": purpose, + "runtime_output_limit": bound.max_output_tokens, + "capacity_provider_key": route.provider_capacity_key, + "capacity_model_key": route.model_capacity_key, + "provider_max_inflight": route.provider_max_inflight, + "model_max_inflight": route.model_max_inflight, + "capacity_queue_timeout_seconds": route.queue_timeout_seconds, + "attempt_timeout_seconds": route.attempt_timeout_seconds, + } + if hasattr(model, "model_copy"): + model = model.model_copy( + update={ + "metadata": {**(getattr(model, "metadata", None) or {}), **metadata} + } + ) + return model + + def _validate_preparation_echo( + self, grant: RoutePreparationGrant, agent_input: AgentInputV3 + ) -> None: + if ( + agent_input.checkpoint_thread_id != grant.checkpoint_thread_id + or self.quote_authority.agent_input_digest(agent_input.projection()) + != grant.gateway_input_digest + or grant.turn_fencing_token < 1 + or grant.reasoning_effort + not in {"disabled", "low", "medium", "high", "max"} + ): + raise EvoRuntimeError("AGENT_INPUT_MISMATCH") + + @staticmethod + def _validate_admission_echo( + quote: PreparedRunQuote, admission: AdmissionGrant + ) -> None: + names = ( + "preparation_id", + "request_id", + "turn_id", + "thread_id", + "subject_id", + "requested_model_ref", + "plan", + "roles", + "requires_vision", + "reasoning_effort", + "title_policy", + "gateway_input_digest", + "prepared_snapshot_digest", + "prepared_input_digest", + "config_revision", + "catalog_revision", + "purpose_attempt_limits", + "total_max_attempts", + "checkpoint_snapshot_id", + "tool_registry_snapshot_id", + "turn_fencing_token", + "provider_run_reserve_microunits", + ) + if any(getattr(quote, name) != getattr(admission, name) for name in names): + raise EvoRuntimeError("ADMISSION_INVALID") + if ( + admission.billing_fencing_token < 1 + or not admission.admission_snapshot_id + or not admission.admission_id + or not admission.hold_id + ): + raise EvoRuntimeError("ADMISSION_INVALID") + + def _validate_still_fresh(self, handle: _PreparedHandle) -> None: + config = self.store.load_revision(handle.snapshot.config_revision) + if self._tool_registry_payload(handle.host) != handle.tool_registry_payload: + raise EvoRuntimeError("TOOL_REGISTRY_STALE") + for routes in handle.snapshot.purpose_routes.values(): + for route in routes: + endpoint = config.providers[route.ref.provider].endpoints[ + route.ref.endpoint + ] + for reference in (endpoint.auth, *endpoint.header_refs.values()): + resolved = resolve_secret( + reference, secret_resolver=self.secret_resolver + ) + if ( + handle.secret_fingerprints.get(reference.ref) + != resolved.runtime_fingerprint + ): + raise EvoRuntimeError("PREPARATION_STALE") + + @staticmethod + def _tool_registry_payload(host: WebHostContext) -> Mapping[str, Any]: + registry = host.tool_registry + revision = host.tool_registry_revision + if host.tool_registry_provider is not None: + registry, revision = host.tool_registry_provider() + tools = [] + for tool in registry: + if isinstance(tool, Mapping): + tools.append( + { + "name": str(tool.get("name") or ""), + "description": str(tool.get("description") or ""), + "schema": dict(tool.get("schema") or {}), + } + ) + continue + schema = getattr(tool, "args_schema", None) + if hasattr(schema, "model_json_schema"): + schema = schema.model_json_schema() + tools.append( + { + "name": str( + getattr( + tool, "name", getattr(tool, "__name__", type(tool).__name__) + ) + ), + "description": str(getattr(tool, "description", "")), + "schema": schema or {}, + } + ) + return {"revision": revision, "tools": tools} + + def _expire_prepared_locked(self) -> None: + current = now_ms() + for handle in tuple(self._prepared.values()): + if handle.state == "PREPARED" and handle.quote.expires_at < current: + handle.state = "EXPIRED" + self._retire_prepared_locked(handle) + for key, tombstone in tuple(self._prepared_tombstones.items()): + if tombstone.retained_until < current: + self._prepared_tombstones.pop(key, None) + + def _prepared_replay_locked( + self, request_key: tuple[str, str, str], request_digest: str + ) -> PreparedRunQuote | None: + tombstone = self._prepared_tombstones.get(request_key) + if tombstone is not None: + if tombstone.request_digest != request_digest: + raise EvoRuntimeError("PREPARATION_CONFLICT") + raise EvoRuntimeError("PREPARATION_STALE") + preparation_id = self._prepared_requests.get(request_key) + if preparation_id is None: + return None + handle = self._prepared.get(preparation_id) + if handle is None: + self._prepared_requests.pop(request_key, None) + return None + if handle.state in {"CANCELLED", "EXPIRED"}: + self._retire_prepared_locked(handle) + if handle.request_digest != request_digest: + raise EvoRuntimeError("PREPARATION_CONFLICT") + raise EvoRuntimeError("PREPARATION_STALE") + if handle.request_digest != request_digest: + raise EvoRuntimeError("PREPARATION_CONFLICT") + return handle.quote + + def _retire_prepared_locked(self, handle: _PreparedHandle) -> None: + request_key = ( + handle.grant.subject_id, + handle.grant.request_id, + handle.grant.turn_id, + ) + self._prepared_tombstones[request_key] = _PreparedTombstone( + handle.request_digest, + handle.state, + now_ms() + 24 * 60 * 60 * 1000, + ) + self._prepared_requests.pop(request_key, None) + self._prepared.pop(handle.quote.preparation_id, None) + + @staticmethod + def _preparation_request_digest( + grant: RoutePreparationGrant, + agent_input: AgentInputV3, + tool_registry_payload: Mapping[str, Any], + ) -> str: + grant_payload = grant.unsigned_payload() + for ephemeral in ("grant_id", "issued_at", "expires_at", "key_id"): + grant_payload.pop(ephemeral, None) + payload = { + "grant": grant_payload, + "agent_input": agent_input.projection(), + "tool_registry": tool_registry_payload, + } + return hashlib.sha256(canonical_json_v1(payload)).hexdigest() + + @staticmethod + def _access_allowed( + model: Any, plan: str, roles: tuple[str, ...], subject_id: str = "" + ) -> bool: + if model.allowed_plans and plan not in model.allowed_plans: + return False + if model.allowed_roles and not bool(set(roles) & set(model.allowed_roles)): + return False + return EvoModelRuntime._access_policy_allowed( + getattr(model, "access", {}), roles, subject_id + ) + + @staticmethod + def _access_policy_allowed( + policy: Mapping[str, Any], roles: tuple[str, ...], subject_id: str + ) -> bool: + if not policy or policy.get("visibility") == "authenticated": + return True + if policy.get("visibility") == "role_based": + role_set = set(roles) + group_set = { + role.removeprefix("group:") + for role in role_set + if role.startswith("group:") + } + return bool( + role_set & set(policy.get("roles") or ()) + or group_set & set(policy.get("groups") or ()) + ) + if policy.get("visibility") == "private": + return subject_id in set(policy.get("users") or ()) + return False + + @staticmethod + def _default_model_factory(**kwargs: Any) -> Any: + if kwargs.get("provider") == "google_interactions": + from .gemini_interactions import create_gemini_interactions_model + + return create_gemini_interactions_model(**kwargs) + from .models import get_chat_model + + return get_chat_model(**kwargs) + + @staticmethod + def _default_agent_factory( + snapshot: ModelRuntimeSnapshot, host: WebHostContext, model_set: AgentModelSet + ) -> Any: + from ..web_runtime import create_web_agent + + return create_web_agent(snapshot=snapshot, host=host, model_set=model_set) + + +class _EvoWebRun: + def __init__( + self, + *, + runtime: EvoModelRuntime, + admission: AdmissionGrant, + prepared: _PreparedHandle, + agent: Any, + model_set: AgentModelSet, + ) -> None: + self._runtime = runtime + self._admission = admission + self._prepared = prepared + self._snapshot = prepared.snapshot + self._input = prepared.input + self._host = prepared.host + self._agent = agent + self._model_set = model_set + self._run_id = str(uuid.uuid4()) + self._state = "ACTIVE" + self._sequence = 0 + self._journal: list[EvoRuntimeEvent] = [] + self._journal_bytes = 0 + self._condition = asyncio.Condition() + self._event_commit_lock = asyncio.Lock() + self._budget_lock = asyncio.Lock() + self._terminal_lock = asyncio.Lock() + self._agent_task: asyncio.Task[None] | None = None + self._stream_wakeup_task: asyncio.Task[None] | None = None + self._attempt_counts: dict[str, int] = defaultdict(int) + self._remaining = admission.provider_run_reserve_microunits + self._callback_attempts: dict[ + str, tuple[ModelAttemptEvent, int, PricingQuote, ResolvedRoute] + ] = {} + self._last_model_failure: tuple[str, Mapping[str, Any]] | None = None + self._force_context_repair_pending = bool( + self._input.metadata.get("force_context_repair", False) + ) + self._attempt_callback = _RuntimeAttemptCallback(self) + self._terminal_event: EvoRuntimeEvent | None = None + + @property + def run_id(self) -> str: + return self._run_id + + async def stream( + self, after_sequence: int | None = None + ) -> AsyncIterator[EvoRuntimeEvent]: + async with self._condition: + if after_sequence is not None and ( + isinstance(after_sequence, bool) or not isinstance(after_sequence, int) + ): + raise EvoRuntimeError("EVENT_CURSOR_INVALID") + cursor = 0 if after_sequence is None else after_sequence + if cursor < 0 or cursor > self._sequence: + raise EvoRuntimeError("EVENT_CURSOR_INVALID") + if self._agent_task is None: + self._agent_task = asyncio.create_task( + self._run_with_timeout(), name=f"evo-run-{self._run_id}" + ) + self._agent_task.add_done_callback(self._wake_stream_waiters) + async for event in self._replay(cursor): + yield event + + def _wake_stream_waiters(self, _task: asyncio.Task[None]) -> None: + self._stream_wakeup_task = asyncio.create_task(self._notify_stream_waiters()) + + async def _notify_stream_waiters(self) -> None: + async with self._condition: + self._condition.notify_all() + + async def cancel(self, reason: str) -> str: + if self._terminal_event is not None: + return str(self._terminal_event.payload.get("outcome") or "terminal") + if self._agent_task is not None: + self._agent_task.cancel() + await self._terminal_locked( + "cancelled", error_code="RUN_CANCELLED", reason=reason + ) + return "cancelled" + + async def _run_with_timeout(self) -> None: + try: + async with asyncio.timeout(self._snapshot.active_run_timeout_seconds): + await self._run_agent() + except TimeoutError: + await self._terminal_locked("failed", error_code="RUN_TIMEOUT") + except asyncio.CancelledError: + return + + async def _run_agent(self) -> None: + try: + await self._append_locked( + "run", + { + "kind": "run_started", + "outcome": "active", + "preparation_id": self._admission.preparation_id, + "admission_snapshot_id": self._admission.admission_snapshot_id, + "admission_id": self._admission.admission_id, + "hold_id": self._admission.hold_id, + "turn_fencing_token": self._admission.turn_fencing_token, + "billing_fencing_token": self._admission.billing_fencing_token, + }, + ) + from ..stream.events import stream_agent_events + + async for source in stream_agent_events( + self._agent, + self._input.message, + self._input.checkpoint_thread_id, + metadata=dict(self._input.metadata), + media=[ + item.locator + if hasattr(item, "locator") + else str(item.get("locator") or "") + if isinstance(item, Mapping) + else str(item) + for item in self._input.media + ], + callbacks=[self._attempt_callback], + configurable={ + "turn_fencing_token": self._admission.turn_fencing_token, + "turn_lease_owner": self._admission.request_id, + }, + error_mode="raise", + ): + if str(source.get("type") or "") == "done": + continue + await self._append_locked("agent", dict(source)) + if self._admission.title_policy == "best_effort": + await self._run_title() + else: + await self._append_locked( + "title", {"kind": "skipped", "reason": "DISABLED"} + ) + await self._terminal_locked("completed") + except asyncio.CancelledError: + return + except Exception as exc: + error_code = _safe_error_code(exc, fallback="AGENT_RUNTIME_ERROR") + error_details = _run_failure_details(exc, error_code) + if self._last_model_failure is not None: + error_code, error_details = self._last_model_failure + await self._terminal_locked( + "failed", + error_code=error_code, + error_details=error_details, + ) + + async def _run_title(self) -> None: + if self._model_set.title is None: + await self._append_locked( + "title", {"kind": "skipped", "reason": "TITLE_RUNTIME_INVARIANT"} + ) + return + try: + async with asyncio.timeout(self._snapshot.title_start_timeout_seconds): + response = await self._model_set.title.ainvoke( + _title_prompt(str(self._input.message)), + config={"callbacks": [self._attempt_callback]}, + ) + title = _response_text(response).strip().strip('"').strip()[:100] + await self._append_locked("title", {"kind": "generated", "title": title}) + except Exception as exc: + await self._append_locked( + "title", {"kind": "skipped", "reason": _safe_error_code(exc)} + ) + + async def _begin_callback_attempt( + self, + *, + callback_run_id: str, + purpose: str, + route: ResolvedRoute, + provider_input_bound_tokens: int, + provider_input_breakdown: Mapping[str, int] | None = None, + ) -> None: + async with self._budget_lock: + attempt_index = self._attempt_counts[purpose] + 1 + bound = next( + item + for item in self._snapshot.purpose_route_call_bounds[purpose] + if item.route_identity == route.identity + ) + provider_health_state = self._runtime._provider_health.status( + route.provider_capacity_key + ) + health_state = self._runtime._route_health.status(route.identity.route_key) + if ( + provider_health_state == "open" + or health_state == "open" + or not self._runtime._route_health.acquire_half_open( + route.identity.route_key + ) + ): + self._attempt_counts[purpose] = attempt_index + rejected = ModelAttemptEvent( + request_id=self._admission.request_id, + turn_id=self._admission.turn_id, + run_id=self._run_id, + preparation_id=self._admission.preparation_id, + admission_snapshot_id=self._admission.admission_snapshot_id, + admission_id=self._admission.admission_id, + hold_id=self._admission.hold_id, + prepared_snapshot_digest=self._admission.prepared_snapshot_digest, + prepared_input_digest=self._admission.prepared_input_digest, + billing_fencing_token=self._admission.billing_fencing_token, + turn_fencing_token=self._admission.turn_fencing_token, + logical_call_id=str(uuid.uuid4()), + attempt_id=str(uuid.uuid4()), + purpose=purpose, + attempt_index=attempt_index, + identity=route.identity, + quote_id=route.quote.quote_id, + billing_intent="user_charge" + if purpose == "main_agent" + else "platform_cost", + outcome="rejected", + provider_request_started=False, + provider_input_bound_tokens=provider_input_bound_tokens, + provider_reserved_microunits=0, + health_state=health_state, + error_code="ROUTE_HEALTH_UNAVAILABLE", + ) + try: + await self._append_locked("model_attempt", event_payload(rejected)) + except Exception: + self._attempt_counts[purpose] -= 1 + raise + raise EvoRuntimeError("ROUTE_HEALTH_UNAVAILABLE") + if provider_input_bound_tokens > bound.payload_input_hard_cap: + raise EvoRuntimeError( + "MODEL_CONTEXT_WINDOW_EXCEEDED", + details=( + { + "provider_input_bound_tokens": provider_input_bound_tokens, + "payload_input_hard_cap": bound.payload_input_hard_cap, + **dict(provider_input_breakdown or {}), + }, + ), + ) + # Usage remains observable and is settled from the terminal event, + # but there is intentionally no per-run reserve or call-count gate. + reserved = 0 + self._attempt_counts[purpose] = attempt_index + logical_call_id = str(uuid.uuid4()) + attempt_id = str(uuid.uuid4()) + attempt = ModelAttemptEvent( + request_id=self._admission.request_id, + turn_id=self._admission.turn_id, + run_id=self._run_id, + preparation_id=self._admission.preparation_id, + admission_snapshot_id=self._admission.admission_snapshot_id, + admission_id=self._admission.admission_id, + hold_id=self._admission.hold_id, + prepared_snapshot_digest=self._admission.prepared_snapshot_digest, + prepared_input_digest=self._admission.prepared_input_digest, + billing_fencing_token=self._admission.billing_fencing_token, + turn_fencing_token=self._admission.turn_fencing_token, + logical_call_id=logical_call_id, + attempt_id=attempt_id, + purpose=purpose, + attempt_index=attempt_index, + identity=route.identity, + quote_id=route.quote.quote_id, + billing_intent="user_charge" + if purpose == "main_agent" + else "platform_cost", + outcome="started", + provider_request_started=True, + provider_input_bound_tokens=provider_input_bound_tokens, + provider_reserved_microunits=reserved, + health_state=health_state, + ) + self._callback_attempts[callback_run_id] = ( + attempt, + reserved, + route.quote, + route, + ) + try: + await self._append_locked("model_attempt", event_payload(attempt)) + except Exception: + self._callback_attempts.pop(callback_run_id, None) + self._remaining += reserved + self._attempt_counts[purpose] -= 1 + self._runtime._route_health.release_half_open(route.identity.route_key) + raise + + async def _finish_callback_attempt( + self, + callback_run_id: str, + *, + usage: Mapping[str, int | str | None] | None, + error_code: str | None, + health_scope: str = "model_route", + ) -> None: + async with self._budget_lock: + record = self._callback_attempts.pop(callback_run_id, None) + if record is None: + return + started, _reserved, quote, route = record + valid_usage = _normalize_usage(usage) + outcome = "failed" if error_code else "succeeded" + if valid_usage is not None: + _actual_cost(valid_usage, quote) + bound = next( + item + for item in self._snapshot.purpose_route_call_bounds[ + started.purpose + ] + if item.route_identity == started.identity + ) + if ( + int(valid_usage["input_tokens"]) > bound.billable_input_cap + or int(valid_usage["output_tokens"]) > bound.max_output_tokens + ): + valid_usage = None + error_code = "USAGE_LIMIT_VIOLATION" + outcome = "failed" + if valid_usage is None: + outcome = "usage_unconfirmed" + terminal = replace( + started, + outcome=outcome, + usage=valid_usage, + usage_available=valid_usage is not None, + error_code=error_code, + timestamp=now_ms(), + ) + if error_code: + if health_scope == "provider_connection": + self._runtime._provider_health.record_failure( + route.provider_capacity_key, error_code + ) + elif health_scope == "model_route": + self._runtime._route_health.record_failure( + started.identity.route_key, error_code + ) + else: + self._runtime._route_health.record_success(started.identity.route_key) + self._runtime._provider_health.record_success( + route.provider_capacity_key + ) + await self._append_locked("model_attempt", event_payload(terminal)) + self._runtime._record_runtime_observation( + route, + purpose=started.purpose, + outcome=outcome, + error_code=error_code, + ) + + async def _append_locked( + self, + kind: str, + payload: Mapping[str, Any], + *, + control_terminal: bool = False, + ) -> EvoRuntimeEvent: + async with self._event_commit_lock: + if self._terminal_event is not None and kind != "run": + return self._terminal_event + sequence = self._sequence + 1 + event = EvoRuntimeEvent( + event_id=str(uuid.uuid4()), + runtime_instance_id=self._runtime.runtime_instance_id, + run_id=self._run_id, + request_id=self._admission.request_id, + turn_id=self._admission.turn_id, + sequence=sequence, + kind=kind, + payload=dict(payload), + ) + event_size = len(canonical_json_v1(asdict(event))) + if not control_terminal and ( + sequence >= self._snapshot.max_run_journal_events + or self._journal_bytes + event_size + > self._snapshot.max_run_journal_bytes + ): + raise EvoRuntimeError("RUN_EVENT_LIMIT_EXCEEDED") + sink = self._host.runtime_event_sink + assert sink is not None + payload_digest = hashlib.sha256( + canonical_json_v1(asdict(event)) + ).hexdigest() + commit_error: Exception | None = None + try: + result = await sink.commit(event) + except EvoRuntimeError: + raise + except Exception as exc: + commit_error = exc + result = "unknown" + if result not in {"committed", "duplicate"}: + try: + confirmation = await sink.confirm(event.event_id, payload_digest) + except Exception as exc: + raise EvoRuntimeError("EVENT_COMMIT_INDETERMINATE") from exc + if confirmation == "committed": + result = "committed" + elif confirmation == "absent": + raise EvoRuntimeError("EVENT_INGRESS_UNAVAILABLE") from commit_error + elif confirmation == "conflict": + raise EvoRuntimeError("EVO_EVENT_CONFLICT") + else: + raise EvoRuntimeError("EVENT_COMMIT_INDETERMINATE") + async with self._condition: + self._sequence = sequence + self._journal_bytes += event_size + self._journal.append(event) + self._condition.notify_all() + return event + + async def _terminal_locked( + self, + outcome: str, + *, + error_code: str | None = None, + reason: str | None = None, + error_details: Mapping[str, Any] | None = None, + ) -> EvoRuntimeEvent: + async with self._terminal_lock: + if self._terminal_event is not None: + return self._terminal_event + event = await self._append_locked( + "run", + { + "kind": "run_terminal", + "preparation_id": self._admission.preparation_id, + "admission_snapshot_id": self._admission.admission_snapshot_id, + "admission_id": self._admission.admission_id, + "hold_id": self._admission.hold_id, + "turn_fencing_token": self._admission.turn_fencing_token, + "billing_fencing_token": self._admission.billing_fencing_token, + "outcome": outcome, + "error_code": error_code, + "reason": reason, + "error_details": dict(error_details or {}), + "purpose_attempt_counts": dict(self._attempt_counts), + "provider_reserve_remaining_microunits": self._remaining, + }, + control_terminal=True, + ) + async with self._condition: + self._state = "TERMINAL" + self._terminal_event = event + self._condition.notify_all() + return event + + async def _replay(self, cursor: int) -> AsyncIterator[EvoRuntimeEvent]: + while True: + async with self._condition: + while self._sequence <= cursor and self._terminal_event is None: + if self._agent_task is not None and self._agent_task.done(): + break + await self._condition.wait() + events = [event for event in self._journal if event.sequence > cursor] + terminal = self._terminal_event is not None + agent_task = self._agent_task + for event in events: + cursor = event.sequence + yield event + if terminal and cursor >= self._sequence: + return + if not events and agent_task is not None and agent_task.done(): + error = agent_task.exception() + if error is not None: + raise error + raise EvoRuntimeError("EVENT_REPLAY_UNAVAILABLE") + + +class _RuntimeAttemptCallback(AsyncCallbackHandler): + # LangChain otherwise logs and suppresses callback failures, which would let + # a Provider request start without a committed STARTED attempt event. + raise_error = True + + def __init__(self, run: _EvoWebRun) -> None: + self.run = run + self._stream_diagnostics: dict[str, _ModelStreamDiagnosticState] = {} + + async def on_chat_model_start( + self, + _serialized: dict[str, Any], + messages: list[list[Any]], + *, + run_id: Any, + metadata: Mapping[str, Any] | None = None, + **_kwargs: Any, + ) -> None: + purpose, route = self._route_for(metadata) + try: + callback_payload = _callback_messages_payload(messages) + input_bound = _provider_input_token_bound(callback_payload) + if purpose == "main_agent" and self.run._force_context_repair_pending: + self.run._force_context_repair_pending = False + raise EvoRuntimeError( + "MODEL_CONTEXT_WINDOW_EXCEEDED", + details=( + { + "reason": "forced_context_repair", + "repair_requested": True, + **input_bound.projection(), + }, + ), + ) + await self.run._begin_callback_attempt( + callback_run_id=str(run_id), + purpose=purpose, + route=route, + provider_input_bound_tokens=input_bound.total_tokens, + provider_input_breakdown=input_bound.projection(), + ) + summary = _callback_payload_debug_summary( + callback_payload, input_bound.total_tokens + ) + plan = route.invocation_plan + plan_params = ",".join( + sorted(str(key)[:64] for key in (plan.sdk_params if plan else {})) + ) + parameter_values = _invocation_parameters_debug(plan) + message_schema = _callback_message_schema_debug(callback_payload) + attempt_record = self.run._callback_attempts.get(str(run_id)) + attempt = attempt_record[0] if attempt_record is not None else None + started_at = time.monotonic() + self._stream_diagnostics[str(run_id)] = _ModelStreamDiagnosticState( + started_at=started_at, + last_checkpoint_at=started_at, + attempt_id=attempt.attempt_id if attempt is not None else "", + purpose=purpose, + route_key=route.identity.route_key, + ) + logger.info( + "model_request_debug call_id=%s attempt_id=%s attempt_index=%s " + "purpose=%s route=%s plan_hash=%s input_bound_tokens=%s " + "input_bytes=%s input_sha256=%s " + "api_mode=%s tool_transport=%s output=%s=%s reasoning=%s " + "sdk_param_keys=%s params=%s message_types=%s content_block_types=%s " + "tool_calls=%s tool_results=%s message_schema=%s", + str(run_id), + attempt.attempt_id if attempt is not None else "", + attempt.attempt_index if attempt is not None else "", + purpose, + route.identity.route_key, + plan.plan_hash if plan else "", + summary["input_bound_tokens"], + summary["input_bytes"], + summary["input_sha256"], + plan.api_mode if plan else route.identity.api_mode, + plan.tool_call_transport + if plan + else route.identity.tool_call_transport, + plan.output_token_parameter if plan else "", + plan.output_token_limit if plan else route.max_output_tokens, + plan.reasoning_effort if plan else "", + plan_params, + parameter_values, + summary["message_types"], + summary["content_block_types"], + summary["tool_calls"], + summary["tool_results"], + message_schema, + ) + logger.info( + "model_stream_checkpoint phase=request_started call_id=%s " + "attempt_id=%s purpose=%s route=%s plan_hash=%s " + "prepared_input_digest=%s visible_chunks=0", + str(run_id), + attempt.attempt_id if attempt is not None else "", + purpose, + route.identity.route_key, + plan.plan_hash if plan else "", + attempt.prepared_input_digest if attempt is not None else "", + ) + except EvoRuntimeError as exc: + self.run._last_model_failure = ( + exc.code, + _callback_start_failure_details(purpose, route, exc), + ) + raise + + async def on_llm_new_token( + self, + token: str | list[str | dict[str, Any]], + *, + chunk: Any = None, + run_id: Any, + **_kwargs: Any, + ) -> None: + state = self._stream_diagnostics.get(str(run_id)) + if state is None: + return + now = time.monotonic() + previous_visible_at = state.last_visible_at + state.visible_chunks += 1 + state.text_chars += _stream_token_chars(token) + kinds = _stream_chunk_kinds(token, chunk) + for kind in kinds: + state.kind_counts[kind] = state.kind_counts.get(kind, 0) + 1 + state.last_visible_at = now + should_log = ( + state.visible_chunks in {1, 10, 100} + or state.visible_chunks % 1_000 == 0 + or now - state.last_checkpoint_at >= 30.0 + ) + if not should_log: + return + phase = ( + "first_visible_chunk" if state.visible_chunks == 1 else "stream_progress" + ) + idle_before_ms = int( + max(0.0, now - (previous_visible_at or state.started_at)) * 1_000 + ) + state.last_checkpoint_at = now + logger.info( + "model_stream_checkpoint phase=%s call_id=%s attempt_id=%s purpose=%s " + "route=%s elapsed_ms=%s idle_before_ms=%s visible_chunks=%s " + "text_chars=%s chunk_kinds=%s", + phase, + str(run_id), + state.attempt_id, + state.purpose, + state.route_key, + int(max(0.0, now - state.started_at) * 1_000), + idle_before_ms, + state.visible_chunks, + state.text_chars, + _stream_kind_counts_debug(state.kind_counts), + ) + + async def on_llm_end(self, response: Any, *, run_id: Any, **_kwargs: Any) -> None: + self._log_terminal_checkpoint(str(run_id), phase="completed") + self.run._last_model_failure = None + await self.run._finish_callback_attempt( + str(run_id), usage=_usage_from_response(response), error_code=None + ) + + async def on_llm_error( + self, error: BaseException, *, run_id: Any, **_kwargs: Any + ) -> None: + record = self.run._callback_attempts.get(str(run_id)) + route = record[3] if record is not None else None + disposition = ( + route.adapter.classify_error(error) + if route is not None and route.adapter is not None + else None + ) + error_code = disposition.error_code if disposition else _safe_error_code(error) + parse_debug = _provider_json_decode_debug(error) + self._log_terminal_checkpoint(str(run_id), phase="error", error=error) + if parse_debug: + logger.warning( + "model_response_json_decode purpose=%s route=%s document_bytes=%s " + "document_sha256=%s line=%s column=%s position=%s " + "invalid_codepoint=%s inside_string=%s event_type=%s " + "window_bytes=%s window_sha256=%s control_codepoints=%s trace=%s", + record[0].purpose if record is not None else "", + route.identity.route_key if route is not None else "", + parse_debug["provider_json_document_bytes"], + parse_debug["provider_json_document_sha256"], + parse_debug["provider_json_line"], + parse_debug["provider_json_column"], + parse_debug["provider_json_position"], + parse_debug.get("provider_json_invalid_codepoint", ""), + parse_debug["provider_json_position_inside_string"], + parse_debug.get("provider_json_event_type", ""), + parse_debug["provider_json_window_bytes"], + parse_debug["provider_json_window_sha256"], + parse_debug.get("provider_json_control_codepoints", ""), + parse_debug["provider_json_trace"], + ) + failure_details = _provider_failure_details(error, route, error_code) + provider_message = _provider_error_message_debug(error, route) + logger.warning( + "model_provider_failure purpose=%s route=%s code=%s status=%s " + "provider_code=%s provider_param=%s api_mode=%s tool_transport=%s " + "output=%s=%s provider_message=%s", + record[0].purpose if record is not None else "", + route.identity.route_key if route is not None else "", + error_code, + failure_details.get("http_status", ""), + failure_details.get("provider_error_code", ""), + failure_details.get("provider_error_parameter", ""), + failure_details.get("api_mode", ""), + failure_details.get("tool_call_transport", ""), + failure_details.get("output_token_parameter", ""), + failure_details.get("output_token_limit", ""), + provider_message, + ) + self.run._last_model_failure = (error_code, failure_details) + await self.run._finish_callback_attempt( + str(run_id), + usage=None, + error_code=error_code, + health_scope=(disposition.health_scope if disposition else "model_route"), + ) + + def _log_terminal_checkpoint( + self, run_id: str, *, phase: str, error: BaseException | None = None + ) -> None: + state = self._stream_diagnostics.pop(run_id, None) + if state is None: + return + now = time.monotonic() + timeout_debug = _provider_stream_timeout_debug(error) + last_activity = state.last_visible_at or state.started_at + log = logger.warning if error is not None else logger.info + log( + "model_stream_checkpoint phase=%s call_id=%s attempt_id=%s purpose=%s " + "route=%s elapsed_ms=%s idle_since_visible_ms=%s visible_chunks=%s " + "text_chars=%s chunk_kinds=%s provider_chunks_received=%s " + "provider_idle_timeout_ms=%s error_type=%s", + phase, + run_id, + state.attempt_id, + state.purpose, + state.route_key, + int(max(0.0, now - state.started_at) * 1_000), + int(max(0.0, now - last_activity) * 1_000), + state.visible_chunks, + state.text_chars, + _stream_kind_counts_debug(state.kind_counts), + timeout_debug.get("provider_stream_chunks_received", ""), + timeout_debug.get("provider_stream_idle_timeout_ms", ""), + type(error).__name__[:128] if error is not None else "", + ) + + def _route_for( + self, metadata: Mapping[str, Any] | None + ) -> tuple[str, ResolvedRoute]: + values = dict(metadata or {}) + purpose = str(values.get("runtime_purpose") or "main_agent") + route_key = str(values.get("route_key") or "") + routes = self.run._snapshot.purpose_routes.get(purpose, ()) + for route in routes: + if route.identity.route_key == route_key or ( + not route_key and len(routes) == 1 + ): + return purpose, route + raise EvoRuntimeError("ADMISSION_STALE") + + +@dataclass(slots=True) +class _ModelStreamDiagnosticState: + started_at: float + last_checkpoint_at: float + attempt_id: str + purpose: str + route_key: str + visible_chunks: int = 0 + text_chars: int = 0 + last_visible_at: float | None = None + kind_counts: dict[str, int] = field(default_factory=dict) + + +def _stream_token_chars(token: Any) -> int: + if isinstance(token, str): + return len(token) + if not isinstance(token, list): + return 0 + return sum( + len(str(item.get("text") or "")) for item in token if isinstance(item, Mapping) + ) + + +def _stream_chunk_kinds(token: Any, chunk: Any) -> tuple[str, ...]: + kinds: set[str] = set() + if _stream_token_chars(token) > 0: + kinds.add("text") + message = getattr(chunk, "message", None) + content = getattr(message, "content", None) + if isinstance(content, list): + for block in content: + if not isinstance(block, Mapping): + continue + block_type = str(block.get("type") or "") + if "reasoning" in block_type: + kinds.add("reasoning") + elif block_type in {"text", "output_text"}: + kinds.add("text") + elif "tool" in block_type or "function" in block_type: + kinds.add("tool_call") + additional = getattr(message, "additional_kwargs", None) + if isinstance(additional, Mapping) and any( + key in additional for key in ("reasoning", "reasoning_content") + ): + kinds.add("reasoning") + if getattr(message, "tool_call_chunks", None): + kinds.add("tool_call") + return tuple(sorted(kinds or {"metadata"})) + + +def _stream_kind_counts_debug(counts: Mapping[str, int]) -> str: + return ",".join( + f"{kind}:{int(counts.get(kind, 0))}" + for kind in ("text", "reasoning", "tool_call", "metadata") + if counts.get(kind, 0) + ) + + +def _callback_start_failure_details( + purpose: str, route: ResolvedRoute, error: EvoRuntimeError +) -> dict[str, str | int]: + """Expose a bounded callback-start failure without provider contents.""" + + plan = route.invocation_plan + details: dict[str, str | int] = { + "failure_stage": "model_attempt_start", + "reason": error.code.lower(), + "purpose": purpose, + "provider": route.identity.provider_id, + "model": route.identity.model_id, + "route_key": route.identity.route_key, + "endpoint": route.identity.endpoint_name, + "api_mode": plan.api_mode if plan else route.identity.api_mode, + "tool_call_transport": ( + plan.tool_call_transport if plan else route.identity.tool_call_transport + ), + } + if plan is not None: + details["output_token_limit"] = plan.output_token_limit + details["output_token_parameter"] = plan.output_token_parameter + details["invocation_plan_hash"] = plan.plan_hash + for item in error.details: + for key, value in item.items(): + if isinstance(value, str | int | bool): + details[str(key)] = value + return details + + +def _checked_ceil_cost( + input_tokens: int, + input_rate: int, + output_tokens: int, + output_rate: int, + unit_scale: int, +) -> int: + numerator = input_tokens * input_rate + output_tokens * output_rate + if numerator < 0 or numerator > _BIGINT_MAX * unit_scale: + raise EvoRuntimeError("COST_OVERFLOW") + result = math.ceil(numerator / unit_scale) + if result > _BIGINT_MAX: + raise EvoRuntimeError("COST_OVERFLOW") + return result + + +def _payload_token_bound(payload: Any) -> int: + try: + return len(canonical_json_v1(payload)) + except Exception as exc: + raise EvoRuntimeError("TOKEN_BOUND_UNAVAILABLE") from exc + + +@dataclass(frozen=True, slots=True) +class _ProviderInputBound: + total_tokens: int + text_tokens: int + media_tokens: int + media_blocks: int + largest_media_bytes: int + + def projection(self) -> dict[str, int]: + return { + "provider_input_bound_tokens": self.total_tokens, + "text_input_bound_tokens": self.text_tokens, + "media_input_bound_tokens": self.media_tokens, + "media_blocks": self.media_blocks, + "largest_media_bytes": self.largest_media_bytes, + } + + +def _decode_media_payload(payload: str, mime: str) -> tuple[bytes, str]: + if payload.startswith("data:"): + try: + header, encoded = payload.split(",", 1) + resolved_mime = header[5:].split(";", 1)[0] or mime + if ";base64" not in header.lower(): + raise ValueError("media data URL is not base64 encoded") + return base64.b64decode(encoded, validate=True), resolved_mime + except (ValueError, binascii.Error) as exc: + raise EvoRuntimeError("TOKEN_BOUND_UNAVAILABLE") from exc + try: + return base64.b64decode(payload, validate=True), mime + except (ValueError, binascii.Error) as exc: + raise EvoRuntimeError("TOKEN_BOUND_UNAVAILABLE") from exc + + +def _bounded_media_block( + value: Mapping[str, Any], +) -> tuple[dict[str, Any], int, int] | None: + block = dict(value) + mime = str(block.get("mime_type") or "application/octet-stream") + payload: str | None = None + location: tuple[str, ...] = () + + if isinstance(block.get("base64"), str): + payload = str(block["base64"]) + location = ("base64",) + elif isinstance(block.get("url"), str) and str(block["url"]).startswith("data:"): + payload = str(block["url"]) + location = ("url",) + elif isinstance(block.get("image_url"), str) and str(block["image_url"]).startswith( + "data:" + ): + payload = str(block["image_url"]) + location = ("image_url",) + elif isinstance(block.get("image_url"), Mapping): + image_url = dict(block["image_url"]) + if isinstance(image_url.get("url"), str) and str(image_url["url"]).startswith( + "data:" + ): + payload = str(image_url["url"]) + location = ("image_url", "url") + elif isinstance(block.get("source"), Mapping): + source = dict(block["source"]) + if source.get("type") == "base64" and isinstance(source.get("data"), str): + payload = str(source["data"]) + mime = str(source.get("media_type") or mime) + location = ("source", "data") + elif isinstance(block.get("inline_data"), Mapping): + inline_data = dict(block["inline_data"]) + if isinstance(inline_data.get("data"), str): + payload = str(inline_data["data"]) + mime = str(inline_data.get("mime_type") or mime) + location = ("inline_data", "data") + + if payload is None: + return None + raw, mime = _decode_media_payload(payload, mime) + digest = hashlib.sha256(raw).hexdigest()[:24] + marker = f"" + if location == ("base64",): + block["base64"] = marker + elif location == ("url",): + block["url"] = marker + elif location == ("image_url",): + block["image_url"] = marker + elif location == ("image_url", "url"): + nested = dict(block["image_url"]) + nested["url"] = marker + block["image_url"] = nested + elif location == ("source", "data"): + nested = dict(block["source"]) + nested["data"] = marker + block["source"] = nested + elif location == ("inline_data", "data"): + nested = dict(block["inline_data"]) + nested["data"] = marker + block["inline_data"] = nested + # This is the existing conservative media bound, now applied to decoded + # media bytes rather than to the larger base64/JSON representation. + return block, (len(raw) + 2) // 3 + 512, len(raw) + + +def _project_media_for_bound(value: Any, metrics: dict[str, int]) -> Any: + if isinstance(value, Mapping): + bounded = _bounded_media_block(value) + if bounded is not None: + projected, media_bound, media_bytes = bounded + metrics["media_tokens"] += media_bound + metrics["media_blocks"] += 1 + metrics["largest_media_bytes"] = max( + metrics["largest_media_bytes"], media_bytes + ) + return { + key: _project_media_for_bound(item, metrics) + for key, item in projected.items() + } + return { + key: _project_media_for_bound(item, metrics) for key, item in value.items() + } + if isinstance(value, list | tuple): + return [_project_media_for_bound(item, metrics) for item in value] + return value + + +def _provider_input_token_bound(payload: Any) -> _ProviderInputBound: + metrics = {"media_tokens": 0, "media_blocks": 0, "largest_media_bytes": 0} + projected = _project_media_for_bound(payload, metrics) + text_tokens = _payload_token_bound(projected) + media_tokens = metrics["media_tokens"] + return _ProviderInputBound( + total_tokens=text_tokens + media_tokens, + text_tokens=text_tokens, + media_tokens=media_tokens, + media_blocks=metrics["media_blocks"], + largest_media_bytes=metrics["largest_media_bytes"], + ) + + +def _callback_messages_payload(messages: Sequence[Sequence[Any]]) -> list[list[Any]]: + """Convert LangChain callback messages to canonical-JSON-safe values.""" + + return [ + [ + message_to_dict(message) if isinstance(message, BaseMessage) else message + for message in batch + ] + for batch in messages + ] + + +def _callback_payload_debug_summary( + payload: Sequence[Sequence[Any]], input_bound_tokens: int +) -> dict[str, str | int]: + """Return an irreversible, content-free summary of a model request.""" + + message_types: dict[str, int] = defaultdict(int) + content_block_types: dict[str, int] = defaultdict(int) + tool_calls = 0 + tool_results = 0 + for batch in payload: + for item in batch: + if not isinstance(item, Mapping): + message_types[type(item).__name__] += 1 + continue + message_type = str(item.get("type") or "unknown")[:64] + message_types[message_type] += 1 + data = item.get("data") + data = data if isinstance(data, Mapping) else {} + if message_type == "tool": + tool_results += 1 + declared_tool_calls = data.get("tool_calls") + has_declared_tool_calls = isinstance(declared_tool_calls, list) + if isinstance(declared_tool_calls, list): + tool_calls += len(declared_tool_calls) + content = data.get("content") + if not isinstance(content, list): + continue + for block in content: + if isinstance(block, Mapping): + block_type = str(block.get("type") or "unknown")[:64] + if not has_declared_tool_calls and block_type in { + "tool_call", + "tool_use", + "function_call", + }: + tool_calls += 1 + else: + block_type = type(block).__name__[:64] + content_block_types[block_type] += 1 + encoded = canonical_json_v1(payload) + return { + "input_bound_tokens": input_bound_tokens, + "input_bytes": len(encoded), + "input_sha256": hashlib.sha256(encoded).hexdigest()[:24], + "message_types": ",".join( + f"{name}:{count}" for name, count in sorted(message_types.items()) + ), + "content_block_types": ",".join( + f"{name}:{count}" for name, count in sorted(content_block_types.items()) + ), + "tool_calls": tool_calls, + "tool_results": tool_results, + } + + +def _safe_invocation_parameter_value(value: Any, *, nested: bool = False) -> Any: + """Project invocation parameters onto a bounded, non-secret debug shape.""" + + if value is None or isinstance(value, (bool, int, float)): + return value + if isinstance(value, str): + return value[:128] + if isinstance(value, Mapping): + allowed = ( + _DEBUG_NESTED_PARAMETER_KEYS if nested else _DEBUG_INVOCATION_PARAMETER_KEYS + ) + return { + str(key): _safe_invocation_parameter_value(item, nested=True) + for key, item in sorted(value.items(), key=lambda pair: str(pair[0])) + if str(key) in allowed + } + if isinstance(value, (list, tuple)): + return [ + _safe_invocation_parameter_value(item, nested=True) for item in value[:16] + ] + return f"<{type(value).__name__}>" + + +def _invocation_parameters_debug(plan: InvocationPlan | None) -> str: + """Serialize only provider-call parameters that cannot contain credentials.""" + + if plan is None: + return "{}" + projected = { + str(key): _safe_invocation_parameter_value(value, nested=True) + for key, value in sorted(plan.sdk_params.items(), key=lambda pair: str(pair[0])) + if str(key) in _DEBUG_INVOCATION_PARAMETER_KEYS + } + return json.dumps( + projected, ensure_ascii=True, sort_keys=True, separators=(",", ":") + ) + + +def _callback_message_schema_debug( + payload: Sequence[Sequence[Any]], +) -> str: + """Describe each callback message without logging IDs or message content.""" + + messages: list[dict[str, Any]] = [] + for batch_index, batch in enumerate(payload[:8]): + for message_index, item in enumerate(batch[:128]): + if not isinstance(item, Mapping): + messages.append( + { + "batch": batch_index, + "index": message_index, + "type": type(item).__name__[:64], + } + ) + continue + data = item.get("data") + data = data if isinstance(data, Mapping) else {} + content = data.get("content") + descriptor: dict[str, Any] = { + "batch": batch_index, + "index": message_index, + "type": str(item.get("type") or "unknown")[:64], + "content_kind": type(content).__name__, + "empty": content in (None, "", []), + } + if isinstance(content, str): + descriptor["content_chars"] = len(content) + elif isinstance(content, list): + descriptor["content_items"] = len(content) + descriptor["content_types"] = [ + str(block.get("type") or "dict")[:64] + if isinstance(block, Mapping) + else type(block).__name__[:64] + for block in content[:64] + ] + descriptor["text_chars"] = sum( + len(str(block.get("text") or "")) + for block in content + if isinstance(block, Mapping) + ) + additional = data.get("additional_kwargs") + if isinstance(additional, Mapping) and additional: + descriptor["additional_keys"] = sorted( + str(key)[:64] for key in additional + ) + response_metadata = data.get("response_metadata") + if isinstance(response_metadata, Mapping) and response_metadata: + descriptor["response_metadata_keys"] = sorted( + str(key)[:64] for key in response_metadata + ) + tool_calls = data.get("tool_calls") + if isinstance(tool_calls, list) and tool_calls: + descriptor["tool_calls"] = len(tool_calls) + if data.get("name") is not None: + descriptor["has_name"] = True + if data.get("tool_call_id") is not None: + descriptor["has_tool_call_id"] = True + messages.append(descriptor) + return json.dumps( + messages, ensure_ascii=True, sort_keys=True, separators=(",", ":") + ) + + +def _provider_error_message_debug( + error: BaseException, route: ResolvedRoute | None +) -> str: + """Return a bounded provider explanation with route credentials removed.""" + + body = getattr(error, "body", None) + candidates: list[Any] = [] + if isinstance(body, Mapping): + nested = body.get("error") + if isinstance(nested, Mapping): + candidates.append(nested.get("message")) + candidates.append(body.get("message")) + candidates.append(getattr(error, "message", None)) + candidates.append(str(error)) + message = next( + ( + str(value) + for value in candidates + if isinstance(value, str) and value.strip() + ), + "", + ) + if not message: + return "" + secrets: list[str] = [] + if route is not None: + secrets.append(str(route.api_key or "")) + secrets.extend(str(value) for value in route.default_headers.values()) + for secret in secrets: + if len(secret) >= 4: + message = message.replace(secret, "") + message = _AUTH_VALUE_PATTERN.sub( + lambda match: match.group(1) + "", message + ) + return " ".join(message.split())[:1_024] + + +def _provider_json_decode_debug(error: BaseException) -> dict[str, str | int]: + """Describe a JSON decoder failure without exposing parsed content.""" + + if not isinstance(error, json.JSONDecodeError): + return {} + document = error.doc if isinstance(error.doc, str) else "" + raw_document = document.encode("utf-8", errors="replace") + position = max(0, min(int(error.pos), len(document))) + window = document[max(0, position - 128) : min(len(document), position + 129)] + raw_window = window.encode("utf-8", errors="replace") + control_counts: dict[int, int] = {} + for character in window: + codepoint = ord(character) + if codepoint < 0x20: + control_counts[codepoint] = control_counts.get(codepoint, 0) + 1 + control_summary = ",".join( + f"U+{codepoint:04X}:{count}" + for codepoint, count in sorted(control_counts.items())[:8] + ) + event_match = _OPENAI_RESPONSE_EVENT_TYPE_PATTERN.search(document[:4_096]) + frames = traceback.extract_tb(error.__traceback__)[-8:] + trace = ">".join( + f"{frame.filename.rsplit('/', 1)[-1]}:{frame.lineno}:{frame.name}" + for frame in frames + )[:1_024] + details: dict[str, str | int] = { + "provider_json_document_bytes": len(raw_document), + "provider_json_document_sha256": hashlib.sha256(raw_document).hexdigest()[:24], + "provider_json_line": int(error.lineno), + "provider_json_column": int(error.colno), + "provider_json_position": position, + "provider_json_position_inside_string": ( + "true" if _json_position_inside_string(document, position) else "false" + ), + "provider_json_window_bytes": len(raw_window), + "provider_json_window_sha256": hashlib.sha256(raw_window).hexdigest()[:24], + "provider_json_trace": trace, + } + if position < len(document): + details["provider_json_invalid_codepoint"] = f"U+{ord(document[position]):04X}" + if control_summary: + details["provider_json_control_codepoints"] = control_summary + if event_match is not None: + details["provider_json_event_type"] = event_match.group(1) + return details + + +def _json_position_inside_string(document: str, position: int) -> bool: + """Return lexical string state at a JSON position without parsing its content.""" + + inside_string = False + escaped = False + for character in document[: max(0, min(position, len(document)))]: + if not inside_string: + if character == '"': + inside_string = True + continue + if escaped: + escaped = False + elif character == "\\": + escaped = True + elif character == '"': + inside_string = False + return inside_string + + +def _provider_stream_timeout_debug( + error: BaseException | None, +) -> dict[str, str | int]: + """Extract structured idle-timeout counters without relying on error text.""" + + if error is None: + return {} + details: dict[str, str | int] = {} + chunks_received = getattr(error, "chunks_received", None) + if ( + isinstance(chunks_received, int) + and not isinstance(chunks_received, bool) + and chunks_received >= 0 + ): + details["provider_stream_chunks_received"] = chunks_received + timeout_seconds = getattr(error, "timeout_s", None) + if ( + isinstance(timeout_seconds, (int, float)) + and not isinstance(timeout_seconds, bool) + and math.isfinite(float(timeout_seconds)) + and float(timeout_seconds) >= 0 + ): + details["provider_stream_idle_timeout_ms"] = round( + float(timeout_seconds) * 1_000 + ) + return details + + +def _normalize_usage( + value: Mapping[str, Any] | None, +) -> Mapping[str, int | str] | None: + if value is None: + return None + try: + raw_input = value.get("input_tokens", value.get("prompt_tokens")) + input_details = value.get("input_token_details") + if not isinstance(input_details, Mapping): + input_details = value.get("prompt_tokens_details") + if not isinstance(input_details, Mapping): + input_details = {} + raw_cached = value.get( + "cached_input_tokens", + value.get( + "cached_tokens", + input_details.get( + "cache_read", + input_details.get( + "cached_tokens", + input_details.get("cache_read_input_tokens"), + ), + ), + ), + ) + raw_output = value.get("output_tokens", value.get("completion_tokens")) + input_tokens = int(raw_input) if raw_input is not None else None + cached_tokens = int(raw_cached) if raw_cached is not None else 0 + output_tokens = int(raw_output) if raw_output is not None else None + output_details = value.get("output_token_details") + if not isinstance(output_details, Mapping): + output_details = value.get("completion_tokens_details") + if not isinstance(output_details, Mapping): + output_details = {} + raw_reasoning = value.get( + "reasoning_tokens", + output_details.get("reasoning", output_details.get("reasoning_tokens")), + ) + reasoning_tokens = int(raw_reasoning) if raw_reasoning is not None else None + raw_total = value.get("total_tokens") + total_tokens = int(raw_total) if raw_total is not None else None + except (TypeError, ValueError): + return None + finality = str(value.get("usage_finality") or "confirmed") + if finality not in {"confirmed", "partial", "unconfirmed"}: + finality = "unconfirmed" + request_hash = str(value.get("provider_request_id_hash") or "") or None + if request_hash is None and value.get("provider_request_id"): + request_hash = ( + "sha256:" + + hashlib.sha256(str(value["provider_request_id"]).encode()).hexdigest() + ) + return NormalizedUsage( + input_tokens=input_tokens, + cached_input_tokens=cached_tokens, + output_tokens=output_tokens, + reasoning_tokens=reasoning_tokens, + total_tokens=total_tokens, + provider_request_id_hash=request_hash, + finality=finality, # type: ignore[arg-type] + ).confirmed_projection() + + +def _actual_cost(usage: Mapping[str, int | str], quote: PricingQuote) -> int: + input_tokens = int(usage["input_tokens"]) + cached_tokens = int(usage["cached_input_tokens"]) + output_tokens = int(usage["output_tokens"]) + uncached = input_tokens - cached_tokens + numerator = ( + uncached * quote.input_microunits_per_million + + cached_tokens * quote.cached_input_microunits_per_million + + output_tokens * quote.output_microunits_per_million + ) + if numerator > _BIGINT_MAX * quote.unit_scale: + raise EvoRuntimeError("COST_OVERFLOW") + return math.ceil(numerator / quote.unit_scale) + + +def _safe_error_code( + exc: BaseException, *, fallback: str = "MODEL_PROVIDER_ERROR" +) -> str: + if isinstance(exc, EvoRuntimeError): + return exc.code + if type(exc).__name__ == "GraphRecursionError": + return "AGENT_RECURSION_LIMIT_EXCEEDED" + stable_code = getattr(exc, "code", None) + if ( + isinstance(stable_code, str) + and 2 <= len(stable_code) <= 64 + and "A" <= stable_code[0] <= "Z" + and all( + character == "_" or character.isdigit() or "A" <= character <= "Z" + for character in stable_code + ) + ): + return stable_code + name = type(exc).__name__.upper() + if "TIMEOUT" in name: + return "PROVIDER_TIMEOUT" if fallback == "MODEL_PROVIDER_ERROR" else fallback + return fallback + + +def _run_failure_details( + exc: BaseException, error_code: str +) -> dict[str, str | int | bool] | None: + """Return content-free diagnostics for failures outside the model callback.""" + + if isinstance(exc, EvoRuntimeError) and exc.details: + merged: dict[str, str | int | bool] = {} + for item in exc.details: + for key, value in item.items(): + if isinstance(value, str | int | bool): + merged[str(key)] = value + return merged or None + if error_code != "AGENT_RUNTIME_ERROR": + return None + return { + "failure_stage": "agent_execution", + "reason": "unclassified_agent_exception", + "agent_error_type": type(exc).__name__[:128], + "agent_error_module": type(exc).__module__[:128], + } + + +def _provider_failure_details( + error: BaseException, + route: ResolvedRoute | None, + error_code: str, +) -> dict[str, str | int]: + """Project a provider failure into safe, actionable runtime diagnostics.""" + + details: dict[str, str | int] = { + "failure_stage": "provider_request", + "reason": _provider_failure_reason(error_code), + "provider_error_type": type(error).__name__[:128], + "provider_error_module": type(error).__module__[:128], + } + details.update(_provider_json_decode_debug(error)) + details.update(_provider_stream_timeout_debug(error)) + if route is not None: + plan = route.invocation_plan + details.update( + { + "provider": route.identity.provider_id, + "model": route.identity.model_id, + "route_key": route.identity.route_key, + "endpoint": route.identity.endpoint_name, + "api_mode": plan.api_mode if plan else route.identity.api_mode, + "tool_call_transport": ( + plan.tool_call_transport + if plan + else route.identity.tool_call_transport + ), + } + ) + if plan is not None: + details["output_token_limit"] = plan.output_token_limit + details["output_token_parameter"] = plan.output_token_parameter + details["invocation_plan_hash"] = plan.plan_hash + elif route.max_output_tokens > 0: + details["output_token_limit"] = route.max_output_tokens + if status_code := _provider_http_status(error): + details["http_status"] = status_code + if provider_code := _provider_error_code(error): + details["provider_error_code"] = provider_code + if provider_parameter := _provider_error_parameter(error): + details["provider_error_parameter"] = provider_parameter + return details + + +def _provider_http_status(error: BaseException) -> int | None: + status = getattr(error, "status_code", None) + if status is None: + status = getattr(getattr(error, "response", None), "status_code", None) + return status if isinstance(status, int) and 100 <= status <= 599 else None + + +def _provider_error_code(error: BaseException) -> str | None: + """Extract only a token-like provider code; never expose provider messages.""" + + values = [getattr(error, "code", None)] + body = getattr(error, "body", None) + if isinstance(body, Mapping): + values.extend((body.get("code"), body.get("type"))) + nested = body.get("error") + if isinstance(nested, Mapping): + values.extend((nested.get("code"), nested.get("type"))) + for value in values: + if not isinstance(value, str): + continue + normalized = value.strip() + if 1 <= len(normalized) <= 128 and all( + character.isascii() and (character.isalnum() or character in "_.:-") + for character in normalized + ): + return normalized + return None + + +def _provider_error_parameter(error: BaseException) -> str | None: + """Return only a known rejected request field, never provider prose.""" + + body = getattr(error, "body", None) + mappings = [body] if isinstance(body, Mapping) else [] + if isinstance(body, Mapping) and isinstance(body.get("error"), Mapping): + mappings.append(body["error"]) + for mapping in mappings: + for key in ("param", "parameter", "field"): + value = mapping.get(key) + if isinstance(value, str): + normalized = value.strip() + if normalized in _KNOWN_PROVIDER_REQUEST_FIELDS: + return normalized + message = mapping.get("message") + if isinstance(message, str): + lowered = message.lower() + for parameter in _KNOWN_PROVIDER_REQUEST_FIELDS: + if parameter in lowered: + return parameter + return None + + +_KNOWN_PROVIDER_REQUEST_FIELDS = ( + "max_output_tokens", + "max_completion_tokens", + "max_tokens", + "reasoning_effort", + "reasoning", + "temperature", + "top_p", + "tool_choice", + "tools", + "response_format", + "stream", + "messages", +) + + +def _provider_failure_reason(error_code: str) -> str: + return { + "MODEL_PROVIDER_REQUEST_REJECTED": "provider_rejected_request", + "MODEL_AUTHENTICATION_FAILED": "provider_authentication_failed", + "MODEL_NOT_FOUND": "provider_model_or_endpoint_not_found", + "MODEL_RATE_LIMITED": "provider_rate_limited", + "MODEL_TIMEOUT": "provider_timeout", + }.get(error_code, "provider_request_failed") + + +def _title_prompt(source_text: str) -> str: + normalized = " ".join(source_text.split())[:500] + return f"Generate a concise title of no more than 30 characters. Return only the title.\n\n{normalized}" + + +def _response_text(response: Any) -> str: + content = getattr(response, "content", response) + if isinstance(content, str): + return content + if isinstance(content, list): + return "".join( + str(item.get("text") or "") if isinstance(item, Mapping) else str(item) + for item in content + ) + return str(content) + + +def _usage_from_response(response: Any) -> Mapping[str, Any] | None: + usage = getattr(response, "usage_metadata", None) + if isinstance(usage, Mapping): + return usage + llm_output = getattr(response, "llm_output", None) + if isinstance(llm_output, Mapping): + candidate = llm_output.get("usage") or llm_output.get("token_usage") + if isinstance(candidate, Mapping): + return candidate + generations = getattr(response, "generations", None) + if isinstance(generations, Sequence) and not isinstance(generations, str | bytes): + for generation_group in generations: + items = ( + generation_group + if isinstance(generation_group, Sequence) + and not isinstance(generation_group, str | bytes) + else (generation_group,) + ) + for generation in items: + message = getattr(generation, "message", None) + for value in (generation, message): + nested_usage = getattr(value, "usage_metadata", None) + if isinstance(nested_usage, Mapping): + return nested_usage + response_metadata = getattr(value, "response_metadata", None) + if isinstance(response_metadata, Mapping): + candidate = response_metadata.get( + "usage" + ) or response_metadata.get("token_usage") + if isinstance(candidate, Mapping): + return candidate + generation_info = getattr(generation, "generation_info", None) + if isinstance(generation_info, Mapping): + candidate = generation_info.get("usage") or generation_info.get( + "token_usage" + ) + if isinstance(candidate, Mapping): + return candidate + return None diff --git a/EvoScientist/llm/secret_store.py b/EvoScientist/llm/secret_store.py new file mode 100644 index 0000000..b9791c7 --- /dev/null +++ b/EvoScientist/llm/secret_store.py @@ -0,0 +1,472 @@ +"""Encrypted, versioned storage for model-provider credentials. + +Only secret references cross the model-route configuration boundary. The +store deliberately has no API for reading a plaintext credential after it was +written; callers receive an opaque reference and masked metadata instead. +""" + +from __future__ import annotations + +import base64 +import hashlib +import os +import re +import sqlite3 +from dataclasses import dataclass +from pathlib import Path + +from cryptography.fernet import Fernet, InvalidToken +from filelock import FileLock + +from ..config.settings import get_config_dir +from .configuration import ResolvedSecret, SecretReference +from .contracts import EvoRuntimeError + +_SECRET_ID_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._/-]{0,127}$") + + +@dataclass(frozen=True, slots=True) +class SecretMetadata: + secret_id: str + version: int + masked_value: str + created_by: str + created_at: str + status: str = "active" + retired_at: str | None = None + revoked_at: str | None = None + revoked_by: str | None = None + revoke_reason: str | None = None + + @property + def ref(self) -> str: + return f"secret://{self.secret_id}#{self.version}" + + +class EncryptedModelSecretStore: + """SQLite-backed Fernet store scoped to the Evo configuration directory.""" + + def __init__( + self, + path: Path | None = None, + *, + master_secret: str | None = None, + ) -> None: + self.path = path or (get_config_dir() / "model_secrets.sqlite") + self.path.parent.mkdir(parents=True, exist_ok=True) + self._lock = FileLock(str(self.path) + ".lock") + material = master_secret or os.environ.get( + "AI4SCI_EVO_MODEL_SECRET_MASTER_KEY", "" + ) + if not material: + # The identity secret is already mandatory for the V3 runtime. + # Operators can configure a dedicated secret to separate rotation. + material = os.environ.get("AI4SCI_EVO_CONFIG_IDENTITY_SECRET", "") + if len(material.encode("utf-8")) < 32: + raise RuntimeError( + "AI4SCI_EVO_MODEL_SECRET_MASTER_KEY or " + "AI4SCI_EVO_CONFIG_IDENTITY_SECRET must contain at least 32 bytes" + ) + key = base64.urlsafe_b64encode(hashlib.sha256(material.encode()).digest()) + self._fernet = Fernet(key) + self._init_schema() + + def _connect(self) -> sqlite3.Connection: + connection = sqlite3.connect(self.path) + connection.row_factory = sqlite3.Row + connection.execute("PRAGMA foreign_keys=ON") + connection.execute("PRAGMA busy_timeout=5000") + connection.execute("PRAGMA synchronous=FULL") + connection.execute("PRAGMA journal_mode=WAL") + return connection + + def _init_schema(self) -> None: + with self._lock, self._connect() as connection: + connection.executescript( + """ + CREATE TABLE IF NOT EXISTS model_secret_versions ( + secret_id TEXT NOT NULL, + version INTEGER NOT NULL CHECK (version > 0), + ciphertext BLOB NOT NULL, + masked_value TEXT NOT NULL, + created_by TEXT NOT NULL, + created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP, + PRIMARY KEY (secret_id, version) + ); + CREATE INDEX IF NOT EXISTS idx_model_secret_versions_latest + ON model_secret_versions (secret_id, version DESC); + """ + ) + columns = { + str(row["name"]) + for row in connection.execute( + "PRAGMA table_info(model_secret_versions)" + ).fetchall() + } + additions = { + "status": "TEXT NOT NULL DEFAULT 'active'", + "retired_at": "TEXT", + "revoked_at": "TEXT", + "revoked_by": "TEXT", + "revoke_reason": "TEXT", + "transition_operation_id": "TEXT", + } + for name, definition in additions.items(): + if name not in columns: + connection.execute( + f"ALTER TABLE model_secret_versions ADD COLUMN {name} {definition}" + ) + # Legacy stores treated every version as active. Keep only the latest + # active version before installing the partial unique index. + connection.execute( + """UPDATE model_secret_versions AS current + SET status='retired', retired_at=COALESCE(retired_at, CURRENT_TIMESTAMP) + WHERE status='active' AND version < ( + SELECT MAX(newer.version) FROM model_secret_versions AS newer + WHERE newer.secret_id=current.secret_id + )""" + ) + connection.execute( + """CREATE UNIQUE INDEX IF NOT EXISTS uq_model_secret_active + ON model_secret_versions(secret_id) WHERE status='active'""" + ) + connection.execute( + """CREATE TABLE IF NOT EXISTS model_secret_operations ( + operation_id TEXT PRIMARY KEY, + action TEXT NOT NULL, + request_digest TEXT NOT NULL, + secret_id TEXT NOT NULL, + version INTEGER NOT NULL, + created_at TEXT NOT NULL DEFAULT CURRENT_TIMESTAMP + )""" + ) + try: + os.chmod(self.path, 0o600) + except OSError: + pass + + @staticmethod + def _validate_secret_id(secret_id: str) -> str: + normalized = str(secret_id or "").strip() + if not _SECRET_ID_RE.fullmatch(normalized): + raise EvoRuntimeError("LLM_SECRET_INVALID") + return normalized + + @staticmethod + def _mask(value: str) -> str: + if len(value) <= 8: + return "*" * len(value) + return f"{value[:4]}...{value[-4:]}" + + def put( + self, + secret_id: str, + value: str, + *, + created_by: str, + status: str = "active", + operation_id: str | None = None, + ) -> SecretMetadata: + secret_id = self._validate_secret_id(secret_id) + value = str(value or "") + if not value or any(char in value for char in "\r\n\0"): + raise EvoRuntimeError("LLM_SECRET_INVALID") + actor = str(created_by or "unknown")[:256] + if status not in {"pending", "active"}: + raise EvoRuntimeError("LLM_SECRET_INVALID") + ciphertext = self._fernet.encrypt(value.encode("utf-8")) + with self._lock, self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + row = connection.execute( + "SELECT COALESCE(MAX(version), 0) AS version " + "FROM model_secret_versions WHERE secret_id=?", + (secret_id,), + ).fetchone() + version = int(row["version"]) + 1 + if status == "active": + connection.execute( + """UPDATE model_secret_versions + SET status='retired', retired_at=CURRENT_TIMESTAMP, + transition_operation_id=? + WHERE secret_id=? AND status='active'""", + (operation_id, secret_id), + ) + connection.execute( + """INSERT INTO model_secret_versions + (secret_id, version, ciphertext, masked_value, created_by, + status, transition_operation_id) + VALUES (?, ?, ?, ?, ?, ?, ?)""", + ( + secret_id, + version, + ciphertext, + self._mask(value), + actor, + status, + operation_id, + ), + ) + saved = connection.execute( + """SELECT secret_id, version, masked_value, created_by, created_at, + status, retired_at, revoked_at, revoked_by, revoke_reason + FROM model_secret_versions WHERE secret_id=? AND version=?""", + (secret_id, version), + ).fetchone() + return SecretMetadata( + secret_id=str(saved["secret_id"]), + version=int(saved["version"]), + masked_value=str(saved["masked_value"]), + created_by=str(saved["created_by"]), + created_at=str(saved["created_at"]), + status=str(saved["status"]), + retired_at=saved["retired_at"], + revoked_at=saved["revoked_at"], + revoked_by=saved["revoked_by"], + revoke_reason=saved["revoke_reason"], + ) + + def create_pending( + self, + provider_id: str, + value: str, + *, + created_by: str, + operation_id: str, + ) -> SecretMetadata: + secret_id = f"model-providers/{self._validate_secret_id(provider_id)}" + digest = hashlib.sha256( + (secret_id + "\0" + str(value)).encode("utf-8") + ).hexdigest() + with self._lock: + with self._connect() as connection: + replay = connection.execute( + "SELECT * FROM model_secret_operations WHERE operation_id=?", + (operation_id,), + ).fetchone() + if replay is not None: + if ( + str(replay["action"]) != "create_pending" + or str(replay["request_digest"]) != digest + or str(replay["secret_id"]) != secret_id + ): + raise EvoRuntimeError("IDEMPOTENCY_CONFLICT") + return self._metadata(secret_id, int(replay["version"])) + result = self.put( + secret_id, + value, + created_by=created_by, + status="pending", + operation_id=operation_id, + ) + with self._connect() as connection: + connection.execute( + """INSERT INTO model_secret_operations + (operation_id, action, request_digest, secret_id, version) + VALUES (?, 'create_pending', ?, ?, ?)""", + (operation_id, digest, secret_id, result.version), + ) + return result + + def list_metadata(self) -> list[SecretMetadata]: + with self._connect() as connection: + rows = connection.execute( + """SELECT secret_id, version, masked_value, created_by, created_at, + status, retired_at, revoked_at, revoked_by, revoke_reason + FROM model_secret_versions + ORDER BY secret_id ASC, version DESC""" + ).fetchall() + return [ + SecretMetadata( + secret_id=str(row["secret_id"]), + version=int(row["version"]), + masked_value=str(row["masked_value"]), + created_by=str(row["created_by"]), + created_at=str(row["created_at"]), + status=str(row["status"]), + retired_at=row["retired_at"], + revoked_at=row["revoked_at"], + revoked_by=row["revoked_by"], + revoke_reason=row["revoke_reason"], + ) + for row in rows + ] + + def current_provider_metadata(self, provider_id: str) -> SecretMetadata | None: + secret_id = f"model-providers/{self._validate_secret_id(provider_id)}" + with self._connect() as connection: + row = connection.execute( + """SELECT secret_id, version, masked_value, created_by, created_at, + status, retired_at, revoked_at, revoked_by, revoke_reason + FROM model_secret_versions + WHERE secret_id=? AND status='active' + ORDER BY version DESC LIMIT 1""", + (secret_id,), + ).fetchone() + if row is None: + return None + return SecretMetadata( + secret_id=str(row["secret_id"]), + version=int(row["version"]), + masked_value=str(row["masked_value"]), + created_by=str(row["created_by"]), + created_at=str(row["created_at"]), + status=str(row["status"]), + retired_at=row["retired_at"], + revoked_at=row["revoked_at"], + revoked_by=row["revoked_by"], + revoke_reason=row["revoke_reason"], + ) + + def resolve(self, reference: SecretReference) -> ResolvedSecret: + if reference.ref.startswith("provider://"): + provider_id = self._validate_secret_id(reference.ref[11:]) + secret_id = f"model-providers/{provider_id}" + with self._connect() as connection: + row = connection.execute( + """SELECT ciphertext, version FROM model_secret_versions + WHERE secret_id=? AND status='active' + ORDER BY version DESC LIMIT 1""", + (secret_id,), + ).fetchone() + if row is None: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + try: + value = self._fernet.decrypt(bytes(row["ciphertext"])).decode("utf-8") + except (InvalidToken, UnicodeDecodeError) as exc: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") from exc + fingerprint = hashlib.sha256(value.encode("utf-8")).hexdigest() + return ResolvedSecret( + value, reference.revision, str(row["version"]), fingerprint + ) + if not reference.ref.startswith("secret://") or "#" not in reference.ref: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + secret_id, version_text = reference.ref[9:].rsplit("#", 1) + try: + version = int(version_text) + except ValueError as exc: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") from exc + if version != reference.revision: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + secret_id = self._validate_secret_id(secret_id) + with self._connect() as connection: + row = connection.execute( + """SELECT ciphertext, status FROM model_secret_versions + WHERE secret_id=? AND version=?""", + (secret_id, version), + ).fetchone() + if row is None: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + if str(row["status"]) in {"revoked", "destroyed"}: + raise EvoRuntimeError("MODEL_CREDENTIAL_REVOKED") + try: + value = self._fernet.decrypt(bytes(row["ciphertext"])).decode("utf-8") + except (InvalidToken, UnicodeDecodeError) as exc: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") from exc + fingerprint = hashlib.sha256(value.encode("utf-8")).hexdigest() + return ResolvedSecret(value, version, str(version), fingerprint) + + def activate( + self, secret_id: str, version: int, *, operation_id: str + ) -> SecretMetadata: + secret_id = self._validate_secret_id(secret_id) + with self._lock, self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + row = connection.execute( + "SELECT status FROM model_secret_versions WHERE secret_id=? AND version=?", + (secret_id, version), + ).fetchone() + if row is None or str(row["status"]) in {"revoked", "destroyed"}: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + connection.execute( + """UPDATE model_secret_versions SET status='retired', + retired_at=COALESCE(retired_at, CURRENT_TIMESTAMP), + transition_operation_id=? + WHERE secret_id=? AND status='active' AND version<>?""", + (operation_id, secret_id, version), + ) + connection.execute( + """UPDATE model_secret_versions SET status='active', retired_at=NULL, + transition_operation_id=? WHERE secret_id=? AND version=?""", + (operation_id, secret_id, version), + ) + return self._metadata(secret_id, version) + + def retire( + self, secret_id: str, version: int, *, operation_id: str + ) -> SecretMetadata: + secret_id = self._validate_secret_id(secret_id) + with self._lock, self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + row = connection.execute( + "SELECT status FROM model_secret_versions WHERE secret_id=? AND version=?", + (secret_id, version), + ).fetchone() + if row is None: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + if str(row["status"]) not in {"revoked", "destroyed"}: + connection.execute( + """UPDATE model_secret_versions SET status='retired', + retired_at=COALESCE(retired_at, CURRENT_TIMESTAMP), + transition_operation_id=? WHERE secret_id=? AND version=?""", + (operation_id, secret_id, version), + ) + return self._metadata(secret_id, version) + + def revoke( + self, + secret_id: str, + version: int, + *, + revoked_by: str, + reason: str, + operation_id: str, + ) -> SecretMetadata: + secret_id = self._validate_secret_id(secret_id) + clean_reason = str(reason or "").strip() + if not clean_reason: + raise EvoRuntimeError("LLM_SECRET_INVALID") + with self._lock, self._connect() as connection: + connection.execute("BEGIN IMMEDIATE") + row = connection.execute( + "SELECT status FROM model_secret_versions WHERE secret_id=? AND version=?", + (secret_id, version), + ).fetchone() + if row is None or str(row["status"]) == "destroyed": + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + if str(row["status"]) != "revoked": + connection.execute( + """UPDATE model_secret_versions SET status='revoked', + revoked_at=CURRENT_TIMESTAMP, revoked_by=?, revoke_reason=?, + transition_operation_id=? WHERE secret_id=? AND version=?""", + ( + revoked_by[:256], + clean_reason[:1024], + operation_id, + secret_id, + version, + ), + ) + return self._metadata(secret_id, version) + + def _metadata(self, secret_id: str, version: int) -> SecretMetadata: + with self._connect() as connection: + row = connection.execute( + """SELECT secret_id, version, masked_value, created_by, created_at, + status, retired_at, revoked_at, revoked_by, revoke_reason + FROM model_secret_versions WHERE secret_id=? AND version=?""", + (secret_id, version), + ).fetchone() + if row is None: + raise EvoRuntimeError("ROUTE_SECRET_UNAVAILABLE") + return SecretMetadata( + secret_id=str(row["secret_id"]), + version=int(row["version"]), + masked_value=str(row["masked_value"]), + created_by=str(row["created_by"]), + created_at=str(row["created_at"]), + status=str(row["status"]), + retired_at=row["retired_at"], + revoked_at=row["revoked_at"], + revoked_by=row["revoked_by"], + revoke_reason=row["revoke_reason"], + ) diff --git a/EvoScientist/llm/user_options.py b/EvoScientist/llm/user_options.py new file mode 100644 index 0000000..785a0aa --- /dev/null +++ b/EvoScientist/llm/user_options.py @@ -0,0 +1,258 @@ +"""Canonical, provider-neutral validation for user-adjustable model options.""" + +from __future__ import annotations + +import hashlib +import math +from collections.abc import Mapping, Sequence +from typing import Any + +from .contracts import EvoRuntimeError +from .crypto import canonical_json_v1 + +_REASONING_VALUES = frozenset({"off", "on", "low", "medium", "high", "max"}) +_MAX_OPTION_COUNT = 16 +_MAX_CANONICAL_BYTES = 4_096 + + +def project_user_options_for_purpose( + *, + values: Mapping[str, Any], + user_options: Mapping[str, Mapping[str, Any]], + purpose: str, +) -> dict[str, Any]: + """Return values whose user-option contract permits the target purpose. + + Parameters outside ``user_options`` are runtime-owned and remain untouched; + downstream adapter validation is still authoritative for those fields. + """ + + return { + name: value + for name, value in values.items() + if name not in user_options + or purpose + in set(user_options[name].get("applies_to") or ("main_agent",)) + } + + +def model_options_schema_hash( + *, + model_profile_id: str, + user_options: Mapping[str, Mapping[str, Any]], + supports_reasoning: bool, + reasoning_mode: str, + allowed_reasoning_efforts: Sequence[str], + parameter_constraints: Sequence[Mapping[str, Any]], + adapter_id: str = "", + adapter_revision: str = "", +) -> str: + """Hash only the public semantics that determine valid user input.""" + + public_options = { + str(name): { + key: ( + sorted(str(item) for item in value) + if key in {"choices", "applies_to"} + else value + ) + for key, value in dict(option).items() + if key + in { + "type", + "applies_to", + "minimum", + "maximum", + "minimum_exclusive", + "maximum_exclusive", + "choices", + } + } + for name, option in sorted(user_options.items()) + } + normalized_constraints = [ + { + key: ( + sorted(str(item) for item in value) + if key == "at_most_one_of" + else value + ) + for key, value in dict(constraint).items() + } + for constraint in parameter_constraints + ] + normalized_constraints.sort(key=canonical_json_v1) + payload = { + "model_profile_id": str(model_profile_id), + "user_options": public_options, + "supports_reasoning": bool(supports_reasoning), + "reasoning_mode": str(reasoning_mode), + "allowed_reasoning_efforts": sorted( + str(value) for value in allowed_reasoning_efforts + ), + "parameter_constraints": normalized_constraints, + "adapter_id": str(adapter_id), + "adapter_revision": str(adapter_revision), + } + return "sha256:" + hashlib.sha256(canonical_json_v1(payload)).hexdigest() + + +def validate_parameter_constraints( + constraints: Sequence[Mapping[str, Any]], *, allowed_names: set[str] +) -> tuple[dict[str, Any], ...]: + """Validate and normalize the supported public constraint grammar.""" + + normalized: list[dict[str, Any]] = [] + for constraint in constraints: + values = dict(constraint) + if set(values) != {"at_most_one_of"}: + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", "unsupported user option constraint" + ) + names = values["at_most_one_of"] + if not isinstance(names, list | tuple) or len(names) < 2: + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", "at_most_one_of must contain two names" + ) + clean = tuple(str(name) for name in names) + if len(set(clean)) != len(clean) or not set(clean) <= allowed_names: + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", "constraint references invalid options" + ) + normalized.append({"at_most_one_of": clean}) + return tuple(normalized) + + +def validate_user_model_options( + *, + supplied: Mapping[str, Any] | None, + user_options: Mapping[str, Mapping[str, Any]], + supports_reasoning: bool, + reasoning_mode: str, + allowed_reasoning_efforts: Sequence[str], + default_reasoning_effort: str | None, + parameter_constraints: Sequence[Mapping[str, Any]], + purpose: str = "main_agent", + allow_reasoning: bool = True, +) -> dict[str, Any]: + """Return canonical user options or raise a stable contract error.""" + + values = dict(supplied or {}) + if len(values) > _MAX_OPTION_COUNT: + raise EvoRuntimeError("MODEL_PARAMETER_INVALID", "too many model options") + try: + encoded = canonical_json_v1(values) + except (TypeError, ValueError) as exc: + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", "model options must be canonical JSON" + ) from exc + if len(encoded) > _MAX_CANONICAL_BYTES: + raise EvoRuntimeError("MODEL_PARAMETER_INVALID", "model options are too large") + + allowed = set(user_options) + if allow_reasoning and supports_reasoning: + allowed.add("reasoning") + unknown = set(values) - allowed + if unknown: + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", + f"unsupported model options: {', '.join(sorted(unknown))}", + ) + + result: dict[str, Any] = {} + for name, value in values.items(): + if name == "reasoning": + result[name] = _normalize_reasoning( + value, + supports_reasoning=supports_reasoning, + reasoning_mode=reasoning_mode, + allowed_reasoning_efforts=allowed_reasoning_efforts, + default_reasoning_effort=default_reasoning_effort, + ) + continue + option = dict(user_options[name]) + applies_to = set(option.get("applies_to") or ("main_agent",)) + if purpose not in applies_to: + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", + f"model option {name} cannot apply to {purpose}", + ) + _validate_option_value(name, value, option) + result[name] = value + + constraints = validate_parameter_constraints( + parameter_constraints, + allowed_names=set(user_options) + | ({"reasoning"} if supports_reasoning else set()), + ) + for constraint in constraints: + names = tuple(constraint["at_most_one_of"]) + if sum(name in result for name in names) > 1: + raise EvoRuntimeError( + "MODEL_PARAMETER_CONFLICT", + f"at most one of {', '.join(names)} may be supplied", + ) + return result + + +def _normalize_reasoning( + value: Any, + *, + supports_reasoning: bool, + reasoning_mode: str, + allowed_reasoning_efforts: Sequence[str], + default_reasoning_effort: str | None, +) -> str: + if not supports_reasoning: + raise EvoRuntimeError("MODEL_PARAMETER_INVALID", "reasoning is unsupported") + clean = str(value or "").strip().lower() + if clean == "disabled": + clean = "off" + if clean not in _REASONING_VALUES: + raise EvoRuntimeError("MODEL_PARAMETER_INVALID", "reasoning value is invalid") + if clean == "off": + return clean + if reasoning_mode == "boolean": + return "on" + if reasoning_mode != "effort": + raise EvoRuntimeError("MODEL_PARAMETER_INVALID", "reasoning is unsupported") + if clean == "on": + clean = str(default_reasoning_effort or "") + allowed = {str(item) for item in allowed_reasoning_efforts} + if clean == "medium" and clean not in allowed and "high" in allowed: + clean = "high" + if clean not in allowed: + raise EvoRuntimeError("MODEL_PARAMETER_INVALID", "reasoning effort is invalid") + return clean + + +def _validate_option_value(name: str, value: Any, option: Mapping[str, Any]) -> None: + kind = str(option.get("type") or "") + valid_type = ( + (kind == "boolean" and isinstance(value, bool)) + or ( + kind == "integer" and isinstance(value, int) and not isinstance(value, bool) + ) + or ( + kind == "number" + and isinstance(value, int | float) + and not isinstance(value, bool) + and math.isfinite(float(value)) + ) + or (kind in {"string", "enum"} and isinstance(value, str)) + or (kind == "object" and isinstance(value, Mapping)) + ) + if not valid_type: + raise EvoRuntimeError("MODEL_PARAMETER_INVALID", f"{name} has invalid type") + if "minimum" in option and value < option["minimum"]: + raise EvoRuntimeError("MODEL_PARAMETER_INVALID", f"{name} is below minimum") + if "minimum_exclusive" in option and value <= option["minimum_exclusive"]: + raise EvoRuntimeError("MODEL_PARAMETER_INVALID", f"{name} is below minimum") + if "maximum" in option and value > option["maximum"]: + raise EvoRuntimeError("MODEL_PARAMETER_INVALID", f"{name} exceeds maximum") + if "maximum_exclusive" in option and value >= option["maximum_exclusive"]: + raise EvoRuntimeError("MODEL_PARAMETER_INVALID", f"{name} exceeds maximum") + if option.get("choices") and value not in option["choices"]: + raise EvoRuntimeError( + "MODEL_PARAMETER_INVALID", f"{name} is not an allowed value" + ) diff --git a/EvoScientist/memory/agents/_factory.py b/EvoScientist/memory/agents/_factory.py index 8373929..b8a4fc6 100644 --- a/EvoScientist/memory/agents/_factory.py +++ b/EvoScientist/memory/agents/_factory.py @@ -98,6 +98,11 @@ def build_memory_agent_graph( workspace_dir=workspace_dir, memory_dir=memory_dir, ) + middleware = list(middleware) + if skills: + from ...middleware import BudgetedSkillsMiddleware + + middleware.append(BudgetedSkillsMiddleware(backend=backend, sources=skills)) agent = create_deep_agent( name=name, @@ -105,9 +110,9 @@ def build_memory_agent_graph( system_prompt=system_prompt, tools=list(tools), backend=backend, - middleware=list(middleware), + middleware=middleware, subagents=[], - skills=skills, + skills=None, **kwargs, ) return agent.with_config({"recursion_limit": recursion_limit}) diff --git a/EvoScientist/memory/agents/memory_worker.py b/EvoScientist/memory/agents/memory_worker.py index 5009d72..4f5c79b 100644 --- a/EvoScientist/memory/agents/memory_worker.py +++ b/EvoScientist/memory/agents/memory_worker.py @@ -428,8 +428,10 @@ def _memory_worker_middleware( enable_observation_memory: bool = True, ): """Build middleware for memory workers, excluding task execution tools.""" + from ...middleware.configurable_model import ConfigurableModelMiddleware from ...middleware.error_normalization import ErrorNormalizationMiddleware from ...middleware.memory import create_memory_middleware + from ...middleware.recoverable_metering import RecoverableMeteringMiddleware memory_controls = MemoryControls( profile_enabled=enable_profile_memory, @@ -444,6 +446,8 @@ def _memory_worker_middleware( # Outermost — normalize provider-SDK exceptions from the # auxiliary model call before any inner middleware sees them. ErrorNormalizationMiddleware(), + ConfigurableModelMiddleware(), + RecoverableMeteringMiddleware(), *memory_agent_middleware( create_memory_middleware( str(memory_dir), diff --git a/EvoScientist/memory/agents/observation_linker.py b/EvoScientist/memory/agents/observation_linker.py index 044158b..8cd13f2 100644 --- a/EvoScientist/memory/agents/observation_linker.py +++ b/EvoScientist/memory/agents/observation_linker.py @@ -71,7 +71,9 @@ def build_observation_linker_graph( workspace_dir: str | Path | None = None, ) -> CompiledStateGraph: """Build the registered LangGraph observation linker.""" + from ...middleware.configurable_model import ConfigurableModelMiddleware from ...middleware.error_normalization import ErrorNormalizationMiddleware + from ...middleware.recoverable_metering import RecoverableMeteringMiddleware agent_paths = resolve_memory_agent_paths( memory_dir=memory_dir, @@ -89,5 +91,10 @@ def build_observation_linker_graph( workspace_dir=agent_paths.workspace_dir, # Outermost — normalize provider-SDK exceptions from the # auxiliary model call before any inner middleware sees them. - middleware=[ErrorNormalizationMiddleware(), *memory_agent_middleware()], + middleware=[ + ErrorNormalizationMiddleware(), + ConfigurableModelMiddleware(), + RecoverableMeteringMiddleware(), + *memory_agent_middleware(), + ], ) diff --git a/EvoScientist/memory/launch.py b/EvoScientist/memory/launch.py index d15ddda..0ab3ace 100644 --- a/EvoScientist/memory/launch.py +++ b/EvoScientist/memory/launch.py @@ -3,7 +3,7 @@ from __future__ import annotations import json -from collections.abc import Callable +from collections.abc import Callable, Mapping from pathlib import Path from typing import cast @@ -42,6 +42,59 @@ OBSERVATION_LINKER_GRAPH_ID = "evomemory-observation-linker" MemoryWorkerFinishedHook = Callable[[BackgroundRun, MemoryOutputDelta | None], None] MemoryWorkerAbortedHook = Callable[[BackgroundRun, MemoryOutputDelta | None], None] +_INHERITED_RUNTIME_KEYS = ( + "model", + "model_provider", + "ai4sci_metering", + "ai4sci_model_proxy", +) + + +def _current_runtime_context() -> tuple[str | None, dict[str, object]]: + """Capture the signed parent Run context before launching a child graph.""" + try: + from langgraph.config import get_config + + config = get_config() + except Exception: + return None, {} + if not isinstance(config, Mapping): + return None, {} + configurable = config.get("configurable") + inherited = { + key: configurable[key] + for key in _INHERITED_RUNTIME_KEYS + if isinstance(configurable, Mapping) and key in configurable + } + metadata = config.get("metadata") + runtime_url = ( + str(metadata.get("langgraph_api_url") or "") + if isinstance(metadata, Mapping) + else "" + ) + return runtime_url or None, inherited + + +def _with_runtime_context( + payload: BackgroundRunPayload, + *, + inherited: Mapping[str, object] | None, + source_type: str, +) -> BackgroundRunPayload: + """Attach a child billing scope without changing the signed envelope.""" + normalized = cast("BackgroundRunPayload", dict(payload)) + config = dict(normalized.get("config") or {}) + configurable = dict(config.get("configurable") or {}) + configurable.update(dict(inherited or {})) + metering = configurable.get("ai4sci_metering") + if isinstance(metering, Mapping): + configurable["ai4sci_metering"] = { + **dict(metering), + "source_type": source_type, + } + config["configurable"] = configurable + return cast("BackgroundRunPayload", {**normalized, "config": config}) + def _observation_linking_enabled() -> bool: return MemoryControls.from_config(get_effective_config()).observations_enabled @@ -104,6 +157,7 @@ def _memory_worker_run_payload( *, context: MemorySourceContext, thread_id: str, + inherited: Mapping[str, object] | None = None, ) -> BackgroundRunPayload: """Build the LangGraph SDK run payload for a memory worker.""" metadata = _memory_worker_metadata(context) @@ -121,7 +175,17 @@ def _memory_worker_run_payload( } }, } - return _runs_create_kwargs(payload) + payload = _runs_create_kwargs(payload) + source_type = ( + "evomemory_turn_worker" + if context.source_type == MemorySourceType.TURN + else "evomemory_subagent_worker" + ) + return _with_runtime_context( + payload, + inherited=inherited, + source_type=source_type, + ) def memory_worker_launch_request( @@ -129,14 +193,20 @@ def memory_worker_launch_request( ) -> BackgroundRunRequest: """Build the background run request for a memory worker.""" metadata = _memory_worker_metadata(context) + runtime_url, inherited = _current_runtime_context() def run_payload(thread_id: str) -> BackgroundRunPayload: - return _memory_worker_run_payload(context=context, thread_id=thread_id) + return _memory_worker_run_payload( + context=context, + thread_id=thread_id, + inherited=inherited, + ) return BackgroundRunRequest( graph_id=_memory_worker_graph_id(context.source_type), run_payload=run_payload, thread_metadata=metadata, + url=runtime_url, name="EvoMemory worker", ) @@ -192,7 +262,12 @@ def _observation_linker_run_payload( } }, } - return _runs_create_kwargs(payload) + payload = _runs_create_kwargs(payload) + return _with_runtime_context( + payload, + inherited=context.runtime_configurable, + source_type="evomemory_linker", + ) def observation_linker_launch_request( @@ -210,6 +285,7 @@ def observation_linker_launch_request( graph_id=OBSERVATION_LINKER_GRAPH_ID, run_payload=run_payload, thread_metadata=_observation_linker_metadata(context), + url=context.runtime_url, name="EvoMemory observation linker", ) diff --git a/EvoScientist/memory/observations/__init__.py b/EvoScientist/memory/observations/__init__.py index 7b1d516..1f456a2 100644 --- a/EvoScientist/memory/observations/__init__.py +++ b/EvoScientist/memory/observations/__init__.py @@ -12,7 +12,7 @@ from .index import ( build_observation_index_context, build_observation_linker_index_context, ) -from .relations import link_observation_files +from .relations import archive_observation_file, link_observation_files from .store import ( OBSERVATION_DIR, ObservationFrontmatter, @@ -51,6 +51,7 @@ __all__ = [ "RecordObservationArgs", "RelatedObservationEntry", "SearchObservationsArgs", + "archive_observation_file", "build_observation_index_context", "build_observation_linker_index_context", "create_link_observations_tool", diff --git a/EvoScientist/memory/observations/relations.py b/EvoScientist/memory/observations/relations.py index 209a957..d86422a 100644 --- a/EvoScientist/memory/observations/relations.py +++ b/EvoScientist/memory/observations/relations.py @@ -2,21 +2,23 @@ from __future__ import annotations -import threading +import os +import posixpath from datetime import UTC, datetime from pathlib import Path +from filelock import FileLock + from ..types import ObservationRelation from .store import ( ObservationFrontmatter, RelatedObservationEntry, observation_document_by_id, + read_observation_document, related_observation_entries, write_observation_document, ) -_link_write_lock = threading.Lock() - def _relation_value(value: ObservationRelation | str) -> str: try: @@ -87,7 +89,8 @@ def link_observation_files( if not reason_text: raise ValueError("reason must not be empty") - with _link_write_lock: + relation_lock = Path(memory_dir).expanduser() / ".relation-write.lock" + with FileLock(str(relation_lock), timeout=30): source_document = observation_document_by_id( memory_dir=memory_dir, project_id=project_id, @@ -154,3 +157,84 @@ def link_observation_files( ], "missing_observation_ids": [], } + + +def archive_observation_file( + *, + memory_dir: str | Path, + observation_id: str, + observation_path: str, +) -> dict[str, object]: + """Archive one observation and remove references to it from live memory.""" + + requested_id = observation_id.strip() + raw_path = observation_path.strip().replace("\\", "/").lstrip("/") + if ".." in raw_path.split("/"): + raise ValueError("invalid observation archive target") + normalized = posixpath.normpath("/" + raw_path) + parts = Path(normalized.lstrip("/")).parts + if ( + not requested_id.startswith("O-") + or len(parts) < 3 + or parts[0] != "observations" + or parts[-1] != f"{requested_id}.md" + or ".." in parts + ): + raise ValueError("invalid observation archive target") + + root = Path(memory_dir).expanduser().resolve() + source = root.joinpath(*parts) + if source.is_symlink() or not source.is_file(): + raise FileNotFoundError(observation_path) + + observation_lock = FileLock(str(root / ".observation-write.lock"), timeout=30) + relation_lock = FileLock(str(root / ".relation-write.lock"), timeout=30) + with observation_lock, relation_lock: + document = read_observation_document(source) + if document is None or document[0].id != requested_id: + raise FileNotFoundError(observation_path) + target_metadata, _target_body = document + + archive_id = datetime.now(UTC).strftime("%Y%m%dT%H%M%S%fZ") + archived = root / "trash" / "observations" / archive_id / Path(*parts[1:]) + archived.parent.mkdir(parents=True, exist_ok=True) + os.replace(source, archived) + + observations_root = root / "observations" + if target_metadata.scope.value == "global": + candidates = observations_root.rglob("*.md") + else: + project_id = str(target_metadata.project_id or "") + candidates = iter( + [ + *(observations_root / "global").glob("*.md"), + *(observations_root / "projects" / project_id).glob("*.md"), + ] + ) + + updated_observation_ids: list[str] = [] + for candidate in sorted(candidates): + if candidate.is_symlink() or not candidate.is_file(): + continue + related_document = read_observation_document(candidate) + if related_document is None: + continue + metadata, body = related_document + retained = [ + entry + for entry in metadata.related_observations + if entry.id != requested_id + ] + if len(retained) == len(metadata.related_observations): + continue + metadata.related_observations = retained + write_observation_document(candidate, metadata=metadata, body=body) + updated_observation_ids.append(metadata.id) + + return { + "removed": True, + "observation_id": requested_id, + "archive_id": archive_id, + "archive_path": archived.relative_to(root).as_posix(), + "updated_observation_ids": updated_observation_ids, + } diff --git a/EvoScientist/memory/observations/store.py b/EvoScientist/memory/observations/store.py index 8f06f31..53538e7 100644 --- a/EvoScientist/memory/observations/store.py +++ b/EvoScientist/memory/observations/store.py @@ -9,11 +9,14 @@ from __future__ import annotations import hashlib import json +import os +import tempfile from dataclasses import replace from datetime import UTC, date, datetime from pathlib import Path import yaml +from filelock import FileLock from pydantic import BaseModel, ConfigDict, Field, ValidationError, field_validator from ..search import ( @@ -227,7 +230,25 @@ def write_observation_document( allow_unicode=True, sort_keys=False, ) - Path(path).write_text(f"---\n{frontmatter}---\n{body}", encoding="utf-8") + _atomic_write_text(Path(path), f"---\n{frontmatter}---\n{body}") + + +def _atomic_write_text(path: Path, content: str) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + fd, temporary = tempfile.mkstemp( + prefix=f".{path.name}.", suffix=".tmp", dir=path.parent + ) + try: + with os.fdopen(fd, "w", encoding="utf-8") as handle: + handle.write(content) + handle.flush() + os.fsync(handle.fileno()) + os.replace(temporary, path) + finally: + try: + os.unlink(temporary) + except FileNotFoundError: + pass def read_observation_id_from_path(path: str | Path) -> str | None: @@ -605,25 +626,26 @@ def record_observation_file( ) path = Path(memory_dir).expanduser() / memory_path.lstrip("/") created = False - if not path.exists(): - created_at = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ") - content = _format_observation_markdown( - observation_id=observation_id, - created_at=created_at, - memory_type=memory_type, - summary=summary_text, - observation=observation_text, - why_it_matters=why_text, - evidence=evidence.strip() if evidence else None, - scope=scope, - source_type=source_type, - source_agent=source_agent, - source_session_id=source_session_id, - project_id=project_id, - ) - path.parent.mkdir(parents=True, exist_ok=True) - path.write_text(content, encoding="utf-8") - created = True + memory_root = Path(memory_dir).expanduser() + with FileLock(str(memory_root / ".observation-write.lock"), timeout=30): + if not path.exists(): + created_at = datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ") + content = _format_observation_markdown( + observation_id=observation_id, + created_at=created_at, + memory_type=memory_type, + summary=summary_text, + observation=observation_text, + why_it_matters=why_text, + evidence=evidence.strip() if evidence else None, + scope=scope, + source_type=source_type, + source_agent=source_agent, + source_session_id=source_session_id, + project_id=project_id, + ) + _atomic_write_text(path, content) + created = True result: ObservationRecordResult = { "observation_id": observation_id, diff --git a/EvoScientist/memory/scheduler.py b/EvoScientist/memory/scheduler.py index 32720fd..e358b07 100644 --- a/EvoScientist/memory/scheduler.py +++ b/EvoScientist/memory/scheduler.py @@ -4,7 +4,7 @@ from __future__ import annotations import logging import threading -from collections.abc import Callable +from collections.abc import Callable, Mapping from dataclasses import dataclass from pathlib import Path from typing import NamedTuple @@ -29,6 +29,8 @@ class ObservationLinkerContext: workspace_dir: Path project_id: str observation_ids: tuple[str, ...] + runtime_url: str | None = None + runtime_configurable: Mapping[str, object] | None = None ObservationLinkerLauncher = Callable[[ObservationLinkerContext], BackgroundRun | None] @@ -77,6 +79,9 @@ class MemoryScheduler: self._launch_linker = launch_linker self._has_active_workers = has_active_workers self._pending: dict[_BatchKey, set[str]] = {} + self._runtime_context: dict[ + _BatchKey, tuple[str | None, Mapping[str, object] | None] + ] = {} self._lock = threading.Lock() def _launch_ready( @@ -121,6 +126,11 @@ class MemoryScheduler: ) with self._lock: self._pending.setdefault(key, set()).update(context.observation_ids) + if context.runtime_url or context.runtime_configurable: + self._runtime_context[key] = ( + context.runtime_url, + context.runtime_configurable, + ) def flush_ready(self) -> None: """Launch any pending linker batches that are no longer blocked.""" @@ -144,19 +154,39 @@ class MemoryScheduler: contexts = self._record_finished_and_drain_ready(run=run, delta=delta) self._launch_ready(contexts) - def _ready_batches_locked(self) -> list[tuple[_BatchKey, set[str]]]: + def _ready_batches_locked( + self, + ) -> list[ + tuple[ + _BatchKey, + set[str], + tuple[str | None, Mapping[str, object] | None], + ] + ]: ready_batches = [] for key in list(self._pending): if not self._has_active_workers(key.memory_dir): - ready_batches.append((key, self._pending.pop(key))) + ready_batches.append( + ( + key, + self._pending.pop(key), + self._runtime_context.pop(key, (None, None)), + ) + ) return ready_batches def _contexts_for_batches( self, - ready_batches: list[tuple[_BatchKey, set[str]]], + ready_batches: list[ + tuple[ + _BatchKey, + set[str], + tuple[str | None, Mapping[str, object] | None], + ] + ], ) -> tuple[ObservationLinkerContext, ...]: ready_contexts = [] - for key, observation_ids in ready_batches: + for key, observation_ids, runtime_context in ready_batches: if not observation_ids: continue ready_contexts.append( @@ -165,6 +195,8 @@ class MemoryScheduler: workspace_dir=Path(key.workspace_dir), project_id=key.project_id, observation_ids=tuple(sorted(observation_ids)), + runtime_url=runtime_context[0], + runtime_configurable=runtime_context[1], ) ) @@ -194,6 +226,10 @@ class MemoryScheduler: with self._lock: if key is not None and observation_ids: self._pending.setdefault(key, set()).update(observation_ids) + if isinstance(run.configurable, Mapping) and isinstance( + run.configurable.get("ai4sci_metering"), Mapping + ): + self._runtime_context[key] = (run.url, run.configurable) ready_batches = self._ready_batches_locked() diff --git a/EvoScientist/memory/search.py b/EvoScientist/memory/search.py index 5955e87..8f5b3c6 100644 --- a/EvoScientist/memory/search.py +++ b/EvoScientist/memory/search.py @@ -23,6 +23,7 @@ DEFAULT_MATCH_LINES = 3 DEFAULT_MATCH_CHARS = 240 _TOKEN_RE = re.compile(r"[a-z0-9_]+") +_CJK_RUN_RE = re.compile(r"[\u3400-\u4dbf\u4e00-\u9fff\uf900-\ufaff]+") def _compile_query_pattern(query: str) -> re.Pattern[str]: @@ -35,11 +36,20 @@ def _compile_query_pattern(query: str) -> re.Pattern[str]: def _tokens(text: str) -> list[str]: """Return simple lowercase search tokens.""" - return [ + normalized = text.casefold() + tokens = [ token - for token in _TOKEN_RE.findall(text.casefold()) + for token in _TOKEN_RE.findall(normalized) if len(token) >= MIN_TOKEN_CHARS ] + for run in _CJK_RUN_RE.findall(normalized): + tokens.append(run) + for size in (2, 3): + tokens.extend( + run[index : index + size] + for index in range(max(0, len(run) - size + 1)) + ) + return tokens def _document_tokens(document: ObservationSearchDocument) -> set[str]: diff --git a/EvoScientist/middleware/__init__.py b/EvoScientist/middleware/__init__.py index d79e0ed..d2aa82e 100644 --- a/EvoScientist/middleware/__init__.py +++ b/EvoScientist/middleware/__init__.py @@ -18,6 +18,7 @@ from .context_editing import ( create_context_editing_middleware, ) from .context_overflow import ContextOverflowMapperMiddleware +from .disable_subagent import DisableSubagentToolMiddleware from .error_normalization import ErrorNormalizationMiddleware from .memory import ( EvoMemoryMiddleware, @@ -29,6 +30,9 @@ from .memory_lifecycle import ( default_memory_scheduler, ) from .model_fallback import ModelFallbackMiddleware, load_fallback_chain +from .provider_context import ProviderContextMediaMiddleware +from .recoverable_metering import RecoverableMeteringMiddleware +from .recoverable_tools import RecoverableToolEffectMiddleware from .repetitive_tool_guard import ( DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS, DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD, @@ -40,6 +44,12 @@ from .scheduler import ( SchedulerMiddleware, create_scheduler_middleware, ) +from .skill_context import ( + DEFAULT_MAX_DESCRIPTION_BYTES, + DEFAULT_MAX_SKILLS, + DEFAULT_MAX_SKILLS_BYTES, + BudgetedSkillsMiddleware, +) from .tool_error_handler import ToolErrorHandlerMiddleware from .tool_protocol_guard import ToolProtocolGuardMiddleware from .tool_selector import create_tool_selector_middleware @@ -47,18 +57,26 @@ from .utils import disable_thinking __all__ = [ "DEFAULT_MAX_CONSECUTIVE_TOOL_ERRORS", + "DEFAULT_MAX_DESCRIPTION_BYTES", + "DEFAULT_MAX_SKILLS", + "DEFAULT_MAX_SKILLS_BYTES", "DEFAULT_REPETITIVE_TOOL_CALL_THRESHOLD", "AskUserMiddleware", "AskUserRequest", "AskUserWidgetResult", + "BudgetedSkillsMiddleware", "Choice", "ConfigurableModelMiddleware", "ContextOverflowMapperMiddleware", + "DisableSubagentToolMiddleware", "ErrorNormalizationMiddleware", "EvoMemoryLifecycleMiddleware", "EvoMemoryMiddleware", "ModelFallbackMiddleware", + "ProviderContextMediaMiddleware", "Question", + "RecoverableMeteringMiddleware", + "RecoverableToolEffectMiddleware", "RepetitiveToolCallGuardMiddleware", "RuntimeContextMiddleware", "SchedulerMiddleware", diff --git a/EvoScientist/middleware/configurable_model.py b/EvoScientist/middleware/configurable_model.py index 4c65f4e..407c208 100644 --- a/EvoScientist/middleware/configurable_model.py +++ b/EvoScientist/middleware/configurable_model.py @@ -74,6 +74,28 @@ def _read_model_override() -> tuple[str | None, str | None]: ) +def _read_proxy_override() -> tuple[dict[str, Any] | None, str, str]: + try: + from langgraph.config import get_config + + cfg = get_config() + except Exception: + return None, "", "" + configurable = cfg.get("configurable") if isinstance(cfg, dict) else None + if not isinstance(configurable, dict): + return None, "", "" + proxy = configurable.get("ai4sci_model_proxy") + metering = configurable.get("ai4sci_metering") + if not isinstance(proxy, dict): + return None, "", "" + metering = metering if isinstance(metering, dict) else {} + return ( + dict(proxy), + str(metering.get("provider_id") or ""), + str(metering.get("model_id") or ""), + ) + + class ConfigurableModelMiddleware(AgentMiddleware): """Re-resolve the chat model from RunnableConfig.configurable on every call. @@ -150,6 +172,17 @@ class ConfigurableModelMiddleware(AgentMiddleware): request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse], ) -> ModelResponse: + proxy, provider_id, model_id = _read_proxy_override() + if proxy is not None: + from ..llm.gateway_proxy import proxy_from_config + + return handler( + request.override( + model=proxy_from_config( + proxy, provider_id=provider_id, model_id=model_id + ) + ) + ) model_name, provider = _read_model_override() if model_name is None: return handler(request) @@ -172,6 +205,17 @@ class ConfigurableModelMiddleware(AgentMiddleware): request: ModelRequest, handler: Callable[[ModelRequest], Awaitable[ModelResponse]], ) -> ModelResponse: + proxy, provider_id, model_id = _read_proxy_override() + if proxy is not None: + from ..llm.gateway_proxy import proxy_from_config + + return await handler( + request.override( + model=proxy_from_config( + proxy, provider_id=provider_id, model_id=model_id + ) + ) + ) model_name, provider = _read_model_override() if model_name is None: return await handler(request) diff --git a/EvoScientist/middleware/context_overflow.py b/EvoScientist/middleware/context_overflow.py index 1976c15..04aa59b 100644 --- a/EvoScientist/middleware/context_overflow.py +++ b/EvoScientist/middleware/context_overflow.py @@ -67,6 +67,8 @@ class ContextOverflowMapperMiddleware(AgentMiddleware): It triggers when there's an error 400 raised and one of specified patterns exists in the error message. """ + if getattr(exc, "code", None) == "MODEL_CONTEXT_WINDOW_EXCEEDED": + return True err_msg = str(exc).lower() patterns = [ diff --git a/EvoScientist/middleware/disable_subagent.py b/EvoScientist/middleware/disable_subagent.py new file mode 100644 index 0000000..ee6aead --- /dev/null +++ b/EvoScientist/middleware/disable_subagent.py @@ -0,0 +1,77 @@ +"""Web profile guard that makes DeepAgents subagent delegation unavailable.""" + +from __future__ import annotations + +from collections.abc import Awaitable, Callable +from typing import Any + +from langchain.agents.middleware.types import ( + AgentMiddleware, + ModelRequest, + ModelResponse, + ToolCallRequest, +) +from langchain_core.messages import ToolMessage +from langgraph.types import Command + + +def _tool_name(tool: Any) -> str: + if isinstance(tool, dict): + return str(tool.get("name") or tool.get("function", {}).get("name") or "") + return str(getattr(tool, "name", "") or "") + + +class DisableSubagentToolMiddleware(AgentMiddleware): + """Hide and reject the built-in ``task`` tool for the Web profile. + + DeepAgents adds a general-purpose task tool even when an empty subagent + list is supplied. Filtering the model request alone is therefore not a + sufficient safety boundary; the tool-call guard protects restored or + malformed checkpoints as well. + """ + + name = "disable_subagent_tool" + + def wrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], ModelResponse], + ) -> ModelResponse: + tools = [tool for tool in request.tools if _tool_name(tool) != "task"] + return handler(request.override(tools=tools)) + + async def awrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], Awaitable[ModelResponse]], + ) -> ModelResponse: + tools = [tool for tool in request.tools if _tool_name(tool) != "task"] + return await handler(request.override(tools=tools)) + + def wrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]], + ) -> ToolMessage | Command[Any]: + if str(request.tool_call.get("name") or "") == "task": + return self._rejected(request) + return handler(request) + + async def awrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command[Any]]], + ) -> ToolMessage | Command[Any]: + if str(request.tool_call.get("name") or "") == "task": + return self._rejected(request) + return await handler(request) + + @staticmethod + def _rejected(request: ToolCallRequest) -> ToolMessage: + return ToolMessage( + content="SUBAGENTS_DISABLED", + tool_call_id=str(request.tool_call.get("id") or "subagents_disabled"), + name="task", + status="error", + ) + diff --git a/EvoScientist/middleware/error_normalization.py b/EvoScientist/middleware/error_normalization.py index 992445c..8dbe6ab 100644 --- a/EvoScientist/middleware/error_normalization.py +++ b/EvoScientist/middleware/error_normalization.py @@ -152,7 +152,6 @@ def _normalize(request: ModelRequest, exc: BaseException) -> ProviderStreamError _extract_provider_code, _extract_status_code, _provider_from_model, - _redact_api_keys, ) # Already normalized (e.g. by ModelFallbackMiddleware wrapping against @@ -191,11 +190,24 @@ def _normalize(request: ModelRequest, exc: BaseException) -> ProviderStreamError else None ) + status_code = _extract_status_code(exc) + safe_message = { + 400: "Provider rejected the request.", + 401: "Provider authentication failed.", + 403: "Provider authorization failed.", + 404: "Provider model or endpoint was not found.", + 408: "Provider request timed out.", + 429: "Provider rate limit was exceeded.", + 500: "Provider request failed.", + 502: "Provider gateway failed.", + 503: "Provider is temporarily unavailable.", + 504: "Provider gateway timed out.", + }.get(status_code, "Provider request failed.") return ProviderStreamError( provider=provider, class_qualname=class_qualname, - message=_redact_api_keys(str(exc)), - status_code=_extract_status_code(exc), + message=safe_message, + status_code=status_code, code=_extract_provider_code(exc), err_type=_extract_error_type(exc), request_id=request_id, diff --git a/EvoScientist/middleware/evo_route_fallback.py b/EvoScientist/middleware/evo_route_fallback.py new file mode 100644 index 0000000..d4d1145 --- /dev/null +++ b/EvoScientist/middleware/evo_route_fallback.py @@ -0,0 +1,455 @@ +"""Evo-owned fallback middleware for frozen Web model-route snapshots.""" + +from __future__ import annotations + +import asyncio +import logging +from collections.abc import Awaitable, Callable, Iterable +from typing import Any + +from langchain.agents.middleware.types import ( + AgentMiddleware, + ModelRequest, + ModelResponse, +) +from langchain_core.messages import AIMessage, SystemMessage, ToolMessage + +from ..llm.contracts import EvoRuntimeError +from ..llm.errors import ModelProviderResponseError, ModelToolProtocolError +from ..llm.invocation import ( + assistant_message_has_output, + project_provider_messages, +) + +logger = logging.getLogger(__name__) +_TOOL_RESULT_CONTEXT_LIMIT = 12_000 +_TOOL_PROTOCOL_METADATA = frozenset( + {"tool_calls", "tool_call_chunks", "function_call", "tool_call_id", "call_id"} +) + + +class EvoRouteFallbackMiddleware(AgentMiddleware): + """Retry only configured, billing-equivalent Evo fallback route models.""" + + name = "evo_route_fallback" + + def __init__( + self, + fallback_models: Iterable[Any], + route_health: Any = None, + capacity: Any = None, + ) -> None: + super().__init__() + self._fallback_models = tuple(fallback_models) + self._route_health = route_health + self._capacity = capacity + + def wrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], ModelResponse], + ) -> ModelResponse: + if not self._available(request.model): + return self._try_sync( + request, + handler, + EvoRuntimeError("MODEL_ROUTE_UNAVAILABLE"), + force_fallback=True, + ) + try: + return self._invoke_sync( + self._request_for_model(request, request.model), handler + ) + except Exception as primary_error: + return self._try_sync(request, handler, primary_error) + + async def awrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], Awaitable[ModelResponse]], + ) -> ModelResponse: + if not self._available(request.model): + return await self._try_async_fallbacks( + request, + handler, + EvoRuntimeError("MODEL_ROUTE_UNAVAILABLE"), + ) + try: + return await self._invoke_async( + self._request_for_model(request, request.model), handler + ) + except Exception as primary_error: + if not _fallbackable(primary_error, request.model): + raise + return await self._try_async_fallbacks(request, handler, primary_error) + + async def _try_async_fallbacks( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], Awaitable[ModelResponse]], + primary_error: Exception, + ) -> ModelResponse: + last_error = primary_error + for model in self._retry_models(request, primary_error): + if not self._available(model): + continue + try: + return await self._invoke_async( + self._request_for_model( + request, + model, + repair_error=last_error, + ), + handler, + ) + except Exception as error: + if not _fallbackable(error, model): + raise + last_error = error + raise last_error from primary_error + + async def _invoke_async( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], Awaitable[ModelResponse]], + ) -> ModelResponse: + lease = await self._capacity.acquire(request.model) if self._capacity else () + try: + metadata = getattr(request.model, "metadata", None) or {} + async with asyncio.timeout( + int(metadata.get("attempt_timeout_seconds") or 600) + ): + response = await handler(request) + _require_valid_response(response, request) + return response + finally: + if self._capacity: + self._capacity.release(lease) + + def _try_sync( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], ModelResponse], + primary_error: Exception, + *, + force_fallback: bool = False, + ) -> ModelResponse: + if not force_fallback and not _fallbackable(primary_error, request.model): + raise primary_error + last_error = primary_error + models = ( + self._fallback_models + if force_fallback + else self._retry_models(request, primary_error) + ) + for model in models: + if not self._available(model): + continue + try: + return self._invoke_sync( + self._request_for_model( + request, + model, + repair_error=last_error, + ), + handler, + ) + except Exception as error: + if not _fallbackable(error, model): + raise + last_error = error + raise last_error + + @staticmethod + def _invoke_sync( + request: ModelRequest, + handler: Callable[[ModelRequest], ModelResponse], + ) -> ModelResponse: + response = handler(request) + _require_valid_response(response, request) + return response + + def _retry_models( + self, request: ModelRequest, error: Exception + ) -> tuple[Any, ...]: + candidates = ( + (request.model, *self._fallback_models) + if isinstance(error, ModelProviderResponseError) + else self._fallback_models + ) + result: list[Any] = [] + for model in candidates: + if all(model is not existing for existing in result): + result.append(model) + return tuple(result) + + def _available(self, model: Any) -> bool: + if self._route_health is None: + return True + metadata = getattr(model, "metadata", None) or {} + route_key = str(metadata.get("route_key") or "") + return not route_key or not self._route_health.is_open(route_key) + + @staticmethod + def _request_for_model( + request: ModelRequest, + model: Any, + *, + repair_error: Exception | None = None, + ) -> ModelRequest: + metadata = getattr(model, "metadata", None) or {} + supports_tools = metadata.get("route_supports_tools") + tools = [] if supports_tools is False else request.tools + messages = ( + _without_tool_protocol(request.messages) + if supports_tools is False + else request.messages + ) + messages, dropped_empty_assistants = project_provider_messages(messages) + if dropped_empty_assistants: + logger.info( + "model_input_repaired route=%s dropped_empty_assistant=%s", + str(metadata.get("route_key") or ""), + dropped_empty_assistants, + ) + if isinstance(repair_error, ModelToolProtocolError) and tools: + tool_names = sorted( + { + name + for tool in tools + if (name := _tool_name(tool)) is not None + } + ) + allowed = ", ".join(tool_names[:64]) + if len(tool_names) > 64: + allowed += ", ..." + repair = SystemMessage( + content=( + "Retry the previous response because its structured tool call " + f"was invalid ({repair_error.reason}). If a tool is needed, use " + "exactly one of the supplied tool names, include a non-empty call " + "ID, and emit arguments as one valid JSON object. Otherwise answer " + "normally." + + (f" Supplied tool names: {allowed}." if allowed else "") + ) + ) + leading_system_messages = 0 + for message in messages: + if getattr(message, "type", "") not in {"system", "developer"}: + break + leading_system_messages += 1 + messages = [ + *messages[:leading_system_messages], + repair, + *messages[leading_system_messages:], + ] + elif isinstance(repair_error, ModelProviderResponseError): + logger.warning( + "model_response_repair_retry route=%s model=%s api_mode=%s " + "tool_transport=%s reason=%s", + str(metadata.get("route_key") or ""), + str(metadata.get("route_model") or ""), + str(metadata.get("route_api_mode") or ""), + str(metadata.get("route_tool_call_transport") or ""), + repair_error.reason, + ) + repair = SystemMessage( + content=( + "Retry the response because the previous attempt completed " + "without final text or a structured tool call. Complete the " + "request with a user-visible final answer, or emit a valid " + "tool call when a tool is required." + ) + ) + leading_system_messages = 0 + for message in messages: + if getattr(message, "type", "") not in {"system", "developer"}: + break + leading_system_messages += 1 + messages = [ + *messages[:leading_system_messages], + repair, + *messages[leading_system_messages:], + ] + overrides: dict[str, Any] = { + "model": model, + "tools": tools, + "messages": messages, + # The invocation plan has already compiled every provider SDK + # parameter. Agent-level settings must not mutate it afterwards. + "model_settings": {}, + } + if supports_tools is False: + # A no-tools route cannot accept an inherited forced tool choice or + # structured-output contract from a prior model invocation. + overrides["tool_choice"] = None + overrides["response_format"] = None + return request.override(**overrides) + + +def _without_tool_protocol(messages: list[Any]) -> list[Any]: + """Make a checkpoint replayable by a route that does not support tools. + + A model switch can leave completed ToolMessage/AI tool-call pairs in the + checkpointer. Sending those pairs while omitting ``tools`` is rejected by + strict OpenAI-compatible providers. Keep normal assistant text, remove + protocol-only fields, and retain a bounded plain-text transcript of tool + results so a text-only model does not lose work completed before a switch. + """ + + sanitized: list[Any] = [] + for message in messages: + if isinstance(message, ToolMessage) or getattr(message, "type", "") == "tool": + transcript = _tool_result_transcript(message) + if transcript: + sanitized.append(AIMessage(content=transcript)) + continue + if not isinstance(message, AIMessage): + sanitized.append(message) + continue + additional = dict(getattr(message, "additional_kwargs", {}) or {}) + has_tool_protocol = bool( + getattr(message, "tool_calls", None) + or getattr(message, "invalid_tool_calls", None) + or additional.get("tool_calls") + or _contains_tool_content(getattr(message, "content", "")) + ) + if not has_tool_protocol: + sanitized.append(message) + continue + portable_content = _portable_assistant_content(getattr(message, "content", "")) + # An empty assistant turn only represented a tool call. Its following + # ToolMessage becomes a transcript entry, so there is nothing useful to + # send for this message itself. + if not portable_content: + continue + for key in _TOOL_PROTOCOL_METADATA: + additional.pop(key, None) + sanitized.append( + message.model_copy( + update={ + "content": portable_content, + "tool_calls": [], + "invalid_tool_calls": [], + "additional_kwargs": additional, + } + ) + ) + return sanitized + + +def _portable_assistant_content(content: Any) -> str: + """Keep textual assistant content while removing embedded tool-call blocks.""" + + if isinstance(content, str): + return content + if not isinstance(content, list): + return str(content or "") + parts: list[str] = [] + for block in content: + if isinstance(block, str): + parts.append(block) + continue + if not isinstance(block, dict): + continue + block_type = str(block.get("type") or "") + if block_type in {"tool_call", "tool_use", "function_call"}: + continue + text = block.get("text") + if isinstance(text, str): + parts.append(text) + return "\n".join(part for part in parts if part) + + +def _contains_tool_content(content: Any) -> bool: + """Identify provider content blocks that encode a tool call.""" + + return isinstance(content, list) and any( + isinstance(block, dict) + and str(block.get("type") or "") + in {"tool_call", "tool_use", "function_call"} + for block in content + ) + + +def _tool_result_transcript(message: Any) -> str: + """Project a completed tool result to bounded text for a no-tools route.""" + + content = getattr(message, "content", "") + if isinstance(content, list): + parts = [] + for block in content: + if isinstance(block, str): + parts.append(block) + elif isinstance(block, dict) and isinstance(block.get("text"), str): + parts.append(block["text"]) + text = "\n".join(parts) + else: + text = str(content or "") + text = text.strip() + if not text: + return "" + if len(text) > _TOOL_RESULT_CONTEXT_LIMIT: + text = text[:_TOOL_RESULT_CONTEXT_LIMIT] + "\n[Tool result truncated]" + name = str(getattr(message, "name", "") or "tool")[:128] + return f"[Completed tool result: {name}]\n{text}" + + +def _tool_name(tool: Any) -> str | None: + if isinstance(tool, dict): + value = tool.get("name") + function = tool.get("function") + if not value and isinstance(function, dict): + value = function.get("name") + else: + value = getattr(tool, "name", None) + return value.strip() if isinstance(value, str) and value.strip() else None + + +def _fallbackable(error: Exception, model: Any) -> bool: + if isinstance(error, EvoRuntimeError): + return False + explicit = getattr(error, "fallbackable", None) + if explicit is not None: + return bool(explicit) + if getattr(error, "non_fallbackable", False): + return False + metadata = getattr(model, "metadata", None) or {} + adapter_id = str(metadata.get("route_adapter_id") or "") + adapter_revision = str(metadata.get("route_adapter_revision") or "") + if not adapter_id or not adapter_revision: + return isinstance(error, (ConnectionError, TimeoutError)) + from ..llm.adapter_registry import get_adapter_registry + + return ( + get_adapter_registry() + .get(adapter_id, adapter_revision) + .classify_error(error) + .retryable + ) + + +def _require_valid_response(response: Any, request: ModelRequest) -> None: + """Reject completed assistant responses that cannot advance the agent.""" + + messages: list[AIMessage] + if isinstance(response, AIMessage): + messages = [response] + else: + result = getattr(response, "result", None) + messages = ( + [message for message in result if isinstance(message, AIMessage)] + if isinstance(result, list | tuple) + else [] + ) + if messages and not any(assistant_message_has_output(message) for message in messages): + metadata = getattr(request.model, "metadata", None) or {} + logger.warning( + "model_provider_response_invalid route=%s model=%s api_mode=%s " + "tool_transport=%s reason=empty_assistant_response", + str(metadata.get("route_key") or ""), + str(metadata.get("route_model") or ""), + str(metadata.get("route_api_mode") or ""), + str(metadata.get("route_tool_call_transport") or ""), + ) + raise ModelProviderResponseError() diff --git a/EvoScientist/middleware/model_fallback.py b/EvoScientist/middleware/model_fallback.py index 42c2222..5371dbf 100644 --- a/EvoScientist/middleware/model_fallback.py +++ b/EvoScientist/middleware/model_fallback.py @@ -14,6 +14,7 @@ can deal with them. from __future__ import annotations import logging +import re import threading from collections.abc import Awaitable, Callable @@ -63,6 +64,20 @@ _AUTH_ERROR_PATTERNS: list[str] = [ These are intentionally *not* treated as non-fallbackable because a different provider in the chain may have valid credentials.""" +_SAFE_ERROR_CODE = re.compile(r"^[A-Z][A-Z0-9_]{1,63}$") + + +def _safe_error_label(exc: BaseException) -> str: + """Describe an error without rendering a provider-controlled response body.""" + label = type(exc).__name__ + status = getattr(exc, "status_code", None) + if isinstance(status, int) and 100 <= status <= 599: + label = f"{label} status={status}" + code = getattr(exc, "code", None) + if isinstance(code, str) and _SAFE_ERROR_CODE.fullmatch(code): + label = f"{label} code={code}" + return label + def set_ui_emit(fn: Callable[[str, str], None] | None) -> None: """Register (or clear) the UI callback for fallback status messages. @@ -260,13 +275,9 @@ async def _try_fallbacks( """ 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 - ) + primary_label = _safe_error_label(primary_exc) + _emit(f"Primary model failed: {primary_label}", style="yellow") + logger.warning("Primary model failed: %s", primary_label) # Track the request whose model actually raised ``last_exc`` so we # can attribute the exception to the failing model, not the @@ -280,7 +291,7 @@ async def _try_fallbacks( 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}", + f"due to: {_safe_error_label(last_exc)}", style="yellow", ) try: @@ -305,15 +316,14 @@ async def _try_fallbacks( last_exc = fb_exc last_failing_request = fb_request _emit( - f" x {model_name} also failed: {type(fb_exc).__name__}: {fb_exc}", + f" x {model_name} also failed: {_safe_error_label(fb_exc)}", style="red", ) logger.warning( - "Fallback %s (provider=%s) failed: %s: %s", + "Fallback %s (provider=%s) failed: %s", model_name, provider, - type(fb_exc).__name__, - fb_exc, + _safe_error_label(fb_exc), ) _emit(" All fallbacks exhausted -- re-raising last error", style="red") diff --git a/EvoScientist/middleware/provider_context.py b/EvoScientist/middleware/provider_context.py new file mode 100644 index 0000000..0e39106 --- /dev/null +++ b/EvoScientist/middleware/provider_context.py @@ -0,0 +1,376 @@ +"""Bound provider context by externalizing assistant-generated inline media.""" + +from __future__ import annotations + +import base64 +import hashlib +import logging +import mimetypes +from collections.abc import Awaitable, Callable, Mapping, Sequence +from dataclasses import replace +from typing import Any + +from langchain.agents.middleware.types import ( + AgentMiddleware, + ExtendedModelResponse, + ModelRequest, + ModelResponse, +) +from langchain_core.messages import AIMessage, BaseMessage +from langgraph.types import Overwrite + +from ..llm.contracts import EvoRuntimeError + +logger = logging.getLogger(__name__) + +_DEFAULT_MAX_INLINE_MEDIA_BYTES = 16_777_216 +_MEDIA_PREFIX = "/artifacts/model-output" + + +def _decode_base64_block(block: Mapping[str, Any]) -> tuple[bytes, str] | None: + payload = block.get("base64") + mime = str(block.get("mime_type") or "application/octet-stream") + if not isinstance(payload, str) or not payload: + return None + try: + return base64.b64decode(payload, validate=True), mime + except (ValueError, TypeError) as exc: + raise EvoRuntimeError( + "MEDIA_PERSIST_FAILED", + details=({"reason": "invalid_base64", "media_type": mime},), + ) from exc + + +def _extension(mime: str) -> str: + return (mimetypes.guess_extension(mime) or ".bin").lstrip(".") + + +def _artifact_reference(path: str, mime: str, digest: str) -> dict[str, Any]: + major = mime.split("/", 1)[0] + if major in {"image", "audio", "video"}: + return { + "type": major, + "url": path, + "mime_type": mime, + "media_id": f"sha256:{digest}", + } + return { + "type": "text", + "text": f'', + } + + +def _provider_reference(block: Mapping[str, Any]) -> dict[str, str] | None: + block_type = str(block.get("type") or "") + path = block.get("url") + if ( + block_type not in {"image", "audio", "video"} + or not isinstance(path, str) + or not path.startswith(f"{_MEDIA_PREFIX}/") + ): + return None + mime = str(block.get("mime_type") or f"{block_type}/unknown") + media_id = str(block.get("media_id") or "") + media_id_attribute = f' media_id="{media_id}"' if media_id else "" + return { + "type": "text", + "text": ( + f'" + ), + } + + +def _assistant_messages(messages: Sequence[BaseMessage]) -> list[AIMessage]: + return [message for message in messages if isinstance(message, AIMessage)] + + +def _collect_inline_media( + messages: Sequence[BaseMessage], + *, + max_inline_media_bytes: int, +) -> dict[str, tuple[str, bytes]]: + collected: dict[str, tuple[str, bytes]] = {} + for message in _assistant_messages(messages): + content = message.content + if not isinstance(content, list): + continue + for value in content: + if not isinstance(value, Mapping): + continue + decoded = _decode_base64_block(value) + if decoded is None: + continue + raw, mime = decoded + if len(raw) > max_inline_media_bytes: + raise EvoRuntimeError( + "MEDIA_PERSIST_FAILED", + details=( + { + "reason": "media_too_large", + "media_type": mime, + "media_bytes": len(raw), + }, + ), + ) + digest = hashlib.sha256(raw).hexdigest() + collected.setdefault(digest, (mime, raw)) + return collected + + +def _rewrite_messages( + messages: Sequence[BaseMessage], + paths: Mapping[str, tuple[str, str]], + *, + for_provider: bool, +) -> list[BaseMessage]: + rewritten: list[BaseMessage] = [] + for message in messages: + if not isinstance(message, AIMessage) or not isinstance(message.content, list): + rewritten.append(message) + continue + modified = False + content: list[Any] = [] + for value in message.content: + if not isinstance(value, Mapping): + content.append(value) + continue + if for_provider: + reference = _provider_reference(value) + if reference is not None: + content.append(reference) + modified = True + continue + decoded = _decode_base64_block(value) + if decoded is None: + content.append(value) + continue + raw, mime = decoded + digest = hashlib.sha256(raw).hexdigest() + path_entry = paths.get(digest) + if path_entry is None: + raise EvoRuntimeError( + "MEDIA_PERSIST_FAILED", + details=({"reason": "artifact_path_missing", "media_type": mime},), + ) + path, stored_mime = path_entry + reference = _artifact_reference(path, stored_mime, digest) + content.append( + _provider_reference(reference) if for_provider else reference + ) + modified = True + if modified: + copy = message.model_copy() + copy.content = content + rewritten.append(copy) + else: + rewritten.append(message) + return rewritten + + +def _response_messages( + response: Any, +) -> tuple[list[BaseMessage], Callable[[list[BaseMessage]], Any]]: + if isinstance(response, ExtendedModelResponse): + return response.model_response.result, lambda result: replace( + response, + model_response=replace(response.model_response, result=result), + ) + if isinstance(response, ModelResponse): + return response.result, lambda result: replace(response, result=result) + if isinstance(response, AIMessage): + return [response], lambda result: result[0] + return [], lambda _result: response + + +class ProviderContextMediaMiddleware(AgentMiddleware): + """Persist assistant media and keep base64 out of later provider calls.""" + + name = "provider_context_media" + + def __init__( + self, + backend: Any, + *, + max_inline_media_bytes: int = _DEFAULT_MAX_INLINE_MEDIA_BYTES, + ) -> None: + self.backend = backend + self.max_inline_media_bytes = max(1, int(max_inline_media_bytes)) + + @staticmethod + def _paths_for( + media: Mapping[str, tuple[str, bytes]], + ) -> dict[str, tuple[str, str]]: + return { + digest: (f"{_MEDIA_PREFIX}/{digest[:24]}.{_extension(mime)}", mime) + for digest, (mime, _raw) in media.items() + } + + def _persist(self, messages: Sequence[BaseMessage]) -> dict[str, tuple[str, str]]: + media = _collect_inline_media( + messages, + max_inline_media_bytes=self.max_inline_media_bytes, + ) + paths = self._paths_for(media) + for digest, (mime, raw) in media.items(): + path = paths[digest][0] + responses = self.backend.upload_files([(path, raw)]) + error = ( + getattr(responses[0], "error", None) + if responses + else "missing upload response" + ) + if error: + raise EvoRuntimeError( + "MEDIA_PERSIST_FAILED", + details=({"reason": "artifact_upload_failed", "media_type": mime},), + ) + logger.info( + "provider_context_media_persisted path=%s media_type=%s bytes=%s", + path, + mime, + len(raw), + ) + return paths + + async def _apersist( + self, messages: Sequence[BaseMessage] + ) -> dict[str, tuple[str, str]]: + media = _collect_inline_media( + messages, + max_inline_media_bytes=self.max_inline_media_bytes, + ) + paths = self._paths_for(media) + for digest, (mime, raw) in media.items(): + path = paths[digest][0] + responses = await self.backend.aupload_files([(path, raw)]) + error = ( + getattr(responses[0], "error", None) + if responses + else "missing upload response" + ) + if error: + raise EvoRuntimeError( + "MEDIA_PERSIST_FAILED", + details=({"reason": "artifact_upload_failed", "media_type": mime},), + ) + logger.info( + "provider_context_media_persisted path=%s media_type=%s bytes=%s", + path, + mime, + len(raw), + ) + return paths + + def _prepare_request(self, request: ModelRequest) -> ModelRequest: + paths = self._persist(request.messages) + messages = _rewrite_messages( + request.messages, + paths, + for_provider=True, + ) + return request.override(messages=messages) + + async def _aprepare_request(self, request: ModelRequest) -> ModelRequest: + paths = await self._apersist(request.messages) + messages = _rewrite_messages( + request.messages, + paths, + for_provider=True, + ) + return request.override(messages=messages) + + def _prepare_response(self, response: Any) -> Any: + messages, rebuild = _response_messages(response) + if not messages: + return response + paths = self._persist(messages) + return rebuild(_rewrite_messages(messages, paths, for_provider=False)) + + async def _aprepare_response(self, response: Any) -> Any: + messages, rebuild = _response_messages(response) + if not messages: + return response + paths = await self._apersist(messages) + return rebuild(_rewrite_messages(messages, paths, for_provider=False)) + + def before_model(self, state: Any, runtime: Any) -> dict[str, Any] | None: + _ = runtime + try: + messages = state.get("messages") if isinstance(state, Mapping) else None + if not isinstance(messages, Sequence) or isinstance(messages, str | bytes): + return None + original = list(messages) + paths = self._persist(original) + rewritten = _rewrite_messages(original, paths, for_provider=False) + if all( + left is right for left, right in zip(original, rewritten, strict=True) + ): + return None + logger.info( + "provider_context_media_checkpoint_repaired messages=%s", + len(rewritten), + ) + return {"messages": Overwrite(rewritten)} + except EvoRuntimeError: + raise + except Exception as exc: + raise self._middleware_failure("before_model", exc) from exc + + async def abefore_model(self, state: Any, runtime: Any) -> dict[str, Any] | None: + _ = runtime + try: + messages = state.get("messages") if isinstance(state, Mapping) else None + if not isinstance(messages, Sequence) or isinstance(messages, str | bytes): + return None + original = list(messages) + paths = await self._apersist(original) + rewritten = _rewrite_messages(original, paths, for_provider=False) + if all( + left is right for left, right in zip(original, rewritten, strict=True) + ): + return None + logger.info( + "provider_context_media_checkpoint_repaired messages=%s", + len(rewritten), + ) + return {"messages": Overwrite(rewritten)} + except EvoRuntimeError: + raise + except Exception as exc: + raise self._middleware_failure("before_model", exc) from exc + + @staticmethod + def _middleware_failure(node: str, exc: Exception) -> EvoRuntimeError: + return EvoRuntimeError( + "AGENT_MIDDLEWARE_FAILED", + details=( + { + "failure_stage": "agent_middleware", + "middleware": ProviderContextMediaMiddleware.name, + "middleware_node": ( + f"{ProviderContextMediaMiddleware.name}.{node}" + ), + "agent_error_type": type(exc).__name__[:128], + "agent_error_module": type(exc).__module__[:128], + }, + ), + ) + + def wrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], ModelResponse], + ) -> ModelResponse: + return self._prepare_response(handler(self._prepare_request(request))) + + async def awrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], Awaitable[ModelResponse]], + ) -> ModelResponse: + prepared = await self._aprepare_request(request) + return await self._aprepare_response(await handler(prepared)) + + +__all__ = ["ProviderContextMediaMiddleware"] diff --git a/EvoScientist/middleware/recoverable_metering.py b/EvoScientist/middleware/recoverable_metering.py new file mode 100644 index 0000000..7418213 --- /dev/null +++ b/EvoScientist/middleware/recoverable_metering.py @@ -0,0 +1,272 @@ +"""Durable model-attempt metering for Graph-native Ai4Sci runs.""" + +from __future__ import annotations + +import asyncio +import logging +import uuid +from collections.abc import Awaitable, Callable, Mapping +from typing import Any + +import httpx +from langchain.agents.middleware.types import ( + AgentMiddleware, + ModelRequest, + ModelResponse, +) +from langchain_core.callbacks import AsyncCallbackHandler + +logger = logging.getLogger(__name__) +_clients: dict[str, httpx.AsyncClient] = {} + + +def _config() -> dict[str, Any] | None: + try: + from langgraph.config import get_config + + value = get_config() + except Exception: + return None + return value if isinstance(value, dict) else None + + +def _metering_config(config: Mapping[str, Any] | None) -> dict[str, str] | None: + if not config: + return None + configurable = config.get("configurable") + if not isinstance(configurable, Mapping): + return None + value = configurable.get("ai4sci_metering") + if not isinstance(value, Mapping): + return None + required = ("gateway_url", "run_id", "envelope_signature") + normalized = { + name: str(value.get(name) or "") + for name in (*required, "provider_id", "model_id", "source_type") + } + return normalized if all(normalized[name] for name in required) else None + + +def _first(mapping: Mapping[str, Any], *keys: str) -> Any: + for key in keys: + value = mapping.get(key) + if value not in (None, ""): + return value + return None + + +def _source_type(metadata: Mapping[str, Any], tags: list[str]) -> str: + explicit = str(metadata.get("metering_scope") or "").lower() + aliases = { + "main": "main_agent", + "main_agent": "main_agent", + "subagent": "subagent", + "tool_selector": "tool_selector", + "summarizer": "summarizer", + "title": "title", + "memory": "evomemory_turn_worker", + "evomemory_turn_worker": "evomemory_turn_worker", + "evomemory_subagent_worker": "evomemory_subagent_worker", + "evomemory_linker": "evomemory_linker", + } + if explicit in aliases: + return aliases[explicit] + hint = " ".join([*tags, *(str(value) for value in metadata.values())]).lower() + if "selector" in hint: + return "tool_selector" + if "summar" in hint or "compact" in hint: + return "summarizer" + if "title" in hint: + return "title" + if "memory" in hint or "observation" in hint: + return "evomemory_turn_worker" + if "subagent" in hint or "sub_agent" in hint or "task:" in hint: + return "subagent" + return "main_agent" + + +def _usage(response: Any) -> dict[str, Any] | None: + llm_output = dict(getattr(response, "llm_output", None) or {}) + candidates: list[Mapping[str, Any]] = [ + dict(llm_output.get("token_usage") or llm_output.get("usage") or {}) + ] + for group in getattr(response, "generations", None) or []: + for generation in group or []: + message = getattr(generation, "message", None) + if message is not None: + candidates.append(dict(getattr(message, "usage_metadata", None) or {})) + for value in candidates: + input_tokens = _first(value, "input_tokens", "prompt_tokens", "input_token_count") + output_tokens = _first(value, "output_tokens", "completion_tokens", "output_token_count") + if input_tokens is None or output_tokens is None: + continue + input_details = dict( + value.get("input_token_details") or value.get("prompt_tokens_details") or {} + ) + output_details = dict( + value.get("output_token_details") or value.get("completion_tokens_details") or {} + ) + normalized_input = max(0, int(input_tokens)) + normalized_output = max(0, int(output_tokens)) + return { + "input_tokens": normalized_input, + "output_tokens": normalized_output, + "cached_input_tokens": max( + 0, + int( + _first( + input_details, + "cache_read", + "cached_tokens", + "cache_read_input_tokens", + ) + or 0 + ), + ), + "reasoning_tokens": max( + 0, int(_first(output_details, "reasoning", "reasoning_tokens") or 0) + ), + "total_tokens": normalized_input + normalized_output, + } + return None + + +async def _post(config: Mapping[str, str], phase: str, payload: dict[str, Any]) -> None: + base_url = config["gateway_url"] + client = _clients.get(base_url) + if client is None: + client = httpx.AsyncClient(timeout=httpx.Timeout(15.0, connect=3.0)) + _clients[base_url] = client + response = await client.post( + f"{base_url}/api/internal/recoverable-runs/metering/{phase}", + json={ + **payload, + "run_id": config["run_id"], + "envelope_signature": config["envelope_signature"], + }, + ) + response.raise_for_status() + + +class RecoverableMeteringCallback(AsyncCallbackHandler): + raise_error = True + + def __init__(self, config: Mapping[str, str]) -> None: + super().__init__() + self.config = dict(config) + self._attempts: set[uuid.UUID] = set() + + async def on_chat_model_start( + self, + serialized: dict[str, Any], + messages: list[list[Any]], + *, + run_id: uuid.UUID, + tags: list[str] | None = None, + metadata: dict[str, Any] | None = None, + **kwargs: Any, + ) -> None: + del messages + metadata = dict(metadata or {}) + invocation = dict(kwargs.get("invocation_params") or {}) + serialized_kwargs = dict(serialized.get("kwargs") or {}) + provider = self.config.get("provider_id") or str( + _first(metadata, "route_provider", "ls_provider", "provider") + or _first(invocation, "provider", "model_provider") + or serialized_kwargs.get("model_provider") + or "unknown" + ) + model = self.config.get("model_id") or str( + _first(metadata, "route_model", "ls_model_name", "model") + or _first(invocation, "model", "model_name") + or _first(serialized_kwargs, "model", "model_name") + or serialized.get("name") + or "" + ) + await _post( + self.config, + "start", + { + "attempt_id": str(run_id), + "source_type": self.config.get("source_type") + or _source_type(metadata, list(tags or [])), + "provider_id": provider, + "model_id": model, + }, + ) + self._attempts.add(run_id) + + async def on_llm_end(self, response: Any, *, run_id: uuid.UUID, **kwargs: Any) -> None: + del kwargs + await self._terminal(run_id, "succeeded", _usage(response)) + + async def on_llm_error( + self, error: BaseException, *, run_id: uuid.UUID, **kwargs: Any + ) -> None: + del error, kwargs + await self._terminal(run_id, "failed", None) + + async def _terminal( + self, run_id: uuid.UUID, outcome: str, usage: dict[str, Any] | None + ) -> None: + if run_id not in self._attempts: + return + last_error: BaseException | None = None + for attempt in range(3): + try: + await _post( + self.config, + "terminal", + {"attempt_id": str(run_id), "outcome": outcome, "usage": usage}, + ) + self._attempts.discard(run_id) + return + except Exception as exc: + last_error = exc + await asyncio.sleep(0.1 * (attempt + 1)) + raise RuntimeError("AI4SCI_METERING_TERMINAL_FAILED") from last_error + + +class RecoverableMeteringMiddleware(AgentMiddleware): + """Attach one durable callback manager to every model invoked by this Run.""" + + name = "recoverable_metering" + + @staticmethod + def _install() -> None: + config = _config() + metering = _metering_config(config) + if config is None or metering is None: + return + callbacks = config.get("callbacks") + if callbacks is None: + config["callbacks"] = [RecoverableMeteringCallback(metering)] + return + if isinstance(callbacks, list): + if not any(isinstance(item, RecoverableMeteringCallback) for item in callbacks): + callbacks.append(RecoverableMeteringCallback(metering)) + return + if hasattr(callbacks, "add_handler"): + handlers = list(getattr(callbacks, "handlers", []) or []) + if not any(isinstance(item, RecoverableMeteringCallback) for item in handlers): + callbacks.add_handler(RecoverableMeteringCallback(metering), inherit=True) + return + logger.error("Unsupported LangChain callback container for recoverable metering") + raise RuntimeError("AI4SCI_METERING_CALLBACKS_UNAVAILABLE") + + def wrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], ModelResponse], + ) -> ModelResponse: + if _metering_config(_config()) is None: + return handler(request) + raise RuntimeError("AI4SCI_RECOVERABLE_RUN_REQUIRES_ASYNC_MODEL_PATH") + + async def awrap_model_call( + self, + request: ModelRequest, + handler: Callable[[ModelRequest], Awaitable[ModelResponse]], + ) -> ModelResponse: + self._install() + return await handler(request) diff --git a/EvoScientist/middleware/recoverable_tools.py b/EvoScientist/middleware/recoverable_tools.py new file mode 100644 index 0000000..b2be0da --- /dev/null +++ b/EvoScientist/middleware/recoverable_tools.py @@ -0,0 +1,205 @@ +"""At-least-once tool-effect recovery for Ai4Sci Graph runs.""" + +from __future__ import annotations + +import hashlib +import json +from collections.abc import Awaitable, Callable, Mapping +from typing import TYPE_CHECKING, Any + +import httpx +from langchain.agents.middleware.types import AgentMiddleware +from langchain_core.messages import ToolMessage, message_to_dict, messages_from_dict +from langgraph.types import Command + +if TYPE_CHECKING: + from langchain.agents.middleware.types import ToolCallRequest + +_READ_ONLY_PREFIXES = ("read_", "get_", "list_", "search_", "find_", "check_") +_READ_ONLY_NAMES = { + "web_search", + "glob", + "grep", + "ls", + "view_image", + "fetch_url", +} +_IDEMPOTENT_NAMES = { + "write_file", + "create_directory", + "mkdir", + "update_file", +} + + +def _canonical(value: Any) -> str: + return json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":"), default=str) + + +def _hash(value: Any) -> str: + return hashlib.sha256(_canonical(value).encode()).hexdigest() + + +def _context() -> tuple[dict[str, str] | None, dict[str, Any]]: + try: + from langgraph.config import get_config + + config = get_config() + except Exception: + return None, {} + configurable = config.get("configurable") if isinstance(config, dict) else None + if not isinstance(configurable, Mapping): + return None, {} + proxy = configurable.get("ai4sci_tool_effect") + metadata = dict(config.get("metadata") or {}) + if not isinstance(proxy, Mapping): + # Compatibility for Runs dispatched before the tool-effect grant was + # split from the model proxy. Detached EvoMemory graphs must never use + # the parent conversation's tool-effect authority. + run_kind = str(metadata.get("run_kind") or "") + if not run_kind.startswith("evomemory_"): + proxy = configurable.get("ai4sci_model_proxy") + if not isinstance(proxy, Mapping): + return None, metadata + normalized = { + name: str(proxy.get(name) or "") + for name in ("gateway_url", "run_id", "envelope_signature") + } + return ( + normalized if all(normalized.values()) else None, + metadata, + ) + + +def _effect_class(tool_name: str) -> str: + lowered = tool_name.lower() + if lowered in _READ_ONLY_NAMES or lowered.startswith(_READ_ONLY_PREFIXES): + return "read_only" + if lowered in _IDEMPOTENT_NAMES: + return "idempotent" + return "non_idempotent" + + +async def _post(proxy: Mapping[str, str], phase: str, payload: dict[str, Any]) -> dict[str, Any]: + async with httpx.AsyncClient(timeout=httpx.Timeout(30.0, connect=3.0)) as client: + response = await client.post( + f"{proxy['gateway_url'].rstrip('/')}/api/internal/recoverable-runs/tool-effect/{phase}", + json={ + **payload, + "run_id": proxy["run_id"], + "attempt_id": proxy["run_id"], + "envelope_signature": proxy["envelope_signature"], + }, + ) + response.raise_for_status() + return dict(response.json()) + + +class RecoverableToolEffectMiddleware(AgentMiddleware): + name = "recoverable_tool_effect" + + def wrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]], + ) -> ToolMessage | Command[Any]: + proxy, _ = _context() + if proxy is None: + return handler(request) + raise RuntimeError("AI4SCI_RECOVERABLE_RUN_REQUIRES_ASYNC_TOOL_PATH") + + async def awrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command[Any]]], + ) -> ToolMessage | Command[Any]: + proxy, metadata = _context() + if proxy is None: + return await handler(request) + tool_call = dict(request.tool_call) + tool_name = str(tool_call.get("name") or "unknown_tool") + tool_call_id = str(tool_call.get("id") or "") + arguments = tool_call.get("args") or {} + request_hash = _hash(arguments) + checkpoint_ns = str(metadata.get("checkpoint_ns") or metadata.get("langgraph_checkpoint_ns") or "") + task_path = ":".join( + str(metadata.get(name) or "") + for name in ("langgraph_step", "langgraph_node", "langgraph_task_idx") + ) + effect_id = _hash( + { + "run_id": proxy["run_id"], + "checkpoint_ns": checkpoint_ns, + "task_path": task_path, + "tool_call_id": tool_call_id, + "tool_name": tool_name, + "request_hash": request_hash, + } + ) + prepared = await _post( + proxy, + "prepare", + { + "effect_id": effect_id, + "checkpoint_ns": checkpoint_ns, + "task_path": task_path, + "tool_call_id": tool_call_id, + "tool_name": tool_name, + "effect_class": _effect_class(tool_name), + "request_hash": request_hash, + }, + ) + if prepared.get("action") == "manual_reconcile": + return ToolMessage( + content=( + f"Tool '{tool_name}' may already have produced an external side effect. " + "Automatic retry is blocked; user confirmation is required." + ), + tool_call_id=tool_call_id, + name=tool_name, + status="error", + ) + if prepared.get("action") == "cached": + result = prepared.get("result") + if isinstance(result, dict) and isinstance(result.get("message"), dict): + messages = messages_from_dict([result["message"]]) + if len(messages) == 1 and isinstance(messages[0], ToolMessage): + return messages[0] + return ToolMessage( + content="Cached tool result is unavailable; user confirmation is required.", + tool_call_id=tool_call_id, + name=tool_name, + status="error", + ) + fencing_token = int(prepared["fencing_token"]) + try: + result = await handler(request) + except BaseException: + await _post( + proxy, + "terminal", + { + "effect_id": effect_id, + "fencing_token": fencing_token, + "outcome": "failed", + "result": {}, + }, + ) + raise + successful = isinstance(result, ToolMessage) and result.status != "error" + payload = ( + {"message": message_to_dict(result)} + if isinstance(result, ToolMessage) + else {"command_result": True} + ) + await _post( + proxy, + "terminal", + { + "effect_id": effect_id, + "fencing_token": fencing_token, + "outcome": "succeeded" if successful else "failed", + "result": payload, + }, + ) + return result diff --git a/EvoScientist/middleware/skill_context.py b/EvoScientist/middleware/skill_context.py new file mode 100644 index 0000000..59860b6 --- /dev/null +++ b/EvoScientist/middleware/skill_context.py @@ -0,0 +1,161 @@ +"""Bounded skill discovery for model prompts. + +DeepAgents' stock ``SkillsMiddleware`` keeps the full skill catalog in agent +state and renders every description into each model request. That is suitable +for a small catalog but makes large global catalogs consume the whole context +window. This subclass preserves catalog loading and file access while exposing +only a relevant, byte-bounded subset in the prompt. +""" + +from __future__ import annotations + +import re +from collections.abc import Iterable, Sequence +from typing import Any + +from deepagents.middleware._utils import append_to_system_message +from deepagents.middleware.skills import SkillsMiddleware +from langchain.agents.middleware.types import ModelRequest +from langchain_core.messages import HumanMessage + +DEFAULT_MAX_SKILLS = 16 +DEFAULT_MAX_SKILLS_BYTES = 12 * 1024 +DEFAULT_MAX_DESCRIPTION_BYTES = 320 + +_WORD_RE = re.compile(r"[a-z0-9][a-z0-9_-]{1,}", re.IGNORECASE) +_CJK_RUN_RE = re.compile(r"[\u4e00-\u9fff]{2,}") + + +def _truncate_utf8(value: str, limit: int) -> str: + """Truncate text on a UTF-8 boundary, reserving room for an ellipsis.""" + + encoded = value.encode("utf-8") + if len(encoded) <= limit: + return value + if limit <= 3: + return "" + return encoded[: limit - 3].decode("utf-8", errors="ignore") + "..." + + +def _query_terms(value: str) -> set[str]: + """Extract ASCII words and CJK n-grams without a tokenizer dependency.""" + + normalized = value.lower() + terms = set(_WORD_RE.findall(normalized)) + for run in _CJK_RUN_RE.findall(normalized): + for width in range(2, min(4, len(run)) + 1): + terms.update( + run[index : index + width] for index in range(len(run) - width + 1) + ) + return terms + + +def _message_text(messages: Iterable[Any]) -> str: + """Return recent user text only; tool output must not drive skill ranking.""" + + for message in reversed(list(messages)): + if not isinstance(message, HumanMessage): + continue + content = message.content + if isinstance(content, str): + return content + if isinstance(content, Sequence): + return " ".join(str(item) for item in content) + return "" + + +class BudgetedSkillsMiddleware(SkillsMiddleware): + """Expose only relevant skills within a fixed system-prompt byte budget.""" + + name = "budgeted_skills" + + def __init__( + self, + *, + backend: Any, + sources: Sequence[str | tuple[str, str]] | str, + max_skills: int = DEFAULT_MAX_SKILLS, + max_skills_bytes: int = DEFAULT_MAX_SKILLS_BYTES, + max_description_bytes: int = DEFAULT_MAX_DESCRIPTION_BYTES, + ) -> None: + resolved_sources = [sources] if isinstance(sources, str) else list(sources) + super().__init__(backend=backend, sources=resolved_sources) + if max_skills < 1 or max_skills_bytes < 1 or max_description_bytes < 1: + raise ValueError("skill context limits must be positive") + self.max_skills = max_skills + self.max_skills_bytes = max_skills_bytes + self.max_description_bytes = max_description_bytes + + def _select_skills( + self, skills: Sequence[dict[str, Any]], query: str + ) -> list[dict[str, Any]]: + terms = _query_terms(query) + if not terms: + return [] + + scored: list[tuple[int, str, dict[str, Any]]] = [] + for skill in skills: + name = str(skill.get("name") or "") + description = str(skill.get("description") or "") + name_text = name.lower() + description_text = description.lower() + score = 0 + for term in terms: + if term == name_text: + score += 100 + elif term in name_text: + score += 24 + if term in description_text: + score += 4 + if score: + scored.append((score, name, skill)) + + scored.sort(key=lambda item: (-item[0], item[1])) + return [skill for _, _, skill in scored[: self.max_skills]] + + def _format_budgeted_skills(self, skills: Sequence[dict[str, Any]]) -> str: + lines: list[str] = [] + used = 0 + for skill in skills: + name = str(skill.get("name") or "unnamed") + path = str(skill.get("path") or "") + description = _truncate_utf8( + str(skill.get("description") or ""), self.max_description_bytes + ) + item = ( + f"- **{name}**: {description}\n -> Read `{path}` for full instructions" + ) + item_bytes = len(item.encode("utf-8")) + separator = 1 if lines else 0 + if used + separator + item_bytes > self.max_skills_bytes: + continue + lines.append(item) + used += separator + item_bytes + return "\n".join(lines) + + def modify_request(self, request: ModelRequest) -> ModelRequest: + if self.system_prompt_template is None: + return request + + state = request.state or {} + skills = state.get("skills_metadata", []) + selected = self._select_skills(skills, _message_text(request.messages)) + skills_list = self._format_budgeted_skills(selected) + if len(selected) < len(skills): + discovery_note = ( + "\n\nThis is a query-relevant subset of the installed skills. " + "Use skill_manager(action='list') to discover the full catalog." + ) + skills_list += discovery_note + skills_section = self.system_prompt_template.format( + skills_locations=self._format_skills_locations(), + skills_load_warnings=self._format_skills_load_warnings( + state.get("skills_load_errors", []) + ), + skills_list=skills_list, + ) + return request.override( + system_message=append_to_system_message( + request.system_message, skills_section + ) + ) diff --git a/EvoScientist/middleware/tool_call_normalizer.py b/EvoScientist/middleware/tool_call_normalizer.py new file mode 100644 index 0000000..bb3312a --- /dev/null +++ b/EvoScientist/middleware/tool_call_normalizer.py @@ -0,0 +1,464 @@ +"""Normalize provider tool-call shapes into LangChain's canonical call form. + +Provider adapters are allowed to disagree about mechanical fields such as +``tool_call.id`` and whether arguments are delivered as a JSON object or a +JSON string. The agent execution layer is not. This module is deliberately +limited to deterministic protocol translation: it never infers a tool name, +repairs incomplete JSON, or extracts calls from ordinary model text. +""" + +from __future__ import annotations + +import copy +import hashlib +import hmac +import json +import os +from collections.abc import Mapping +from dataclasses import dataclass, replace +from typing import Any + +from langchain.agents.middleware.types import ExtendedModelResponse, ModelResponse +from langchain_core.messages import AIMessage + +_TOOL_BLOCK_TYPES = frozenset({"tool_call", "function_call", "tool_use"}) +_ID_SECRET_ENVS = ( + "AI4SCI_EVO_TOOL_CALL_ID_SECRET", + "EVOSCI_TOOL_CALL_ID_SECRET", + "AI4SCI_EVO_RUNTIME_GRANT_SECRET", +) +_DEFAULT_ID_SECRET = b"evoscientist-tool-call-normalizer-v1" +_RAW_PROVIDER_CALL_FIELDS = ( + "tool_calls", + "function_call", + "functionCall", + "tool_use", +) + + +class ToolCallNormalizationError(ValueError): + """A provider response could not be translated without guessing.""" + + def __init__( + self, + reason: str, + *, + call_index: int, + source: str, + ) -> None: + super().__init__(reason) + self.reason = reason + self.call_index = call_index + self.source = source + + def diagnostic(self) -> dict[str, Any]: + return { + "source": self.source, + "call_index": self.call_index, + "normalization_failure": self.reason, + } + + +@dataclass(frozen=True, slots=True) +class CanonicalToolCall: + """Provider-neutral tool-call representation accepted by the ToolNode.""" + + id: str + name: str + arguments: Any + index: int + source: str + id_origin: str + + def as_langchain_call(self) -> dict[str, Any]: + args = dict(self.arguments) if isinstance(self.arguments, Mapping) else self.arguments + return {"id": self.id, "name": self.name, "args": args, "type": "tool_call"} + + +@dataclass(frozen=True, slots=True) +class _DecodedToolCall: + call_id: str + name: str + arguments: Any + arguments_present: bool + index: int + source: str + + +def _text(value: Any) -> str: + return value.strip() if isinstance(value, str) else "" + + +def _call_fields(call: Mapping[str, Any]) -> tuple[str, str, Any, bool]: + """Read the common OpenAI, Anthropic and Gemini call field variants.""" + function = call.get("function") + function = function if isinstance(function, Mapping) else {} + # Responses function-call items expose two identifiers: ``id`` identifies + # the output item (for example ``fc_*``), while ``call_id`` links the + # eventual tool result (for example ``call_*``). Chat Completions only + # has ``id``. Prefer the transport-level call identifier when present so + # the parsed LangChain call and its content block describe the same call. + call_id = _text(call.get("call_id") or call.get("id")) + name = _text(call.get("name") or call.get("tool_name") or function.get("name")) + if "args" in call: + return call_id, name, call.get("args"), True + if "arguments" in call: + return call_id, name, call.get("arguments"), True + if "input" in call: + return call_id, name, call.get("input"), True + if "arguments" in function: + return call_id, name, function.get("arguments"), True + return call_id, name, None, False + + +def _tool_blocks(message: AIMessage) -> list[Mapping[str, Any]]: + content = getattr(message, "content", None) + if not isinstance(content, list): + return [] + return [ + block + for block in content + if isinstance(block, Mapping) and block.get("type") in _TOOL_BLOCK_TYPES + ] + + +class AdapterToolCallDecoder: + """Decode provider-shaped calls without exposing provider data downstream. + + LangChain has already decoded most Provider wire formats to ``AIMessage``. + This adapter deliberately accepts those canonical calls plus the three + remaining lossless sources: OpenAI-compatible raw calls, Anthropic/Gemini + content blocks and old ``function_call`` fields. Adapter-specific + decoders can replace this class later without changing the normalizer or + guard contract. + """ + + def __init__(self, adapter_id: str | None = None) -> None: + self.adapter_id = adapter_id or "generic" + + def decode(self, message: AIMessage) -> list[_DecodedToolCall]: + sources = self._sources(message) + if not sources: + return [] + primary_name, primary_calls = sources[0] + decoded = [ + self._decode_one(call, index=index, source=primary_name) + for index, call in enumerate(primary_calls) + ] + for source_name, source_calls in sources[1:]: + # A second source can only be used as field-level evidence when it + # preserves the same call ordering. Anything else is ambiguous. + if len(source_calls) != len(decoded): + raise ToolCallNormalizationError( + "inconsistent_source_count", + call_index=0, + source=source_name, + ) + decoded = [ + self._merge( + primary, + self._decode_one(raw, index=index, source=source_name), + ) + for index, (primary, raw) in enumerate( + zip(decoded, source_calls, strict=True) + ) + ] + return decoded + + @staticmethod + def _sources(message: AIMessage) -> list[tuple[str, list[Mapping[str, Any]]]]: + sources: list[tuple[str, list[Mapping[str, Any]]]] = [] + + def add_source(source: str, value: Any) -> None: + if value is None: + return + calls = list(value) if isinstance(value, list | tuple) else [value] + if not calls: + return + if any(not isinstance(call, Mapping) for call in calls): + raise ToolCallNormalizationError( + "invalid_call_shape", call_index=0, source=source + ) + sources.append((source, calls)) + + add_source("parsed_tool_calls", getattr(message, "tool_calls", None) or []) + + additional = getattr(message, "additional_kwargs", None) + additional = additional if isinstance(additional, Mapping) else {} + add_source("provider_raw_tool_calls", additional.get("tool_calls")) + for legacy_field in _RAW_PROVIDER_CALL_FIELDS[1:]: + add_source(f"provider_{legacy_field}", additional.get(legacy_field)) + + blocks = _tool_blocks(message) + if blocks: + sources.append(("content_blocks", blocks)) + return sources + + @staticmethod + def _decode_one( + call: Mapping[str, Any], *, index: int, source: str + ) -> _DecodedToolCall: + call_id, name, arguments, arguments_present = _call_fields(call) + return _DecodedToolCall( + call_id=call_id, + name=name, + arguments=arguments, + arguments_present=arguments_present, + index=index, + source=source, + ) + + @staticmethod + def _merge( + primary: _DecodedToolCall, evidence: _DecodedToolCall + ) -> _DecodedToolCall: + def choose(field: str, first: Any, second: Any) -> Any: + first_present = bool(first) if field in {"call_id", "name"} else first is not None + second_present = bool(second) if field in {"call_id", "name"} else second is not None + if first_present and second_present and first != second: + raise ToolCallNormalizationError( + "inconsistent_source", + call_index=primary.index, + source=evidence.source, + ) + return first if first_present else second + + call_id = choose("call_id", primary.call_id, evidence.call_id) + name = choose("name", primary.name, evidence.name) + if primary.arguments_present and evidence.arguments_present: + if _parse_json_object(primary.arguments) != _parse_json_object( + evidence.arguments + ): + raise ToolCallNormalizationError( + "inconsistent_source", + call_index=primary.index, + source=evidence.source, + ) + arguments = primary.arguments + elif primary.arguments_present: + arguments = primary.arguments + else: + arguments = evidence.arguments + return _DecodedToolCall( + call_id=call_id, + name=name, + arguments=arguments, + arguments_present=( + primary.arguments_present or evidence.arguments_present + ), + index=primary.index, + source=primary.source, + ) + + +def _parse_json_object(value: Any) -> Any: + if not isinstance(value, str): + return value + try: + return json.loads(value) + except (TypeError, ValueError, json.JSONDecodeError): + return value + + +class ToolCallNormalizer: + """Create canonical calls and immutable normalized AI messages.""" + + version = "v1" + + def __init__(self, *, id_secret: bytes | None = None) -> None: + env_secret = next( + (os.environ[name] for name in _ID_SECRET_ENVS if os.environ.get(name)), + None, + ) + self._id_secret = id_secret or ( + env_secret.encode("utf-8") if env_secret else _DEFAULT_ID_SECRET + ) + + def normalize_message( + self, + message: AIMessage, + *, + adapter_id: str | None, + request_scope: str, + ) -> AIMessage: + decoded = AdapterToolCallDecoder(adapter_id).decode(message) + if not decoded: + return message + + canonical_calls = [ + self._canonicalize(call, request_scope=request_scope) for call in decoded + ] + canonical_dicts = [call.as_langchain_call() for call in canonical_calls] + if self._already_canonical(message, canonical_dicts): + return message + return self._copy_message(message, canonical_dicts) + + def _canonicalize( + self, call: _DecodedToolCall, *, request_scope: str + ) -> CanonicalToolCall: + arguments = self._arguments(call) + call_id = call.call_id + id_origin = "provider" + if not call_id and call.name and _is_json_object(arguments): + call_id = self._generated_id( + request_scope=request_scope, + call=call, + arguments=arguments, + ) + id_origin = "gateway_generated" + return CanonicalToolCall( + id=call_id, + name=call.name, + arguments=arguments, + index=call.index, + source=call.source, + id_origin=id_origin, + ) + + @staticmethod + def _arguments(call: _DecodedToolCall) -> Any: + if not call.arguments_present: + return None + return _parse_json_object(call.arguments) + + def _generated_id( + self, + *, + request_scope: str, + call: _DecodedToolCall, + arguments: Mapping[str, Any], + ) -> str: + canonical_args = json.dumps( + arguments, + ensure_ascii=False, + sort_keys=True, + separators=(",", ":"), + ) + material = "\x1f".join( + (self.version, request_scope, str(call.index), call.name, canonical_args) + ).encode("utf-8") + return "call_" + hmac.new( + self._id_secret, material, hashlib.sha256 + ).hexdigest()[:32] + + @staticmethod + def _already_canonical( + message: AIMessage, canonical_calls: list[dict[str, Any]] + ) -> bool: + current_calls = list(getattr(message, "tool_calls", None) or []) + if current_calls != canonical_calls: + return False + if getattr(message, "invalid_tool_calls", None): + return False + additional = getattr(message, "additional_kwargs", None) or {} + if any(additional.get(key) for key in _RAW_PROVIDER_CALL_FIELDS): + return False + blocks = _tool_blocks(message) + if not blocks: + return True + if len(blocks) != len(canonical_calls): + return False + return all( + _text(block.get("call_id") or block.get("id")) == call["id"] + and _text( + block.get("name") + or block.get("tool_name") + or ( + block.get("function", {}).get("name") + if isinstance(block.get("function"), Mapping) + else "" + ) + ) + == call["name"] + for block, call in zip(blocks, canonical_calls, strict=True) + ) + + @staticmethod + def _copy_message( + message: AIMessage, canonical_calls: list[dict[str, Any]] + ) -> AIMessage: + copied = copy.copy(message) + copied.tool_calls = canonical_calls + # A repaired call cannot make a second malformed call safe. Preserve + # invalid calls so the Guard still rejects the entire response. + copied.invalid_tool_calls = list( + getattr(message, "invalid_tool_calls", None) or [] + ) + additional = dict(getattr(message, "additional_kwargs", None) or {}) + # Parsed calls are now canonical; retaining a raw provider copy can + # reintroduce the missing ID when history is replayed. + for field in _RAW_PROVIDER_CALL_FIELDS: + additional.pop(field, None) + copied.additional_kwargs = additional + + content = getattr(message, "content", None) + if isinstance(content, list): + call_index = 0 + normalized_content: list[Any] = [] + for value in content: + if not isinstance(value, Mapping) or value.get("type") not in _TOOL_BLOCK_TYPES: + normalized_content.append(value) + continue + if call_index >= len(canonical_calls): + normalized_content.append(value) + continue + call = canonical_calls[call_index] + block = dict(value) + block["id"] = call["id"] + block["name"] = call["name"] + if isinstance(block.get("function"), Mapping): + block["function"] = { + **block["function"], + "name": call["name"], + } + normalized_content.append(block) + call_index += 1 + copied.content = normalized_content + return copied + + def normalize_response( + self, + response: Any, + *, + adapter_id: str | None, + request_scope: str, + ) -> Any: + if isinstance(response, AIMessage): + return self.normalize_message( + response, adapter_id=adapter_id, request_scope=request_scope + ) + if isinstance(response, ExtendedModelResponse): + normalized = self.normalize_response( + response.model_response, + adapter_id=adapter_id, + request_scope=request_scope, + ) + return ( + response + if normalized is response.model_response + else replace(response, model_response=normalized) + ) + if isinstance(response, ModelResponse): + result = list(response.result) + normalized_result = [ + self.normalize_message( + message, + adapter_id=adapter_id, + request_scope=request_scope, + ) + if isinstance(message, AIMessage) + else message + for message in result + ] + return response if normalized_result == result else replace(response, result=normalized_result) + return response + + +def _is_json_object(value: Any) -> bool: + if not isinstance(value, Mapping): + return False + try: + json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":")) + except (TypeError, ValueError): + return False + return True diff --git a/EvoScientist/middleware/tool_protocol_guard.py b/EvoScientist/middleware/tool_protocol_guard.py index eb149d4..2ec4283 100644 --- a/EvoScientist/middleware/tool_protocol_guard.py +++ b/EvoScientist/middleware/tool_protocol_guard.py @@ -4,6 +4,7 @@ from __future__ import annotations import hashlib import json +import logging from collections.abc import Awaitable, Callable, Mapping, Sequence from typing import Any @@ -17,8 +18,11 @@ from langchain_core.messages import AIMessage from langchain_core.tools import BaseTool from ..llm.errors import ModelToolProtocolError, _provider_from_model +from .tool_call_normalizer import ToolCallNormalizationError, ToolCallNormalizer -_TOOL_BLOCK_TYPES = frozenset({"tool_call", "function_call"}) +logger = logging.getLogger(__name__) + +_TOOL_BLOCK_TYPES = frozenset({"tool_call", "function_call", "tool_use"}) _MAX_DIAGNOSTIC_KEYS = 16 _MAX_DIAGNOSTIC_KEY_CHARS = 64 @@ -54,7 +58,10 @@ def _ai_messages(response: Any) -> list[AIMessage]: def _block_identity(block: Mapping[str, Any]) -> tuple[str, str]: - call_id = str(block.get("id") or block.get("call_id") or "").strip() + # Responses output items use ``id`` for the item and ``call_id`` for the + # tool-result correlation key. The latter is the identity ToolMessage + # must carry; Chat Completions continues to fall back to ``id``. + call_id = str(block.get("call_id") or block.get("id") or "").strip() name = block.get("name") or block.get("tool_name") function = block.get("function") if not name and isinstance(function, Mapping): @@ -104,12 +111,22 @@ def _argument_diagnostic(value: Any, *, present: bool) -> dict[str, Any]: } +def _is_json_object(value: Any) -> bool: + if not isinstance(value, Mapping): + return False + try: + json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":")) + except (TypeError, ValueError): + return False + return True + + def _summarize_call(call: Any) -> dict[str, Any]: if not isinstance(call, Mapping): return {"call_type": type(call).__name__} function = call.get("function") function = function if isinstance(function, Mapping) else {} - call_id = str(call.get("id") or call.get("call_id") or "").strip() + call_id = str(call.get("call_id") or call.get("id") or "").strip() name = call.get("name") or call.get("tool_name") or function.get("name") name = str(name or "").strip() if "args" in call: @@ -205,6 +222,28 @@ def _route_metadata(request: ModelRequest) -> dict[str, Any]: } +def _normalization_scope(request: ModelRequest) -> str: + """Return non-sensitive per-turn material for generated call identifiers.""" + route = _route_metadata(request) + runtime = getattr(request, "runtime", None) + config = getattr(runtime, "config", None) + config = config if isinstance(config, Mapping) else {} + configurable = config.get("configurable") + configurable = configurable if isinstance(configurable, Mapping) else {} + thread_id = str(configurable.get("thread_id") or "") + messages = getattr(request, "messages", None) + message_count = len(messages) if isinstance(messages, Sequence) else 0 + return "|".join( + ( + str(route.get("route_key") or ""), + str(route.get("config_generation") or ""), + str(route.get("api_mode") or ""), + thread_id, + str(message_count), + ) + ) + + def _raise_protocol_error( request: ModelRequest, reason: str, @@ -212,11 +251,24 @@ def _raise_protocol_error( call_id: str | None = None, call_diagnostic: dict[str, Any] | None = None, ) -> None: + route = _route_metadata(request) + logger.warning( + "model_tool_protocol_invalid reason=%s provider=%s model=%s route_key=%s " + "api_mode=%s transport=%s call_id_present=%s diagnostic=%s", + reason, + route["provider"], + route["model"], + route["route_key"], + route["api_mode"], + route["tool_call_transport"], + bool(call_id), + json.dumps(call_diagnostic or {}, ensure_ascii=True, sort_keys=True), + ) raise ModelToolProtocolError( reason, call_id=call_id or None, call_diagnostic=call_diagnostic, - **_route_metadata(request), + **route, ) @@ -282,7 +334,7 @@ def _validate_message( call_diagnostic=diagnostic, ) args = raw_call.get("args") - if not isinstance(args, Mapping): + if not _is_json_object(args): _raise_protocol_error( request, "invalid_args", @@ -330,10 +382,30 @@ def _validate_message( class ToolProtocolGuardMiddleware(AgentMiddleware): - """Fail closed on malformed final tool calls using the actual request tools.""" + """Normalize then fail closed on malformed final tool calls.""" name = "tool_protocol_guard" + def __init__(self, *, normalizer: ToolCallNormalizer | None = None) -> None: + super().__init__() + self._normalizer = normalizer or ToolCallNormalizer() + + def _normalize(self, response: Any, request: ModelRequest) -> Any: + metadata = getattr(request.model, "metadata", None) + metadata = metadata if isinstance(metadata, Mapping) else {} + try: + return self._normalizer.normalize_response( + response, + adapter_id=str(metadata.get("route_adapter_id") or "") or None, + request_scope=_normalization_scope(request), + ) + except ToolCallNormalizationError as error: + _raise_protocol_error( + request, + error.reason, + call_diagnostic=error.diagnostic(), + ) + @staticmethod def _validate(response: Any, request: ModelRequest) -> None: allowed_names = frozenset( @@ -348,6 +420,7 @@ class ToolProtocolGuardMiddleware(AgentMiddleware): handler: Callable[[ModelRequest], ModelResponse], ) -> ModelResponse: response = handler(request) + response = self._normalize(response, request) self._validate(response, request) return response @@ -357,5 +430,6 @@ class ToolProtocolGuardMiddleware(AgentMiddleware): handler: Callable[[ModelRequest], Awaitable[ModelResponse]], ) -> ModelResponse: response = await handler(request) + response = self._normalize(response, request) self._validate(response, request) return response diff --git a/EvoScientist/runtime_integrations.py b/EvoScientist/runtime_integrations.py index bc0ae2f..0ff8b86 100644 --- a/EvoScientist/runtime_integrations.py +++ b/EvoScientist/runtime_integrations.py @@ -16,7 +16,6 @@ from typing import Any AsyncProvider = Callable[[], Awaitable[Any]] AsyncFileHandler = Callable[[Path], Awaitable[Any]] AsyncUsageRecorder = Callable[[str, str], Awaitable[Any]] -ModelResolver = Callable[[str | None, str | None], Any | None] class RuntimeIntegrationUnavailable(RuntimeError): @@ -33,7 +32,6 @@ class RuntimeIntegrations: knowledge_file_handler: AsyncFileHandler | None = None usage_recorder: AsyncUsageRecorder | None = None image_backend_factory: Callable[[], Any] | None = None - model_resolver: ModelResolver | None = None _integrations = RuntimeIntegrations() @@ -89,12 +87,6 @@ def resolve_user_storage_root(user_id: str) -> Path | None: return provider(user_id) if provider is not None else None -def resolve_runtime_model(model: str | None, provider: str | None = None) -> Any | None: - """Resolve a host-managed model configuration when one is registered.""" - resolver = _integrations.model_resolver - return resolver(model, provider) if resolver is not None else None - - async def handle_knowledge_file(path: Path) -> None: handler = _integrations.knowledge_file_handler if handler is not None: diff --git a/EvoScientist/scope_registry.py b/EvoScientist/scope_registry.py new file mode 100644 index 0000000..7b2117a --- /dev/null +++ b/EvoScientist/scope_registry.py @@ -0,0 +1,1279 @@ +"""Persistent ownership registry for conversation workspaces. + +The LangGraph checkpoint store is not an authority for filesystem ownership: +threads, runs and cron records can be created independently and the filesystem +cannot participate in their transactions. This module keeps the small, +deployment-local registry that binds all of them to one conversation scope. + +The v1 implementation intentionally uses SQLite. It is safe for the supported +single-host deployment topology and keeps the registry outside every agent +workspace. Callers must use the public methods below rather than addressing the +database directly so a future PostgreSQL adapter has one replacement point. +""" + +from __future__ import annotations + +import os +import secrets +import sqlite3 +import threading +import uuid +from collections.abc import Iterator +from contextlib import contextmanager +from dataclasses import dataclass +from datetime import UTC, datetime, timedelta +from pathlib import Path +from typing import Literal + +from . import paths + +ScopeState = Literal["provisioning", "draft", "active", "deleting", "deleted"] +OwnerState = Literal[ + "reserved", "active", "draining", "terminal", "failed", "quarantined" +] + + +class ScopeRegistryError(RuntimeError): + """Base class for registry failures.""" + + +class ScopeNotFoundError(ScopeRegistryError): + """Raised when no matching deployment scope exists.""" + + +class ScopeConflictError(ScopeRegistryError): + """Raised for stale revisions or incompatible ownership mappings.""" + + +class ScopeIdempotencyConflictError(ScopeConflictError): + """Raised when one run request id is reused for another payload.""" + + +class ScopeInterruptResolvedError(ScopeConflictError): + """Raised when an interrupt already has a persisted resolution.""" + + +class ScopeAccessError(ScopeRegistryError): + """Raised when a runtime does not own the requested scope.""" + + +@dataclass(frozen=True, slots=True) +class ScopeRecord: + deployment_id: str + scope_id: str + primary_thread_id: str + state: ScopeState + revision: int + primary_owner_id: str + created_at: str + updated_at: str + + +@dataclass(frozen=True, slots=True) +class OwnerRecord: + deployment_id: str + owner_id: str + scope_id: str + owner_type: str + resource_id: str | None + parent_owner_id: str | None + state: OwnerState + created_at: str + updated_at: str + + +@dataclass(frozen=True, slots=True) +class DeploymentLock: + deployment_id: str + lock_name: str + operation_id: str + expires_at: str + + +@dataclass(frozen=True, slots=True) +class ScopeOperation: + deployment_id: str + operation_id: str + scope_id: str | None + kind: str + state: str + result_sha256: str | None + last_error_code: str | None + + +@dataclass(frozen=True, slots=True) +class RunReservation: + deployment_id: str + scope_id: str + run_request_id: str + turn_id: str + interrupt_key: str | None + request_hash: str + run_owner_id: str + run_id: str | None + state: str + + +# Kept as a source-compatible type name for callers that have not yet switched +# to the run-request terminology. A turn is a logical conversation unit; a +# reservation belongs to one concrete run request. +TurnReservation = RunReservation + + +_TERMINAL_SCOPE_STATES = frozenset({"deleted"}) +_TERMINAL_OWNER_STATES = frozenset({"terminal", "quarantined"}) +_SCOPE_TRANSITIONS: dict[str, frozenset[str]] = { + "provisioning": frozenset({"draft", "deleting", "deleted"}), + "draft": frozenset({"active", "deleting", "deleted"}), + "active": frozenset({"deleting"}), + "deleting": frozenset({"deleted"}), + "deleted": frozenset(), +} + + +def _utc_now() -> str: + return datetime.now(UTC).isoformat() + + +def _ensure_uuid(value: str, field: str) -> str: + try: + return str(uuid.UUID(value)) + except (TypeError, ValueError) as exc: + raise ScopeRegistryError(f"{field} must be a UUID") from exc + + +def deployment_id_for_workspace(workspace_root: Path | str | None = None) -> str: + """Return the stable deployment identifier for a workspace root.""" + + configured = os.getenv("EVOSCIENTIST_DEPLOYMENT_ID", "").strip() + if configured: + return configured + # Deploy resolves the workspace before starting LangGraph. Avoid resolve() + # here because this function also runs on the agent's async execution path. + root = Path(workspace_root or paths.WORKSPACE_ROOT).expanduser() + return str(uuid.uuid5(uuid.NAMESPACE_URL, f"evoscientist:{root}")) + + +def default_registry_path(workspace_root: Path | str | None = None) -> Path: + root = Path(workspace_root or paths.WORKSPACE_ROOT).expanduser() + return root / ".evoscientist" / "control" / "scope-registry.sqlite3" + + +def scope_service_token_path(workspace_root: Path | str | None = None) -> Path: + """Location of the local backend-to-BFF registry credential.""" + configured = os.getenv("EVOSCIENTIST_CONTROL_DIR", "").strip() + control_dir = ( + Path(configured).expanduser() + if configured + else Path.home() / ".evoscientist" / "control" + ) + return control_dir / "scope-service-token" + + +def get_scope_service_token(workspace_root: Path | str | None = None) -> str: + """Return the durable host-local token shared by deploy and the WebUI. + + It lives beside the control-plane database, never in a conversation scope + or browser-delivered configuration. Creation is atomic so a concurrent + launcher cannot replace a running backend's credential. + """ + + path = scope_service_token_path(workspace_root) + path.parent.mkdir(mode=0o700, parents=True, exist_ok=True) + try: + path.parent.chmod(0o700) + except OSError: + pass + try: + token = path.read_text(encoding="utf-8").strip() + except FileNotFoundError: + token = "" + if token: + return token + generated = secrets.token_urlsafe(32) + try: + fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600) + except FileExistsError: + token = path.read_text(encoding="utf-8").strip() + if token: + return token + raise ScopeRegistryError("workspace scope service token is empty") from None + with os.fdopen(fd, "w", encoding="utf-8") as token_file: + token_file.write(generated) + token_file.write("\n") + try: + path.chmod(0o600) + except OSError: + pass + return generated + + +class ScopeRegistry: + """SQLite-backed registry with revision-checked state transitions.""" + + def __init__(self, database_path: Path | str): + self.path = Path(database_path).expanduser() + self._init_lock = threading.Lock() + self._initialized = False + + def initialize(self) -> None: + with self._init_lock: + if self._initialized: + return + self.path.parent.mkdir(mode=0o700, parents=True, exist_ok=True) + try: + self.path.parent.chmod(0o700) + except OSError: + pass + with self._connect() as conn: + self._migrate(conn) + try: + self.path.chmod(0o600) + except OSError: + pass + self._initialized = True + + def _connect(self) -> sqlite3.Connection: + conn = sqlite3.connect(self.path, timeout=30, isolation_level=None) + conn.row_factory = sqlite3.Row + conn.execute("PRAGMA foreign_keys = ON") + conn.execute("PRAGMA journal_mode = WAL") + conn.execute("PRAGMA busy_timeout = 30000") + return conn + + @staticmethod + def _migrate(conn: sqlite3.Connection) -> None: + version = int(conn.execute("PRAGMA user_version").fetchone()[0]) + if version > 2: + raise ScopeRegistryError("scope registry schema is newer than this binary") + if version == 0: + conn.executescript( + """ + CREATE TABLE scopes ( + deployment_id TEXT NOT NULL, + scope_id TEXT NOT NULL, + primary_thread_id TEXT NOT NULL, + state TEXT NOT NULL, + revision INTEGER NOT NULL, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + deleted_at TEXT, + PRIMARY KEY (deployment_id, scope_id), + UNIQUE (deployment_id, primary_thread_id) + ); + + CREATE TABLE scope_owners ( + deployment_id TEXT NOT NULL, + owner_id TEXT NOT NULL, + scope_id TEXT NOT NULL, + owner_type TEXT NOT NULL, + resource_id TEXT, + parent_owner_id TEXT, + state TEXT NOT NULL, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + terminal_at TEXT, + PRIMARY KEY (deployment_id, owner_id), + FOREIGN KEY (deployment_id, scope_id) + REFERENCES scopes(deployment_id, scope_id) + ); + CREATE UNIQUE INDEX scope_owner_resource_unique + ON scope_owners(deployment_id, owner_type, resource_id) + WHERE resource_id IS NOT NULL; + CREATE INDEX scope_owners_scope_state + ON scope_owners(deployment_id, scope_id, state); + + CREATE TABLE scope_run_requests ( + deployment_id TEXT NOT NULL, + scope_id TEXT NOT NULL, + run_request_id TEXT NOT NULL, + turn_id TEXT NOT NULL, + interrupt_key TEXT, + request_hash TEXT NOT NULL, + run_owner_id TEXT NOT NULL, + run_id TEXT, + state TEXT NOT NULL, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + PRIMARY KEY (deployment_id, scope_id, run_request_id), + FOREIGN KEY (deployment_id, scope_id) + REFERENCES scopes(deployment_id, scope_id) + ); + CREATE INDEX scope_run_requests_turn + ON scope_run_requests(deployment_id, scope_id, turn_id); + CREATE UNIQUE INDEX scope_run_requests_interrupt_unique + ON scope_run_requests(deployment_id, scope_id, interrupt_key) + WHERE interrupt_key IS NOT NULL; + + CREATE TABLE scope_operations ( + deployment_id TEXT NOT NULL, + operation_id TEXT NOT NULL, + scope_id TEXT, + kind TEXT NOT NULL, + expected_revision INTEGER, + state TEXT NOT NULL, + external_resource_id TEXT, + result_sha256 TEXT, + last_error_code TEXT, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + PRIMARY KEY (deployment_id, operation_id) + ); + CREATE INDEX scope_operations_scope_state + ON scope_operations(deployment_id, scope_id, state); + + CREATE TABLE deployment_locks ( + deployment_id TEXT NOT NULL, + lock_name TEXT NOT NULL, + operation_id TEXT NOT NULL, + expires_at TEXT NOT NULL, + created_at TEXT NOT NULL, + PRIMARY KEY (deployment_id, lock_name) + ); + """ + ) + conn.execute("PRAGMA user_version = 2") + return + if version == 1: + # SQLite cannot change a composite primary key in place. Historical + # reservations used turn_id as the idempotency key, so copy each row + # with run_request_id = turn_id. New resume runs may then retain the + # logical turn while receiving distinct request ids. + conn.executescript( + """ + CREATE TABLE scope_run_requests ( + deployment_id TEXT NOT NULL, + scope_id TEXT NOT NULL, + run_request_id TEXT NOT NULL, + turn_id TEXT NOT NULL, + interrupt_key TEXT, + request_hash TEXT NOT NULL, + run_owner_id TEXT NOT NULL, + run_id TEXT, + state TEXT NOT NULL, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL, + PRIMARY KEY (deployment_id, scope_id, run_request_id), + FOREIGN KEY (deployment_id, scope_id) + REFERENCES scopes(deployment_id, scope_id) + ); + INSERT INTO scope_run_requests( + deployment_id, scope_id, run_request_id, turn_id, + interrupt_key, request_hash, run_owner_id, run_id, state, + created_at, updated_at + ) + SELECT deployment_id, scope_id, turn_id, turn_id, + NULL, request_hash, run_owner_id, run_id, state, + created_at, updated_at + FROM scope_turns; + CREATE INDEX scope_run_requests_turn + ON scope_run_requests(deployment_id, scope_id, turn_id); + CREATE UNIQUE INDEX scope_run_requests_interrupt_unique + ON scope_run_requests(deployment_id, scope_id, interrupt_key) + WHERE interrupt_key IS NOT NULL; + DROP TABLE scope_turns; + """ + ) + conn.execute("PRAGMA user_version = 2") + + @contextmanager + def _transaction(self) -> Iterator[sqlite3.Connection]: + self.initialize() + with self._connect() as conn: + conn.execute("BEGIN IMMEDIATE") + try: + yield conn + except Exception: + conn.rollback() + raise + else: + conn.commit() + + @staticmethod + def _scope_from_row(row: sqlite3.Row, owner_id: str) -> ScopeRecord: + return ScopeRecord( + deployment_id=str(row["deployment_id"]), + scope_id=str(row["scope_id"]), + primary_thread_id=str(row["primary_thread_id"]), + state=str(row["state"]), # type: ignore[arg-type] + revision=int(row["revision"]), + primary_owner_id=owner_id, + created_at=str(row["created_at"]), + updated_at=str(row["updated_at"]), + ) + + @staticmethod + def _owner_from_row(row: sqlite3.Row) -> OwnerRecord: + return OwnerRecord( + deployment_id=str(row["deployment_id"]), + owner_id=str(row["owner_id"]), + scope_id=str(row["scope_id"]), + owner_type=str(row["owner_type"]), + resource_id=(str(row["resource_id"]) if row["resource_id"] else None), + parent_owner_id=( + str(row["parent_owner_id"]) if row["parent_owner_id"] else None + ), + state=str(row["state"]), # type: ignore[arg-type] + created_at=str(row["created_at"]), + updated_at=str(row["updated_at"]), + ) + + @staticmethod + def _primary_owner( + conn: sqlite3.Connection, deployment_id: str, scope_id: str + ) -> str: + row = conn.execute( + """ + SELECT owner_id FROM scope_owners + WHERE deployment_id = ? AND scope_id = ? AND owner_type = 'primary_thread' + LIMIT 1 + """, + (deployment_id, scope_id), + ).fetchone() + if row is None: + raise ScopeRegistryError("scope has no primary owner") + return str(row["owner_id"]) + + @staticmethod + def _workspace_mutation_lock_active( + conn: sqlite3.Connection, + deployment_id: str, + operation_id: str | None = None, + ) -> bool: + row = conn.execute( + """ + SELECT 1 FROM deployment_locks + WHERE deployment_id = ? + AND lock_name IN ('workspace-cutover', 'workspace-lifecycle') + AND expires_at > ? + AND (? IS NULL OR operation_id != ?) + """, + (deployment_id, _utc_now(), operation_id, operation_id), + ).fetchone() + return row is not None + + def provision( + self, + deployment_id: str, + primary_thread_id: str, + *, + scope_id: str | None = None, + state: ScopeState = "draft", + operation_id: str | None = None, + lock_operation_id: str | None = None, + ) -> ScopeRecord: + """Reserve one immutable scope for a primary thread. + + Retrying the same primary thread returns its existing mapping. Passing a + different explicit scope for an existing thread is a conflict rather than + an opportunity to silently remap its files. + """ + + if state not in {"provisioning", "draft"}: + raise ScopeRegistryError("new scopes must start provisioning or draft") + requested_scope = ( + _ensure_uuid(scope_id, "scope_id") if scope_id else str(uuid.uuid4()) + ) + now = _utc_now() + operation_id = ( + _ensure_uuid(operation_id, "operation_id") + if operation_id + else str(uuid.uuid4()) + ) + if lock_operation_id is not None: + lock_operation_id = _ensure_uuid(lock_operation_id, "lock_operation_id") + with self._transaction() as conn: + existing = conn.execute( + """ + SELECT * FROM scopes WHERE deployment_id = ? AND primary_thread_id = ? + """, + (deployment_id, primary_thread_id), + ).fetchone() + if existing is not None: + if scope_id and str(existing["scope_id"]) != requested_scope: + raise ScopeConflictError("thread already belongs to another scope") + primary_owner = self._primary_owner( + conn, deployment_id, str(existing["scope_id"]) + ) + return self._scope_from_row(existing, primary_owner) + if self._workspace_mutation_lock_active( + conn, deployment_id, lock_operation_id + ): + raise ScopeAccessError( + "workspace cutover or maintenance is in progress" + ) + conn.execute( + """ + INSERT INTO scopes( + deployment_id, scope_id, primary_thread_id, state, revision, + created_at, updated_at + ) VALUES (?, ?, ?, ?, 1, ?, ?) + """, + (deployment_id, requested_scope, primary_thread_id, state, now, now), + ) + primary_owner = str(uuid.uuid4()) + conn.execute( + """ + INSERT INTO scope_owners( + deployment_id, owner_id, scope_id, owner_type, resource_id, + parent_owner_id, state, created_at, updated_at + ) VALUES (?, ?, ?, 'primary_thread', ?, NULL, 'active', ?, ?) + """, + ( + deployment_id, + primary_owner, + requested_scope, + primary_thread_id, + now, + now, + ), + ) + conn.execute( + """ + INSERT OR REPLACE INTO scope_operations( + deployment_id, operation_id, scope_id, kind, state, created_at, updated_at + ) VALUES (?, ?, ?, 'provision', 'completed', ?, ?) + """, + (deployment_id, operation_id, requested_scope, now, now), + ) + return ScopeRecord( + deployment_id=deployment_id, + scope_id=requested_scope, + primary_thread_id=primary_thread_id, + state=state, + revision=1, + primary_owner_id=primary_owner, + created_at=now, + updated_at=now, + ) + + def get_by_thread(self, deployment_id: str, thread_id: str) -> ScopeRecord: + self.initialize() + with self._connect() as conn: + row = conn.execute( + "SELECT * FROM scopes WHERE deployment_id = ? AND primary_thread_id = ?", + (deployment_id, thread_id), + ).fetchone() + if row is None: + raise ScopeNotFoundError("no workspace scope for thread") + return self._scope_from_row( + row, self._primary_owner(conn, deployment_id, str(row["scope_id"])) + ) + + def get(self, deployment_id: str, scope_id: str) -> ScopeRecord: + scope_id = _ensure_uuid(scope_id, "scope_id") + self.initialize() + with self._connect() as conn: + row = conn.execute( + "SELECT * FROM scopes WHERE deployment_id = ? AND scope_id = ?", + (deployment_id, scope_id), + ).fetchone() + if row is None: + raise ScopeNotFoundError("workspace scope not found") + return self._scope_from_row( + row, self._primary_owner(conn, deployment_id, scope_id) + ) + + def list_scopes( + self, deployment_id: str, *, state: ScopeState | None = None + ) -> list[ScopeRecord]: + """List deployment scopes for administrative maintenance only.""" + + if state is not None and state not in _SCOPE_TRANSITIONS: + raise ScopeRegistryError("invalid workspace scope state") + self.initialize() + with self._connect() as conn: + if state is None: + rows = conn.execute( + """ + SELECT * FROM scopes WHERE deployment_id = ? ORDER BY created_at ASC + """, + (deployment_id,), + ).fetchall() + else: + rows = conn.execute( + """ + SELECT * FROM scopes WHERE deployment_id = ? AND state = ? + ORDER BY created_at ASC + """, + (deployment_id, state), + ).fetchall() + return [ + self._scope_from_row( + row, + self._primary_owner(conn, deployment_id, str(row["scope_id"])), + ) + for row in rows + ] + + def transition_scope( + self, + deployment_id: str, + scope_id: str, + *, + expected_revision: int, + state: ScopeState, + ) -> ScopeRecord: + scope_id = _ensure_uuid(scope_id, "scope_id") + now = _utc_now() + with self._transaction() as conn: + current = conn.execute( + "SELECT * FROM scopes WHERE deployment_id = ? AND scope_id = ?", + (deployment_id, scope_id), + ).fetchone() + if current is None: + raise ScopeNotFoundError("workspace scope not found") + current_state = str(current["state"]) + if state not in _SCOPE_TRANSITIONS.get(current_state, frozenset()): + raise ScopeConflictError( + f"cannot transition scope from {current_state} to {state}" + ) + deleted_at = now if state == "deleted" else None + updated = conn.execute( + """ + UPDATE scopes + SET state = ?, revision = revision + 1, updated_at = ?, + deleted_at = COALESCE(?, deleted_at) + WHERE deployment_id = ? AND scope_id = ? AND revision = ? + """, + (state, now, deleted_at, deployment_id, scope_id, expected_revision), + ) + if updated.rowcount != 1: + raise ScopeConflictError("scope revision changed") + row = conn.execute( + "SELECT * FROM scopes WHERE deployment_id = ? AND scope_id = ?", + (deployment_id, scope_id), + ).fetchone() + assert row is not None + return self._scope_from_row( + row, self._primary_owner(conn, deployment_id, scope_id) + ) + + def register_owner( + self, + deployment_id: str, + scope_id: str, + *, + owner_type: str, + resource_id: str | None = None, + parent_owner_id: str | None = None, + owner_id: str | None = None, + state: OwnerState = "reserved", + ) -> OwnerRecord: + scope_id = _ensure_uuid(scope_id, "scope_id") + owner_id = _ensure_uuid(owner_id, "owner_id") if owner_id else str(uuid.uuid4()) + now = _utc_now() + with self._transaction() as conn: + scope = conn.execute( + "SELECT state FROM scopes WHERE deployment_id = ? AND scope_id = ?", + (deployment_id, scope_id), + ).fetchone() + if scope is None: + raise ScopeNotFoundError("workspace scope not found") + if self._workspace_mutation_lock_active(conn, deployment_id): + raise ScopeAccessError( + "workspace cutover or maintenance is in progress" + ) + if ( + str(scope["state"]) in _TERMINAL_SCOPE_STATES + or str(scope["state"]) == "deleting" + ): + raise ScopeAccessError("scope does not accept new owners") + if parent_owner_id: + parent = conn.execute( + """ + SELECT state FROM scope_owners + WHERE deployment_id = ? AND owner_id = ? AND scope_id = ? + """, + (deployment_id, parent_owner_id, scope_id), + ).fetchone() + if parent is None or str(parent["state"]) != "active": + raise ScopeAccessError("parent owner is not active") + try: + conn.execute( + """ + INSERT INTO scope_owners( + deployment_id, owner_id, scope_id, owner_type, resource_id, + parent_owner_id, state, created_at, updated_at + ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) + """, + ( + deployment_id, + owner_id, + scope_id, + owner_type, + resource_id, + parent_owner_id, + state, + now, + now, + ), + ) + except sqlite3.IntegrityError as exc: + raise ScopeConflictError( + "owner or external resource already exists" + ) from exc + return OwnerRecord( + deployment_id=deployment_id, + owner_id=owner_id, + scope_id=scope_id, + owner_type=owner_type, + resource_id=resource_id, + parent_owner_id=parent_owner_id, + state=state, + created_at=now, + updated_at=now, + ) + + def bind_owner( + self, + deployment_id: str, + scope_id: str, + owner_id: str, + resource_id: str, + *, + state: OwnerState = "active", + ) -> OwnerRecord: + scope_id = _ensure_uuid(scope_id, "scope_id") + owner_id = _ensure_uuid(owner_id, "owner_id") + now = _utc_now() + with self._transaction() as conn: + try: + result = conn.execute( + """ + UPDATE scope_owners + SET resource_id = ?, state = ?, updated_at = ? + WHERE deployment_id = ? AND scope_id = ? AND owner_id = ? + AND state NOT IN ('terminal', 'quarantined') + """, + (resource_id, state, now, deployment_id, scope_id, owner_id), + ) + except sqlite3.IntegrityError as exc: + raise ScopeConflictError( + "external resource belongs to another scope" + ) from exc + if result.rowcount != 1: + raise ScopeAccessError("owner cannot be bound") + row = conn.execute( + """ + SELECT * FROM scope_owners + WHERE deployment_id = ? AND scope_id = ? AND owner_id = ? + """, + (deployment_id, scope_id, owner_id), + ).fetchone() + assert row is not None + return self._owner_from_row(row) + + def assert_runtime( + self, + deployment_id: str, + scope_id: str, + thread_id: str, + owner_id: str, + ) -> ScopeRecord: + """Verify a graph/tool runtime is an active owner of this scope.""" + + scope_id = _ensure_uuid(scope_id, "workspace_scope_id") + owner_id = _ensure_uuid(owner_id, "workspace_scope_owner_id") + self.initialize() + with self._connect() as conn: + scope = conn.execute( + "SELECT * FROM scopes WHERE deployment_id = ? AND scope_id = ?", + (deployment_id, scope_id), + ).fetchone() + if scope is None: + raise ScopeNotFoundError("workspace scope not found") + if str(scope["state"]) not in {"draft", "active"}: + raise ScopeAccessError("workspace scope is not active") + if self._workspace_mutation_lock_active(conn, deployment_id): + raise ScopeAccessError( + "workspace cutover or maintenance is in progress" + ) + owner = conn.execute( + """ + SELECT * FROM scope_owners + WHERE deployment_id = ? AND scope_id = ? AND owner_id = ? + """, + (deployment_id, scope_id, owner_id), + ).fetchone() + if owner is None or ( + str(owner["state"]) != "active" + and not ( + str(owner["owner_type"]) == "primary_run" + and str(owner["state"]) == "reserved" + ) + ): + raise ScopeAccessError("workspace owner is not active") + if str(owner["owner_type"]) == "primary_thread": + if ( + str(scope["primary_thread_id"]) != thread_id + or str(owner["resource_id"]) != thread_id + ): + raise ScopeAccessError("primary thread does not own this scope") + elif str(owner["owner_type"]) == "primary_run": + if str(scope["primary_thread_id"]) != thread_id or str( + owner["parent_owner_id"] + ) != self._primary_owner(conn, deployment_id, scope_id): + raise ScopeAccessError("primary run does not own this scope") + elif str(owner["owner_type"]) != "schedule" and str( + owner["resource_id"] + ) not in {None, thread_id}: + raise ScopeAccessError("derived thread does not own this scope") + return self._scope_from_row( + scope, self._primary_owner(conn, deployment_id, scope_id) + ) + + def owners(self, deployment_id: str, scope_id: str) -> list[OwnerRecord]: + scope_id = _ensure_uuid(scope_id, "scope_id") + self.initialize() + with self._connect() as conn: + rows = conn.execute( + """ + SELECT * FROM scope_owners + WHERE deployment_id = ? AND scope_id = ? ORDER BY created_at ASC + """, + (deployment_id, scope_id), + ).fetchall() + return [self._owner_from_row(row) for row in rows] + + def get_owner_by_resource( + self, deployment_id: str, resource_id: str + ) -> OwnerRecord: + """Return the durable owner for one external child resource.""" + + self.initialize() + with self._connect() as conn: + row = conn.execute( + """ + SELECT * FROM scope_owners + WHERE deployment_id = ? AND resource_id = ? + """, + (deployment_id, resource_id), + ).fetchone() + if row is None: + raise ScopeNotFoundError("workspace owner not found") + return self._owner_from_row(row) + + def reserve_run( + self, + deployment_id: str, + scope_id: str, + run_request_id: str, + turn_id: str, + request_hash: str, + *, + interrupt_key: str | None = None, + ) -> RunReservation: + """Reserve one idempotent run request before creating it remotely. + + ``turn_id`` groups a user message and all of its approval resumes. + ``run_request_id`` identifies one external ``runs.create`` call, so a + lost response can be retried without turning a valid resume into a + conflict with the initial run. + """ + + scope_id = _ensure_uuid(scope_id, "scope_id") + run_request_id = _ensure_uuid(run_request_id, "run_request_id") + turn_id = _ensure_uuid(turn_id, "turn_id") + if interrupt_key is not None: + if ( + not isinstance(interrupt_key, str) + or not interrupt_key + or len(interrupt_key) > 256 + ): + raise ScopeRegistryError( + "interrupt_key must be a non-empty string up to 256 characters" + ) + now = _utc_now() + with self._transaction() as conn: + existing = conn.execute( + """ + SELECT * FROM scope_run_requests + WHERE deployment_id = ? AND scope_id = ? AND run_request_id = ? + """, + (deployment_id, scope_id, run_request_id), + ).fetchone() + if existing is not None: + existing_interrupt_key = ( + str(existing["interrupt_key"]) + if existing["interrupt_key"] + else None + ) + if ( + str(existing["request_hash"]) != request_hash + or str(existing["turn_id"]) != turn_id + or existing_interrupt_key != interrupt_key + ): + raise ScopeIdempotencyConflictError( + "run_request_id was reused with another request" + ) + return RunReservation( + deployment_id=deployment_id, + scope_id=scope_id, + run_request_id=str(existing["run_request_id"]), + turn_id=turn_id, + interrupt_key=existing_interrupt_key, + request_hash=request_hash, + run_owner_id=str(existing["run_owner_id"]), + run_id=str(existing["run_id"]) if existing["run_id"] else None, + state=str(existing["state"]), + ) + if interrupt_key is not None: + resolved_interrupt = conn.execute( + """ + SELECT run_request_id FROM scope_run_requests + WHERE deployment_id = ? AND scope_id = ? AND interrupt_key = ? + """, + (deployment_id, scope_id, interrupt_key), + ).fetchone() + if resolved_interrupt is not None: + raise ScopeInterruptResolvedError( + "interrupt already has a resume request" + ) + scope = conn.execute( + "SELECT state FROM scopes WHERE deployment_id = ? AND scope_id = ?", + (deployment_id, scope_id), + ).fetchone() + if scope is None or str(scope["state"]) not in {"draft", "active"}: + raise ScopeAccessError("workspace scope does not accept a run") + if self._workspace_mutation_lock_active(conn, deployment_id): + raise ScopeAccessError( + "workspace cutover or maintenance is in progress" + ) + parent_owner = self._primary_owner(conn, deployment_id, scope_id) + run_owner_id = str(uuid.uuid4()) + conn.execute( + """ + INSERT INTO scope_owners( + deployment_id, owner_id, scope_id, owner_type, resource_id, + parent_owner_id, state, created_at, updated_at + ) VALUES (?, ?, ?, 'primary_run', NULL, ?, 'reserved', ?, ?) + """, + (deployment_id, run_owner_id, scope_id, parent_owner, now, now), + ) + conn.execute( + """ + INSERT INTO scope_run_requests( + deployment_id, scope_id, run_request_id, turn_id, interrupt_key, + request_hash, run_owner_id, run_id, state, created_at, updated_at + ) VALUES (?, ?, ?, ?, ?, ?, ?, NULL, 'reserved', ?, ?) + """, + ( + deployment_id, + scope_id, + run_request_id, + turn_id, + interrupt_key, + request_hash, + run_owner_id, + now, + now, + ), + ) + return RunReservation( + deployment_id=deployment_id, + scope_id=scope_id, + run_request_id=run_request_id, + turn_id=turn_id, + interrupt_key=interrupt_key, + request_hash=request_hash, + run_owner_id=run_owner_id, + run_id=None, + state="reserved", + ) + + def bind_run( + self, + deployment_id: str, + scope_id: str, + run_request_id: str, + run_id: str, + ) -> RunReservation: + """Attach the remote run id to a prior reservation exactly once.""" + + scope_id = _ensure_uuid(scope_id, "scope_id") + run_request_id = _ensure_uuid(run_request_id, "run_request_id") + now = _utc_now() + with self._transaction() as conn: + row = conn.execute( + """ + SELECT * FROM scope_run_requests + WHERE deployment_id = ? AND scope_id = ? AND run_request_id = ? + """, + (deployment_id, scope_id, run_request_id), + ).fetchone() + if row is None: + raise ScopeNotFoundError("workspace run reservation not found") + if row["run_id"] and str(row["run_id"]) != run_id: + raise ScopeConflictError("run request is already bound to another run") + owner_id = str(row["run_owner_id"]) + conn.execute( + """ + UPDATE scope_owners + SET resource_id = ?, state = 'active', updated_at = ? + WHERE deployment_id = ? AND scope_id = ? AND owner_id = ? + AND state IN ('reserved', 'active') + """, + (run_id, now, deployment_id, scope_id, owner_id), + ) + conn.execute( + """ + UPDATE scope_run_requests SET run_id = ?, state = 'active', updated_at = ? + WHERE deployment_id = ? AND scope_id = ? AND run_request_id = ? + """, + (run_id, now, deployment_id, scope_id, run_request_id), + ) + return RunReservation( + deployment_id=deployment_id, + scope_id=scope_id, + run_request_id=run_request_id, + turn_id=str(row["turn_id"]), + interrupt_key=( + str(row["interrupt_key"]) if row["interrupt_key"] else None + ), + request_hash=str(row["request_hash"]), + run_owner_id=owner_id, + run_id=run_id, + state="active", + ) + + def reserve_turn( + self, + deployment_id: str, + scope_id: str, + turn_id: str, + request_hash: str, + ) -> RunReservation: + """Backward-compatible reservation for pre-run-request callers.""" + + return self.reserve_run( + deployment_id, + scope_id, + turn_id, + turn_id, + request_hash, + ) + + def bind_turn( + self, + deployment_id: str, + scope_id: str, + turn_id: str, + run_id: str, + ) -> RunReservation: + """Backward-compatible binding for pre-run-request callers.""" + + return self.bind_run(deployment_id, scope_id, turn_id, run_id) + + def acquire_lock( + self, + deployment_id: str, + lock_name: str, + operation_id: str, + *, + lease_seconds: int = 60, + ) -> DeploymentLock: + operation_id = _ensure_uuid(operation_id, "operation_id") + now = datetime.now(UTC) + expires = now + timedelta(seconds=max(1, lease_seconds)) + now_text, expires_text = now.isoformat(), expires.isoformat() + with self._transaction() as conn: + row = conn.execute( + """ + SELECT operation_id, expires_at FROM deployment_locks + WHERE deployment_id = ? AND lock_name = ? + """, + (deployment_id, lock_name), + ).fetchone() + if ( + row is not None + and str(row["operation_id"]) != operation_id + and str(row["expires_at"]) > now_text + ): + raise ScopeConflictError("deployment lock is held") + conn.execute( + """ + INSERT INTO deployment_locks( + deployment_id, lock_name, operation_id, expires_at, created_at + ) VALUES (?, ?, ?, ?, ?) + ON CONFLICT(deployment_id, lock_name) DO UPDATE SET + operation_id = excluded.operation_id, + expires_at = excluded.expires_at, + created_at = excluded.created_at + """, + (deployment_id, lock_name, operation_id, expires_text, now_text), + ) + return DeploymentLock(deployment_id, lock_name, operation_id, expires_text) + + def release_lock( + self, deployment_id: str, lock_name: str, operation_id: str + ) -> None: + with self._transaction() as conn: + result = conn.execute( + """ + DELETE FROM deployment_locks + WHERE deployment_id = ? AND lock_name = ? AND operation_id = ? + """, + (deployment_id, lock_name, operation_id), + ) + if result.rowcount != 1: + raise ScopeAccessError("deployment lock is not held by this operation") + + def renew_lock( + self, + deployment_id: str, + lock_name: str, + operation_id: str, + *, + lease_seconds: int = 60, + ) -> DeploymentLock: + """Renew a lease only while this operation still owns an active lock.""" + + operation_id = _ensure_uuid(operation_id, "operation_id") + now = datetime.now(UTC) + now_text = now.isoformat() + expires_text = (now + timedelta(seconds=max(1, lease_seconds))).isoformat() + with self._transaction() as conn: + result = conn.execute( + """ + UPDATE deployment_locks SET expires_at = ? + WHERE deployment_id = ? AND lock_name = ? AND operation_id = ? + AND expires_at > ? + """, + (expires_text, deployment_id, lock_name, operation_id, now_text), + ) + if result.rowcount != 1: + raise ScopeAccessError( + "deployment lock expired or belongs to another operation" + ) + return DeploymentLock(deployment_id, lock_name, operation_id, expires_text) + + def active_lock(self, deployment_id: str, lock_name: str) -> DeploymentLock | None: + """Return an unexpired deployment lock without mutating its lease.""" + + self.initialize() + now = _utc_now() + with self._connect() as conn: + row = conn.execute( + """ + SELECT operation_id, expires_at FROM deployment_locks + WHERE deployment_id = ? AND lock_name = ? AND expires_at > ? + """, + (deployment_id, lock_name, now), + ).fetchone() + if row is None: + return None + return DeploymentLock( + deployment_id=deployment_id, + lock_name=lock_name, + operation_id=str(row["operation_id"]), + expires_at=str(row["expires_at"]), + ) + + def begin_operation( + self, + deployment_id: str, + operation_id: str, + *, + kind: str, + scope_id: str | None = None, + ) -> None: + operation_id = _ensure_uuid(operation_id, "operation_id") + if scope_id is not None: + scope_id = _ensure_uuid(scope_id, "scope_id") + now = _utc_now() + with self._transaction() as conn: + existing = conn.execute( + """ + SELECT kind, state FROM scope_operations + WHERE deployment_id = ? AND operation_id = ? + """, + (deployment_id, operation_id), + ).fetchone() + if existing is not None: + if str(existing["kind"]) != kind: + raise ScopeConflictError("operation id belongs to another kind") + return + conn.execute( + """ + INSERT INTO scope_operations( + deployment_id, operation_id, scope_id, kind, state, created_at, updated_at + ) VALUES (?, ?, ?, ?, 'running', ?, ?) + """, + (deployment_id, operation_id, scope_id, kind, now, now), + ) + + def finish_operation( + self, + deployment_id: str, + operation_id: str, + *, + state: str, + result_sha256: str | None = None, + last_error_code: str | None = None, + ) -> None: + if state not in {"completed", "failed"}: + raise ScopeRegistryError( + "operation terminal state must be completed or failed" + ) + operation_id = _ensure_uuid(operation_id, "operation_id") + with self._transaction() as conn: + result = conn.execute( + """ + UPDATE scope_operations + SET state = ?, result_sha256 = ?, last_error_code = ?, updated_at = ? + WHERE deployment_id = ? AND operation_id = ? AND state = 'running' + """, + ( + state, + result_sha256, + last_error_code, + _utc_now(), + deployment_id, + operation_id, + ), + ) + if result.rowcount != 1: + raise ScopeAccessError("operation is not running") + + def get_operation(self, deployment_id: str, operation_id: str) -> ScopeOperation: + operation_id = _ensure_uuid(operation_id, "operation_id") + self.initialize() + with self._connect() as conn: + row = conn.execute( + """ + SELECT scope_id, kind, state, result_sha256, last_error_code + FROM scope_operations WHERE deployment_id = ? AND operation_id = ? + """, + (deployment_id, operation_id), + ).fetchone() + if row is None: + raise ScopeNotFoundError("scope operation not found") + return ScopeOperation( + deployment_id=deployment_id, + operation_id=operation_id, + scope_id=str(row["scope_id"]) if row["scope_id"] else None, + kind=str(row["kind"]), + state=str(row["state"]), + result_sha256=( + str(row["result_sha256"]) if row["result_sha256"] else None + ), + last_error_code=( + str(row["last_error_code"]) if row["last_error_code"] else None + ), + ) + + +_registry_cache: dict[Path, ScopeRegistry] = {} +_registry_cache_lock = threading.Lock() + + +def get_scope_registry(workspace_root: Path | str | None = None) -> ScopeRegistry: + path = default_registry_path(workspace_root) + with _registry_cache_lock: + registry = _registry_cache.get(path) + if registry is None: + registry = ScopeRegistry(path) + _registry_cache[path] = registry + registry.initialize() + return registry diff --git a/EvoScientist/sessions.py b/EvoScientist/sessions.py index b968e58..122d56d 100644 --- a/EvoScientist/sessions.py +++ b/EvoScientist/sessions.py @@ -29,8 +29,10 @@ WebUI / langgraph-dev checkpointer: import asyncio import atexit +import hashlib import logging import math +import time import uuid from collections.abc import AsyncIterator, Awaitable, Callable from contextlib import asynccontextmanager @@ -76,6 +78,19 @@ MAIN_THREAD_FILTER_SQL = ( " OR json_extract(metadata, '$.graph_id') = ?)" ) MAIN_THREAD_FILTER_PARAMS = (AGENT_NAME, AGENT_NAME) +_CHECKPOINT_MSGPACK_MODULES = frozenset( + { + ("EvoScientist.llm.errors", "AgentControlError"), + ("EvoScientist.llm.errors", "ModelToolProtocolError"), + ("EvoScientist.llm.errors", "ProviderStreamError"), + } +) + + +def _checkpoint_serde() -> JsonPlusSerializer: + """Return the checkpoint serializer with app-owned types allowlisted.""" + + return JsonPlusSerializer(allowed_msgpack_modules=_CHECKPOINT_MSGPACK_MODULES) # --------------------------------------------------------------------------- @@ -182,7 +197,9 @@ class PruningCheckpointer(AsyncSqliteSaver): keep_per_ns: int = _DEFAULT_KEEP_PER_NS, serde: Any = None, ) -> None: - super().__init__(conn, serde=serde) + super().__init__( + conn, serde=serde if serde is not None else _checkpoint_serde() + ) self._keep_per_ns = max(0, int(keep_per_ns)) # Outer lock guarantees ``super().aput()`` and ``_prune_after_put()`` # are atomic *as a pair*. Without this, a concurrent ``aput()`` on a @@ -321,12 +338,7 @@ class PruningCheckpointer(AsyncSqliteSaver): checkpoint_ns: str, oldest_anchor_id: str, ) -> set[str]: - """Walk parent chain until hitting a ``messages`` seed. - - Returns the set of ancestor ids to preserve (inclusive of the - snapshot ancestor). On chain-break or deserialization failure, - returns what was visited so far — the safe side is over-preserve. - """ + """Walk the parent chain until a messages seed is found.""" extra: set[str] = set() cursor = await self._fetch_parent_checkpoint_id( thread_id, checkpoint_ns, oldest_anchor_id @@ -336,7 +348,7 @@ class PruningCheckpointer(AsyncSqliteSaver): steps += 1 blob = await self._fetch_checkpoint_blob(thread_id, checkpoint_ns, cursor) if blob is None: - break # chain broken (legacy DB); preserve what we have + break extra.add(cursor) try: ck = self.serde.loads_typed(blob) @@ -348,10 +360,10 @@ class PruningCheckpointer(AsyncSqliteSaver): thread_id, exc, ) - break # safe-side: preserve everything visited so far + break cv = ck.get("channel_values") or {} if _unwrap_messages_seed(cv.get("messages")) is not None: - break # found seed; this ancestor anchors reconstruction + break cursor = await self._fetch_parent_checkpoint_id( thread_id, checkpoint_ns, cursor ) @@ -392,25 +404,12 @@ class PruningCheckpointer(AsyncSqliteSaver): agent: str, kept_ids: set[str], ) -> None: - """DELETE rows whose ``checkpoint_id`` is NOT in ``kept_ids``. - - Writes deleted first to preserve referential ordering — if we - dropped checkpoints first, surviving writes' ``checkpoint_id`` - would become orphans. - - Empty ``kept_ids`` is a no-op rather than "delete everything" — - a defensive check; the caller always passes anchor_ids which is - non-empty by construction (already checked ``len >= keep`` in - the caller). - """ + """Delete checkpoints outside the retained head and seed chain.""" if not kept_ids: return kept_list = list(kept_ids) placeholders = ",".join("?" * len(kept_list)) - # Writes DELETE only runs if the ``writes`` table exists. Legacy - # DBs from pre-DeltaChannel builds may have only ``checkpoints`` — - # we still want to prune those, just skipping the writes step. if await _table_exists(self.conn, "writes"): del_writes = ( "DELETE FROM writes " @@ -446,6 +445,226 @@ class PruningCheckpointer(AsyncSqliteSaver): ) +@dataclass(frozen=True, slots=True) +class TurnLease: + thread_id: str + owner_id: str + fencing_token: int + expires_at_ms: int + checkpoint_snapshot_id: str + + +class FencedPruningCheckpointer(PruningCheckpointer): + """Single-worker linearizable turn lease around LangGraph checkpoint writes.""" + + def __init__(self, *args: Any, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self._turn_fence_lock = asyncio.Lock() + + @classmethod + @asynccontextmanager + async def from_conn_string_with_keep( + cls, conn_string: str, keep_per_ns: int = _DEFAULT_KEEP_PER_NS + ) -> AsyncIterator["FencedPruningCheckpointer"]: + async with aiosqlite.connect(conn_string) as conn: + saver = cls(conn, keep_per_ns=keep_per_ns) + await saver.setup_fencing() + yield saver + + async def setup_fencing(self) -> None: + await super().setup() + async with self.lock: + await self.conn.executescript( + """ + CREATE TABLE IF NOT EXISTS thread_turn_fences ( + thread_id TEXT PRIMARY KEY, + generation INTEGER NOT NULL CHECK (generation > 0), + owner_id TEXT, + expires_at_ms INTEGER, + updated_at_ms INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS thread_checkpoint_versions ( + thread_id TEXT PRIMARY KEY, + sequence INTEGER NOT NULL DEFAULT 0 CHECK (sequence >= 0), + checkpoint_id TEXT, + updated_at_ms INTEGER NOT NULL + ); + """ + ) + await self.conn.commit() + + async def acquire_turn_lease( + self, + thread_id: str, + owner_id: str, + *, + ttl_seconds: int, + ) -> TurnLease: + clean_thread = str(thread_id or "").strip() + clean_owner = str(owner_id or "").strip() + if not clean_thread or not clean_owner or ttl_seconds < 1: + raise ValueError("thread, owner and positive TTL are required") + async with self._turn_fence_lock, self.lock: + now = time.time_ns() // 1_000_000 + await self.conn.execute("BEGIN IMMEDIATE") + try: + row = await ( + await self.conn.execute( + "SELECT generation, owner_id, expires_at_ms FROM thread_turn_fences WHERE thread_id=?", + (clean_thread,), + ) + ).fetchone() + if ( + row is not None + and row[1] + and row[1] != clean_owner + and int(row[2] or 0) >= now + ): + raise RuntimeError("TURN_LEASE_BUSY") + generation = int(row[0]) + 1 if row is not None else 1 + expires_at = now + ttl_seconds * 1000 + await self.conn.execute( + """INSERT INTO thread_turn_fences + (thread_id, generation, owner_id, expires_at_ms, updated_at_ms) + VALUES (?, ?, ?, ?, ?) + ON CONFLICT(thread_id) DO UPDATE SET + generation=excluded.generation, + owner_id=excluded.owner_id, + expires_at_ms=excluded.expires_at_ms, + updated_at_ms=excluded.updated_at_ms""", + (clean_thread, generation, clean_owner, expires_at, now), + ) + version = await ( + await self.conn.execute( + "SELECT sequence, checkpoint_id FROM thread_checkpoint_versions WHERE thread_id=?", + (clean_thread,), + ) + ).fetchone() + sequence = int(version[0]) if version is not None else 0 + checkpoint_id = ( + str(version[1] or "root") if version is not None else "root" + ) + if version is None: + await self.conn.execute( + """INSERT INTO thread_checkpoint_versions + (thread_id, sequence, checkpoint_id, updated_at_ms) + VALUES (?, 0, NULL, ?)""", + (clean_thread, now), + ) + await self.conn.commit() + except Exception: + await self.conn.rollback() + raise + snapshot = hashlib.sha256( + f"{clean_thread}\0{sequence}\0{checkpoint_id}".encode() + ).hexdigest() + return TurnLease( + clean_thread, clean_owner, generation, expires_at, f"sha256:{snapshot}" + ) + + async def renew_turn_lease( + self, lease: TurnLease, *, ttl_seconds: int + ) -> TurnLease: + async with self._turn_fence_lock, self.lock: + now = time.time_ns() // 1_000_000 + expires_at = now + ttl_seconds * 1000 + cursor = await self.conn.execute( + """UPDATE thread_turn_fences + SET expires_at_ms=?, updated_at_ms=? + WHERE thread_id=? AND generation=? AND owner_id=? + AND expires_at_ms>=?""", + ( + expires_at, + now, + lease.thread_id, + lease.fencing_token, + lease.owner_id, + now, + ), + ) + await self.conn.commit() + if cursor.rowcount != 1: + raise RuntimeError("TURN_LEASE_LOST") + return TurnLease( + lease.thread_id, + lease.owner_id, + lease.fencing_token, + expires_at, + lease.checkpoint_snapshot_id, + ) + + async def release_turn_lease(self, lease: TurnLease) -> bool: + async with self._turn_fence_lock, self.lock: + now = time.time_ns() // 1_000_000 + cursor = await self.conn.execute( + """UPDATE thread_turn_fences + SET owner_id=NULL, expires_at_ms=NULL, updated_at_ms=? + WHERE thread_id=? AND generation=? AND owner_id=?""", + (now, lease.thread_id, lease.fencing_token, lease.owner_id), + ) + await self.conn.commit() + return cursor.rowcount == 1 + + async def _require_write_lease(self, config: Any) -> tuple[str, int]: + configurable = dict(config.get("configurable") or {}) + thread_id = str(configurable.get("thread_id") or "") + if not thread_id.startswith("web:"): + return thread_id, 0 + owner_id = str(configurable.get("turn_lease_owner") or "") + token = int(configurable.get("turn_fencing_token") or 0) + now = time.time_ns() // 1_000_000 + row = await ( + await self.conn.execute( + """SELECT 1 FROM thread_turn_fences + WHERE thread_id=? AND generation=? AND owner_id=? + AND expires_at_ms>=?""", + (thread_id, token, owner_id, now), + ) + ).fetchone() + if row is None: + raise RuntimeError("TURN_FENCED") + return thread_id, token + + async def aput( + self, config: Any, checkpoint: Any, metadata: Any, new_versions: Any + ) -> Any: + async with self._turn_fence_lock: + async with self.lock: + thread_id, token = await self._require_write_lease(config) + result = await super().aput(config, checkpoint, metadata, new_versions) + if token == 0: + return result + checkpoint_id = str( + result.get("configurable", {}).get("checkpoint_id") or "" + ) + async with self.lock: + now = time.time_ns() // 1_000_000 + await self.conn.execute( + """INSERT INTO thread_checkpoint_versions + (thread_id, sequence, checkpoint_id, updated_at_ms) + VALUES (?, 1, ?, ?) + ON CONFLICT(thread_id) DO UPDATE SET + sequence=thread_checkpoint_versions.sequence+1, + checkpoint_id=excluded.checkpoint_id, + updated_at_ms=excluded.updated_at_ms""", + (thread_id, checkpoint_id, now), + ) + await self.conn.commit() + return result + + async def aput_writes( + self, + config: Any, + writes: Any, + task_id: str, + task_path: str = "", + ) -> None: + async with self._turn_fence_lock: + async with self.lock: + await self._require_write_lease(config) + await super().aput_writes(config, writes, task_id, task_path) + + # --------------------------------------------------------------------------- # Checkpointer context manager # --------------------------------------------------------------------------- @@ -462,7 +681,7 @@ def _resolve_keep_per_ns() -> int: @asynccontextmanager -async def get_checkpointer() -> AsyncIterator[PruningCheckpointer]: +async def get_checkpointer() -> AsyncIterator[FencedPruningCheckpointer]: """Yield a pruning-enabled checkpointer connected to the sessions DB. Wraps ``AsyncSqliteSaver`` with ``PruningCheckpointer`` so every @@ -479,7 +698,7 @@ async def get_checkpointer() -> AsyncIterator[PruningCheckpointer]: On failure ``user_version`` is NOT bumped, so the next launch retries. """ keep = _resolve_keep_per_ns() - async with PruningCheckpointer.from_conn_string_with_keep( + async with FencedPruningCheckpointer.from_conn_string_with_keep( str(get_db_path()), keep_per_ns=keep ) as saver: # The whole gate is wrapped in a broad try/except: any unexpected @@ -919,7 +1138,7 @@ async def list_threads( if (include_message_count or include_preview) and threads: # Share one saver across all threads so ``setup()`` runs once. - serde = JsonPlusSerializer() + serde = _checkpoint_serde() saver = AsyncSqliteSaver(conn, serde=serde) for t in threads: msgs = await _load_checkpoint_messages(saver, t["thread_id"]) @@ -1086,7 +1305,7 @@ async def get_thread_messages(thread_id: str) -> list: async with conn.execute(check, (thread_id, *MAIN_THREAD_FILTER_PARAMS)) as cur: if not await cur.fetchone(): return [] - serde = JsonPlusSerializer() + serde = _checkpoint_serde() saver = AsyncSqliteSaver(conn, serde=serde) return await _load_checkpoint_messages(saver, thread_id) @@ -1646,7 +1865,7 @@ async def _restore_webui_threads_to_global_store() -> None: # message (stubs carry values=None, so the WebUI would otherwise # render every restored thread as "Untitled Thread"). if sqlite_data: - saver = AsyncSqliteSaver(conn, serde=JsonPlusSerializer()) + saver = AsyncSqliteSaver(conn, serde=_checkpoint_serde()) for thread_uuid in sqlite_data: try: msgs = await _load_checkpoint_messages(saver, str(thread_uuid)) diff --git a/EvoScientist/stream/events.py b/EvoScientist/stream/events.py index 2e63eef..410505e 100644 --- a/EvoScientist/stream/events.py +++ b/EvoScientist/stream/events.py @@ -541,7 +541,10 @@ class _V3EventProcessor: ) -> str: matches = [ call_id - for (candidate_scope, call_id), candidate in self._pending_tool_calls.items() + for ( + candidate_scope, + call_id, + ), candidate in self._pending_tool_calls.items() if candidate_scope == scope and candidate == (name, args) ] return matches[0] if len(matches) == 1 else "" @@ -849,7 +852,9 @@ class _V3EventProcessor: call_id = str( request_map.get("id") or request_map.get("tool_call_id") or "" ) - name = str(request_map.get("name") or request_map.get("tool_name") or "") + name = str( + request_map.get("name") or request_map.get("tool_name") or "" + ) args_map = _as_raw_map( request_map.get("args") if "args" in request_map @@ -1007,6 +1012,7 @@ async def stream_agent_events( metadata: dict[str, Any] | None = None, media: list[str] | None = None, callbacks: list[Any] | None = None, + configurable: dict[str, Any] | None = None, error_mode: str = "emit", ) -> AsyncGenerator[dict[str, Any], None]: """Stream events from a DeepAgents/LangGraph v3 run. @@ -1032,7 +1038,9 @@ async def stream_agent_events( subagent_start, subagent_tool_call, subagent_tool_result, subagent_end, done, error """ - config: dict[str, Any] = {"configurable": {"thread_id": thread_id}} + config_values = dict(configurable or {}) + config_values["thread_id"] = thread_id + config: dict[str, Any] = {"configurable": config_values} if metadata: config["metadata"] = metadata if callbacks: diff --git a/EvoScientist/subagents/_factory.py b/EvoScientist/subagents/_factory.py index 21aeaca..0c29746 100644 --- a/EvoScientist/subagents/_factory.py +++ b/EvoScientist/subagents/_factory.py @@ -47,6 +47,7 @@ def build_async_subagent_graph(name: str) -> Any: _get_default_middleware, _inject_subagent_middleware, ) + from EvoScientist.middleware import BudgetedSkillsMiddleware from EvoScientist.tools import skill_manager, tavily_search, think_tool from EvoScientist.utils import load_subagents @@ -113,13 +114,19 @@ def build_async_subagent_graph(name: str) -> Any: _ensure_auxiliary_chat_model() if name == "scheduler" else _ensure_chat_model() ) + backend = _get_default_backend() + if skill_sources := spec.get("skills"): + middleware.append( + BudgetedSkillsMiddleware(backend=backend, sources=skill_sources) + ) + return create_deep_agent( name=name, model=model, system_prompt=spec.get("system_prompt", ""), tools=spec.get("tools", []) + agent_mcp_tools, - skills=spec.get("skills"), - backend=_get_default_backend(), + skills=None, + backend=backend, middleware=middleware, subagents=subagents, ).with_config({"recursion_limit": cfg.recursion_limit}) diff --git a/EvoScientist/web_runtime.py b/EvoScientist/web_runtime.py new file mode 100644 index 0000000..5cf01dd --- /dev/null +++ b/EvoScientist/web_runtime.py @@ -0,0 +1,170 @@ +"""Web-specific agent construction owned by EvoScientist. + +Gateway supplies a workspace/checkpointer host context only. Model routes, +provider clients, and the profile semantics remain entirely inside the Evo +runtime. +""" + +from __future__ import annotations + +import copy +import hashlib +import json +import os +from collections.abc import Awaitable, Callable +from typing import Any + +from langchain.agents.middleware.types import AgentMiddleware, ToolCallRequest +from langchain_core.messages import ToolMessage +from langgraph.types import Command + +from .config.settings import load_config +from .llm.contracts import ( + AgentExecutionProfile, + AgentModelSet, + EvoRuntimeError, + WebHostContext, +) + + +def web_tool_registry_manifest() -> tuple[tuple[dict[str, Any], ...], str]: + """Return the current bounded Web profile and main-agent MCP tools.""" + + names = ( + "think_tool", + "execute", + "read_file", + "write_file", + "edit_file", + "ls", + "glob", + "grep", + "write_todos", + "web_search", + "parse_documents", + "use_skill", + ) + if os.environ.get("TAVILY_API_KEY"): + names = (*names, "tavily_search") + schema = { + "type": "object", + "additionalProperties": True, + "maxProperties": 32, + } + manifest: tuple[dict[str, Any], ...] = tuple( + { + "name": name, + "description": "EvoScientist Web runtime tool", + "schema": schema, + } + for name in names + ) + from .EvoScientist import _load_mcp_config_once, _load_mcp_tools_cached + + mcp_tools = _load_mcp_tools_cached().get("main", []) + dynamic = [] + for tool in mcp_tools: + args_schema = getattr(tool, "args_schema", None) + if hasattr(args_schema, "model_json_schema"): + args_schema = args_schema.model_json_schema() + dynamic.append( + { + "name": str(getattr(tool, "name", type(tool).__name__)), + "description": str(getattr(tool, "description", "")), + "schema": args_schema or {}, + } + ) + static_names = {item["name"] for item in manifest} + dynamic_names = [item["name"] for item in dynamic] + if static_names & set(dynamic_names) or len(dynamic_names) != len( + set(dynamic_names) + ): + raise EvoRuntimeError("TOOL_REGISTRY_CONFLICT") + manifest = tuple(sorted((*manifest, *dynamic), key=lambda item: item["name"])) + mcp_config_signature, _mcp_config = _load_mcp_config_once() + mcp_config_revision = hashlib.sha256( + mcp_config_signature.encode("utf-8") + ).hexdigest() + encoded = json.dumps( + { + "tools": manifest, + "mcp_config_revision": mcp_config_revision, + "builtin_revision": "evoscientist-web-tools-v1", + }, + sort_keys=True, + separators=(",", ":"), + ).encode() + return manifest, f"sha256:{hashlib.sha256(encoded).hexdigest()}" + + +class _ToolRegistryFenceMiddleware(AgentMiddleware): + name = "web_tool_registry_fence" + + def __init__(self, expected_revision: str) -> None: + super().__init__() + self.expected_revision = expected_revision + + def _require_current(self) -> None: + _manifest, revision = web_tool_registry_manifest() + if revision != self.expected_revision: + raise EvoRuntimeError("TOOL_REGISTRY_STALE") + + def wrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[[ToolCallRequest], ToolMessage | Command[Any]], + ) -> ToolMessage | Command[Any]: + self._require_current() + return handler(request) + + async def awrap_tool_call( + self, + request: ToolCallRequest, + handler: Callable[ + [ToolCallRequest], Awaitable[ToolMessage | Command[Any]] + ], + ) -> ToolMessage | Command[Any]: + self._require_current() + return await handler(request) + + +def create_web_agent( + *, snapshot: Any, host: WebHostContext, model_set: AgentModelSet +) -> Any: + """Create a `web_v3` agent without importing Gateway types.""" + + from .EvoScientist import create_cli_agent + from .middleware.evo_route_fallback import EvoRouteFallbackMiddleware + + config = copy.copy(load_config()) + config.auto_approve = True + config.auto_mode = True + config.enable_ask_user = False + config.enable_async_subagents = False + config.enable_scheduler = False + config.memory_workers_enabled = False + route_middleware = EvoRouteFallbackMiddleware( + model_set.main_fallbacks, + route_health=model_set.route_health, + capacity=model_set.capacity, + ) + tool_registry_fence = _ToolRegistryFenceMiddleware( + host.tool_registry_revision + ) + return create_cli_agent( + workspace_dir=host.workspace_dir, + memory_dir=host.memory_dir, + workspace_backend=host.workspace_backend, + checkpointer=host.checkpointer, + config=config, + chat_model=model_set.main_agent, + on_mcp_progress=host.on_mcp_progress, + tool_selector_threshold=host.tool_selector_threshold, + memory_max_inline_profile_chars=host.memory_max_inline_profile_chars, + enable_subagents=False, + enable_background_execution=False, + main_agent_outer_middlewares=[tool_registry_fence], + main_agent_route_middleware=route_middleware, + execution_profile=AgentExecutionProfile.web_v3(), + agent_model_set=model_set, + ) diff --git a/EvoScientist/workspace_scope.py b/EvoScientist/workspace_scope.py new file mode 100644 index 0000000..95d7fc9 --- /dev/null +++ b/EvoScientist/workspace_scope.py @@ -0,0 +1,541 @@ +"""Conversation-scoped workspace resolution and DeepAgents backend factory.""" + +from __future__ import annotations + +import os +import shutil +import subprocess +import threading +import uuid +from collections.abc import Callable +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +from deepagents.backends.protocol import ( + EditResult, + ExecuteResponse, + FileDownloadResponse, + FileUploadResponse, + GlobResult, + GrepResult, + LsResult, + ReadResult, + SandboxBackendProtocol, + WriteResult, +) +from langchain.tools import ToolRuntime + +from . import paths +from .scope_registry import ( + ScopeAccessError, + ScopeRecord, + deployment_id_for_workspace, + get_scope_registry, +) + +IsolationMode = str + + +def workspace_isolation_mode() -> IsolationMode: + value = os.getenv("EVOSCIENTIST_WORKSPACE_ISOLATION", "optional").strip().lower() + if value not in {"legacy", "optional", "required"}: + raise RuntimeError("EVOSCIENTIST_WORKSPACE_ISOLATION must be legacy, optional or required") + return value + + +def is_required() -> bool: + return workspace_isolation_mode() == "required" + + +def verify_required_executor() -> None: + """Fail startup unless the pinned scope executor is locally usable.""" + + docker = shutil.which("docker") + if not docker: + raise RuntimeError("required workspace isolation needs the docker OCI runtime") + image = os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_IMAGE", "").strip() + if "@sha256:" not in image: + raise RuntimeError("required workspace isolation needs an OCI image pinned by digest") + try: + probe = subprocess.run( + [docker, "image", "inspect", image], + check=False, + capture_output=True, + text=True, + timeout=10, + ) + except (OSError, subprocess.TimeoutExpired) as exc: + raise RuntimeError("required workspace isolation cannot verify the OCI executor") from exc + if probe.returncode != 0: + raise RuntimeError( + f"required workspace isolation needs local OCI image {image!r}" + ) + + +def current_deployment_id() -> str: + return deployment_id_for_workspace(paths.WORKSPACE_ROOT) + + +def conversation_root(scope_id: str, workspace_root: Path | None = None) -> Path: + scope = str(uuid.UUID(scope_id)) + # The deploy process supplies an absolute workspace root. This helper is + # called from synchronous DeepAgents backend factories on the ASGI loop. + root = (workspace_root or paths.WORKSPACE_ROOT).expanduser() + return root / ".evoscientist" / "conversations" / scope + + +def conversation_files_dir(scope_id: str, workspace_root: Path | None = None) -> Path: + return conversation_root(scope_id, workspace_root) / "files" + + +@dataclass(frozen=True, slots=True) +class ScopeContext: + deployment_id: str + scope_id: str + owner_id: str + thread_id: str + revision: int + files_dir: Path + runtime_dir: Path + primary_thread_id: str + + +@dataclass(frozen=True, slots=True) +class _RuntimeScopeConfig: + """Untrusted runtime identifiers parsed without filesystem or Registry I/O.""" + + scope_id: str + owner_id: str + thread_id: str + deployment_id: str | None + + +class ScopedContainerBackend: + """Filesystem backend whose shell commands execute in a scope-only OCI container.""" + + def __init__(self, root_dir: Path, *, timeout: int) -> None: + from .backends import CustomSandboxBackend + + # Reuse the hardened filesystem operations; only ``execute`` is + # replaced so no agent shell runs in the host process. + self._filesystem = CustomSandboxBackend( + root_dir=str(root_dir), virtual_mode=True, timeout=timeout, dangerous=False + ) + self._root_dir = root_dir + self._timeout = timeout + + def __getattr__(self, name: str) -> Any: + return getattr(self._filesystem, name) + + def execute(self, command: str, *, timeout: int | None = None) -> Any: + from .backends import ExecuteResponse, prepare_sandbox_command + + command, error = prepare_sandbox_command( + command, self._filesystem.cwd, virtual_mode=True, dangerous=False + ) + if error: + return ExecuteResponse(output=error, exit_code=1, truncated=False) + image = os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_IMAGE", "").strip() + if "@sha256:" not in image: + return ExecuteResponse( + output="Required workspace isolation needs an OCI image pinned by digest.", + exit_code=125, + truncated=False, + ) + effective_timeout = max(1, min(timeout or self._timeout, 3600)) + invocation = [ + "docker", + "run", + "--rm", + "--network", + "none", + "--read-only", + "--tmpfs", + "/tmp:rw,noexec,nosuid,size=64m", + "--cap-drop", + "ALL", + "--security-opt", + "no-new-privileges", + "--pids-limit", + os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_PIDS", "128"), + "--memory", + os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_MEMORY", "1g"), + "--cpus", + os.getenv("EVOSCIENTIST_STRICT_EXECUTOR_CPUS", "1"), + "--mount", + f"type=bind,src={self._root_dir},dst=/workspace", + "--workdir", + "/workspace", + image, + "sh", + "-lc", + command, + ] + try: + completed = subprocess.run( + invocation, + check=False, + capture_output=True, + text=True, + timeout=effective_timeout, + ) + except FileNotFoundError: + return ExecuteResponse( + output="Required workspace isolation needs an OCI runtime (docker was not found).", + exit_code=127, + truncated=False, + ) + except subprocess.TimeoutExpired as exc: + output = (exc.stdout or "") + (exc.stderr or "") + return ExecuteResponse(output=output, exit_code=124, truncated=False) + output = completed.stdout + completed.stderr + return ExecuteResponse(output=output, exit_code=completed.returncode, truncated=False) + + +def _configurable(runtime: ToolRuntime[Any, Any] | Any | None) -> dict[str, Any]: + """Return the active runnable config, with a non-graph fallback. + + ``ToolRuntime`` deliberately does not expose ``RunnableConfig`` during a + graph execution. LangGraph keeps it in a context variable instead. The + fallback preserves direct callers and unit tests that supply a lightweight + runtime object outside a runnable context. + """ + + config: Any = None + try: + from langgraph.config import get_config + + config = get_config() + except (ImportError, LookupError, RuntimeError): + pass + if not isinstance(config, dict) and runtime is not None: + config = getattr(runtime, "config", None) or {} + if not isinstance(config, dict): + return {} + configurable = config.get("configurable") or {} + return dict(configurable) if isinstance(configurable, dict) else {} + + +def _required_string(configurable: dict[str, Any], key: str) -> str: + value = configurable.get(key) + if not isinstance(value, str) or not value: + raise ScopeAccessError(f"missing {key}") + return value + + +def _runtime_scope_config( + runtime: ToolRuntime[Any, Any] | Any | None, + *, + kind: str, +) -> _RuntimeScopeConfig | None: + """Parse scope identifiers without treating config as an authorization grant.""" + + configurable = _configurable(runtime) + scope_id = configurable.get("workspace_scope_id") + owner_id = configurable.get("workspace_scope_owner_id") + thread_id = configurable.get("thread_id") + + if scope_id is None and owner_id is None: + if workspace_isolation_mode() == "required": + raise ScopeAccessError(f"{kind} requires a workspace scope") + return None + if ( + not isinstance(scope_id, str) + or not isinstance(owner_id, str) + or not isinstance(thread_id, str) + ): + raise ScopeAccessError(f"{kind} has an incomplete workspace scope") + try: + canonical_scope_id = str(uuid.UUID(scope_id)) + canonical_owner_id = str(uuid.UUID(owner_id)) + except ValueError as exc: + raise ScopeAccessError(f"{kind} has an invalid workspace scope") from exc + deployment_id = configurable.get("workspace_deployment_id") + if deployment_id is not None and ( + not isinstance(deployment_id, str) or not deployment_id + ): + raise ScopeAccessError(f"{kind} has an invalid workspace deployment") + return _RuntimeScopeConfig( + scope_id=canonical_scope_id, + owner_id=canonical_owner_id, + thread_id=thread_id, + deployment_id=deployment_id, + ) + + +def _validated_scope_directories(scope_id: str) -> tuple[Path, Path]: + """Return canonical private directories after preventing symlink escape. + + This function intentionally resolves paths and must run only from a + filesystem-operation worker, never from the runtime backend factory. + """ + + conversations_dir = ( + paths.WORKSPACE_ROOT.expanduser() / ".evoscientist" / "conversations" + ).resolve(strict=True) + scope_root = (conversations_dir / scope_id).resolve(strict=True) + files_dir = (scope_root / "files").resolve(strict=True) + runtime_dir = (scope_root / "runtime").resolve(strict=True) + if ( + scope_root.parent != conversations_dir + or files_dir.parent != scope_root + or runtime_dir.parent != scope_root + ): + raise ScopeAccessError("workspace directory escapes its scope") + if not files_dir.is_dir() or not runtime_dir.is_dir(): + raise ScopeAccessError("workspace directory is missing") + return files_dir, runtime_dir + + +def _resolve_scope_context(config: _RuntimeScopeConfig | None) -> ScopeContext | None: + """Validate parsed scope identifiers against the active registry.""" + + if config is None: + return None + deployment_id = config.deployment_id or current_deployment_id() + registry = get_scope_registry(paths.WORKSPACE_ROOT) + if registry.active_lock(deployment_id, "workspace-cutover") is not None: + raise ScopeAccessError("workspace cutover is in progress") + record = registry.assert_runtime( + deployment_id, config.scope_id, config.thread_id, config.owner_id + ) + files_dir, runtime_dir = _validated_scope_directories(record.scope_id) + return ScopeContext( + deployment_id=deployment_id, + scope_id=record.scope_id, + owner_id=config.owner_id, + thread_id=config.thread_id, + revision=record.revision, + files_dir=files_dir, + runtime_dir=runtime_dir, + primary_thread_id=record.primary_thread_id, + ) + + +def require_scoped_runtime( + runtime: ToolRuntime[Any, Any] | Any | None, + *, + kind: str = "tool", +) -> ScopeContext | None: + """Resolve and validate a runtime scope. + + ``optional`` retains legacy CLI compatibility when no scope has been + injected. ``required`` never falls back to ``WORKSPACE_ROOT``. + """ + + config = _runtime_scope_config(runtime, kind=kind) + return _resolve_scope_context(config) + + +def provision_conversation_scope( + thread_id: str, + *, + deployment_id: str | None = None, + scope_id: str | None = None, + workspace_root: Path | None = None, + lock_operation_id: str | None = None, +) -> ScopeRecord: + """Create the registry mapping and private directory for a primary thread.""" + + root = (workspace_root or paths.WORKSPACE_ROOT).expanduser() + deployment_id = deployment_id or deployment_id_for_workspace(root) + registry = get_scope_registry(root) + record = registry.provision( + deployment_id, + thread_id, + scope_id=scope_id, + lock_operation_id=lock_operation_id, + ) + root = conversation_root(record.scope_id, root) + try: + (root / "files").mkdir(mode=0o700, parents=True, exist_ok=True) + (root / "runtime").mkdir(mode=0o700, parents=True, exist_ok=True) + for directory in (root, root / "files", root / "runtime"): + try: + directory.chmod(0o700) + except OSError: + pass + except OSError: + # Keep the durable reservation for the recovery job; it is safer than + # silently falling back to the shared deployment root. + raise + return record + + +def _build_backend(root_dir: Path, *, dangerous: bool) -> Any: + from deepagents.backends import CompositeBackend + + from .backends import ( + CustomSandboxBackend, + MemoryFilesystemBackend, + MergedSkillsBackend, + ) + from .EvoScientist import SKILLS_DIR + + cfg_timeout = int(os.getenv("EVOSCIENTIST_SANDBOX_EXECUTE_TIMEOUT", "300")) + ws_backend: Any + if is_required(): + ws_backend = ScopedContainerBackend(root_dir, timeout=cfg_timeout) + else: + ws_backend = CustomSandboxBackend( + root_dir=str(root_dir), + virtual_mode=True, + timeout=cfg_timeout, + dangerous=dangerous, + ) + return CompositeBackend( + default=ws_backend, + routes={ + "/skills/": MergedSkillsBackend( + primary_dir=str(paths.USER_SKILLS_DIR), + global_dir=str(paths.GLOBAL_SKILLS_DIR), + secondary_dir=SKILLS_DIR, + ), + "/memories/": MemoryFilesystemBackend( + root_dir=str(paths.MEMORIES_DIR), virtual_mode=True + ), + }, + ) + + +class DeferredScopedBackend(SandboxBackendProtocol): + """Resolve the scoped filesystem backend only from a worker thread. + + DeepAgents invokes its deprecated backend factory from async middleware. + Its concrete filesystem backends synchronously call ``Path.resolve()`` in + their constructors, so doing that work in the factory makes every run fail + under LangGraph's blocking-call detector. This proxy itself is I/O-free; + the inherited async methods dispatch the synchronous operations to a + thread, where Registry validation and concrete backend construction occur. + """ + + def __init__( + self, + config: _RuntimeScopeConfig, + *, + dangerous: bool, + ) -> None: + self._config = config + self._dangerous = dangerous + self._backend: Any | None = None + self._backend_key: tuple[str, str, str, int] | None = None + self._lock = threading.RLock() + + @property + def id(self) -> str: + # This is queried while composing the model request; do not initialize + # the real backend or touch the Registry here. + return f"scope-{self._config.scope_id[:8]}-{self._config.owner_id[:8]}" + + def _delegate(self) -> Any: + """Validate the current scope and return a concrete backend. + + Every operation enters here, so a deleted scope or stale owner cannot + keep using a backend constructed before the lifecycle transition. + """ + + # Async backend methods run this code in a worker thread. LangGraph's + # RunnableConfig context variable is not available there, so validate + # the immutable scope parsed by the factory on the graph thread. + context = _resolve_scope_context(self._config) + if context is None: + raise ScopeAccessError("scoped backend lost its workspace scope") + if is_required() and self._dangerous: + raise ScopeAccessError( + "dangerous_mode is incompatible with required isolation" + ) + key = (context.scope_id, context.owner_id, context.thread_id, context.revision) + with self._lock: + if self._backend is None or self._backend_key != key: + self._backend = _build_backend( + context.files_dir, dangerous=self._dangerous + ) + self._backend_key = key + return self._backend + + def ls(self, path: str) -> LsResult: + return self._delegate().ls(path) + + def read(self, file_path: str, offset: int = 0, limit: int = 2000) -> ReadResult: + return self._delegate().read(file_path, offset, limit) + + def grep( + self, pattern: str, path: str | None = None, glob: str | None = None + ) -> GrepResult: + return self._delegate().grep(pattern, path, glob) + + def glob(self, pattern: str, path: str | None = None) -> GlobResult: + return self._delegate().glob(pattern, path) + + def write(self, file_path: str, content: str) -> WriteResult: + return self._delegate().write(file_path, content) + + def edit( + self, + file_path: str, + old_string: str, + new_string: str, + replace_all: bool = False, + ) -> EditResult: + return self._delegate().edit(file_path, old_string, new_string, replace_all) + + def upload_files(self, files: list[tuple[str, bytes]]) -> list[FileUploadResponse]: + return self._delegate().upload_files(files) + + def download_files(self, paths: list[str]) -> list[FileDownloadResponse]: + return self._delegate().download_files(paths) + + def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse: + return self._delegate().execute(command, timeout=timeout) + + +def create_workspace_backend( + runtime: ToolRuntime[Any, Any], + *, + legacy_backend: Callable[[], Any], + dangerous: bool = False, + allow_unscoped_legacy: bool = True, +) -> Any: + """Return a backend handle without blocking the Agent event loop.""" + + config = _runtime_scope_config(runtime, kind="filesystem backend") + if config is None: + if not allow_unscoped_legacy: + raise ScopeAccessError( + "deployed graph runs require a workspace scope" + ) + return legacy_backend() + if is_required() and dangerous: + raise ScopeAccessError("dangerous_mode is incompatible with required isolation") + return DeferredScopedBackend(config, dangerous=dangerous) + + +def create_workspace_backend_factory( + legacy_backend: Callable[[], Any], + *, + dangerous: bool = False, + allow_unscoped_legacy: bool = True, +) -> Callable[[ToolRuntime[Any, Any]], Any]: + def factory(runtime: ToolRuntime[Any, Any]) -> Any: + return create_workspace_backend( + runtime, + legacy_backend=legacy_backend, + dangerous=dangerous, + allow_unscoped_legacy=allow_unscoped_legacy, + ) + + return factory + + +def workspace_metadata(record: ScopeRecord) -> dict[str, str | int]: + """Metadata mirrored onto the LangGraph primary thread by trusted callers.""" + + return { + "workspace_schema_version": 1, + "workspace_scope_id": record.scope_id, + "workspace_status": record.state, + "workspace_scope_owner_id": record.primary_owner_id, + "workspace_scope_revision": record.revision, + "workspace_deployment_id": record.deployment_id, + } diff --git a/pyproject.toml b/pyproject.toml index ccd4890..b83f298 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -26,6 +26,7 @@ dependencies = [ "langchain-openrouter>=0.2.5", "tavily-python>=0.7", "pyyaml>=6.0", + "rfc8785==0.1.4", "rich>=15.0", "prompt-toolkit>=3.0", "questionary>=2.1", diff --git a/start-langgraph.sh b/start-langgraph.sh new file mode 100755 index 0000000..c8455ac --- /dev/null +++ b/start-langgraph.sh @@ -0,0 +1,30 @@ +#!/usr/bin/env bash + +set -euo pipefail + +PROJECT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)" +LANGGRAPH_CONFIG="${PROJECT_DIR}/EvoScientist/langgraph_dev/langgraph.json" +HOST="${EVOSCIENTIST_LANGGRAPH_HOST:-127.0.0.1}" +PORT="${EVOSCIENTIST_LANGGRAPH_DEV_PORT:-3076}" +WEB_ENV="${PROJECT_DIR}/../Ai4Sci-Web/.env" + +if [[ ! -x "${PROJECT_DIR}/.venv/bin/langgraph" ]]; then + echo "LangGraph executable not found: ${PROJECT_DIR}/.venv/bin/langgraph" >&2 + echo "Run 'uv sync' in ${PROJECT_DIR} first." >&2 + exit 1 +fi + +cd "${PROJECT_DIR}" + +# The Web Gateway always supplies a verified conversation workspace scope. +export EVOSCIENTIST_DEPLOY_MODE="${EVOSCIENTIST_DEPLOY_MODE:-full}" +export EVOSCIENTIST_WORKSPACE_DIR="${EVOSCIENTIST_WORKSPACE_DIR:-${PROJECT_DIR}/../.ai4sci/workspace}" + +exec uv run --env-file "${WEB_ENV}" langgraph dev \ + --config "${LANGGRAPH_CONFIG}" \ + --host "${HOST}" \ + --port "${PORT}" \ + --no-browser \ + --no-reload \ + --allow-blocking \ + --n-jobs-per-worker 1 diff --git a/tests/__init__.py b/tests/__init__.py index e69de29..0e14beb 100644 --- a/tests/__init__.py +++ b/tests/__init__.py @@ -0,0 +1 @@ +"""EvoScientist test package.""" diff --git a/tests/test_admin_control_v2.py b/tests/test_admin_control_v2.py new file mode 100644 index 0000000..13a32e5 --- /dev/null +++ b/tests/test_admin_control_v2.py @@ -0,0 +1,162 @@ +from __future__ import annotations + +import sqlite3 +from dataclasses import fields + +import pytest + +from EvoScientist.llm.config_admin import EvoModelConfigAdminService +from EvoScientist.llm.contracts import ( + ADMIN_CONTROL_VERSION, + CommitProposalRequest, + CreateProposalRequest, + HmacGrantAuthority, + ProbeProposalRequest, + UpdateProposalRequest, + ValidateProposalRequest, +) +from EvoScientist.llm.crypto import HmacKeyRing, KeyMaterial, sha256_id +from EvoScientist.llm.model_config import FileEvoModelConfigStore +from EvoScientist.llm.secret_store import EncryptedModelSecretStore +from tests.test_provider_model_config_v3 import v3_payload + + +def _request(authority, cls, action: str, **values): + payload = {"admin_control_version": ADMIN_CONTROL_VERSION, **values} + operation_id = str(values["operation_id"]) + grant = authority.sign_admin( + subject_id="admin-1", + action=action, + operation_id=operation_id, + request_digest=sha256_id(payload), + ) + valid = {item.name for item in fields(cls)} + complete = {**values, "admin_grant": grant, "admin_control_version": 2} + return cls(**{key: complete[key] for key in valid}) + + +@pytest.mark.asyncio +async def test_admin_control_v2_proposal_commit(tmp_path) -> None: + authority = HmacGrantAuthority("g" * 32, key_id="grant-v1") + ring = HmacKeyRing(KeyMaterial.create("identity-v1", "i" * 32)) + store = FileEvoModelConfigStore( + tmp_path / "model_routes.yaml", + admin_verifier=authority, + ops_path=tmp_path / "model_config_ops.sqlite", + ) + secrets = EncryptedModelSecretStore( + tmp_path / "model_secrets.sqlite", master_secret="s" * 32 + ) + payload = v3_payload() + for provider in payload["providers"]: + item = secrets.create_pending( + provider["provider_id"], + "sk-" + provider["provider_id"], + created_by="admin-1", + operation_id="secret-" + provider["provider_id"], + ) + provider["connection"]["credential_ref"] = item.ref + service = EvoModelConfigAdminService( + store, + grant_authority=authority, + identity_key_ring=ring, + secret_resolver=secrets.resolve, + secret_store=secrets, + probe_runner=lambda _config, _route, _kind: True, + ) + + created = service.create_proposal( + _request( + authority, + CreateProposalRequest, + "model_config:proposal:create", + operation_id="create-1", + expected_active_revision=0, + ) + ) + updated = service.update_proposal( + _request( + authority, + UpdateProposalRequest, + "model_config:proposal:update", + operation_id="update-1", + proposal_id=created.proposal_id, + expected_state_version=created.state_version, + expected_draft_etag=created.draft_etag, + draft_payload=payload, + ) + ) + validated = service.validate_proposal( + _request( + authority, + ValidateProposalRequest, + "model_config:proposal:validate", + operation_id="validate-1", + proposal_id=created.proposal_id, + expected_state_version=updated.state_version, + expected_draft_etag=updated.draft_etag, + ) + ) + current = validated + for index, route in enumerate(validated.routes): + for kind in route["required_probe_kinds"]: + current = await service.probe_proposal( + _request( + authority, + ProbeProposalRequest, + "model_config:proposal:probe", + operation_id=f"probe-{index}-{kind}", + proposal_id=created.proposal_id, + validated_digest=validated.validated_digest, + route_semantics_hash=route["route_semantics_hash"], + probe_kind=kind, + ) + ) + assert current.state == "READY" + committed = service.commit_proposal( + _request( + authority, + CommitProposalRequest, + "model_config:proposal:commit", + operation_id="commit-1", + proposal_id=created.proposal_id, + expected_active_revision=0, + expected_state_version=current.state_version, + expected_draft_etag=current.draft_etag, + validated_digest=validated.validated_digest, + evidence_ids=current.evidence_ids, + ) + ) + assert committed.state == "COMMITTED" + assert committed.active_revision == 1 + active = store.load() + assert active.schema_version == 3 + assert len(active.providers) == 4 + assert all(item.status == "active" for item in secrets.list_metadata()) + + with sqlite3.connect(store.ops_path) as connection: + connection.execute( + "UPDATE config_commit_operations SET stage='CONFIG_COMMITTED' WHERE operation_id='commit-1'" + ) + connection.execute( + "UPDATE config_proposals SET state='COMMITTING' WHERE proposal_id=?", + (created.proposal_id,), + ) + EvoModelConfigAdminService( + store, + grant_authority=authority, + identity_key_ring=ring, + secret_resolver=secrets.resolve, + secret_store=secrets, + probe_runner=lambda _config, _route, _kind: True, + ) + with sqlite3.connect(store.ops_path) as connection: + stage = connection.execute( + "SELECT stage FROM config_commit_operations WHERE operation_id='commit-1'" + ).fetchone()[0] + state = connection.execute( + "SELECT state FROM config_proposals WHERE proposal_id=?", + (created.proposal_id,), + ).fetchone()[0] + assert stage == "COMPLETED" + assert state == "COMMITTED" diff --git a/tests/test_agent_factory_extensions.py b/tests/test_agent_factory_extensions.py index e05dce5..eebd5d2 100644 --- a/tests/test_agent_factory_extensions.py +++ b/tests/test_agent_factory_extensions.py @@ -68,10 +68,13 @@ def test_create_cli_agent_accepts_host_backend_and_memory_options( assert calls["middleware_kwargs"]["tool_selector_threshold"] == 8 assert calls["middleware_kwargs"]["memory_max_inline_profile_chars"] == 1000 assert calls["middleware_kwargs"]["enable_background_execution"] is False + assert calls["middleware_kwargs"]["enable_legacy_model_fallback"] is True assert calls["agent_config"] == {"recursion_limit": 321} -def test_create_cli_agent_installs_route_middleware_in_fixed_slot(monkeypatch, tmp_path): +def test_create_cli_agent_installs_route_middleware_in_fixed_slot( + monkeypatch, tmp_path +): import EvoScientist.EvoScientist as agent_module from EvoScientist.config.settings import EvoScientistConfig @@ -106,11 +109,14 @@ def test_create_cli_agent_installs_route_middleware_in_fixed_slot(monkeypatch, t monkeypatch.setattr("EvoScientist.backends.MemoryFilesystemBackend", _Backend) monkeypatch.setattr("EvoScientist.backends.MergedSkillsBackend", _Backend) monkeypatch.setattr(agent_module, "set_active_workspace", lambda _path: None) + def fake_default_middleware(**kwargs): calls["middleware_kwargs"] = kwargs return list(default_chain) - monkeypatch.setattr(agent_module, "_get_default_middleware", fake_default_middleware) + monkeypatch.setattr( + agent_module, "_get_default_middleware", fake_default_middleware + ) def fake_load(_backend, middleware, **_kwargs): calls["middleware"] = middleware @@ -128,10 +134,68 @@ def test_create_cli_agent_installs_route_middleware_in_fixed_slot(monkeypatch, t ) assert calls["middleware_kwargs"]["enable_legacy_model_fallback"] is False - assert [middleware.name for middleware in calls["middleware"][:5]] == [ + assert [middleware.name for middleware in calls["middleware"][:6]] == [ "error_normalization", + "provider_context_media", "configurable_model", "gateway_route_fallback", "context_editing", "tool_protocol_guard", ] + + +def test_create_cli_agent_replaces_framework_skill_injection(monkeypatch, tmp_path): + import EvoScientist.EvoScientist as agent_module + from EvoScientist.config.settings import EvoScientistConfig + from EvoScientist.middleware import BudgetedSkillsMiddleware + + calls = {} + + class _Backend: + def __init__(self, **_kwargs): + pass + + class _CompositeBackend: + def __init__(self, **_kwargs): + pass + + class _Agent: + def with_config(self, _config): + return self + + def _create_deep_agent(**kwargs): + calls["kwargs"] = kwargs + return _Agent() + + monkeypatch.setattr("deepagents.backends.CompositeBackend", _CompositeBackend) + monkeypatch.setattr("deepagents.create_deep_agent", _create_deep_agent) + monkeypatch.setattr("EvoScientist.backends.MemoryFilesystemBackend", _Backend) + monkeypatch.setattr("EvoScientist.backends.MergedSkillsBackend", _Backend) + monkeypatch.setattr(agent_module, "set_active_workspace", lambda _path: None) + monkeypatch.setattr(agent_module, "_get_default_middleware", lambda **_kwargs: []) + monkeypatch.setattr( + agent_module, + "load_mcp_and_build_kwargs", + lambda *_args, **_kwargs: { + "skills": ["/skills/"], + "middleware": [], + "subagents": [{"name": "research", "skills": ["/skills/"]}], + }, + ) + + agent_module.create_cli_agent( + workspace_dir=str(tmp_path), + checkpointer=object(), + config=EvoScientistConfig(auto_approve=True), + chat_model=object(), + workspace_backend=object(), + ) + + assert calls["kwargs"]["skills"] is None + assert any( + isinstance(middleware, BudgetedSkillsMiddleware) + for middleware in calls["kwargs"]["middleware"] + ) + subagent = calls["kwargs"]["subagents"][0] + assert subagent["skills"] is None + assert isinstance(subagent["middleware"][0], BudgetedSkillsMiddleware) diff --git a/tests/test_async_subagent_factory.py b/tests/test_async_subagent_factory.py index f8e0c3c..72b5d5e 100644 --- a/tests/test_async_subagent_factory.py +++ b/tests/test_async_subagent_factory.py @@ -78,7 +78,7 @@ def test_factory_requests_async_safe_middleware( "name": "writing-agent", "system_prompt": "", "tools": [], - "skills": None, + "skills": ["/skills/"], } ] # ``create_deep_agent(...).with_config({...})`` chain — return something @@ -94,6 +94,12 @@ def test_factory_requests_async_safe_middleware( for_async_subagent=True, memory_source_agent="writing-agent", ) + from EvoScientist.middleware import BudgetedSkillsMiddleware + + assert mock_create.call_args.kwargs["skills"] is None + assert isinstance( + mock_create.call_args.kwargs["middleware"][-1], BudgetedSkillsMiddleware + ) subagents = mock_create.call_args.kwargs["subagents"] assert subagents[0]["name"] == "general-purpose" _assert_subagent_memory_middleware( diff --git a/tests/test_config.py b/tests/test_config.py index 07a040e..5d8c4a8 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -25,6 +25,10 @@ from EvoScientist.config import ( set_config_value, ) + +def test_langgraph_dev_port_defaults_to_ai4sci_runtime_port(): + assert EvoScientistConfig().langgraph_dev_port == 3076 + # ============================================================================= # Fixtures # ============================================================================= @@ -60,6 +64,7 @@ def temp_config_dir(tmp_path, monkeypatch): "ANTHROPIC_API_KEY", "OPENAI_API_KEY", "TAVILY_API_KEY", + "S2_API_KEY", "EVOSCIENTIST_DEFAULT_MODE", "EVOSCIENTIST_WORKSPACE_DIR", "EVOSCIENTIST_UI_BACKEND", @@ -90,6 +95,7 @@ def clean_env(monkeypatch): "ANTHROPIC_API_KEY", "OPENAI_API_KEY", "TAVILY_API_KEY", + "S2_API_KEY", "EVOSCIENTIST_DEFAULT_MODE", "EVOSCIENTIST_WORKSPACE_DIR", "EVOSCIENTIST_UI_BACKEND", @@ -125,6 +131,7 @@ class TestEvoScientistConfig: assert config.anthropic_api_key == "" assert config.openai_api_key == "" assert config.tavily_api_key == "" + assert config.semantic_scholar_api_key == "" assert config.provider == "anthropic" assert config.model == "claude-sonnet-4-6" assert config.default_mode == "daemon" @@ -132,7 +139,6 @@ class TestEvoScientistConfig: assert config.show_thinking is True assert config.ui_backend == "tui" assert config.log_level == "warning" - assert config.reasoning_effort == "" assert config.openrouter_anthropic_prompt_cache is True assert config.openrouter_http_referer == ( "https://github.com/EvoScientist/EvoScientist" @@ -599,12 +605,12 @@ class TestPriorityChain: config = get_effective_config() assert config.log_level == "DEBUG" - def test_env_reasoning_effort_override(self, temp_config_dir, monkeypatch): - """Reasoning effort can be selected via environment variable.""" - save_config(EvoScientistConfig(reasoning_effort="medium")) + def test_env_reasoning_effort_is_not_a_config_override(self, temp_config_dir, monkeypatch): + """Reasoning is selected by the invocation plan, never deployment env.""" + save_config(EvoScientistConfig()) monkeypatch.setenv("EVOSCIENTIST_REASONING_EFFORT", "high") config = get_effective_config() - assert config.reasoning_effort == "high" + assert not hasattr(config, "reasoning_effort") def test_env_channel_debug_tracing_override(self, temp_config_dir, monkeypatch): """Channel tracing can be enabled via environment variable.""" @@ -750,6 +756,7 @@ class TestApplyConfigToEnv: anthropic_api_key="config-ant-key", openai_api_key="config-oai-key", tavily_api_key="config-tav-key", + semantic_scholar_api_key="config-s2-key", ) apply_config_to_env(config) @@ -757,6 +764,7 @@ class TestApplyConfigToEnv: assert os.environ.get("ANTHROPIC_API_KEY") == "config-ant-key" assert os.environ.get("OPENAI_API_KEY") == "config-oai-key" assert os.environ.get("TAVILY_API_KEY") == "config-tav-key" + assert os.environ.get("S2_API_KEY") == "config-s2-key" def test_does_not_override_existing_env(self, monkeypatch): """Test that existing env vars are not overridden.""" diff --git a/tests/test_context_overflow_middleware.py b/tests/test_context_overflow_middleware.py index 09574bc..36e55f1 100644 --- a/tests/test_context_overflow_middleware.py +++ b/tests/test_context_overflow_middleware.py @@ -7,6 +7,7 @@ from langchain.agents.middleware.types import ModelRequest from langchain_core.exceptions import ContextOverflowError from langchain_core.messages import HumanMessage +from EvoScientist.llm.contracts import EvoRuntimeError from EvoScientist.middleware.context_overflow import ContextOverflowMapperMiddleware @@ -24,6 +25,15 @@ def test_is_context_limit_error_anthropic(): assert mw._is_context_limit_error(exc) is True +def test_is_context_limit_error_for_runtime_admission_guard(): + mw = ContextOverflowMapperMiddleware() + + assert ( + mw._is_context_limit_error(EvoRuntimeError("MODEL_CONTEXT_WINDOW_EXCEEDED")) + is True + ) + + def test_is_not_context_limit_error_without_400(): mw = ContextOverflowMapperMiddleware() exc = Exception("context_length_exceeded, but no status code") diff --git a/tests/test_error_normalization_middleware.py b/tests/test_error_normalization_middleware.py index efd1799..1bd014d 100644 --- a/tests/test_error_normalization_middleware.py +++ b/tests/test_error_normalization_middleware.py @@ -295,6 +295,8 @@ class TestMiddleware: self._run_awrap(mw, req, handler) assert excinfo.value.provider == "openrouter" assert excinfo.value.__cause__ is raised + assert excinfo.value.message == "Provider request failed." + assert "boom" not in str(excinfo.value.model_dump()) def test_awrap_passes_through_non_provider_model_exception(self): """If the model isn't a recognized provider SDK, the exception diff --git a/tests/test_evo_route_fallback.py b/tests/test_evo_route_fallback.py new file mode 100644 index 0000000..4fc8fdc --- /dev/null +++ b/tests/test_evo_route_fallback.py @@ -0,0 +1,328 @@ +from __future__ import annotations + +import pytest +from langchain.agents.middleware.types import ModelRequest +from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage + +from EvoScientist.llm.contracts import EvoRuntimeError +from EvoScientist.llm.errors import ( + AgentControlError, + ModelProviderResponseError, + ModelToolProtocolError, +) +from EvoScientist.middleware.evo_route_fallback import EvoRouteFallbackMiddleware + + +class _Model: + def __init__(self, route_key: str, *, supports_tools: bool | None = None) -> None: + self.metadata = {"route_key": route_key} + if supports_tools is not None: + self.metadata["route_supports_tools"] = supports_tools + + +class _Health: + def __init__(self, open_routes: set[str] | None = None) -> None: + self.open_routes = open_routes or set() + + def is_open(self, route_key: str) -> bool: + return route_key in self.open_routes + + +def _request(model: _Model) -> ModelRequest: + return ModelRequest(model=model, messages=[], tools=[]) + + +def test_route_strips_tools_when_frozen_model_capability_is_false(): + model = _Model("text-only", supports_tools=False) + request = ModelRequest( + model=model, + messages=[], + tools=[{"type": "function", "function": {"name": "search"}}], + ) + + def handler(routed: ModelRequest): + return routed.tools + + assert EvoRouteFallbackMiddleware([]).wrap_model_call(request, handler) == [] + + +def test_tools_disabled_route_removes_checkpoint_tool_protocol(): + model = _Model("text-only", supports_tools=False) + request = ModelRequest( + model=model, + messages=[ + SystemMessage("Keep the answer concise."), + AIMessage( + content="", + tool_calls=[{"name": "search", "args": {"q": "K3"}, "id": "call_1"}], + ), + ToolMessage(content="tool result", tool_call_id="call_1", name="search"), + AIMessage( + content="The tool found a result.", + additional_kwargs={"tool_calls": [{"id": "call_2"}]}, + ), + HumanMessage("Continue."), + ], + tools=[{"type": "function", "function": {"name": "search"}}], + tool_choice="required", + response_format={"type": "json_object"}, + model_settings={"temperature": 0, "max_tokens": 1}, + ) + + def handler(routed: ModelRequest): + return routed + + routed = EvoRouteFallbackMiddleware([]).wrap_model_call(request, handler) + + assert routed.tools == [] + assert routed.tool_choice is None + assert routed.response_format is None + assert routed.model_settings == {} + assert [message.type for message in routed.messages] == [ + "system", + "ai", + "ai", + "human", + ] + assistant = routed.messages[1] + assert isinstance(assistant, AIMessage) + assert assistant.content == "[Completed tool result: search]\ntool result" + assert routed.messages[2].content == "The tool found a result." + assert routed.messages[2].tool_calls == [] + assert "tool_calls" not in routed.messages[2].additional_kwargs + + +def test_tools_disabled_route_projects_tool_result_and_content_blocks_to_text(): + model = _Model("text-only", supports_tools=False) + request = ModelRequest( + model=model, + messages=[ + AIMessage( + content=[ + {"type": "text", "text": "I checked the source."}, + {"type": "tool_call", "id": "call_1", "name": "read"}, + ], + ), + ToolMessage( + content="x" * 12_100, + tool_call_id="call_1", + name="read", + ), + ], + tools=[{"type": "function", "function": {"name": "read"}}], + ) + + routed = EvoRouteFallbackMiddleware([]).wrap_model_call(request, lambda item: item) + + assert routed.tools == [] + assert len(routed.messages) == 2 + assert routed.messages[0].content == "I checked the source." + assert routed.messages[0].tool_calls == [] + assert routed.messages[1].content.startswith("[Completed tool result: read]\n") + assert routed.messages[1].content.endswith("[Tool result truncated]") + + +def test_route_keeps_tools_when_frozen_model_capability_is_true(): + model = _Model("tools", supports_tools=True) + tools = [{"type": "function", "function": {"name": "search"}}] + request = ModelRequest( + model=model, + messages=[], + tools=tools, + model_settings={"temperature": 0}, + ) + + def handler(routed: ModelRequest): + return routed + + routed = EvoRouteFallbackMiddleware([]).wrap_model_call(request, handler) + assert routed.tools == tools + assert routed.model_settings == {} + + +def test_route_drops_empty_assistant_history_but_keeps_tool_calls(): + model = _Model("tools", supports_tools=True) + tool_message = AIMessage( + content="", + tool_calls=[{"name": "search", "args": {"q": "K3"}, "id": "call_1"}], + ) + request = ModelRequest( + model=model, + messages=[ + HumanMessage("First"), + AIMessage(content="", additional_kwargs={"reasoning_content": "hidden"}), + HumanMessage("Continue"), + tool_message, + ], + tools=[], + ) + + routed = EvoRouteFallbackMiddleware([]).wrap_model_call(request, lambda item: item) + + assert routed.messages == [request.messages[0], request.messages[2], tool_message] + + +@pytest.mark.asyncio +async def test_empty_provider_response_retries_same_route_once(): + primary = _Model("primary", supports_tools=True) + request = ModelRequest(model=primary, messages=[HumanMessage("Answer")], tools=[]) + seen: list[ModelRequest] = [] + + async def handler(routed: ModelRequest): + seen.append(routed) + if len(seen) == 1: + return AIMessage( + content="", additional_kwargs={"reasoning_content": "hidden"} + ) + return AIMessage(content="Final answer") + + response = await EvoRouteFallbackMiddleware([]).awrap_model_call(request, handler) + + assert isinstance(response, AIMessage) + assert response.content == "Final answer" + assert len(seen) == 2 + assert isinstance(seen[1].messages[0], SystemMessage) + assert "without final text" in str(seen[1].messages[0].content) + + +@pytest.mark.asyncio +async def test_repeated_empty_provider_response_fails_with_specific_code(): + primary = _Model("primary", supports_tools=True) + attempts = 0 + + async def handler(_routed: ModelRequest): + nonlocal attempts + attempts += 1 + return AIMessage(content=[{"type": "reasoning", "summary": []}]) + + with pytest.raises(ModelProviderResponseError) as captured: + await EvoRouteFallbackMiddleware([]).awrap_model_call( + ModelRequest(model=primary, messages=[], tools=[]), handler + ) + + assert captured.value.code == "MODEL_PROVIDER_RESPONSE_INVALID" + assert attempts == 2 + + +@pytest.mark.asyncio +async def test_fallback_applies_its_own_tool_capability(): + primary = _Model("primary", supports_tools=True) + fallback = _Model("fallback", supports_tools=False) + tools = [{"type": "function", "function": {"name": "search"}}] + request = ModelRequest(model=primary, messages=[], tools=tools) + seen = [] + + async def handler(routed: ModelRequest): + seen.append((routed.model.metadata["route_key"], list(routed.tools))) + if routed.model is primary: + raise ConnectionError("upstream unavailable") + return "ok" + + assert await EvoRouteFallbackMiddleware([fallback]).awrap_model_call( + request, handler + ) == "ok" + assert seen == [("primary", tools), ("fallback", [])] + + +@pytest.mark.asyncio +async def test_route_fallback_uses_evo_frozen_fallback_model(): + primary = _Model("primary") + fallback = _Model("fallback") + middleware = EvoRouteFallbackMiddleware([fallback], _Health()) + seen: list[str] = [] + + async def handler(request: ModelRequest): + route = request.model.metadata["route_key"] + seen.append(route) + if route == "primary": + raise ConnectionError("upstream unavailable") + return "fallback-response" + + result = await middleware.awrap_model_call(_request(primary), handler) + + assert result == "fallback-response" + assert seen == ["primary", "fallback"] + + +@pytest.mark.asyncio +async def test_protocol_error_retries_with_a_bounded_repair_instruction(): + primary = _Model("primary", supports_tools=True) + request = ModelRequest( + model=primary, + messages=[], + tools=[{"type": "function", "function": {"name": "search"}}], + ) + seen: list[ModelRequest] = [] + + async def handler(routed: ModelRequest): + seen.append(routed) + if len(seen) == 1: + raise ModelToolProtocolError("unknown_name") + return "repaired-response" + + result = await EvoRouteFallbackMiddleware([primary]).awrap_model_call( + request, handler + ) + + assert result == "repaired-response" + assert len(seen) == 2 + assert seen[0].messages == [] + assert isinstance(seen[1].messages[0], SystemMessage) + assert "unknown_name" in str(seen[1].messages[0].content) + assert "search" in str(seen[1].messages[0].content) + + +@pytest.mark.asyncio +async def test_other_agent_control_errors_remain_non_fallbackable(): + primary = _Model("primary") + fallback = _Model("fallback") + seen: list[str] = [] + + async def handler(routed: ModelRequest): + seen.append(routed.model.metadata["route_key"]) + raise AgentControlError("MODEL_REQUEST_REJECTED", "rejected") + + with pytest.raises(AgentControlError): + await EvoRouteFallbackMiddleware([fallback]).awrap_model_call( + _request(primary), handler + ) + + assert seen == ["primary"] + + +@pytest.mark.asyncio +async def test_open_primary_is_skipped_and_control_errors_do_not_fallback(): + primary = _Model("primary") + fallback = _Model("fallback") + middleware = EvoRouteFallbackMiddleware([fallback], _Health({"primary"})) + seen: list[str] = [] + + async def handler(request: ModelRequest): + seen.append(request.model.metadata["route_key"]) + return "fallback-response" + + assert await middleware.awrap_model_call(_request(primary), handler) == "fallback-response" + assert seen == ["fallback"] + + async def controlled(_request: ModelRequest): + raise EvoRuntimeError("ADMISSION_EXHAUSTED") + + with pytest.raises(EvoRuntimeError, match="ADMISSION_EXHAUSTED"): + await EvoRouteFallbackMiddleware([fallback]).awrap_model_call( + _request(primary), + controlled, + ) + + +def test_sync_open_primary_is_skipped(): + primary = _Model("primary") + fallback = _Model("fallback") + middleware = EvoRouteFallbackMiddleware([fallback], _Health({"primary"})) + seen: list[str] = [] + + def handler(request: ModelRequest): + seen.append(request.model.metadata["route_key"]) + return "fallback-response" + + assert middleware.wrap_model_call(_request(primary), handler) == "fallback-response" + assert seen == ["fallback"] diff --git a/tests/test_gateway_background_runs.py b/tests/test_gateway_background_runs.py index ae7ea78..b0a691a 100644 --- a/tests/test_gateway_background_runs.py +++ b/tests/test_gateway_background_runs.py @@ -127,6 +127,7 @@ def test_launch_background_run_submits_run_and_invokes_hooks(monkeypatch): assert handle is not None assert handle.thread_id == "thread-1" assert handle.run_id == "run-1" + assert handle.configurable == {"thread_id": "thread-1"} assert payload_calls == ["thread-1"] assert before_calls == ["thread-1"] assert started == [handle] diff --git a/tests/test_host_metering_extensions.py b/tests/test_host_metering_extensions.py index 7123e04..afbaa07 100644 --- a/tests/test_host_metering_extensions.py +++ b/tests/test_host_metering_extensions.py @@ -1,5 +1,6 @@ from EvoScientist.llm.errors import AgentControlError from EvoScientist.middleware.model_fallback import _is_non_fallbackable +from EvoScientist.middleware.recoverable_metering import _metering_config, _source_type def test_agent_control_error_is_non_fallbackable(): @@ -11,3 +12,24 @@ def test_agent_control_error_is_non_fallbackable(): assert "platform control error" in (_is_non_fallbackable(error) or "") assert error.model_dump()["code"] == "INSUFFICIENT_BALANCE" + + +def test_recoverable_metering_reads_explicit_evomemory_scope(): + metering = _metering_config( + { + "configurable": { + "ai4sci_metering": { + "gateway_url": "http://gateway", + "run_id": "run-parent", + "envelope_signature": "signed-parent", + "source_type": "evomemory_linker", + } + } + } + ) + + assert metering is not None + assert metering["source_type"] == "evomemory_linker" + assert _source_type({"metering_scope": "evomemory_subagent_worker"}, []) == ( + "evomemory_subagent_worker" + ) diff --git a/tests/test_invocation_contract.py b/tests/test_invocation_contract.py new file mode 100644 index 0000000..6180bbb --- /dev/null +++ b/tests/test_invocation_contract.py @@ -0,0 +1,119 @@ +from __future__ import annotations + +import pytest + +from EvoScientist.llm.configuration import ( + EndpointConfig, + ModelConfig, + ProviderConfig, +) +from EvoScientist.llm.contracts import EvoRuntimeError +from EvoScientist.llm.invocation import ( + compile_invocation_plan, + derive_runtime_invocation, + derive_tool_call_transport, +) +from EvoScientist.llm.model_config import ( + EndpointConfig as LegacyEndpointConfig, +) +from EvoScientist.llm.model_config import ModelConfig as LegacyModelConfig +from EvoScientist.llm.model_config import ProviderConfig as LegacyProviderConfig + + +@pytest.mark.parametrize( + ("capabilities", "expected"), + [ + ({"text": True, "tools": True}, "native"), + ({"text": True, "tools": False}, "disabled"), + ({"text": True}, "disabled"), + ], +) +def test_tool_transport_is_derived_only_from_model_capabilities( + capabilities, expected +): + assert derive_tool_call_transport(capabilities) == expected + + +def test_runtime_invocation_combines_model_api_mode_and_derived_transport(): + assert derive_runtime_invocation( + "chat_completions", {"text": True, "tools": False} + ) == { + "api_mode": "chat_completions", + "tool_call_transport": "disabled", + } + + +def test_k3_chat_plan_freezes_adapter_compiled_parameters(): + plan = compile_invocation_plan( + api_mode="chat_completions", + declared_tool_call_transport="disabled", + supports_tools=False, + purpose="main_agent", + output_token_limit=65_000, + reasoning_effort="high", + runtime_provider="openai", + sdk_params={ + "max_completion_tokens": 65_000, + "reasoning_effort": "high", + "use_responses_api": False, + }, + ) + + assert plan.output_token_parameter == "max_completion_tokens" + assert plan.tool_call_transport == "disabled" + assert plan.streaming is True + assert plan.sdk_params["streaming"] is True + with pytest.raises(TypeError): + plan.sdk_params["max_completion_tokens"] = 1 + + +def test_responses_plan_accepts_native_tools_when_capability_is_enabled(): + plan = compile_invocation_plan( + api_mode="responses", + declared_tool_call_transport="native", + supports_tools=True, + purpose="title", + output_token_limit=1_024, + reasoning_effort="disabled", + runtime_provider="openai", + sdk_params={"max_output_tokens": 1_024, "use_responses_api": True}, + ) + + assert plan.tool_call_transport == "native" + assert plan.streaming is False + + +def test_plan_rejects_runtime_projection_that_disagrees_with_capabilities(): + with pytest.raises(EvoRuntimeError, match="MODEL_ADAPTER_COMPILE_FAILED"): + compile_invocation_plan( + api_mode="chat_completions", + declared_tool_call_transport="native", + supports_tools=False, + purpose="main_agent", + output_token_limit=1_024, + reasoning_effort="disabled", + runtime_provider="openai", + sdk_params={"max_tokens": 1_024, "use_responses_api": False}, + ) + + +def test_provider_and_model_contracts_have_separate_ownership(): + model_fields = ModelConfig.__dataclass_fields__ + endpoint_fields = EndpointConfig.__dataclass_fields__ + provider_fields = ProviderConfig.__dataclass_fields__ + + assert "capabilities" in model_fields + assert "max_output_tokens" in model_fields + assert "base_url" not in model_fields + assert "auth" not in model_fields + assert "tool_call_transport" not in model_fields + assert "base_url" in endpoint_fields + assert "auth" in endpoint_fields + assert "adapter_id" in provider_fields + assert "connection_defaults" in provider_fields + + +def test_legacy_model_config_imports_reexport_canonical_contracts(): + assert LegacyEndpointConfig is EndpointConfig + assert LegacyModelConfig is ModelConfig + assert LegacyProviderConfig is ProviderConfig diff --git a/tests/test_langgraph_dev_http.py b/tests/test_langgraph_dev_http.py index 88cd3ca..4b83250 100644 --- a/tests/test_langgraph_dev_http.py +++ b/tests/test_langgraph_dev_http.py @@ -6,7 +6,9 @@ langgraph dev. from __future__ import annotations from unittest.mock import patch +from uuid import uuid4 +import pytest from starlette.testclient import TestClient from EvoScientist.config import EvoScientistConfig @@ -161,3 +163,55 @@ def test_ollama_discovery_skipped_when_base_url_absent(): {"name": n, "model_id": m, "provider": p} for n, m, p in list_models_by_provider() ] + + +def test_recoverable_capabilities_include_resume_and_workspace(monkeypatch): + monkeypatch.setenv("EVOSCIENTIST_DEPLOY_MODE", "full") + response = client.get("/api/ai4sci/recoverable-runs/capabilities") + assert response.status_code == 200 + body = response.json() + assert body["interrupt_resume"] is True + assert body["pending_interrupt_state"] is True + assert body["workspace_scope_v1"] is True + + +def test_recoverable_capabilities_fail_closed_without_full_deploy(monkeypatch): + monkeypatch.delenv("EVOSCIENTIST_DEPLOY_MODE", raising=False) + response = client.get("/api/ai4sci/recoverable-runs/capabilities") + assert response.status_code == 200 + assert response.json()["workspace_scope_v1"] is False + + +def test_recoverable_resume_rejects_human_input(monkeypatch): + run_id = str(uuid4()) + response = client.post( + "/api/ai4sci/recoverable-runs/create", + headers={"x-auth-scheme": "langsmith"}, + json={ + "operation": "resume", + "thread_id": str(uuid4()), + "run_id": run_id, + "run_request_id": run_id, + "request_hash": "a" * 64, + "assistant_id": "EvoScientist", + "input": {"messages": [{"role": "user", "content": "continue"}]}, + "command": {"resume": {"decisions": [{"type": "approve"}]}}, + }, + ) + assert response.status_code == 400 + assert response.json()["code"] == "INVALID_RESUME_REQUEST" + + +def test_workspace_scope_routes_require_service_token(): + response = client.post( + "/internal/workspace-scopes/provision", + json={"thread_id": str(uuid4())}, + ) + assert response.status_code == 401 + + +def test_workspace_materialize_rejects_path_escape(): + from EvoScientist.langgraph_dev.http import _materialize_target + + with pytest.raises(ValueError, match="uploads"): + _materialize_target(str(uuid4()), "../secret.txt") diff --git a/tests/test_llm.py b/tests/test_llm.py index c4ec1a9..dc604bb 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -1,6 +1,5 @@ """Tests for EvoScientist LLM module.""" -from types import SimpleNamespace from unittest.mock import patch import pytest @@ -161,67 +160,24 @@ class TestGetModelInfo: class TestGetChatModel: - @patch("EvoScientist.llm.models.init_chat_model") - def test_uses_host_model_resolver(self, mock_init): - """An embedding host can provide model routing without a core dependency.""" - from EvoScientist.runtime_integrations import ( - configure_runtime_integrations, - reset_runtime_integrations, - ) - - mock_init.return_value = "mock_model" - resolved = SimpleNamespace( - provider_name="relay-a", - model_id="model-a", - protocol="openai", - api_key="sk-host", - base_url="https://relay.example/v1/", - params={"max_tokens": 8192, "_default_headers": {"X-Relay": "a"}}, - supports_reasoning=False, - ) - configure_runtime_integrations(model_resolver=lambda model, provider: resolved) - try: - assert get_chat_model("alias-a") == "mock_model" - finally: - reset_runtime_integrations() - - call_kwargs = mock_init.call_args.kwargs - assert call_kwargs["model"] == "model-a" - assert call_kwargs["model_provider"] == "openai" - assert call_kwargs["api_key"] == "sk-host" - assert call_kwargs["base_url"] == "https://relay.example/v1" - assert call_kwargs["max_tokens"] == 8192 - assert call_kwargs["default_headers"] == {"X-Relay": "a"} - assert "reasoning" not in call_kwargs - @patch("EvoScientist.llm.models._patch_openai_compat_content") @patch("EvoScientist.llm.models.init_chat_model") - def test_host_openai_provider_with_custom_base_uses_compat_patch( - self, mock_init, mock_compat - ): - from EvoScientist.runtime_integrations import ( - configure_runtime_integrations, - reset_runtime_integrations, - ) - + def test_openai_custom_base_url_uses_compat_patch(self, mock_init, mock_compat): model_instance = object() mock_init.return_value = model_instance - resolved = SimpleNamespace( - provider_name="openai", - model_id="gpt-5.5", - protocol="openai", + + get_chat_model( + "gpt-5.5", + provider="openai", api_key="sk-host", base_url="https://relay.example/v1", - params={}, - supports_reasoning=True, ) - configure_runtime_integrations(model_resolver=lambda model, provider: resolved) - try: - get_chat_model("gpt-5.5", provider="openai") - finally: - reset_runtime_integrations() - mock_compat.assert_called_once_with(model_instance, hoist_tool_media=True) + mock_compat.assert_called_once_with( + model_instance, + hoist_tool_media=True, + drop_reasoning_metadata=True, + ) @patch("EvoScientist.llm.models.init_chat_model") def test_uses_default_model_when_none(self, mock_init): @@ -498,9 +454,6 @@ class TestThirdPartyRouting: """OpenRouter should use native 'openrouter' provider via init_chat_model.""" mock_init.return_value = "mock_model" monkeypatch.setenv("OPENROUTER_API_KEY", "or-key-456") - # Assert the DEFAULT effort, so isolate from any leaked env override. - monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False) - get_chat_model("x-ai/grok-4.3", provider="openrouter") call_kwargs = mock_init.call_args[1] @@ -524,8 +477,10 @@ class TestThirdPartyRouting: assert call_kwargs["reasoning"] == {"effort": "low"} @patch("EvoScientist.llm.models.init_chat_model") - def test_openrouter_reasoning_effort_from_env(self, mock_init, monkeypatch): - """Reasoning effort should be configurable via env var.""" + def test_openrouter_reasoning_effort_environment_is_ignored( + self, mock_init, monkeypatch + ): + """The deployment environment cannot alter an invocation parameter.""" mock_init.return_value = "mock_model" monkeypatch.setenv("OPENROUTER_API_KEY", "or-key") monkeypatch.setenv("EVOSCIENTIST_REASONING_EFFORT", "medium") @@ -533,7 +488,7 @@ class TestThirdPartyRouting: get_chat_model("x-ai/grok-4.3", provider="openrouter") call_kwargs = mock_init.call_args[1] - assert call_kwargs["reasoning"] == {"effort": "medium", "summary": "auto"} + assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"} # --- OpenRouter app attribution (issue #339) --- @@ -634,7 +589,6 @@ class TestThirdPartyRouting: """Attribution must not disturb reasoning or Anthropic prompt caching.""" mock_init.return_value = "mock_model" monkeypatch.setenv("OPENROUTER_API_KEY", "or-key") - monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False) monkeypatch.delenv( "EVOSCIENTIST_OPENROUTER_ANTHROPIC_PROMPT_CACHE", raising=False ) @@ -1233,6 +1187,57 @@ class TestFlattenMessageContent: # ============================================================================= +class TestOpenAIEmptySSEKeepalivePatch: + def test_blank_sse_keepalive_is_skipped(self): + from EvoScientist.llm.patches import _is_blank_sse_keepalive + + class Event: + def __init__(self, data): + self.data = data + + assert _is_blank_sse_keepalive(Event("")) + assert _is_blank_sse_keepalive(Event(" \t\n")) + assert _is_blank_sse_keepalive(Event(None)) + assert not _is_blank_sse_keepalive(Event('{"type":"response.output_text"}')) + assert not _is_blank_sse_keepalive(Event("[DONE]")) + + @pytest.mark.asyncio + async def test_async_stream_filters_blank_keepalive_before_json_parse(self): + from openai._streaming import AsyncStream, ServerSentEvent + + class Decoder: + async def aiter_bytes(self, _bytes): + yield ServerSentEvent(data="") + yield ServerSentEvent(data='{"type":"response.created"}') + + class Response: + async def aiter_bytes(self): + if False: + yield b"" + + stream = type("Stream", (), {"_decoder": Decoder(), "response": Response()})() + events = [event async for event in AsyncStream._iter_events(stream)] + + assert [event.data for event in events] == ['{"type":"response.created"}'] + + def test_sync_stream_filters_blank_keepalive_before_json_parse(self): + from openai._streaming import ServerSentEvent, Stream + + class Decoder: + def iter_bytes(self, _bytes): + yield ServerSentEvent(data="") + yield ServerSentEvent(data='{"type":"response.created"}') + + class Response: + def iter_bytes(self): + return iter(()) + + stream = type("Stream", (), {"_decoder": Decoder(), "response": Response()})() + events = list(Stream._iter_events(stream)) + + assert [event.data for event in events] == ['{"type":"response.created"}'] + + class TestPatchOpenAICompatContent: """Verify content flattening covers _generate, _agenerate, _stream, _astream.""" @@ -1305,6 +1310,7 @@ class TestPatchOpenAICompatContent: { "type": "tool_call", "id": "wrong-id", + "call_id": "wrong-call-id", "name": "wrong-name", "args": {}, } @@ -1316,6 +1322,7 @@ class TestPatchOpenAICompatContent: ) assert normalized[0].content[0]["id"] == "call-1" + assert normalized[0].content[0]["call_id"] == "call-1" assert normalized[0].content[0]["name"] == "execute" def test_invalid_tool_call_is_not_replayed_to_responses_api(self): @@ -1394,6 +1401,25 @@ class TestPatchOpenAICompatContent: assert [message.type for message in normalized] == ["ai", "human"] assert normalized[0].tool_calls == [] + def test_nonportable_reasoning_metadata_is_removed_for_cross_model_replay(self): + from langchain_core.messages import AIMessage + + from EvoScientist.llm.patches import _sanitize_messages + + message = AIMessage( + content="portable answer", + additional_kwargs={ + "reasoning_content": "provider-specific trace", + "reasoning_details": [{"type": "reasoning"}], + "safe_field": "preserved", + }, + ) + + normalized = _sanitize_messages([message], drop_reasoning_metadata=True) + + assert normalized[0].additional_kwargs == {"safe_field": "preserved"} + assert "reasoning_content" in message.additional_kwargs + def test_generate_flattened(self): from langchain_core.messages import HumanMessage @@ -2675,11 +2701,6 @@ class TestPatchOpenrouterStripResponsesReasoning: class TestAutoConfig: - @pytest.fixture(autouse=True) - def _clear_reasoning_effort_env(self, monkeypatch): - """Keep auto-config tests independent of the developer environment.""" - monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False) - @patch("EvoScientist.llm.models.init_chat_model") def test_internal_sentinels_disable_auto_reasoning(self, mock_init, monkeypatch): """Internal callers can disable reasoning without leaking sentinels.""" @@ -2796,8 +2817,6 @@ class TestAutoConfig: """gpt-5.4+ and codex models get xhigh reasoning.""" mock_init.return_value = "mock_model" monkeypatch.delenv("OPENAI_BASE_URL", raising=False) - monkeypatch.delenv("EVOSCIENTIST_REASONING_EFFORT", raising=False) - get_chat_model("gpt-5.4", provider="openai") assert mock_init.call_args[1]["reasoning"] == { "effort": "xhigh", @@ -2823,8 +2842,10 @@ class TestAutoConfig: } @patch("EvoScientist.llm.models.init_chat_model") - def test_openai_reasoning_effort_from_env(self, mock_init, monkeypatch): - """Native OpenAI reasoning effort should be configurable via env var.""" + def test_openai_reasoning_effort_environment_is_ignored( + self, mock_init, monkeypatch + ): + """The deployment environment cannot alter an invocation parameter.""" mock_init.return_value = "mock_model" monkeypatch.delenv("OPENAI_BASE_URL", raising=False) monkeypatch.setenv("EVOSCIENTIST_REASONING_EFFORT", "high") @@ -2832,7 +2853,7 @@ class TestAutoConfig: get_chat_model("gpt-5.5", provider="openai") assert mock_init.call_args[1]["reasoning"] == { - "effort": "high", + "effort": "xhigh", "summary": "auto", } @@ -2867,10 +2888,10 @@ class TestAutoConfig: assert call_kwargs["model_provider"] == "openai" assert call_kwargs["base_url"] == "http://127.0.0.1:8000/codex/v1" assert call_kwargs["api_key"] == "ccproxy-oauth" - # ccproxy uses the Responses API, so reasoning configuration is valid. + # Endpoint detection may add compatible client headers, but cannot + # select an API protocol. The compiled invocation plan owns that. assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"} - # Proxy mode: Responses API (bypasses format chain), streaming ON - assert call_kwargs["use_responses_api"] is True + assert "use_responses_api" not in call_kwargs assert "streaming" not in call_kwargs @patch("EvoScientist.llm.models.init_chat_model") @@ -3054,72 +3075,23 @@ class TestAutoConfig: call_kwargs = mock_init.call_args[1] assert call_kwargs["include_thoughts"] is True + @pytest.mark.parametrize("env_value", ["false", "true", " TRUE "]) @patch("EvoScientist.llm.models.init_chat_model") - def test_use_responses_api_false(self, mock_init, monkeypatch): - """use_responses_api=false forces Chat Completions and drops reasoning.""" - mock_init.return_value = "mock_model" - monkeypatch.delenv("OPENAI_BASE_URL", raising=False) - monkeypatch.setenv("EVOSCIENTIST_USE_RESPONSES_API", "false") - - get_chat_model("gpt-5-nano", provider="openai") - - call_kwargs = mock_init.call_args[1] - assert call_kwargs["use_responses_api"] is False - assert "reasoning" not in call_kwargs - - @patch("EvoScientist.llm.models.init_chat_model") - def test_use_responses_api_true(self, mock_init, monkeypatch): - """use_responses_api=true explicitly enables the Responses API.""" - mock_init.return_value = "mock_model" - monkeypatch.delenv("OPENAI_BASE_URL", raising=False) - monkeypatch.setenv("EVOSCIENTIST_USE_RESPONSES_API", "true") - - get_chat_model("gpt-5-nano", provider="openai") - - call_kwargs = mock_init.call_args[1] - assert call_kwargs["use_responses_api"] is True - assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"} - - @patch("EvoScientist.llm.models.init_chat_model") - def test_use_responses_api_default_unchanged(self, mock_init, monkeypatch): - """Empty use_responses_api preserves default behavior (no kwarg set).""" - mock_init.return_value = "mock_model" - monkeypatch.delenv("OPENAI_BASE_URL", raising=False) - monkeypatch.delenv("EVOSCIENTIST_USE_RESPONSES_API", raising=False) - - get_chat_model("gpt-5-nano", provider="openai") - - call_kwargs = mock_init.call_args[1] - assert "use_responses_api" not in call_kwargs - assert call_kwargs["reasoning"] == {"effort": "high", "summary": "auto"} - - @pytest.mark.parametrize("env_value", ["FALSE", " false ", "False"]) - @patch("EvoScientist.llm.models.init_chat_model") - def test_use_responses_api_false_normalization( + def test_response_api_environment_cannot_override_an_explicit_call_plan( self, mock_init, monkeypatch, env_value ): - """Case/whitespace variants of 'false' are normalized correctly.""" mock_init.return_value = "mock_model" monkeypatch.delenv("OPENAI_BASE_URL", raising=False) monkeypatch.setenv("EVOSCIENTIST_USE_RESPONSES_API", env_value) - get_chat_model("gpt-5-nano", provider="openai") + get_chat_model( + "gpt-5-nano", + provider="openai", + use_responses_api=False, + reasoning_effort="high", + ) call_kwargs = mock_init.call_args[1] assert call_kwargs["use_responses_api"] is False + assert call_kwargs["reasoning_effort"] == "high" assert "reasoning" not in call_kwargs - - @pytest.mark.parametrize("env_value", ["TRUE", " true ", "True"]) - @patch("EvoScientist.llm.models.init_chat_model") - def test_use_responses_api_true_normalization( - self, mock_init, monkeypatch, env_value - ): - """Case/whitespace variants of 'true' are normalized correctly.""" - mock_init.return_value = "mock_model" - monkeypatch.delenv("OPENAI_BASE_URL", raising=False) - monkeypatch.setenv("EVOSCIENTIST_USE_RESPONSES_API", env_value) - - get_chat_model("gpt-5-nano", provider="openai") - - call_kwargs = mock_init.call_args[1] - assert call_kwargs["use_responses_api"] is True diff --git a/tests/test_memory_agent_factory.py b/tests/test_memory_agent_factory.py new file mode 100644 index 0000000..9d9085d --- /dev/null +++ b/tests/test_memory_agent_factory.py @@ -0,0 +1,37 @@ +"""Tests for bounded skill context in background memory-agent graphs.""" + +from __future__ import annotations + + +def test_memory_agent_factory_replaces_framework_skill_injection(monkeypatch, tmp_path): + from EvoScientist.memory.agents import _factory + from EvoScientist.middleware import BudgetedSkillsMiddleware + + captured = {} + + class _Agent: + def with_config(self, _config): + return self + + def _create_deep_agent(**kwargs): + captured.update(kwargs) + return _Agent() + + monkeypatch.setattr("deepagents.create_deep_agent", _create_deep_agent) + monkeypatch.setattr( + "EvoScientist.EvoScientist._ensure_auxiliary_chat_model", lambda: object() + ) + + _factory.build_memory_agent_graph( + name="test-memory-agent", + system_prompt="Test prompt", + memory_dir=tmp_path / "memory", + workspace_dir=tmp_path / "workspace", + tools=[], + middleware=[], + skills=["/skills/"], + backend=object(), + ) + + assert captured["skills"] is None + assert isinstance(captured["middleware"][-1], BudgetedSkillsMiddleware) diff --git a/tests/test_model_config_v3.py b/tests/test_model_config_v3.py new file mode 100644 index 0000000..5449575 --- /dev/null +++ b/tests/test_model_config_v3.py @@ -0,0 +1,460 @@ +from __future__ import annotations + +import sqlite3 +import uuid + +import pytest +import yaml + +from EvoScientist.llm.config_admin import EvoModelConfigAdminService +from EvoScientist.llm.contracts import ( + CommitModelConfigRequest, + EvoRuntimeError, + HmacGrantAuthority, + ProbeCandidateRouteRequest, + ValidateCandidateConfigRequest, +) +from EvoScientist.llm.crypto import canonical_json_v1, sha256_id +from EvoScientist.llm.model_config import EvoModelConfig, FileEvoModelConfigStore +from tests.v3_fixtures import ( + RUNTIME_KEY_ID, + RUNTIME_SECRET, + identity_ring, + v3_payload, +) + + +def _grant(authority, *, action, operation_id, payload): + return authority.sign_admin( + subject_id="admin", + action=action, + operation_id=operation_id, + request_digest=sha256_id(payload), + ttl_ms=60_000, + ) + + +@pytest.mark.asyncio +async def test_validate_probe_commit_bootstraps_v2_config(monkeypatch, tmp_path): + monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "test-secret") + authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID) + store = FileEvoModelConfigStore( + tmp_path / "model_routes.yaml", admin_verifier=authority + ) + service = EvoModelConfigAdminService( + store, + grant_authority=authority, + identity_key_ring=identity_ring(), + probe_runner=lambda *_args: True, + ) + candidate = v3_payload() + candidate.pop("capability_evidence") + validate_operation = str(uuid.uuid4()) + validate_payload = { + "operation_id": validate_operation, + "expected_revision": 0, + "payload": candidate, + } + result = service.validate_candidate( + ValidateCandidateConfigRequest( + validate_operation, + 0, + candidate, + _grant( + authority, + action="model_config:validate", + operation_id=validate_operation, + payload=validate_payload, + ), + ) + ) + evidence_ids = [] + for route in result.concrete_routes: + for probe_kind in route.required_probe_kinds: + operation_id = str(uuid.uuid4()) + payload = { + "operation_id": operation_id, + "proposal_hash": result.proposal_hash, + "route_semantics_hash": route.route_semantics_hash, + "probe_kind": probe_kind, + } + probe = await service.probe_candidate( + ProbeCandidateRouteRequest( + operation_id, + result.proposal_hash, + route.route_semantics_hash, + probe_kind, + _grant( + authority, + action="model_config:probe", + operation_id=operation_id, + payload=payload, + ), + ) + ) + evidence_ids.append(probe.evidence_id) + commit_operation = str(uuid.uuid4()) + commit_payload = { + "operation_id": commit_operation, + "expected_revision": 0, + "proposal_hash": result.proposal_hash, + "payload": candidate, + "evidence_ids": sorted(set(evidence_ids)), + } + committed = service.commit( + CommitModelConfigRequest( + commit_operation, + 0, + result.proposal_hash, + candidate, + tuple(commit_payload["evidence_ids"]), + _grant( + authority, + action="model_config:commit", + operation_id=commit_operation, + payload=commit_payload, + ), + ) + ) + + assert committed.config_revision == 1 + assert store.load().config_revision == 1 + + +def test_v2_parser_rejects_old_alias_and_manual_capability_boundary(): + payload = v3_payload() + payload["providers"]["custom-openai"]["models"][0]["alias"] = "legacy" + with pytest.raises(EvoRuntimeError, match="unknown fields"): + EvoModelConfig.parse(payload) + + payload = v3_payload() + payload["capability_requirements"] = [] + with pytest.raises(EvoRuntimeError, match="unknown fields"): + EvoModelConfig.parse(payload) + + +def test_model_output_capability_rejects_oversized_output_token_limit(): + payload = v3_payload() + payload["providers"]["custom-openai"]["models"][0]["params"] = { + "output_token_limit": 4096 + } + + with pytest.raises( + EvoRuntimeError, match="output_token_limit exceeds model capability" + ): + EvoModelConfig.parse(payload) + + +def test_admin_grant_is_bound_to_action_and_payload(monkeypatch, tmp_path): + monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "test-secret") + authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID) + service = EvoModelConfigAdminService( + FileEvoModelConfigStore( + tmp_path / "model_routes.yaml", admin_verifier=authority + ), + grant_authority=authority, + identity_key_ring=identity_ring(), + ) + candidate = v3_payload() + candidate.pop("capability_evidence") + operation_id = str(uuid.uuid4()) + wrong_payload = { + "operation_id": operation_id, + "expected_revision": 1, + "payload": candidate, + } + with pytest.raises(EvoRuntimeError, match="ADMIN_CONFIG_FORBIDDEN"): + service.validate_candidate( + ValidateCandidateConfigRequest( + operation_id, + 0, + candidate, + _grant( + authority, + action="model_config:validate", + operation_id=operation_id, + payload=wrong_payload, + ), + ) + ) + + +def test_validate_operation_replays_and_rejects_changed_payload(monkeypatch, tmp_path): + monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "test-secret") + authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID) + service = EvoModelConfigAdminService( + FileEvoModelConfigStore(tmp_path / "model_routes.yaml"), + grant_authority=authority, + identity_key_ring=identity_ring(), + ) + candidate = v3_payload() + candidate.pop("capability_evidence") + operation_id = str(uuid.uuid4()) + payload = { + "operation_id": operation_id, + "expected_revision": 0, + "payload": candidate, + } + request = ValidateCandidateConfigRequest( + operation_id, + 0, + candidate, + _grant( + authority, + action="model_config:validate", + operation_id=operation_id, + payload=payload, + ), + ) + first = service.validate_candidate(request) + replay = service.validate_candidate(request) + assert replay == first + + changed = {**candidate, "runtime_defaults": {"max_retries": 1}} + changed_payload = {**payload, "payload": changed} + with pytest.raises(EvoRuntimeError, match="IDEMPOTENCY_CONFLICT"): + service.validate_candidate( + ValidateCandidateConfigRequest( + operation_id, + 0, + changed, + _grant( + authority, + action="model_config:validate", + operation_id=operation_id, + payload=changed_payload, + ), + ) + ) + + +def test_provider_defaults_accept_values_up_to_v3_bounds(): + from EvoScientist.llm.model_config import _parse_v3_provider_defaults + + runtime = { + "connect_timeout_seconds": 10, + "first_event_timeout_seconds": 60, + "stream_idle_timeout_seconds": 60, + "attempt_timeout_seconds": 600, + } + defaults = _parse_v3_provider_defaults( + {"connect_timeout_seconds": 60, "attempt_timeout_seconds": 600}, runtime + ) + assert defaults["connect_timeout_seconds"] == 60 + assert defaults["attempt_timeout_seconds"] == 600 + + +def test_provider_defaults_reject_values_beyond_v3_bounds(): + from EvoScientist.llm.model_config import _parse_v3_provider_defaults + + runtime = { + "connect_timeout_seconds": 10, + "first_event_timeout_seconds": 60, + "stream_idle_timeout_seconds": 60, + "attempt_timeout_seconds": 600, + } + with pytest.raises(EvoRuntimeError): + _parse_v3_provider_defaults({"connect_timeout_seconds": 61}, runtime) + + +def test_store_recovers_renamed_target_from_preparing_journal(tmp_path): + path = tmp_path / "model_routes.yaml" + ops_path = tmp_path / "model_config_ops.sqlite" + store = FileEvoModelConfigStore(path, ops_path=ops_path) + store.bootstrap_for_development(v3_payload(), operation_id="bootstrap") + + target = v3_payload(revision=2) + payload_hash = sha256_id(target) + path.write_text( + yaml.safe_dump(target, allow_unicode=True, sort_keys=False), + encoding="utf-8", + ) + with sqlite3.connect(ops_path) as connection: + connection.execute( + """INSERT INTO config_operations + (subject_id, action, operation_id, request_digest, status, + expected_revision, target_revision, payload_hash, + canonical_payload, created_at, updated_at) + VALUES ('admin', 'model_config:commit', 'recover-op', ?, + 'PREPARING', 1, 2, ?, ?, 1, 1)""", + ( + payload_hash, + payload_hash, + canonical_json_v1(target).decode("utf-8"), + ), + ) + + recovered = FileEvoModelConfigStore(path, ops_path=ops_path) + assert recovered.load().config_revision == 2 + with sqlite3.connect(ops_path) as connection: + status = connection.execute( + "SELECT status FROM config_operations WHERE operation_id='recover-op'" + ).fetchone()[0] + assert status == "COMMITTED" + + +@pytest.mark.asyncio +async def test_probe_dispatch_is_durable_and_operation_replay_is_free( + monkeypatch, tmp_path +): + monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "test-secret") + authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID) + sink_events = [] + provider_calls = 0 + + async def sink(payload): + sink_events.append(dict(payload)) + if payload["outcome"] == "started": + return "committed" + return f"terminal:{payload['probe_result']}" + + async def runner(*_args): + nonlocal provider_calls + provider_calls += 1 + return True + + service = EvoModelConfigAdminService( + FileEvoModelConfigStore(tmp_path / "model_routes.yaml"), + grant_authority=authority, + identity_key_ring=identity_ring(), + probe_runner=runner, + probe_event_sink=sink, + ) + candidate = v3_payload() + candidate.pop("capability_evidence") + validate_id = str(uuid.uuid4()) + validate_payload = { + "operation_id": validate_id, + "expected_revision": 0, + "payload": candidate, + } + validated = service.validate_candidate( + ValidateCandidateConfigRequest( + validate_id, + 0, + candidate, + _grant( + authority, + action="model_config:validate", + operation_id=validate_id, + payload=validate_payload, + ), + ) + ) + route = validated.concrete_routes[0] + operation_id = str(uuid.uuid4()) + payload = { + "operation_id": operation_id, + "proposal_hash": validated.proposal_hash, + "route_semantics_hash": route.route_semantics_hash, + "probe_kind": route.required_probe_kinds[0], + } + request = ProbeCandidateRouteRequest( + operation_id, + validated.proposal_hash, + route.route_semantics_hash, + route.required_probe_kinds[0], + _grant( + authority, + action="model_config:probe", + operation_id=operation_id, + payload=payload, + ), + ) + first = await service.probe_candidate(request) + replay = await service.probe_candidate(request) + + assert replay == first + assert provider_calls == 1 + assert [event["outcome"] for event in sink_events] == [ + "started", + "usage_unconfirmed", + ] + assert all(event["billing_intent"] == "platform_cost" for event in sink_events) + + +@pytest.mark.asyncio +async def test_commit_rejects_probe_evidence_after_runtime_secret_changes( + monkeypatch, tmp_path +): + monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "secret-before-probe") + authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID) + store = FileEvoModelConfigStore(tmp_path / "model_routes.yaml") + service = EvoModelConfigAdminService( + store, + grant_authority=authority, + identity_key_ring=identity_ring(), + probe_runner=lambda *_args: True, + ) + candidate = v3_payload() + candidate.pop("capability_evidence") + validate_id = str(uuid.uuid4()) + validate_payload = { + "operation_id": validate_id, + "expected_revision": 0, + "payload": candidate, + } + validated = service.validate_candidate( + ValidateCandidateConfigRequest( + validate_id, + 0, + candidate, + _grant( + authority, + action="model_config:validate", + operation_id=validate_id, + payload=validate_payload, + ), + ) + ) + evidence_ids = [] + for route in validated.concrete_routes: + for probe_kind in route.required_probe_kinds: + operation_id = str(uuid.uuid4()) + probe_payload = { + "operation_id": operation_id, + "proposal_hash": validated.proposal_hash, + "route_semantics_hash": route.route_semantics_hash, + "probe_kind": probe_kind, + } + result = await service.probe_candidate( + ProbeCandidateRouteRequest( + operation_id, + validated.proposal_hash, + route.route_semantics_hash, + probe_kind, + _grant( + authority, + action="model_config:probe", + operation_id=operation_id, + payload=probe_payload, + ), + ) + ) + evidence_ids.append(result.evidence_id) + + monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "secret-after-probe") + commit_id = str(uuid.uuid4()) + commit_payload = { + "operation_id": commit_id, + "expected_revision": 0, + "proposal_hash": validated.proposal_hash, + "payload": candidate, + "evidence_ids": sorted(evidence_ids), + } + with pytest.raises(EvoRuntimeError, match="CAPABILITY_EVIDENCE_STALE"): + service.commit( + CommitModelConfigRequest( + commit_id, + 0, + validated.proposal_hash, + candidate, + tuple(commit_payload["evidence_ids"]), + _grant( + authority, + action="model_config:commit", + operation_id=commit_id, + payload=commit_payload, + ), + ) + ) diff --git a/tests/test_model_config_v4.py b/tests/test_model_config_v4.py new file mode 100644 index 0000000..179687a --- /dev/null +++ b/tests/test_model_config_v4.py @@ -0,0 +1,822 @@ +from __future__ import annotations + +import json +import sqlite3 +import uuid +from datetime import UTC, datetime, timedelta + +import pytest + +from EvoScientist.llm.contracts import EvoRuntimeError, HmacGrantAuthority +from EvoScientist.llm.crypto import HmacKeyRing, KeyMaterial, sha256_id +from EvoScientist.llm.model_config_v4 import ( + EvoModelConfig, + UnifiedModelConfigStore, + build_supported_v3_evidence, + model_profile_identity_map, + normalize_v4_config, + project_v4_to_v3, +) +from EvoScientist.llm.runtime import EvoModelRuntime + + +def _payload(*, provider_id: str | None = None, profile_id: str | None = None) -> dict: + provider_id = provider_id or str(uuid.uuid4()) + profile_id = profile_id or str(uuid.uuid4()) + return { + "schema_version": 4, + "default_model_profile_id": profile_id, + "purpose_defaults": {}, + "purpose_call_limits": {}, + "providers": [ + { + "provider_id": provider_id, + "display_name": "OpenAI Production", + "adapter_id": "openai", + "adapter_revision": "openai-v1", + "enabled": True, + "connection": {"base_url": "https://api.openai.com/v1"}, + "models": [ + { + "model_profile_id": profile_id, + "provider_model_id": "gpt-test", + "display_name": "General", + "enabled": True, + "version_policy": "rolling", + "invocation": { + "api_mode": "responses", + "tool_call_transport": "native", + }, + "capabilities": {"text": True, "tools": True}, + "limits": { + "context_tokens": 128_000, + "max_output_tokens": 8_192, + }, + "parameters": {}, + "billing": { + "sku": "internal/gpt-test", + "pricing_revision": "configured-v1", + "currency": "CNY", + "unit_scale": 1_000_000, + "input_microunits_per_million": 1_000_000, + "cached_microunits_per_million": 200_000, + "output_microunits_per_million": 4_000_000, + "multiplier": 1, + }, + } + ], + } + ], + } + + +def _ring() -> HmacKeyRing: + return HmacKeyRing(KeyMaterial.create("identity-v1", "i" * 32)) + + +def test_v4_normalization_uses_stable_profile_identity(): + payload = normalize_v4_config(_payload()) + identities = model_profile_identity_map(payload) + + assert payload["schema_version"] == 4 + assert payload["default_model_profile_id"] in identities + assert identities[payload["default_model_profile_id"]][1] == "gpt-test" + + +def test_v4_rejects_enabled_model_without_billing_configuration(): + raw = _payload() + raw["providers"][0]["models"][0].pop("billing") + + with pytest.raises(EvoRuntimeError) as raised: + normalize_v4_config(raw, lenient=True, warnings=[]) + + assert raised.value.details[0]["code"] == "CONFIG_PRICING_REQUIRED" + + +def test_v4_accepts_explicit_zero_price_for_free_model(): + raw = _payload() + raw["providers"][0]["models"][0]["billing"].update( + input_microunits_per_million=0, + cached_microunits_per_million=0, + output_microunits_per_million=0, + ) + + payload = normalize_v4_config(raw, lenient=True, warnings=[]) + + assert payload["providers"][0]["models"][0]["billing"][ + "input_microunits_per_million" + ] == 0 + + +def test_v4_normalization_preserves_model_output_token_limit(): + raw = _payload() + raw["providers"][0]["models"][0]["parameters"] = { + "defaults": {"output_token_limit": 2048, "temperature": 0.2}, + "purpose_overrides": { + "main_agent": {"output_token_limit": 1024, "temperature": 0.3} + }, + "user_options": { + "output_token_limit": { + "default": 512, + "applies_to": ["main_agent"], + } + }, + } + + payload = normalize_v4_config(raw) + parameters = payload["providers"][0]["models"][0]["parameters"] + + assert parameters["defaults"] == {"output_token_limit": 2048, "temperature": 0.2} + assert parameters["purpose_overrides"]["main_agent"] == { + "output_token_limit": 1024, + "temperature": 0.3, + } + assert parameters["user_options"] == {} + for limit in payload["purpose_call_limits"].values(): + assert "max_output_tokens" not in limit + assert payload["purpose_call_limits"]["main_agent"] == { + "max_attempts_per_run": 4 + } + + +def test_v4_rejects_output_token_limit_above_model_capability(): + raw = _payload() + raw["providers"][0]["models"][0]["parameters"] = { + "defaults": {"output_token_limit": 16_384} + } + + with pytest.raises(EvoRuntimeError) as raised: + normalize_v4_config(raw) + + assert raised.value.code == "LLM_ROUTE_CONFIGURATION_REQUIRED" + assert raised.value.details[0]["code"] == "CONFIG_LIMIT_EXCEEDED" + + +def test_v4_lenient_clamps_output_token_limit_to_model_capability(): + raw = _payload() + raw["providers"][0]["models"][0]["parameters"] = { + "defaults": {"output_token_limit": 16_384} + } + warnings: list[dict] = [] + + payload = normalize_v4_config(raw, lenient=True, warnings=warnings) + + defaults = payload["providers"][0]["models"][0]["parameters"]["defaults"] + assert defaults["output_token_limit"] == 8_192 + assert any(item["code"] == "CONFIG_LIMIT_EXCEEDED" for item in warnings) + + +def test_v4_removes_legacy_tool_transport_from_admin_config(): + raw = _payload() + raw["providers"][0]["models"][0]["capabilities"]["tools"] = False + + payload = normalize_v4_config(raw) + + model = payload["providers"][0]["models"][0] + assert model["invocation"] == {"api_mode": "responses"} + + +def test_v4_runtime_projection_derives_tool_transport_from_capabilities(): + raw = _payload() + raw["providers"][0]["models"][0]["capabilities"]["tools"] = False + warnings: list[dict] = [] + + payload = normalize_v4_config(raw, lenient=True, warnings=warnings) + + model = payload["providers"][0]["models"][0] + assert "tool_call_transport" not in model["invocation"] + assert not any( + item["code"] == "CONFIG_INVOCATION_CONTRACT_INVALID" for item in warnings + ) + + projection = project_v4_to_v3( + payload, + revision=1, + identity_key_id=_ring().current.key_id, + bindings={payload["providers"][0]["provider_id"]: 1}, + ) + config = EvoModelConfig.parse(projection, require_evidence=False) + route = config.concrete_routes(next(iter(config.route_selectors)))[0] + assert route.tool_call_transport == "disabled" + + +def test_v4_rejects_non_uuid_business_ids(): + raw = _payload(provider_id="openai-main") + + with pytest.raises(EvoRuntimeError) as raised: + normalize_v4_config(raw) + + assert raised.value.code == "LLM_ROUTE_CONFIGURATION_REQUIRED" + + +@pytest.mark.parametrize( + "base_url", + [ + "", + "not a url", + "ftp://user:pass@example.test/path?key=value#fragment", + "http://127.0.0.1:11434/v1", + "https://169.254.169.254/latest/meta-data", + "custom://model.internal:70000/path", + ], +) +def test_v4_base_url_is_not_validated(base_url): + raw = _payload() + raw["providers"][0]["connection"]["base_url"] = base_url + + payload = normalize_v4_config(raw) + provider_id = payload["providers"][0]["provider_id"] + projection = project_v4_to_v3( + payload, + revision=1, + identity_key_id=_ring().current.key_id, + bindings={provider_id: 1}, + ) + + assert payload["providers"][0]["connection"]["base_url"] == base_url + assert projection["providers"][0]["connection"]["base_url"] == base_url + + +def test_unified_store_commits_config_and_secret_without_evidence_gate(tmp_path): + ring = _ring() + store = UnifiedModelConfigStore( + tmp_path / "model-config.sqlite", + identity_key_ring=ring, + encryption_keys={"enc-v1": "e" * 32}, + current_encryption_key_id="enc-v1", + ) + payload = normalize_v4_config(_payload()) + provider_id = payload["providers"][0]["provider_id"] + profile_id = payload["default_model_profile_id"] + plan = store.plan_credentials(payload, {provider_id: "sk-secret-value"}) + projection = project_v4_to_v3( + payload, + revision=1, + identity_key_id=ring.current.key_id, + bindings=plan.bindings, + ) + evidence = build_supported_v3_evidence( + projection, + identity_key_ring=ring, + verified_profiles={provider_id: frozenset({profile_id})}, + ) + request_hash = sha256_id({"payload": payload, "fingerprints": plan.fingerprints}) + store.begin_operation( + "save-1", + request_hash=request_hash, + expected_revision=0, + actor="admin", + lease_owner="worker-1", + ) + result = store.commit( + payload, + expected_revision=0, + operation_id="save-1", + actor="admin", + request_hash=request_hash, + credential_plan=plan, + evidence=evidence, + ) + + assert result.config_revision == 1 + assert result.changed is True + runtime = store.load() + assert runtime.main_routes.default_alias == profile_id + assert store.resolve( + runtime.providers[provider_id].endpoints[provider_id].auth + ).value == "sk-secret-value" + revision, admin_payload, credentials = store.get_admin_config() + assert revision == 1 + assert admin_payload == payload + assert credentials[provider_id].configured is True + assert "sk-secret-value" not in str(credentials) + with sqlite3.connect(store.path) as connection: + assert connection.execute("SELECT COUNT(*) FROM model_config_revisions").fetchone()[0] == 1 + assert connection.execute("SELECT COUNT(*) FROM provider_secret_versions").fetchone()[0] == 1 + assert connection.execute("SELECT COUNT(*) FROM capability_evidence").fetchone()[0] == 0 + store.record_runtime_observation( + config_revision=1, + provider_id=provider_id, + model_profile_id=profile_id, + provider_model_id="gpt-test", + api_mode="responses", + purpose="main_agent", + outcome="succeeded", + error_code=None, + strategy={"structured_output": False}, + ) + summary = store.runtime_observation_summary() + assert len(summary) == 1 + assert summary[0]["model_profile_id"] == profile_id + assert summary[0]["provider_id"] == provider_id + assert summary[0]["latest_outcome"] == "succeeded" + assert summary[0]["latest_error_code"] is None + assert summary[0]["total_calls"] == 1 + assert summary[0]["successful_calls"] == 1 + assert summary[0]["strategy"] == {"structured_output": False} + assert summary[0]["latest_at"] + + +def test_unified_store_idempotency_rejects_same_id_different_request(tmp_path): + store = UnifiedModelConfigStore( + tmp_path / "model-config.sqlite", + identity_key_ring=_ring(), + encryption_keys={"enc-v1": "e" * 32}, + current_encryption_key_id="enc-v1", + ) + store.begin_operation( + "save-1", + request_hash="hmac:one", + expected_revision=0, + actor="admin", + lease_owner="worker-1", + ) + + with pytest.raises(EvoRuntimeError) as raised: + store.begin_operation( + "save-1", + request_hash="hmac:two", + expected_revision=0, + actor="admin", + lease_owner="worker-2", + ) + + assert raised.value.code == "IDEMPOTENCY_CONFLICT" + + +def test_operation_failure_preserves_last_progress(tmp_path): + store = UnifiedModelConfigStore( + tmp_path / "model-config.sqlite", + identity_key_ring=_ring(), + encryption_keys={"enc-v1": "e" * 32}, + current_encryption_key_id="enc-v1", + ) + store.begin_operation( + "save-1", request_hash="request-1", expected_revision=0, + actor="admin", lease_owner="worker", + ) + store.set_operation_status( + "save-1", "PROBING", progress={"completed": 2, "total": 3} + ) + store.set_operation_status( + "save-1", + "FAILED_RETRYABLE", + stage="PROBING", + error_code="MODEL_PROBE_FAILED", + error_details=({"path": "models.test", "code": "MODEL_TIMEOUT"},), + ) + + operation = store.get_operation("save-1") + assert operation is not None + assert operation["progress"] == {"completed": 2, "total": 3} + assert operation["error_details"] == [ + {"path": "models.test", "code": "MODEL_TIMEOUT"} + ] + + + +def test_unified_store_load_ignores_historical_evidence(tmp_path): + ring = _ring() + store = UnifiedModelConfigStore( + tmp_path / "model-config.sqlite", + identity_key_ring=ring, + encryption_keys={"enc-v1": "e" * 32}, + current_encryption_key_id="enc-v1", + ) + payload = normalize_v4_config(_payload()) + provider_id = payload["providers"][0]["provider_id"] + profile_id = payload["default_model_profile_id"] + plan = store.plan_credentials(payload, {provider_id: "sk-secret-value"}) + projection = project_v4_to_v3( + payload, revision=1, identity_key_id=ring.current.key_id, bindings=plan.bindings + ) + evidence = build_supported_v3_evidence( + projection, + identity_key_ring=ring, + verified_profiles={provider_id: frozenset({profile_id})}, + ) + store.begin_operation( + "save-1", request_hash="request-1", expected_revision=0, + actor="admin", lease_owner="worker", + ) + store.commit( + payload, expected_revision=0, operation_id="save-1", actor="admin", + request_hash="request-1", credential_plan=plan, evidence=evidence, + ) + # The supplied evidence is intentionally not persisted or read by load. + # A legacy/tampered evidence row therefore cannot make the revision stale. + assert store.load_revision(1).config_revision == 1 + + +def test_rollback_publishes_a_new_monotonic_revision(tmp_path): + ring = _ring() + store = UnifiedModelConfigStore( + tmp_path / "model-config.sqlite", + identity_key_ring=ring, + encryption_keys={"enc-v1": "e" * 32}, + current_encryption_key_id="enc-v1", + ) + first = normalize_v4_config(_payload()) + provider_id = first["providers"][0]["provider_id"] + profile_id = first["default_model_profile_id"] + def publish(payload: dict, revision: int, operation_id: str, request_hash: str): + plan = store.plan_credentials( + payload, {provider_id: "sk-secret-value"} if revision == 1 else {} + ) + projection = project_v4_to_v3( + payload, + revision=revision, + identity_key_id=ring.current.key_id, + bindings=plan.bindings, + ) + evidence = build_supported_v3_evidence( + projection, + identity_key_ring=ring, + verified_profiles={provider_id: frozenset({profile_id})}, + ) + store.begin_operation( + operation_id, request_hash=request_hash, expected_revision=revision - 1, + actor="admin", lease_owner="worker", + ) + return store.commit( + payload, expected_revision=revision - 1, operation_id=operation_id, + actor="admin", request_hash=request_hash, credential_plan=plan, + evidence=evidence, + ) + + publish(first, 1, "save-1", "request-1") + second = json.loads(json.dumps(first)) + second["providers"][0]["display_name"] = "Changed display name" + publish(second, 2, "save-2", "request-2") + store.begin_operation( + "rollback-1", request_hash="rollback-request", expected_revision=2, + actor="admin", lease_owner="worker", + ) + result = store.rollback( + target_revision=1, + expected_revision=2, + operation_id="rollback-1", + actor="admin", + request_hash="rollback-request", + ) + + assert result.config_revision == 3 + assert store.current_revision() == 3 + assert store.get_admin_config()[1]["providers"][0]["display_name"] == "OpenAI Production" + assert store.load().config_revision == 3 + + +def test_identity_key_rotation_does_not_rotate_unchanged_credential(tmp_path): + first_ring = _ring() + path = tmp_path / "model-config.sqlite" + store = UnifiedModelConfigStore( + path, + identity_key_ring=first_ring, + encryption_keys={"enc-v1": "e" * 32}, + current_encryption_key_id="enc-v1", + ) + payload = normalize_v4_config(_payload()) + provider_id = payload["providers"][0]["provider_id"] + profile_id = payload["default_model_profile_id"] + plan = store.plan_credentials(payload, {provider_id: "sk-same-secret"}) + projection = project_v4_to_v3( + payload, revision=1, identity_key_id=first_ring.current.key_id, + bindings=plan.bindings, + ) + evidence = build_supported_v3_evidence( + projection, + identity_key_ring=first_ring, + verified_profiles={provider_id: frozenset({profile_id})}, + ) + store.begin_operation( + "save-1", request_hash="request-1", expected_revision=0, + actor="admin", lease_owner="worker", + ) + store.commit( + payload, expected_revision=0, operation_id="save-1", actor="admin", + request_hash="request-1", credential_plan=plan, evidence=evidence, + ) + rotated_ring = HmacKeyRing( + KeyMaterial.create("identity-v2", "j" * 32), retained=(first_ring.current,) + ) + rotated_store = UnifiedModelConfigStore( + path, + identity_key_ring=rotated_ring, + encryption_keys={"enc-v1": "e" * 32}, + current_encryption_key_id="enc-v1", + ) + + rotated_plan = rotated_store.plan_credentials( + payload, {provider_id: "sk-same-secret"} + ) + + assert rotated_plan.new_versions == frozenset() + assert rotated_plan.bindings[provider_id] == 1 + assert rotated_plan.fingerprints[provider_id] == plan.fingerprints[provider_id] + + +def test_evidence_requires_observed_results_for_every_declared_capability(): + ring = _ring() + raw = _payload() + raw["providers"][0]["models"][0]["capabilities"] = { + "text": True, + "vision": True, + "tools": True, + } + payload = normalize_v4_config(raw) + provider_id = payload["providers"][0]["provider_id"] + profile_id = payload["default_model_profile_id"] + projection = project_v4_to_v3( + payload, + revision=1, + identity_key_id=ring.current.key_id, + bindings={provider_id: 1}, + ) + + with pytest.raises(EvoRuntimeError) as raised: + build_supported_v3_evidence( + projection, + identity_key_ring=ring, + probe_results={ + provider_id: {profile_id: {"connectivity": "supported"}} + }, + ) + + assert raised.value.code == "CAPABILITY_EVIDENCE_STALE" + + +def test_historical_evidence_is_not_reused_for_route_admission(tmp_path): + ring = _ring() + store = UnifiedModelConfigStore( + tmp_path / "model-config.sqlite", + identity_key_ring=ring, + encryption_keys={"enc-v1": "e" * 32}, + current_encryption_key_id="enc-v1", + ) + payload = normalize_v4_config(_payload()) + provider_id = payload["providers"][0]["provider_id"] + profile_id = payload["default_model_profile_id"] + plan = store.plan_credentials(payload, {provider_id: "sk-secret-value"}) + projection = project_v4_to_v3( + payload, + revision=1, + identity_key_id=ring.current.key_id, + bindings=plan.bindings, + ) + evidence = build_supported_v3_evidence( + projection, + identity_key_ring=ring, + verified_profiles={provider_id: frozenset({profile_id})}, + ) + store.begin_operation( + "save-1", request_hash="request-1", expected_revision=0, + actor="admin", lease_owner="worker", + ) + store.commit( + payload, expected_revision=0, operation_id="save-1", actor="admin", + request_hash="request-1", credential_plan=plan, evidence=evidence, + ) + + next_payload = json.loads(json.dumps(payload)) + next_payload["providers"][0]["display_name"] = "Renamed Provider" + next_plan = store.plan_credentials(next_payload, {}) + next_projection = project_v4_to_v3( + next_payload, + revision=2, + identity_key_id=ring.current.key_id, + bindings=next_plan.bindings, + ) + observed, windows, missing = store.reusable_probe_results( + next_projection, + fingerprints=next_plan.fingerprints, + versions=next_plan.bindings, + ) + + assert observed == {} + assert windows == {} + assert missing == frozenset({(provider_id, profile_id)}) + + +def test_availability_reflects_enabled_state(tmp_path): + ring = _ring() + store = UnifiedModelConfigStore( + tmp_path / "model-config.sqlite", + identity_key_ring=ring, + encryption_keys={"enc-v1": "e" * 32}, + current_encryption_key_id="enc-v1", + ) + raw = _payload() + disabled_profile_id = str(uuid.uuid4()) + disabled = json.loads(json.dumps(raw["providers"][0]["models"][0])) + disabled.update( + { + "model_profile_id": disabled_profile_id, + "provider_model_id": "disabled-model", + "display_name": "Disabled", + "enabled": False, + } + ) + raw["providers"][0]["models"].append(disabled) + payload = normalize_v4_config(raw) + provider_id = payload["providers"][0]["provider_id"] + profile_id = payload["default_model_profile_id"] + plan = store.plan_credentials(payload, {provider_id: "sk-secret-value"}) + store.begin_operation( + "save-1", request_hash="request-1", expected_revision=0, + actor="admin", lease_owner="worker", + ) + store.commit( + payload, expected_revision=0, operation_id="save-1", actor="admin", + request_hash="request-1", credential_plan=plan, evidence=[], + ) + + availability = store.get_model_availability() + + assert availability[profile_id] == { + "model_profile_id": profile_id, + "enabled": True, + "selectable": True, + } + assert availability[disabled_profile_id] == { + "model_profile_id": disabled_profile_id, + "enabled": False, + "selectable": False, + } + + +def test_schema_upgrade_adds_availability_and_operation_columns(tmp_path): + path = tmp_path / "model-config.sqlite" + UnifiedModelConfigStore( + path, + identity_key_ring=_ring(), + encryption_keys={"enc-v1": "e" * 32}, + current_encryption_key_id="enc-v1", + ) + with sqlite3.connect(path) as connection: + connection.execute("DROP INDEX idx_capability_evidence_semantics") + connection.execute("ALTER TABLE active_model_config DROP COLUMN evidence_epoch") + connection.execute("ALTER TABLE capability_evidence DROP COLUMN credential_version") + connection.execute("ALTER TABLE model_config_operations DROP COLUMN stage") + connection.execute("ALTER TABLE model_config_operations DROP COLUMN progress_json") + connection.execute("ALTER TABLE model_config_operations DROP COLUMN error_details_json") + + UnifiedModelConfigStore( + path, + identity_key_ring=_ring(), + encryption_keys={"enc-v1": "e" * 32}, + current_encryption_key_id="enc-v1", + ) + + with sqlite3.connect(path) as connection: + assert "evidence_epoch" in { + row[1] for row in connection.execute("PRAGMA table_info(active_model_config)") + } + assert "credential_version" in { + row[1] for row in connection.execute("PRAGMA table_info(capability_evidence)") + } + operation_columns = { + row[1] for row in connection.execute("PRAGMA table_info(model_config_operations)") + } + assert {"stage", "progress_json", "error_details_json"} <= operation_columns + + +def test_v3_runtime_accepts_expired_historical_evidence(tmp_path): + ring = _ring() + store = UnifiedModelConfigStore( + tmp_path / "model-config.sqlite", + identity_key_ring=ring, + encryption_keys={"enc-v1": "e" * 32}, + current_encryption_key_id="enc-v1", + ) + payload = normalize_v4_config(_payload()) + provider_id = payload["providers"][0]["provider_id"] + profile_id = payload["default_model_profile_id"] + plan = store.plan_credentials(payload, {provider_id: "sk-secret-value"}) + projection = project_v4_to_v3( + payload, revision=1, identity_key_id=ring.current.key_id, + bindings=plan.bindings, + ) + now = datetime.now(UTC) + evidence = build_supported_v3_evidence( + projection, + identity_key_ring=ring, + verified_profiles={provider_id: frozenset({profile_id})}, + evidence_windows={ + provider_id: { + profile_id: ( + (now - timedelta(hours=2)).isoformat(timespec="microseconds"), + (now - timedelta(hours=1)).isoformat(timespec="microseconds"), + ) + } + }, + ) + store.begin_operation( + "save-1", request_hash="request-1", expected_revision=0, + actor="admin", lease_owner="worker", + ) + store.commit( + payload, expected_revision=0, operation_id="save-1", actor="admin", + request_hash="request-1", credential_plan=plan, evidence=evidence, + ) + config = store.load() + authority = HmacGrantAuthority("r" * 32) + runtime = EvoModelRuntime( + store, + admission_verifier=authority, + quote_authority=authority, + identity_key_ring=ring, + secret_resolver=store.resolve, + ) + selector = config.resolve_main_selector(None) + route = config.concrete_routes(selector.selector_id)[0] + + resolved = runtime._resolve_route(config, route, "main_agent") + + assert resolved.identity.model_id == "gpt-test" + + +def test_lenient_normalization_ignores_unknown_fields_with_warning(): + raw = _payload() + raw["unexpected_top"] = True + raw["providers"][0]["unexpected_provider"] = "x" + raw["providers"][0]["models"][0]["unexpected_model"] = 1 + warnings: list[dict] = [] + + payload = normalize_v4_config(raw, lenient=True, warnings=warnings) + + assert payload["schema_version"] == 4 + assert [item["code"] for item in warnings] == ["CONFIG_UNKNOWN_FIELDS"] * 3 + + +def test_lenient_normalization_clamps_out_of_range_and_timeout_conflict(): + raw = _payload() + raw["providers"][0]["connection"] = { + "base_url": "https://api.openai.com/v1", + "connect_timeout_seconds": 120, + "attempt_timeout_seconds": 30, + } + warnings: list[dict] = [] + + payload = normalize_v4_config(raw, lenient=True, warnings=warnings) + connection = payload["providers"][0]["connection"] + + assert connection["connect_timeout_seconds"] == 60 + assert connection["attempt_timeout_seconds"] == 60 + codes = {item["code"] for item in warnings} + assert "CONFIG_VALUE_OUT_OF_RANGE" in codes + assert "CONFIG_VALUE_CONFLICT" in codes + + +def test_lenient_normalization_forces_text_capability_and_keeps_unknown_purpose(): + raw = _payload() + raw["providers"][0]["models"][0]["capabilities"] = {"text": False, "vision": True} + raw["providers"][0]["models"][0]["parameters"] = { + "purpose_overrides": {"custom_purpose": {"temperature": 0.5}} + } + warnings: list[dict] = [] + + payload = normalize_v4_config(raw, lenient=True, warnings=warnings) + model = payload["providers"][0]["models"][0] + + assert model["capabilities"]["text"] is True + assert model["parameters"]["purpose_overrides"]["custom_purpose"] == {"temperature": 0.5} + assert any(item["code"] == "CONFIG_REQUIRED" for item in warnings) + assert any(item["code"] == "CONFIG_REFERENCE_INVALID" for item in warnings) + + +def test_strict_normalization_still_rejects_unknown_fields(): + raw = _payload() + raw["unexpected_top"] = True + + with pytest.raises(EvoRuntimeError): + normalize_v4_config(raw) + + +def test_unverified_enabled_model_is_selectable(tmp_path): + ring = _ring() + store = UnifiedModelConfigStore( + tmp_path / "model-config.sqlite", + identity_key_ring=ring, + encryption_keys={"enc-v1": "e" * 32}, + current_encryption_key_id="enc-v1", + ) + payload = normalize_v4_config(_payload()) + provider_id = payload["providers"][0]["provider_id"] + profile_id = payload["default_model_profile_id"] + plan = store.plan_credentials(payload, {provider_id: "sk-secret-value"}) + store.begin_operation( + "save-1", request_hash="request-1", expected_revision=0, + actor="admin", lease_owner="worker", + ) + store.commit( + payload, expected_revision=0, operation_id="save-1", actor="admin", + request_hash="request-1", credential_plan=plan, evidence=[], + ) + + availability = store.get_model_availability()[profile_id] + + assert availability == { + "model_profile_id": profile_id, + "enabled": True, + "selectable": True, + } diff --git a/tests/test_model_fallback.py b/tests/test_model_fallback.py index 36ea9d6..6e65af2 100644 --- a/tests/test_model_fallback.py +++ b/tests/test_model_fallback.py @@ -161,6 +161,26 @@ class TestTryFallbacks: invoke.assert_awaited_once() mock_gcm.assert_called_once_with(model="fb-model", provider="fb-provider") + async def test_provider_error_details_are_not_emitted(self): + """Fallback diagnostics must not expose provider bodies or credentials.""" + add_fallback("fb-model", "fb-provider") + req = _fake_request() + invoke = AsyncMock(return_value=AI_RESPONSE) + emitted: list[str] = [] + set_ui_emit(lambda message, _style: emitted.append(message)) + + with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: + mock_gcm.return_value = MagicMock() + await _try_fallbacks( + req, + invoke, + Exception("provider body contains sk-live-do-not-log"), + ) + + output = "\n".join(emitted) + assert "sk-live-do-not-log" not in output + assert "Exception" in output + async def test_skips_failing_fallback_tries_next(self): """When the first fallback fails, try the second.""" add_fallback("fb-bad", "prov-a") @@ -270,7 +290,8 @@ class TestTryFallbacks: # Attribution flipped to moonshot (the failing fallback), not # openai (the original request's model). assert exc_info.value.provider == "moonshot" - assert "quota exceeded" in exc_info.value.message + assert exc_info.value.message == "Provider request failed." + assert "quota exceeded" not in exc_info.value.message async def test_langgraph_error_at_fallback_raise_point_passes_through(self): """Regression: ``_raise_normalized`` calls ``_normalize`` diff --git a/tests/test_model_secret_store.py b/tests/test_model_secret_store.py new file mode 100644 index 0000000..cf20970 --- /dev/null +++ b/tests/test_model_secret_store.py @@ -0,0 +1,48 @@ +from __future__ import annotations + +import pytest + +from EvoScientist.llm.contracts import EvoRuntimeError +from EvoScientist.llm.model_config import SecretReference +from EvoScientist.llm.secret_store import EncryptedModelSecretStore + + +def test_secret_store_versions_masks_and_resolves(tmp_path): + store = EncryptedModelSecretStore( + tmp_path / "secrets.sqlite", + master_secret="test-master-secret-that-is-at-least-32-bytes", + ) + + first = store.put("dashscope/primary", "sk-first-secret-value", created_by="admin") + second = store.put( + "dashscope/primary", "sk-second-secret-value", created_by="admin" + ) + + assert first.version == 1 + assert second.version == 2 + assert "second-secret" not in second.masked_value + assert second.ref == "secret://dashscope/primary#2" + metadata = store.list_metadata() + assert [item.version for item in metadata] == [2, 1] + resolved = store.resolve(SecretReference(second.ref, second.version)) + assert resolved.value == "sk-second-secret-value" + assert resolved.authoritative_version == "2" + + +def test_secret_store_rejects_wrong_revision_and_master_key(tmp_path): + path = tmp_path / "secrets.sqlite" + store = EncryptedModelSecretStore( + path, + master_secret="test-master-secret-that-is-at-least-32-bytes", + ) + item = store.put("openai/primary", "sk-secret", created_by="admin") + + with pytest.raises(EvoRuntimeError, match="ROUTE_SECRET_UNAVAILABLE"): + store.resolve(SecretReference(item.ref, item.version + 1)) + + wrong_key = EncryptedModelSecretStore( + path, + master_secret="different-master-secret-that-is-at-least-32-bytes", + ) + with pytest.raises(EvoRuntimeError, match="ROUTE_SECRET_UNAVAILABLE"): + wrong_key.resolve(SecretReference(item.ref, item.version)) diff --git a/tests/test_observation_memory.py b/tests/test_observation_memory.py index 3db0d79..2c07517 100644 --- a/tests/test_observation_memory.py +++ b/tests/test_observation_memory.py @@ -38,6 +38,7 @@ from EvoScientist.memory.observations import ( MemorySourceType, MemoryType, ObservationSearchMode, + archive_observation_file, create_link_observations_tool, create_read_memory_tool, create_search_observations_tool, @@ -145,6 +146,7 @@ def _memory_worker_run( source_agent: str = "EvoScientist", source_session_id: str = "thread-1", trajectory_digest: str = "digest-1", + configurable: dict[str, object] | None = None, ) -> background_runs.BackgroundRun: return background_runs.BackgroundRun( name="EvoMemory worker", @@ -160,6 +162,7 @@ def _memory_worker_run( "source_session_id": source_session_id, "trajectory_digest": trajectory_digest, }, + configurable=configurable, ) @@ -345,6 +348,26 @@ def test_record_observation_file_writes_contract_and_dedupes(tmp_path): } +def test_record_observation_file_serializes_concurrent_deduplication(tmp_path): + memories = tmp_path / "memories" + barrier = threading.Barrier(2) + results: list[dict[str, Any]] = [] + + def record() -> None: + barrier.wait() + results.append(_record_test_observation(memories)) + + threads = [threading.Thread(target=record) for _ in range(2)] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + assert sorted(result["created"] for result in results) == [False, True] + observation_path = memories / _memory_relative_path(results[0]) + assert read_observation_document(observation_path) is not None + + def test_link_observation_files_writes_frontmatter_and_dedupes(tmp_path): memories = tmp_path / "memories" first = record_observation_file( @@ -384,6 +407,7 @@ def test_link_observation_files_writes_frontmatter_and_dedupes(tmp_path): relation=ObservationRelation.COMPLEMENTS, reason="Both observations describe the durable background-memory flow.", ) + duplicate = link_observation_files( memory_dir=memories, project_id="P-project", @@ -432,6 +456,59 @@ def test_link_observation_files_writes_frontmatter_and_dedupes(tmp_path): datetime.strptime(first_links[0]["linked_at"], "%Y-%m-%dT%H:%M:%SZ") +def test_archive_observation_removes_live_file_and_related_links(tmp_path): + memories = tmp_path / "memories" + source = _record_test_observation( + memories, + summary="Global source memory.", + observation="A global memory linked to one project observation.", + ) + target = _record_test_observation( + memories, + summary="Project target memory.", + observation="A project memory that can be archived safely.", + scope=MemoryScope.PROJECT, + ) + link_observation_files( + memory_dir=memories, + project_id="P-project", + source_observation_id=source["observation_id"], + target_observation_id=target["observation_id"], + reason="Archive cleanup regression test.", + ) + + result = archive_observation_file( + memory_dir=memories, + observation_id=target["observation_id"], + observation_path=_memory_relative_path(target), + ) + + target_path = memories / _memory_relative_path(target) + source_document = read_observation_document( + memories / _memory_relative_path(source) + ) + assert result["removed"] is True + assert not target_path.exists() + assert (memories / str(result["archive_path"])).is_file() + assert source_document is not None + assert all( + relation.id != target["observation_id"] + for relation in source_document[0].related_observations + ) + + +def test_archive_observation_rejects_path_mismatch(tmp_path): + memories = tmp_path / "memories" + observation = _record_test_observation(memories) + + with pytest.raises(ValueError, match="invalid observation archive target"): + archive_observation_file( + memory_dir=memories, + observation_id=observation["observation_id"], + observation_path="../observations/global/other.md", + ) + + def test_link_observation_files_serializes_concurrent_frontmatter_updates(tmp_path): memories = tmp_path / "memories" source = record_observation_file( @@ -1333,6 +1410,43 @@ def test_search_observation_files_returns_no_low_confidence_fallback(tmp_path): assert hits == [] +def test_search_observation_files_ranks_chinese_queries(tmp_path): + memories = tmp_path / "memories" + relevant = record_observation_file( + memory_dir=memories, + project_id="P-project", + memory_type=MemoryType.PROCEDURAL, + summary="切换对话时保留用户级长期记忆", + observation="用户画像跨对话共享,项目记忆跟随当前工作区。", + why_it_matters="不同对话需要共享稳定偏好,但不能混合项目约束。", + scope=MemoryScope.GLOBAL, + source_type=MemorySourceType.TURN, + source_session_id="thread-1", + source_agent="EvoScientist", + ) + record_observation_file( + memory_dir=memories, + project_id="P-project", + memory_type=MemoryType.SEMANTIC, + summary="模型计费快照", + observation="模型调用使用不可变计费快照。", + why_it_matters="结算需要稳定证据。", + scope=MemoryScope.GLOBAL, + source_type=MemorySourceType.TURN, + source_session_id="thread-1", + source_agent="EvoScientist", + ) + + hits = search_observation_files( + memory_dir=memories, + project_id="P-project", + query="跨对话共享记忆", + ) + + assert hits + assert hits[0]["observation_id"] == relevant["observation_id"] + + def test_record_observation_tool_can_use_worker_config_source(tmp_path): from EvoScientist.middleware.memory import create_memory_middleware @@ -1825,6 +1939,54 @@ def test_memory_worker_run_payload_use_server_thread_id_and_source_metadata( } +def test_memory_worker_inherits_signed_runtime_and_scopes_metering(monkeypatch): + metering = { + "gateway_url": "http://gateway", + "run_id": "run-parent", + "envelope_signature": "signed-parent", + "provider_id": "provider-1", + "model_id": "model-1", + } + proxy = { + "gateway_url": "http://gateway", + "run_id": "run-parent", + "envelope_signature": "signed-parent", + } + monkeypatch.setattr( + "langgraph.config.get_config", + lambda: { + "configurable": { + "model": "model-1", + "model_provider": "provider-1", + "ai4sci_metering": metering, + "ai4sci_model_proxy": proxy, + "untrusted_extra": "do-not-copy", + }, + "metadata": {"langgraph_api_url": "http://127.0.0.1:6176"}, + }, + ) + context = _memory_source_context( + memory_dir="/memories", + workspace_dir="/workspace", + source_type=MemorySourceType.TURN, + trajectory=[{"role": "human", "content": "hi"}], + ) + + request = memory_launch.memory_worker_launch_request(context) + payload = request.run_payload("worker-thread") + configurable = payload["config"]["configurable"] + + assert request.url == "http://127.0.0.1:6176" + assert configurable["ai4sci_model_proxy"] == proxy + assert "ai4sci_tool_effect" not in configurable + assert configurable["ai4sci_metering"] == { + **metering, + "source_type": "evomemory_turn_worker", + } + assert "untrusted_extra" not in configurable + assert "source_type" not in metering + + def test_memory_worker_finish_launches_linker_for_new_observations( tmp_path, ): @@ -2154,6 +2316,66 @@ def test_observation_linker_launch_request_encodes_batch_context(tmp_path): ] +def test_observation_linker_inherits_worker_runtime_scope(tmp_path): + context = memory_scheduler.ObservationLinkerContext( + memory_dir=tmp_path / "memories", + workspace_dir=tmp_path / "workspace", + project_id="P-project", + observation_ids=("O-1",), + runtime_url="http://127.0.0.1:6176", + runtime_configurable={ + "ai4sci_metering": { + "gateway_url": "http://gateway", + "run_id": "run-parent", + "envelope_signature": "signed-parent", + "source_type": "evomemory_turn_worker", + }, + "ai4sci_model_proxy": { + "gateway_url": "http://gateway", + "run_id": "run-parent", + "envelope_signature": "signed-parent", + }, + }, + ) + + request = memory_launch.observation_linker_launch_request(context) + configurable = request.run_payload("linker-thread")["config"]["configurable"] + + assert request.url == "http://127.0.0.1:6176" + assert configurable["ai4sci_metering"]["source_type"] == "evomemory_linker" + assert configurable["ai4sci_model_proxy"]["run_id"] == "run-parent" + + +def test_memory_scheduler_propagates_signed_runtime_to_linker(tmp_path): + launched: list[memory_scheduler.ObservationLinkerContext] = [] + scheduler = memory_scheduler.MemoryScheduler(launch_linker=launched.append) + memory_dir = tmp_path / "memories" + observation = _record_test_observation(memory_dir) + delta = worker_activity.MemoryOutputDelta( + memory_dir=memory_dir, + observation_paths=(_memory_relative_path(observation),), + ) + configurable = { + "ai4sci_metering": { + "gateway_url": "http://gateway", + "run_id": "run-parent", + "envelope_signature": "signed-parent", + } + } + + scheduler.record_worker_finished( + _memory_worker_run( + workspace_dir=str(tmp_path / "workspace"), + configurable=configurable, + ), + delta, + ) + + assert len(launched) == 1 + assert launched[0].runtime_url == "http://x" + assert launched[0].runtime_configurable == configurable + + def test_observation_linker_does_not_launch_when_observations_disabled( tmp_path, monkeypatch, @@ -2298,7 +2520,12 @@ def test_memory_worker_observation_writer_modes( # exceptions from the auxiliary model call get normalized before # the tool-error handler sees them. assert type(middleware[0]).__name__ == "ErrorNormalizationMiddleware" - assert type(middleware[1]).__name__ == "ToolErrorHandlerMiddleware" + assert [type(item).__name__ for item in middleware[:4]] == [ + "ErrorNormalizationMiddleware", + "ConfigurableModelMiddleware", + "RecoverableMeteringMiddleware", + "ToolErrorHandlerMiddleware", + ] assert _memory_tool_names(middleware) == expected_tools diff --git a/tests/test_provider_context_middleware.py b/tests/test_provider_context_middleware.py new file mode 100644 index 0000000..11c82ab --- /dev/null +++ b/tests/test_provider_context_middleware.py @@ -0,0 +1,176 @@ +from __future__ import annotations + +import base64 +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +from langchain.agents import create_agent +from langchain.agents.middleware.types import ModelRequest, ModelResponse +from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel +from langchain_core.messages import AIMessage, HumanMessage +from langgraph.types import Overwrite + +from EvoScientist.llm.contracts import EvoRuntimeError +from EvoScientist.middleware.provider_context import ProviderContextMediaMiddleware + + +class _Backend: + def __init__(self, *, error: str | None = None) -> None: + self.error = error + self.uploads: list[tuple[str, bytes]] = [] + + def upload_files(self, files): + self.uploads.extend(files) + return [SimpleNamespace(path=path, error=self.error) for path, _data in files] + + async def aupload_files(self, files): + return self.upload_files(files) + + +class _FakeModel(FakeMessagesListChatModel): + def bind_tools(self, _tools, *, tool_choice=None, **_kwargs): + return self + + +def _request(messages): + return ModelRequest( + messages=list(messages), + model=MagicMock(), + state={}, + runtime=MagicMock(), + system_message=MagicMock(), + ) + + +def _image_block(raw: bytes) -> dict: + return { + "type": "image", + "base64": base64.b64encode(raw).decode("ascii"), + "mime_type": "image/png", + } + + +def test_externalizes_historical_and_new_assistant_media() -> None: + backend = _Backend() + middleware = ProviderContextMediaMiddleware(backend) + old_raw = b"old-png" + new_raw = b"new-png" + historical = AIMessage(content=[_image_block(old_raw)]) + captured = {} + + def handler(request): + captured["messages"] = request.messages + return ModelResponse( + result=[ + AIMessage( + content=[ + {"type": "text", "text": "done"}, + _image_block(new_raw), + ] + ) + ] + ) + + response = middleware.wrap_model_call(_request([historical]), handler) + + provider_content = captured["messages"][0].content + assert all("base64" not in block for block in provider_content) + assert "generated_image" in provider_content[0]["text"] + stored_content = response.result[0].content + assert stored_content[1]["type"] == "image" + assert stored_content[1]["url"].startswith("/artifacts/model-output/") + assert "base64" not in stored_content[1] + assert {data for _path, data in backend.uploads} == {old_raw, new_raw} + + +def test_before_model_durably_replaces_historical_inline_media() -> None: + middleware = ProviderContextMediaMiddleware(_Backend()) + + update = middleware.before_model( + {"messages": [AIMessage(content=[_image_block(b"old-png")])]}, None + ) + + assert update is not None + assert isinstance(update["messages"], Overwrite) + content = update["messages"].value[0].content + assert content[0]["url"].startswith("/artifacts/model-output/") + assert "base64" not in content[0] + + +@pytest.mark.asyncio +async def test_real_langgraph_injects_runtime_and_repairs_checkpoint_media() -> None: + backend = _Backend() + raw = b"x" * 1_349_952 + agent = create_agent( + model=_FakeModel(responses=[AIMessage(content="done")]), + tools=[], + middleware=[ProviderContextMediaMiddleware(backend)], + ) + + result = await agent.ainvoke( + { + "messages": [ + AIMessage(content=[_image_block(raw)]), + HumanMessage(content="continue"), + ] + } + ) + + repaired = result["messages"][0].content[0] + assert repaired["url"].startswith("/artifacts/model-output/") + assert "base64" not in repaired + assert any(data == raw for _path, data in backend.uploads) + + +@pytest.mark.asyncio +async def test_before_model_normalizes_internal_middleware_failure() -> None: + class _BrokenBackend(_Backend): + async def aupload_files(self, _files): + raise TypeError("sensitive internal detail") + + middleware = ProviderContextMediaMiddleware(_BrokenBackend()) + + with pytest.raises(EvoRuntimeError, match="AGENT_MIDDLEWARE_FAILED") as exc: + await middleware.abefore_model( + {"messages": [AIMessage(content=[_image_block(b"png")])]}, None + ) + + assert exc.value.details == ( + { + "failure_stage": "agent_middleware", + "middleware": "provider_context_media", + "middleware_node": "provider_context_media.before_model", + "agent_error_type": "TypeError", + "agent_error_module": "builtins", + }, + ) + assert "sensitive internal detail" not in str(exc.value.details) + + +@pytest.mark.asyncio +async def test_async_externalization_is_content_addressed() -> None: + backend = _Backend() + middleware = ProviderContextMediaMiddleware(backend) + raw = b"same-png" + + async def handler(_request): + return ModelResponse(result=[AIMessage(content=[_image_block(raw)])]) + + first = await middleware.awrap_model_call(_request([]), handler) + second = await middleware.awrap_model_call(_request(first.result), handler) + + assert first.result[0].content[0]["url"] == second.result[0].content[0]["url"] + assert all(len(data) < 100 for _path, data in backend.uploads) + + +def test_upload_failure_is_terminal_and_does_not_return_inline_media() -> None: + middleware = ProviderContextMediaMiddleware(_Backend(error="disk full")) + + with pytest.raises(EvoRuntimeError, match="MEDIA_PERSIST_FAILED"): + middleware.wrap_model_call( + _request([]), + lambda _request: ModelResponse( + result=[AIMessage(content=[_image_block(b"png")])] + ), + ) diff --git a/tests/test_provider_model_config_v3.py b/tests/test_provider_model_config_v3.py new file mode 100644 index 0000000..0d07d6d --- /dev/null +++ b/tests/test_provider_model_config_v3.py @@ -0,0 +1,1037 @@ +from __future__ import annotations + +from pathlib import Path + +import httpx +import pytest + +import EvoScientist.llm.adapter_registry as adapter_registry +from EvoScientist.llm.adapter_registry import NormalizedUsage, get_adapter_registry +from EvoScientist.llm.contracts import EvoRuntimeError +from EvoScientist.llm.model_config import ( + EvoModelConfig, + convert_v2_to_v3_draft, + invocation_fingerprint, + route_semantics_hash, +) +from EvoScientist.llm.secret_store import EncryptedModelSecretStore + + +def _model(model_key: str, model_id: str, *, api_mode: str) -> dict: + return { + "model_key": model_key, + "provider_model_id": model_id, + "version_policy": "rolling", + "resolved_model_revision": None, + "display_name": model_key, + "description": "fixture", + "enabled": True, + "tags": ["fixture"], + "invocation": {"api_mode": api_mode, "tool_call_transport": "native"}, + "capabilities": {"text": True}, + "limits": {"context_tokens": 128000, "max_output_tokens": 8192}, + "parameters": { + "defaults": {}, + "purpose_overrides": {}, + "user_options": {}, + "constraints": [], + }, + "access": {"visibility": "authenticated", "roles": []}, + "billing": { + "sku": f"internal/{model_key}", + "pricing_revision": "internal-unmetered-v1", + "currency": "CNY", + "unit_scale": 1000000, + "input_microunits_per_million": 0, + "output_microunits_per_million": 0, + "cached_microunits_per_million": 0, + "multiplier": 1.25, + }, + } + + +def v3_payload() -> dict: + specs = ( + ( + "anthropic-prod", + "anthropic", + "anthropic-v1", + "anthropic_native", + "messages", + "claude-fixture", + ), + ( + "openai-prod", + "openai", + "openai-v1", + "openai_native", + "responses", + "gpt-fixture", + ), + ( + "gemini-prod", + "google-gemini", + "google-gemini-v1", + "gemini_native", + "interactions", + "gemini-fixture", + ), + ("xai-prod", "xai", "xai-v1", "openai_compatible", "responses", "grok-fixture"), + ) + providers = [] + aliases = [] + for provider_id, adapter_id, revision, wire, mode, model_id in specs: + providers.append( + { + "provider_id": provider_id, + "display_name": provider_id, + "adapter_id": adapter_id, + "adapter_revision": revision, + "wire_protocol": wire, + "enabled": True, + "connection": { + "base_url": get_adapter_registry() + .get(adapter_id, revision) + .recommended_base_url, + "credential_ref": f"secret://model-providers/{provider_id}#1", + }, + "defaults": {}, + "models": [ + _model("general", model_id, api_mode=mode), + _model("fast", model_id + "-fast", api_mode=mode), + ], + } + ) + for model_key in ("general", "fast"): + aliases.append( + { + "alias": f"{provider_id}-{model_key}", + "display_name": f"{provider_id} {model_key}", + "provider_ref": provider_id, + "model_ref": model_key, + "enabled": True, + "access": {"visibility": "authenticated", "roles": []}, + "defaults": {}, + } + ) + return { + "schema_version": 3, + "config_revision": 1, + "config_identity_key_id": "identity-v1", + "runtime_defaults": {}, + "providers": providers, + "aliases": aliases, + "purpose_defaults": { + name: {} + for name in ( + "main_agent", + "tool_selector", + "deepagents_summarizer", + "title", + ) + }, + "purpose_routes": { + "main_agent": {"default_alias": "openai-prod-general"}, + "tool_selector": "inherit_main", + "deepagents_summarizer": "inherit_main", + "title": {"default_alias": "openai-prod-fast"}, + }, + "purpose_call_limits": { + "main_agent": {"max_output_tokens": 8192, "max_attempts_per_run": 2}, + "tool_selector": {"max_output_tokens": 4096, "max_attempts_per_run": 2}, + "deepagents_summarizer": { + "max_output_tokens": 4096, + "max_attempts_per_run": 2, + }, + "title": {"max_output_tokens": 256, "max_attempts_per_run": 1}, + }, + "health_policy": {"provider_connection": {}, "model_route": {}}, + "web_runtime": {}, + "capability_evidence": [], + } + + +def test_v3_supports_four_native_adapters_and_multiple_models() -> None: + config = EvoModelConfig.parse(v3_payload(), require_evidence=False) + assert config.schema_version == 3 + assert len(config.providers) == 4 + assert all(len(provider.models) == 2 for provider in config.providers.values()) + assert config.endpoint_pools == {} + assert config.tool_protocol_fallbacks == {} + assert all( + model.quote.multiplier == "1.25" + for provider in config.providers.values() + for model in provider.models.values() + ) + + +@pytest.mark.parametrize( + "base_url", + [ + "", + "not a url", + "ftp://user:pass@example.test/path?key=value#fragment", + "https://169.254.169.254/latest/meta-data", + "custom://model.internal:70000/path", + ], +) +def test_v3_provider_base_url_is_not_validated(base_url) -> None: + payload = v3_payload() + payload["providers"][0]["connection"]["base_url"] = base_url + + config = EvoModelConfig.parse(payload, require_evidence=False) + + endpoints = config.providers["anthropic-prod"].endpoints + assert len(endpoints) == 1 + assert next(iter(endpoints.values())).base_url == base_url + + +@pytest.mark.asyncio +async def test_model_discovery_does_not_prevalidate_base_url(monkeypatch) -> None: + requested = {} + + class Response: + def raise_for_status(self): + return None + + def json(self): + return {"data": [{"id": "model-a"}]} + + class Client: + def __init__(self, **_kwargs): + pass + + async def __aenter__(self): + return self + + async def __aexit__(self, *_args): + return None + + async def get(self, url, *, headers): + requested["url"] = url + requested["headers"] = headers + return Response() + + monkeypatch.setattr(httpx, "AsyncClient", Client) + registration = get_adapter_registry().get("openai", "openai-v1") + + result = await registration.discover_models( + base_url="custom://model.internal:70000/path?key=value#fragment", + api_key="sk-test", + ) + + assert requested["url"] == ( + "custom://model.internal:70000/path?key=value#fragment/models" + ) + assert result == ({"provider_model_id": "model-a", "display_name": "model-a"},) + + +def test_v3_provider_can_use_current_credential_without_secret_ref( + tmp_path: Path, +) -> None: + payload = v3_payload() + for provider in payload["providers"]: + provider["connection"].pop("credential_ref") + config = EvoModelConfig.parse(payload, require_evidence=False) + assert config.providers["openai-prod"].endpoints["openai-prod"].auth.ref == ( + "provider://openai-prod" + ) + + store = EncryptedModelSecretStore( + tmp_path / "model_secrets.sqlite", master_secret="x" * 32 + ) + store.put( + "model-providers/openai-prod", + "sk-current", + created_by="admin", + status="active", + ) + resolved = store.resolve( + config.providers["openai-prod"].endpoints["openai-prod"].auth + ) + assert resolved.value == "sk-current" + assert resolved.authoritative_version == "1" + + +def test_v3_rejects_removed_endpoint_field() -> None: + payload = v3_payload() + payload["providers"][0]["endpoints"] = [] + with pytest.raises(EvoRuntimeError, match="unknown fields"): + EvoModelConfig.parse(payload, require_evidence=False) + + +def test_qwen_37_descriptor_has_correct_bounds_and_parameter_range() -> None: + registration = get_adapter_registry().get("dashscope", "dashscope-v1") + descriptor = registration.resolve_model_descriptor( + "qwen3.7-plus", + "chat_completions", + context_tokens=1_000_000, + max_output_tokens=65_536, + declared_capabilities={ + "text": True, + "tools": True, + "thinking": True, + "structured_output": True, + "vision": True, + }, + ) + assert descriptor.context_tokens == 1_000_000 + assert descriptor.max_output_tokens == 65_536 + assert descriptor.capabilities == frozenset( + {"text", "vision", "tools", "thinking", "structured_output"} + ) + # Adapter catalogs supply defaults only. Product policy may enable a + # capability before the static catalog has been updated. + registration.resolve_model_descriptor( + "qwen3.7-plus", + "chat_completions", + context_tokens=1_000_000, + max_output_tokens=65_536, + declared_capabilities={"text": True, "documents": True}, + ) + with pytest.raises(EvoRuntimeError) as exc: + registration.validate_parameters({"temperature": 2}, path="parameters") + assert exc.value.code == "MODEL_PARAMETER_INVALID" + + +def test_dashscope_chat_compilation_sends_explicit_thinking_and_output_bounds() -> None: + registration = get_adapter_registry().get("dashscope", "dashscope-v1") + + disabled = registration.compile_runtime_parameters( + "chat_completions", {"reasoning": "off"}, 65_536 + ) + enabled = registration.compile_runtime_parameters( + "chat_completions", + {"reasoning": "high", "reasoning_budget_tokens": 32_768}, + 65_536, + ) + + assert disabled["max_completion_tokens"] == 65_536 + assert disabled["extra_body"] == {"enable_thinking": False} + assert enabled["extra_body"] == { + "enable_thinking": True, + "thinking_budget": 32_768, + } + + +def test_dashscope_request_level_policy_disables_thinking_for_json_and_forced_tools() -> None: + registration = get_adapter_registry().get("dashscope", "dashscope-v1") + + structured = registration.compile_runtime_parameters( + "chat_completions", {"structured_output": True}, 65_536 + ) + forced_tool = registration.compile_runtime_parameters( + "chat_completions", {"tool_choice": "required"}, 65_536 + ) + + assert structured["extra_body"] == {"enable_thinking": False} + assert structured["response_format"] == {"type": "json_object"} + assert forced_tool["extra_body"] == {"enable_thinking": False} + + with pytest.raises(EvoRuntimeError) as structured_conflict: + registration.compile_runtime_parameters( + "chat_completions", + {"structured_output": True, "reasoning": "high"}, + 65_536, + ) + assert structured_conflict.value.code == "MODEL_PARAMETER_CONFLICT" + + with pytest.raises(EvoRuntimeError) as forced_tool_conflict: + registration.compile_runtime_parameters( + "chat_completions", + {"tool_choice": "required", "reasoning": "high"}, + 65_536, + ) + assert forced_tool_conflict.value.code == "MODEL_PARAMETER_CONFLICT" + + +def test_openai_gpt5_chat_uses_max_completion_tokens_for_runtime_and_connection() -> None: + registration = get_adapter_registry().get("openai", "openai-v1") + + params = registration.compile_runtime_parameters( + "chat_completions", {}, 16_384, provider_model_id="gpt-5.6" + ) + _, _, body = registration.build_probe_request( + base_url="https://api.openai.com/v1", + api_key="sk-test", + provider_model_id="gpt-5.6", + api_mode="chat_completions", + probe_kind="connectivity", + ) + + assert params["max_completion_tokens"] == 16_384 + assert "max_tokens" not in params + assert body["max_completion_tokens"] == 16 + assert "max_tokens" not in body + + +def test_kimi_k3_chat_uses_max_completion_tokens_for_runtime_and_connection() -> None: + registration = get_adapter_registry().get("openai", "openai-v1") + + params = registration.compile_runtime_parameters( + "chat_completions", {"reasoning": "high"}, 65_000, provider_model_id="k3" + ) + _, _, body = registration.build_probe_request( + base_url="https://api.kimi.com/coding/v1", + api_key="sk-test", + provider_model_id="k3", + api_mode="chat_completions", + probe_kind="reasoning", + ) + + assert params == { + "max_completion_tokens": 65_000, + "use_responses_api": False, + "reasoning_effort": "high", + } + assert body["max_completion_tokens"] == 16 + assert body["reasoning_effort"] == "low" + assert "max_tokens" not in body + + +@pytest.mark.parametrize( + ("adapter_id", "revision", "api_mode", "probe_kind", "path", "auth_header"), + [ + ("anthropic", "anthropic-v1", "messages", "tools", "/v1/messages", "x-api-key"), + ("openai", "openai-v1", "responses", "structured_output", "/responses", "Authorization"), + ("google-gemini", "google-gemini-v1", "interactions", "reasoning", "/v1beta/interactions", "x-goog-api-key"), + ("xai", "xai-v1", "chat_completions", "vision", "/chat/completions", "Authorization"), + ("dashscope", "dashscope-v1", "chat_completions", "reasoning", "/chat/completions", "Authorization"), + ], +) +def test_adapter_capability_probe_requests_are_provider_specific( + adapter_id: str, + revision: str, + api_mode: str, + probe_kind: str, + path: str, + auth_header: str, +) -> None: + registration = get_adapter_registry().get(adapter_id, revision) + + url, headers, body = registration.build_probe_request( + base_url=registration.recommended_base_url, + api_key="sk-probe", + provider_model_id="probe-model", + api_mode=api_mode, + probe_kind=probe_kind, + ) + + assert url.endswith(path) + assert auth_header in headers + assert body["model"] == "probe-model" if "model" in body else True + assert "sk-probe" not in str(body) + + +def test_adapter_rejects_capability_without_controlled_probe_fixture() -> None: + registration = get_adapter_registry().get("openai", "openai-v1") + + with pytest.raises(EvoRuntimeError) as raised: + registration.build_probe_request( + base_url=registration.recommended_base_url, + api_key="sk-probe", + provider_model_id="probe-model", + api_mode="responses", + probe_kind="video", + ) + + assert raised.value.code == "MODEL_CAPABILITY_PROBE_UNSUPPORTED" + + +def test_probe_requires_semantic_tool_call_not_only_http_success() -> None: + registration = get_adapter_registry().get("openai", "openai-v1") + payloads = registration._decode_probe_payloads( + b'{"id":"resp_1","model":"gpt-test","output":[]}' + ) + + with pytest.raises(EvoRuntimeError) as raised: + registration._validate_probe_payloads("tools", payloads) + + assert raised.value.code == "MODEL_CAPABILITY_PROBE_FAILED" + + +def test_probe_accepts_bounded_sse_reasoning_evidence() -> None: + registration = get_adapter_registry().get("dashscope", "dashscope-v1") + payloads = registration._decode_probe_payloads( + b'data: {"id":"1","model":"qwen-test","choices":[{"delta":{"reasoning_content":"x"}}]}\n\n' + b'data: [DONE]\n\n' + ) + + revision = registration._validate_probe_payloads("reasoning", payloads) + + assert revision == "qwen-test" + + +def test_probe_validates_structured_output_shape() -> None: + registration = get_adapter_registry().get("openai", "openai-v1") + payloads = registration._decode_probe_payloads( + b'{"id":"1","choices":[{"message":{"content":"{\\"ok\\":true}"}}]}' + ) + + assert registration._validate_probe_payloads("structured_output", payloads) == "" + + +@pytest.mark.asyncio +async def test_probe_retries_one_transient_provider_failure(monkeypatch) -> None: + registration = get_adapter_registry().get("openai", "openai-v1") + requests = 0 + + def respond(request: httpx.Request) -> httpx.Response: + nonlocal requests + requests += 1 + if requests == 1: + return httpx.Response(503, json={"error": {"code": "unavailable"}}) + return httpx.Response( + 200, + json={ + "model": "k3", + "choices": [{"message": {"content": "OK"}, "finish_reason": "stop"}], + }, + ) + + original_client = httpx.AsyncClient + transport = httpx.MockTransport(respond) + + def client(**kwargs): + return original_client(transport=transport, **kwargs) + + monkeypatch.setattr(httpx, "AsyncClient", client) + monkeypatch.setattr(adapter_registry, "_PROBE_RETRY_BASE_SECONDS", 0.0) + + result = await registration.probe_model( + base_url="https://provider.example/v1", + api_key="sk-probe", + provider_model_id="k3", + api_mode="chat_completions", + probe_kinds=("connectivity",), + timeout_seconds=5, + max_attempts=2, + ) + + assert result == {"connectivity": "supported", "resolved_model_revision": "k3"} + assert requests == 2 + + +@pytest.mark.asyncio +async def test_probe_reports_exhausted_attempt_details(monkeypatch) -> None: + registration = get_adapter_registry().get("openai", "openai-v1") + requests = 0 + + def respond(request: httpx.Request) -> httpx.Response: + nonlocal requests + requests += 1 + return httpx.Response(503, json={"error": {"code": "unavailable"}}) + + original_client = httpx.AsyncClient + transport = httpx.MockTransport(respond) + + def client(**kwargs): + return original_client(transport=transport, **kwargs) + + monkeypatch.setattr(httpx, "AsyncClient", client) + monkeypatch.setattr(adapter_registry, "_PROBE_RETRY_BASE_SECONDS", 0.0) + + with pytest.raises(EvoRuntimeError) as raised: + await registration.probe_model( + base_url="https://provider.example/v1", + api_key="sk-probe", + provider_model_id="k3", + api_mode="chat_completions", + probe_kinds=("reasoning",), + timeout_seconds=5, + max_attempts=2, + ) + + assert raised.value.code == "MODEL_PROVIDER_ERROR" + assert raised.value.details == ( + { + "path": "probe.reasoning", + "code": "MODEL_PROVIDER_ERROR", + "probe_kind": "reasoning", + "attempts": 2, + "retryable": True, + }, + ) + assert requests == 2 + + +def test_secret_lifecycle_revocation_is_immediate(tmp_path: Path) -> None: + store = EncryptedModelSecretStore( + tmp_path / "model_secrets.sqlite", master_secret="x" * 32 + ) + pending = store.create_pending( + "openai-prod", "sk-secret", created_by="admin", operation_id="create-1" + ) + assert pending.status == "pending" + active = store.activate(pending.secret_id, pending.version, operation_id="commit-1") + assert active.status == "active" + revoked = store.revoke( + pending.secret_id, + pending.version, + revoked_by="admin", + reason="compromised", + operation_id="revoke-1", + ) + assert revoked.status == "revoked" + from EvoScientist.llm.model_config import SecretReference + + with pytest.raises(EvoRuntimeError) as exc: + store.resolve(SecretReference(pending.ref, pending.version)) + assert exc.value.code == "MODEL_CREDENTIAL_REVOKED" + + +def test_retired_secret_remains_resolvable_for_frozen_runs(tmp_path: Path) -> None: + from EvoScientist.llm.model_config import SecretReference + + store = EncryptedModelSecretStore( + tmp_path / "model_secrets.sqlite", master_secret="x" * 32 + ) + first = store.create_pending( + "openai-prod", "sk-old", created_by="admin", operation_id="create-old" + ) + store.activate(first.secret_id, first.version, operation_id="activate-old") + second = store.create_pending( + "openai-prod", "sk-new", created_by="admin", operation_id="create-new" + ) + store.activate(second.secret_id, second.version, operation_id="activate-new") + + metadata = {item.version: item for item in store.list_metadata()} + assert metadata[first.version].status == "retired" + assert store.resolve(SecretReference(first.ref, first.version)).value == "sk-old" + + +def test_v2_converter_refuses_to_guess_multiple_endpoints() -> None: + from tests.v3_fixtures import v3_payload as legacy_v2_payload + + payload = legacy_v2_payload(revision=1) + provider = payload["providers"]["custom-openai"] + provider["endpoints"].append( + { + **provider["endpoints"][0], + "name": "secondary", + "base_url": "https://secondary.example/v1", + } + ) + payload["endpoint_pools"]["default"]["endpoints"].append( + {"name": "secondary", "weight": 1} + ) + draft, report = convert_v2_to_v3_draft( + payload, + target_revision=2, + config_identity_key_id="identity-v1", + ) + assert draft["schema_version"] == 3 + assert not draft["providers"] + assert report.blocking_issues[0]["code"] == "MULTIPLE_ENDPOINTS_REQUIRE_SPLIT" + + +def test_v2_converter_projects_single_endpoint_for_direct_editing() -> None: + from tests.v3_fixtures import v3_payload as legacy_v2_payload + + payload = legacy_v2_payload(revision=4) + draft, report = convert_v2_to_v3_draft( + payload, + target_revision=4, + config_identity_key_id="identity-v1", + ) + + assert not report.blocking_issues + assert draft["schema_version"] == 3 + assert draft["providers"][0]["provider_id"] == "custom-openai" + assert draft["providers"][0]["models"][0]["provider_model_id"] == "model-id" + assert draft["purpose_routes"]["main_agent"]["default_alias"] == "visible-model" + EvoModelConfig.parse(draft, require_evidence=False) + + +def test_stateless_adapter_invariants_and_partial_usage() -> None: + registry = get_adapter_registry() + for adapter_id, revision in (("openai", "openai-v1"), ("xai", "xai-v1")): + params = registry.get(adapter_id, revision).compile_runtime_parameters( + "responses", {"reasoning_effort": "high"}, 4096 + ) + assert params["store"] is False + assert "previous_response_id" not in params + gemini = registry.get("google-gemini", "google-gemini-v1") + assert gemini.compile_runtime_parameters("interactions", {}, 4096)["store"] is False + assert ( + NormalizedUsage(10, 5, None, finality="partial").confirmed_projection() is None + ) + + +def test_openai_chat_compilation_preserves_reasoning_effort() -> None: + registration = get_adapter_registry().get("openai", "openai-v1") + + params = registration.compile_runtime_parameters( + "chat_completions", {"reasoning": "max"}, 8192 + ) + + assert params == { + "max_tokens": 8192, + "use_responses_api": False, + "reasoning_effort": "max", + } + + +@pytest.mark.parametrize( + ("adapter_id", "adapter_revision", "model_id", "api_mode", "limit", "expected"), + [ + ( + "generic-openai-compatible", + "generic-openai-compatible-v1", + "qwen3.7-plus", + "chat_completions", + 65_536, + {"max_tokens": 65_536, "use_responses_api": False}, + ), + ( + "openai", + "openai-v1", + "kimi-for-coding", + "chat_completions", + 65_000, + {"max_completion_tokens": 65_000, "use_responses_api": False}, + ), + ( + "openai", + "openai-v1", + "kimi-for-coding-highspeed", + "chat_completions", + 65_000, + {"max_completion_tokens": 65_000, "use_responses_api": False}, + ), + ( + "openai", + "openai-v1", + "k3", + "chat_completions", + 65_000, + {"max_completion_tokens": 65_000, "use_responses_api": False}, + ), + ( + "openai", + "openai-v1", + "gpt-5.6", + "responses", + 65_000, + { + "max_output_tokens": 65_000, + "use_responses_api": True, + "store": False, + }, + ), + ( + "openai", + "openai-v1", + "gpt-5.6-sol", + "responses", + 65_000, + { + "max_output_tokens": 65_000, + "use_responses_api": True, + "store": False, + }, + ), + ( + "openai", + "openai-v1", + "gpt-5.6-terra", + "responses", + 65_000, + { + "max_output_tokens": 65_000, + "use_responses_api": True, + "store": False, + }, + ), + ], +) +def test_current_model_call_plans_compile_to_one_api_envelope( + adapter_id: str, + adapter_revision: str, + model_id: str, + api_mode: str, + limit: int, + expected: dict[str, int | bool], +) -> None: + params = get_adapter_registry().get( + adapter_id, adapter_revision + ).compile_runtime_parameters( + api_mode, {}, limit, provider_model_id=model_id + ) + + assert params == expected + + +def test_openai_gpt_chat_plan_uses_completion_tokens_not_responses_tokens() -> None: + params = get_adapter_registry().get("openai", "openai-v1").compile_runtime_parameters( + "chat_completions", {}, 65_000, provider_model_id="gpt-5.6" + ) + + assert params == {"max_completion_tokens": 65_000, "use_responses_api": False} + + +def test_kimi_discovery_descriptor_is_partial_and_has_official_reasoning_policy() -> None: + registration = get_adapter_registry().get("openai", "openai-v1") + + descriptor = registration.resolve_discovery_descriptor("k3") + + assert descriptor is not None + assert descriptor.context_tokens == 1_048_576 + assert descriptor.max_output_tokens is None + assert descriptor.reasoning_efforts == ("low", "high", "max") + assert descriptor.default_reasoning_effort == "high" + + +def test_model_reasoning_policy_can_restrict_efforts_and_enable_max() -> None: + payload = v3_payload() + provider = next(item for item in payload["providers"] if item["adapter_id"] == "openai") + model = provider["models"][0] + model["capabilities"]["thinking"] = True + model["parameters"]["reasoning_policy"] = { + "mode": "effort", + "allowed_efforts": ["low", "high", "max"], + "default_effort": "high", + } + + config = EvoModelConfig.parse(payload, require_evidence=False) + parsed = config.providers[provider["provider_id"]].models[model["model_key"]] + + assert parsed.supports_reasoning is True + assert parsed.allowed_reasoning_efforts == ("high", "low", "max") + assert parsed.reasoning_mode == "effort" + assert parsed.reasoning_enabled_params == {"reasoning": "high"} + + +def test_generic_openai_compatible_adapter_supports_standard_model_discovery() -> None: + registration = get_adapter_registry().get( + "generic-openai-compatible", "generic-openai-compatible-v1" + ) + + assert registration.discovery_capability is True + + +def test_generic_adapter_does_not_publish_unimplemented_reasoning_capability() -> None: + payload = v3_payload() + provider = payload["providers"][1] + provider.update( + { + "adapter_id": "generic-openai-compatible", + "adapter_revision": "generic-openai-compatible-v1", + "wire_protocol": "openai_compatible", + } + ) + provider["connection"]["base_url"] = "https://provider.example/v1" + for model in provider["models"]: + model["invocation"]["api_mode"] = "chat_completions" + model["capabilities"]["thinking"] = True + + config = EvoModelConfig.parse(payload, require_evidence=False) + model = config.providers["openai-prod"].models["general"] + + assert model.capabilities["thinking"] is True + assert model.reasoning_mode == "none" + assert model.supports_reasoning is False + assert model.allowed_reasoning_efforts == () + + +def test_aliases_share_capability_evidence_but_not_invocation_identity() -> None: + payload = v3_payload() + openai_model = payload["providers"][1]["models"][0] + openai_model["parameters"]["user_options"]["temperature"] = { + "default": 0.2, + "applies_to": ["main_agent"], + "minimum": 0, + "maximum_exclusive": 2, + } + payload["aliases"][2]["defaults"] = {"temperature": 0.2} + payload["aliases"].append( + { + **payload["aliases"][2], + "alias": "openai-prod-creative", + "display_name": "OpenAI creative", + "defaults": {"temperature": 0.8}, + } + ) + config = EvoModelConfig.parse(payload, require_evidence=False) + general = config.concrete_routes( + config.main_routes.selectable["openai-prod-general"] + )[0] + creative = config.concrete_routes( + config.main_routes.selectable["openai-prod-creative"] + )[0] + key = b"identity-test-key-32-bytes-long!" + + assert route_semantics_hash(config, general, key) == route_semantics_hash( + config, creative, key + ) + assert invocation_fingerprint( + config, general, "main_agent", {"temperature": 0.2}, key + ) != invocation_fingerprint( + config, creative, "main_agent", {"temperature": 0.8}, key + ) + option = ( + config.providers["openai-prod"].models["general"].user_options["temperature"] + ) + assert option["minimum"] == 0 + assert option["maximum_exclusive"] == 2 + + +def test_user_option_cannot_loosen_adapter_bounds() -> None: + payload = v3_payload() + payload["providers"][1]["models"][0]["parameters"]["user_options"][ + "temperature" + ] = {"default": 2, "maximum": 2} + with pytest.raises(EvoRuntimeError) as exc: + EvoModelConfig.parse(payload, require_evidence=False) + assert exc.value.code == "MODEL_PARAMETER_INVALID" + + +def test_explicit_system_purpose_alias_is_preserved() -> None: + payload = v3_payload() + payload["purpose_routes"]["tool_selector"] = { + "default_alias": "anthropic-prod-fast" + } + config = EvoModelConfig.parse(payload, require_evidence=False) + selector_id = config.purpose_selector_ids["tool_selector"] + assert config.route_selectors[selector_id].alias == "anthropic-prod-fast" + + +def test_validate_rejects_conflict_after_alias_parameter_merge() -> None: + payload = v3_payload() + model = payload["providers"][1]["models"][0] + model["parameters"]["user_options"]["temperature"] = {"default": 0.2} + payload["aliases"][2]["defaults"] = {"temperature": 0.8} + payload["purpose_defaults"]["main_agent"] = {"top_p": 0.9} + with pytest.raises(EvoRuntimeError) as exc: + EvoModelConfig.parse(payload, require_evidence=False) + assert exc.value.code == "MODEL_PARAMETER_CONFLICT" + + +def test_anthropic_thinking_budget_is_strictly_below_output_limit() -> None: + registration = get_adapter_registry().get("anthropic", "anthropic-v1") + params = registration.compile_runtime_parameters( + "messages", {"thinking_enabled": True}, 4096 + ) + assert 1024 <= params["thinking"]["budget_tokens"] < params["max_tokens"] + + +def test_gemini_interactions_preserves_thought_signature() -> None: + from types import SimpleNamespace + + from EvoScientist.llm.gemini_interactions import _chat_result, _compile_messages + + class Block: + def model_dump(self, **_kwargs): + return {"type": "thought", "signature": "signed-opaque", "summary": []} + + response = SimpleNamespace( + outputs=[Block()], + usage=SimpleNamespace( + total_input_tokens=2, + total_cached_tokens=0, + total_output_tokens=3, + total_thought_tokens=1, + total_tokens=5, + ), + id="provider-id", + status="completed", + model=SimpleNamespace(id="gemini-fixture"), + ) + message = _chat_result(response).generations[0].message + turns, _ = _compile_messages([message]) + assert turns[0]["content"][0]["signature"] == "signed-opaque" + + +@pytest.mark.asyncio +async def test_gemini_interactions_streams_and_replays_signed_blocks( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from langchain_core.messages import HumanMessage + + from EvoScientist.llm.gemini_interactions import ( + GeminiInteractionsChatModel, + _compile_messages, + create_gemini_interactions_model, + ) + + events = [ + { + "event_type": "content.start", + "index": 0, + "content": {"type": "text", "text": ""}, + }, + { + "event_type": "content.delta", + "index": 0, + "delta": {"type": "text", "text": "hello"}, + }, + {"event_type": "content.stop", "index": 0}, + { + "event_type": "content.start", + "index": 1, + "content": {"type": "thought", "summary": []}, + }, + { + "event_type": "content.delta", + "index": 1, + "delta": {"type": "thought_signature", "signature": "signed-stream"}, + }, + {"event_type": "content.stop", "index": 1}, + { + "event_type": "content.start", + "index": 2, + "content": { + "type": "function_call", + "id": "call-1", + "name": "probe", + "arguments": {"value": "ok"}, + }, + }, + {"event_type": "content.stop", "index": 2}, + { + "event_type": "interaction.complete", + "interaction": { + "id": "request-1", + "status": "completed", + "model": {"id": "gemini-fixture"}, + "usage": { + "total_input_tokens": 2, + "total_cached_tokens": 0, + "total_output_tokens": 3, + }, + }, + }, + ] + + class Stream: + def __aiter__(self): + self.iterator = iter(events) + return self + + async def __anext__(self): + try: + return next(self.iterator) + except StopIteration as exc: + raise StopAsyncIteration from exc + + class Interactions: + async def create(self, **request): + assert request["stream"] is True + assert request["store"] is False + return Stream() + + class Client: + aio = type("AsyncClient", (), {"interactions": Interactions()})() + + monkeypatch.setattr(GeminiInteractionsChatModel, "_client", lambda self: Client()) + model = create_gemini_interactions_model(model="gemini-fixture", api_key="secret") + chunks = [chunk async for chunk in model._astream([HumanMessage("hi")])] + combined = chunks[0].message + for chunk in chunks[1:]: + combined += chunk.message + + assert combined.text == "hello" + assert combined.tool_calls[0]["name"] == "probe" + assert combined.usage_metadata["cached_input_tokens"] == 0 + turns, _ = _compile_messages([combined]) + assert turns[0]["content"][1]["signature"] == "signed-stream" diff --git a/tests/test_recoverable_tools.py b/tests/test_recoverable_tools.py new file mode 100644 index 0000000..b42e0d2 --- /dev/null +++ b/tests/test_recoverable_tools.py @@ -0,0 +1,53 @@ +from __future__ import annotations + +from EvoScientist.middleware import recoverable_tools + + +def test_tool_effect_context_prefers_dedicated_grant(monkeypatch): + monkeypatch.setattr( + "langgraph.config.get_config", + lambda: { + "configurable": { + "ai4sci_model_proxy": { + "gateway_url": "http://gateway", + "run_id": "model-run", + "envelope_signature": "model-signature", + }, + "ai4sci_tool_effect": { + "gateway_url": "http://gateway", + "run_id": "tool-run", + "envelope_signature": "tool-signature", + }, + }, + "metadata": {}, + }, + ) + + proxy, _ = recoverable_tools._context() + + assert proxy == { + "gateway_url": "http://gateway", + "run_id": "tool-run", + "envelope_signature": "tool-signature", + } + + +def test_evomemory_never_falls_back_to_parent_model_proxy(monkeypatch): + monkeypatch.setattr( + "langgraph.config.get_config", + lambda: { + "configurable": { + "ai4sci_model_proxy": { + "gateway_url": "http://gateway", + "run_id": "parent-run", + "envelope_signature": "parent-signature", + }, + }, + "metadata": {"run_kind": "evomemory_turn_worker"}, + }, + ) + + proxy, metadata = recoverable_tools._context() + + assert proxy is None + assert metadata["run_kind"] == "evomemory_turn_worker" diff --git a/tests/test_runtime_integrations.py b/tests/test_runtime_integrations.py index 9e3ffa7..37c6a9e 100644 --- a/tests/test_runtime_integrations.py +++ b/tests/test_runtime_integrations.py @@ -15,7 +15,6 @@ from EvoScientist.runtime_integrations import ( handle_knowledge_file, record_service_usage, reset_runtime_integrations, - resolve_runtime_model, ) @@ -101,17 +100,3 @@ async def test_host_can_register_runtime_integrations(tmp_path): assert get_image_backend() is image_backend assert knowledge_paths == [path] assert usage == [("mineru", "parse")] - - -def test_host_can_register_model_resolver(): - resolved = object() - calls = [] - - def resolve_model(model, provider): - calls.append((model, provider)) - return resolved - - configure_runtime_integrations(model_resolver=resolve_model) - - assert resolve_runtime_model("model-a", "provider-a") is resolved - assert calls == [("model-a", "provider-a")] diff --git a/tests/test_sessions.py b/tests/test_sessions.py index 60839e1..12b5b9a 100644 --- a/tests/test_sessions.py +++ b/tests/test_sessions.py @@ -16,6 +16,7 @@ from langgraph.graph.message import REMOVE_ALL_MESSAGES from EvoScientist.sessions import ( AGENT_NAME, + _checkpoint_serde, _format_relative_time, _reduce_messages_delta, delete_thread, @@ -76,6 +77,59 @@ class TestGetDbPath(unittest.TestCase): assert ".evoscientist" in long_form or "evoscientist" in long_form.lower() +def test_checkpoint_serde_allows_app_owned_errors(): + allowed = _checkpoint_serde()._allowed_msgpack_modules + assert ("EvoScientist.llm.errors", "AgentControlError") in allowed + assert ("EvoScientist.llm.errors", "ModelToolProtocolError") in allowed + assert ("EvoScientist.llm.errors", "ProviderStreamError") in allowed + + +def test_checkpoint_serde_roundtrips_model_tool_protocol_error(): + from EvoScientist.llm.errors import ModelToolProtocolError + + error = ModelToolProtocolError( + "missing_name", + provider="openai", + model="gpt-example", + route_key="openai:primary:gpt-example", + config_generation=7, + call_id="call-1", + call_diagnostic={"raw": "must-not-be-checkpointed"}, + ) + serde = _checkpoint_serde() + restored = serde.loads_typed(serde.dumps_typed({"error": error}))["error"] + + assert isinstance(restored, ModelToolProtocolError) + assert restored.code == "MODEL_TOOL_PROTOCOL_INVALID" + assert restored.reason == "missing_name" + assert restored.provider == "openai" + assert restored.config_generation == 7 + assert restored.call_id == "call-1" + assert restored.call_diagnostic == {} + + +def test_checkpoint_serde_roundtrips_provider_stream_error(): + from EvoScientist.llm.errors import ProviderStreamError + + error = ProviderStreamError( + provider="openai", + class_qualname="openai.BadRequestError", + message="Provider rejected the request.", + status_code=400, + code="invalid_request", + request_id="request-1", + ) + serde = _checkpoint_serde() + restored = serde.loads_typed(serde.dumps_typed({"error": error}))["error"] + + assert isinstance(restored, ProviderStreamError) + assert restored.provider == "openai" + assert restored.class_qualname == "openai.BadRequestError" + assert restored.status_code == 400 + assert restored.code == "invalid_request" + assert restored.request_id == "request-1" + + class TestFormatRelativeTime(unittest.TestCase): def test_none(self): assert _format_relative_time(None) == "" diff --git a/tests/test_skill_context_middleware.py b/tests/test_skill_context_middleware.py new file mode 100644 index 0000000..37c744e --- /dev/null +++ b/tests/test_skill_context_middleware.py @@ -0,0 +1,84 @@ +from __future__ import annotations + +from langchain.agents.middleware.types import ModelRequest +from langchain_core.messages import HumanMessage + +from EvoScientist.middleware.skill_context import BudgetedSkillsMiddleware + + +def _skill(name: str, description: str) -> dict[str, object]: + return { + "name": name, + "description": description, + "path": f"/skills/{name}/SKILL.md", + "allowed_tools": [], + } + + +def _request(query: str, skills: list[dict[str, object]]) -> ModelRequest: + return ModelRequest( + model=object(), + messages=[HumanMessage(content=query)], + system_prompt="Base system prompt.", + state={"skills_metadata": skills}, + ) + + +def test_skill_context_prefers_relevant_skill_and_omits_irrelevant_catalog(): + middleware = BudgetedSkillsMiddleware( + backend=object(), sources=["/skills/"], max_skills=2, max_skills_bytes=1024 + ) + result = middleware.modify_request( + _request( + "请分析蛋白质结构预测结果", + [ + _skill("protein-structure", "蛋白质结构预测和结果分析。"), + _skill("frontend-design", "Build polished user interfaces."), + ], + ) + ) + + prompt = str(result.system_message.content) + assert "protein-structure" in prompt + assert "frontend-design" not in prompt + assert "query-relevant subset" in prompt + + +def test_skill_context_enforces_count_and_utf8_budget(): + middleware = BudgetedSkillsMiddleware( + backend=object(), + sources=["/skills/"], + max_skills=16, + max_skills_bytes=2048, + max_description_bytes=128, + ) + skills = [ + _skill(f"analysis-{index}", "analysis " + "x" * 1_000) for index in range(300) + ] + + selected = middleware._select_skills(skills, "analysis") + rendered = middleware._format_budgeted_skills(selected) + + assert len(selected) == 16 + assert len(rendered.encode("utf-8")) <= 2048 + assert "analysis-299" not in rendered + + +def test_skill_context_does_not_fall_back_to_all_skills_without_a_match(): + middleware = BudgetedSkillsMiddleware(backend=object(), sources=["/skills/"]) + result = middleware.modify_request( + _request( + "unrelated request", + [_skill("protein-structure", "Protein folding workflow.")], + ) + ) + + prompt = str(result.system_message.content) + assert "protein-structure" not in prompt + assert "query-relevant subset" in prompt + + +def test_skill_context_accepts_a_single_skill_source(): + middleware = BudgetedSkillsMiddleware(backend=object(), sources="/skills/") + + assert middleware.sources == ["/skills/"] diff --git a/tests/test_tool_protocol_guard.py b/tests/test_tool_protocol_guard.py index a82e5a5..6dc01b0 100644 --- a/tests/test_tool_protocol_guard.py +++ b/tests/test_tool_protocol_guard.py @@ -35,7 +35,6 @@ def _call(call_id: str = "call-1", name: str = "search", args: Any = None): (_call(name=""), "missing_name"), (_call(name=" "), "missing_name"), (_call(name="missing"), "unknown_name"), - (_call(call_id=""), "missing_id"), ], ) def test_invalid_final_tool_call_fails_closed(call, reason): @@ -46,14 +45,28 @@ def test_invalid_final_tool_call_fails_closed(call, reason): middleware.wrap_model_call(request, lambda _request: _response(call)) assert caught.value.reason == reason - assert caught.value.retryable is False + assert caught.value.retryable is True assert caught.value.fallbackable is True + assert caught.value.non_fallbackable is False -def test_non_mapping_args_are_rejected_if_adapter_bypasses_message_validation(): +def test_json_string_args_are_normalized_to_an_object(): message = AIMessage(content="", tool_calls=[_call()]) message.tool_calls[0]["args"] = "{}" + result = ToolProtocolGuardMiddleware().wrap_model_call( + _Request(tools=[{"name": "search"}]), + lambda _request: ModelResponse(result=[message]), + ) + + assert result.result[0].tool_calls[0]["args"] == {} + assert message.tool_calls[0]["args"] == "{}" + + +def test_malformed_json_args_are_rejected(): + message = AIMessage(content="", tool_calls=[_call()]) + message.tool_calls[0]["args"] = "{" + with pytest.raises(ModelToolProtocolError) as caught: ToolProtocolGuardMiddleware().wrap_model_call( _Request(tools=[{"name": "search"}]), @@ -121,7 +134,30 @@ def test_content_block_must_match_parsed_call(): lambda _request: response, ) - assert caught.value.reason == "inconsistent_block" + assert caught.value.reason == "inconsistent_source" + + +def test_responses_content_block_uses_call_id_over_output_item_id(): + """Responses item IDs are not the identifiers used for tool results.""" + + response = _response( + _call(call_id="call-result-1", name="search", args={}), + content=[ + { + "type": "function_call", + "id": "fc-output-item-1", + "call_id": "call-result-1", + "name": "search", + "arguments": "{}", + } + ], + ) + + result = ToolProtocolGuardMiddleware().wrap_model_call( + _Request(tools=[{"name": "search"}]), lambda _request: response + ) + + assert result.result[0].tool_calls[0]["id"] == "call-result-1" def test_parsed_only_valid_call_and_extended_response_pass(): @@ -172,37 +208,50 @@ def test_error_carries_safe_route_metadata(): assert "args" not in payload -def test_missing_id_carries_redacted_call_diagnostic_only_for_internal_logging(): +def test_missing_id_is_generated_without_mutating_the_provider_message(): call = _call(call_id="", name="search", args={"query": "private search text"}) - with pytest.raises(ModelToolProtocolError) as caught: - ToolProtocolGuardMiddleware().wrap_model_call( - _Request(tools=[{"name": "search"}]), - lambda _request: _response(call), - ) + provider_response = _response(call) + result = ToolProtocolGuardMiddleware().wrap_model_call( + _Request(tools=[{"name": "search"}]), + lambda _request: provider_response, + ) - diagnostic = caught.value.call_diagnostic - assert diagnostic == { - "source": "parsed_tool_calls", - "call_index": 0, - "call_count": 1, - "call_type": "object", - "name": "search", - "id_present": False, - "args_present": True, - "args_type": "object", - "args_key_count": 1, - "args_keys": ["query"], - "args_keys_truncated": False, - "args_digest": diagnostic["args_digest"], - "raw_openai_call_available": False, - } - assert diagnostic["args_digest"].startswith("sha256:") - assert "private search text" not in str(diagnostic) - assert "call_diagnostic" not in caught.value.model_dump() + normalized = result.result[0].tool_calls[0] + assert normalized["id"].startswith("call_") + assert normalized["args"] == {"query": "private search text"} + assert provider_response.result[0].tool_calls[0]["id"] == "" -def test_diagnostic_compares_parsed_and_preserved_raw_openai_call_shapes(): +def test_generic_openai_compatible_parallel_calls_without_ids_are_stable(): + model = SimpleNamespace( + metadata={ + "route_adapter_id": "generic-openai-compatible", + "route_key": "qwen-primary", + "route_config_generation": 7, + "route_api_mode": "chat_completions", + } + ) + response = _response( + _call(call_id="", name="think_tool", args={"reflection": "plan"}), + _call(call_id="", name="execute", args={"command": "pwd"}), + ) + + result = ToolProtocolGuardMiddleware().wrap_model_call( + _Request( + tools=[{"name": "think_tool"}, {"name": "execute"}], model=model + ), + lambda _request: response, + ) + + calls = result.result[0].tool_calls + assert [call["name"] for call in calls] == ["think_tool", "execute"] + assert all(call["id"].startswith("call_") for call in calls) + assert calls[0]["id"] != calls[1]["id"] + assert response.result[0].tool_calls[0]["id"] == "" + + +def test_raw_openai_id_is_merged_into_the_canonical_call(): parsed = _call(call_id="", name="search", args={"query": "secret"}) raw = { "id": "provider-call-id", @@ -215,19 +264,62 @@ def test_diagnostic_compares_parsed_and_preserved_raw_openai_call_shapes(): additional_kwargs={"tool_calls": [raw]}, ) - with pytest.raises(ModelToolProtocolError) as caught: - ToolProtocolGuardMiddleware().wrap_model_call( - _Request(tools=[{"name": "search"}]), - lambda _request: ModelResponse(result=[message]), - ) + result = ToolProtocolGuardMiddleware().wrap_model_call( + _Request(tools=[{"name": "search"}]), + lambda _request: ModelResponse(result=[message]), + ) - diagnostic = caught.value.call_diagnostic - assert diagnostic["id_present"] is False - assert diagnostic["raw_openai_call_available"] is True - assert diagnostic["raw_openai_call"]["id_present"] is True - assert diagnostic["raw_openai_call"]["name"] == "search" - assert "provider-call-id" not in str(diagnostic) - assert "secret" not in str(diagnostic) + normalized = result.result[0] + assert normalized.tool_calls == [ + { + "id": "provider-call-id", + "name": "search", + "args": {"query": "secret"}, + "type": "tool_call", + } + ] + assert "tool_calls" not in normalized.additional_kwargs + assert message.additional_kwargs["tool_calls"] == [raw] + + +def test_content_only_function_call_is_decoded_and_normalized(): + message = AIMessage( + content=[ + { + "type": "function_call", + "id": "", + "name": "search", + "arguments": '{"query":"x"}', + } + ] + ) + + result = ToolProtocolGuardMiddleware().wrap_model_call( + _Request(tools=[{"name": "search"}]), + lambda _request: ModelResponse(result=[message]), + ) + + normalized = result.result[0] + call_id = normalized.tool_calls[0]["id"] + assert call_id.startswith("call_") + assert normalized.tool_calls[0]["args"] == {"query": "x"} + assert normalized.content[0]["id"] == call_id + + +def test_legacy_function_call_is_decoded_and_removed_from_replay_metadata(): + legacy = {"name": "search", "arguments": '{"query":"x"}'} + message = AIMessage(content="", additional_kwargs={"function_call": legacy}) + + result = ToolProtocolGuardMiddleware().wrap_model_call( + _Request(tools=[{"name": "search"}]), + lambda _request: ModelResponse(result=[message]), + ) + + normalized = result.result[0] + assert normalized.tool_calls[0]["id"].startswith("call_") + assert normalized.tool_calls[0]["args"] == {"query": "x"} + assert "function_call" not in normalized.additional_kwargs + assert message.additional_kwargs["function_call"] == legacy def test_diagnostic_failure_cannot_mask_the_protocol_error(): @@ -244,3 +336,18 @@ def test_diagnostic_failure_cannot_mask_the_protocol_error(): assert caught.value.reason == "missing_id" assert caught.value.call_diagnostic["args_digest"].startswith("sha256:") + + +def test_protocol_failure_logs_only_redacted_call_diagnostic(caplog): + call = _call(name="", args={"query": "private search text"}) + + with caplog.at_level("WARNING"), pytest.raises(ModelToolProtocolError): + ToolProtocolGuardMiddleware().wrap_model_call( + _Request(tools=[{"name": "search"}]), + lambda _request: _response(call), + ) + + record = caplog.records[-1].getMessage() + assert "reason=missing_name" in record + assert '"args_keys": ["query"]' in record + assert "private search text" not in record diff --git a/tests/test_user_model_options.py b/tests/test_user_model_options.py new file mode 100644 index 0000000..4d279b2 --- /dev/null +++ b/tests/test_user_model_options.py @@ -0,0 +1,103 @@ +from __future__ import annotations + +import pytest + +from EvoScientist.llm.contracts import EvoRuntimeError +from EvoScientist.llm.user_options import ( + model_options_schema_hash, + project_user_options_for_purpose, + validate_user_model_options, +) + + +def _validate(options): + return validate_user_model_options( + supplied=options, + user_options={ + "temperature": { + "type": "number", + "minimum": 0, + "maximum": 2, + "applies_to": ["main_agent"], + }, + "top_p": { + "type": "number", + "minimum_exclusive": 0, + "maximum": 1, + "applies_to": ["main_agent"], + }, + }, + supports_reasoning=True, + reasoning_mode="effort", + allowed_reasoning_efforts=("low", "high"), + default_reasoning_effort="high", + parameter_constraints=({"at_most_one_of": ("temperature", "top_p")},), + ) + + +def test_user_options_validate_canonical_values(): + assert _validate({"temperature": 0.4, "reasoning": "high"}) == { + "temperature": 0.4, + "reasoning": "high", + } + + +def test_purpose_projection_filters_user_options_but_preserves_runtime_parameters(): + projected = project_user_options_for_purpose( + values={"temperature": 1, "structured_output": True}, + user_options={ + "temperature": { + "type": "number", + "applies_to": ["main_agent"], + } + }, + purpose="tool_selector", + ) + + assert projected == {"structured_output": True} + + +@pytest.mark.parametrize( + "options", + [ + {"temperature": 3}, + {"top_p": 0}, + {"unknown": True}, + {"temperature": 0.4, "top_p": 0.8}, + ], +) +def test_user_options_reject_invalid_or_conflicting_values(options): + with pytest.raises(EvoRuntimeError): + _validate(options) + + +def test_options_schema_hash_tracks_semantics_not_defaults(): + common = { + "model_profile_id": "profile-1", + "supports_reasoning": False, + "reasoning_mode": "none", + "allowed_reasoning_efforts": (), + "parameter_constraints": (), + "adapter_id": "openai-compatible", + "adapter_revision": "3", + } + first = model_options_schema_hash( + **common, + user_options={"temperature": {"type": "number", "maximum": 2, "default": 1}}, + ) + default_changed = model_options_schema_hash( + **common, + user_options={"temperature": {"type": "number", "maximum": 2, "default": 0}}, + ) + constraint_changed = model_options_schema_hash( + **common, + user_options={"temperature": {"type": "number", "maximum": 1, "default": 0}}, + ) + adapter_changed = model_options_schema_hash( + **{**common, "adapter_revision": "4"}, + user_options={"temperature": {"type": "number", "maximum": 2, "default": 0}}, + ) + + assert first == default_changed + assert first != constraint_changed + assert first != adapter_changed diff --git a/tests/test_user_options_runtime.py b/tests/test_user_options_runtime.py new file mode 100644 index 0000000..985ca40 --- /dev/null +++ b/tests/test_user_options_runtime.py @@ -0,0 +1,146 @@ +from __future__ import annotations + +from dataclasses import replace +from types import SimpleNamespace + +import pytest + +from EvoScientist.llm.configuration.provider import ResolvedSecret +from EvoScientist.llm.contracts import ( + AgentInputV3, + EvoRuntimeError, + HmacGrantAuthority, + WebHostContext, +) +from EvoScientist.llm.model_config import EvoModelConfig +from EvoScientist.llm.runtime import EvoModelRuntime +from EvoScientist.llm.user_options import model_options_schema_hash +from tests.test_provider_model_config_v3 import v3_payload as provider_v3_payload +from tests.test_web_model_runtime import ( + _input, + _preparation, + _runtime, + _Sink, +) +from tests.v3_fixtures import ( + IDENTITY_KEY_ID, + RUNTIME_KEY_ID, + RUNTIME_SECRET, + identity_ring, +) + + +@pytest.mark.asyncio +async def test_prepare_rejects_gateway_catalog_schema_that_became_stale( + monkeypatch, tmp_path +): + runtime, authority = _runtime(tmp_path, monkeypatch) + original = _input() + agent_input = replace( + original, + metadata={ + **dict(original.metadata), + "model_options_schema_hash": "sha256:stale", + }, + ) + host = WebHostContext( + "/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink() + ) + + with pytest.raises(EvoRuntimeError, match="MODEL_OPTIONS_STALE"): + await runtime.prepare_model_run( + _preparation(authority, agent_input), + agent_input, + host, + ) + + +@pytest.mark.asyncio +async def test_prepare_filters_main_only_options_from_inherited_auxiliary_routes(): + payload = provider_v3_payload() + payload["config_identity_key_id"] = IDENTITY_KEY_ID + provider = next( + item for item in payload["providers"] if item["adapter_id"] == "openai" + ) + model_payload = provider["models"][0] + model_payload["capabilities"]["thinking"] = True + model_payload["parameters"]["reasoning_policy"] = { + "mode": "effort", + "allowed_efforts": ["low", "medium", "high"], + "default_effort": "high", + } + model_payload["parameters"]["user_options"]["temperature"] = { + "default": 1, + "applies_to": ["main_agent"], + "minimum": 0, + "maximum_exclusive": 2, + } + alias_payload = next( + item for item in payload["aliases"] if item["alias"] == "openai-prod-general" + ) + alias_payload["defaults"] = {"temperature": 1} + config = EvoModelConfig.parse(payload, require_evidence=False) + authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID) + + def resolve_secret(reference): + return ResolvedSecret( + "test-secret", reference.revision, "1", "test-fingerprint" + ) + + runtime = EvoModelRuntime( + SimpleNamespace(load=lambda: config), + admission_verifier=authority, + quote_authority=authority, + identity_key_ring=identity_ring(), + secret_resolver=resolve_secret, + ) + selector = config.resolve_main_selector("openai-prod-general") + selected_model = config.providers[selector.provider].models[selector.model] + selected_provider = config.providers[selector.provider] + schema_hash = model_options_schema_hash( + model_profile_id="openai-prod-general", + user_options=selected_model.user_options, + supports_reasoning=selected_model.supports_reasoning, + reasoning_mode=selected_model.reasoning_mode, + allowed_reasoning_efforts=selected_model.allowed_reasoning_efforts, + parameter_constraints=selected_model.parameter_constraints, + adapter_id=selected_provider.adapter_id, + adapter_revision=selected_provider.adapter_revision, + ) + agent_input = AgentInputV3( + "hello", + "web:user:thread", + metadata={ + "source": "web", + "model_options": {"temperature": 0.7}, + "model_options_schema_hash": schema_hash, + }, + ) + grant = authority.sign_preparation( + request_id="11111111-1111-4111-8111-111111111111", + turn_id="22222222-2222-4222-8222-222222222222", + thread_id="thread", + subject_id="user", + requested_model_ref="openai-prod-general", + plan="starter", + roles=("user",), + requires_vision=False, + reasoning_effort="high", + title_policy="best_effort", + gateway_input_digest=authority.agent_input_digest(agent_input.projection()), + checkpoint_thread_id=agent_input.checkpoint_thread_id, + checkpoint_snapshot_id="sha256:checkpoint", + turn_fencing_token=1, + ttl_ms=60_000, + ) + + await runtime.prepare_model_run( + grant, + agent_input, + WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink()), + ) + + snapshot = next(iter(runtime._prepared.values())).snapshot + assert snapshot.purpose_routes["main_agent"][0].params["temperature"] == 0.7 + assert "temperature" not in snapshot.purpose_routes["tool_selector"][0].params + assert "temperature" not in snapshot.purpose_routes["deepagents_summarizer"][0].params diff --git a/tests/test_v3_contracts_and_fencing.py b/tests/test_v3_contracts_and_fencing.py new file mode 100644 index 0000000..6ab5b07 --- /dev/null +++ b/tests/test_v3_contracts_and_fencing.py @@ -0,0 +1,64 @@ +from __future__ import annotations + +from dataclasses import replace + +import pytest + +from EvoScientist.llm.contracts import EvoRuntimeError, HmacGrantAuthority +from EvoScientist.llm.crypto import canonical_json_v1 +from EvoScientist.sessions import FencedPruningCheckpointer +from tests.v3_fixtures import RUNTIME_KEY_ID, RUNTIME_SECRET + + +def test_canonical_json_normalizes_unicode_and_orders_keys(): + assert canonical_json_v1({"z": "e\u0301", "a": 1}) == canonical_json_v1( + {"a": 1, "z": "\u00e9"} + ) + + +def test_previous_runtime_key_verifies_during_rotation(): + old_secret = "old-runtime-secret-for-tests-000000000000000000000" + old = HmacGrantAuthority(old_secret, "old-key") + grant = old.sign_subject( + subject_id="user", plan="starter", roles=("user",), ttl_ms=60_000 + ) + rotated = HmacGrantAuthority( + RUNTIME_SECRET, + RUNTIME_KEY_ID, + previous_secret=old_secret, + previous_key_id="old-key", + ) + assert rotated.verify_subject(grant) + with pytest.raises(EvoRuntimeError, match="CONTRACT_SIGNATURE_INVALID"): + rotated.require_admin(replace(grant, audience="evoscientist-runtime")) + + +@pytest.mark.asyncio +async def test_stale_turn_token_cannot_write_or_release_new_lease(tmp_path): + database = tmp_path / "sessions.sqlite" + async with FencedPruningCheckpointer.from_conn_string_with_keep( + str(database), keep_per_ns=10 + ) as saver: + first = await saver.acquire_turn_lease( + "web:user:thread", "owner-1", ttl_seconds=30 + ) + assert await saver.release_turn_lease(first) + current = await saver.acquire_turn_lease( + "web:user:thread", "owner-2", ttl_seconds=30 + ) + assert current.fencing_token == first.fencing_token + 1 + assert not await saver.release_turn_lease(first) + with pytest.raises(RuntimeError, match="TURN_FENCED"): + await saver.aput_writes( + { + "configurable": { + "thread_id": first.thread_id, + "checkpoint_id": "checkpoint", + "turn_lease_owner": first.owner_id, + "turn_fencing_token": first.fencing_token, + } + }, + [("channel", "value")], + "task", + ) + assert await saver.release_turn_lease(current) diff --git a/tests/test_web_model_runtime.py b/tests/test_web_model_runtime.py new file mode 100644 index 0000000..8cdb7a7 --- /dev/null +++ b/tests/test_web_model_runtime.py @@ -0,0 +1,1024 @@ +from __future__ import annotations + +import asyncio +import base64 +import json +from pathlib import Path +from types import SimpleNamespace + +import pytest +from langchain_core.messages import HumanMessage + +from EvoScientist.llm.adapter_registry import get_adapter_registry +from EvoScientist.llm.contracts import ( + AgentInputV3, + EvoRuntimeError, + HmacGrantAuthority, + WebHostContext, +) +from EvoScientist.llm.errors import ModelToolProtocolError +from EvoScientist.llm.model_config import EvoModelConfig, FileEvoModelConfigStore +from EvoScientist.llm.runtime import ( + _PROTOCOL_MARGIN_TOKENS, + EvoModelRuntime, + _callback_message_schema_debug, + _callback_payload_debug_summary, + _invocation_parameters_debug, + _provider_error_message_debug, + _provider_failure_details, + _provider_input_token_bound, + _run_failure_details, + _safe_error_code, +) +from tests.v3_fixtures import RUNTIME_KEY_ID, RUNTIME_SECRET, identity_ring, v3_payload + + +class _Sink: + def __init__(self) -> None: + self.events = [] + + async def commit(self, event): + self.events.append(event) + return "committed" + + async def confirm(self, _event_id, _payload_digest): + return "committed" + + +class _UncertainSink(_Sink): + def __init__(self, confirmation: str) -> None: + super().__init__() + self.confirmation = confirmation + self.confirmed = [] + + async def commit(self, event): + self.events.append(event) + raise ConnectionError("commit result unavailable") + + async def confirm(self, event_id, payload_digest): + self.confirmed.append((event_id, payload_digest)) + return self.confirmation + + +class _TerminalFailingSink(_Sink): + async def commit(self, event): + if event.kind == "run" and event.payload.get("kind") == "run_terminal": + raise EvoRuntimeError("EVO_EVENT_CONFLICT") + return await super().commit(event) + + +class _Model: + def __init__(self) -> None: + self.metadata = {} + + def model_copy(self, *, update): + copy = _Model() + copy.metadata = update.get("metadata", {}) + return copy + + async def ainvoke(self, _input, **_kwargs): + return SimpleNamespace( + content="Test title", usage_metadata={"input_tokens": 1, "output_tokens": 1} + ) + + +class _Agent: + pass + + +def _runtime(tmp_path: Path, monkeypatch): + monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "test-secret") + authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID) + store = FileEvoModelConfigStore( + tmp_path / "model_routes.yaml", admin_verifier=authority + ) + store.bootstrap_for_development(v3_payload()) + runtime = EvoModelRuntime( + store, + admission_verifier=authority, + quote_authority=authority, + identity_key_ring=identity_ring(), + model_factory=lambda **_kwargs: _Model(), + agent_factory=lambda *_args: _Agent(), + ) + return runtime, authority + + +def _input() -> AgentInputV3: + return AgentInputV3( + "hello", + "web:user:thread", + metadata={"source": "web", "ignored": "not-digested"}, + ) + + +def test_provider_input_bound_counts_decoded_media_separately() -> None: + raw = b"x" * 30_000 + payload = [ + { + "type": "image", + "mime_type": "image/png", + "base64": base64.b64encode(raw).decode("ascii"), + } + ] + + bound = _provider_input_token_bound(payload) + + assert bound.media_blocks == 1 + assert bound.largest_media_bytes >= len(raw) - 2 + assert bound.media_tokens == (len(raw) + 2) // 3 + 512 + assert bound.total_tokens < len(base64.b64encode(raw)) + + +def test_agent_input_digest_includes_forced_context_repair() -> None: + normal = AgentInputV3( + "hello", + "web:user:thread", + metadata={"source": "web", "force_context_repair": False}, + ) + repair = AgentInputV3( + "hello", + "web:user:thread", + metadata={"source": "web", "force_context_repair": True}, + ) + + assert normal.canonical_bytes() != repair.canonical_bytes() + + +def _preparation( + authority, + agent_input, + *, + reasoning_effort="disabled", + title_policy="best_effort", +): + return authority.sign_preparation( + request_id="11111111-1111-4111-8111-111111111111", + turn_id="22222222-2222-4222-8222-222222222222", + thread_id="thread", + subject_id="user", + requested_model_ref="visible-model", + plan="starter", + roles=("user",), + requires_vision=False, + reasoning_effort=reasoning_effort, + title_policy=title_policy, + gateway_input_digest=authority.agent_input_digest(agent_input.projection()), + checkpoint_thread_id=agent_input.checkpoint_thread_id, + checkpoint_snapshot_id="sha256:checkpoint", + turn_fencing_token=1, + ttl_ms=60_000, + ) + + +def test_runtime_preserves_model_tool_protocol_error_code() -> None: + error = ModelToolProtocolError("missing_name", provider="openai") + + assert _safe_error_code(error) == "MODEL_TOOL_PROTOCOL_INVALID" + + +def test_provider_failure_details_are_diagnostic_but_do_not_expose_messages() -> None: + class ProviderBadRequest(RuntimeError): + status_code = 400 + code = "unsupported_parameter" + + def __init__(self) -> None: + super().__init__("request failed with api_key=sk-secret") + self.body = { + "error": { + "code": "unsupported_parameter", + "message": ( + "invalid temperature: only 1 is allowed; api_key=sk-secret" + ), + } + } + + details = _provider_failure_details( + ProviderBadRequest(), + route=None, + error_code="MODEL_PROVIDER_REQUEST_REJECTED", + ) + + assert details == { + "failure_stage": "provider_request", + "reason": "provider_rejected_request", + "provider_error_type": "ProviderBadRequest", + "provider_error_module": __name__, + "http_status": 400, + "provider_error_code": "unsupported_parameter", + "provider_error_parameter": "temperature", + } + assert "sk-secret" not in str(details) + + +def test_request_debug_summary_reports_plan_safe_tool_structure_only() -> None: + payload = [ + [ + { + "type": "ai", + "data": { + "content": [ + {"type": "text", "text": "secret assistant content"}, + {"type": "tool_call", "name": "read_file"}, + ], + "tool_calls": [{"id": "call_1", "name": "read_file"}], + }, + }, + { + "type": "tool", + "data": {"content": "secret tool result"}, + }, + ] + ] + + summary = _callback_payload_debug_summary(payload, 123) + + assert summary["tool_calls"] == 1 + assert summary["tool_results"] == 1 + assert summary["content_block_types"] == "text:1,tool_call:1" + assert "secret" not in str(summary) + + +def test_request_debug_log_projects_parameter_values_without_secrets() -> None: + plan = SimpleNamespace( + sdk_params={ + "max_completion_tokens": 65_000, + "reasoning_effort": "high", + "streaming": True, + "use_responses_api": False, + "extra_body": { + "enable_thinking": True, + "thinking_budget": 8_192, + "private_extension": "secret extension", + }, + "api_key": "sk-secret", + "default_headers": {"Authorization": "Bearer sk-secret"}, + } + ) + + debug = _invocation_parameters_debug(plan) + + assert '"max_completion_tokens":65000' in debug + assert '"reasoning_effort":"high"' in debug + assert '"enable_thinking":true' in debug + assert "api_key" not in debug + assert "Authorization" not in debug + assert "secret" not in debug + + +def test_request_debug_message_schema_identifies_empty_and_metadata_without_content() -> ( + None +): + payload = [ + [ + { + "type": "ai", + "data": { + "content": "secret assistant content", + "additional_kwargs": {"reasoning_content": "secret reasoning"}, + "response_metadata": {"model_name": "secret model"}, + }, + }, + {"type": "ai", "data": {"content": ""}}, + ] + ] + + debug = _callback_message_schema_debug(payload) + + assert '"additional_keys":["reasoning_content"]' in debug + assert '"response_metadata_keys":["model_name"]' in debug + assert '"content_chars":24' in debug + assert '"empty":true' in debug + assert "secret assistant content" not in debug + assert "secret reasoning" not in debug + assert "secret model" not in debug + + +def test_provider_error_debug_message_is_bounded_and_redacts_route_secrets() -> None: + class ProviderBadRequest(RuntimeError): + def __init__(self) -> None: + super().__init__("fallback includes api_key=sk-route-secret") + self.body = { + "error": { + "message": ( + "Invalid max_completion_tokens; " + "Authorization: Bearer sk-route-secret" + ) + } + } + + route = SimpleNamespace( + api_key="sk-route-secret", + default_headers={"X-Provider-Secret": "header-secret-value"}, + ) + + debug = _provider_error_message_debug(ProviderBadRequest(), route) + + assert "Invalid max_completion_tokens" in debug + assert "sk-route-secret" not in debug + assert "header-secret-value" not in debug + assert "" in debug + assert len(debug) <= 1_024 + + +def test_json_decode_diagnostics_record_only_structure_hash_and_trace() -> None: + document = '{"type":"response.output_text.delta","delta":"secret model\noutput"}' + position = document.index("\n") + try: + raise json.JSONDecodeError("Invalid control character", document, position) + except json.JSONDecodeError as caught: + error = caught + + details = _provider_failure_details( + error, + route=None, + error_code="MODEL_PROVIDER_RESPONSE_INVALID", + ) + + assert details["provider_json_document_bytes"] == len(document) + assert details["provider_json_line"] == 1 + assert details["provider_json_column"] == position + 1 + assert details["provider_json_position"] == position + assert details["provider_json_invalid_codepoint"] == "U+000A" + assert details["provider_json_position_inside_string"] == "true" + assert details["provider_json_event_type"] == "response.output_text.delta" + assert details["provider_json_control_codepoints"] == "U+000A:1" + assert details["provider_json_document_sha256"] + assert details["provider_json_window_sha256"] + assert details["provider_json_window_bytes"] == len(document) + assert details["provider_json_trace"] + assert "secret model" not in str(details) + + +def test_stream_timeout_diagnostics_use_structured_exception_attributes() -> None: + class StreamChunkTimeoutError(TimeoutError): + chunks_received = 1 + timeout_s = 120.0 + + details = _provider_failure_details( + StreamChunkTimeoutError("secret upstream message"), + route=None, + error_code="MODEL_TIMEOUT", + ) + + assert details["provider_stream_chunks_received"] == 1 + assert details["provider_stream_idle_timeout_ms"] == 120_000 + assert "secret upstream message" not in str(details) + + +def test_openai_compatible_adapter_classifies_bad_request_as_contract_rejection() -> ( + None +): + class ProviderBadRequest(RuntimeError): + status_code = 400 + + registration = get_adapter_registry().get( + "generic-openai-compatible", "generic-openai-compatible-v1" + ) + + assert registration.classify_error(ProviderBadRequest()).error_code == ( + "MODEL_PROVIDER_REQUEST_REJECTED" + ) + + +def test_openai_compatible_adapter_classifies_invalid_json_as_provider_response_error() -> ( + None +): + registration = get_adapter_registry().get( + "generic-openai-compatible", "generic-openai-compatible-v1" + ) + + assert registration.classify_error( + json.JSONDecodeError("bad", "", 0) + ).error_code == ("MODEL_PROVIDER_RESPONSE_INVALID") + + +def test_generic_openai_compatible_chat_has_a_protocol_margin() -> None: + assert ( + _PROTOCOL_MARGIN_TOKENS[("generic-openai-compatible", "chat_completions")] == 64 + ) + + +def test_runtime_rejects_reasoning_duplicated_in_model_options() -> None: + config = EvoModelConfig.parse(v3_payload(), require_evidence=False) + model = config.providers["custom-openai"].models["model-id"] + registration = get_adapter_registry().get( + "generic-openai-compatible", "generic-openai-compatible-v1" + ) + + with pytest.raises(EvoRuntimeError, match="AGENT_INPUT_MISMATCH"): + EvoModelRuntime._validated_user_options( + model, + registration, + {"reasoning": "medium", "reasoning_effort": "high"}, + "main_agent", + ) + + +def _admission(authority, quote): + return authority.sign_admission( + preparation_id=quote.preparation_id, + request_id=quote.request_id, + turn_id=quote.turn_id, + thread_id=quote.thread_id, + subject_id=quote.subject_id, + requested_model_ref=quote.requested_model_ref, + plan=quote.plan, + roles=quote.roles, + requires_vision=quote.requires_vision, + reasoning_effort=quote.reasoning_effort, + title_policy=quote.title_policy, + gateway_input_digest=quote.gateway_input_digest, + prepared_snapshot_digest=quote.prepared_snapshot_digest, + prepared_input_digest=quote.prepared_input_digest, + config_revision=quote.config_revision, + catalog_revision=quote.catalog_revision, + purpose_attempt_limits=quote.purpose_attempt_limits, + total_max_attempts=quote.total_max_attempts, + checkpoint_snapshot_id=quote.checkpoint_snapshot_id, + tool_registry_snapshot_id=quote.tool_registry_snapshot_id, + turn_fencing_token=quote.turn_fencing_token, + admission_snapshot_id="33333333-3333-4333-8333-333333333333", + admission_id="44444444-4444-4444-8444-444444444444", + hold_id="55555555-5555-4555-8555-555555555555", + billing_fencing_token=1, + provider_run_reserve_microunits=quote.provider_run_reserve_microunits, + billing_policy_version="test-v3", + expires_at=quote.expires_at, + ) + + +@pytest.mark.asyncio +async def test_prepare_freezes_input_and_start_reuses_it(monkeypatch, tmp_path): + runtime, authority = _runtime(tmp_path, monkeypatch) + sink = _Sink() + agent_input = _input() + host = WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=sink) + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), agent_input, host + ) + + assert quote.contract_version == 3 + assert quote.prepared_input_digest.startswith("hmac-sha256:") + assert quote.total_max_attempts == sum(quote.purpose_attempt_limits.values()) + assert quote.provider_run_reserve_microunits == 0 + run = await runtime.start_web_run(_admission(authority, quote)) + assert run.run_id + with pytest.raises(EvoRuntimeError, match="EVENT_CURSOR_INVALID"): + await anext(run.stream(agent_input)) + + +@pytest.mark.asyncio +async def test_disabled_title_is_omitted_from_quote_snapshot_and_model_set( + monkeypatch, tmp_path +): + runtime, authority = _runtime(tmp_path, monkeypatch) + agent_input = _input() + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input, title_policy="disabled"), + agent_input, + WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink()), + ) + + expected = {"main_agent", "tool_selector", "deepagents_summarizer"} + assert set(quote.enabled_purposes) == expected + assert set(quote.purpose_routes) == expected + assert set(quote.purpose_route_call_bounds) == expected + assert set(quote.purpose_attempt_limits) == expected + + run = await runtime.start_web_run(_admission(authority, quote)) + + assert set(run._snapshot.purpose_routes) == expected + assert set(run._snapshot.purpose_route_call_bounds) == expected + assert run._model_set.title is None + + +@pytest.mark.asyncio +async def test_best_effort_title_remains_available(monkeypatch, tmp_path): + runtime, authority = _runtime(tmp_path, monkeypatch) + agent_input = _input() + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input, title_policy="best_effort"), + agent_input, + WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink()), + ) + + assert "title" in quote.enabled_purposes + assert "title" in quote.purpose_routes + assert "title" in quote.purpose_route_call_bounds + + run = await runtime.start_web_run(_admission(authority, quote)) + + assert run._model_set.title is not None + + +@pytest.mark.asyncio +async def test_forced_context_repair_is_consumed_once_before_provider_start( + monkeypatch, tmp_path +): + runtime, authority = _runtime(tmp_path, monkeypatch) + agent_input = AgentInputV3( + "hello", + "web:user:thread", + metadata={"source": "web", "force_context_repair": True}, + ) + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), + agent_input, + WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink()), + ) + run = await runtime.start_web_run(_admission(authority, quote)) + + with pytest.raises(EvoRuntimeError, match="MODEL_CONTEXT_WINDOW_EXCEEDED") as exc: + await run._attempt_callback.on_chat_model_start( + {}, [[HumanMessage(content="hello")]], run_id="repair-trigger" + ) + + assert exc.value.details[0]["reason"] == "forced_context_repair" + assert exc.value.details[0]["repair_requested"] is True + assert run._force_context_repair_pending is False + assert "repair-trigger" not in run._callback_attempts + + await run._attempt_callback.on_chat_model_start( + {}, [[HumanMessage(content="hello")]], run_id="repair-retry" + ) + assert "repair-retry" in run._callback_attempts + await run._finish_callback_attempt( + "repair-retry", + usage={"input_tokens": 1, "output_tokens": 1}, + error_code=None, + ) + + +@pytest.mark.asyncio +async def test_provider_input_hard_cap_reports_media_aware_safe_breakdown( + monkeypatch, tmp_path +): + runtime, authority = _runtime(tmp_path, monkeypatch) + agent_input = _input() + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), + agent_input, + WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink()), + ) + run = await runtime.start_web_run(_admission(authority, quote)) + route = run._snapshot.main_routes[0] + bound = run._snapshot.purpose_route_call_bounds["main_agent"][0] + + with pytest.raises(EvoRuntimeError, match="MODEL_CONTEXT_WINDOW_EXCEEDED") as exc: + await run._begin_callback_attempt( + callback_run_id="too-large", + purpose="main_agent", + route=route, + provider_input_bound_tokens=bound.payload_input_hard_cap + 1, + provider_input_breakdown={ + "text_input_bound_tokens": 123, + "media_input_bound_tokens": 456, + "media_blocks": 2, + "largest_media_bytes": 1024, + }, + ) + + assert exc.value.details == ( + { + "provider_input_bound_tokens": bound.payload_input_hard_cap + 1, + "payload_input_hard_cap": bound.payload_input_hard_cap, + "text_input_bound_tokens": 123, + "media_input_bound_tokens": 456, + "media_blocks": 2, + "largest_media_bytes": 1024, + }, + ) + + +@pytest.mark.asyncio +async def test_route_parameters_compile_once_during_prepare(monkeypatch, tmp_path): + runtime, authority = _runtime(tmp_path, monkeypatch) + original = EvoModelRuntime._compile_route + calls = [] + + def compile_once(route, purpose, bound, reasoning_effort): + calls.append((route.identity.route_key, purpose, bound.max_output_tokens)) + return original(route, purpose, bound, reasoning_effort) + + monkeypatch.setattr(EvoModelRuntime, "_compile_route", staticmethod(compile_once)) + agent_input = _input() + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), + agent_input, + WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink()), + ) + prepare_call_count = len(calls) + + assert prepare_call_count == 4 + await runtime.start_web_run(_admission(authority, quote)) + assert len(calls) == prepare_call_count + + +@pytest.mark.asyncio +async def test_prepare_freezes_an_immutable_invocation_plan(monkeypatch, tmp_path): + runtime, authority = _runtime(tmp_path, monkeypatch) + agent_input = _input() + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), + agent_input, + WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink()), + ) + + route = runtime._prepared[quote.preparation_id].snapshot.main_routes[0] + plan = route.invocation_plan + + assert plan is not None + assert plan.api_mode == "chat_completions" + assert plan.output_token_parameter == "max_tokens" + assert plan.output_token_limit == 512 + assert plan.tool_call_transport == "disabled" + assert plan.sdk_params["use_responses_api"] is False + with pytest.raises(TypeError): + plan.sdk_params["max_tokens"] = 1 + + +@pytest.mark.asyncio +async def test_boolean_reasoning_adapter_maps_to_provider_extra_body( + monkeypatch, tmp_path +): + monkeypatch.setenv("WEB_RUNTIME_TEST_KEY", "test-secret") + authority = HmacGrantAuthority(RUNTIME_SECRET, RUNTIME_KEY_ID) + payload = v3_payload() + payload["providers"]["custom-openai"]["models"][0]["reasoning"] = { + "mode": "boolean", + "enabled_params": {"extra_body": {"enable_thinking": True}}, + "disabled_params": {"extra_body": {"enable_thinking": False}}, + } + # Rebuild evidence after changing route semantics. + payload.pop("capability_evidence") + from tests.v3_fixtures import v3_payload as build_payload + + rebuilt = build_payload() + rebuilt["providers"]["custom-openai"]["models"][0]["reasoning"] = payload[ + "providers" + ]["custom-openai"]["models"][0]["reasoning"] + from EvoScientist.llm.model_config import ( + EvoModelConfig, + endpoint_fingerprint, + route_semantics_hash, + ) + + candidate = EvoModelConfig.parse(rebuilt, require_evidence=False) + ring = identity_ring() + semantics_key = ring.derive_current("ai4sci/route-semantics-hash/v3")[1] + endpoint_key = ring.derive_current("ai4sci/endpoint-fingerprint/v3")[1] + route = candidate.concrete_routes("visible-main")[0] + rebuilt["capability_evidence"][0]["probe"]["route_semantics_hash"] = ( + route_semantics_hash(candidate, route, semantics_key) + ) + rebuilt["capability_evidence"][0]["probe"]["endpoint_fingerprint"] = ( + endpoint_fingerprint(candidate, route, endpoint_key) + ) + + store = FileEvoModelConfigStore(tmp_path / "routes.yaml", admin_verifier=authority) + store.bootstrap_for_development(rebuilt) + calls = [] + runtime = EvoModelRuntime( + store, + admission_verifier=authority, + quote_authority=authority, + identity_key_ring=identity_ring(), + model_factory=lambda **kwargs: calls.append(kwargs) or _Model(), + agent_factory=lambda *_args: _Agent(), + ) + agent_input = _input() + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input, reasoning_effort="high"), + agent_input, + WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink()), + ) + await runtime.start_web_run(_admission(authority, quote)) + + main = next(item for item in calls if item["max_tokens"] == 512) + assert main["extra_body"]["enable_thinking"] is True + assert "reasoning_effort" not in main + assert all( + item["extra_body"]["enable_thinking"] is False + for item in calls + if item["max_tokens"] != 512 + ) + + +@pytest.mark.asyncio +async def test_prepare_is_idempotent_by_subject_request_and_turn(monkeypatch, tmp_path): + runtime, authority = _runtime(tmp_path, monkeypatch) + agent_input = _input() + host = WebHostContext( + "/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink() + ) + + first = await runtime.prepare_model_run( + _preparation(authority, agent_input), agent_input, host + ) + replay = await runtime.prepare_model_run( + _preparation(authority, agent_input), agent_input, host + ) + + assert replay == first + assert len(runtime._prepared) == 1 + + +@pytest.mark.asyncio +async def test_prepare_rejects_semantic_conflict_for_same_request( + monkeypatch, tmp_path +): + runtime, authority = _runtime(tmp_path, monkeypatch) + agent_input = _input() + host = WebHostContext( + "/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink() + ) + await runtime.prepare_model_run( + _preparation(authority, agent_input), agent_input, host + ) + + with pytest.raises(EvoRuntimeError, match="PREPARATION_CONFLICT"): + await runtime.prepare_model_run( + _preparation(authority, agent_input, reasoning_effort="high"), + agent_input, + host, + ) + + +@pytest.mark.asyncio +async def test_cancelled_prepare_leaves_bounded_idempotency_tombstone( + monkeypatch, tmp_path +): + runtime, authority = _runtime(tmp_path, monkeypatch) + agent_input = _input() + host = WebHostContext( + "/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink() + ) + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), agent_input, host + ) + assert await runtime.cancel_prepared_run(quote.preparation_id, reason="test") + + with pytest.raises(EvoRuntimeError, match="PREPARATION_STALE"): + await runtime.prepare_model_run( + _preparation(authority, agent_input), agent_input, host + ) + assert not runtime._prepared + assert len(runtime._prepared_tombstones) == 1 + + +@pytest.mark.asyncio +async def test_tampered_admission_is_rejected_before_agent_creation( + monkeypatch, tmp_path +): + runtime, authority = _runtime(tmp_path, monkeypatch) + agent_input = _input() + host = WebHostContext( + "/tmp", "/tmp", object(), object(), runtime_event_sink=_Sink() + ) + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), agent_input, host + ) + admission = _admission(authority, quote) + tampered = admission.__class__( + **{ + **admission.unsigned_payload(), + "plan": "enterprise", + "signature": admission.signature, + } + ) + with pytest.raises(EvoRuntimeError, match="CONTRACT_SIGNATURE_INVALID"): + await runtime.start_web_run(tampered) + + +@pytest.mark.asyncio +async def test_start_rejects_changed_tool_registry_revision(monkeypatch, tmp_path): + runtime, authority = _runtime(tmp_path, monkeypatch) + agent_input = _input() + revision = {"value": "registry-v1"} + + def registry_provider(): + return (), revision["value"] + + host = WebHostContext( + "/tmp", + "/tmp", + object(), + object(), + runtime_event_sink=_Sink(), + tool_registry_provider=registry_provider, + ) + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), agent_input, host + ) + revision["value"] = "registry-v2" + + with pytest.raises(EvoRuntimeError, match="TOOL_REGISTRY_STALE"): + await runtime.start_web_run(_admission(authority, quote)) + + +@pytest.mark.asyncio +async def test_started_event_is_committed_before_provider_boundary( + monkeypatch, tmp_path +): + runtime, authority = _runtime(tmp_path, monkeypatch) + sink = _Sink() + agent_input = _input() + host = WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=sink) + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), agent_input, host + ) + run = await runtime.start_web_run(_admission(authority, quote)) + route = run._snapshot.purpose_routes["main_agent"][0] + await run._begin_callback_attempt( + callback_run_id="callback", + purpose="main_agent", + route=route, + provider_input_bound_tokens=10, + ) + assert sink.events[0].payload["outcome"] == "started" + + +@pytest.mark.asyncio +async def test_normal_agent_model_rounds_are_not_limited_by_quote_attempt_counts( + monkeypatch, tmp_path +): + runtime, authority = _runtime(tmp_path, monkeypatch) + sink = _Sink() + agent_input = _input() + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), + agent_input, + WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=sink), + ) + run = await runtime.start_web_run(_admission(authority, quote)) + route = run._snapshot.purpose_routes["main_agent"][0] + configured_estimate = quote.purpose_attempt_limits["main_agent"] + + for index in range(configured_estimate + 2): + callback_run_id = f"callback-{index}" + await run._begin_callback_attempt( + callback_run_id=callback_run_id, + purpose="main_agent", + route=route, + provider_input_bound_tokens=10, + ) + await run._finish_callback_attempt( + callback_run_id, + usage={"input_tokens": 1, "cached_input_tokens": 0, "output_tokens": 1}, + error_code=None, + ) + + assert run._attempt_counts["main_agent"] == configured_estimate + 2 + starts = [ + event.payload + for event in sink.events + if event.kind == "model_attempt" and event.payload["outcome"] == "started" + ] + assert [event["attempt_index"] for event in starts] == list( + range(1, configured_estimate + 3) + ) + assert all(event["provider_reserved_microunits"] == 0 for event in starts) + + +def test_graph_recursion_error_has_a_stable_runtime_code() -> None: + GraphRecursionError = type("GraphRecursionError", (Exception,), {}) + + assert _safe_error_code(GraphRecursionError()) == "AGENT_RECURSION_LIMIT_EXCEEDED" + + +def test_agent_execution_fallback_is_not_misattributed_to_provider() -> None: + error = TypeError("sensitive graph detail") + code = _safe_error_code(error, fallback="AGENT_RUNTIME_ERROR") + + assert code == "AGENT_RUNTIME_ERROR" + assert _safe_error_code(error) == "MODEL_PROVIDER_ERROR" + assert _run_failure_details(error, code) == { + "failure_stage": "agent_execution", + "reason": "unclassified_agent_exception", + "agent_error_type": "TypeError", + "agent_error_module": "builtins", + } + assert "sensitive graph detail" not in str(_run_failure_details(error, code)) + + +@pytest.mark.asyncio +async def test_callback_accepts_langchain_messages_for_token_bound( + monkeypatch, tmp_path, caplog +): + runtime, authority = _runtime(tmp_path, monkeypatch) + sink = _Sink() + agent_input = _input() + host = WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=sink) + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), agent_input, host + ) + run = await runtime.start_web_run(_admission(authority, quote)) + route = run._snapshot.purpose_routes["main_agent"][0] + + with caplog.at_level("INFO", logger="EvoScientist.llm.runtime"): + await run._attempt_callback.on_chat_model_start( + {}, + [[HumanMessage(content="hello")]], + run_id="callback", + metadata={ + "runtime_purpose": "main_agent", + "route_key": route.identity.route_key, + }, + ) + await run._attempt_callback.on_llm_new_token( + "secret streamed text", + chunk=SimpleNamespace( + message=SimpleNamespace( + content="secret streamed text", + additional_kwargs={}, + tool_call_chunks=[], + ) + ), + run_id="callback", + ) + + assert sink.events[0].payload["provider_input_bound_tokens"] > 0 + assert run._attempt_callback.raise_error is True + assert "phase=request_started" in caplog.text + assert "phase=first_visible_chunk" in caplog.text + assert "visible_chunks=1" in caplog.text + assert "text_chars=20" in caplog.text + assert "secret streamed text" not in caplog.text + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("confirmation", "expected_error"), + [ + ("absent", "EVENT_INGRESS_UNAVAILABLE"), + ("conflict", "EVO_EVENT_CONFLICT"), + ("unknown", "EVENT_COMMIT_INDETERMINATE"), + ], +) +async def test_uncertain_event_commit_is_confirmed_before_dispatch( + monkeypatch, tmp_path, confirmation, expected_error +): + runtime, authority = _runtime(tmp_path, monkeypatch) + sink = _UncertainSink(confirmation) + agent_input = _input() + host = WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=sink) + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), agent_input, host + ) + run = await runtime.start_web_run(_admission(authority, quote)) + route = run._snapshot.purpose_routes["main_agent"][0] + + with pytest.raises(EvoRuntimeError, match=expected_error): + await run._begin_callback_attempt( + callback_run_id="callback", + purpose="main_agent", + route=route, + provider_input_bound_tokens=10, + ) + assert len(sink.confirmed) == 1 + assert run._sequence == 0 + + +@pytest.mark.asyncio +async def test_uncertain_but_committed_event_advances_once(monkeypatch, tmp_path): + runtime, authority = _runtime(tmp_path, monkeypatch) + sink = _UncertainSink("committed") + agent_input = _input() + host = WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=sink) + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), agent_input, host + ) + run = await runtime.start_web_run(_admission(authority, quote)) + route = run._snapshot.purpose_routes["main_agent"][0] + + await run._begin_callback_attempt( + callback_run_id="callback", + purpose="main_agent", + route=route, + provider_input_bound_tokens=10, + ) + + assert run._sequence == 1 + assert len(sink.confirmed) == 1 + + +@pytest.mark.asyncio +async def test_stream_raises_when_terminal_commit_fails(monkeypatch, tmp_path): + runtime, authority = _runtime(tmp_path, monkeypatch) + sink = _TerminalFailingSink() + agent_input = _input() + quote = await runtime.prepare_model_run( + _preparation(authority, agent_input), + agent_input, + WebHostContext("/tmp", "/tmp", object(), object(), runtime_event_sink=sink), + ) + run = await runtime.start_web_run(_admission(authority, quote)) + + async def consume() -> None: + async for _event in run.stream(): + pass + + with pytest.raises(EvoRuntimeError, match="EVO_EVENT_CONFLICT"): + async with asyncio.timeout(1): + await consume() diff --git a/tests/test_web_tool_registry.py b/tests/test_web_tool_registry.py new file mode 100644 index 0000000..76ec011 --- /dev/null +++ b/tests/test_web_tool_registry.py @@ -0,0 +1,17 @@ +from __future__ import annotations + +import pytest + +from EvoScientist.llm.contracts import EvoRuntimeError +from EvoScientist.web_runtime import _ToolRegistryFenceMiddleware + + +def test_tool_dispatch_fence_rejects_changed_registry(monkeypatch): + monkeypatch.setattr( + "EvoScientist.web_runtime.web_tool_registry_manifest", + lambda: ((), "registry-v2"), + ) + middleware = _ToolRegistryFenceMiddleware("registry-v1") + + with pytest.raises(EvoRuntimeError, match="TOOL_REGISTRY_STALE"): + middleware._require_current() diff --git a/tests/v3_fixtures.py b/tests/v3_fixtures.py new file mode 100644 index 0000000..7cbe061 --- /dev/null +++ b/tests/v3_fixtures.py @@ -0,0 +1,161 @@ +from __future__ import annotations + +from typing import Any + +from EvoScientist.llm.crypto import HmacKeyRing, KeyMaterial +from EvoScientist.llm.model_config import ( + EvoModelConfig, + adapter_revision, + endpoint_fingerprint, + route_semantics_hash, +) + +RUNTIME_SECRET = "runtime-secret-for-tests-000000000000000000000000" +IDENTITY_SECRET = "identity-secret-for-tests-00000000000000000000000" +RUNTIME_KEY_ID = "runtime-test" +IDENTITY_KEY_ID = "identity-test" + + +def identity_ring() -> HmacKeyRing: + return HmacKeyRing(KeyMaterial.create(IDENTITY_KEY_ID, IDENTITY_SECRET)) + + +def v3_payload(*, revision: int = 1) -> dict[str, Any]: + payload: dict[str, Any] = { + "schema_version": 2, + "config_revision": revision, + "config_identity_key_id": IDENTITY_KEY_ID, + "runtime_defaults": {"max_retries": 0}, + "purpose_defaults": { + "main_agent": {"reasoning_effort": "medium"}, + "tool_selector": {"reasoning_effort": "disabled"}, + "deepagents_summarizer": {"reasoning_effort": "disabled"}, + "title": {"reasoning_effort": "disabled"}, + }, + "providers": { + "custom-openai": { + "protocol": "custom-openai", + "params": {}, + "endpoints": [ + { + "name": "primary", + "base_url": "https://provider.example/v1", + "auth": {"ref": "env://WEB_RUNTIME_TEST_KEY", "revision": 1}, + "headers": {"User-Agent": "Ai4Sci-Test"}, + "header_refs": {}, + "params": {"extra_body": {}}, + } + ], + "models": [ + { + "id": "model-id", + "params": {"output_token_limit": 512}, + "supports_vision": True, + "supports_reasoning": True, + "allowed_reasoning_efforts": ["low", "medium", "high"], + "context_window": 8192, + "max_output_tokens": 2048, + "access": {"allowed_plans": [], "allowed_roles": []}, + "billing": { + "sku": "visible-model", + "pricing_revision": "test-1", + "currency": "CNY", + "unit_scale": 1_000_000, + "input_microunits_per_million": 1_000_000, + "output_microunits_per_million": 2_000_000, + "cached_microunits_per_million": 500_000, + }, + } + ], + } + }, + "endpoint_pools": { + "default": { + "provider": "custom-openai", + "strategy": "smooth_weighted_round_robin", + "endpoints": [{"name": "primary", "weight": 1}], + } + }, + "route_health": { + "failure_threshold": 3, + "cooldown_seconds": 30, + "half_open_max_inflight": 1, + "counted_error_codes": ["PROVIDER_5XX", "PROVIDER_TIMEOUT"], + "open_immediately_error_codes": ["PROVIDER_AUTH_INVALID"], + }, + "route_selectors": { + "visible-main": { + "provider": "custom-openai", + "endpoint_pool": "default", + "model": "model-id", + "api_mode": "chat_completions", + "tool_call_transport": "non_streaming", + } + }, + "purpose_routes": { + "main_agent": { + "default_alias": "visible-model", + "selectable": {"visible-model": "visible-main"}, + }, + "title": {"default": "visible-main"}, + }, + "purpose_call_limits": { + "main_agent": {"max_attempts_per_run": 4}, + "tool_selector": {"max_attempts_per_run": 2}, + "deepagents_summarizer": {"max_attempts_per_run": 1}, + "title": {"max_attempts_per_run": 1}, + }, + "web_runtime": { + "title_start_timeout_seconds": 30, + "prepare_ttl_seconds": 30, + "turn_lease_grace_seconds": 10, + "active_run_timeout_seconds": 60, + "max_run_journal_events": 1000, + "max_run_journal_bytes": 1_048_576, + "max_prepared_runs_per_subject": 2, + "max_prepared_runs_total": 100, + }, + "capability_evidence": [], + "tool_protocol_fallbacks": [{"primary": "visible-main", "fallbacks": []}], + } + candidate = EvoModelConfig.parse(payload, require_evidence=False) + ring = identity_ring() + semantics_key = ring.derive_current("ai4sci/route-semantics-hash/v3")[1] + endpoint_key = ring.derive_current("ai4sci/endpoint-fingerprint/v3")[1] + evidence = [] + seen = set() + for selector_id in candidate.route_selectors: + for route in candidate.concrete_routes(selector_id): + if route.key() in seen: + continue + seen.add(route.key()) + evidence.append( + { + "route": { + "provider": route.provider, + "endpoint": route.endpoint, + "model": route.model, + "api_mode": route.api_mode, + "tool_call_transport": route.tool_call_transport, + }, + "connectivity": "supported", + "tool_capability": "supported", + "probe": { + "route_semantics_hash": route_semantics_hash( + candidate, route, semantics_key + ), + "endpoint_fingerprint": endpoint_fingerprint( + candidate, route, endpoint_key + ), + "config_identity_key_id": IDENTITY_KEY_ID, + "adapter_revision": adapter_revision( + candidate.providers[route.provider].protocol, + route.api_mode, + ), + "fixture_digest": "sha256:test-fixture-v3", + "verified_at": "2026-07-20T00:00:00+00:00", + }, + } + ) + payload["capability_evidence"] = evidence + return payload diff --git a/uv.lock b/uv.lock index 947e9fa..116eb3c 100644 --- a/uv.lock +++ b/uv.lock @@ -968,6 +968,7 @@ dependencies = [ { name = "python-dotenv" }, { name = "pyyaml" }, { name = "questionary" }, + { name = "rfc8785" }, { name = "rich" }, { name = "tavily-python" }, { name = "textual" }, @@ -1090,6 +1091,7 @@ requires-dist = [ { name = "qrcode", marker = "extra == 'qq'", specifier = ">=7.4" }, { name = "qrcode", marker = "extra == 'wechat'", specifier = ">=7.4" }, { name = "questionary", specifier = ">=2.1" }, + { name = "rfc8785", specifier = "==0.1.4" }, { name = "rich", specifier = ">=15.0" }, { name = "ruff", marker = "extra == 'dev'", specifier = ">=0.5" }, { name = "slack-sdk", marker = "extra == 'all-channels'", specifier = ">=3.27" }, @@ -3871,6 +3873,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/3f/51/d4db610ef29373b879047326cbf6fa98b6c1969d6f6dc423279de2b1be2c/requests_toolbelt-1.0.0-py2.py3-none-any.whl", hash = "sha256:cccfdd665f0a24fcf4726e690f65639d272bb0637b9b92dfd91a5568ccf6bd06", size = 54481, upload-time = "2023-05-01T04:11:28.427Z" }, ] +[[package]] +name = "rfc8785" +version = "0.1.4" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/ef/2f/fa1d2e740c490191b572d33dbca5daa180cb423c24396b856f5886371d8b/rfc8785-0.1.4.tar.gz", hash = "sha256:e545841329fe0eee4f6a3b44e7034343100c12b4ec566dc06ca9735681deb4da", size = 14321, upload-time = "2024-09-27T16:33:31.206Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/4d/78/119878110660b2ad709888c8a1614fce7e2fab39080ab960656dc8605bf6/rfc8785-0.1.4-py3-none-any.whl", hash = "sha256:520d690b448ecf0703691c76e1a34a24ddcd4fc5bc41d589cb7c58ec651bcd48", size = 9240, upload-time = "2024-09-27T16:33:29.683Z" }, +] + [[package]] name = "rich" version = "15.0.0"