mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-08 13:28:43 +03:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b8c39da238 | ||
|
|
cc7e0d7244 | ||
|
|
e015f469ad | ||
|
|
8a51059bdf | ||
|
|
15964fb466 |
+35
-171
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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,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"))
|
||||
|
||||
@@ -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 = (
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user