Compare commits

...
15 changed files with 480 additions and 264 deletions
+35 -171
View File
@@ -24,12 +24,15 @@ from nanobot.agent.hook import AgentHook, CompositeHook
from nanobot.agent.memory import Consolidator
from nanobot.agent.progress_hook import AgentProgressHook
from nanobot.agent.runner import _MAX_INJECTIONS_PER_TURN, AgentRunner, AgentRunSpec
from nanobot.agent.runtime_model import RuntimeModelCoordinator
from nanobot.agent.streaming import StreamingCoordinator
from nanobot.agent.subagent import SubagentManager
from nanobot.agent.tools.context import RequestContext, bind_request_context, reset_request_context
from nanobot.agent.tools.file_state import FileStateStore, bind_file_states, reset_file_states
from nanobot.agent.tools.message import MessageTool
from nanobot.agent.tools.registry import ToolRegistry
from nanobot.agent.tools.self import MyTool
from nanobot.agent.turn_session import TurnSessionCoordinator
from nanobot.bus.events import InboundMessage, OutboundMessage
from nanobot.bus.progress import build_bus_progress_callback
from nanobot.bus.queue import MessageBus
@@ -160,11 +163,10 @@ class AgentLoop:
def llm_runtime(self) -> LLMRuntime:
"""Return the current provider/model pair owned by this loop."""
self._refresh_provider_snapshot()
return LLMRuntime(self.provider, self.model)
return self.runtime_models.llm_runtime()
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
_PENDING_USER_TURN_KEY = "pending_user_turn"
_RUNTIME_CHECKPOINT_KEY = TurnSessionCoordinator.RUNTIME_CHECKPOINT_KEY
_PENDING_USER_TURN_KEY = TurnSessionCoordinator.PENDING_USER_TURN_KEY
# Event-driven state transition table.
# Handlers return an event string; the driver looks up the next state here.
@@ -220,6 +222,7 @@ class AgentLoop:
_tc = tools_config or ToolsConfig()
defaults = AgentDefaults()
self.bus = bus
self.streaming = StreamingCoordinator(bus)
self.runtime_events = runtime_events or RuntimeEventBus()
self.runtime_event_publisher = RuntimeEventPublisher(self.runtime_events)
self.channels_config = channels_config
@@ -271,6 +274,7 @@ class AgentLoop:
self.context = ContextBuilder(workspace, timezone=timezone, disabled_skills=disabled_skills)
self.sessions = session_manager or SessionManager(workspace)
self.turn_sessions = TurnSessionCoordinator(self.sessions)
self.tools = ToolRegistry()
# One file-read/write tracker per logical session. The tool registry is
# shared by this loop, so tools resolve the active state via contextvars.
@@ -332,6 +336,7 @@ class AgentLoop:
)
self.model_presets: dict[str, ModelPresetConfig] = model_presets or {}
self._active_preset: str | None = None
self.runtime_models = RuntimeModelCoordinator(self)
if model_preset:
self.set_model_preset(model_preset, publish_update=False)
self._register_default_tools()
@@ -398,7 +403,7 @@ class AgentLoop:
def _sync_subagent_runtime_limits(self) -> None:
"""Keep subagent runtime limits aligned with mutable loop settings."""
self.subagents.max_iterations = self.max_iterations
self.runtime_models.sync_subagent_runtime_limits()
def _apply_provider_snapshot(
self,
@@ -408,52 +413,14 @@ class AgentLoop:
model_preset: str | None = None,
) -> None:
"""Swap model/provider for future turns without disturbing an active one."""
provider = snapshot.provider
model = snapshot.model
context_window_tokens = snapshot.context_window_tokens
old_model = self.model
self.provider = provider
self.model = model
self.context_window_tokens = context_window_tokens
self.runner.provider = provider
self.subagents.set_provider(provider, model)
self.consolidator.set_provider(provider, model, context_window_tokens)
self._provider_signature = snapshot.signature
if publish_update and self._runtime_model_publisher is not None:
self._runtime_model_publisher(
self.model,
model_preset if model_preset is not None else self.model_preset,
)
if publish_update:
self._runtime_events().runtime_model_changed(
self.model,
model_preset if model_preset is not None else self.model_preset,
)
logger.info("Runtime model switched for next turn: {} -> {}", old_model, model)
self.runtime_models.apply_provider_snapshot(
snapshot,
publish_update=publish_update,
model_preset=model_preset,
)
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)
self.runtime_models.refresh_provider_snapshot()
@property
def model_preset(self) -> str | None:
@@ -464,19 +431,11 @@ class AgentLoop:
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,
)
return self.runtime_models.build_model_preset_snapshot(name)
def set_model_preset(self, name: str | None, *, publish_update: bool = True) -> None:
"""Resolve a preset by name and apply all runtime model dependents."""
name = preset_helpers.normalize_preset_name(name, self.model_presets)
snapshot = self._build_model_preset_snapshot(name)
self._apply_provider_snapshot(snapshot, publish_update=publish_update, model_preset=name)
self._active_preset = name
self.runtime_models.set_model_preset(name, publish_update=publish_update)
def _register_default_tools(self) -> None:
"""Register the default set of tools via plugin loader."""
@@ -969,35 +928,9 @@ class AgentLoop:
try:
on_stream = on_stream_end = None
if msg.metadata.get("_wants_stream"):
# Split one answer into distinct stream segments.
stream_base_id = f"{msg.session_key}:{time.time_ns()}"
stream_segment = 0
def _current_stream_id() -> str:
return f"{stream_base_id}:{stream_segment}"
async def on_stream(delta: str) -> None:
meta = dict(msg.metadata or {})
meta["_stream_delta"] = True
meta["_stream_id"] = _current_stream_id()
await self.bus.publish_outbound(OutboundMessage(
channel=msg.channel, chat_id=msg.chat_id,
content=delta,
metadata=meta,
))
async def on_stream_end(*, resuming: bool = False) -> None:
nonlocal stream_segment
meta = dict(msg.metadata or {})
meta["_stream_end"] = True
meta["_resuming"] = resuming
meta["_stream_id"] = _current_stream_id()
await self.bus.publish_outbound(OutboundMessage(
channel=msg.channel, chat_id=msg.chat_id,
content="",
metadata=meta,
))
stream_segment += 1
callbacks = self.streaming.build_callbacks(msg)
on_stream = callbacks.on_stream
on_stream_end = callbacks.on_stream_end
response = await self._process_message(
msg, on_stream=on_stream, on_stream_end=on_stream_end,
@@ -1697,104 +1630,35 @@ class AgentLoop:
def _set_runtime_checkpoint(self, session: Session, payload: dict[str, Any]) -> None:
"""Persist the latest in-flight turn state into session metadata."""
session.metadata[self._RUNTIME_CHECKPOINT_KEY] = payload
self.sessions.save(session)
self._turn_session_coordinator().set_runtime_checkpoint(session, payload)
def _turn_session_coordinator(self) -> TurnSessionCoordinator:
coordinator = getattr(self, "turn_sessions", None)
if coordinator is None:
coordinator = TurnSessionCoordinator(getattr(self, "sessions", None))
self.turn_sessions = coordinator
return coordinator
def _mark_pending_user_turn(self, session: Session) -> None:
session.metadata[self._PENDING_USER_TURN_KEY] = True
self._turn_session_coordinator().mark_pending_user_turn(session)
def _clear_pending_user_turn(self, session: Session) -> None:
session.metadata.pop(self._PENDING_USER_TURN_KEY, None)
self._turn_session_coordinator().clear_pending_user_turn(session)
def _clear_runtime_checkpoint(self, session: Session) -> None:
if self._RUNTIME_CHECKPOINT_KEY in session.metadata:
session.metadata.pop(self._RUNTIME_CHECKPOINT_KEY, None)
self._turn_session_coordinator().clear_runtime_checkpoint(session)
@staticmethod
def _checkpoint_message_key(message: dict[str, Any]) -> tuple[Any, ...]:
return (
message.get("role"),
message.get("content"),
message.get("tool_call_id"),
message.get("name"),
message.get("tool_calls"),
message.get("reasoning_content"),
message.get("thinking_blocks"),
)
return TurnSessionCoordinator.checkpoint_message_key(message)
def _restore_runtime_checkpoint(self, session: Session) -> bool:
"""Materialize an unfinished turn into session history before a new request."""
from datetime import datetime
checkpoint = session.metadata.get(self._RUNTIME_CHECKPOINT_KEY)
if not isinstance(checkpoint, dict):
return False
assistant_message = checkpoint.get("assistant_message")
completed_tool_results = checkpoint.get("completed_tool_results") or []
pending_tool_calls = checkpoint.get("pending_tool_calls") or []
restored_messages: list[dict[str, Any]] = []
if isinstance(assistant_message, dict):
restored = dict(assistant_message)
restored.setdefault("timestamp", datetime.now().isoformat())
restored_messages.append(restored)
for message in completed_tool_results:
if isinstance(message, dict):
restored = dict(message)
restored.setdefault("timestamp", datetime.now().isoformat())
restored_messages.append(restored)
for tool_call in pending_tool_calls:
if not isinstance(tool_call, dict):
continue
tool_id = tool_call.get("id")
name = ((tool_call.get("function") or {}).get("name")) or "tool"
restored_messages.append(
{
"role": "tool",
"tool_call_id": tool_id,
"name": name,
"content": "Error: Task interrupted before this tool finished.",
"timestamp": datetime.now().isoformat(),
}
)
overlap = 0
max_overlap = min(len(session.messages), len(restored_messages))
for size in range(max_overlap, 0, -1):
existing = session.messages[-size:]
restored = restored_messages[:size]
if all(
self._checkpoint_message_key(left) == self._checkpoint_message_key(right)
for left, right in zip(existing, restored)
):
overlap = size
break
session.messages.extend(restored_messages[overlap:])
self._clear_pending_user_turn(session)
self._clear_runtime_checkpoint(session)
return True
return self._turn_session_coordinator().restore_runtime_checkpoint(session)
def _restore_pending_user_turn(self, session: Session) -> bool:
"""Close a turn that only persisted the user message before crashing."""
from datetime import datetime
if not session.metadata.get(self._PENDING_USER_TURN_KEY):
return False
if session.messages and session.messages[-1].get("role") == "user":
session.messages.append(
{
"role": "assistant",
"content": "Error: Task interrupted before a response was generated.",
"timestamp": datetime.now().isoformat(),
}
)
session.updated_at = datetime.now()
self._clear_pending_user_turn(session)
return True
return self._turn_session_coordinator().restore_pending_user_turn(session)
async def process_direct(
self,
+107
View File
@@ -0,0 +1,107 @@
"""Runtime model/provider coordination for the agent loop."""
from __future__ import annotations
from typing import TYPE_CHECKING
from loguru import logger
from nanobot.agent import model_presets as preset_helpers
from nanobot.providers.factory import ProviderSnapshot
from nanobot.utils.llm_runtime import LLMRuntime
if TYPE_CHECKING:
from nanobot.agent.loop import AgentLoop
class RuntimeModelCoordinator:
"""Owns mutable provider/model state transitions for an AgentLoop."""
def __init__(self, loop: AgentLoop) -> None:
self._loop = loop
def llm_runtime(self) -> LLMRuntime:
"""Return the current provider/model pair owned by the loop."""
self.refresh_provider_snapshot()
return LLMRuntime(self._loop.provider, self._loop.model)
def sync_subagent_runtime_limits(self) -> None:
"""Keep subagent runtime limits aligned with mutable loop settings."""
self._loop.subagents.max_iterations = self._loop.max_iterations
def apply_provider_snapshot(
self,
snapshot: ProviderSnapshot,
*,
publish_update: bool = True,
model_preset: str | None = None,
) -> None:
"""Swap model/provider for future turns without disturbing an active one."""
loop = self._loop
provider = snapshot.provider
model = snapshot.model
context_window_tokens = snapshot.context_window_tokens
old_model = loop.model
loop.provider = provider
loop.model = model
loop.context_window_tokens = context_window_tokens
loop.runner.provider = provider
loop.subagents.set_provider(provider, model)
loop.consolidator.set_provider(provider, model, context_window_tokens)
loop._provider_signature = snapshot.signature
active_preset = model_preset if model_preset is not None else loop.model_preset
if publish_update and loop._runtime_model_publisher is not None:
loop._runtime_model_publisher(loop.model, active_preset)
if publish_update:
loop._runtime_events().runtime_model_changed(loop.model, active_preset)
logger.info("Runtime model switched for next turn: {} -> {}", old_model, model)
def refresh_provider_snapshot(self) -> None:
"""Refresh runtime provider state from the configured snapshot loader."""
loop = self._loop
if loop._provider_snapshot_loader is None:
return
try:
snapshot = loop._provider_snapshot_loader()
except Exception:
logger.exception("Failed to refresh provider config")
return
default_selection = preset_helpers.default_selection_signature(snapshot.signature)
if loop._active_preset and loop._default_selection_signature in (None, default_selection):
loop._default_selection_signature = default_selection
try:
snapshot = self.build_model_preset_snapshot(loop._active_preset)
except Exception:
logger.exception("Failed to refresh active model preset")
return
else:
loop._active_preset = None
loop._default_selection_signature = default_selection
if snapshot.signature == loop._provider_signature:
return
loop._default_selection_signature = preset_helpers.default_selection_signature(
snapshot.signature
)
self.apply_provider_snapshot(snapshot)
def build_model_preset_snapshot(self, name: str) -> ProviderSnapshot:
"""Resolve a preset into a provider snapshot."""
loop = self._loop
return preset_helpers.build_runtime_preset_snapshot(
name=name,
presets=loop.model_presets,
provider=loop.provider,
loader=loop._preset_snapshot_loader,
)
def set_model_preset(self, name: str | None, *, publish_update: bool = True) -> None:
"""Resolve a preset by name and apply all runtime model dependents."""
loop = self._loop
normalized = preset_helpers.normalize_preset_name(name, loop.model_presets)
snapshot = self.build_model_preset_snapshot(normalized)
self.apply_provider_snapshot(
snapshot,
publish_update=publish_update,
model_preset=normalized,
)
loop._active_preset = normalized
+62
View File
@@ -0,0 +1,62 @@
"""Streaming response callbacks for bus-backed channels."""
from __future__ import annotations
import time
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from nanobot.bus.events import InboundMessage, OutboundMessage
from nanobot.bus.queue import MessageBus
@dataclass(frozen=True)
class StreamCallbacks:
on_stream: Callable[[str], Awaitable[None]]
on_stream_end: Callable[..., Awaitable[None]]
class StreamingCoordinator:
"""Builds outbound bus callbacks for segmented response streaming."""
def __init__(self, bus: MessageBus) -> None:
self._bus = bus
def build_callbacks(self, msg: InboundMessage) -> StreamCallbacks:
"""Split one answer into stream segments and publish deltas to the bus."""
stream_base_id = f"{msg.session_key}:{time.time_ns()}"
stream_segment = 0
def _current_stream_id() -> str:
return f"{stream_base_id}:{stream_segment}"
async def on_stream(delta: str) -> None:
meta = dict(msg.metadata or {})
meta["_stream_delta"] = True
meta["_stream_id"] = _current_stream_id()
await self._bus.publish_outbound(
OutboundMessage(
channel=msg.channel,
chat_id=msg.chat_id,
content=delta,
metadata=meta,
)
)
async def on_stream_end(*, resuming: bool = False) -> None:
nonlocal stream_segment
meta = dict(msg.metadata or {})
meta["_stream_end"] = True
meta["_resuming"] = resuming
meta["_stream_id"] = _current_stream_id()
await self._bus.publish_outbound(
OutboundMessage(
channel=msg.channel,
chat_id=msg.chat_id,
content="",
metadata=meta,
)
)
stream_segment += 1
return StreamCallbacks(on_stream=on_stream, on_stream_end=on_stream_end)
+77
View File
@@ -0,0 +1,77 @@
"""Tool-owned configuration parsing helpers."""
from __future__ import annotations
from typing import Any
from loguru import logger
from pydantic import BaseModel
_CONFIG_CLASSES_BY_KEY: dict[str, type[BaseModel]] | None = None
def _extra_values(config: BaseModel) -> dict[str, Any]:
return getattr(config, "__pydantic_extra__", None) or {}
def _set_extra_value(config: BaseModel, key: str, value: Any) -> None:
setattr(config, key, value)
def _config_classes_by_key() -> dict[str, type[BaseModel]]:
global _CONFIG_CLASSES_BY_KEY
if _CONFIG_CLASSES_BY_KEY is not None:
return _CONFIG_CLASSES_BY_KEY
from nanobot.agent.tools.loader import ToolLoader
classes: dict[str, type[BaseModel]] = {}
for tool_cls in ToolLoader().discover_config_classes():
key = getattr(tool_cls, "config_key", "")
config_cls = tool_cls.config_cls()
if not key or config_cls is None:
continue
previous = classes.get(key)
if previous is not None and previous is not config_cls:
logger.warning(
"Tool config key collision for %s: %s replaces %s",
key,
config_cls.__name__,
previous.__name__,
)
classes[key] = config_cls
_CONFIG_CLASSES_BY_KEY = classes
return classes
def _materialize_config(config: BaseModel, key: str, config_cls: type[BaseModel]) -> BaseModel:
raw = _extra_values(config).get(key, None)
if isinstance(raw, config_cls):
return raw
if raw is None:
parsed = config_cls()
elif isinstance(raw, BaseModel):
parsed = config_cls.model_validate(raw.model_dump(mode="python"))
else:
parsed = config_cls.model_validate(raw)
_set_extra_value(config, key, parsed)
return parsed
def tool_config_by_key(config: Any, key: str) -> Any:
"""Return the parsed config section for a tool config key."""
if not isinstance(config, BaseModel):
return getattr(config, key)
config_cls = _config_classes_by_key().get(key)
if config_cls is None:
raise KeyError(key)
return _materialize_config(config, key, config_cls)
def materialize_tool_configs(config: Any) -> Any:
"""Parse all discoverable tool config sections on a ToolsConfig object."""
if not isinstance(config, BaseModel):
return config
for key, config_cls in _config_classes_by_key().items():
_materialize_config(config, key, config_cls)
return config
+15 -2
View File
@@ -28,10 +28,15 @@ class ToolLoader:
self._plugins: dict[str, type[Tool]] | None = None
def discover(self) -> list[type[Tool]]:
"""Discover concrete tools that should be registered automatically."""
if self._test_classes is not None:
return list(self._test_classes)
if self._discovered is not None:
return self._discovered
self._discovered = self._discover_package_tools(include_non_discoverable=False)
return self._discovered
def _discover_package_tools(self, *, include_non_discoverable: bool) -> list[type[Tool]]:
seen: set[int] = set()
results: list[type[Tool]] = []
for _importer, module_name, _ispkg in pkgutil.iter_modules(self._package.__path__):
@@ -50,15 +55,23 @@ class ToolLoader:
and attr is not Tool
and not attr_name.startswith("_")
and not getattr(attr, "__abstractmethods__", None)
and getattr(attr, "_plugin_discoverable", True)
and (include_non_discoverable or getattr(attr, "_plugin_discoverable", True))
and id(attr) not in seen
):
seen.add(id(attr))
results.append(attr)
results.sort(key=lambda cls: cls.__name__)
self._discovered = results
return results
def discover_config_classes(self) -> list[type[Tool]]:
"""Discover tool classes that declare owned config models."""
classes = self._discover_package_tools(include_non_discoverable=True)
classes.extend(self._discover_plugins().values())
return [
cls for cls in classes
if getattr(cls, "config_key", "") and cls.config_cls() is not None
]
def _discover_plugins(self) -> dict[str, type[Tool]]:
"""Discover external tool plugins registered via entry_points."""
if self._plugins is not None:
+115
View File
@@ -0,0 +1,115 @@
"""Turn-scoped session metadata coordination."""
from __future__ import annotations
from datetime import datetime
from typing import Any
from nanobot.session.manager import Session, SessionManager
class TurnSessionCoordinator:
"""Owns in-flight turn metadata stored on a session."""
RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
PENDING_USER_TURN_KEY = "pending_user_turn"
def __init__(self, sessions: SessionManager | None) -> None:
self._sessions = sessions
def set_runtime_checkpoint(self, session: Session, payload: dict[str, Any]) -> None:
"""Persist the latest in-flight turn state into session metadata."""
session.metadata[self.RUNTIME_CHECKPOINT_KEY] = payload
if self._sessions is not None:
self._sessions.save(session)
def mark_pending_user_turn(self, session: Session) -> None:
session.metadata[self.PENDING_USER_TURN_KEY] = True
def clear_pending_user_turn(self, session: Session) -> None:
session.metadata.pop(self.PENDING_USER_TURN_KEY, None)
def clear_runtime_checkpoint(self, session: Session) -> None:
session.metadata.pop(self.RUNTIME_CHECKPOINT_KEY, None)
@staticmethod
def checkpoint_message_key(message: dict[str, Any]) -> tuple[Any, ...]:
return (
message.get("role"),
message.get("content"),
message.get("tool_call_id"),
message.get("name"),
message.get("tool_calls"),
message.get("reasoning_content"),
message.get("thinking_blocks"),
)
def restore_runtime_checkpoint(self, session: Session) -> bool:
"""Materialize an unfinished turn into session history before a new request."""
checkpoint = session.metadata.get(self.RUNTIME_CHECKPOINT_KEY)
if not isinstance(checkpoint, dict):
return False
assistant_message = checkpoint.get("assistant_message")
completed_tool_results = checkpoint.get("completed_tool_results") or []
pending_tool_calls = checkpoint.get("pending_tool_calls") or []
restored_messages: list[dict[str, Any]] = []
if isinstance(assistant_message, dict):
restored = dict(assistant_message)
restored.setdefault("timestamp", datetime.now().isoformat())
restored_messages.append(restored)
for message in completed_tool_results:
if isinstance(message, dict):
restored = dict(message)
restored.setdefault("timestamp", datetime.now().isoformat())
restored_messages.append(restored)
for tool_call in pending_tool_calls:
if not isinstance(tool_call, dict):
continue
tool_id = tool_call.get("id")
name = ((tool_call.get("function") or {}).get("name")) or "tool"
restored_messages.append(
{
"role": "tool",
"tool_call_id": tool_id,
"name": name,
"content": "Error: Task interrupted before this tool finished.",
"timestamp": datetime.now().isoformat(),
}
)
overlap = 0
max_overlap = min(len(session.messages), len(restored_messages))
for size in range(max_overlap, 0, -1):
existing = session.messages[-size:]
restored = restored_messages[:size]
if all(
self.checkpoint_message_key(left) == self.checkpoint_message_key(right)
for left, right in zip(existing, restored)
):
overlap = size
break
session.messages.extend(restored_messages[overlap:])
self.clear_pending_user_turn(session)
self.clear_runtime_checkpoint(session)
return True
def restore_pending_user_turn(self, session: Session) -> bool:
"""Close a turn that only persisted the user message before crashing."""
if not session.metadata.get(self.PENDING_USER_TURN_KEY):
return False
if session.messages and session.messages[-1].get("role") == "user":
session.messages.append(
{
"role": "assistant",
"content": "Error: Task interrupted before a response was generated.",
"timestamp": datetime.now().isoformat(),
}
)
session.updated_at = datetime.now()
self.clear_pending_user_turn(session)
return True
+10 -7
View File
@@ -9,11 +9,10 @@ from typing import Any
import pydantic
from pydantic import BaseModel
from nanobot.config.schema import Config, _resolve_tool_config_refs
from nanobot.config.schema import Config
# Global variable to store current config path (for multi-instance support)
_current_config_path: Path | None = None
_schema_refs_ready = False
def set_config_path(path: Path) -> None:
@@ -39,11 +38,6 @@ def load_config(config_path: Path | None = None) -> Config:
Returns:
Loaded configuration object.
"""
global _schema_refs_ready
if not _schema_refs_ready:
_resolve_tool_config_refs()
_schema_refs_ready = True
path = config_path or get_config_path()
config = Config()
@@ -56,6 +50,12 @@ def load_config(config_path: Path | None = None) -> Config:
except (json.JSONDecodeError, ValueError, pydantic.ValidationError) as e:
raise ValueError(f"Failed to load config from {path}: {e}") from e
from nanobot.agent.tools.config import materialize_tool_configs
try:
materialize_tool_configs(config.tools)
except (ValueError, pydantic.ValidationError) as e:
raise ValueError(f"Failed to load config from {path}: {e}") from e
_apply_ssrf_whitelist(config)
return config
@@ -78,6 +78,9 @@ def save_config(config: Config, config_path: Path | None = None) -> None:
path = config_path or get_config_path()
path.parent.mkdir(parents=True, exist_ok=True)
from nanobot.agent.tools.config import materialize_tool_configs
materialize_tool_configs(config.tools)
data = config.model_dump(mode="json", by_alias=True)
with open(path, "w", encoding="utf-8") as f:
+19 -72
View File
@@ -2,7 +2,7 @@
from __future__ import annotations
from pathlib import Path
from typing import TYPE_CHECKING, Any, Literal
from typing import Any, Literal
from pydantic import AliasChoices, ConfigDict, Field, model_validator
from pydantic_settings import BaseSettings
@@ -10,14 +10,6 @@ from pydantic_settings import BaseSettings
from nanobot.config_base import Base
from nanobot.cron.types import CronSchedule
if TYPE_CHECKING:
from nanobot.agent.tools.cli_apps import CliAppsToolConfig
from nanobot.agent.tools.filesystem import FileToolsConfig
from nanobot.agent.tools.image_generation import ImageGenerationToolConfig
from nanobot.agent.tools.self import MyToolConfig
from nanobot.agent.tools.shell import ExecToolConfig
from nanobot.agent.tools.web import WebToolsConfig
class ChannelsConfig(Base):
"""Configuration for chat channels.
@@ -304,29 +296,16 @@ class MCPServerConfig(Base):
enabled_tools: list[str] = Field(default_factory=lambda: ["*"]) # Only register these tools; accepts raw MCP names or wrapped mcp_<server>_<tool> names; ["*"] = all tools; [] = no tools
def _lazy_default(module_path: str, class_name: str) -> Any:
"""Deferred import helper for ToolsConfig default factories."""
import importlib
module = importlib.import_module(module_path)
return getattr(module, class_name)()
class ToolsConfig(Base):
"""Tools configuration.
Field types for tool-specific sub-configs are resolved via model_rebuild()
at the bottom of this file so tool config classes can stay next to their
tool implementations.
Concrete tool sub-configs are stored as extra fields and parsed by the
owning tool module when tools are loaded. This keeps the root schema from
importing or naming concrete tool configuration classes.
"""
web: WebToolsConfig = Field(default_factory=lambda: _lazy_default("nanobot.agent.tools.web", "WebToolsConfig"))
exec: ExecToolConfig = Field(default_factory=lambda: _lazy_default("nanobot.agent.tools.shell", "ExecToolConfig"))
file: FileToolsConfig = Field(default_factory=lambda: _lazy_default("nanobot.agent.tools.filesystem", "FileToolsConfig"))
cli_apps: CliAppsToolConfig = Field(default_factory=lambda: _lazy_default("nanobot.agent.tools.cli_apps", "CliAppsToolConfig"))
my: MyToolConfig = Field(default_factory=lambda: _lazy_default("nanobot.agent.tools.self", "MyToolConfig"))
image_generation: ImageGenerationToolConfig = Field(
default_factory=lambda: _lazy_default("nanobot.agent.tools.image_generation", "ImageGenerationToolConfig"),
)
model_config = ConfigDict(extra="allow")
restrict_to_workspace: bool = False # policy intent: keep tool access inside workspace when possible
webui_allow_local_service_access: bool = Field(
default=True,
@@ -340,6 +319,19 @@ class ToolsConfig(Base):
mcp_servers: dict[str, MCPServerConfig] = Field(default_factory=dict)
ssrf_whitelist: list[str] = Field(default_factory=list) # CIDR ranges to exempt from SSRF blocking (e.g. ["100.64.0.0/10"] for Tailscale)
def __getattr__(self, name: str) -> Any:
try:
return super().__getattr__(name)
except AttributeError as exc:
if name.startswith("_"):
raise
from nanobot.agent.tools.config import tool_config_by_key
try:
return tool_config_by_key(self, name)
except KeyError:
raise exc from None
class Config(BaseSettings):
"""Root configuration for nanobot."""
@@ -356,11 +348,6 @@ class Config(BaseSettings):
validation_alias=AliasChoices("modelPresets", "model_presets"),
)
def __init__(self, **values: Any) -> None:
if not type(self).__pydantic_complete__:
_resolve_tool_config_refs()
super().__init__(**values)
@model_validator(mode="after")
def _validate_model_preset(self) -> "Config":
if "default" in self.model_presets:
@@ -548,43 +535,3 @@ class Config(BaseSettings):
return None
model_config = ConfigDict(env_prefix="NANOBOT_", env_nested_delimiter="__")
def _resolve_tool_config_refs() -> None:
"""Resolve forward references in ToolsConfig by importing tool config classes.
Must be called after all modules are loaded (breaks circular imports).
Re-exports the classes into this module's namespace so existing imports
like ``from nanobot.config.schema import ExecToolConfig`` continue to work.
"""
import sys
from nanobot.agent.tools.cli_apps import CliAppsToolConfig
from nanobot.agent.tools.filesystem import FileToolsConfig
from nanobot.agent.tools.image_generation import ImageGenerationToolConfig
from nanobot.agent.tools.self import MyToolConfig
from nanobot.agent.tools.shell import ExecToolConfig
from nanobot.agent.tools.web import WebFetchConfig, WebSearchConfig, WebToolsConfig
# Re-export into this module's namespace
mod = sys.modules[__name__]
mod.ExecToolConfig = ExecToolConfig # type: ignore[attr-defined]
mod.FileToolsConfig = FileToolsConfig # type: ignore[attr-defined]
mod.CliAppsToolConfig = CliAppsToolConfig # type: ignore[attr-defined]
mod.WebToolsConfig = WebToolsConfig # type: ignore[attr-defined]
mod.WebSearchConfig = WebSearchConfig # type: ignore[attr-defined]
mod.WebFetchConfig = WebFetchConfig # type: ignore[attr-defined]
mod.MyToolConfig = MyToolConfig # type: ignore[attr-defined]
mod.ImageGenerationToolConfig = ImageGenerationToolConfig # type: ignore[attr-defined]
ToolsConfig.model_rebuild()
Config.model_rebuild()
# Eagerly resolve when the import chain allows it (no circular deps at this
# point). If it fails (first import triggers a cycle), the rebuild will
# happen lazily when Config/ToolsConfig is first used at runtime.
try:
_resolve_tool_config_refs()
except ImportError:
pass
@@ -7,10 +7,11 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools.image_generation import ImageGenerationToolConfig
from nanobot.bus.events import InboundMessage
from nanobot.bus.queue import MessageBus
from nanobot.config.loader import set_config_path
from nanobot.config.schema import ImageGenerationToolConfig, ProviderConfig, ToolsConfig
from nanobot.config.schema import ProviderConfig, ToolsConfig
from nanobot.providers.base import LLMResponse, ToolCallRequest
from nanobot.providers.image_generation import GeneratedImageResponse
+7 -3
View File
@@ -7,10 +7,16 @@ import pytest
from nanobot.agent.tools.cli_apps import CliAppsTool
from nanobot.agent.tools.filesystem import ReadFileTool
from nanobot.agent.tools.image_generation import ImageGenerationError, ImageGenerationTool
from nanobot.agent.tools.image_generation import (
ImageGenerationError,
ImageGenerationTool,
ImageGenerationToolConfig,
)
from nanobot.agent.tools.message import MessageTool
from nanobot.agent.tools.shell import ExecTool
from nanobot.agent.tools.spawn import SpawnTool
from nanobot.apps.cli.service import CliAppManager, CliAppsRuntimeConfig
from nanobot.config.schema import ProviderConfig
from nanobot.security.workspace_access import (
WORKSPACE_SCOPE_METADATA_KEY,
WorkspaceScopeError,
@@ -20,8 +26,6 @@ from nanobot.security.workspace_access import (
validate_workspace_scope_payload,
workspace_scope_from_metadata,
)
from nanobot.apps.cli.service import CliAppManager, CliAppsRuntimeConfig
from nanobot.config.schema import ImageGenerationToolConfig, ProviderConfig
PNG_BYTES = (
b"\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01"
@@ -20,6 +20,32 @@ print("nanobot.config.schema" in sys.modules)
assert result.stdout.strip() == "False"
def test_config_schema_import_does_not_load_builtin_tool_modules():
code = """
import sys
import nanobot.config.schema
print(any(
name in sys.modules
for name in (
"nanobot.agent.tools.cli_apps",
"nanobot.agent.tools.filesystem",
"nanobot.agent.tools.image_generation",
"nanobot.agent.tools.self",
"nanobot.agent.tools.shell",
"nanobot.agent.tools.web",
)
))
"""
result = subprocess.run(
[sys.executable, "-c", code],
check=True,
capture_output=True,
text=True,
)
assert result.stdout.strip() == "False"
def test_builtin_tool_configs_do_not_depend_on_config_schema_base():
repo = Path(__file__).resolve().parents[2]
tool_paths = sorted((repo / "nanobot/agent/tools").glob("*.py"))
+2 -2
View File
@@ -6,9 +6,9 @@ from typing import Any
import pytest
from nanobot.agent.tools.image_generation import ImageGenerationTool
from nanobot.agent.tools.image_generation import ImageGenerationTool, ImageGenerationToolConfig
from nanobot.config.loader import set_config_path
from nanobot.config.schema import ImageGenerationToolConfig, ProviderConfig
from nanobot.config.schema import ProviderConfig
from nanobot.providers.image_generation import GeneratedImageResponse
PNG_BYTES = (
+1 -2
View File
@@ -13,9 +13,8 @@ import pytest
from nanobot.agent.loop import AgentLoop
from nanobot.agent.subagent import SubagentManager, SubagentStatus
from nanobot.agent.tools.search import FindFilesTool, GrepTool
from nanobot.agent.tools.web import WebSearchTool
from nanobot.agent.tools.web import WebSearchConfig, WebSearchTool
from nanobot.bus.queue import MessageBus
from nanobot.config.schema import WebSearchConfig
@pytest.mark.asyncio
+1 -2
View File
@@ -10,8 +10,7 @@ import httpx
import pytest
from nanobot.agent.tools import web as web_module
from nanobot.agent.tools.web import WebFetchTool
from nanobot.config.schema import WebFetchConfig
from nanobot.agent.tools.web import WebFetchConfig, WebFetchTool
from nanobot.security.workspace_access import (
bind_workspace_scope,
build_workspace_scope,
+1 -2
View File
@@ -3,8 +3,7 @@
import httpx
import pytest
from nanobot.agent.tools.web import WebSearchTool
from nanobot.config.schema import WebSearchConfig
from nanobot.agent.tools.web import WebSearchConfig, WebSearchTool
def _tool(