refactor(agent): capture runtime at turn admission

This commit is contained in:
chengyongru 2026-07-10 14:30:09 +08:00 committed by Xubin Ren
parent 45ace23580
commit 198fd9f869
13 changed files with 209 additions and 54 deletions

View File

@ -23,6 +23,7 @@ from nanobot.agent.context import ContextBuilder
from nanobot.agent.cron_turns import CronTurnCoordinator from nanobot.agent.cron_turns import CronTurnCoordinator
from nanobot.agent.hook import AgentHook, AgentTurnHookFactory from nanobot.agent.hook import AgentHook, AgentTurnHookFactory
from nanobot.agent.memory import Consolidator from nanobot.agent.memory import Consolidator
from nanobot.agent.model_runtime import ModelRuntimeResolver
from nanobot.agent.runner import _MAX_INJECTIONS_PER_TURN, AgentRunner, AgentRunSpec from nanobot.agent.runner import _MAX_INJECTIONS_PER_TURN, AgentRunner, AgentRunSpec
from nanobot.agent.subagent import SubagentManager from nanobot.agent.subagent import SubagentManager
from nanobot.agent.tools.context import RequestContext, bind_request_context, reset_request_context from nanobot.agent.tools.context import RequestContext, bind_request_context, reset_request_context
@ -113,6 +114,7 @@ class TurnContext:
session_key: str session_key: str
state: TurnState state: TurnState
turn_id: str turn_id: str
runtime: LLMRuntime
original_user_text: str | None = None original_user_text: str | None = None
session: Session | None = None session: Session | None = None
@ -173,15 +175,37 @@ class AgentLoop:
return self.tools.tool_names return self.tools.tool_names
def llm_runtime(self) -> LLMRuntime: def llm_runtime(self) -> LLMRuntime:
"""Capture the current provider/model settings owned by this loop.""" """Resolve the immutable default used to admit the next turn."""
self._refresh_provider_snapshot() self._refresh_provider_snapshot()
return LLMRuntime.capture( runtime = self.runtime_resolver.current()
captured = LLMRuntime.capture(
self.provider, self.provider,
self.model, self.model,
context_window_tokens=self.context_window_tokens, context_window_tokens=self.context_window_tokens,
model_preset=self.model_preset, model_preset=self._active_preset,
snapshot_signature=self._provider_signature, snapshot_signature=self._provider_signature,
) )
# Temporary compatibility for MyTool's legacy direct mutations. Round 9
# moves those writes behind the resolver and deletes these projections.
if (
runtime.provider is not self.provider
or runtime.model != self.model
or runtime.generation != captured.generation
or runtime.context_window_tokens != self.context_window_tokens
or runtime.model_preset != self._active_preset
):
snapshot = ProviderSnapshot(
provider=self.provider,
model=self.model,
context_window_tokens=self.context_window_tokens,
signature=self._provider_signature or ("legacy_loop_runtime", self.model),
generation=captured.generation,
)
runtime = self.runtime_resolver.adopt_snapshot(
snapshot,
model_preset=self._active_preset,
)
return runtime
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint" _RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
_PENDING_USER_TURN_KEY = "pending_user_turn" _PENDING_USER_TURN_KEY = "pending_user_turn"
@ -263,6 +287,19 @@ class AgentLoop:
if context_window_tokens is not None if context_window_tokens is not None
else defaults.context_window_tokens else defaults.context_window_tokens
) )
self.model_presets: dict[str, ModelPresetConfig] = model_presets or {}
self._active_preset: str | None = None
self.runtime_resolver = ModelRuntimeResolver(
LLMRuntime.capture(
provider,
self.model,
context_window_tokens=self.context_window_tokens,
snapshot_signature=provider_signature,
),
model_presets=self.model_presets,
provider_snapshot_loader=provider_snapshot_loader,
preset_snapshot_loader=preset_snapshot_loader,
)
self.context_block_limit = context_block_limit self.context_block_limit = context_block_limit
self.max_tool_result_chars = ( self.max_tool_result_chars = (
max_tool_result_chars max_tool_result_chars
@ -369,8 +406,6 @@ class AgentLoop:
consolidator=self.consolidator, consolidator=self.consolidator,
session_ttl_minutes=session_ttl_minutes, session_ttl_minutes=session_ttl_minutes,
) )
self.model_presets: dict[str, ModelPresetConfig] = model_presets or {}
self._active_preset: str | None = None
if model_preset: if model_preset:
self.set_model_preset(model_preset, publish_update=False) self.set_model_preset(model_preset, publish_update=False)
self._register_default_tools() self._register_default_tools()
@ -448,11 +483,14 @@ class AgentLoop:
model_preset: str | None = None, model_preset: str | None = None,
) -> None: ) -> None:
"""Swap model/provider for future turns without disturbing an active one.""" """Swap model/provider for future turns without disturbing an active one."""
provider = snapshot.provider runtime = self.runtime_resolver.adopt_snapshot(
model = snapshot.model snapshot,
context_window_tokens = snapshot.context_window_tokens model_preset=model_preset,
if snapshot.generation is not None: )
provider.generation = snapshot.generation provider = runtime.provider
model = runtime.model
context_window_tokens = runtime.context_window_tokens
provider.generation = runtime.generation
old_model = self.model old_model = self.model
self.provider = provider self.provider = provider
self.model = model self.model = model
@ -519,7 +557,14 @@ class AgentLoop:
def set_model_preset(self, name: str | None, *, publish_update: bool = True) -> None: def set_model_preset(self, name: str | None, *, publish_update: bool = True) -> None:
"""Resolve a preset by name and apply all runtime model dependents.""" """Resolve a preset by name and apply all runtime model dependents."""
name = preset_helpers.normalize_preset_name(name, self.model_presets) name = preset_helpers.normalize_preset_name(name, self.model_presets)
snapshot = self._build_model_preset_snapshot(name) runtime = self.runtime_resolver.select_preset(name)
snapshot = ProviderSnapshot(
provider=runtime.provider,
model=runtime.model,
context_window_tokens=runtime.context_window_tokens,
signature=runtime.snapshot_signature or ("model_preset", name),
generation=runtime.generation,
)
self._apply_provider_snapshot(snapshot, publish_update=publish_update, model_preset=name) self._apply_provider_snapshot(snapshot, publish_update=publish_update, model_preset=name)
self._active_preset = name self._active_preset = name
@ -696,17 +741,18 @@ class AgentLoop:
return UNIFIED_SESSION_KEY return UNIFIED_SESSION_KEY
return msg.session_key return msg.session_key
def _replay_token_budget(self) -> int: @staticmethod
def _replay_token_budget(runtime: LLMRuntime) -> int:
"""Derive a token budget for session history replay from the context window.""" """Derive a token budget for session history replay from the context window."""
if self.context_window_tokens <= 0: if runtime.context_window_tokens <= 0:
return 0 return 0
max_output = getattr(getattr(self.provider, "generation", None), "max_tokens", 4096) max_output = runtime.generation.max_tokens
try: try:
reserved_output = int(max_output) reserved_output = int(max_output)
except (TypeError, ValueError): except (TypeError, ValueError):
reserved_output = 4096 reserved_output = 4096
budget = self.context_window_tokens - max(1, reserved_output) - 1024 budget = runtime.context_window_tokens - max(1, reserved_output) - 1024
return budget if budget > 0 else max(128, self.context_window_tokens // 2) return budget if budget > 0 else max(128, runtime.context_window_tokens // 2)
async def _run_agent_loop( async def _run_agent_loop(
self, self,
@ -716,6 +762,7 @@ class AgentLoop:
on_stream_end: Callable[..., Awaitable[None]] | None = None, on_stream_end: Callable[..., Awaitable[None]] | None = None,
on_retry_wait: Callable[[str], Awaitable[None]] | None = None, on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
*, *,
runtime: LLMRuntime,
session: Session | None = None, session: Session | None = None,
channel: str = "cli", channel: str = "cli",
chat_id: str = "direct", chat_id: str = "direct",
@ -823,6 +870,7 @@ class AgentLoop:
message_id=message_id, message_id=message_id,
session_key=active_session_key, session_key=active_session_key,
original_user_text=original_user_text, original_user_text=original_user_text,
runtime=runtime,
metadata=dict(metadata or {}), metadata=dict(metadata or {}),
) )
file_state_token = bind_file_states(self._file_state_store.for_session(active_session_key)) file_state_token = bind_file_states(self._file_state_store.for_session(active_session_key))
@ -864,7 +912,7 @@ class AgentLoop:
result = await self.runner.run(AgentRunSpec( result = await self.runner.run(AgentRunSpec(
initial_messages=initial_messages, initial_messages=initial_messages,
tools=effective_tools, tools=effective_tools,
runtime=self.llm_runtime(), runtime=runtime,
max_iterations=self.max_iterations, max_iterations=self.max_iterations,
max_tool_result_chars=self.max_tool_result_chars, max_tool_result_chars=self.max_tool_result_chars,
hook=hook, hook=hook,
@ -1200,6 +1248,8 @@ class AgentLoop:
async def _process_system_message( async def _process_system_message(
self, self,
msg: InboundMessage, msg: InboundMessage,
*,
runtime: LLMRuntime,
session_key: str | None = None, session_key: str | None = None,
on_progress: Callable[..., Awaitable[None]] | None = None, on_progress: Callable[..., Awaitable[None]] | None = None,
on_stream: Callable[[str], Awaitable[None]] | None = None, on_stream: Callable[[str], Awaitable[None]] | None = None,
@ -1214,6 +1264,7 @@ class AgentLoop:
logger.info("Processing system message from {}", msg.sender_id) logger.info("Processing system message from {}", msg.sender_id)
key = msg.session_key_override or f"{channel}:{chat_id}" key = msg.session_key_override or f"{channel}:{chat_id}"
session = self.sessions.get_or_create(key) session = self.sessions.get_or_create(key)
self._runtime_events().record_turn_runtime(key, runtime)
if self._restore_runtime_checkpoint(session): if self._restore_runtime_checkpoint(session):
self.sessions.save(session) self.sessions.save(session)
if self._restore_pending_user_turn(session): if self._restore_pending_user_turn(session):
@ -1225,7 +1276,9 @@ class AgentLoop:
await self.consolidator.maybe_consolidate_by_tokens( await self.consolidator.maybe_consolidate_by_tokens(
session, session,
replay_max_messages=self._max_messages, replay_max_messages=replay_max_messages_for_context(
runtime.context_window_tokens
),
) )
is_subagent = msg.sender_id == "subagent" is_subagent = msg.sender_id == "subagent"
if is_subagent and self._persist_subagent_followup(session, msg): if is_subagent and self._persist_subagent_followup(session, msg):
@ -1233,8 +1286,8 @@ class AgentLoop:
self.sessions.save(session) self.sessions.save(session)
current_role = "assistant" if is_subagent else "user" current_role = "assistant" if is_subagent else "user"
_hist_kwargs: dict[str, Any] = { _hist_kwargs: dict[str, Any] = {
"max_messages": self._max_messages, "max_messages": replay_max_messages_for_context(runtime.context_window_tokens),
"max_tokens": self._replay_token_budget(), "max_tokens": self._replay_token_budget(runtime),
"extend_to_user": is_subagent, "extend_to_user": is_subagent,
} }
history = session.get_history(**_hist_kwargs) history = session.get_history(**_hist_kwargs)
@ -1259,6 +1312,7 @@ class AgentLoop:
t_wall = time.time() t_wall = time.time()
final_content, _, all_msgs, stop_reason, _ = await self._run_agent_loop( final_content, _, all_msgs, stop_reason, _ = await self._run_agent_loop(
messages, session=session, channel=channel, chat_id=chat_id, messages, session=session, channel=channel, chat_id=chat_id,
runtime=runtime,
message_id=msg.metadata.get("message_id"), message_id=msg.metadata.get("message_id"),
metadata=msg.metadata, metadata=msg.metadata,
session_key=key, session_key=key,
@ -1278,7 +1332,9 @@ class AgentLoop:
self._schedule_background( self._schedule_background(
self.consolidator.maybe_consolidate_by_tokens( self.consolidator.maybe_consolidate_by_tokens(
session, session,
replay_max_messages=self._max_messages, replay_max_messages=replay_max_messages_for_context(
runtime.context_window_tokens
),
) )
) )
content = final_content or "Background task completed." content = final_content or "Background task completed."
@ -1307,13 +1363,16 @@ class AgentLoop:
hooks: list[AgentHook] | None = None, hooks: list[AgentHook] | None = None,
hook_factories: list[AgentTurnHookFactory] | None = None, hook_factories: list[AgentTurnHookFactory] | None = None,
tools: ToolRegistry | None = None, tools: ToolRegistry | None = None,
runtime: LLMRuntime | None = None,
) -> OutboundMessage | None: ) -> OutboundMessage | None:
"""Process a single inbound message and return the response.""" """Process a single inbound message and return the response."""
self._refresh_provider_snapshot() if runtime is None:
runtime = self.llm_runtime()
if msg.channel == "system": if msg.channel == "system":
return await self._process_system_message( return await self._process_system_message(
msg, msg,
runtime=runtime,
session_key=session_key, session_key=session_key,
on_progress=on_progress, on_progress=on_progress,
on_stream=on_stream, on_stream=on_stream,
@ -1330,6 +1389,7 @@ class AgentLoop:
session_key=key, session_key=key,
state=TurnState.RESTORE, state=TurnState.RESTORE,
turn_id=f"{key}:{time.time_ns()}", turn_id=f"{key}:{time.time_ns()}",
runtime=runtime,
original_user_text=( original_user_text=(
None None
if turn_continuation.internal_continuation_inbound(msg.metadata) if turn_continuation.internal_continuation_inbound(msg.metadata)
@ -1506,24 +1566,27 @@ class AgentLoop:
return "dispatch" return "dispatch"
async def _state_build(self, ctx: TurnContext) -> str: async def _state_build(self, ctx: TurnContext) -> str:
replay_max_messages = replay_max_messages_for_context(
ctx.runtime.context_window_tokens
)
if not ctx.ephemeral: if not ctx.ephemeral:
await self.consolidator.maybe_consolidate_by_tokens( await self.consolidator.maybe_consolidate_by_tokens(
ctx.session, ctx.session,
replay_max_messages=self._max_messages, replay_max_messages=replay_max_messages,
) )
if message_tool := self.tools.get("message"): if message_tool := self.tools.get("message"):
if isinstance(message_tool, MessageTool): if isinstance(message_tool, MessageTool):
message_tool.start_turn() message_tool.start_turn()
_hist_kwargs: dict[str, Any] = { _hist_kwargs: dict[str, Any] = {
"max_messages": self._max_messages, "max_messages": replay_max_messages,
"max_tokens": self._replay_token_budget(), "max_tokens": self._replay_token_budget(ctx.runtime),
"extend_to_user": False, "extend_to_user": False,
} }
ctx.history = ctx.session.get_history(**_hist_kwargs) ctx.history = ctx.session.get_history(**_hist_kwargs)
self._runtime_events().record_turn_runtime( self._runtime_events().record_turn_runtime(
ctx.session_key, ctx.session_key,
self.llm_runtime(), ctx.runtime,
) )
ctx.initial_messages = self._build_initial_messages( ctx.initial_messages = self._build_initial_messages(
@ -1555,6 +1618,7 @@ class AgentLoop:
) )
result = await self._run_agent_loop( result = await self._run_agent_loop(
ctx.initial_messages, ctx.initial_messages,
runtime=ctx.runtime,
on_progress=ctx.on_progress, on_progress=ctx.on_progress,
on_stream=ctx.on_stream, on_stream=ctx.on_stream,
on_stream_end=ctx.on_stream_end, on_stream_end=ctx.on_stream_end,
@ -1613,7 +1677,9 @@ class AgentLoop:
self._schedule_background( self._schedule_background(
self.consolidator.maybe_consolidate_by_tokens( self.consolidator.maybe_consolidate_by_tokens(
ctx.session, ctx.session,
replay_max_messages=self._max_messages, replay_max_messages=replay_max_messages_for_context(
ctx.runtime.context_window_tokens
),
) )
) )
self._clear_pending_user_turn(ctx.session) self._clear_pending_user_turn(ctx.session)
@ -1891,6 +1957,7 @@ class AgentLoop:
hook_factories: list[AgentTurnHookFactory] | None = None, hook_factories: list[AgentTurnHookFactory] | None = None,
tools: ToolRegistry | None = None, tools: ToolRegistry | None = None,
persist_user_message: bool = True, persist_user_message: bool = True,
runtime: LLMRuntime | None = None,
) -> OutboundMessage | None: ) -> OutboundMessage | None:
"""Process a message directly and return the outbound payload.""" """Process a message directly and return the outbound payload."""
await self._connect_mcp() await self._connect_mcp()
@ -1920,6 +1987,8 @@ class AgentLoop:
kwargs["hook_factories"] = hook_factories kwargs["hook_factories"] = hook_factories
if tools is not None: if tools is not None:
kwargs["tools"] = tools kwargs["tools"] = tools
if runtime is not None:
kwargs["runtime"] = runtime
return await self._process_message( return await self._process_message(
msg, msg,
**kwargs, **kwargs,

View File

@ -4,7 +4,10 @@ from __future__ import annotations
from contextlib import contextmanager from contextlib import contextmanager
from contextvars import ContextVar, Token from contextvars import ContextVar, Token
from dataclasses import dataclass, field from dataclasses import dataclass, field
from typing import Any, Callable, Protocol, runtime_checkable from typing import TYPE_CHECKING, Any, Callable, Protocol, runtime_checkable
if TYPE_CHECKING:
from nanobot.utils.llm_runtime import LLMRuntime
_CURRENT_REQUEST_CONTEXT: ContextVar["RequestContext | None"] = ContextVar( _CURRENT_REQUEST_CONTEXT: ContextVar["RequestContext | None"] = ContextVar(
"nanobot_tool_request_context", "nanobot_tool_request_context",
@ -20,6 +23,7 @@ class RequestContext:
message_id: str | None = None message_id: str | None = None
session_key: str | None = None session_key: str | None = None
original_user_text: str | None = None original_user_text: str | None = None
runtime: LLMRuntime | None = None
metadata: dict[str, Any] = field(default_factory=dict) metadata: dict[str, Any] = field(default_factory=dict)

View File

@ -53,6 +53,7 @@ async def test_state_restore_extracts_documents_by_default(
session_key="cli:c", session_key="cli:c",
state=TurnState.RESTORE, state=TurnState.RESTORE,
turn_id="turn-1", turn_id="turn-1",
runtime=loop.llm_runtime(),
) )
assert await loop._state_restore(ctx) == "ok" assert await loop._state_restore(ctx) == "ok"
@ -87,6 +88,7 @@ async def test_state_restore_references_documents_when_extraction_disabled(
session_key="cli:c", session_key="cli:c",
state=TurnState.RESTORE, state=TurnState.RESTORE,
turn_id="turn-1", turn_id="turn-1",
runtime=loop.llm_runtime(),
) )
assert await loop._state_restore(ctx) == "ok" assert await loop._state_restore(ctx) == "ok"
@ -133,6 +135,7 @@ async def test_pending_followup_references_documents_when_extraction_disabled(
final_content, _, _, _, had_injections = await loop._run_agent_loop( final_content, _, _, _, had_injections = await loop._run_agent_loop(
[{"role": "user", "content": "hello"}], [{"role": "user", "content": "hello"}],
runtime=loop.llm_runtime(),
channel="cli", channel="cli",
chat_id="c", chat_id="c",
pending_queue=pending_queue, pending_queue=pending_queue,

View File

@ -458,7 +458,8 @@ async def test_agent_loop_extra_hook_receives_calls(tmp_path):
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
content, tools_used, messages, _, _ = await loop._run_agent_loop( content, tools_used, messages, _, _ = await loop._run_agent_loop(
[{"role": "user", "content": "hi"}] [{"role": "user", "content": "hi"}],
runtime=loop.llm_runtime(),
) )
assert content == "done" assert content == "done"
@ -502,6 +503,7 @@ async def test_agent_loop_turn_hook_factories_receive_context(tmp_path):
await loop._run_agent_loop( await loop._run_agent_loop(
[{"role": "user", "content": "hi"}], [{"role": "user", "content": "hi"}],
runtime=loop.llm_runtime(),
on_progress=on_progress, on_progress=on_progress,
channel="websocket", channel="websocket",
chat_id="chat-1", chat_id="chat-1",
@ -544,7 +546,8 @@ async def test_agent_loop_extra_hook_error_isolation(tmp_path):
loop.tools.get_definitions = MagicMock(return_value=[]) loop.tools.get_definitions = MagicMock(return_value=[])
content, _, _, _, _ = await loop._run_agent_loop( content, _, _, _, _ = await loop._run_agent_loop(
[{"role": "user", "content": "hi"}] [{"role": "user", "content": "hi"}],
runtime=loop.llm_runtime(),
) )
assert content == "still works" assert content == "still works"
@ -568,7 +571,9 @@ async def test_agent_loop_extra_hooks_do_not_swallow_loop_hook_errors(tmp_path):
raise RuntimeError("progress failed") raise RuntimeError("progress failed")
with pytest.raises(RuntimeError, match="progress failed"): with pytest.raises(RuntimeError, match="progress failed"):
await loop._run_agent_loop([], on_progress=bad_progress) await loop._run_agent_loop(
[], runtime=loop.llm_runtime(), on_progress=bad_progress
)
@pytest.mark.asyncio @pytest.mark.asyncio
@ -585,7 +590,9 @@ async def test_agent_loop_no_hooks_backward_compat(tmp_path):
loop.tools.execute = AsyncMock(return_value="ok") loop.tools.execute = AsyncMock(return_value="ok")
loop.max_iterations = 2 loop.max_iterations = 2
content, tools_used, _, _, _ = await loop._run_agent_loop([]) content, tools_used, _, _, _ = await loop._run_agent_loop(
[], runtime=loop.llm_runtime()
)
assert content == ( assert content == (
"I reached the maximum number of tool call iterations (2) " "I reached the maximum number of tool call iterations (2) "
"without completing the task. You can try breaking the task into smaller steps." "without completing the task. You can try breaking the task into smaller steps."

View File

@ -76,7 +76,9 @@ class TestToolEventProgress:
) -> None: ) -> None:
progress.append((content, tool_hint, tool_events)) progress.append((content, tool_hint, tool_events))
final_content, _, _, _, _ = await loop._run_agent_loop([], on_progress=on_progress) final_content, _, _, _, _ = await loop._run_agent_loop(
[], runtime=loop.llm_runtime(), on_progress=on_progress
)
assert final_content == "Done" assert final_content == "Done"
assert progress == [ assert progress == [
@ -145,7 +147,9 @@ class TestToolEventProgress:
if file_edit_events: if file_edit_events:
file_events.extend(file_edit_events) file_events.extend(file_edit_events)
final_content, _, _, _, _ = await loop._run_agent_loop([], on_progress=on_progress) final_content, _, _, _, _ = await loop._run_agent_loop(
[], runtime=loop.llm_runtime(), on_progress=on_progress
)
assert final_content == "Done" assert final_content == "Done"
assert [event["phase"] for event in file_events] == ["start", "end"] assert [event["phase"] for event in file_events] == ["start", "end"]
@ -213,7 +217,9 @@ class TestToolEventProgress:
prepare_file_edit_trackers, prepare_file_edit_trackers,
) )
final_content, _, _, _, _ = await loop._run_agent_loop([], on_progress=on_progress) final_content, _, _, _, _ = await loop._run_agent_loop(
[], runtime=loop.llm_runtime(), on_progress=on_progress
)
assert final_content == "Done" assert final_content == "Done"
assert target.read_text(encoding="utf-8") == "new\n" assert target.read_text(encoding="utf-8") == "new\n"
@ -249,7 +255,9 @@ class TestToolEventProgress:
if file_edit_events: if file_edit_events:
file_events.extend(file_edit_events) file_events.extend(file_edit_events)
await loop._run_agent_loop([], on_progress=on_progress) await loop._run_agent_loop(
[], runtime=loop.llm_runtime(), on_progress=on_progress
)
assert file_events == [] assert file_events == []
@ -623,6 +631,7 @@ class TestToolEventProgress:
final_content, _, _, _, _ = await loop._run_agent_loop( final_content, _, _, _, _ = await loop._run_agent_loop(
[], [],
runtime=loop.llm_runtime(),
on_progress=on_progress, on_progress=on_progress,
on_stream=on_stream, on_stream=on_stream,
) )

View File

@ -40,7 +40,9 @@ async def test_loop_max_iterations_message_stays_stable(tmp_path):
loop.tools.execute = AsyncMock(return_value="ok") loop.tools.execute = AsyncMock(return_value="ok")
loop.max_iterations = 2 loop.max_iterations = 2
final_content, _, _, _, _ = await loop._run_agent_loop([]) final_content, _, _, _, _ = await loop._run_agent_loop(
[], runtime=loop.llm_runtime()
)
assert final_content == ( assert final_content == (
"I reached the maximum number of tool call iterations (2) " "I reached the maximum number of tool call iterations (2) "
@ -61,6 +63,7 @@ async def test_loop_goal_turn_uses_standard_iteration_budget(tmp_path):
final_content, _, _, stop_reason, _ = await loop._run_agent_loop( final_content, _, _, stop_reason, _ = await loop._run_agent_loop(
[], [],
runtime=loop.llm_runtime(),
metadata={"original_command": "/goal"}, metadata={"original_command": "/goal"},
) )
@ -94,6 +97,7 @@ async def test_loop_stream_filter_handles_think_only_prefix_without_crashing(tmp
final_content, _, _, _, _ = await loop._run_agent_loop( final_content, _, _, _, _ = await loop._run_agent_loop(
[], [],
runtime=loop.llm_runtime(),
on_stream=on_stream, on_stream=on_stream,
on_stream_end=on_stream_end, on_stream_end=on_stream_end,
) )
@ -118,7 +122,9 @@ async def test_loop_stream_filter_hides_partial_trailing_think_prefix(tmp_path):
async def on_stream(delta: str) -> None: async def on_stream(delta: str) -> None:
deltas.append(delta) deltas.append(delta)
final_content, _, _, _, _ = await loop._run_agent_loop([], on_stream=on_stream) final_content, _, _, _, _ = await loop._run_agent_loop(
[], runtime=loop.llm_runtime(), on_stream=on_stream
)
assert final_content == "Hello World" assert final_content == "Hello World"
assert deltas == ["Hello", " World"] assert deltas == ["Hello", " World"]
@ -139,7 +145,9 @@ async def test_loop_stream_filter_hides_complete_trailing_think_tag(tmp_path):
async def on_stream(delta: str) -> None: async def on_stream(delta: str) -> None:
deltas.append(delta) deltas.append(delta)
final_content, _, _, _, _ = await loop._run_agent_loop([], on_stream=on_stream) final_content, _, _, _, _ = await loop._run_agent_loop(
[], runtime=loop.llm_runtime(), on_stream=on_stream
)
assert final_content == "Hello World" assert final_content == "Hello World"
assert deltas == ["Hello", " World"] assert deltas == ["Hello", " World"]
@ -158,7 +166,9 @@ async def test_loop_retries_think_only_final_response(tmp_path):
loop.provider.chat_with_retry = chat_with_retry loop.provider.chat_with_retry = chat_with_retry
final_content, _, _, _, _ = await loop._run_agent_loop([]) final_content, _, _, _, _ = await loop._run_agent_loop(
[], runtime=loop.llm_runtime()
)
assert final_content == "Recovered answer" assert final_content == "Recovered answer"
assert call_count["n"] == 2 assert call_count["n"] == 2

View File

@ -316,7 +316,7 @@ def test_webui_title_update_uses_captured_llm_runtime(
coordinator.capture_title_context( coordinator.capture_title_context(
"websocket:chat1", "websocket:chat1",
msg, msg,
LLMRuntime(provider, "turn-model"), LLMRuntime.capture(provider, "turn-model", context_window_tokens=32_768),
) )
asyncio.run(coordinator.handle_turn_end( asyncio.run(coordinator.handle_turn_end(
msg, msg,
@ -1055,6 +1055,7 @@ async def test_run_agent_loop_goal_continue_message_reads_latest_metadata(
await loop._run_agent_loop( await loop._run_agent_loop(
[], [],
runtime=loop.llm_runtime(),
session=session, session=session,
channel="websocket", channel="websocket",
chat_id="late-goal", chat_id="late-goal",
@ -1273,10 +1274,14 @@ async def test_system_subagent_followup_is_persisted_before_prompt_assembly(tmp_
session.add_message("assistant", "working") session.add_message("assistant", "working")
loop.sessions.save(session) loop.sessions.save(session)
seen: dict[str, list[dict]] = {} runtime = loop.llm_runtime()
seen: dict[str, object] = {}
record_runtime = MagicMock(wraps=loop._runtime_events().record_turn_runtime)
loop.runtime_event_publisher.record_turn_runtime = record_runtime
async def fake_run_agent_loop(initial_messages, **_kwargs): async def fake_run_agent_loop(initial_messages, **kwargs):
seen["initial_messages"] = initial_messages seen["initial_messages"] = initial_messages
seen["runtime"] = kwargs["runtime"]
return ( return (
"done", "done",
[], [],
@ -1294,10 +1299,15 @@ async def test_system_subagent_followup_is_persisted_before_prompt_assembly(tmp_
chat_id="cli:test", chat_id="cli:test",
content="subagent result", content="subagent result",
metadata={"subagent_task_id": "sub-1"}, metadata={"subagent_task_id": "sub-1"},
) ),
runtime=runtime,
) )
non_system = [m for m in seen["initial_messages"] if m.get("role") != "system"] assert seen["runtime"] is runtime
record_runtime.assert_called_once_with("cli:test", runtime)
initial_messages = seen["initial_messages"]
assert isinstance(initial_messages, list)
non_system = [m for m in initial_messages if m.get("role") != "system"]
assert "question" in non_system[0]["content"] assert "question" in non_system[0]["content"]
assert "working" in non_system[1]["content"] assert "working" in non_system[1]["content"]
# Persisted timestamps stay in session records, but replay content is not # Persisted timestamps stay in session records, but replay content is not

View File

@ -23,10 +23,12 @@ class _ContextRecordingTool:
def __init__(self) -> None: def __init__(self) -> None:
self.contexts: list[dict] = [] self.contexts: list[dict] = []
self.runtimes: list[object] = []
async def execute(self, **_kwargs) -> str: async def execute(self, **_kwargs) -> str:
ctx = current_request_context() ctx = current_request_context()
assert ctx is not None assert ctx is not None
self.runtimes.append(ctx.runtime)
self.contexts.append({ self.contexts.append({
"channel": ctx.channel, "channel": ctx.channel,
"chat_id": ctx.chat_id, "chat_id": ctx.chat_id,
@ -81,8 +83,10 @@ async def test_loop_binds_request_context_for_tool_execution(tmp_path: Path) ->
loop.tools = _Tools(cron) loop.tools = _Tools(cron)
metadata = {"slack": {"thread_ts": "111.222", "channel_type": "channel"}} metadata = {"slack": {"thread_ts": "111.222", "channel_type": "channel"}}
runtime = loop.llm_runtime()
await loop._run_agent_loop( await loop._run_agent_loop(
[], [],
runtime=runtime,
channel="slack", channel="slack",
chat_id="C123", chat_id="C123",
metadata=metadata, metadata=metadata,
@ -95,6 +99,7 @@ async def test_loop_binds_request_context_for_tool_execution(tmp_path: Path) ->
"metadata": metadata, "metadata": metadata,
"session_key": "slack:C123:111.222", "session_key": "slack:C123:111.222",
} }
assert cron.runtimes[-1] is runtime
def test_request_context_nested_bind_restores_outer_context() -> None: def test_request_context_nested_bind_restores_outer_context() -> None:
@ -160,10 +165,13 @@ async def test_agent_loop_restores_outer_request_context_after_runner_exception(
model="test-model", model="test-model",
) )
outer = RequestContext(channel="test", chat_id="outer", session_key="test:outer") outer = RequestContext(channel="test", chat_id="outer", session_key="test:outer")
runtime = loop.llm_runtime()
async def fail_run(_spec): async def fail_run(spec):
current = current_request_context() current = current_request_context()
assert current is not None assert current is not None
assert spec.runtime is runtime
assert current.runtime is runtime
assert current.channel == "slack" assert current.channel == "slack"
assert current.chat_id == "C123" assert current.chat_id == "C123"
assert current.session_key == "slack:C123:111.222" assert current.session_key == "slack:C123:111.222"
@ -176,6 +184,7 @@ async def test_agent_loop_restores_outer_request_context_after_runner_exception(
with pytest.raises(RuntimeError, match="runner failed"): with pytest.raises(RuntimeError, match="runner failed"):
await loop._run_agent_loop( await loop._run_agent_loop(
[], [],
runtime=runtime,
channel="slack", channel="slack",
chat_id="C123", chat_id="C123",
session_key="slack:C123:111.222", session_key="slack:C123:111.222",
@ -209,10 +218,11 @@ async def test_process_message_captures_original_text_before_restore(
workspace=tmp_path, workspace=tmp_path,
model="test-model", model="test-model",
) )
seen: list[str | None] = [] runtime = loop.llm_runtime()
seen: list[tuple[str | None, object]] = []
async def stop_after_capture(ctx) -> str: async def stop_after_capture(ctx) -> str:
seen.append(ctx.original_user_text) seen.append((ctx.original_user_text, ctx.runtime))
raise RuntimeError("captured before restore") raise RuntimeError("captured before restore")
loop._state_restore = stop_after_capture # type: ignore[method-assign] loop._state_restore = stop_after_capture # type: ignore[method-assign]
@ -225,7 +235,8 @@ async def test_process_message_captures_original_text_before_restore(
chat_id="C123", chat_id="C123",
content=" original user text ", content=" original user text ",
metadata=metadata, metadata=metadata,
) ),
runtime=runtime,
) )
assert seen == [expected] assert seen == [(expected, runtime)]

View File

@ -446,6 +446,7 @@ async def test_loop_injected_followup_preserves_image_media(tmp_path):
final_content, _, _, _, had_injections = await loop._run_agent_loop( final_content, _, _, _, had_injections = await loop._run_agent_loop(
[{"role": "user", "content": "hello"}], [{"role": "user", "content": "hello"}],
runtime=loop.llm_runtime(),
channel="cli", channel="cli",
chat_id="c", chat_id="c",
pending_queue=pending_queue, pending_queue=pending_queue,
@ -511,6 +512,7 @@ async def test_subagent_pending_injection_is_hidden_history_and_not_merged(tmp_p
final_content, _, all_msgs, _, had_injections = await loop._run_agent_loop( final_content, _, all_msgs, _, had_injections = await loop._run_agent_loop(
[{"role": "user", "content": "hello"}], [{"role": "user", "content": "hello"}],
runtime=loop.llm_runtime(),
channel="cli", channel="cli",
chat_id="c", chat_id="c",
pending_queue=pending_queue, pending_queue=pending_queue,
@ -993,6 +995,7 @@ async def test_pending_queue_preserves_overflow_for_next_injection_cycle(tmp_pat
final_content, _, _, _, had_injections = await loop._run_agent_loop( final_content, _, _, _, had_injections = await loop._run_agent_loop(
[{"role": "user", "content": "hello"}], [{"role": "user", "content": "hello"}],
runtime=loop.llm_runtime(),
channel="cli", channel="cli",
chat_id="c", chat_id="c",
pending_queue=pending_queue, pending_queue=pending_queue,

View File

@ -6,6 +6,7 @@ from nanobot.agent.loop import AgentLoop
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.config.loader import save_config from nanobot.config.loader import save_config
from nanobot.config.schema import Config from nanobot.config.schema import Config
from nanobot.providers.base import GenerationSettings
from nanobot.providers.factory import ProviderSnapshot, load_provider_snapshot from nanobot.providers.factory import ProviderSnapshot, load_provider_snapshot
from nanobot.webui.settings_api import update_agent_settings from nanobot.webui.settings_api import update_agent_settings
@ -74,6 +75,29 @@ def test_llm_runtime_refreshes_provider_snapshot(tmp_path: Path) -> None:
assert not hasattr(loop.runner, "provider") assert not hasattr(loop.runner, "provider")
def test_next_turn_captures_generation_changed_after_previous_admission(
tmp_path: Path,
) -> None:
provider = _provider("test-model")
provider.generation = GenerationSettings(temperature=0.2, max_tokens=1024)
loop = AgentLoop(
bus=MessageBus(),
provider=provider,
workspace=tmp_path,
model="test-model",
context_window_tokens=16_384,
)
first = loop.llm_runtime()
provider.generation = GenerationSettings(temperature=0.8, max_tokens=512)
second = loop.llm_runtime()
assert first.generation.temperature == 0.2
assert first.generation.max_tokens == 1024
assert second.generation.temperature == 0.8
assert second.generation.max_tokens == 512
def test_settings_context_window_refreshes_runtime_state( def test_settings_context_window_refreshes_runtime_state(
tmp_path: Path, tmp_path: Path,
monkeypatch, monkeypatch,

View File

@ -272,7 +272,7 @@ async def test_agent_loop_syncs_updated_max_iterations_before_run(tmp_path):
loop.runner.run = AsyncMock(side_effect=fake_run) loop.runner.run = AsyncMock(side_effect=fake_run)
loop.max_iterations = 55 loop.max_iterations = 55
await loop._run_agent_loop([]) await loop._run_agent_loop([], runtime=loop.llm_runtime())
loop.runner.run.assert_awaited_once() loop.runner.run.assert_awaited_once()
@ -327,6 +327,7 @@ async def test_drain_pending_blocks_while_subagents_running(tmp_path):
# Run _run_agent_loop — this defines the _drain_pending closure # Run _run_agent_loop — this defines the _drain_pending closure
await loop._run_agent_loop( await loop._run_agent_loop(
[{"role": "user", "content": "test"}], [{"role": "user", "content": "test"}],
runtime=loop.llm_runtime(),
session=session, session=session,
channel="test", channel="test",
chat_id="c1", chat_id="c1",
@ -402,6 +403,7 @@ async def test_drain_pending_no_block_when_no_subagents(tmp_path):
await loop._run_agent_loop( await loop._run_agent_loop(
[{"role": "user", "content": "test"}], [{"role": "user", "content": "test"}],
runtime=loop.llm_runtime(),
session=None, session=None,
channel="test", channel="test",
chat_id="c1", chat_id="c1",
@ -458,6 +460,7 @@ async def test_drain_pending_timeout(tmp_path):
await loop._run_agent_loop( await loop._run_agent_loop(
[{"role": "user", "content": "test"}], [{"role": "user", "content": "test"}],
runtime=loop.llm_runtime(),
session=session, session=session,
channel="test", channel="test",
chat_id="c1", chat_id="c1",

View File

@ -296,11 +296,11 @@ class TestRestartCommand:
LLMResponse(content="second", usage={}), LLMResponse(content="second", usage={}),
]) ])
await loop._run_agent_loop([]) await loop._run_agent_loop([], runtime=loop.llm_runtime())
assert loop._last_usage["prompt_tokens"] == 9 assert loop._last_usage["prompt_tokens"] == 9
assert loop._last_usage["completion_tokens"] == 4 assert loop._last_usage["completion_tokens"] == 4
await loop._run_agent_loop([]) await loop._run_agent_loop([], runtime=loop.llm_runtime())
assert loop._last_usage["prompt_tokens"] == 123 assert loop._last_usage["prompt_tokens"] == 123
assert loop._last_usage["completion_tokens"] == 7 assert loop._last_usage["completion_tokens"] == 7
assert loop._last_usage["estimated_tokens"] == 130 assert loop._last_usage["estimated_tokens"] == 130

View File

@ -144,7 +144,9 @@ class TestMessageToolSuppressLogic:
async def on_progress(content: str, *, tool_hint: bool = False) -> None: async def on_progress(content: str, *, tool_hint: bool = False) -> None:
progress.append((content, tool_hint)) progress.append((content, tool_hint))
final_content, _, _, _, _ = await loop._run_agent_loop([], on_progress=on_progress) final_content, _, _, _, _ = await loop._run_agent_loop(
[], runtime=loop.llm_runtime(), on_progress=on_progress
)
assert final_content == "Done" assert final_content == "Done"
assert progress == [ assert progress == [