fix(code): persist resume model state privately (#4400)

Deep Agents Code now records resume model state privately after
successful model calls, including model params, without exposing the
internal `effective_model` runtime context field in LangSmith traces.

```dcode -r 019f1a63-1fca-70a3-90bb-e94458e61f29```

---

Model resume bookkeeping now stays in private checkpoint state instead of being carried through runtime context as `effective_model`. The model middleware records the actual model request that completed, which also avoids persisting a requested override when server-side resolution falls back to the previous model.

## Changes
- Remove `effective_model` from `CLIContext` so it no longer appears as LangSmith trace metadata.
- Move `_model_spec` persistence into `ConfigurableModelMiddleware`, emitted after a successful model response using `ExtendedModelResponse` and a private state update.
- Persist `_model_params` alongside `_model_spec` so resumed sessions restore invocation params such as `temperature` or `max_tokens`.
- Keep subagent model middleware from writing parent-thread resume metadata.
- Leave `ResumeStateMiddleware` focused on `_context_tokens`, with model metadata owned by the middleware that resolves the actual request.
This commit is contained in:
Mason Daugherty
2026-07-01 18:13:17 -04:00
committed by GitHub
parent 434f29e5cb
commit bbd0f0d1dd
9 changed files with 297 additions and 137 deletions
-13
View File
@@ -35,8 +35,6 @@ class CLIContextSchema:
model_params: dict[str, Any] = field(default_factory=dict)
effective_model: str | None = None
auto_approve: bool = False
approval_mode_key: str | None = None
@@ -63,17 +61,6 @@ class CLIContext(TypedDict, total=False):
"""Invocation params (e.g. `temperature`, `max_tokens`) to merge
into `model_settings`."""
effective_model: str | None
"""Resolved `provider:model` spec actually in use for this invocation.
Unlike `model` (a swap *instruction* that is `None` when the base model is
used), this carries the model in effect — whether from a `/model` override
or the startup default — so `ResumeStateMiddleware` can record it to
checkpoint state for restore on resume. `None` when no usable spec is
resolved yet (e.g. credentials not configured), in which case nothing is
recorded rather than a malformed spec.
"""
auto_approve: bool
"""Whether gated tool calls should skip the human-approval interrupt.
+5 -5
View File
@@ -1496,7 +1496,7 @@ def create_cli_agent(
def _subagent_cli_middleware(*, has_explicit_model: bool) -> list[AgentMiddleware]:
middleware: list[AgentMiddleware] = []
if not has_explicit_model:
middleware.append(ConfigurableModelMiddleware())
middleware.append(ConfigurableModelMiddleware(persist_model_state=False))
if restrictive_shell_allow_list is not None:
middleware.append(ShellAllowListMiddleware(restrictive_shell_allow_list))
# Subagents share the on-disk filesystem backend and can edit the user
@@ -1558,10 +1558,10 @@ def create_cli_agent(
ConfigurableModelMiddleware(),
]
# Resume state: declares the `_context_tokens` and `_model_spec` channels
# and writes them from `after_model` (token count from the latest
# `AIMessage.usage_metadata`, model spec from `context["effective_model"]`).
# The CLI reads them back from `state_values` on thread resume.
# Resume state: declares private checkpoint channels used on resume.
# `ResumeStateMiddleware.after_model` writes `_context_tokens`; model metadata
# is written by `ConfigurableModelMiddleware` from the actual completed model
# request. The CLI reads them back from `state_values` on thread resume.
# Goal tools: exposes the read-only `get_goal`/`get_rubric` tools and the
# constrained `update_goal` tool, and injects goal guidance into the prompt.
from deepagents_code.goal_tools import GoalToolsMiddleware
+20 -3
View File
@@ -1353,6 +1353,9 @@ class _ThreadHistoryPayload:
"""Persisted `_model_spec` from the checkpoint, or `""` for legacy threads
saved before model persistence existed."""
model_params: dict[str, Any] | None = None
"""Persisted `_model_params` from the checkpoint, if any."""
rubric: str | None = None
"""Legacy persisted rubric or graph rubric input, if any."""
@@ -7959,6 +7962,7 @@ class DeepAgentsApp(App):
messages: list[MessageData],
context_tokens: int,
model_spec: str,
model_params: dict[str, Any] | None = None,
) -> _ThreadHistoryPayload:
"""Build a thread payload from raw checkpoint channel values.
@@ -7971,6 +7975,7 @@ class DeepAgentsApp(App):
messages: Converted message data (empty for metadata-only reads).
context_tokens: Persisted context-token count.
model_spec: Persisted model spec, or `""` for legacy threads.
model_params: Persisted model params, or `None` when absent.
Returns:
Payload with goal/rubric channels coerced to known types.
@@ -7984,6 +7989,7 @@ class DeepAgentsApp(App):
messages,
context_tokens,
model_spec,
model_params,
rubric=_as_str(state_values.get("rubric")),
sticky_rubric=_as_str(state_values.get("_sticky_rubric")),
sticky_rubric_recorded="_sticky_rubric" in state_values,
@@ -10232,7 +10238,6 @@ class DeepAgentsApp(App):
context=CLIContext(
model=self._model_override,
model_params=self._model_params_override or {},
effective_model=self._effective_model_spec(),
),
turn_stats=turn_stats,
)
@@ -10523,6 +10528,8 @@ class DeepAgentsApp(App):
)
raw_spec = state_values.get("_model_spec")
model_spec = raw_spec if isinstance(raw_spec, str) else ""
raw_params = state_values.get("_model_params")
model_params = dict(raw_params) if isinstance(raw_params, dict) else None
if _warn_discarded_goal_channels(state_values):
self.notify(
"Some saved goal/rubric state was corrupted and was not restored.",
@@ -10533,6 +10540,7 @@ class DeepAgentsApp(App):
messages=[],
context_tokens=context_tokens,
model_spec=model_spec,
model_params=model_params,
)
messages = state_values.get("messages", [])
@@ -10554,27 +10562,33 @@ class DeepAgentsApp(App):
self,
*,
model_spec: str | None = None,
model_params: dict[str, Any] | None = None,
thread_id: str | None = None,
) -> None:
"""Adopt a resumed thread's persisted model for this session only.
Args:
model_spec: Already-fetched `_model_spec`, when available.
thread_id: Thread ID to fetch `_model_spec` from if needed.
model_params: Already-fetched `_model_params`, when available.
thread_id: Thread ID to fetch `_model_spec`/`_model_params` from if needed.
"""
if not self._should_adopt_resumed_model:
return
self._should_adopt_resumed_model = False
spec = model_spec
params = model_params
if spec is None and thread_id:
state_values = await self._get_thread_state_values(thread_id)
raw_spec = state_values.get("_model_spec")
spec = raw_spec if isinstance(raw_spec, str) else ""
raw_params = state_values.get("_model_params")
params = dict(raw_params) if isinstance(raw_params, dict) else None
if spec:
await self._switch_model(
spec,
extra_kwargs=params,
announce_unchanged=False,
persist=False,
from_resume=True,
@@ -10691,7 +10705,10 @@ class DeepAgentsApp(App):
# consumed on this first load — otherwise a legacy thread (no
# persisted spec) could leave it armed for a later in-session
# `/threads` switch.
await self._adopt_resumed_model_if_needed(model_spec=payload.model_spec)
await self._adopt_resumed_model_if_needed(
model_spec=payload.model_spec,
model_params=payload.model_params,
)
await self._remount_pending_goal_rubric_review()
if not payload.messages:
+133 -25
View File
@@ -10,26 +10,51 @@ from __future__ import annotations
import asyncio
import logging
from collections.abc import Mapping
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from deepagents._models import model_matches_spec # noqa: PLC2701
from deepagents._models import ( # noqa: PLC2701
get_model_identifier,
model_matches_spec,
)
from langchain.agents.middleware.types import (
AgentMiddleware,
ExtendedModelResponse,
ModelRequest,
ModelResponse,
)
from langgraph.types import Command
from deepagents_code._cli_context import CLIContextSchema
if TYPE_CHECKING:
from collections.abc import Awaitable, Callable
from langchain_core.language_models import BaseChatModel
from deepagents_code.config import ModelResult
logger = logging.getLogger(__name__)
@dataclass(frozen=True)
class _ResolvedModelRequest:
"""Model request plus the checkpoint metadata it should persist."""
request: ModelRequest
"""Request to pass to the downstream model handler."""
model_spec: str | None
"""Resolved `provider:model` spec to persist for resume, when known."""
model_params: dict[str, Any] | None = None
"""Invocation params to persist, or `None` to clear checkpointed params."""
model_params_known: bool = False
"""Whether `model_params` is known and should be written to the checkpoint."""
def _get_ls_provider(model: object) -> str | None:
"""Return the LangSmith provider name reported by a chat model.
@@ -40,7 +65,7 @@ def _get_ls_provider(model: object) -> str | None:
"""
try:
ls_params = model._get_ls_params() # ty: ignore[unresolved-attribute]
except (AttributeError, TypeError, RuntimeError):
except (AttributeError, TypeError, RuntimeError, NotImplementedError):
logger.debug("_get_ls_params raised for %s", type(model).__name__)
return None
if isinstance(ls_params, dict):
@@ -152,7 +177,6 @@ def _get_context(request: ModelRequest) -> CLIContextSchema | None:
return CLIContextSchema(
model=ctx.get("model"),
model_params=ctx.get("model_params") or {},
effective_model=ctx.get("effective_model"),
auto_approve=bool(ctx.get("auto_approve", False)),
approval_mode_key=raw_key if isinstance(raw_key, str) else None,
thread_id=raw_thread_id if isinstance(raw_thread_id, str) else None,
@@ -160,6 +184,31 @@ def _get_context(request: ModelRequest) -> CLIContextSchema | None:
return None
def _model_spec_from_model(model: BaseChatModel) -> str | None:
"""Return a resumable `provider:model` spec for a model object."""
provider = _get_ls_provider(model)
model_name = get_model_identifier(model)
if provider and model_name:
return f"{provider}:{model_name}"
from deepagents_code.config import settings
settings_provider = settings.model_provider or ""
settings_model = settings.model_name or ""
if settings_provider and settings_model:
return f"{settings_provider}:{settings_model}"
return None
def _model_spec_from_result(
model_result: ModelResult | None, model: BaseChatModel
) -> str | None:
"""Return the resolved spec from `create_model`, falling back to model metadata."""
if model_result is not None and model_result.provider and model_result.model_name:
return f"{model_result.provider}:{model_result.model_name}"
return _model_spec_from_model(model)
def _build_overrides(
request: ModelRequest, ctx: CLIContextSchema, model_result: ModelResult | None
) -> ModelRequest:
@@ -256,8 +305,8 @@ def _build_overrides(
return request.override(**overrides)
def _apply_overrides(request: ModelRequest) -> ModelRequest:
"""Apply model/param overrides from `CLIContext` on the runtime.
def _apply_overrides(request: ModelRequest) -> _ResolvedModelRequest:
"""Apply model/param overrides and return checkpoint persistence metadata.
Reads `'model'` and `'model_params'` from `runtime.context` and, when
present, swaps the model and/or merges extra settings into the request.
@@ -269,13 +318,12 @@ def _apply_overrides(request: ModelRequest) -> ModelRequest:
request: The incoming model request from the middleware chain.
Returns:
The original request unchanged when no `CLIContext` is present or it
contains no overrides, otherwise a new request with overrides
applied via `request.override()`.
The request to send downstream plus the actual model spec and user-supplied
model params that should be recorded for resume.
"""
ctx = _get_context(request)
if ctx is None:
return request
return _ResolvedModelRequest(request, _model_spec_from_model(request.model))
model_result = None
model = ctx.model
@@ -292,21 +340,32 @@ def _apply_overrides(request: ModelRequest) -> ModelRequest:
"continuing with current model",
model,
)
return request
return _ResolvedModelRequest(
request,
_model_spec_from_model(request.model),
model_params_known=True,
)
return _build_overrides(request, ctx, model_result)
updated = _build_overrides(request, ctx, model_result)
params = dict(ctx.model_params) if ctx.model_params else None
return _ResolvedModelRequest(
updated,
_model_spec_from_result(model_result, updated.model),
params,
model_params_known=True,
)
async def _apply_overrides_async(request: ModelRequest) -> ModelRequest:
async def _apply_overrides_async(request: ModelRequest) -> _ResolvedModelRequest:
"""Async variant of `_apply_overrides` that offloads model construction.
Returns:
The original request when no async override applies, otherwise a request
with the runtime model or settings override applied.
The request to send downstream plus the actual model spec and user-supplied
model params that should be recorded for resume.
"""
ctx = _get_context(request)
if ctx is None:
return request
return _ResolvedModelRequest(request, _model_spec_from_model(request.model))
model_result = None
model = ctx.model
@@ -323,9 +382,36 @@ async def _apply_overrides_async(request: ModelRequest) -> ModelRequest:
"continuing with current model",
model,
)
return request
return _ResolvedModelRequest(
request,
_model_spec_from_model(request.model),
model_params_known=True,
)
return _build_overrides(request, ctx, model_result)
updated = _build_overrides(request, ctx, model_result)
params = dict(ctx.model_params) if ctx.model_params else None
return _ResolvedModelRequest(
updated,
_model_spec_from_result(model_result, updated.model),
params,
model_params_known=True,
)
def _checkpoint_command(resolved: _ResolvedModelRequest) -> Command[Any] | None:
"""Build the private resume-state update for a completed model call.
Returns:
Command with private checkpoint updates, or `None` when nothing is known.
"""
update: dict[str, Any] = {}
if resolved.model_spec:
update["_model_spec"] = resolved.model_spec
if resolved.model_params_known:
update["_model_params"] = resolved.model_params
if not update:
return None
return Command(update=update)
class ConfigurableModelMiddleware(AgentMiddleware):
@@ -345,26 +431,48 @@ class ConfigurableModelMiddleware(AgentMiddleware):
`AnthropicPromptCachingMiddleware`) runs.
"""
def wrap_model_call( # noqa: PLR6301
def __init__(self, *, persist_model_state: bool = True) -> None:
"""Initialize the middleware.
Args:
persist_model_state: Whether completed calls should write private
resume metadata. Subagent instances disable this because they do
not own the parent thread's resume state.
"""
self._persist_model_state = persist_model_state
def wrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], ModelResponse],
) -> ModelResponse:
) -> ModelResponse | ExtendedModelResponse:
"""Apply runtime overrides and delegate to the next handler.
Returns:
The `ModelResponse` produced by the downstream handler.
The downstream response plus a private resume-state update when the
completed call has model metadata to checkpoint.
"""
return handler(_apply_overrides(request))
resolved = _apply_overrides(request)
response = handler(resolved.request)
command = _checkpoint_command(resolved) if self._persist_model_state else None
if command is None:
return response
return ExtendedModelResponse(model_response=response, command=command)
async def awrap_model_call( # noqa: PLR6301
async def awrap_model_call(
self,
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse:
) -> ModelResponse | ExtendedModelResponse:
"""Apply runtime overrides and delegate to the next async handler.
Returns:
The `ModelResponse` produced by the downstream handler.
The downstream response plus a private resume-state update when the
completed call has model metadata to checkpoint.
"""
return await handler(await _apply_overrides_async(request))
resolved = await _apply_overrides_async(request)
response = await handler(resolved.request)
command = _checkpoint_command(resolved) if self._persist_model_state else None
if command is None:
return response
return ExtendedModelResponse(model_response=response, command=command)
+22 -42
View File
@@ -3,14 +3,15 @@
`ResumeState` declares several checkpointed, schema-private channels. They fall
into two groups with *different* write paths:
Written by `ResumeStateMiddleware.after_model`, from inside the graph:
Written from inside the graph on successful model turns:
- `_context_tokens` — total context tokens from the latest
`AIMessage.usage_metadata`. Powers `/tokens` and the status bar.
- `_model_spec` — the `provider:model` spec that was effectively in use for
the turn, read from `runtime.context["effective_model"]`. Lets `dcode -r`
restore the model the resumed thread was actually using instead of falling
back to the user's global default.
`AIMessage.usage_metadata`, written by `ResumeStateMiddleware.after_model`.
Powers `/tokens` and the status bar.
- `_model_spec` / `_model_params` — the model and invocation params effectively
in use for the turn, written by `ConfigurableModelMiddleware` after a
successful model call. Lets `dcode -r` restore the model the resumed thread
was actually using instead of falling back to the user's global default.
Written primarily by the TUI client, via `aupdate_state` (see
`DeepAgentsApp._persist_goal_rubric_state`) — these are user/agent-owned. Most
@@ -31,10 +32,10 @@ have no model-node write site; the two exceptions are called out below:
All of these are facts the CLI reads back from `state_values` on thread resume
so it can rehydrate the session without replaying or re-tokenizing history.
The `after_model` channels are persisted from inside the graph (rather than via
a separate client-side `aupdate_state` call) so the write rides the same
checkpoint as the model response and avoids creating a standalone `UpdateState`
run in LangSmith. Because they are versioned channel state, resuming a specific
The model-turn channels are persisted from inside the graph (rather than via a
separate client-side `aupdate_state` call) so the write rides the same checkpoint
as the model response and avoids creating a standalone `UpdateState` run in
LangSmith. Because they are versioned channel state, resuming a specific
checkpoint yields the values as of *that* checkpoint — not a thread-level
aggregate. The goal/rubric channels are client-written because the user sets
them outside any model turn (except `_goal_status`/`_goal_status_note`, which the
@@ -62,8 +63,6 @@ from langchain.agents.middleware.types import (
)
from langchain_core.messages import AIMessage
from deepagents_code._cli_context import CLIContextSchema
if TYPE_CHECKING:
from langgraph.runtime import Runtime
@@ -145,6 +144,9 @@ class ResumeState(GoalRubricChannels):
_model_spec: Annotated[NotRequired[str], PrivateStateAttr]
"""`provider:model` spec effectively in use for the latest turn."""
_model_params: Annotated[NotRequired[dict[str, Any] | None], PrivateStateAttr]
"""Invocation params effectively in use for the latest turn."""
_pending_goal_objective: Annotated[NotRequired[str | None], PrivateStateAttr]
"""Goal objective awaiting acceptance of proposed criteria."""
@@ -169,25 +171,6 @@ def _extract_context_tokens(message: AIMessage) -> int | None:
return total or None
def _extract_model_spec(runtime: Runtime[ContextT]) -> str | None:
"""Return the effective `provider:model` spec from the runtime context.
The CLI passes the resolved spec in `context["effective_model"]` on every
invocation. Returns `None` when no context is present (e.g. non-CLI
callers) or the field is unset/blank.
"""
ctx = getattr(runtime, "context", None)
if isinstance(ctx, CLIContextSchema):
spec = ctx.effective_model
elif isinstance(ctx, dict):
spec = ctx.get("effective_model")
else:
return None
if isinstance(spec, str) and spec:
return spec
return None
class ResumeStateMiddleware(AgentMiddleware[ResumeState, ContextT]):
"""Persists per-checkpoint resume facts after each model call.
@@ -201,20 +184,21 @@ class ResumeStateMiddleware(AgentMiddleware[ResumeState, ContextT]):
def after_model( # noqa: PLR6301 # AgentMiddleware hook must be an instance method.
self,
state: ResumeState,
runtime: Runtime[ContextT],
runtime: Runtime[ContextT], # noqa: ARG002
) -> dict[str, Any] | None:
"""Write `_context_tokens` and `_model_spec` for the latest turn.
"""Write `_context_tokens` for the latest turn.
Token count comes from the most recent `AIMessage.usage_metadata`; the
model spec comes from `runtime.context["effective_model"]`.
Model metadata is written by `ConfigurableModelMiddleware` from the
actual request that completed successfully; this hook only records token
usage from the most recent `AIMessage.usage_metadata`.
Args:
state: Current agent state; only `messages` is inspected.
runtime: LangGraph runtime; `context["effective_model"]` is read.
runtime: LangGraph runtime required by the middleware interface.
Returns:
State update with whichever of `_context_tokens` / `_model_spec`
could be resolved, or `None` when neither is available.
State update with `_context_tokens`, or `None` when no token count is
available.
"""
update: dict[str, Any] = {}
@@ -225,8 +209,4 @@ class ResumeStateMiddleware(AgentMiddleware[ResumeState, ContextT]):
update["_context_tokens"] = tokens
break
spec = _extract_model_spec(runtime)
if spec is not None:
update["_model_spec"] = spec
return update or None
+2
View File
@@ -573,11 +573,13 @@ class TestStartupSequence:
async def capture_switch( # noqa: RUF029
model_spec: str,
*,
extra_kwargs: dict[str, Any] | None = None,
announce_unchanged: bool = True,
persist: bool = True,
from_resume: bool = False,
) -> None:
assert model_spec == "anthropic:claude-sonnet-4-5"
assert extra_kwargs is None
assert announce_unchanged is False
assert persist is False
assert from_resume is True
@@ -8,7 +8,11 @@ from typing import Any, cast
from unittest.mock import MagicMock, patch
import pytest
from langchain.agents.middleware.types import ModelRequest, ModelResponse
from langchain.agents.middleware.types import (
ExtendedModelResponse,
ModelRequest,
ModelResponse,
)
from langchain_core.language_models import BaseChatModel
from langchain_core.messages import AIMessage, HumanMessage
@@ -56,6 +60,16 @@ def _make_response() -> ModelResponse[Any]:
return ModelResponse(result=[AIMessage(content="response")])
def _checkpoint_update(
result: ModelResponse[Any] | ExtendedModelResponse[Any],
) -> dict[str, Any]:
"""Return the checkpoint update emitted by the middleware."""
assert isinstance(result, ExtendedModelResponse)
assert result.command is not None
assert isinstance(result.command.update, dict)
return result.command.update
def _make_model_result(
model: MagicMock,
*,
@@ -79,24 +93,41 @@ _PATCH_CREATE = "deepagents_code.config.create_model"
_mw = ConfigurableModelMiddleware()
class TestCheckpointPersistence:
"""Tests for private resume-state checkpoint updates."""
def test_can_disable_model_state_persistence(self) -> None:
middleware = ConfigurableModelMiddleware(persist_model_state=False)
request = _make_request(_make_model("gpt-5.5"))
result = middleware.wrap_model_call(request, lambda _request: _make_response())
assert isinstance(result, ModelResponse)
class TestNoOverride:
"""Cases where the middleware should pass the request through unchanged."""
def test_no_context(self) -> None:
request = _make_request(_make_model("claude-sonnet-4-6"), context=None)
captured: list[ModelRequest] = []
_mw.wrap_model_call(
result = _mw.wrap_model_call(
request, lambda r: (captured.append(r), _make_response())[1]
)
assert captured[0].model is request.model
assert _checkpoint_update(result) == {"_model_spec": "openai:claude-sonnet-4-6"}
def test_empty_context(self) -> None:
request = _make_request(_make_model("claude-sonnet-4-6"), context=CLIContext())
captured: list[ModelRequest] = []
_mw.wrap_model_call(
result = _mw.wrap_model_call(
request, lambda r: (captured.append(r), _make_response())[1]
)
assert captured[0] is request
assert _checkpoint_update(result) == {
"_model_spec": "openai:claude-sonnet-4-6",
"_model_params": None,
}
def test_dict_context_reconstructs_approval_fields(self) -> None:
request = _make_request(
@@ -217,10 +248,14 @@ class TestNoOverride:
context=CLIContext(model_params={}),
)
captured: list[ModelRequest] = []
_mw.wrap_model_call(
result = _mw.wrap_model_call(
request, lambda r: (captured.append(r), _make_response())[1]
)
assert captured[0] is request
assert _checkpoint_update(result) == {
"_model_spec": "openai:claude-sonnet-4-6",
"_model_params": None,
}
class TestModelSwap:
@@ -294,17 +329,46 @@ class TestModelSwap:
from deepagents_code.model_config import ModelConfigError
original = _make_model("claude-sonnet-4-6")
original._get_ls_params.return_value = {"ls_provider": "anthropic"}
request = _make_request(
original,
context=CLIContext(model="unknown:bad-model"),
context=CLIContext(
model="unknown:bad-model",
model_params={"temperature": 0.7},
),
)
captured: list[ModelRequest] = []
with patch(_PATCH_CREATE, side_effect=ModelConfigError("no such provider")):
_mw.wrap_model_call(
result = _mw.wrap_model_call(
request, lambda r: (captured.append(r), _make_response())[1]
)
assert captured[0].model is original
assert captured[0].model_settings == {}
assert _checkpoint_update(result) == {
"_model_spec": "anthropic:claude-sonnet-4-6",
"_model_params": None,
}
def test_successful_swap_records_resolved_model_spec(self) -> None:
original = _make_model("claude-sonnet-4-6")
override = _make_model("gpt-5.5")
request = _make_request(original, context=CLIContext(model="openai:gpt-5.5"))
with patch(
_PATCH_CREATE,
return_value=_make_model_result(
override,
model_name="gpt-5.5",
provider="openai",
),
):
result = _mw.wrap_model_call(request, lambda _request: _make_response())
assert _checkpoint_update(result) == {
"_model_spec": "openai:gpt-5.5",
"_model_params": None,
}
class TestAnthropicSettingsStripped:
@@ -715,12 +779,16 @@ class TestModelParams:
context=CLIContext(model_params={"temperature": 0.7}),
)
captured: list[ModelRequest] = []
_mw.wrap_model_call(
result = _mw.wrap_model_call(
request, lambda r: (captured.append(r), _make_response())[1]
)
assert captured[0].model is request.model
assert captured[0].model_settings == {"temperature": 0.7}
assert _checkpoint_update(result) == {
"_model_spec": "openai:claude-sonnet-4-6",
"_model_params": {"temperature": 0.7},
}
def test_params_merge_preserves_existing(self) -> None:
request = _make_request(
@@ -11,7 +11,6 @@ from deepagents_code.resume_state import (
ResumeState,
ResumeStateMiddleware,
_extract_context_tokens,
_extract_model_spec,
coerce_goal_status,
)
@@ -30,6 +29,10 @@ class TestResumeState:
"""ResumeState declares the `_model_spec` channel."""
assert "_model_spec" in ResumeState.__annotations__
def test_state_has_model_params_field(self):
"""ResumeState declares the `_model_params` channel."""
assert "_model_params" in ResumeState.__annotations__
def test_sticky_rubric_field_is_private(self):
"""Persistent TUI rubrics must not leak through the public schema."""
# `_sticky_rubric` is inherited from `GoalRubricChannels`, so resolve the
@@ -105,27 +108,6 @@ class TestExtractContextTokens:
assert _extract_context_tokens(msg) is None
class TestExtractModelSpec:
"""Tests for `_extract_model_spec`."""
def test_returns_effective_model_from_context(self) -> None:
runtime = _runtime({"effective_model": "anthropic:claude-sonnet-4-5"})
assert _extract_model_spec(runtime) == "anthropic:claude-sonnet-4-5" # ty: ignore
def test_returns_none_when_context_missing(self) -> None:
assert _extract_model_spec(_runtime(None)) is None # ty: ignore
def test_returns_none_when_field_absent(self) -> None:
assert _extract_model_spec(_runtime({"model": "x"})) is None # ty: ignore
def test_returns_none_for_blank_or_nonstring(self) -> None:
assert _extract_model_spec(_runtime({"effective_model": ""})) is None # ty: ignore
assert _extract_model_spec(_runtime({"effective_model": None})) is None # ty: ignore
def test_returns_none_when_runtime_is_none(self) -> None:
assert _extract_model_spec(None) is None # ty: ignore
class TestAfterModelHook:
"""Tests for the `after_model` persistence hook."""
@@ -147,7 +129,8 @@ class TestAfterModelHook:
result = middleware.after_model(state, _runtime(None)) # ty: ignore
assert result == {"_context_tokens": 1700}
async def test_writes_model_spec_from_context(self) -> None:
async def test_does_not_write_model_spec_from_context(self) -> None:
"""Model metadata is written by ConfigurableModelMiddleware."""
middleware = ResumeStateMiddleware()
state: dict[str, Any] = {
"messages": [
@@ -162,25 +145,9 @@ class TestAfterModelHook:
),
],
}
runtime = _runtime({"effective_model": "openai:gpt-5.1"})
runtime = _runtime({"model": "openai:gpt-5.1"})
result = middleware.after_model(state, runtime) # ty: ignore
assert result == {
"_context_tokens": 1700,
"_model_spec": "openai:gpt-5.1",
}
async def test_writes_model_spec_without_token_usage(self) -> None:
"""Model spec is recorded even when the AI message reports no usage."""
middleware = ResumeStateMiddleware()
state: dict[str, Any] = {
"messages": [
HumanMessage(content="hi"),
AIMessage(content="no usage info"),
],
}
runtime = _runtime({"effective_model": "openai:gpt-5.1"})
result = middleware.after_model(state, runtime) # ty: ignore
assert result == {"_model_spec": "openai:gpt-5.1"}
assert result == {"_context_tokens": 1700}
async def test_returns_none_when_no_ai_message(self) -> None:
middleware = ResumeStateMiddleware()
@@ -3673,7 +3673,10 @@ class TestResumeModelAdoption:
@staticmethod
def _payload(
model_spec: str, *, with_messages: bool = True
model_spec: str,
*,
with_messages: bool = True,
model_params: dict[str, Any] | None = None,
) -> _ThreadHistoryPayload:
from deepagents_code.widgets.message_store import MessageData, MessageType
@@ -3686,6 +3689,7 @@ class TestResumeModelAdoption:
messages=messages,
context_tokens=0,
model_spec=model_spec,
model_params=model_params,
)
async def test_adopts_persisted_model_session_only(self) -> None:
@@ -3704,12 +3708,39 @@ class TestResumeModelAdoption:
call = switch_mock.await_args
assert call is not None
assert call.args[0] == "anthropic:claude-sonnet-4-5"
assert call.kwargs["extra_kwargs"] is None
assert call.kwargs["persist"] is False
assert call.kwargs["announce_unchanged"] is False
assert call.kwargs["from_resume"] is True
# One-shot: the flag is consumed so later loads don't re-adopt.
assert app._should_adopt_resumed_model is False
async def test_adopts_persisted_model_params(self) -> None:
"""Resume restores the invocation params saved with the model spec."""
app = self._make_app()
switch_mock = AsyncMock()
_app_test_double(app)._switch_model = switch_mock
app._should_adopt_resumed_model = True
await app._load_thread_history(
thread_id="tid-1",
preloaded_payload=self._payload(
"anthropic:claude-sonnet-4-5",
model_params={"temperature": 0.7, "max_tokens": 1024},
),
)
switch_mock.assert_awaited_once()
call = switch_mock.await_args
assert call is not None
assert call.args[0] == "anthropic:claude-sonnet-4-5"
assert call.kwargs["extra_kwargs"] == {
"temperature": 0.7,
"max_tokens": 1024,
}
assert call.kwargs["persist"] is False
assert call.kwargs["from_resume"] is True
async def test_no_adoption_when_flag_unset(self) -> None:
"""Without the armed flag (e.g. in-session switch), model is untouched."""
app = self._make_app()