refactor(providers): move provider snapshot creation into factory
This commit is contained in:
+17
-17
@@ -18,7 +18,6 @@ from nanobot.agent.context import ContextBuilder
|
|||||||
from nanobot.agent.hook import AgentHook, AgentHookContext, CompositeHook
|
from nanobot.agent.hook import AgentHook, AgentHookContext, CompositeHook
|
||||||
from nanobot.agent.memory import Consolidator, Dream
|
from nanobot.agent.memory import Consolidator, Dream
|
||||||
from nanobot.agent.runner import _MAX_INJECTIONS_PER_TURN, AgentRunner, AgentRunSpec
|
from nanobot.agent.runner import _MAX_INJECTIONS_PER_TURN, AgentRunner, AgentRunSpec
|
||||||
from nanobot.agent.runtime import AgentRuntime
|
|
||||||
from nanobot.agent.skills import BUILTIN_SKILLS_DIR
|
from nanobot.agent.skills import BUILTIN_SKILLS_DIR
|
||||||
from nanobot.agent.subagent import SubagentManager
|
from nanobot.agent.subagent import SubagentManager
|
||||||
from nanobot.agent.tools.ask import (
|
from nanobot.agent.tools.ask import (
|
||||||
@@ -43,6 +42,7 @@ from nanobot.bus.queue import MessageBus
|
|||||||
from nanobot.command import CommandContext, CommandRouter, register_builtin_commands
|
from nanobot.command import CommandContext, CommandRouter, register_builtin_commands
|
||||||
from nanobot.config.schema import AgentDefaults
|
from nanobot.config.schema import AgentDefaults
|
||||||
from nanobot.providers.base import LLMProvider
|
from nanobot.providers.base import LLMProvider
|
||||||
|
from nanobot.providers.factory import ProviderSnapshot
|
||||||
from nanobot.session.manager import Session, SessionManager
|
from nanobot.session.manager import Session, SessionManager
|
||||||
from nanobot.utils.document import extract_documents
|
from nanobot.utils.document import extract_documents
|
||||||
from nanobot.utils.helpers import image_placeholder_text
|
from nanobot.utils.helpers import image_placeholder_text
|
||||||
@@ -196,8 +196,8 @@ class AgentLoop:
|
|||||||
unified_session: bool = False,
|
unified_session: bool = False,
|
||||||
disabled_skills: list[str] | None = None,
|
disabled_skills: list[str] | None = None,
|
||||||
tools_config: ToolsConfig | None = None,
|
tools_config: ToolsConfig | None = None,
|
||||||
runtime_loader: Callable[[], AgentRuntime] | None = None,
|
provider_snapshot_loader: Callable[[], ProviderSnapshot] | None = None,
|
||||||
runtime_signature: tuple[object, ...] | None = None,
|
provider_signature: tuple[object, ...] | None = None,
|
||||||
):
|
):
|
||||||
from nanobot.config.schema import ExecToolConfig, ToolsConfig, WebToolsConfig
|
from nanobot.config.schema import ExecToolConfig, ToolsConfig, WebToolsConfig
|
||||||
|
|
||||||
@@ -206,8 +206,8 @@ class AgentLoop:
|
|||||||
self.bus = bus
|
self.bus = bus
|
||||||
self.channels_config = channels_config
|
self.channels_config = channels_config
|
||||||
self.provider = provider
|
self.provider = provider
|
||||||
self._runtime_loader = runtime_loader
|
self._provider_snapshot_loader = provider_snapshot_loader
|
||||||
self._runtime_signature = runtime_signature
|
self._provider_signature = provider_signature
|
||||||
self.workspace = workspace
|
self.workspace = workspace
|
||||||
self.model = model or provider.get_default_model()
|
self.model = model or provider.get_default_model()
|
||||||
self.max_iterations = (
|
self.max_iterations = (
|
||||||
@@ -295,11 +295,11 @@ class AgentLoop:
|
|||||||
self.commands = CommandRouter()
|
self.commands = CommandRouter()
|
||||||
register_builtin_commands(self.commands)
|
register_builtin_commands(self.commands)
|
||||||
|
|
||||||
def _apply_runtime(self, runtime: AgentRuntime) -> None:
|
def _apply_provider_snapshot(self, snapshot: ProviderSnapshot) -> None:
|
||||||
"""Swap model/provider for future turns without disturbing an active one."""
|
"""Swap model/provider for future turns without disturbing an active one."""
|
||||||
provider = runtime.provider
|
provider = snapshot.provider
|
||||||
model = runtime.model
|
model = snapshot.model
|
||||||
context_window_tokens = runtime.context_window_tokens
|
context_window_tokens = snapshot.context_window_tokens
|
||||||
if self.provider is provider and self.model == model:
|
if self.provider is provider and self.model == model:
|
||||||
return
|
return
|
||||||
old_model = self.model
|
old_model = self.model
|
||||||
@@ -317,20 +317,20 @@ class AgentLoop:
|
|||||||
self.dream.provider = provider
|
self.dream.provider = provider
|
||||||
self.dream.model = model
|
self.dream.model = model
|
||||||
self.dream._runner.provider = provider
|
self.dream._runner.provider = provider
|
||||||
self._runtime_signature = runtime.signature
|
self._provider_signature = snapshot.signature
|
||||||
logger.info("Runtime model switched for next turn: {} -> {}", old_model, model)
|
logger.info("Runtime model switched for next turn: {} -> {}", old_model, model)
|
||||||
|
|
||||||
def _refresh_runtime(self) -> None:
|
def _refresh_provider_snapshot(self) -> None:
|
||||||
if self._runtime_loader is None:
|
if self._provider_snapshot_loader is None:
|
||||||
return
|
return
|
||||||
try:
|
try:
|
||||||
runtime = self._runtime_loader()
|
snapshot = self._provider_snapshot_loader()
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.exception("Failed to refresh runtime config")
|
logger.exception("Failed to refresh provider config")
|
||||||
return
|
return
|
||||||
if runtime.signature == self._runtime_signature:
|
if snapshot.signature == self._provider_signature:
|
||||||
return
|
return
|
||||||
self._apply_runtime(runtime)
|
self._apply_provider_snapshot(snapshot)
|
||||||
|
|
||||||
def _register_default_tools(self) -> None:
|
def _register_default_tools(self) -> None:
|
||||||
"""Register the default set of tools."""
|
"""Register the default set of tools."""
|
||||||
@@ -810,7 +810,7 @@ class AgentLoop:
|
|||||||
pending_queue: asyncio.Queue | None = None,
|
pending_queue: asyncio.Queue | None = None,
|
||||||
) -> OutboundMessage | None:
|
) -> OutboundMessage | None:
|
||||||
"""Process a single inbound message and return the response."""
|
"""Process a single inbound message and return the response."""
|
||||||
self._refresh_runtime()
|
self._refresh_provider_snapshot()
|
||||||
# System messages: parse origin from chat_id ("channel:chat_id")
|
# System messages: parse origin from chat_id ("channel:chat_id")
|
||||||
if msg.channel == "system":
|
if msg.channel == "system":
|
||||||
channel, chat_id = (
|
channel, chat_id = (
|
||||||
|
|||||||
@@ -412,7 +412,7 @@ def _make_provider(config: Config):
|
|||||||
|
|
||||||
Routing is driven by ``ProviderSpec.backend`` in the registry.
|
Routing is driven by ``ProviderSpec.backend`` in the registry.
|
||||||
"""
|
"""
|
||||||
from nanobot.agent.runtime import make_provider
|
from nanobot.providers.factory import make_provider
|
||||||
|
|
||||||
try:
|
try:
|
||||||
return make_provider(config)
|
return make_provider(config)
|
||||||
@@ -597,7 +597,6 @@ def _run_gateway(
|
|||||||
) -> None:
|
) -> None:
|
||||||
"""Shared gateway runtime; ``open_browser_url`` opens a tab once channels are up."""
|
"""Shared gateway runtime; ``open_browser_url`` opens a tab once channels are up."""
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
from nanobot.agent.runtime import build_agent_runtime, load_agent_runtime
|
|
||||||
from nanobot.agent.tools.cron import CronTool
|
from nanobot.agent.tools.cron import CronTool
|
||||||
from nanobot.agent.tools.message import MessageTool
|
from nanobot.agent.tools.message import MessageTool
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
@@ -605,6 +604,7 @@ def _run_gateway(
|
|||||||
from nanobot.cron.service import CronService
|
from nanobot.cron.service import CronService
|
||||||
from nanobot.cron.types import CronJob
|
from nanobot.cron.types import CronJob
|
||||||
from nanobot.heartbeat.service import HeartbeatService
|
from nanobot.heartbeat.service import HeartbeatService
|
||||||
|
from nanobot.providers.factory import build_provider_snapshot, load_provider_snapshot
|
||||||
from nanobot.session.manager import SessionManager
|
from nanobot.session.manager import SessionManager
|
||||||
|
|
||||||
port = port if port is not None else config.gateway.port
|
port = port if port is not None else config.gateway.port
|
||||||
@@ -613,11 +613,11 @@ def _run_gateway(
|
|||||||
sync_workspace_templates(config.workspace_path)
|
sync_workspace_templates(config.workspace_path)
|
||||||
bus = MessageBus()
|
bus = MessageBus()
|
||||||
try:
|
try:
|
||||||
runtime = build_agent_runtime(config)
|
provider_snapshot = build_provider_snapshot(config)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
console.print(f"[red]Error: {exc}[/red]")
|
console.print(f"[red]Error: {exc}[/red]")
|
||||||
raise typer.Exit(1) from exc
|
raise typer.Exit(1) from exc
|
||||||
provider = runtime.provider
|
provider = provider_snapshot.provider
|
||||||
session_manager = SessionManager(config.workspace_path)
|
session_manager = SessionManager(config.workspace_path)
|
||||||
|
|
||||||
# Preserve existing single-workspace installs, but keep custom workspaces clean.
|
# Preserve existing single-workspace installs, but keep custom workspaces clean.
|
||||||
@@ -633,9 +633,9 @@ def _run_gateway(
|
|||||||
bus=bus,
|
bus=bus,
|
||||||
provider=provider,
|
provider=provider,
|
||||||
workspace=config.workspace_path,
|
workspace=config.workspace_path,
|
||||||
model=runtime.model,
|
model=provider_snapshot.model,
|
||||||
max_iterations=config.agents.defaults.max_tool_iterations,
|
max_iterations=config.agents.defaults.max_tool_iterations,
|
||||||
context_window_tokens=runtime.context_window_tokens,
|
context_window_tokens=provider_snapshot.context_window_tokens,
|
||||||
web_config=config.tools.web,
|
web_config=config.tools.web,
|
||||||
context_block_limit=config.agents.defaults.context_block_limit,
|
context_block_limit=config.agents.defaults.context_block_limit,
|
||||||
max_tool_result_chars=config.agents.defaults.max_tool_result_chars,
|
max_tool_result_chars=config.agents.defaults.max_tool_result_chars,
|
||||||
@@ -652,8 +652,8 @@ def _run_gateway(
|
|||||||
session_ttl_minutes=config.agents.defaults.session_ttl_minutes,
|
session_ttl_minutes=config.agents.defaults.session_ttl_minutes,
|
||||||
consolidation_ratio=config.agents.defaults.consolidation_ratio,
|
consolidation_ratio=config.agents.defaults.consolidation_ratio,
|
||||||
tools_config=config.tools,
|
tools_config=config.tools,
|
||||||
runtime_loader=load_agent_runtime,
|
provider_snapshot_loader=load_provider_snapshot,
|
||||||
runtime_signature=runtime.signature,
|
provider_signature=provider_snapshot.signature,
|
||||||
)
|
)
|
||||||
|
|
||||||
from nanobot.agent.loop import UNIFIED_SESSION_KEY
|
from nanobot.agent.loop import UNIFIED_SESSION_KEY
|
||||||
|
|||||||
+1
-1
@@ -120,6 +120,6 @@ class Nanobot:
|
|||||||
|
|
||||||
def _make_provider(config: Any) -> Any:
|
def _make_provider(config: Any) -> Any:
|
||||||
"""Create the LLM provider from config (extracted from CLI)."""
|
"""Create the LLM provider from config (extracted from CLI)."""
|
||||||
from nanobot.agent.runtime import make_provider
|
from nanobot.providers.factory import make_provider
|
||||||
|
|
||||||
return make_provider(config)
|
return make_provider(config)
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
"""Runtime model/provider resolution for agent turns."""
|
"""Create LLM providers from config."""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
@@ -11,7 +11,7 @@ from nanobot.providers.registry import find_by_name
|
|||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class AgentRuntime:
|
class ProviderSnapshot:
|
||||||
provider: LLMProvider
|
provider: LLMProvider
|
||||||
model: str
|
model: str
|
||||||
context_window_tokens: int
|
context_window_tokens: int
|
||||||
@@ -80,8 +80,8 @@ def make_provider(config: Config) -> LLMProvider:
|
|||||||
return provider
|
return provider
|
||||||
|
|
||||||
|
|
||||||
def runtime_signature(config: Config) -> tuple[object, ...]:
|
def provider_signature(config: Config) -> tuple[object, ...]:
|
||||||
"""Return the config fields that affect the primary LLM runtime."""
|
"""Return the config fields that affect the primary LLM provider."""
|
||||||
model = config.agents.defaults.model
|
model = config.agents.defaults.model
|
||||||
defaults = config.agents.defaults
|
defaults = config.agents.defaults
|
||||||
return (
|
return (
|
||||||
@@ -97,16 +97,16 @@ def runtime_signature(config: Config) -> tuple[object, ...]:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def build_agent_runtime(config: Config) -> AgentRuntime:
|
def build_provider_snapshot(config: Config) -> ProviderSnapshot:
|
||||||
return AgentRuntime(
|
return ProviderSnapshot(
|
||||||
provider=make_provider(config),
|
provider=make_provider(config),
|
||||||
model=config.agents.defaults.model,
|
model=config.agents.defaults.model,
|
||||||
context_window_tokens=config.agents.defaults.context_window_tokens,
|
context_window_tokens=config.agents.defaults.context_window_tokens,
|
||||||
signature=runtime_signature(config),
|
signature=provider_signature(config),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def load_agent_runtime(config_path: Path | None = None) -> AgentRuntime:
|
def load_provider_snapshot(config_path: Path | None = None) -> ProviderSnapshot:
|
||||||
from nanobot.config.loader import load_config, resolve_config_env_vars
|
from nanobot.config.loader import load_config, resolve_config_env_vars
|
||||||
|
|
||||||
return build_agent_runtime(resolve_config_env_vars(load_config(config_path)))
|
return build_provider_snapshot(resolve_config_env_vars(load_config(config_path)))
|
||||||
@@ -3,8 +3,8 @@ from types import SimpleNamespace
|
|||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
from nanobot.agent.runtime import AgentRuntime
|
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.providers.factory import ProviderSnapshot
|
||||||
|
|
||||||
|
|
||||||
def _provider(default_model: str, max_tokens: int = 123) -> MagicMock:
|
def _provider(default_model: str, max_tokens: int = 123) -> MagicMock:
|
||||||
@@ -23,7 +23,7 @@ def test_runtime_refresh_updates_loop_dependents(tmp_path: Path) -> None:
|
|||||||
workspace=tmp_path,
|
workspace=tmp_path,
|
||||||
model="old-model",
|
model="old-model",
|
||||||
context_window_tokens=1000,
|
context_window_tokens=1000,
|
||||||
runtime_loader=lambda: AgentRuntime(
|
provider_snapshot_loader=lambda: ProviderSnapshot(
|
||||||
provider=new_provider,
|
provider=new_provider,
|
||||||
model="new-model",
|
model="new-model",
|
||||||
context_window_tokens=2000,
|
context_window_tokens=2000,
|
||||||
@@ -31,7 +31,7 @@ def test_runtime_refresh_updates_loop_dependents(tmp_path: Path) -> None:
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
loop._refresh_runtime()
|
loop._refresh_provider_snapshot()
|
||||||
|
|
||||||
assert loop.provider is new_provider
|
assert loop.provider is new_provider
|
||||||
assert loop.model == "new-model"
|
assert loop.model == "new-model"
|
||||||
|
|||||||
+15
-15
@@ -8,11 +8,11 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
|||||||
import pytest
|
import pytest
|
||||||
from typer.testing import CliRunner
|
from typer.testing import CliRunner
|
||||||
|
|
||||||
from nanobot.agent.runtime import AgentRuntime
|
|
||||||
from nanobot.bus.events import OutboundMessage
|
from nanobot.bus.events import OutboundMessage
|
||||||
from nanobot.cli.commands import _make_provider, app
|
from nanobot.cli.commands import _make_provider, app
|
||||||
from nanobot.config.schema import Config
|
from nanobot.config.schema import Config
|
||||||
from nanobot.cron.types import CronJob, CronPayload
|
from nanobot.cron.types import CronJob, CronPayload
|
||||||
|
from nanobot.providers.factory import ProviderSnapshot
|
||||||
from nanobot.providers.openai_codex_provider import _strip_model_prefix
|
from nanobot.providers.openai_codex_provider import _strip_model_prefix
|
||||||
from nanobot.providers.registry import find_by_name
|
from nanobot.providers.registry import find_by_name
|
||||||
|
|
||||||
@@ -777,8 +777,8 @@ def _stop_gateway_provider(_config) -> object:
|
|||||||
raise _StopGatewayError("stop")
|
raise _StopGatewayError("stop")
|
||||||
|
|
||||||
|
|
||||||
def _test_agent_runtime(provider: object, config: Config) -> AgentRuntime:
|
def _test_provider_snapshot(provider: object, config: Config) -> ProviderSnapshot:
|
||||||
return AgentRuntime(
|
return ProviderSnapshot(
|
||||||
provider=provider,
|
provider=provider,
|
||||||
model=config.agents.defaults.model,
|
model=config.agents.defaults.model,
|
||||||
context_window_tokens=config.agents.defaults.context_window_tokens,
|
context_window_tokens=config.agents.defaults.context_window_tokens,
|
||||||
@@ -815,12 +815,12 @@ def _patch_cli_command_runtime(
|
|||||||
provider_factory,
|
provider_factory,
|
||||||
)
|
)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"nanobot.agent.runtime.build_agent_runtime",
|
"nanobot.providers.factory.build_provider_snapshot",
|
||||||
lambda _config: _test_agent_runtime(provider_factory(_config), _config),
|
lambda _config: _test_provider_snapshot(provider_factory(_config), _config),
|
||||||
)
|
)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"nanobot.agent.runtime.load_agent_runtime",
|
"nanobot.providers.factory.load_provider_snapshot",
|
||||||
lambda _config_path=None: _test_agent_runtime(provider_factory(config), config),
|
lambda _config_path=None: _test_provider_snapshot(provider_factory(config), config),
|
||||||
)
|
)
|
||||||
|
|
||||||
if message_bus is not None:
|
if message_bus is not None:
|
||||||
@@ -962,12 +962,12 @@ def test_gateway_cron_evaluator_receives_scheduled_reminder_context(
|
|||||||
monkeypatch.setattr("nanobot.cli.commands.sync_workspace_templates", lambda _path: None)
|
monkeypatch.setattr("nanobot.cli.commands.sync_workspace_templates", lambda _path: None)
|
||||||
monkeypatch.setattr("nanobot.cli.commands._make_provider", lambda _config: provider)
|
monkeypatch.setattr("nanobot.cli.commands._make_provider", lambda _config: provider)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"nanobot.agent.runtime.build_agent_runtime",
|
"nanobot.providers.factory.build_provider_snapshot",
|
||||||
lambda _config: _test_agent_runtime(provider, _config),
|
lambda _config: _test_provider_snapshot(provider, _config),
|
||||||
)
|
)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"nanobot.agent.runtime.load_agent_runtime",
|
"nanobot.providers.factory.load_provider_snapshot",
|
||||||
lambda _config_path=None: _test_agent_runtime(provider, config),
|
lambda _config_path=None: _test_provider_snapshot(provider, config),
|
||||||
)
|
)
|
||||||
monkeypatch.setattr("nanobot.bus.queue.MessageBus", lambda: bus)
|
monkeypatch.setattr("nanobot.bus.queue.MessageBus", lambda: bus)
|
||||||
|
|
||||||
@@ -1111,12 +1111,12 @@ def test_gateway_cron_job_suppresses_intermediate_progress(
|
|||||||
monkeypatch.setattr("nanobot.cli.commands.sync_workspace_templates", lambda _path: None)
|
monkeypatch.setattr("nanobot.cli.commands.sync_workspace_templates", lambda _path: None)
|
||||||
monkeypatch.setattr("nanobot.cli.commands._make_provider", lambda _config: object())
|
monkeypatch.setattr("nanobot.cli.commands._make_provider", lambda _config: object())
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"nanobot.agent.runtime.build_agent_runtime",
|
"nanobot.providers.factory.build_provider_snapshot",
|
||||||
lambda _config: _test_agent_runtime(object(), _config),
|
lambda _config: _test_provider_snapshot(object(), _config),
|
||||||
)
|
)
|
||||||
monkeypatch.setattr(
|
monkeypatch.setattr(
|
||||||
"nanobot.agent.runtime.load_agent_runtime",
|
"nanobot.providers.factory.load_provider_snapshot",
|
||||||
lambda _config_path=None: _test_agent_runtime(object(), config),
|
lambda _config_path=None: _test_provider_snapshot(object(), config),
|
||||||
)
|
)
|
||||||
monkeypatch.setattr("nanobot.bus.queue.MessageBus", lambda: bus)
|
monkeypatch.setattr("nanobot.bus.queue.MessageBus", lambda: bus)
|
||||||
monkeypatch.setattr("nanobot.session.manager.SessionManager", lambda _workspace: object())
|
monkeypatch.setattr("nanobot.session.manager.SessionManager", lambda _workspace: object())
|
||||||
|
|||||||
Reference in New Issue
Block a user