diff --git a/nanobot/agent/loop.py b/nanobot/agent/loop.py index 734b0789..d5351c69 100644 --- a/nanobot/agent/loop.py +++ b/nanobot/agent/loop.py @@ -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, diff --git a/nanobot/agent/tools/context.py b/nanobot/agent/tools/context.py index bf80493b..7973f1d8 100644 --- a/nanobot/agent/tools/context.py +++ b/nanobot/agent/tools/context.py @@ -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) diff --git a/tests/agent/test_document_extraction_toggle.py b/tests/agent/test_document_extraction_toggle.py index 67e566cf..fc839b73 100644 --- a/tests/agent/test_document_extraction_toggle.py +++ b/tests/agent/test_document_extraction_toggle.py @@ -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, diff --git a/tests/agent/test_hook_composite.py b/tests/agent/test_hook_composite.py index e0720e1b..51bca697 100644 --- a/tests/agent/test_hook_composite.py +++ b/tests/agent/test_hook_composite.py @@ -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." diff --git a/tests/agent/test_loop_progress.py b/tests/agent/test_loop_progress.py index 54ef783e..33590405 100644 --- a/tests/agent/test_loop_progress.py +++ b/tests/agent/test_loop_progress.py @@ -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, ) diff --git a/tests/agent/test_loop_runner_integration.py b/tests/agent/test_loop_runner_integration.py index bcd727a6..f39523ec 100644 --- a/tests/agent/test_loop_runner_integration.py +++ b/tests/agent/test_loop_runner_integration.py @@ -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 diff --git a/tests/agent/test_loop_save_turn.py b/tests/agent/test_loop_save_turn.py index 7be00b48..80711270 100644 --- a/tests/agent/test_loop_save_turn.py +++ b/tests/agent/test_loop_save_turn.py @@ -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 diff --git a/tests/agent/test_loop_tool_context.py b/tests/agent/test_loop_tool_context.py index 101f580a..83cfcd10 100644 --- a/tests/agent/test_loop_tool_context.py +++ b/tests/agent/test_loop_tool_context.py @@ -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)] diff --git a/tests/agent/test_runner_injections.py b/tests/agent/test_runner_injections.py index e78e8dd6..e6522671 100644 --- a/tests/agent/test_runner_injections.py +++ b/tests/agent/test_runner_injections.py @@ -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, diff --git a/tests/agent/test_runtime_refresh.py b/tests/agent/test_runtime_refresh.py index ede30d36..a4b8a61d 100644 --- a/tests/agent/test_runtime_refresh.py +++ b/tests/agent/test_runtime_refresh.py @@ -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, diff --git a/tests/agent/tools/test_subagent_tools.py b/tests/agent/tools/test_subagent_tools.py index 3b1aaa64..bc9ceb06 100644 --- a/tests/agent/tools/test_subagent_tools.py +++ b/tests/agent/tools/test_subagent_tools.py @@ -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", diff --git a/tests/cli/test_restart_command.py b/tests/cli/test_restart_command.py index f75a92f3..9b5012cb 100644 --- a/tests/cli/test_restart_command.py +++ b/tests/cli/test_restart_command.py @@ -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 diff --git a/tests/tools/test_message_tool_suppress.py b/tests/tools/test_message_tool_suppress.py index b06f5413..4e1542cc 100644 --- a/tests/tools/test_message_tool_suppress.py +++ b/tests/tools/test_message_tool_suppress.py @@ -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 == [