mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-04 16:38:49 +00:00
refactor(agent): make resolver sole runtime owner
This commit is contained in:
parent
c9d3e74342
commit
21f58cbabf
@ -6,6 +6,7 @@ import asyncio
|
|||||||
import dataclasses
|
import dataclasses
|
||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
|
from collections.abc import Mapping
|
||||||
from contextlib import AsyncExitStack, nullcontext, suppress
|
from contextlib import AsyncExitStack, nullcontext, suppress
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from enum import Enum, auto
|
from enum import Enum, auto
|
||||||
@ -174,37 +175,48 @@ class AgentLoop:
|
|||||||
def tool_names(self) -> list[str]:
|
def tool_names(self) -> list[str]:
|
||||||
return self.tools.tool_names
|
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:
|
def llm_runtime(self) -> LLMRuntime:
|
||||||
"""Resolve the immutable default used to admit the next turn."""
|
"""Resolve the immutable default used to admit the next turn."""
|
||||||
self._refresh_provider_snapshot()
|
previous = self.runtime_resolver.runtime
|
||||||
runtime = self.runtime_resolver.current()
|
try:
|
||||||
captured = LLMRuntime.capture(
|
runtime = self.runtime_resolver.current(refresh=True)
|
||||||
self.provider,
|
except Exception:
|
||||||
self.model,
|
logger.exception("Failed to refresh model runtime")
|
||||||
context_window_tokens=self.context_window_tokens,
|
return previous
|
||||||
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 (
|
if (
|
||||||
runtime.provider is not self.provider
|
runtime.model != previous.model
|
||||||
or runtime.model != self.model
|
or runtime.model_preset != previous.model_preset
|
||||||
or runtime.generation != captured.generation
|
or runtime.snapshot_signature != previous.snapshot_signature
|
||||||
or runtime.context_window_tokens != self.context_window_tokens
|
|
||||||
or runtime.model_preset != self._active_preset
|
|
||||||
):
|
):
|
||||||
snapshot = ProviderSnapshot(
|
self._publish_runtime_selection(runtime)
|
||||||
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
|
return runtime
|
||||||
|
|
||||||
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
|
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
|
||||||
@ -271,32 +283,26 @@ class AgentLoop:
|
|||||||
self.runtime_event_publisher = RuntimeEventPublisher(self.runtime_events)
|
self.runtime_event_publisher = RuntimeEventPublisher(self.runtime_events)
|
||||||
self.channels_config = channels_config
|
self.channels_config = channels_config
|
||||||
self.restart_mode = restart_mode
|
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._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.workspace = workspace
|
||||||
self.model = model or provider.get_default_model()
|
initial_model = model or provider.get_default_model()
|
||||||
self.max_iterations = (
|
self.max_iterations = (
|
||||||
max_iterations if max_iterations is not None else defaults.max_tool_iterations
|
max_iterations if max_iterations is not None else defaults.max_tool_iterations
|
||||||
)
|
)
|
||||||
self.context_window_tokens = (
|
initial_context_window = (
|
||||||
context_window_tokens
|
context_window_tokens
|
||||||
if context_window_tokens is not None
|
if context_window_tokens is not None
|
||||||
else defaults.context_window_tokens
|
else defaults.context_window_tokens
|
||||||
)
|
)
|
||||||
self.model_presets: dict[str, ModelPresetConfig] = model_presets or {}
|
configured_presets = model_presets or {}
|
||||||
self._active_preset: str | None = None
|
|
||||||
self.runtime_resolver = ModelRuntimeResolver(
|
self.runtime_resolver = ModelRuntimeResolver(
|
||||||
LLMRuntime.capture(
|
LLMRuntime.capture(
|
||||||
provider,
|
provider,
|
||||||
self.model,
|
initial_model,
|
||||||
context_window_tokens=self.context_window_tokens,
|
context_window_tokens=initial_context_window,
|
||||||
snapshot_signature=provider_signature,
|
snapshot_signature=provider_signature,
|
||||||
),
|
),
|
||||||
model_presets=self.model_presets,
|
model_presets=configured_presets,
|
||||||
provider_snapshot_loader=provider_snapshot_loader,
|
provider_snapshot_loader=provider_snapshot_loader,
|
||||||
preset_snapshot_loader=preset_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),
|
llm_wall_timeout_for_session=lambda sk: runner_wall_llm_timeout_s(self.sessions, sk),
|
||||||
)
|
)
|
||||||
self._unified_session = unified_session
|
self._unified_session = unified_session
|
||||||
self._max_messages = replay_max_messages_for_context(self.context_window_tokens)
|
|
||||||
self._running = False
|
self._running = False
|
||||||
self._mcp_servers = mcp_servers or {}
|
self._mcp_servers = mcp_servers or {}
|
||||||
self._mcp_stacks: dict[str, AsyncExitStack] = {}
|
self._mcp_stacks: dict[str, AsyncExitStack] = {}
|
||||||
@ -401,7 +406,7 @@ class AgentLoop:
|
|||||||
)
|
)
|
||||||
if model_preset:
|
if model_preset:
|
||||||
self.set_model_preset(model_preset, publish_update=False)
|
self.set_model_preset(model_preset, publish_update=False)
|
||||||
self._register_default_tools()
|
self._register_default_tools(provider_snapshot_loader=provider_snapshot_loader)
|
||||||
self._runtime_vars: dict[str, Any] = {}
|
self._runtime_vars: dict[str, Any] = {}
|
||||||
self._current_iteration: int = 0
|
self._current_iteration: int = 0
|
||||||
self.commands = CommandRouter()
|
self.commands = CommandRouter()
|
||||||
@ -468,98 +473,51 @@ class AgentLoop:
|
|||||||
"""Keep subagent runtime limits aligned with mutable loop settings."""
|
"""Keep subagent runtime limits aligned with mutable loop settings."""
|
||||||
self.subagents.max_iterations = self.max_iterations
|
self.subagents.max_iterations = self.max_iterations
|
||||||
|
|
||||||
def _apply_provider_snapshot(
|
def _publish_runtime_selection(
|
||||||
self,
|
self,
|
||||||
snapshot: ProviderSnapshot,
|
runtime: LLMRuntime,
|
||||||
*,
|
*,
|
||||||
publish_update: bool = True,
|
publish_update: bool = True,
|
||||||
model_preset: str | None = None,
|
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Swap model/provider for future turns without disturbing an active one."""
|
if not publish_update:
|
||||||
runtime = self.runtime_resolver.adopt_snapshot(
|
return
|
||||||
snapshot,
|
if self._runtime_model_publisher is not None:
|
||||||
model_preset=model_preset,
|
self._runtime_model_publisher(runtime.model, runtime.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
|
|
||||||
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._runtime_events().runtime_model_changed(
|
||||||
self.model,
|
runtime.model,
|
||||||
model_preset if model_preset is not None else self.model_preset,
|
runtime.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:
|
def set_model_preset(
|
||||||
"""Resolve a preset by name and apply all runtime model dependents."""
|
self,
|
||||||
name = preset_helpers.normalize_preset_name(name, self.model_presets)
|
name: str | None,
|
||||||
|
*,
|
||||||
|
publish_update: bool = True,
|
||||||
|
) -> LLMRuntime:
|
||||||
|
"""Select a named default runtime for future turns."""
|
||||||
|
old_model = self.model
|
||||||
runtime = self.runtime_resolver.select_preset(name)
|
runtime = self.runtime_resolver.select_preset(name)
|
||||||
snapshot = ProviderSnapshot(
|
self._publish_runtime_selection(runtime, publish_update=publish_update)
|
||||||
provider=runtime.provider,
|
logger.info(
|
||||||
model=runtime.model,
|
"Runtime model switched for next turn: {} -> {}",
|
||||||
context_window_tokens=runtime.context_window_tokens,
|
old_model,
|
||||||
signature=runtime.snapshot_signature or ("model_preset", name),
|
runtime.model,
|
||||||
generation=runtime.generation,
|
|
||||||
)
|
)
|
||||||
self._apply_provider_snapshot(snapshot, publish_update=publish_update, model_preset=name)
|
return runtime
|
||||||
self._active_preset = name
|
|
||||||
|
|
||||||
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."""
|
"""Register the default set of tools via plugin loader."""
|
||||||
from nanobot.agent.tools.context import ToolContext
|
from nanobot.agent.tools.context import ToolContext
|
||||||
from nanobot.agent.tools.loader import ToolLoader
|
from nanobot.agent.tools.loader import ToolLoader
|
||||||
@ -571,7 +529,7 @@ class AgentLoop:
|
|||||||
subagent_manager=self.subagents,
|
subagent_manager=self.subagents,
|
||||||
cron_service=self.cron_service,
|
cron_service=self.cron_service,
|
||||||
sessions=self.sessions,
|
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,
|
image_generation_provider_configs=self._image_generation_provider_configs,
|
||||||
timezone=self.context.timezone or "UTC",
|
timezone=self.context.timezone or "UTC",
|
||||||
workspace_sandbox=self.workspace_scopes.sandbox_status,
|
workspace_sandbox=self.workspace_scopes.sandbox_status,
|
||||||
|
|||||||
@ -56,7 +56,9 @@ class RuntimeState(Protocol):
|
|||||||
|
|
||||||
def _sync_subagent_runtime_limits(self) -> None: ...
|
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
|
@property
|
||||||
def model_preset(self) -> str | None: ...
|
def model_preset(self) -> str | None: ...
|
||||||
|
|
||||||
_active_preset: str | None
|
|
||||||
|
|||||||
@ -57,7 +57,7 @@ class MyTool(Tool):
|
|||||||
|
|
||||||
BLOCKED = frozenset({
|
BLOCKED = frozenset({
|
||||||
# Core infrastructure
|
# Core infrastructure
|
||||||
"bus", "provider", "_running", "tools",
|
"bus", "provider", "runtime_resolver", "_running", "tools",
|
||||||
# Config management
|
# Config management
|
||||||
"_runtime_vars",
|
"_runtime_vars",
|
||||||
# Subsystems
|
# Subsystems
|
||||||
@ -107,6 +107,11 @@ class MyTool(Tool):
|
|||||||
}
|
}
|
||||||
|
|
||||||
_MAX_RUNTIME_KEYS = 64
|
_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:
|
def __init__(self, runtime_state: RuntimeState, modify_allowed: bool = True) -> None:
|
||||||
self._runtime_state = runtime_state
|
self._runtime_state = runtime_state
|
||||||
@ -325,9 +330,20 @@ class MyTool(Tool):
|
|||||||
|
|
||||||
# -- inspect --
|
# -- 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:
|
def _inspect(self, key: str | None) -> str:
|
||||||
if not key:
|
if not key:
|
||||||
return self._inspect_all()
|
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]
|
top = key.split(".")[0]
|
||||||
if top in self._DENIED_ATTRS or top.startswith("__"):
|
if top in self._DENIED_ATTRS or top.startswith("__"):
|
||||||
return ToolResult.error(f"Error: '{top}' is not accessible")
|
return ToolResult.error(f"Error: '{top}' is not accessible")
|
||||||
@ -353,8 +369,13 @@ class MyTool(Tool):
|
|||||||
parts: list[str] = []
|
parts: list[str] = []
|
||||||
# RESTRICTED keys
|
# RESTRICTED keys
|
||||||
for k in self.RESTRICTED:
|
for k in self.RESTRICTED:
|
||||||
parts.append(self._format_value(getattr(state, k, None), k))
|
found, value = self._current_runtime_value(k)
|
||||||
parts.append(self._format_value(state.model_preset, "model_preset"))
|
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
|
# 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"):
|
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):
|
if _has_real_attr(state, k):
|
||||||
@ -432,13 +453,16 @@ class MyTool(Tool):
|
|||||||
return ToolResult.error(f"Error: '{key}' must be <= {spec['max']}")
|
return ToolResult.error(f"Error: '{key}' must be <= {spec['max']}")
|
||||||
if "min_len" in spec and len(str(value)) < spec["min_len"]:
|
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")
|
return ToolResult.error(f"Error: '{key}' must be at least {spec['min_len']} characters")
|
||||||
setattr(self._runtime_state, key, value)
|
|
||||||
if key == "model":
|
if key == "model":
|
||||||
self._runtime_state._active_preset = None
|
self._runtime_state.set_runtime_model(value)
|
||||||
sync_replay = getattr(self._runtime_state, "_sync_replay_max_messages", None)
|
elif key == "context_window_tokens":
|
||||||
if key == "context_window_tokens" and callable(sync_replay):
|
self._runtime_state.set_runtime_context_window(value)
|
||||||
sync_replay()
|
else:
|
||||||
if key == "max_iterations" and hasattr(self._runtime_state, "_sync_subagent_runtime_limits"):
|
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._runtime_state._sync_subagent_runtime_limits()
|
||||||
self._audit("modify", f"{key}: {old!r} -> {value!r}")
|
self._audit("modify", f"{key}: {old!r} -> {value!r}")
|
||||||
return f"Set {key} = {value!r} (was {old!r})"
|
return f"Set {key} = {value!r} (was {old!r})"
|
||||||
|
|||||||
@ -349,7 +349,7 @@ async def cmd_model(ctx: CommandContext) -> OutboundMessage:
|
|||||||
|
|
||||||
name = parts[0]
|
name = parts[0]
|
||||||
try:
|
try:
|
||||||
loop.set_model_preset(name)
|
runtime = loop.set_model_preset(name)
|
||||||
except (KeyError, ValueError) as exc:
|
except (KeyError, ValueError) as exc:
|
||||||
names = _model_preset_names(loop)
|
names = _model_preset_names(loop)
|
||||||
return OutboundMessage(
|
return OutboundMessage(
|
||||||
@ -362,11 +362,11 @@ async def cmd_model(ctx: CommandContext) -> OutboundMessage:
|
|||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
)
|
)
|
||||||
|
|
||||||
max_tokens = getattr(getattr(loop.provider, "generation", None), "max_tokens", None)
|
max_tokens = runtime.generation.max_tokens
|
||||||
lines = [
|
lines = [
|
||||||
f"Switched model preset to `{loop.model_preset}`.",
|
f"Switched model preset to `{runtime.model_preset}`.",
|
||||||
f"- Model: `{loop.model}`",
|
f"- Model: `{runtime.model}`",
|
||||||
f"- Context window: {loop.context_window_tokens}",
|
f"- Context window: {runtime.context_window_tokens}",
|
||||||
]
|
]
|
||||||
if max_tokens is not None:
|
if max_tokens is not None:
|
||||||
lines.append(f"- Max output tokens: {max_tokens}")
|
lines.append(f"- Max output tokens: {max_tokens}")
|
||||||
|
|||||||
@ -14,7 +14,6 @@ from nanobot.config.schema import Config
|
|||||||
from nanobot.providers.image_generation import image_gen_provider_configs
|
from nanobot.providers.image_generation import image_gen_provider_configs
|
||||||
from nanobot.sdk.clients import MemoryClient, RuntimeClient, SessionClient
|
from nanobot.sdk.clients import MemoryClient, RuntimeClient, SessionClient
|
||||||
from nanobot.sdk.runtime import (
|
from nanobot.sdk.runtime import (
|
||||||
SDKRuntimeController,
|
|
||||||
build_process_direct_kwargs,
|
build_process_direct_kwargs,
|
||||||
ensure_single_model_selector,
|
ensure_single_model_selector,
|
||||||
)
|
)
|
||||||
@ -74,7 +73,6 @@ class Nanobot:
|
|||||||
def __init__(self, loop: AgentLoop, *, config: Config | None = None) -> None:
|
def __init__(self, loop: AgentLoop, *, config: Config | None = None) -> None:
|
||||||
self._loop = loop
|
self._loop = loop
|
||||||
self._config = config
|
self._config = config
|
||||||
self._runtime_overrides = SDKRuntimeController(loop, config=config)
|
|
||||||
self.sessions = SessionClient(loop)
|
self.sessions = SessionClient(loop)
|
||||||
self.memory = MemoryClient(loop)
|
self.memory = MemoryClient(loop)
|
||||||
self.runtime = RuntimeClient(loop)
|
self.runtime = RuntimeClient(loop)
|
||||||
@ -156,7 +154,11 @@ class Nanobot:
|
|||||||
"""
|
"""
|
||||||
capture = SDKCaptureHook()
|
capture = SDKCaptureHook()
|
||||||
per_run_hooks = [capture, *(hooks or [])]
|
per_run_hooks = [capture, *(hooks or [])]
|
||||||
async with self._runtime_overrides.override(model=model, model_preset=model_preset):
|
runtime = self._loop.runtime_resolver.resolve_override(
|
||||||
|
model=model,
|
||||||
|
model_preset=model_preset,
|
||||||
|
config=self._config,
|
||||||
|
)
|
||||||
kwargs = build_process_direct_kwargs(
|
kwargs = build_process_direct_kwargs(
|
||||||
session_key=session_key,
|
session_key=session_key,
|
||||||
channel=channel,
|
channel=channel,
|
||||||
@ -165,6 +167,8 @@ class Nanobot:
|
|||||||
media=media,
|
media=media,
|
||||||
ephemeral=ephemeral,
|
ephemeral=ephemeral,
|
||||||
)
|
)
|
||||||
|
if runtime is not None:
|
||||||
|
kwargs["runtime"] = runtime
|
||||||
response = await self._loop.process_direct(
|
response = await self._loop.process_direct(
|
||||||
message,
|
message,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
@ -188,7 +192,11 @@ class Nanobot:
|
|||||||
model_preset: str | None = None,
|
model_preset: str | None = None,
|
||||||
) -> RunStream:
|
) -> RunStream:
|
||||||
"""Start a streamed run and return a handle for events and final result."""
|
"""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)
|
queue: asyncio.Queue[StreamEvent | object] = asyncio.Queue(maxsize=256)
|
||||||
emitter = SDKStreamEmitter(queue)
|
emitter = SDKStreamEmitter(queue)
|
||||||
stream_hook = SDKStreamingHook(emitter)
|
stream_hook = SDKStreamingHook(emitter)
|
||||||
@ -202,7 +210,6 @@ class Nanobot:
|
|||||||
await emitter.text_completed(resuming=resuming)
|
await emitter.text_completed(resuming=resuming)
|
||||||
|
|
||||||
async def _run() -> RunResult:
|
async def _run() -> RunResult:
|
||||||
async with self._runtime_overrides.override(model=model, model_preset=model_preset):
|
|
||||||
kwargs = build_process_direct_kwargs(
|
kwargs = build_process_direct_kwargs(
|
||||||
session_key=session_key,
|
session_key=session_key,
|
||||||
channel=channel,
|
channel=channel,
|
||||||
@ -213,6 +220,7 @@ class Nanobot:
|
|||||||
on_stream=_on_stream,
|
on_stream=_on_stream,
|
||||||
on_stream_end=_on_stream_end,
|
on_stream_end=_on_stream_end,
|
||||||
)
|
)
|
||||||
|
kwargs["runtime"] = runtime
|
||||||
await emitter.emit(StreamEvent(
|
await emitter.emit(StreamEvent(
|
||||||
type=STREAM_EVENT_RUN_STARTED,
|
type=STREAM_EVENT_RUN_STARTED,
|
||||||
metadata={
|
metadata={
|
||||||
@ -220,10 +228,8 @@ class Nanobot:
|
|||||||
"channel": channel,
|
"channel": channel,
|
||||||
"chat_id": chat_id,
|
"chat_id": chat_id,
|
||||||
"sender_id": sender_id,
|
"sender_id": sender_id,
|
||||||
"model": self._loop.model,
|
"model": runtime.model,
|
||||||
"model_preset": (
|
"model_preset": runtime.model_preset,
|
||||||
model_preset if model_preset is not None else self._loop.model_preset
|
|
||||||
),
|
|
||||||
},
|
},
|
||||||
))
|
))
|
||||||
try:
|
try:
|
||||||
|
|||||||
@ -2,16 +2,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
from typing import Any
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
def ensure_single_model_selector(
|
def ensure_single_model_selector(
|
||||||
@ -51,142 +42,3 @@ def build_process_direct_kwargs(
|
|||||||
if on_stream_end is not None:
|
if on_stream_end is not None:
|
||||||
kwargs["on_stream_end"] = on_stream_end
|
kwargs["on_stream_end"] = on_stream_end
|
||||||
return kwargs
|
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)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
call_count = 0
|
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
|
nonlocal call_count
|
||||||
|
assert runtime is expected_runtime
|
||||||
assert session_key == "cli:test"
|
assert session_key == "cli:test"
|
||||||
call_count += 1
|
call_count += 1
|
||||||
return False
|
return False
|
||||||
@ -528,7 +530,7 @@ class TestNewCommandArchival:
|
|||||||
loop.consolidator.archive = _failing_summarize # type: ignore[method-assign]
|
loop.consolidator.archive = _failing_summarize # type: ignore[method-assign]
|
||||||
|
|
||||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
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 response is not None
|
||||||
assert "new session started" in response.content.lower()
|
assert "new session started" in response.content.lower()
|
||||||
@ -553,9 +555,11 @@ class TestNewCommandArchival:
|
|||||||
|
|
||||||
archived_count = -1
|
archived_count = -1
|
||||||
archived_session_key = None
|
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
|
nonlocal archived_count, archived_session_key
|
||||||
|
assert runtime is expected_runtime
|
||||||
archived_count = len(messages)
|
archived_count = len(messages)
|
||||||
archived_session_key = session_key
|
archived_session_key = session_key
|
||||||
return True
|
return True
|
||||||
@ -563,7 +567,7 @@ class TestNewCommandArchival:
|
|||||||
loop.consolidator.archive = _fake_summarize # type: ignore[method-assign]
|
loop.consolidator.archive = _fake_summarize # type: ignore[method-assign]
|
||||||
|
|
||||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
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 response is not None
|
||||||
assert "new session started" in response.content.lower()
|
assert "new session started" in response.content.lower()
|
||||||
@ -582,15 +586,17 @@ class TestNewCommandArchival:
|
|||||||
session.add_message("user", f"msg{i}")
|
session.add_message("user", f"msg{i}")
|
||||||
session.add_message("assistant", f"resp{i}")
|
session.add_message("assistant", f"resp{i}")
|
||||||
loop.sessions.save(session)
|
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"
|
assert session_key == "cli:test"
|
||||||
return True
|
return True
|
||||||
|
|
||||||
loop.consolidator.archive = _ok_summarize # type: ignore[method-assign]
|
loop.consolidator.archive = _ok_summarize # type: ignore[method-assign]
|
||||||
|
|
||||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
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 response is not None
|
||||||
assert "new session started" in response.content.lower()
|
assert "new session started" in response.content.lower()
|
||||||
@ -610,8 +616,10 @@ class TestNewCommandArchival:
|
|||||||
|
|
||||||
archived = asyncio.Event()
|
archived = asyncio.Event()
|
||||||
release_archive = 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"
|
assert session_key == "cli:test"
|
||||||
await release_archive.wait()
|
await release_archive.wait()
|
||||||
archived.set()
|
archived.set()
|
||||||
@ -620,7 +628,7 @@ class TestNewCommandArchival:
|
|||||||
loop.consolidator.archive = _slow_summarize # type: ignore[method-assign]
|
loop.consolidator.archive = _slow_summarize # type: ignore[method-assign]
|
||||||
|
|
||||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
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()
|
assert not archived.is_set()
|
||||||
release_archive.set()
|
release_archive.set()
|
||||||
|
|||||||
@ -6,6 +6,7 @@ import nanobot.agent.memory as memory_module
|
|||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.providers.base import LLMResponse
|
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:
|
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(
|
loop.consolidator.maybe_consolidate_by_tokens.assert_any_await(
|
||||||
session,
|
session,
|
||||||
runtime=runtime,
|
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 len(loop.consolidator.maybe_consolidate_by_tokens.call_args_list) == 2
|
||||||
assert all(
|
assert all(
|
||||||
|
|||||||
@ -21,6 +21,7 @@ from nanobot.bus.outbound_events import (
|
|||||||
)
|
)
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||||
|
from nanobot.providers.factory import ProviderSnapshot
|
||||||
from nanobot.session.webui_turns import WebuiTurnCoordinator
|
from nanobot.session.webui_turns import WebuiTurnCoordinator
|
||||||
from nanobot.utils.progress_events import (
|
from nanobot.utils.progress_events import (
|
||||||
invoke_file_edit_progress,
|
invoke_file_edit_progress,
|
||||||
@ -815,8 +816,14 @@ class TestToolEventProgress:
|
|||||||
))
|
))
|
||||||
|
|
||||||
assert len(scheduled_title) == 1
|
assert len(scheduled_title) == 1
|
||||||
loop.provider = MagicMock()
|
next_provider = MagicMock()
|
||||||
loop.model = "switched-after-turn"
|
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]
|
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.bus.queue import MessageBus
|
||||||
from nanobot.cron.session_turns import CRON_HISTORY_META, CRON_TRIGGER_META
|
from nanobot.cron.session_turns import CRON_HISTORY_META, CRON_TRIGGER_META
|
||||||
from nanobot.providers.base import LLMResponse
|
from nanobot.providers.base import LLMResponse
|
||||||
|
from nanobot.providers.factory import ProviderSnapshot
|
||||||
from nanobot.session.automation_turns import AUTOMATION_HISTORY_META
|
from nanobot.session.automation_turns import AUTOMATION_HISTORY_META
|
||||||
from nanobot.session.goal_state import GOAL_STATE_KEY
|
from nanobot.session.goal_state import GOAL_STATE_KEY
|
||||||
from nanobot.session.manager import Session, SessionManager
|
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"
|
assert runtime.model == "test-model"
|
||||||
|
|
||||||
next_provider = MagicMock()
|
next_provider = MagicMock()
|
||||||
loop.provider = next_provider
|
next_provider.generation = SimpleNamespace(
|
||||||
loop.model = "next-model"
|
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()
|
runtime = loop.llm_runtime()
|
||||||
|
|
||||||
assert runtime.provider is next_provider
|
assert runtime.provider is next_provider
|
||||||
|
|||||||
@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import replace
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
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:
|
def test_default_for_200k_context_reaches_file_cap(self, tmp_path: Path) -> None:
|
||||||
loop = _make_loop(tmp_path)
|
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:
|
def test_default_scales_with_context_window(self, tmp_path: Path) -> None:
|
||||||
loop = _make_loop(tmp_path, context_window_tokens=32_768)
|
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:
|
def test_provider_refresh_resyncs_context_derived_limit(self, tmp_path: Path) -> None:
|
||||||
old_provider = MagicMock()
|
old_provider = MagicMock()
|
||||||
@ -93,9 +96,10 @@ class TestMaxMessagesInit:
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
assert loop._max_messages == 327
|
initial = loop.runtime_resolver.runtime
|
||||||
loop._refresh_provider_snapshot()
|
assert replay_max_messages_for_context(initial.context_window_tokens) == 327
|
||||||
assert loop._max_messages == FILE_MAX_MESSAGES
|
refreshed = loop.llm_runtime()
|
||||||
|
assert replay_max_messages_for_context(refreshed.context_window_tokens) == FILE_MAX_MESSAGES
|
||||||
|
|
||||||
|
|
||||||
class TestGetHistoryWithMaxMessages:
|
class TestGetHistoryWithMaxMessages:
|
||||||
@ -136,7 +140,7 @@ class TestMaxMessagesIntegration:
|
|||||||
async def test_process_message_passes_limit_to_history_call(self, tmp_path: Path) -> None:
|
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."""
|
"""The real message path should pass max_messages into session history replay."""
|
||||||
loop = _make_loop(tmp_path)
|
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(
|
loop.provider.chat_with_retry = AsyncMock(
|
||||||
return_value=LLMResponse(content="ok", tool_calls=[], usage={})
|
return_value=LLMResponse(content="ok", tool_calls=[], usage={})
|
||||||
)
|
)
|
||||||
@ -146,12 +150,13 @@ class TestMaxMessagesIntegration:
|
|||||||
session = loop.sessions.get_or_create("cli:test")
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
with patch.object(session, "get_history", wraps=session.get_history) as mock_hist:
|
with patch.object(session, "get_history", wraps=session.get_history) as mock_hist:
|
||||||
result = await loop._process_message(
|
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 result is not None
|
||||||
assert mock_hist.call_count == 1
|
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
|
assert mock_hist.call_args.kwargs["extend_to_user"] is False
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@ -182,8 +187,7 @@ class TestMaxMessagesIntegration:
|
|||||||
tmp_path: Path,
|
tmp_path: Path,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""A live user turn should not extend history to an older long tool turn."""
|
"""A live user turn should not extend history to an older long tool turn."""
|
||||||
loop = _make_loop(tmp_path)
|
loop = _make_loop(tmp_path, context_window_tokens=8_000)
|
||||||
loop._max_messages = 6
|
|
||||||
loop.provider.chat_with_retry = AsyncMock(
|
loop.provider.chat_with_retry = AsyncMock(
|
||||||
return_value=LLMResponse(content="ok", tool_calls=[], usage={})
|
return_value=LLMResponse(content="ok", tool_calls=[], usage={})
|
||||||
)
|
)
|
||||||
@ -194,7 +198,7 @@ class TestMaxMessagesIntegration:
|
|||||||
session.add_message("user", "old")
|
session.add_message("user", "old")
|
||||||
session.add_message("assistant", "old answer")
|
session.add_message("assistant", "old answer")
|
||||||
session.add_message("user", "long older turn")
|
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.messages.extend(_tool_round(f"older-{i}"))
|
||||||
session.add_message("assistant", "older final")
|
session.add_message("assistant", "older final")
|
||||||
|
|
||||||
|
|||||||
@ -18,7 +18,7 @@ def _provider(default_model: str, max_tokens: int = 123) -> MagicMock:
|
|||||||
return provider
|
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")
|
old_provider = _provider("old-model")
|
||||||
new_provider = _provider("new-model", max_tokens=456)
|
new_provider = _provider("new-model", max_tokens=456)
|
||||||
loop = AgentLoop(
|
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.provider is new_provider
|
||||||
assert loop.model == "new-model"
|
assert loop.model == "new-model"
|
||||||
assert loop.context_window_tokens == 2000
|
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")
|
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:
|
def test_llm_runtime_refreshes_provider_snapshot(tmp_path: Path) -> None:
|
||||||
old_provider = _provider("old-model")
|
old_provider = _provider("old-model")
|
||||||
new_provider = _provider("new-model", max_tokens=456)
|
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)
|
loop = AgentLoop.from_config(config, provider_snapshot_loader=loader)
|
||||||
|
|
||||||
payload = update_agent_settings({"context_window_tokens": ["262144"]})
|
payload = update_agent_settings({"context_window_tokens": ["262144"]})
|
||||||
loop._refresh_provider_snapshot()
|
loop.llm_runtime()
|
||||||
|
|
||||||
assert payload["requires_restart"] is False
|
assert payload["requires_restart"] is False
|
||||||
assert loop.context_window_tokens == 262_144
|
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_preset == "fast"
|
||||||
assert loop.model == "openai/gpt-4.1"
|
assert loop.model == "openai/gpt-4.1"
|
||||||
assert loop.context_window_tokens == 32_768
|
assert loop.context_window_tokens == 32_768
|
||||||
assert loop.provider.generation.temperature == 0.5
|
runtime = loop.llm_runtime()
|
||||||
assert loop.provider.generation.max_tokens == 4096
|
assert runtime.generation.temperature == 0.5
|
||||||
assert loop.provider.generation.reasoning_effort == "low"
|
assert runtime.generation.max_tokens == 4096
|
||||||
|
assert runtime.generation.reasoning_effort == "low"
|
||||||
assert not hasattr(loop.subagents, "model")
|
assert not hasattr(loop.subagents, "model")
|
||||||
assert not hasattr(loop.consolidator, "model")
|
assert not hasattr(loop.consolidator, "model")
|
||||||
assert not hasattr(loop.consolidator, "context_window_tokens")
|
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.set_model_preset("fast")
|
||||||
loop._refresh_provider_snapshot()
|
loop.llm_runtime()
|
||||||
|
|
||||||
assert loop.model_preset == "fast"
|
assert loop.model_preset == "fast"
|
||||||
assert loop.provider is fast_provider
|
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.set_model_preset("fast")
|
||||||
loop._refresh_provider_snapshot()
|
loop.llm_runtime()
|
||||||
|
|
||||||
assert loop.model_preset is None
|
assert loop.model_preset is None
|
||||||
assert loop.provider is webui_provider
|
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)
|
tool = MyTool(runtime_state=loop, modify_allowed=True)
|
||||||
result = tool._modify("model", "anthropic/claude-opus-4-5")
|
result = tool._modify("model", "anthropic/claude-opus-4-5")
|
||||||
assert "Error" not in result
|
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"
|
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")
|
fake_provider = _provider("openai/gpt-4.1")
|
||||||
with patch("nanobot.providers.factory.make_provider", return_value=fake_provider):
|
with patch("nanobot.providers.factory.make_provider", return_value=fake_provider):
|
||||||
loop = AgentLoop.from_config(config)
|
loop = AgentLoop.from_config(config)
|
||||||
assert loop._provider_snapshot_loader is None
|
default_runtime = loop.runtime_resolver.runtime
|
||||||
assert loop._preset_snapshot_loader is not None
|
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._concurrency_gate = None
|
||||||
loop._unified_session = False
|
loop._unified_session = False
|
||||||
loop._extra_hooks = []
|
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
|
# web_config mock — needed for check tests
|
||||||
loop.web_config = MagicMock()
|
loop.web_config = MagicMock()
|
||||||
@ -237,12 +243,12 @@ class TestModifyRestricted:
|
|||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_modify_context_window_valid(self):
|
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)
|
tool = _make_tool(runtime_state=loop)
|
||||||
result = await tool.execute(action="set", key="context_window_tokens", value=131072)
|
result = await tool.execute(action="set", key="context_window_tokens", value=131072)
|
||||||
assert "Set context_window_tokens" in result
|
assert "Set context_window_tokens" in result
|
||||||
assert loop.context_window_tokens == 131072
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_modify_none_value_for_restricted_int(self):
|
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")
|
msg = InboundMessage(channel="telegram", sender_id="u1", chat_id="c1", content="/status")
|
||||||
runtime = loop.llm_runtime()
|
runtime = loop.llm_runtime()
|
||||||
loop.model = "replacement-model"
|
loop.set_runtime_model("replacement-model")
|
||||||
loop.context_window_tokens = 10
|
loop.set_runtime_context_window(10)
|
||||||
loop.provider.generation = SimpleNamespace(
|
loop.provider.generation = SimpleNamespace(
|
||||||
temperature=1.0,
|
temperature=1.0,
|
||||||
max_tokens=1,
|
max_tokens=1,
|
||||||
|
|||||||
@ -30,6 +30,7 @@ from nanobot.nanobot import (
|
|||||||
StreamEvent,
|
StreamEvent,
|
||||||
StreamEventType,
|
StreamEventType,
|
||||||
)
|
)
|
||||||
|
from nanobot.utils.llm_runtime import runtime_from_provider_snapshot
|
||||||
|
|
||||||
|
|
||||||
def _write_config(tmp_path: Path, overrides: dict | None = None) -> Path:
|
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
|
@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.bus.events import OutboundMessage
|
||||||
from nanobot.providers.factory import ProviderSnapshot
|
from nanobot.providers.factory import ProviderSnapshot
|
||||||
|
|
||||||
config_path = _write_config(tmp_path)
|
config_path = _write_config(tmp_path)
|
||||||
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||||
original_model = bot._loop.model
|
assert not hasattr(bot, "_runtime_overrides")
|
||||||
active_models: list[str] = []
|
original_runtime = bot._loop.runtime_resolver.runtime
|
||||||
snapshot_base_models: list[str] = []
|
active_runs: list[tuple[str, str]] = []
|
||||||
first_entered = asyncio.Event()
|
first_entered = asyncio.Event()
|
||||||
|
both_entered = asyncio.Event()
|
||||||
release_first = 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 is not None
|
||||||
assert model_preset is None
|
assert model_preset is None
|
||||||
snapshot_base_models.append(bot._loop.model)
|
assert config is bot._config
|
||||||
return ProviderSnapshot(
|
return runtime_from_provider_snapshot(ProviderSnapshot(
|
||||||
provider=_fake_provider(model, max_tokens=2048),
|
provider=_fake_provider(model, max_tokens=2048),
|
||||||
model=model,
|
model=model,
|
||||||
context_window_tokens=4096,
|
context_window_tokens=4096,
|
||||||
signature=("sdk", model),
|
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):
|
async def fake_process_direct(message, *, session_key, hooks, runtime):
|
||||||
active_models.append(bot._loop.model)
|
active_runs.append((session_key, runtime.model))
|
||||||
|
assert bot._loop.runtime_resolver.runtime is original_runtime
|
||||||
if message == "first":
|
if message == "first":
|
||||||
first_entered.set()
|
first_entered.set()
|
||||||
|
if len(active_runs) == 2:
|
||||||
|
both_entered.set()
|
||||||
await asyncio.wait_for(release_first.wait(), timeout=1)
|
await asyncio.wait_for(release_first.wait(), timeout=1)
|
||||||
return OutboundMessage(channel="cli", chat_id="direct", content=message)
|
return OutboundMessage(channel="cli", chat_id="direct", content=message)
|
||||||
|
|
||||||
bot._loop.process_direct = fake_process_direct
|
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)
|
await asyncio.wait_for(first_entered.wait(), timeout=1)
|
||||||
|
|
||||||
second = asyncio.create_task(bot.run("second", model="model:second"))
|
second = asyncio.create_task(bot.run(
|
||||||
await asyncio.sleep(0)
|
"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()
|
assert not second.done()
|
||||||
|
|
||||||
release_first.set()
|
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 first_result.content == "first"
|
||||||
assert second_result.content == "second"
|
assert second_result.content == "second"
|
||||||
assert active_models == ["model:first", "model:second"]
|
assert set(active_runs) == {
|
||||||
assert snapshot_base_models == [original_model, original_model]
|
("sdk:first", "model:first"),
|
||||||
assert bot._loop.model == original_model
|
("sdk:second", "model:second"),
|
||||||
|
}
|
||||||
|
assert bot._loop.runtime_resolver.runtime is original_runtime
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@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.bus.events import OutboundMessage
|
||||||
from nanobot.providers.factory import ProviderSnapshot
|
from nanobot.providers.factory import ProviderSnapshot
|
||||||
|
|
||||||
config_path = _write_config(tmp_path)
|
config_path = _write_config(tmp_path)
|
||||||
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||||
original_provider = bot._loop.provider
|
original_runtime = bot._loop.runtime_resolver.runtime
|
||||||
original_model = bot._loop.model
|
|
||||||
original_signature = bot._loop._provider_signature
|
|
||||||
override_provider = _fake_provider("override-provider", max_tokens=2048)
|
override_provider = _fake_provider("override-provider", max_tokens=2048)
|
||||||
override = ProviderSnapshot(
|
override = ProviderSnapshot(
|
||||||
provider=override_provider,
|
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,
|
context_window_tokens=4096,
|
||||||
signature=("sdk", "override"),
|
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):
|
async def fake_process_direct(message, *, session_key, hooks, runtime):
|
||||||
assert bot._loop.provider is override_provider
|
assert runtime is override_runtime
|
||||||
assert not hasattr(bot._loop.runner, "provider")
|
assert not hasattr(bot._loop.runner, "provider")
|
||||||
assert bot._loop.model == "openai/gpt-4.1-mini"
|
assert runtime.model == "openai/gpt-4.1-mini"
|
||||||
assert bot._loop.context_window_tokens == 4096
|
assert runtime.context_window_tokens == 4096
|
||||||
|
assert bot._loop.runtime_resolver.runtime is original_runtime
|
||||||
return OutboundMessage(channel="cli", chat_id="direct", content="ok")
|
return OutboundMessage(channel="cli", chat_id="direct", content="ok")
|
||||||
|
|
||||||
bot._loop.process_direct = fake_process_direct
|
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")
|
result = await bot.run("hi", model="openai/gpt-4.1-mini")
|
||||||
|
|
||||||
assert result.content == "ok"
|
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="openai/gpt-4.1-mini",
|
||||||
model_preset=None,
|
model_preset=None,
|
||||||
|
config=bot._config,
|
||||||
)
|
)
|
||||||
assert bot._loop.provider is original_provider
|
|
||||||
assert not hasattr(bot._loop.runner, "provider")
|
assert not hasattr(bot._loop.runner, "provider")
|
||||||
assert bot._loop.model == original_model
|
assert bot._loop.runtime_resolver.runtime is original_runtime
|
||||||
assert bot._loop._provider_signature == original_signature
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@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)
|
config_path = _write_config(tmp_path)
|
||||||
bot = Nanobot.from_config(config_path, workspace=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_provider = _fake_provider("preset-provider", max_tokens=1024)
|
||||||
override = ProviderSnapshot(
|
override = ProviderSnapshot(
|
||||||
provider=override_provider,
|
provider=override_provider,
|
||||||
@ -611,19 +628,25 @@ async def test_run_model_preset_override_is_per_run(tmp_path):
|
|||||||
context_window_tokens=2048,
|
context_window_tokens=2048,
|
||||||
signature=("preset", "fast"),
|
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):
|
async def fake_process_direct(message, *, session_key, hooks, runtime):
|
||||||
assert bot._loop.provider is override_provider
|
assert runtime is override_runtime
|
||||||
assert bot._loop.model == "openai/gpt-4.1-mini"
|
|
||||||
return OutboundMessage(channel="cli", chat_id="direct", content="ok")
|
return OutboundMessage(channel="cli", chat_id="direct", content="ok")
|
||||||
|
|
||||||
bot._loop.process_direct = fake_process_direct
|
bot._loop.process_direct = fake_process_direct
|
||||||
|
|
||||||
await bot.run("hi", model_preset="fast")
|
await bot.run("hi", model_preset="fast")
|
||||||
|
|
||||||
bot._loop._build_model_preset_snapshot.assert_called_once_with("fast")
|
bot._loop.runtime_resolver.resolve_override.assert_called_once_with(
|
||||||
assert bot._loop.model == original_model
|
model=None,
|
||||||
|
model_preset="fast",
|
||||||
|
config=bot._config,
|
||||||
|
)
|
||||||
|
assert bot._loop.runtime_resolver.runtime is original_runtime
|
||||||
assert bot._loop.model_preset is None
|
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)
|
config_path = _write_config(tmp_path)
|
||||||
bot = Nanobot.from_config(config_path, workspace=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 message == "hi"
|
||||||
assert session_key == "sdk:default"
|
assert session_key == "sdk:default"
|
||||||
await on_stream("Hel")
|
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)
|
config_path = _write_config(tmp_path)
|
||||||
bot = Nanobot.from_config(config_path, workspace=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("done")
|
||||||
await on_stream_end(resuming=False)
|
await on_stream_end(resuming=False)
|
||||||
ctx = AgentRunHookContext(
|
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)
|
config_path = _write_config(tmp_path)
|
||||||
bot = Nanobot.from_config(config_path, workspace=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):
|
for i in range(400):
|
||||||
await on_stream(str(i))
|
await on_stream(str(i))
|
||||||
await on_stream_end(resuming=False)
|
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"])
|
||||||
assert callable(kwargs["on_stream_end"])
|
assert callable(kwargs["on_stream_end"])
|
||||||
assert kwargs["hooks"]
|
assert kwargs["hooks"]
|
||||||
|
assert kwargs["runtime"] is bot._loop.runtime_resolver.runtime
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@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.bus.events import OutboundMessage
|
||||||
from nanobot.providers.factory import ProviderSnapshot
|
from nanobot.providers.factory import ProviderSnapshot
|
||||||
|
|
||||||
config_path = _write_config(tmp_path)
|
config_path = _write_config(tmp_path)
|
||||||
bot = Nanobot.from_config(config_path, workspace=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_provider = _fake_provider("stream-provider", max_tokens=2048)
|
||||||
override = ProviderSnapshot(
|
override = ProviderSnapshot(
|
||||||
provider=override_provider,
|
provider=override_provider,
|
||||||
@ -903,11 +933,22 @@ async def test_run_streamed_model_override_reports_model_and_restores(tmp_path):
|
|||||||
context_window_tokens=4096,
|
context_window_tokens=4096,
|
||||||
signature=("sdk", "stream"),
|
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):
|
async def fake_process_direct(
|
||||||
assert bot._loop.provider is override_provider
|
message,
|
||||||
assert bot._loop.model == "openai/gpt-4.1-mini"
|
*,
|
||||||
|
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("ok")
|
||||||
await on_stream_end(resuming=False)
|
await on_stream_end(resuming=False)
|
||||||
return OutboundMessage(channel="cli", chat_id="direct", content="ok")
|
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].type == "run.started"
|
||||||
assert events[0].metadata["model"] == "openai/gpt-4.1-mini"
|
assert events[0].metadata["model"] == "openai/gpt-4.1-mini"
|
||||||
assert events[0].metadata["model_preset"] is None
|
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
|
@pytest.mark.asyncio
|
||||||
@ -947,7 +988,9 @@ async def test_run_streamed_emits_tool_events(tmp_path):
|
|||||||
config_path = _write_config(tmp_path)
|
config_path = _write_config(tmp_path)
|
||||||
bot = Nanobot.from_config(config_path, workspace=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 = [
|
calls = [
|
||||||
ToolCallRequest(id="call_ok", name="read_file", arguments={"path": "README.md"}),
|
ToolCallRequest(id="call_ok", name="read_file", arguments={"path": "README.md"}),
|
||||||
ToolCallRequest(id="call_bad", name="exec", arguments={"cmd": "false"}),
|
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)
|
config_path = _write_config(tmp_path)
|
||||||
bot = Nanobot.from_config(config_path, workspace=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:
|
for hook in hooks:
|
||||||
await hook.emit_reasoning("thinking")
|
await hook.emit_reasoning("thinking")
|
||||||
await hook.emit_reasoning_end()
|
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)
|
bot = Nanobot.from_config(config_path, workspace=tmp_path)
|
||||||
cancelled = asyncio.Event()
|
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:
|
try:
|
||||||
await on_stream("first")
|
await on_stream("first")
|
||||||
await asyncio.sleep(10)
|
await asyncio.sleep(10)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user