refactor(agent): capture runtime at turn admission
This commit is contained in:
+96
-27
@@ -23,6 +23,7 @@ from nanobot.agent.context import ContextBuilder
|
||||
from nanobot.agent.cron_turns import CronTurnCoordinator
|
||||
from nanobot.agent.hook import AgentHook, AgentTurnHookFactory
|
||||
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.subagent import SubagentManager
|
||||
from nanobot.agent.tools.context import RequestContext, bind_request_context, reset_request_context
|
||||
@@ -113,6 +114,7 @@ class TurnContext:
|
||||
session_key: str
|
||||
state: TurnState
|
||||
turn_id: str
|
||||
runtime: LLMRuntime
|
||||
original_user_text: str | None = None
|
||||
session: Session | None = None
|
||||
|
||||
@@ -173,15 +175,37 @@ class AgentLoop:
|
||||
return self.tools.tool_names
|
||||
|
||||
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()
|
||||
return LLMRuntime.capture(
|
||||
runtime = self.runtime_resolver.current()
|
||||
captured = LLMRuntime.capture(
|
||||
self.provider,
|
||||
self.model,
|
||||
context_window_tokens=self.context_window_tokens,
|
||||
model_preset=self.model_preset,
|
||||
model_preset=self._active_preset,
|
||||
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"
|
||||
_PENDING_USER_TURN_KEY = "pending_user_turn"
|
||||
@@ -263,6 +287,19 @@ class AgentLoop:
|
||||
if context_window_tokens is not None
|
||||
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.max_tool_result_chars = (
|
||||
max_tool_result_chars
|
||||
@@ -369,8 +406,6 @@ class AgentLoop:
|
||||
consolidator=self.consolidator,
|
||||
session_ttl_minutes=session_ttl_minutes,
|
||||
)
|
||||
self.model_presets: dict[str, ModelPresetConfig] = model_presets or {}
|
||||
self._active_preset: str | None = None
|
||||
if model_preset:
|
||||
self.set_model_preset(model_preset, publish_update=False)
|
||||
self._register_default_tools()
|
||||
@@ -448,11 +483,14 @@ class AgentLoop:
|
||||
model_preset: str | None = None,
|
||||
) -> None:
|
||||
"""Swap model/provider for future turns without disturbing an active one."""
|
||||
provider = snapshot.provider
|
||||
model = snapshot.model
|
||||
context_window_tokens = snapshot.context_window_tokens
|
||||
if snapshot.generation is not None:
|
||||
provider.generation = snapshot.generation
|
||||
runtime = self.runtime_resolver.adopt_snapshot(
|
||||
snapshot,
|
||||
model_preset=model_preset,
|
||||
)
|
||||
provider = runtime.provider
|
||||
model = runtime.model
|
||||
context_window_tokens = runtime.context_window_tokens
|
||||
provider.generation = runtime.generation
|
||||
old_model = self.model
|
||||
self.provider = provider
|
||||
self.model = model
|
||||
@@ -519,7 +557,14 @@ class AgentLoop:
|
||||
def set_model_preset(self, name: str | None, *, publish_update: bool = True) -> None:
|
||||
"""Resolve a preset by name and apply all runtime model dependents."""
|
||||
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._active_preset = name
|
||||
|
||||
@@ -696,17 +741,18 @@ class AgentLoop:
|
||||
return UNIFIED_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."""
|
||||
if self.context_window_tokens <= 0:
|
||||
if runtime.context_window_tokens <= 0:
|
||||
return 0
|
||||
max_output = getattr(getattr(self.provider, "generation", None), "max_tokens", 4096)
|
||||
max_output = runtime.generation.max_tokens
|
||||
try:
|
||||
reserved_output = int(max_output)
|
||||
except (TypeError, ValueError):
|
||||
reserved_output = 4096
|
||||
budget = self.context_window_tokens - max(1, reserved_output) - 1024
|
||||
return budget if budget > 0 else max(128, self.context_window_tokens // 2)
|
||||
budget = runtime.context_window_tokens - max(1, reserved_output) - 1024
|
||||
return budget if budget > 0 else max(128, runtime.context_window_tokens // 2)
|
||||
|
||||
async def _run_agent_loop(
|
||||
self,
|
||||
@@ -716,6 +762,7 @@ class AgentLoop:
|
||||
on_stream_end: Callable[..., Awaitable[None]] | None = None,
|
||||
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
|
||||
*,
|
||||
runtime: LLMRuntime,
|
||||
session: Session | None = None,
|
||||
channel: str = "cli",
|
||||
chat_id: str = "direct",
|
||||
@@ -823,6 +870,7 @@ class AgentLoop:
|
||||
message_id=message_id,
|
||||
session_key=active_session_key,
|
||||
original_user_text=original_user_text,
|
||||
runtime=runtime,
|
||||
metadata=dict(metadata or {}),
|
||||
)
|
||||
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(
|
||||
initial_messages=initial_messages,
|
||||
tools=effective_tools,
|
||||
runtime=self.llm_runtime(),
|
||||
runtime=runtime,
|
||||
max_iterations=self.max_iterations,
|
||||
max_tool_result_chars=self.max_tool_result_chars,
|
||||
hook=hook,
|
||||
@@ -1200,6 +1248,8 @@ class AgentLoop:
|
||||
async def _process_system_message(
|
||||
self,
|
||||
msg: InboundMessage,
|
||||
*,
|
||||
runtime: LLMRuntime,
|
||||
session_key: str | None = None,
|
||||
on_progress: Callable[..., 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)
|
||||
key = msg.session_key_override or f"{channel}:{chat_id}"
|
||||
session = self.sessions.get_or_create(key)
|
||||
self._runtime_events().record_turn_runtime(key, runtime)
|
||||
if self._restore_runtime_checkpoint(session):
|
||||
self.sessions.save(session)
|
||||
if self._restore_pending_user_turn(session):
|
||||
@@ -1225,7 +1276,9 @@ class AgentLoop:
|
||||
|
||||
await self.consolidator.maybe_consolidate_by_tokens(
|
||||
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"
|
||||
if is_subagent and self._persist_subagent_followup(session, msg):
|
||||
@@ -1233,8 +1286,8 @@ class AgentLoop:
|
||||
self.sessions.save(session)
|
||||
current_role = "assistant" if is_subagent else "user"
|
||||
_hist_kwargs: dict[str, Any] = {
|
||||
"max_messages": self._max_messages,
|
||||
"max_tokens": self._replay_token_budget(),
|
||||
"max_messages": replay_max_messages_for_context(runtime.context_window_tokens),
|
||||
"max_tokens": self._replay_token_budget(runtime),
|
||||
"extend_to_user": is_subagent,
|
||||
}
|
||||
history = session.get_history(**_hist_kwargs)
|
||||
@@ -1259,6 +1312,7 @@ class AgentLoop:
|
||||
t_wall = time.time()
|
||||
final_content, _, all_msgs, stop_reason, _ = await self._run_agent_loop(
|
||||
messages, session=session, channel=channel, chat_id=chat_id,
|
||||
runtime=runtime,
|
||||
message_id=msg.metadata.get("message_id"),
|
||||
metadata=msg.metadata,
|
||||
session_key=key,
|
||||
@@ -1278,7 +1332,9 @@ class AgentLoop:
|
||||
self._schedule_background(
|
||||
self.consolidator.maybe_consolidate_by_tokens(
|
||||
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."
|
||||
@@ -1307,13 +1363,16 @@ class AgentLoop:
|
||||
hooks: list[AgentHook] | None = None,
|
||||
hook_factories: list[AgentTurnHookFactory] | None = None,
|
||||
tools: ToolRegistry | None = None,
|
||||
runtime: LLMRuntime | None = None,
|
||||
) -> OutboundMessage | None:
|
||||
"""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":
|
||||
return await self._process_system_message(
|
||||
msg,
|
||||
runtime=runtime,
|
||||
session_key=session_key,
|
||||
on_progress=on_progress,
|
||||
on_stream=on_stream,
|
||||
@@ -1330,6 +1389,7 @@ class AgentLoop:
|
||||
session_key=key,
|
||||
state=TurnState.RESTORE,
|
||||
turn_id=f"{key}:{time.time_ns()}",
|
||||
runtime=runtime,
|
||||
original_user_text=(
|
||||
None
|
||||
if turn_continuation.internal_continuation_inbound(msg.metadata)
|
||||
@@ -1506,24 +1566,27 @@ class AgentLoop:
|
||||
return "dispatch"
|
||||
|
||||
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:
|
||||
await self.consolidator.maybe_consolidate_by_tokens(
|
||||
ctx.session,
|
||||
replay_max_messages=self._max_messages,
|
||||
replay_max_messages=replay_max_messages,
|
||||
)
|
||||
if message_tool := self.tools.get("message"):
|
||||
if isinstance(message_tool, MessageTool):
|
||||
message_tool.start_turn()
|
||||
|
||||
_hist_kwargs: dict[str, Any] = {
|
||||
"max_messages": self._max_messages,
|
||||
"max_tokens": self._replay_token_budget(),
|
||||
"max_messages": replay_max_messages,
|
||||
"max_tokens": self._replay_token_budget(ctx.runtime),
|
||||
"extend_to_user": False,
|
||||
}
|
||||
ctx.history = ctx.session.get_history(**_hist_kwargs)
|
||||
self._runtime_events().record_turn_runtime(
|
||||
ctx.session_key,
|
||||
self.llm_runtime(),
|
||||
ctx.runtime,
|
||||
)
|
||||
|
||||
ctx.initial_messages = self._build_initial_messages(
|
||||
@@ -1555,6 +1618,7 @@ class AgentLoop:
|
||||
)
|
||||
result = await self._run_agent_loop(
|
||||
ctx.initial_messages,
|
||||
runtime=ctx.runtime,
|
||||
on_progress=ctx.on_progress,
|
||||
on_stream=ctx.on_stream,
|
||||
on_stream_end=ctx.on_stream_end,
|
||||
@@ -1613,7 +1677,9 @@ class AgentLoop:
|
||||
self._schedule_background(
|
||||
self.consolidator.maybe_consolidate_by_tokens(
|
||||
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)
|
||||
@@ -1891,6 +1957,7 @@ class AgentLoop:
|
||||
hook_factories: list[AgentTurnHookFactory] | None = None,
|
||||
tools: ToolRegistry | None = None,
|
||||
persist_user_message: bool = True,
|
||||
runtime: LLMRuntime | None = None,
|
||||
) -> OutboundMessage | None:
|
||||
"""Process a message directly and return the outbound payload."""
|
||||
await self._connect_mcp()
|
||||
@@ -1920,6 +1987,8 @@ class AgentLoop:
|
||||
kwargs["hook_factories"] = hook_factories
|
||||
if tools is not None:
|
||||
kwargs["tools"] = tools
|
||||
if runtime is not None:
|
||||
kwargs["runtime"] = runtime
|
||||
return await self._process_message(
|
||||
msg,
|
||||
**kwargs,
|
||||
|
||||
@@ -4,7 +4,10 @@ from __future__ import annotations
|
||||
from contextlib import contextmanager
|
||||
from contextvars import ContextVar, Token
|
||||
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(
|
||||
"nanobot_tool_request_context",
|
||||
@@ -20,6 +23,7 @@ class RequestContext:
|
||||
message_id: str | None = None
|
||||
session_key: str | None = None
|
||||
original_user_text: str | None = None
|
||||
runtime: LLMRuntime | None = None
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
|
||||
@@ -53,6 +53,7 @@ async def test_state_restore_extracts_documents_by_default(
|
||||
session_key="cli:c",
|
||||
state=TurnState.RESTORE,
|
||||
turn_id="turn-1",
|
||||
runtime=loop.llm_runtime(),
|
||||
)
|
||||
|
||||
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",
|
||||
state=TurnState.RESTORE,
|
||||
turn_id="turn-1",
|
||||
runtime=loop.llm_runtime(),
|
||||
)
|
||||
|
||||
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(
|
||||
[{"role": "user", "content": "hello"}],
|
||||
runtime=loop.llm_runtime(),
|
||||
channel="cli",
|
||||
chat_id="c",
|
||||
pending_queue=pending_queue,
|
||||
|
||||
@@ -458,7 +458,8 @@ async def test_agent_loop_extra_hook_receives_calls(tmp_path):
|
||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||
|
||||
content, tools_used, messages, _, _ = await loop._run_agent_loop(
|
||||
[{"role": "user", "content": "hi"}]
|
||||
[{"role": "user", "content": "hi"}],
|
||||
runtime=loop.llm_runtime(),
|
||||
)
|
||||
|
||||
assert content == "done"
|
||||
@@ -502,6 +503,7 @@ async def test_agent_loop_turn_hook_factories_receive_context(tmp_path):
|
||||
|
||||
await loop._run_agent_loop(
|
||||
[{"role": "user", "content": "hi"}],
|
||||
runtime=loop.llm_runtime(),
|
||||
on_progress=on_progress,
|
||||
channel="websocket",
|
||||
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=[])
|
||||
|
||||
content, _, _, _, _ = await loop._run_agent_loop(
|
||||
[{"role": "user", "content": "hi"}]
|
||||
[{"role": "user", "content": "hi"}],
|
||||
runtime=loop.llm_runtime(),
|
||||
)
|
||||
|
||||
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")
|
||||
|
||||
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
|
||||
@@ -585,7 +590,9 @@ async def test_agent_loop_no_hooks_backward_compat(tmp_path):
|
||||
loop.tools.execute = AsyncMock(return_value="ok")
|
||||
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 == (
|
||||
"I reached the maximum number of tool call iterations (2) "
|
||||
"without completing the task. You can try breaking the task into smaller steps."
|
||||
|
||||
@@ -76,7 +76,9 @@ class TestToolEventProgress:
|
||||
) -> None:
|
||||
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 progress == [
|
||||
@@ -145,7 +147,9 @@ class TestToolEventProgress:
|
||||
if 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 [event["phase"] for event in file_events] == ["start", "end"]
|
||||
@@ -213,7 +217,9 @@ class TestToolEventProgress:
|
||||
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 target.read_text(encoding="utf-8") == "new\n"
|
||||
@@ -249,7 +255,9 @@ class TestToolEventProgress:
|
||||
if 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 == []
|
||||
|
||||
@@ -623,6 +631,7 @@ class TestToolEventProgress:
|
||||
|
||||
final_content, _, _, _, _ = await loop._run_agent_loop(
|
||||
[],
|
||||
runtime=loop.llm_runtime(),
|
||||
on_progress=on_progress,
|
||||
on_stream=on_stream,
|
||||
)
|
||||
|
||||
@@ -40,7 +40,9 @@ async def test_loop_max_iterations_message_stays_stable(tmp_path):
|
||||
loop.tools.execute = AsyncMock(return_value="ok")
|
||||
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 == (
|
||||
"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(
|
||||
[],
|
||||
runtime=loop.llm_runtime(),
|
||||
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(
|
||||
[],
|
||||
runtime=loop.llm_runtime(),
|
||||
on_stream=on_stream,
|
||||
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:
|
||||
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 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:
|
||||
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 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
|
||||
|
||||
final_content, _, _, _, _ = await loop._run_agent_loop([])
|
||||
final_content, _, _, _, _ = await loop._run_agent_loop(
|
||||
[], runtime=loop.llm_runtime()
|
||||
)
|
||||
|
||||
assert final_content == "Recovered answer"
|
||||
assert call_count["n"] == 2
|
||||
|
||||
@@ -316,7 +316,7 @@ def test_webui_title_update_uses_captured_llm_runtime(
|
||||
coordinator.capture_title_context(
|
||||
"websocket:chat1",
|
||||
msg,
|
||||
LLMRuntime(provider, "turn-model"),
|
||||
LLMRuntime.capture(provider, "turn-model", context_window_tokens=32_768),
|
||||
)
|
||||
asyncio.run(coordinator.handle_turn_end(
|
||||
msg,
|
||||
@@ -1055,6 +1055,7 @@ async def test_run_agent_loop_goal_continue_message_reads_latest_metadata(
|
||||
|
||||
await loop._run_agent_loop(
|
||||
[],
|
||||
runtime=loop.llm_runtime(),
|
||||
session=session,
|
||||
channel="websocket",
|
||||
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")
|
||||
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["runtime"] = kwargs["runtime"]
|
||||
return (
|
||||
"done",
|
||||
[],
|
||||
@@ -1294,10 +1299,15 @@ async def test_system_subagent_followup_is_persisted_before_prompt_assembly(tmp_
|
||||
chat_id="cli:test",
|
||||
content="subagent result",
|
||||
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 "working" in non_system[1]["content"]
|
||||
# Persisted timestamps stay in session records, but replay content is not
|
||||
|
||||
@@ -23,10 +23,12 @@ class _ContextRecordingTool:
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.contexts: list[dict] = []
|
||||
self.runtimes: list[object] = []
|
||||
|
||||
async def execute(self, **_kwargs) -> str:
|
||||
ctx = current_request_context()
|
||||
assert ctx is not None
|
||||
self.runtimes.append(ctx.runtime)
|
||||
self.contexts.append({
|
||||
"channel": ctx.channel,
|
||||
"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)
|
||||
|
||||
metadata = {"slack": {"thread_ts": "111.222", "channel_type": "channel"}}
|
||||
runtime = loop.llm_runtime()
|
||||
await loop._run_agent_loop(
|
||||
[],
|
||||
runtime=runtime,
|
||||
channel="slack",
|
||||
chat_id="C123",
|
||||
metadata=metadata,
|
||||
@@ -95,6 +99,7 @@ async def test_loop_binds_request_context_for_tool_execution(tmp_path: Path) ->
|
||||
"metadata": metadata,
|
||||
"session_key": "slack:C123:111.222",
|
||||
}
|
||||
assert cron.runtimes[-1] is runtime
|
||||
|
||||
|
||||
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",
|
||||
)
|
||||
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()
|
||||
assert current is not None
|
||||
assert spec.runtime is runtime
|
||||
assert current.runtime is runtime
|
||||
assert current.channel == "slack"
|
||||
assert current.chat_id == "C123"
|
||||
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"):
|
||||
await loop._run_agent_loop(
|
||||
[],
|
||||
runtime=runtime,
|
||||
channel="slack",
|
||||
chat_id="C123",
|
||||
session_key="slack:C123:111.222",
|
||||
@@ -209,10 +218,11 @@ async def test_process_message_captures_original_text_before_restore(
|
||||
workspace=tmp_path,
|
||||
model="test-model",
|
||||
)
|
||||
seen: list[str | None] = []
|
||||
runtime = loop.llm_runtime()
|
||||
seen: list[tuple[str | None, object]] = []
|
||||
|
||||
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")
|
||||
|
||||
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",
|
||||
content=" original user text ",
|
||||
metadata=metadata,
|
||||
)
|
||||
),
|
||||
runtime=runtime,
|
||||
)
|
||||
|
||||
assert seen == [expected]
|
||||
assert seen == [(expected, runtime)]
|
||||
|
||||
@@ -446,6 +446,7 @@ async def test_loop_injected_followup_preserves_image_media(tmp_path):
|
||||
|
||||
final_content, _, _, _, had_injections = await loop._run_agent_loop(
|
||||
[{"role": "user", "content": "hello"}],
|
||||
runtime=loop.llm_runtime(),
|
||||
channel="cli",
|
||||
chat_id="c",
|
||||
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(
|
||||
[{"role": "user", "content": "hello"}],
|
||||
runtime=loop.llm_runtime(),
|
||||
channel="cli",
|
||||
chat_id="c",
|
||||
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(
|
||||
[{"role": "user", "content": "hello"}],
|
||||
runtime=loop.llm_runtime(),
|
||||
channel="cli",
|
||||
chat_id="c",
|
||||
pending_queue=pending_queue,
|
||||
|
||||
@@ -6,6 +6,7 @@ from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.config.loader import save_config
|
||||
from nanobot.config.schema import Config
|
||||
from nanobot.providers.base import GenerationSettings
|
||||
from nanobot.providers.factory import ProviderSnapshot, load_provider_snapshot
|
||||
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")
|
||||
|
||||
|
||||
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(
|
||||
tmp_path: Path,
|
||||
monkeypatch,
|
||||
|
||||
@@ -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.max_iterations = 55
|
||||
|
||||
await loop._run_agent_loop([])
|
||||
await loop._run_agent_loop([], runtime=loop.llm_runtime())
|
||||
|
||||
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
|
||||
await loop._run_agent_loop(
|
||||
[{"role": "user", "content": "test"}],
|
||||
runtime=loop.llm_runtime(),
|
||||
session=session,
|
||||
channel="test",
|
||||
chat_id="c1",
|
||||
@@ -402,6 +403,7 @@ async def test_drain_pending_no_block_when_no_subagents(tmp_path):
|
||||
|
||||
await loop._run_agent_loop(
|
||||
[{"role": "user", "content": "test"}],
|
||||
runtime=loop.llm_runtime(),
|
||||
session=None,
|
||||
channel="test",
|
||||
chat_id="c1",
|
||||
@@ -458,6 +460,7 @@ async def test_drain_pending_timeout(tmp_path):
|
||||
|
||||
await loop._run_agent_loop(
|
||||
[{"role": "user", "content": "test"}],
|
||||
runtime=loop.llm_runtime(),
|
||||
session=session,
|
||||
channel="test",
|
||||
chat_id="c1",
|
||||
|
||||
@@ -296,11 +296,11 @@ class TestRestartCommand:
|
||||
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["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["completion_tokens"] == 7
|
||||
assert loop._last_usage["estimated_tokens"] == 130
|
||||
|
||||
@@ -144,7 +144,9 @@ class TestMessageToolSuppressLogic:
|
||||
async def on_progress(content: str, *, tool_hint: bool = False) -> None:
|
||||
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 progress == [
|
||||
|
||||
Reference in New Issue
Block a user