diff --git a/.gitignore b/.gitignore index ab4bb6a..04c1361 100644 --- a/.gitignore +++ b/.gitignore @@ -49,3 +49,6 @@ conversation_history/ botpy.log large_tool_results/ runs/ + +# local runtime artifacts (scope tokens, control DBs) +.evoscientist/ diff --git a/EvoScientist/EvoScientist.py b/EvoScientist/EvoScientist.py index f928208..ad41a99 100644 --- a/EvoScientist/EvoScientist.py +++ b/EvoScientist/EvoScientist.py @@ -187,6 +187,31 @@ def _ensure_auxiliary_chat_model(): return _auxiliary_chat_model +def _compile_time_role_model(role: str = "primary"): + """Return the compile-time model binding for graph construction. + + Unlike ``_ensure_chat_model()`` this never raises on a bootstrap + registry: graphs must still materialize so the Config API can serve + (design doc section 10 — only run creation is forbidden in bootstrap), + so a ``RegistryNotReadyChatModel`` placeholder is bound instead. It + raises ``MODEL_REGISTRY_NOT_READY`` on the first model call; per-run + resolution still comes from the run snapshot via + ``ConfigurableModelMiddleware``. Nothing is cached in module globals — + once the registry becomes active, rebuilt graphs get the real model. + """ + from .model_registry.errors import MODEL_REGISTRY_NOT_READY, ModelRegistryError + from .model_registry.placeholder import RegistryNotReadyChatModel + from .model_registry.runtime import get_snapshot_runtime + + runtime = get_snapshot_runtime() + try: + return runtime.build_default_role_model(role) + except ModelRegistryError as exc: + if exc.code != MODEL_REGISTRY_NOT_READY: + raise + return RegistryNotReadyChatModel(detail=str(exc)) + + # ============================================================================= # MCP caching # ============================================================================= @@ -449,6 +474,7 @@ def _build_base_kwargs( from .utils import load_subagents cfg = cfg if cfg is not None else _ensure_config() + model = chat_model if chat_model is not None else _compile_time_role_model("primary") tool_registry = {"think_tool": think_tool} if os.environ.get("TAVILY_API_KEY"): tool_registry["tavily_search"] = tavily_search @@ -462,12 +488,12 @@ def _build_base_kwargs( ) _ensure_general_purpose_subagent(subs) _inject_subagent_middleware( - subs, workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model + subs, workspace_dir=workspace_dir, cfg=cfg, chat_model=model ) subs = _maybe_swap_async_subagents(subs, base_middleware, cfg=cfg) return { "name": "EvoScientist", - "model": chat_model if chat_model is not None else _ensure_chat_model(), + "model": model, "tools": list(base_tools), "backend": base_backend, "subagents": subs, @@ -503,13 +529,14 @@ def load_mcp_and_build_kwargs( from .utils import load_subagents cfg = cfg if cfg is not None else _ensure_config() + model = chat_model if chat_model is not None else _compile_time_role_model("primary") mcp_by_agent = _load_mcp_tools_cached(on_progress=on_mcp_progress) if not mcp_by_agent: return _build_base_kwargs( base_backend, base_middleware, cfg=cfg, - chat_model=chat_model, + chat_model=model, workspace_dir=workspace_dir, ) @@ -535,7 +562,7 @@ def load_mcp_and_build_kwargs( _ensure_general_purpose_subagent(subs) _inject_subagent_middleware( - subs, workspace_dir=workspace_dir, cfg=cfg, chat_model=chat_model + subs, workspace_dir=workspace_dir, cfg=cfg, chat_model=model ) # Inject MCP tools into subagents by name @@ -549,7 +576,7 @@ def load_mcp_and_build_kwargs( return { "name": "EvoScientist", - "model": chat_model if chat_model is not None else _ensure_chat_model(), + "model": model, "tools": base_tools + mcp_main, "backend": base_backend, "subagents": subs, @@ -686,7 +713,7 @@ def _get_default_middleware( ) cfg = cfg if cfg is not None else _ensure_config() - model = chat_model if chat_model is not None else _ensure_chat_model() + model = chat_model if chat_model is not None else _compile_time_role_model("primary") if backend is None: # Preserve the factory's pure path for callers that provide an # explicit model/configuration (notably tests and subagent assembly). @@ -726,7 +753,7 @@ def _get_default_middleware( if for_async_subagent: tool_selector_model = model elif chat_model is None: - tool_selector_model = _ensure_auxiliary_chat_model() + tool_selector_model = _compile_time_role_model("auxiliary") else: # Pure path (explicit model + config): the threaded model stands in # for tool selection at build time; per-call resolution still comes diff --git a/EvoScientist/memory/agents/_factory.py b/EvoScientist/memory/agents/_factory.py index 8373929..e0987cb 100644 --- a/EvoScientist/memory/agents/_factory.py +++ b/EvoScientist/memory/agents/_factory.py @@ -87,7 +87,7 @@ def build_memory_agent_graph( from deepagents import create_deep_agent from ...backends import build_memory_agent_backend - from ...EvoScientist import _ensure_auxiliary_chat_model + from ...EvoScientist import _compile_time_role_model kwargs: dict[str, Any] = {} if response_format is not None: @@ -101,7 +101,7 @@ def build_memory_agent_graph( agent = create_deep_agent( name=name, - model=_ensure_auxiliary_chat_model(), + model=_compile_time_role_model("auxiliary"), system_prompt=system_prompt, tools=list(tools), backend=backend, diff --git a/EvoScientist/middleware/configurable_model.py b/EvoScientist/middleware/configurable_model.py index f29aa0d..dda1a27 100644 --- a/EvoScientist/middleware/configurable_model.py +++ b/EvoScientist/middleware/configurable_model.py @@ -92,31 +92,28 @@ def check_no_outside_snapshot_model_config(configurable: Mapping[str, Any]) -> N def read_snapshot_binding( configurable: Mapping[str, Any], -) -> tuple[str, str | None, str | None] | None: - """Extract ``(snapshot_id, deployment_id, thread_id)`` from configurable. +) -> tuple[str, str | None] | None: + """Extract ``(snapshot_id, thread_id)`` from configurable. - Returns ``None`` when the run carries no ``runtime_snapshot_id``. The - deployment ID falls back to ``None`` (caller substitutes the platform's - local deployment ID); a missing thread ID fails closed later because it - can never match the snapshot's binding. + Returns ``None`` when the run carries no ``runtime_snapshot_id``. A + missing thread ID fails closed later because it can never match the + snapshot's binding. The deployment is deliberately NOT taken from + ``workspace_deployment_id``: that key names the workspace-isolation + scope, not the snapshot issuer — issuer verification happens against + the platform-registered deployment set in ``get_snapshot_for_run``. """ snapshot_id = configurable.get("runtime_snapshot_id") if not isinstance(snapshot_id, str) or not snapshot_id: return None - deployment_id = configurable.get("workspace_deployment_id") thread_id = configurable.get("thread_id") - return ( - snapshot_id, - deployment_id if isinstance(deployment_id, str) and deployment_id else None, - thread_id if isinstance(thread_id, str) else None, - ) + return (snapshot_id, thread_id if isinstance(thread_id, str) else None) def ensure_snapshot_binding( configurable: Mapping[str, Any], runtime: SnapshotRuntime, -) -> tuple[str, str, str]: - """Return ``(snapshot_id, deployment_id, thread_id)`` for the active run. +) -> tuple[str, str]: + """Return ``(snapshot_id, thread_id)`` for the active run. The run's explicit ``runtime_snapshot_id`` wins (section 8.2 binding). Without one, the run is a local entry that could not inject a snapshot @@ -136,8 +133,8 @@ def ensure_snapshot_binding( """ binding = read_snapshot_binding(configurable) if binding is not None: - snapshot_id, deployment_id, thread_id = binding - return (snapshot_id, deployment_id or runtime.local_deployment_id, thread_id or "") + snapshot_id, thread_id = binding + return (snapshot_id, thread_id or "") thread_id = configurable.get("thread_id") if not isinstance(thread_id, str) or not thread_id: raise ModelRegistryError( @@ -150,7 +147,7 @@ def ensure_snapshot_binding( snapshot = runtime.create_local_snapshot( thread_id, run_request_id=f"auto:{thread_id}" ) - return (snapshot.snapshot_id, snapshot.deployment_id, snapshot.thread_id) + return (snapshot.snapshot_id, snapshot.thread_id) class ConfigurableModelMiddleware(AgentMiddleware): @@ -230,14 +227,8 @@ class ConfigurableModelMiddleware(AgentMiddleware): configurable = _current_configurable() check_no_outside_snapshot_model_config(configurable) runtime = self._snapshot_runtime() - snapshot_id, deployment_id, thread_id = ensure_snapshot_binding( - configurable, runtime - ) - return runtime.get_snapshot( - snapshot_id, - deployment_id=deployment_id, - thread_id=thread_id, - ) + snapshot_id, thread_id = ensure_snapshot_binding(configurable, runtime) + return runtime.get_snapshot_for_run(snapshot_id, thread_id=thread_id) def _resolve(self) -> Any: """Return a cached or freshly-built chat model for the run's snapshot.""" diff --git a/EvoScientist/middleware/message_budget.py b/EvoScientist/middleware/message_budget.py index 1ef303a..06b43cc 100644 --- a/EvoScientist/middleware/message_budget.py +++ b/EvoScientist/middleware/message_budget.py @@ -41,6 +41,7 @@ tool/attachment modes), so the middleware does not re-check it per call. from __future__ import annotations +import asyncio import threading from collections.abc import Iterable, Mapping from contextvars import ContextVar @@ -195,6 +196,7 @@ class MessageBudgetMiddleware: def __init__(self) -> None: self._snapshot_model_cache: dict[str, Any] = {} + self._snapshot_cache: dict[str, RuntimeSnapshot] = {} self._model_cache_lock = threading.RLock() self._has_tools = has_tools self._snapshot_role = snapshot_role @@ -233,14 +235,23 @@ class MessageBudgetMiddleware: ``MODEL_REGISTRY_NOT_READY`` instead of a 32K fallback. """ runtime = self._snapshot_runtime() - snapshot_id, deployment_id, thread_id = ensure_snapshot_binding( + snapshot_id, thread_id = ensure_snapshot_binding( _current_configurable(), runtime ) - return runtime.get_snapshot( - snapshot_id, - deployment_id=deployment_id, - thread_id=thread_id, + with self._model_cache_lock: + cached = self._snapshot_cache.get(snapshot_id) + if cached is not None: + # Snapshots are immutable; the per-run read already + # verified the binding, and repeat reads must not hit + # SQLite from the event loop (e.g. the ``model`` + # property during async summarization). + return cached + snapshot = runtime.get_snapshot_for_run( + snapshot_id, thread_id=thread_id ) + with self._model_cache_lock: + self._snapshot_cache[snapshot_id] = snapshot + return snapshot @property def model(self) -> Any: # type: ignore[override] @@ -283,7 +294,10 @@ class MessageBudgetMiddleware: _ACTIVE_BUDGET.reset(token) async def awrap_model_call(self, request: Any, handler: Any) -> Any: - token = _ACTIVE_BUDGET.set(self._budget_for_request(request)) + # The snapshot read hits SQLite — blocking I/O that must stay + # off the event loop (langgraph dev's blockbuster rejects it). + budget = await asyncio.to_thread(self._budget_for_request, request) + token = _ACTIVE_BUDGET.set(budget) try: return await super().awrap_model_call(request, handler) finally: diff --git a/EvoScientist/model_registry/http_api.py b/EvoScientist/model_registry/http_api.py index 7adbd82..ec4b9d3 100644 --- a/EvoScientist/model_registry/http_api.py +++ b/EvoScientist/model_registry/http_api.py @@ -481,7 +481,7 @@ class ModelRegistryHttpApi: INVALID_REQUEST, "The request body must be valid JSON." ) from None - def _authenticate( + def _authenticate_sync( self, request: Request, *, required_scope: str, require_thread_id: bool ) -> tuple[ApiServices, ActorContext]: services = self._services() @@ -492,12 +492,25 @@ class ModelRegistryHttpApi: ) return services, actor + async def _authenticate( + self, request: Request, *, required_scope: str, require_thread_id: bool + ) -> tuple[ApiServices, ActorContext]: + # Services build (config.yaml read, store mkdir/chmod) and jti + # registration (SQLite) are blocking I/O — keep them off the event + # loop (langgraph dev's blockbuster rejects them there). + return await asyncio.to_thread( + self._authenticate_sync, + request, + required_scope=required_scope, + require_thread_id=require_thread_id, + ) + # --- Config API ---------------------------------------------------------- async def get_model_registry(self, request: Request) -> Response: request_id = uuid.uuid4().hex try: - services, _actor = self._authenticate( + services, _actor = await self._authenticate( request, required_scope="model_config:read", require_thread_id=False ) response = await asyncio.to_thread(self._registry_response, services) @@ -508,7 +521,7 @@ class ModelRegistryHttpApi: async def put_model_registry(self, request: Request) -> Response: request_id = uuid.uuid4().hex try: - services, _actor = self._authenticate( + services, _actor = await self._authenticate( request, required_scope="model_config:write", require_thread_id=False ) body = await self._body(request) @@ -522,7 +535,7 @@ class ModelRegistryHttpApi: async def put_credential(self, request: Request) -> Response: request_id = uuid.uuid4().hex try: - services, _actor = self._authenticate( + services, _actor = await self._authenticate( request, required_scope="model_config:write", require_thread_id=False ) body = await self._body(request) @@ -544,7 +557,7 @@ class ModelRegistryHttpApi: async def test_provider(self, request: Request) -> Response: request_id = uuid.uuid4().hex try: - services, _actor = self._authenticate( + services, _actor = await self._authenticate( request, required_scope="model_config:test", require_thread_id=False ) body = await self._body(request) @@ -558,7 +571,7 @@ class ModelRegistryHttpApi: async def get_selectable_models(self, request: Request) -> Response: request_id = uuid.uuid4().hex try: - services, _actor = self._authenticate( + services, _actor = await self._authenticate( request, required_scope="model:select", require_thread_id=False ) response = await asyncio.to_thread(self._selectable_models, services) @@ -571,7 +584,7 @@ class ModelRegistryHttpApi: async def create_snapshot(self, request: Request) -> Response: request_id = uuid.uuid4().hex try: - services, actor = self._authenticate( + services, actor = await self._authenticate( request, required_scope="run:create", require_thread_id=True ) body = await self._body(request) @@ -591,7 +604,7 @@ class ModelRegistryHttpApi: async def bind_snapshot(self, request: Request) -> Response: request_id = uuid.uuid4().hex try: - services, actor = self._authenticate( + services, actor = await self._authenticate( request, required_scope="run:create", require_thread_id=True ) body = await self._body(request) @@ -611,7 +624,7 @@ class ModelRegistryHttpApi: async def delete_snapshot(self, request: Request) -> Response: request_id = uuid.uuid4().hex try: - services, actor = self._authenticate( + services, actor = await self._authenticate( request, required_scope="run:create", require_thread_id=True ) await asyncio.to_thread( diff --git a/EvoScientist/model_registry/placeholder.py b/EvoScientist/model_registry/placeholder.py new file mode 100644 index 0000000..5fb9f35 --- /dev/null +++ b/EvoScientist/model_registry/placeholder.py @@ -0,0 +1,37 @@ +"""Compile-time chat-model placeholder for a bootstrap registry. + +Graph construction must tolerate a bootstrap registry — the Config API has +to serve so operators can configure the first model (design doc section +10), and only run creation is forbidden. Every graph therefore binds this +placeholder when no registry default exists; it supports the build-time +surface (``bind_tools``, profile inspection, middleware wiring) but raises +``MODEL_REGISTRY_NOT_READY`` on the first model call, so no run can ever +silently fall back to an implicit model. +""" + +from typing import Any + +from langchain_core.language_models.chat_models import BaseChatModel +from langchain_core.messages import BaseMessage +from langchain_core.outputs import ChatResult + +from .errors import MODEL_REGISTRY_NOT_READY, ModelRegistryError + + +class RegistryNotReadyChatModel(BaseChatModel): + """Placeholder that fails every call with MODEL_REGISTRY_NOT_READY.""" + + detail: str + + @property + def _llm_type(self) -> str: + return "evoscientist-registry-not-ready" + + def _generate( + self, + messages: list[BaseMessage], + stop: list[str] | None = None, + run_manager: Any = None, + **kwargs: Any, + ) -> ChatResult: + raise ModelRegistryError(MODEL_REGISTRY_NOT_READY, self.detail) diff --git a/EvoScientist/model_registry/runtime.py b/EvoScientist/model_registry/runtime.py index 7fedae9..c08da98 100644 --- a/EvoScientist/model_registry/runtime.py +++ b/EvoScientist/model_registry/runtime.py @@ -50,6 +50,7 @@ class SnapshotRuntime: *, endpoint_policy: EndpointPolicy | None = None, local_deployment_id: str = DEFAULT_LOCAL_DEPLOYMENT_ID, + allowed_deployment_ids: tuple[str, ...] | None = None, ) -> None: self._store = store self._resolver = ModelRegistryResolver(store) @@ -58,6 +59,14 @@ class SnapshotRuntime: endpoint_policy if endpoint_policy is not None else EndpointPolicy(()) ) self._local_deployment_id = local_deployment_id + # Deployment ids allowed to issue snapshots a run may consume: the + # local entry points plus every registered WebUI delegation + # deployment (design doc 7.3/8.2). Defaults to just the local id. + self._allowed_deployment_ids = frozenset( + allowed_deployment_ids + if allowed_deployment_ids is not None + else (local_deployment_id,) + ) @property def store(self) -> ModelRuntimeStore: @@ -75,6 +84,10 @@ class SnapshotRuntime: def local_deployment_id(self) -> str: return self._local_deployment_id + @property + def allowed_deployment_ids(self) -> frozenset[str]: + return self._allowed_deployment_ids + # --- snapshot-driven (per-call) resolution ------------------------------ def get_snapshot( @@ -85,6 +98,16 @@ class SnapshotRuntime: snapshot_id, deployment_id=deployment_id, thread_id=thread_id ) + def get_snapshot_for_run( + self, snapshot_id: str, *, thread_id: str + ) -> RuntimeSnapshot: + """Read a run's snapshot, verifying thread binding + issuer set.""" + return self._snapshots.get_for_run( + snapshot_id, + thread_id=thread_id, + allowed_deployment_ids=self._allowed_deployment_ids, + ) + def build_role_model( self, snapshot: RuntimeSnapshot, role: ModelRole ) -> BaseChatModel: @@ -203,19 +226,20 @@ _default_runtime_lock = threading.Lock() def _read_local_platform_fields() -> tuple[ - str, Path | None, tuple[DevelopmentEndpoint, ...] + str, Path | None, tuple[DevelopmentEndpoint, ...], tuple[str, ...] ]: """Read the runtime-relevant platform fields without the BFF auth gates. Unlike ``load_platform_security_config`` (which rightly refuses to serve the Config API without a BFF token and delegation keys), the runtime - link only needs ``local_deployment_id``, ``model_runtime_db``, and - ``development_endpoints`` — all safe to read with the documented + link only needs ``local_deployment_id``, ``model_runtime_db``, + ``development_endpoints``, and the deployment ids of + ``webui_delegation_public_keys`` — all safe to read with the documented defaults when ``config.yaml`` is absent. """ path = get_config_path() if not path.exists(): - return DEFAULT_LOCAL_DEPLOYMENT_ID, None, () + return DEFAULT_LOCAL_DEPLOYMENT_ID, None, (), () try: with open(path, encoding="utf-8") as handle: data = yaml.safe_load(handle) or {} @@ -251,15 +275,37 @@ def _read_local_platform_fields() -> tuple[ raise PlatformConfigError( f"Invalid development_endpoints entry: {exc}." ) from exc + raw_keys = data.get("webui_delegation_public_keys") + webui_deployment_ids: tuple[str, ...] = () + if raw_keys is not None: + if not isinstance(raw_keys, list): + raise PlatformConfigError( + "config.yaml field 'webui_delegation_public_keys' must be a list." + ) + ids: list[str] = [] + for index, entry in enumerate(raw_keys): + if not isinstance(entry, dict): + raise PlatformConfigError( + f"webui_delegation_public_keys[{index}] must be an object." + ) + deployment = entry.get("deployment_id") + if not isinstance(deployment, str) or not deployment.strip(): + raise PlatformConfigError( + f"webui_delegation_public_keys[{index}] requires a non-empty " + "'deployment_id'." + ) + ids.append(deployment.strip()) + webui_deployment_ids = tuple(ids) return ( deployment_id.strip() if deployment_id else DEFAULT_LOCAL_DEPLOYMENT_ID, Path(database).expanduser() if database else None, endpoints, + webui_deployment_ids, ) def _build_default_runtime() -> SnapshotRuntime: - deployment_id, database_path, endpoints = _read_local_platform_fields() + deployment_id, database_path, endpoints, webui_ids = _read_local_platform_fields() store = ( ModelRuntimeStore(database_path=database_path) if database_path is not None @@ -269,6 +315,7 @@ def _build_default_runtime() -> SnapshotRuntime: store, endpoint_policy=EndpointPolicy(endpoints), local_deployment_id=deployment_id, + allowed_deployment_ids=(deployment_id, *webui_ids), ) diff --git a/EvoScientist/model_registry/snapshots.py b/EvoScientist/model_registry/snapshots.py index a0b67d2..5f94c08 100644 --- a/EvoScientist/model_registry/snapshots.py +++ b/EvoScientist/model_registry/snapshots.py @@ -361,6 +361,37 @@ class SnapshotService: self._ensure_frozen_specs_available(snapshot) return snapshot + def get_for_run( + self, + snapshot_id: str, + *, + thread_id: str, + allowed_deployment_ids: frozenset[str], + ) -> RuntimeSnapshot: + """Read a snapshot for an executing run (section 8.2). + + A run cannot self-assert its issuing deployment: the LangGraph + ``configurable`` carries ``workspace_deployment_id``, which names + the workspace-isolation scope, not the snapshot issuer. The binding + is therefore verified as ``thread_id`` equality plus membership of + the snapshot's ``deployment_id`` in the platform-registered set + (``local_deployment_id`` and every ``webui_delegation_public_keys`` + entry), so a snapshot can never be reused by another thread or by + an unregistered deployment. + """ + row = self._store.get_run_snapshot(snapshot_id) + if ( + row is None + or row["thread_id"] != thread_id + or row["deployment_id"] not in allowed_deployment_ids + ): + raise _snapshot_not_found() + if row["status"] in _TERMINAL_STATUSES: + raise _snapshot_expired() + snapshot = RuntimeSnapshot.from_row(row) + self._ensure_frozen_specs_available(snapshot) + return snapshot + def cleanup_expired(self, now: int | None = None) -> list[str]: """Mark due snapshots ``expired``; returns the transitioned IDs.""" return self._store.expire_due_run_snapshots( diff --git a/EvoScientist/subagents/_factory.py b/EvoScientist/subagents/_factory.py index 797add2..34c5b47 100644 --- a/EvoScientist/subagents/_factory.py +++ b/EvoScientist/subagents/_factory.py @@ -114,32 +114,11 @@ def build_async_subagent_graph(name: str) -> Any: # placeholder that raises MODEL_REGISTRY_NOT_READY on the first model # call. Every run fails with that clear structured error until a # primary model is configured and enabled; no 32K/implicit fallback. - from langchain_core.language_models.chat_models import BaseChatModel - from langchain_core.messages import BaseMessage - from langchain_core.outputs import ChatResult - from EvoScientist.model_registry.errors import ( MODEL_REGISTRY_NOT_READY, ModelRegistryError, ) - - class _RegistryNotReadyChatModel(BaseChatModel): - """Placeholder that fails every call with MODEL_REGISTRY_NOT_READY.""" - - detail: str - - @property - def _llm_type(self) -> str: - return "evoscientist-registry-not-ready" - - def _generate( - self, - messages: list[BaseMessage], - stop: list[str] | None = None, - run_manager: Any = None, - **kwargs: Any, - ) -> ChatResult: - raise ModelRegistryError(MODEL_REGISTRY_NOT_READY, self.detail) + from EvoScientist.model_registry.placeholder import RegistryNotReadyChatModel runtime = get_snapshot_runtime() snapshot_role = "auxiliary" if name == "scheduler" else "primary" @@ -148,7 +127,7 @@ def build_async_subagent_graph(name: str) -> Any: except ModelRegistryError as exc: if exc.code != MODEL_REGISTRY_NOT_READY: raise - model = _RegistryNotReadyChatModel(detail=str(exc)) + model = RegistryNotReadyChatModel(detail=str(exc)) subagents = [] _ensure_general_purpose_subagent(subagents) diff --git a/tests/test_async_subagent_factory.py b/tests/test_async_subagent_factory.py index 6b328b3..793f7be 100644 --- a/tests/test_async_subagent_factory.py +++ b/tests/test_async_subagent_factory.py @@ -222,7 +222,7 @@ def test_bootstrap_registry_binds_not_ready_placeholder( build_async_subagent_graph("writing-agent") # must not raise model = mock_create.call_args.kwargs["model"] - assert type(model).__name__ == "_RegistryNotReadyChatModel" + assert type(model).__name__ == "RegistryNotReadyChatModel" with pytest.raises(ModelRegistryError) as excinfo: model.invoke([]) assert excinfo.value.code == MODEL_REGISTRY_NOT_READY diff --git a/tests/test_auxiliary_model.py b/tests/test_auxiliary_model.py index 77068e9..19ab866 100644 --- a/tests/test_auxiliary_model.py +++ b/tests/test_auxiliary_model.py @@ -112,8 +112,13 @@ class TestAuxiliaryMiddlewareScope: main_model, aux_model = object(), object() with ( patch.object(E, "_ensure_config", return_value=_mock_cfg()), - patch.object(E, "_ensure_chat_model", return_value=main_model), - patch.object(E, "_ensure_auxiliary_chat_model", return_value=aux_model), + patch.object( + E, + "_compile_time_role_model", + side_effect=lambda role="primary": aux_model + if role == "auxiliary" + else main_model, + ), patch( "EvoScientist.middleware.create_tool_selector_middleware", side_effect=fake_ts, @@ -133,8 +138,13 @@ class TestAuxiliaryMiddlewareScope: main_model, aux_model = object(), object() with ( patch.object(E, "_ensure_config", return_value=_mock_cfg()), - patch.object(E, "_ensure_chat_model", return_value=main_model), - patch.object(E, "_ensure_auxiliary_chat_model", return_value=aux_model), + patch.object( + E, + "_compile_time_role_model", + side_effect=lambda role="primary": aux_model + if role == "auxiliary" + else main_model, + ), patch( "EvoScientist.middleware.create_tool_selector_middleware", side_effect=fake_ts, @@ -197,7 +207,7 @@ class TestAuxiliaryMiddlewareScope: def test_memory_agent_factory_uses_auxiliary_role(monkeypatch): """Memory workers bind the auxiliary role model at graph build.""" sentinel = object() - monkeypatch.setattr(E, "_ensure_auxiliary_chat_model", lambda: sentinel) + monkeypatch.setattr(E, "_compile_time_role_model", lambda role="primary": sentinel) captured: dict[str, object] = {} def fake_create_deep_agent(**kwargs): diff --git a/tests/test_configurable_model_middleware.py b/tests/test_configurable_model_middleware.py index 2d8335c..94810c1 100644 --- a/tests/test_configurable_model_middleware.py +++ b/tests/test_configurable_model_middleware.py @@ -159,6 +159,8 @@ class TestOutsideSnapshotConfig: class TestReadSnapshotBinding: def test_full_binding(self): + # ``workspace_deployment_id`` names the workspace-isolation scope and + # must NOT be picked up as the snapshot issuer (section 8.2). binding = read_snapshot_binding( { "runtime_snapshot_id": "snap-1", @@ -166,18 +168,17 @@ class TestReadSnapshotBinding: "thread_id": "thread-1", } ) - assert binding == ("snap-1", "deploy-1", "thread-1") + assert binding == ("snap-1", "thread-1") def test_missing_snapshot_id_returns_none(self): assert read_snapshot_binding({}) is None assert read_snapshot_binding({"runtime_snapshot_id": ""}) is None assert read_snapshot_binding({"runtime_snapshot_id": 42}) is None - def test_missing_deployment_and_thread_fall_back(self): + def test_missing_thread_falls_back(self): assert read_snapshot_binding({"runtime_snapshot_id": "snap-1"}) == ( "snap-1", None, - None, ) @@ -269,7 +270,7 @@ class TestLazyLocalSnapshot: first = ensure_snapshot_binding({"thread_id": "t-1"}, runtime) second = ensure_snapshot_binding({"thread_id": "t-1"}, runtime) assert first == second - assert first[1] == runtime.local_deployment_id + assert first[1] == "t-1" def test_explicit_snapshot_wins_over_lazy_creation(self, store, runtime): snapshot = make_snapshot(store) @@ -279,6 +280,37 @@ class TestLazyLocalSnapshot: ) assert binding[0] == snapshot.snapshot_id + def test_webui_snapshot_survives_scope_deployment_mismatch(self, store): + """BFF regression: the run's ``workspace_deployment_id`` names the + workspace-isolation scope, not the snapshot issuer. A snapshot issued + by a registered WebUI deployment must resolve even when the two ids + differ (the mismatch previously failed with SNAPSHOT_NOT_FOUND).""" + runtime = SnapshotRuntime( + store, allowed_deployment_ids=("local", "webui-local") + ) + snapshot = make_snapshot(store, deployment_id="webui-local") + mw = ConfigurableModelMiddleware(runtime=runtime) + req = _make_request() + handler = MagicMock(return_value="response") + configurable = _configurable_for( + snapshot, workspace_deployment_id="61d1b61b-078b-4bda-84c3-6c68a4afcd38" + ) + with _patched_config(configurable): + assert mw.wrap_model_call(req, handler) == "response" + assert handler.call_args[0][0].model is not req.model + + def test_unregistered_deployment_snapshot_rejected(self, store, runtime): + """A snapshot from an unregistered deployment fails closed.""" + snapshot = make_snapshot(store, deployment_id="rogue-webui") + mw = ConfigurableModelMiddleware(runtime=runtime) + req = _make_request() + with ( + _patched_config(_configurable_for(snapshot)), + pytest.raises(ModelRegistryError) as excinfo, + ): + mw.wrap_model_call(req, MagicMock()) + assert excinfo.value.code == SNAPSHOT_NOT_FOUND + # ============================================================================= # 4. Snapshot-driven model construction (full chain) @@ -368,19 +400,20 @@ class TestSnapshotDrivenConstruction: mw.wrap_model_call(req, MagicMock()) assert excinfo.value.code == SNAPSHOT_NOT_FOUND - def test_wrong_deployment_binding_fails_closed(self, store, runtime): + def test_scope_deployment_id_is_not_the_snapshot_issuer(self, store, runtime): + """``workspace_deployment_id`` names the workspace-isolation scope, so + any value — matching or not — is irrelevant to the snapshot binding; + the issuer check is the registered-deployment membership (see + test_unregistered_deployment_snapshot_rejected).""" mw = ConfigurableModelMiddleware(runtime=runtime) snapshot = make_snapshot(store) req = _make_request() + handler = MagicMock(return_value="response") - with ( - _patched_config( - _configurable_for(snapshot, workspace_deployment_id="other-deploy") - ), - pytest.raises(ModelRegistryError) as excinfo, + with _patched_config( + _configurable_for(snapshot, workspace_deployment_id="other-deploy") ): - mw.wrap_model_call(req, MagicMock()) - assert excinfo.value.code == SNAPSHOT_NOT_FOUND + assert mw.wrap_model_call(req, handler) == "response" def test_expired_snapshot_fails_loudly(self, store, runtime): mw = ConfigurableModelMiddleware(runtime=runtime) diff --git a/tests/test_snapshot_runtime.py b/tests/test_snapshot_runtime.py index 89efd7b..29a048a 100644 --- a/tests/test_snapshot_runtime.py +++ b/tests/test_snapshot_runtime.py @@ -166,19 +166,24 @@ class TestLocalPlatformFields: return _read_local_platform_fields() def test_missing_config_uses_defaults(self, tmp_path, monkeypatch): - deployment_id, db_path, endpoints = self._read(tmp_path, monkeypatch, None) + deployment_id, db_path, endpoints, webui_ids = self._read( + tmp_path, monkeypatch, None + ) assert deployment_id == DEFAULT_LOCAL_DEPLOYMENT_ID assert db_path is None assert endpoints == () + assert webui_ids == () def test_reads_runtime_fields(self, tmp_path, monkeypatch): - deployment_id, db_path, endpoints = self._read( + deployment_id, db_path, endpoints, webui_ids = self._read( tmp_path, monkeypatch, "local_deployment_id: dev-deploy\n" "model_runtime_db: /tmp/mr.sqlite3\n" "development_endpoints:\n" - " - {id: ollama, url: 'http://localhost:11434', label: Local}\n", + " - {id: ollama, url: 'http://localhost:11434', label: Local}\n" + "webui_delegation_public_keys:\n" + " - {deployment_id: webui-local, public_key: '-----BEGIN PUBLIC KEY-----\\nMFkwEwYHKoZIzj0CAQYIKoZIzj0DAQcDQgAE\\n-----END PUBLIC KEY-----\\n'}\n", ) assert deployment_id == "dev-deploy" assert db_path == __import__("pathlib").Path("/tmp/mr.sqlite3") @@ -187,6 +192,7 @@ class TestLocalPlatformFields: id="ollama", url="http://localhost:11434", label="Local" ), ) + assert webui_ids == ("webui-local",) def test_invalid_yaml_raises(self, tmp_path, monkeypatch): with pytest.raises(PlatformConfigError):