refactor(agent): make resolver sole runtime owner
This commit is contained in:
+83
-125
@@ -6,6 +6,7 @@ import asyncio
|
||||
import dataclasses
|
||||
import os
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from contextlib import AsyncExitStack, nullcontext, suppress
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum, auto
|
||||
@@ -174,37 +175,48 @@ class AgentLoop:
|
||||
def tool_names(self) -> list[str]:
|
||||
return self.tools.tool_names
|
||||
|
||||
@property
|
||||
def provider(self) -> LLMProvider:
|
||||
"""Provider selected for future turn admissions."""
|
||||
return self.runtime_resolver.runtime.provider
|
||||
|
||||
@property
|
||||
def model(self) -> str:
|
||||
"""Model selected for future turn admissions."""
|
||||
return self.runtime_resolver.runtime.model
|
||||
|
||||
@property
|
||||
def context_window_tokens(self) -> int:
|
||||
"""Context limit selected for future turn admissions."""
|
||||
return self.runtime_resolver.runtime.context_window_tokens
|
||||
|
||||
@property
|
||||
def model_presets(self) -> Mapping[str, ModelPresetConfig]:
|
||||
"""Configured model presets exposed for selection and display."""
|
||||
return self.runtime_resolver.model_presets
|
||||
|
||||
@property
|
||||
def model_preset(self) -> str | None:
|
||||
return self.runtime_resolver.model_preset
|
||||
|
||||
@model_preset.setter
|
||||
def model_preset(self, name: str | None) -> None:
|
||||
self.set_model_preset(name)
|
||||
|
||||
def llm_runtime(self) -> LLMRuntime:
|
||||
"""Resolve the immutable default used to admit the next turn."""
|
||||
self._refresh_provider_snapshot()
|
||||
runtime = self.runtime_resolver.current()
|
||||
captured = LLMRuntime.capture(
|
||||
self.provider,
|
||||
self.model,
|
||||
context_window_tokens=self.context_window_tokens,
|
||||
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.
|
||||
previous = self.runtime_resolver.runtime
|
||||
try:
|
||||
runtime = self.runtime_resolver.current(refresh=True)
|
||||
except Exception:
|
||||
logger.exception("Failed to refresh model runtime")
|
||||
return previous
|
||||
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
|
||||
runtime.model != previous.model
|
||||
or runtime.model_preset != previous.model_preset
|
||||
or runtime.snapshot_signature != previous.snapshot_signature
|
||||
):
|
||||
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,
|
||||
)
|
||||
self._publish_runtime_selection(runtime)
|
||||
return runtime
|
||||
|
||||
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
|
||||
@@ -271,32 +283,26 @@ class AgentLoop:
|
||||
self.runtime_event_publisher = RuntimeEventPublisher(self.runtime_events)
|
||||
self.channels_config = channels_config
|
||||
self.restart_mode = restart_mode
|
||||
self.provider = provider
|
||||
self._provider_snapshot_loader = provider_snapshot_loader
|
||||
self._preset_snapshot_loader = preset_snapshot_loader
|
||||
self._runtime_model_publisher = runtime_model_publisher
|
||||
self._provider_signature = provider_signature
|
||||
self._default_selection_signature = preset_helpers.default_selection_signature(provider_signature)
|
||||
self.workspace = workspace
|
||||
self.model = model or provider.get_default_model()
|
||||
initial_model = model or provider.get_default_model()
|
||||
self.max_iterations = (
|
||||
max_iterations if max_iterations is not None else defaults.max_tool_iterations
|
||||
)
|
||||
self.context_window_tokens = (
|
||||
initial_context_window = (
|
||||
context_window_tokens
|
||||
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
|
||||
configured_presets = model_presets or {}
|
||||
self.runtime_resolver = ModelRuntimeResolver(
|
||||
LLMRuntime.capture(
|
||||
provider,
|
||||
self.model,
|
||||
context_window_tokens=self.context_window_tokens,
|
||||
initial_model,
|
||||
context_window_tokens=initial_context_window,
|
||||
snapshot_signature=provider_signature,
|
||||
),
|
||||
model_presets=self.model_presets,
|
||||
model_presets=configured_presets,
|
||||
provider_snapshot_loader=provider_snapshot_loader,
|
||||
preset_snapshot_loader=preset_snapshot_loader,
|
||||
)
|
||||
@@ -352,7 +358,6 @@ class AgentLoop:
|
||||
llm_wall_timeout_for_session=lambda sk: runner_wall_llm_timeout_s(self.sessions, sk),
|
||||
)
|
||||
self._unified_session = unified_session
|
||||
self._max_messages = replay_max_messages_for_context(self.context_window_tokens)
|
||||
self._running = False
|
||||
self._mcp_servers = mcp_servers or {}
|
||||
self._mcp_stacks: dict[str, AsyncExitStack] = {}
|
||||
@@ -401,7 +406,7 @@ class AgentLoop:
|
||||
)
|
||||
if model_preset:
|
||||
self.set_model_preset(model_preset, publish_update=False)
|
||||
self._register_default_tools()
|
||||
self._register_default_tools(provider_snapshot_loader=provider_snapshot_loader)
|
||||
self._runtime_vars: dict[str, Any] = {}
|
||||
self._current_iteration: int = 0
|
||||
self.commands = CommandRouter()
|
||||
@@ -468,98 +473,51 @@ class AgentLoop:
|
||||
"""Keep subagent runtime limits aligned with mutable loop settings."""
|
||||
self.subagents.max_iterations = self.max_iterations
|
||||
|
||||
def _apply_provider_snapshot(
|
||||
def _publish_runtime_selection(
|
||||
self,
|
||||
snapshot: ProviderSnapshot,
|
||||
runtime: LLMRuntime,
|
||||
*,
|
||||
publish_update: bool = True,
|
||||
model_preset: str | None = None,
|
||||
) -> None:
|
||||
"""Swap model/provider for future turns without disturbing an active one."""
|
||||
runtime = self.runtime_resolver.adopt_snapshot(
|
||||
snapshot,
|
||||
model_preset=model_preset,
|
||||
if not publish_update:
|
||||
return
|
||||
if self._runtime_model_publisher is not None:
|
||||
self._runtime_model_publisher(runtime.model, runtime.model_preset)
|
||||
self._runtime_events().runtime_model_changed(
|
||||
runtime.model,
|
||||
runtime.model_preset,
|
||||
)
|
||||
provider = runtime.provider
|
||||
model = runtime.model
|
||||
context_window_tokens = runtime.context_window_tokens
|
||||
provider.generation = runtime.generation
|
||||
|
||||
def set_model_preset(
|
||||
self,
|
||||
name: str | None,
|
||||
*,
|
||||
publish_update: bool = True,
|
||||
) -> LLMRuntime:
|
||||
"""Select a named default runtime for future turns."""
|
||||
old_model = self.model
|
||||
self.provider = provider
|
||||
self.model = model
|
||||
self.context_window_tokens = context_window_tokens
|
||||
self._sync_replay_max_messages()
|
||||
self._provider_signature = snapshot.signature
|
||||
if publish_update and self._runtime_model_publisher is not None:
|
||||
self._runtime_model_publisher(
|
||||
self.model,
|
||||
model_preset if model_preset is not None else self.model_preset,
|
||||
)
|
||||
if publish_update:
|
||||
self._runtime_events().runtime_model_changed(
|
||||
self.model,
|
||||
model_preset if model_preset is not None else self.model_preset,
|
||||
)
|
||||
logger.info("Runtime model switched for next turn: {} -> {}", old_model, model)
|
||||
|
||||
def _sync_replay_max_messages(self) -> None:
|
||||
self._max_messages = replay_max_messages_for_context(self.context_window_tokens)
|
||||
|
||||
def _refresh_provider_snapshot(self) -> None:
|
||||
if self._provider_snapshot_loader is None:
|
||||
return
|
||||
try:
|
||||
snapshot = self._provider_snapshot_loader()
|
||||
except Exception:
|
||||
logger.exception("Failed to refresh provider config")
|
||||
return
|
||||
default_selection = preset_helpers.default_selection_signature(snapshot.signature)
|
||||
if self._active_preset and self._default_selection_signature in (None, default_selection):
|
||||
self._default_selection_signature = default_selection
|
||||
try:
|
||||
snapshot = self._build_model_preset_snapshot(self._active_preset)
|
||||
except Exception:
|
||||
logger.exception("Failed to refresh active model preset")
|
||||
return
|
||||
else:
|
||||
self._active_preset = None
|
||||
self._default_selection_signature = default_selection
|
||||
if snapshot.signature == self._provider_signature:
|
||||
return
|
||||
self._default_selection_signature = preset_helpers.default_selection_signature(snapshot.signature)
|
||||
self._apply_provider_snapshot(snapshot)
|
||||
|
||||
@property
|
||||
def model_preset(self) -> str | None:
|
||||
return self._active_preset
|
||||
|
||||
@model_preset.setter
|
||||
def model_preset(self, name: str | None) -> None:
|
||||
self.set_model_preset(name)
|
||||
|
||||
def _build_model_preset_snapshot(self, name: str) -> ProviderSnapshot:
|
||||
return preset_helpers.build_runtime_preset_snapshot(
|
||||
name=name,
|
||||
presets=self.model_presets,
|
||||
provider=self.provider,
|
||||
loader=self._preset_snapshot_loader,
|
||||
)
|
||||
|
||||
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)
|
||||
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._publish_runtime_selection(runtime, publish_update=publish_update)
|
||||
logger.info(
|
||||
"Runtime model switched for next turn: {} -> {}",
|
||||
old_model,
|
||||
runtime.model,
|
||||
)
|
||||
self._apply_provider_snapshot(snapshot, publish_update=publish_update, model_preset=name)
|
||||
self._active_preset = name
|
||||
return runtime
|
||||
|
||||
def _register_default_tools(self) -> None:
|
||||
def set_runtime_model(self, model: str) -> LLMRuntime:
|
||||
"""Select a model on the current provider for future turns."""
|
||||
return self.runtime_resolver.select_model(model)
|
||||
|
||||
def set_runtime_context_window(self, context_window_tokens: int) -> LLMRuntime:
|
||||
"""Select a context limit for future turns."""
|
||||
return self.runtime_resolver.select_context_window(context_window_tokens)
|
||||
|
||||
def _register_default_tools(
|
||||
self,
|
||||
*,
|
||||
provider_snapshot_loader: Callable[..., ProviderSnapshot] | None,
|
||||
) -> None:
|
||||
"""Register the default set of tools via plugin loader."""
|
||||
from nanobot.agent.tools.context import ToolContext
|
||||
from nanobot.agent.tools.loader import ToolLoader
|
||||
@@ -571,7 +529,7 @@ class AgentLoop:
|
||||
subagent_manager=self.subagents,
|
||||
cron_service=self.cron_service,
|
||||
sessions=self.sessions,
|
||||
provider_snapshot_loader=self._provider_snapshot_loader,
|
||||
provider_snapshot_loader=provider_snapshot_loader,
|
||||
image_generation_provider_configs=self._image_generation_provider_configs,
|
||||
timezone=self.context.timezone or "UTC",
|
||||
workspace_sandbox=self.workspace_scopes.sandbox_status,
|
||||
|
||||
@@ -56,7 +56,9 @@ class RuntimeState(Protocol):
|
||||
|
||||
def _sync_subagent_runtime_limits(self) -> None: ...
|
||||
|
||||
def set_runtime_model(self, model: str) -> Any: ...
|
||||
|
||||
def set_runtime_context_window(self, context_window_tokens: int) -> Any: ...
|
||||
|
||||
@property
|
||||
def model_preset(self) -> str | None: ...
|
||||
|
||||
_active_preset: str | None
|
||||
|
||||
@@ -57,7 +57,7 @@ class MyTool(Tool):
|
||||
|
||||
BLOCKED = frozenset({
|
||||
# Core infrastructure
|
||||
"bus", "provider", "_running", "tools",
|
||||
"bus", "provider", "runtime_resolver", "_running", "tools",
|
||||
# Config management
|
||||
"_runtime_vars",
|
||||
# Subsystems
|
||||
@@ -107,6 +107,11 @@ class MyTool(Tool):
|
||||
}
|
||||
|
||||
_MAX_RUNTIME_KEYS = 64
|
||||
_MODEL_RUNTIME_FIELDS = frozenset({
|
||||
"model",
|
||||
"model_preset",
|
||||
"context_window_tokens",
|
||||
})
|
||||
|
||||
def __init__(self, runtime_state: RuntimeState, modify_allowed: bool = True) -> None:
|
||||
self._runtime_state = runtime_state
|
||||
@@ -325,9 +330,20 @@ class MyTool(Tool):
|
||||
|
||||
# -- inspect --
|
||||
|
||||
def _current_runtime_value(self, key: str) -> tuple[bool, Any]:
|
||||
request_ctx = current_request_context()
|
||||
runtime = request_ctx.runtime if request_ctx is not None else None
|
||||
if runtime is None or key not in self._MODEL_RUNTIME_FIELDS:
|
||||
return False, None
|
||||
return True, getattr(runtime, key)
|
||||
|
||||
def _inspect(self, key: str | None) -> str:
|
||||
if not key:
|
||||
return self._inspect_all()
|
||||
if "." not in key:
|
||||
found, value = self._current_runtime_value(key)
|
||||
if found:
|
||||
return self._format_value(value, key)
|
||||
top = key.split(".")[0]
|
||||
if top in self._DENIED_ATTRS or top.startswith("__"):
|
||||
return ToolResult.error(f"Error: '{top}' is not accessible")
|
||||
@@ -353,8 +369,13 @@ class MyTool(Tool):
|
||||
parts: list[str] = []
|
||||
# RESTRICTED keys
|
||||
for k in self.RESTRICTED:
|
||||
parts.append(self._format_value(getattr(state, k, None), k))
|
||||
parts.append(self._format_value(state.model_preset, "model_preset"))
|
||||
found, value = self._current_runtime_value(k)
|
||||
parts.append(self._format_value(value if found else getattr(state, k, None), k))
|
||||
found, value = self._current_runtime_value("model_preset")
|
||||
parts.append(self._format_value(
|
||||
value if found else state.model_preset,
|
||||
"model_preset",
|
||||
))
|
||||
# Other useful top-level keys shown in description
|
||||
for k in ("workspace", "provider_retry_mode", "max_tool_result_chars", "_current_iteration", "web_config", "exec_config", "workspace_sandbox", "subagents"):
|
||||
if _has_real_attr(state, k):
|
||||
@@ -432,13 +453,16 @@ class MyTool(Tool):
|
||||
return ToolResult.error(f"Error: '{key}' must be <= {spec['max']}")
|
||||
if "min_len" in spec and len(str(value)) < spec["min_len"]:
|
||||
return ToolResult.error(f"Error: '{key}' must be at least {spec['min_len']} characters")
|
||||
setattr(self._runtime_state, key, value)
|
||||
if key == "model":
|
||||
self._runtime_state._active_preset = None
|
||||
sync_replay = getattr(self._runtime_state, "_sync_replay_max_messages", None)
|
||||
if key == "context_window_tokens" and callable(sync_replay):
|
||||
sync_replay()
|
||||
if key == "max_iterations" and hasattr(self._runtime_state, "_sync_subagent_runtime_limits"):
|
||||
self._runtime_state.set_runtime_model(value)
|
||||
elif key == "context_window_tokens":
|
||||
self._runtime_state.set_runtime_context_window(value)
|
||||
else:
|
||||
setattr(self._runtime_state, key, value)
|
||||
if key == "max_iterations" and hasattr(
|
||||
self._runtime_state,
|
||||
"_sync_subagent_runtime_limits",
|
||||
):
|
||||
self._runtime_state._sync_subagent_runtime_limits()
|
||||
self._audit("modify", f"{key}: {old!r} -> {value!r}")
|
||||
return f"Set {key} = {value!r} (was {old!r})"
|
||||
|
||||
@@ -349,7 +349,7 @@ async def cmd_model(ctx: CommandContext) -> OutboundMessage:
|
||||
|
||||
name = parts[0]
|
||||
try:
|
||||
loop.set_model_preset(name)
|
||||
runtime = loop.set_model_preset(name)
|
||||
except (KeyError, ValueError) as exc:
|
||||
names = _model_preset_names(loop)
|
||||
return OutboundMessage(
|
||||
@@ -362,11 +362,11 @@ async def cmd_model(ctx: CommandContext) -> OutboundMessage:
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
max_tokens = getattr(getattr(loop.provider, "generation", None), "max_tokens", None)
|
||||
max_tokens = runtime.generation.max_tokens
|
||||
lines = [
|
||||
f"Switched model preset to `{loop.model_preset}`.",
|
||||
f"- Model: `{loop.model}`",
|
||||
f"- Context window: {loop.context_window_tokens}",
|
||||
f"Switched model preset to `{runtime.model_preset}`.",
|
||||
f"- Model: `{runtime.model}`",
|
||||
f"- Context window: {runtime.context_window_tokens}",
|
||||
]
|
||||
if max_tokens is not None:
|
||||
lines.append(f"- Max output tokens: {max_tokens}")
|
||||
|
||||
+69
-63
@@ -14,7 +14,6 @@ from nanobot.config.schema import Config
|
||||
from nanobot.providers.image_generation import image_gen_provider_configs
|
||||
from nanobot.sdk.clients import MemoryClient, RuntimeClient, SessionClient
|
||||
from nanobot.sdk.runtime import (
|
||||
SDKRuntimeController,
|
||||
build_process_direct_kwargs,
|
||||
ensure_single_model_selector,
|
||||
)
|
||||
@@ -74,7 +73,6 @@ class Nanobot:
|
||||
def __init__(self, loop: AgentLoop, *, config: Config | None = None) -> None:
|
||||
self._loop = loop
|
||||
self._config = config
|
||||
self._runtime_overrides = SDKRuntimeController(loop, config=config)
|
||||
self.sessions = SessionClient(loop)
|
||||
self.memory = MemoryClient(loop)
|
||||
self.runtime = RuntimeClient(loop)
|
||||
@@ -156,20 +154,26 @@ class Nanobot:
|
||||
"""
|
||||
capture = SDKCaptureHook()
|
||||
per_run_hooks = [capture, *(hooks or [])]
|
||||
async with self._runtime_overrides.override(model=model, model_preset=model_preset):
|
||||
kwargs = build_process_direct_kwargs(
|
||||
session_key=session_key,
|
||||
channel=channel,
|
||||
chat_id=chat_id,
|
||||
sender_id=sender_id,
|
||||
media=media,
|
||||
ephemeral=ephemeral,
|
||||
)
|
||||
response = await self._loop.process_direct(
|
||||
message,
|
||||
**kwargs,
|
||||
hooks=per_run_hooks,
|
||||
)
|
||||
runtime = self._loop.runtime_resolver.resolve_override(
|
||||
model=model,
|
||||
model_preset=model_preset,
|
||||
config=self._config,
|
||||
)
|
||||
kwargs = build_process_direct_kwargs(
|
||||
session_key=session_key,
|
||||
channel=channel,
|
||||
chat_id=chat_id,
|
||||
sender_id=sender_id,
|
||||
media=media,
|
||||
ephemeral=ephemeral,
|
||||
)
|
||||
if runtime is not None:
|
||||
kwargs["runtime"] = runtime
|
||||
response = await self._loop.process_direct(
|
||||
message,
|
||||
**kwargs,
|
||||
hooks=per_run_hooks,
|
||||
)
|
||||
|
||||
return result_from_response(response, capture)
|
||||
|
||||
@@ -188,7 +192,11 @@ class Nanobot:
|
||||
model_preset: str | None = None,
|
||||
) -> RunStream:
|
||||
"""Start a streamed run and return a handle for events and final result."""
|
||||
ensure_single_model_selector(model=model, model_preset=model_preset)
|
||||
runtime = self._loop.runtime_resolver.resolve_override(
|
||||
model=model,
|
||||
model_preset=model_preset,
|
||||
config=self._config,
|
||||
) or self._loop.llm_runtime()
|
||||
queue: asyncio.Queue[StreamEvent | object] = asyncio.Queue(maxsize=256)
|
||||
emitter = SDKStreamEmitter(queue)
|
||||
stream_hook = SDKStreamingHook(emitter)
|
||||
@@ -202,55 +210,53 @@ class Nanobot:
|
||||
await emitter.text_completed(resuming=resuming)
|
||||
|
||||
async def _run() -> RunResult:
|
||||
async with self._runtime_overrides.override(model=model, model_preset=model_preset):
|
||||
kwargs = build_process_direct_kwargs(
|
||||
session_key=session_key,
|
||||
channel=channel,
|
||||
chat_id=chat_id,
|
||||
sender_id=sender_id,
|
||||
media=media,
|
||||
ephemeral=ephemeral,
|
||||
on_stream=_on_stream,
|
||||
on_stream_end=_on_stream_end,
|
||||
kwargs = build_process_direct_kwargs(
|
||||
session_key=session_key,
|
||||
channel=channel,
|
||||
chat_id=chat_id,
|
||||
sender_id=sender_id,
|
||||
media=media,
|
||||
ephemeral=ephemeral,
|
||||
on_stream=_on_stream,
|
||||
on_stream_end=_on_stream_end,
|
||||
)
|
||||
kwargs["runtime"] = runtime
|
||||
await emitter.emit(StreamEvent(
|
||||
type=STREAM_EVENT_RUN_STARTED,
|
||||
metadata={
|
||||
"session_key": session_key,
|
||||
"channel": channel,
|
||||
"chat_id": chat_id,
|
||||
"sender_id": sender_id,
|
||||
"model": runtime.model,
|
||||
"model_preset": runtime.model_preset,
|
||||
},
|
||||
))
|
||||
try:
|
||||
response = await self._loop.process_direct(
|
||||
message,
|
||||
**kwargs,
|
||||
hooks=per_run_hooks,
|
||||
)
|
||||
await emitter.text_completed(resuming=False, force=False)
|
||||
result = result_from_response(response, capture)
|
||||
await emitter.emit(StreamEvent(
|
||||
type=STREAM_EVENT_RUN_STARTED,
|
||||
metadata={
|
||||
"session_key": session_key,
|
||||
"channel": channel,
|
||||
"chat_id": chat_id,
|
||||
"sender_id": sender_id,
|
||||
"model": self._loop.model,
|
||||
"model_preset": (
|
||||
model_preset if model_preset is not None else self._loop.model_preset
|
||||
),
|
||||
},
|
||||
type=STREAM_EVENT_RUN_COMPLETED,
|
||||
content=result.content,
|
||||
result=result,
|
||||
usage=dict(result.usage),
|
||||
metadata=dict(result.metadata),
|
||||
))
|
||||
try:
|
||||
response = await self._loop.process_direct(
|
||||
message,
|
||||
**kwargs,
|
||||
hooks=per_run_hooks,
|
||||
)
|
||||
await emitter.text_completed(resuming=False, force=False)
|
||||
result = result_from_response(response, capture)
|
||||
await emitter.emit(StreamEvent(
|
||||
type=STREAM_EVENT_RUN_COMPLETED,
|
||||
content=result.content,
|
||||
result=result,
|
||||
usage=dict(result.usage),
|
||||
metadata=dict(result.metadata),
|
||||
))
|
||||
return result
|
||||
except Exception as exc:
|
||||
await emitter.emit(StreamEvent(
|
||||
type=STREAM_EVENT_RUN_FAILED,
|
||||
error=str(exc),
|
||||
metadata={"exception_type": type(exc).__name__},
|
||||
))
|
||||
raise
|
||||
finally:
|
||||
emitter.close()
|
||||
return result
|
||||
except Exception as exc:
|
||||
await emitter.emit(StreamEvent(
|
||||
type=STREAM_EVENT_RUN_FAILED,
|
||||
error=str(exc),
|
||||
metadata={"exception_type": type(exc).__name__},
|
||||
))
|
||||
raise
|
||||
finally:
|
||||
emitter.close()
|
||||
|
||||
task = asyncio.create_task(_run())
|
||||
return RunStream(task, queue)
|
||||
|
||||
+1
-149
@@ -2,16 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from nanobot.config.schema import Config, ModelPresetConfig
|
||||
from nanobot.providers.factory import ProviderSnapshot, build_provider_snapshot
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from typing import Any
|
||||
|
||||
|
||||
def ensure_single_model_selector(
|
||||
@@ -51,142 +42,3 @@ def build_process_direct_kwargs(
|
||||
if on_stream_end is not None:
|
||||
kwargs["on_stream_end"] = on_stream_end
|
||||
return kwargs
|
||||
|
||||
|
||||
class SDKRuntimeGate:
|
||||
"""Allow normal SDK runs to overlap while model overrides stay exclusive."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._condition = asyncio.Condition()
|
||||
self._readers = 0
|
||||
self._writer_active = False
|
||||
self._writers_waiting = 0
|
||||
|
||||
def slot(self, *, exclusive: bool) -> SDKRuntimeGateSlot:
|
||||
return SDKRuntimeGateSlot(self, exclusive=exclusive)
|
||||
|
||||
async def _acquire(self, *, exclusive: bool) -> None:
|
||||
async with self._condition:
|
||||
if exclusive:
|
||||
self._writers_waiting += 1
|
||||
try:
|
||||
await self._condition.wait_for(
|
||||
lambda: not self._writer_active and self._readers == 0
|
||||
)
|
||||
self._writer_active = True
|
||||
finally:
|
||||
self._writers_waiting -= 1
|
||||
self._condition.notify_all()
|
||||
return
|
||||
|
||||
await self._condition.wait_for(
|
||||
lambda: not self._writer_active and self._writers_waiting == 0
|
||||
)
|
||||
self._readers += 1
|
||||
|
||||
async def _release(self, *, exclusive: bool) -> None:
|
||||
async with self._condition:
|
||||
if exclusive:
|
||||
self._writer_active = False
|
||||
else:
|
||||
self._readers = max(0, self._readers - 1)
|
||||
self._condition.notify_all()
|
||||
|
||||
|
||||
class SDKRuntimeGateSlot:
|
||||
def __init__(self, gate: SDKRuntimeGate, *, exclusive: bool) -> None:
|
||||
self._gate = gate
|
||||
self._exclusive = exclusive
|
||||
|
||||
async def __aenter__(self) -> None:
|
||||
await self._gate._acquire(exclusive=self._exclusive)
|
||||
|
||||
async def __aexit__(self, *exc: object) -> None:
|
||||
await self._gate._release(exclusive=self._exclusive)
|
||||
|
||||
|
||||
class SDKRuntimeController:
|
||||
"""Apply per-run SDK model overrides without leaking global runtime state."""
|
||||
|
||||
def __init__(self, loop: AgentLoop, *, config: Config | None = None) -> None:
|
||||
self._loop = loop
|
||||
self._config = config
|
||||
self._gate = SDKRuntimeGate()
|
||||
|
||||
@asynccontextmanager
|
||||
async def override(
|
||||
self,
|
||||
*,
|
||||
model: str | None,
|
||||
model_preset: str | None,
|
||||
) -> AsyncIterator[None]:
|
||||
ensure_single_model_selector(model=model, model_preset=model_preset)
|
||||
exclusive = model is not None or model_preset is not None
|
||||
async with self._gate.slot(exclusive=exclusive):
|
||||
override = self.model_override_snapshot(model=model, model_preset=model_preset)
|
||||
restore = self._current_snapshot() if override is not None else None
|
||||
restore_signature = self._loop._provider_signature
|
||||
if override is not None:
|
||||
self._loop._apply_provider_snapshot(
|
||||
override,
|
||||
publish_update=False,
|
||||
model_preset=model_preset,
|
||||
)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
if restore is not None:
|
||||
self._restore_snapshot(
|
||||
restore,
|
||||
provider_signature=restore_signature,
|
||||
)
|
||||
|
||||
def model_override_snapshot(
|
||||
self,
|
||||
*,
|
||||
model: str | None,
|
||||
model_preset: str | None,
|
||||
) -> ProviderSnapshot | None:
|
||||
ensure_single_model_selector(model=model, model_preset=model_preset)
|
||||
if model_preset is not None:
|
||||
return self._loop._build_model_preset_snapshot(model_preset)
|
||||
if model is None:
|
||||
return None
|
||||
|
||||
if self._config is not None:
|
||||
base = self._config.resolve_preset(self._loop.model_preset)
|
||||
preset = base.model_copy(update={"model": model, "provider": "auto"})
|
||||
return build_provider_snapshot(self._config, preset=preset)
|
||||
|
||||
generation = getattr(self._loop.provider, "generation", None)
|
||||
preset = ModelPresetConfig(
|
||||
model=model,
|
||||
provider="auto",
|
||||
max_tokens=getattr(generation, "max_tokens", 8192),
|
||||
context_window_tokens=self._loop.context_window_tokens,
|
||||
temperature=getattr(generation, "temperature", 0.1),
|
||||
reasoning_effort=getattr(generation, "reasoning_effort", None),
|
||||
)
|
||||
from nanobot.agent.model_presets import build_static_preset_snapshot
|
||||
|
||||
return build_static_preset_snapshot(self._loop.provider, "sdk:override", preset)
|
||||
|
||||
def _current_snapshot(self) -> ProviderSnapshot:
|
||||
signature = self._loop._provider_signature
|
||||
if signature is None:
|
||||
signature = ("sdk:runtime", id(self._loop.provider), self._loop.model)
|
||||
return ProviderSnapshot(
|
||||
provider=self._loop.provider,
|
||||
model=self._loop.model,
|
||||
context_window_tokens=self._loop.context_window_tokens,
|
||||
signature=signature,
|
||||
)
|
||||
|
||||
def _restore_snapshot(
|
||||
self,
|
||||
snapshot: ProviderSnapshot,
|
||||
*,
|
||||
provider_signature: tuple[object, ...] | None,
|
||||
) -> None:
|
||||
self._loop._apply_provider_snapshot(snapshot, publish_update=False)
|
||||
self._loop._provider_signature = provider_signature
|
||||
|
||||
@@ -518,9 +518,11 @@ class TestNewCommandArchival:
|
||||
loop.sessions.save(session)
|
||||
|
||||
call_count = 0
|
||||
expected_runtime = loop.llm_runtime()
|
||||
|
||||
async def _failing_summarize(_messages, *, session_key=None) -> bool:
|
||||
async def _failing_summarize(_messages, *, runtime, session_key=None) -> bool:
|
||||
nonlocal call_count
|
||||
assert runtime is expected_runtime
|
||||
assert session_key == "cli:test"
|
||||
call_count += 1
|
||||
return False
|
||||
@@ -528,7 +530,7 @@ class TestNewCommandArchival:
|
||||
loop.consolidator.archive = _failing_summarize # type: ignore[method-assign]
|
||||
|
||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||
response = await loop._process_message(new_msg)
|
||||
response = await loop._process_message(new_msg, runtime=expected_runtime)
|
||||
|
||||
assert response is not None
|
||||
assert "new session started" in response.content.lower()
|
||||
@@ -553,9 +555,11 @@ class TestNewCommandArchival:
|
||||
|
||||
archived_count = -1
|
||||
archived_session_key = None
|
||||
expected_runtime = loop.llm_runtime()
|
||||
|
||||
async def _fake_summarize(messages, *, session_key=None) -> bool:
|
||||
async def _fake_summarize(messages, *, runtime, session_key=None) -> bool:
|
||||
nonlocal archived_count, archived_session_key
|
||||
assert runtime is expected_runtime
|
||||
archived_count = len(messages)
|
||||
archived_session_key = session_key
|
||||
return True
|
||||
@@ -563,7 +567,7 @@ class TestNewCommandArchival:
|
||||
loop.consolidator.archive = _fake_summarize # type: ignore[method-assign]
|
||||
|
||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||
response = await loop._process_message(new_msg)
|
||||
response = await loop._process_message(new_msg, runtime=expected_runtime)
|
||||
|
||||
assert response is not None
|
||||
assert "new session started" in response.content.lower()
|
||||
@@ -582,15 +586,17 @@ class TestNewCommandArchival:
|
||||
session.add_message("user", f"msg{i}")
|
||||
session.add_message("assistant", f"resp{i}")
|
||||
loop.sessions.save(session)
|
||||
expected_runtime = loop.llm_runtime()
|
||||
|
||||
async def _ok_summarize(_messages, *, session_key=None) -> bool:
|
||||
async def _ok_summarize(_messages, *, runtime, session_key=None) -> bool:
|
||||
assert runtime is expected_runtime
|
||||
assert session_key == "cli:test"
|
||||
return True
|
||||
|
||||
loop.consolidator.archive = _ok_summarize # type: ignore[method-assign]
|
||||
|
||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||
response = await loop._process_message(new_msg)
|
||||
response = await loop._process_message(new_msg, runtime=expected_runtime)
|
||||
|
||||
assert response is not None
|
||||
assert "new session started" in response.content.lower()
|
||||
@@ -610,8 +616,10 @@ class TestNewCommandArchival:
|
||||
|
||||
archived = asyncio.Event()
|
||||
release_archive = asyncio.Event()
|
||||
expected_runtime = loop.llm_runtime()
|
||||
|
||||
async def _slow_summarize(_messages, *, session_key=None) -> bool:
|
||||
async def _slow_summarize(_messages, *, runtime, session_key=None) -> bool:
|
||||
assert runtime is expected_runtime
|
||||
assert session_key == "cli:test"
|
||||
await release_archive.wait()
|
||||
archived.set()
|
||||
@@ -620,7 +628,7 @@ class TestNewCommandArchival:
|
||||
loop.consolidator.archive = _slow_summarize # type: ignore[method-assign]
|
||||
|
||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||
await loop._process_message(new_msg)
|
||||
await loop._process_message(new_msg, runtime=expected_runtime)
|
||||
|
||||
assert not archived.is_set()
|
||||
release_archive.set()
|
||||
|
||||
@@ -6,6 +6,7 @@ import nanobot.agent.memory as memory_module
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import LLMResponse
|
||||
from nanobot.session.manager import replay_max_messages_for_context
|
||||
|
||||
|
||||
def _make_loop(tmp_path, *, estimated_tokens: int, context_window_tokens: int) -> AgentLoop:
|
||||
@@ -222,7 +223,7 @@ async def test_preflight_consolidation_receives_pending_summary(tmp_path) -> Non
|
||||
loop.consolidator.maybe_consolidate_by_tokens.assert_any_await(
|
||||
session,
|
||||
runtime=runtime,
|
||||
replay_max_messages=loop._max_messages,
|
||||
replay_max_messages=replay_max_messages_for_context(runtime.context_window_tokens),
|
||||
)
|
||||
assert len(loop.consolidator.maybe_consolidate_by_tokens.call_args_list) == 2
|
||||
assert all(
|
||||
|
||||
@@ -21,6 +21,7 @@ from nanobot.bus.outbound_events import (
|
||||
)
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||
from nanobot.providers.factory import ProviderSnapshot
|
||||
from nanobot.session.webui_turns import WebuiTurnCoordinator
|
||||
from nanobot.utils.progress_events import (
|
||||
invoke_file_edit_progress,
|
||||
@@ -815,8 +816,14 @@ class TestToolEventProgress:
|
||||
))
|
||||
|
||||
assert len(scheduled_title) == 1
|
||||
loop.provider = MagicMock()
|
||||
loop.model = "switched-after-turn"
|
||||
next_provider = MagicMock()
|
||||
next_provider.generation = loop.llm_runtime().generation
|
||||
loop.runtime_resolver.adopt_snapshot(ProviderSnapshot(
|
||||
provider=next_provider,
|
||||
model="switched-after-turn",
|
||||
context_window_tokens=loop.context_window_tokens,
|
||||
signature=("switched-after-turn",),
|
||||
))
|
||||
|
||||
await scheduled_title[0] # type: ignore[misc]
|
||||
|
||||
|
||||
@@ -19,6 +19,7 @@ from nanobot.bus.outbound_events import (
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.cron.session_turns import CRON_HISTORY_META, CRON_TRIGGER_META
|
||||
from nanobot.providers.base import LLMResponse
|
||||
from nanobot.providers.factory import ProviderSnapshot
|
||||
from nanobot.session.automation_turns import AUTOMATION_HISTORY_META
|
||||
from nanobot.session.goal_state import GOAL_STATE_KEY
|
||||
from nanobot.session.manager import Session, SessionManager
|
||||
@@ -69,8 +70,17 @@ def test_agent_loop_llm_runtime_reflects_current_provider_and_model(tmp_path: Pa
|
||||
assert runtime.model == "test-model"
|
||||
|
||||
next_provider = MagicMock()
|
||||
loop.provider = next_provider
|
||||
loop.model = "next-model"
|
||||
next_provider.generation = SimpleNamespace(
|
||||
temperature=0.1,
|
||||
max_tokens=4096,
|
||||
reasoning_effort=None,
|
||||
)
|
||||
loop.runtime_resolver.adopt_snapshot(ProviderSnapshot(
|
||||
provider=next_provider,
|
||||
model="next-model",
|
||||
context_window_tokens=runtime.context_window_tokens,
|
||||
signature=("next-model",),
|
||||
))
|
||||
runtime = loop.llm_runtime()
|
||||
|
||||
assert runtime.provider is next_provider
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import replace
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
@@ -67,11 +68,13 @@ class TestMaxMessagesInit:
|
||||
|
||||
def test_default_for_200k_context_reaches_file_cap(self, tmp_path: Path) -> None:
|
||||
loop = _make_loop(tmp_path)
|
||||
assert loop._max_messages == FILE_MAX_MESSAGES
|
||||
runtime = loop.runtime_resolver.runtime
|
||||
assert replay_max_messages_for_context(runtime.context_window_tokens) == FILE_MAX_MESSAGES
|
||||
|
||||
def test_default_scales_with_context_window(self, tmp_path: Path) -> None:
|
||||
loop = _make_loop(tmp_path, context_window_tokens=32_768)
|
||||
assert loop._max_messages == 327
|
||||
runtime = loop.runtime_resolver.runtime
|
||||
assert replay_max_messages_for_context(runtime.context_window_tokens) == 327
|
||||
|
||||
def test_provider_refresh_resyncs_context_derived_limit(self, tmp_path: Path) -> None:
|
||||
old_provider = MagicMock()
|
||||
@@ -93,9 +96,10 @@ class TestMaxMessagesInit:
|
||||
),
|
||||
)
|
||||
|
||||
assert loop._max_messages == 327
|
||||
loop._refresh_provider_snapshot()
|
||||
assert loop._max_messages == FILE_MAX_MESSAGES
|
||||
initial = loop.runtime_resolver.runtime
|
||||
assert replay_max_messages_for_context(initial.context_window_tokens) == 327
|
||||
refreshed = loop.llm_runtime()
|
||||
assert replay_max_messages_for_context(refreshed.context_window_tokens) == FILE_MAX_MESSAGES
|
||||
|
||||
|
||||
class TestGetHistoryWithMaxMessages:
|
||||
@@ -136,7 +140,7 @@ class TestMaxMessagesIntegration:
|
||||
async def test_process_message_passes_limit_to_history_call(self, tmp_path: Path) -> None:
|
||||
"""The real message path should pass max_messages into session history replay."""
|
||||
loop = _make_loop(tmp_path)
|
||||
loop._max_messages = 25
|
||||
runtime = replace(loop.llm_runtime(), context_window_tokens=32_768)
|
||||
loop.provider.chat_with_retry = AsyncMock(
|
||||
return_value=LLMResponse(content="ok", tool_calls=[], usage={})
|
||||
)
|
||||
@@ -146,12 +150,13 @@ class TestMaxMessagesIntegration:
|
||||
session = loop.sessions.get_or_create("cli:test")
|
||||
with patch.object(session, "get_history", wraps=session.get_history) as mock_hist:
|
||||
result = await loop._process_message(
|
||||
InboundMessage(channel="cli", sender_id="user", chat_id="test", content="hello")
|
||||
InboundMessage(channel="cli", sender_id="user", chat_id="test", content="hello"),
|
||||
runtime=runtime,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert mock_hist.call_count == 1
|
||||
assert mock_hist.call_args.kwargs["max_messages"] == 25
|
||||
assert mock_hist.call_args.kwargs["max_messages"] == 327
|
||||
assert mock_hist.call_args.kwargs["extend_to_user"] is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -182,8 +187,7 @@ class TestMaxMessagesIntegration:
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""A live user turn should not extend history to an older long tool turn."""
|
||||
loop = _make_loop(tmp_path)
|
||||
loop._max_messages = 6
|
||||
loop = _make_loop(tmp_path, context_window_tokens=8_000)
|
||||
loop.provider.chat_with_retry = AsyncMock(
|
||||
return_value=LLMResponse(content="ok", tool_calls=[], usage={})
|
||||
)
|
||||
@@ -194,7 +198,7 @@ class TestMaxMessagesIntegration:
|
||||
session.add_message("user", "old")
|
||||
session.add_message("assistant", "old answer")
|
||||
session.add_message("user", "long older turn")
|
||||
for i in range(8):
|
||||
for i in range(70):
|
||||
session.messages.extend(_tool_round(f"older-{i}"))
|
||||
session.add_message("assistant", "older final")
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ def _provider(default_model: str, max_tokens: int = 123) -> MagicMock:
|
||||
return provider
|
||||
|
||||
|
||||
def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None:
|
||||
def test_provider_refresh_updates_only_runtime_resolver(tmp_path: Path) -> None:
|
||||
old_provider = _provider("old-model")
|
||||
new_provider = _provider("new-model", max_tokens=456)
|
||||
loop = AgentLoop(
|
||||
@@ -35,8 +35,9 @@ def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None:
|
||||
),
|
||||
)
|
||||
|
||||
loop._refresh_provider_snapshot()
|
||||
runtime = loop.llm_runtime()
|
||||
|
||||
assert runtime is loop.runtime_resolver.runtime
|
||||
assert loop.provider is new_provider
|
||||
assert loop.model == "new-model"
|
||||
assert loop.context_window_tokens == 2000
|
||||
@@ -50,6 +51,29 @@ def test_provider_refresh_updates_all_model_dependents(tmp_path: Path) -> None:
|
||||
assert not hasattr(loop.consolidator, "max_completion_tokens")
|
||||
|
||||
|
||||
def test_loop_has_no_mutable_runtime_mirrors_or_legacy_snapshot_api(tmp_path: Path) -> None:
|
||||
loop = AgentLoop(
|
||||
bus=MessageBus(),
|
||||
provider=_provider("test-model"),
|
||||
workspace=tmp_path,
|
||||
model="test-model",
|
||||
context_window_tokens=1000,
|
||||
)
|
||||
|
||||
assert {
|
||||
"provider",
|
||||
"model",
|
||||
"context_window_tokens",
|
||||
"model_presets",
|
||||
"_active_preset",
|
||||
"_provider_signature",
|
||||
"_max_messages",
|
||||
}.isdisjoint(loop.__dict__)
|
||||
assert not hasattr(loop, "_apply_provider_snapshot")
|
||||
assert not hasattr(loop, "_build_model_preset_snapshot")
|
||||
assert not hasattr(loop, "_sync_replay_max_messages")
|
||||
|
||||
|
||||
def test_llm_runtime_refreshes_provider_snapshot(tmp_path: Path) -> None:
|
||||
old_provider = _provider("old-model")
|
||||
new_provider = _provider("new-model", max_tokens=456)
|
||||
@@ -118,7 +142,7 @@ def test_settings_context_window_refreshes_runtime_state(
|
||||
loop = AgentLoop.from_config(config, provider_snapshot_loader=loader)
|
||||
|
||||
payload = update_agent_settings({"context_window_tokens": ["262144"]})
|
||||
loop._refresh_provider_snapshot()
|
||||
loop.llm_runtime()
|
||||
|
||||
assert payload["requires_restart"] is False
|
||||
assert loop.context_window_tokens == 262_144
|
||||
|
||||
@@ -54,9 +54,10 @@ def test_model_preset_setter_updates_state(tmp_path) -> None:
|
||||
assert loop.model_preset == "fast"
|
||||
assert loop.model == "openai/gpt-4.1"
|
||||
assert loop.context_window_tokens == 32_768
|
||||
assert loop.provider.generation.temperature == 0.5
|
||||
assert loop.provider.generation.max_tokens == 4096
|
||||
assert loop.provider.generation.reasoning_effort == "low"
|
||||
runtime = loop.llm_runtime()
|
||||
assert runtime.generation.temperature == 0.5
|
||||
assert runtime.generation.max_tokens == 4096
|
||||
assert runtime.generation.reasoning_effort == "low"
|
||||
assert not hasattr(loop.subagents, "model")
|
||||
assert not hasattr(loop.consolidator, "model")
|
||||
assert not hasattr(loop.consolidator, "context_window_tokens")
|
||||
@@ -174,7 +175,7 @@ def test_active_model_preset_survives_unchanged_config_refresh(tmp_path) -> None
|
||||
)
|
||||
|
||||
loop.set_model_preset("fast")
|
||||
loop._refresh_provider_snapshot()
|
||||
loop.llm_runtime()
|
||||
|
||||
assert loop.model_preset == "fast"
|
||||
assert loop.provider is fast_provider
|
||||
@@ -210,7 +211,7 @@ def test_config_model_refresh_clears_active_model_preset(tmp_path) -> None:
|
||||
)
|
||||
|
||||
loop.set_model_preset("fast")
|
||||
loop._refresh_provider_snapshot()
|
||||
loop.llm_runtime()
|
||||
|
||||
assert loop.model_preset is None
|
||||
assert loop.provider is webui_provider
|
||||
@@ -292,7 +293,7 @@ def test_self_tool_set_model_clears_active_preset(tmp_path) -> None:
|
||||
tool = MyTool(runtime_state=loop, modify_allowed=True)
|
||||
result = tool._modify("model", "anthropic/claude-opus-4-5")
|
||||
assert "Error" not in result
|
||||
assert loop._active_preset is None
|
||||
assert loop.model_preset is None
|
||||
assert loop.model == "anthropic/claude-opus-4-5"
|
||||
|
||||
|
||||
@@ -323,5 +324,7 @@ def test_from_config_static_preset_loader_does_not_enable_hot_reload(tmp_path) -
|
||||
fake_provider = _provider("openai/gpt-4.1")
|
||||
with patch("nanobot.providers.factory.make_provider", return_value=fake_provider):
|
||||
loop = AgentLoop.from_config(config)
|
||||
assert loop._provider_snapshot_loader is None
|
||||
assert loop._preset_snapshot_loader is not None
|
||||
default_runtime = loop.runtime_resolver.runtime
|
||||
resolved = loop.runtime_resolver.resolve_preset("fast")
|
||||
assert resolved.model == "openai/gpt-4.1-mini"
|
||||
assert loop.runtime_resolver.runtime is default_runtime
|
||||
|
||||
@@ -35,6 +35,12 @@ def _make_mock_loop(**overrides):
|
||||
loop._concurrency_gate = None
|
||||
loop._unified_session = False
|
||||
loop._extra_hooks = []
|
||||
loop.set_runtime_model.side_effect = lambda value: setattr(loop, "model", value)
|
||||
loop.set_runtime_context_window.side_effect = lambda value: setattr(
|
||||
loop,
|
||||
"context_window_tokens",
|
||||
value,
|
||||
)
|
||||
|
||||
# web_config mock — needed for check tests
|
||||
loop.web_config = MagicMock()
|
||||
@@ -237,12 +243,12 @@ class TestModifyRestricted:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_modify_context_window_valid(self):
|
||||
loop = _make_mock_loop(_sync_replay_max_messages=MagicMock())
|
||||
loop = _make_mock_loop()
|
||||
tool = _make_tool(runtime_state=loop)
|
||||
result = await tool.execute(action="set", key="context_window_tokens", value=131072)
|
||||
assert "Set context_window_tokens" in result
|
||||
assert loop.context_window_tokens == 131072
|
||||
loop._sync_replay_max_messages.assert_called_once_with()
|
||||
loop.set_runtime_context_window.assert_called_once_with(131072)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_modify_none_value_for_restricted_int(self):
|
||||
|
||||
@@ -245,8 +245,8 @@ class TestRestartCommand:
|
||||
|
||||
msg = InboundMessage(channel="telegram", sender_id="u1", chat_id="c1", content="/status")
|
||||
runtime = loop.llm_runtime()
|
||||
loop.model = "replacement-model"
|
||||
loop.context_window_tokens = 10
|
||||
loop.set_runtime_model("replacement-model")
|
||||
loop.set_runtime_context_window(10)
|
||||
loop.provider.generation = SimpleNamespace(
|
||||
temperature=1.0,
|
||||
max_tokens=1,
|
||||
|
||||
@@ -30,6 +30,7 @@ from nanobot.nanobot import (
|
||||
StreamEvent,
|
||||
StreamEventType,
|
||||
)
|
||||
from nanobot.utils.llm_runtime import runtime_from_provider_snapshot
|
||||
|
||||
|
||||
def _write_config(tmp_path: Path, overrides: dict | None = None) -> Path:
|
||||
@@ -504,45 +505,58 @@ async def test_run_allows_parallel_sessions_without_model_override(tmp_path):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_model_overrides_are_serialized_before_snapshot_build(tmp_path):
|
||||
async def test_run_model_overrides_can_overlap_without_default_mutation(tmp_path):
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.providers.factory import ProviderSnapshot
|
||||
|
||||
config_path = _write_config(tmp_path)
|
||||
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||
original_model = bot._loop.model
|
||||
active_models: list[str] = []
|
||||
snapshot_base_models: list[str] = []
|
||||
assert not hasattr(bot, "_runtime_overrides")
|
||||
original_runtime = bot._loop.runtime_resolver.runtime
|
||||
active_runs: list[tuple[str, str]] = []
|
||||
first_entered = asyncio.Event()
|
||||
both_entered = asyncio.Event()
|
||||
release_first = asyncio.Event()
|
||||
|
||||
def fake_snapshot(*, model, model_preset):
|
||||
def fake_resolve(*, model, model_preset, config):
|
||||
assert model is not None
|
||||
assert model_preset is None
|
||||
snapshot_base_models.append(bot._loop.model)
|
||||
return ProviderSnapshot(
|
||||
assert config is bot._config
|
||||
return runtime_from_provider_snapshot(ProviderSnapshot(
|
||||
provider=_fake_provider(model, max_tokens=2048),
|
||||
model=model,
|
||||
context_window_tokens=4096,
|
||||
signature=("sdk", model),
|
||||
)
|
||||
))
|
||||
|
||||
bot._runtime_overrides.model_override_snapshot = MagicMock(side_effect=fake_snapshot)
|
||||
bot._loop.runtime_resolver.resolve_override = MagicMock(side_effect=fake_resolve)
|
||||
|
||||
async def fake_process_direct(message, *, session_key, hooks):
|
||||
active_models.append(bot._loop.model)
|
||||
async def fake_process_direct(message, *, session_key, hooks, runtime):
|
||||
active_runs.append((session_key, runtime.model))
|
||||
assert bot._loop.runtime_resolver.runtime is original_runtime
|
||||
if message == "first":
|
||||
first_entered.set()
|
||||
await asyncio.wait_for(release_first.wait(), timeout=1)
|
||||
if len(active_runs) == 2:
|
||||
both_entered.set()
|
||||
await asyncio.wait_for(release_first.wait(), timeout=1)
|
||||
return OutboundMessage(channel="cli", chat_id="direct", content=message)
|
||||
|
||||
bot._loop.process_direct = fake_process_direct
|
||||
|
||||
first = asyncio.create_task(bot.run("first", model="model:first"))
|
||||
first = asyncio.create_task(bot.run(
|
||||
"first",
|
||||
session_key="sdk:first",
|
||||
model="model:first",
|
||||
))
|
||||
await asyncio.wait_for(first_entered.wait(), timeout=1)
|
||||
|
||||
second = asyncio.create_task(bot.run("second", model="model:second"))
|
||||
await asyncio.sleep(0)
|
||||
second = asyncio.create_task(bot.run(
|
||||
"second",
|
||||
session_key="sdk:second",
|
||||
model="model:second",
|
||||
))
|
||||
await asyncio.wait_for(both_entered.wait(), timeout=1)
|
||||
assert not first.done()
|
||||
assert not second.done()
|
||||
|
||||
release_first.set()
|
||||
@@ -550,21 +564,21 @@ async def test_run_model_overrides_are_serialized_before_snapshot_build(tmp_path
|
||||
|
||||
assert first_result.content == "first"
|
||||
assert second_result.content == "second"
|
||||
assert active_models == ["model:first", "model:second"]
|
||||
assert snapshot_base_models == [original_model, original_model]
|
||||
assert bot._loop.model == original_model
|
||||
assert set(active_runs) == {
|
||||
("sdk:first", "model:first"),
|
||||
("sdk:second", "model:second"),
|
||||
}
|
||||
assert bot._loop.runtime_resolver.runtime is original_runtime
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_model_override_is_per_run_and_restores_default(tmp_path):
|
||||
async def test_run_model_override_is_per_run_without_default_mutation(tmp_path):
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.providers.factory import ProviderSnapshot
|
||||
|
||||
config_path = _write_config(tmp_path)
|
||||
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||
original_provider = bot._loop.provider
|
||||
original_model = bot._loop.model
|
||||
original_signature = bot._loop._provider_signature
|
||||
original_runtime = bot._loop.runtime_resolver.runtime
|
||||
override_provider = _fake_provider("override-provider", max_tokens=2048)
|
||||
override = ProviderSnapshot(
|
||||
provider=override_provider,
|
||||
@@ -572,13 +586,17 @@ async def test_run_model_override_is_per_run_and_restores_default(tmp_path):
|
||||
context_window_tokens=4096,
|
||||
signature=("sdk", "override"),
|
||||
)
|
||||
bot._runtime_overrides.model_override_snapshot = MagicMock(return_value=override)
|
||||
override_runtime = runtime_from_provider_snapshot(override)
|
||||
bot._loop.runtime_resolver.resolve_override = MagicMock(
|
||||
return_value=override_runtime
|
||||
)
|
||||
|
||||
async def fake_process_direct(message, *, session_key, hooks):
|
||||
assert bot._loop.provider is override_provider
|
||||
async def fake_process_direct(message, *, session_key, hooks, runtime):
|
||||
assert runtime is override_runtime
|
||||
assert not hasattr(bot._loop.runner, "provider")
|
||||
assert bot._loop.model == "openai/gpt-4.1-mini"
|
||||
assert bot._loop.context_window_tokens == 4096
|
||||
assert runtime.model == "openai/gpt-4.1-mini"
|
||||
assert runtime.context_window_tokens == 4096
|
||||
assert bot._loop.runtime_resolver.runtime is original_runtime
|
||||
return OutboundMessage(channel="cli", chat_id="direct", content="ok")
|
||||
|
||||
bot._loop.process_direct = fake_process_direct
|
||||
@@ -586,14 +604,13 @@ async def test_run_model_override_is_per_run_and_restores_default(tmp_path):
|
||||
result = await bot.run("hi", model="openai/gpt-4.1-mini")
|
||||
|
||||
assert result.content == "ok"
|
||||
bot._runtime_overrides.model_override_snapshot.assert_called_once_with(
|
||||
bot._loop.runtime_resolver.resolve_override.assert_called_once_with(
|
||||
model="openai/gpt-4.1-mini",
|
||||
model_preset=None,
|
||||
config=bot._config,
|
||||
)
|
||||
assert bot._loop.provider is original_provider
|
||||
assert not hasattr(bot._loop.runner, "provider")
|
||||
assert bot._loop.model == original_model
|
||||
assert bot._loop._provider_signature == original_signature
|
||||
assert bot._loop.runtime_resolver.runtime is original_runtime
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -603,7 +620,7 @@ async def test_run_model_preset_override_is_per_run(tmp_path):
|
||||
|
||||
config_path = _write_config(tmp_path)
|
||||
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||
original_model = bot._loop.model
|
||||
original_runtime = bot._loop.runtime_resolver.runtime
|
||||
override_provider = _fake_provider("preset-provider", max_tokens=1024)
|
||||
override = ProviderSnapshot(
|
||||
provider=override_provider,
|
||||
@@ -611,19 +628,25 @@ async def test_run_model_preset_override_is_per_run(tmp_path):
|
||||
context_window_tokens=2048,
|
||||
signature=("preset", "fast"),
|
||||
)
|
||||
bot._loop._build_model_preset_snapshot = MagicMock(return_value=override)
|
||||
override_runtime = runtime_from_provider_snapshot(override, model_preset="fast")
|
||||
bot._loop.runtime_resolver.resolve_override = MagicMock(
|
||||
return_value=override_runtime
|
||||
)
|
||||
|
||||
async def fake_process_direct(message, *, session_key, hooks):
|
||||
assert bot._loop.provider is override_provider
|
||||
assert bot._loop.model == "openai/gpt-4.1-mini"
|
||||
async def fake_process_direct(message, *, session_key, hooks, runtime):
|
||||
assert runtime is override_runtime
|
||||
return OutboundMessage(channel="cli", chat_id="direct", content="ok")
|
||||
|
||||
bot._loop.process_direct = fake_process_direct
|
||||
|
||||
await bot.run("hi", model_preset="fast")
|
||||
|
||||
bot._loop._build_model_preset_snapshot.assert_called_once_with("fast")
|
||||
assert bot._loop.model == original_model
|
||||
bot._loop.runtime_resolver.resolve_override.assert_called_once_with(
|
||||
model=None,
|
||||
model_preset="fast",
|
||||
config=bot._config,
|
||||
)
|
||||
assert bot._loop.runtime_resolver.runtime is original_runtime
|
||||
assert bot._loop.model_preset is None
|
||||
|
||||
|
||||
@@ -745,7 +768,9 @@ async def test_stream_yields_text_events_in_order(tmp_path):
|
||||
config_path = _write_config(tmp_path)
|
||||
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||
|
||||
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
|
||||
async def fake_process_direct(
|
||||
message, *, session_key, on_stream, on_stream_end, hooks, runtime
|
||||
):
|
||||
assert message == "hi"
|
||||
assert session_key == "sdk:default"
|
||||
await on_stream("Hel")
|
||||
@@ -780,7 +805,9 @@ async def test_run_streamed_wait_returns_full_result_without_consuming_events(tm
|
||||
config_path = _write_config(tmp_path)
|
||||
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||
|
||||
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
|
||||
async def fake_process_direct(
|
||||
message, *, session_key, on_stream, on_stream_end, hooks, runtime
|
||||
):
|
||||
await on_stream("done")
|
||||
await on_stream_end(resuming=False)
|
||||
ctx = AgentRunHookContext(
|
||||
@@ -822,7 +849,9 @@ async def test_run_streamed_cancel_releases_full_queue_without_consuming(tmp_pat
|
||||
config_path = _write_config(tmp_path)
|
||||
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||
|
||||
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
|
||||
async def fake_process_direct(
|
||||
message, *, session_key, on_stream, on_stream_end, hooks, runtime
|
||||
):
|
||||
for i in range(400):
|
||||
await on_stream(str(i))
|
||||
await on_stream_end(resuming=False)
|
||||
@@ -886,16 +915,17 @@ async def test_run_streamed_forwards_runtime_options(tmp_path):
|
||||
assert callable(kwargs["on_stream"])
|
||||
assert callable(kwargs["on_stream_end"])
|
||||
assert kwargs["hooks"]
|
||||
assert kwargs["runtime"] is bot._loop.runtime_resolver.runtime
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_streamed_model_override_reports_model_and_restores(tmp_path):
|
||||
async def test_run_streamed_model_override_reports_admitted_runtime(tmp_path):
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.providers.factory import ProviderSnapshot
|
||||
|
||||
config_path = _write_config(tmp_path)
|
||||
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||
original_model = bot._loop.model
|
||||
original_runtime = bot._loop.runtime_resolver.runtime
|
||||
override_provider = _fake_provider("stream-provider", max_tokens=2048)
|
||||
override = ProviderSnapshot(
|
||||
provider=override_provider,
|
||||
@@ -903,11 +933,22 @@ async def test_run_streamed_model_override_reports_model_and_restores(tmp_path):
|
||||
context_window_tokens=4096,
|
||||
signature=("sdk", "stream"),
|
||||
)
|
||||
bot._runtime_overrides.model_override_snapshot = MagicMock(return_value=override)
|
||||
override_runtime = runtime_from_provider_snapshot(override)
|
||||
bot._loop.runtime_resolver.resolve_override = MagicMock(
|
||||
return_value=override_runtime
|
||||
)
|
||||
|
||||
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
|
||||
assert bot._loop.provider is override_provider
|
||||
assert bot._loop.model == "openai/gpt-4.1-mini"
|
||||
async def fake_process_direct(
|
||||
message,
|
||||
*,
|
||||
session_key,
|
||||
on_stream,
|
||||
on_stream_end,
|
||||
hooks,
|
||||
runtime,
|
||||
):
|
||||
assert runtime is override_runtime
|
||||
assert bot._loop.runtime_resolver.runtime is original_runtime
|
||||
await on_stream("ok")
|
||||
await on_stream_end(resuming=False)
|
||||
return OutboundMessage(channel="cli", chat_id="direct", content="ok")
|
||||
@@ -922,7 +963,7 @@ async def test_run_streamed_model_override_reports_model_and_restores(tmp_path):
|
||||
assert events[0].type == "run.started"
|
||||
assert events[0].metadata["model"] == "openai/gpt-4.1-mini"
|
||||
assert events[0].metadata["model_preset"] is None
|
||||
assert bot._loop.model == original_model
|
||||
assert bot._loop.runtime_resolver.runtime is original_runtime
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -947,7 +988,9 @@ async def test_run_streamed_emits_tool_events(tmp_path):
|
||||
config_path = _write_config(tmp_path)
|
||||
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||
|
||||
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
|
||||
async def fake_process_direct(
|
||||
message, *, session_key, on_stream, on_stream_end, hooks, runtime
|
||||
):
|
||||
calls = [
|
||||
ToolCallRequest(id="call_ok", name="read_file", arguments={"path": "README.md"}),
|
||||
ToolCallRequest(id="call_bad", name="exec", arguments={"cmd": "false"}),
|
||||
@@ -993,7 +1036,9 @@ async def test_run_streamed_emits_reasoning_events(tmp_path):
|
||||
config_path = _write_config(tmp_path)
|
||||
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||
|
||||
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
|
||||
async def fake_process_direct(
|
||||
message, *, session_key, on_stream, on_stream_end, hooks, runtime
|
||||
):
|
||||
for hook in hooks:
|
||||
await hook.emit_reasoning("thinking")
|
||||
await hook.emit_reasoning_end()
|
||||
@@ -1020,7 +1065,9 @@ async def test_stream_generator_break_cancels_underlying_run(tmp_path):
|
||||
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||
cancelled = asyncio.Event()
|
||||
|
||||
async def fake_process_direct(message, *, session_key, on_stream, on_stream_end, hooks):
|
||||
async def fake_process_direct(
|
||||
message, *, session_key, on_stream, on_stream_end, hooks, runtime
|
||||
):
|
||||
try:
|
||||
await on_stream("first")
|
||||
await asyncio.sleep(10)
|
||||
|
||||
Reference in New Issue
Block a user