refactor: enforce BasedPyright strict type checking (#5158)
This commit is contained in:
+27
-1
@@ -6,6 +6,32 @@ import tomllib
|
||||
from importlib.metadata import PackageNotFoundError
|
||||
from importlib.metadata import version as _pkg_version
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .agent.tools.context import RequestContext
|
||||
from .bus.runtime_events import SessionTurnPersisted
|
||||
from .nanobot import (
|
||||
STREAM_EVENT_REASONING_COMPLETED,
|
||||
STREAM_EVENT_REASONING_DELTA,
|
||||
STREAM_EVENT_RUN_COMPLETED,
|
||||
STREAM_EVENT_RUN_FAILED,
|
||||
STREAM_EVENT_RUN_STARTED,
|
||||
STREAM_EVENT_TEXT_COMPLETED,
|
||||
STREAM_EVENT_TEXT_DELTA,
|
||||
STREAM_EVENT_TOOL_COMPLETED,
|
||||
STREAM_EVENT_TOOL_FAILED,
|
||||
STREAM_EVENT_TOOL_STARTED,
|
||||
STREAM_EVENT_TYPES,
|
||||
Nanobot,
|
||||
RunResult,
|
||||
RunStream,
|
||||
SessionInfo,
|
||||
SessionSnapshot,
|
||||
StreamEvent,
|
||||
StreamEventType,
|
||||
)
|
||||
from .runtime_context import RuntimeContextBlock, RuntimeContextProvider
|
||||
|
||||
|
||||
def _read_pyproject_version() -> str | None:
|
||||
@@ -54,7 +80,7 @@ _LAZY_EXPORTS = {
|
||||
}
|
||||
|
||||
|
||||
def __getattr__(name: str):
|
||||
def __getattr__(name: str) -> Any:
|
||||
module_path = _LAZY_EXPORTS.get(name)
|
||||
if module_path is None:
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
from collections.abc import Collection
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Callable, Coroutine
|
||||
from typing import TYPE_CHECKING, Any, Callable, Coroutine, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
@@ -65,7 +65,7 @@ class AutoCompact:
|
||||
|
||||
def check_expired(
|
||||
self,
|
||||
schedule_background: Callable[[Coroutine], None],
|
||||
schedule_background: Callable[[Coroutine[Any, Any, None]], None],
|
||||
resolve_runtime: Callable[[Session], LLMRuntime],
|
||||
active_session_keys: Collection[str] = (),
|
||||
) -> None:
|
||||
@@ -103,8 +103,8 @@ class AutoCompact:
|
||||
meta = session.metadata.get("_last_summary")
|
||||
if isinstance(meta, dict):
|
||||
self._summaries[key] = (
|
||||
meta["text"],
|
||||
datetime.fromisoformat(meta["last_active"]),
|
||||
cast(str, meta["text"]),
|
||||
datetime.fromisoformat(cast(str, meta["last_active"])),
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Auto-compact: failed for {}", key)
|
||||
@@ -126,5 +126,8 @@ class AutoCompact:
|
||||
# Cold path: summary persisted in session metadata (process restarted).
|
||||
meta = session.metadata.get("_last_summary")
|
||||
if isinstance(meta, dict):
|
||||
return session, self._format_summary(meta["text"], datetime.fromisoformat(meta["last_active"]))
|
||||
return session, self._format_summary(
|
||||
cast(str, meta["text"]),
|
||||
datetime.fromisoformat(cast(str, meta["last_active"])),
|
||||
)
|
||||
return session, None
|
||||
|
||||
@@ -4,7 +4,7 @@ import base64
|
||||
import mimetypes
|
||||
import platform
|
||||
from pathlib import Path
|
||||
from typing import Any, Mapping, Sequence
|
||||
from typing import Any, Mapping, Sequence, cast
|
||||
|
||||
from nanobot.agent.memory import MemoryStore
|
||||
from nanobot.agent.skills import SkillsLoader
|
||||
@@ -148,7 +148,12 @@ class ContextBuilder:
|
||||
|
||||
def _to_blocks(value: Any) -> list[dict[str, Any]]:
|
||||
if isinstance(value, list):
|
||||
return [item if isinstance(item, dict) else {"type": "text", "text": str(item)} for item in value]
|
||||
return [
|
||||
cast(dict[str, Any], item)
|
||||
if isinstance(item, dict)
|
||||
else {"type": "text", "text": str(item)}
|
||||
for item in cast(list[Any], value)
|
||||
]
|
||||
if value is None:
|
||||
return []
|
||||
return [{"type": "text", "text": str(value)}]
|
||||
@@ -157,7 +162,7 @@ class ContextBuilder:
|
||||
|
||||
def _load_bootstrap_files(self, workspace: Path | None = None) -> str:
|
||||
"""Load project instructions plus the agent's global profile files."""
|
||||
parts = []
|
||||
parts: list[str] = []
|
||||
project_root = workspace or self.workspace
|
||||
sources = [
|
||||
("AGENTS.md", project_root),
|
||||
@@ -212,7 +217,7 @@ class ContextBuilder:
|
||||
user_content = self.build_user_content(current_message, image_paths=media)
|
||||
blocks = list(runtime_context_blocks or ()) if current_role == "user" else []
|
||||
merged, runtime_context_meta = append_runtime_context(user_content, blocks)
|
||||
messages = [
|
||||
messages: list[dict[str, Any]] = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": self.build_system_prompt(
|
||||
@@ -235,7 +240,7 @@ class ContextBuilder:
|
||||
last["_meta"] = internal_meta
|
||||
messages[-1] = last
|
||||
return messages
|
||||
current = {"role": current_role, "content": merged}
|
||||
current: dict[str, Any] = {"role": current_role, "content": merged}
|
||||
if current_role == "user" and runtime_context_meta is not None:
|
||||
current["_meta"] = {RUNTIME_CONTEXT_MESSAGE_META: runtime_context_meta}
|
||||
messages.append(current)
|
||||
@@ -250,7 +255,7 @@ class ContextBuilder:
|
||||
if not image_paths:
|
||||
return text
|
||||
|
||||
image_blocks = []
|
||||
image_blocks: list[dict[str, Any]] = []
|
||||
for path in image_paths:
|
||||
p = Path(path)
|
||||
if not p.is_file():
|
||||
|
||||
@@ -9,7 +9,7 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
@@ -23,6 +23,7 @@ from nanobot.utils.helpers import (
|
||||
from nanobot.utils.runtime import ensure_nonempty_tool_result
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
from nanobot.providers.base import LLMProvider
|
||||
|
||||
SNIP_SAFETY_BUFFER = 1024
|
||||
@@ -49,8 +50,9 @@ def _tool_call_name_is_valid(tool_call: Any) -> bool:
|
||||
"""
|
||||
if not isinstance(tool_call, dict):
|
||||
return False
|
||||
fn = tool_call.get("function")
|
||||
name = fn.get("name") if isinstance(fn, dict) else tool_call.get("name")
|
||||
tool_call_data = cast(dict[str, Any], tool_call)
|
||||
fn = tool_call_data.get("function")
|
||||
name = cast(dict[str, Any], fn).get("name") if isinstance(fn, dict) else tool_call_data.get("name")
|
||||
return isinstance(name, str) and bool(name)
|
||||
|
||||
|
||||
@@ -58,7 +60,7 @@ def _tool_call_name_is_valid(tool_call: Any) -> bool:
|
||||
class ContextGovernanceConfig:
|
||||
provider: LLMProvider
|
||||
model: str
|
||||
tools: Any
|
||||
tools: ToolRegistry
|
||||
workspace: Path | None
|
||||
session_key: str | None
|
||||
max_tool_result_chars: int
|
||||
@@ -199,7 +201,7 @@ class ContextGovernor:
|
||||
if updated is not None:
|
||||
updated.append(msg)
|
||||
continue
|
||||
kept = [tc for tc in calls if _tool_call_name_is_valid(tc)]
|
||||
kept = [tc for tc in cast(list[Any], calls) if _tool_call_name_is_valid(tc)]
|
||||
if len(kept) == len(calls):
|
||||
if updated is not None:
|
||||
updated.append(msg)
|
||||
@@ -238,9 +240,11 @@ class ContextGovernor:
|
||||
for idx, msg in enumerate(messages):
|
||||
role = msg.get("role")
|
||||
if role == "assistant":
|
||||
for tc in msg.get("tool_calls") or []:
|
||||
if isinstance(tc, dict) and tc.get("id"):
|
||||
declared.add(str(tc["id"]))
|
||||
for tc in cast(list[Any], msg.get("tool_calls") or []):
|
||||
if isinstance(tc, dict):
|
||||
tool_call = cast(dict[str, Any], tc)
|
||||
if tool_call.get("id"):
|
||||
declared.add(str(tool_call["id"]))
|
||||
if role == "tool":
|
||||
tid = msg.get("tool_call_id")
|
||||
tid_str = str(tid) if tid else ""
|
||||
@@ -266,13 +270,17 @@ class ContextGovernor:
|
||||
for idx, msg in enumerate(messages):
|
||||
role = msg.get("role")
|
||||
if role == "assistant":
|
||||
for tc in msg.get("tool_calls") or []:
|
||||
if isinstance(tc, dict) and tc.get("id"):
|
||||
for tc in cast(list[Any], msg.get("tool_calls") or []):
|
||||
if isinstance(tc, dict):
|
||||
name = ""
|
||||
func = tc.get("function")
|
||||
if isinstance(func, dict):
|
||||
name = func.get("name", "")
|
||||
declared.append((idx, str(tc["id"]), name))
|
||||
tool_call = cast(dict[str, Any], tc)
|
||||
if tool_call.get("id"):
|
||||
func = tool_call.get("function")
|
||||
if isinstance(func, dict):
|
||||
func_data = cast(dict[str, Any], func)
|
||||
raw_name = func_data.get("name", "")
|
||||
name = raw_name if isinstance(raw_name, str) else str(raw_name)
|
||||
declared.append((idx, str(tool_call["id"]), name))
|
||||
elif role == "tool":
|
||||
tid = msg.get("tool_call_id")
|
||||
if tid:
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
from nanobot.agent.hook import (
|
||||
AgentHook,
|
||||
@@ -56,17 +56,21 @@ class FileEditActivityHook(AgentHook):
|
||||
) -> None:
|
||||
if self._on_progress is None or not isinstance(params, dict):
|
||||
return
|
||||
typed_params = cast(dict[str, Any], params)
|
||||
trackers = prepare_file_edit_trackers(
|
||||
call_id=tool_call.id,
|
||||
tool_name=tool_call.name,
|
||||
tool=tool,
|
||||
workspace=self._workspace,
|
||||
params=params,
|
||||
params=typed_params,
|
||||
)
|
||||
if not trackers:
|
||||
return
|
||||
self._trackers_by_call[self._tool_call_key(tool_call)] = trackers
|
||||
await self._emit([build_file_edit_start_event(tracker, params) for tracker in trackers])
|
||||
await self._emit([
|
||||
build_file_edit_start_event(tracker, typed_params)
|
||||
for tracker in trackers
|
||||
])
|
||||
|
||||
async def after_execute_tool(
|
||||
self,
|
||||
|
||||
+165
-81
@@ -1,5 +1,7 @@
|
||||
"""Agent loop: the core processing engine."""
|
||||
|
||||
# pyright: reportPrivateUsage=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
@@ -7,13 +9,13 @@ import dataclasses
|
||||
import inspect
|
||||
import os
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Coroutine, Iterable, Mapping
|
||||
from contextlib import AbstractContextManager, ExitStack, nullcontext, suppress
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum, auto
|
||||
from functools import partial
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable, TypeVar
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable, TypeVar, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
@@ -94,10 +96,13 @@ if TYPE_CHECKING:
|
||||
from nanobot.agent.tools.mcp import MCPConnection
|
||||
from nanobot.config.schema import (
|
||||
ChannelsConfig,
|
||||
Config,
|
||||
MCPServerConfig,
|
||||
ProviderConfig,
|
||||
ToolsConfig,
|
||||
)
|
||||
from nanobot.cron.service import CronService
|
||||
from nanobot.triggers.local_store import LocalTriggerStore
|
||||
|
||||
_T = TypeVar("_T")
|
||||
|
||||
@@ -142,7 +147,7 @@ class TurnContext:
|
||||
on_runtime_admitted: Callable[[LLMRuntime], Awaitable[None]] | None = None
|
||||
on_retry_wait: Callable[[str], Awaitable[None]] | None = None
|
||||
|
||||
pending_queue: asyncio.Queue | None = None
|
||||
pending_queue: asyncio.Queue[InboundMessage] | None = None
|
||||
pending_summary: str | None = None
|
||||
|
||||
ephemeral: bool = False
|
||||
@@ -156,6 +161,18 @@ class TurnContext:
|
||||
visible_run_started_at: float | None = None
|
||||
turn_latency_ms: int | None = None
|
||||
|
||||
def require_runtime(self) -> LLMRuntime:
|
||||
"""Return the runtime established by the BUILD stage."""
|
||||
if self.runtime is None:
|
||||
raise RuntimeError("turn runtime is not initialized; BUILD must run before this stage")
|
||||
return self.runtime
|
||||
|
||||
def require_session(self) -> Session:
|
||||
"""Return the session established by the RESTORE stage."""
|
||||
if self.session is None:
|
||||
raise RuntimeError("turn session is not initialized; RESTORE must run before this stage")
|
||||
return self.session
|
||||
|
||||
|
||||
class AgentLoop:
|
||||
"""
|
||||
@@ -243,7 +260,7 @@ class AgentLoop:
|
||||
cron_service: CronService | None = None,
|
||||
restrict_to_workspace: bool = False,
|
||||
session_manager: SessionManager | None = None,
|
||||
mcp_servers: dict | None = None,
|
||||
mcp_servers: dict[str, MCPServerConfig] | None = None,
|
||||
channels_config: ChannelsConfig | None = None,
|
||||
timezone: str | None = None,
|
||||
session_ttl_minutes: int = 0,
|
||||
@@ -266,7 +283,7 @@ class AgentLoop:
|
||||
turn_delivery_factory: TurnDeliveryFactory | None = None,
|
||||
runtime_model_publisher: Callable[[str, str | None], None] | None = None,
|
||||
restart_mode: str = "auto",
|
||||
local_trigger_store: Any | None = None,
|
||||
local_trigger_store: LocalTriggerStore | None = None,
|
||||
idle_compact_check_interval_seconds: int = 0,
|
||||
):
|
||||
from nanobot.config.schema import ToolsConfig
|
||||
@@ -381,7 +398,7 @@ class AgentLoop:
|
||||
# Per-session pending queues for mid-turn message injection.
|
||||
# When a session has an active task, new messages for that session
|
||||
# are routed here instead of creating a new task.
|
||||
self._pending_queues: dict[str, asyncio.Queue] = {}
|
||||
self._pending_queues: dict[str, asyncio.Queue[InboundMessage]] = {}
|
||||
self._deferred_automation_turns: dict[str, list[InboundMessage]] = {}
|
||||
self._cron_turns = CronTurnCoordinator(
|
||||
publish_inbound=self.bus.publish_inbound,
|
||||
@@ -430,7 +447,7 @@ class AgentLoop:
|
||||
@classmethod
|
||||
def from_config(
|
||||
cls,
|
||||
config: Any,
|
||||
config: Config,
|
||||
bus: MessageBus | None = None,
|
||||
**extra: Any,
|
||||
) -> AgentLoop:
|
||||
@@ -657,12 +674,17 @@ class AgentLoop:
|
||||
"""
|
||||
if not turn_continuation.should_persist_user_message(msg.metadata):
|
||||
return False
|
||||
media_paths = [p for p in (msg.media or []) if isinstance(p, str) and p]
|
||||
has_text = isinstance(msg.content, str) and msg.content.strip()
|
||||
media_paths = [
|
||||
path
|
||||
for path in (msg.media or [])
|
||||
if isinstance(cast(object, path), str) and path
|
||||
]
|
||||
content_value = cast(object, msg.content)
|
||||
has_text = isinstance(content_value, str) and content_value.strip()
|
||||
if has_text or media_paths or runtime_context_blocks:
|
||||
extra: dict[str, Any] = ({"media": list(media_paths)} if media_paths else {}) | agent_context.session_extra(msg.metadata)
|
||||
extra.update(kwargs)
|
||||
text = msg.content if isinstance(msg.content, str) else ""
|
||||
text = content_value if isinstance(content_value, str) else ""
|
||||
text_override, automation_extra = automation_history_overrides(msg.metadata)
|
||||
if text_override is not None:
|
||||
text = text_override
|
||||
@@ -810,7 +832,7 @@ class AgentLoop:
|
||||
|
||||
async def _run_agent_loop(
|
||||
self,
|
||||
initial_messages: list[dict],
|
||||
initial_messages: list[dict[str, Any]],
|
||||
on_progress: Callable[..., Awaitable[None]] | None = None,
|
||||
on_stream: Callable[[str], Awaitable[None]] | None = None,
|
||||
on_stream_end: Callable[..., Awaitable[None]] | None = None,
|
||||
@@ -824,7 +846,7 @@ class AgentLoop:
|
||||
metadata: dict[str, Any] | None = None,
|
||||
session_key: str | None = None,
|
||||
original_user_text: str | None = None,
|
||||
pending_queue: asyncio.Queue | None = None,
|
||||
pending_queue: asyncio.Queue[InboundMessage] | None = None,
|
||||
ephemeral: bool = False,
|
||||
run_extra_hooks_for_ephemeral: bool = False,
|
||||
hooks: list[AgentHook] | None = None,
|
||||
@@ -832,7 +854,7 @@ class AgentLoop:
|
||||
turn_scopes: list[AbstractContextManager[Any]] | None = None,
|
||||
tools: ToolRegistry | None = None,
|
||||
request_context: RequestContext | None = None,
|
||||
) -> tuple[str | None, list[str], list[dict], str, bool]:
|
||||
) -> tuple[str | None, list[str], list[dict[str, Any]], str, bool]:
|
||||
"""Run the agent iteration loop.
|
||||
|
||||
*on_stream*: called with each content delta during streaming.
|
||||
@@ -875,7 +897,12 @@ class AgentLoop:
|
||||
image_paths=image_paths,
|
||||
)
|
||||
row: dict[str, Any] = {"role": "user", "content": user_content}
|
||||
metadata = pending_msg.metadata if isinstance(pending_msg.metadata, dict) else {}
|
||||
metadata_value = cast(object, pending_msg.metadata)
|
||||
metadata = (
|
||||
pending_msg.metadata
|
||||
if isinstance(metadata_value, dict)
|
||||
else {}
|
||||
)
|
||||
if pending_msg.channel != "system":
|
||||
scope = self.workspace_scopes.for_turn(
|
||||
channel=pending_msg.channel,
|
||||
@@ -899,19 +926,24 @@ class AgentLoop:
|
||||
pending_request,
|
||||
effective_tools,
|
||||
)
|
||||
row["content"], marker = append_runtime_context(user_content, blocks)
|
||||
if marker is not None:
|
||||
row["_meta"] = {RUNTIME_CONTEXT_MESSAGE_META: marker}
|
||||
row["content"], runtime_marker = append_runtime_context(
|
||||
user_content,
|
||||
blocks,
|
||||
)
|
||||
if runtime_marker is not None:
|
||||
row["_meta"] = {
|
||||
RUNTIME_CONTEXT_MESSAGE_META: runtime_marker,
|
||||
}
|
||||
if (
|
||||
pending_msg.sender_id == "subagent"
|
||||
and metadata.get("injected_event") == "subagent_result"
|
||||
):
|
||||
marker: dict[str, Any] = {"kind": "subagent_result"}
|
||||
subagent_marker: dict[str, Any] = {"kind": "subagent_result"}
|
||||
task_id = metadata.get("subagent_task_id")
|
||||
if isinstance(task_id, str) and task_id:
|
||||
marker["subagent_task_id"] = task_id
|
||||
subagent_marker["subagent_task_id"] = task_id
|
||||
row["subagent_task_id"] = task_id
|
||||
row[HIDDEN_HISTORY_META] = marker
|
||||
row[HIDDEN_HISTORY_META] = subagent_marker
|
||||
row["injected_event"] = "subagent_result"
|
||||
return row
|
||||
|
||||
@@ -1178,7 +1210,7 @@ class AgentLoop:
|
||||
gate = self._concurrency_gate or nullcontext()
|
||||
|
||||
delivery = self.turn_delivery_factory.unrouted(msg, session_key)
|
||||
pending: asyncio.Queue | None = None
|
||||
pending: asyncio.Queue[InboundMessage] | None = None
|
||||
try:
|
||||
async with lock, gate:
|
||||
# Only the task that owns the session lock may publish the
|
||||
@@ -1304,7 +1336,7 @@ class AgentLoop:
|
||||
if errors:
|
||||
raise BaseExceptionGroup("failed to close agent resources", errors)
|
||||
|
||||
def _schedule_background(self, coro) -> None:
|
||||
def _schedule_background(self, coro: Coroutine[Any, Any, Any]) -> None:
|
||||
"""Schedule a coroutine as a tracked background task (drained on shutdown)."""
|
||||
task = asyncio.create_task(coro)
|
||||
self._background_tasks.add(task)
|
||||
@@ -1322,7 +1354,7 @@ class AgentLoop:
|
||||
on_progress: Callable[..., Awaitable[None]] | None = None,
|
||||
on_stream: Callable[[str], Awaitable[None]] | None = None,
|
||||
on_stream_end: Callable[..., Awaitable[None]] | None = None,
|
||||
pending_queue: asyncio.Queue | None = None,
|
||||
pending_queue: asyncio.Queue[InboundMessage] | None = None,
|
||||
ephemeral: bool = False,
|
||||
run_extra_hooks_for_ephemeral: bool = False,
|
||||
hooks: list[AgentHook] | None = None,
|
||||
@@ -1518,27 +1550,33 @@ class AgentLoop:
|
||||
# ensure it exists in case this handler is invoked independently.
|
||||
if ctx.session is None:
|
||||
ctx.session = self.sessions.get_or_create(ctx.session_key)
|
||||
session = ctx.session
|
||||
self._remember_unified_session_route(
|
||||
ctx.session,
|
||||
session,
|
||||
msg,
|
||||
is_user_turn=ctx.original_user_text is not None,
|
||||
)
|
||||
await ctx.delivery.started()
|
||||
if ctx.kind is TurnKind.USER:
|
||||
self.workspace_scopes.persist_message_scope(ctx.session, msg)
|
||||
self.workspace_scopes.persist_message_scope(session, msg)
|
||||
|
||||
if self._restore_runtime_checkpoint(ctx.session):
|
||||
self.sessions.save(ctx.session)
|
||||
if self._restore_pending_user_turn(ctx.session):
|
||||
self.sessions.save(ctx.session)
|
||||
if self._restore_runtime_checkpoint(session):
|
||||
self.sessions.save(session)
|
||||
if self._restore_pending_user_turn(session):
|
||||
self.sessions.save(session)
|
||||
|
||||
async def _compact_session(self, ctx: TurnContext) -> None:
|
||||
ctx.session, pending = self.auto_compact.prepare_session(ctx.session, ctx.session_key)
|
||||
session = ctx.require_session()
|
||||
ctx.session, pending = self.auto_compact.prepare_session(
|
||||
session,
|
||||
ctx.session_key,
|
||||
)
|
||||
ctx.pending_summary = pending
|
||||
|
||||
async def _dispatch_command(self, ctx: TurnContext) -> bool:
|
||||
if ctx.kind is TurnKind.SYSTEM:
|
||||
return False
|
||||
session = ctx.require_session()
|
||||
raw = ctx.msg.content.strip()
|
||||
_, automation_metadata = automation_history_overrides(ctx.msg.metadata)
|
||||
is_user_turn = (
|
||||
@@ -1549,7 +1587,7 @@ class AgentLoop:
|
||||
)
|
||||
cmd_ctx = CommandContext(
|
||||
msg=ctx.msg,
|
||||
session=ctx.session,
|
||||
session=session,
|
||||
key=ctx.session_key,
|
||||
raw=raw,
|
||||
loop=self,
|
||||
@@ -1567,13 +1605,13 @@ class AgentLoop:
|
||||
# intentionally clears the session.
|
||||
if cmd_ctx.raw.lower() != "/new":
|
||||
ctx.input_persisted_early = self._persist_user_message_early(
|
||||
ctx.msg, ctx.session, _command=True
|
||||
ctx.msg, session, _command=True
|
||||
)
|
||||
ctx.session.add_message(
|
||||
session.add_message(
|
||||
"assistant", result.content, _command=True
|
||||
)
|
||||
self._clear_pending_user_turn(ctx.session)
|
||||
self.sessions.save(ctx.session)
|
||||
self._clear_pending_user_turn(session)
|
||||
self.sessions.save(session)
|
||||
if not ctx.ephemeral:
|
||||
await self.runtime_event_publisher.session_turn_persisted(
|
||||
ctx.msg,
|
||||
@@ -1585,9 +1623,10 @@ class AgentLoop:
|
||||
return False
|
||||
|
||||
async def _build_turn(self, ctx: TurnContext) -> None:
|
||||
session = ctx.require_session()
|
||||
runtime = ctx.runtime
|
||||
if runtime is None:
|
||||
runtime = self.runtime_for_session(ctx.session)
|
||||
runtime = self.runtime_for_session(session)
|
||||
ctx.runtime = runtime
|
||||
if ctx.session_key.startswith("dream:"):
|
||||
logger.info(
|
||||
@@ -1602,7 +1641,7 @@ class AgentLoop:
|
||||
)
|
||||
if not ctx.ephemeral:
|
||||
await self.consolidator.maybe_consolidate_by_tokens(
|
||||
ctx.session,
|
||||
session,
|
||||
runtime=runtime,
|
||||
replay_max_messages=replay_max_messages,
|
||||
)
|
||||
@@ -1617,18 +1656,18 @@ class AgentLoop:
|
||||
"max_tokens": self._replay_token_budget(runtime),
|
||||
"extend_to_user": is_subagent,
|
||||
}
|
||||
ctx.history = ctx.session.get_history(**_hist_kwargs)
|
||||
ctx.history = session.get_history(**_hist_kwargs)
|
||||
if is_subagent:
|
||||
# Keep the durable internal delivery as an assistant record, but
|
||||
# present this completion to the model as fresh follow-up input.
|
||||
# Providers without assistant-prefill support drop trailing
|
||||
# assistant messages, so using the persisted record as the current
|
||||
# prompt would hide an independently dispatched subagent result.
|
||||
if self._persist_subagent_followup(ctx.session, ctx.msg):
|
||||
if self._persist_subagent_followup(session, ctx.msg):
|
||||
logger.debug("Subagent result persisted for session {}", ctx.session_key)
|
||||
self.sessions.save(ctx.session)
|
||||
self.sessions.save(session)
|
||||
ctx.input_persisted_early = True
|
||||
ctx.delivery.record_runtime(ctx.runtime)
|
||||
ctx.delivery.record_runtime(runtime)
|
||||
|
||||
ctx.request_context = self._request_context_for_turn(ctx)
|
||||
if ctx.kind is TurnKind.USER:
|
||||
@@ -1637,7 +1676,7 @@ class AgentLoop:
|
||||
if ctx.kind is TurnKind.USER:
|
||||
ctx.input_persisted_early = self._persist_user_message_early(
|
||||
ctx.msg,
|
||||
ctx.session,
|
||||
session,
|
||||
runtime_context_blocks=ctx.runtime_context_blocks,
|
||||
)
|
||||
|
||||
@@ -1647,12 +1686,13 @@ class AgentLoop:
|
||||
ctx.on_retry_wait = ctx.delivery.retry_wait_callback()
|
||||
|
||||
async def _run_turn(self, ctx: TurnContext) -> None:
|
||||
runtime = ctx.require_runtime()
|
||||
if ctx.visible_run_started_at is None:
|
||||
ctx.visible_run_started_at = time.time()
|
||||
await ctx.delivery.running(started_at=ctx.visible_run_started_at)
|
||||
result = await self._run_agent_loop(
|
||||
ctx.initial_messages,
|
||||
runtime=ctx.runtime,
|
||||
runtime=runtime,
|
||||
on_progress=ctx.on_progress,
|
||||
on_stream=ctx.on_stream,
|
||||
on_stream_end=ctx.on_stream_end,
|
||||
@@ -1682,6 +1722,8 @@ class AgentLoop:
|
||||
await turn_continuation.maybe_continue_turn(ctx)
|
||||
|
||||
async def _persist_turn(self, ctx: TurnContext) -> None:
|
||||
runtime = ctx.require_runtime()
|
||||
session = ctx.require_session()
|
||||
turn_continuation.prepare_save_boundary(ctx)
|
||||
|
||||
if (
|
||||
@@ -1702,26 +1744,26 @@ class AgentLoop:
|
||||
)
|
||||
ctx.turn_latency_ms = max(0, int((time.time() - latency_started_at) * 1000))
|
||||
self._save_turn(
|
||||
ctx.session, ctx.all_messages, ctx.save_skip,
|
||||
session, ctx.all_messages, ctx.save_skip,
|
||||
turn_latency_ms=ctx.turn_latency_ms,
|
||||
)
|
||||
ctx.delivery.record_latency(ctx.turn_latency_ms)
|
||||
if not ctx.ephemeral:
|
||||
ctx.session.enforce_file_cap(
|
||||
session.enforce_file_cap(
|
||||
on_archive=partial(self.context.memory.raw_archive, session_key=ctx.session_key)
|
||||
)
|
||||
self._schedule_background(
|
||||
self.consolidator.maybe_consolidate_by_tokens(
|
||||
ctx.session,
|
||||
runtime=ctx.runtime,
|
||||
session,
|
||||
runtime=runtime,
|
||||
replay_max_messages=replay_max_messages_for_context(
|
||||
ctx.runtime.context_window_tokens
|
||||
runtime.context_window_tokens
|
||||
),
|
||||
)
|
||||
)
|
||||
self._clear_pending_user_turn(ctx.session)
|
||||
self._clear_runtime_checkpoint(ctx.session)
|
||||
self.sessions.save(ctx.session)
|
||||
self._clear_pending_user_turn(session)
|
||||
self._clear_runtime_checkpoint(session)
|
||||
self.sessions.save(session)
|
||||
if not ctx.ephemeral:
|
||||
await self.runtime_event_publisher.session_turn_persisted(
|
||||
ctx.msg,
|
||||
@@ -1744,7 +1786,7 @@ class AgentLoop:
|
||||
return
|
||||
ctx.outbound = self._assemble_outbound(
|
||||
ctx.msg,
|
||||
ctx.final_content,
|
||||
cast(str, ctx.final_content),
|
||||
ctx.stop_reason,
|
||||
ctx.had_injections,
|
||||
ctx.streamed_content,
|
||||
@@ -1755,39 +1797,47 @@ class AgentLoop:
|
||||
|
||||
def _sanitize_persisted_blocks(
|
||||
self,
|
||||
content: list[dict[str, Any]],
|
||||
content: list[object],
|
||||
*,
|
||||
should_truncate_text: bool = False,
|
||||
) -> list[dict[str, Any]]:
|
||||
) -> list[object]:
|
||||
"""Strip volatile multimodal payloads before writing session history."""
|
||||
filtered: list[dict[str, Any]] = []
|
||||
filtered: list[object] = []
|
||||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
filtered.append(block)
|
||||
continue
|
||||
|
||||
if block.get("type") == "image_url" and block.get("image_url", {}).get(
|
||||
"url", ""
|
||||
block_data = cast(dict[str, Any], block)
|
||||
image_url = cast(dict[str, Any], block_data.get("image_url", {}))
|
||||
if block_data.get("type") == "image_url" and str(
|
||||
image_url.get("url", "")
|
||||
).startswith("data:image/"):
|
||||
path = (block.get("_meta") or {}).get("path", "")
|
||||
filtered.append({"type": "text", "text": image_placeholder_text(path)})
|
||||
internal_meta = cast(dict[str, Any], block_data.get("_meta") or {})
|
||||
path = cast(str, internal_meta.get("path", ""))
|
||||
filtered.append(
|
||||
{"type": "text", "text": image_placeholder_text(path)}
|
||||
)
|
||||
continue
|
||||
|
||||
if block.get("type") == "text" and isinstance(block.get("text"), str):
|
||||
text = block["text"]
|
||||
if block_data.get("type") == "text" and isinstance(
|
||||
block_data.get("text"),
|
||||
str,
|
||||
):
|
||||
text = cast(str, block_data["text"])
|
||||
if should_truncate_text and len(text) > self.max_tool_result_chars:
|
||||
text = truncate_text_fn(text, self.max_tool_result_chars)
|
||||
filtered.append({**block, "text": text})
|
||||
filtered.append({**block_data, "text": text})
|
||||
continue
|
||||
|
||||
filtered.append(block)
|
||||
filtered.append(block_data)
|
||||
|
||||
return filtered
|
||||
|
||||
def _save_turn(
|
||||
self,
|
||||
session: Session,
|
||||
messages: list[dict],
|
||||
messages: list[dict[str, Any]],
|
||||
skip: int,
|
||||
*,
|
||||
turn_latency_ms: int | None = None,
|
||||
@@ -1799,8 +1849,10 @@ class AgentLoop:
|
||||
str(tc["id"])
|
||||
for m in session.messages
|
||||
if m.get("role") == "assistant"
|
||||
for tc in m.get("tool_calls") or []
|
||||
if isinstance(tc, dict) and tc.get("id")
|
||||
for tc_value in cast(Iterable[object], m.get("tool_calls") or [])
|
||||
if isinstance(tc_value, dict)
|
||||
for tc in (cast(dict[str, Any], tc_value),)
|
||||
if tc.get("id")
|
||||
}
|
||||
fulfilled_tool_call_ids = {
|
||||
str(m["tool_call_id"])
|
||||
@@ -1810,9 +1862,11 @@ class AgentLoop:
|
||||
last_assistant_idx: int | None = None
|
||||
for m in messages[skip:]:
|
||||
entry = dict(m)
|
||||
internal_meta = entry.pop("_meta", None)
|
||||
internal_meta = cast(object, entry.pop("_meta", None))
|
||||
runtime_context_meta = (
|
||||
internal_meta.get(RUNTIME_CONTEXT_MESSAGE_META)
|
||||
cast(dict[str, Any], internal_meta).get(
|
||||
RUNTIME_CONTEXT_MESSAGE_META
|
||||
)
|
||||
if isinstance(internal_meta, dict)
|
||||
else None
|
||||
)
|
||||
@@ -1838,7 +1892,10 @@ class AgentLoop:
|
||||
if isinstance(content, str) and len(content) > self.max_tool_result_chars:
|
||||
entry["content"] = truncate_text_fn(content, self.max_tool_result_chars)
|
||||
elif isinstance(content, list):
|
||||
filtered = self._sanitize_persisted_blocks(content, should_truncate_text=True)
|
||||
filtered = self._sanitize_persisted_blocks(
|
||||
cast(list[object], content),
|
||||
should_truncate_text=True,
|
||||
)
|
||||
if not filtered:
|
||||
# Preserve the tool_call/result pair after block filtering.
|
||||
filtered = [
|
||||
@@ -1847,7 +1904,9 @@ class AgentLoop:
|
||||
entry["content"] = filtered
|
||||
elif role == "user":
|
||||
if isinstance(content, list):
|
||||
filtered = self._sanitize_persisted_blocks(content)
|
||||
filtered = self._sanitize_persisted_blocks(
|
||||
cast(list[object], content),
|
||||
)
|
||||
if not filtered:
|
||||
continue
|
||||
entry["content"] = filtered
|
||||
@@ -1859,8 +1918,13 @@ class AgentLoop:
|
||||
last_assistant_idx = len(session.messages) - 1
|
||||
declared_tool_call_ids.update(
|
||||
str(tc["id"])
|
||||
for tc in entry.get("tool_calls") or []
|
||||
if isinstance(tc, dict) and tc.get("id")
|
||||
for tc_value in cast(
|
||||
Iterable[object],
|
||||
entry.get("tool_calls") or [],
|
||||
)
|
||||
if isinstance(tc_value, dict)
|
||||
for tc in (cast(dict[str, Any], tc_value),)
|
||||
if tc.get("id")
|
||||
)
|
||||
if turn_latency_ms is not None and last_assistant_idx is not None:
|
||||
session.messages[last_assistant_idx]["latency_ms"] = int(turn_latency_ms)
|
||||
@@ -1875,7 +1939,12 @@ class AgentLoop:
|
||||
"""
|
||||
if not msg.content:
|
||||
return False
|
||||
task_id = msg.metadata.get("subagent_task_id") if isinstance(msg.metadata, dict) else None
|
||||
metadata_value = cast(object, msg.metadata)
|
||||
task_id = (
|
||||
msg.metadata.get("subagent_task_id")
|
||||
if isinstance(metadata_value, dict)
|
||||
else None
|
||||
)
|
||||
if task_id and any(
|
||||
m.get("injected_event") == "subagent_result" and m.get("subagent_task_id") == task_id
|
||||
for m in session.messages
|
||||
@@ -1921,29 +1990,44 @@ class AgentLoop:
|
||||
"""Materialize an unfinished turn into session history before a new request."""
|
||||
from datetime import datetime
|
||||
|
||||
checkpoint = session.metadata.get(self._RUNTIME_CHECKPOINT_KEY)
|
||||
checkpoint = cast(
|
||||
object,
|
||||
session.metadata.get(self._RUNTIME_CHECKPOINT_KEY),
|
||||
)
|
||||
if not isinstance(checkpoint, dict):
|
||||
return False
|
||||
checkpoint_data = cast(dict[str, Any], checkpoint)
|
||||
|
||||
assistant_message = checkpoint.get("assistant_message")
|
||||
completed_tool_results = checkpoint.get("completed_tool_results") or []
|
||||
pending_tool_calls = checkpoint.get("pending_tool_calls") or []
|
||||
assistant_message = cast(object, checkpoint_data.get("assistant_message"))
|
||||
completed_tool_results = cast(
|
||||
Iterable[object],
|
||||
checkpoint_data.get("completed_tool_results") or [],
|
||||
)
|
||||
pending_tool_calls = cast(
|
||||
Iterable[object],
|
||||
checkpoint_data.get("pending_tool_calls") or [],
|
||||
)
|
||||
|
||||
restored_messages: list[dict[str, Any]] = []
|
||||
if isinstance(assistant_message, dict):
|
||||
restored = dict(assistant_message)
|
||||
restored = dict(cast(dict[str, Any], 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 = dict(cast(dict[str, Any], 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"
|
||||
tool_call_data = cast(dict[str, Any], tool_call)
|
||||
tool_id = tool_call_data.get("id")
|
||||
function_data = cast(
|
||||
dict[str, Any],
|
||||
tool_call_data.get("function") or {},
|
||||
)
|
||||
name = function_data.get("name") or "tool"
|
||||
restored_messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
|
||||
+40
-24
@@ -1,5 +1,10 @@
|
||||
"""Memory system: pure file I/O store and lightweight Consolidator."""
|
||||
|
||||
# Tool schemas are installed by the ``@tool_parameters`` class decorator at
|
||||
# runtime; static analyzers cannot observe that it clears ``parameters`` from
|
||||
# ``__abstractmethods__`` before these classes are instantiated.
|
||||
# pyright: reportAbstractUsage=false, reportPrivateUsage=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
@@ -11,7 +16,7 @@ import weakref
|
||||
from contextlib import suppress
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, Callable, Iterator
|
||||
from typing import TYPE_CHECKING, Any, Callable, Iterator, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
@@ -38,6 +43,7 @@ from nanobot.utils.workspace_prompts import (
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
from nanobot.utils.llm_runtime import LLMRuntime
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -58,7 +64,7 @@ class DreamRunProgress:
|
||||
**_kwargs: Any,
|
||||
) -> None:
|
||||
if any(
|
||||
isinstance(event, dict) and event.get("phase") == "error"
|
||||
isinstance(cast(object, event), dict) and event.get("phase") == "error"
|
||||
for event in tool_events or ()
|
||||
):
|
||||
self.had_tool_errors = True
|
||||
@@ -474,11 +480,11 @@ class MemoryStore:
|
||||
line = line.strip()
|
||||
if line:
|
||||
try:
|
||||
parsed = json.loads(line)
|
||||
parsed: object = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if isinstance(parsed, dict):
|
||||
entries.append(parsed)
|
||||
entries.append(cast(dict[str, Any], parsed))
|
||||
|
||||
return entries
|
||||
|
||||
@@ -496,8 +502,8 @@ class MemoryStore:
|
||||
lines = [line for line in data.split("\n") if line.strip()]
|
||||
if not lines:
|
||||
return None
|
||||
parsed = json.loads(lines[-1])
|
||||
return parsed if isinstance(parsed, dict) else None
|
||||
parsed: object = json.loads(lines[-1])
|
||||
return cast(dict[str, Any], parsed) if isinstance(parsed, dict) else None
|
||||
except (FileNotFoundError, json.JSONDecodeError, UnicodeDecodeError):
|
||||
return None
|
||||
|
||||
@@ -612,7 +618,7 @@ class MemoryStore:
|
||||
("USER.md", self.user_file),
|
||||
("memory/MEMORY.md", self.memory_file),
|
||||
]
|
||||
blocks = []
|
||||
blocks: list[str] = []
|
||||
for label, path in files:
|
||||
try:
|
||||
content = path.read_text(encoding="utf-8") if path.exists() else ""
|
||||
@@ -633,7 +639,7 @@ class MemoryStore:
|
||||
return ""
|
||||
return self._git.summarize_working_tree(list(self._DREAM_CONTENT_PATHS))
|
||||
|
||||
def build_dream_tools(self):
|
||||
def build_dream_tools(self) -> ToolRegistry:
|
||||
"""Build the restricted tool registry used by Dream runs."""
|
||||
from nanobot.agent.skills import BUILTIN_SKILLS_DIR
|
||||
from nanobot.agent.tools.apply_patch import ApplyPatchTool
|
||||
@@ -684,17 +690,15 @@ class MemoryStore:
|
||||
) -> bool:
|
||||
"""Return True only when a Dream turn completed without tool failures."""
|
||||
metadata = getattr(resp, "metadata", None)
|
||||
return (
|
||||
not had_tool_errors
|
||||
and isinstance(metadata, dict)
|
||||
and metadata.get("_stop_reason") == "completed"
|
||||
)
|
||||
if had_tool_errors or not isinstance(metadata, dict):
|
||||
return False
|
||||
return cast(dict[str, Any], metadata).get("_stop_reason") == "completed"
|
||||
|
||||
# -- message formatting utility ------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _format_messages(messages: list[dict]) -> str:
|
||||
lines = []
|
||||
def _format_messages(messages: list[dict[str, Any]]) -> str:
|
||||
lines: list[str] = []
|
||||
for message in messages:
|
||||
content = content_with_media_breadcrumbs(
|
||||
message.get("role"),
|
||||
@@ -703,16 +707,22 @@ class MemoryStore:
|
||||
)
|
||||
if not content:
|
||||
continue
|
||||
tools = f" [tools: {', '.join(message['tools_used'])}]" if message.get("tools_used") else ""
|
||||
tools_used = message.get("tools_used")
|
||||
tools = (
|
||||
f" [tools: {', '.join(cast(list[str], tools_used))}]"
|
||||
if tools_used
|
||||
else ""
|
||||
)
|
||||
timestamp = cast(str, message.get("timestamp", "?"))
|
||||
role = cast(str, message["role"])
|
||||
lines.append(
|
||||
f"[{message.get('timestamp', '?')[:16]}] "
|
||||
f"{message['role'].upper()}{tools}: {content}"
|
||||
f"[{timestamp[:16]}] {role.upper()}{tools}: {content}"
|
||||
)
|
||||
return "\n".join(lines)
|
||||
|
||||
def raw_archive(
|
||||
self,
|
||||
messages: list[dict],
|
||||
messages: list[dict[str, Any]],
|
||||
*,
|
||||
max_chars: int | None = None,
|
||||
session_key: str | None = None,
|
||||
@@ -766,9 +776,9 @@ class MemoryStore:
|
||||
Only current base64url-encoded Dream session keys are considered.
|
||||
Non-dream session files are never touched.
|
||||
"""
|
||||
dream_files = []
|
||||
dream_files: list[Path] = []
|
||||
for path in sessions_dir.glob("*.jsonl"):
|
||||
decoded_key = SessionManager._decode_storage_key(path.stem)
|
||||
decoded_key = SessionManager.decode_storage_key(path.stem)
|
||||
if decoded_key is not None and decoded_key.startswith("dream:"):
|
||||
dream_files.append(path)
|
||||
dream_files.sort(key=lambda p: p.stat().st_mtime)
|
||||
@@ -943,7 +953,13 @@ class Consolidator:
|
||||
channel = session.key.split(":", 1)[0] if ":" in session.key else None
|
||||
# Include archived summary in estimation so the budget accounts for it.
|
||||
meta = session.metadata.get("_last_summary")
|
||||
summary = meta.get("text") if isinstance(meta, dict) else (meta if isinstance(meta, str) else None)
|
||||
summary = (
|
||||
cast(dict[str, Any], meta).get("text")
|
||||
if isinstance(meta, dict)
|
||||
else meta
|
||||
if isinstance(meta, str)
|
||||
else None
|
||||
)
|
||||
probe_messages = self._build_messages(
|
||||
history=history,
|
||||
current_message="[token-probe]",
|
||||
@@ -976,11 +992,11 @@ class Consolidator:
|
||||
|
||||
async def archive(
|
||||
self,
|
||||
messages: list[dict],
|
||||
messages: list[dict[str, Any]],
|
||||
*,
|
||||
runtime: LLMRuntime,
|
||||
session_key: str | None = None,
|
||||
summary_messages: list[dict] | None = None,
|
||||
summary_messages: list[dict[str, Any]] | None = None,
|
||||
) -> str | None:
|
||||
"""Summarize messages via LLM and append to history.jsonl.
|
||||
|
||||
|
||||
@@ -5,9 +5,8 @@ from __future__ import annotations
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import replace
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from nanobot.config.schema import ModelPresetConfig
|
||||
from nanobot.config.schema import Config, ModelPresetConfig
|
||||
from nanobot.providers.base import LLMProvider
|
||||
from nanobot.providers.factory import ProviderSnapshot, build_provider_snapshot
|
||||
|
||||
@@ -22,7 +21,7 @@ def default_selection_signature(
|
||||
return (model_preset, *signature[:2]) if signature else None
|
||||
|
||||
|
||||
def configured_model_presets(config: Any) -> dict[str, ModelPresetConfig]:
|
||||
def configured_model_presets(config: Config) -> dict[str, ModelPresetConfig]:
|
||||
return {**config.model_presets, "default": config.resolve_default_preset()}
|
||||
|
||||
|
||||
@@ -41,7 +40,7 @@ def load_model_preset_catalog(
|
||||
|
||||
|
||||
def make_preset_snapshot_loader(
|
||||
config: Any,
|
||||
config: Config,
|
||||
provider_snapshot_loader: Callable[..., ProviderSnapshot] | None,
|
||||
) -> PresetSnapshotLoader:
|
||||
if provider_snapshot_loader is not None:
|
||||
|
||||
@@ -5,6 +5,7 @@ from __future__ import annotations
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import replace
|
||||
from types import MappingProxyType
|
||||
from typing import cast
|
||||
|
||||
from nanobot.agent import model_presets as preset_helpers
|
||||
from nanobot.config.schema import Config, ModelPresetConfig
|
||||
@@ -139,7 +140,7 @@ class ModelRuntimeResolver:
|
||||
|
||||
def select_model(self, model: str) -> LLMRuntime:
|
||||
"""Change the default model without reconstructing downstream consumers."""
|
||||
if not isinstance(model, str) or not model.strip():
|
||||
if not isinstance(cast(object, model), str) or not model.strip():
|
||||
raise ValueError("model must be a non-empty string")
|
||||
self._runtime = replace(
|
||||
self._runtime,
|
||||
@@ -150,8 +151,9 @@ class ModelRuntimeResolver:
|
||||
|
||||
def select_context_window(self, context_window_tokens: int) -> LLMRuntime:
|
||||
"""Change the default context limit for future admissions."""
|
||||
if not isinstance(context_window_tokens, int) or isinstance(
|
||||
context_window_tokens,
|
||||
raw_context_window = cast(object, context_window_tokens)
|
||||
if not isinstance(raw_context_window, int) or isinstance(
|
||||
raw_context_window,
|
||||
bool,
|
||||
):
|
||||
raise TypeError("context_window_tokens must be an integer")
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import json
|
||||
from typing import Any, Awaitable, Callable
|
||||
from typing import Any, Awaitable, Callable, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
@@ -124,7 +124,7 @@ class AgentProgressHook(AgentHook):
|
||||
arguments = event.get("arguments")
|
||||
if not isinstance(arguments, dict):
|
||||
arguments = {}
|
||||
payload = {
|
||||
payload: dict[str, Any] = {
|
||||
"version": 1,
|
||||
"phase": phase,
|
||||
"call_id": str(call_id),
|
||||
@@ -169,7 +169,7 @@ class AgentProgressHook(AgentHook):
|
||||
tool_events = [build_tool_event_start_payload(tc) for tc in context.tool_calls]
|
||||
await invoke_on_progress(
|
||||
self._on_progress,
|
||||
tool_hint,
|
||||
cast(str, tool_hint),
|
||||
tool_hint=True,
|
||||
tool_events=tool_events,
|
||||
)
|
||||
|
||||
+66
-39
@@ -5,10 +5,11 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
import inspect
|
||||
import os
|
||||
from collections.abc import Awaitable, Callable, Iterable
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable
|
||||
from typing import Any, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
@@ -48,6 +49,10 @@ from nanobot.utils.runtime import (
|
||||
)
|
||||
|
||||
GoalContinueMessage = str | Callable[[], str | None]
|
||||
ProgressCallback = Callable[[str], Awaitable[None]]
|
||||
RetryWaitCallback = Callable[[str], Awaitable[None]]
|
||||
CheckpointCallback = Callable[[dict[str, Any]], Awaitable[None]]
|
||||
InjectionCallback = Callable[..., Awaitable[Iterable[Any] | None]]
|
||||
|
||||
_DEFAULT_ERROR_MESSAGE = "Sorry, I encountered an error calling the AI model."
|
||||
_ARREARAGE_ERROR_MESSAGE = (
|
||||
@@ -90,11 +95,11 @@ class AgentRunSpec:
|
||||
session_key: str | None = None
|
||||
context_block_limit: int | None = None
|
||||
provider_retry_mode: str = "standard"
|
||||
progress_callback: Any | None = None
|
||||
progress_callback: ProgressCallback | None = None
|
||||
stream_progress_deltas: bool = True
|
||||
retry_wait_callback: Any | None = None
|
||||
checkpoint_callback: Any | None = None
|
||||
injection_callback: Any | None = None
|
||||
retry_wait_callback: RetryWaitCallback | None = None
|
||||
checkpoint_callback: CheckpointCallback | None = None
|
||||
injection_callback: InjectionCallback | None = None
|
||||
llm_timeout_s: float | None = None
|
||||
goal_active_predicate: Callable[[], bool] | None = None
|
||||
goal_continue_message: GoalContinueMessage | None = None
|
||||
@@ -131,8 +136,10 @@ class AgentRunner:
|
||||
def _to_blocks(value: Any) -> list[dict[str, Any]]:
|
||||
if isinstance(value, list):
|
||||
return [
|
||||
item if isinstance(item, dict) else {"type": "text", "text": str(item)}
|
||||
for item in value
|
||||
cast(dict[str, Any], item)
|
||||
if isinstance(item, dict)
|
||||
else {"type": "text", "text": str(item)}
|
||||
for item in cast(list[Any], value)
|
||||
]
|
||||
if value is None:
|
||||
return []
|
||||
@@ -158,25 +165,37 @@ class AgentRunner:
|
||||
merged = dict(messages[-1])
|
||||
left_meta = merged.get("_meta")
|
||||
right_meta = injection.get("_meta")
|
||||
left_meta_dict = cast(dict[str, Any], left_meta) if isinstance(left_meta, dict) else None
|
||||
right_meta_dict = (
|
||||
cast(dict[str, Any], right_meta) if isinstance(right_meta, dict) else None
|
||||
)
|
||||
left_marker = (
|
||||
left_meta.get(RUNTIME_CONTEXT_MESSAGE_META)
|
||||
if isinstance(left_meta, dict)
|
||||
left_meta_dict.get(RUNTIME_CONTEXT_MESSAGE_META)
|
||||
if left_meta_dict is not None
|
||||
else None
|
||||
)
|
||||
right_marker = (
|
||||
right_meta.get(RUNTIME_CONTEXT_MESSAGE_META)
|
||||
if isinstance(right_meta, dict)
|
||||
right_meta_dict.get(RUNTIME_CONTEXT_MESSAGE_META)
|
||||
if right_meta_dict is not None
|
||||
else None
|
||||
)
|
||||
left_marker_dict = (
|
||||
cast(dict[str, Any], left_marker) if isinstance(left_marker, dict) else None
|
||||
)
|
||||
right_marker_dict = (
|
||||
cast(dict[str, Any], right_marker) if isinstance(right_marker, dict) else None
|
||||
)
|
||||
empty_sources: list[str] = []
|
||||
empty_blocks: list[dict[str, Any]] = []
|
||||
detached_left = (
|
||||
detach_runtime_context(merged.get("content"), left_marker)
|
||||
if isinstance(left_marker, dict)
|
||||
else (merged.get("content"), [], [])
|
||||
detach_runtime_context(merged.get("content"), left_marker_dict)
|
||||
if left_marker_dict is not None
|
||||
else (merged.get("content"), empty_sources, empty_blocks)
|
||||
)
|
||||
detached_right = (
|
||||
detach_runtime_context(injection.get("content"), right_marker)
|
||||
if isinstance(right_marker, dict)
|
||||
else (injection.get("content"), [], [])
|
||||
detach_runtime_context(injection.get("content"), right_marker_dict)
|
||||
if right_marker_dict is not None
|
||||
else (injection.get("content"), empty_sources, empty_blocks)
|
||||
)
|
||||
if detached_left is not None and detached_right is not None:
|
||||
left_content, left_sources, left_blocks = detached_left
|
||||
@@ -189,9 +208,9 @@ class AgentRunner:
|
||||
[*left_sources, *right_sources],
|
||||
context_blocks,
|
||||
)
|
||||
internal_meta = dict(left_meta) if isinstance(left_meta, dict) else {}
|
||||
if isinstance(right_meta, dict):
|
||||
for key, value in right_meta.items():
|
||||
internal_meta = dict(left_meta_dict) if left_meta_dict is not None else {}
|
||||
if right_meta_dict is not None:
|
||||
for key, value in right_meta_dict.items():
|
||||
internal_meta.setdefault(key, value)
|
||||
internal_meta[RUNTIME_CONTEXT_MESSAGE_META] = marker
|
||||
merged["_meta"] = internal_meta
|
||||
@@ -302,11 +321,11 @@ class AgentRunner:
|
||||
for item in items:
|
||||
if item is None:
|
||||
continue
|
||||
if isinstance(item, dict) and item.get("role") == "user" and "content" in item:
|
||||
if self._has_injection_content(item.get("content")):
|
||||
injected_messages.append(item)
|
||||
continue
|
||||
if isinstance(item, dict):
|
||||
message_item = cast(dict[str, Any], item)
|
||||
if message_item.get("role") == "user" and "content" in message_item:
|
||||
if self._has_injection_content(message_item.get("content")):
|
||||
injected_messages.append(message_item)
|
||||
continue
|
||||
content = getattr(item, "content") if hasattr(item, "content") else str(item)
|
||||
if self._has_injection_content(content):
|
||||
@@ -327,7 +346,7 @@ class AgentRunner:
|
||||
if isinstance(content, str):
|
||||
return bool(content.strip())
|
||||
if isinstance(content, list):
|
||||
return bool(content)
|
||||
return bool(cast(list[Any], content))
|
||||
return True
|
||||
|
||||
async def run(self, spec: AgentRunSpec) -> AgentRunResult:
|
||||
@@ -592,7 +611,7 @@ class AgentRunner:
|
||||
if response.finish_reason == "length" and not is_blank_text(clean):
|
||||
if len(length_recovery_parts) < _MAX_LENGTH_RECOVERIES:
|
||||
length_recovery_parts.append(
|
||||
_restore_outer_whitespace(clean, original_content)
|
||||
_restore_outer_whitespace(clean or "", original_content)
|
||||
)
|
||||
logger.info(
|
||||
"Output truncated on turn {} for {} ({}/{}); continuing",
|
||||
@@ -609,7 +628,7 @@ class AgentRunner:
|
||||
reasoning_content=response.reasoning_content,
|
||||
thinking_blocks=response.thinking_blocks,
|
||||
))
|
||||
messages.append(build_length_recovery_message(clean))
|
||||
messages.append(build_length_recovery_message(clean or ""))
|
||||
await hook.after_iteration(context)
|
||||
continue
|
||||
|
||||
@@ -626,7 +645,7 @@ class AgentRunner:
|
||||
):
|
||||
await hook.on_stream(
|
||||
context,
|
||||
_restore_outer_whitespace(clean, original_content),
|
||||
_restore_outer_whitespace(clean or "", original_content),
|
||||
)
|
||||
context.streamed_content = True
|
||||
|
||||
@@ -717,7 +736,7 @@ class AgentRunner:
|
||||
if length_recovery_parts:
|
||||
final_content = (
|
||||
"".join(length_recovery_parts)
|
||||
+ _restore_outer_whitespace(clean, original_content)
|
||||
+ _restore_outer_whitespace(clean or "", original_content)
|
||||
).strip()
|
||||
else:
|
||||
final_content = clean
|
||||
@@ -798,7 +817,7 @@ class AgentRunner:
|
||||
context: AgentHookContext,
|
||||
*,
|
||||
malformed_retry: bool = False,
|
||||
):
|
||||
) -> LLMResponse:
|
||||
timeout_s: float | None = spec.llm_timeout_s
|
||||
if timeout_s is None:
|
||||
# Default to a finite timeout to avoid per-session lock starvation when an LLM
|
||||
@@ -809,7 +828,7 @@ class AgentRunner:
|
||||
timeout_s = float(raw)
|
||||
except (TypeError, ValueError):
|
||||
timeout_s = 300.0
|
||||
if timeout_s is not None and timeout_s <= 0:
|
||||
if timeout_s <= 0:
|
||||
timeout_s = None
|
||||
|
||||
kwargs = self._build_request_kwargs(
|
||||
@@ -818,10 +837,11 @@ class AgentRunner:
|
||||
tools=spec.tools.get_definitions(),
|
||||
)
|
||||
wants_streaming = hook.wants_streaming()
|
||||
progress_callback = spec.progress_callback
|
||||
wants_progress_streaming = (
|
||||
not wants_streaming
|
||||
and spec.stream_progress_deltas
|
||||
and spec.progress_callback is not None
|
||||
and progress_callback is not None
|
||||
and getattr(spec.runtime.provider, "supports_progress_deltas", False) is True
|
||||
)
|
||||
|
||||
@@ -894,7 +914,9 @@ class AgentRunner:
|
||||
await hook.emit_reasoning_end()
|
||||
progress_state["reasoning_open"] = False
|
||||
context.streamed_content = True
|
||||
await spec.progress_callback(incremental)
|
||||
callback = progress_callback
|
||||
if callback is not None:
|
||||
await callback(incremental)
|
||||
|
||||
coro = spec.runtime.provider.chat_stream_with_retry(
|
||||
**kwargs,
|
||||
@@ -1038,7 +1060,7 @@ class AgentRunner:
|
||||
self,
|
||||
spec: AgentRunSpec,
|
||||
messages: list[dict[str, Any]],
|
||||
):
|
||||
) -> LLMResponse:
|
||||
retry_messages = self._finalization_retry_messages(messages)
|
||||
return await self._request_no_tools(spec, retry_messages)
|
||||
|
||||
@@ -1224,7 +1246,7 @@ class AgentRunner:
|
||||
))
|
||||
tool_results.extend(batch_results)
|
||||
else:
|
||||
batch_results = []
|
||||
batch_results: list[tuple[Any, dict[str, str], BaseException | None]] = []
|
||||
for tool_call in batch:
|
||||
result = await self._run_tool(
|
||||
spec,
|
||||
@@ -1273,12 +1295,17 @@ class AgentRunner:
|
||||
if spec.fail_on_tool_error:
|
||||
return lookup_error + hint, event, RuntimeError(lookup_error)
|
||||
return lookup_error + hint, event, None
|
||||
prepare_call = getattr(spec.tools, "prepare_call", None)
|
||||
prepare_call = cast(
|
||||
Callable[[str, Any], object] | None,
|
||||
getattr(spec.tools, "prepare_call", None),
|
||||
)
|
||||
tool, params, prep_error = None, tool_call.arguments, None
|
||||
if callable(prepare_call):
|
||||
prepared = prepare_call(tool_call.name, tool_call.arguments)
|
||||
if isinstance(prepared, tuple) and len(prepared) == 3:
|
||||
tool, params, prep_error = prepared
|
||||
if isinstance(prepared, tuple):
|
||||
prepared_tuple = cast(tuple[object, ...], prepared)
|
||||
if len(prepared_tuple) == 3:
|
||||
tool, params, prep_error = cast(tuple[Any, Any, str | None], prepared_tuple)
|
||||
if prep_error:
|
||||
event = {
|
||||
"name": tool_call.name,
|
||||
@@ -1490,7 +1517,7 @@ class AgentRunner:
|
||||
batches: list[list[ToolCallRequest]] = []
|
||||
current: list[ToolCallRequest] = []
|
||||
for tool_call in tool_calls:
|
||||
get_tool = getattr(spec.tools, "get", None)
|
||||
get_tool = cast(Callable[[str], Any] | None, getattr(spec.tools, "get", None))
|
||||
tool = get_tool(tool_call.name) if callable(get_tool) else None
|
||||
can_batch = bool(tool and tool.concurrency_safe)
|
||||
if can_batch:
|
||||
|
||||
+23
-20
@@ -5,6 +5,7 @@ import os
|
||||
import re
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
|
||||
import yaml
|
||||
|
||||
@@ -144,7 +145,7 @@ class SkillsLoader:
|
||||
skill_name = entry["name"]
|
||||
meta = self._get_skill_meta(skill_name)
|
||||
available = self._check_requirements(meta)
|
||||
desc = self._get_skill_description(skill_name)
|
||||
desc = self.get_skill_description(skill_name)
|
||||
suffix = ""
|
||||
if not available:
|
||||
missing = self._get_missing_requirements(meta)
|
||||
@@ -155,18 +156,18 @@ class SkillsLoader:
|
||||
return "\n\n".join(sections)
|
||||
|
||||
@staticmethod
|
||||
def _requirement_lists(skill_meta: dict) -> tuple[list[str], list[str]]:
|
||||
def _requirement_lists(skill_meta: dict[str, Any]) -> tuple[list[str], list[str]]:
|
||||
"""Return (bins, env) lists from skill metadata, tolerating null/wrong shapes."""
|
||||
requires = skill_meta.get("requires") or {}
|
||||
if not isinstance(requires, dict):
|
||||
requires = cast(dict[str, Any], skill_meta.get("requires") or {})
|
||||
if not isinstance(skill_meta.get("requires") or {}, dict):
|
||||
return [], []
|
||||
bins_raw = requires.get("bins") or []
|
||||
env_raw = requires.get("env") or []
|
||||
bins = [str(v) for v in bins_raw if isinstance(v, str) and v.strip()] if isinstance(bins_raw, list) else []
|
||||
env = [str(v) for v in env_raw if isinstance(v, str) and v.strip()] if isinstance(env_raw, list) else []
|
||||
bins_raw: object = requires.get("bins") or []
|
||||
env_raw: object = requires.get("env") or []
|
||||
bins = [value for value in cast(list[object], bins_raw) if isinstance(value, str) and value.strip()] if isinstance(bins_raw, list) else []
|
||||
env = [value for value in cast(list[object], env_raw) if isinstance(value, str) and value.strip()] if isinstance(env_raw, list) else []
|
||||
return bins, env
|
||||
|
||||
def _get_missing_requirements(self, skill_meta: dict) -> str:
|
||||
def _get_missing_requirements(self, skill_meta: dict[str, Any]) -> str:
|
||||
"""Get a description of missing requirements."""
|
||||
required_bins, required_env_vars = self._requirement_lists(skill_meta)
|
||||
return ", ".join(
|
||||
@@ -190,11 +191,12 @@ class SkillsLoader:
|
||||
"missing_env": [value for value in env if not os.environ.get(value)],
|
||||
}
|
||||
|
||||
def _get_skill_description(self, name: str) -> str:
|
||||
def get_skill_description(self, name: str) -> str:
|
||||
"""Get the description of a skill from its frontmatter."""
|
||||
meta = self.get_skill_metadata(name)
|
||||
if meta and meta.get("description"):
|
||||
return meta["description"]
|
||||
description = meta.get("description") if meta else None
|
||||
if isinstance(description, str) and description:
|
||||
return description
|
||||
return name # Fallback to skill name
|
||||
|
||||
def _strip_frontmatter(self, content: str) -> str:
|
||||
@@ -206,13 +208,13 @@ class SkillsLoader:
|
||||
return content[match.end():].strip()
|
||||
return content
|
||||
|
||||
def _parse_nanobot_metadata(self, raw: object) -> dict:
|
||||
def _parse_nanobot_metadata(self, raw: object) -> dict[str, Any]:
|
||||
"""Extract nanobot/openclaw metadata from a frontmatter field.
|
||||
|
||||
``raw`` may be a dict (already parsed by yaml.safe_load) or a JSON str.
|
||||
"""
|
||||
if isinstance(raw, dict):
|
||||
data = raw
|
||||
data = cast(dict[str, Any], raw)
|
||||
elif isinstance(raw, str):
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
@@ -222,17 +224,18 @@ class SkillsLoader:
|
||||
return {}
|
||||
if not isinstance(data, dict):
|
||||
return {}
|
||||
payload = data.get("nanobot", data.get("openclaw", {}))
|
||||
return payload if isinstance(payload, dict) else {}
|
||||
data_object = cast(dict[str, Any], data)
|
||||
payload = data_object.get("nanobot", data_object.get("openclaw", {}))
|
||||
return cast(dict[str, Any], payload) if isinstance(payload, dict) else {}
|
||||
|
||||
def _check_requirements(self, skill_meta: dict) -> bool:
|
||||
def _check_requirements(self, skill_meta: dict[str, Any]) -> bool:
|
||||
"""Check if skill requirements are met (bins, env vars)."""
|
||||
required_bins, required_env_vars = self._requirement_lists(skill_meta)
|
||||
return all(shutil.which(cmd) for cmd in required_bins) and all(
|
||||
os.environ.get(var) for var in required_env_vars
|
||||
)
|
||||
|
||||
def _get_skill_meta(self, name: str) -> dict:
|
||||
def _get_skill_meta(self, name: str) -> dict[str, Any]:
|
||||
"""Get nanobot metadata for a skill (cached in frontmatter)."""
|
||||
raw_meta = self.get_skill_metadata(name) or {}
|
||||
return self._parse_nanobot_metadata(raw_meta.get("metadata"))
|
||||
@@ -249,7 +252,7 @@ class SkillsLoader:
|
||||
)
|
||||
]
|
||||
|
||||
def get_skill_metadata(self, name: str) -> dict | None:
|
||||
def get_skill_metadata(self, name: str) -> dict[str, object] | None:
|
||||
"""
|
||||
Get metadata from a skill's frontmatter.
|
||||
|
||||
@@ -274,6 +277,6 @@ class SkillsLoader:
|
||||
# yaml.safe_load returns native types (int, bool, list, etc.);
|
||||
# keep values as-is so downstream consumers get correct types.
|
||||
metadata: dict[str, object] = {}
|
||||
for key, value in parsed.items():
|
||||
for key, value in cast(dict[object, object], parsed).items():
|
||||
metadata[str(key)] = value
|
||||
return metadata
|
||||
|
||||
+21
-11
@@ -7,12 +7,12 @@ import uuid
|
||||
import warnings
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable
|
||||
from typing import Any, Callable, TypedDict
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.agent.hook import AgentHook, AgentHookContext
|
||||
from nanobot.agent.runner import AgentRunner, AgentRunSpec
|
||||
from nanobot.agent.runner import AgentRunner, AgentRunResult, AgentRunSpec
|
||||
from nanobot.agent.tools.base import ToolResult
|
||||
from nanobot.agent.tools.context import (
|
||||
RequestContext,
|
||||
@@ -38,6 +38,12 @@ from nanobot.utils.llm_runtime import LLMRuntime
|
||||
from nanobot.utils.prompt_templates import render_template
|
||||
|
||||
|
||||
class _SubagentOrigin(TypedDict):
|
||||
channel: str
|
||||
chat_id: str
|
||||
session_key: str | None
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SubagentStatus:
|
||||
"""Real-time status of a running subagent."""
|
||||
@@ -48,8 +54,8 @@ class SubagentStatus:
|
||||
started_at: float # time.monotonic()
|
||||
phase: str = "initializing" # initializing | awaiting_tools | tools_completed | final_response | done | error
|
||||
iteration: int = 0
|
||||
tool_events: list = field(default_factory=list) # [{name, status, detail}, ...]
|
||||
usage: dict = field(default_factory=dict) # token usage
|
||||
tool_events: list[dict[str, str]] = field(default_factory=list)
|
||||
usage: dict[str, int] = field(default_factory=dict)
|
||||
stop_reason: str | None = None
|
||||
error: str | None = None
|
||||
|
||||
@@ -237,7 +243,11 @@ class SubagentManager:
|
||||
runtime = runtime.with_generation_overrides(temperature=temperature)
|
||||
task_id = str(uuid.uuid4())[:8]
|
||||
display_label = label or task[:30] + ("..." if len(task) > 30 else "")
|
||||
origin = {"channel": origin_channel, "chat_id": origin_chat_id, "session_key": session_key}
|
||||
origin: _SubagentOrigin = {
|
||||
"channel": origin_channel,
|
||||
"chat_id": origin_chat_id,
|
||||
"session_key": session_key,
|
||||
}
|
||||
|
||||
status = SubagentStatus(
|
||||
task_id=task_id,
|
||||
@@ -263,7 +273,7 @@ class SubagentManager:
|
||||
if session_key:
|
||||
self._session_tasks.setdefault(session_key, set()).add(task_id)
|
||||
|
||||
def _cleanup(_: asyncio.Task) -> None:
|
||||
def _cleanup(_: asyncio.Task[str]) -> None:
|
||||
self._running_tasks.pop(task_id, None)
|
||||
self._task_statuses.pop(task_id, None)
|
||||
if session_key and (ids := self._session_tasks.get(session_key)):
|
||||
@@ -296,7 +306,7 @@ class SubagentManager:
|
||||
runtime = runtime.with_generation_overrides(temperature=temperature)
|
||||
task_id = str(uuid.uuid4())[:8]
|
||||
display_label = label or task[:30] + ("..." if len(task) > 30 else "")
|
||||
origin = {
|
||||
origin: _SubagentOrigin = {
|
||||
"channel": origin_channel,
|
||||
"chat_id": origin_chat_id,
|
||||
"session_key": session_key,
|
||||
@@ -343,7 +353,7 @@ class SubagentManager:
|
||||
task_id: str,
|
||||
task: str,
|
||||
label: str,
|
||||
origin: dict[str, str],
|
||||
origin: _SubagentOrigin,
|
||||
status: SubagentStatus,
|
||||
runtime: LLMRuntime,
|
||||
origin_message_id: str | None = None,
|
||||
@@ -354,7 +364,7 @@ class SubagentManager:
|
||||
"""Execute the subagent task and announce the result."""
|
||||
logger.info("Subagent [{}] starting task: {}", task_id, label)
|
||||
|
||||
async def _on_checkpoint(payload: dict) -> None:
|
||||
async def _on_checkpoint(payload: dict[str, Any]) -> None:
|
||||
status.phase = payload.get("phase", status.phase)
|
||||
status.iteration = payload.get("iteration", status.iteration)
|
||||
|
||||
@@ -456,7 +466,7 @@ class SubagentManager:
|
||||
label: str,
|
||||
task: str,
|
||||
result: str,
|
||||
origin: dict[str, str],
|
||||
origin: _SubagentOrigin,
|
||||
status: str,
|
||||
origin_message_id: str | None = None,
|
||||
) -> None:
|
||||
@@ -496,7 +506,7 @@ class SubagentManager:
|
||||
logger.debug("Subagent [{}] announced result to {}:{}", task_id, origin['channel'], origin['chat_id'])
|
||||
|
||||
@staticmethod
|
||||
def _format_partial_progress(result) -> str:
|
||||
def _format_partial_progress(result: AgentRunResult) -> str:
|
||||
completed = [e for e in result.tool_events if e["status"] == "ok"]
|
||||
failure = next((e for e in reversed(result.tool_events) if e["status"] == "error"), None)
|
||||
lines: list[str] = []
|
||||
|
||||
@@ -5,10 +5,10 @@ from __future__ import annotations
|
||||
import difflib
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
from nanobot.agent.tools.base import ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.filesystem import _FsTool
|
||||
from nanobot.agent.tools.filesystem import _FsTool # pyright: ignore[reportPrivateUsage]
|
||||
from nanobot.agent.tools.schema import (
|
||||
ArraySchema,
|
||||
BooleanSchema,
|
||||
@@ -134,7 +134,7 @@ class ApplyPatchTool(_FsTool):
|
||||
|
||||
async def execute(
|
||||
self,
|
||||
edits: list[dict] | None = None,
|
||||
edits: list[object] | None = None,
|
||||
dry_run: bool = False,
|
||||
**kwargs: Any,
|
||||
) -> str:
|
||||
@@ -145,9 +145,10 @@ class ApplyPatchTool(_FsTool):
|
||||
writes: dict[Path, str] = {}
|
||||
summaries: list[_PatchSummary] = []
|
||||
|
||||
for edit in edits:
|
||||
if not isinstance(edit, dict):
|
||||
for edit_value in edits:
|
||||
if not isinstance(edit_value, dict):
|
||||
raise _PatchError("each edit must be an object")
|
||||
edit = cast(dict[str, Any], edit_value)
|
||||
raw_path = edit.get("path")
|
||||
if not isinstance(raw_path, str):
|
||||
raise _PatchError("path required for edit")
|
||||
@@ -161,6 +162,7 @@ class ApplyPatchTool(_FsTool):
|
||||
new_text = edit.get("new_text")
|
||||
if new_text is None:
|
||||
raise _PatchError(f"new_text required for add: {path}")
|
||||
new_text = cast(str, new_text)
|
||||
|
||||
pending = writes.get(source)
|
||||
if pending is not None:
|
||||
@@ -204,9 +206,11 @@ class ApplyPatchTool(_FsTool):
|
||||
old_text = edit.get("old_text") or ""
|
||||
if not old_text:
|
||||
raise _PatchError(f"old_text required for replace: {path}")
|
||||
old_text = cast(str, old_text)
|
||||
new_text = edit.get("new_text")
|
||||
if new_text is None:
|
||||
raise _PatchError(f"new_text required for replace: {path}")
|
||||
new_text = cast(str, new_text)
|
||||
|
||||
pending = writes.get(source)
|
||||
if pending is not None:
|
||||
|
||||
+31
-20
@@ -5,7 +5,7 @@ import typing
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Callable
|
||||
from copy import deepcopy
|
||||
from typing import Any, TypeVar
|
||||
from typing import Any, TypeVar, cast
|
||||
|
||||
if typing.TYPE_CHECKING:
|
||||
from pydantic import BaseModel
|
||||
@@ -38,8 +38,9 @@ class Schema(ABC):
|
||||
def resolve_json_schema_type(t: Any) -> str | None:
|
||||
"""Resolve the non-null type name from JSON Schema ``type`` (e.g. ``['string','null']`` -> ``'string'``)."""
|
||||
if isinstance(t, list):
|
||||
return next((x for x in t if x != "null"), None)
|
||||
return t # type: ignore[return-value]
|
||||
types = cast(list[Any], t)
|
||||
return cast(str | None, next((x for x in types if x != "null"), None))
|
||||
return cast(str | None, t)
|
||||
|
||||
@staticmethod
|
||||
def subpath(path: str, key: str) -> str:
|
||||
@@ -76,33 +77,41 @@ class Schema(ABC):
|
||||
if "maximum" in schema and val > schema["maximum"]:
|
||||
errors.append(f"{label} must be <= {schema['maximum']}")
|
||||
if t == "string":
|
||||
if "minLength" in schema and len(val) < schema["minLength"]:
|
||||
string_value = cast(str, val)
|
||||
if "minLength" in schema and len(string_value) < schema["minLength"]:
|
||||
errors.append(f"{label} must be at least {schema['minLength']} chars")
|
||||
if "maxLength" in schema and len(val) > schema["maxLength"]:
|
||||
if "maxLength" in schema and len(string_value) > schema["maxLength"]:
|
||||
errors.append(f"{label} must be at most {schema['maxLength']} chars")
|
||||
if t == "object":
|
||||
props = schema.get("properties", {})
|
||||
for k in schema.get("required", []):
|
||||
if k not in val:
|
||||
object_value = cast(dict[str, Any], val)
|
||||
props = cast(dict[str, Any], schema.get("properties", {}))
|
||||
required = cast(list[Any], schema.get("required", []))
|
||||
for k in required:
|
||||
if k not in object_value:
|
||||
errors.append(f"missing required {Schema.subpath(path, k)}")
|
||||
additional = schema.get("additionalProperties", True)
|
||||
for k, v in val.items():
|
||||
for k, v in object_value.items():
|
||||
if k in props:
|
||||
errors.extend(Schema.validate_json_schema_value(v, props[k], Schema.subpath(path, k)))
|
||||
elif additional is False:
|
||||
errors.append(f"unexpected parameter {Schema.subpath(path, k)}")
|
||||
elif isinstance(additional, dict):
|
||||
errors.extend(
|
||||
Schema.validate_json_schema_value(v, additional, Schema.subpath(path, k))
|
||||
Schema.validate_json_schema_value(
|
||||
v,
|
||||
cast(dict[str, Any], additional),
|
||||
Schema.subpath(path, k),
|
||||
)
|
||||
)
|
||||
if t == "array":
|
||||
if "minItems" in schema and len(val) < schema["minItems"]:
|
||||
array_value = cast(list[Any], val)
|
||||
if "minItems" in schema and len(array_value) < schema["minItems"]:
|
||||
errors.append(f"{label} must have at least {schema['minItems']} items")
|
||||
if "maxItems" in schema and len(val) > schema["maxItems"]:
|
||||
if "maxItems" in schema and len(array_value) > schema["maxItems"]:
|
||||
errors.append(f"{label} must be at most {schema['maxItems']} items")
|
||||
if "items" in schema:
|
||||
prefix = f"{path}[{{}}]" if path else "[{}]"
|
||||
for i, item in enumerate(val):
|
||||
for i, item in enumerate(array_value):
|
||||
errors.extend(
|
||||
Schema.validate_json_schema_value(item, schema["items"], prefix.format(i))
|
||||
)
|
||||
@@ -114,9 +123,9 @@ class Schema(ABC):
|
||||
# Try to_json_schema first: Schema instances must be distinguished from dicts that are already JSON Schema
|
||||
to_js = getattr(value, "to_json_schema", None)
|
||||
if callable(to_js):
|
||||
return to_js()
|
||||
return cast(dict[str, Any], to_js())
|
||||
if isinstance(value, dict):
|
||||
return value
|
||||
return cast(dict[str, Any], value)
|
||||
raise TypeError(f"Expected schema object or dict, got {type(value).__name__}")
|
||||
|
||||
@abstractmethod
|
||||
@@ -223,14 +232,15 @@ class Tool(ABC):
|
||||
def _cast_object(self, obj: Any, schema: dict[str, Any]) -> dict[str, Any]:
|
||||
if not isinstance(obj, dict):
|
||||
return obj
|
||||
props = schema.get("properties", {})
|
||||
props = cast(dict[str, Any], schema.get("properties", {}))
|
||||
additional = schema.get("additionalProperties")
|
||||
casted: dict[str, Any] = {}
|
||||
for k, v in obj.items():
|
||||
object_value = cast(dict[str, Any], obj)
|
||||
for k, v in object_value.items():
|
||||
if k in props:
|
||||
casted[k] = self._cast_value(v, props[k])
|
||||
elif isinstance(additional, dict):
|
||||
casted[k] = self._cast_value(v, additional)
|
||||
casted[k] = self._cast_value(v, cast(dict[str, Any], additional))
|
||||
else:
|
||||
casted[k] = v
|
||||
return casted
|
||||
@@ -273,7 +283,8 @@ class Tool(ABC):
|
||||
|
||||
if t == "array" and isinstance(val, list):
|
||||
items = schema.get("items")
|
||||
return [self._cast_value(x, items) for x in val] if items else val
|
||||
array_value = cast(list[Any], val)
|
||||
return [self._cast_value(x, items) for x in array_value] if items else array_value
|
||||
|
||||
if t == "object" and isinstance(val, dict):
|
||||
return self._cast_object(val, schema)
|
||||
@@ -282,7 +293,7 @@ class Tool(ABC):
|
||||
|
||||
def validate_params(self, params: dict[str, Any]) -> list[str]:
|
||||
"""Validate against JSON schema; empty list means valid."""
|
||||
if not isinstance(params, dict):
|
||||
if not isinstance(cast(object, params), dict):
|
||||
return [f"parameters must be an object, got {type(params).__name__}"]
|
||||
schema = self.parameters or {}
|
||||
if schema.get("type", "object") != "object":
|
||||
|
||||
@@ -1,14 +1,15 @@
|
||||
"""Controlled runner for installed CLI Apps."""
|
||||
|
||||
# pyright: reportIncompatibleMethodOverride=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import RequestContext
|
||||
from nanobot.agent.tools.context import RequestContext, ToolContext
|
||||
from nanobot.agent.tools.schema import (
|
||||
ArraySchema,
|
||||
BooleanSchema,
|
||||
@@ -66,11 +67,11 @@ class CliAppsTool(Tool):
|
||||
return CliAppsToolConfig
|
||||
|
||||
@classmethod
|
||||
def enabled(cls, ctx: Any) -> bool:
|
||||
def enabled(cls, ctx: ToolContext) -> bool:
|
||||
return ctx.config.cli_apps.enable
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
def create(cls, ctx: ToolContext) -> Tool:
|
||||
cfg = ctx.config.cli_apps
|
||||
return cls(
|
||||
workspace=Path(ctx.workspace),
|
||||
|
||||
@@ -8,6 +8,16 @@ from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, Callable, Protocol, runtime_checkable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.subagent import SubagentManager
|
||||
from nanobot.agent.tools.exec_session import ExecSessionManager
|
||||
from nanobot.agent.tools.file_state import FileStates
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.bus.runtime_events import RuntimeEventBus
|
||||
from nanobot.config.schema import ProviderConfig, ToolsConfig
|
||||
from nanobot.cron.service import CronService
|
||||
from nanobot.providers.factory import ProviderSnapshot
|
||||
from nanobot.security.workspace_access import WorkspaceSandboxStatus
|
||||
from nanobot.session.manager import SessionManager
|
||||
from nanobot.utils.llm_runtime import LLMRuntime
|
||||
|
||||
_CURRENT_REQUEST_CONTEXT: ContextVar["RequestContext | None"] = ContextVar(
|
||||
@@ -67,16 +77,16 @@ def current_request_session_key() -> str | None:
|
||||
|
||||
@dataclass
|
||||
class ToolContext:
|
||||
config: Any
|
||||
config: ToolsConfig
|
||||
workspace: str
|
||||
bus: Any | None = None
|
||||
subagent_manager: Any | None = None
|
||||
cron_service: Any | None = None
|
||||
exec_session_manager: Any | None = None
|
||||
sessions: Any | None = None
|
||||
file_state_store: Any = field(default=None)
|
||||
provider_snapshot_loader: Callable[[], Any] | None = None
|
||||
image_generation_provider_configs: dict[str, Any] | None = None
|
||||
bus: MessageBus | None = None
|
||||
subagent_manager: SubagentManager | None = None
|
||||
cron_service: CronService | None = None
|
||||
exec_session_manager: ExecSessionManager | None = None
|
||||
sessions: SessionManager | None = None
|
||||
file_state_store: FileStates | None = None
|
||||
provider_snapshot_loader: Callable[..., ProviderSnapshot] | None = None
|
||||
image_generation_provider_configs: dict[str, ProviderConfig] | None = None
|
||||
timezone: str = "UTC"
|
||||
workspace_sandbox: Any | None = None
|
||||
runtime_events: Any | None = None
|
||||
workspace_sandbox: WorkspaceSandboxStatus | None = None
|
||||
runtime_events: RuntimeEventBus | None = None
|
||||
|
||||
@@ -1,13 +1,15 @@
|
||||
"""Cron tool for scheduling reminders and tasks."""
|
||||
|
||||
# pyright: reportIncompatibleMethodOverride=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from contextvars import ContextVar
|
||||
from contextvars import ContextVar, Token
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import current_request_context
|
||||
from nanobot.agent.tools.context import ToolContext, current_request_context
|
||||
from nanobot.agent.tools.schema import (
|
||||
IntegerSchema,
|
||||
StringSchema,
|
||||
@@ -60,12 +62,15 @@ class CronTool(Tool):
|
||||
self._in_cron_context: ContextVar[bool] = ContextVar("cron_in_context", default=False)
|
||||
|
||||
@classmethod
|
||||
def enabled(cls, ctx: Any) -> bool:
|
||||
def enabled(cls, ctx: ToolContext) -> bool:
|
||||
return ctx.cron_service is not None
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
return cls(cron_service=ctx.cron_service, default_timezone=ctx.timezone)
|
||||
def create(cls, ctx: ToolContext) -> Tool:
|
||||
cron_service = ctx.cron_service
|
||||
if cron_service is None:
|
||||
raise RuntimeError("CronTool requires an initialized cron service")
|
||||
return cls(cron_service=cron_service, default_timezone=ctx.timezone)
|
||||
|
||||
@staticmethod
|
||||
def _request_route() -> tuple[str, str, str, dict[str, Any]]:
|
||||
@@ -79,11 +84,11 @@ class CronTool(Tool):
|
||||
)
|
||||
return session_key, ctx.channel or "", ctx.chat_id or "", dict(ctx.metadata or {})
|
||||
|
||||
def set_cron_context(self, active: bool):
|
||||
def set_cron_context(self, active: bool) -> Token[bool]:
|
||||
"""Mark whether the tool is executing inside a cron job callback."""
|
||||
return self._in_cron_context.set(active)
|
||||
|
||||
def reset_cron_context(self, token) -> None:
|
||||
def reset_cron_context(self, token: Token[bool]) -> None:
|
||||
"""Restore previous cron context."""
|
||||
self._in_cron_context.reset(token)
|
||||
|
||||
@@ -257,7 +262,7 @@ class CronTool(Tool):
|
||||
jobs = self._cron.list_jobs()
|
||||
if not jobs:
|
||||
return "No scheduled jobs."
|
||||
lines = []
|
||||
lines: list[str] = []
|
||||
for j in jobs:
|
||||
timing = self._format_timing(j.schedule)
|
||||
parts = [f"- {j.name} (id: {j.id}, {timing})"]
|
||||
|
||||
@@ -10,7 +10,7 @@ from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import current_request_session_key
|
||||
from nanobot.agent.tools.context import ToolContext, current_request_session_key
|
||||
from nanobot.agent.tools.schema import (
|
||||
BooleanSchema,
|
||||
IntegerSchema,
|
||||
@@ -151,8 +151,8 @@ class _ExecSession:
|
||||
timeout=2.0,
|
||||
)
|
||||
# Safety-net reap after normal exit.
|
||||
from nanobot.agent.tools.shell import _reap_pid
|
||||
_reap_pid(self.process.pid)
|
||||
from nanobot.agent.tools.shell import _reap_pid # pyright: ignore[reportPrivateUsage]
|
||||
_reap_pid(self.process.pid) # pyright: ignore[reportPrivateUsage]
|
||||
elif yield_time_ms > 0:
|
||||
await self._wait_for_buffered_output()
|
||||
|
||||
@@ -177,9 +177,9 @@ class _ExecSession:
|
||||
|
||||
try:
|
||||
if self._process_tree:
|
||||
await ExecTool._kill_process_tree(self.process)
|
||||
await ExecTool._kill_process_tree(self.process) # pyright: ignore[reportPrivateUsage]
|
||||
else:
|
||||
await ExecTool._kill_process(self.process)
|
||||
await ExecTool._kill_process(self.process) # pyright: ignore[reportPrivateUsage]
|
||||
finally:
|
||||
with suppress(asyncio.TimeoutError):
|
||||
await asyncio.wait_for(
|
||||
@@ -311,13 +311,13 @@ class ExecSessionManager:
|
||||
"""Terminate and remove all active sessions during shutdown."""
|
||||
async with self._lock:
|
||||
self._closed = True
|
||||
sessions = list(self._sessions.values())
|
||||
sessions: list[_ExecSession] = list(self._sessions.values())
|
||||
self._sessions.clear()
|
||||
results = await asyncio.gather(
|
||||
results: list[None | BaseException] = list(await asyncio.gather(
|
||||
*(session.kill() for session in sessions),
|
||||
return_exceptions=True,
|
||||
)
|
||||
failures = [
|
||||
))
|
||||
failures: list[tuple[_ExecSession, BaseException]] = [
|
||||
(session, result)
|
||||
for session, result in zip(sessions, results, strict=True)
|
||||
if isinstance(result, BaseException)
|
||||
@@ -337,15 +337,15 @@ class ExecSessionManager:
|
||||
async def terminate_by_owner(self, owner_session_key: str) -> int:
|
||||
"""Terminate all sessions owned by owner_session_key. Returns count."""
|
||||
async with self._lock:
|
||||
victims = []
|
||||
victims: list[_ExecSession] = []
|
||||
for sid, s in list(self._sessions.items()):
|
||||
if s.owner_session_key == owner_session_key:
|
||||
victims.append(self._sessions.pop(sid))
|
||||
results = await asyncio.gather(
|
||||
results: list[None | BaseException] = list(await asyncio.gather(
|
||||
*(s.kill() for s in victims),
|
||||
return_exceptions=True,
|
||||
)
|
||||
failures = [
|
||||
))
|
||||
failures: list[tuple[_ExecSession, BaseException]] = [
|
||||
(session, result)
|
||||
for session, result in zip(victims, results, strict=True)
|
||||
if isinstance(result, BaseException)
|
||||
@@ -384,7 +384,7 @@ class ExecSessionManager:
|
||||
) -> asyncio.subprocess.Process:
|
||||
from nanobot.agent.tools.shell import ExecTool
|
||||
|
||||
return await ExecTool._spawn(
|
||||
return await ExecTool._spawn( # pyright: ignore[reportPrivateUsage]
|
||||
command, cwd, env, shell_program, login,
|
||||
stdin=asyncio.subprocess.PIPE,
|
||||
process_tree=True,
|
||||
@@ -489,7 +489,7 @@ class WriteStdinTool(Tool):
|
||||
return ExecToolConfig
|
||||
|
||||
@classmethod
|
||||
def enabled(cls, ctx: Any) -> bool:
|
||||
def enabled(cls, ctx: ToolContext) -> bool:
|
||||
return ctx.config.exec.enable
|
||||
|
||||
def __init__(
|
||||
@@ -500,8 +500,8 @@ class WriteStdinTool(Tool):
|
||||
self._manager = manager or DEFAULT_EXEC_SESSION_MANAGER
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
return cls(manager=getattr(ctx, "exec_session_manager", None))
|
||||
def create(cls, ctx: ToolContext) -> Tool:
|
||||
return cls(manager=ctx.exec_session_manager)
|
||||
|
||||
@property
|
||||
def exclusive(self) -> bool:
|
||||
@@ -522,7 +522,7 @@ class WriteStdinTool(Tool):
|
||||
"Do not use this to start new commands; start them with exec."
|
||||
)
|
||||
|
||||
async def execute(
|
||||
async def execute( # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
self,
|
||||
session_id: str,
|
||||
chars: str | None = None,
|
||||
@@ -633,7 +633,7 @@ class ListExecSessionsTool(Tool):
|
||||
return ExecToolConfig
|
||||
|
||||
@classmethod
|
||||
def enabled(cls, ctx: Any) -> bool:
|
||||
def enabled(cls, ctx: ToolContext) -> bool:
|
||||
return ctx.config.exec.enable
|
||||
|
||||
def __init__(
|
||||
@@ -644,8 +644,8 @@ class ListExecSessionsTool(Tool):
|
||||
self._manager = manager or DEFAULT_EXEC_SESSION_MANAGER
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
return cls(manager=getattr(ctx, "exec_session_manager", None))
|
||||
def create(cls, ctx: ToolContext) -> Tool:
|
||||
return cls(manager=ctx.exec_session_manager)
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
@@ -671,7 +671,7 @@ class ListExecSessionsTool(Tool):
|
||||
)
|
||||
if not sessions:
|
||||
return "No active exec sessions."
|
||||
lines = []
|
||||
lines: list[str] = []
|
||||
for info in sessions:
|
||||
command = " ".join(info.command.split())
|
||||
if len(command) > 120:
|
||||
|
||||
@@ -125,6 +125,10 @@ class FileStates:
|
||||
"""Return the raw ReadState entry for a path, or None."""
|
||||
return self._state.get(str(Path(path).resolve()))
|
||||
|
||||
def raw_state(self) -> dict[str, ReadState]:
|
||||
"""Return the mutable backing map for legacy compatibility."""
|
||||
return self._state
|
||||
|
||||
def clear(self) -> None:
|
||||
"""Clear all tracked state (useful for testing)."""
|
||||
self._state.clear()
|
||||
@@ -201,5 +205,5 @@ def clear() -> None:
|
||||
# so existing imports keep working.
|
||||
def __getattr__(name: str):
|
||||
if name == "_state":
|
||||
return _default._state
|
||||
return _default.raw_state()
|
||||
raise AttributeError(name)
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""File system tools: read, write, edit, list."""
|
||||
|
||||
# pyright: reportPrivateUsage=false, reportUnusedFunction=false
|
||||
|
||||
import difflib
|
||||
import mimetypes
|
||||
import os
|
||||
@@ -8,6 +10,7 @@ from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import ToolContext
|
||||
from nanobot.agent.tools.file_state import FileStates, _hash_file, current_file_states
|
||||
from nanobot.agent.tools.path_utils import resolve_workspace_path
|
||||
from nanobot.agent.tools.schema import (
|
||||
@@ -37,7 +40,7 @@ class _FsTool(Tool):
|
||||
return FileToolsConfig
|
||||
|
||||
@classmethod
|
||||
def enabled(cls, ctx: Any) -> bool:
|
||||
def enabled(cls, ctx: ToolContext) -> bool:
|
||||
return ctx.config.file.enable
|
||||
|
||||
def __init__(
|
||||
@@ -77,7 +80,7 @@ class _FsTool(Tool):
|
||||
self._fallback_file_states = FileStates()
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
def create(cls, ctx: ToolContext) -> Tool:
|
||||
from nanobot.agent.skills import BUILTIN_SKILLS_DIR
|
||||
|
||||
agent_workspace = Path(ctx.workspace)
|
||||
@@ -408,7 +411,8 @@ class ReadFileTool(_FsTool):
|
||||
result = "\n".join(numbered)
|
||||
|
||||
if len(result) > self._MAX_CHARS:
|
||||
trimmed, chars = [], 0
|
||||
trimmed: list[str] = []
|
||||
chars = 0
|
||||
for line in numbered:
|
||||
chars += len(line) + 1
|
||||
if chars > self._MAX_CHARS:
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from loguru import logger
|
||||
from pydantic import Field
|
||||
@@ -23,6 +23,7 @@ from nanobot.bus.events import (
|
||||
RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD,
|
||||
InboundMessage,
|
||||
)
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.config.paths import get_media_dir
|
||||
from nanobot.config_base import Base
|
||||
from nanobot.providers.image_generation import (
|
||||
@@ -41,6 +42,7 @@ from nanobot.utils.artifacts import (
|
||||
from nanobot.utils.helpers import detect_image_mime
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.tools.context import ToolContext
|
||||
from nanobot.config.schema import ProviderConfig
|
||||
|
||||
|
||||
@@ -89,11 +91,11 @@ class ImageGenerationTool(Tool):
|
||||
return ImageGenerationToolConfig
|
||||
|
||||
@classmethod
|
||||
def enabled(cls, ctx: Any) -> bool:
|
||||
def enabled(cls, ctx: ToolContext) -> bool:
|
||||
return ctx.config.image_generation.enabled
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
def create(cls, ctx: ToolContext) -> Tool:
|
||||
return cls(
|
||||
workspace=ctx.workspace,
|
||||
config=ctx.config.image_generation,
|
||||
@@ -134,12 +136,14 @@ class ImageGenerationTool(Tool):
|
||||
cls = get_image_gen_provider(self.config.provider)
|
||||
if cls is None:
|
||||
return None
|
||||
kwargs = {
|
||||
"api_key": provider.api_key if provider else None,
|
||||
"api_base": provider.api_base if provider else None,
|
||||
"extra_headers": provider.extra_headers if provider else None,
|
||||
"extra_body": provider.extra_body if provider else None,
|
||||
"proxy": provider.proxy if provider else None,
|
||||
kwargs: dict[str, Any] = {
|
||||
"api_key": provider.api_key if provider and isinstance(provider.api_key, str) else None,
|
||||
"api_base": provider.api_base if provider and isinstance(provider.api_base, str) else None,
|
||||
"extra_headers": provider.extra_headers
|
||||
if provider and isinstance(provider.extra_headers, dict) else None,
|
||||
"extra_body": provider.extra_body
|
||||
if provider and isinstance(provider.extra_body, dict) else None,
|
||||
"proxy": provider.proxy if provider and isinstance(provider.proxy, str) else None,
|
||||
}
|
||||
return cls(**kwargs)
|
||||
|
||||
@@ -172,7 +176,7 @@ class ImageGenerationTool(Tool):
|
||||
return []
|
||||
return [self._resolve_reference_image(value) for value in values if value]
|
||||
|
||||
async def execute(
|
||||
async def execute( # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
self,
|
||||
prompt: str,
|
||||
reference_images: list[str] | None = None,
|
||||
@@ -238,7 +242,7 @@ async def reload_image_generation_tool(state: Any, registry: ToolRegistry) -> di
|
||||
}
|
||||
|
||||
next_tool = (
|
||||
ImageGenerationTool(
|
||||
ImageGenerationTool( # pyright: ignore[reportAbstractUsage]
|
||||
workspace=state.workspace,
|
||||
config=tool_config,
|
||||
provider_configs=provider_configs,
|
||||
@@ -271,7 +275,7 @@ async def reload_image_generation_tool(state: Any, registry: ToolRegistry) -> di
|
||||
|
||||
|
||||
async def request_image_generation_reload(
|
||||
bus: Any,
|
||||
bus: MessageBus,
|
||||
*,
|
||||
timeout: float = 5.0,
|
||||
) -> dict[str, Any]:
|
||||
@@ -298,11 +302,13 @@ async def request_image_generation_reload(
|
||||
"message": "Image generation hot reload timed out.",
|
||||
"requires_restart": True,
|
||||
}
|
||||
return result if isinstance(result, dict) else {
|
||||
"ok": False,
|
||||
"message": "Image generation hot reload returned an unexpected response.",
|
||||
"requires_restart": True,
|
||||
}
|
||||
if not isinstance(cast(object, result), dict):
|
||||
return {
|
||||
"ok": False,
|
||||
"message": "Image generation hot reload returned an unexpected response.",
|
||||
"requires_restart": True,
|
||||
}
|
||||
return result
|
||||
|
||||
|
||||
async def handle_runtime_control(
|
||||
@@ -311,7 +317,7 @@ async def handle_runtime_control(
|
||||
registry: ToolRegistry,
|
||||
) -> bool:
|
||||
"""Handle an in-process image generation reload request."""
|
||||
metadata = msg.metadata if isinstance(msg.metadata, dict) else {}
|
||||
metadata = msg.metadata
|
||||
if metadata.get(INBOUND_META_RUNTIME_CONTROL) != RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD:
|
||||
return False
|
||||
|
||||
@@ -327,5 +333,5 @@ async def handle_runtime_control(
|
||||
"error": str(exc),
|
||||
}
|
||||
if isinstance(ack, asyncio.Future) and not ack.done():
|
||||
ack.set_result(result)
|
||||
cast(asyncio.Future[Any], ack).set_result(result)
|
||||
return True
|
||||
|
||||
@@ -1,16 +1,22 @@
|
||||
"""Tool discovery and registration via package scanning."""
|
||||
|
||||
# pyright: reportIncompatibleVariableOverride=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import pkgutil
|
||||
from importlib.metadata import entry_points
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.agent.tools.base import Tool, ToolResult
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.tools.context import RequestContext, ToolContext
|
||||
|
||||
_SKIP_MODULES = frozenset({
|
||||
"base", "schema", "registry", "context", "loader", "config",
|
||||
"file_state", "sandbox", "mcp", "__init__", "runtime_state",
|
||||
@@ -83,7 +89,7 @@ class ToolLoader:
|
||||
self._plugins = plugins
|
||||
return plugins
|
||||
|
||||
def load(self, ctx: Any, registry: ToolRegistry, *, scope: str = "core") -> list[str]:
|
||||
def load(self, ctx: ToolContext, registry: ToolRegistry, *, scope: str = "core") -> list[str]:
|
||||
registered: list[str] = []
|
||||
builtin_names: set[str] = set()
|
||||
sources = [(self.discover(), False), (self._discover_plugins().values(), True)]
|
||||
@@ -157,7 +163,7 @@ class _LegacyErrorPrefixTool(Tool):
|
||||
def config_key(self) -> str:
|
||||
return getattr(self._wrapped, "config_key", "")
|
||||
|
||||
def set_context(self, ctx: Any) -> None:
|
||||
def set_context(self, ctx: RequestContext) -> None:
|
||||
set_context = getattr(self._wrapped, "set_context", None)
|
||||
if callable(set_context):
|
||||
set_context(ctx)
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Sustained-goal tools with explicit user opt-in at the execution boundary."""
|
||||
|
||||
# pyright: reportIncompatibleMethodOverride=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from copy import deepcopy
|
||||
@@ -11,7 +13,7 @@ from nanobot.agent.goal_permission import (
|
||||
revoke_goal_mutation_permission,
|
||||
)
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import RequestContext, current_request_context
|
||||
from nanobot.agent.tools.context import RequestContext, ToolContext, current_request_context
|
||||
from nanobot.agent.tools.schema import StringSchema, tool_parameters_schema
|
||||
from nanobot.bus.runtime_events import GoalStateChanged, RuntimeEventBus, RuntimeEventContext
|
||||
from nanobot.runtime_context import RuntimeContextBlock, wrap_runtime_context_lines
|
||||
@@ -132,23 +134,24 @@ class CreateGoalTool(Tool, _GoalToolsMixin):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
sessions: Any,
|
||||
sessions: SessionManager,
|
||||
runtime_events: RuntimeEventBus | None = None,
|
||||
) -> None:
|
||||
_GoalToolsMixin.__init__(self, sessions, runtime_events)
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
sess = getattr(ctx, "sessions", None)
|
||||
assert sess is not None
|
||||
def create(cls, ctx: ToolContext) -> Tool:
|
||||
sess = ctx.sessions
|
||||
if sess is None:
|
||||
raise RuntimeError("CreateGoalTool requires an initialized session manager")
|
||||
return cls(
|
||||
sessions=sess,
|
||||
runtime_events=getattr(ctx, "runtime_events", None),
|
||||
runtime_events=ctx.runtime_events,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def enabled(cls, ctx: Any) -> bool:
|
||||
return getattr(ctx, "sessions", None) is not None
|
||||
def enabled(cls, ctx: ToolContext) -> bool:
|
||||
return ctx.sessions is not None
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
@@ -262,23 +265,24 @@ class UpdateGoalTool(Tool, _GoalToolsMixin):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
sessions: Any,
|
||||
sessions: SessionManager,
|
||||
runtime_events: RuntimeEventBus | None = None,
|
||||
) -> None:
|
||||
_GoalToolsMixin.__init__(self, sessions, runtime_events)
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
sess = getattr(ctx, "sessions", None)
|
||||
assert sess is not None
|
||||
def create(cls, ctx: ToolContext) -> Tool:
|
||||
sess = ctx.sessions
|
||||
if sess is None:
|
||||
raise RuntimeError("UpdateGoalTool requires an initialized session manager")
|
||||
return cls(
|
||||
sessions=sess,
|
||||
runtime_events=getattr(ctx, "runtime_events", None),
|
||||
runtime_events=ctx.runtime_events,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def enabled(cls, ctx: Any) -> bool:
|
||||
return getattr(ctx, "sessions", None) is not None
|
||||
def enabled(cls, ctx: ToolContext) -> bool:
|
||||
return ctx.sessions is not None
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
|
||||
+99
-46
@@ -7,9 +7,9 @@ import os
|
||||
import re
|
||||
import shutil
|
||||
import urllib.parse
|
||||
from collections.abc import Awaitable, Callable
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable
|
||||
from contextlib import AsyncExitStack, suppress
|
||||
from typing import Any, Mapping, Protocol
|
||||
from typing import TYPE_CHECKING, Any, Mapping, Protocol, cast
|
||||
from weakref import WeakKeyDictionary
|
||||
|
||||
import httpx
|
||||
@@ -23,6 +23,7 @@ from nanobot.bus.events import (
|
||||
RUNTIME_CONTROL_MCP_RELOAD,
|
||||
InboundMessage,
|
||||
)
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.security.network import (
|
||||
PinnedDNSAsyncTransport,
|
||||
env_proxy_applies_to_url,
|
||||
@@ -32,6 +33,13 @@ from nanobot.security.network import (
|
||||
)
|
||||
from nanobot.utils.cancellation import task_is_cancelling
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mcp import ClientSession
|
||||
from mcp.types import Prompt, Resource
|
||||
from mcp.types import Tool as MCPToolDefinition
|
||||
|
||||
from nanobot.config.schema import MCPServerConfig
|
||||
|
||||
# Transient connection errors that warrant a single retry.
|
||||
# These typically happen when an MCP server restarts or a network
|
||||
# connection is interrupted between calls.
|
||||
@@ -92,7 +100,7 @@ def _mcp_jsonrpc_payload(message: Any) -> Any:
|
||||
|
||||
def _payload_value(payload: Any, key: str) -> Any:
|
||||
if isinstance(payload, Mapping):
|
||||
return payload.get(key)
|
||||
return cast(Mapping[str, Any], payload).get(key)
|
||||
return getattr(payload, key, None)
|
||||
|
||||
|
||||
@@ -106,7 +114,7 @@ class _MalformedProgressNotificationFilter:
|
||||
def __init__(self, read_stream: Any, server_name: str) -> None:
|
||||
self._read_stream = read_stream
|
||||
self._server_name = server_name
|
||||
self._iterator: Any | None = None
|
||||
self._iterator: AsyncIterator[Any] | None = None
|
||||
|
||||
async def __aenter__(self) -> "_MalformedProgressNotificationFilter":
|
||||
await self._read_stream.__aenter__()
|
||||
@@ -120,11 +128,13 @@ class _MalformedProgressNotificationFilter:
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> Any:
|
||||
if self._iterator is None:
|
||||
self._iterator = self._read_stream.__aiter__()
|
||||
iterator = self._iterator
|
||||
if iterator is None:
|
||||
iterator = self._read_stream.__aiter__()
|
||||
self._iterator = iterator
|
||||
|
||||
while True:
|
||||
message = await self._iterator.__anext__()
|
||||
message = await anext(iterator)
|
||||
if _is_malformed_mcp_progress_notification(message):
|
||||
logger.debug(
|
||||
"MCP server '{}': dropped progress notification without progressToken",
|
||||
@@ -241,8 +251,8 @@ def _redact_url(url: str) -> str:
|
||||
return "<redacted-url>"
|
||||
|
||||
|
||||
def _pinned_transport_kwargs() -> dict[str, object]:
|
||||
kwargs: dict[str, object] = {"transport": PinnedDNSAsyncTransport()}
|
||||
def _pinned_transport_kwargs() -> dict[str, Any]:
|
||||
kwargs: dict[str, Any] = {"transport": PinnedDNSAsyncTransport()}
|
||||
mounts = httpx_env_proxy_mounts()
|
||||
if mounts:
|
||||
kwargs["mounts"] = mounts
|
||||
@@ -302,13 +312,14 @@ def _extract_nullable_branch(options: Any) -> tuple[dict[str, Any], bool] | None
|
||||
|
||||
non_null: list[dict[str, Any]] = []
|
||||
saw_null = False
|
||||
for option in options:
|
||||
for option in cast(list[object], options):
|
||||
if not isinstance(option, dict):
|
||||
return None
|
||||
if option.get("type") == "null":
|
||||
option_schema = cast(dict[str, Any], option)
|
||||
if option_schema.get("type") == "null":
|
||||
saw_null = True
|
||||
continue
|
||||
non_null.append(option)
|
||||
non_null.append(option_schema)
|
||||
|
||||
if saw_null and len(non_null) == 1:
|
||||
return non_null[0], True
|
||||
@@ -330,9 +341,9 @@ def _resolve_local_schema_ref(root: dict[str, Any], ref: str) -> Any:
|
||||
for raw_part in pointer[1:].split("/"):
|
||||
part = raw_part.replace("~1", "/").replace("~0", "~")
|
||||
if isinstance(current, dict):
|
||||
current = current[part]
|
||||
current = cast(dict[str, Any], current)[part]
|
||||
elif isinstance(current, list):
|
||||
current = current[int(part)]
|
||||
current = cast(list[Any], current)[int(part)]
|
||||
else:
|
||||
raise KeyError(part)
|
||||
return current
|
||||
@@ -345,14 +356,15 @@ def _rewrite_local_schema_refs(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
|
||||
def rewrite(value: Any) -> Any:
|
||||
if isinstance(value, list):
|
||||
return [rewrite(item) for item in value]
|
||||
return [rewrite(item) for item in cast(list[Any], value)]
|
||||
if not isinstance(value, dict):
|
||||
return value
|
||||
|
||||
rewritten = dict(value)
|
||||
ref = rewritten.get("$ref")
|
||||
rewritten = dict(cast(dict[str, Any], value))
|
||||
raw_ref = rewritten.get("$ref")
|
||||
ref = raw_ref if isinstance(raw_ref, str) else None
|
||||
is_rewritable_ref = False
|
||||
if isinstance(ref, str) and not ref.startswith("#/$defs/"):
|
||||
if ref is not None and not ref.startswith("#/$defs/"):
|
||||
try:
|
||||
pointer = urllib.parse.unquote(ref[1:], errors="strict")
|
||||
except (UnicodeDecodeError, ValueError):
|
||||
@@ -362,6 +374,7 @@ def _rewrite_local_schema_refs(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
not pointer or pointer.startswith("/")
|
||||
)
|
||||
if is_rewritable_ref:
|
||||
assert ref is not None
|
||||
name = rewritten_refs.get(ref)
|
||||
if name is None:
|
||||
try:
|
||||
@@ -369,7 +382,6 @@ def _rewrite_local_schema_refs(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
except (KeyError, IndexError, TypeError, UnicodeDecodeError, ValueError):
|
||||
logger.warning("MCP tool schema contains an unresolved local $ref: {}", ref)
|
||||
else:
|
||||
assert isinstance(ref, str)
|
||||
name = f"ref_{hashlib.sha256(ref.encode()).hexdigest()[:12]}"
|
||||
existing_defs = schema.get("$defs")
|
||||
while isinstance(existing_defs, dict) and name in existing_defs:
|
||||
@@ -383,7 +395,7 @@ def _rewrite_local_schema_refs(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
|
||||
return {key: rewrite(item) for key, item in rewritten.items()}
|
||||
|
||||
result = rewrite(schema)
|
||||
result = cast(dict[str, Any], rewrite(schema))
|
||||
if generated_defs:
|
||||
existing_defs = result.get("$defs")
|
||||
result["$defs"] = {
|
||||
@@ -398,8 +410,9 @@ def _normalize_nullable_schema(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
normalized = dict(schema)
|
||||
raw_type = normalized.get("type")
|
||||
if isinstance(raw_type, list):
|
||||
non_null = [item for item in raw_type if item != "null"]
|
||||
if "null" in raw_type and len(non_null) == 1:
|
||||
type_values = cast(list[Any], raw_type)
|
||||
non_null = [item for item in type_values if item != "null"]
|
||||
if "null" in type_values and len(non_null) == 1:
|
||||
normalized["type"] = non_null[0]
|
||||
normalized["nullable"] = True
|
||||
|
||||
@@ -413,19 +426,28 @@ def _normalize_nullable_schema(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
normalized["nullable"] = True
|
||||
break
|
||||
|
||||
if isinstance(normalized.get("properties"), dict):
|
||||
properties = normalized.get("properties")
|
||||
if isinstance(properties, dict):
|
||||
property_schemas = cast(dict[str, Any], properties)
|
||||
normalized["properties"] = {
|
||||
name: _normalize_nullable_schema(prop) if isinstance(prop, dict) else prop
|
||||
for name, prop in normalized["properties"].items()
|
||||
name: (
|
||||
_normalize_nullable_schema(cast(dict[str, Any], prop))
|
||||
if isinstance(prop, dict)
|
||||
else prop
|
||||
)
|
||||
for name, prop in property_schemas.items()
|
||||
}
|
||||
if isinstance(normalized.get("items"), dict):
|
||||
normalized["items"] = _normalize_nullable_schema(normalized["items"])
|
||||
if isinstance(normalized.get("$defs"), dict):
|
||||
items = normalized.get("items")
|
||||
if isinstance(items, dict):
|
||||
normalized["items"] = _normalize_nullable_schema(cast(dict[str, Any], items))
|
||||
definitions = normalized.get("$defs")
|
||||
if isinstance(definitions, dict):
|
||||
definition_schemas = cast(dict[str, Any], definitions)
|
||||
normalized["$defs"] = {
|
||||
name: _normalize_nullable_schema(definition)
|
||||
name: _normalize_nullable_schema(cast(dict[str, Any], definition))
|
||||
if isinstance(definition, dict)
|
||||
else definition
|
||||
for name, definition in normalized["$defs"].items()
|
||||
for name, definition in definition_schemas.items()
|
||||
}
|
||||
|
||||
if normalized.get("type") == "object":
|
||||
@@ -438,15 +460,19 @@ def _normalize_schema_for_openai(schema: Any) -> dict[str, Any]:
|
||||
"""Normalize MCP JSON Schema patterns for tool definitions."""
|
||||
if not isinstance(schema, dict):
|
||||
return {"type": "object", "properties": {}}
|
||||
return _normalize_nullable_schema(_rewrite_local_schema_refs(schema))
|
||||
schema_mapping = cast(dict[str, Any], schema)
|
||||
return _normalize_nullable_schema(_rewrite_local_schema_refs(schema_mapping))
|
||||
|
||||
|
||||
class _MCPWrapperBase(Tool):
|
||||
"""Common reconnect handling for wrappers bound to one MCP server session."""
|
||||
|
||||
_plugin_discoverable = False
|
||||
_session: "ClientSession"
|
||||
_server_name: str
|
||||
_name: str
|
||||
|
||||
def _set_mcp_connection(self, session: Any, server_name: str) -> None:
|
||||
def _set_mcp_connection(self, session: "ClientSession", server_name: str) -> None:
|
||||
self._session = session
|
||||
self._server_name = server_name
|
||||
self._reconnect: _ReconnectCallback | None = None
|
||||
@@ -500,9 +526,10 @@ def _image_block_data_url(block: Any, types: Any) -> str | None:
|
||||
if embedded_cls is not None and isinstance(block, embedded_cls):
|
||||
resource = getattr(block, "resource", None)
|
||||
if blob_cls is not None and isinstance(resource, blob_cls):
|
||||
mime = getattr(resource, "mimeType", None) or ""
|
||||
blob_resource = cast(Any, resource)
|
||||
mime = getattr(blob_resource, "mimeType", None) or ""
|
||||
if isinstance(mime, str) and mime.startswith("image/"):
|
||||
return f"data:{mime};base64,{resource.blob}"
|
||||
return f"data:{mime};base64,{blob_resource.blob}"
|
||||
return None
|
||||
|
||||
|
||||
@@ -533,7 +560,13 @@ class MCPToolWrapper(_MCPWrapperBase):
|
||||
|
||||
_plugin_discoverable = False
|
||||
|
||||
def __init__(self, session, server_name: str, tool_def, tool_timeout: int = 30):
|
||||
def __init__(
|
||||
self,
|
||||
session: "ClientSession",
|
||||
server_name: str,
|
||||
tool_def: "MCPToolDefinition",
|
||||
tool_timeout: int = 30,
|
||||
):
|
||||
self._set_mcp_connection(session, server_name)
|
||||
self._original_name = tool_def.name
|
||||
self._name = _sanitize_mcp_tool_name(f"mcp_{server_name}_{tool_def.name}")
|
||||
@@ -689,7 +722,13 @@ class MCPResourceWrapper(_MCPWrapperBase):
|
||||
|
||||
_plugin_discoverable = False
|
||||
|
||||
def __init__(self, session, server_name: str, resource_def, resource_timeout: int = 30):
|
||||
def __init__(
|
||||
self,
|
||||
session: "ClientSession",
|
||||
server_name: str,
|
||||
resource_def: "Resource",
|
||||
resource_timeout: int = 30,
|
||||
):
|
||||
self._set_mcp_connection(session, server_name)
|
||||
self._uri = resource_def.uri
|
||||
self._name = _sanitize_mcp_tool_name(f"mcp_{server_name}_resource_{resource_def.name}")
|
||||
@@ -775,7 +814,7 @@ class MCPResourceWrapper(_MCPWrapperBase):
|
||||
for block in result.contents:
|
||||
if isinstance(block, types.TextResourceContents):
|
||||
parts.append(block.text)
|
||||
elif isinstance(block, types.BlobResourceContents):
|
||||
elif isinstance(cast(object, block), types.BlobResourceContents):
|
||||
parts.append(f"[Binary resource: {len(block.blob)} bytes]")
|
||||
else:
|
||||
parts.append(str(block))
|
||||
@@ -787,7 +826,13 @@ class MCPPromptWrapper(_MCPWrapperBase):
|
||||
|
||||
_plugin_discoverable = False
|
||||
|
||||
def __init__(self, session, server_name: str, prompt_def, prompt_timeout: int = 30):
|
||||
def __init__(
|
||||
self,
|
||||
session: "ClientSession",
|
||||
server_name: str,
|
||||
prompt_def: "Prompt",
|
||||
prompt_timeout: int = 30,
|
||||
):
|
||||
self._set_mcp_connection(session, server_name)
|
||||
self._prompt_name = prompt_def.name
|
||||
self._name = _sanitize_mcp_tool_name(f"mcp_{server_name}_prompt_{prompt_def.name}")
|
||||
@@ -916,7 +961,7 @@ class MCPPromptWrapper(_MCPWrapperBase):
|
||||
|
||||
|
||||
async def connect_mcp_servers(
|
||||
mcp_servers: dict, registry: ToolRegistry
|
||||
mcp_servers: "dict[str, MCPServerConfig]", registry: ToolRegistry
|
||||
) -> dict[str, MCPConnection]:
|
||||
"""Connect to configured MCP servers and register their tools, resources, prompts.
|
||||
|
||||
@@ -929,7 +974,9 @@ async def connect_mcp_servers(
|
||||
from mcp.client.stdio import stdio_client
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
|
||||
async def open_single_server(name: str, cfg) -> tuple[str, AsyncExitStack | None]:
|
||||
async def open_single_server(
|
||||
name: str, cfg: "MCPServerConfig"
|
||||
) -> tuple[str, AsyncExitStack | None]:
|
||||
server_stack = AsyncExitStack()
|
||||
await server_stack.__aenter__()
|
||||
|
||||
@@ -1148,7 +1195,9 @@ async def connect_mcp_servers(
|
||||
await server_stack.aclose()
|
||||
return name, None
|
||||
|
||||
async def connect_single_server(name: str, cfg) -> tuple[str, MCPConnection | None]:
|
||||
async def connect_single_server(
|
||||
name: str, cfg: "MCPServerConfig"
|
||||
) -> tuple[str, MCPConnection | None]:
|
||||
loop = asyncio.get_running_loop()
|
||||
ready: asyncio.Future[bool] = loop.create_future()
|
||||
close_requested = asyncio.Event()
|
||||
@@ -1192,7 +1241,7 @@ async def connect_mcp_servers(
|
||||
except Exception as e:
|
||||
logger.exception("MCP server '{}' connection failed: {}", name, e)
|
||||
continue
|
||||
if result is not None and result[1] is not None:
|
||||
if result[1] is not None:
|
||||
server_stacks[result[0]] = result[1]
|
||||
|
||||
return server_stacks
|
||||
@@ -1335,7 +1384,11 @@ async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
async def request_mcp_reload(bus: Any, *, timeout: float = 15.0) -> dict[str, Any]:
|
||||
async def request_mcp_reload(
|
||||
bus: MessageBus,
|
||||
*,
|
||||
timeout: float = 15.0,
|
||||
) -> dict[str, Any]:
|
||||
"""Ask the running agent loop to reconcile live MCP connections."""
|
||||
loop = asyncio.get_running_loop()
|
||||
ack: asyncio.Future[dict[str, Any]] = loop.create_future()
|
||||
@@ -1359,7 +1412,7 @@ async def request_mcp_reload(bus: Any, *, timeout: float = 15.0) -> dict[str, An
|
||||
"message": "MCP hot reload timed out. Restart nanobot to pick up changes.",
|
||||
"requires_restart": True,
|
||||
}
|
||||
return result if isinstance(result, dict) else {
|
||||
return result if isinstance(cast(object, result), dict) else {
|
||||
"ok": False,
|
||||
"message": "MCP hot reload returned an unexpected response.",
|
||||
"requires_restart": True,
|
||||
@@ -1367,7 +1420,7 @@ async def request_mcp_reload(bus: Any, *, timeout: float = 15.0) -> dict[str, An
|
||||
|
||||
|
||||
async def handle_runtime_control(state: Any, msg: InboundMessage, registry: ToolRegistry) -> bool:
|
||||
metadata = msg.metadata if isinstance(msg.metadata, dict) else {}
|
||||
metadata = msg.metadata if isinstance(cast(object, msg.metadata), dict) else {}
|
||||
control = metadata.get(INBOUND_META_RUNTIME_CONTROL)
|
||||
if control != RUNTIME_CONTROL_MCP_RELOAD:
|
||||
return False
|
||||
@@ -1384,7 +1437,7 @@ async def handle_runtime_control(state: Any, msg: InboundMessage, registry: Tool
|
||||
"error": str(exc),
|
||||
}
|
||||
if isinstance(ack, asyncio.Future) and not ack.done():
|
||||
ack.set_result(result)
|
||||
cast(asyncio.Future[dict[str, Any]], ack).set_result(result)
|
||||
return True
|
||||
|
||||
|
||||
|
||||
@@ -1,13 +1,15 @@
|
||||
"""Message tool for sending messages to users."""
|
||||
|
||||
from contextvars import ContextVar
|
||||
# pyright: reportIncompatibleMethodOverride=false
|
||||
|
||||
from contextvars import ContextVar, Token
|
||||
from pathlib import Path
|
||||
from typing import Any, Awaitable, Callable
|
||||
from typing import Any, Awaitable, Callable, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import current_request_context
|
||||
from nanobot.agent.tools.context import ToolContext, current_request_context
|
||||
from nanobot.agent.tools.path_utils import resolve_workspace_path
|
||||
from nanobot.agent.tools.schema import ArraySchema, StringSchema, tool_parameters_schema
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
@@ -73,7 +75,7 @@ class MessageTool(Tool):
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
def create(cls, ctx: ToolContext) -> Tool:
|
||||
send_callback = ctx.bus.publish_outbound if ctx.bus else None
|
||||
return cls(
|
||||
send_callback=send_callback,
|
||||
@@ -89,11 +91,11 @@ class MessageTool(Tool):
|
||||
"""Reset per-turn send tracking."""
|
||||
self._sent_in_turn = False
|
||||
|
||||
def set_suppress_delivery(self, active: bool):
|
||||
def set_suppress_delivery(self, active: bool) -> Token[bool]:
|
||||
"""Acknowledge but don't deliver tool sends (heartbeat internal check)."""
|
||||
return self._suppress_delivery_var.set(active)
|
||||
|
||||
def reset_suppress_delivery(self, token) -> None:
|
||||
def reset_suppress_delivery(self, token: Token[bool]) -> None:
|
||||
"""Restore previous delivery-suppression state."""
|
||||
self._suppress_delivery_var.reset(token)
|
||||
|
||||
@@ -148,19 +150,23 @@ class MessageTool(Tool):
|
||||
chat_id: str | None = None,
|
||||
message_id: str | None = None,
|
||||
media: list[str] | None = None,
|
||||
buttons: list[list[str]] | None = None,
|
||||
buttons: Any = None,
|
||||
**kwargs: Any,
|
||||
) -> str:
|
||||
) -> str: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
from nanobot.utils.helpers import strip_think
|
||||
|
||||
content = strip_think(content)
|
||||
|
||||
button_rows: list[list[str]] | None = None
|
||||
if buttons is not None:
|
||||
if not isinstance(buttons, list) or any(
|
||||
not isinstance(row, list) or any(not isinstance(label, str) for label in row)
|
||||
for row in buttons
|
||||
raw_buttons = cast(list[Any], buttons) if isinstance(buttons, list) else None
|
||||
if raw_buttons is None or any(
|
||||
not isinstance(row, list)
|
||||
or any(not isinstance(label, str) for label in cast(list[Any], row))
|
||||
for row in raw_buttons
|
||||
):
|
||||
return ToolResult.error("Error: buttons must be a list of list of strings")
|
||||
button_rows = cast(list[list[str]], raw_buttons)
|
||||
request_ctx = current_request_context()
|
||||
default_channel = (
|
||||
request_ctx.channel if request_ctx is not None else self._fallback_channel
|
||||
@@ -228,7 +234,7 @@ class MessageTool(Tool):
|
||||
chat_id=chat_id,
|
||||
content=content,
|
||||
media=media or [],
|
||||
buttons=buttons or [],
|
||||
buttons=button_rows or [],
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
@@ -241,7 +247,11 @@ class MessageTool(Tool):
|
||||
if channel == default_channel and chat_id == default_chat_id:
|
||||
self._sent_in_turn = True
|
||||
media_info = f" with {len(media)} attachments" if media else ""
|
||||
button_info = f" with {sum(len(row) for row in buttons)} button(s)" if buttons else ""
|
||||
button_info = (
|
||||
f" with {sum(len(row) for row in button_rows)} button(s)"
|
||||
if button_rows
|
||||
else ""
|
||||
)
|
||||
return f"Message sent to {channel}:{chat_id}{media_info}{button_info}"
|
||||
except Exception as e:
|
||||
return ToolResult.error(f"Error sending message: {str(e)}")
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from nanobot.agent.tools.base import Tool, ToolResult
|
||||
from nanobot.agent.tools.context import ContextAware, current_request_context
|
||||
@@ -77,7 +77,7 @@ class ToolRegistry:
|
||||
"""Extract a normalized tool name from either OpenAI or flat schemas."""
|
||||
fn = schema.get("function")
|
||||
if isinstance(fn, dict):
|
||||
name = fn.get("name")
|
||||
name = cast(dict[str, Any], fn).get("name")
|
||||
if isinstance(name, str):
|
||||
return name
|
||||
name = schema.get("name")
|
||||
@@ -140,7 +140,7 @@ class ToolRegistry:
|
||||
)
|
||||
)
|
||||
|
||||
cast_params = tool.cast_params(params)
|
||||
cast_params = tool.cast_params(cast(dict[str, Any], params))
|
||||
errors = tool.validate_params(cast_params)
|
||||
if errors:
|
||||
return tool, cast_params, (
|
||||
@@ -176,12 +176,15 @@ class ToolRegistry:
|
||||
|
||||
@classmethod
|
||||
def _unwrap_arguments_payload(cls, tool: Tool, params: Any) -> Any:
|
||||
if not isinstance(params, dict) or set(params) != {"arguments"}:
|
||||
if not isinstance(params, dict):
|
||||
return params
|
||||
arguments_payload = cast(dict[str, Any], params)
|
||||
if set(arguments_payload) != {"arguments"}:
|
||||
return arguments_payload
|
||||
properties = (tool.parameters or {}).get("properties", {})
|
||||
if isinstance(properties, dict) and "arguments" in properties:
|
||||
return params
|
||||
return cls._coerce_argument_value(params.get("arguments"))
|
||||
return arguments_payload
|
||||
return cls._coerce_argument_value(arguments_payload.get("arguments"))
|
||||
|
||||
async def execute(self, name: str, params: Any) -> Any:
|
||||
"""Execute a tool by name with given parameters."""
|
||||
|
||||
@@ -1,6 +1,15 @@
|
||||
"""RuntimeState protocol: agent loop state exposed to MyTool."""
|
||||
|
||||
from typing import Any, Protocol
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, Protocol
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.subagent import SubagentManager
|
||||
from nanobot.agent.tools.shell import ExecToolConfig
|
||||
from nanobot.agent.tools.web import WebToolsConfig
|
||||
from nanobot.utils.llm_runtime import LLMRuntime
|
||||
|
||||
|
||||
class RuntimeState(Protocol):
|
||||
@@ -25,7 +34,7 @@ class RuntimeState(Protocol):
|
||||
def tool_names(self) -> list[str]: ...
|
||||
|
||||
@property
|
||||
def workspace(self) -> str: ...
|
||||
def workspace(self) -> Path: ...
|
||||
|
||||
@property
|
||||
def provider_retry_mode(self) -> str: ...
|
||||
@@ -37,34 +46,31 @@ class RuntimeState(Protocol):
|
||||
def context_window_tokens(self) -> int: ...
|
||||
|
||||
@property
|
||||
def web_config(self) -> Any: ...
|
||||
def web_config(self) -> WebToolsConfig: ...
|
||||
|
||||
@property
|
||||
def exec_config(self) -> Any: ...
|
||||
def exec_config(self) -> ExecToolConfig: ...
|
||||
|
||||
@property
|
||||
def workspace_sandbox(self) -> Any: ...
|
||||
|
||||
@property
|
||||
def subagents(self) -> Any: ...
|
||||
def subagents(self) -> SubagentManager: ...
|
||||
|
||||
@property
|
||||
def _runtime_vars(self) -> dict[str, Any]: ...
|
||||
|
||||
@property
|
||||
def _last_usage(self) -> Any: ...
|
||||
def _last_usage(self) -> dict[str, int]: ...
|
||||
|
||||
def _sync_subagent_runtime_limits(self) -> None: ...
|
||||
|
||||
def set_runtime_model(self, model: str) -> Any: ...
|
||||
def set_runtime_model(self, model: str) -> LLMRuntime: ...
|
||||
|
||||
def set_runtime_context_window(self, context_window_tokens: int) -> Any: ...
|
||||
def set_runtime_context_window(self, context_window_tokens: int) -> LLMRuntime: ...
|
||||
|
||||
def set_session_model_preset(
|
||||
self,
|
||||
session_key: str,
|
||||
name: str,
|
||||
) -> Any: ...
|
||||
) -> LLMRuntime: ...
|
||||
|
||||
@property
|
||||
def model_preset(self) -> str | None: ...
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Search tools: file discovery and grep."""
|
||||
|
||||
# pyright: reportIncompatibleMethodOverride=false, reportPrivateUsage=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import fnmatch
|
||||
|
||||
+53
-30
@@ -1,10 +1,14 @@
|
||||
"""MyTool: runtime state inspection and configuration for the agent loop."""
|
||||
|
||||
# RuntimeState intentionally exposes a narrow set of AgentLoop internals to
|
||||
# this manually registered tool. Tool.execute accepts heterogeneous schemas.
|
||||
# pyright: reportPrivateUsage=false, reportIncompatibleMethodOverride=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Mapping
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, TypeGuard, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
@@ -15,6 +19,7 @@ from nanobot.config_base import Base
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.subagent import SubagentStatus
|
||||
from nanobot.agent.tools.context import ToolContext
|
||||
|
||||
|
||||
class MyToolConfig(Base):
|
||||
@@ -36,7 +41,7 @@ def _has_real_attr(obj: Any, key: str) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _is_subagent_status(value: Any) -> bool:
|
||||
def _is_subagent_status(value: object) -> TypeGuard[SubagentStatus]:
|
||||
from nanobot.agent.subagent import SubagentStatus
|
||||
|
||||
return isinstance(value, SubagentStatus)
|
||||
@@ -53,7 +58,7 @@ class MyTool(Tool):
|
||||
return MyToolConfig
|
||||
|
||||
@classmethod
|
||||
def enabled(cls, ctx: Any) -> bool:
|
||||
def enabled(cls, ctx: ToolContext) -> bool:
|
||||
return ctx.config.my.enable
|
||||
|
||||
BLOCKED = frozenset({
|
||||
@@ -205,7 +210,7 @@ class MyTool(Tool):
|
||||
|
||||
def _resolve_path(self, path: str) -> tuple[Any, str | None]:
|
||||
parts = path.split(".")
|
||||
obj = self._runtime_state
|
||||
obj: Any = self._runtime_state
|
||||
for part in parts:
|
||||
if part in self._DENIED_ATTRS or part.startswith("__"):
|
||||
return None, f"'{part}' is not accessible"
|
||||
@@ -215,8 +220,9 @@ class MyTool(Tool):
|
||||
return None, f"'{part}' is not accessible"
|
||||
try:
|
||||
if isinstance(obj, Mapping):
|
||||
if part in obj:
|
||||
obj = obj[part]
|
||||
mapping = cast(Mapping[str, Any], obj)
|
||||
if part in mapping:
|
||||
obj = mapping[part]
|
||||
else:
|
||||
return None, f"'{part}' not found in mapping"
|
||||
else:
|
||||
@@ -259,28 +265,40 @@ class MyTool(Tool):
|
||||
detail = MyTool._format_status(val, " ")
|
||||
return f"{header}\n task: {val.task_description}\n{detail}"
|
||||
# SubagentManager: delegate to its _task_statuses dict
|
||||
if hasattr(val, "_task_statuses") and isinstance(val._task_statuses, dict):
|
||||
return MyTool._format_value(val._task_statuses, key)
|
||||
if isinstance(val, Mapping) and val and _is_subagent_status(next(iter(val.values()))):
|
||||
task_statuses = getattr(val, "_task_statuses", None)
|
||||
if isinstance(task_statuses, dict):
|
||||
return MyTool._format_value(task_statuses, key)
|
||||
if isinstance(val, Mapping):
|
||||
mapping = cast(Mapping[object, object], val)
|
||||
else:
|
||||
mapping = None
|
||||
if (
|
||||
mapping
|
||||
and _is_subagent_status(next(iter(mapping.values())))
|
||||
):
|
||||
status_mapping: Mapping[object, SubagentStatus] = cast(Any, mapping)
|
||||
prefix = f"{key}: " if key else ""
|
||||
lines = [f"{prefix}{len(val)} subagent(s):"]
|
||||
for tid, st in val.items():
|
||||
lines = [f"{prefix}{len(status_mapping)} subagent(s):"]
|
||||
for tid, st in status_mapping.items():
|
||||
detail = MyTool._format_status(st, " ")
|
||||
lines.append(f" [{tid}] '{st.label}'\n{detail}")
|
||||
return "\n".join(lines)
|
||||
if hasattr(val, "tool_names"):
|
||||
return f"tools: {len(val.tool_names)} registered — {val.tool_names}"
|
||||
dynamic_value = cast(Any, val)
|
||||
if hasattr(dynamic_value, "tool_names"):
|
||||
tool_names: Any = getattr(dynamic_value, "tool_names")
|
||||
return f"tools: {len(tool_names)} registered — {tool_names}"
|
||||
# Scalar types — repr is fine
|
||||
if isinstance(val, (str, int, float, bool, type(None))):
|
||||
r = repr(val)
|
||||
return f"{key}: {r}" if key else r
|
||||
# Mapping — small: show content; large: show keys for dot-path navigation
|
||||
if isinstance(val, Mapping):
|
||||
ks = list(val.keys())
|
||||
value_mapping = cast(Mapping[object, object], val)
|
||||
ks = list(value_mapping.keys())
|
||||
if not ks:
|
||||
return f"{key}: {{}}" if key else "{}"
|
||||
if len(ks) <= 5:
|
||||
r = repr(val)
|
||||
r = repr(value_mapping)
|
||||
if len(r) <= 200:
|
||||
return f"{key}: {r}" if key else r
|
||||
preview = ", ".join(str(k) for k in ks[:15])
|
||||
@@ -288,18 +306,20 @@ class MyTool(Tool):
|
||||
return f"{key}: {{{preview}{suffix}}}" if key else f"{{{preview}{suffix}}}"
|
||||
# List/tuple — count for large, repr for small
|
||||
if isinstance(val, (list, tuple)):
|
||||
if len(val) > 20:
|
||||
return f"{key}: [{len(val)} items]" if key else f"[{len(val)} items]"
|
||||
r = repr(val)
|
||||
sequence = cast(list[object] | tuple[object, ...], val)
|
||||
if len(sequence) > 20:
|
||||
return f"{key}: [{len(sequence)} items]" if key else f"[{len(sequence)} items]"
|
||||
r = repr(sequence)
|
||||
return f"{key}: {r}" if key else r
|
||||
# Complex object — small Pydantic models: show values; others: show field names for navigation
|
||||
cls_name = type(val).__name__
|
||||
model_fields = getattr(type(val), "model_fields", None)
|
||||
if model_fields:
|
||||
fields = list(model_fields.keys())
|
||||
value_type = type(cast(object, val))
|
||||
cls_name = value_type.__name__
|
||||
model_fields = cast(object, getattr(value_type, "model_fields", None))
|
||||
if isinstance(model_fields, Mapping) and model_fields:
|
||||
fields = list(cast(Mapping[str, object], model_fields).keys())
|
||||
if len(fields) <= 8:
|
||||
# Small config objects: show field=value pairs
|
||||
pairs = []
|
||||
pairs: list[str] = []
|
||||
for f in fields:
|
||||
fv = getattr(val, f, "?")
|
||||
if MyTool._is_sensitive_field_name(f):
|
||||
@@ -311,7 +331,8 @@ class MyTool(Tool):
|
||||
preview = ", ".join(pairs)
|
||||
return f"{key}: {preview}" if key else preview
|
||||
else:
|
||||
fields = [a for a in getattr(val, "__dict__", {}) if not a.startswith("__")]
|
||||
attributes = cast(dict[str, Any], getattr(val, "__dict__", {}))
|
||||
fields = [name for name in attributes if not name.startswith("__")]
|
||||
if fields:
|
||||
preview = ", ".join(str(f) for f in fields[:20])
|
||||
suffix = ", ..." if len(fields) > 20 else ""
|
||||
@@ -417,6 +438,7 @@ class MyTool(Tool):
|
||||
def _modify(self, key: str | None, value: Any) -> str:
|
||||
if err := self._validate_key(key):
|
||||
return err
|
||||
key = cast(str, key)
|
||||
top = key.split(".")[0]
|
||||
if top in self.BLOCKED or top in self._DENIED_ATTRS or top.startswith("__") or top.lower() in self._SENSITIVE_NAMES:
|
||||
self._audit("modify", f"BLOCKED {key}")
|
||||
@@ -478,7 +500,7 @@ class MyTool(Tool):
|
||||
|
||||
def _modify_restricted(self, key: str, value: Any) -> str:
|
||||
spec = self.RESTRICTED[key]
|
||||
expected = spec["type"]
|
||||
expected = cast(type[Any], spec["type"])
|
||||
if expected is int and isinstance(value, bool):
|
||||
return ToolResult.error(f"Error: '{key}' must be {expected.__name__}, got bool")
|
||||
if not isinstance(value, expected):
|
||||
@@ -499,9 +521,9 @@ class MyTool(Tool):
|
||||
"during an active session; use a configured model_preset"
|
||||
)
|
||||
if key == "model":
|
||||
self._runtime_state.set_runtime_model(value)
|
||||
self._runtime_state.set_runtime_model(cast(str, value))
|
||||
elif key == "context_window_tokens":
|
||||
self._runtime_state.set_runtime_context_window(value)
|
||||
self._runtime_state.set_runtime_context_window(cast(int, value))
|
||||
else:
|
||||
setattr(self._runtime_state, key, value)
|
||||
if key == "max_iterations" and hasattr(
|
||||
@@ -516,7 +538,8 @@ class MyTool(Tool):
|
||||
if _has_real_attr(self._runtime_state, key):
|
||||
old = getattr(self._runtime_state, key)
|
||||
if isinstance(old, (str, int, float, bool)):
|
||||
old_t, new_t = type(old), type(value)
|
||||
old_t: type[Any] = type(old)
|
||||
new_t = cast(type[Any], type(value))
|
||||
if old_t is float and new_t is int:
|
||||
pass # int → float coercion allowed
|
||||
elif old_t is not new_t:
|
||||
@@ -555,12 +578,12 @@ class MyTool(Tool):
|
||||
if isinstance(value, (str, int, float, bool, type(None))):
|
||||
return None
|
||||
if isinstance(value, list):
|
||||
for i, item in enumerate(value):
|
||||
for i, item in enumerate(cast(list[Any], value)):
|
||||
if err := cls._validate_json_safe(item, depth + 1):
|
||||
return f"list[{i}] contains {err}"
|
||||
return None
|
||||
if isinstance(value, dict):
|
||||
for k, v in value.items():
|
||||
for k, v in cast(dict[Any, Any], value).items():
|
||||
if not isinstance(k, str):
|
||||
return f"dict key must be str, got {type(k).__name__}"
|
||||
if err := cls._validate_json_safe(v, depth + 1):
|
||||
|
||||
@@ -18,13 +18,14 @@ from loguru import logger
|
||||
from pydantic import Field
|
||||
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import current_request_session_key
|
||||
from nanobot.agent.tools.context import ToolContext, current_request_session_key
|
||||
from nanobot.agent.tools.exec_session import (
|
||||
DEFAULT_EXEC_SESSION_MANAGER,
|
||||
DEFAULT_MAX_OUTPUT_CHARS,
|
||||
DEFAULT_YIELD_MS,
|
||||
MAX_OUTPUT_CHARS,
|
||||
MAX_YIELD_MS,
|
||||
ExecSessionManager,
|
||||
clamp_session_int,
|
||||
format_session_poll,
|
||||
)
|
||||
@@ -174,11 +175,11 @@ class ExecTool(Tool):
|
||||
return ExecToolConfig
|
||||
|
||||
@classmethod
|
||||
def enabled(cls, ctx: Any) -> bool:
|
||||
def enabled(cls, ctx: ToolContext) -> bool:
|
||||
return ctx.config.exec.enable
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
def create(cls, ctx: ToolContext) -> Tool:
|
||||
cfg = ctx.config.exec
|
||||
return cls(
|
||||
working_dir=ctx.workspace,
|
||||
@@ -193,7 +194,7 @@ class ExecTool(Tool):
|
||||
allowed_env_keys=cfg.allowed_env_keys,
|
||||
allow_patterns=cfg.allow_patterns,
|
||||
deny_patterns=cfg.deny_patterns,
|
||||
session_manager=getattr(ctx, "exec_session_manager", None),
|
||||
session_manager=ctx.exec_session_manager,
|
||||
)
|
||||
|
||||
def __init__(
|
||||
@@ -211,7 +212,7 @@ class ExecTool(Tool):
|
||||
sandbox_ro_binds: list[str] | None = None,
|
||||
sandbox_rw_binds: list[str] | None = None,
|
||||
allowed_env_keys: list[str] | None = None,
|
||||
session_manager: Any | None = None,
|
||||
session_manager: ExecSessionManager | None = None,
|
||||
):
|
||||
self.timeout = timeout
|
||||
self.working_dir = working_dir
|
||||
@@ -344,7 +345,7 @@ class ExecTool(Tool):
|
||||
# misses it, leaving a zombie.
|
||||
_reap_pid(process.pid)
|
||||
|
||||
output_parts = []
|
||||
output_parts: list[str] = []
|
||||
|
||||
if stdout:
|
||||
output_parts.append(stdout.decode("utf-8", errors="replace"))
|
||||
@@ -504,7 +505,7 @@ class ExecTool(Tool):
|
||||
)
|
||||
|
||||
def _compose_path(self, current_path: str) -> str:
|
||||
parts = []
|
||||
parts: list[str] = []
|
||||
if self.path_prepend:
|
||||
parts.append(self.path_prepend)
|
||||
if current_path:
|
||||
@@ -514,7 +515,7 @@ class ExecTool(Tool):
|
||||
return os.pathsep.join(parts)
|
||||
|
||||
def _wrap_path_export(self, command: str, env: dict[str, str]) -> str:
|
||||
segments = []
|
||||
segments: list[str] = []
|
||||
if self.path_prepend:
|
||||
env["NANOBOT_PATH_PREPEND"] = self.path_prepend
|
||||
segments.append("$NANOBOT_PATH_PREPEND")
|
||||
@@ -568,11 +569,21 @@ class ExecTool(Tool):
|
||||
env=env,
|
||||
)
|
||||
shell_program = shell_program or shutil.which("bash") or "/bin/bash"
|
||||
args = [shell_program]
|
||||
args: list[str] = [shell_program]
|
||||
shell_name = Path(shell_program).name.lower()
|
||||
if login and shell_name in {"bash", "bash.exe", "zsh", "zsh.exe"}:
|
||||
args.append("-l")
|
||||
args.extend(["-c", command])
|
||||
if process_tree:
|
||||
return await asyncio.create_subprocess_exec(
|
||||
*args,
|
||||
stdin=stdin,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
cwd=cwd,
|
||||
env=env,
|
||||
start_new_session=True,
|
||||
)
|
||||
return await asyncio.create_subprocess_exec(
|
||||
*args,
|
||||
stdin=stdin,
|
||||
@@ -580,7 +591,6 @@ class ExecTool(Tool):
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
cwd=cwd,
|
||||
env=env,
|
||||
**({"start_new_session": True} if process_tree else {}),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Spawn tool for creating background subagents."""
|
||||
|
||||
# pyright: reportIncompatibleMethodOverride=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
@@ -16,6 +18,7 @@ from nanobot.security.workspace_access import current_workspace_scope
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.subagent import SubagentManager
|
||||
from nanobot.agent.tools.context import ToolContext
|
||||
|
||||
|
||||
@tool_parameters(
|
||||
@@ -49,8 +52,11 @@ class SpawnTool(Tool):
|
||||
self._manager = manager
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
return cls(manager=ctx.subagent_manager)
|
||||
def create(cls, ctx: ToolContext) -> Tool:
|
||||
manager = ctx.subagent_manager
|
||||
if manager is None:
|
||||
raise RuntimeError("SpawnTool requires an initialized subagent manager")
|
||||
return cls(manager=manager)
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
|
||||
+100
-55
@@ -1,5 +1,7 @@
|
||||
"""Web tools: web_search and web_fetch."""
|
||||
|
||||
# pyright: reportIncompatibleMethodOverride=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
@@ -7,7 +9,8 @@ import html
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from typing import Any, Callable
|
||||
from collections.abc import Callable
|
||||
from typing import Any, cast
|
||||
from urllib.parse import quote, urljoin, urlparse
|
||||
|
||||
import httpx
|
||||
@@ -15,6 +18,7 @@ from loguru import logger
|
||||
from pydantic import Field
|
||||
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import ToolContext
|
||||
from nanobot.agent.tools.schema import (
|
||||
BooleanSchema,
|
||||
IntegerSchema,
|
||||
@@ -291,8 +295,8 @@ class WebSearchTool(Tool):
|
||||
"""Search the web using configured provider."""
|
||||
_scopes = {"core", "subagent"}
|
||||
|
||||
name = "web_search"
|
||||
description = (
|
||||
name = "web_search" # pyright: ignore[reportIncompatibleMethodOverride, reportAssignmentType]
|
||||
description = ( # pyright: ignore[reportIncompatibleMethodOverride, reportAssignmentType]
|
||||
"Search the web. Returns titles, URLs, and snippets. "
|
||||
"count defaults to 5 (max 10). "
|
||||
"Some providers support timeRange, authLevel, and queryRewrite. "
|
||||
@@ -302,20 +306,21 @@ class WebSearchTool(Tool):
|
||||
config_key = "web"
|
||||
|
||||
@classmethod
|
||||
def config_cls(cls):
|
||||
def config_cls(cls) -> type[WebToolsConfig]:
|
||||
return WebToolsConfig
|
||||
|
||||
@classmethod
|
||||
def enabled(cls, ctx: Any) -> bool:
|
||||
def enabled(cls, ctx: ToolContext) -> bool:
|
||||
return ctx.config.web.enable
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
config_loader = None
|
||||
def create(cls, ctx: ToolContext) -> Tool:
|
||||
config_loader: Callable[[], WebSearchConfig] | None = None
|
||||
if ctx.provider_snapshot_loader is not None:
|
||||
def config_loader():
|
||||
def _load_search_config() -> WebSearchConfig:
|
||||
from nanobot.config.loader import load_config, resolve_config_env_vars
|
||||
return resolve_config_env_vars(load_config()).tools.web.search
|
||||
config_loader = _load_search_config
|
||||
return cls(
|
||||
config=ctx.config.web.search,
|
||||
proxy=ctx.config.web.proxy,
|
||||
@@ -404,7 +409,7 @@ class WebSearchTool(Tool):
|
||||
auth_level: int | None = None,
|
||||
query_rewrite: bool | None = None,
|
||||
**kwargs: Any,
|
||||
) -> str:
|
||||
) -> str: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
self._refresh_config()
|
||||
provider = self.config.provider.strip().lower() or "brave"
|
||||
n = min(max(count or self.config.max_results, 1), 10)
|
||||
@@ -448,15 +453,20 @@ class WebSearchTool(Tool):
|
||||
|
||||
async def _search_olostep(self, query: str, n: int) -> str:
|
||||
try:
|
||||
from olostep import AsyncOlostep, Olostep_BaseError
|
||||
from olostep import ( # pyright: ignore[reportMissingImports]
|
||||
AsyncOlostep, # pyright: ignore[reportUnknownVariableType]
|
||||
Olostep_BaseError, # pyright: ignore[reportUnknownVariableType]
|
||||
)
|
||||
except ImportError:
|
||||
return ToolResult.error("Error: olostep package not installed. Run: pip install olostep")
|
||||
async_olostep = cast(Any, AsyncOlostep)
|
||||
olostep_base_error = cast(type[Exception], Olostep_BaseError)
|
||||
api_key = self.config.api_key or os.environ.get("OLOSTEP_API_KEY", "")
|
||||
if not api_key:
|
||||
logger.warning("OLOSTEP_API_KEY not set, falling back to DuckDuckGo")
|
||||
return await self._search_duckduckgo(query, n)
|
||||
try:
|
||||
async with AsyncOlostep(api_key=api_key) as client:
|
||||
async with async_olostep(api_key=api_key) as client:
|
||||
if self.proxy:
|
||||
transport = getattr(client, "_transport", None)
|
||||
http_client = getattr(transport, "_client", None)
|
||||
@@ -472,14 +482,16 @@ class WebSearchTool(Tool):
|
||||
),
|
||||
http2=True,
|
||||
)
|
||||
result = await client.answers.create(task=query)
|
||||
result: Any = await client.answers.create(task=query)
|
||||
|
||||
sources = getattr(result, "sources", None) or []
|
||||
source_lines = []
|
||||
for i, source in enumerate(sources[:n], 1):
|
||||
sources = cast(list[Any], getattr(result, "sources", None) or [])
|
||||
source_lines: list[str] = []
|
||||
for i, source_value in enumerate(sources[:n], 1):
|
||||
source: Any = source_value
|
||||
if isinstance(source, dict):
|
||||
title = source.get("title", "")
|
||||
url = source.get("url", "")
|
||||
source_dict = cast(dict[str, Any], source)
|
||||
title = source_dict.get("title", "")
|
||||
url = source_dict.get("url", "")
|
||||
else:
|
||||
title = getattr(source, "title", "")
|
||||
url = getattr(source, "url", "")
|
||||
@@ -493,7 +505,7 @@ class WebSearchTool(Tool):
|
||||
answer_text = getattr(result, "answer", "") or ""
|
||||
items = [{"title": answer_text or "Olostep answer", "url": "", "content": "\n".join(source_lines)}]
|
||||
return _format_results(query, items, n)
|
||||
except Olostep_BaseError as e:
|
||||
except olostep_base_error as e:
|
||||
return ToolResult.error(f"Error: Olostep search error: {type(e).__name__}: {e}")
|
||||
except Exception as e:
|
||||
return ToolResult.error(f"Error: Olostep search error: {type(e).__name__}: {e}")
|
||||
@@ -510,6 +522,7 @@ class WebSearchTool(Tool):
|
||||
"User-Agent": self.user_agent,
|
||||
}
|
||||
async with httpx.AsyncClient(proxy=self.proxy) as client:
|
||||
r: httpx.Response | None = None
|
||||
for attempt in range(2):
|
||||
r = await client.get(
|
||||
"https://api.search.brave.com/res/v1/web/search",
|
||||
@@ -522,6 +535,7 @@ class WebSearchTool(Tool):
|
||||
if attempt == 0:
|
||||
logger.warning("Brave search rate limited; retrying once in 1.0s")
|
||||
await asyncio.sleep(1.0)
|
||||
assert r is not None
|
||||
r.raise_for_status()
|
||||
items = [
|
||||
{"title": x.get("title", ""), "url": x.get("url", ""), "content": x.get("description", "")}
|
||||
@@ -691,13 +705,19 @@ class WebSearchTool(Tool):
|
||||
timeout=float(self.config.timeout),
|
||||
)
|
||||
r.raise_for_status()
|
||||
items = []
|
||||
for result in r.json().get("results", []):
|
||||
if not isinstance(result, dict):
|
||||
data = cast(dict[str, Any], r.json())
|
||||
items: list[dict[str, Any]] = []
|
||||
for result_value in cast(list[object], data.get("results", [])):
|
||||
if not isinstance(result_value, dict):
|
||||
continue
|
||||
highlights = result.get("highlights") or []
|
||||
result = cast(dict[str, Any], result_value)
|
||||
highlights: Any = result.get("highlights") or []
|
||||
if isinstance(highlights, list):
|
||||
content = "\n".join(str(highlight) for highlight in highlights if highlight)
|
||||
content = "\n".join(
|
||||
str(highlight)
|
||||
for highlight in cast(list[object], highlights)
|
||||
if highlight
|
||||
)
|
||||
else:
|
||||
content = str(highlights)
|
||||
if not content:
|
||||
@@ -737,14 +757,17 @@ class WebSearchTool(Tool):
|
||||
timeout=float(self.config.timeout),
|
||||
)
|
||||
r.raise_for_status()
|
||||
items = [
|
||||
data = cast(dict[str, Any], r.json())
|
||||
organic = cast(list[object], data.get("organic", []))
|
||||
items: list[dict[str, Any]] = [
|
||||
{
|
||||
"title": result.get("title", ""),
|
||||
"url": result.get("link", ""),
|
||||
"content": result.get("snippet", ""),
|
||||
}
|
||||
for result in r.json().get("organic", [])
|
||||
if isinstance(result, dict)
|
||||
for result_value in organic
|
||||
if isinstance(result_value, dict)
|
||||
for result in (cast(dict[str, Any], result_value),)
|
||||
]
|
||||
return _format_results(query, items, n)
|
||||
except httpx.HTTPStatusError as e:
|
||||
@@ -806,7 +829,7 @@ class WebSearchTool(Tool):
|
||||
timeout=float(self.config.timeout),
|
||||
)
|
||||
r.raise_for_status()
|
||||
data = r.json()
|
||||
data = cast(dict[str, Any], r.json())
|
||||
except httpx.HTTPStatusError as e:
|
||||
if e.response.status_code == 429:
|
||||
return ToolResult.error("Error: Volcengine search rate limited. Try again later or reduce search frequency.")
|
||||
@@ -814,20 +837,36 @@ class WebSearchTool(Tool):
|
||||
except Exception as e:
|
||||
return ToolResult.error(f"Error: Volcengine search failed: {e}")
|
||||
|
||||
error = (data.get("ResponseMetadata") or {}).get("Error") or data.get("Error") or data.get("error")
|
||||
response_metadata = cast(
|
||||
dict[str, Any],
|
||||
data.get("ResponseMetadata") or {},
|
||||
)
|
||||
error = (
|
||||
response_metadata.get("Error")
|
||||
or data.get("Error")
|
||||
or data.get("error")
|
||||
)
|
||||
if error:
|
||||
if isinstance(error, dict):
|
||||
error = cast(dict[str, Any], error)
|
||||
code = error.get("Code") or error.get("code") or "unknown"
|
||||
message = error.get("Message") or error.get("message") or error
|
||||
return ToolResult.error(f"Error: Volcengine search error {code}: {message}")
|
||||
return ToolResult.error(f"Error: Volcengine search error: {error}")
|
||||
|
||||
result = data.get("Result") or data
|
||||
web_results = result.get("WebResults") or result.get("webResults") or result.get("results") or []
|
||||
result = cast(dict[str, Any], data.get("Result") or data)
|
||||
web_results = cast(
|
||||
list[object],
|
||||
result.get("WebResults")
|
||||
or result.get("webResults")
|
||||
or result.get("results")
|
||||
or [],
|
||||
)
|
||||
items: list[dict[str, Any]] = []
|
||||
for item in web_results:
|
||||
if not isinstance(item, dict):
|
||||
for item_value in web_results:
|
||||
if not isinstance(item_value, dict):
|
||||
continue
|
||||
item = cast(dict[str, Any], item_value)
|
||||
meta_parts = [
|
||||
str(part)
|
||||
for part in (
|
||||
@@ -837,7 +876,7 @@ class WebSearchTool(Tool):
|
||||
)
|
||||
if part
|
||||
]
|
||||
summary = (
|
||||
summary = cast(str, (
|
||||
item.get("Summary")
|
||||
or item.get("summary")
|
||||
or item.get("Snippet")
|
||||
@@ -845,7 +884,7 @@ class WebSearchTool(Tool):
|
||||
or item.get("Content")
|
||||
or item.get("content")
|
||||
or ""
|
||||
)
|
||||
))
|
||||
content = "\n".join(part for part in (" | ".join(meta_parts), summary) if part)
|
||||
items.append(
|
||||
{
|
||||
@@ -861,18 +900,20 @@ class WebSearchTool(Tool):
|
||||
try:
|
||||
# Note: duckduckgo_search is synchronous and does its own requests
|
||||
# We run it in a thread to avoid blocking the loop
|
||||
from ddgs import DDGS
|
||||
from ddgs import DDGS # pyright: ignore[reportUnknownVariableType]
|
||||
|
||||
ddgs = DDGS(timeout=10, proxy=self.proxy)
|
||||
ddgs_type = cast(Any, DDGS)
|
||||
ddgs = ddgs_type(timeout=10, proxy=self.proxy)
|
||||
raw = await asyncio.wait_for(
|
||||
asyncio.to_thread(ddgs.text, query, max_results=n),
|
||||
timeout=self.config.timeout,
|
||||
)
|
||||
if not raw:
|
||||
return f"No results for: {query}"
|
||||
items = [
|
||||
raw_items = cast(list[dict[str, Any]], raw)
|
||||
items: list[dict[str, Any]] = [
|
||||
{"title": r.get("title", ""), "url": r.get("href", ""), "content": r.get("body", "")}
|
||||
for r in raw
|
||||
for r in raw_items
|
||||
]
|
||||
return _format_results(query, items, n)
|
||||
except Exception as e:
|
||||
@@ -907,15 +948,19 @@ class WebSearchTool(Tool):
|
||||
if r.status_code == 429:
|
||||
return ToolResult.error("Error: Bocha search rate-limited (HTTP 429). Wait and retry.")
|
||||
r.raise_for_status()
|
||||
data = r.json()
|
||||
wrapped_data = data.get("data") if isinstance(data, dict) else None
|
||||
result_data = wrapped_data if isinstance(wrapped_data, dict) else data
|
||||
web_pages = (
|
||||
result_data.get("webPages", {}).get("value", [])
|
||||
if isinstance(result_data, dict)
|
||||
else []
|
||||
data = cast(dict[str, Any], r.json())
|
||||
wrapped_data = data.get("data")
|
||||
result_data = (
|
||||
cast(dict[str, Any], wrapped_data)
|
||||
if isinstance(wrapped_data, dict)
|
||||
else data
|
||||
)
|
||||
items = [
|
||||
web_pages_data = cast(
|
||||
dict[str, Any],
|
||||
result_data.get("webPages", {}),
|
||||
)
|
||||
web_pages = cast(list[dict[str, Any]], web_pages_data.get("value", []))
|
||||
items: list[dict[str, Any]] = [
|
||||
{
|
||||
"title": x.get("name", ""),
|
||||
"url": x.get("url", ""),
|
||||
@@ -946,8 +991,8 @@ class WebFetchTool(Tool):
|
||||
"""Fetch and extract content from a URL."""
|
||||
_scopes = {"core", "subagent"}
|
||||
|
||||
name = "web_fetch"
|
||||
description = (
|
||||
name = "web_fetch" # pyright: ignore[reportIncompatibleMethodOverride, reportAssignmentType]
|
||||
description = ( # pyright: ignore[reportIncompatibleMethodOverride, reportAssignmentType]
|
||||
"Fetch a URL and extract readable content (HTML → markdown/text). "
|
||||
"Output is capped at maxChars (default 50 000). "
|
||||
"Works for most web pages and docs; may fail on login-walled or JS-heavy sites."
|
||||
@@ -956,15 +1001,15 @@ class WebFetchTool(Tool):
|
||||
config_key = "web"
|
||||
|
||||
@classmethod
|
||||
def config_cls(cls):
|
||||
def config_cls(cls) -> type[WebToolsConfig]:
|
||||
return WebToolsConfig
|
||||
|
||||
@classmethod
|
||||
def enabled(cls, ctx: Any) -> bool:
|
||||
def enabled(cls, ctx: ToolContext) -> bool:
|
||||
return ctx.config.web.enable
|
||||
|
||||
@classmethod
|
||||
def create(cls, ctx: Any) -> Tool:
|
||||
def create(cls, ctx: ToolContext) -> Tool:
|
||||
return cls(
|
||||
config=ctx.config.web.fetch,
|
||||
proxy=ctx.config.web.proxy,
|
||||
@@ -987,10 +1032,10 @@ class WebFetchTool(Tool):
|
||||
extract_mode: str = "markdown",
|
||||
max_chars: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
) -> Any: # pyright: ignore[reportIncompatibleMethodOverride]
|
||||
url = url.strip(" \t\r\n`\"'")
|
||||
extract_mode = kwargs.pop("extractMode", extract_mode)
|
||||
max_chars = kwargs.pop("maxChars", max_chars) or self.max_chars
|
||||
max_chars = cast(int, kwargs.pop("maxChars", max_chars) or self.max_chars)
|
||||
is_valid, error_msg = _validate_url_safe(url)
|
||||
if not is_valid:
|
||||
return json.dumps({"error": f"URL validation failed: {error_msg}", "url": url}, ensure_ascii=False)
|
||||
@@ -1119,10 +1164,10 @@ class WebFetchTool(Tool):
|
||||
return json.dumps({"error": str(e), "url": url}, ensure_ascii=False)
|
||||
|
||||
def _extract_readable_html(self, html_content: str, extract_mode: str) -> str:
|
||||
from readability import Document
|
||||
from readability import Document # pyright: ignore[reportMissingTypeStubs]
|
||||
|
||||
doc = Document(html_content)
|
||||
summary = doc.summary()
|
||||
summary = cast(str, doc.summary())
|
||||
content = self._to_markdown(summary) if extract_mode == "markdown" else _strip_tags(summary)
|
||||
return f"# {doc.title()}\n\n{content}" if doc.title() else content
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ import dataclasses
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
from nanobot.bus.outbound_events import (
|
||||
@@ -20,6 +20,9 @@ from nanobot.bus.progress import build_bus_progress_callback
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.bus.runtime_events import RuntimeEventBus, RuntimeEventPublisher
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.utils.llm_runtime import LLMRuntime
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TurnRoute:
|
||||
@@ -62,7 +65,7 @@ class TurnDeliveryFactory:
|
||||
route = self._default_route(msg, session_key)
|
||||
if self.route_policy is not None:
|
||||
route = self.route_policy(msg, session_key, route)
|
||||
if not isinstance(route, TurnRoute):
|
||||
if not isinstance(cast(object, route), TurnRoute):
|
||||
raise TypeError("turn route policy must return TurnRoute")
|
||||
return TurnDelivery(
|
||||
bus=self.bus,
|
||||
@@ -186,7 +189,7 @@ class TurnDelivery:
|
||||
started_at=started_at,
|
||||
)
|
||||
|
||||
def record_runtime(self, runtime: Any) -> None:
|
||||
def record_runtime(self, runtime: LLMRuntime) -> None:
|
||||
self.runtime_event_publisher.record_turn_runtime(self.session_key, runtime)
|
||||
|
||||
def record_latency(self, latency_ms: int | None) -> None:
|
||||
|
||||
@@ -35,7 +35,7 @@ def api_runtime_paths(config_path: Path) -> ProcessRuntimePaths:
|
||||
)
|
||||
|
||||
|
||||
class ApiRuntime(ManagedProcessRuntime):
|
||||
class ApiRuntime(ManagedProcessRuntime[ApiStartOptions]):
|
||||
"""Manage a WebUI-controlled OpenAI-compatible API process."""
|
||||
|
||||
service_name = "api"
|
||||
|
||||
+64
-18
@@ -12,7 +12,7 @@ import hmac
|
||||
import json as _json
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable, cast
|
||||
|
||||
from aiohttp import web
|
||||
from loguru import logger
|
||||
@@ -30,6 +30,9 @@ from nanobot.utils.media_decode import (
|
||||
)
|
||||
from nanobot.utils.runtime import EMPTY_FINAL_RESPONSE_MESSAGE
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
|
||||
__all__ = (
|
||||
"MAX_FILE_SIZE",
|
||||
"_FileSizeExceeded",
|
||||
@@ -44,7 +47,7 @@ API_CHAT_ID = "default"
|
||||
_AGENT_LOOP_KEY = web.AppKey[Any]("agent_loop")
|
||||
_MODEL_NAME_KEY = web.AppKey[str]("model_name")
|
||||
_REQUEST_TIMEOUT_KEY = web.AppKey[float]("request_timeout")
|
||||
_SESSION_LOCKS_KEY = web.AppKey[dict]("session_locks")
|
||||
_SESSION_LOCKS_KEY = web.AppKey[dict[str, asyncio.Lock]]("session_locks")
|
||||
_MISSING = object()
|
||||
|
||||
|
||||
@@ -111,6 +114,26 @@ def _response_text(value: Any) -> str:
|
||||
return str(getattr(value, "content") or "")
|
||||
return str(value)
|
||||
|
||||
|
||||
def _as_str(value: object) -> str:
|
||||
"""Return *value* when it is text, otherwise an empty string."""
|
||||
return value if isinstance(value, str) else ""
|
||||
|
||||
|
||||
def _require_json_object(value: object, field: str) -> dict[str, Any]:
|
||||
"""Validate an object-valued field from an untrusted JSON request."""
|
||||
if not isinstance(value, dict):
|
||||
raise TypeError(f"{field} must be an object")
|
||||
return cast(dict[str, Any], value)
|
||||
|
||||
|
||||
def _require_json_string(value: object, field: str) -> str:
|
||||
"""Validate a string-valued field from an untrusted JSON request."""
|
||||
if not isinstance(value, str):
|
||||
raise TypeError(f"{field} must be a string")
|
||||
return value
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SSE helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -141,13 +164,19 @@ _SSE_DONE = b"data: [DONE]\n\n"
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _parse_json_content(body: dict) -> tuple[str, list[str]]:
|
||||
def _parse_json_content(body: dict[str, Any]) -> tuple[str, list[str]]:
|
||||
"""Parse JSON request body. Returns (text, media_paths)."""
|
||||
messages = body.get("messages")
|
||||
if not isinstance(messages, list) or len(messages) != 1:
|
||||
messages_value = cast(object, body.get("messages"))
|
||||
if not isinstance(messages_value, list):
|
||||
raise ValueError("Only a single user message is supported")
|
||||
message = messages[0]
|
||||
if not isinstance(message, dict) or message.get("role") != "user":
|
||||
messages = cast(list[object], messages_value)
|
||||
if len(messages) != 1:
|
||||
raise ValueError("Only a single user message is supported")
|
||||
message_value: object = messages[0]
|
||||
if not isinstance(message_value, dict):
|
||||
raise ValueError("Only a single user message is supported")
|
||||
message = cast(dict[str, Any], message_value)
|
||||
if message.get("role") != "user":
|
||||
raise ValueError("Only a single user message is supported")
|
||||
|
||||
user_content = message.get("content", "")
|
||||
@@ -156,13 +185,26 @@ def _parse_json_content(body: dict) -> tuple[str, list[str]]:
|
||||
|
||||
if isinstance(user_content, list):
|
||||
text_parts: list[str] = []
|
||||
for part in user_content:
|
||||
if not isinstance(part, dict):
|
||||
for part_value in cast(list[object], user_content):
|
||||
if not isinstance(part_value, dict):
|
||||
continue
|
||||
part = cast(dict[str, Any], part_value)
|
||||
if part.get("type") == "text":
|
||||
text_parts.append(part.get("text", ""))
|
||||
text_parts.append(
|
||||
_require_json_string(
|
||||
cast(object, part.get("text", "")),
|
||||
"messages[0].content[].text",
|
||||
)
|
||||
)
|
||||
elif part.get("type") == "image_url":
|
||||
url = part.get("image_url", {}).get("url", "")
|
||||
image_url = _require_json_object(
|
||||
cast(object, part.get("image_url", {})),
|
||||
"messages[0].content[].image_url",
|
||||
)
|
||||
url = _require_json_string(
|
||||
cast(object, image_url.get("url", "")),
|
||||
"messages[0].content[].image_url.url",
|
||||
)
|
||||
if url.startswith("data:"):
|
||||
saved = _save_base64_data_url(url, media_dir)
|
||||
if saved:
|
||||
@@ -191,7 +233,7 @@ async def _parse_multipart(request: web.Request) -> tuple[str, list[str], str |
|
||||
media_paths: list[str] = []
|
||||
|
||||
while True:
|
||||
part = await reader.next()
|
||||
part: Any = await reader.next()
|
||||
if part is None:
|
||||
break
|
||||
if part.name == "message":
|
||||
@@ -223,11 +265,9 @@ async def _parse_multipart(request: web.Request) -> tuple[str, list[str], str |
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def handle_chat_completions(request: web.Request) -> web.Response:
|
||||
async def handle_chat_completions(request: web.Request) -> web.Response | web.StreamResponse:
|
||||
"""POST /v1/chat/completions — supports JSON and multipart/form-data."""
|
||||
content_type = request.content_type or ""
|
||||
if not isinstance(content_type, str):
|
||||
content_type = ""
|
||||
content_type = _as_str(cast(object, request.content_type or ""))
|
||||
|
||||
agent_loop = _app_value(request.app, _AGENT_LOOP_KEY, "agent_loop")
|
||||
timeout_s: float = _app_value(
|
||||
@@ -247,6 +287,9 @@ async def handle_chat_completions(request: web.Request) -> web.Response:
|
||||
body = await request.json()
|
||||
except Exception:
|
||||
return _error_json(400, "Invalid JSON body")
|
||||
if not isinstance(body, dict):
|
||||
return _error_json(400, "Invalid JSON body")
|
||||
body = cast(dict[str, Any], body)
|
||||
stream = body.get("stream", False)
|
||||
requested_model = body.get("model")
|
||||
text, media_paths = _parse_json_content(body)
|
||||
@@ -405,7 +448,7 @@ async def handle_health(request: web.Request) -> web.Response:
|
||||
|
||||
|
||||
def create_app(
|
||||
agent_loop,
|
||||
agent_loop: "AgentLoop",
|
||||
model_name: str = "nanobot",
|
||||
request_timeout: float = 120.0,
|
||||
api_key: str = "",
|
||||
@@ -425,7 +468,10 @@ def create_app(
|
||||
app[_SESSION_LOCKS_KEY] = {} # per-user locks, keyed by session_key
|
||||
|
||||
@web.middleware
|
||||
async def auth_middleware(request: web.Request, handler) -> web.StreamResponse:
|
||||
async def auth_middleware(
|
||||
request: web.Request,
|
||||
handler: Callable[[web.Request], Awaitable[web.StreamResponse]],
|
||||
) -> web.StreamResponse:
|
||||
# Allow unauthenticated health checks.
|
||||
if request.path == "/health":
|
||||
return await handler(request)
|
||||
|
||||
+35
-23
@@ -10,10 +10,11 @@ import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass
|
||||
from importlib import metadata as importlib_metadata
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
@@ -204,6 +205,11 @@ def _now() -> float:
|
||||
return time.time()
|
||||
|
||||
|
||||
def _as_object_dict(value: object) -> dict[str, Any] | None:
|
||||
"""Narrow a JSON-like object to the string-keyed mapping used by this module."""
|
||||
return cast(dict[str, Any], value) if isinstance(value, dict) else None
|
||||
|
||||
|
||||
def _safe_skill_name(name: str) -> str:
|
||||
clean = _SAFE_NAME_RE.sub("-", name.lower()).strip("-")
|
||||
return f"cli-app-{clean or 'app'}"
|
||||
@@ -277,10 +283,11 @@ def _console_script_distribution(entry_point: str) -> str | None:
|
||||
if item.group != "console_scripts" or item.name != entry_point:
|
||||
continue
|
||||
try:
|
||||
name = distribution.metadata.get("Name")
|
||||
name: object = cast(Any, distribution.metadata).get("Name")
|
||||
except Exception:
|
||||
name = None
|
||||
return str(name or getattr(distribution, "name", "") or "").strip() or None
|
||||
fallback_name = cast(object, getattr(distribution, "name", ""))
|
||||
return str(name or fallback_name or "").strip() or None
|
||||
return None
|
||||
|
||||
|
||||
@@ -335,10 +342,10 @@ def _brand_payload(app: dict[str, Any]) -> tuple[str | None, str | None]:
|
||||
|
||||
def _read_json(path: Path) -> dict[str, Any] | None:
|
||||
try:
|
||||
data = json.loads(path.read_text(encoding="utf-8"))
|
||||
data: object = json.loads(path.read_text(encoding="utf-8"))
|
||||
except (OSError, json.JSONDecodeError):
|
||||
return None
|
||||
return data if isinstance(data, dict) else None
|
||||
return _as_object_dict(data)
|
||||
|
||||
|
||||
def _write_json(path: Path, data: dict[str, Any]) -> None:
|
||||
@@ -414,8 +421,8 @@ class CliAppManager:
|
||||
cached = _read_json(cache_path)
|
||||
if not cached:
|
||||
return None, 0.0
|
||||
data = cached.get("data")
|
||||
if not isinstance(data, dict):
|
||||
data = _as_object_dict(cached.get("data"))
|
||||
if data is None:
|
||||
return None, 0.0
|
||||
try:
|
||||
cached_at = float(cached.get("_cached_at", 0))
|
||||
@@ -425,8 +432,8 @@ class CliAppManager:
|
||||
|
||||
def _load_installed(self) -> dict[str, Any]:
|
||||
data = _read_json(self.installed_path) or {}
|
||||
apps = data.get("apps") if isinstance(data.get("apps"), dict) else data
|
||||
return apps if isinstance(apps, dict) else {}
|
||||
apps = _as_object_dict(data.get("apps"))
|
||||
return apps if apps is not None else data
|
||||
|
||||
def _save_installed(self, installed: dict[str, Any]) -> None:
|
||||
_write_json(self.installed_path, {"schema_version": 1, "apps": installed})
|
||||
@@ -453,8 +460,8 @@ class CliAppManager:
|
||||
try:
|
||||
response = httpx.get(url, timeout=15.0, follow_redirects=True)
|
||||
response.raise_for_status()
|
||||
fetched = response.json()
|
||||
if not isinstance(fetched, dict):
|
||||
fetched = _as_object_dict(response.json())
|
||||
if fetched is None:
|
||||
raise ValueError("registry response must be an object")
|
||||
except Exception:
|
||||
if data is not None:
|
||||
@@ -483,8 +490,8 @@ class CliAppManager:
|
||||
async with httpx.AsyncClient(timeout=15.0, follow_redirects=True) as client:
|
||||
response = await client.get(url)
|
||||
response.raise_for_status()
|
||||
fetched = response.json()
|
||||
if not isinstance(fetched, dict):
|
||||
fetched = _as_object_dict(response.json())
|
||||
if fetched is None:
|
||||
raise ValueError("registry response must be an object")
|
||||
except Exception:
|
||||
if data is not None:
|
||||
@@ -534,13 +541,14 @@ class CliAppManager:
|
||||
apps_by_name: dict[str, dict[str, Any]] = {}
|
||||
updated_values: list[str] = []
|
||||
for source, raw_base, registry in registries:
|
||||
meta = registry.get("meta")
|
||||
if isinstance(meta, dict) and isinstance(meta.get("updated"), str):
|
||||
meta = _as_object_dict(registry.get("meta"))
|
||||
if meta is not None and isinstance(meta.get("updated"), str):
|
||||
updated_values.append(meta["updated"])
|
||||
for row in registry.get("clis", []):
|
||||
if not isinstance(row, dict) or not row.get("name"):
|
||||
for row in cast(Iterable[object], registry.get("clis", [])):
|
||||
entry = _as_object_dict(row)
|
||||
if entry is None or not entry.get("name"):
|
||||
continue
|
||||
entry = dict(row)
|
||||
entry = dict(entry)
|
||||
entry["_source"] = source
|
||||
entry["_raw_base"] = raw_base
|
||||
key = str(entry["name"]).lower()
|
||||
@@ -588,7 +596,7 @@ class CliAppManager:
|
||||
if not installed:
|
||||
return []
|
||||
installed_by_name = {
|
||||
str(name).lower(): (str(name), data if isinstance(data, dict) else {})
|
||||
str(name).lower(): (str(name), _as_object_dict(data) or {})
|
||||
for name, data in installed.items()
|
||||
}
|
||||
seen: set[str] = set()
|
||||
@@ -769,12 +777,14 @@ class CliAppManager:
|
||||
for app in cached_apps
|
||||
if app.get("name")
|
||||
}
|
||||
rows = []
|
||||
rows: list[dict[str, Any]] = []
|
||||
for name, raw_entry in sorted(installed.items()):
|
||||
entry = raw_entry if isinstance(raw_entry, dict) else {}
|
||||
entry = _as_object_dict(raw_entry)
|
||||
if entry is None:
|
||||
entry = {}
|
||||
strategy = str(entry.get("strategy") or "bundled")
|
||||
cached_app = cached_by_name.get(str(name).lower(), {})
|
||||
app = {
|
||||
app: dict[str, Any] = {
|
||||
"name": str(name),
|
||||
"display_name": str(
|
||||
cached_app.get("display_name") or entry.get("display_name") or name
|
||||
@@ -1165,7 +1175,9 @@ Use the `run_cli_app` tool with `name="{name}"` for command execution. Do not in
|
||||
if str(app["name"]) not in installed:
|
||||
raise CliAppError("CLI app is not installed")
|
||||
raw_installed_entry = installed.get(str(app["name"]))
|
||||
installed_entry = raw_installed_entry if isinstance(raw_installed_entry, dict) else {}
|
||||
installed_entry = _as_object_dict(raw_installed_entry)
|
||||
if installed_entry is None:
|
||||
installed_entry = {}
|
||||
strategy = self._strategy(app)
|
||||
entry_point = str(app.get("entry_point") or "").strip()
|
||||
managed_entry_path = str(installed_entry.get("entry_point_path") or "").strip()
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any, Mapping
|
||||
from typing import Any, Mapping, cast
|
||||
|
||||
|
||||
def session_extra(metadata: Mapping[str, Any] | None) -> dict[str, Any]:
|
||||
@@ -29,9 +29,11 @@ def runtime_lines_for_request(
|
||||
"""Return CLI App annotations from an immutable request snapshot."""
|
||||
structured = metadata.get("cli_apps") if isinstance(metadata, Mapping) else None
|
||||
if isinstance(structured, list):
|
||||
structured_items = cast(list[Any], structured)
|
||||
mentions = [
|
||||
item for item in structured
|
||||
if isinstance(item, Mapping) and isinstance(item.get("name"), str)
|
||||
cast(Mapping[str, Any], item) for item in structured_items
|
||||
if isinstance(item, Mapping)
|
||||
and isinstance(cast(Mapping[str, Any], item).get("name"), str)
|
||||
]
|
||||
if mentions:
|
||||
return [
|
||||
@@ -49,7 +51,10 @@ def runtime_lines_for_request(
|
||||
try:
|
||||
from nanobot.apps.cli import CliAppManager
|
||||
|
||||
mentions = CliAppManager(workspace=workspace).mentioned_installed_apps(text)
|
||||
mentions = cast(
|
||||
list[dict[str, Any]],
|
||||
CliAppManager(workspace=workspace).mentioned_installed_apps(text),
|
||||
)
|
||||
except Exception:
|
||||
return []
|
||||
return [
|
||||
|
||||
@@ -22,6 +22,7 @@ from nanobot.audio.transcription_registry import (
|
||||
)
|
||||
from nanobot.config.loader import resolve_env_refs
|
||||
from nanobot.config.paths import get_media_dir
|
||||
from nanobot.config.schema import Config, ProviderConfig
|
||||
from nanobot.providers.registry import find_by_name
|
||||
from nanobot.utils.media_decode import FileSizeExceeded, save_base64_data_url
|
||||
|
||||
@@ -73,8 +74,9 @@ def _as_provider(value: Any) -> TranscriptionProviderName | None:
|
||||
return spec.name if spec else None
|
||||
|
||||
|
||||
def _provider_config(config: Any, provider: str) -> Any:
|
||||
return getattr(getattr(config, "providers", None), provider, None)
|
||||
def _provider_config(config: Config, provider: str) -> ProviderConfig | None:
|
||||
value = getattr(config.providers, provider, None)
|
||||
return value if isinstance(value, ProviderConfig) else None
|
||||
|
||||
|
||||
def _provider_default_api_base(provider: str) -> str | None:
|
||||
@@ -82,7 +84,10 @@ def _provider_default_api_base(provider: str) -> str | None:
|
||||
return spec.default_api_base if spec else None
|
||||
|
||||
|
||||
def _resolve_transcription_api_key(provider: str, provider_cfg: Any) -> str:
|
||||
def _resolve_transcription_api_key(
|
||||
provider: str,
|
||||
provider_cfg: ProviderConfig | None,
|
||||
) -> str:
|
||||
api_key = resolve_env_refs(getattr(provider_cfg, "api_key", None) or "") if provider_cfg else ""
|
||||
if api_key:
|
||||
return api_key
|
||||
@@ -94,10 +99,13 @@ def _resolve_transcription_api_key(provider: str, provider_cfg: Any) -> str:
|
||||
return env_key
|
||||
|
||||
env_key = spec.env_key if spec else ""
|
||||
return os.environ.get(env_key) if env_key else ""
|
||||
return os.environ.get(env_key, "") if env_key else ""
|
||||
|
||||
|
||||
def _resolve_transcription_api_base(provider: str, provider_cfg: Any) -> str:
|
||||
def _resolve_transcription_api_base(
|
||||
provider: str,
|
||||
provider_cfg: ProviderConfig | None,
|
||||
) -> str:
|
||||
api_base = resolve_env_refs(getattr(provider_cfg, "api_base", None) or "") if provider_cfg else ""
|
||||
if api_base:
|
||||
return api_base
|
||||
@@ -111,7 +119,7 @@ def _extract_data_url_mime(url: str) -> str | None:
|
||||
return header[5:].split(";", 1)[0].strip().lower() or None
|
||||
|
||||
|
||||
def resolve_transcription_config(config: Any) -> EffectiveTranscriptionConfig:
|
||||
def resolve_transcription_config(config: Config) -> EffectiveTranscriptionConfig:
|
||||
"""Resolve top-level transcription settings with legacy channel fallback."""
|
||||
top = getattr(config, "transcription", None)
|
||||
channels = getattr(config, "channels", None)
|
||||
|
||||
@@ -9,7 +9,7 @@ from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
|
||||
@@ -153,7 +153,11 @@ def _legacy_event_from_metadata(msg: OutboundMessage) -> OutboundEvent | None:
|
||||
)
|
||||
if meta.get("_goal_state_sync"):
|
||||
goal_state = meta.get("goal_state")
|
||||
return GoalStateSyncEvent(goal_state if isinstance(goal_state, dict) else {"active": False})
|
||||
return GoalStateSyncEvent(
|
||||
cast(dict[str, Any], goal_state)
|
||||
if isinstance(goal_state, dict)
|
||||
else {"active": False}
|
||||
)
|
||||
if meta.get("_goal_status"):
|
||||
status = meta.get("goal_status")
|
||||
if not isinstance(status, str) or not status:
|
||||
@@ -166,7 +170,7 @@ def _legacy_event_from_metadata(msg: OutboundMessage) -> OutboundEvent | None:
|
||||
goal_state = meta.get("goal_state")
|
||||
return TurnEndEvent(
|
||||
latency_ms=_metadata_int(meta, "latency_ms"),
|
||||
goal_state=goal_state if isinstance(goal_state, dict) else None,
|
||||
goal_state=cast(dict[str, Any], goal_state) if isinstance(goal_state, dict) else None,
|
||||
)
|
||||
if meta.get("_session_updated"):
|
||||
return SessionUpdatedEvent(scope=_metadata_str(meta, "_session_update_scope"))
|
||||
@@ -203,8 +207,12 @@ def _legacy_event_from_metadata(msg: OutboundMessage) -> OutboundEvent | None:
|
||||
reasoning_delta=bool(meta.get("_reasoning_delta")),
|
||||
reasoning_end=bool(meta.get("_reasoning_end")),
|
||||
stream_id=_metadata_str(meta, "_stream_id"),
|
||||
tool_events=tool_events if isinstance(tool_events, list) else None,
|
||||
file_edit_events=file_edit_events if isinstance(file_edit_events, list) else None,
|
||||
tool_events=cast(list[dict[str, Any]], tool_events)
|
||||
if isinstance(tool_events, list)
|
||||
else None,
|
||||
file_edit_events=cast(list[dict[str, Any]], file_edit_events)
|
||||
if isinstance(file_edit_events, list)
|
||||
else None,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
@@ -12,12 +12,15 @@ import contextlib
|
||||
import inspect
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.bus.events import InboundMessage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.utils.llm_runtime import LLMRuntime
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RuntimeEventContext:
|
||||
@@ -52,7 +55,7 @@ class TurnCompleted:
|
||||
|
||||
context: RuntimeEventContext
|
||||
latency_ms: int | None = None
|
||||
runtime: Any | None = None
|
||||
runtime: LLMRuntime | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -155,7 +158,7 @@ class RuntimeEventPublisher:
|
||||
def __init__(self, bus: RuntimeEventBus | None = None) -> None:
|
||||
self.bus = bus or RuntimeEventBus()
|
||||
self._turn_latency_ms: dict[str, int] = {}
|
||||
self._turn_runtime: dict[str, Any] = {}
|
||||
self._turn_runtime: dict[str, LLMRuntime] = {}
|
||||
|
||||
@staticmethod
|
||||
def _context(
|
||||
@@ -174,7 +177,7 @@ class RuntimeEventPublisher:
|
||||
attributes=dict(attributes or {}),
|
||||
)
|
||||
|
||||
def record_turn_runtime(self, session_key: str, runtime: Any) -> None:
|
||||
def record_turn_runtime(self, session_key: str, runtime: LLMRuntime) -> None:
|
||||
self._turn_runtime[session_key] = runtime
|
||||
|
||||
def record_turn_latency(self, session_key: str, latency_ms: int | None) -> None:
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
@@ -201,13 +201,21 @@ class BaseChannel(ABC):
|
||||
def supports_streaming(self) -> bool:
|
||||
"""True when config enables streaming AND this subclass implements send_delta."""
|
||||
cfg = self.config
|
||||
streaming = cfg.get("streaming", False) if isinstance(cfg, dict) else getattr(cfg, "streaming", False)
|
||||
config_mapping = cast(dict[str, Any], cfg) if isinstance(cfg, dict) else None
|
||||
streaming: Any = (
|
||||
config_mapping.get("streaming", False)
|
||||
if config_mapping is not None
|
||||
else getattr(cast(Any, cfg), "streaming", False)
|
||||
)
|
||||
return bool(streaming) and type(self).send_delta is not BaseChannel.send_delta
|
||||
|
||||
def is_allowed(self, sender_id: str) -> bool:
|
||||
"""Check sender permission: star > allowlist > pairing store > deny."""
|
||||
if isinstance(self.config, dict):
|
||||
allow_list = self.config.get("allow_from") or self.config.get("allowFrom") or []
|
||||
config_mapping = cast(dict[str, Any], self.config)
|
||||
allow_list: Any = (
|
||||
config_mapping.get("allow_from") or config_mapping.get("allowFrom") or []
|
||||
)
|
||||
else:
|
||||
allow_list = getattr(self.config, "allow_from", None) or []
|
||||
if "*" in allow_list:
|
||||
|
||||
@@ -6,7 +6,7 @@ from collections.abc import Iterable
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, Callable, Literal
|
||||
from typing import TYPE_CHECKING, Any, Callable, Literal, TypeGuard, cast
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.channels.plugin import ChannelPlugin
|
||||
@@ -22,6 +22,8 @@ class ChannelValidationContext:
|
||||
allow_local_service_access: bool = False
|
||||
|
||||
|
||||
# Keep callback contracts precise for static consumers. The public adapters below
|
||||
# still validate third-party implementations at runtime.
|
||||
SetupValidator = Callable[[dict[str, Any], ChannelValidationContext], dict[str, Any]]
|
||||
DefaultConfigFactory = Callable[[], dict[str, Any]]
|
||||
InstanceSpecsFactory = Callable[..., Iterable["ChannelInstanceSpec"]]
|
||||
@@ -87,7 +89,7 @@ class ChannelActivation:
|
||||
instances = (
|
||||
tuple(
|
||||
cls.from_config(item, include_instances=True)
|
||||
for item in raw_instances
|
||||
for item in cast(list[Any], raw_instances)
|
||||
if _config_mapping(item) is not None
|
||||
)
|
||||
if isinstance(raw_instances, list)
|
||||
@@ -193,7 +195,7 @@ class ChannelSetupSpec:
|
||||
def to_public_dict(self, channel_name: str) -> dict[str, Any]:
|
||||
"""Serialize the writable setup contract for generic WebUI consumers."""
|
||||
simple_required = set(self.simple_required_fields)
|
||||
fields = []
|
||||
fields: list[dict[str, Any]] = []
|
||||
for name, field in self.fields.items():
|
||||
if not field.writable:
|
||||
continue
|
||||
@@ -268,35 +270,37 @@ def channel_default_config(plugin: ChannelPlugin) -> dict[str, Any]:
|
||||
defaults: dict[str, Any] = {"enabled": plugin.default_enabled}
|
||||
if plugin.setup is not None:
|
||||
for name, field in plugin.setup.fields.items():
|
||||
value = field.default
|
||||
value: Any = field.default
|
||||
if value is None:
|
||||
value = {
|
||||
fallback_defaults: dict[str, Any] = {
|
||||
"string": "",
|
||||
"secret": "",
|
||||
"list": [],
|
||||
"bool": False,
|
||||
}.get(field.kind, _MISSING)
|
||||
}
|
||||
value = fallback_defaults.get(field.kind, _MISSING)
|
||||
if value is not _MISSING:
|
||||
_assign_channel_field(defaults, name, deepcopy(value))
|
||||
|
||||
factory = plugin.management.default_config
|
||||
if factory is None:
|
||||
return defaults
|
||||
values = factory()
|
||||
if not isinstance(values, dict):
|
||||
values_raw = cast(object, factory())
|
||||
if not isinstance(values_raw, dict):
|
||||
raise TypeError(f"ChannelPlugin.management.default_config for '{plugin.name}' must return a dict")
|
||||
return merge_missing_defaults(values, defaults)
|
||||
values = cast(dict[str, Any], values_raw)
|
||||
return cast(dict[str, Any], merge_missing_defaults(values, defaults))
|
||||
|
||||
|
||||
def _assign_channel_field(values: dict[str, Any], field: str, value: Any) -> None:
|
||||
target = values
|
||||
parts = field.split(".")
|
||||
for part in parts[:-1]:
|
||||
nested = target.get(part)
|
||||
nested: object = target.get(part)
|
||||
if not isinstance(nested, dict):
|
||||
nested = {}
|
||||
target[part] = nested
|
||||
target = nested
|
||||
target = cast(dict[str, Any], nested)
|
||||
target[parts[-1]] = value
|
||||
|
||||
|
||||
@@ -327,27 +331,28 @@ def channel_instance_specs(
|
||||
factory = plugin.management.instance_specs
|
||||
if factory is None:
|
||||
activation = ChannelActivation.from_config(section)
|
||||
raw_specs: Iterable[ChannelInstanceSpec] = (
|
||||
raw_specs: object = (
|
||||
[]
|
||||
if enabled_only and not activation.resolve(default=plugin.default_enabled)
|
||||
else [ChannelInstanceSpec(instance_id="default", config=section)]
|
||||
)
|
||||
else:
|
||||
raw_specs = factory(section, enabled_only=enabled_only)
|
||||
raw_specs = cast(object, factory(section, enabled_only=enabled_only))
|
||||
if not isinstance(raw_specs, Iterable):
|
||||
raise TypeError(
|
||||
f"ChannelPlugin.management.instance_specs for '{plugin.name}' must return an iterable"
|
||||
)
|
||||
specs = list(raw_specs)
|
||||
specs = list(cast(Iterable[object], raw_specs))
|
||||
if not _all_channel_instance_specs(specs):
|
||||
raise TypeError(
|
||||
f"ChannelPlugin.management.instance_specs for '{plugin.name}' returned an invalid item"
|
||||
)
|
||||
|
||||
instance_ids: set[str] = set()
|
||||
runtime_names: set[str] = set()
|
||||
for spec in specs:
|
||||
if not isinstance(spec, ChannelInstanceSpec):
|
||||
raise TypeError(
|
||||
f"ChannelPlugin.management.instance_specs for '{plugin.name}' returned an invalid item"
|
||||
)
|
||||
if not isinstance(spec.instance_id, str) or not spec.instance_id.strip():
|
||||
instance_id = cast(object, spec.instance_id)
|
||||
if not isinstance(instance_id, str) or not instance_id.strip():
|
||||
raise ValueError(
|
||||
f"ChannelPlugin.management.instance_specs for '{plugin.name}' returned an empty instance id"
|
||||
)
|
||||
@@ -367,6 +372,12 @@ def channel_instance_specs(
|
||||
return specs
|
||||
|
||||
|
||||
def _all_channel_instance_specs(
|
||||
values: list[object],
|
||||
) -> TypeGuard[list[ChannelInstanceSpec]]:
|
||||
return all(isinstance(value, ChannelInstanceSpec) for value in values)
|
||||
|
||||
|
||||
def resolve_channel_action_target(
|
||||
requested_instance_id: str | None,
|
||||
) -> str:
|
||||
@@ -393,8 +404,17 @@ def channel_instance_config(
|
||||
return {}
|
||||
config = selected.config
|
||||
if hasattr(config, "model_dump"):
|
||||
return dict(config.model_dump(mode="json", by_alias=True))
|
||||
return dict(config) if isinstance(config, dict) else {}
|
||||
dumped: dict[str, Any] = config.model_dump(mode="json", by_alias=True)
|
||||
copied: dict[str, Any] = {}
|
||||
for key in dumped:
|
||||
copied[key] = dumped[key]
|
||||
return copied
|
||||
if not isinstance(config, dict):
|
||||
return {}
|
||||
copied_config: dict[str, Any] = {}
|
||||
for key, value in cast(dict[object, Any], config).items():
|
||||
copied_config[cast(str, key)] = value
|
||||
return copied_config
|
||||
|
||||
|
||||
def channel_update_instance_config(
|
||||
@@ -409,7 +429,10 @@ def channel_update_instance_config(
|
||||
if instance_id not in {"", "default"}:
|
||||
raise ValueError(f"{plugin.name} does not support multiple instances")
|
||||
return values
|
||||
return updater(section, values, instance_id=instance_id)
|
||||
updated = cast(object, updater(section, values, instance_id=instance_id))
|
||||
if not isinstance(updated, dict):
|
||||
raise TypeError(f"ChannelPlugin.management.update_instance_config for '{plugin.name}' must return a dict")
|
||||
return cast(dict[str, Any], updated)
|
||||
|
||||
|
||||
def channel_set_config_enabled(
|
||||
@@ -423,7 +446,7 @@ def channel_set_config_enabled(
|
||||
from nanobot.config.loader import merge_missing_defaults
|
||||
|
||||
values = channel_instance_config(plugin, section, instance_id=instance_id)
|
||||
values = merge_missing_defaults(values, channel_default_config(plugin))
|
||||
values = cast(dict[str, Any], merge_missing_defaults(values, channel_default_config(plugin)))
|
||||
values["enabled"] = enabled
|
||||
return channel_update_instance_config(
|
||||
plugin,
|
||||
@@ -440,12 +463,16 @@ def channel_feature_instances(
|
||||
setup_spec: ChannelSetupSpec | None = None,
|
||||
) -> list[dict[str, Any]] | None:
|
||||
factory = plugin.management.feature_instances
|
||||
overrides = factory(section, setup_spec=setup_spec) if factory is not None else None
|
||||
overrides = (
|
||||
cast(object, factory(section, setup_spec=setup_spec))
|
||||
if factory is not None
|
||||
else None
|
||||
)
|
||||
if overrides is None and not plugin.management.multi_instance:
|
||||
return None
|
||||
if overrides is not None and (
|
||||
not isinstance(overrides, list)
|
||||
or any(not isinstance(instance, dict) for instance in overrides)
|
||||
or any(not isinstance(instance, dict) for instance in cast(list[object], overrides))
|
||||
):
|
||||
raise TypeError(
|
||||
f"ChannelPlugin.management.feature_instances for '{plugin.name}' "
|
||||
@@ -470,7 +497,8 @@ def channel_feature_instances(
|
||||
|
||||
by_id = {instance["id"]: instance for instance in instances}
|
||||
seen: set[str] = set()
|
||||
for override in overrides:
|
||||
for override_value in cast(list[object], overrides):
|
||||
override = cast(dict[str, Any], override_value)
|
||||
instance_id = override.get("id")
|
||||
if not isinstance(instance_id, str) or instance_id not in by_id:
|
||||
raise ValueError(
|
||||
@@ -514,20 +542,21 @@ def _validate_runtime_name(plugin: ChannelPlugin, runtime_name: Any) -> None:
|
||||
|
||||
|
||||
def channel_field_value(values: Any, field_path: str) -> Any:
|
||||
current = values
|
||||
current: Any = values
|
||||
for part in field_path.split("."):
|
||||
candidates = (part, _camel_to_snake(part))
|
||||
if isinstance(current, dict):
|
||||
for candidate in candidates:
|
||||
if candidate in current:
|
||||
current = current[candidate]
|
||||
current = cast(Any, current)[candidate]
|
||||
break
|
||||
else:
|
||||
return None
|
||||
continue
|
||||
for candidate in candidates:
|
||||
if hasattr(current, candidate):
|
||||
current = getattr(current, candidate)
|
||||
current_value = current
|
||||
if hasattr(current_value, candidate):
|
||||
current = getattr(current_value, candidate)
|
||||
break
|
||||
else:
|
||||
return None
|
||||
@@ -542,7 +571,7 @@ def stringify_channel_value(value: Any) -> str:
|
||||
if isinstance(value, bool):
|
||||
return "true" if value else "false"
|
||||
if isinstance(value, list):
|
||||
return ", ".join(str(item) for item in value)
|
||||
return ", ".join(str(item) for item in cast(list[Any], value))
|
||||
return str(value)
|
||||
|
||||
|
||||
@@ -586,8 +615,8 @@ def _channel_feature_instance(
|
||||
def _config_mapping(value: Any) -> dict[str, Any] | None:
|
||||
if hasattr(value, "model_dump"):
|
||||
dumped = value.model_dump(mode="json", by_alias=True)
|
||||
return dumped if isinstance(dumped, dict) else None
|
||||
return value if isinstance(value, dict) else None
|
||||
return cast(dict[str, Any], dumped) if isinstance(dumped, dict) else None
|
||||
return cast(dict[str, Any], value) if isinstance(value, dict) else None
|
||||
|
||||
|
||||
def _camel_to_snake(value: str) -> str:
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportConstantRedefinition=false, reportMissingTypeStubs=false, reportPrivateUsage=false
|
||||
"""DingTalk/DingDing channel implementation using Stream Mode."""
|
||||
|
||||
import asyncio
|
||||
@@ -10,7 +11,7 @@ from contextlib import suppress
|
||||
from inspect import isawaitable
|
||||
from io import BytesIO
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
from urllib.parse import unquote, urljoin, urlparse
|
||||
|
||||
import httpx
|
||||
@@ -36,11 +37,17 @@ def _escape_markdown_sender_name(value: str) -> str:
|
||||
for char in normalized
|
||||
)
|
||||
|
||||
DINGTALK_AVAILABLE = False
|
||||
AckMessage: Any = None
|
||||
CallbackHandler: Any = object
|
||||
Credential: Any = None
|
||||
DingTalkStreamClient: Any = None
|
||||
ChatbotMessage: Any = None
|
||||
|
||||
try:
|
||||
from dingtalk_stream import (
|
||||
AckMessage,
|
||||
CallbackHandler,
|
||||
CallbackMessage,
|
||||
Credential,
|
||||
DingTalkStreamClient,
|
||||
)
|
||||
@@ -48,41 +55,41 @@ try:
|
||||
|
||||
DINGTALK_AVAILABLE = True
|
||||
except ImportError:
|
||||
DINGTALK_AVAILABLE = False
|
||||
# Fallback so class definitions don't crash at module level
|
||||
CallbackHandler = object # type: ignore[assignment,misc]
|
||||
CallbackMessage = None # type: ignore[assignment,misc]
|
||||
AckMessage = None # type: ignore[assignment,misc]
|
||||
ChatbotMessage = None # type: ignore[assignment,misc]
|
||||
pass
|
||||
|
||||
|
||||
class NanobotDingTalkHandler(CallbackHandler):
|
||||
_CallbackHandlerBase = CallbackHandler
|
||||
|
||||
|
||||
class NanobotDingTalkHandler(_CallbackHandlerBase):
|
||||
"""
|
||||
Standard DingTalk Stream SDK Callback Handler.
|
||||
Parses incoming messages and forwards them to the Nanobot channel.
|
||||
"""
|
||||
|
||||
def __init__(self, channel: "DingTalkChannel"):
|
||||
super().__init__()
|
||||
super().__init__() # pyright: ignore[reportUnknownMemberType]
|
||||
self.channel = channel
|
||||
|
||||
async def process(self, message: CallbackMessage):
|
||||
async def process(self, message: Any) -> tuple[Any, str]:
|
||||
"""Process incoming stream message."""
|
||||
try:
|
||||
# Parse using SDK's ChatbotMessage for robust handling
|
||||
chatbot_msg = ChatbotMessage.from_dict(message.data)
|
||||
chatbot_msg: Any = ChatbotMessage.from_dict(message.data)
|
||||
message_data = cast(dict[str, Any], message.data)
|
||||
|
||||
# Extract text content; fall back to raw dict if SDK object is empty
|
||||
content = ""
|
||||
if chatbot_msg.text:
|
||||
content = chatbot_msg.text.content.strip()
|
||||
content = cast(str, chatbot_msg.text.content).strip()
|
||||
elif chatbot_msg.extensions.get("content", {}).get("recognition"):
|
||||
content = chatbot_msg.extensions["content"]["recognition"].strip()
|
||||
content = cast(str, chatbot_msg.extensions["content"]["recognition"]).strip()
|
||||
if not content:
|
||||
content = message.data.get("text", {}).get("content", "").strip()
|
||||
text_data = cast(dict[str, Any], message_data.get("text", {}))
|
||||
content = cast(str, text_data.get("content", "")).strip()
|
||||
|
||||
# Handle file/image messages
|
||||
file_paths = []
|
||||
file_paths: list[str] = []
|
||||
if chatbot_msg.message_type == "picture" and chatbot_msg.image_content:
|
||||
download_code = chatbot_msg.image_content.download_code
|
||||
if download_code:
|
||||
@@ -93,8 +100,18 @@ class NanobotDingTalkHandler(CallbackHandler):
|
||||
content = content or "[Image]"
|
||||
|
||||
elif chatbot_msg.message_type == "file":
|
||||
download_code = message.data.get("content", {}).get("downloadCode") or message.data.get("downloadCode")
|
||||
fname = message.data.get("content", {}).get("fileName") or message.data.get("fileName") or "file"
|
||||
message_content = cast(dict[str, Any], message_data.get("content", {}))
|
||||
download_code = cast(
|
||||
str,
|
||||
message_content.get("downloadCode")
|
||||
or message_data.get("downloadCode"),
|
||||
)
|
||||
fname = cast(
|
||||
str,
|
||||
message_content.get("fileName")
|
||||
or message_data.get("fileName")
|
||||
or "file",
|
||||
)
|
||||
if download_code:
|
||||
sender_uid = chatbot_msg.sender_staff_id or chatbot_msg.sender_id or "unknown"
|
||||
fp = await self.channel._download_dingtalk_file(download_code, fname, sender_uid)
|
||||
@@ -103,13 +120,17 @@ class NanobotDingTalkHandler(CallbackHandler):
|
||||
content = content or "[File]"
|
||||
|
||||
elif chatbot_msg.message_type == "richText" and chatbot_msg.rich_text_content:
|
||||
rich_list = chatbot_msg.rich_text_content.rich_text_list or []
|
||||
for item in rich_list:
|
||||
if not isinstance(item, dict):
|
||||
rich_list = cast(
|
||||
list[object],
|
||||
chatbot_msg.rich_text_content.rich_text_list or [],
|
||||
)
|
||||
for item_value in rich_list:
|
||||
if not isinstance(item_value, dict):
|
||||
continue
|
||||
item = cast(dict[str, Any], item_value)
|
||||
# A rich-text item may carry text and/or a downloadCode; the
|
||||
# DingTalk SDK treats them independently, so handle both.
|
||||
t = item.get("text", "").strip()
|
||||
t = cast(str, item.get("text", "")).strip()
|
||||
if t:
|
||||
fmt = item.get("type", "")
|
||||
if fmt == "bold":
|
||||
@@ -124,8 +145,8 @@ class NanobotDingTalkHandler(CallbackHandler):
|
||||
formatted = t
|
||||
content = (content + " " + formatted).strip() if content else formatted
|
||||
if item.get("downloadCode"):
|
||||
dc = item["downloadCode"]
|
||||
fname = item.get("fileName") or "file"
|
||||
dc = cast(str, item["downloadCode"])
|
||||
fname = cast(str, item.get("fileName") or "file")
|
||||
sender_uid = chatbot_msg.sender_staff_id or chatbot_msg.sender_id or "unknown"
|
||||
fp = await self.channel._download_dingtalk_file(dc, fname, sender_uid)
|
||||
if fp:
|
||||
@@ -143,13 +164,22 @@ class NanobotDingTalkHandler(CallbackHandler):
|
||||
)
|
||||
return AckMessage.STATUS_OK, "OK"
|
||||
|
||||
sender_id = chatbot_msg.sender_staff_id or chatbot_msg.sender_id
|
||||
sender_name = chatbot_msg.sender_nick or "Unknown"
|
||||
sender_id = cast(
|
||||
str | None,
|
||||
chatbot_msg.sender_staff_id or chatbot_msg.sender_id,
|
||||
)
|
||||
sender_name = cast(str, chatbot_msg.sender_nick or "Unknown")
|
||||
|
||||
conversation_type = message.data.get("conversationType")
|
||||
conversation_type = cast(
|
||||
str | None,
|
||||
message_data.get("conversationType"),
|
||||
)
|
||||
conversation_id = (
|
||||
message.data.get("conversationId")
|
||||
or message.data.get("openConversationId")
|
||||
cast(
|
||||
str | None,
|
||||
message_data.get("conversationId")
|
||||
or message_data.get("openConversationId"),
|
||||
)
|
||||
)
|
||||
|
||||
self.channel.logger.info("Received message from {} ({}): {}", sender_name, sender_id, content)
|
||||
@@ -218,14 +248,14 @@ class DingTalkChannel(BaseChannel):
|
||||
self.config: DingTalkConfig = config
|
||||
self._client: Any = None
|
||||
self._http: httpx.AsyncClient | None = None
|
||||
self._start_task: asyncio.Task | None = None
|
||||
self._start_task: asyncio.Task[Any] | None = None
|
||||
|
||||
# Access Token management for sending messages
|
||||
self._access_token: str | None = None
|
||||
self._token_expiry: float = 0
|
||||
|
||||
# Hold references to background tasks to prevent GC
|
||||
self._background_tasks: set[asyncio.Task] = set()
|
||||
self._background_tasks: set[asyncio.Task[None]] = set()
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Start the DingTalk bot with Stream Mode."""
|
||||
@@ -575,7 +605,11 @@ class DingTalkChannel(BaseChannel):
|
||||
try:
|
||||
resp = await self._http.post(url, files=files)
|
||||
text = resp.text
|
||||
result = resp.json() if resp.headers.get("content-type", "").startswith("application/json") else {}
|
||||
result = (
|
||||
cast(dict[str, Any], resp.json())
|
||||
if resp.headers.get("content-type", "").startswith("application/json")
|
||||
else {}
|
||||
)
|
||||
if resp.status_code >= 400:
|
||||
self.logger.error("media upload failed status={} type={} body={}", resp.status_code, media_type, text[:500])
|
||||
return None
|
||||
@@ -583,7 +617,7 @@ class DingTalkChannel(BaseChannel):
|
||||
if errcode != 0:
|
||||
self.logger.error("media upload api error type={} errcode={} body={}", media_type, errcode, text[:500])
|
||||
return None
|
||||
sub = result.get("result") or {}
|
||||
sub = cast(dict[str, Any], result.get("result") or {})
|
||||
media_id = result.get("media_id") or result.get("mediaId") or sub.get("media_id") or sub.get("mediaId")
|
||||
if not media_id:
|
||||
self.logger.error("media upload missing media_id body={}", text[:500])
|
||||
@@ -634,7 +668,7 @@ class DingTalkChannel(BaseChannel):
|
||||
self.logger.error("send failed msgKey={} status={} body={}", msg_key, resp.status_code, body[:500])
|
||||
return False
|
||||
try:
|
||||
result = resp.json()
|
||||
result = cast(dict[str, Any], resp.json())
|
||||
except Exception:
|
||||
result = {}
|
||||
errcode = result.get("errcode")
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
"""Discord channel implementation using discord.py."""
|
||||
# pyright: reportPrivateUsage=false, reportUnusedFunction=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -8,7 +9,7 @@ import time
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, Literal
|
||||
from typing import TYPE_CHECKING, Any, Literal, cast
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
@@ -43,7 +44,7 @@ class _StreamBuf:
|
||||
"""Per-chat streaming accumulator for progressive Discord message edits."""
|
||||
|
||||
text: str = ""
|
||||
message: Any | None = None
|
||||
message: discord.Message | None = None
|
||||
last_edit: float = 0.0
|
||||
stream_id: str | None = None
|
||||
|
||||
@@ -266,13 +267,14 @@ if DISCORD_AVAILABLE:
|
||||
self._channel.logger.warning("channel {} unavailable: {}", msg.chat_id, e)
|
||||
raise
|
||||
|
||||
reference, mention_settings = self._build_reply_context(channel, msg.reply_to)
|
||||
messageable_channel = cast(Messageable, channel)
|
||||
reference, mention_settings = self._build_reply_context(messageable_channel, msg.reply_to)
|
||||
sent_media = False
|
||||
failed_media: list[str] = []
|
||||
|
||||
for index, media_path in enumerate(msg.media or []):
|
||||
if await self._send_file(
|
||||
channel,
|
||||
messageable_channel,
|
||||
media_path,
|
||||
reference=reference if index == 0 else None,
|
||||
mention_settings=mention_settings,
|
||||
@@ -288,7 +290,7 @@ if DISCORD_AVAILABLE:
|
||||
if index == 0 and reference is not None and not sent_media:
|
||||
kwargs["reference"] = reference
|
||||
kwargs["allowed_mentions"] = mention_settings
|
||||
await channel.send(**kwargs)
|
||||
await messageable_channel.send(**kwargs)
|
||||
|
||||
async def _send_file(
|
||||
self,
|
||||
@@ -344,7 +346,7 @@ if DISCORD_AVAILABLE:
|
||||
self._channel.logger.warning("Invalid reply target: {}", reply_to)
|
||||
return None, mention_settings
|
||||
|
||||
return channel.get_partial_message(message_id), mention_settings
|
||||
return cast(Any, channel).get_partial_message(message_id), mention_settings
|
||||
|
||||
|
||||
class DiscordChannel(BaseChannel):
|
||||
@@ -423,8 +425,8 @@ class DiscordChannel(BaseChannel):
|
||||
import aiohttp
|
||||
|
||||
proxy_auth = aiohttp.BasicAuth(
|
||||
login=self.config.proxy_username,
|
||||
password=self.config.proxy_password,
|
||||
login=cast(str, self.config.proxy_username),
|
||||
password=cast(str, self.config.proxy_password),
|
||||
)
|
||||
elif has_user != has_pass:
|
||||
self.logger.warning(
|
||||
@@ -507,7 +509,7 @@ class DiscordChannel(BaseChannel):
|
||||
return
|
||||
if stream_id is not None and buf.stream_id is not None and buf.stream_id != stream_id:
|
||||
return
|
||||
await self._finalize_stream(chat_id, buf)
|
||||
await self._finalize_stream(chat_id, buf, buf.message)
|
||||
return
|
||||
|
||||
buf = self._stream_bufs.get(chat_id)
|
||||
@@ -635,7 +637,12 @@ class DiscordChannel(BaseChannel):
|
||||
self.logger.warning("channel {} unavailable: {}", chat_id, e)
|
||||
return None
|
||||
|
||||
async def _finalize_stream(self, chat_id: str, buf: _StreamBuf) -> None:
|
||||
async def _finalize_stream(
|
||||
self,
|
||||
chat_id: str,
|
||||
buf: _StreamBuf,
|
||||
message: discord.Message,
|
||||
) -> None:
|
||||
"""Commit the final streamed content and flush overflow chunks."""
|
||||
chunks = DiscordBotClient._build_chunks(buf.text, [], False)
|
||||
if not chunks:
|
||||
@@ -643,16 +650,12 @@ class DiscordChannel(BaseChannel):
|
||||
return
|
||||
|
||||
try:
|
||||
await buf.message.edit(content=chunks[0])
|
||||
await message.edit(content=chunks[0])
|
||||
except Exception as e:
|
||||
self.logger.warning("final stream edit failed: {}", e)
|
||||
raise
|
||||
|
||||
target = getattr(buf.message, "channel", None) or await self._resolve_channel(chat_id)
|
||||
if target is None:
|
||||
self.logger.warning("stream follow-up target {} unavailable", chat_id)
|
||||
self._stream_bufs.pop(chat_id, None)
|
||||
return
|
||||
target = message.channel
|
||||
|
||||
for extra_chunk in chunks[1:]:
|
||||
await target.send(content=extra_chunk)
|
||||
|
||||
@@ -17,7 +17,7 @@ from email.parser import BytesParser
|
||||
from email.utils import parseaddr
|
||||
from fnmatch import fnmatch
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
from typing import Any, Literal, cast
|
||||
|
||||
from loguru import logger
|
||||
from pydantic import Field
|
||||
@@ -188,7 +188,9 @@ class EmailChannel(BaseChannel):
|
||||
self.logger.exception("Error delivering email from {}", sender)
|
||||
continue
|
||||
|
||||
uid = str((item.get("metadata") or {}).get("uid") or "")
|
||||
metadata = item.get("metadata")
|
||||
metadata_data = cast(dict[str, Any], metadata) if isinstance(metadata, dict) else {}
|
||||
uid = str(metadata_data.get("uid") or "")
|
||||
if uid and should_apply_post_action:
|
||||
post_actions_uids.add(uid)
|
||||
|
||||
@@ -312,7 +314,7 @@ class EmailChannel(BaseChannel):
|
||||
raise
|
||||
|
||||
def _validate_config(self) -> bool:
|
||||
missing = []
|
||||
missing: list[str] = []
|
||||
if not self.config.imap_host:
|
||||
missing.append("imap_host")
|
||||
if not self.config.imap_username:
|
||||
@@ -427,7 +429,7 @@ class EmailChannel(BaseChannel):
|
||||
messages: list[dict[str, Any]],
|
||||
skipped_uids: set[str],
|
||||
cycle_uids: set[str],
|
||||
) -> None:
|
||||
) -> list[dict[str, Any]] | None:
|
||||
"""Fetch messages by arbitrary IMAP search criteria."""
|
||||
mailbox = self.config.imap_mailbox or "INBOX"
|
||||
|
||||
@@ -765,8 +767,10 @@ class EmailChannel(BaseChannel):
|
||||
@staticmethod
|
||||
def _extract_message_bytes(fetched: list[Any]) -> bytes | None:
|
||||
for item in fetched:
|
||||
if isinstance(item, tuple) and len(item) >= 2 and isinstance(item[1], (bytes, bytearray)):
|
||||
return bytes(item[1])
|
||||
if isinstance(item, tuple):
|
||||
fetched_item = cast(tuple[Any, ...], item)
|
||||
if len(fetched_item) >= 2 and isinstance(fetched_item[1], (bytes, bytearray)):
|
||||
return bytes(fetched_item[1])
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
@@ -837,8 +841,8 @@ class EmailChannel(BaseChannel):
|
||||
"""
|
||||
spf_pass = False
|
||||
dkim_pass = False
|
||||
for ar_header in parsed_msg.get_all("Authentication-Results") or []:
|
||||
ar_lower = ar_header.lower()
|
||||
for ar_header in cast(list[Any], parsed_msg.get_all("Authentication-Results") or []):
|
||||
ar_lower = str(ar_header).lower()
|
||||
if re.search(r"\bspf\s*=\s*pass\b", ar_lower):
|
||||
spf_pass = True
|
||||
if re.search(r"\bdkim\s*=\s*pass\b", ar_lower):
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Short-lived WebUI channel connection sessions."""
|
||||
|
||||
# pyright: reportPrivateUsage=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
@@ -46,7 +46,7 @@ def update_managed_feishu_instance(
|
||||
*,
|
||||
instance_id: str = DEFAULT_INSTANCE_ID,
|
||||
) -> dict[str, Any]:
|
||||
existing = section if isinstance(section, dict) else {}
|
||||
existing = cast(dict[str, Any], section) if isinstance(section, dict) else {}
|
||||
return upsert_feishu_instance(
|
||||
existing,
|
||||
feishu_default_config(),
|
||||
@@ -69,8 +69,8 @@ def _normalize_feishu_instance(
|
||||
inherited: dict[str, Any] | None = None,
|
||||
fallback_id: str = DEFAULT_INSTANCE_ID,
|
||||
) -> dict[str, Any]:
|
||||
config = merge_missing_defaults(inherited or {}, defaults)
|
||||
config = merge_missing_defaults(raw, config)
|
||||
config = cast(dict[str, Any], merge_missing_defaults(inherited or {}, defaults))
|
||||
config = cast(dict[str, Any], merge_missing_defaults(raw, config))
|
||||
|
||||
raw_id = raw.get("id") or raw.get("instanceId") or raw.get("instance_id") or fallback_id
|
||||
instance_id = validate_instance_id(str(raw_id))
|
||||
@@ -97,12 +97,13 @@ def _feishu_instance_inputs(
|
||||
section = section.model_dump(mode="json", by_alias=True)
|
||||
if not isinstance(section, dict):
|
||||
section = {}
|
||||
section_data = cast(dict[str, Any], section)
|
||||
|
||||
instances = section.get("instances")
|
||||
instances = section_data.get("instances")
|
||||
if isinstance(instances, list):
|
||||
inherited = {key: value for key, value in section.items() if key != "instances"}
|
||||
return list(instances), inherited
|
||||
return ([section] if section else [_base_feishu_instance_config(defaults)]), None
|
||||
inherited = {key: value for key, value in section_data.items() if key != "instances"}
|
||||
return list(cast(list[Any], instances)), inherited
|
||||
return ([section_data] if section_data else [_base_feishu_instance_config(defaults)]), None
|
||||
|
||||
|
||||
def feishu_instance_specs(
|
||||
@@ -124,7 +125,7 @@ def feishu_instance_specs(
|
||||
fallback_id = DEFAULT_INSTANCE_ID if index == 0 else f"assistant-{index + 1}"
|
||||
try:
|
||||
config = _normalize_feishu_instance(
|
||||
raw,
|
||||
cast(dict[str, Any], raw),
|
||||
defaults,
|
||||
inherited=inherited,
|
||||
fallback_id=fallback_id,
|
||||
@@ -179,7 +180,7 @@ def canonical_feishu_section(section: Any, defaults: dict[str, Any]) -> dict[str
|
||||
fallback_id = DEFAULT_INSTANCE_ID if index == 0 else f"assistant-{index + 1}"
|
||||
try:
|
||||
config = _normalize_feishu_instance(
|
||||
raw,
|
||||
cast(dict[str, Any], raw),
|
||||
defaults,
|
||||
inherited=inherited,
|
||||
fallback_id=fallback_id,
|
||||
@@ -238,9 +239,9 @@ def update_feishu_instance_preserving_shape(
|
||||
if (
|
||||
instance_id == DEFAULT_INSTANCE_ID
|
||||
and isinstance(section, dict)
|
||||
and not isinstance(section.get("instances"), list)
|
||||
and not isinstance(cast(dict[str, Any], section).get("instances"), list)
|
||||
):
|
||||
return {**section, **values}
|
||||
return {**cast(dict[str, Any], section), **values}
|
||||
|
||||
return upsert_feishu_instance(section, defaults, instance_id, values)
|
||||
|
||||
|
||||
+218
-132
@@ -1,4 +1,5 @@
|
||||
"""Feishu/Lark channel implementation using lark-oapi SDK with WebSocket long connection."""
|
||||
# pyright: reportMissingModuleSource=false, reportMissingTypeStubs=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -14,8 +15,9 @@ from collections import OrderedDict
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
from functools import partial
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, TypedDict, cast
|
||||
|
||||
from rich.console import Console
|
||||
from rich.markup import escape
|
||||
@@ -44,7 +46,10 @@ from nanobot.utils.helpers import safe_filename
|
||||
from nanobot.utils.logging_bridge import redirect_lib_logging
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from lark_oapi.api.im.v1.model import MentionEvent, P2ImMessageReceiveV1
|
||||
from lark_oapi.api.im.v1.model import ( # pyright: ignore[reportMissingTypeStubs]
|
||||
MentionEvent,
|
||||
P2ImMessageReceiveV1,
|
||||
)
|
||||
|
||||
FEISHU_AVAILABLE = importlib.util.find_spec("lark_oapi") is not None
|
||||
_LOGIN_CONSOLE = Console()
|
||||
@@ -55,6 +60,20 @@ def _identity_timestamp() -> str:
|
||||
return datetime.now(UTC).isoformat(timespec="seconds").replace("+00:00", "Z")
|
||||
|
||||
|
||||
def _as_json_object(value: Any) -> dict[str, Any] | None:
|
||||
"""Narrow untyped SDK/JSON objects at the channel boundary."""
|
||||
return cast(dict[str, Any], value) if isinstance(value, dict) else None
|
||||
|
||||
|
||||
def _as_json_list(value: Any) -> list[Any] | None:
|
||||
"""Narrow untyped SDK/JSON arrays at the channel boundary."""
|
||||
return cast(list[Any], value) if isinstance(value, list) else None
|
||||
|
||||
|
||||
def _ignore_event(_: Any) -> None:
|
||||
"""Consume SDK events that intentionally have no channel action."""
|
||||
|
||||
|
||||
def _load_lark_runtime() -> tuple[Any, str, str]:
|
||||
"""Import the heavy Feishu SDK lazily.
|
||||
|
||||
@@ -69,9 +88,12 @@ def _load_lark_runtime() -> tuple[Any, str, str]:
|
||||
# close the same loop.
|
||||
with _LARK_RUNTIME_LOCK:
|
||||
ws_client_already_imported = "lark_oapi.ws.client" in sys.modules
|
||||
import lark_oapi as lark
|
||||
import lark_oapi.ws.client as lark_ws_client
|
||||
from lark_oapi.core.const import FEISHU_DOMAIN, LARK_DOMAIN
|
||||
import lark_oapi as lark # pyright: ignore[reportMissingTypeStubs]
|
||||
import lark_oapi.ws.client as lark_ws_client # pyright: ignore[reportMissingTypeStubs]
|
||||
from lark_oapi.core.const import ( # pyright: ignore[reportMissingTypeStubs]
|
||||
FEISHU_DOMAIN,
|
||||
LARK_DOMAIN,
|
||||
)
|
||||
|
||||
if (
|
||||
not ws_client_already_imported
|
||||
@@ -106,7 +128,7 @@ def fetch_feishu_app_identity(
|
||||
|
||||
try:
|
||||
lark, feishu_domain, lark_domain = _load_lark_runtime()
|
||||
from lark_oapi.api.application.v6.model.get_application_request import (
|
||||
from lark_oapi.api.application.v6.model.get_application_request import ( # pyright: ignore[reportMissingTypeStubs]
|
||||
GetApplicationRequest,
|
||||
)
|
||||
|
||||
@@ -151,9 +173,9 @@ MSG_TYPE_MAP = {
|
||||
}
|
||||
|
||||
|
||||
def _extract_share_card_content(content_json: dict, msg_type: str) -> str:
|
||||
def _extract_share_card_content(content_json: dict[str, Any], msg_type: str) -> str:
|
||||
"""Extract text representation from share cards and interactive messages."""
|
||||
parts = []
|
||||
parts: list[str] = []
|
||||
|
||||
if msg_type == "share_chat":
|
||||
parts.append(f"[shared chat: {content_json.get('chat_id', '')}]")
|
||||
@@ -171,9 +193,9 @@ def _extract_share_card_content(content_json: dict, msg_type: str) -> str:
|
||||
return "\n".join(parts) if parts else f"[{msg_type}]"
|
||||
|
||||
|
||||
def _extract_interactive_content(content: dict) -> list[str]:
|
||||
def _extract_interactive_content(content: str | dict[str, Any]) -> list[str]:
|
||||
"""Recursively extract text and links from interactive card content."""
|
||||
parts = []
|
||||
parts: list[str] = []
|
||||
|
||||
if isinstance(content, str):
|
||||
try:
|
||||
@@ -189,8 +211,9 @@ def _extract_interactive_content(content: dict) -> list[str]:
|
||||
if isinstance(user_dsl, str) and user_dsl.strip():
|
||||
try:
|
||||
dsl = json.loads(user_dsl)
|
||||
if isinstance(dsl, dict):
|
||||
parts.extend(_extract_interactive_content(dsl))
|
||||
dsl_object = _as_json_object(dsl)
|
||||
if dsl_object is not None:
|
||||
parts.extend(_extract_interactive_content(dsl_object))
|
||||
if parts:
|
||||
return parts
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
@@ -198,8 +221,9 @@ def _extract_interactive_content(content: dict) -> list[str]:
|
||||
|
||||
if "title" in content:
|
||||
title = content["title"]
|
||||
if isinstance(title, dict):
|
||||
title_content = title.get("content", "") or title.get("text", "")
|
||||
title_object = _as_json_object(title)
|
||||
if title_object is not None:
|
||||
title_content = title_object.get("content", "") or title_object.get("text", "")
|
||||
if title_content:
|
||||
parts.append(f"title: {title_content}")
|
||||
elif isinstance(title, str):
|
||||
@@ -207,34 +231,39 @@ def _extract_interactive_content(content: dict) -> list[str]:
|
||||
|
||||
# Top-level elements: flat list or nested list format
|
||||
elements = content.get("elements")
|
||||
if isinstance(elements, list):
|
||||
if elements and isinstance(elements[0], list):
|
||||
elements_list = _as_json_list(elements)
|
||||
if elements_list is not None:
|
||||
if elements_list and isinstance(elements_list[0], list):
|
||||
# Nested list: [[{tag:"text",text:"..."}], ...]
|
||||
for row in elements:
|
||||
if isinstance(row, list):
|
||||
for element in row:
|
||||
for row in elements_list:
|
||||
row_list = _as_json_list(row)
|
||||
if row_list is not None:
|
||||
for element in row_list:
|
||||
parts.extend(_extract_element_content(element))
|
||||
else:
|
||||
# Flat list: [{tag:"markdown",content:"..."}, ...]
|
||||
for element in elements:
|
||||
for element in elements_list:
|
||||
parts.extend(_extract_element_content(element))
|
||||
|
||||
# Body elements (schema 2.0)
|
||||
body = content.get("body", {})
|
||||
if isinstance(body, dict):
|
||||
body_elements = body.get("elements")
|
||||
if isinstance(body_elements, list):
|
||||
body_object = _as_json_object(body)
|
||||
if body_object is not None:
|
||||
body_elements = _as_json_list(body_object.get("elements"))
|
||||
if body_elements is not None:
|
||||
for element in body_elements:
|
||||
parts.extend(_extract_element_content(element))
|
||||
|
||||
card = content.get("card", {})
|
||||
if card:
|
||||
parts.extend(_extract_interactive_content(card))
|
||||
card_object = _as_json_object(card)
|
||||
if card_object:
|
||||
parts.extend(_extract_interactive_content(card_object))
|
||||
|
||||
header = content.get("header", {})
|
||||
if header:
|
||||
header_title = header.get("title", {})
|
||||
if isinstance(header_title, dict):
|
||||
header_object = _as_json_object(header)
|
||||
if header_object is not None:
|
||||
header_title = _as_json_object(header_object.get("title", {}))
|
||||
if header_title is not None:
|
||||
header_text = header_title.get("content", "") or header_title.get("text", "")
|
||||
if header_text:
|
||||
parts.append(f"title: {header_text}")
|
||||
@@ -242,13 +271,16 @@ def _extract_interactive_content(content: dict) -> list[str]:
|
||||
return parts
|
||||
|
||||
|
||||
def _extract_element_content(element: dict) -> list[str]:
|
||||
def _extract_element_content(element: Any) -> list[str]:
|
||||
"""Extract content from a single card element."""
|
||||
parts = []
|
||||
parts: list[str] = []
|
||||
|
||||
if not isinstance(element, dict):
|
||||
element_object = _as_json_object(element)
|
||||
if element_object is None:
|
||||
return parts
|
||||
|
||||
element = element_object
|
||||
|
||||
tag = element.get("tag", "")
|
||||
|
||||
if tag in ("markdown", "lark_md"):
|
||||
@@ -263,16 +295,18 @@ def _extract_element_content(element: dict) -> list[str]:
|
||||
|
||||
elif tag == "div":
|
||||
text = element.get("text", {})
|
||||
if isinstance(text, dict):
|
||||
text_content = text.get("content", "") or text.get("text", "")
|
||||
text_object = _as_json_object(text)
|
||||
if text_object is not None:
|
||||
text_content = text_object.get("content", "") or text_object.get("text", "")
|
||||
if text_content:
|
||||
parts.append(text_content)
|
||||
elif isinstance(text, str):
|
||||
parts.append(text)
|
||||
for field in element.get("fields") or []:
|
||||
if isinstance(field, dict):
|
||||
field_text = field.get("text", {})
|
||||
if isinstance(field_text, dict):
|
||||
for field in _as_json_list(element.get("fields")) or []:
|
||||
field_object = _as_json_object(field)
|
||||
if field_object is not None:
|
||||
field_text = _as_json_object(field_object.get("text", {}))
|
||||
if field_text is not None:
|
||||
c = field_text.get("content", "")
|
||||
if c:
|
||||
parts.append(c)
|
||||
@@ -287,30 +321,33 @@ def _extract_element_content(element: dict) -> list[str]:
|
||||
|
||||
elif tag == "button":
|
||||
text = element.get("text", {})
|
||||
if isinstance(text, dict):
|
||||
c = text.get("content", "")
|
||||
text_object = _as_json_object(text)
|
||||
if text_object is not None:
|
||||
c = text_object.get("content", "")
|
||||
if c:
|
||||
parts.append(c)
|
||||
multi_url = element.get("multi_url") or {}
|
||||
multi_url: Any = element.get("multi_url") or {}
|
||||
multi_url_object = _as_json_object(multi_url)
|
||||
url = element.get("url", "") or (
|
||||
multi_url.get("url", "") if isinstance(multi_url, dict) else ""
|
||||
multi_url_object.get("url", "") if multi_url_object is not None else ""
|
||||
)
|
||||
if url:
|
||||
parts.append(f"link: {url}")
|
||||
|
||||
elif tag == "img":
|
||||
alt = element.get("alt", {})
|
||||
parts.append(alt.get("content", "[image]") if isinstance(alt, dict) else "[image]")
|
||||
alt = _as_json_object(element.get("alt", {}))
|
||||
parts.append(alt.get("content", "[image]") if alt is not None else "[image]")
|
||||
|
||||
elif tag == "note":
|
||||
for ne in element.get("elements") or []:
|
||||
for ne in _as_json_list(element.get("elements")) or []:
|
||||
parts.extend(_extract_element_content(ne))
|
||||
|
||||
elif tag == "column_set":
|
||||
for col in element.get("columns") or []:
|
||||
if not isinstance(col, dict):
|
||||
for col in _as_json_list(element.get("columns")) or []:
|
||||
col_object = _as_json_object(col)
|
||||
if col_object is None:
|
||||
continue
|
||||
for ce in col.get("elements") or []:
|
||||
for ce in _as_json_list(col_object.get("elements")) or []:
|
||||
parts.extend(_extract_element_content(ce))
|
||||
|
||||
elif tag == "plain_text":
|
||||
@@ -319,36 +356,44 @@ def _extract_element_content(element: dict) -> list[str]:
|
||||
parts.append(content)
|
||||
|
||||
elif tag == "table":
|
||||
columns = [
|
||||
(column["name"], str(column.get("display_name") or column["name"]))
|
||||
for column in (element.get("columns") or [])
|
||||
if isinstance(column, dict) and column.get("name")
|
||||
]
|
||||
rows = element.get("rows") or []
|
||||
columns: list[tuple[str, str]] = []
|
||||
for column in _as_json_list(element.get("columns")) or []:
|
||||
column_object = _as_json_object(column)
|
||||
if column_object is None:
|
||||
continue
|
||||
name = column_object.get("name")
|
||||
if isinstance(name, str) and name:
|
||||
columns.append((name, str(column_object.get("display_name") or name)))
|
||||
rows = _as_json_list(element.get("rows")) or []
|
||||
if columns:
|
||||
parts.append(" | ".join(header for _, header in columns))
|
||||
if isinstance(rows, list):
|
||||
if rows:
|
||||
for row in rows:
|
||||
if not isinstance(row, dict):
|
||||
row_object = _as_json_object(row)
|
||||
if row_object is None:
|
||||
continue
|
||||
values = []
|
||||
values: list[str] = []
|
||||
for name, _ in columns:
|
||||
value = row.get(name)
|
||||
value = row_object.get(name)
|
||||
if isinstance(value, list):
|
||||
value = " ".join(str(item).strip() for item in value if item is not None)
|
||||
value = " ".join(
|
||||
str(item).strip()
|
||||
for item in cast(list[Any], value)
|
||||
if item is not None
|
||||
)
|
||||
values.append("" if value is None else str(value).strip())
|
||||
row_text = " | ".join(values).strip()
|
||||
if row_text:
|
||||
parts.append(row_text)
|
||||
|
||||
else:
|
||||
for ne in element.get("elements") or []:
|
||||
for ne in _as_json_list(element.get("elements")) or []:
|
||||
parts.extend(_extract_element_content(ne))
|
||||
|
||||
return parts
|
||||
|
||||
|
||||
def _extract_post_content(content_json: dict) -> tuple[str, list[str]]:
|
||||
def _extract_post_content(content_json: dict[str, Any]) -> tuple[str, list[str]]:
|
||||
"""Extract text and image keys from Feishu post (rich text) message.
|
||||
|
||||
Handles three payload shapes:
|
||||
@@ -357,45 +402,48 @@ def _extract_post_content(content_json: dict) -> tuple[str, list[str]]:
|
||||
- Wrapped: {"post": {"zh_cn": {"title": "...", "content": [...]}}}
|
||||
"""
|
||||
|
||||
def _parse_block(block: dict) -> tuple[str | None, list[str]]:
|
||||
if not isinstance(block, dict) or not isinstance(block.get("content"), list):
|
||||
def _parse_block(block: dict[str, Any]) -> tuple[str | None, list[str]]:
|
||||
content = _as_json_list(block.get("content"))
|
||||
if content is None:
|
||||
return None, []
|
||||
texts, images = [], []
|
||||
texts: list[str] = []
|
||||
images: list[str] = []
|
||||
title = block.get("title")
|
||||
if isinstance(title, str) and title:
|
||||
texts.append(title)
|
||||
for row in block["content"]:
|
||||
if not isinstance(row, list):
|
||||
for row in content:
|
||||
row_items = _as_json_list(row)
|
||||
if row_items is None:
|
||||
continue
|
||||
for el in row:
|
||||
if not isinstance(el, dict):
|
||||
for el in row_items:
|
||||
element = _as_json_object(el)
|
||||
if element is None:
|
||||
continue
|
||||
tag = el.get("tag")
|
||||
tag = element.get("tag")
|
||||
if tag in ("text", "a"):
|
||||
text = el.get("text", "")
|
||||
text = element.get("text", "")
|
||||
if isinstance(text, str):
|
||||
texts.append(text)
|
||||
elif tag == "at":
|
||||
user = el.get("user_name", "user")
|
||||
user = element.get("user_name", "user")
|
||||
texts.append(f"@{user if isinstance(user, str) and user else 'user'}")
|
||||
elif tag == "code_block":
|
||||
lang = el.get("language", "")
|
||||
code_text = el.get("text", "")
|
||||
lang = element.get("language", "")
|
||||
code_text = element.get("text", "")
|
||||
if not isinstance(lang, str):
|
||||
lang = ""
|
||||
if not isinstance(code_text, str):
|
||||
code_text = ""
|
||||
texts.append(f"\n```{lang}\n{code_text}\n```\n")
|
||||
elif tag == "img" and (key := el.get("image_key")):
|
||||
elif tag == "img" and isinstance((key := element.get("image_key")), str):
|
||||
images.append(key)
|
||||
return (" ".join(texts).strip() or None), images
|
||||
|
||||
# Unwrap optional {"post": ...} envelope
|
||||
root = content_json
|
||||
if isinstance(root, dict) and isinstance(root.get("post"), dict):
|
||||
root = root["post"]
|
||||
if not isinstance(root, dict):
|
||||
return "", []
|
||||
post = _as_json_object(root.get("post"))
|
||||
if post is not None:
|
||||
root = post
|
||||
|
||||
# Direct format
|
||||
if "content" in root:
|
||||
@@ -406,19 +454,23 @@ def _extract_post_content(content_json: dict) -> tuple[str, list[str]]:
|
||||
# Localized: prefer known locales, then fall back to any dict child
|
||||
for key in ("zh_cn", "en_us", "ja_jp"):
|
||||
if key in root:
|
||||
text, imgs = _parse_block(root[key])
|
||||
block = _as_json_object(root[key])
|
||||
if block is None:
|
||||
continue
|
||||
text, imgs = _parse_block(block)
|
||||
if text or imgs:
|
||||
return text or "", imgs
|
||||
for val in root.values():
|
||||
if isinstance(val, dict):
|
||||
text, imgs = _parse_block(val)
|
||||
block = _as_json_object(val)
|
||||
if block is not None:
|
||||
text, imgs = _parse_block(block)
|
||||
if text or imgs:
|
||||
return text or "", imgs
|
||||
|
||||
return "", []
|
||||
|
||||
|
||||
def _extract_post_text(content_json: dict) -> str:
|
||||
def _extract_post_text(content_json: dict[str, Any]) -> str: # pyright: ignore[reportUnusedFunction]
|
||||
"""Extract plain text from Feishu post (rich text) message content.
|
||||
|
||||
Legacy wrapper for _extract_post_content, returns only text.
|
||||
@@ -442,11 +494,18 @@ _REGISTRATION_PATH = "/oauth/v1/app/registration"
|
||||
_ONBOARD_REQUEST_TIMEOUT_S = 10
|
||||
|
||||
|
||||
class _RegistrationStart(TypedDict):
|
||||
device_code: str
|
||||
qr_url: str
|
||||
interval: int
|
||||
expire_in: int
|
||||
|
||||
|
||||
def _accounts_base_url(domain: str) -> str:
|
||||
return _ONBOARD_ACCOUNTS_URLS.get(domain, _ONBOARD_ACCOUNTS_URLS["feishu"])
|
||||
|
||||
|
||||
def _post_registration(base_url: str, body: dict[str, str]) -> dict:
|
||||
def _post_registration(base_url: str, body: dict[str, str]) -> dict[str, Any]:
|
||||
"""POST form-encoded data to the registration endpoint, return parsed JSON.
|
||||
|
||||
The registration endpoint returns JSON even on HTTP errors (e.g. poll
|
||||
@@ -462,7 +521,8 @@ def _post_registration(base_url: str, body: dict[str, str]) -> dict:
|
||||
headers={"Content-Type": "application/x-www-form-urlencoded"},
|
||||
)
|
||||
try:
|
||||
return resp.json()
|
||||
parsed = resp.json()
|
||||
return _as_json_object(parsed) or {}
|
||||
except json.JSONDecodeError:
|
||||
resp.raise_for_status()
|
||||
return {}
|
||||
@@ -472,7 +532,7 @@ def _init_registration(domain: str = "feishu") -> None:
|
||||
"""Verify the environment supports client_secret auth. Raises RuntimeError if not."""
|
||||
base_url = _accounts_base_url(domain)
|
||||
res = _post_registration(base_url, {"action": "init"})
|
||||
methods = res.get("supported_auth_methods") or []
|
||||
methods = _as_json_list(res.get("supported_auth_methods")) or []
|
||||
if "client_secret" not in methods:
|
||||
raise RuntimeError(
|
||||
f"Feishu / Lark registration does not support client_secret auth. "
|
||||
@@ -480,7 +540,7 @@ def _init_registration(domain: str = "feishu") -> None:
|
||||
)
|
||||
|
||||
|
||||
def _begin_registration(domain: str = "feishu") -> dict:
|
||||
def _begin_registration(domain: str = "feishu") -> _RegistrationStart:
|
||||
"""Start the device-code flow. Returns device_code, qr_url, interval, expire_in."""
|
||||
base_url = _accounts_base_url(domain)
|
||||
res = _post_registration(base_url, {
|
||||
@@ -490,16 +550,18 @@ def _begin_registration(domain: str = "feishu") -> dict:
|
||||
"request_user_info": "open_id",
|
||||
})
|
||||
device_code = res.get("device_code")
|
||||
if not device_code:
|
||||
if not isinstance(device_code, str) or not device_code:
|
||||
raise RuntimeError("Feishu / Lark registration did not return a device_code")
|
||||
qr_url = res.get("verification_uri_complete", "")
|
||||
if not qr_url:
|
||||
if not isinstance(qr_url, str) or not qr_url:
|
||||
raise RuntimeError("Feishu / Lark registration did not return a login URL")
|
||||
interval = res.get("interval")
|
||||
expire_in = res.get("expire_in")
|
||||
return {
|
||||
"device_code": device_code,
|
||||
"qr_url": qr_url,
|
||||
"interval": res.get("interval") or 5,
|
||||
"expire_in": res.get("expire_in") or 600,
|
||||
"interval": interval if isinstance(interval, int) else 5,
|
||||
"expire_in": expire_in if isinstance(expire_in, int) else 600,
|
||||
}
|
||||
|
||||
|
||||
@@ -509,7 +571,7 @@ def _poll_registration(
|
||||
interval: int,
|
||||
expire_in: int,
|
||||
domain: str = "feishu",
|
||||
) -> dict | None:
|
||||
) -> dict[str, Any] | None:
|
||||
"""Poll until the user scans the QR code, or timeout/denial.
|
||||
|
||||
Returns dict with app_id, app_secret, domain on success, None on failure.
|
||||
@@ -548,7 +610,7 @@ def poll_registration_once(
|
||||
*,
|
||||
device_code: str,
|
||||
domain: str = "feishu",
|
||||
) -> dict:
|
||||
) -> dict[str, Any]:
|
||||
"""Poll the Feishu/Lark device-code flow once.
|
||||
|
||||
This non-blocking shape is used by WebUI. The CLI keeps using
|
||||
@@ -562,7 +624,7 @@ def poll_registration_once(
|
||||
"tp": "ob_app",
|
||||
})
|
||||
|
||||
user_info = res.get("user_info") or {}
|
||||
user_info = _as_json_object(res.get("user_info")) or {}
|
||||
tenant_brand = user_info.get("tenant_brand")
|
||||
if tenant_brand == "lark":
|
||||
current_domain = "lark"
|
||||
@@ -641,9 +703,7 @@ def sync_saved_feishu_identity_boundary(
|
||||
from nanobot.config.loader import load_config, save_config
|
||||
|
||||
full_config = load_config()
|
||||
feishu_cfg = getattr(full_config.channels, "feishu", None) or {}
|
||||
if not isinstance(feishu_cfg, dict):
|
||||
feishu_cfg = {}
|
||||
feishu_cfg = _as_json_object(getattr(full_config.channels, "feishu", None)) or {}
|
||||
|
||||
defaults = feishu_default_config()
|
||||
previous_identity_key = ""
|
||||
@@ -675,7 +735,7 @@ def sync_saved_feishu_identity_boundary(
|
||||
|
||||
|
||||
def save_registration_result(
|
||||
result: dict,
|
||||
result: dict[str, Any],
|
||||
*,
|
||||
instance_id: str = DEFAULT_INSTANCE_ID,
|
||||
name: str | None = None,
|
||||
@@ -684,9 +744,7 @@ def save_registration_result(
|
||||
from nanobot.config.loader import load_config, save_config
|
||||
|
||||
full_config = load_config()
|
||||
feishu_cfg = getattr(full_config.channels, "feishu", None) or {}
|
||||
if not isinstance(feishu_cfg, dict):
|
||||
feishu_cfg = {}
|
||||
feishu_cfg = _as_json_object(getattr(full_config.channels, "feishu", None)) or {}
|
||||
defaults = feishu_default_config()
|
||||
app_id = str(result["app_id"]).strip()
|
||||
domain = str(result.get("domain", "feishu") or "feishu").strip().lower()
|
||||
@@ -809,7 +867,7 @@ def refresh_saved_feishu_identities(
|
||||
def qr_register(
|
||||
*,
|
||||
initial_domain: str = "feishu",
|
||||
) -> dict | None:
|
||||
) -> dict[str, Any] | None:
|
||||
"""Run the Feishu / Lark scan-to-create QR registration flow.
|
||||
|
||||
Returns on success:
|
||||
@@ -853,7 +911,7 @@ def _print_qr_code(url: str) -> None:
|
||||
def _qr_register_inner(
|
||||
*,
|
||||
initial_domain: str,
|
||||
) -> dict | None:
|
||||
) -> dict[str, Any] | None:
|
||||
"""Run init → begin → poll. Raises on network/protocol errors."""
|
||||
_LOGIN_CONSOLE.print("[cyan]Preparing Feishu/Lark login...[/cyan]")
|
||||
_init_registration(initial_domain)
|
||||
@@ -935,7 +993,7 @@ class FeishuChannel(BaseChannel):
|
||||
self._loop: asyncio.AbstractEventLoop | None = None
|
||||
self._stream_bufs: dict[str, _FeishuStreamBuf] = {}
|
||||
self._bot_open_id: str | None = None
|
||||
self._background_tasks: set[asyncio.Task] = set()
|
||||
self._background_tasks: set[asyncio.Task[Any]] = set()
|
||||
self._reaction_ids: dict[str, str] = {} # message_id → reaction_id
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -1062,12 +1120,12 @@ class FeishuChannel(BaseChannel):
|
||||
builder = self._register_optional_event(
|
||||
builder,
|
||||
"register_p2_im_chat_member_bot_added_v1",
|
||||
lambda _: None,
|
||||
_ignore_event,
|
||||
)
|
||||
builder = self._register_optional_event(
|
||||
builder,
|
||||
"register_p2_im_chat_member_bot_deleted_v1",
|
||||
lambda _: None,
|
||||
_ignore_event,
|
||||
)
|
||||
event_handler = builder.build()
|
||||
|
||||
@@ -1126,9 +1184,11 @@ class FeishuChannel(BaseChannel):
|
||||
if response.success():
|
||||
import json
|
||||
|
||||
data = json.loads(response.raw.content)
|
||||
bot = (data.get("data") or data).get("bot") or data.get("bot") or {}
|
||||
return bot.get("open_id")
|
||||
data = _as_json_object(json.loads(response.raw.content)) or {}
|
||||
wrapped = _as_json_object(data.get("data")) or data
|
||||
bot = _as_json_object(wrapped.get("bot")) or _as_json_object(data.get("bot")) or {}
|
||||
open_id = bot.get("open_id")
|
||||
return open_id if isinstance(open_id, str) else None
|
||||
self.logger.warning("Failed to get bot info: code={}, msg={}", response.code, response.msg)
|
||||
return None
|
||||
except Exception as e:
|
||||
@@ -1218,7 +1278,7 @@ class FeishuChannel(BaseChannel):
|
||||
if "@_all" in raw_content:
|
||||
return True
|
||||
|
||||
for mention in getattr(message, "mentions", None) or []:
|
||||
for mention in cast(list[Any], getattr(message, "mentions", None) or []):
|
||||
if self._is_bot_mention_event(mention):
|
||||
return True
|
||||
return False
|
||||
@@ -1312,7 +1372,7 @@ class FeishuChannel(BaseChannel):
|
||||
loop = asyncio.get_running_loop()
|
||||
await loop.run_in_executor(None, self._remove_reaction_sync, message_id, reaction_id)
|
||||
|
||||
def _on_background_task_done(self, task: asyncio.Task) -> None:
|
||||
def _on_background_task_done(self, task: asyncio.Task[Any]) -> None:
|
||||
"""Callback: remove from tracking set and log unhandled exceptions."""
|
||||
self._background_tasks.discard(task)
|
||||
if task.cancelled():
|
||||
@@ -1322,7 +1382,7 @@ class FeishuChannel(BaseChannel):
|
||||
except Exception as exc:
|
||||
self.logger.warning("Background task failed: {}", exc)
|
||||
|
||||
def _on_reaction_added(self, message_id: str, task: asyncio.Task) -> None:
|
||||
def _on_reaction_added(self, message_id: str, task: asyncio.Task[Any]) -> None:
|
||||
"""Callback: store reaction_id after background add-reaction completes."""
|
||||
if task.cancelled():
|
||||
return
|
||||
@@ -1375,7 +1435,7 @@ class FeishuChannel(BaseChannel):
|
||||
return text
|
||||
|
||||
@classmethod
|
||||
def _parse_md_table(cls, table_text: str) -> dict | None:
|
||||
def _parse_md_table(cls, table_text: str) -> dict[str, Any] | None:
|
||||
"""Parse a markdown table into a Feishu table element."""
|
||||
lines = [_line.strip() for _line in table_text.strip().split("\n") if _line.strip()]
|
||||
if len(lines) < 3:
|
||||
@@ -1399,7 +1459,7 @@ class FeishuChannel(BaseChannel):
|
||||
],
|
||||
}
|
||||
|
||||
def _build_card_elements(self, content: str) -> list[dict]:
|
||||
def _build_card_elements(self, content: str) -> list[dict[str, Any]]:
|
||||
"""Split content into div/markdown + table elements for Feishu card."""
|
||||
protected = content
|
||||
code_blocks: list[str] = []
|
||||
@@ -1407,7 +1467,8 @@ class FeishuChannel(BaseChannel):
|
||||
code_blocks.append(m.group(1))
|
||||
protected = protected.replace(m.group(1), f"\x00CODE{len(code_blocks) - 1}\x00", 1)
|
||||
|
||||
elements, last_end = [], 0
|
||||
elements: list[dict[str, Any]] = []
|
||||
last_end = 0
|
||||
for m in self._TABLE_RE.finditer(protected):
|
||||
before = protected[last_end : m.start()]
|
||||
if before.strip():
|
||||
@@ -1429,8 +1490,8 @@ class FeishuChannel(BaseChannel):
|
||||
|
||||
@staticmethod
|
||||
def _split_elements_by_table_limit(
|
||||
elements: list[dict], max_tables: int = 1
|
||||
) -> list[list[dict]]:
|
||||
elements: list[dict[str, Any]], max_tables: int = 1
|
||||
) -> list[list[dict[str, Any]]]:
|
||||
"""Split card elements into groups with at most *max_tables* table elements each.
|
||||
|
||||
Feishu cards have a hard limit of one table per card (API error 11310).
|
||||
@@ -1439,8 +1500,8 @@ class FeishuChannel(BaseChannel):
|
||||
"""
|
||||
if not elements:
|
||||
return [[]]
|
||||
groups: list[list[dict]] = []
|
||||
current: list[dict] = []
|
||||
groups: list[list[dict[str, Any]]] = []
|
||||
current: list[dict[str, Any]] = []
|
||||
table_count = 0
|
||||
for el in elements:
|
||||
if el.get("tag") == "table":
|
||||
@@ -1457,15 +1518,15 @@ class FeishuChannel(BaseChannel):
|
||||
groups.append(current)
|
||||
return groups or [[]]
|
||||
|
||||
def _split_headings(self, content: str) -> list[dict]:
|
||||
def _split_headings(self, content: str) -> list[dict[str, Any]]:
|
||||
"""Split content by headings, converting headings to div elements."""
|
||||
protected = content
|
||||
code_blocks = []
|
||||
code_blocks: list[str] = []
|
||||
for m in self._CODE_BLOCK_RE.finditer(content):
|
||||
code_blocks.append(m.group(1))
|
||||
protected = protected.replace(m.group(1), f"\x00CODE{len(code_blocks) - 1}\x00", 1)
|
||||
|
||||
elements = []
|
||||
elements: list[dict[str, Any]] = []
|
||||
last_end = 0
|
||||
for m in self._HEADING_RE.finditer(protected):
|
||||
before = protected[last_end : m.start()].strip()
|
||||
@@ -1573,10 +1634,10 @@ class FeishuChannel(BaseChannel):
|
||||
Each line becomes a paragraph (row) in the post body.
|
||||
"""
|
||||
lines = content.strip().split("\n")
|
||||
paragraphs: list[list[dict]] = []
|
||||
paragraphs: list[list[dict[str, Any]]] = []
|
||||
|
||||
for line in lines:
|
||||
elements: list[dict] = []
|
||||
elements: list[dict[str, Any]] = []
|
||||
last_end = 0
|
||||
|
||||
for m in cls._MD_LINK_RE.finditer(line):
|
||||
@@ -1768,7 +1829,7 @@ class FeishuChannel(BaseChannel):
|
||||
return candidate
|
||||
|
||||
async def _download_and_save_media(
|
||||
self, msg_type: str, content_json: dict, message_id: str | None = None
|
||||
self, msg_type: str, content_json: dict[str, Any], message_id: str | None = None
|
||||
) -> tuple[str | None, str]:
|
||||
"""
|
||||
Download media from Feishu and save to local disk.
|
||||
@@ -2306,8 +2367,11 @@ class FeishuChannel(BaseChannel):
|
||||
fallback_msg_id = self._thread_reply_target(meta)
|
||||
if fallback_msg_id:
|
||||
await loop.run_in_executor(
|
||||
None, lambda: self._reply_message_sync(
|
||||
fallback_msg_id, "interactive", card,
|
||||
None, partial(
|
||||
self._reply_message_sync,
|
||||
fallback_msg_id,
|
||||
"interactive",
|
||||
card,
|
||||
reply_in_thread=self._should_use_reply_in_thread(meta),
|
||||
),
|
||||
)
|
||||
@@ -2563,6 +2627,9 @@ class FeishuChannel(BaseChannel):
|
||||
return
|
||||
try:
|
||||
event = data.event
|
||||
if event is None or event.message is None or event.sender is None:
|
||||
self.logger.warning("Ignoring incomplete Feishu message event")
|
||||
return
|
||||
message = event.message
|
||||
sender = event.sender
|
||||
|
||||
@@ -2579,6 +2646,20 @@ class FeishuChannel(BaseChannel):
|
||||
chat_id = message.chat_id
|
||||
chat_type = message.chat_type
|
||||
msg_type = message.message_type
|
||||
if not all(isinstance(value, str) and value for value in (
|
||||
message_id,
|
||||
sender_id,
|
||||
chat_id,
|
||||
chat_type,
|
||||
msg_type,
|
||||
)):
|
||||
self.logger.warning("Ignoring Feishu message event with missing routing fields")
|
||||
return
|
||||
message_id = cast(str, message_id)
|
||||
sender_id = cast(str, sender_id)
|
||||
chat_id = cast(str, chat_id)
|
||||
chat_type = cast(str, chat_type)
|
||||
msg_type = cast(str, msg_type)
|
||||
|
||||
if chat_type == "group" and not self._is_group_message_for_bot(message):
|
||||
self.logger.debug("skipping group message (not mentioned)")
|
||||
@@ -2616,17 +2697,19 @@ class FeishuChannel(BaseChannel):
|
||||
task.add_done_callback(lambda t: self._on_reaction_added(message_id, t))
|
||||
|
||||
# Parse content
|
||||
content_parts = []
|
||||
media_paths = []
|
||||
content_parts: list[str] = []
|
||||
media_paths: list[str] = []
|
||||
|
||||
try:
|
||||
content_json = json.loads(message.content) if message.content else {}
|
||||
raw_content = message.content if isinstance(message.content, str) else ""
|
||||
content_json = _as_json_object(json.loads(raw_content)) if raw_content else {}
|
||||
except json.JSONDecodeError:
|
||||
content_json = {}
|
||||
content_json = content_json or {}
|
||||
|
||||
if msg_type == "text":
|
||||
text = content_json.get("text", "")
|
||||
if text:
|
||||
if isinstance(text, str) and text:
|
||||
mentions = getattr(message, "mentions", None)
|
||||
text = self._strip_leading_bot_mention(text, mentions)
|
||||
text = self._resolve_mentions(text, mentions)
|
||||
@@ -2676,9 +2759,12 @@ class FeishuChannel(BaseChannel):
|
||||
content_parts.append(MSG_TYPE_MAP.get(msg_type, f"[{msg_type}]"))
|
||||
|
||||
# Extract reply context (parent/root message IDs)
|
||||
parent_id = getattr(message, "parent_id", None) or None
|
||||
root_id = getattr(message, "root_id", None) or None
|
||||
thread_id = getattr(message, "thread_id", None) or None
|
||||
parent_id = getattr(message, "parent_id", None)
|
||||
root_id = getattr(message, "root_id", None)
|
||||
thread_id = getattr(message, "thread_id", None)
|
||||
parent_id = parent_id if isinstance(parent_id, str) else None
|
||||
root_id = root_id if isinstance(root_id, str) else None
|
||||
thread_id = thread_id if isinstance(thread_id, str) else None
|
||||
|
||||
# Prepend quoted message text when the user replied to another message
|
||||
if parent_id and self._client:
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportMissingTypeStubs=false, reportPrivateUsage=false
|
||||
"""Shared Feishu/Lark WebSocket runtime.
|
||||
|
||||
The official lark_oapi websocket client stores an asyncio loop in a module-level
|
||||
@@ -148,7 +149,7 @@ class FeishuWsRunner:
|
||||
async def _client_main(
|
||||
self, key: str, client: _LarkWsClient, stop_event: asyncio.Event
|
||||
) -> None:
|
||||
ping_task: asyncio.Task | None = None
|
||||
ping_task: asyncio.Task[None] | None = None
|
||||
while not stop_event.is_set():
|
||||
try:
|
||||
await client._connect()
|
||||
@@ -171,12 +172,12 @@ class FeishuWsRunner:
|
||||
await client._disconnect()
|
||||
|
||||
|
||||
_RUNNER: FeishuWsRunner | None = None
|
||||
_runner: FeishuWsRunner | None = None
|
||||
|
||||
|
||||
def get_feishu_ws_runner() -> FeishuWsRunner:
|
||||
"""Return the process-wide Feishu WebSocket runner."""
|
||||
global _RUNNER
|
||||
if _RUNNER is None:
|
||||
_RUNNER = FeishuWsRunner()
|
||||
return _RUNNER
|
||||
global _runner
|
||||
if _runner is None:
|
||||
_runner = FeishuWsRunner()
|
||||
return _runner
|
||||
|
||||
+18
-13
@@ -8,7 +8,7 @@ import inspect
|
||||
from collections.abc import Callable, Iterable
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
@@ -41,7 +41,9 @@ from nanobot.utils.restart import (
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.cron.service import CronService
|
||||
from nanobot.session.manager import SessionManager
|
||||
from nanobot.triggers.local_store import LocalTriggerStore
|
||||
|
||||
|
||||
def _default_webui_dist() -> Path | None:
|
||||
@@ -90,8 +92,8 @@ class ChannelManager:
|
||||
bus: MessageBus,
|
||||
*,
|
||||
session_manager: "SessionManager | None" = None,
|
||||
cron_service: Any | None = None,
|
||||
local_trigger_store: Any | None = None,
|
||||
cron_service: CronService | None = None,
|
||||
local_trigger_store: LocalTriggerStore | None = None,
|
||||
webui_runtime_model_name: Callable[[], str | None] | None = None,
|
||||
webui_cron_pending_job_ids: Callable[[str], set[str]] | None = None,
|
||||
webui_local_trigger_pending_ids: Callable[[str], set[str]] | None = None,
|
||||
@@ -114,8 +116,8 @@ class ChannelManager:
|
||||
self._channel_owners: dict[str, str] = {}
|
||||
self._channel_runtime_specs: dict[str, tuple[str, str]] = {}
|
||||
self._channel_errors: dict[str, str] = {}
|
||||
self._channel_tasks: dict[str, asyncio.Task] = {}
|
||||
self._dispatch_task: asyncio.Task | None = None
|
||||
self._channel_tasks: dict[str, asyncio.Task[None]] = {}
|
||||
self._dispatch_task: asyncio.Task[None] | None = None
|
||||
self._started = False
|
||||
self._origin_reply_fingerprints: dict[tuple[str, str, str], str] = {}
|
||||
|
||||
@@ -291,10 +293,11 @@ class ChannelManager:
|
||||
for name, ch in self.channels.items():
|
||||
cfg = ch.config
|
||||
if isinstance(cfg, dict):
|
||||
if "allow_from" in cfg:
|
||||
allow = cfg.get("allow_from")
|
||||
config_data = cast(dict[str, Any], cfg)
|
||||
if "allow_from" in config_data:
|
||||
allow = config_data.get("allow_from")
|
||||
else:
|
||||
allow = cfg.get("allowFrom")
|
||||
allow = config_data.get("allowFrom")
|
||||
else:
|
||||
allow = getattr(cfg, "allow_from", None)
|
||||
if allow is None:
|
||||
@@ -321,11 +324,12 @@ class ChannelManager:
|
||||
Pydantic models.
|
||||
"""
|
||||
if isinstance(section, dict):
|
||||
value = section.get(key)
|
||||
section_data = cast(dict[str, Any], section)
|
||||
value = section_data.get(key)
|
||||
if value is None:
|
||||
camel = _BOOL_CAMEL_ALIASES.get(key)
|
||||
if camel:
|
||||
value = section.get(camel)
|
||||
value = section_data.get(camel)
|
||||
return value if isinstance(value, bool) else default
|
||||
value = getattr(section, key, None)
|
||||
return value if isinstance(value, bool) else default
|
||||
@@ -344,7 +348,7 @@ class ChannelManager:
|
||||
errors[name] = "Channel failed to start. Check gateway logs."
|
||||
logger.exception("Failed to start channel {}", name)
|
||||
|
||||
def _start_channel_task(self, name: str, channel: BaseChannel) -> asyncio.Task:
|
||||
def _start_channel_task(self, name: str, channel: BaseChannel) -> asyncio.Task[None]:
|
||||
logger.info("Starting {} channel...", name)
|
||||
task = asyncio.create_task(self._start_channel(name, channel))
|
||||
self._channel_tasks[name] = task
|
||||
@@ -361,7 +365,8 @@ class ChannelManager:
|
||||
await channel.stop()
|
||||
logger.info("Stopped {} channel", name)
|
||||
except asyncio.CancelledError:
|
||||
if asyncio.current_task() and asyncio.current_task().cancelling():
|
||||
current_task = asyncio.current_task()
|
||||
if current_task is not None and current_task.cancelling():
|
||||
raise
|
||||
logger.debug("Channel {} stop task was already cancelled", name)
|
||||
except Exception:
|
||||
@@ -553,7 +558,7 @@ class ChannelManager:
|
||||
self._dispatch_task = asyncio.create_task(self._dispatch_outbound())
|
||||
|
||||
# Start channels
|
||||
tasks = []
|
||||
tasks: list[asyncio.Task[None]] = []
|
||||
for name, channel in self.channels.items():
|
||||
tasks.append(self._start_channel_task(name, channel))
|
||||
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Matrix (Element) channel — inbound sync + outbound message/media delivery."""
|
||||
|
||||
# pyright: reportMissingTypeStubs=false
|
||||
|
||||
import asyncio
|
||||
import html
|
||||
import json
|
||||
@@ -10,7 +12,7 @@ import time
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal, TypeAlias
|
||||
from typing import Any, Callable, Literal, Protocol, TypeAlias, cast
|
||||
from urllib.parse import quote, unquote, urlparse
|
||||
|
||||
from pydantic import Field
|
||||
@@ -75,6 +77,18 @@ MATRIX_MEDIA_EVENT_FILTER = (RoomMessageMedia, RoomEncryptedMedia)
|
||||
MatrixMediaEvent: TypeAlias = RoomMessageMedia | RoomEncryptedMedia
|
||||
|
||||
|
||||
class _MatrixCallbackRegistrar(Protocol):
|
||||
"""Runtime callback surface whose upstream stubs reject valid filtered handlers."""
|
||||
|
||||
def add_event_callback(self, callback: Callable[..., Any], event_filter: Any) -> None: ...
|
||||
def add_to_device_callback(
|
||||
self,
|
||||
callback: Callable[..., Any],
|
||||
event_filter: Any,
|
||||
) -> None: ...
|
||||
def add_response_callback(self, callback: Callable[..., Any], response_filter: Any) -> None: ...
|
||||
|
||||
|
||||
class _MediaTooLargeError(Exception):
|
||||
"""Raised when an inbound Matrix media download exceeds the configured cap."""
|
||||
|
||||
@@ -187,7 +201,7 @@ def _render_markdown_html(text: str) -> str | None:
|
||||
"""Render markdown to sanitized HTML; returns None for plain text."""
|
||||
try:
|
||||
masked_text = _mask_mxc_markdown_image_sources(text)
|
||||
rendered = _mask_mxc_image_sources(MATRIX_MARKDOWN(masked_text))
|
||||
rendered = _mask_mxc_image_sources(cast(str, MATRIX_MARKDOWN(masked_text)))
|
||||
formatted = _unmask_mxc_image_sources(MATRIX_HTML_CLEANER.clean(rendered).strip())
|
||||
except Exception:
|
||||
return None
|
||||
@@ -229,16 +243,17 @@ def _build_matrix_text_content(
|
||||
content["format"] = MATRIX_HTML_FORMAT
|
||||
content["formatted_body"] = html
|
||||
if event_id:
|
||||
content["m.new_content"] = {
|
||||
new_content: dict[str, object] = {
|
||||
"body": text,
|
||||
"msgtype": "m.text",
|
||||
}
|
||||
content["m.new_content"] = new_content
|
||||
content["m.relates_to"] = {
|
||||
"rel_type": "m.replace",
|
||||
"event_id": event_id,
|
||||
}
|
||||
if thread_relates_to:
|
||||
content["m.new_content"]["m.relates_to"] = thread_relates_to
|
||||
new_content["m.relates_to"] = thread_relates_to
|
||||
elif thread_relates_to:
|
||||
content["m.relates_to"] = thread_relates_to
|
||||
|
||||
@@ -276,7 +291,7 @@ class MatrixChannel(BaseChannel):
|
||||
name = "matrix"
|
||||
display_name = "Matrix"
|
||||
_STREAM_EDIT_INTERVAL = 2 # min seconds between edit_message_text calls
|
||||
monotonic_time = time.monotonic
|
||||
monotonic_time: Callable[[], float] = staticmethod(time.monotonic)
|
||||
|
||||
@classmethod
|
||||
def default_config(cls) -> dict[str, Any]:
|
||||
@@ -294,8 +309,8 @@ class MatrixChannel(BaseChannel):
|
||||
config = MatrixConfig.model_validate(config)
|
||||
super().__init__(config, bus)
|
||||
self.client: AsyncClient | None = None
|
||||
self._sync_task: asyncio.Task | None = None
|
||||
self._typing_tasks: dict[str, asyncio.Task] = {}
|
||||
self._sync_task: asyncio.Task[None] | None = None
|
||||
self._typing_tasks: dict[str, asyncio.Task[None]] = {}
|
||||
self._restrict_to_workspace = bool(restrict_to_workspace)
|
||||
self._workspace = (
|
||||
Path(workspace).expanduser().resolve(strict=False) if workspace is not None else None
|
||||
@@ -325,7 +340,7 @@ class MatrixChannel(BaseChannel):
|
||||
self.client = AsyncClient(
|
||||
homeserver=self.config.homeserver,
|
||||
user=self.config.user_id,
|
||||
store_path=self.store_path,
|
||||
store_path=str(self.store_path),
|
||||
config=AsyncClientConfig(
|
||||
store_sync_tokens=True,
|
||||
encryption_enabled=self.config.e2ee_enabled,
|
||||
@@ -386,6 +401,16 @@ class MatrixChannel(BaseChannel):
|
||||
|
||||
self._sync_task = asyncio.create_task(self._sync_loop())
|
||||
|
||||
def _require_client(self) -> AsyncClient:
|
||||
if self.client is None:
|
||||
raise RuntimeError("Matrix client is not started")
|
||||
return self.client
|
||||
|
||||
def _callback_registrar(self) -> _MatrixCallbackRegistrar:
|
||||
# matrix-nio's callback annotations do not model filtered subtype or
|
||||
# async handlers, although the runtime API supports both.
|
||||
return cast(_MatrixCallbackRegistrar, self._require_client())
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""Stop the Matrix channel with graceful sync shutdown."""
|
||||
self._running = False
|
||||
@@ -428,9 +453,10 @@ class MatrixChannel(BaseChannel):
|
||||
seen: set[str] = set()
|
||||
candidates: list[Path] = []
|
||||
for raw in media:
|
||||
if not isinstance(raw, str) or not raw.strip():
|
||||
raw_value = cast(object, raw)
|
||||
if not isinstance(raw_value, str) or not raw_value.strip():
|
||||
continue
|
||||
path = Path(raw.strip()).expanduser()
|
||||
path = Path(raw_value.strip()).expanduser()
|
||||
try:
|
||||
key = str(path.resolve(strict=False))
|
||||
except OSError:
|
||||
@@ -535,8 +561,13 @@ class MatrixChannel(BaseChannel):
|
||||
self.logger.error("Matrix media upload failed for %s", filename, exc_info=True)
|
||||
return fail
|
||||
|
||||
upload_response = upload_result[0] if isinstance(upload_result, tuple) else upload_result
|
||||
encryption_info = upload_result[1] if isinstance(upload_result, tuple) and isinstance(upload_result[1], dict) else None
|
||||
is_tuple_result = isinstance(cast(object, upload_result), tuple)
|
||||
upload_response = upload_result[0] if is_tuple_result else upload_result
|
||||
encryption_info = (
|
||||
upload_result[1]
|
||||
if is_tuple_result and isinstance(cast(object, upload_result[1]), dict)
|
||||
else None
|
||||
)
|
||||
if isinstance(upload_response, UploadError):
|
||||
return fail
|
||||
mxc_url = getattr(upload_response, "content_uri", None)
|
||||
@@ -645,28 +676,31 @@ class MatrixChannel(BaseChannel):
|
||||
buf.last_edit = now
|
||||
if not buf.event_id:
|
||||
# we are editing the same message all the time, so only the first time the event id needs to be set
|
||||
buf.event_id = response.event_id
|
||||
buf.event_id = cast(RoomSendResponse, response).event_id
|
||||
except Exception:
|
||||
self.logger.error("Stream send/edit failed for chat_id=%s", chat_id, exc_info=True)
|
||||
await self._stop_typing_keepalive(chat_id, clear_typing=True)
|
||||
|
||||
|
||||
def _register_event_callbacks(self) -> None:
|
||||
self.client.add_event_callback(self._on_message, RoomMessageText)
|
||||
self.client.add_event_callback(self._on_media_message, MATRIX_MEDIA_EVENT_FILTER)
|
||||
self.client.add_event_callback(self._on_room_invite, InviteEvent)
|
||||
client = self._callback_registrar()
|
||||
client.add_event_callback(self._on_message, RoomMessageText)
|
||||
client.add_event_callback(self._on_media_message, MATRIX_MEDIA_EVENT_FILTER)
|
||||
client.add_event_callback(self._on_room_invite, InviteEvent)
|
||||
|
||||
def _register_to_device_callbacks(self) -> None:
|
||||
if self.config.e2ee_enabled and self.config.sas_verification:
|
||||
self.client.add_to_device_callback(
|
||||
client = self._callback_registrar()
|
||||
client.add_to_device_callback(
|
||||
self._on_key_verification_event,
|
||||
(KeyVerificationEvent,),
|
||||
)
|
||||
|
||||
def _register_response_callbacks(self) -> None:
|
||||
self.client.add_response_callback(self._on_sync_error, SyncError)
|
||||
self.client.add_response_callback(self._on_join_error, JoinError)
|
||||
self.client.add_response_callback(self._on_send_error, RoomSendError)
|
||||
client = self._callback_registrar()
|
||||
client.add_response_callback(self._on_sync_error, SyncError)
|
||||
client.add_response_callback(self._on_join_error, JoinError)
|
||||
client.add_response_callback(self._on_send_error, RoomSendError)
|
||||
|
||||
def _is_sas_sender_allowed(self, sender: str) -> bool:
|
||||
return bool(sender and self.is_allowed(sender))
|
||||
@@ -791,7 +825,8 @@ class MatrixChannel(BaseChannel):
|
||||
backoff = 2.0
|
||||
while self._running:
|
||||
try:
|
||||
await self.client.sync_forever(timeout=30000, full_state=True)
|
||||
client = self._require_client()
|
||||
await client.sync_forever(timeout=30000, full_state=True)
|
||||
backoff = 2.0
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
@@ -803,7 +838,8 @@ class MatrixChannel(BaseChannel):
|
||||
|
||||
async def _on_room_invite(self, room: MatrixRoom, event: InviteEvent) -> None:
|
||||
if self.is_allowed(event.sender):
|
||||
await self.client.join(room.room_id)
|
||||
client = self._require_client()
|
||||
await client.join(room.room_id)
|
||||
|
||||
def _is_direct_room(self, room: MatrixRoom) -> bool:
|
||||
count = getattr(room, "member_count", None)
|
||||
@@ -814,13 +850,19 @@ class MatrixChannel(BaseChannel):
|
||||
source = getattr(event, "source", None)
|
||||
if not isinstance(source, dict):
|
||||
return False
|
||||
mentions = (source.get("content") or {}).get("m.mentions")
|
||||
source_data = cast(dict[str, Any], source)
|
||||
content = cast(dict[str, Any], source_data.get("content") or {})
|
||||
mentions = cast(object, content.get("m.mentions"))
|
||||
if not isinstance(mentions, dict):
|
||||
return False
|
||||
user_ids = mentions.get("user_ids")
|
||||
mentions_data = cast(dict[str, Any], mentions)
|
||||
user_ids = cast(object, mentions_data.get("user_ids"))
|
||||
if isinstance(user_ids, list) and self.config.user_id in user_ids:
|
||||
return True
|
||||
return bool(self.config.allow_room_mentions and mentions.get("room") is True)
|
||||
return bool(
|
||||
self.config.allow_room_mentions
|
||||
and mentions_data.get("room") is True
|
||||
)
|
||||
|
||||
def _is_pre_startup_event(self, event: RoomMessage) -> bool:
|
||||
"""Skip events that landed in the timeline before this process started.
|
||||
@@ -855,14 +897,21 @@ class MatrixChannel(BaseChannel):
|
||||
source = getattr(event, "source", None)
|
||||
if not isinstance(source, dict):
|
||||
return {}
|
||||
content = source.get("content")
|
||||
return content if isinstance(content, dict) else {}
|
||||
source_data = cast(dict[str, Any], source)
|
||||
content = cast(object, source_data.get("content"))
|
||||
return cast(dict[str, Any], content) if isinstance(content, dict) else {}
|
||||
|
||||
def _event_thread_root_id(self, event: RoomMessage) -> str | None:
|
||||
relates_to = self._event_source_content(event).get("m.relates_to")
|
||||
if not isinstance(relates_to, dict) or relates_to.get("rel_type") != "m.thread":
|
||||
relates_to = cast(
|
||||
object,
|
||||
self._event_source_content(event).get("m.relates_to"),
|
||||
)
|
||||
if not isinstance(relates_to, dict):
|
||||
return None
|
||||
root_id = relates_to.get("event_id")
|
||||
relation = cast(dict[str, Any], relates_to)
|
||||
if relation.get("rel_type") != "m.thread":
|
||||
return None
|
||||
root_id = cast(object, relation.get("event_id"))
|
||||
return root_id if isinstance(root_id, str) and root_id else None
|
||||
|
||||
def _thread_metadata(self, event: RoomMessage) -> dict[str, str] | None:
|
||||
@@ -888,7 +937,7 @@ class MatrixChannel(BaseChannel):
|
||||
|
||||
def _event_attachment_type(self, event: MatrixMediaEvent) -> str:
|
||||
msgtype = self._event_source_content(event).get("msgtype")
|
||||
return _MSGTYPE_MAP.get(msgtype, "file")
|
||||
return _MSGTYPE_MAP.get(cast(str, msgtype), "file")
|
||||
|
||||
@staticmethod
|
||||
def _is_encrypted_media_event(event: MatrixMediaEvent) -> bool:
|
||||
@@ -897,16 +946,27 @@ class MatrixChannel(BaseChannel):
|
||||
and isinstance(getattr(event, "iv", None), str))
|
||||
|
||||
def _event_declared_size_bytes(self, event: MatrixMediaEvent) -> int | None:
|
||||
info = self._event_source_content(event).get("info")
|
||||
size = info.get("size") if isinstance(info, dict) else None
|
||||
info = cast(object, self._event_source_content(event).get("info"))
|
||||
size = (
|
||||
cast(dict[str, Any], info).get("size")
|
||||
if isinstance(info, dict)
|
||||
else None
|
||||
)
|
||||
return size if type(size) is int and size >= 0 else None # noqa: E721
|
||||
|
||||
def _event_mime(self, event: MatrixMediaEvent) -> str | None:
|
||||
info = self._event_source_content(event).get("info")
|
||||
if isinstance(info, dict) and isinstance(m := info.get("mimetype"), str) and m:
|
||||
return m
|
||||
m = getattr(event, "mimetype", None)
|
||||
return m if isinstance(m, str) and m else None
|
||||
info = cast(object, self._event_source_content(event).get("info"))
|
||||
if (
|
||||
isinstance(info, dict)
|
||||
and isinstance(
|
||||
mime := cast(dict[str, Any], info).get("mimetype"),
|
||||
str,
|
||||
)
|
||||
and mime
|
||||
):
|
||||
return mime
|
||||
mime = getattr(event, "mimetype", None)
|
||||
return mime if isinstance(mime, str) and mime else None
|
||||
|
||||
def _event_filename(self, event: MatrixMediaEvent, attachment_type: str) -> str:
|
||||
body = getattr(event, "body", None)
|
||||
@@ -973,9 +1033,21 @@ class MatrixChannel(BaseChannel):
|
||||
|
||||
def _decrypt_media_bytes(self, event: MatrixMediaEvent, ciphertext: bytes) -> bytes | None:
|
||||
key_obj, hashes, iv = getattr(event, "key", None), getattr(event, "hashes", None), getattr(event, "iv", None)
|
||||
key = key_obj.get("k") if isinstance(key_obj, dict) else None
|
||||
sha256 = hashes.get("sha256") if isinstance(hashes, dict) else None
|
||||
if not all(isinstance(v, str) for v in (key, sha256, iv)):
|
||||
key = (
|
||||
cast(dict[str, Any], key_obj).get("k")
|
||||
if isinstance(key_obj, dict)
|
||||
else None
|
||||
)
|
||||
sha256 = (
|
||||
cast(dict[str, Any], hashes).get("sha256")
|
||||
if isinstance(hashes, dict)
|
||||
else None
|
||||
)
|
||||
if (
|
||||
not isinstance(key, str)
|
||||
or not isinstance(sha256, str)
|
||||
or not isinstance(iv, str)
|
||||
):
|
||||
return None
|
||||
try:
|
||||
return decrypt_attachment(ciphertext, key, sha256, iv)
|
||||
|
||||
@@ -6,7 +6,7 @@ import asyncio
|
||||
import json
|
||||
import re
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
import httpx
|
||||
from pydantic import Field
|
||||
@@ -86,7 +86,7 @@ class MattermostChannel(BaseChannel):
|
||||
self._server_url = config.server_url.rstrip("/")
|
||||
self._ws_url = _server_url_to_ws_url(self._server_url)
|
||||
self._http_client: httpx.AsyncClient | None = None
|
||||
self._ws_task: asyncio.Task | None = None
|
||||
self._ws_task: asyncio.Task[None] | None = None
|
||||
self._self_id: str | None = None
|
||||
self._self_username: str | None = None
|
||||
self._self_email: str | None = None
|
||||
@@ -118,7 +118,7 @@ class MattermostChannel(BaseChannel):
|
||||
try:
|
||||
resp = await self._http_client.get("/api/v4/users/me")
|
||||
resp.raise_for_status()
|
||||
me = resp.json()
|
||||
me = cast(dict[str, Any], resp.json())
|
||||
self._self_id = me.get("id")
|
||||
self._self_username = me.get("username")
|
||||
self._self_email = me.get("email", "")
|
||||
@@ -169,7 +169,7 @@ class MattermostChannel(BaseChannel):
|
||||
self.logger.debug("websocket connected")
|
||||
delay = MATTERMOST_WS_RECONNECT_BASE_DELAY
|
||||
async for raw in ws:
|
||||
await self._handle_ws_message(json.loads(raw))
|
||||
await self._handle_ws_message(cast(dict[str, Any], json.loads(raw)))
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
@@ -191,12 +191,15 @@ class MattermostChannel(BaseChannel):
|
||||
# Event: posted ------------------------------------------------------------
|
||||
|
||||
async def _handle_posted_event(self, msg: dict[str, Any]) -> None:
|
||||
data = msg.get("data", {})
|
||||
broadcast = msg.get("broadcast", {})
|
||||
data = cast(dict[str, Any], msg.get("data", {}))
|
||||
broadcast = cast(dict[str, Any], msg.get("broadcast", {}))
|
||||
|
||||
raw_post = data.get("post", "{}")
|
||||
try:
|
||||
post = json.loads(raw_post) if isinstance(raw_post, str) else raw_post
|
||||
post = cast(
|
||||
dict[str, Any],
|
||||
json.loads(raw_post) if isinstance(raw_post, str) else raw_post,
|
||||
)
|
||||
except json.JSONDecodeError:
|
||||
self.logger.warning("failed to parse post json")
|
||||
return
|
||||
@@ -206,7 +209,7 @@ class MattermostChannel(BaseChannel):
|
||||
message_text = post.get("message", "")
|
||||
root_id = post.get("root_id", "") or ""
|
||||
post_id = post.get("id", "")
|
||||
file_ids: list[str] = post.get("file_ids", [])
|
||||
file_ids = cast(list[str], post.get("file_ids", []))
|
||||
|
||||
if self._self_id and sender_id == self._self_id:
|
||||
return
|
||||
@@ -292,11 +295,11 @@ class MattermostChannel(BaseChannel):
|
||||
# Event: action ------------------------------------------------------------
|
||||
|
||||
async def _handle_action_event(self, msg: dict[str, Any]) -> None:
|
||||
data = msg.get("data", {})
|
||||
data = cast(dict[str, Any], msg.get("data", {}))
|
||||
sender_id = data.get("user_id", "")
|
||||
channel_id = data.get("channel_id", "")
|
||||
context = data.get("context", {}) or {}
|
||||
value = context.get("selected_option", "")
|
||||
context = cast(dict[str, Any], data.get("context", {}) or {})
|
||||
value = cast(str, context.get("selected_option", ""))
|
||||
|
||||
if not sender_id or not channel_id or not value:
|
||||
return
|
||||
@@ -319,10 +322,13 @@ class MattermostChannel(BaseChannel):
|
||||
# Event: post_deleted ------------------------------------------------------
|
||||
|
||||
async def _handle_post_deleted_event(self, msg: dict[str, Any]) -> None:
|
||||
data = msg.get("data", {})
|
||||
data = cast(dict[str, Any], msg.get("data", {}))
|
||||
raw_post = data.get("post", "{}")
|
||||
try:
|
||||
post = json.loads(raw_post) if isinstance(raw_post, str) else raw_post
|
||||
post = cast(
|
||||
dict[str, Any],
|
||||
json.loads(raw_post) if isinstance(raw_post, str) else raw_post,
|
||||
)
|
||||
except json.JSONDecodeError:
|
||||
return
|
||||
post_id = post.get("id", "")
|
||||
@@ -363,15 +369,15 @@ class MattermostChannel(BaseChannel):
|
||||
return chat_id in self.config.group_allow_from
|
||||
return False
|
||||
|
||||
_BOT_MENTION_RE: re.Pattern | None = None
|
||||
_bot_mention_re: re.Pattern[str] | None = None
|
||||
|
||||
def _is_mentioned(self, text: str) -> bool:
|
||||
if not self._self_username:
|
||||
return False
|
||||
if self._BOT_MENTION_RE is None:
|
||||
if self._bot_mention_re is None:
|
||||
pat = r"(?<![@\w])@" + re.escape(self._self_username) + r"(?![@\w])"
|
||||
self._BOT_MENTION_RE = re.compile(pat)
|
||||
return bool(self._BOT_MENTION_RE.search(text))
|
||||
self._bot_mention_re = re.compile(pat)
|
||||
return bool(self._bot_mention_re.search(text))
|
||||
|
||||
def _strip_bot_mention(self, text: str) -> str:
|
||||
if not text or not self._self_username:
|
||||
@@ -432,8 +438,8 @@ class MattermostChannel(BaseChannel):
|
||||
self.logger.warning("thread context unavailable for {}: {}", key, e)
|
||||
return text
|
||||
|
||||
posts = data.get("posts", {})
|
||||
order = data.get("order", [])
|
||||
posts = cast(dict[str, dict[str, Any]], data.get("posts", {}))
|
||||
order = cast(list[str], data.get("order", []))
|
||||
if not order:
|
||||
return text
|
||||
|
||||
@@ -467,8 +473,11 @@ class MattermostChannel(BaseChannel):
|
||||
try:
|
||||
chat_id = msg.chat_id
|
||||
meta = msg.metadata or {}
|
||||
mm_meta = meta.get("mattermost", {}) or {}
|
||||
root_id = mm_meta.get("root_id") or mm_meta.get("thread_ts") or meta.get("root_id")
|
||||
mm_meta = cast(dict[str, Any], meta.get("mattermost", {}) or {})
|
||||
root_id = cast(
|
||||
str | None,
|
||||
mm_meta.get("root_id") or mm_meta.get("thread_ts") or meta.get("root_id"),
|
||||
)
|
||||
|
||||
file_ids: list[str] = []
|
||||
for media_path in msg.media or []:
|
||||
@@ -521,7 +530,7 @@ class MattermostChannel(BaseChannel):
|
||||
return
|
||||
|
||||
meta = metadata or {}
|
||||
stream_id = stream_id or meta.get("_stream_id") or chat_id
|
||||
stream_id = cast(str, stream_id or meta.get("_stream_id") or chat_id)
|
||||
stream_end = stream_end or bool(meta.get("_stream_end"))
|
||||
resuming = resuming or bool(meta.get("_resuming"))
|
||||
|
||||
@@ -541,13 +550,17 @@ class MattermostChannel(BaseChannel):
|
||||
return
|
||||
|
||||
if final and not meta.get("_progress"):
|
||||
mm_meta = (meta.get("mattermost", {}) or {}) if isinstance(meta.get("mattermost"), dict) else {}
|
||||
root_id = (
|
||||
mm_meta = (
|
||||
cast(dict[str, Any], meta.get("mattermost", {}) or {})
|
||||
if isinstance(meta.get("mattermost"), dict)
|
||||
else {}
|
||||
)
|
||||
root_id = cast(str | None, (
|
||||
mm_meta.get("root_id")
|
||||
or mm_meta.get("thread_ts")
|
||||
or meta.get("root_id")
|
||||
or self._stream_root_ids.get(stream_id)
|
||||
)
|
||||
))
|
||||
chunks = split_message(final, MATTERMOST_MAX_MESSAGE_LEN)
|
||||
first_post_id: str | None = None
|
||||
try:
|
||||
@@ -579,8 +592,15 @@ class MattermostChannel(BaseChannel):
|
||||
if not delta.strip():
|
||||
return
|
||||
|
||||
mm_meta = (meta.get("mattermost", {}) or {}) if isinstance(meta.get("mattermost"), dict) else {}
|
||||
root_id = mm_meta.get("root_id") or mm_meta.get("thread_ts") or meta.get("root_id")
|
||||
mm_meta = (
|
||||
cast(dict[str, Any], meta.get("mattermost", {}) or {})
|
||||
if isinstance(meta.get("mattermost"), dict)
|
||||
else {}
|
||||
)
|
||||
root_id = cast(
|
||||
str | None,
|
||||
mm_meta.get("root_id") or mm_meta.get("thread_ts") or meta.get("root_id"),
|
||||
)
|
||||
if root_id:
|
||||
self._stream_root_ids[stream_id] = root_id
|
||||
committed = self._stream_committed.get(stream_id, "")
|
||||
@@ -598,20 +618,25 @@ class MattermostChannel(BaseChannel):
|
||||
|
||||
# API helpers ---------------------------------------------------------------
|
||||
|
||||
def _require_http_client(self) -> httpx.AsyncClient:
|
||||
if self._http_client is None:
|
||||
raise RuntimeError("Mattermost client is not started")
|
||||
return self._http_client
|
||||
|
||||
async def _api_get(self, path: str) -> dict[str, Any]:
|
||||
resp = await self._http_client.get(path)
|
||||
resp = await self._require_http_client().get(path)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
return cast(dict[str, Any], resp.json())
|
||||
|
||||
async def _api_post(self, path: str, json_data: dict[str, Any]) -> dict[str, Any]:
|
||||
resp = await self._http_client.post(path, json=json_data)
|
||||
resp = await self._require_http_client().post(path, json=json_data)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
return cast(dict[str, Any], resp.json())
|
||||
|
||||
async def _api_put(self, path: str, json_data: dict[str, Any]) -> dict[str, Any]:
|
||||
resp = await self._http_client.put(path, json=json_data)
|
||||
resp = await self._require_http_client().put(path, json=json_data)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
return cast(dict[str, Any], resp.json())
|
||||
|
||||
async def _create_post(
|
||||
self,
|
||||
@@ -642,14 +667,14 @@ class MattermostChannel(BaseChannel):
|
||||
|
||||
try:
|
||||
files = {"files": (path.name, path.read_bytes())}
|
||||
resp = await self._http_client.post(
|
||||
resp = await self._require_http_client().post(
|
||||
"/api/v4/files",
|
||||
data={"channel_id": channel_id},
|
||||
files=files,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
infos = data.get("file_infos", [])
|
||||
data = cast(dict[str, Any], resp.json())
|
||||
infos = cast(list[dict[str, Any]], data.get("file_infos", []))
|
||||
if infos:
|
||||
return infos[0].get("id")
|
||||
except Exception as e:
|
||||
@@ -658,14 +683,15 @@ class MattermostChannel(BaseChannel):
|
||||
|
||||
async def _download_file(self, file_id: str) -> str | None:
|
||||
try:
|
||||
info_resp = await self._http_client.get(f"/api/v4/files/{file_id}/info")
|
||||
client = self._require_http_client()
|
||||
info_resp = await client.get(f"/api/v4/files/{file_id}/info")
|
||||
info_resp.raise_for_status()
|
||||
info = info_resp.json()
|
||||
info = cast(dict[str, Any], info_resp.json())
|
||||
name = Path(info.get("name", file_id)).name
|
||||
out = Path(get_media_dir("mattermost")) / safe_filename(f"{file_id}_{name}")
|
||||
out.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
dl = await self._http_client.get(f"/api/v4/files/{file_id}")
|
||||
dl = await client.get(f"/api/v4/files/{file_id}")
|
||||
dl.raise_for_status()
|
||||
out.write_bytes(dl.content)
|
||||
return str(out)
|
||||
@@ -685,7 +711,7 @@ class MattermostChannel(BaseChannel):
|
||||
async def _remove_reaction(self, post_id: str, emoji: str) -> None:
|
||||
if not self._self_id or not emoji:
|
||||
return
|
||||
resp = await self._http_client.delete(
|
||||
resp = await self._require_http_client().delete(
|
||||
f"/api/v4/users/{self._self_id}/posts/{post_id}/reactions/{emoji}",
|
||||
)
|
||||
if resp.status_code >= 400:
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportConstantRedefinition=false, reportMissingTypeStubs=false
|
||||
"""Mochat channel implementation using Socket.IO with HTTP polling fallback."""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -5,10 +6,11 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
import json
|
||||
from collections import deque
|
||||
from collections.abc import Awaitable, Callable
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
import httpx
|
||||
from pydantic import Field
|
||||
@@ -27,7 +29,7 @@ except ImportError:
|
||||
SOCKETIO_AVAILABLE = False
|
||||
|
||||
try:
|
||||
import msgpack # noqa: F401
|
||||
import msgpack # noqa: F401 # pyright: ignore[reportUnusedImport]
|
||||
MSGPACK_AVAILABLE = True
|
||||
except ImportError:
|
||||
MSGPACK_AVAILABLE = False
|
||||
@@ -57,7 +59,7 @@ class DelayState:
|
||||
"""Per-target delayed message state."""
|
||||
entries: list[MochatBufferedEntry] = field(default_factory=list)
|
||||
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
|
||||
timer: asyncio.Task | None = None
|
||||
timer: asyncio.Task[None] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -71,12 +73,12 @@ class MochatTarget:
|
||||
# Pure helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _safe_dict(value: Any) -> dict:
|
||||
def _safe_dict(value: Any) -> dict[str, Any]:
|
||||
"""Return *value* if it's a dict, else empty dict."""
|
||||
return value if isinstance(value, dict) else {}
|
||||
return cast(dict[str, Any], value) if isinstance(value, dict) else {}
|
||||
|
||||
|
||||
def _str_field(src: dict, *keys: str) -> str:
|
||||
def _str_field(src: dict[str, Any], *keys: str) -> str:
|
||||
"""Return the first non-empty str value found for *keys*, stripped."""
|
||||
for k in keys:
|
||||
v = src.get(k)
|
||||
@@ -100,7 +102,7 @@ def _make_synthetic_event(
|
||||
payload["authorInfo"] = _safe_dict(author_info)
|
||||
return {
|
||||
"type": "message.add",
|
||||
"timestamp": timestamp or datetime.utcnow().isoformat(),
|
||||
"timestamp": timestamp or datetime.utcnow().isoformat(), # pyright: ignore[reportDeprecated]
|
||||
"payload": payload,
|
||||
}
|
||||
|
||||
@@ -141,11 +143,12 @@ def extract_mention_ids(value: Any) -> list[str]:
|
||||
if not isinstance(value, list):
|
||||
return []
|
||||
ids: list[str] = []
|
||||
for item in value:
|
||||
for item in cast(list[object], value):
|
||||
if isinstance(item, str):
|
||||
if item.strip():
|
||||
ids.append(item.strip())
|
||||
elif isinstance(item, dict):
|
||||
item = cast(dict[str, Any], item)
|
||||
for key in ("id", "userId", "_id"):
|
||||
candidate = item.get(key)
|
||||
if isinstance(candidate, str) and candidate.strip():
|
||||
@@ -158,6 +161,7 @@ def resolve_was_mentioned(payload: dict[str, Any], agent_user_id: str) -> bool:
|
||||
"""Resolve mention state from payload metadata and text fallback."""
|
||||
meta = payload.get("meta")
|
||||
if isinstance(meta, dict):
|
||||
meta = cast(dict[str, Any], meta)
|
||||
if meta.get("mentioned") is True or meta.get("wasMentioned") is True:
|
||||
return True
|
||||
for f in ("mentions", "mentionIds", "mentionedUserIds", "mentionedUsers"):
|
||||
@@ -278,7 +282,7 @@ class MochatChannel(BaseChannel):
|
||||
self._state_dir = get_runtime_subdir("mochat")
|
||||
self._cursor_path = self._state_dir / "session_cursors.json"
|
||||
self._session_cursor: dict[str, int] = {}
|
||||
self._cursor_save_task: asyncio.Task | None = None
|
||||
self._cursor_save_task: asyncio.Task[None] | None = None
|
||||
|
||||
self._session_set: set[str] = set()
|
||||
self._panel_set: set[str] = set()
|
||||
@@ -292,9 +296,9 @@ class MochatChannel(BaseChannel):
|
||||
self._delay_states: dict[str, DelayState] = {}
|
||||
|
||||
self._fallback_mode = False
|
||||
self._session_fallback_tasks: dict[str, asyncio.Task] = {}
|
||||
self._panel_fallback_tasks: dict[str, asyncio.Task] = {}
|
||||
self._refresh_task: asyncio.Task | None = None
|
||||
self._session_fallback_tasks: dict[str, asyncio.Task[None]] = {}
|
||||
self._panel_fallback_tasks: dict[str, asyncio.Task[None]] = {}
|
||||
self._refresh_task: asyncio.Task[None] | None = None
|
||||
self._target_locks: dict[str, asyncio.Lock] = {}
|
||||
|
||||
# ---- lifecycle ---------------------------------------------------------
|
||||
@@ -352,7 +356,11 @@ class MochatChannel(BaseChannel):
|
||||
|
||||
parts = ([msg.content.strip()] if msg.content and msg.content.strip() else [])
|
||||
if msg.media:
|
||||
parts.extend(m for m in msg.media if isinstance(m, str) and m.strip())
|
||||
parts.extend(
|
||||
m
|
||||
for m in msg.media
|
||||
if isinstance(cast(object, m), str) and m.strip()
|
||||
)
|
||||
content = "\n".join(parts).strip()
|
||||
if not content:
|
||||
return
|
||||
@@ -404,7 +412,8 @@ class MochatChannel(BaseChannel):
|
||||
else:
|
||||
self.logger.warning("msgpack not installed but socket_disable_msgpack=false; using JSON")
|
||||
|
||||
client = socketio.AsyncClient(
|
||||
socketio_module = cast(Any, socketio)
|
||||
client: Any = socketio_module.AsyncClient(
|
||||
reconnection=True,
|
||||
reconnection_attempts=self.config.max_retry_attempts or None,
|
||||
reconnection_delay=max(0.1, self.config.socket_reconnect_delay_ms / 1000.0),
|
||||
@@ -412,7 +421,6 @@ class MochatChannel(BaseChannel):
|
||||
logger=False, engineio_logger=False, serializer=serializer,
|
||||
)
|
||||
|
||||
@client.event
|
||||
async def connect() -> None:
|
||||
self._ws_connected, self._ws_ready = True, False
|
||||
self.logger.info("websocket connected")
|
||||
@@ -420,7 +428,6 @@ class MochatChannel(BaseChannel):
|
||||
self._ws_ready = subscribed
|
||||
await (self._stop_fallback_workers() if subscribed else self._ensure_fallback_workers())
|
||||
|
||||
@client.event
|
||||
async def disconnect() -> None:
|
||||
if not self._running:
|
||||
return
|
||||
@@ -428,18 +435,21 @@ class MochatChannel(BaseChannel):
|
||||
self.logger.warning("websocket disconnected")
|
||||
await self._ensure_fallback_workers()
|
||||
|
||||
@client.event
|
||||
async def connect_error(data: Any) -> None:
|
||||
self.logger.error("websocket connect error: {}", data)
|
||||
|
||||
@client.on("claw.session.events")
|
||||
async def on_session_events(payload: dict[str, Any]) -> None:
|
||||
await self._handle_watch_payload(payload, "session")
|
||||
|
||||
@client.on("claw.panel.events")
|
||||
async def on_panel_events(payload: dict[str, Any]) -> None:
|
||||
await self._handle_watch_payload(payload, "panel")
|
||||
|
||||
client.event(connect)
|
||||
client.event(disconnect)
|
||||
client.event(connect_error)
|
||||
client.on("claw.session.events", on_session_events)
|
||||
client.on("claw.panel.events", on_panel_events)
|
||||
|
||||
for ev in ("notify:chat.inbox.append", "notify:chat.message.add",
|
||||
"notify:chat.message.update", "notify:chat.message.recall",
|
||||
"notify:chat.message.delete"):
|
||||
@@ -463,7 +473,10 @@ class MochatChannel(BaseChannel):
|
||||
self._socket = None
|
||||
return False
|
||||
|
||||
def _build_notify_handler(self, event_name: str):
|
||||
def _build_notify_handler(
|
||||
self,
|
||||
event_name: str,
|
||||
) -> Callable[[Any], Awaitable[None]]:
|
||||
async def handler(payload: Any) -> None:
|
||||
if event_name == "notify:chat.inbox.append":
|
||||
await self._handle_notify_inbox_append(payload)
|
||||
@@ -498,11 +511,20 @@ class MochatChannel(BaseChannel):
|
||||
data = ack.get("data")
|
||||
items: list[dict[str, Any]] = []
|
||||
if isinstance(data, list):
|
||||
items = [i for i in data if isinstance(i, dict)]
|
||||
items = [
|
||||
cast(dict[str, Any], item)
|
||||
for item in cast(list[object], data)
|
||||
if isinstance(item, dict)
|
||||
]
|
||||
elif isinstance(data, dict):
|
||||
data = cast(dict[str, Any], data)
|
||||
sessions = data.get("sessions")
|
||||
if isinstance(sessions, list):
|
||||
items = [i for i in sessions if isinstance(i, dict)]
|
||||
items = [
|
||||
cast(dict[str, Any], item)
|
||||
for item in cast(list[object], sessions)
|
||||
if isinstance(item, dict)
|
||||
]
|
||||
elif "sessionId" in data:
|
||||
items = [data]
|
||||
for p in items:
|
||||
@@ -525,7 +547,11 @@ class MochatChannel(BaseChannel):
|
||||
raw = await self._socket.call(event_name, payload, timeout=10)
|
||||
except Exception as e:
|
||||
return {"result": False, "message": str(e)}
|
||||
return raw if isinstance(raw, dict) else {"result": True, "data": raw}
|
||||
return (
|
||||
cast(dict[str, Any], raw)
|
||||
if isinstance(raw, dict)
|
||||
else {"result": True, "data": raw}
|
||||
)
|
||||
|
||||
# ---- refresh / discovery -----------------------------------------------
|
||||
|
||||
@@ -558,10 +584,11 @@ class MochatChannel(BaseChannel):
|
||||
return
|
||||
|
||||
new_ids: list[str] = []
|
||||
for s in sessions:
|
||||
if not isinstance(s, dict):
|
||||
for session_value in cast(list[object], sessions):
|
||||
if not isinstance(session_value, dict):
|
||||
continue
|
||||
sid = _str_field(s, "sessionId")
|
||||
session = cast(dict[str, Any], session_value)
|
||||
sid = _str_field(session, "sessionId")
|
||||
if not sid:
|
||||
continue
|
||||
if sid not in self._session_set:
|
||||
@@ -569,7 +596,7 @@ class MochatChannel(BaseChannel):
|
||||
new_ids.append(sid)
|
||||
if sid not in self._session_cursor:
|
||||
self._cold_sessions.add(sid)
|
||||
cid = _str_field(s, "converseId")
|
||||
cid = _str_field(session, "converseId")
|
||||
if cid:
|
||||
self._session_by_converse[cid] = sid
|
||||
|
||||
@@ -592,13 +619,14 @@ class MochatChannel(BaseChannel):
|
||||
return
|
||||
|
||||
new_ids: list[str] = []
|
||||
for p in raw_panels:
|
||||
if not isinstance(p, dict):
|
||||
for panel_value in cast(list[object], raw_panels):
|
||||
if not isinstance(panel_value, dict):
|
||||
continue
|
||||
pt = p.get("type")
|
||||
panel = cast(dict[str, Any], panel_value)
|
||||
pt = panel.get("type")
|
||||
if isinstance(pt, int) and pt != 0:
|
||||
continue
|
||||
pid = _str_field(p, "id", "_id")
|
||||
pid = _str_field(panel, "id", "_id")
|
||||
if pid and pid not in self._panel_set:
|
||||
self._panel_set.add(pid)
|
||||
new_ids.append(pid)
|
||||
@@ -658,16 +686,19 @@ class MochatChannel(BaseChannel):
|
||||
})
|
||||
msgs = resp.get("messages")
|
||||
if isinstance(msgs, list):
|
||||
for m in reversed(msgs):
|
||||
if not isinstance(m, dict):
|
||||
for message_value in reversed(cast(list[object], msgs)):
|
||||
if not isinstance(message_value, dict):
|
||||
continue
|
||||
message = cast(dict[str, Any], message_value)
|
||||
evt = _make_synthetic_event(
|
||||
message_id=str(m.get("messageId") or ""),
|
||||
author=str(m.get("author") or ""),
|
||||
content=m.get("content"),
|
||||
meta=m.get("meta"), group_id=str(resp.get("groupId") or ""),
|
||||
converse_id=panel_id, timestamp=m.get("createdAt"),
|
||||
author_info=m.get("authorInfo"),
|
||||
message_id=str(message.get("messageId") or ""),
|
||||
author=str(message.get("author") or ""),
|
||||
content=message.get("content"),
|
||||
meta=message.get("meta"),
|
||||
group_id=str(resp.get("groupId") or ""),
|
||||
converse_id=panel_id,
|
||||
timestamp=message.get("createdAt"),
|
||||
author_info=message.get("authorInfo"),
|
||||
)
|
||||
await self._process_inbound_event(panel_id, evt, "panel")
|
||||
except asyncio.CancelledError:
|
||||
@@ -679,7 +710,7 @@ class MochatChannel(BaseChannel):
|
||||
# ---- inbound event processing ------------------------------------------
|
||||
|
||||
async def _handle_watch_payload(self, payload: dict[str, Any], target_kind: str) -> None:
|
||||
if not isinstance(payload, dict):
|
||||
if not isinstance(cast(object, payload), dict):
|
||||
return
|
||||
target_id = _str_field(payload, "sessionId")
|
||||
if not target_id:
|
||||
@@ -699,9 +730,10 @@ class MochatChannel(BaseChannel):
|
||||
self._cold_sessions.discard(target_id)
|
||||
return
|
||||
|
||||
for event in raw_events:
|
||||
if not isinstance(event, dict):
|
||||
for event_value in cast(list[object], raw_events):
|
||||
if not isinstance(event_value, dict):
|
||||
continue
|
||||
event = cast(dict[str, Any], event_value)
|
||||
seq = event.get("seq")
|
||||
if target_kind == "session" and isinstance(seq, int) and seq > self._session_cursor.get(target_id, prev):
|
||||
self._mark_session_cursor(target_id, seq)
|
||||
@@ -712,6 +744,7 @@ class MochatChannel(BaseChannel):
|
||||
payload = event.get("payload")
|
||||
if not isinstance(payload, dict):
|
||||
return
|
||||
payload = cast(dict[str, Any], payload)
|
||||
|
||||
author = _str_field(payload, "author")
|
||||
if not author or (self.config.agent_user_id and author == self.config.agent_user_id):
|
||||
@@ -821,6 +854,7 @@ class MochatChannel(BaseChannel):
|
||||
async def _handle_notify_chat_message(self, payload: Any) -> None:
|
||||
if not isinstance(payload, dict):
|
||||
return
|
||||
payload = cast(dict[str, Any], payload)
|
||||
group_id = _str_field(payload, "groupId")
|
||||
panel_id = _str_field(payload, "converseId", "panelId")
|
||||
if not group_id or not panel_id:
|
||||
@@ -838,11 +872,15 @@ class MochatChannel(BaseChannel):
|
||||
await self._process_inbound_event(panel_id, evt, "panel")
|
||||
|
||||
async def _handle_notify_inbox_append(self, payload: Any) -> None:
|
||||
if not isinstance(payload, dict) or payload.get("type") != "message":
|
||||
if not isinstance(payload, dict):
|
||||
return
|
||||
payload = cast(dict[str, Any], payload)
|
||||
if payload.get("type") != "message":
|
||||
return
|
||||
detail = payload.get("payload")
|
||||
if not isinstance(detail, dict):
|
||||
return
|
||||
detail = cast(dict[str, Any], detail)
|
||||
if _str_field(detail, "groupId"):
|
||||
return
|
||||
converse_id = _str_field(detail, "converseId")
|
||||
@@ -886,9 +924,14 @@ class MochatChannel(BaseChannel):
|
||||
except Exception as e:
|
||||
self.logger.warning("Failed to read cursor file: {}", e)
|
||||
return
|
||||
cursors = data.get("cursors") if isinstance(data, dict) else None
|
||||
data_object = cast(object, data)
|
||||
cursors = (
|
||||
cast(dict[str, Any], data_object).get("cursors")
|
||||
if isinstance(data_object, dict)
|
||||
else None
|
||||
)
|
||||
if isinstance(cursors, dict):
|
||||
for sid, cur in cursors.items():
|
||||
for sid, cur in cast(dict[object, object], cursors).items():
|
||||
if isinstance(sid, str) and isinstance(cur, int) and cur >= 0:
|
||||
self._session_cursor[sid] = cur
|
||||
|
||||
@@ -896,7 +939,8 @@ class MochatChannel(BaseChannel):
|
||||
try:
|
||||
self._state_dir.mkdir(parents=True, exist_ok=True)
|
||||
self._cursor_path.write_text(json.dumps({
|
||||
"schemaVersion": 1, "updatedAt": datetime.utcnow().isoformat(),
|
||||
"schemaVersion": 1,
|
||||
"updatedAt": datetime.utcnow().isoformat(), # pyright: ignore[reportDeprecated]
|
||||
"cursors": self._session_cursor,
|
||||
}, ensure_ascii=False, indent=2) + "\n", "utf-8")
|
||||
except Exception as e:
|
||||
@@ -917,13 +961,22 @@ class MochatChannel(BaseChannel):
|
||||
parsed = response.json()
|
||||
except Exception:
|
||||
parsed = response.text
|
||||
if isinstance(parsed, dict) and isinstance(parsed.get("code"), int):
|
||||
if parsed["code"] != 200:
|
||||
msg = str(parsed.get("message") or parsed.get("name") or "request failed")
|
||||
raise RuntimeError(f"Mochat API error: {msg} (code={parsed['code']})")
|
||||
data = parsed.get("data")
|
||||
return data if isinstance(data, dict) else {}
|
||||
return parsed if isinstance(parsed, dict) else {}
|
||||
if isinstance(parsed, dict):
|
||||
parsed_dict = cast(dict[str, Any], parsed)
|
||||
if isinstance(parsed_dict.get("code"), int):
|
||||
if parsed_dict["code"] != 200:
|
||||
msg = str(
|
||||
parsed_dict.get("message")
|
||||
or parsed_dict.get("name")
|
||||
or "request failed"
|
||||
)
|
||||
raise RuntimeError(
|
||||
f"Mochat API error: {msg} (code={parsed_dict['code']})"
|
||||
)
|
||||
data = parsed_dict.get("data")
|
||||
return cast(dict[str, Any], data) if isinstance(data, dict) else {}
|
||||
return parsed_dict
|
||||
return {}
|
||||
|
||||
async def _api_send(self, path: str, id_key: str, id_val: str,
|
||||
content: str, reply_to: str | None, group_id: str | None = None) -> dict[str, Any]:
|
||||
@@ -937,7 +990,7 @@ class MochatChannel(BaseChannel):
|
||||
|
||||
@staticmethod
|
||||
def _read_group_id(metadata: dict[str, Any]) -> str | None:
|
||||
if not isinstance(metadata, dict):
|
||||
if not isinstance(cast(object, metadata), dict):
|
||||
return None
|
||||
value = metadata.get("group_id") or metadata.get("groupId")
|
||||
return value.strip() if isinstance(value, str) and value.strip() else None
|
||||
|
||||
@@ -23,7 +23,8 @@ import time
|
||||
from contextlib import contextmanager, suppress
|
||||
from dataclasses import dataclass
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, Generator, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
try: # pragma: no cover - Windows fallback path
|
||||
@@ -47,9 +48,11 @@ MSTEAMS_AVAILABLE = (
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import jwt
|
||||
from jwt.algorithms import RSAAlgorithm
|
||||
|
||||
if MSTEAMS_AVAILABLE:
|
||||
import jwt
|
||||
from jwt.algorithms import RSAAlgorithm
|
||||
|
||||
MSTEAMS_REF_TTL_DAYS = 30
|
||||
MSTEAMS_WEBCHAT_HOST = "webchat.botframework.com"
|
||||
@@ -182,9 +185,10 @@ class MSTeamsChannel(BaseChannel):
|
||||
auth_header = self.headers.get("Authorization", "")
|
||||
if channel.config.validate_inbound_auth:
|
||||
try:
|
||||
loop = cast(asyncio.AbstractEventLoop, channel._loop)
|
||||
fut = asyncio.run_coroutine_threadsafe(
|
||||
channel._validate_inbound_auth(auth_header, payload),
|
||||
channel._loop,
|
||||
loop,
|
||||
)
|
||||
fut.result(timeout=15)
|
||||
except Exception as e:
|
||||
@@ -195,9 +199,10 @@ class MSTeamsChannel(BaseChannel):
|
||||
self.wfile.write(b'{"error":"unauthorized"}')
|
||||
return
|
||||
try:
|
||||
loop = cast(asyncio.AbstractEventLoop, channel._loop)
|
||||
fut = asyncio.run_coroutine_threadsafe(
|
||||
channel._handle_activity(payload),
|
||||
channel._loop,
|
||||
loop,
|
||||
)
|
||||
fut.result(timeout=15)
|
||||
except Exception as e:
|
||||
@@ -269,7 +274,7 @@ class MSTeamsChannel(BaseChannel):
|
||||
"text": msg.content or " ",
|
||||
}
|
||||
if use_thread_reply:
|
||||
payload["replyToId"] = ref.activity_id
|
||||
payload["replyToId"] = cast(str, ref.activity_id)
|
||||
|
||||
try:
|
||||
resp = await self._http.post(base_url, headers=headers, json=payload)
|
||||
@@ -285,10 +290,10 @@ class MSTeamsChannel(BaseChannel):
|
||||
if activity.get("type") != "message":
|
||||
return
|
||||
|
||||
conversation = activity.get("conversation") or {}
|
||||
from_user = activity.get("from") or {}
|
||||
recipient = activity.get("recipient") or {}
|
||||
channel_data = activity.get("channelData") or {}
|
||||
conversation = cast(dict[str, Any], activity.get("conversation") or {})
|
||||
from_user = cast(dict[str, Any], activity.get("from") or {})
|
||||
recipient = cast(dict[str, Any], activity.get("recipient") or {})
|
||||
channel_data = cast(dict[str, Any], activity.get("channelData") or {})
|
||||
|
||||
sender_id = str(from_user.get("aadObjectId") or from_user.get("id") or "").strip()
|
||||
conversation_id = str(conversation.get("id") or "").strip()
|
||||
@@ -336,7 +341,16 @@ class MSTeamsChannel(BaseChannel):
|
||||
bot_id=str(recipient.get("id") or "") or None,
|
||||
activity_id=activity_id or None,
|
||||
conversation_type=conversation_type or None,
|
||||
tenant_id=str((channel_data.get("tenant") or {}).get("id") or "") or None,
|
||||
tenant_id=(
|
||||
str(
|
||||
cast(
|
||||
dict[str, Any],
|
||||
channel_data.get("tenant") or {},
|
||||
).get("id")
|
||||
or ""
|
||||
)
|
||||
or None
|
||||
),
|
||||
updated_at=time.time(),
|
||||
)
|
||||
self._save_refs_locked()
|
||||
@@ -361,7 +375,7 @@ class MSTeamsChannel(BaseChannel):
|
||||
text = self._strip_possible_bot_mention(text)
|
||||
text = self._normalize_html_whitespace(text)
|
||||
|
||||
channel_data = activity.get("channelData") or {}
|
||||
channel_data = cast(dict[str, Any], activity.get("channelData") or {})
|
||||
reply_to_id = str(activity.get("replyToId") or "").strip()
|
||||
normalized_preview = html.unescape(text).replace("&rsquo", "’").strip()
|
||||
normalized_preview = normalized_preview.replace("\xa0", " ")
|
||||
@@ -473,15 +487,15 @@ class MSTeamsChannel(BaseChannel):
|
||||
raise ValueError("missing token kid")
|
||||
|
||||
jwks = await self._get_botframework_jwks()
|
||||
keys = jwks.get("keys") or []
|
||||
keys = cast(list[dict[str, Any]], jwks.get("keys") or [])
|
||||
jwk = next((key for key in keys if key.get("kid") == kid), None)
|
||||
if not jwk:
|
||||
raise ValueError(f"signing key not found for kid={kid}")
|
||||
|
||||
public_key = jwt.algorithms.RSAAlgorithm.from_jwk(json.dumps(jwk))
|
||||
public_key = RSAAlgorithm.from_jwk(json.dumps(jwk))
|
||||
claims = jwt.decode(
|
||||
token,
|
||||
key=public_key,
|
||||
key=cast(Any, public_key),
|
||||
algorithms=["RS256"],
|
||||
audience=self.config.app_id,
|
||||
issuer="https://api.botframework.com",
|
||||
@@ -509,9 +523,10 @@ class MSTeamsChannel(BaseChannel):
|
||||
|
||||
resp = await self._http.get(self._botframework_openid_config_url)
|
||||
resp.raise_for_status()
|
||||
self._botframework_openid_config = resp.json()
|
||||
openid_config = cast(dict[str, Any], resp.json())
|
||||
self._botframework_openid_config = openid_config
|
||||
self._botframework_openid_config_expires_at = now + 3600
|
||||
return self._botframework_openid_config
|
||||
return openid_config
|
||||
|
||||
async def _get_botframework_jwks(self) -> dict[str, Any]:
|
||||
"""Fetch and cache Bot Framework JWKS."""
|
||||
@@ -530,36 +545,38 @@ class MSTeamsChannel(BaseChannel):
|
||||
|
||||
resp = await self._http.get(jwks_uri)
|
||||
resp.raise_for_status()
|
||||
self._botframework_jwks = resp.json()
|
||||
jwks = cast(dict[str, Any], resp.json())
|
||||
self._botframework_jwks = jwks
|
||||
self._botframework_jwks_expires_at = now + 3600
|
||||
return self._botframework_jwks
|
||||
return jwks
|
||||
|
||||
@staticmethod
|
||||
def _safe_float(value: Any) -> float | None:
|
||||
def _safe_float(value: object) -> float | None:
|
||||
try:
|
||||
out = float(value)
|
||||
out = float(cast(Any, value))
|
||||
if out > 0:
|
||||
return out
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
return None
|
||||
|
||||
def _normalize_ref_record(self, value: Any) -> ConversationRef | None:
|
||||
def _normalize_ref_record(self, value: object) -> ConversationRef | None:
|
||||
"""Normalize a stored ref record from legacy/current schema."""
|
||||
if not isinstance(value, dict):
|
||||
return None
|
||||
service_url = str(value.get("service_url") or "").strip()
|
||||
conversation_id = str(value.get("conversation_id") or "").strip()
|
||||
record = cast(dict[str, Any], value)
|
||||
service_url = str(record.get("service_url") or "").strip()
|
||||
conversation_id = str(record.get("conversation_id") or "").strip()
|
||||
if not service_url or not conversation_id:
|
||||
return None
|
||||
return ConversationRef(
|
||||
service_url=service_url,
|
||||
conversation_id=conversation_id,
|
||||
bot_id=str(value.get("bot_id") or "") or None,
|
||||
activity_id=str(value.get("activity_id") or "") or None,
|
||||
conversation_type=str(value.get("conversation_type") or "") or None,
|
||||
tenant_id=str(value.get("tenant_id") or "") or None,
|
||||
updated_at=self._safe_float(value.get("updated_at")),
|
||||
bot_id=str(record.get("bot_id") or "") or None,
|
||||
activity_id=str(record.get("activity_id") or "") or None,
|
||||
conversation_type=str(record.get("conversation_type") or "") or None,
|
||||
tenant_id=str(record.get("tenant_id") or "") or None,
|
||||
updated_at=self._safe_float(cast(object, record.get("updated_at"))),
|
||||
)
|
||||
|
||||
def _load_refs_raw(self) -> tuple[dict[str, Any], dict[str, Any], bool]:
|
||||
@@ -570,17 +587,19 @@ class MSTeamsChannel(BaseChannel):
|
||||
|
||||
if self._refs_path.exists():
|
||||
try:
|
||||
loaded = json.loads(self._refs_path.read_text(encoding="utf-8"))
|
||||
loaded: object = json.loads(self._refs_path.read_text(encoding="utf-8"))
|
||||
if isinstance(loaded, dict):
|
||||
main_data = loaded
|
||||
main_data = cast(dict[str, Any], loaded)
|
||||
except Exception as e:
|
||||
self.logger.warning("Failed to load conversation refs: {}", e)
|
||||
|
||||
if meta_exists:
|
||||
try:
|
||||
loaded_meta = json.loads(self._refs_meta_path.read_text(encoding="utf-8"))
|
||||
loaded_meta: object = json.loads(
|
||||
self._refs_meta_path.read_text(encoding="utf-8")
|
||||
)
|
||||
if isinstance(loaded_meta, dict):
|
||||
meta_data = loaded_meta
|
||||
meta_data = cast(dict[str, Any], loaded_meta)
|
||||
except Exception as e:
|
||||
self.logger.warning("Failed to load conversation refs metadata: {}", e)
|
||||
|
||||
@@ -599,10 +618,11 @@ class MSTeamsChannel(BaseChannel):
|
||||
if not ref:
|
||||
continue
|
||||
|
||||
meta_entry = meta_data.get(key) if isinstance(meta_data, dict) else None
|
||||
meta_ts = None
|
||||
meta_entry = cast(object, meta_data.get(key))
|
||||
meta_ts: float | None = None
|
||||
if isinstance(meta_entry, dict):
|
||||
meta_ts = self._safe_float(meta_entry.get("updated_at"))
|
||||
meta_record = cast(dict[str, Any], meta_entry)
|
||||
meta_ts = self._safe_float(cast(object, meta_record.get("updated_at")))
|
||||
elif meta_entry is not None:
|
||||
meta_ts = self._safe_float(meta_entry)
|
||||
|
||||
@@ -623,7 +643,7 @@ class MSTeamsChannel(BaseChannel):
|
||||
return self._load_refs_from_disk()
|
||||
|
||||
@contextmanager
|
||||
def _refs_file_lock(self):
|
||||
def _refs_file_lock(self) -> Generator[None, None, None]:
|
||||
"""Cross-process lock while merging and writing refs state."""
|
||||
self._refs_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
lock_fp = self._refs_lock_path.open("a+", encoding="utf-8")
|
||||
@@ -742,7 +762,7 @@ class MSTeamsChannel(BaseChannel):
|
||||
if persist:
|
||||
self._save_refs_locked()
|
||||
|
||||
def _write_json_atomically(self, path, data: dict[str, Any]) -> None:
|
||||
def _write_json_atomically(self, path: Path, data: dict[str, Any]) -> None:
|
||||
"""Write refs JSON atomically to reduce corruption risk during crashes."""
|
||||
payload = json.dumps(data, indent=2)
|
||||
tmp_path: str | None = None
|
||||
@@ -816,7 +836,8 @@ class MSTeamsChannel(BaseChannel):
|
||||
}
|
||||
resp = await self._http.post(token_url, data=data)
|
||||
resp.raise_for_status()
|
||||
payload = resp.json()
|
||||
self._token = payload["access_token"]
|
||||
payload = cast(dict[str, Any], resp.json())
|
||||
token = cast(str, payload["access_token"])
|
||||
self._token = token
|
||||
self._token_expires_at = now + int(payload.get("expires_in", 3600))
|
||||
return self._token
|
||||
return token
|
||||
|
||||
@@ -11,7 +11,7 @@ import time
|
||||
import uuid
|
||||
from collections import deque
|
||||
from pathlib import Path
|
||||
from typing import Annotated, Any, Literal
|
||||
from typing import Annotated, Any, Literal, cast
|
||||
|
||||
import aiohttp
|
||||
from loguru import logger
|
||||
@@ -103,7 +103,7 @@ class NapcatChannel(BaseChannel):
|
||||
await asyncio.sleep(next(backoff, 30))
|
||||
|
||||
async def _run_once(self) -> None:
|
||||
headers = []
|
||||
headers: list[tuple[str, str]] = []
|
||||
if self.config.access_token:
|
||||
headers.append(("Authorization", f"Bearer {self.config.access_token}"))
|
||||
|
||||
@@ -132,12 +132,17 @@ class NapcatChannel(BaseChannel):
|
||||
payload = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if isinstance(payload, dict) and payload.get("echo") == echo:
|
||||
data = payload.get("data") or {}
|
||||
if isinstance(payload, dict):
|
||||
login_payload = cast(dict[str, Any], payload)
|
||||
else:
|
||||
login_payload = None
|
||||
if login_payload is not None and login_payload.get("echo") == echo:
|
||||
data = login_payload.get("data")
|
||||
login_data = cast(dict[str, Any], data) if isinstance(data, dict) else {}
|
||||
logger.info(
|
||||
"napcat: logged in as {} (user_id={})",
|
||||
data.get("nickname"),
|
||||
data.get("user_id"),
|
||||
login_data.get("nickname"),
|
||||
login_data.get("user_id"),
|
||||
)
|
||||
break
|
||||
await self._dispatch_frame(raw)
|
||||
@@ -189,26 +194,27 @@ class NapcatChannel(BaseChannel):
|
||||
return
|
||||
if not isinstance(payload, dict):
|
||||
return
|
||||
frame = cast(dict[str, Any], payload)
|
||||
|
||||
# Action response: identified by `echo` and absence of post_type.
|
||||
if "echo" in payload and payload.get("post_type") is None:
|
||||
echo = payload.get("echo")
|
||||
if "echo" in frame and frame.get("post_type") is None:
|
||||
echo = frame.get("echo")
|
||||
fut = self._pending.pop(echo, None) if isinstance(echo, str) else None
|
||||
if fut and not fut.done():
|
||||
fut.set_result(payload)
|
||||
fut.set_result(frame)
|
||||
return
|
||||
|
||||
if (sid := payload.get("self_id")) is not None:
|
||||
if (sid := frame.get("self_id")) is not None:
|
||||
try:
|
||||
self._self_id = int(sid)
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
|
||||
post_type = payload.get("post_type")
|
||||
post_type = frame.get("post_type")
|
||||
if post_type == "message":
|
||||
self._create_background_task(self._on_message(payload), "message")
|
||||
self._create_background_task(self._on_message(frame), "message")
|
||||
elif post_type == "notice":
|
||||
self._create_background_task(self._on_notice(payload), "notice")
|
||||
self._create_background_task(self._on_notice(frame), "notice")
|
||||
|
||||
def _create_background_task(self, coro: Any, kind: str) -> None:
|
||||
task = asyncio.create_task(coro)
|
||||
@@ -249,7 +255,8 @@ class NapcatChannel(BaseChannel):
|
||||
if local := await self._download_image(info):
|
||||
media_paths.append(local)
|
||||
|
||||
sender = ev.get("sender") or {}
|
||||
sender_raw = ev.get("sender")
|
||||
sender = cast(dict[str, Any], sender_raw) if isinstance(sender_raw, dict) else {}
|
||||
nickname = sender.get("card") or sender.get("nickname")
|
||||
|
||||
if message_type == "group":
|
||||
@@ -270,7 +277,7 @@ class NapcatChannel(BaseChannel):
|
||||
chat_id = f"group:{group_id}"
|
||||
content = self._format_group_content(
|
||||
text=text,
|
||||
nickname=nickname,
|
||||
nickname=cast(str, nickname),
|
||||
user_id=user_id,
|
||||
)
|
||||
else:
|
||||
@@ -299,7 +306,7 @@ class NapcatChannel(BaseChannel):
|
||||
# segment rather than parsing CQ codes — that path is fragile and
|
||||
# users can configure napcat to emit arrays.
|
||||
if isinstance(message, list):
|
||||
return [seg for seg in message if isinstance(seg, dict)]
|
||||
return [cast(dict[str, Any], seg) for seg in cast(list[Any], message) if isinstance(seg, dict)]
|
||||
if isinstance(message, str) and message:
|
||||
return [{"type": "text", "data": {"text": message}}]
|
||||
return []
|
||||
@@ -315,7 +322,8 @@ class NapcatChannel(BaseChannel):
|
||||
|
||||
for seg in segments:
|
||||
stype = seg.get("type")
|
||||
data = seg.get("data") or {}
|
||||
raw_data = seg.get("data")
|
||||
data = cast(dict[str, Any], raw_data) if isinstance(raw_data, dict) else {}
|
||||
if stype == "text":
|
||||
if txt := data.get("text"):
|
||||
parts.append(str(txt))
|
||||
@@ -455,7 +463,8 @@ class NapcatChannel(BaseChannel):
|
||||
params["user_id"] = int(target)
|
||||
|
||||
resp = await self._call_action("send_msg", params)
|
||||
data = resp.get("data") or {}
|
||||
raw_data = resp.get("data")
|
||||
data = cast(dict[str, Any], raw_data) if isinstance(raw_data, dict) else {}
|
||||
if (mid := data.get("message_id")) is not None:
|
||||
self._bot_outbound_ids.append(int(mid))
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ import re
|
||||
from dataclasses import dataclass
|
||||
from functools import lru_cache
|
||||
from importlib.resources import files
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from packaging.requirements import InvalidRequirement, Requirement
|
||||
|
||||
@@ -49,12 +49,12 @@ class ChannelPlugin:
|
||||
_target_parts(self.runtime, label="runtime")
|
||||
if self.connector is not None:
|
||||
_target_parts(self.connector, label="connector")
|
||||
if self.setup is not None and not isinstance(self.setup, ChannelSetupSpec):
|
||||
if self.setup is not None and not isinstance(cast(object, self.setup), ChannelSetupSpec):
|
||||
raise TypeError("channel plugin setup must be a ChannelSetupSpec or None")
|
||||
if not isinstance(self.management, ChannelManagementSpec):
|
||||
if not isinstance(cast(object, self.management), ChannelManagementSpec):
|
||||
raise TypeError("channel plugin management must be a ChannelManagementSpec")
|
||||
if not isinstance(self.dependencies, tuple) or not all(
|
||||
isinstance(requirement, str) and requirement.strip()
|
||||
if not isinstance(cast(object, self.dependencies), tuple) or not all(
|
||||
isinstance(cast(object, requirement), str) and requirement.strip()
|
||||
for requirement in self.dependencies
|
||||
):
|
||||
raise TypeError("channel plugin dependencies must be a tuple of requirements")
|
||||
|
||||
@@ -16,6 +16,8 @@ Notes:
|
||||
- Attachment structures differ across botpy versions; we try multiple field candidates.
|
||||
"""
|
||||
|
||||
# pyright: reportConstantRedefinition=false, reportMissingTypeStubs=false, reportPrivateUsage=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
@@ -27,7 +29,7 @@ import time
|
||||
from collections import deque
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, Literal
|
||||
from typing import Any, BinaryIO, Literal, cast
|
||||
from urllib.parse import unquote, urlparse
|
||||
|
||||
import aiohttp
|
||||
@@ -58,11 +60,6 @@ except ImportError: # pragma: no cover
|
||||
BotWebSocket = None
|
||||
Route = None
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from botpy.message import BaseMessage, C2CMessage, GroupMessage
|
||||
from botpy.types.message import Media
|
||||
|
||||
|
||||
# QQ rich media file_type: 1=image, 4=file
|
||||
# (2=voice, 3=video are restricted; we only use image vs file)
|
||||
QQ_FILE_TYPE_IMAGE = 1
|
||||
@@ -118,30 +115,34 @@ def _is_network_error(exc: BaseException) -> bool:
|
||||
)
|
||||
|
||||
|
||||
def _make_bot_class(channel: QQChannel) -> type[botpy.Client]:
|
||||
def _make_bot_class(channel: QQChannel) -> type[Any]:
|
||||
"""Create a botpy client with per-session reconnect backoff."""
|
||||
intents = botpy.Intents(public_messages=True, direct_message=True)
|
||||
botpy_sdk = cast(Any, botpy)
|
||||
intents = botpy_sdk.Intents(public_messages=True, direct_message=True)
|
||||
|
||||
class _Bot(botpy.Client):
|
||||
class _Bot(botpy_sdk.Client):
|
||||
def __init__(self):
|
||||
# Disable botpy's file log — nanobot uses loguru; default "botpy.log" fails on read-only fs
|
||||
super().__init__(intents=intents, ext_handlers=False)
|
||||
super().__init__( # pyright: ignore[reportUnknownMemberType]
|
||||
intents=intents,
|
||||
ext_handlers=False,
|
||||
)
|
||||
self._ws_backoff: dict[int, int] = {}
|
||||
self._ws_retry_at: dict[int, float] = {}
|
||||
|
||||
async def on_ready(self):
|
||||
logger.info("QQ bot ready: {}", self.robot.name)
|
||||
|
||||
async def on_c2c_message_create(self, message: C2CMessage):
|
||||
async def on_c2c_message_create(self, message: object) -> None:
|
||||
await channel._on_message(message, is_group=False)
|
||||
|
||||
async def on_group_at_message_create(self, message: GroupMessage):
|
||||
async def on_group_at_message_create(self, message: object) -> None:
|
||||
await channel._on_message(message, is_group=True)
|
||||
|
||||
async def on_direct_message_create(self, message):
|
||||
async def on_direct_message_create(self, message: object) -> None:
|
||||
await channel._on_message(message, is_group=False)
|
||||
|
||||
async def bot_connect(self, session):
|
||||
async def bot_connect(self, session: object) -> None:
|
||||
"""Connect a botpy session with exponential retry backoff."""
|
||||
session_id = id(session)
|
||||
retry_at = self._ws_retry_at.pop(session_id, None)
|
||||
@@ -150,7 +151,8 @@ def _make_bot_class(channel: QQChannel) -> type[botpy.Client]:
|
||||
if remaining > 0:
|
||||
await asyncio.sleep(remaining)
|
||||
|
||||
client = BotWebSocket(session, self._connection)
|
||||
websocket_class = cast(Any, BotWebSocket)
|
||||
client = websocket_class(session, self._connection)
|
||||
backoff = self._ws_backoff.get(session_id, _RECONNECT_BACKOFF_START)
|
||||
try:
|
||||
await client.ws_connect()
|
||||
@@ -207,7 +209,7 @@ class QQChannel(BaseChannel):
|
||||
super().__init__(config, bus)
|
||||
self.config: QQConfig = config
|
||||
|
||||
self._client: botpy.Client | None = None
|
||||
self._client: Any | None = None
|
||||
self._http: aiohttp.ClientSession | None = None
|
||||
|
||||
self._processed_ids: deque[str] = deque(maxlen=1000)
|
||||
@@ -260,7 +262,8 @@ class QQChannel(BaseChannel):
|
||||
max_backoff = 300
|
||||
while self._running:
|
||||
try:
|
||||
await self._client.start(appid=self.config.app_id, secret=self.config.secret)
|
||||
client = cast(Any, self._client)
|
||||
await client.start(appid=self.config.app_id, secret=self.config.secret)
|
||||
backoff = 5
|
||||
except Exception as e:
|
||||
if _is_network_error(e):
|
||||
@@ -490,7 +493,7 @@ class QQChannel(BaseChannel):
|
||||
file_data: str,
|
||||
file_name: str | None = None,
|
||||
srv_send_msg: bool = False,
|
||||
) -> Media:
|
||||
) -> dict[str, Any]:
|
||||
"""Upload base64-encoded file and return Media object."""
|
||||
if not self._client:
|
||||
raise RuntimeError("QQ client not initialized")
|
||||
@@ -514,39 +517,44 @@ class QQChannel(BaseChannel):
|
||||
if file_type != QQ_FILE_TYPE_IMAGE and file_name:
|
||||
payload["file_name"] = file_name
|
||||
|
||||
route = Route("POST", endpoint, **{id_key: chat_id})
|
||||
result = await self._client.api._http.request(route, json=payload)
|
||||
route_class = cast(Any, Route)
|
||||
route = route_class("POST", endpoint, **{id_key: chat_id})
|
||||
client = self._client
|
||||
result: object = await client.api._http.request(route, json=payload)
|
||||
|
||||
# Extract only the file_info field to avoid extra fields (file_uuid, ttl, etc.)
|
||||
# that may confuse QQ client when sending the media object.
|
||||
if isinstance(result, dict) and "file_info" in result:
|
||||
return {"file_info": result["file_info"]}
|
||||
return result
|
||||
result_data = cast(dict[str, Any], result)
|
||||
return {"file_info": result_data["file_info"]}
|
||||
return cast(dict[str, Any], result)
|
||||
|
||||
# ---------------------------
|
||||
# Inbound (receive)
|
||||
# ---------------------------
|
||||
|
||||
async def _on_message(self, data: C2CMessage | GroupMessage, is_group: bool = False) -> None:
|
||||
async def _on_message(self, data: object, is_group: bool = False) -> None:
|
||||
"""Parse inbound message, download attachments, and publish to the bus."""
|
||||
try:
|
||||
message = cast(Any, data)
|
||||
if is_group:
|
||||
chat_id = data.group_openid
|
||||
user_id = data.author.member_openid
|
||||
chat_id = cast(str, message.group_openid)
|
||||
user_id = cast(str, message.author.member_openid)
|
||||
chat_type = "group"
|
||||
else:
|
||||
chat_id = str(
|
||||
getattr(data.author, "id", None)
|
||||
or getattr(data.author, "user_openid", "unknown")
|
||||
getattr(message.author, "id", None)
|
||||
or getattr(message.author, "user_openid", "unknown")
|
||||
)
|
||||
user_id = chat_id
|
||||
chat_type = "c2c"
|
||||
|
||||
content = (data.content or "").strip()
|
||||
content = str(message.content or "").strip()
|
||||
|
||||
if data.id in self._processed_ids:
|
||||
message_id = cast(str, message.id)
|
||||
if message_id in self._processed_ids:
|
||||
return
|
||||
self._processed_ids.append(data.id)
|
||||
self._processed_ids.append(message_id)
|
||||
self._chat_type_cache[chat_id] = chat_type
|
||||
|
||||
# Early permission check — avoid attachment downloads and ack side effects
|
||||
@@ -564,7 +572,10 @@ class QQChannel(BaseChannel):
|
||||
|
||||
# the data used by tests don't contain attachments property
|
||||
# so we use getattr with a default of [] to avoid AttributeError in tests
|
||||
attachments = getattr(data, "attachments", None) or []
|
||||
attachments = cast(
|
||||
list[object],
|
||||
getattr(message, "attachments", None) or [],
|
||||
)
|
||||
media_paths, recv_lines, att_meta = await self._handle_attachments(attachments)
|
||||
|
||||
# Compose content that always contains actionable saved paths
|
||||
@@ -587,7 +598,7 @@ class QQChannel(BaseChannel):
|
||||
await self._send_text_only(
|
||||
chat_id=chat_id,
|
||||
is_group=is_group,
|
||||
msg_id=data.id,
|
||||
msg_id=message_id,
|
||||
content=self.config.ack_message,
|
||||
)
|
||||
except Exception:
|
||||
@@ -599,17 +610,20 @@ class QQChannel(BaseChannel):
|
||||
content=content,
|
||||
media=media_paths if media_paths else None,
|
||||
metadata={
|
||||
"message_id": data.id,
|
||||
"message_id": message_id,
|
||||
"attachments": att_meta,
|
||||
},
|
||||
is_dm=not is_group,
|
||||
)
|
||||
except Exception:
|
||||
self.logger.exception("Error handling inbound message id={}", getattr(data, "id", "?"))
|
||||
self.logger.exception(
|
||||
"Error handling inbound message id={}",
|
||||
getattr(data, "id", "?"),
|
||||
)
|
||||
|
||||
async def _handle_attachments(
|
||||
self,
|
||||
attachments: list[BaseMessage._Attachments],
|
||||
attachments: list[object],
|
||||
) -> tuple[list[str], list[str], list[dict[str, Any]]]:
|
||||
"""Extract, download (chunked), and format attachments for agent consumption."""
|
||||
media_paths: list[str] = []
|
||||
@@ -718,9 +732,11 @@ class QQChannel(BaseChannel):
|
||||
1024 * 1024, int(self.config.download_max_bytes or (200 * 1024 * 1024))
|
||||
)
|
||||
|
||||
def _open_tmp():
|
||||
tmp_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
return open(tmp_path, "wb") # noqa: SIM115
|
||||
active_tmp_path = tmp_path
|
||||
|
||||
def _open_tmp() -> BinaryIO:
|
||||
active_tmp_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
return active_tmp_path.open("wb") # noqa: SIM115
|
||||
|
||||
f = await asyncio.to_thread(_open_tmp)
|
||||
try:
|
||||
@@ -740,7 +756,7 @@ class QQChannel(BaseChannel):
|
||||
await asyncio.to_thread(f.close)
|
||||
|
||||
# Atomic rename
|
||||
await asyncio.to_thread(os.replace, tmp_path, target)
|
||||
await asyncio.to_thread(os.replace, active_tmp_path, target)
|
||||
tmp_path = None # mark as moved
|
||||
self.logger.info("file saved: {}", str(target))
|
||||
return str(target)
|
||||
|
||||
@@ -12,7 +12,7 @@ from collections.abc import AsyncIterator, Callable
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, TypedDict, cast
|
||||
|
||||
import httpx
|
||||
from pydantic import Field, computed_field, field_validator
|
||||
@@ -53,7 +53,7 @@ _SIG_TOKEN_RE = re.compile(r"\x00C(\d+)\x00")
|
||||
# stripper needs a fixed, narrow subset (no single-asterisk italic, no
|
||||
# single-tilde strikethrough) and benefits from each pattern's group 1 being
|
||||
# the content directly.
|
||||
_SIG_CELL_STRIP_PATTERNS: tuple[tuple[re.Pattern, str], ...] = (
|
||||
_SIG_CELL_STRIP_PATTERNS: tuple[tuple[re.Pattern[str], str], ...] = (
|
||||
(re.compile(r"\*\*(.+?)\*\*"), r"\1"),
|
||||
(re.compile(r"__(.+?)__"), r"\1"),
|
||||
(re.compile(r"~~(.+?)~~"), r"\1"),
|
||||
@@ -61,6 +61,27 @@ _SIG_CELL_STRIP_PATTERNS: tuple[tuple[re.Pattern, str], ...] = (
|
||||
)
|
||||
|
||||
|
||||
def _as_json_object(value: object) -> dict[str, Any] | None:
|
||||
"""Return an untrusted JSON value only when it is an object."""
|
||||
if isinstance(value, dict):
|
||||
return cast(dict[str, Any], value)
|
||||
return None
|
||||
|
||||
|
||||
def _as_json_object_list(value: object) -> list[dict[str, Any]]:
|
||||
"""Return the object members of an untrusted JSON array."""
|
||||
if not isinstance(value, list):
|
||||
return []
|
||||
return [cast(dict[str, Any], item) for item in cast(list[object], value) if isinstance(item, dict)]
|
||||
|
||||
|
||||
class _BufferedMessage(TypedDict):
|
||||
sender_name: str
|
||||
sender_number: str
|
||||
content: str
|
||||
timestamp: int | None
|
||||
|
||||
|
||||
def _utf16_len(s: str) -> int:
|
||||
"""UTF-16 code-unit length, matching Signal BodyRange semantics."""
|
||||
return len(s.encode("utf-16-le")) // 2
|
||||
@@ -118,7 +139,7 @@ def _markdown_to_signal(text: str) -> tuple[str, list[str]]:
|
||||
# so they're protected from inline-style processing.
|
||||
protected: list[str] = []
|
||||
|
||||
def save_code(m: re.Match) -> str:
|
||||
def save_code(m: re.Match[str]) -> str:
|
||||
protected.append(m.group(1))
|
||||
return f"\x00C{len(protected) - 1}\x00"
|
||||
|
||||
@@ -149,8 +170,8 @@ def _markdown_to_signal(text: str) -> tuple[str, list[str]]:
|
||||
runs: list[_Run] = [_Run(text)]
|
||||
|
||||
def transform(
|
||||
pattern: re.Pattern,
|
||||
make_runs: Callable[[re.Match, frozenset[str]], list[_Run]],
|
||||
pattern: re.Pattern[str],
|
||||
make_runs: Callable[[re.Match[str], frozenset[str]], list[_Run]],
|
||||
) -> None:
|
||||
new_runs: list[_Run] = []
|
||||
for run in runs:
|
||||
@@ -189,7 +210,7 @@ def _markdown_to_signal(text: str) -> tuple[str, list[str]]:
|
||||
transform(_SIG_OLIST_RE, lambda m, s: [_Run(m.group(1) + ". ", s)])
|
||||
|
||||
# Links → "text (url)" or bare url when text equals url.
|
||||
def _link_runs(m: re.Match, s: frozenset) -> list[_Run]:
|
||||
def _link_runs(m: re.Match[str], s: frozenset[str]) -> list[_Run]:
|
||||
link_text, url = m.group(1), m.group(2)
|
||||
|
||||
def _norm(u: str) -> str:
|
||||
@@ -357,15 +378,15 @@ class SignalChannel(BaseChannel):
|
||||
self.config: SignalConfig = config
|
||||
self._http: httpx.AsyncClient | None = None
|
||||
self._request_id = 0
|
||||
self._sse_task: asyncio.Task | None = None
|
||||
self._typing_tasks: dict[str, asyncio.Task] = {}
|
||||
self._sse_task: asyncio.Task[None] | None = None
|
||||
self._typing_tasks: dict[str, asyncio.Task[None]] = {}
|
||||
self._typing_uuid_warnings: set[str] = set()
|
||||
self._account_id_aliases: set[str] = set()
|
||||
self._remember_account_id_alias(self.config.phone_number)
|
||||
|
||||
# Rolling message buffer for group context (group_id -> deque of messages)
|
||||
# Each message is a dict with: sender_name, sender_number, content, timestamp
|
||||
self._group_buffers: dict[str, deque] = {}
|
||||
self._group_buffers: dict[str, deque[_BufferedMessage]] = {}
|
||||
|
||||
def is_allowed(self, sender_id: str) -> bool:
|
||||
"""Override base check to normalize and split pipe-joined identifiers.
|
||||
@@ -409,6 +430,7 @@ class SignalChannel(BaseChannel):
|
||||
metadata: dict[str, Any] | None = None,
|
||||
session_key: str | None = None,
|
||||
is_dm: bool = False,
|
||||
authorization_id: str | None = None,
|
||||
) -> None:
|
||||
"""Handle an inbound message whose policy has already been checked.
|
||||
|
||||
@@ -418,6 +440,7 @@ class SignalChannel(BaseChannel):
|
||||
``super()._handle_message`` instead, which goes through
|
||||
``is_allowed`` and issues a pairing code.
|
||||
"""
|
||||
del authorization_id
|
||||
meta = metadata or {}
|
||||
if self.supports_streaming:
|
||||
meta = {**meta, "_wants_stream": True}
|
||||
@@ -594,7 +617,7 @@ class SignalChannel(BaseChannel):
|
||||
self.logger.info("Subscribed to Signal messages via SSE")
|
||||
|
||||
# Buffer for accumulating SSE data across multiple lines
|
||||
event_buffer = []
|
||||
event_buffer: list[str] = []
|
||||
|
||||
async for line in response.aiter_lines():
|
||||
if not self._running:
|
||||
@@ -605,7 +628,7 @@ class SignalChannel(BaseChannel):
|
||||
self.logger.debug("SSE line received: {}", line[:200])
|
||||
|
||||
# SSE format handling
|
||||
if isinstance(line, str):
|
||||
if isinstance(line, str): # pyright: ignore[reportUnnecessaryIsInstance]
|
||||
# Empty line signals end of event
|
||||
if not line or line == ":":
|
||||
if event_buffer:
|
||||
@@ -613,7 +636,10 @@ class SignalChannel(BaseChannel):
|
||||
data_str = ""
|
||||
try:
|
||||
data_str = "\n".join(event_buffer)
|
||||
data = json.loads(data_str)
|
||||
data = _as_json_object(json.loads(data_str))
|
||||
if data is None:
|
||||
self.logger.warning("Ignoring non-object SSE event: {}", data_str[:200])
|
||||
continue
|
||||
self.logger.debug("SSE event parsed: {}", data)
|
||||
await self._handle_receive_notification(data)
|
||||
except json.JSONDecodeError as e:
|
||||
@@ -644,7 +670,7 @@ class SignalChannel(BaseChannel):
|
||||
self.logger.error("Error in SSE receive loop: {}", e)
|
||||
raise
|
||||
|
||||
@asynccontextmanager
|
||||
@asynccontextmanager # pyright: ignore[reportDeprecated]
|
||||
async def _safe_handle(self, action: str, payload: Any = None) -> AsyncIterator[None]:
|
||||
"""Swallow and log any exception from a top-level handler block.
|
||||
|
||||
@@ -666,17 +692,18 @@ class SignalChannel(BaseChannel):
|
||||
self.logger.debug("_handle_receive_notification called with: {}", params)
|
||||
async with self._safe_handle("receive notification", params):
|
||||
# Extract envelope from SSE notification: {"envelope": {...}}
|
||||
envelope = params.get("envelope", {})
|
||||
envelope = _as_json_object(params.get("envelope"))
|
||||
|
||||
self.logger.debug("Extracted envelope: {}", envelope)
|
||||
|
||||
if not envelope:
|
||||
if envelope is None:
|
||||
self.logger.debug("No envelope found in params")
|
||||
return
|
||||
|
||||
# Extract sender information
|
||||
sender_parts = self._collect_sender_id_parts(envelope)
|
||||
source_name = envelope.get("sourceName")
|
||||
source_name_value = envelope.get("sourceName")
|
||||
source_name = source_name_value if isinstance(source_name_value, str) else None
|
||||
|
||||
if not sender_parts:
|
||||
self.logger.debug("Received message without source, skipping")
|
||||
@@ -691,10 +718,10 @@ class SignalChannel(BaseChannel):
|
||||
self._remember_account_id_alias(part)
|
||||
|
||||
# Check different message types
|
||||
data_message = envelope.get("dataMessage")
|
||||
sync_message = envelope.get("syncMessage")
|
||||
typing_message = envelope.get("typingMessage")
|
||||
receipt_message = envelope.get("receiptMessage")
|
||||
data_message = _as_json_object(envelope.get("dataMessage"))
|
||||
sync_message = _as_json_object(envelope.get("syncMessage"))
|
||||
typing_message = _as_json_object(envelope.get("typingMessage"))
|
||||
receipt_message = _as_json_object(envelope.get("receiptMessage"))
|
||||
|
||||
# Ignore receipt messages (delivery/read receipts)
|
||||
if receipt_message:
|
||||
@@ -705,8 +732,7 @@ class SignalChannel(BaseChannel):
|
||||
await self._handle_data_message(sender_id, sender_number, data_message, source_name)
|
||||
|
||||
# Handle sync messages (messages sent from another device)
|
||||
elif sync_message and sync_message.get("sentMessage"):
|
||||
sent_msg = sync_message["sentMessage"]
|
||||
elif sync_message and (sent_msg := _as_json_object(sync_message.get("sentMessage"))):
|
||||
destination = sent_msg.get("destination") or sent_msg.get("destinationNumber")
|
||||
if destination:
|
||||
self.logger.debug(
|
||||
@@ -725,10 +751,12 @@ class SignalChannel(BaseChannel):
|
||||
sender_name: str | None,
|
||||
) -> None:
|
||||
"""Handle a data message (text, attachments, etc.)."""
|
||||
message_text = data_message.get("message") or ""
|
||||
attachments = data_message.get("attachments", [])
|
||||
mentions = data_message.get("mentions", [])
|
||||
timestamp = data_message.get("timestamp")
|
||||
message_value = data_message.get("message")
|
||||
message_text = message_value if isinstance(message_value, str) else ""
|
||||
attachments = _as_json_object_list(data_message.get("attachments"))
|
||||
mentions = _as_json_object_list(data_message.get("mentions"))
|
||||
timestamp_value = data_message.get("timestamp")
|
||||
timestamp = timestamp_value if isinstance(timestamp_value, int) else None
|
||||
|
||||
self.logger.info(
|
||||
"Data message from {}: groupInfo={}, groupV2={}, keys={}",
|
||||
@@ -815,7 +843,7 @@ class SignalChannel(BaseChannel):
|
||||
group_id: str | None,
|
||||
is_group_message: bool,
|
||||
message_text: str,
|
||||
mentions: list,
|
||||
mentions: list[dict[str, Any]],
|
||||
sender_name: str | None,
|
||||
timestamp: int | None,
|
||||
) -> tuple[bool, str]:
|
||||
@@ -877,8 +905,8 @@ class SignalChannel(BaseChannel):
|
||||
sender_name: str | None,
|
||||
sender_number: str,
|
||||
message_text: str,
|
||||
attachments: list,
|
||||
mentions: list,
|
||||
attachments: list[dict[str, Any]],
|
||||
mentions: list[dict[str, Any]],
|
||||
is_group_message: bool,
|
||||
chat_id: str,
|
||||
) -> tuple[str, list[str]]:
|
||||
@@ -952,7 +980,9 @@ class SignalChannel(BaseChannel):
|
||||
"""
|
||||
# Create buffer for this group if it doesn't exist
|
||||
if group_id not in self._group_buffers:
|
||||
self._group_buffers[group_id] = deque(maxlen=self.config.group_message_buffer_size)
|
||||
self._group_buffers[group_id] = deque[_BufferedMessage](
|
||||
maxlen=self.config.group_message_buffer_size
|
||||
)
|
||||
|
||||
# Add message to buffer (deque will automatically drop oldest when full)
|
||||
self._group_buffers[group_id].append(
|
||||
@@ -992,7 +1022,7 @@ class SignalChannel(BaseChannel):
|
||||
# We want to show context BEFORE the mention
|
||||
context_messages = list(buffer)[:-1] # Exclude the last (current) message
|
||||
|
||||
lines = []
|
||||
lines: list[str] = []
|
||||
for msg in context_messages:
|
||||
sender = msg["sender_name"]
|
||||
content = msg["content"][:200] # Limit to 200 chars per message
|
||||
@@ -1053,8 +1083,6 @@ class SignalChannel(BaseChannel):
|
||||
"""Remember known bot identifiers for mention matching."""
|
||||
if not value:
|
||||
return
|
||||
if not isinstance(value, str):
|
||||
return
|
||||
for candidate in self._normalize_signal_id(value):
|
||||
self._account_id_aliases.add(candidate)
|
||||
|
||||
@@ -1062,8 +1090,6 @@ class SignalChannel(BaseChannel):
|
||||
"""Return True when an identifier refers to the bot account."""
|
||||
if not value:
|
||||
return False
|
||||
if not isinstance(value, str):
|
||||
return False
|
||||
return any(
|
||||
candidate in self._account_id_aliases for candidate in self._normalize_signal_id(value)
|
||||
)
|
||||
@@ -1097,13 +1123,14 @@ class SignalChannel(BaseChannel):
|
||||
return sender_parts[0] if sender_parts else ""
|
||||
|
||||
@staticmethod
|
||||
def _extract_group_id(group_info: Any, group_v2: Any) -> str | None:
|
||||
def _extract_group_id(group_info: object, group_v2: object) -> str | None:
|
||||
"""Extract group ID from groupInfo/groupV2 payloads across signal-cli variants."""
|
||||
for group_obj in (group_info, group_v2):
|
||||
if not isinstance(group_obj, dict):
|
||||
continue
|
||||
group = cast(dict[str, Any], group_obj)
|
||||
for key in ("groupId", "id", "groupID"):
|
||||
value = group_obj.get(key)
|
||||
value = group.get(key)
|
||||
if isinstance(value, str) and value:
|
||||
return value
|
||||
return None
|
||||
@@ -1113,18 +1140,19 @@ class SignalChannel(BaseChannel):
|
||||
"""Extract possible identifier fields from a mention payload."""
|
||||
ids: list[str] = []
|
||||
|
||||
def _walk(value: dict[str, Any] | Any, depth: int = 0) -> None:
|
||||
def _walk(value: object, depth: int = 0) -> None:
|
||||
if depth > 2:
|
||||
return
|
||||
if not isinstance(value, dict):
|
||||
return
|
||||
for key, child in value.items():
|
||||
key_lower = str(key).lower()
|
||||
object_value = cast(dict[str, Any], value)
|
||||
for key, child in object_value.items():
|
||||
key_lower = key.lower()
|
||||
if isinstance(child, str) and child:
|
||||
if any(token in key_lower for token in ("number", "uuid", "serviceid", "aci")):
|
||||
ids.append(child)
|
||||
elif isinstance(child, dict):
|
||||
_walk(child, depth + 1)
|
||||
_walk(cast(object, child), depth + 1)
|
||||
|
||||
_walk(mention)
|
||||
return list(dict.fromkeys(ids))
|
||||
@@ -1187,8 +1215,6 @@ class SignalChannel(BaseChannel):
|
||||
|
||||
# If mention is required, check if bot was mentioned.
|
||||
for mention in mentions:
|
||||
if not isinstance(mention, dict):
|
||||
continue
|
||||
for mention_id in self._mention_id_candidates(mention):
|
||||
if self._id_matches_account(mention_id):
|
||||
return True
|
||||
@@ -1197,15 +1223,13 @@ class SignalChannel(BaseChannel):
|
||||
# (for handle-style mentions). Accept a leading identifier-less mention
|
||||
# as a mention of the bot to avoid false negatives.
|
||||
for mention in mentions:
|
||||
if not isinstance(mention, dict):
|
||||
continue
|
||||
if self._mention_id_candidates(mention):
|
||||
continue
|
||||
span = self._mention_span(mention)
|
||||
if not span:
|
||||
continue
|
||||
start, _ = span
|
||||
if message_text is not None and not message_text[:start].strip():
|
||||
if not message_text[:start].strip():
|
||||
self.logger.debug("Accepting identifier-less leading mention as bot mention")
|
||||
return True
|
||||
|
||||
@@ -1241,10 +1265,8 @@ class SignalChannel(BaseChannel):
|
||||
return text
|
||||
|
||||
# Build a list of (start, length) tuples for our bot's mentions
|
||||
bot_mentions = []
|
||||
bot_mentions: list[tuple[int, int]] = []
|
||||
for mention in mentions:
|
||||
if not isinstance(mention, dict):
|
||||
continue
|
||||
mention_ids = self._mention_id_candidates(mention)
|
||||
span = self._mention_span(mention)
|
||||
if not span:
|
||||
@@ -1382,7 +1404,7 @@ class SignalChannel(BaseChannel):
|
||||
request_id = self._request_id
|
||||
|
||||
# Build JSON-RPC request
|
||||
request = {"jsonrpc": "2.0", "method": method, "id": request_id}
|
||||
request: dict[str, Any] = {"jsonrpc": "2.0", "method": method, "id": request_id}
|
||||
|
||||
if params:
|
||||
request["params"] = params
|
||||
@@ -1397,7 +1419,10 @@ class SignalChannel(BaseChannel):
|
||||
try:
|
||||
response = await self._http.post("/api/v1/rpc", json=request)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
response_json = _as_json_object(response.json())
|
||||
if response_json is None:
|
||||
return {"error": {"message": "signal-cli returned a non-object JSON-RPC response"}}
|
||||
return response_json
|
||||
except Exception as e:
|
||||
self.logger.error("HTTP request failed: {}", e)
|
||||
return {"error": {"message": str(e)}}
|
||||
|
||||
@@ -3,15 +3,16 @@
|
||||
import asyncio
|
||||
import re
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, Protocol, cast
|
||||
|
||||
import httpx
|
||||
from pydantic import Field
|
||||
from slack_sdk.socket_mode.async_client import AsyncBaseSocketModeClient
|
||||
from slack_sdk.socket_mode.request import SocketModeRequest
|
||||
from slack_sdk.socket_mode.response import SocketModeResponse
|
||||
from slack_sdk.socket_mode.websockets import SocketModeClient
|
||||
from slack_sdk.web.async_client import AsyncWebClient
|
||||
from slackify_markdown import slackify_markdown
|
||||
from slackify_markdown import slackify_markdown # pyright: ignore[reportMissingTypeStubs]
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
@@ -23,6 +24,30 @@ from nanobot.pairing import is_approved
|
||||
from nanobot.utils.helpers import safe_filename, split_message
|
||||
|
||||
|
||||
def _as_json_object(value: Any) -> dict[str, Any] | None:
|
||||
"""Narrow Slack's untyped Socket Mode payloads at the boundary."""
|
||||
return cast(dict[str, Any], value) if isinstance(value, dict) else None
|
||||
|
||||
|
||||
def _as_json_list(value: Any) -> list[Any] | None:
|
||||
"""Narrow Slack's untyped Socket Mode arrays at the boundary."""
|
||||
return cast(list[Any], value) if isinstance(value, list) else None
|
||||
|
||||
|
||||
class _SlackWebAPI(Protocol):
|
||||
"""Subset of slack-sdk's dynamically typed Web API used by this channel."""
|
||||
|
||||
async def auth_test(self, **kwargs: Any) -> Any: ...
|
||||
async def chat_postMessage(self, **kwargs: Any) -> Any: ... # noqa: N802
|
||||
async def conversations_list(self, **kwargs: Any) -> Any: ...
|
||||
async def conversations_open(self, **kwargs: Any) -> Any: ...
|
||||
async def conversations_replies(self, **kwargs: Any) -> Any: ...
|
||||
async def files_upload_v2(self, **kwargs: Any) -> Any: ...
|
||||
async def reactions_add(self, **kwargs: Any) -> Any: ...
|
||||
async def reactions_remove(self, **kwargs: Any) -> Any: ...
|
||||
async def users_list(self, **kwargs: Any) -> Any: ...
|
||||
|
||||
|
||||
class SlackDMConfig(Base):
|
||||
"""Slack DM policy configuration."""
|
||||
|
||||
@@ -90,6 +115,13 @@ class SlackChannel(BaseChannel):
|
||||
self._target_cache: dict[str, str] = {}
|
||||
self._thread_context_attempted: set[str] = set()
|
||||
|
||||
def _require_web_api(self) -> _SlackWebAPI:
|
||||
if self._web_client is None:
|
||||
raise RuntimeError("Slack Web API client is not started")
|
||||
# slack-sdk's public methods are runtime-stable but its annotations do
|
||||
# not expose a useful shared interface, so narrow once at the SDK edge.
|
||||
return cast(_SlackWebAPI, self._web_client)
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Start the Slack Socket Mode client."""
|
||||
if not self.config.bot_token or not self.config.app_token:
|
||||
@@ -111,7 +143,8 @@ class SlackChannel(BaseChannel):
|
||||
|
||||
# Resolve bot user ID for mention handling
|
||||
try:
|
||||
auth = await self._web_client.auth_test()
|
||||
web_api = self._require_web_api()
|
||||
auth = await web_api.auth_test()
|
||||
self._bot_user_id = auth.get("user_id")
|
||||
self.logger.info("bot connected as {}", self._bot_user_id)
|
||||
except Exception as e:
|
||||
@@ -155,10 +188,17 @@ class SlackChannel(BaseChannel):
|
||||
self.logger.warning("client not running")
|
||||
return
|
||||
try:
|
||||
web_api = self._require_web_api()
|
||||
target_chat_id = await self._resolve_target_chat_id(msg.chat_id)
|
||||
slack_meta = msg.metadata.get("slack", {}) if msg.metadata else {}
|
||||
raw_slack_meta: Any = msg.metadata.get("slack", {}) if msg.metadata else {}
|
||||
slack_meta: dict[str, Any] = (
|
||||
cast(dict[str, Any], raw_slack_meta)
|
||||
if isinstance(raw_slack_meta, dict)
|
||||
else {}
|
||||
)
|
||||
thread_ts = slack_meta.get("thread_ts")
|
||||
origin_chat_id = str((slack_meta.get("event", {}) or {}).get("channel") or msg.chat_id)
|
||||
event_meta = cast(dict[str, Any], slack_meta.get("event", {}) or {})
|
||||
origin_chat_id = str(event_meta.get("channel") or msg.chat_id)
|
||||
# Reply in the same thread the inbound message belongs to (works
|
||||
# for both real channel threads and DM threads). When the agent
|
||||
# is forwarding to a different channel, drop thread_ts because it
|
||||
@@ -170,7 +210,20 @@ class SlackChannel(BaseChannel):
|
||||
pass # skip empty progress messages (e.g. tool-event-only updates)
|
||||
elif msg.content or not (msg.media or []):
|
||||
mrkdwn = self._to_mrkdwn(msg.content) if msg.content else " "
|
||||
buttons = getattr(msg, "buttons", None) or []
|
||||
raw_buttons = getattr(msg, "buttons", None)
|
||||
buttons: list[list[str]] = (
|
||||
cast(list[list[str]], raw_buttons)
|
||||
if isinstance(raw_buttons, list)
|
||||
and all(
|
||||
isinstance(row, list)
|
||||
and all(
|
||||
isinstance(label, str)
|
||||
for label in cast(list[object], row)
|
||||
)
|
||||
for row in cast(list[object], raw_buttons)
|
||||
)
|
||||
else []
|
||||
)
|
||||
chunks = split_message(mrkdwn, SLACK_MAX_MESSAGE_LEN)
|
||||
for index, chunk in enumerate(chunks):
|
||||
kwargs: dict[str, Any] = dict(
|
||||
@@ -178,11 +231,11 @@ class SlackChannel(BaseChannel):
|
||||
)
|
||||
if buttons and index == len(chunks) - 1:
|
||||
kwargs["blocks"] = self._build_button_blocks(chunk, buttons)
|
||||
await self._web_client.chat_postMessage(**kwargs)
|
||||
await web_api.chat_postMessage(**kwargs)
|
||||
|
||||
for media_path in msg.media or []:
|
||||
try:
|
||||
await self._web_client.files_upload_v2(
|
||||
await web_api.files_upload_v2(
|
||||
channel=target_chat_id,
|
||||
file=media_path,
|
||||
thread_ts=thread_ts_param,
|
||||
@@ -192,8 +245,16 @@ class SlackChannel(BaseChannel):
|
||||
|
||||
# Update reaction emoji when the final (non-progress) response is sent
|
||||
if not is_progress:
|
||||
event = slack_meta.get("event", {})
|
||||
await self._update_react_emoji(origin_chat_id, event.get("ts"))
|
||||
raw_event = slack_meta.get("event", {})
|
||||
event = (
|
||||
cast(dict[str, Any], raw_event)
|
||||
if isinstance(raw_event, dict)
|
||||
else {}
|
||||
)
|
||||
await self._update_react_emoji(
|
||||
origin_chat_id,
|
||||
cast(str | None, event.get("ts")),
|
||||
)
|
||||
|
||||
except Exception:
|
||||
self.logger.exception("Error sending message")
|
||||
@@ -237,20 +298,26 @@ class SlackChannel(BaseChannel):
|
||||
return self._target_cache[cache_key]
|
||||
|
||||
cursor: str | None = None
|
||||
web_api = self._require_web_api()
|
||||
while True:
|
||||
response = await self._web_client.conversations_list(
|
||||
response = cast(dict[str, Any], await web_api.conversations_list(
|
||||
types="public_channel,private_channel",
|
||||
exclude_archived=True,
|
||||
limit=200,
|
||||
cursor=cursor,
|
||||
)
|
||||
for channel in response.get("channels", []):
|
||||
))
|
||||
for channel_value in cast(list[object], response.get("channels", [])):
|
||||
channel = cast(dict[str, Any], channel_value)
|
||||
if self._normalize_target_name(str(channel.get("name") or "")) == normalized:
|
||||
channel_id = str(channel.get("id") or "")
|
||||
if channel_id:
|
||||
self._target_cache[cache_key] = channel_id
|
||||
return channel_id
|
||||
cursor = ((response.get("response_metadata") or {}).get("next_cursor") or "").strip()
|
||||
response_metadata = cast(
|
||||
dict[str, Any],
|
||||
response.get("response_metadata") or {},
|
||||
)
|
||||
cursor = str(response_metadata.get("next_cursor") or "").strip()
|
||||
if not cursor:
|
||||
break
|
||||
|
||||
@@ -269,9 +336,14 @@ class SlackChannel(BaseChannel):
|
||||
return self._target_cache[cache_key]
|
||||
|
||||
cursor: str | None = None
|
||||
web_api = self._require_web_api()
|
||||
while True:
|
||||
response = await self._web_client.users_list(limit=200, cursor=cursor)
|
||||
for member in response.get("members", []):
|
||||
response = cast(
|
||||
dict[str, Any],
|
||||
await web_api.users_list(limit=200, cursor=cursor),
|
||||
)
|
||||
for member_value in cast(list[object], response.get("members", [])):
|
||||
member = cast(dict[str, Any], member_value)
|
||||
if self._member_matches_handle(member, normalized):
|
||||
user_id = str(member.get("id") or "")
|
||||
if not user_id:
|
||||
@@ -279,7 +351,11 @@ class SlackChannel(BaseChannel):
|
||||
dm_id = await self._open_dm_for_user(user_id)
|
||||
self._target_cache[cache_key] = dm_id
|
||||
return dm_id
|
||||
cursor = ((response.get("response_metadata") or {}).get("next_cursor") or "").strip()
|
||||
response_metadata = cast(
|
||||
dict[str, Any],
|
||||
response.get("response_metadata") or {},
|
||||
)
|
||||
cursor = str(response_metadata.get("next_cursor") or "").strip()
|
||||
if not cursor:
|
||||
break
|
||||
|
||||
@@ -288,8 +364,13 @@ class SlackChannel(BaseChannel):
|
||||
)
|
||||
|
||||
async def _open_dm_for_user(self, user_id: str) -> str:
|
||||
response = await self._web_client.conversations_open(users=user_id)
|
||||
channel_id = str(((response.get("channel") or {}).get("id")) or "")
|
||||
web_api = self._require_web_api()
|
||||
response = cast(
|
||||
dict[str, Any],
|
||||
await web_api.conversations_open(users=user_id),
|
||||
)
|
||||
channel = cast(dict[str, Any], response.get("channel") or {})
|
||||
channel_id = str(channel.get("id") or "")
|
||||
if not channel_id:
|
||||
raise ValueError(f"Slack DM target for user '{user_id}' could not be opened.")
|
||||
return channel_id
|
||||
@@ -300,7 +381,7 @@ class SlackChannel(BaseChannel):
|
||||
|
||||
@classmethod
|
||||
def _member_matches_handle(cls, member: dict[str, Any], normalized: str) -> bool:
|
||||
profile = member.get("profile") or {}
|
||||
profile = cast(dict[str, Any], member.get("profile") or {})
|
||||
candidates = {
|
||||
str(member.get("name") or ""),
|
||||
str(profile.get("display_name") or ""),
|
||||
@@ -312,7 +393,7 @@ class SlackChannel(BaseChannel):
|
||||
|
||||
async def _on_socket_request(
|
||||
self,
|
||||
client: SocketModeClient,
|
||||
client: AsyncBaseSocketModeClient,
|
||||
req: SocketModeRequest,
|
||||
) -> None:
|
||||
"""Handle incoming Socket Mode requests."""
|
||||
@@ -327,8 +408,8 @@ class SlackChannel(BaseChannel):
|
||||
SocketModeResponse(envelope_id=req.envelope_id)
|
||||
)
|
||||
|
||||
payload = req.payload or {}
|
||||
event = payload.get("event") or {}
|
||||
payload = _as_json_object(cast(Any, req).payload) or {}
|
||||
event = _as_json_object(payload.get("event")) or {}
|
||||
event_type = event.get("type")
|
||||
|
||||
# Handle app mentions or plain messages
|
||||
@@ -349,6 +430,8 @@ class SlackChannel(BaseChannel):
|
||||
# Avoid double-processing: Slack sends both `message` and `app_mention`
|
||||
# for mentions in channels. Prefer `app_mention`.
|
||||
text = event.get("text") or ""
|
||||
if not isinstance(text, str):
|
||||
return
|
||||
if event_type == "message" and self._bot_user_id and f"<@{self._bot_user_id}>" in text:
|
||||
return
|
||||
|
||||
@@ -362,10 +445,12 @@ class SlackChannel(BaseChannel):
|
||||
event.get("channel_type"),
|
||||
text[:80],
|
||||
)
|
||||
if not sender_id or not chat_id:
|
||||
if not isinstance(sender_id, str) or not sender_id or not isinstance(chat_id, str) or not chat_id:
|
||||
return
|
||||
|
||||
channel_type = event.get("channel_type") or ""
|
||||
if not isinstance(channel_type, str):
|
||||
channel_type = ""
|
||||
|
||||
if not self._is_allowed(sender_id, chat_id, channel_type):
|
||||
if channel_type == "im" and self.config.dm.enabled:
|
||||
@@ -383,7 +468,9 @@ class SlackChannel(BaseChannel):
|
||||
text = self._strip_bot_mention(text)
|
||||
|
||||
event_ts = event.get("ts")
|
||||
event_ts = event_ts if isinstance(event_ts, str) else None
|
||||
raw_thread_ts = event.get("thread_ts")
|
||||
raw_thread_ts = raw_thread_ts if isinstance(raw_thread_ts, str) else None
|
||||
thread_ts = raw_thread_ts
|
||||
# In DMs we don't auto-open a thread on top-level messages (it would
|
||||
# bury replies under "1 reply"). But if the user explicitly opened a
|
||||
@@ -396,11 +483,12 @@ class SlackChannel(BaseChannel):
|
||||
thread_ts = event_ts
|
||||
# Add :eyes: reaction to the triggering message (best-effort)
|
||||
try:
|
||||
if self._web_client and event.get("ts"):
|
||||
await self._web_client.reactions_add(
|
||||
if self._web_client and event_ts:
|
||||
web_api = self._require_web_api()
|
||||
await web_api.reactions_add(
|
||||
channel=chat_id,
|
||||
name=self.config.react_emoji,
|
||||
timestamp=event.get("ts"),
|
||||
timestamp=event_ts,
|
||||
)
|
||||
except Exception as e:
|
||||
self.logger.debug("reactions_add failed: {}", e)
|
||||
@@ -413,10 +501,11 @@ class SlackChannel(BaseChannel):
|
||||
)
|
||||
media_paths: list[str] = []
|
||||
file_markers: list[str] = []
|
||||
for file_info in event.get("files") or []:
|
||||
if not isinstance(file_info, dict):
|
||||
for file_info in _as_json_list(event.get("files")) or []:
|
||||
file_info_object = _as_json_object(file_info)
|
||||
if file_info_object is None:
|
||||
continue
|
||||
file_path, marker = await self._download_slack_file(file_info)
|
||||
file_path, marker = await self._download_slack_file(file_info_object)
|
||||
if file_path:
|
||||
media_paths.append(file_path)
|
||||
if marker:
|
||||
@@ -503,22 +592,30 @@ class SlackChannel(BaseChannel):
|
||||
preview = response.content[:256].lstrip().lower()
|
||||
return preview.startswith(_HTML_DOWNLOAD_PREFIXES)
|
||||
|
||||
async def _on_block_action(self, client: SocketModeClient, req: SocketModeRequest) -> None:
|
||||
async def _on_block_action(
|
||||
self,
|
||||
client: AsyncBaseSocketModeClient,
|
||||
req: SocketModeRequest,
|
||||
) -> None:
|
||||
"""Handle button clicks from inline action buttons."""
|
||||
await client.send_socket_mode_response(SocketModeResponse(envelope_id=req.envelope_id))
|
||||
payload = req.payload or {}
|
||||
actions = payload.get("actions") or []
|
||||
payload = cast(dict[str, Any], cast(Any, req).payload or {})
|
||||
actions = cast(list[Any], payload.get("actions") or [])
|
||||
if not actions:
|
||||
return
|
||||
value = str(actions[0].get("value") or "")
|
||||
user_info = payload.get("user") or {}
|
||||
action = cast(dict[str, Any], actions[0])
|
||||
value = str(action.get("value") or "")
|
||||
user_info = cast(dict[str, Any], payload.get("user") or {})
|
||||
sender_id = str(user_info.get("id") or "")
|
||||
channel_info = payload.get("channel") or {}
|
||||
channel_info = cast(dict[str, Any], payload.get("channel") or {})
|
||||
chat_id = str(channel_info.get("id") or "")
|
||||
if not sender_id or not chat_id or not value:
|
||||
return
|
||||
message_info = payload.get("message") or {}
|
||||
thread_ts = message_info.get("thread_ts") or message_info.get("ts")
|
||||
message_info = cast(dict[str, Any], payload.get("message") or {})
|
||||
thread_ts = cast(
|
||||
str | None,
|
||||
message_info.get("thread_ts") or message_info.get("ts"),
|
||||
)
|
||||
channel_type = self._infer_channel_type(chat_id)
|
||||
if not self._is_allowed(sender_id, chat_id, channel_type):
|
||||
return
|
||||
@@ -563,17 +660,18 @@ class SlackChannel(BaseChannel):
|
||||
self._thread_context_attempted.add(key)
|
||||
|
||||
try:
|
||||
response = await self._web_client.conversations_replies(
|
||||
web_api = self._require_web_api()
|
||||
response = cast(dict[str, Any], await web_api.conversations_replies(
|
||||
channel=chat_id,
|
||||
ts=thread_ts,
|
||||
limit=max(1, self.config.thread_context_limit),
|
||||
)
|
||||
))
|
||||
except Exception as e:
|
||||
self.logger.warning("thread context unavailable for {}: {}", key, e)
|
||||
return text
|
||||
|
||||
lines = self._format_thread_context(
|
||||
response.get("messages", []),
|
||||
cast(list[dict[str, Any]], response.get("messages", [])),
|
||||
current_ts=current_ts,
|
||||
)
|
||||
if not lines:
|
||||
@@ -605,7 +703,7 @@ class SlackChannel(BaseChannel):
|
||||
blocks: list[dict[str, Any]] = [
|
||||
{"type": "section", "text": {"type": "mrkdwn", "text": text[:3000]}},
|
||||
]
|
||||
elements = []
|
||||
elements: list[dict[str, Any]] = []
|
||||
for row in buttons:
|
||||
for label in row:
|
||||
elements.append({
|
||||
@@ -622,8 +720,9 @@ class SlackChannel(BaseChannel):
|
||||
"""Remove the in-progress reaction and optionally add a done reaction."""
|
||||
if not self._web_client or not ts:
|
||||
return
|
||||
web_api = self._require_web_api()
|
||||
try:
|
||||
await self._web_client.reactions_remove(
|
||||
await web_api.reactions_remove(
|
||||
channel=chat_id,
|
||||
name=self.config.react_emoji,
|
||||
timestamp=ts,
|
||||
@@ -632,7 +731,7 @@ class SlackChannel(BaseChannel):
|
||||
self.logger.debug("reactions_remove failed: {}", e)
|
||||
if self.config.done_emoji:
|
||||
try:
|
||||
await self._web_client.reactions_add(
|
||||
await web_api.reactions_add(
|
||||
channel=chat_id,
|
||||
name=self.config.done_emoji,
|
||||
timestamp=ts,
|
||||
@@ -703,7 +802,7 @@ class SlackChannel(BaseChannel):
|
||||
return ""
|
||||
code_blocks: list[str] = []
|
||||
|
||||
def _save_fence(m: re.Match) -> str:
|
||||
def _save_fence(m: re.Match[str]) -> str:
|
||||
code_blocks.append(m.group(0))
|
||||
return f"\x00CB{len(code_blocks) - 1}\x00"
|
||||
|
||||
@@ -718,7 +817,7 @@ class SlackChannel(BaseChannel):
|
||||
"""Fix markdown artifacts that slackify_markdown misses."""
|
||||
code_blocks: list[str] = []
|
||||
|
||||
def _save_code(m: re.Match) -> str:
|
||||
def _save_code(m: re.Match[str]) -> str:
|
||||
code_blocks.append(m.group(0))
|
||||
return f"\x00CB{len(code_blocks) - 1}\x00"
|
||||
|
||||
@@ -726,14 +825,17 @@ class SlackChannel(BaseChannel):
|
||||
text = cls._INLINE_CODE_RE.sub(_save_code, text)
|
||||
text = cls._LEFTOVER_BOLD_RE.sub(r"*\1*", text)
|
||||
text = cls._LEFTOVER_HEADER_RE.sub(r"*\1*", text)
|
||||
text = cls._BARE_URL_RE.sub(lambda m: m.group(0).replace("&", "&"), text)
|
||||
text = cls._BARE_URL_RE.sub(
|
||||
lambda m: m.group(0).replace("&", "&"),
|
||||
text,
|
||||
)
|
||||
|
||||
for i, block in enumerate(code_blocks):
|
||||
text = text.replace(f"\x00CB{i}\x00", block)
|
||||
return text
|
||||
|
||||
@staticmethod
|
||||
def _convert_table(match: re.Match) -> str:
|
||||
def _convert_table(match: re.Match[str]) -> str:
|
||||
"""Convert a Markdown table to a Slack-readable list."""
|
||||
lines = [ln.strip() for ln in match.group(0).strip().splitlines() if ln.strip()]
|
||||
if len(lines) < 2:
|
||||
|
||||
@@ -8,8 +8,9 @@ import time
|
||||
import unicodedata
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from datetime import timedelta
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
from typing import Any, Awaitable, Callable, Literal, TypeAlias, TypeVar, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from pydantic import Field, field_validator, model_validator
|
||||
@@ -17,9 +18,12 @@ from telegram import (
|
||||
BotCommand,
|
||||
InlineKeyboardButton,
|
||||
InlineKeyboardMarkup,
|
||||
Message,
|
||||
MessageEntity,
|
||||
ReactionTypeEmoji,
|
||||
ReplyParameters,
|
||||
Update,
|
||||
User,
|
||||
)
|
||||
from telegram.error import BadRequest, NetworkError, TimedOut
|
||||
from telegram.ext import Application, CallbackQueryHandler, ContextTypes, MessageHandler, filters
|
||||
@@ -43,6 +47,12 @@ TELEGRAM_MAX_MESSAGE_LEN = 4000 # Telegram message character limit
|
||||
TELEGRAM_HTML_MAX_LEN = 4096
|
||||
TELEGRAM_REPLY_CONTEXT_MAX_LEN = TELEGRAM_MAX_MESSAGE_LEN # Max length for reply context in user message
|
||||
|
||||
# python-telegram-bot exposes a six-parameter Application generic. Nanobot
|
||||
# doesn't customize its context/data/job-queue types, so keep that SDK boundary
|
||||
# explicit rather than allowing unspecialized generics to spread Unknown.
|
||||
TelegramApplication: TypeAlias = Application[Any, Any, Any, Any, Any, Any]
|
||||
_T = TypeVar("_T")
|
||||
|
||||
|
||||
def _split_telegram_markdown(content: str, max_len: int) -> list[str]:
|
||||
"""Split raw Telegram Markdown without leaving fenced code blocks unbalanced."""
|
||||
@@ -218,7 +228,7 @@ def _markdown_to_telegram_html(text: str) -> str:
|
||||
|
||||
# 1. Extract and protect code blocks (preserve content from other processing)
|
||||
code_blocks: list[str] = []
|
||||
def save_code_block(m: re.Match) -> str:
|
||||
def save_code_block(m: re.Match[str]) -> str:
|
||||
code_blocks.append(m.group(1))
|
||||
return f"\x00CB{len(code_blocks) - 1}\x00"
|
||||
|
||||
@@ -247,7 +257,7 @@ def _markdown_to_telegram_html(text: str) -> str:
|
||||
|
||||
# 2. Extract and protect inline code
|
||||
inline_codes: list[str] = []
|
||||
def save_inline_code(m: re.Match) -> str:
|
||||
def save_inline_code(m: re.Match[str]) -> str:
|
||||
inline_codes.append(m.group(1))
|
||||
return f"\x00IC{len(inline_codes) - 1}\x00"
|
||||
|
||||
@@ -350,7 +360,7 @@ class _QueuedTelegramUpdate:
|
||||
|
||||
kind: Literal["command", "message"]
|
||||
update: Update
|
||||
context: Any
|
||||
context: ContextTypes.DEFAULT_TYPE
|
||||
sort_key: tuple[int, int]
|
||||
|
||||
|
||||
@@ -421,7 +431,7 @@ class TelegramChannel(BaseChannel):
|
||||
display_name = "Telegram"
|
||||
|
||||
# Commands registered with Telegram's command menu
|
||||
BOT_COMMANDS = [
|
||||
BOT_COMMANDS: list[BotCommand] = [
|
||||
BotCommand("start", "Start the bot"),
|
||||
BotCommand("new", "Start a new conversation"),
|
||||
BotCommand("stop", "Stop the current task"),
|
||||
@@ -455,19 +465,24 @@ class TelegramChannel(BaseChannel):
|
||||
config = TelegramConfig.model_validate(config)
|
||||
super().__init__(config, bus)
|
||||
self.config: TelegramConfig = config
|
||||
self._app: Application | None = None
|
||||
self._app: TelegramApplication | None = None
|
||||
self._chat_ids: dict[str, int] = {} # Map sender_id to chat_id for replies
|
||||
self._typing_tasks: dict[str, asyncio.Task] = {} # chat_id -> typing loop task
|
||||
self._media_group_buffers: dict[str, dict] = {}
|
||||
self._media_group_tasks: dict[str, asyncio.Task] = {}
|
||||
self._typing_tasks: dict[str, asyncio.Task[None]] = {} # chat_id -> typing loop task
|
||||
self._media_group_buffers: dict[str, dict[str, Any]] = {}
|
||||
self._media_group_tasks: dict[str, asyncio.Task[None]] = {}
|
||||
self._message_threads: dict[tuple[str, int], int] = {}
|
||||
self._bot_user_id: int | None = None
|
||||
self._bot_username: str | None = None
|
||||
self._stream_bufs: dict[str, _StreamBuf] = {} # chat_id -> streaming state
|
||||
self._inbound_buffers: dict[str, list[_QueuedTelegramUpdate]] = {}
|
||||
self._inbound_workers: dict[str, asyncio.Task] = {}
|
||||
self._inbound_workers: dict[str, asyncio.Task[None]] = {}
|
||||
self._rich_send_disabled: bool = False # Latch off if Bot API < 10.1
|
||||
|
||||
def _require_app(self) -> TelegramApplication:
|
||||
if self._app is None:
|
||||
raise RuntimeError("Telegram application is not started")
|
||||
return self._app
|
||||
|
||||
def is_allowed(self, sender_id: str) -> bool:
|
||||
"""Preserve Telegram's legacy id|username allowlist matching."""
|
||||
if super().is_allowed(sender_id):
|
||||
@@ -595,7 +610,7 @@ class TelegramChannel(BaseChannel):
|
||||
if self.config.mode == "webhook":
|
||||
# ``url_path`` is the local HTTP route. ``webhook_url`` is the
|
||||
# public HTTPS URL Telegram calls; reverse proxies may rewrite it.
|
||||
await self._app.updater.start_webhook(
|
||||
await cast(Any, self._app.updater).start_webhook(
|
||||
listen=self.config.webhook_listen_host,
|
||||
port=self.config.webhook_listen_port,
|
||||
url_path=self.config.webhook_path.lstrip("/"),
|
||||
@@ -607,7 +622,7 @@ class TelegramChannel(BaseChannel):
|
||||
)
|
||||
else:
|
||||
# Start polling (this runs until stopped)
|
||||
await self._app.updater.start_polling(
|
||||
await cast(Any, self._app.updater).start_polling(
|
||||
allowed_updates=allowed_updates,
|
||||
drop_pending_updates=False, # Process pending messages on startup
|
||||
error_callback=self._on_polling_error,
|
||||
@@ -637,7 +652,7 @@ class TelegramChannel(BaseChannel):
|
||||
|
||||
if self._app:
|
||||
self.logger.info("Stopping bot...")
|
||||
await self._app.updater.stop()
|
||||
await cast(Any, self._app.updater).stop()
|
||||
await self._app.stop()
|
||||
await self._app.shutdown()
|
||||
self._app = None
|
||||
@@ -674,9 +689,9 @@ class TelegramChannel(BaseChannel):
|
||||
self,
|
||||
chat_id: int,
|
||||
content: str,
|
||||
reply_params=None,
|
||||
thread_kwargs: dict | None = None,
|
||||
reply_markup=None,
|
||||
reply_params: ReplyParameters | dict[str, int | bool] | None = None,
|
||||
thread_kwargs: dict[str, int] | None = None,
|
||||
reply_markup: InlineKeyboardMarkup | None = None,
|
||||
) -> bool:
|
||||
"""Attempt sendRichMessage (Bot API 10.1). Returns True on success."""
|
||||
if not self._app:
|
||||
@@ -692,13 +707,17 @@ class TelegramChannel(BaseChannel):
|
||||
# sendRichMessage uses reply_parameters (object), not reply_to_message_id.
|
||||
if hasattr(reply_params, "message_id"):
|
||||
payload["reply_parameters"] = {
|
||||
"message_id": reply_params.message_id,
|
||||
"message_id": cast(ReplyParameters, reply_params).message_id,
|
||||
"allow_sending_without_reply": True,
|
||||
}
|
||||
else:
|
||||
payload["reply_parameters"] = reply_params
|
||||
if thread_kwargs:
|
||||
payload.update({k: v for k, v in thread_kwargs.items() if v is not None})
|
||||
payload.update({
|
||||
k: v
|
||||
for k, v in thread_kwargs.items()
|
||||
if v is not None # pyright: ignore[reportUnnecessaryComparison]
|
||||
})
|
||||
if reply_markup is not None:
|
||||
payload["reply_markup"] = reply_markup
|
||||
|
||||
@@ -749,7 +768,7 @@ class TelegramChannel(BaseChannel):
|
||||
message_thread_id = msg.metadata.get("message_thread_id")
|
||||
if message_thread_id is None and reply_to_message_id is not None:
|
||||
message_thread_id = self._message_threads.get((msg.chat_id, reply_to_message_id))
|
||||
thread_kwargs = {}
|
||||
thread_kwargs: dict[str, int] = {}
|
||||
if message_thread_id is not None:
|
||||
thread_kwargs["message_thread_id"] = message_thread_id
|
||||
|
||||
@@ -820,7 +839,7 @@ class TelegramChannel(BaseChannel):
|
||||
# Send text content
|
||||
if msg.content and msg.content != "[empty message]":
|
||||
render_as_blockquote = bool(progress_event and progress_event.tool_hint)
|
||||
buttons = getattr(msg, "buttons", None) or []
|
||||
buttons = cast(list[list[str]], getattr(msg, "buttons", None) or [])
|
||||
reply_markup = self._build_keyboard(buttons) if buttons else None
|
||||
text = msg.content
|
||||
# Fallback: no native keyboard → splice labels into the message so the choices survive.
|
||||
@@ -850,7 +869,12 @@ class TelegramChannel(BaseChannel):
|
||||
reply_markup=reply_markup if is_last else None,
|
||||
)
|
||||
|
||||
async def _call_with_retry(self, fn, *args, **kwargs):
|
||||
async def _call_with_retry(
|
||||
self,
|
||||
fn: Callable[..., Awaitable[_T]],
|
||||
*args: Any,
|
||||
**kwargs: Any,
|
||||
) -> _T:
|
||||
"""Call an async Telegram API function with retry on pool/network timeout and RetryAfter."""
|
||||
from telegram.error import RetryAfter
|
||||
|
||||
@@ -869,27 +893,34 @@ class TelegramChannel(BaseChannel):
|
||||
except RetryAfter as e:
|
||||
if attempt == _SEND_MAX_RETRIES:
|
||||
raise
|
||||
delay = float(e.retry_after)
|
||||
retry_after = e.retry_after
|
||||
delay = (
|
||||
retry_after.total_seconds()
|
||||
if isinstance(retry_after, timedelta)
|
||||
else float(retry_after)
|
||||
)
|
||||
self.logger.warning(
|
||||
"Flood Control (attempt {}/{}), retrying in {:.1f}s",
|
||||
attempt, _SEND_MAX_RETRIES, delay,
|
||||
)
|
||||
await asyncio.sleep(delay)
|
||||
raise RuntimeError("Telegram retry loop exited unexpectedly")
|
||||
|
||||
async def _send_text(
|
||||
self,
|
||||
chat_id: int,
|
||||
text: str,
|
||||
reply_params=None,
|
||||
thread_kwargs: dict | None = None,
|
||||
reply_params: ReplyParameters | None = None,
|
||||
thread_kwargs: dict[str, int] | None = None,
|
||||
render_as_blockquote: bool = False,
|
||||
reply_markup=None,
|
||||
reply_markup: InlineKeyboardMarkup | None = None,
|
||||
) -> None:
|
||||
"""Send a plain text message with HTML fallback."""
|
||||
app = self._require_app()
|
||||
try:
|
||||
html = _tool_hint_to_telegram_blockquote(text) if render_as_blockquote else _markdown_to_telegram_html(text)
|
||||
await self._call_with_retry(
|
||||
self._app.bot.send_message,
|
||||
app.bot.send_message,
|
||||
chat_id=chat_id, text=html, parse_mode="HTML",
|
||||
reply_parameters=reply_params,
|
||||
reply_markup=reply_markup,
|
||||
@@ -899,7 +930,7 @@ class TelegramChannel(BaseChannel):
|
||||
self.logger.warning("HTML parse failed, falling back to plain text: {}", e)
|
||||
try:
|
||||
await self._call_with_retry(
|
||||
self._app.bot.send_message,
|
||||
app.bot.send_message,
|
||||
chat_id=chat_id,
|
||||
text=text,
|
||||
reply_parameters=reply_params,
|
||||
@@ -945,7 +976,7 @@ class TelegramChannel(BaseChannel):
|
||||
if reply_to_message_id := meta.get("message_id"):
|
||||
with suppress(ValueError):
|
||||
await self._remove_reaction(chat_id, int(reply_to_message_id))
|
||||
thread_kwargs = {}
|
||||
thread_kwargs: dict[str, int] = {}
|
||||
if message_thread_id := meta.get("message_thread_id"):
|
||||
thread_kwargs["message_thread_id"] = message_thread_id
|
||||
raw_text = buf.text
|
||||
@@ -1032,16 +1063,16 @@ class TelegramChannel(BaseChannel):
|
||||
return
|
||||
|
||||
now = time.monotonic()
|
||||
thread_kwargs = {}
|
||||
stream_thread_kwargs: dict[str, int] = {}
|
||||
if message_thread_id := meta.get("message_thread_id"):
|
||||
thread_kwargs["message_thread_id"] = message_thread_id
|
||||
stream_thread_kwargs["message_thread_id"] = message_thread_id
|
||||
if buf.message_id is None:
|
||||
preview = _strip_md_block(buf.text)
|
||||
try:
|
||||
sent = await self._call_with_retry(
|
||||
self._app.bot.send_message,
|
||||
chat_id=int_chat_id, text=preview,
|
||||
**thread_kwargs,
|
||||
**stream_thread_kwargs,
|
||||
)
|
||||
buf.message_id = sent.message_id
|
||||
buf.last_edit = now
|
||||
@@ -1050,7 +1081,7 @@ class TelegramChannel(BaseChannel):
|
||||
raise # Let ChannelManager handle retry
|
||||
elif (now - buf.last_edit) >= self.config.stream_edit_interval:
|
||||
if len(buf.text) > TELEGRAM_MAX_MESSAGE_LEN:
|
||||
await self._flush_stream_overflow(int_chat_id, buf, thread_kwargs)
|
||||
await self._flush_stream_overflow(int_chat_id, buf, stream_thread_kwargs)
|
||||
buf.last_edit = now
|
||||
return
|
||||
preview = _strip_md_block(buf.text)
|
||||
@@ -1072,7 +1103,7 @@ class TelegramChannel(BaseChannel):
|
||||
self,
|
||||
chat_id: int,
|
||||
buf: "_StreamBuf",
|
||||
thread_kwargs: dict,
|
||||
thread_kwargs: dict[str, int],
|
||||
) -> None:
|
||||
"""Split an oversized stream buffer mid-flight.
|
||||
|
||||
@@ -1083,10 +1114,11 @@ class TelegramChannel(BaseChannel):
|
||||
chunks = _split_telegram_markdown_html_chunks(buf.text, TELEGRAM_HTML_MAX_LEN)
|
||||
if len(chunks) <= 1:
|
||||
return
|
||||
app = self._require_app()
|
||||
first_markdown, first_html = chunks[0]
|
||||
try:
|
||||
await self._call_with_retry(
|
||||
self._app.bot.edit_message_text,
|
||||
app.bot.edit_message_text,
|
||||
chat_id=chat_id, message_id=buf.message_id,
|
||||
text=first_html,
|
||||
parse_mode="HTML",
|
||||
@@ -1098,7 +1130,7 @@ class TelegramChannel(BaseChannel):
|
||||
)
|
||||
try:
|
||||
await self._call_with_retry(
|
||||
self._app.bot.edit_message_text,
|
||||
app.bot.edit_message_text,
|
||||
chat_id=chat_id, message_id=buf.message_id,
|
||||
text=first_markdown,
|
||||
)
|
||||
@@ -1113,7 +1145,7 @@ class TelegramChannel(BaseChannel):
|
||||
async def send_chunk(markdown: str, html: str) -> Any:
|
||||
try:
|
||||
return await self._call_with_retry(
|
||||
self._app.bot.send_message,
|
||||
app.bot.send_message,
|
||||
chat_id=chat_id, text=html, parse_mode="HTML", **thread_kwargs,
|
||||
)
|
||||
except BadRequest as e:
|
||||
@@ -1121,7 +1153,7 @@ class TelegramChannel(BaseChannel):
|
||||
"Stream overflow HTML send failed, falling back to plain text: {}", e
|
||||
)
|
||||
return await self._call_with_retry(
|
||||
self._app.bot.send_message,
|
||||
app.bot.send_message,
|
||||
chat_id=chat_id, text=markdown, **thread_kwargs,
|
||||
)
|
||||
|
||||
@@ -1160,12 +1192,14 @@ class TelegramChannel(BaseChannel):
|
||||
await update.message.reply_text(build_help_text())
|
||||
|
||||
@staticmethod
|
||||
def _sender_id(user) -> str:
|
||||
def _sender_id(user: User) -> str:
|
||||
"""Build sender_id with username for allowlist matching."""
|
||||
sid = str(user.id)
|
||||
return f"{sid}|{user.username}" if user.username else sid
|
||||
|
||||
async def _send_pairing_code_if_private(self, sender_id: str, message, user) -> None:
|
||||
async def _send_pairing_code_if_private(
|
||||
self, sender_id: str, message: Message, user: User
|
||||
) -> None:
|
||||
if message.chat.type != "private":
|
||||
return
|
||||
await self._handle_message(
|
||||
@@ -1177,7 +1211,7 @@ class TelegramChannel(BaseChannel):
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _derive_topic_session_key(message) -> str | None:
|
||||
def _derive_topic_session_key(message: Message) -> str | None:
|
||||
"""Derive topic-scoped session key for Telegram chats with threads."""
|
||||
message_thread_id = getattr(message, "message_thread_id", None)
|
||||
if message_thread_id is None:
|
||||
@@ -1185,7 +1219,7 @@ class TelegramChannel(BaseChannel):
|
||||
return f"telegram:{message.chat_id}:topic:{message_thread_id}"
|
||||
|
||||
@staticmethod
|
||||
def _build_message_metadata(message, user) -> dict:
|
||||
def _build_message_metadata(message: Message, user: User) -> dict[str, Any]:
|
||||
"""Build common Telegram inbound metadata payload."""
|
||||
reply_to = getattr(message, "reply_to_message", None)
|
||||
return {
|
||||
@@ -1199,7 +1233,7 @@ class TelegramChannel(BaseChannel):
|
||||
"reply_to_message_id": getattr(reply_to, "message_id", None) if reply_to else None,
|
||||
}
|
||||
|
||||
async def _extract_reply_context(self, message) -> str | None:
|
||||
async def _extract_reply_context(self, message: Message) -> str | None:
|
||||
"""Extract text from the message being replied to, if any."""
|
||||
reply = getattr(message, "reply_to_message", None)
|
||||
if not reply:
|
||||
@@ -1224,7 +1258,7 @@ class TelegramChannel(BaseChannel):
|
||||
return f"[Reply to: {text}]"
|
||||
|
||||
async def _download_message_media(
|
||||
self, msg, *, add_failure_content: bool = False
|
||||
self, msg: Message, *, add_failure_content: bool = False
|
||||
) -> tuple[list[str], list[str]]:
|
||||
"""Download media from a message (current or reply). Returns (media_paths, content_parts)."""
|
||||
media_file = None
|
||||
@@ -1255,7 +1289,7 @@ class TelegramChannel(BaseChannel):
|
||||
try:
|
||||
file = await self._app.bot.get_file(media_file.file_id)
|
||||
ext = self._get_extension(
|
||||
media_type,
|
||||
cast(str, media_type),
|
||||
getattr(media_file, "mime_type", None),
|
||||
getattr(media_file, "file_name", None),
|
||||
)
|
||||
@@ -1291,7 +1325,7 @@ class TelegramChannel(BaseChannel):
|
||||
@staticmethod
|
||||
def _has_mention_entity(
|
||||
text: str,
|
||||
entities,
|
||||
entities: list[MessageEntity] | None,
|
||||
bot_username: str,
|
||||
bot_id: int | None,
|
||||
) -> bool:
|
||||
@@ -1314,7 +1348,7 @@ class TelegramChannel(BaseChannel):
|
||||
return True
|
||||
return handle in text.lower()
|
||||
|
||||
async def _is_group_message_for_bot(self, message) -> bool:
|
||||
async def _is_group_message_for_bot(self, message: Message) -> bool:
|
||||
"""Allow group messages when policy is open, @mentioned, or replying to the bot."""
|
||||
if message.chat.type == "private" or self.config.group_policy == "open":
|
||||
return True
|
||||
@@ -1341,7 +1375,7 @@ class TelegramChannel(BaseChannel):
|
||||
reply_user = getattr(getattr(message, "reply_to_message", None), "from_user", None)
|
||||
return bool(bot_id and reply_user and reply_user.id == bot_id)
|
||||
|
||||
def _remember_thread_context(self, message) -> None:
|
||||
def _remember_thread_context(self, message: Message) -> None:
|
||||
"""Cache Telegram thread context by chat/message id for follow-up replies."""
|
||||
message_thread_id = getattr(message, "message_thread_id", None)
|
||||
if message_thread_id is None:
|
||||
@@ -1352,7 +1386,7 @@ class TelegramChannel(BaseChannel):
|
||||
self._message_threads.pop(next(iter(self._message_threads)))
|
||||
|
||||
@staticmethod
|
||||
def _queue_key_for_message(message) -> str:
|
||||
def _queue_key_for_message(message: Message) -> str:
|
||||
"""Return the final nanobot session key used for ordered Telegram ingress."""
|
||||
return TelegramChannel._derive_topic_session_key(message) or f"telegram:{message.chat_id}"
|
||||
|
||||
@@ -1373,6 +1407,8 @@ class TelegramChannel(BaseChannel):
|
||||
) -> None:
|
||||
"""Stage a Telegram update behind a short per-session reorder window."""
|
||||
message = update.message
|
||||
if message is None:
|
||||
return
|
||||
key = self._queue_key_for_message(message)
|
||||
self._inbound_buffers.setdefault(key, []).append(
|
||||
_QueuedTelegramUpdate(
|
||||
@@ -1432,6 +1468,8 @@ class TelegramChannel(BaseChannel):
|
||||
"""Process a queued slash command."""
|
||||
message = update.message
|
||||
user = update.effective_user
|
||||
if message is None or user is None:
|
||||
return
|
||||
sender_id = self._sender_id(user)
|
||||
if not self.is_allowed(sender_id):
|
||||
await self._send_pairing_code_if_private(sender_id, message, user)
|
||||
@@ -1469,6 +1507,8 @@ class TelegramChannel(BaseChannel):
|
||||
|
||||
message = update.message
|
||||
user = update.effective_user
|
||||
if message is None or user is None:
|
||||
return
|
||||
chat_id = message.chat_id
|
||||
sender_id = self._sender_id(user)
|
||||
if not self.is_allowed(sender_id):
|
||||
@@ -1483,8 +1523,8 @@ class TelegramChannel(BaseChannel):
|
||||
return
|
||||
|
||||
# Build content from text and/or media
|
||||
content_parts = []
|
||||
media_paths = []
|
||||
content_parts: list[str] = []
|
||||
media_paths: list[str] = []
|
||||
|
||||
# Text content
|
||||
if message.text:
|
||||
@@ -1625,8 +1665,10 @@ class TelegramChannel(BaseChannel):
|
||||
self.logger.debug("Typing indicator stopped for {}: {}", chat_id, e)
|
||||
|
||||
@staticmethod
|
||||
def _format_telegram_error(exc: Exception) -> str:
|
||||
def _format_telegram_error(exc: Exception | None) -> str:
|
||||
"""Return a short, readable error summary for logs."""
|
||||
if exc is None:
|
||||
return "None"
|
||||
text = str(exc).strip()
|
||||
if text:
|
||||
return text
|
||||
@@ -1682,7 +1724,7 @@ class TelegramChannel(BaseChannel):
|
||||
|
||||
return ""
|
||||
|
||||
def _build_keyboard(self, buttons: list) -> InlineKeyboardMarkup | None:
|
||||
def _build_keyboard(self, buttons: list[list[str]]) -> InlineKeyboardMarkup | None:
|
||||
"""Build inline keyboard markup if inline_keyboards is enabled."""
|
||||
if not buttons or not self.config.inline_keyboards:
|
||||
return None
|
||||
@@ -1711,7 +1753,8 @@ class TelegramChannel(BaseChannel):
|
||||
return
|
||||
query = update.callback_query
|
||||
user = update.effective_user
|
||||
chat_id = query.message.chat_id if query.message else None
|
||||
query_message = query.message
|
||||
chat_id = query_message.chat.id if query_message else None
|
||||
sender_id = self._sender_id(user)
|
||||
if not chat_id:
|
||||
self.logger.warning("Callback query without chat_id")
|
||||
@@ -1720,9 +1763,9 @@ class TelegramChannel(BaseChannel):
|
||||
return
|
||||
button_label = query.data or ""
|
||||
await query.answer()
|
||||
if query.message:
|
||||
if isinstance(query_message, Message):
|
||||
with suppress(Exception):
|
||||
await query.message.edit_reply_markup(reply_markup=None)
|
||||
await query_message.edit_reply_markup(reply_markup=None)
|
||||
self.logger.debug("Inline button tap from {}: {}", sender_id, button_label)
|
||||
self._start_typing(str(chat_id))
|
||||
await self._handle_message(
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import asyncio
|
||||
from datetime import timedelta
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
@@ -1911,6 +1912,36 @@ async def test_on_message_location_with_text() -> None:
|
||||
# Tests for retry amplification fix (issue #3050)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_with_retry_accepts_timedelta_retry_after(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from telegram.error import RetryAfter
|
||||
|
||||
channel = TelegramChannel(
|
||||
TelegramConfig(enabled=True, token="123:abc", allow_from=["*"]),
|
||||
MessageBus(),
|
||||
)
|
||||
attempts = 0
|
||||
|
||||
async def retry_once() -> str:
|
||||
nonlocal attempts
|
||||
attempts += 1
|
||||
if attempts == 1:
|
||||
raise RetryAfter(timedelta(seconds=1.5))
|
||||
return "ok"
|
||||
|
||||
sleep = AsyncMock()
|
||||
monkeypatch.setenv("PTB_TIMEDELTA", "1")
|
||||
monkeypatch.setattr(
|
||||
"nanobot.channels.telegram.runtime.asyncio.sleep",
|
||||
sleep,
|
||||
)
|
||||
|
||||
assert await channel._call_with_retry(retry_once) == "ok"
|
||||
sleep.assert_awaited_once_with(1.5)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_text_does_not_fallback_on_network_timeout() -> None:
|
||||
"""TimedOut should propagate immediately, NOT trigger plain-text fallback.
|
||||
@@ -2318,7 +2349,7 @@ async def test_callback_query_ignores_unauthorized_user_before_side_effects() ->
|
||||
data="Yes",
|
||||
answer=AsyncMock(),
|
||||
message=SimpleNamespace(
|
||||
chat_id=123,
|
||||
chat=SimpleNamespace(id=123),
|
||||
edit_reply_markup=AsyncMock(),
|
||||
),
|
||||
)
|
||||
@@ -2332,3 +2363,35 @@ async def test_callback_query_ignores_unauthorized_user_before_side_effects() ->
|
||||
query.answer.assert_not_awaited()
|
||||
query.message.edit_reply_markup.assert_not_awaited()
|
||||
channel._handle_message.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_callback_query_handles_inaccessible_message() -> None:
|
||||
from telegram import Chat, InaccessibleMessage
|
||||
|
||||
channel = TelegramChannel(
|
||||
TelegramConfig(enabled=True, token="123:abc", allow_from=["*"], inline_keyboards=True),
|
||||
MessageBus(),
|
||||
)
|
||||
channel._handle_message = AsyncMock()
|
||||
channel._start_typing = lambda _chat_id: None
|
||||
|
||||
query = SimpleNamespace(
|
||||
id="cb_inaccessible",
|
||||
data="Yes",
|
||||
answer=AsyncMock(),
|
||||
message=InaccessibleMessage(
|
||||
chat=Chat(id=123, type="private"),
|
||||
message_id=456,
|
||||
),
|
||||
)
|
||||
update = SimpleNamespace(
|
||||
callback_query=query,
|
||||
effective_user=SimpleNamespace(id=12345, username="alice", first_name="Alice"),
|
||||
)
|
||||
|
||||
await channel._on_callback_query(update, None)
|
||||
|
||||
query.answer.assert_awaited_once()
|
||||
channel._handle_message.assert_awaited_once()
|
||||
assert channel._handle_message.await_args.kwargs["chat_id"] == "123"
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
"""Telegram setup validation owned by the channel package."""
|
||||
|
||||
import re
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
@@ -39,7 +39,7 @@ def _get_me(token: str, proxy: str | None) -> dict[str, Any]:
|
||||
response = client.get(f"https://api.telegram.org/bot{token}/getMe")
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
return data if isinstance(data, dict) else {}
|
||||
return cast(dict[str, Any], data) if isinstance(data, dict) else {}
|
||||
|
||||
|
||||
def validate(values: dict[str, Any], _context: ChannelValidationContext) -> dict[str, Any]:
|
||||
|
||||
@@ -11,7 +11,7 @@ import re
|
||||
import socket
|
||||
import ssl
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
import httpx
|
||||
|
||||
@@ -76,7 +76,7 @@ def validate_channel_config(
|
||||
allow_local_service_access=config.tools.webui_allow_local_service_access,
|
||||
)
|
||||
custom_payload = setup_spec.validator(values, context)
|
||||
if custom_payload is not None:
|
||||
if cast(object, custom_payload) is not None:
|
||||
payload = dict(custom_payload)
|
||||
payload.setdefault("checks", [])
|
||||
payload.setdefault("missing_fields", [])
|
||||
@@ -116,7 +116,7 @@ def _channel_config(
|
||||
if hasattr(section, "model_dump"):
|
||||
return dict(section.model_dump(mode="json", by_alias=True))
|
||||
if isinstance(section, dict):
|
||||
return dict(section)
|
||||
return dict(cast(dict[str, Any], section))
|
||||
return {}
|
||||
|
||||
|
||||
@@ -130,9 +130,9 @@ def _merge_form_values(
|
||||
merged = dict(values)
|
||||
prefix = f"channels.{name}."
|
||||
spec = setup_spec
|
||||
secrets = spec.secrets if spec is not None else frozenset()
|
||||
secrets: frozenset[str] = spec.secrets if spec is not None else frozenset()
|
||||
for raw_key, raw_value in raw_values.items():
|
||||
if not isinstance(raw_key, str) or not raw_key:
|
||||
if not raw_key:
|
||||
continue
|
||||
field = raw_key[len(prefix):] if raw_key.startswith(prefix) else raw_key
|
||||
if field in secrets and not _str(raw_value):
|
||||
@@ -281,7 +281,7 @@ def _assign(values: dict[str, Any], field: str, value: Any) -> None:
|
||||
if not isinstance(current, dict):
|
||||
current = {}
|
||||
target[part] = current
|
||||
target = current
|
||||
target = cast(dict[str, Any], current)
|
||||
target[parts[-1]] = value
|
||||
|
||||
|
||||
@@ -290,7 +290,7 @@ def _get(values: dict[str, Any], field: str) -> Any:
|
||||
for part in field.split("."):
|
||||
if not isinstance(target, dict):
|
||||
return None
|
||||
target = target.get(part)
|
||||
target = cast(dict[str, Any], target).get(part)
|
||||
return target
|
||||
|
||||
|
||||
@@ -346,7 +346,7 @@ def _http_get(url: str, *, headers: dict[str, str] | None = None) -> dict[str, A
|
||||
response = client.get(url, headers=headers)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
return data if isinstance(data, dict) else {}
|
||||
return cast(dict[str, Any], data) if isinstance(data, dict) else {}
|
||||
|
||||
|
||||
def _http_post(url: str, *, headers: dict[str, str] | None = None) -> dict[str, Any]:
|
||||
@@ -354,7 +354,7 @@ def _http_post(url: str, *, headers: dict[str, str] | None = None) -> dict[str,
|
||||
response = client.post(url, headers=headers)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
return data if isinstance(data, dict) else {}
|
||||
return cast(dict[str, Any], data) if isinstance(data, dict) else {}
|
||||
|
||||
|
||||
def _probe_tcp(host: str, port: int, *, allow_loopback: bool = False) -> None:
|
||||
|
||||
@@ -11,7 +11,7 @@ import uuid
|
||||
from collections.abc import Callable
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import Any, Self
|
||||
from typing import Any, Self, TypeGuard, cast
|
||||
|
||||
from pydantic import Field, field_validator, model_validator
|
||||
from websockets.asyncio.server import ServerConnection, serve, unix_serve
|
||||
@@ -191,12 +191,13 @@ def _parse_inbound_payload(raw: str) -> str | None:
|
||||
return None
|
||||
if text.startswith("{"):
|
||||
try:
|
||||
data = json.loads(text)
|
||||
data = cast(object, json.loads(text))
|
||||
except json.JSONDecodeError:
|
||||
return text
|
||||
if isinstance(data, dict):
|
||||
payload = cast(dict[str, Any], data)
|
||||
for key in ("content", "text", "message"):
|
||||
value = data.get(key)
|
||||
value = payload.get(key)
|
||||
if isinstance(value, str) and value.strip():
|
||||
return value
|
||||
return None
|
||||
@@ -209,7 +210,7 @@ def _parse_inbound_payload(raw: str) -> str | None:
|
||||
_CHAT_ID_RE = re.compile(r"^[A-Za-z0-9_:-]{1,64}$")
|
||||
|
||||
|
||||
def _is_valid_chat_id(value: Any) -> bool:
|
||||
def _is_valid_chat_id(value: Any) -> TypeGuard[str]:
|
||||
return isinstance(value, str) and _CHAT_ID_RE.match(value) is not None
|
||||
|
||||
|
||||
@@ -224,15 +225,16 @@ def _parse_envelope(raw: str) -> dict[str, Any] | None:
|
||||
if not text.startswith("{"):
|
||||
return None
|
||||
try:
|
||||
data = json.loads(text)
|
||||
data = cast(object, json.loads(text))
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
if not isinstance(data, dict):
|
||||
return None
|
||||
t = data.get("type")
|
||||
envelope = cast(dict[str, Any], data)
|
||||
t = envelope.get("type")
|
||||
if not isinstance(t, str):
|
||||
return None
|
||||
return data
|
||||
return envelope
|
||||
|
||||
|
||||
def _is_websocket_upgrade(request: WsRequest) -> bool:
|
||||
@@ -264,13 +266,13 @@ class WebSocketChannel(BaseChannel):
|
||||
super().__init__(config, bus)
|
||||
self.config: WebSocketConfig = config
|
||||
# chat_id -> connections subscribed to it (fan-out target).
|
||||
self._subs: dict[str, set[Any]] = {}
|
||||
self._subs: dict[str, set[ServerConnection]] = {}
|
||||
# connection -> chat_ids it is subscribed to (O(1) cleanup on disconnect).
|
||||
self._conn_chats: dict[Any, set[str]] = {}
|
||||
self._conn_chats: dict[ServerConnection, set[str]] = {}
|
||||
# connection -> default chat_id for legacy frames that omit routing.
|
||||
self._conn_default: dict[Any, str] = {}
|
||||
self._conn_default: dict[ServerConnection, str] = {}
|
||||
# Connections authenticated with a one-time token from /webui/bootstrap.
|
||||
self._webui_connections: set[Any] = set()
|
||||
self._webui_connections: set[ServerConnection] = set()
|
||||
self._stop_event: asyncio.Event | None = None
|
||||
self._server_task: asyncio.Task[None] | None = None
|
||||
|
||||
@@ -286,15 +288,43 @@ class WebSocketChannel(BaseChannel):
|
||||
|
||||
# -- Subscription bookkeeping -------------------------------------------
|
||||
|
||||
def _workspace_controls_available(self, connection: Any) -> bool:
|
||||
def _workspace_controls_available(self, connection: ServerConnection) -> bool:
|
||||
return self._http_router.workspace_controls_available(connection)
|
||||
|
||||
def _attach(self, connection: Any, chat_id: str) -> None:
|
||||
def _attach(self, connection: ServerConnection, chat_id: str) -> None:
|
||||
"""Idempotently subscribe *connection* to *chat_id*."""
|
||||
self._subs.setdefault(chat_id, set()).add(connection)
|
||||
self._conn_chats.setdefault(connection, set()).add(chat_id)
|
||||
|
||||
def _cleanup_connection(self, connection: Any) -> None:
|
||||
async def send_webui_protocol_error(
|
||||
self,
|
||||
connection: ServerConnection,
|
||||
detail: str,
|
||||
) -> None:
|
||||
"""Send a stable protocol error from a WebUI-owned orchestration helper."""
|
||||
await self._send_event(connection, "error", detail=detail)
|
||||
|
||||
async def attach_webui_fork(
|
||||
self,
|
||||
connection: ServerConnection,
|
||||
*,
|
||||
fork_id: str,
|
||||
fork_key: str,
|
||||
) -> None:
|
||||
"""Attach and hydrate a newly created WebUI chat fork."""
|
||||
scope = self._workspaces.scope_for_session_key(fork_key)
|
||||
self._attach(connection, fork_id)
|
||||
await self._send_event(connection, "attached", chat_id=fork_id)
|
||||
await self._send_event(
|
||||
connection,
|
||||
"session_updated",
|
||||
chat_id=fork_id,
|
||||
scope="metadata",
|
||||
workspace_scope=scope.payload(),
|
||||
)
|
||||
await self._hydrate_after_subscribe(fork_id)
|
||||
|
||||
def _cleanup_connection(self, connection: ServerConnection) -> None:
|
||||
"""Remove *connection* from every subscription set; safe to call multiple times."""
|
||||
chat_ids = self._conn_chats.pop(connection, set())
|
||||
for cid in chat_ids:
|
||||
@@ -317,10 +347,11 @@ class WebSocketChannel(BaseChannel):
|
||||
if self.gateway.session_manager is None:
|
||||
return
|
||||
row = self.gateway.session_manager.read_session_file(f"websocket:{chat_id}")
|
||||
meta = row.get("metadata", {}) if isinstance(row, dict) else {}
|
||||
row_data = row if isinstance(row, dict) else {}
|
||||
meta = row_data.get("metadata", {})
|
||||
if not isinstance(meta, dict):
|
||||
meta = {}
|
||||
blob = goal_state_ws_blob(meta)
|
||||
blob = goal_state_ws_blob(cast(dict[str, Any], meta))
|
||||
if not blob.get("active"):
|
||||
return
|
||||
await self.send_goal_state(chat_id, blob)
|
||||
@@ -342,7 +373,12 @@ class WebSocketChannel(BaseChannel):
|
||||
await self._maybe_push_active_goal_state(chat_id)
|
||||
await self._maybe_push_turn_run_wall_clock(chat_id)
|
||||
|
||||
async def _send_event(self, connection: Any, event: str, **fields: Any) -> None:
|
||||
async def _send_event(
|
||||
self,
|
||||
connection: ServerConnection,
|
||||
event: str,
|
||||
**fields: Any,
|
||||
) -> None:
|
||||
"""Send a control event (attached, error, ...) to a single connection."""
|
||||
payload: dict[str, Any] = {"event": event}
|
||||
payload.update(fields)
|
||||
@@ -377,7 +413,7 @@ class WebSocketChannel(BaseChannel):
|
||||
|
||||
# -- HTTP dispatch ------------------------------------------------------
|
||||
|
||||
async def _dispatch_http(self, connection: Any, request: WsRequest) -> Any:
|
||||
async def _dispatch_http(self, connection: ServerConnection, request: WsRequest) -> Any:
|
||||
"""Route an inbound HTTP request to the HTTP handler or WS upgrade."""
|
||||
got, query = _parse_request_path(request.path)
|
||||
|
||||
@@ -394,7 +430,11 @@ class WebSocketChannel(BaseChannel):
|
||||
# Everything else goes to the HTTP handler
|
||||
return await self._http_router.dispatch(connection, request)
|
||||
|
||||
def _authorize_websocket_handshake(self, connection: Any, query: dict[str, list[str]]) -> Any:
|
||||
def _authorize_websocket_handshake(
|
||||
self,
|
||||
connection: ServerConnection,
|
||||
query: dict[str, list[str]],
|
||||
) -> Any:
|
||||
supplied = _query_first(query, "token")
|
||||
static_token = self.config.token.strip()
|
||||
|
||||
@@ -414,7 +454,7 @@ class WebSocketChannel(BaseChannel):
|
||||
self._consume_issued_token(connection, supplied)
|
||||
return None
|
||||
|
||||
def _consume_issued_token(self, connection: Any, token: str) -> bool:
|
||||
def _consume_issued_token(self, connection: ServerConnection, token: str) -> bool:
|
||||
audience = self._tokens.take_issued_token_audience(token)
|
||||
if audience == "webui":
|
||||
self._webui_connections.add(connection)
|
||||
@@ -509,7 +549,7 @@ class WebSocketChannel(BaseChannel):
|
||||
self._server_task = asyncio.create_task(runner())
|
||||
await self._server_task
|
||||
|
||||
async def _connection_loop(self, connection: Any) -> None:
|
||||
async def _connection_loop(self, connection: ServerConnection) -> None:
|
||||
request = connection.request
|
||||
path_part = request.path if request else "/"
|
||||
_, query = _parse_request_path(path_part)
|
||||
@@ -574,7 +614,7 @@ class WebSocketChannel(BaseChannel):
|
||||
|
||||
async def _dispatch_envelope(
|
||||
self,
|
||||
connection: Any,
|
||||
connection: ServerConnection,
|
||||
client_id: str,
|
||||
envelope: dict[str, Any],
|
||||
) -> None:
|
||||
@@ -700,7 +740,7 @@ class WebSocketChannel(BaseChannel):
|
||||
**rejection_fields,
|
||||
)
|
||||
return
|
||||
media_paths, reason = self._media.store_inbound_attachments(raw_media)
|
||||
media_paths, reason = self._media.store_inbound_attachments(cast(list[Any], raw_media))
|
||||
if reason is not None:
|
||||
await self._send_event(
|
||||
connection,
|
||||
@@ -810,7 +850,7 @@ class WebSocketChannel(BaseChannel):
|
||||
|
||||
async def _workspace_scope_or_error(
|
||||
self,
|
||||
connection: Any,
|
||||
connection: ServerConnection,
|
||||
resolver: Callable[[], Any],
|
||||
*,
|
||||
chat_id: str | None = None,
|
||||
@@ -841,7 +881,8 @@ class WebSocketChannel(BaseChannel):
|
||||
try:
|
||||
await self._server_task
|
||||
except asyncio.CancelledError:
|
||||
if asyncio.current_task() and asyncio.current_task().cancelling():
|
||||
current_task = asyncio.current_task()
|
||||
if current_task is not None and current_task.cancelling():
|
||||
raise
|
||||
self.logger.debug("server task was already cancelled during shutdown")
|
||||
except Exception as e:
|
||||
@@ -853,7 +894,13 @@ class WebSocketChannel(BaseChannel):
|
||||
self._webui_connections.clear()
|
||||
self._tokens.clear()
|
||||
|
||||
async def _safe_send_to(self, connection: Any, raw: str, *, label: str = "") -> None:
|
||||
async def _safe_send_to(
|
||||
self,
|
||||
connection: ServerConnection,
|
||||
raw: str,
|
||||
*,
|
||||
label: str = "",
|
||||
) -> None:
|
||||
"""Send a raw frame to one connection, cleaning up on ConnectionClosed."""
|
||||
try:
|
||||
await connection.send(raw)
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportMissingTypeStubs=false
|
||||
"""WeCom (Enterprise WeChat) channel implementation using wecom_aibot_sdk."""
|
||||
|
||||
import asyncio
|
||||
@@ -7,8 +8,9 @@ import importlib.util
|
||||
import os
|
||||
import re
|
||||
from collections import OrderedDict
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
@@ -96,7 +98,7 @@ class WecomChannel(BaseChannel):
|
||||
self._client: Any = None
|
||||
self._processed_message_ids: OrderedDict[str, None] = OrderedDict()
|
||||
self._loop: asyncio.AbstractEventLoop | None = None
|
||||
self._generate_req_id = None
|
||||
self._generate_req_id: Callable[[str], str] | None = None
|
||||
# Store frame headers for each chat to enable replies
|
||||
self._chat_frames: dict[str, Any] = {}
|
||||
|
||||
@@ -117,7 +119,8 @@ class WecomChannel(BaseChannel):
|
||||
self._generate_req_id = generate_req_id
|
||||
|
||||
# Create WebSocket client
|
||||
self._client = WSClient({
|
||||
ws_client = cast(Any, WSClient)
|
||||
self._client = ws_client({
|
||||
"bot_id": self.config.bot_id,
|
||||
"secret": self.config.secret,
|
||||
"reconnect_interval": 1000,
|
||||
@@ -195,14 +198,16 @@ class WecomChannel(BaseChannel):
|
||||
"""Handle enter_chat event (user opens chat with bot)."""
|
||||
try:
|
||||
# Extract body from WsFrame dataclass or dict
|
||||
if hasattr(frame, 'body'):
|
||||
body = frame.body or {}
|
||||
if hasattr(frame, "body"):
|
||||
body: Any = frame.body or {}
|
||||
elif isinstance(frame, dict):
|
||||
body = frame.get("body", frame)
|
||||
frame_dict = cast(dict[str, Any], frame)
|
||||
body = frame_dict.get("body", frame_dict)
|
||||
else:
|
||||
body = {}
|
||||
|
||||
chat_id = body.get("chatid", "") if isinstance(body, dict) else ""
|
||||
body_dict = cast(dict[str, Any], body) if isinstance(body, dict) else {}
|
||||
chat_id = cast(str, body_dict.get("chatid", ""))
|
||||
|
||||
if chat_id and not self.is_allowed(chat_id):
|
||||
return
|
||||
@@ -219,26 +224,32 @@ class WecomChannel(BaseChannel):
|
||||
"""Process incoming message and forward to bus."""
|
||||
try:
|
||||
# Extract body from WsFrame dataclass or dict
|
||||
if hasattr(frame, 'body'):
|
||||
body = frame.body or {}
|
||||
if hasattr(frame, "body"):
|
||||
body: Any = frame.body or {}
|
||||
elif isinstance(frame, dict):
|
||||
body = frame.get("body", frame)
|
||||
frame_dict = cast(dict[str, Any], frame)
|
||||
body = frame_dict.get("body", frame_dict)
|
||||
else:
|
||||
body = {}
|
||||
|
||||
# Ensure body is a dict
|
||||
if not isinstance(body, dict):
|
||||
self.logger.warning("Invalid body type: {}", type(body))
|
||||
self.logger.warning("Invalid body type: {}", type(cast(object, body)))
|
||||
return
|
||||
body = cast(dict[str, Any], body)
|
||||
|
||||
# Extract message info
|
||||
msg_id = body.get("msgid", "")
|
||||
msg_id = cast(str, body.get("msgid", ""))
|
||||
if not msg_id:
|
||||
msg_id = f"{body.get('chatid', '')}_{body.get('sendertime', '')}"
|
||||
|
||||
# Extract sender info from "from" field (SDK format)
|
||||
from_info = body.get("from", {})
|
||||
sender_id = from_info.get("userid", "unknown") if isinstance(from_info, dict) else "unknown"
|
||||
sender_id = (
|
||||
cast(str, cast(dict[str, Any], from_info).get("userid", "unknown"))
|
||||
if isinstance(from_info, dict)
|
||||
else "unknown"
|
||||
)
|
||||
if not self.is_allowed(sender_id):
|
||||
return
|
||||
|
||||
@@ -253,21 +264,22 @@ class WecomChannel(BaseChannel):
|
||||
|
||||
# For single chat, chatid is the sender's userid
|
||||
# For group chat, chatid is provided in body
|
||||
chat_type = body.get("chattype", "single")
|
||||
chat_id = body.get("chatid", sender_id)
|
||||
chat_type = cast(str, body.get("chattype", "single"))
|
||||
chat_id = cast(str, body.get("chatid", sender_id))
|
||||
|
||||
content_parts = []
|
||||
content_parts: list[str] = []
|
||||
media_paths: list[str] = []
|
||||
|
||||
if msg_type == "text":
|
||||
text = body.get("text", {}).get("content", "")
|
||||
text_info = cast(dict[str, Any], body.get("text", {}))
|
||||
text = cast(str, text_info.get("content", ""))
|
||||
if text:
|
||||
content_parts.append(text)
|
||||
|
||||
elif msg_type == "image":
|
||||
image_info = body.get("image", {})
|
||||
file_url = image_info.get("url", "")
|
||||
aes_key = image_info.get("aeskey", "")
|
||||
image_info = cast(dict[str, Any], body.get("image", {}))
|
||||
file_url = cast(str, image_info.get("url", ""))
|
||||
aes_key = cast(str, image_info.get("aeskey", ""))
|
||||
|
||||
if file_url and aes_key:
|
||||
file_path = await self._download_and_save_media(file_url, aes_key, "image")
|
||||
@@ -281,19 +293,19 @@ class WecomChannel(BaseChannel):
|
||||
content_parts.append("[image: download failed]")
|
||||
|
||||
elif msg_type == "voice":
|
||||
voice_info = body.get("voice", {})
|
||||
voice_info = cast(dict[str, Any], body.get("voice", {}))
|
||||
# Voice message already contains transcribed content from WeCom
|
||||
voice_content = voice_info.get("content", "")
|
||||
voice_content = cast(str, voice_info.get("content", ""))
|
||||
if voice_content:
|
||||
content_parts.append(f"[voice] {voice_content}")
|
||||
else:
|
||||
content_parts.append("[voice]")
|
||||
|
||||
elif msg_type == "file":
|
||||
file_info = body.get("file", {})
|
||||
file_url = file_info.get("url", "")
|
||||
aes_key = file_info.get("aeskey", "")
|
||||
file_name = file_info.get("name") or None
|
||||
file_info = cast(dict[str, Any], body.get("file", {}))
|
||||
file_url = cast(str, file_info.get("url", ""))
|
||||
aes_key = cast(str, file_info.get("aeskey", ""))
|
||||
file_name = cast(str | None, file_info.get("name") or None)
|
||||
|
||||
if file_url and aes_key:
|
||||
file_path = await self._download_and_save_media(file_url, aes_key, "file", file_name)
|
||||
@@ -308,16 +320,20 @@ class WecomChannel(BaseChannel):
|
||||
|
||||
elif msg_type == "mixed":
|
||||
# Mixed content contains multiple message items
|
||||
msg_items = body.get("mixed", {}).get("msg_item", [])
|
||||
for item in msg_items:
|
||||
item_type = item.get("msgtype", "")
|
||||
mixed_info = cast(dict[str, Any], body.get("mixed", {}))
|
||||
msg_items = cast(list[Any], mixed_info.get("msg_item", []))
|
||||
for raw_item in msg_items:
|
||||
item = cast(dict[str, Any], raw_item)
|
||||
item_type = cast(str, item.get("msgtype", ""))
|
||||
if item_type == "text":
|
||||
text = item.get("text", {}).get("content", "")
|
||||
text_info = cast(dict[str, Any], item.get("text", {}))
|
||||
text = cast(str, text_info.get("content", ""))
|
||||
if text:
|
||||
content_parts.append(text)
|
||||
elif item_type == "image":
|
||||
file_url = item.get("image", {}).get("url", "")
|
||||
aes_key = item.get("image", {}).get("aeskey", "")
|
||||
image_info = cast(dict[str, Any], item.get("image", {}))
|
||||
file_url = cast(str, image_info.get("url", ""))
|
||||
aes_key = cast(str, image_info.get("aeskey", ""))
|
||||
if file_url and aes_key:
|
||||
file_path = await self._download_and_save_media(file_url, aes_key, "image")
|
||||
if file_path:
|
||||
@@ -385,7 +401,7 @@ class WecomChannel(BaseChannel):
|
||||
media_dir = get_media_dir("wecom")
|
||||
if not filename:
|
||||
filename = fname or f"{media_type}_{hash(file_url) % 100000}"
|
||||
filename = _sanitize_filename(filename)
|
||||
filename = _sanitize_filename(cast(str, filename))
|
||||
|
||||
file_path = media_dir / filename
|
||||
await asyncio.to_thread(file_path.write_bytes, data)
|
||||
@@ -397,8 +413,10 @@ class WecomChannel(BaseChannel):
|
||||
return None
|
||||
|
||||
async def _upload_media_ws(
|
||||
self, client: Any, file_path: str,
|
||||
) -> "tuple[str, str] | tuple[None, None]":
|
||||
self,
|
||||
client: Any,
|
||||
file_path: str,
|
||||
) -> tuple[str, str] | tuple[None, None]:
|
||||
"""Upload a local file to WeCom via WebSocket 3-step protocol (base64).
|
||||
|
||||
Uses the WeCom WebSocket upload commands directly via
|
||||
@@ -417,7 +435,7 @@ class WecomChannel(BaseChannel):
|
||||
media_type = _guess_wecom_media_type(fname)
|
||||
|
||||
# Read file size and data in a thread to avoid blocking the event loop
|
||||
def _read_file():
|
||||
def _read_file() -> tuple[int, bytes]:
|
||||
file_size = os.path.getsize(file_path)
|
||||
if file_size > WECOM_UPLOAD_MAX_BYTES:
|
||||
raise ValueError(
|
||||
@@ -530,7 +548,10 @@ class WecomChannel(BaseChannel):
|
||||
# Both progress and final messages must use reply_stream (cmd="aibot_respond_msg").
|
||||
# The plain reply() uses cmd="reply" which does not support "text" msgtype
|
||||
# and causes errcode=40008 from WeCom API.
|
||||
stream_id = self._generate_req_id("stream")
|
||||
generate_req_id = self._generate_req_id
|
||||
if generate_req_id is None:
|
||||
raise RuntimeError("WeCom request-id generator is not initialized")
|
||||
stream_id = generate_req_id("stream")
|
||||
await self._client.reply_stream(
|
||||
frame,
|
||||
stream_id,
|
||||
|
||||
@@ -4,22 +4,22 @@ from __future__ import annotations
|
||||
|
||||
import secrets
|
||||
import time
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
from nanobot.channels.connect import ChannelConnectError, QueryParams, query_first
|
||||
from nanobot.config.loader import load_config
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.channels.weixin.runtime import WeixinChannel
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class WeixinConnectSession:
|
||||
id: str
|
||||
qrcode_id: str
|
||||
qr_url: str
|
||||
channel: Any
|
||||
channel: WeixinChannel
|
||||
current_poll_base_url: str
|
||||
refresh_count: int
|
||||
created_wall: float
|
||||
@@ -58,9 +58,8 @@ class WeixinConnectStore:
|
||||
channel = self._build_channel()
|
||||
if force:
|
||||
# Preserve the working account until a replacement scan succeeds.
|
||||
channel._token = ""
|
||||
channel._get_updates_buf = ""
|
||||
elif channel._load_state():
|
||||
channel.connect_reset_pending_credentials()
|
||||
elif channel.connect_load_state():
|
||||
return {
|
||||
"session_id": "",
|
||||
"status": "succeeded",
|
||||
@@ -68,13 +67,9 @@ class WeixinConnectStore:
|
||||
"interval_ms": 2000,
|
||||
}
|
||||
|
||||
channel._client = httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(60, connect=30),
|
||||
follow_redirects=True,
|
||||
)
|
||||
channel._running = True
|
||||
channel.connect_open_client()
|
||||
try:
|
||||
qrcode_id, qr_url = await channel._fetch_qr_code()
|
||||
qrcode_id, qr_url = await channel.connect_fetch_qr_code()
|
||||
except Exception as exc:
|
||||
await self._close_channel(channel)
|
||||
raise ChannelConnectError(
|
||||
@@ -89,7 +84,7 @@ class WeixinConnectStore:
|
||||
qrcode_id=qrcode_id,
|
||||
qr_url=qr_url,
|
||||
channel=channel,
|
||||
current_poll_base_url=channel.config.base_url,
|
||||
current_poll_base_url=channel.connect_base_url,
|
||||
refresh_count=0,
|
||||
created_wall=now_wall,
|
||||
deadline=time.monotonic() + 600,
|
||||
@@ -107,14 +102,12 @@ class WeixinConnectStore:
|
||||
}
|
||||
|
||||
try:
|
||||
status_data = await session.channel._api_get_with_base(
|
||||
status_data = await session.channel.connect_poll_qr_code(
|
||||
base_url=session.current_poll_base_url,
|
||||
endpoint="ilink/bot/get_qrcode_status",
|
||||
params={"qrcode": session.qrcode_id},
|
||||
auth=False,
|
||||
qrcode_id=session.qrcode_id,
|
||||
)
|
||||
except Exception as exc:
|
||||
if session.channel._is_retryable_qr_poll_error(exc):
|
||||
if session.channel.connect_poll_error_is_retryable(exc):
|
||||
session.last_error = str(exc)
|
||||
return self._pending_payload(session)
|
||||
self._sessions.pop(session_id, None)
|
||||
@@ -125,10 +118,8 @@ class WeixinConnectStore:
|
||||
"message": f"WeChat QR login failed: {exc}",
|
||||
}
|
||||
|
||||
if not isinstance(status_data, dict):
|
||||
return self._pending_payload(session)
|
||||
|
||||
status = status_data.get("status", "")
|
||||
status_payload = status_data
|
||||
status = status_payload.get("status", "")
|
||||
if status == "confirmed":
|
||||
if self._sessions.get(session_id) is not session:
|
||||
return {
|
||||
@@ -136,7 +127,7 @@ class WeixinConnectStore:
|
||||
"status": "cancelled",
|
||||
"message": "WeChat login cancelled.",
|
||||
}
|
||||
token = str(status_data.get("bot_token", "") or "")
|
||||
token = str(status_payload.get("bot_token", "") or "")
|
||||
if not token:
|
||||
self._sessions.pop(session_id, None)
|
||||
await self._close_channel(session.channel)
|
||||
@@ -145,22 +136,19 @@ class WeixinConnectStore:
|
||||
"status": "failed",
|
||||
"message": "WeChat confirmed the scan but returned no token.",
|
||||
}
|
||||
base_url = str(status_data.get("baseurl", "") or "")
|
||||
session.channel._token = token
|
||||
if base_url:
|
||||
session.channel.config.base_url = base_url
|
||||
session.channel._save_state()
|
||||
base_url = str(status_payload.get("baseurl", "") or "")
|
||||
session.channel.connect_commit_account(token=token, base_url=base_url)
|
||||
self._sessions.pop(session_id, None)
|
||||
await self._close_channel(session.channel)
|
||||
return {
|
||||
"session_id": session_id,
|
||||
"status": "succeeded",
|
||||
"message": "WeChat is connected.",
|
||||
"account": str(status_data.get("ilink_user_id", "") or ""),
|
||||
"account": str(status_payload.get("ilink_user_id", "") or ""),
|
||||
}
|
||||
|
||||
if status == "scaned_but_redirect":
|
||||
redirect_host = str(status_data.get("redirect_host", "") or "").strip()
|
||||
redirect_host = str(status_payload.get("redirect_host", "") or "").strip()
|
||||
if redirect_host:
|
||||
session.current_poll_base_url = (
|
||||
redirect_host
|
||||
@@ -182,7 +170,9 @@ class WeixinConnectStore:
|
||||
"message": "This WeChat QR code expired. Start again.",
|
||||
}
|
||||
try:
|
||||
session.qrcode_id, session.qr_url = await session.channel._fetch_qr_code()
|
||||
session.qrcode_id, session.qr_url = (
|
||||
await session.channel.connect_fetch_qr_code()
|
||||
)
|
||||
except Exception as exc:
|
||||
self._sessions.pop(session_id, None)
|
||||
await self._close_channel(session.channel)
|
||||
@@ -191,7 +181,7 @@ class WeixinConnectStore:
|
||||
"status": "failed",
|
||||
"message": f"Could not refresh WeChat QR code: {exc}",
|
||||
}
|
||||
session.current_poll_base_url = session.channel.config.base_url
|
||||
session.current_poll_base_url = session.channel.connect_base_url
|
||||
return self._pending_payload(session)
|
||||
|
||||
return self._pending_payload(session)
|
||||
@@ -219,27 +209,22 @@ class WeixinConnectStore:
|
||||
await self._close_channel(session.channel)
|
||||
|
||||
@staticmethod
|
||||
def _build_channel() -> Any:
|
||||
def _build_channel() -> WeixinChannel:
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.weixin.runtime import WeixinChannel
|
||||
|
||||
section = getattr(load_config().channels, "weixin", None)
|
||||
if hasattr(section, "model_dump"):
|
||||
if section is not None and hasattr(section, "model_dump"):
|
||||
config = section.model_dump(mode="json", by_alias=True)
|
||||
elif isinstance(section, dict):
|
||||
config = dict(section)
|
||||
config = dict(cast(dict[str, Any], section))
|
||||
else:
|
||||
config = {}
|
||||
return WeixinChannel(config, MessageBus())
|
||||
|
||||
@staticmethod
|
||||
async def _close_channel(channel: Any) -> None:
|
||||
channel._running = False
|
||||
client = getattr(channel, "_client", None)
|
||||
if client is not None:
|
||||
with suppress(Exception):
|
||||
await client.aclose()
|
||||
channel._client = None
|
||||
async def _close_channel(channel: WeixinChannel) -> None:
|
||||
await channel.connect_close_client()
|
||||
|
||||
@staticmethod
|
||||
def _start_payload(session: WeixinConnectSession) -> dict[str, Any]:
|
||||
|
||||
@@ -21,7 +21,7 @@ import uuid
|
||||
from collections import OrderedDict
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
@@ -168,10 +168,10 @@ class WeixinChannel(BaseChannel):
|
||||
self._processed_ids: OrderedDict[str, None] = OrderedDict()
|
||||
self._state_dir: Path | None = None
|
||||
self._token: str = ""
|
||||
self._poll_task: asyncio.Task | None = None
|
||||
self._poll_task: asyncio.Task[None] | None = None
|
||||
self._next_poll_timeout_s: int = DEFAULT_LONG_POLL_TIMEOUT_S
|
||||
self._session_pause_until: float = 0.0
|
||||
self._typing_tasks: dict[str, asyncio.Task] = {}
|
||||
self._typing_tasks: dict[str, asyncio.Task[None]] = {}
|
||||
self._typing_tickets: dict[str, dict[str, Any]] = {}
|
||||
self._context_token_at: dict[str, float] = {}
|
||||
self._pending_tool_hints: dict[str, list[str]] = {}
|
||||
@@ -201,14 +201,14 @@ class WeixinChannel(BaseChannel):
|
||||
if not state_file.exists():
|
||||
return False
|
||||
try:
|
||||
data = json.loads(state_file.read_text())
|
||||
data = cast(dict[str, Any], json.loads(state_file.read_text()))
|
||||
self._token = data.get("token", "")
|
||||
self._get_updates_buf = data.get("get_updates_buf", "")
|
||||
context_tokens = data.get("context_tokens", {})
|
||||
if isinstance(context_tokens, dict):
|
||||
self._context_tokens = {
|
||||
str(user_id): str(token)
|
||||
for user_id, token in context_tokens.items()
|
||||
for user_id, token in cast(dict[object, object], context_tokens).items()
|
||||
if str(user_id).strip() and str(token).strip()
|
||||
}
|
||||
else:
|
||||
@@ -216,8 +216,8 @@ class WeixinChannel(BaseChannel):
|
||||
typing_tickets = data.get("typing_tickets", {})
|
||||
if isinstance(typing_tickets, dict):
|
||||
self._typing_tickets = {
|
||||
str(user_id): ticket
|
||||
for user_id, ticket in typing_tickets.items()
|
||||
str(user_id): cast(dict[str, Any], ticket)
|
||||
for user_id, ticket in cast(dict[object, object], typing_tickets).items()
|
||||
if str(user_id).strip() and isinstance(ticket, dict)
|
||||
}
|
||||
else:
|
||||
@@ -276,18 +276,22 @@ class WeixinChannel(BaseChannel):
|
||||
if isinstance(err, httpx.TimeoutException | httpx.TransportError):
|
||||
return True
|
||||
if isinstance(err, httpx.HTTPStatusError):
|
||||
status_code = err.response.status_code if err.response is not None else 0
|
||||
status_code = (
|
||||
err.response.status_code
|
||||
if cast(object, err.response) is not None
|
||||
else 0
|
||||
)
|
||||
return status_code >= 500
|
||||
return False
|
||||
|
||||
async def _api_get(
|
||||
self,
|
||||
endpoint: str,
|
||||
params: dict | None = None,
|
||||
params: dict[str, Any] | None = None,
|
||||
*,
|
||||
auth: bool = True,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> dict:
|
||||
) -> dict[str, Any]:
|
||||
assert self._client is not None
|
||||
url = f"{self.config.base_url}/{endpoint}"
|
||||
hdrs = self._make_headers(auth=auth)
|
||||
@@ -295,17 +299,17 @@ class WeixinChannel(BaseChannel):
|
||||
hdrs.update(extra_headers)
|
||||
resp = await self._client.get(url, params=params, headers=hdrs)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
return cast(dict[str, Any], resp.json())
|
||||
|
||||
async def _api_get_with_base(
|
||||
self,
|
||||
*,
|
||||
base_url: str,
|
||||
endpoint: str,
|
||||
params: dict | None = None,
|
||||
params: dict[str, Any] | None = None,
|
||||
auth: bool = True,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
) -> dict:
|
||||
) -> dict[str, Any]:
|
||||
"""GET helper that allows overriding base_url for QR redirect polling."""
|
||||
assert self._client is not None
|
||||
url = f"{base_url.rstrip('/')}/{endpoint}"
|
||||
@@ -314,15 +318,15 @@ class WeixinChannel(BaseChannel):
|
||||
hdrs.update(extra_headers)
|
||||
resp = await self._client.get(url, params=params, headers=hdrs)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
return cast(dict[str, Any], resp.json())
|
||||
|
||||
async def _api_post(
|
||||
self,
|
||||
endpoint: str,
|
||||
body: dict | None = None,
|
||||
body: dict[str, Any] | None = None,
|
||||
*,
|
||||
auth: bool = True,
|
||||
) -> dict:
|
||||
) -> dict[str, Any]:
|
||||
assert self._client is not None
|
||||
url = f"{self.config.base_url}/{endpoint}"
|
||||
payload = body or {}
|
||||
@@ -330,7 +334,7 @@ class WeixinChannel(BaseChannel):
|
||||
payload["base_info"] = BASE_INFO
|
||||
resp = await self._client.post(url, json=payload, headers=self._make_headers(auth=auth))
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
return cast(dict[str, Any], resp.json())
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# QR Code Login (matches login-qr.ts)
|
||||
@@ -343,8 +347,8 @@ class WeixinChannel(BaseChannel):
|
||||
params={"bot_type": "3"},
|
||||
auth=False,
|
||||
)
|
||||
qrcode_img_content = data.get("qrcode_img_content", "")
|
||||
qrcode_id = data.get("qrcode", "")
|
||||
qrcode_img_content = cast(str, data.get("qrcode_img_content", ""))
|
||||
qrcode_id = cast(str, data.get("qrcode", ""))
|
||||
if not qrcode_id:
|
||||
raise RuntimeError(f"Failed to get QR code from WeChat API: {data}")
|
||||
return qrcode_id, (qrcode_img_content or qrcode_id)
|
||||
@@ -371,7 +375,7 @@ class WeixinChannel(BaseChannel):
|
||||
continue
|
||||
raise
|
||||
|
||||
if not isinstance(status_data, dict):
|
||||
if not isinstance(cast(object, status_data), dict):
|
||||
await asyncio.sleep(1)
|
||||
continue
|
||||
|
||||
@@ -431,15 +435,73 @@ class WeixinChannel(BaseChannel):
|
||||
if isinstance(err, httpx.TimeoutException | httpx.TransportError):
|
||||
return True
|
||||
if isinstance(err, httpx.HTTPStatusError):
|
||||
status_code = err.response.status_code if err.response is not None else 0
|
||||
status_code = (
|
||||
err.response.status_code
|
||||
if cast(object, err.response) is not None
|
||||
else 0
|
||||
)
|
||||
if status_code >= 500:
|
||||
return True
|
||||
return False
|
||||
|
||||
@property
|
||||
def connect_base_url(self) -> str:
|
||||
"""Base URL currently selected for the interactive connection flow."""
|
||||
return self.config.base_url
|
||||
|
||||
def connect_reset_pending_credentials(self) -> None:
|
||||
"""Clear only in-memory credentials while a replacement QR login is pending."""
|
||||
self._token = ""
|
||||
self._get_updates_buf = ""
|
||||
|
||||
def connect_load_state(self) -> bool:
|
||||
"""Load an existing account for the interactive connection flow."""
|
||||
return self._load_state()
|
||||
|
||||
def connect_open_client(self) -> None:
|
||||
"""Open the short-lived HTTP client used by WebUI QR login."""
|
||||
self._client = httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(60, connect=30),
|
||||
follow_redirects=True,
|
||||
)
|
||||
self._running = True
|
||||
|
||||
async def connect_fetch_qr_code(self) -> tuple[str, str]:
|
||||
return await self._fetch_qr_code()
|
||||
|
||||
async def connect_poll_qr_code(
|
||||
self,
|
||||
*,
|
||||
base_url: str,
|
||||
qrcode_id: str,
|
||||
) -> dict[str, Any]:
|
||||
return await self._api_get_with_base(
|
||||
base_url=base_url,
|
||||
endpoint="ilink/bot/get_qrcode_status",
|
||||
params={"qrcode": qrcode_id},
|
||||
auth=False,
|
||||
)
|
||||
|
||||
def connect_poll_error_is_retryable(self, err: Exception) -> bool:
|
||||
return self._is_retryable_qr_poll_error(err)
|
||||
|
||||
def connect_commit_account(self, *, token: str, base_url: str) -> None:
|
||||
self._token = token
|
||||
if base_url:
|
||||
self.config.base_url = base_url
|
||||
self._save_state()
|
||||
|
||||
async def connect_close_client(self) -> None:
|
||||
self._running = False
|
||||
if self._client is not None:
|
||||
with suppress(Exception):
|
||||
await self._client.aclose()
|
||||
self._client = None
|
||||
|
||||
@staticmethod
|
||||
def _print_qr_code(url: str) -> None:
|
||||
try:
|
||||
import qrcode as qr_lib
|
||||
import qrcode as qr_lib # pyright: ignore[reportMissingModuleSource]
|
||||
|
||||
qr = qr_lib.QRCode(border=1)
|
||||
qr.add_data(url)
|
||||
@@ -596,7 +658,7 @@ class WeixinChannel(BaseChannel):
|
||||
self._save_state()
|
||||
|
||||
# Process messages (WeixinMessage[] from types.ts)
|
||||
msgs: list[dict] = data.get("msgs", []) or []
|
||||
msgs = cast(list[dict[str, Any]], data.get("msgs", []) or [])
|
||||
for msg in msgs:
|
||||
try:
|
||||
await self._process_message(msg)
|
||||
@@ -607,7 +669,7 @@ class WeixinChannel(BaseChannel):
|
||||
# Inbound message processing (matches inbound.ts + process-message.ts)
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _process_message(self, msg: dict) -> None:
|
||||
async def _process_message(self, msg: dict[str, Any]) -> None:
|
||||
"""Process a single WeixinMessage from getUpdates."""
|
||||
# Skip bot's own messages (message_type 2 = BOT)
|
||||
if msg.get("message_type") == MESSAGE_TYPE_BOT:
|
||||
@@ -679,7 +741,7 @@ class WeixinChannel(BaseChannel):
|
||||
self._save_state()
|
||||
|
||||
# Parse item_list (WeixinMessage.item_list — types.ts:161)
|
||||
item_list: list[dict] = msg.get("item_list") or []
|
||||
item_list = cast(list[dict[str, Any]], msg.get("item_list") or [])
|
||||
content_parts: list[str] = []
|
||||
media_paths: list[str] = []
|
||||
has_top_level_downloadable_media = False
|
||||
@@ -688,12 +750,16 @@ class WeixinChannel(BaseChannel):
|
||||
item_type = item.get("type", 0)
|
||||
|
||||
if item_type == ITEM_TEXT:
|
||||
text = (item.get("text_item") or {}).get("text", "")
|
||||
text_item = cast(dict[str, Any], item.get("text_item") or {})
|
||||
text = cast(str, text_item.get("text", ""))
|
||||
if text:
|
||||
# Handle quoted/ref messages (inbound.ts:86-98)
|
||||
ref = item.get("ref_msg")
|
||||
ref = cast(dict[str, Any] | None, item.get("ref_msg"))
|
||||
if ref:
|
||||
ref_item = ref.get("message_item")
|
||||
ref_item = cast(
|
||||
dict[str, Any] | None,
|
||||
ref.get("message_item"),
|
||||
)
|
||||
# If quoted message is media, just pass the text
|
||||
if ref_item and ref_item.get("type", 0) in (
|
||||
ITEM_IMAGE,
|
||||
@@ -705,9 +771,13 @@ class WeixinChannel(BaseChannel):
|
||||
else:
|
||||
parts: list[str] = []
|
||||
if ref.get("title"):
|
||||
parts.append(ref["title"])
|
||||
parts.append(cast(str, ref["title"]))
|
||||
if ref_item:
|
||||
ref_text = (ref_item.get("text_item") or {}).get("text", "")
|
||||
ref_text_item = cast(
|
||||
dict[str, Any],
|
||||
ref_item.get("text_item") or {},
|
||||
)
|
||||
ref_text = cast(str, ref_text_item.get("text", ""))
|
||||
if ref_text:
|
||||
parts.append(ref_text)
|
||||
if parts:
|
||||
@@ -718,7 +788,7 @@ class WeixinChannel(BaseChannel):
|
||||
content_parts.append(text)
|
||||
|
||||
elif item_type == ITEM_IMAGE:
|
||||
image_item = item.get("image_item") or {}
|
||||
image_item = cast(dict[str, Any], item.get("image_item") or {})
|
||||
if _has_downloadable_media_locator(image_item.get("media")):
|
||||
has_top_level_downloadable_media = True
|
||||
file_path = await self._download_media_item(image_item, "image")
|
||||
@@ -729,9 +799,9 @@ class WeixinChannel(BaseChannel):
|
||||
content_parts.append("[image]")
|
||||
|
||||
elif item_type == ITEM_VOICE:
|
||||
voice_item = item.get("voice_item") or {}
|
||||
voice_item = cast(dict[str, Any], item.get("voice_item") or {})
|
||||
# Voice-to-text provided by WeChat (inbound.ts:101-103)
|
||||
voice_text = voice_item.get("text", "")
|
||||
voice_text = cast(str, voice_item.get("text", ""))
|
||||
if voice_text:
|
||||
content_parts.append(f"[voice] {voice_text}")
|
||||
else:
|
||||
@@ -749,10 +819,10 @@ class WeixinChannel(BaseChannel):
|
||||
content_parts.append("[voice]")
|
||||
|
||||
elif item_type == ITEM_FILE:
|
||||
file_item = item.get("file_item") or {}
|
||||
file_item = cast(dict[str, Any], item.get("file_item") or {})
|
||||
if _has_downloadable_media_locator(file_item.get("media")):
|
||||
has_top_level_downloadable_media = True
|
||||
file_name = file_item.get("file_name", "unknown")
|
||||
file_name = cast(str, file_item.get("file_name", "unknown"))
|
||||
file_path = await self._download_media_item(
|
||||
file_item,
|
||||
"file",
|
||||
@@ -765,7 +835,7 @@ class WeixinChannel(BaseChannel):
|
||||
content_parts.append(f"[file: {file_name}]")
|
||||
|
||||
elif item_type == ITEM_VIDEO:
|
||||
video_item = item.get("video_item") or {}
|
||||
video_item = cast(dict[str, Any], item.get("video_item") or {})
|
||||
if _has_downloadable_media_locator(video_item.get("media")):
|
||||
has_top_level_downloadable_media = True
|
||||
file_path = await self._download_media_item(video_item, "video")
|
||||
@@ -783,8 +853,8 @@ class WeixinChannel(BaseChannel):
|
||||
for item in item_list:
|
||||
if item.get("type", 0) != ITEM_TEXT:
|
||||
continue
|
||||
ref = item.get("ref_msg") or {}
|
||||
candidate = ref.get("message_item") or {}
|
||||
ref = cast(dict[str, Any], item.get("ref_msg") or {})
|
||||
candidate = cast(dict[str, Any], ref.get("message_item") or {})
|
||||
if candidate.get("type", 0) in (ITEM_IMAGE, ITEM_VOICE, ITEM_FILE, ITEM_VIDEO):
|
||||
ref_media_item = candidate
|
||||
break
|
||||
@@ -792,13 +862,19 @@ class WeixinChannel(BaseChannel):
|
||||
if ref_media_item:
|
||||
ref_type = ref_media_item.get("type", 0)
|
||||
if ref_type == ITEM_IMAGE:
|
||||
image_item = ref_media_item.get("image_item") or {}
|
||||
image_item = cast(
|
||||
dict[str, Any],
|
||||
ref_media_item.get("image_item") or {},
|
||||
)
|
||||
file_path = await self._download_media_item(image_item, "image")
|
||||
if file_path:
|
||||
content_parts.append(f"[image]\n[Image: source: {file_path}]")
|
||||
media_paths.append(file_path)
|
||||
elif ref_type == ITEM_VOICE:
|
||||
voice_item = ref_media_item.get("voice_item") or {}
|
||||
voice_item = cast(
|
||||
dict[str, Any],
|
||||
ref_media_item.get("voice_item") or {},
|
||||
)
|
||||
file_path = await self._download_media_item(voice_item, "voice")
|
||||
if file_path:
|
||||
transcription = await self.transcribe_audio(file_path)
|
||||
@@ -808,14 +884,20 @@ class WeixinChannel(BaseChannel):
|
||||
content_parts.append(f"[voice]\n[Audio: source: {file_path}]")
|
||||
media_paths.append(file_path)
|
||||
elif ref_type == ITEM_FILE:
|
||||
file_item = ref_media_item.get("file_item") or {}
|
||||
file_name = file_item.get("file_name", "unknown")
|
||||
file_item = cast(
|
||||
dict[str, Any],
|
||||
ref_media_item.get("file_item") or {},
|
||||
)
|
||||
file_name = cast(str, file_item.get("file_name", "unknown"))
|
||||
file_path = await self._download_media_item(file_item, "file", file_name)
|
||||
if file_path:
|
||||
content_parts.append(f"[file: {file_name}]\n[File: source: {file_path}]")
|
||||
media_paths.append(file_path)
|
||||
elif ref_type == ITEM_VIDEO:
|
||||
video_item = ref_media_item.get("video_item") or {}
|
||||
video_item = cast(
|
||||
dict[str, Any],
|
||||
ref_media_item.get("video_item") or {},
|
||||
)
|
||||
file_path = await self._download_media_item(video_item, "video")
|
||||
if file_path:
|
||||
content_parts.append(f"[video]\n[Video: source: {file_path}]")
|
||||
@@ -848,13 +930,13 @@ class WeixinChannel(BaseChannel):
|
||||
|
||||
async def _download_media_item(
|
||||
self,
|
||||
typed_item: dict,
|
||||
typed_item: dict[str, Any],
|
||||
media_type: str,
|
||||
filename: str | None = None,
|
||||
) -> str | None:
|
||||
"""Download + AES-decrypt a media item. Returns local path or None."""
|
||||
try:
|
||||
media = typed_item.get("media") or {}
|
||||
media = cast(dict[str, Any], typed_item.get("media") or {})
|
||||
encrypt_query_param = str(media.get("encrypt_query_param", "") or "")
|
||||
full_url = str(media.get("full_url", "") or "").strip()
|
||||
|
||||
@@ -865,8 +947,8 @@ class WeixinChannel(BaseChannel):
|
||||
# image_item.aeskey is a raw hex string (16 bytes as 32 hex chars).
|
||||
# media.aes_key is always base64-encoded.
|
||||
# For images, prefer image_item.aeskey; for others use media.aes_key.
|
||||
raw_aeskey_hex = typed_item.get("aeskey", "")
|
||||
media_aes_key_b64 = media.get("aes_key", "")
|
||||
raw_aeskey_hex = cast(str, typed_item.get("aeskey", ""))
|
||||
media_aes_key_b64 = cast(str, media.get("aes_key", ""))
|
||||
|
||||
aes_key_b64: str = ""
|
||||
if raw_aeskey_hex:
|
||||
@@ -1160,7 +1242,7 @@ class WeixinChannel(BaseChannel):
|
||||
await self._send_typing(msg.chat_id, typing_ticket, TYPING_STATUS_TYPING)
|
||||
|
||||
typing_keepalive_stop = asyncio.Event()
|
||||
typing_keepalive_task: asyncio.Task | None = None
|
||||
typing_keepalive_task: asyncio.Task[None] | None = None
|
||||
if typing_ticket:
|
||||
typing_keepalive_task = asyncio.create_task(
|
||||
self._typing_keepalive_loop(msg.chat_id, typing_ticket, typing_keepalive_stop)
|
||||
@@ -1183,7 +1265,7 @@ class WeixinChannel(BaseChannel):
|
||||
except httpx.HTTPStatusError as http_err:
|
||||
status_code = (
|
||||
http_err.response.status_code
|
||||
if http_err.response is not None
|
||||
if cast(object, http_err.response) is not None
|
||||
else 0
|
||||
)
|
||||
if status_code >= 500:
|
||||
@@ -1192,7 +1274,7 @@ class WeixinChannel(BaseChannel):
|
||||
"Server error ({} {}) sending media {}",
|
||||
status_code,
|
||||
http_err.response.reason_phrase
|
||||
if http_err.response is not None
|
||||
if cast(object, http_err.response) is not None
|
||||
else "",
|
||||
media_path,
|
||||
)
|
||||
@@ -1342,7 +1424,7 @@ class WeixinChannel(BaseChannel):
|
||||
"""Send a text message matching the exact protocol from send.ts."""
|
||||
client_id = f"nanobot-{uuid.uuid4().hex[:12]}"
|
||||
|
||||
item_list: list[dict] = []
|
||||
item_list: list[dict[str, Any]] = []
|
||||
if text:
|
||||
item_list.append({"type": ITEM_TEXT, "text_item": {"text": text}})
|
||||
|
||||
@@ -1496,7 +1578,9 @@ class WeixinChannel(BaseChannel):
|
||||
|
||||
# Send each media item as its own message (matching reference plugin)
|
||||
client_id = f"nanobot-{uuid.uuid4().hex[:12]}"
|
||||
item_list: list[dict] = [{"type": item_type, item_key: media_item}]
|
||||
item_list: list[dict[str, Any]] = [
|
||||
{"type": item_type, item_key: media_item}
|
||||
]
|
||||
|
||||
weixin_msg: dict[str, Any] = {
|
||||
"from_user_id": "",
|
||||
@@ -1565,7 +1649,8 @@ def _encrypt_aes_ecb(data: bytes, aes_key_b64: str) -> bytes:
|
||||
with suppress(ImportError):
|
||||
from Crypto.Cipher import AES
|
||||
|
||||
cipher = AES.new(key, AES.MODE_ECB)
|
||||
aes_module = cast(Any, AES)
|
||||
cipher = aes_module.new(key, aes_module.MODE_ECB)
|
||||
return cipher.encrypt(padded)
|
||||
|
||||
try:
|
||||
@@ -1595,7 +1680,8 @@ def _decrypt_aes_ecb(data: bytes, aes_key_b64: str) -> bytes:
|
||||
with suppress(ImportError):
|
||||
from Crypto.Cipher import AES
|
||||
|
||||
cipher = AES.new(key, AES.MODE_ECB)
|
||||
aes_module = cast(Any, AES)
|
||||
cipher = aes_module.new(key, aes_module.MODE_ECB)
|
||||
decrypted = cipher.decrypt(data)
|
||||
|
||||
if decrypted is None:
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportConstantRedefinition=false, reportMissingTypeStubs=false, reportUnusedFunction=false
|
||||
"""WhatsApp channel implementation using neonize."""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -10,7 +11,7 @@ import time
|
||||
from collections import OrderedDict
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal, NamedTuple
|
||||
from typing import Any, Literal, NamedTuple, cast
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
@@ -100,7 +101,8 @@ def _has_field(message: Any, name: str) -> bool:
|
||||
list_fields = getattr(message, "ListFields", None)
|
||||
if callable(list_fields):
|
||||
try:
|
||||
return any(getattr(field, "name", "") == name for field, _ in list_fields())
|
||||
fields = cast(list[tuple[Any, Any]], list_fields())
|
||||
return any(getattr(field, "name", "") == name for field, _ in fields)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@@ -277,7 +279,10 @@ class WhatsAppChannel(BaseChannel):
|
||||
return WhatsAppConfig().model_dump(by_alias=True)
|
||||
|
||||
def __init__(self, config: Any, bus: MessageBus):
|
||||
legacy_bridge_fields = _legacy_bridge_config_fields(config) if isinstance(config, dict) else []
|
||||
legacy_bridge_fields = (
|
||||
_legacy_bridge_config_fields(cast(dict[str, Any], config))
|
||||
if isinstance(config, dict) else []
|
||||
)
|
||||
if isinstance(config, dict):
|
||||
config = WhatsAppConfig.model_validate(config)
|
||||
super().__init__(config, bus)
|
||||
@@ -649,12 +654,13 @@ class WhatsAppChannel(BaseChannel):
|
||||
if not self._self_jids:
|
||||
return False
|
||||
for context in _context_infos(message):
|
||||
mentioned = (
|
||||
raw_mentioned: Any = (
|
||||
_safe_attr(context, "mentionedJID")
|
||||
or _safe_attr(context, "mentionedJid")
|
||||
or _safe_attr(context, "mentioned_jid")
|
||||
or []
|
||||
)
|
||||
mentioned: list[Any] = cast(list[Any], raw_mentioned)
|
||||
for jid in mentioned:
|
||||
normalized = _normalize_jid(jid)
|
||||
if normalized in self._self_jids or _bare_jid(normalized) in self._self_jids:
|
||||
|
||||
+135
-106
@@ -1,15 +1,23 @@
|
||||
"""CLI commands for nanobot."""
|
||||
|
||||
# pyright: reportConstantRedefinition=false, reportMissingTypeStubs=false, reportPrivateUsage=false, reportUnusedFunction=false
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import select
|
||||
import signal
|
||||
import sys
|
||||
import time
|
||||
from collections.abc import Callable, Iterable
|
||||
from collections.abc import Awaitable, Callable, Coroutine, Iterable
|
||||
from contextlib import nullcontext, suppress
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from types import FrameType
|
||||
from typing import TYPE_CHECKING, Any, Literal, cast
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.gateway.runtime import GatewayRuntime
|
||||
from nanobot.providers.registry import ProviderSpec
|
||||
|
||||
|
||||
# Force UTF-8 encoding for Windows console
|
||||
if sys.platform == "win32":
|
||||
@@ -17,8 +25,10 @@ if sys.platform == "win32":
|
||||
os.environ["PYTHONIOENCODING"] = "utf-8"
|
||||
# Re-open stdout/stderr with UTF-8 encoding
|
||||
with suppress(Exception):
|
||||
sys.stdout.reconfigure(encoding="utf-8", errors="replace")
|
||||
sys.stderr.reconfigure(encoding="utf-8", errors="replace")
|
||||
for stream in (sys.stdout, sys.stderr):
|
||||
reconfigure = getattr(stream, "reconfigure", None)
|
||||
if callable(reconfigure):
|
||||
reconfigure(encoding="utf-8", errors="replace")
|
||||
|
||||
# Keep console encoding setup before importing CLI UI/logging libraries.
|
||||
import typer # noqa: E402
|
||||
@@ -52,6 +62,7 @@ from prompt_toolkit.application import run_in_terminal # noqa: E402
|
||||
from prompt_toolkit.formatted_text import ANSI, HTML # noqa: E402
|
||||
from prompt_toolkit.history import FileHistory # noqa: E402
|
||||
from prompt_toolkit.key_binding import KeyBindings # noqa: E402
|
||||
from prompt_toolkit.key_binding.key_processor import KeyPressEvent # noqa: E402
|
||||
from prompt_toolkit.keys import Keys # noqa: E402
|
||||
from prompt_toolkit.patch_stdout import patch_stdout # noqa: E402
|
||||
from pydantic import ValidationError # noqa: E402
|
||||
@@ -139,7 +150,7 @@ def _ensure_interactive_tty_mode() -> None:
|
||||
def _install_gateway_shutdown_handlers(
|
||||
loop: asyncio.AbstractEventLoop,
|
||||
shutdown_event: asyncio.Event,
|
||||
tasks: list[asyncio.Task],
|
||||
tasks: list[asyncio.Task[Any]],
|
||||
print_status: Callable[[str], None],
|
||||
) -> Callable[[], None]:
|
||||
"""Install foreground gateway signal handlers and return a restore callback."""
|
||||
@@ -298,8 +309,8 @@ def _pick_heartbeat_target_from_sessions(
|
||||
# CLI input: prompt_toolkit for editing, paste, history, and display
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_PROMPT_SESSION: PromptSession | None = None
|
||||
_SAVED_TERM_ATTRS = None # original termios settings, restored on exit
|
||||
_PROMPT_SESSION: PromptSession[str] | None = None
|
||||
_saved_term_attrs: list[Any] | None = None # original termios settings, restored on exit
|
||||
|
||||
|
||||
def _flush_pending_tty_input() -> None:
|
||||
@@ -328,12 +339,12 @@ def _flush_pending_tty_input() -> None:
|
||||
|
||||
def _restore_terminal() -> None:
|
||||
"""Restore terminal to its original state (echo, line buffering, etc.)."""
|
||||
if _SAVED_TERM_ATTRS is None:
|
||||
if _saved_term_attrs is None:
|
||||
return
|
||||
with suppress(Exception):
|
||||
import termios
|
||||
|
||||
termios.tcsetattr(sys.stdin.fileno(), termios.TCSADRAIN, _SAVED_TERM_ATTRS)
|
||||
termios.tcsetattr(sys.stdin.fileno(), termios.TCSADRAIN, _saved_term_attrs)
|
||||
|
||||
|
||||
def _build_cli_key_bindings() -> KeyBindings:
|
||||
@@ -357,20 +368,20 @@ def _build_cli_key_bindings() -> KeyBindings:
|
||||
kb = KeyBindings()
|
||||
|
||||
@kb.add("enter")
|
||||
def _(event):
|
||||
def _(event: KeyPressEvent) -> None:
|
||||
event.current_buffer.validate_and_handle()
|
||||
|
||||
@kb.add("escape", "enter") # Alt+Enter / Meta+Enter (ESC + CR, "\x1b\r")
|
||||
def _(event):
|
||||
def _(event: KeyPressEvent) -> None:
|
||||
event.current_buffer.insert_text("\n")
|
||||
|
||||
# LF-as-Enter terminals send Alt+Enter as ESC + LF rather than ESC + CR.
|
||||
@kb.add("escape", Keys.ControlJ) # Alt+Enter on LF-as-Enter terminals
|
||||
def _(event):
|
||||
def _(event: KeyPressEvent) -> None:
|
||||
event.current_buffer.insert_text("\n")
|
||||
|
||||
@kb.add(Keys.ControlF3) # Shift+Enter on CSI-u capable terminals
|
||||
def _(event):
|
||||
def _(event: KeyPressEvent) -> None:
|
||||
event.current_buffer.insert_text("\n")
|
||||
|
||||
return kb
|
||||
@@ -378,13 +389,13 @@ def _build_cli_key_bindings() -> KeyBindings:
|
||||
|
||||
def _init_prompt_session() -> None:
|
||||
"""Create the prompt_toolkit session with persistent file history."""
|
||||
global _PROMPT_SESSION, _SAVED_TERM_ATTRS
|
||||
global _PROMPT_SESSION, _saved_term_attrs
|
||||
|
||||
# Save terminal state so we can restore it on exit
|
||||
with suppress(Exception):
|
||||
import termios
|
||||
|
||||
_SAVED_TERM_ATTRS = termios.tcgetattr(sys.stdin.fileno())
|
||||
_saved_term_attrs = termios.tcgetattr(sys.stdin.fileno())
|
||||
|
||||
from nanobot.config.paths import get_cli_history_path
|
||||
|
||||
@@ -405,11 +416,14 @@ def _make_console() -> Console:
|
||||
return Console(file=sys.stdout)
|
||||
|
||||
|
||||
def _render_interactive_ansi(render_fn) -> str:
|
||||
def _render_interactive_ansi(render_fn: Callable[[Console], None]) -> str:
|
||||
"""Render Rich output to ANSI so prompt_toolkit can print it safely."""
|
||||
ansi_console = Console(
|
||||
force_terminal=sys.stdout.isatty(),
|
||||
color_system=console.color_system or "standard",
|
||||
color_system=cast(
|
||||
Literal["auto", "standard", "256", "truecolor", "windows"],
|
||||
console.color_system or "standard",
|
||||
),
|
||||
width=console.width,
|
||||
)
|
||||
with ansi_console.capture() as capture:
|
||||
@@ -420,7 +434,7 @@ def _render_interactive_ansi(render_fn) -> str:
|
||||
def _print_agent_response(
|
||||
response: str,
|
||||
render_markdown: bool,
|
||||
metadata: dict | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
show_header: bool = True,
|
||||
) -> None:
|
||||
"""Render assistant response with consistent terminal styling."""
|
||||
@@ -434,7 +448,9 @@ def _print_agent_response(
|
||||
console.print()
|
||||
|
||||
|
||||
def _response_renderable(content: str, render_markdown: bool, metadata: dict | None = None):
|
||||
def _response_renderable(
|
||||
content: str, render_markdown: bool, metadata: dict[str, Any] | None = None
|
||||
) -> Text | Markdown:
|
||||
"""Render plain-text command output without markdown collapsing newlines."""
|
||||
if not render_markdown:
|
||||
return Text(content)
|
||||
@@ -457,19 +473,19 @@ async def _print_interactive_line(text: str) -> None:
|
||||
async def _print_interactive_response(
|
||||
response: str,
|
||||
render_markdown: bool,
|
||||
metadata: dict | None = None,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""Print async interactive replies with prompt_toolkit-safe Rich styling."""
|
||||
def _write() -> None:
|
||||
content = response or ""
|
||||
ansi = _render_interactive_ansi(
|
||||
lambda c: (
|
||||
c.print(),
|
||||
c.print(f"[cyan]{__logo__} nanobot[/cyan]"),
|
||||
c.print(_response_renderable(content, render_markdown, metadata)),
|
||||
c.print(),
|
||||
)
|
||||
)
|
||||
|
||||
def _render(target: Console) -> None:
|
||||
target.print()
|
||||
target.print(f"[cyan]{__logo__} nanobot[/cyan]")
|
||||
target.print(_response_renderable(content, render_markdown, metadata))
|
||||
target.print()
|
||||
|
||||
ansi = _render_interactive_ansi(_render)
|
||||
print_formatted_text(ANSI(ansi), end="")
|
||||
|
||||
await run_in_terminal(_write)
|
||||
@@ -663,10 +679,11 @@ def onboard(
|
||||
loaded.agents.defaults.workspace = workspace
|
||||
return loaded
|
||||
|
||||
loaded_config: Config | None = None
|
||||
# Create or update config
|
||||
if config_path.exists():
|
||||
if wizard:
|
||||
config = _apply_workspace_override(load_config(config_path))
|
||||
loaded_config = _apply_workspace_override(load_config(config_path))
|
||||
else:
|
||||
should_refresh = non_interactive_refresh
|
||||
if not non_interactive_refresh:
|
||||
@@ -678,37 +695,39 @@ def onboard(
|
||||
" [bold]N[/bold] = refresh config, keeping existing values and adding new fields"
|
||||
)
|
||||
if typer.confirm("Overwrite?"):
|
||||
config = _apply_workspace_override(Config())
|
||||
save_config(config, config_path)
|
||||
loaded_config = _apply_workspace_override(Config())
|
||||
save_config(loaded_config, config_path)
|
||||
console.print(f"[green]✓[/green] Config reset to defaults at {config_path}")
|
||||
else:
|
||||
should_refresh = True
|
||||
|
||||
if should_refresh:
|
||||
config = _apply_workspace_override(load_config(config_path))
|
||||
save_config(config, config_path)
|
||||
loaded_config = _apply_workspace_override(load_config(config_path))
|
||||
save_config(loaded_config, config_path)
|
||||
console.print(
|
||||
f"[green]✓[/green] Config refreshed at {config_path} (existing values preserved)"
|
||||
)
|
||||
else:
|
||||
config = _apply_workspace_override(Config())
|
||||
loaded_config = _apply_workspace_override(Config())
|
||||
# In wizard mode, don't save yet - the wizard will handle saving if should_save=True
|
||||
if not wizard:
|
||||
save_config(config, config_path)
|
||||
save_config(loaded_config, config_path)
|
||||
console.print(f"[green]✓[/green] Created config at {config_path}")
|
||||
|
||||
assert loaded_config is not None
|
||||
|
||||
# Run interactive wizard if enabled
|
||||
if wizard:
|
||||
from nanobot.cli.onboard import run_onboard
|
||||
|
||||
try:
|
||||
result = run_onboard(initial_config=config)
|
||||
result = run_onboard(initial_config=loaded_config)
|
||||
if not result.should_save:
|
||||
console.print("[yellow]Configuration discarded. No changes were saved.[/yellow]")
|
||||
return
|
||||
|
||||
config = result.config
|
||||
save_config(config, config_path)
|
||||
loaded_config = result.config
|
||||
save_config(loaded_config, config_path)
|
||||
console.print(f"[green]✓[/green] Config saved at {config_path}")
|
||||
except Exception as e:
|
||||
console.print(f"[red]✗[/red] Error during configuration: {e}")
|
||||
@@ -717,7 +736,7 @@ def onboard(
|
||||
_onboard_plugins(config_path)
|
||||
|
||||
# Create workspace, preferring the configured workspace path.
|
||||
workspace_path = get_workspace_path(config.workspace_path)
|
||||
workspace_path = get_workspace_path(loaded_config.workspace_path)
|
||||
if not workspace_path.exists():
|
||||
workspace_path.mkdir(parents=True, exist_ok=True)
|
||||
console.print(f"[green]✓[/green] Created workspace at {workspace_path}")
|
||||
@@ -1000,7 +1019,7 @@ def _webui_config_dict(config: Config) -> dict[str, Any]:
|
||||
"""Return the current WebSocket config as a mutable alias-key dictionary."""
|
||||
from nanobot.channels.websocket.runtime import WebSocketConfig
|
||||
|
||||
current = getattr(config.channels, "websocket", None) or {}
|
||||
current: Any = getattr(config.channels, "websocket", None) or {}
|
||||
model = WebSocketConfig.model_validate(current)
|
||||
return model.model_dump(by_alias=True, exclude_none=True)
|
||||
|
||||
@@ -1008,7 +1027,7 @@ def _webui_config_dict(config: Config) -> dict[str, Any]:
|
||||
def _webui_channel_enabled(config: Config) -> bool:
|
||||
from nanobot.channels.websocket.runtime import WebSocketConfig
|
||||
|
||||
current = getattr(config.channels, "websocket", None) or {}
|
||||
current: Any = getattr(config.channels, "websocket", None) or {}
|
||||
return bool(WebSocketConfig.model_validate(current).enabled)
|
||||
|
||||
|
||||
@@ -1167,7 +1186,7 @@ def _ensure_local_webui_channel(config: Config, *, port: int | None, yes: bool)
|
||||
"""Enable the local WebUI channel with safe localhost defaults."""
|
||||
from nanobot.channels.websocket.runtime import WebSocketConfig
|
||||
|
||||
current = getattr(config.channels, "websocket", None) or {}
|
||||
current: Any = getattr(config.channels, "websocket", None) or {}
|
||||
model = WebSocketConfig.model_validate(current)
|
||||
changed = False
|
||||
generated_secret = False
|
||||
@@ -1329,7 +1348,7 @@ def _print_webui_foreground_lifecycle(*, attached: bool) -> None:
|
||||
console.print("[dim]Press Ctrl+C here to stop nanobot.[/dim]")
|
||||
|
||||
|
||||
def _attach_to_background_gateway(runtime: Any) -> None:
|
||||
def _attach_to_background_gateway(runtime: "GatewayRuntime") -> None:
|
||||
"""Keep a foreground WebUI command attached to a managed gateway."""
|
||||
_print_webui_foreground_lifecycle(attached=True)
|
||||
try:
|
||||
@@ -1512,16 +1531,19 @@ def serve(
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
async def on_startup(_app):
|
||||
async def on_startup(_app: Any) -> None:
|
||||
await agent_loop._connect_mcp()
|
||||
|
||||
async def on_cleanup(_app):
|
||||
async def on_cleanup(_app: Any) -> None:
|
||||
await agent_loop.close_mcp()
|
||||
|
||||
api_app.on_startup.append(on_startup)
|
||||
api_app.on_cleanup.append(on_cleanup)
|
||||
|
||||
web.run_app(api_app, host=host, port=port, print=lambda msg: logger.info(msg))
|
||||
def _log_aiohttp(message: object) -> None:
|
||||
logger.info("{}", message)
|
||||
|
||||
web.run_app(api_app, host=host, port=port, print=_log_aiohttp)
|
||||
|
||||
|
||||
# ============================================================================
|
||||
@@ -1778,6 +1800,7 @@ def _run_gateway(
|
||||
from nanobot.cron.session_turns import is_bound_cron_job
|
||||
from nanobot.cron.types import CronJob
|
||||
from nanobot.providers.factory import (
|
||||
ProviderSnapshot,
|
||||
build_provider_snapshot,
|
||||
build_unconfigured_provider_snapshot,
|
||||
load_provider_snapshot,
|
||||
@@ -1823,12 +1846,15 @@ def _run_gateway(
|
||||
runtime_events = RuntimeEventBus()
|
||||
fallback_model_observer = build_webui_fallback_model_observer(bus)
|
||||
|
||||
def _observe_fallback_models(snapshot):
|
||||
def _observe_fallback_models(snapshot: ProviderSnapshot) -> ProviderSnapshot:
|
||||
if isinstance(snapshot.provider, FallbackProvider):
|
||||
snapshot.provider.set_fallback_model_observer(fallback_model_observer)
|
||||
return snapshot
|
||||
|
||||
def _load_gateway_provider_snapshot(*args: Any, **kwargs: Any):
|
||||
def _load_gateway_provider_snapshot(
|
||||
*args: Any,
|
||||
**kwargs: Any,
|
||||
) -> ProviderSnapshot:
|
||||
try:
|
||||
return _observe_fallback_models(load_provider_snapshot(*args, **kwargs))
|
||||
except ValueError as exc:
|
||||
@@ -1896,10 +1922,13 @@ def _run_gateway(
|
||||
local_trigger_store=trigger_store,
|
||||
hook_factories=[create_file_edit_activity_hook],
|
||||
)
|
||||
def _schedule_webui_background(awaitable: Awaitable[None]) -> None:
|
||||
agent._schedule_background(cast(Coroutine[Any, Any, None], awaitable))
|
||||
|
||||
webui_turn_coordinator = WebuiTurnCoordinator(
|
||||
bus=bus,
|
||||
sessions=session_manager,
|
||||
schedule_background=lambda coro: agent._schedule_background(coro),
|
||||
schedule_background=_schedule_webui_background,
|
||||
)
|
||||
webui_turn_coordinator.subscribe(runtime_events)
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
@@ -1944,14 +1973,14 @@ def _run_gateway(
|
||||
session_manager.save(session)
|
||||
await bus.publish_outbound(msg)
|
||||
|
||||
message_tool = getattr(agent, "tools", {}).get("message")
|
||||
message_tool = agent.tools.get("message")
|
||||
if isinstance(message_tool, MessageTool):
|
||||
message_tool.set_send_callback(_deliver_to_channel)
|
||||
|
||||
# Set cron callback (needs agent)
|
||||
async def on_cron_job(job: CronJob) -> str | None:
|
||||
"""Execute a cron job through the agent."""
|
||||
async def _silent(*_args, **_kwargs):
|
||||
async def _silent(*_args: Any, **_kwargs: Any) -> None:
|
||||
pass
|
||||
|
||||
# Dream is an internal job — run directly, not through the agent loop.
|
||||
@@ -1972,10 +2001,7 @@ def _run_gateway(
|
||||
return None
|
||||
prompt, last_cursor = result
|
||||
key = dream_session_key()
|
||||
resolve_dream_runtime = getattr(agent, "dream_runtime", None)
|
||||
dream_runtime = (
|
||||
resolve_dream_runtime() if callable(resolve_dream_runtime) else None
|
||||
)
|
||||
dream_runtime = agent.dream_runtime()
|
||||
resp = await agent.process_direct(
|
||||
prompt,
|
||||
session_key=key,
|
||||
@@ -2111,11 +2137,7 @@ def _run_gateway(
|
||||
cron.on_job = on_cron_job
|
||||
|
||||
def _webui_runtime_model_name() -> str | None:
|
||||
model = getattr(agent, "model", None)
|
||||
if isinstance(model, str):
|
||||
stripped = model.strip()
|
||||
return stripped or None
|
||||
return None
|
||||
return agent.model.strip() or None
|
||||
|
||||
# Create channel manager (forwards SessionManager so the WebSocket channel
|
||||
# can serve the embedded webui's REST surface).
|
||||
@@ -2126,12 +2148,8 @@ def _run_gateway(
|
||||
cron_service=cron,
|
||||
local_trigger_store=trigger_store,
|
||||
webui_runtime_model_name=_webui_runtime_model_name,
|
||||
webui_cron_pending_job_ids=getattr(agent, "pending_cron_job_ids_for_session", None),
|
||||
webui_local_trigger_pending_ids=getattr(
|
||||
agent,
|
||||
"pending_local_trigger_ids_for_session",
|
||||
None,
|
||||
),
|
||||
webui_cron_pending_job_ids=agent.pending_cron_job_ids_for_session,
|
||||
webui_local_trigger_pending_ids=agent.pending_local_trigger_ids_for_session,
|
||||
webui_static_dist=webui_static_dist,
|
||||
webui_runtime_surface=webui_runtime_surface,
|
||||
webui_runtime_capabilities=webui_runtime_capabilities,
|
||||
@@ -2158,8 +2176,9 @@ def _run_gateway(
|
||||
console.print("[yellow]Warning: No channels enabled[/yellow]")
|
||||
|
||||
cron_status = cron.status()
|
||||
if cron_status["jobs"] > 0:
|
||||
console.print(f"[green]✓[/green] Cron: {cron_status['jobs']} scheduled jobs")
|
||||
cron_job_count = cast(int, cron_status["jobs"])
|
||||
if cron_job_count > 0:
|
||||
console.print(f"[green]✓[/green] Cron: {cron_job_count} scheduled jobs")
|
||||
|
||||
hb_cfg = config.gateway.heartbeat
|
||||
if hb_cfg.enabled:
|
||||
@@ -2167,13 +2186,16 @@ def _run_gateway(
|
||||
else:
|
||||
console.print("[yellow]✗[/yellow] Heartbeat: disabled")
|
||||
|
||||
async def _health_server(host: str, health_port: int):
|
||||
async def _health_server(host: str, health_port: int) -> None:
|
||||
"""Lightweight HTTP health endpoint on the gateway port."""
|
||||
import json as _json
|
||||
|
||||
connection_slots = asyncio.Semaphore(_GATEWAY_HEALTH_MAX_CONNECTIONS)
|
||||
|
||||
async def handle(reader, writer):
|
||||
async def handle(
|
||||
reader: asyncio.StreamReader,
|
||||
writer: asyncio.StreamWriter,
|
||||
) -> None:
|
||||
if connection_slots.locked():
|
||||
writer.close()
|
||||
return
|
||||
@@ -2260,7 +2282,7 @@ def _run_gateway(
|
||||
# Channels start asynchronously; a short poll lets us avoid racing the bind.
|
||||
for _ in range(40): # ~4s max
|
||||
try:
|
||||
reader, writer = await asyncio.open_connection(
|
||||
_reader, writer = await asyncio.open_connection(
|
||||
target_host,
|
||||
target_port,
|
||||
)
|
||||
@@ -2276,10 +2298,10 @@ def _run_gateway(
|
||||
except Exception as e:
|
||||
console.print(f"[yellow]Could not open browser ({e}); visit {open_browser_url}[/yellow]")
|
||||
|
||||
async def run():
|
||||
tasks: list[asyncio.Task] = []
|
||||
shutdown_task: asyncio.Task | None = None
|
||||
runtime_tasks: asyncio.Future | None = None
|
||||
async def run() -> None:
|
||||
tasks: list[asyncio.Task[Any]] = []
|
||||
shutdown_task: asyncio.Task[Any] | None = None
|
||||
runtime_tasks: asyncio.Future[list[Any]] | None = None
|
||||
runtime_tasks_drained = False
|
||||
shutdown_event = asyncio.Event()
|
||||
_ensure_interactive_tty_mode()
|
||||
@@ -2306,7 +2328,7 @@ def _run_gateway(
|
||||
asyncio.create_task(
|
||||
run_local_trigger_queue(
|
||||
store=trigger_store,
|
||||
submit_turn=getattr(agent, "submit_local_trigger_turn", None),
|
||||
submit_turn=agent.submit_local_trigger_turn,
|
||||
is_channel_enabled=lambda name: channels.get_channel(name) is not None,
|
||||
),
|
||||
name="nanobot-local-triggers",
|
||||
@@ -2334,7 +2356,7 @@ def _run_gateway(
|
||||
if runtime_tasks in done:
|
||||
runtime_tasks_drained = True
|
||||
await runtime_tasks
|
||||
elif runtime_tasks is not None:
|
||||
else:
|
||||
runtime_tasks.cancel()
|
||||
except KeyboardInterrupt:
|
||||
console.print("\nShutting down...")
|
||||
@@ -2410,33 +2432,33 @@ def agent(
|
||||
from nanobot.providers.factory import make_provider
|
||||
from nanobot.providers.image_generation import image_gen_provider_configs
|
||||
|
||||
config = _load_runtime_config(config, workspace)
|
||||
runtime_config = _load_runtime_config(config, workspace)
|
||||
try:
|
||||
provider = make_provider(config)
|
||||
provider = make_provider(runtime_config)
|
||||
except ValueError as exc:
|
||||
_print_agent_start_error(exc)
|
||||
raise typer.Exit(1) from exc
|
||||
|
||||
sync_workspace_templates(config.workspace_path)
|
||||
sync_workspace_templates(runtime_config.workspace_path)
|
||||
|
||||
bus = MessageBus()
|
||||
|
||||
# Preserve existing single-workspace installs, but keep custom workspaces clean.
|
||||
if is_default_workspace(config.workspace_path):
|
||||
_migrate_cron_store(config)
|
||||
if is_default_workspace(runtime_config.workspace_path):
|
||||
_migrate_cron_store(runtime_config)
|
||||
|
||||
# Create cron service with workspace-scoped store
|
||||
cron_store_path = config.workspace_path / "cron" / "jobs.json"
|
||||
cron_store_path = runtime_config.workspace_path / "cron" / "jobs.json"
|
||||
cron = CronService(cron_store_path)
|
||||
|
||||
_set_nanobot_logs(logs)
|
||||
|
||||
try:
|
||||
agent_loop = AgentLoop.from_config(
|
||||
config, bus,
|
||||
runtime_config, bus,
|
||||
provider=provider,
|
||||
cron_service=cron,
|
||||
image_generation_provider_configs=image_gen_provider_configs(config),
|
||||
image_generation_provider_configs=image_gen_provider_configs(runtime_config),
|
||||
hook_factories=[create_file_edit_activity_hook],
|
||||
)
|
||||
except ValueError as exc:
|
||||
@@ -2452,7 +2474,9 @@ def agent(
|
||||
# Shared reference for progress callbacks
|
||||
_thinking: ThinkingSpinner | None = None
|
||||
|
||||
def _make_progress(renderer: StreamRenderer | None = None):
|
||||
def _make_progress(
|
||||
renderer: StreamRenderer | None = None,
|
||||
) -> Callable[..., Awaitable[None]]:
|
||||
reasoning_buffer = _ReasoningBuffer()
|
||||
|
||||
async def _cli_progress(content: str, *, tool_hint: bool = False, reasoning: bool = False, **_kwargs: Any) -> None:
|
||||
@@ -2482,11 +2506,11 @@ def agent(
|
||||
|
||||
if message:
|
||||
# Single message mode — direct call, no bus needed
|
||||
async def run_once():
|
||||
async def run_once() -> None:
|
||||
renderer = StreamRenderer(
|
||||
render_markdown=markdown,
|
||||
bot_name=config.agents.defaults.bot_name,
|
||||
bot_icon=config.agents.defaults.bot_icon,
|
||||
bot_name=runtime_config.agents.defaults.bot_name,
|
||||
bot_icon=runtime_config.agents.defaults.bot_icon,
|
||||
)
|
||||
response = await agent_loop.process_direct(
|
||||
message, session_id,
|
||||
@@ -2512,8 +2536,8 @@ def agent(
|
||||
# Interactive mode — route through bus like other channels
|
||||
from nanobot.bus.events import InboundMessage
|
||||
_init_prompt_session()
|
||||
_model, _preset_tag = _model_display(config)
|
||||
_icon = config.agents.defaults.bot_icon or __logo__
|
||||
_model, _preset_tag = _model_display(runtime_config)
|
||||
_icon = runtime_config.agents.defaults.bot_icon or __logo__
|
||||
console.print(f"{_icon} Interactive mode [bold blue]({_model})[/bold blue]{_preset_tag} — type [bold]exit[/bold] or [bold]Ctrl+C[/bold] to quit\n")
|
||||
|
||||
if ":" in session_id:
|
||||
@@ -2521,7 +2545,7 @@ def agent(
|
||||
else:
|
||||
cli_channel, cli_chat_id = "cli", session_id
|
||||
|
||||
def _handle_signal(signum, frame):
|
||||
def _handle_signal(signum: int, _frame: FrameType | None) -> None:
|
||||
sig_name = signal.Signals(signum).name
|
||||
_restore_terminal()
|
||||
console.print(f"\nReceived {sig_name}, goodbye!")
|
||||
@@ -2537,7 +2561,7 @@ def agent(
|
||||
if hasattr(signal, 'SIGPIPE'):
|
||||
signal.signal(signal.SIGPIPE, signal.SIG_IGN)
|
||||
|
||||
async def run_interactive():
|
||||
async def run_interactive() -> None:
|
||||
bus_task = asyncio.create_task(agent_loop.run())
|
||||
turn_done = asyncio.Event()
|
||||
turn_done.set()
|
||||
@@ -2545,7 +2569,7 @@ def agent(
|
||||
renderer: StreamRenderer | None = None
|
||||
reasoning_buffer = _ReasoningBuffer()
|
||||
|
||||
async def _consume_outbound():
|
||||
async def _consume_outbound() -> None:
|
||||
while True:
|
||||
try:
|
||||
msg = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
||||
@@ -2578,7 +2602,7 @@ def agent(
|
||||
|
||||
if await _maybe_print_interactive_progress(
|
||||
msg,
|
||||
renderer,
|
||||
None,
|
||||
agent_loop.channels_config,
|
||||
renderer,
|
||||
reasoning_buffer,
|
||||
@@ -2625,8 +2649,8 @@ def agent(
|
||||
reasoning_buffer.clear()
|
||||
renderer = StreamRenderer(
|
||||
render_markdown=markdown,
|
||||
bot_name=config.agents.defaults.bot_name,
|
||||
bot_icon=config.agents.defaults.bot_icon,
|
||||
bot_name=runtime_config.agents.defaults.bot_name,
|
||||
bot_icon=runtime_config.agents.defaults.bot_icon,
|
||||
)
|
||||
|
||||
await bus.publish_inbound(InboundMessage(
|
||||
@@ -2701,7 +2725,7 @@ def channels_status(
|
||||
if section is None:
|
||||
enabled = False
|
||||
elif isinstance(section, dict):
|
||||
enabled = section.get("enabled", False)
|
||||
enabled = cast(dict[str, Any], section).get("enabled", False)
|
||||
else:
|
||||
enabled = getattr(section, "enabled", False)
|
||||
table.add_row(
|
||||
@@ -2719,10 +2743,11 @@ def channels_login(
|
||||
config: str | None = typer.Option(None, "--config", "-c", help="Path to config file"),
|
||||
):
|
||||
"""Authenticate with a channel via QR code or other interactive login."""
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.registry import discover_all
|
||||
|
||||
_, loaded = _load_inspection_config(config=config)
|
||||
channel_cfg = getattr(loaded.channels, channel_name, None) or {}
|
||||
channel_cfg: Any = getattr(loaded.channels, channel_name, None) or {}
|
||||
|
||||
# Validate channel exists
|
||||
all_channels = discover_all()
|
||||
@@ -2733,8 +2758,8 @@ def channels_login(
|
||||
|
||||
console.print(f"{__logo__} {all_channels[channel_name].display_name} Login\n")
|
||||
|
||||
channel_cls = all_channels[channel_name]
|
||||
channel = channel_cls(channel_cfg, bus=None)
|
||||
channel_factory = all_channels[channel_name]
|
||||
channel = channel_factory(channel_cfg, bus=MessageBus())
|
||||
|
||||
success = asyncio.run(channel.login(force=force))
|
||||
|
||||
@@ -2923,24 +2948,28 @@ _OAUTH_PROVIDER_DEFAULT_MODELS: dict[str, str] = {
|
||||
}
|
||||
|
||||
|
||||
def _register_login(name: str):
|
||||
def _register_login(
|
||||
name: str,
|
||||
) -> Callable[[Callable[[], None]], Callable[[], None]]:
|
||||
"""Register an OAuth login handler."""
|
||||
def decorator(fn):
|
||||
def decorator(fn: Callable[[], None]) -> Callable[[], None]:
|
||||
_LOGIN_HANDLERS[name] = fn
|
||||
return fn
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def _register_logout(name: str):
|
||||
def _register_logout(
|
||||
name: str,
|
||||
) -> Callable[[Callable[[], None]], Callable[[], None]]:
|
||||
"""Register an OAuth logout handler."""
|
||||
def decorator(fn):
|
||||
def decorator(fn: Callable[[], None]) -> Callable[[], None]:
|
||||
_LOGOUT_HANDLERS[name] = fn
|
||||
return fn
|
||||
return decorator
|
||||
|
||||
|
||||
def _resolve_oauth_provider(provider: str):
|
||||
def _resolve_oauth_provider(provider: str) -> "ProviderSpec":
|
||||
"""Resolve and validate an OAuth provider configuration."""
|
||||
from nanobot.providers.registry import PROVIDERS
|
||||
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Typer commands for foreground and background gateway control."""
|
||||
|
||||
# pyright: reportUnusedFunction=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import subprocess
|
||||
|
||||
+116
-69
@@ -1,19 +1,27 @@
|
||||
"""Interactive onboarding questionnaire for nanobot."""
|
||||
|
||||
# pyright: reportMissingTypeStubs=false, reportUnusedFunction=false
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import types
|
||||
from collections.abc import Callable, Iterable, Sized
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from functools import lru_cache
|
||||
from typing import Any, Literal, NamedTuple, get_args, get_origin
|
||||
from typing import Any, Literal, NamedTuple, TypeVar, cast, get_args, get_origin
|
||||
|
||||
try:
|
||||
import questionary
|
||||
except ModuleNotFoundError: # pragma: no cover - exercised in environments without wizard deps
|
||||
questionary = None
|
||||
from loguru import logger
|
||||
from prompt_toolkit.completion import CompleteEvent, Completer, Completion
|
||||
from prompt_toolkit.document import Document
|
||||
from prompt_toolkit.key_binding import KeyBindings
|
||||
from prompt_toolkit.key_binding.key_processor import KeyPressEvent
|
||||
from pydantic import BaseModel
|
||||
from pydantic.fields import FieldInfo
|
||||
from rich.console import Console
|
||||
from rich.markup import escape
|
||||
from rich.panel import Panel
|
||||
@@ -29,6 +37,8 @@ from nanobot.config.schema import Config, ModelPresetConfig
|
||||
|
||||
console = Console()
|
||||
|
||||
_ModelT = TypeVar("_ModelT", bound=BaseModel)
|
||||
|
||||
|
||||
@dataclass
|
||||
class OnboardResult:
|
||||
@@ -119,14 +129,14 @@ _CHANNEL_LOGIN_CHOICE = "Login with QR/link"
|
||||
_CHANNEL_ADVANCED_CHOICE = "Edit advanced settings"
|
||||
|
||||
|
||||
def _get_questionary():
|
||||
def _get_questionary() -> Any:
|
||||
"""Return questionary or raise a clear error when wizard deps are unavailable."""
|
||||
if questionary is None:
|
||||
raise RuntimeError(
|
||||
"Interactive onboarding requires the optional 'questionary' dependency. "
|
||||
"Install project dependencies and rerun with --wizard."
|
||||
)
|
||||
return questionary
|
||||
return cast(Any, questionary)
|
||||
|
||||
|
||||
def _select_with_back(
|
||||
@@ -147,7 +157,6 @@ def _select_with_back(
|
||||
import shutil
|
||||
|
||||
from prompt_toolkit.application import Application
|
||||
from prompt_toolkit.key_binding import KeyBindings
|
||||
from prompt_toolkit.keys import Keys
|
||||
from prompt_toolkit.layout import Layout
|
||||
from prompt_toolkit.layout.containers import HSplit, Window
|
||||
@@ -170,8 +179,8 @@ def _select_with_back(
|
||||
visible_count = min(len(choices), max(1, terminal_lines - 3))
|
||||
|
||||
# Build menu items (uses closure over selected_index)
|
||||
def get_menu_text():
|
||||
items = []
|
||||
def get_menu_text() -> list[tuple[str, str]]:
|
||||
items: list[tuple[str, str]] = []
|
||||
start, end = _choice_viewport(selected_index, len(choices), visible_count)
|
||||
for i in range(start, end):
|
||||
choice = choices[i]
|
||||
@@ -182,14 +191,14 @@ def _select_with_back(
|
||||
return items
|
||||
|
||||
# Create layout
|
||||
menu_control = FormattedTextControl(get_menu_text, show_cursor=False)
|
||||
menu_control = FormattedTextControl(cast(Any, get_menu_text), show_cursor=False)
|
||||
menu_window = Window(content=menu_control, height=visible_count, always_hide_cursor=True)
|
||||
|
||||
def get_prompt_text():
|
||||
def get_prompt_text() -> list[tuple[str, str]]:
|
||||
suffix = f" ({selected_index + 1}/{len(choices)})" if len(choices) > visible_count else ""
|
||||
return [("class:question", f"{prompt}{suffix}")]
|
||||
|
||||
prompt_control = FormattedTextControl(get_prompt_text, show_cursor=False)
|
||||
prompt_control = FormattedTextControl(cast(Any, get_prompt_text), show_cursor=False)
|
||||
prompt_window = Window(content=prompt_control, height=1, always_hide_cursor=True)
|
||||
|
||||
layout = Layout(HSplit([prompt_window, menu_window]))
|
||||
@@ -198,34 +207,34 @@ def _select_with_back(
|
||||
bindings = KeyBindings()
|
||||
|
||||
@bindings.add(Keys.Up)
|
||||
def _up(event):
|
||||
def _up(event: KeyPressEvent) -> None:
|
||||
nonlocal selected_index
|
||||
selected_index = (selected_index - 1) % len(choices)
|
||||
event.app.invalidate()
|
||||
|
||||
@bindings.add(Keys.Down)
|
||||
def _down(event):
|
||||
def _down(event: KeyPressEvent) -> None:
|
||||
nonlocal selected_index
|
||||
selected_index = (selected_index + 1) % len(choices)
|
||||
event.app.invalidate()
|
||||
|
||||
@bindings.add(Keys.Enter)
|
||||
def _enter(event):
|
||||
def _enter(event: KeyPressEvent) -> None:
|
||||
state["result"] = choices[selected_index]
|
||||
event.app.exit()
|
||||
|
||||
@bindings.add("escape")
|
||||
def _escape(event):
|
||||
def _escape(event: KeyPressEvent) -> None:
|
||||
state["result"] = _BACK_PRESSED
|
||||
event.app.exit()
|
||||
|
||||
@bindings.add(Keys.Left)
|
||||
def _left(event):
|
||||
def _left(event: KeyPressEvent) -> None:
|
||||
state["result"] = _BACK_PRESSED
|
||||
event.app.exit()
|
||||
|
||||
@bindings.add(Keys.ControlC)
|
||||
def _ctrl_c(event):
|
||||
def _ctrl_c(event: KeyPressEvent) -> None:
|
||||
state["result"] = None
|
||||
event.app.exit()
|
||||
|
||||
@@ -235,7 +244,7 @@ def _select_with_back(
|
||||
"question": f"fg:{_UI_TEXT}",
|
||||
})
|
||||
|
||||
app = Application(layout=layout, key_bindings=bindings, style=style)
|
||||
app = Application[object](layout=layout, key_bindings=bindings, style=style)
|
||||
app.ttimeoutlen = 0.05
|
||||
app.timeoutlen = 0.05
|
||||
try:
|
||||
@@ -268,7 +277,7 @@ class FieldTypeInfo(NamedTuple):
|
||||
inner_type: Any
|
||||
|
||||
|
||||
def _get_field_type_info(field_info) -> FieldTypeInfo:
|
||||
def _get_field_type_info(field_info: FieldInfo) -> FieldTypeInfo:
|
||||
"""Extract field type info from Pydantic field."""
|
||||
annotation = field_info.annotation
|
||||
if annotation is None:
|
||||
@@ -285,10 +294,11 @@ def _get_field_type_info(field_info) -> FieldTypeInfo:
|
||||
args = get_args(annotation)
|
||||
|
||||
_simple_types: dict[type, str] = {bool: "bool", int: "int", float: "float"}
|
||||
origin_name = getattr(origin, "__name__", None)
|
||||
|
||||
if origin is list or (hasattr(origin, "__name__") and origin.__name__ == "List"):
|
||||
if origin is list or origin_name == "List":
|
||||
return FieldTypeInfo("list", args[0] if args else str)
|
||||
if origin is dict or (hasattr(origin, "__name__") and origin.__name__ == "Dict"):
|
||||
if origin is dict or origin_name == "Dict":
|
||||
return FieldTypeInfo("dict", None)
|
||||
for py_type, name in _simple_types.items():
|
||||
if annotation is py_type:
|
||||
@@ -300,7 +310,7 @@ def _get_field_type_info(field_info) -> FieldTypeInfo:
|
||||
return FieldTypeInfo("str", None)
|
||||
|
||||
|
||||
def _get_field_display_name(field_key: str, field_info) -> str:
|
||||
def _get_field_display_name(field_key: str, field_info: FieldInfo | None) -> str:
|
||||
"""Get display name for a field."""
|
||||
if field_info and field_info.description:
|
||||
return field_info.description
|
||||
@@ -349,22 +359,30 @@ def _format_value(value: Any, rich: bool = True, field_name: str = "") -> str:
|
||||
masked = _mask_value(value)
|
||||
return f"[dim]{masked}[/dim]" if rich else masked
|
||||
if isinstance(value, BaseModel):
|
||||
parts = []
|
||||
model_parts: list[str] = []
|
||||
for fname, _finfo in type(value).model_fields.items():
|
||||
fval = getattr(value, fname, None)
|
||||
formatted = _format_value(fval, rich=False, field_name=fname)
|
||||
if formatted != "[not set]":
|
||||
parts.append(f"{fname}={formatted}")
|
||||
return ", ".join(parts) if parts else ("[dim]not set[/dim]" if rich else "[not set]")
|
||||
model_parts.append(f"{fname}={formatted}")
|
||||
return (
|
||||
", ".join(model_parts)
|
||||
if model_parts
|
||||
else ("[dim]not set[/dim]" if rich else "[not set]")
|
||||
)
|
||||
if isinstance(value, list):
|
||||
return ", ".join(str(v) for v in value)
|
||||
return ", ".join(str(v) for v in cast(list[Any], value))
|
||||
if isinstance(value, dict):
|
||||
# Handle dicts containing BaseModel instances
|
||||
parts = []
|
||||
for k, v in value.items():
|
||||
mapping_parts: list[str] = []
|
||||
for k, v in cast(dict[Any, Any], value).items():
|
||||
formatted = _format_value(v, rich=False, field_name=str(k))
|
||||
parts.append(f"{k}: {formatted}")
|
||||
return ", ".join(parts) if parts else ("[dim]not set[/dim]" if rich else "[not set]")
|
||||
mapping_parts.append(f"{k}: {formatted}")
|
||||
return (
|
||||
", ".join(mapping_parts)
|
||||
if mapping_parts
|
||||
else ("[dim]not set[/dim]" if rich else "[not set]")
|
||||
)
|
||||
return str(value)
|
||||
|
||||
|
||||
@@ -373,13 +391,13 @@ def _format_value_for_input(value: Any, field_type: str) -> str:
|
||||
if value is None or value == "":
|
||||
return ""
|
||||
if field_type == "list" and isinstance(value, list):
|
||||
return ",".join(str(v) for v in value)
|
||||
return ",".join(str(v) for v in cast(list[Any], value))
|
||||
if field_type == "dict" and isinstance(value, dict):
|
||||
return json.dumps(value)
|
||||
return str(value)
|
||||
|
||||
|
||||
def _validate_field_constraint(value: Any, field_info) -> str | None:
|
||||
def _validate_field_constraint(value: Any, field_info: FieldInfo | None) -> str | None:
|
||||
"""Validate a value against Pydantic Field constraints.
|
||||
|
||||
Returns an error message string if validation fails, None if valid.
|
||||
@@ -388,7 +406,8 @@ def _validate_field_constraint(value: Any, field_info) -> str | None:
|
||||
if field_info is None or not hasattr(field_info, "metadata"):
|
||||
return None
|
||||
|
||||
for m in field_info.metadata:
|
||||
for metadata in field_info.metadata:
|
||||
m = metadata
|
||||
if hasattr(m, "ge") and isinstance(value, (int, float)):
|
||||
if value < m.ge:
|
||||
return f"Value must be >= {m.ge}"
|
||||
@@ -402,16 +421,16 @@ def _validate_field_constraint(value: Any, field_info) -> str | None:
|
||||
if value >= m.lt:
|
||||
return f"Value must be < {m.lt}"
|
||||
if hasattr(m, "min_length") and hasattr(value, "__len__"):
|
||||
if len(value) < m.min_length:
|
||||
if len(cast(Sized, value)) < m.min_length:
|
||||
return f"Length must be >= {m.min_length}"
|
||||
if hasattr(m, "max_length") and hasattr(value, "__len__"):
|
||||
if len(value) > m.max_length:
|
||||
if len(cast(Sized, value)) > m.max_length:
|
||||
return f"Length must be <= {m.max_length}"
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _get_constraint_hint(field_info) -> str:
|
||||
def _get_constraint_hint(field_info: FieldInfo | None) -> str:
|
||||
"""Derive a human-readable constraint hint from field metadata.
|
||||
|
||||
Returns a string like " - 0-10" or " - >= 0" to append to field display names.
|
||||
@@ -421,7 +440,8 @@ def _get_constraint_hint(field_info) -> str:
|
||||
|
||||
ge_val = None
|
||||
le_val = None
|
||||
for m in field_info.metadata:
|
||||
for metadata in field_info.metadata:
|
||||
m = metadata
|
||||
if hasattr(m, "ge"):
|
||||
ge_val = m.ge
|
||||
if hasattr(m, "le"):
|
||||
@@ -439,7 +459,11 @@ def _get_constraint_hint(field_info) -> str:
|
||||
# --- Rich UI Components ---
|
||||
|
||||
|
||||
def _show_config_panel(display_name: str, model: BaseModel, fields: list) -> None:
|
||||
def _show_config_panel(
|
||||
display_name: str,
|
||||
model: BaseModel,
|
||||
fields: list[tuple[str, FieldInfo]],
|
||||
) -> None:
|
||||
"""Display current configuration as a rich table."""
|
||||
table = Table(show_header=False, box=None, padding=(0, 2))
|
||||
table.add_column("Field", style=_UI_ACCENT)
|
||||
@@ -504,20 +528,18 @@ def _input_bool(display_name: str, current: bool | None) -> bool | None:
|
||||
).ask()
|
||||
|
||||
|
||||
def _input_back_key_bindings():
|
||||
def _input_back_key_bindings() -> KeyBindings:
|
||||
"""Return key bindings that make Escape behave like a local back action."""
|
||||
from prompt_toolkit.key_binding import KeyBindings
|
||||
|
||||
bindings = KeyBindings()
|
||||
|
||||
@bindings.add("escape")
|
||||
def _escape(event):
|
||||
def _escape(event: KeyPressEvent) -> None:
|
||||
event.app.exit(result=_BACK_PRESSED)
|
||||
|
||||
return bindings
|
||||
|
||||
|
||||
def _ask_prompt(prompt):
|
||||
def _ask_prompt(prompt: Any) -> Any:
|
||||
"""Ask a questionary prompt with responsive Escape handling."""
|
||||
app = getattr(prompt, "application", None)
|
||||
if app is not None:
|
||||
@@ -528,7 +550,12 @@ def _ask_prompt(prompt):
|
||||
return prompt.ask()
|
||||
|
||||
|
||||
def _input_text(display_name: str, current: Any, field_type: str, field_info=None) -> Any:
|
||||
def _input_text(
|
||||
display_name: str,
|
||||
current: Any,
|
||||
field_type: str,
|
||||
field_info: FieldInfo | None = None,
|
||||
) -> Any:
|
||||
"""Get text input and parse based on field type."""
|
||||
default = _format_value_for_input(current, field_type)
|
||||
|
||||
@@ -591,7 +618,10 @@ def _input_secret(display_name: str) -> str | None | object:
|
||||
|
||||
|
||||
def _input_with_existing(
|
||||
display_name: str, current: Any, field_type: str, field_info=None
|
||||
display_name: str,
|
||||
current: Any,
|
||||
field_type: str,
|
||||
field_info: FieldInfo | None = None,
|
||||
) -> Any:
|
||||
"""Handle input with 'keep existing' option for non-empty values."""
|
||||
has_existing = current is not None and current != "" and current != {} and current != []
|
||||
@@ -624,8 +654,6 @@ def _input_model_with_autocomplete(
|
||||
"""Get model input with autocomplete suggestions.
|
||||
|
||||
"""
|
||||
from prompt_toolkit.completion import Completer, Completion
|
||||
|
||||
default = str(current) if current else ""
|
||||
|
||||
class DynamicModelCompleter(Completer):
|
||||
@@ -634,7 +662,12 @@ def _input_model_with_autocomplete(
|
||||
def __init__(self, provider_name: str):
|
||||
self.provider = provider_name
|
||||
|
||||
def get_completions(self, document, _complete_event):
|
||||
def get_completions(
|
||||
self,
|
||||
document: Document,
|
||||
complete_event: CompleteEvent,
|
||||
) -> Iterable[Completion]:
|
||||
_ = complete_event
|
||||
text = document.text_before_cursor
|
||||
suggestions = get_model_suggestions(text, provider=self.provider, limit=50)
|
||||
for model in suggestions:
|
||||
@@ -735,7 +768,7 @@ def _handle_model_field(
|
||||
return
|
||||
if new_value is not None and new_value != current_value:
|
||||
setattr(working_model, field_name, new_value)
|
||||
_try_auto_fill_context_window(working_model, new_value)
|
||||
_try_auto_fill_context_window(working_model, cast(str, new_value))
|
||||
|
||||
|
||||
def _handle_context_window_field(
|
||||
@@ -794,7 +827,11 @@ def _handle_fallback_models_field(
|
||||
"""Handle the 'fallback_models' field with preset-aware list management."""
|
||||
from nanobot.config.schema import InlineFallbackConfig
|
||||
|
||||
items: list[Any] = list(current_value) if isinstance(current_value, list) else []
|
||||
items: list[Any] = (
|
||||
list(cast(list[Any], current_value))
|
||||
if isinstance(current_value, list)
|
||||
else []
|
||||
)
|
||||
preset_names = sorted(_MODEL_PRESET_CACHE)
|
||||
|
||||
while True:
|
||||
@@ -888,11 +925,11 @@ def _is_str_or_none(annotation: Any) -> bool:
|
||||
|
||||
|
||||
def _configure_pydantic_model(
|
||||
model: BaseModel,
|
||||
model: _ModelT,
|
||||
display_name: str,
|
||||
*,
|
||||
skip_fields: set[str] | None = None,
|
||||
) -> BaseModel | None:
|
||||
) -> _ModelT | None:
|
||||
"""Configure a Pydantic model interactively.
|
||||
|
||||
Returns the updated model when the user selects "Done" or navigates back.
|
||||
@@ -901,7 +938,7 @@ def _configure_pydantic_model(
|
||||
skip_fields = skip_fields or set()
|
||||
working_model = model.model_copy(deep=True)
|
||||
|
||||
fields = [
|
||||
fields: list[tuple[str, FieldInfo]] = [
|
||||
(name, info)
|
||||
for name, info in type(working_model).model_fields.items()
|
||||
if name not in skip_fields
|
||||
@@ -911,7 +948,7 @@ def _configure_pydantic_model(
|
||||
return working_model
|
||||
|
||||
def get_choices() -> list[str]:
|
||||
items = []
|
||||
items: list[str] = []
|
||||
for fname, finfo in fields:
|
||||
value = getattr(working_model, fname, None)
|
||||
display = _get_field_display_name(fname, finfo)
|
||||
@@ -1057,6 +1094,10 @@ def _sync_preset_cache(config: Config) -> None:
|
||||
_MODEL_PRESET_CACHE.update(config.model_presets.keys())
|
||||
|
||||
|
||||
def _validate_nonempty_name(text: str) -> bool | str:
|
||||
return True if text and text.strip() else "Name cannot be empty"
|
||||
|
||||
|
||||
def _configure_model_presets(config: Config) -> None:
|
||||
"""Configure model presets (CRUD)."""
|
||||
_sync_preset_cache(config)
|
||||
@@ -1099,7 +1140,7 @@ def _configure_model_presets(config: Config) -> None:
|
||||
if answer == "[+] Add new preset":
|
||||
name_input = _get_questionary().text(
|
||||
"Preset name:",
|
||||
validate=lambda t: True if t and t.strip() else "Name cannot be empty",
|
||||
validate=_validate_nonempty_name,
|
||||
).ask()
|
||||
if not name_input:
|
||||
continue
|
||||
@@ -1218,7 +1259,7 @@ def _configure_providers(config: Config) -> None:
|
||||
|
||||
def get_provider_choices() -> list[str]:
|
||||
"""Build provider choices with config status indicators."""
|
||||
choices = []
|
||||
choices: list[str] = []
|
||||
for name, display in _get_provider_names().items():
|
||||
provider = getattr(config.providers, name, None)
|
||||
if provider and provider.api_key:
|
||||
@@ -1427,7 +1468,7 @@ _SETTINGS_SECTIONS: dict[str, tuple[str, str, set[str] | None]] = {
|
||||
"Tools": ("Tools Settings", "Configure web search, shell exec, and other tools", {"mcp_servers"}),
|
||||
}
|
||||
|
||||
_SETTINGS_GETTER = {
|
||||
_SETTINGS_GETTER: dict[str, Callable[[Config], BaseModel]] = {
|
||||
"Agent Settings": lambda c: c.agents.defaults,
|
||||
"Channel Common": lambda c: c.channels,
|
||||
"API Server": lambda c: c.api,
|
||||
@@ -1435,7 +1476,7 @@ _SETTINGS_GETTER = {
|
||||
"Tools": lambda c: c.tools,
|
||||
}
|
||||
|
||||
_SETTINGS_SETTER = {
|
||||
_SETTINGS_SETTER: dict[str, Callable[[Config, BaseModel], None]] = {
|
||||
"Agent Settings": lambda c, v: setattr(c.agents, "defaults", v),
|
||||
"Channel Common": lambda c, v: setattr(c, "channels", v),
|
||||
"API Server": lambda c, v: setattr(c, "api", v),
|
||||
@@ -1449,7 +1490,7 @@ def _configure_general_settings(config: Config, section: str) -> None:
|
||||
meta = _SETTINGS_SECTIONS.get(section)
|
||||
if not meta:
|
||||
return
|
||||
display_name, subtitle, skip = meta
|
||||
display_name, _subtitle, skip = meta
|
||||
model = _SETTINGS_GETTER[section](config)
|
||||
updated = _configure_pydantic_model(model, display_name, skip_fields=skip)
|
||||
if updated is not None:
|
||||
@@ -1495,7 +1536,7 @@ def _show_summary(config: Config) -> None:
|
||||
console.print()
|
||||
|
||||
# Providers
|
||||
provider_rows = []
|
||||
provider_rows: list[tuple[str, str]] = []
|
||||
for name, display in _get_provider_names().items():
|
||||
provider = getattr(config.providers, name, None)
|
||||
status = (
|
||||
@@ -1507,12 +1548,12 @@ def _show_summary(config: Config) -> None:
|
||||
_print_summary_panel(provider_rows, "LLM Providers")
|
||||
|
||||
# Channels
|
||||
channel_rows = []
|
||||
channel_rows: list[tuple[str, str]] = []
|
||||
for name, display in _get_channel_names().items():
|
||||
channel = getattr(config.channels, name, None)
|
||||
if channel:
|
||||
enabled = (
|
||||
channel.get("enabled", False)
|
||||
cast(dict[str, Any], channel).get("enabled", False)
|
||||
if isinstance(channel, dict)
|
||||
else getattr(channel, "enabled", False)
|
||||
)
|
||||
@@ -1523,7 +1564,7 @@ def _show_summary(config: Config) -> None:
|
||||
_print_summary_panel(channel_rows, "Chat Channels")
|
||||
|
||||
# Model Presets
|
||||
preset_rows = []
|
||||
preset_rows: list[tuple[str, str]] = []
|
||||
for name, preset in config.model_presets.items():
|
||||
preset_rows.append((name, f"{preset.model} - ctx {preset.context_window_tokens}"))
|
||||
_print_summary_panel(preset_rows, "Model Presets")
|
||||
@@ -1562,7 +1603,7 @@ def _set_primary_quick_start_preset(config: Config, provider_name: str, model: s
|
||||
|
||||
def _show_quick_start_progress(active_step: int) -> None:
|
||||
"""Render a compact step tracker for Quick Start."""
|
||||
parts = []
|
||||
parts: list[str] = []
|
||||
for idx, label in enumerate(_QUICK_START_STEPS, 1):
|
||||
if idx < active_step:
|
||||
parts.append(f"[{_UI_SUCCESS}]{idx}. {label}[/]")
|
||||
@@ -1755,7 +1796,10 @@ def _configure_quick_start_provider(config: Config) -> bool | object:
|
||||
continue
|
||||
if api_base_result is None:
|
||||
return False
|
||||
api_base, base_was_prompted = api_base_result
|
||||
api_base, base_was_prompted = cast(
|
||||
tuple[str, bool],
|
||||
api_base_result,
|
||||
)
|
||||
|
||||
api_key: str | None = None
|
||||
if _quick_start_requires_api_key(provider_name, provider_info):
|
||||
@@ -1778,7 +1822,10 @@ def _configure_quick_start_provider(config: Config) -> bool | object:
|
||||
continue
|
||||
if api_base_result is None:
|
||||
return False
|
||||
api_base, base_was_prompted = api_base_result
|
||||
api_base, base_was_prompted = cast(
|
||||
tuple[str, bool],
|
||||
api_base_result,
|
||||
)
|
||||
|
||||
provider_config = getattr(config.providers, provider_name, None)
|
||||
if provider_config is None:
|
||||
@@ -1792,7 +1839,7 @@ def _configure_quick_start_provider(config: Config) -> bool | object:
|
||||
)
|
||||
if model is _BACK_PRESSED:
|
||||
continue
|
||||
model = (model or "").strip()
|
||||
model = cast(str, model or "").strip()
|
||||
if not model:
|
||||
console.print("[yellow]! Model ID is required for Quick Start[/yellow]")
|
||||
return False
|
||||
@@ -1850,7 +1897,7 @@ def _enable_quick_start_websocket_defaults(config: Config) -> bool:
|
||||
console.print("[red]No configuration class found for websocket[/red]")
|
||||
return False
|
||||
|
||||
current = getattr(config.channels, "websocket", None) or {}
|
||||
current: Any = getattr(config.channels, "websocket", None) or {}
|
||||
model = config_cls.model_validate(current)
|
||||
if hasattr(model, "enabled"):
|
||||
setattr(model, "enabled", True)
|
||||
@@ -1997,7 +2044,7 @@ def _configure_advanced_settings(config: Config) -> None:
|
||||
if answer is _BACK_PRESSED or answer is None or answer == "<- Back":
|
||||
break
|
||||
|
||||
_advanced_dispatch = {
|
||||
_advanced_dispatch: dict[str, Callable[[], None]] = {
|
||||
"[P] LLM Provider": lambda: _configure_providers(config),
|
||||
"[M] Model Presets": lambda: _configure_model_presets(config),
|
||||
"[C] Chat Channel": lambda: _configure_channels(config),
|
||||
@@ -2008,9 +2055,9 @@ def _configure_advanced_settings(config: Config) -> None:
|
||||
"[T] Tools": lambda: _configure_general_settings(config, "Tools"),
|
||||
"[V] View Configuration Summary": lambda: _show_summary(config),
|
||||
}
|
||||
action_fn = _advanced_dispatch.get(answer)
|
||||
action_fn = _advanced_dispatch.get(cast(str, answer))
|
||||
if action_fn:
|
||||
last_choice = answer
|
||||
last_choice = cast(str, answer)
|
||||
action_fn()
|
||||
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from contextlib import contextmanager, nullcontext
|
||||
from typing import Literal
|
||||
|
||||
from rich.console import Console
|
||||
from rich.live import Live
|
||||
@@ -51,12 +52,12 @@ class ThinkingSpinner:
|
||||
self._spinner = c.status(f"[dim]{bot_name} is thinking...[/dim]", spinner="dots")
|
||||
self._active = False
|
||||
|
||||
def __enter__(self):
|
||||
def __enter__(self) -> ThinkingSpinner:
|
||||
self._spinner.start()
|
||||
self._active = True
|
||||
return self
|
||||
|
||||
def __exit__(self, *exc):
|
||||
def __exit__(self, *exc: object) -> Literal[False]:
|
||||
self._active = False
|
||||
self._spinner.stop()
|
||||
_clear_current_line(self._console)
|
||||
@@ -110,7 +111,7 @@ class StreamRenderer:
|
||||
self._header_printed = False
|
||||
self._start_spinner()
|
||||
|
||||
def _renderable(self):
|
||||
def _renderable(self) -> Markdown | Text:
|
||||
"""Create a renderable from the current buffer."""
|
||||
if self._md and self._buf:
|
||||
return Markdown(self._buf)
|
||||
|
||||
+48
-39
@@ -9,16 +9,20 @@ import sys
|
||||
import time
|
||||
from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from typing import Literal
|
||||
from typing import TYPE_CHECKING, Any, Literal, cast
|
||||
|
||||
from nanobot import __version__
|
||||
from nanobot.agent.goal_permission import goal_mutation_permission
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.command.router import CommandContext, CommandRouter, normalize_command_text
|
||||
from nanobot.utils.helpers import build_status_content
|
||||
from nanobot.utils.restart import set_restart_notice_to_env
|
||||
from nanobot.utils.workspace_prompts import initialize_workspace_prompt
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.session.manager import Session
|
||||
from nanobot.utils.gitstore import CommitInfo
|
||||
|
||||
# WebUI protocol contract for how a slash command participates in turn state:
|
||||
# - side_channel: returns control text without starting or ending an agent turn.
|
||||
# - finalize_active_turn: side-channel command that also closes the active UI turn.
|
||||
@@ -199,9 +203,9 @@ async def cmd_stop(ctx: CommandContext) -> OutboundMessage:
|
||||
"""Cancel all active tasks and subagents for the session."""
|
||||
loop = ctx.loop
|
||||
msg = ctx.msg
|
||||
total = await loop._cancel_active_tasks(ctx.key)
|
||||
total = await loop._cancel_active_tasks(ctx.key) # pyright: ignore[reportPrivateUsage]
|
||||
# Also drain pending queue to prevent mid-turn injection deadlock
|
||||
pending = loop._pending_queues.pop(ctx.key, None)
|
||||
pending = loop._pending_queues.pop(ctx.key, None) # pyright: ignore[reportPrivateUsage]
|
||||
if pending is not None:
|
||||
while not pending.empty():
|
||||
try:
|
||||
@@ -228,14 +232,14 @@ async def cmd_restart(ctx: CommandContext) -> OutboundMessage:
|
||||
async def _do_restart():
|
||||
await asyncio.sleep(1)
|
||||
argv = [sys.executable, "-m", "nanobot"] + sys.argv[1:]
|
||||
mode = getattr(ctx.loop, "restart_mode", "auto") or "auto"
|
||||
mode = ctx.loop.restart_mode or "auto"
|
||||
if mode == "auto":
|
||||
mode = "spawn" if sys.platform == "win32" else "exec"
|
||||
if mode == "exec":
|
||||
os.execv(sys.executable, argv)
|
||||
return
|
||||
if mode == "spawn":
|
||||
kwargs = {}
|
||||
kwargs: dict[str, Any] = {}
|
||||
if sys.platform == "win32":
|
||||
kwargs["creationflags"] = subprocess.CREATE_NEW_PROCESS_GROUP
|
||||
subprocess.Popen(argv, **kwargs)
|
||||
@@ -260,21 +264,20 @@ async def cmd_status(ctx: CommandContext) -> OutboundMessage:
|
||||
runtime=runtime,
|
||||
)
|
||||
if ctx_est <= 0:
|
||||
ctx_est = loop._last_usage.get("prompt_tokens", 0)
|
||||
ctx_est = loop._last_usage.get("prompt_tokens", 0) # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
# Fetch web search provider usage (best-effort, never blocks the response)
|
||||
search_usage_text: str | None = None
|
||||
# Never let usage fetch break /status
|
||||
with suppress(Exception):
|
||||
from nanobot.utils.searchusage import fetch_search_usage
|
||||
web_cfg = getattr(loop, "web_config", None)
|
||||
search_cfg = getattr(web_cfg, "search", None) if web_cfg else None
|
||||
if search_cfg is not None:
|
||||
provider = getattr(search_cfg, "provider", "duckduckgo")
|
||||
api_key = getattr(search_cfg, "api_key", "") or None
|
||||
usage = await fetch_search_usage(provider=provider, api_key=api_key)
|
||||
search_usage_text = usage.format()
|
||||
active_tasks = loop._active_tasks.get(ctx.key, [])
|
||||
search_cfg = loop.web_config.search
|
||||
usage = await fetch_search_usage(
|
||||
provider=search_cfg.provider,
|
||||
api_key=search_cfg.api_key or None,
|
||||
)
|
||||
search_usage_text = usage.format()
|
||||
active_tasks = loop._active_tasks.get(ctx.key, []) # pyright: ignore[reportPrivateUsage]
|
||||
task_count = sum(1 for t in active_tasks if not t.done())
|
||||
with suppress(Exception):
|
||||
task_count += loop.subagents.get_running_count_by_session(ctx.key)
|
||||
@@ -283,7 +286,7 @@ async def cmd_status(ctx: CommandContext) -> OutboundMessage:
|
||||
chat_id=ctx.msg.chat_id,
|
||||
content=build_status_content(
|
||||
version=__version__, model=runtime.model,
|
||||
start_time=loop._start_time, last_usage=loop._last_usage,
|
||||
start_time=loop._start_time, last_usage=loop._last_usage, # pyright: ignore[reportPrivateUsage]
|
||||
context_window_tokens=runtime.context_window_tokens,
|
||||
session_msg_count=len(session.get_history(max_messages=0)),
|
||||
context_tokens_estimate=ctx_est,
|
||||
@@ -298,17 +301,18 @@ async def cmd_status(ctx: CommandContext) -> OutboundMessage:
|
||||
async def cmd_new(ctx: CommandContext) -> OutboundMessage:
|
||||
"""Stop active task and start a fresh session."""
|
||||
loop = ctx.loop
|
||||
await loop._cancel_active_tasks(ctx.key)
|
||||
await loop._cancel_active_tasks(ctx.key) # pyright: ignore[reportPrivateUsage]
|
||||
session = ctx.session or loop.sessions.get_or_create(ctx.key)
|
||||
snapshot = session.messages[session.last_consolidated:]
|
||||
runtime = None
|
||||
if snapshot:
|
||||
runtime = ctx.runtime or loop.runtime_for_session(session)
|
||||
session.clear()
|
||||
loop.sessions.save(session)
|
||||
loop.sessions.invalidate(session.key)
|
||||
if snapshot:
|
||||
loop._schedule_background(
|
||||
loop.consolidator.archive(
|
||||
if snapshot and runtime is not None:
|
||||
loop._schedule_background( # pyright: ignore[reportPrivateUsage]
|
||||
loop.consolidator.archive( # pyright: ignore[reportUnknownMemberType]
|
||||
snapshot,
|
||||
runtime=runtime,
|
||||
session_key=ctx.key,
|
||||
@@ -325,7 +329,7 @@ def _format_preset_names(names: list[str]) -> str:
|
||||
return ", ".join(f"`{name}`" for name in names) if names else "(none configured)"
|
||||
|
||||
|
||||
def _model_preset_names(loop) -> list[str]:
|
||||
def _model_preset_names(loop: AgentLoop) -> list[str]:
|
||||
names = set(loop.model_presets)
|
||||
names.add("default")
|
||||
return ["default", *sorted(name for name in names if name != "default")]
|
||||
@@ -335,7 +339,7 @@ def _command_error_message(exc: Exception) -> str:
|
||||
return str(exc.args[0]) if isinstance(exc, KeyError) and exc.args else str(exc)
|
||||
|
||||
|
||||
def _model_command_status(loop, session) -> str:
|
||||
def _model_command_status(loop: AgentLoop, session: Session) -> str:
|
||||
names = _model_preset_names(loop)
|
||||
try:
|
||||
runtime = loop.runtime_for_session(session, recover_removed=False)
|
||||
@@ -401,8 +405,7 @@ async def cmd_model(ctx: CommandContext) -> OutboundMessage:
|
||||
f"- Model: `{runtime.model}`",
|
||||
f"- Context window: {runtime.context_window_tokens}",
|
||||
]
|
||||
if max_tokens is not None:
|
||||
lines.append(f"- Max output tokens: {max_tokens}")
|
||||
lines.append(f"- Max output tokens: {max_tokens}")
|
||||
return OutboundMessage(
|
||||
channel=ctx.msg.channel,
|
||||
chat_id=ctx.msg.chat_id,
|
||||
@@ -442,8 +445,7 @@ async def cmd_dream(ctx: CommandContext) -> OutboundMessage:
|
||||
return
|
||||
prompt, last_cursor = result
|
||||
key = dream_session_key()
|
||||
resolve_dream_runtime = getattr(loop, "dream_runtime", None)
|
||||
dream_runtime = resolve_dream_runtime() if callable(resolve_dream_runtime) else None
|
||||
dream_runtime = loop.dream_runtime()
|
||||
resp = await loop.process_direct(
|
||||
prompt,
|
||||
session_key=key,
|
||||
@@ -640,7 +642,12 @@ def _format_changed_files(diff: str) -> str:
|
||||
_DREAM_COMMIT_PREFIX = "dream:"
|
||||
|
||||
|
||||
def _format_dream_log_content(commit, diff: str, *, requested_sha: str | None = None) -> str:
|
||||
def _format_dream_log_content(
|
||||
commit: CommitInfo,
|
||||
diff: str,
|
||||
*,
|
||||
requested_sha: str | None = None,
|
||||
) -> str:
|
||||
files_line = _format_changed_files(diff)
|
||||
lines = [
|
||||
"## Dream Update",
|
||||
@@ -668,7 +675,7 @@ def _format_dream_log_content(commit, diff: str, *, requested_sha: str | None =
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _format_dream_restore_list(commits: list) -> str:
|
||||
def _format_dream_restore_list(commits: list[CommitInfo]) -> str:
|
||||
lines = [
|
||||
"## Dream Restore",
|
||||
"",
|
||||
@@ -806,14 +813,20 @@ _HISTORY_MAX_COUNT = 50
|
||||
_HISTORY_MAX_CONTENT_CHARS = 200
|
||||
|
||||
|
||||
def _format_history_message(msg: dict) -> str | None:
|
||||
def _format_history_message(msg: dict[str, Any]) -> str | None:
|
||||
"""Format a single history message for display. Returns None to skip."""
|
||||
role = msg.get("role")
|
||||
if role not in ("user", "assistant"):
|
||||
return None
|
||||
content = msg.get("content") or ""
|
||||
if isinstance(content, list):
|
||||
parts = [b.get("text", "") for b in content if isinstance(b, dict) and b.get("type") == "text"]
|
||||
parts = [
|
||||
text
|
||||
for block in cast(list[object], content)
|
||||
if (item := cast(dict[str, Any], block) if isinstance(block, dict) else None)
|
||||
and item.get("type") == "text"
|
||||
and isinstance(text := item.get("text"), str)
|
||||
]
|
||||
content = " ".join(parts)
|
||||
content = str(content).strip()
|
||||
if not content:
|
||||
@@ -863,6 +876,8 @@ async def cmd_history(ctx: CommandContext) -> OutboundMessage:
|
||||
|
||||
async def cmd_goal(ctx: CommandContext) -> OutboundMessage | None:
|
||||
"""Mark this turn as an explicit sustained-goal request."""
|
||||
from nanobot.agent.goal_permission import goal_mutation_permission
|
||||
|
||||
goal = ctx.args.strip()
|
||||
if not goal:
|
||||
return OutboundMessage(
|
||||
@@ -923,7 +938,7 @@ async def cmd_skill(ctx: CommandContext) -> OutboundMessage:
|
||||
else:
|
||||
lines = [f"Available skills ({len(skills)}):", ""]
|
||||
for entry in skills:
|
||||
desc = loop.context.skills._get_skill_description(entry["name"])
|
||||
desc = loop.context.skills.get_skill_description(entry["name"])
|
||||
lines.append(f"- **{entry['name']}** — {desc}")
|
||||
content = "\n".join(lines)
|
||||
return OutboundMessage(
|
||||
@@ -951,15 +966,9 @@ async def cmd_trigger(ctx: CommandContext) -> OutboundMessage:
|
||||
from nanobot.triggers.local_store import LocalTriggerStore
|
||||
|
||||
loop = ctx.loop
|
||||
workspace = getattr(loop, "workspace", None)
|
||||
if workspace is None:
|
||||
workspace = getattr(getattr(loop, "context", None), "workspace", None)
|
||||
if workspace is None:
|
||||
raise RuntimeError("workspace unavailable for trigger creation")
|
||||
|
||||
store = getattr(loop, "local_trigger_store", None)
|
||||
store = loop.local_trigger_store
|
||||
if store is None:
|
||||
store = LocalTriggerStore(workspace)
|
||||
store = LocalTriggerStore(loop.workspace)
|
||||
|
||||
from nanobot.session.keys import UNIFIED_SESSION_KEY
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
from nanobot.session.manager import Session
|
||||
from nanobot.utils.llm_runtime import LLMRuntime
|
||||
@@ -44,7 +45,7 @@ class CommandContext:
|
||||
key: str
|
||||
raw: str
|
||||
args: str = ""
|
||||
loop: Any = None
|
||||
loop: AgentLoop = field(kw_only=True)
|
||||
runtime: LLMRuntime | None = None
|
||||
is_user_turn: bool = False
|
||||
turn_scopes: list[AbstractContextManager[Any]] = field(default_factory=list)
|
||||
|
||||
+58
-24
@@ -4,20 +4,28 @@ import json
|
||||
import os
|
||||
import re
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast, overload
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
from pydantic_settings import SettingsError
|
||||
|
||||
from nanobot.config.errors import ConfigIssue, ConfigLoadError, validation_issues
|
||||
from nanobot.config.schema import Config, _resolve_tool_config_refs
|
||||
from nanobot.utils.helpers import _write_text_atomic
|
||||
from nanobot.config.schema import (
|
||||
Config,
|
||||
_resolve_tool_config_refs, # pyright: ignore[reportPrivateUsage]
|
||||
)
|
||||
from nanobot.utils.helpers import _write_text_atomic # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
# Global variable to store current config path (for multi-instance support)
|
||||
_current_config_path: Path | None = None
|
||||
_schema_refs_ready = False
|
||||
|
||||
|
||||
def _as_config_object(value: object) -> dict[str, Any] | None:
|
||||
"""Narrow an untrusted JSON configuration value to an object."""
|
||||
return cast(dict[str, Any], value) if isinstance(value, dict) else None
|
||||
|
||||
|
||||
def set_config_path(path: Path) -> None:
|
||||
"""Set the current config path (used to derive data directory)."""
|
||||
global _current_config_path
|
||||
@@ -110,7 +118,7 @@ def load_config(config_path: Path | None = None) -> Config:
|
||||
),
|
||||
)
|
||||
|
||||
data = _migrate_config(data)
|
||||
data = _migrate_config(cast(dict[str, Any], data))
|
||||
try:
|
||||
config = Config.model_validate(data)
|
||||
except ValidationError as exc:
|
||||
@@ -164,13 +172,15 @@ def save_config(config: Config, config_path: Path | None = None) -> None:
|
||||
_write_text_atomic(path, json.dumps(data, indent=2, ensure_ascii=False))
|
||||
|
||||
|
||||
def merge_missing_defaults(existing: Any, defaults: Any) -> Any:
|
||||
def merge_missing_defaults(existing: object, defaults: object) -> object:
|
||||
"""Recursively add missing defaults without replacing configured values."""
|
||||
if not isinstance(existing, dict) or not isinstance(defaults, dict):
|
||||
return existing
|
||||
return cast(object, existing)
|
||||
|
||||
merged = dict(existing)
|
||||
for key, value in defaults.items():
|
||||
existing_dict = cast(dict[str, object], existing)
|
||||
defaults_dict = cast(dict[str, object], defaults)
|
||||
merged = dict(existing_dict)
|
||||
for key, value in defaults_dict.items():
|
||||
if key not in merged:
|
||||
merged[key] = value
|
||||
else:
|
||||
@@ -203,7 +213,15 @@ def resolve_config_env_vars(
|
||||
return _resolve_in_place(config)
|
||||
|
||||
|
||||
def resolve_env_refs(value: str) -> str:
|
||||
@overload
|
||||
def resolve_env_refs(value: str) -> str: ...
|
||||
|
||||
|
||||
@overload
|
||||
def resolve_env_refs(value: object) -> object: ...
|
||||
|
||||
|
||||
def resolve_env_refs(value: object) -> object:
|
||||
"""Resolve ``${VAR}`` references in a single string, leniently.
|
||||
|
||||
Unlike :func:`resolve_config_env_vars` (which walks a whole ``Config`` and
|
||||
@@ -245,11 +263,21 @@ def _resolve_in_place(obj: Any) -> Any:
|
||||
copy.__pydantic_extra__ = new_extras
|
||||
return copy
|
||||
if isinstance(obj, dict):
|
||||
resolved = {k: _resolve_in_place(v) for k, v in obj.items()}
|
||||
return resolved if any(resolved[k] is not obj[k] for k in obj) else obj
|
||||
object_dict = cast(dict[str, Any], obj)
|
||||
resolved = {key: _resolve_in_place(value) for key, value in object_dict.items()}
|
||||
return (
|
||||
resolved
|
||||
if any(resolved[key] is not object_dict[key] for key in object_dict)
|
||||
else cast(object, obj)
|
||||
)
|
||||
if isinstance(obj, list):
|
||||
resolved = [_resolve_in_place(v) for v in obj]
|
||||
return resolved if any(nv is not ov for nv, ov in zip(resolved, obj)) else obj
|
||||
object_list = cast(list[Any], obj)
|
||||
resolved = [_resolve_in_place(value) for value in object_list]
|
||||
return (
|
||||
resolved
|
||||
if any(new is not old for new, old in zip(resolved, object_list))
|
||||
else cast(object, obj)
|
||||
)
|
||||
return obj
|
||||
|
||||
|
||||
@@ -270,20 +298,21 @@ def _missing_env_issues(
|
||||
issues: list[ConfigIssue] = []
|
||||
for name, field in type(obj).model_fields.items():
|
||||
alias = field.serialization_alias or field.alias or name
|
||||
part = alias if isinstance(alias, str) else name
|
||||
part = alias
|
||||
issues.extend(_missing_env_issues(getattr(obj, name), (*path, part)))
|
||||
for name, value in (obj.__pydantic_extra__ or {}).items():
|
||||
issues.extend(_missing_env_issues(value, (*path, name)))
|
||||
return issues
|
||||
if isinstance(obj, dict):
|
||||
object_dict = cast(dict[str | int, Any], obj)
|
||||
issues = []
|
||||
for name, value in obj.items():
|
||||
part = name if isinstance(name, (str, int)) else str(name)
|
||||
for name, value in object_dict.items():
|
||||
part = name
|
||||
issues.extend(_missing_env_issues(value, (*path, part)))
|
||||
return issues
|
||||
if isinstance(obj, list):
|
||||
issues = []
|
||||
for index, value in enumerate(obj):
|
||||
for index, value in enumerate(cast(list[Any], obj)):
|
||||
issues.extend(_missing_env_issues(value, (*path, index)))
|
||||
return issues
|
||||
return []
|
||||
@@ -294,9 +323,12 @@ def _resolve_env_vars(obj: object) -> object:
|
||||
if isinstance(obj, str):
|
||||
return _ENV_REF_PATTERN.sub(_env_replace, obj)
|
||||
if isinstance(obj, dict):
|
||||
return {k: _resolve_env_vars(v) for k, v in obj.items()}
|
||||
return {
|
||||
key: _resolve_env_vars(value)
|
||||
for key, value in cast(dict[str, object], obj).items()
|
||||
}
|
||||
if isinstance(obj, list):
|
||||
return [_resolve_env_vars(v) for v in obj]
|
||||
return [_resolve_env_vars(value) for value in cast(list[object], obj)]
|
||||
return obj
|
||||
|
||||
|
||||
@@ -310,15 +342,16 @@ def _env_replace(match: re.Match[str]) -> str:
|
||||
return value
|
||||
|
||||
|
||||
def _migrate_config(data: dict) -> dict:
|
||||
def _migrate_config(data: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Migrate old config formats to current."""
|
||||
# Move tools.exec.restrictToWorkspace → tools.restrictToWorkspace
|
||||
tools = data.get("tools", {})
|
||||
if not isinstance(tools, dict):
|
||||
tools_value = data.get("tools", {})
|
||||
if not isinstance(tools_value, dict):
|
||||
return data
|
||||
exec_cfg = tools.get("exec", {})
|
||||
tools = cast(dict[str, Any], tools_value)
|
||||
exec_cfg = _as_config_object(tools.get("exec", {}))
|
||||
if (
|
||||
isinstance(exec_cfg, dict)
|
||||
exec_cfg is not None
|
||||
and "restrictToWorkspace" in exec_cfg
|
||||
and "restrictToWorkspace" not in tools
|
||||
):
|
||||
@@ -334,6 +367,7 @@ def _migrate_config(data: dict) -> dict:
|
||||
tools["my"] = my_cfg
|
||||
if not isinstance(my_cfg, dict):
|
||||
return data
|
||||
my_cfg = cast(dict[str, Any], my_cfg)
|
||||
if "myEnabled" in tools and "enable" not in my_cfg:
|
||||
my_cfg["enable"] = tools.pop("myEnabled")
|
||||
else:
|
||||
|
||||
@@ -48,7 +48,7 @@ def get_webui_dir() -> Path:
|
||||
return get_runtime_subdir("webui")
|
||||
|
||||
|
||||
def get_workspace_path(workspace: str | None = None) -> Path:
|
||||
def get_workspace_path(workspace: str | Path | None = None) -> Path:
|
||||
"""Resolve and ensure the agent workspace path."""
|
||||
path = Path(workspace).expanduser() if workspace else Path.home() / ".nanobot" / "workspace"
|
||||
return ensure_dir(path)
|
||||
|
||||
@@ -5,7 +5,7 @@ from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Literal
|
||||
|
||||
from pydantic import AliasChoices, ConfigDict, Field, field_validator, model_validator
|
||||
from pydantic_settings import BaseSettings
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
from nanobot.config_base import Base
|
||||
from nanobot.cron.types import CronSchedule
|
||||
@@ -618,7 +618,10 @@ class Config(BaseSettings):
|
||||
return spec.default_api_base
|
||||
return None
|
||||
|
||||
model_config = ConfigDict(env_prefix="NANOBOT_", env_nested_delimiter="__")
|
||||
model_config = SettingsConfigDict(
|
||||
env_prefix="NANOBOT_",
|
||||
env_nested_delimiter="__",
|
||||
)
|
||||
|
||||
|
||||
def _resolve_tool_config_refs() -> None:
|
||||
|
||||
@@ -5,7 +5,7 @@ from __future__ import annotations
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
|
||||
from watchfiles import Change, awatch
|
||||
from watchfiles import Change, awatch # pyright: ignore[reportUnknownVariableType]
|
||||
|
||||
|
||||
async def watch_config_file(config_path: Path, on_change: Callable[[], None]) -> None:
|
||||
|
||||
@@ -1,13 +1,18 @@
|
||||
"""Cron service for scheduled agent tasks."""
|
||||
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from nanobot.cron.types import CronJob, CronSchedule
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.cron.service import CronService
|
||||
|
||||
__all__ = ["CronService", "CronJob", "CronSchedule"]
|
||||
|
||||
_LAZY = {"CronService": ".service"}
|
||||
|
||||
|
||||
def __getattr__(name: str):
|
||||
def __getattr__(name: str) -> Any:
|
||||
module_path = _LAZY.get(name)
|
||||
if module_path is None:
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
@@ -6,7 +6,7 @@ import asyncio
|
||||
import hashlib
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any, Protocol
|
||||
from typing import TYPE_CHECKING, Any, Protocol
|
||||
|
||||
from nanobot.agent.tools.cron import CronTool
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
@@ -16,9 +16,12 @@ from nanobot.cron.types import CronJob
|
||||
from nanobot.cron.webui_metadata import cron_proactive_delivery_metadata
|
||||
from nanobot.utils.prompt_templates import render_template
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
|
||||
|
||||
class BoundCronAgent(Protocol):
|
||||
tools: Any
|
||||
tools: ToolRegistry
|
||||
|
||||
async def submit_cron_turn(self, msg: InboundMessage) -> OutboundMessage | None:
|
||||
...
|
||||
|
||||
+23
-15
@@ -10,6 +10,7 @@ from contextlib import suppress
|
||||
from dataclasses import asdict
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from types import EllipsisType
|
||||
from typing import Any, Callable, Coroutine, Literal
|
||||
|
||||
from filelock import FileLock
|
||||
@@ -160,7 +161,7 @@ class CronService:
|
||||
self._lock = FileLock(str(self._action_path.parent) + ".lock")
|
||||
self.on_job = on_job
|
||||
self._store: CronStore | None = None
|
||||
self._timer_task: asyncio.Task | None = None
|
||||
self._timer_task: asyncio.Task[None] | None = None
|
||||
self._running = False
|
||||
self._timer_active = False
|
||||
self.max_sleep_ms = max_sleep_ms
|
||||
@@ -243,19 +244,21 @@ class CronService:
|
||||
return None
|
||||
return jobs, version
|
||||
|
||||
def _merge_action(self):
|
||||
def _merge_action(self) -> None:
|
||||
if not self._action_path.exists():
|
||||
return
|
||||
|
||||
jobs_map = {j.id: j for j in self._store.jobs}
|
||||
def _update(params: dict):
|
||||
jobs_map = {job.id: job for job in self._store.jobs} # pyright: ignore[reportOptionalMemberAccess]
|
||||
|
||||
def _update(params: dict[str, Any]) -> None:
|
||||
j = CronJob.from_dict(params)
|
||||
_normalize_agent_turn_job(j)
|
||||
jobs_map[j.id] = j
|
||||
|
||||
def _del(params: dict):
|
||||
if job_id := params.get("job_id"):
|
||||
jobs_map.pop(job_id)
|
||||
def _del(params: dict[str, Any]) -> None:
|
||||
job_id = params.get("job_id")
|
||||
if isinstance(job_id, str) and job_id:
|
||||
jobs_map.pop(job_id, None)
|
||||
|
||||
with self._lock:
|
||||
with open(self._action_path, "r", encoding="utf-8") as f:
|
||||
@@ -274,7 +277,7 @@ class CronService:
|
||||
except Exception:
|
||||
logger.exception("load action line error")
|
||||
continue
|
||||
self._store.jobs = list(jobs_map.values())
|
||||
self._store.jobs = list(jobs_map.values()) # pyright: ignore[reportOptionalMemberAccess]
|
||||
if self._running and changed:
|
||||
self._action_path.write_text("", encoding="utf-8")
|
||||
self._save_store()
|
||||
@@ -569,7 +572,8 @@ class CronService:
|
||||
# Handle one-shot jobs
|
||||
if job.schedule.kind == "at":
|
||||
if job.delete_after_run:
|
||||
self._store.jobs = [j for j in self._store.jobs if j.id != job.id]
|
||||
store = self._require_store()
|
||||
store.jobs = [item for item in store.jobs if item.id != job.id]
|
||||
else:
|
||||
job.enabled = False
|
||||
job.state.next_run_at_ms = None
|
||||
@@ -577,7 +581,11 @@ class CronService:
|
||||
# Compute next run
|
||||
job.state.next_run_at_ms = _compute_next_run(job.schedule, _now_ms())
|
||||
|
||||
def _append_action(self, action: Literal["add", "del", "update"], params: dict):
|
||||
def _append_action(
|
||||
self,
|
||||
action: Literal["add", "del", "update"],
|
||||
params: dict[str, Any],
|
||||
) -> None:
|
||||
self.store_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with self._lock:
|
||||
with open(self._action_path, "a", encoding="utf-8") as f:
|
||||
@@ -615,11 +623,11 @@ class CronService:
|
||||
channel: str | None = None,
|
||||
to: str | None = None,
|
||||
delete_after_run: bool = False,
|
||||
channel_meta: dict | None = None,
|
||||
channel_meta: dict[str, Any] | None = None,
|
||||
session_key: str | None = None,
|
||||
origin_channel: str | None = None,
|
||||
origin_chat_id: str | None = None,
|
||||
origin_metadata: dict | None = None,
|
||||
origin_metadata: dict[str, Any] | None = None,
|
||||
) -> CronJob:
|
||||
"""Add a new job."""
|
||||
_validate_schedule_for_add(schedule)
|
||||
@@ -727,8 +735,8 @@ class CronService:
|
||||
schedule: CronSchedule | None = None,
|
||||
message: str | None = None,
|
||||
deliver: bool | None = None,
|
||||
channel: str | None = ...,
|
||||
to: str | None = ...,
|
||||
channel: str | None | EllipsisType = ...,
|
||||
to: str | None | EllipsisType = ...,
|
||||
delete_after_run: bool | None = None,
|
||||
) -> CronJob | Literal["not_found", "protected"]:
|
||||
"""Update mutable fields of an existing job. System jobs cannot be updated.
|
||||
@@ -804,7 +812,7 @@ class CronService:
|
||||
store = self._require_store()
|
||||
return next((j for j in store.jobs if j.id == job_id), None)
|
||||
|
||||
def status(self) -> dict:
|
||||
def status(self) -> dict[str, object]:
|
||||
"""Get service status."""
|
||||
store = self._require_store()
|
||||
return {
|
||||
|
||||
+25
-10
@@ -3,11 +3,19 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Literal
|
||||
from typing import Any, Literal, cast, overload
|
||||
|
||||
from nanobot.utils.dict_keys import get_camel_snake
|
||||
|
||||
|
||||
@overload
|
||||
def _store_int(value: Any, default: Literal[None]) -> int | None: ...
|
||||
|
||||
|
||||
@overload
|
||||
def _store_int(value: Any, default: int = 0) -> int: ...
|
||||
|
||||
|
||||
def _store_int(value: Any, default: int | None = 0) -> int | None:
|
||||
"""Coerce JSON numerics to int; treat null/blank like a missing key."""
|
||||
if value is None or value == "":
|
||||
@@ -103,7 +111,10 @@ class CronJobState:
|
||||
|
||||
@classmethod
|
||||
def from_store_dict(cls, data: dict[str, Any]) -> CronJobState:
|
||||
history = get_camel_snake(data, "runHistory", "run_history", []) or []
|
||||
history = cast(
|
||||
list[object],
|
||||
get_camel_snake(data, "runHistory", "run_history", []) or [],
|
||||
)
|
||||
return cls(
|
||||
next_run_at_ms=_store_int(
|
||||
get_camel_snake(data, "nextRunAtMs", "next_run_at_ms"), None
|
||||
@@ -116,7 +127,7 @@ class CronJobState:
|
||||
run_history=[
|
||||
record
|
||||
if isinstance(record, CronRunRecord)
|
||||
else CronRunRecord.from_store_dict(record)
|
||||
else CronRunRecord.from_store_dict(cast(dict[str, Any], record))
|
||||
for record in history
|
||||
if isinstance(record, (dict, CronRunRecord))
|
||||
],
|
||||
@@ -137,16 +148,20 @@ class CronJob:
|
||||
delete_after_run: bool = False
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, kwargs: dict):
|
||||
state_kwargs = dict(kwargs.get("state", {}))
|
||||
def from_dict(cls, kwargs: dict[str, Any]) -> CronJob:
|
||||
state_kwargs = dict(cast(dict[str, Any], kwargs.get("state", {})))
|
||||
state_kwargs["run_history"] = [
|
||||
record if isinstance(record, CronRunRecord) else CronRunRecord(**record)
|
||||
for record in state_kwargs.get("run_history", [])
|
||||
record
|
||||
if isinstance(record, CronRunRecord)
|
||||
else CronRunRecord(**cast(dict[str, Any], record))
|
||||
for record in cast(list[object], state_kwargs.get("run_history", []))
|
||||
]
|
||||
kwargs["schedule"] = CronSchedule(**kwargs.get("schedule", {"kind": "every"}))
|
||||
kwargs["payload"] = CronPayload(**kwargs.get("payload", {}))
|
||||
kwargs["schedule"] = CronSchedule(
|
||||
**cast(dict[str, Any], kwargs.get("schedule", {"kind": "every"}))
|
||||
)
|
||||
kwargs["payload"] = CronPayload(**cast(dict[str, Any], kwargs.get("payload", {})))
|
||||
kwargs["state"] = CronJobState(**state_kwargs)
|
||||
return cls(**kwargs)
|
||||
return cls(**cast(Any, kwargs))
|
||||
|
||||
@classmethod
|
||||
def from_store_dict(cls, data: dict[str, Any]) -> CronJob:
|
||||
|
||||
@@ -69,7 +69,7 @@ class GatewayRuntimePaths(ProcessRuntimePaths):
|
||||
)
|
||||
|
||||
|
||||
class GatewayRuntime(ManagedProcessRuntime):
|
||||
class GatewayRuntime(ManagedProcessRuntime[ProcessStartOptions]):
|
||||
"""Manage a background ``nanobot gateway`` process."""
|
||||
|
||||
service_name = "gateway"
|
||||
|
||||
@@ -7,7 +7,7 @@ import sys
|
||||
from dataclasses import dataclass
|
||||
from importlib.metadata import PackageNotFoundError, distribution
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
from loguru import logger
|
||||
from packaging.requirements import Requirement
|
||||
@@ -87,8 +87,8 @@ def optional_dependency_groups() -> dict[str, list[str] | None]:
|
||||
deps = project.get("optional-dependencies", {})
|
||||
if isinstance(deps, dict) and deps:
|
||||
return {
|
||||
name: list(values)
|
||||
for name, values in deps.items()
|
||||
name: list(cast(list[str], values))
|
||||
for name, values in cast(dict[str, object], deps).items()
|
||||
if name != "dev" and name not in _HIDDEN_OPTIONAL_FEATURES and isinstance(values, list)
|
||||
}
|
||||
return {
|
||||
@@ -153,13 +153,13 @@ def _extra_dependencies_installed(
|
||||
normalized = canonicalize_name(requested_extra)
|
||||
provided = {
|
||||
canonicalize_name(value)
|
||||
for value in (dist.metadata.get_all("Provides-Extra") or [])
|
||||
for value in cast(list[str], dist.metadata.get_all("Provides-Extra") or [])
|
||||
}
|
||||
if provided and normalized not in provided:
|
||||
return False
|
||||
|
||||
matched = False
|
||||
for raw in dist.requires or []:
|
||||
for raw in cast(list[str], dist.requires or []):
|
||||
req = Requirement(raw)
|
||||
if req.marker and not req.marker.evaluate({"extra": requested_extra}):
|
||||
continue
|
||||
@@ -259,7 +259,7 @@ def read_config_data(path: Path) -> dict[str, Any]:
|
||||
if not path.exists():
|
||||
return {}
|
||||
with open(path, encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
return cast(dict[str, Any], json.load(f))
|
||||
|
||||
|
||||
def write_config_data(path: Path, data: dict[str, Any]) -> None:
|
||||
@@ -312,7 +312,7 @@ def channel_enabled(
|
||||
if default_enabled is None:
|
||||
default_enabled = plugin.default_enabled if plugin is not None else channel_default_enabled(name)
|
||||
if section is None:
|
||||
return default_enabled
|
||||
return bool(default_enabled)
|
||||
if plugin is None:
|
||||
from nanobot.channels.registry import load_channel_plugin
|
||||
|
||||
@@ -421,7 +421,7 @@ def optional_features_payload(
|
||||
dependencies = _feature_dependencies(name, channel_plugin, extras)
|
||||
has_dependencies = bool(dependencies)
|
||||
installed = extra_installed(name, dependencies) if has_dependencies else True
|
||||
feature = {
|
||||
feature: dict[str, Any] = {
|
||||
"name": name,
|
||||
"display_name": (
|
||||
channel_plugin.display_name
|
||||
@@ -502,7 +502,7 @@ def optional_features_payload(
|
||||
})
|
||||
features.append(feature)
|
||||
|
||||
payload = {
|
||||
payload: dict[str, Any] = {
|
||||
"features": features,
|
||||
"enabled_count": sum(1 for feature in features if feature["enabled"]),
|
||||
}
|
||||
@@ -520,13 +520,16 @@ def with_channel_runtime_status(
|
||||
for status in runtime_status.values():
|
||||
if not isinstance(status, dict):
|
||||
continue
|
||||
owner = status.get("owner")
|
||||
status_object = cast(dict[str, Any], status)
|
||||
owner = status_object.get("owner")
|
||||
if isinstance(owner, str):
|
||||
statuses_by_owner.setdefault(owner, []).append(status)
|
||||
statuses_by_owner.setdefault(owner, []).append(status_object)
|
||||
|
||||
features: list[dict[str, Any]] = []
|
||||
for original in payload.get("features", []):
|
||||
feature = dict(original)
|
||||
for raw_feature in cast(list[object], payload.get("features", [])):
|
||||
if not isinstance(raw_feature, dict):
|
||||
continue
|
||||
feature = cast(dict[str, Any], raw_feature).copy()
|
||||
if feature.get("type") != "channel":
|
||||
features.append(feature)
|
||||
continue
|
||||
@@ -546,9 +549,11 @@ def with_channel_runtime_status(
|
||||
str(status.get("instance_id", "default")): status
|
||||
for status in owner_statuses
|
||||
}
|
||||
decorated_instances = []
|
||||
for original_instance in instances:
|
||||
instance = dict(original_instance)
|
||||
decorated_instances: list[dict[str, Any]] = []
|
||||
for original_instance in cast(list[object], instances):
|
||||
if not isinstance(original_instance, dict):
|
||||
continue
|
||||
instance = cast(dict[str, Any], original_instance).copy()
|
||||
desired_instance = bool(instance.get("enabled"))
|
||||
status = by_instance.get(str(instance.get("id", "default")))
|
||||
if desired_instance and status is None:
|
||||
|
||||
+32
-35
@@ -13,12 +13,12 @@ import string
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.config.paths import get_data_dir
|
||||
from nanobot.utils.helpers import _write_text_atomic
|
||||
from nanobot.utils.helpers import _write_text_atomic # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
# threading.Lock is used so store functions remain callable from both sync CLI
|
||||
# and async channel handlers. At private-assistant scale (small JSON file,
|
||||
@@ -47,21 +47,20 @@ def _load() -> dict[str, Any]:
|
||||
logger.warning("Corrupted pairing store, resetting")
|
||||
return {"approved": {}, "pending": {}}
|
||||
|
||||
# JSON stores may contain null maps after partial edits; treat like {}.
|
||||
approved = data.get("approved") or {}
|
||||
if not isinstance(approved, dict):
|
||||
approved = {}
|
||||
# JSON stores may contain null or malformed maps after partial edits; treat like {}.
|
||||
data = cast(dict[str, Any], data)
|
||||
raw_approved = data.get("approved")
|
||||
approved = cast(dict[str, Any], raw_approved) if isinstance(raw_approved, dict) else {}
|
||||
data["approved"] = approved
|
||||
pending = data.get("pending") or {}
|
||||
if not isinstance(pending, dict):
|
||||
pending = {}
|
||||
raw_pending = data.get("pending")
|
||||
pending = cast(dict[str, Any], raw_pending) if isinstance(raw_pending, dict) else {}
|
||||
data["pending"] = pending
|
||||
|
||||
# Convert approved lists to str sets for O(1) lookup.
|
||||
for channel, users in approved.items():
|
||||
if not isinstance(users, list):
|
||||
users = []
|
||||
data["approved"][channel] = {str(u) for u in users}
|
||||
data["approved"][channel] = {str(user) for user in cast(list[object], users)}
|
||||
return data
|
||||
|
||||
|
||||
@@ -69,14 +68,12 @@ def _save(data: dict[str, Any]) -> None:
|
||||
path = _store_path()
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
# Convert sets back to lists for JSON serialization
|
||||
approved = data.get("approved") or {}
|
||||
pending = data.get("pending") or {}
|
||||
if not isinstance(approved, dict):
|
||||
approved = {}
|
||||
if not isinstance(pending, dict):
|
||||
pending = {}
|
||||
payload = {
|
||||
"approved": {ch: sorted(list(users)) for ch, users in approved.items()},
|
||||
raw_approved = data.get("approved")
|
||||
approved = cast(dict[str, Any], raw_approved) if isinstance(raw_approved, dict) else {}
|
||||
raw_pending = data.get("pending")
|
||||
pending = cast(dict[str, Any], raw_pending) if isinstance(raw_pending, dict) else {}
|
||||
payload: dict[str, Any] = {
|
||||
"approved": {ch: sorted(list(cast(set[str], users))) for ch, users in approved.items()},
|
||||
"pending": dict(pending),
|
||||
}
|
||||
_write_text_atomic(path, json.dumps(payload, indent=2, ensure_ascii=False))
|
||||
@@ -86,22 +83,22 @@ def _gc_pending(data: dict[str, Any]) -> None:
|
||||
"""Remove expired pending entries in-place."""
|
||||
now = time.time()
|
||||
pending: dict[str, Any] = data.get("pending") or {}
|
||||
if not isinstance(pending, dict):
|
||||
data["pending"] = {}
|
||||
return
|
||||
expired = [
|
||||
code
|
||||
for code, info in pending.items()
|
||||
expired: list[str] = []
|
||||
for code, info in pending.items():
|
||||
if not isinstance(info, dict):
|
||||
expired.append(code)
|
||||
continue
|
||||
entry = cast(dict[str, Any], info)
|
||||
expires_at = entry.get("expires_at")
|
||||
if (
|
||||
not isinstance(info, dict)
|
||||
or not isinstance(info.get("channel"), str)
|
||||
or not info.get("channel")
|
||||
or info.get("sender_id") is None
|
||||
or isinstance(info.get("expires_at"), bool)
|
||||
or not isinstance(info.get("expires_at"), (int, float))
|
||||
or info["expires_at"] < now
|
||||
)
|
||||
]
|
||||
not isinstance(entry.get("channel"), str)
|
||||
or not entry["channel"]
|
||||
or entry.get("sender_id") is None
|
||||
or isinstance(expires_at, bool)
|
||||
or not isinstance(expires_at, (int, float))
|
||||
or expires_at < now
|
||||
):
|
||||
expired.append(code)
|
||||
for code in expired:
|
||||
del pending[code]
|
||||
data["pending"] = pending
|
||||
@@ -322,13 +319,13 @@ def handle_pairing_command(channel: str, subcommand_text: str) -> str:
|
||||
if len(parts) == 2:
|
||||
return (
|
||||
f"Revoked {arg} from {channel}"
|
||||
if revoke(channel, arg)
|
||||
if revoke(channel, parts[1])
|
||||
else f"{arg} was not in the approved list for {channel}"
|
||||
)
|
||||
if len(parts) == 3:
|
||||
return (
|
||||
f"Revoked {parts[2]} from {arg}"
|
||||
if revoke(arg, parts[2])
|
||||
if revoke(parts[1], parts[2])
|
||||
else f"{parts[2]} was not in the approved list for {arg}"
|
||||
)
|
||||
return "Usage: `/pairing revoke <user_id>` or `/pairing revoke <channel> <user_id>`"
|
||||
|
||||
@@ -15,7 +15,7 @@ from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from datetime import UTC, datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, Generic, TypeVar, cast
|
||||
|
||||
from filelock import FileLock
|
||||
|
||||
@@ -63,7 +63,10 @@ class ProcessRuntimePaths:
|
||||
log_path: Path
|
||||
|
||||
|
||||
class ManagedProcessRuntime:
|
||||
_StartOptionsT = TypeVar("_StartOptionsT", bound=ProcessStartOptions)
|
||||
|
||||
|
||||
class ManagedProcessRuntime(Generic[_StartOptionsT]):
|
||||
"""Manage a detached child process without service-specific policy."""
|
||||
|
||||
service_name = "process"
|
||||
@@ -100,12 +103,12 @@ class ManagedProcessRuntime:
|
||||
state["started_at"] = _utc_now()
|
||||
runtime._write_state(state)
|
||||
|
||||
def start_background(self, options: ProcessStartOptions) -> ProcessResult:
|
||||
def start_background(self, options: _StartOptionsT) -> ProcessResult:
|
||||
"""Start the configured command as a detached process."""
|
||||
with self._lifecycle_lock():
|
||||
return self._start_background(options)
|
||||
|
||||
def _start_background(self, options: ProcessStartOptions) -> ProcessResult:
|
||||
def _start_background(self, options: _StartOptionsT) -> ProcessResult:
|
||||
current = self.status()
|
||||
if current.running:
|
||||
return ProcessResult(False, self._message("already_running"), current)
|
||||
@@ -174,7 +177,7 @@ class ManagedProcessRuntime:
|
||||
self._clear_state()
|
||||
return ProcessResult(True, self._message("stopped"), self.status(reason="stopped"))
|
||||
|
||||
def restart(self, options: ProcessStartOptions, *, timeout_s: int = 20) -> ProcessResult:
|
||||
def restart(self, options: _StartOptionsT, *, timeout_s: int = 20) -> ProcessResult:
|
||||
"""Restart the managed process."""
|
||||
with self._lifecycle_lock():
|
||||
stop_result = self._stop(timeout_s=timeout_s)
|
||||
@@ -195,6 +198,7 @@ class ManagedProcessRuntime:
|
||||
log_path=self.paths.log_path,
|
||||
reason=reason or "not_started",
|
||||
)
|
||||
assert state is not None
|
||||
|
||||
if not self._is_pid_running(pid) or not self._record_matches_process(state, pid):
|
||||
self._clear_state()
|
||||
@@ -214,7 +218,7 @@ class ManagedProcessRuntime:
|
||||
log_path=self.paths.log_path,
|
||||
started_at=_as_str(state.get("started_at")),
|
||||
port=_as_int(state.get("port")),
|
||||
command=tuple(command) if isinstance(command, list) else (),
|
||||
command=tuple(cast(list[str], command)) if isinstance(command, list) else (),
|
||||
reason=reason or "running",
|
||||
)
|
||||
|
||||
@@ -253,7 +257,7 @@ class ManagedProcessRuntime:
|
||||
lock_path = self.paths.state_path.with_name(f"{self.paths.state_path.name}.lock")
|
||||
return FileLock(str(lock_path))
|
||||
|
||||
def _build_child_command(self, options: ProcessStartOptions) -> list[str]:
|
||||
def _build_child_command(self, options: _StartOptionsT) -> list[str]:
|
||||
raise NotImplementedError
|
||||
|
||||
def _popen_platform_kwargs(self) -> dict[str, Any]:
|
||||
@@ -365,7 +369,7 @@ class ManagedProcessRuntime:
|
||||
payload = json.load(handle)
|
||||
except (OSError, json.JSONDecodeError, ValueError):
|
||||
return None
|
||||
return payload if isinstance(payload, dict) else None
|
||||
return cast(dict[str, Any], payload) if isinstance(payload, dict) else None
|
||||
|
||||
def _write_state(self, payload: dict[str, Any]) -> None:
|
||||
self.paths.run_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
@@ -9,8 +9,8 @@ import re
|
||||
import secrets
|
||||
import string
|
||||
from collections import deque
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
from collections.abc import Awaitable, Callable, Iterable
|
||||
from typing import Any, cast
|
||||
|
||||
from loguru import logger
|
||||
|
||||
@@ -198,7 +198,13 @@ class AnthropicProvider(LLMProvider):
|
||||
content = msg.get("content")
|
||||
|
||||
if role == "system":
|
||||
system = content if isinstance(content, (str, list)) else str(content or "")
|
||||
system = (
|
||||
cast(list[dict[str, Any]], content)
|
||||
if isinstance(content, list)
|
||||
else content
|
||||
if isinstance(content, str)
|
||||
else str(content or "")
|
||||
)
|
||||
continue
|
||||
|
||||
if role == "tool":
|
||||
@@ -206,7 +212,7 @@ class AnthropicProvider(LLMProvider):
|
||||
if raw and raw[-1]["role"] == "user":
|
||||
prev_c = raw[-1]["content"]
|
||||
if isinstance(prev_c, list):
|
||||
prev_c.append(block)
|
||||
cast(list[Any], prev_c).append(block)
|
||||
else:
|
||||
raw[-1]["content"] = [
|
||||
{"type": "text", "text": prev_c or ""}, block,
|
||||
@@ -264,41 +270,49 @@ class AnthropicProvider(LLMProvider):
|
||||
blocks: list[dict[str, Any]] = []
|
||||
content = msg.get("content")
|
||||
|
||||
for tb in msg.get("thinking_blocks") or []:
|
||||
if isinstance(tb, dict) and tb.get("type") == "thinking":
|
||||
blocks.append({
|
||||
"type": "thinking",
|
||||
"thinking": tb.get("thinking", ""),
|
||||
"signature": tb.get("signature", ""),
|
||||
})
|
||||
for tb in cast(Iterable[object], msg.get("thinking_blocks") or []):
|
||||
if isinstance(tb, dict):
|
||||
thinking_block = cast(dict[str, Any], tb)
|
||||
if thinking_block.get("type") == "thinking":
|
||||
blocks.append({
|
||||
"type": "thinking",
|
||||
"thinking": thinking_block.get("thinking", ""),
|
||||
"signature": thinking_block.get("signature", ""),
|
||||
})
|
||||
|
||||
if isinstance(content, str) and content:
|
||||
blocks.append({"type": "text", "text": content})
|
||||
elif isinstance(content, list):
|
||||
for item in content:
|
||||
for item in cast(list[object], content):
|
||||
if isinstance(item, dict):
|
||||
if not item.get("type"):
|
||||
content_block = cast(dict[str, Any], item)
|
||||
if not content_block.get("type"):
|
||||
# Anthropic requires every content block to declare a "type".
|
||||
# A tool that returned a bare dict lands here; coerce it to
|
||||
# a text block instead of emitting one that the API rejects.
|
||||
blocks.append({
|
||||
"type": "text",
|
||||
"text": AnthropicProvider._stringify_typeless_block(item),
|
||||
"text": AnthropicProvider._stringify_typeless_block(content_block),
|
||||
})
|
||||
else:
|
||||
blocks.append(item)
|
||||
blocks.append(content_block)
|
||||
else:
|
||||
blocks.append({"type": "text", "text": str(item)})
|
||||
|
||||
for tc in msg.get("tool_calls") or []:
|
||||
for tc in cast(Iterable[object], msg.get("tool_calls") or []):
|
||||
if not isinstance(tc, dict):
|
||||
continue
|
||||
func = tc.get("function", {})
|
||||
tool_call = cast(dict[str, Any], tc)
|
||||
func = cast(dict[str, Any], tool_call.get("function", {}))
|
||||
args = func.get("arguments", "{}")
|
||||
raw_id = tc.get("id") or _gen_tool_id()
|
||||
raw_id = tool_call.get("id") or _gen_tool_id()
|
||||
blocks.append({
|
||||
"type": "tool_use",
|
||||
"id": map_tool_id(raw_id) if map_tool_id is not None else _sanitize_tool_id(raw_id),
|
||||
"id": (
|
||||
map_tool_id(raw_id)
|
||||
if map_tool_id is not None
|
||||
else _sanitize_tool_id(cast(str, raw_id))
|
||||
),
|
||||
"name": func.get("name", ""),
|
||||
"input": tool_arguments_object_for_replay(args),
|
||||
})
|
||||
@@ -314,26 +328,27 @@ class AnthropicProvider(LLMProvider):
|
||||
return str(content)
|
||||
|
||||
result: list[dict[str, Any]] = []
|
||||
for item in content:
|
||||
for item in cast(list[object], content):
|
||||
if not isinstance(item, dict):
|
||||
result.append({"type": "text", "text": str(item)})
|
||||
continue
|
||||
if item.get("type") == "image_url":
|
||||
converted = AnthropicProvider._convert_image_block(item)
|
||||
content_block = cast(dict[str, Any], item)
|
||||
if content_block.get("type") == "image_url":
|
||||
converted = AnthropicProvider._convert_image_block(content_block)
|
||||
if converted:
|
||||
result.append(converted)
|
||||
continue
|
||||
if not item.get("type"):
|
||||
if not content_block.get("type"):
|
||||
# Anthropic requires every content block to declare a "type".
|
||||
# A tool that returned a bare dict (or a list of dicts) lands
|
||||
# here; coerce it to a text block instead of emitting a block
|
||||
# the API rejects with "content.0.type: Field required".
|
||||
result.append({
|
||||
"type": "text",
|
||||
"text": AnthropicProvider._stringify_typeless_block(item),
|
||||
"text": AnthropicProvider._stringify_typeless_block(content_block),
|
||||
})
|
||||
continue
|
||||
result.append(item)
|
||||
result.append(content_block)
|
||||
return result or "(empty)"
|
||||
|
||||
@staticmethod
|
||||
@@ -343,7 +358,8 @@ class AnthropicProvider(LLMProvider):
|
||||
@staticmethod
|
||||
def _convert_image_block(block: dict[str, Any]) -> dict[str, Any] | None:
|
||||
"""Convert OpenAI image_url block to Anthropic image block."""
|
||||
url = (block.get("image_url") or {}).get("url", "")
|
||||
image_url = cast(dict[str, Any], block.get("image_url") or {})
|
||||
url = cast(str, image_url.get("url", ""))
|
||||
if not url:
|
||||
return None
|
||||
m = re.match(r"data:(image/\w+);base64,(.+)", url, re.DOTALL)
|
||||
@@ -367,10 +383,13 @@ class AnthropicProvider(LLMProvider):
|
||||
content = msg.get("content")
|
||||
if not isinstance(content, list):
|
||||
return False
|
||||
return any(
|
||||
isinstance(block, dict) and block.get("type") == "tool_use"
|
||||
for block in content
|
||||
)
|
||||
for block in cast(list[object], content):
|
||||
if (
|
||||
isinstance(block, dict)
|
||||
and cast(dict[str, Any], block).get("type") == "tool_use"
|
||||
):
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _merge_consecutive(msgs: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
@@ -402,7 +421,7 @@ class AnthropicProvider(LLMProvider):
|
||||
if isinstance(cur_c, str):
|
||||
cur_c = [{"type": "text", "text": cur_c}]
|
||||
if isinstance(cur_c, list):
|
||||
prev_c.extend(cur_c)
|
||||
cast(list[Any], prev_c).extend(cast(list[Any], cur_c))
|
||||
merged[-1]["content"] = prev_c
|
||||
else:
|
||||
merged.append(msg)
|
||||
@@ -446,7 +465,7 @@ class AnthropicProvider(LLMProvider):
|
||||
def _convert_tools(tools: list[dict[str, Any]] | None) -> list[dict[str, Any]] | None:
|
||||
if not tools:
|
||||
return None
|
||||
result = []
|
||||
result: list[dict[str, Any]] = []
|
||||
for tool in tools:
|
||||
func = tool.get("function", tool)
|
||||
entry: dict[str, Any] = {
|
||||
@@ -506,7 +525,7 @@ class AnthropicProvider(LLMProvider):
|
||||
if isinstance(c, str):
|
||||
new_msgs[-2] = {**m, "content": [{"type": "text", "text": c, "cache_control": marker}]}
|
||||
elif isinstance(c, list) and c:
|
||||
nc = list(c)
|
||||
nc = list(cast(list[dict[str, Any]], c))
|
||||
nc[-1] = {**nc[-1], "cache_control": marker}
|
||||
new_msgs[-2] = {**m, "content": nc}
|
||||
|
||||
@@ -570,7 +589,7 @@ class AnthropicProvider(LLMProvider):
|
||||
kwargs["temperature"] = 1.0
|
||||
elif thinking_enabled:
|
||||
budget_map = {"low": 1024, "medium": 4096, "high": max(8192, max_tokens)}
|
||||
budget = budget_map.get(reasoning_effort.lower(), 4096)
|
||||
budget = budget_map.get(cast(str, reasoning_effort).lower(), 4096)
|
||||
kwargs["thinking"] = {"type": "enabled", "budget_tokens": budget}
|
||||
kwargs["max_tokens"] = max(max_tokens, budget + 4096)
|
||||
if not omit_temperature:
|
||||
@@ -683,7 +702,7 @@ class AnthropicProvider(LLMProvider):
|
||||
reasoning_effort, tool_choice,
|
||||
)
|
||||
try:
|
||||
response = await self._client.messages.create(**kwargs)
|
||||
response = cast(Any, await self._client.messages.create(**kwargs))
|
||||
return self._parse_response(response)
|
||||
except Exception as e:
|
||||
if self._is_streaming_required_error(e):
|
||||
|
||||
@@ -21,7 +21,7 @@ from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
@@ -208,7 +208,7 @@ class AzureOpenAIProvider(LLMProvider):
|
||||
reasoning_effort, tool_choice,
|
||||
)
|
||||
try:
|
||||
response = await self._client.responses.create(**body)
|
||||
response = cast(Any, await self._client.responses.create(**body))
|
||||
return parse_response_output(response)
|
||||
except Exception as e:
|
||||
return self._handle_error(e)
|
||||
@@ -234,7 +234,7 @@ class AzureOpenAIProvider(LLMProvider):
|
||||
body["stream"] = True
|
||||
|
||||
try:
|
||||
stream = await self._client.responses.create(**body)
|
||||
stream = cast(Any, await self._client.responses.create(**body))
|
||||
content, tool_calls, finish_reason, usage, reasoning_content = (
|
||||
await consume_sdk_stream(stream, on_content_delta, on_tool_call_delta)
|
||||
)
|
||||
|
||||
+35
-29
@@ -10,7 +10,7 @@ from contextlib import suppress
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from email.utils import parsedate_to_datetime
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
import json_repair
|
||||
from loguru import logger
|
||||
@@ -67,7 +67,8 @@ class ToolCallRequest:
|
||||
``messages.content.N.tool_use.name: Input should be a valid string``),
|
||||
which permanently wedges the session.
|
||||
"""
|
||||
return isinstance(self.name, str) and bool(self.name)
|
||||
runtime_name = cast(object, self.name)
|
||||
return isinstance(runtime_name, str) and bool(runtime_name)
|
||||
|
||||
def to_openai_tool_call(self) -> dict[str, Any]:
|
||||
"""Serialize to an OpenAI-style tool_call payload."""
|
||||
@@ -76,7 +77,7 @@ class ToolCallRequest:
|
||||
if isinstance(self.arguments, str)
|
||||
else json.dumps(self.arguments, ensure_ascii=False)
|
||||
)
|
||||
tool_call = {
|
||||
tool_call: dict[str, Any] = {
|
||||
"id": self.id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
@@ -126,7 +127,7 @@ def tool_arguments_object_for_replay(arguments: Any) -> dict[str, Any]:
|
||||
if arguments is None:
|
||||
return {}
|
||||
if isinstance(arguments, dict):
|
||||
return arguments
|
||||
return cast(dict[str, Any], arguments)
|
||||
if not isinstance(arguments, str):
|
||||
return {}
|
||||
|
||||
@@ -141,7 +142,7 @@ def tool_arguments_object_for_replay(arguments: Any) -> dict[str, Any]:
|
||||
parsed = json_repair.loads(stripped)
|
||||
except Exception:
|
||||
return {}
|
||||
return parsed if isinstance(parsed, dict) else {}
|
||||
return cast(dict[str, Any], parsed) if isinstance(parsed, dict) else {}
|
||||
|
||||
|
||||
def tool_arguments_json_for_replay(arguments: Any) -> str:
|
||||
@@ -158,7 +159,7 @@ class LLMResponse:
|
||||
usage: dict[str, int] = field(default_factory=dict)
|
||||
retry_after: float | None = None # Provider supplied retry wait in seconds.
|
||||
reasoning_content: str | None = None # Kimi, DeepSeek-R1, MiMo etc.
|
||||
thinking_blocks: list[dict] | None = None # Anthropic extended thinking
|
||||
thinking_blocks: list[dict[str, Any]] | None = None # Anthropic extended thinking
|
||||
# Structured error metadata used by retry policy when finish_reason == "error".
|
||||
error_status_code: int | None = None
|
||||
error_kind: str | None = None # e.g. "timeout", "connection"
|
||||
@@ -298,19 +299,20 @@ class LLMProvider(ABC):
|
||||
if isinstance(content, list):
|
||||
new_items: list[Any] = []
|
||||
changed = False
|
||||
for item in content:
|
||||
for raw_item in cast(list[object], content):
|
||||
item = cast(dict[str, Any], raw_item) if isinstance(raw_item, dict) else None
|
||||
if (
|
||||
isinstance(item, dict)
|
||||
item is not None
|
||||
and item.get("type") in ("text", "input_text", "output_text")
|
||||
and not item.get("text")
|
||||
):
|
||||
changed = True
|
||||
continue
|
||||
if isinstance(item, dict) and "_meta" in item:
|
||||
if item is not None and "_meta" in item:
|
||||
new_items.append({k: v for k, v in item.items() if k != "_meta"})
|
||||
changed = True
|
||||
else:
|
||||
new_items.append(item)
|
||||
new_items.append(raw_item)
|
||||
if changed:
|
||||
clean = dict(msg)
|
||||
if new_items:
|
||||
@@ -332,7 +334,7 @@ class LLMProvider(ABC):
|
||||
# Defense-in-depth: scrub lone UTF-16 surrogates from every string leaf.
|
||||
# This is idempotent and no-op when messages are already clean.
|
||||
sanitized = sanitize_surrogates_deep(result)
|
||||
return sanitized if isinstance(sanitized, list) else result
|
||||
return cast(list[dict[str, Any]], sanitized) if isinstance(sanitized, list) else result
|
||||
|
||||
@staticmethod
|
||||
def _tool_name(tool: dict[str, Any]) -> str:
|
||||
@@ -341,8 +343,9 @@ class LLMProvider(ABC):
|
||||
if isinstance(name, str):
|
||||
return name
|
||||
fn = tool.get("function")
|
||||
if isinstance(fn, dict):
|
||||
fname = fn.get("name")
|
||||
fn_object = cast(dict[str, Any], fn) if isinstance(fn, dict) else None
|
||||
if fn_object is not None:
|
||||
fname = fn_object.get("name")
|
||||
if isinstance(fname, str):
|
||||
return fname
|
||||
return ""
|
||||
@@ -372,7 +375,7 @@ class LLMProvider(ABC):
|
||||
allowed_keys: frozenset[str],
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Keep only provider-safe message keys and normalize assistant content."""
|
||||
sanitized = []
|
||||
sanitized: list[dict[str, Any]] = []
|
||||
for msg in messages:
|
||||
clean = {k: v for k, v in msg.items() if k in allowed_keys}
|
||||
if clean.get("role") == "assistant" and "content" not in clean:
|
||||
@@ -465,7 +468,7 @@ class LLMProvider(ABC):
|
||||
def _extract_error_type_code(cls, payload: Any) -> tuple[str | None, str | None]:
|
||||
data: dict[str, Any] | None = None
|
||||
if isinstance(payload, dict):
|
||||
data = payload
|
||||
data = cast(dict[str, Any], payload)
|
||||
elif isinstance(payload, str):
|
||||
text = payload.strip()
|
||||
if text:
|
||||
@@ -474,16 +477,17 @@ class LLMProvider(ABC):
|
||||
except Exception:
|
||||
parsed = None
|
||||
if isinstance(parsed, dict):
|
||||
data = parsed
|
||||
if not isinstance(data, dict):
|
||||
data = cast(dict[str, Any], parsed)
|
||||
if data is None:
|
||||
return None, None
|
||||
|
||||
error_obj = data.get("error")
|
||||
type_value = data.get("type")
|
||||
code_value = data.get("code")
|
||||
if isinstance(error_obj, dict):
|
||||
type_value = error_obj.get("type") or type_value
|
||||
code_value = error_obj.get("code") or code_value
|
||||
error_object = cast(dict[str, Any], error_obj) if isinstance(error_obj, dict) else None
|
||||
if error_object is not None:
|
||||
type_value = error_object.get("type") or type_value
|
||||
code_value = error_object.get("code") or code_value
|
||||
|
||||
return cls._normalize_error_token(type_value), cls._normalize_error_token(code_value)
|
||||
|
||||
@@ -582,13 +586,14 @@ class LLMProvider(ABC):
|
||||
def _strip_image_content(messages: list[dict[str, Any]]) -> list[dict[str, Any]] | None:
|
||||
"""Replace image_url blocks with text placeholder. Returns None if no images found."""
|
||||
found = False
|
||||
result = []
|
||||
result: list[dict[str, Any]] = []
|
||||
for msg in messages:
|
||||
content = msg.get("content")
|
||||
if isinstance(content, list):
|
||||
new_content = []
|
||||
for b in content:
|
||||
if isinstance(b, dict) and b.get("type") == "image_url":
|
||||
new_content: list[Any] = []
|
||||
for raw_block in cast(list[object], content):
|
||||
block = cast(dict[str, Any], raw_block) if isinstance(raw_block, dict) else None
|
||||
if block is not None and block.get("type") == "image_url":
|
||||
placeholder = (
|
||||
"[Image not delivered to model — "
|
||||
"do not describe or reference it]"
|
||||
@@ -596,7 +601,7 @@ class LLMProvider(ABC):
|
||||
new_content.append({"type": "text", "text": placeholder})
|
||||
found = True
|
||||
else:
|
||||
new_content.append(b)
|
||||
new_content.append(raw_block)
|
||||
result.append({**msg, "content": new_content})
|
||||
else:
|
||||
result.append(msg)
|
||||
@@ -614,8 +619,9 @@ class LLMProvider(ABC):
|
||||
for msg in messages:
|
||||
content = msg.get("content")
|
||||
if isinstance(content, list):
|
||||
for i, b in enumerate(content):
|
||||
if isinstance(b, dict) and b.get("type") == "image_url":
|
||||
for i, raw_block in enumerate(cast(list[object], content)):
|
||||
block = cast(dict[str, Any], raw_block) if isinstance(raw_block, dict) else None
|
||||
if block is not None and block.get("type") == "image_url":
|
||||
placeholder = (
|
||||
"[Image not delivered to model — "
|
||||
"do not describe or reference it]"
|
||||
@@ -815,7 +821,7 @@ class LLMProvider(ABC):
|
||||
if value is not None:
|
||||
return value
|
||||
if isinstance(headers, dict):
|
||||
for key, value in headers.items():
|
||||
for key, value in cast(dict[object, Any], headers).items():
|
||||
if isinstance(key, str) and key.lower() == name.lower():
|
||||
return value
|
||||
return None
|
||||
@@ -986,7 +992,7 @@ class LLMProvider(ABC):
|
||||
on_retry_wait=on_retry_wait,
|
||||
)
|
||||
|
||||
return last_response if last_response is not None else await call(**kw)
|
||||
return last_response if last_response is not None else await call(**kw) # pyright: ignore[reportUnnecessaryComparison]
|
||||
|
||||
@abstractmethod
|
||||
def get_default_model(self) -> str:
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportMissingTypeStubs=false
|
||||
"""AWS Bedrock Converse provider."""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -8,7 +9,7 @@ import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Awaitable, Callable, Iterator
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
from nanobot.providers.base import (
|
||||
LLMProvider,
|
||||
@@ -30,7 +31,10 @@ def _deep_merge(base: dict[str, Any], override: dict[str, Any]) -> dict[str, Any
|
||||
merged = dict(base)
|
||||
for key, value in override.items():
|
||||
if key in merged and isinstance(merged[key], dict) and isinstance(value, dict):
|
||||
merged[key] = _deep_merge(merged[key], value)
|
||||
merged[key] = _deep_merge(
|
||||
cast(dict[str, Any], merged[key]),
|
||||
cast(dict[str, Any], value),
|
||||
)
|
||||
else:
|
||||
merged[key] = value
|
||||
return merged
|
||||
@@ -77,7 +81,8 @@ class BedrockProvider(LLMProvider):
|
||||
session_kwargs: dict[str, Any] = {}
|
||||
if self.profile:
|
||||
session_kwargs["profile_name"] = self.profile
|
||||
session = boto3.Session(**session_kwargs)
|
||||
boto3_module = cast(Any, boto3)
|
||||
session = boto3_module.Session(**session_kwargs)
|
||||
|
||||
client_kwargs: dict[str, Any] = {}
|
||||
if self.region:
|
||||
@@ -107,7 +112,8 @@ class BedrockProvider(LLMProvider):
|
||||
|
||||
@staticmethod
|
||||
def _image_url_block(block: dict[str, Any]) -> dict[str, Any] | None:
|
||||
url = (block.get("image_url") or {}).get("url", "")
|
||||
image_url = cast(dict[str, Any], block.get("image_url") or {})
|
||||
url = image_url.get("url", "")
|
||||
if not isinstance(url, str) or not url:
|
||||
return None
|
||||
match = _IMAGE_DATA_URL.match(url)
|
||||
@@ -132,10 +138,11 @@ class BedrockProvider(LLMProvider):
|
||||
return [{"text": str(content)}]
|
||||
|
||||
blocks: list[dict[str, Any]] = []
|
||||
for item in content:
|
||||
if not isinstance(item, dict):
|
||||
blocks.append({"text": str(item)})
|
||||
for raw_item in cast(list[object], content):
|
||||
if not isinstance(raw_item, dict):
|
||||
blocks.append({"text": str(raw_item)})
|
||||
continue
|
||||
item = cast(dict[str, Any], raw_item)
|
||||
|
||||
item_type = item.get("type")
|
||||
if item_type in _TEXT_BLOCK_TYPES or "text" in item:
|
||||
@@ -181,6 +188,7 @@ class BedrockProvider(LLMProvider):
|
||||
function = tool_call.get("function")
|
||||
if not isinstance(function, dict):
|
||||
return None
|
||||
function = cast(dict[str, Any], function)
|
||||
args = tool_arguments_object_for_replay(function.get("arguments", {}))
|
||||
return {
|
||||
"toolUse": {
|
||||
@@ -216,8 +224,10 @@ class BedrockProvider(LLMProvider):
|
||||
def _assistant_blocks(cls, msg: dict[str, Any]) -> list[dict[str, Any]]:
|
||||
blocks: list[dict[str, Any]] = []
|
||||
|
||||
for thinking in msg.get("thinking_blocks") or []:
|
||||
if isinstance(thinking, dict):
|
||||
thinking_values = cast(list[object], msg.get("thinking_blocks") or [])
|
||||
for thinking_value in thinking_values:
|
||||
if isinstance(thinking_value, dict):
|
||||
thinking = cast(dict[str, Any], thinking_value)
|
||||
reasoning = cls._reasoning_block(thinking)
|
||||
if reasoning:
|
||||
blocks.append(reasoning)
|
||||
@@ -228,8 +238,10 @@ class BedrockProvider(LLMProvider):
|
||||
elif isinstance(content, list):
|
||||
blocks.extend(block for block in cls._content_blocks(content) if "text" in block)
|
||||
|
||||
for tool_call in msg.get("tool_calls") or []:
|
||||
if isinstance(tool_call, dict):
|
||||
tool_call_values = cast(list[object], msg.get("tool_calls") or [])
|
||||
for tool_call_value in tool_call_values:
|
||||
if isinstance(tool_call_value, dict):
|
||||
tool_call = cast(dict[str, Any], tool_call_value)
|
||||
block = cls._tool_use_block(tool_call)
|
||||
if block:
|
||||
blocks.append(block)
|
||||
@@ -240,7 +252,8 @@ class BedrockProvider(LLMProvider):
|
||||
def _has_tool_use(msg: dict[str, Any]) -> bool:
|
||||
content = msg.get("content")
|
||||
return isinstance(content, list) and any(
|
||||
isinstance(block, dict) and "toolUse" in block for block in content
|
||||
isinstance(block, dict) and "toolUse" in block
|
||||
for block in cast(list[object], content)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
@@ -249,12 +262,14 @@ class BedrockProvider(LLMProvider):
|
||||
for msg in messages:
|
||||
if merged and merged[-1].get("role") == msg.get("role"):
|
||||
prev = merged[-1].setdefault("content", [])
|
||||
cur = msg.get("content") or []
|
||||
cur: Any = msg.get("content") or []
|
||||
if not isinstance(prev, list):
|
||||
prev = [{"text": str(prev)}]
|
||||
merged[-1]["content"] = prev
|
||||
else:
|
||||
prev = cast(list[Any], prev)
|
||||
if isinstance(cur, list):
|
||||
prev.extend(cur)
|
||||
prev.extend(cast(list[Any], cur))
|
||||
else:
|
||||
prev.append({"text": str(cur)})
|
||||
else:
|
||||
@@ -303,9 +318,12 @@ class BedrockProvider(LLMProvider):
|
||||
return None
|
||||
result: list[dict[str, Any]] = []
|
||||
for tool in tools:
|
||||
func = tool.get("function") if isinstance(tool.get("function"), dict) else tool
|
||||
if not isinstance(func, dict):
|
||||
continue
|
||||
function_value = tool.get("function")
|
||||
func = (
|
||||
cast(dict[str, Any], function_value)
|
||||
if isinstance(function_value, dict)
|
||||
else tool
|
||||
)
|
||||
name = str(func.get("name") or "")
|
||||
if not name:
|
||||
continue
|
||||
@@ -330,9 +348,11 @@ class BedrockProvider(LLMProvider):
|
||||
content = msg.get("content")
|
||||
if not isinstance(content, list):
|
||||
continue
|
||||
for block in content:
|
||||
if isinstance(block, dict) and ("toolUse" in block or "toolResult" in block):
|
||||
return True
|
||||
for block_value in cast(list[object], content):
|
||||
if isinstance(block_value, dict):
|
||||
block = cast(dict[str, Any], block_value)
|
||||
if "toolUse" in block or "toolResult" in block:
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
@@ -356,7 +376,8 @@ class BedrockProvider(LLMProvider):
|
||||
if tool_choice == "none":
|
||||
return None
|
||||
if isinstance(tool_choice, dict):
|
||||
name = tool_choice.get("function", {}).get("name")
|
||||
function = cast(dict[str, Any], tool_choice.get("function", {}))
|
||||
name = function.get("name")
|
||||
if name:
|
||||
return {"tool": {"name": str(name)}}
|
||||
return {"auto": {}}
|
||||
@@ -457,8 +478,10 @@ class BedrockProvider(LLMProvider):
|
||||
reasoning = block.get("reasoningContent")
|
||||
if not isinstance(reasoning, dict):
|
||||
return None, None
|
||||
reasoning = cast(dict[str, Any], reasoning)
|
||||
text_obj = reasoning.get("reasoningText")
|
||||
if isinstance(text_obj, dict):
|
||||
text_obj = cast(dict[str, Any], text_obj)
|
||||
text = text_obj.get("text")
|
||||
if isinstance(text, str):
|
||||
return text, {
|
||||
@@ -480,15 +503,19 @@ class BedrockProvider(LLMProvider):
|
||||
reasoning_parts: list[str] = []
|
||||
tool_calls: list[ToolCallRequest] = []
|
||||
thinking_blocks: list[dict[str, Any]] = []
|
||||
message = (response.get("output") or {}).get("message") or {}
|
||||
output = cast(dict[str, Any], response.get("output") or {})
|
||||
message = cast(dict[str, Any], output.get("message") or {})
|
||||
|
||||
for block in message.get("content") or []:
|
||||
if not isinstance(block, dict):
|
||||
content_blocks = cast(list[object], message.get("content") or [])
|
||||
for block_value in content_blocks:
|
||||
if not isinstance(block_value, dict):
|
||||
continue
|
||||
block = cast(dict[str, Any], block_value)
|
||||
if isinstance(block.get("text"), str):
|
||||
content_parts.append(block["text"])
|
||||
content_parts.append(cast(str, block["text"]))
|
||||
tool_use = block.get("toolUse")
|
||||
if isinstance(tool_use, dict):
|
||||
tool_use = cast(dict[str, Any], tool_use)
|
||||
arguments = tool_use.get("input", {})
|
||||
tool_calls.append(ToolCallRequest(
|
||||
id=str(tool_use.get("toolUseId") or ""),
|
||||
@@ -504,8 +531,8 @@ class BedrockProvider(LLMProvider):
|
||||
return LLMResponse(
|
||||
content="".join(content_parts) or None,
|
||||
tool_calls=tool_calls,
|
||||
finish_reason=cls._finish_reason(response.get("stopReason")),
|
||||
usage=cls._usage(response.get("usage")),
|
||||
finish_reason=cls._finish_reason(cast(str | None, response.get("stopReason"))),
|
||||
usage=cls._usage(cast(dict[str, Any] | None, response.get("usage"))),
|
||||
reasoning_content="".join(reasoning_parts) or None,
|
||||
thinking_blocks=thinking_blocks or None,
|
||||
)
|
||||
@@ -522,11 +549,12 @@ class BedrockProvider(LLMProvider):
|
||||
state: dict[str, Any],
|
||||
) -> str | None:
|
||||
if "contentBlockStart" in event:
|
||||
data = event["contentBlockStart"]
|
||||
data = cast(dict[str, Any], event["contentBlockStart"])
|
||||
idx = int(data.get("contentBlockIndex") or 0)
|
||||
start = data.get("start") or {}
|
||||
start = cast(dict[str, Any], data.get("start") or {})
|
||||
tool_use = start.get("toolUse")
|
||||
if isinstance(tool_use, dict):
|
||||
tool_use = cast(dict[str, Any], tool_use)
|
||||
tool_buffers[idx] = {
|
||||
"id": str(tool_use.get("toolUseId") or ""),
|
||||
"name": str(tool_use.get("name") or ""),
|
||||
@@ -535,21 +563,27 @@ class BedrockProvider(LLMProvider):
|
||||
return None
|
||||
|
||||
if "contentBlockDelta" in event:
|
||||
data = event["contentBlockDelta"]
|
||||
data = cast(dict[str, Any], event["contentBlockDelta"])
|
||||
idx = int(data.get("contentBlockIndex") or 0)
|
||||
delta = data.get("delta") or {}
|
||||
delta = cast(dict[str, Any], data.get("delta") or {})
|
||||
text = delta.get("text")
|
||||
if isinstance(text, str):
|
||||
content_parts.append(text)
|
||||
return text
|
||||
tool_delta = delta.get("toolUse")
|
||||
if isinstance(tool_delta, dict):
|
||||
tool_delta = cast(dict[str, Any], tool_delta)
|
||||
buf = tool_buffers.setdefault(idx, {"id": "", "name": "", "input": ""})
|
||||
if isinstance(tool_delta.get("input"), str):
|
||||
buf["input"] += tool_delta["input"]
|
||||
reasoning = delta.get("reasoningContent")
|
||||
if isinstance(reasoning, dict):
|
||||
buf = state.setdefault("reasoning_buffers", {}).setdefault(
|
||||
reasoning = cast(dict[str, Any], reasoning)
|
||||
reasoning_buffers = cast(
|
||||
dict[int, dict[str, Any]],
|
||||
state.setdefault("reasoning_buffers", {}),
|
||||
)
|
||||
buf = reasoning_buffers.setdefault(
|
||||
idx, {"text": "", "signature": "", "redactedContent": None}
|
||||
)
|
||||
if isinstance(reasoning.get("text"), str):
|
||||
@@ -562,8 +596,13 @@ class BedrockProvider(LLMProvider):
|
||||
return None
|
||||
|
||||
if "contentBlockStop" in event:
|
||||
idx = int((event["contentBlockStop"] or {}).get("contentBlockIndex") or 0)
|
||||
reasoning_buf = state.setdefault("reasoning_buffers", {}).pop(idx, None)
|
||||
stop = cast(dict[str, Any], event["contentBlockStop"] or {})
|
||||
idx = int(stop.get("contentBlockIndex") or 0)
|
||||
reasoning_buffers = cast(
|
||||
dict[int, dict[str, Any]],
|
||||
state.setdefault("reasoning_buffers", {}),
|
||||
)
|
||||
reasoning_buf = reasoning_buffers.pop(idx, None)
|
||||
if reasoning_buf:
|
||||
if reasoning_buf.get("text"):
|
||||
thinking_blocks.append({
|
||||
@@ -589,11 +628,12 @@ class BedrockProvider(LLMProvider):
|
||||
return None
|
||||
|
||||
if "messageStop" in event:
|
||||
state["stop_reason"] = (event["messageStop"] or {}).get("stopReason")
|
||||
message_stop = cast(dict[str, Any], event["messageStop"] or {})
|
||||
state["stop_reason"] = message_stop.get("stopReason")
|
||||
return None
|
||||
|
||||
if "metadata" in event:
|
||||
metadata = event["metadata"] or {}
|
||||
metadata = cast(dict[str, Any], event["metadata"] or {})
|
||||
if isinstance(metadata.get("usage"), dict):
|
||||
state["usage"] = metadata["usage"]
|
||||
return None
|
||||
@@ -631,14 +671,29 @@ class BedrockProvider(LLMProvider):
|
||||
|
||||
@classmethod
|
||||
def _handle_error(cls, e: Exception) -> LLMResponse:
|
||||
response = getattr(e, "response", None)
|
||||
metadata = response.get("ResponseMetadata", {}) if isinstance(response, dict) else {}
|
||||
headers = metadata.get("HTTPHeaders") if isinstance(metadata, dict) else None
|
||||
error_obj = response.get("Error", {}) if isinstance(response, dict) else {}
|
||||
message = error_obj.get("Message") if isinstance(error_obj, dict) else None
|
||||
code = error_obj.get("Code") if isinstance(error_obj, dict) else None
|
||||
status_code = metadata.get("HTTPStatusCode") if isinstance(metadata, dict) else None
|
||||
body = message or str(e)
|
||||
response_value = getattr(e, "response", None)
|
||||
response = (
|
||||
cast(dict[str, Any], response_value)
|
||||
if isinstance(response_value, dict)
|
||||
else {}
|
||||
)
|
||||
metadata_value = response.get("ResponseMetadata", {})
|
||||
metadata = (
|
||||
cast(dict[str, Any], metadata_value)
|
||||
if isinstance(metadata_value, dict)
|
||||
else {}
|
||||
)
|
||||
headers = metadata.get("HTTPHeaders")
|
||||
error_value = response.get("Error", {})
|
||||
error_obj = (
|
||||
cast(dict[str, Any], error_value)
|
||||
if isinstance(error_value, dict)
|
||||
else {}
|
||||
)
|
||||
message = error_obj.get("Message")
|
||||
code = error_obj.get("Code")
|
||||
status_code = metadata.get("HTTPStatusCode")
|
||||
body = cast(str, message or str(e))
|
||||
retry_after = cls._extract_retry_after_from_headers(headers)
|
||||
if retry_after is None:
|
||||
retry_after = cls._extract_retry_after(body)
|
||||
@@ -683,7 +738,10 @@ class BedrockProvider(LLMProvider):
|
||||
kwargs = self._build_kwargs(
|
||||
messages, tools, model, max_tokens, temperature, reasoning_effort, tool_choice
|
||||
)
|
||||
response = await asyncio.to_thread(self._client.converse, **kwargs)
|
||||
response = cast(
|
||||
dict[str, Any],
|
||||
await asyncio.to_thread(self._client.converse, **kwargs),
|
||||
)
|
||||
return self._parse_response(response)
|
||||
except Exception as e:
|
||||
return self._handle_error(e)
|
||||
@@ -713,8 +771,11 @@ class BedrockProvider(LLMProvider):
|
||||
kwargs = self._build_kwargs(
|
||||
messages, tools, model, max_tokens, temperature, reasoning_effort, tool_choice
|
||||
)
|
||||
response = await asyncio.to_thread(self._client.converse_stream, **kwargs)
|
||||
stream = iter(response.get("stream") or [])
|
||||
response = cast(
|
||||
dict[str, Any],
|
||||
await asyncio.to_thread(self._client.converse_stream, **kwargs),
|
||||
)
|
||||
stream = cast(Iterator[dict[str, Any]], iter(response.get("stream") or []))
|
||||
while True:
|
||||
event = await asyncio.wait_for(
|
||||
asyncio.to_thread(_next_or_none, stream),
|
||||
|
||||
@@ -160,6 +160,8 @@ def _make_provider_core(
|
||||
elif backend == "azure_openai":
|
||||
from nanobot.providers.azure_openai_provider import AzureOpenAIProvider
|
||||
|
||||
if p is None or p.api_base is None:
|
||||
raise RuntimeError("validated Azure provider setup is missing api_base")
|
||||
provider = AzureOpenAIProvider(
|
||||
api_key=p.api_key or "",
|
||||
api_base=p.api_base,
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Provider wrapper that transparently fails over to fallback models on error."""
|
||||
|
||||
# pyright: reportIncompatibleMethodOverride=false, reportIncompatibleVariableOverride=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
@@ -8,7 +10,7 @@ from typing import Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.providers.base import LLMProvider, LLMResponse
|
||||
from nanobot.providers.base import GenerationSettings, LLMProvider, LLMResponse
|
||||
|
||||
# Circuit breaker tuned to match OpenAICompatProvider's Responses API breaker.
|
||||
_PRIMARY_FAILURE_THRESHOLD = 3
|
||||
@@ -121,11 +123,11 @@ class FallbackProvider(LLMProvider):
|
||||
self._primary_tripped_at: float | None = None
|
||||
|
||||
@property
|
||||
def generation(self):
|
||||
def generation(self) -> GenerationSettings:
|
||||
return self._primary.generation
|
||||
|
||||
@generation.setter
|
||||
def generation(self, value):
|
||||
def generation(self, value: GenerationSettings) -> None:
|
||||
self._primary.generation = value
|
||||
|
||||
def get_default_model(self) -> str:
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""GitHub Copilot OAuth-backed provider."""
|
||||
|
||||
# pyright: reportMissingTypeStubs=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
@@ -8,11 +10,13 @@ import time
|
||||
import webbrowser
|
||||
from collections.abc import Awaitable, Callable
|
||||
from contextlib import suppress
|
||||
from typing import Any, cast
|
||||
|
||||
import httpx
|
||||
from oauth_cli_kit.models import OAuthToken
|
||||
from oauth_cli_kit.storage import FileTokenStorage
|
||||
|
||||
from nanobot.providers.base import LLMResponse
|
||||
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
|
||||
|
||||
DEFAULT_GITHUB_DEVICE_CODE_URL = "https://github.com/login/device/code"
|
||||
@@ -232,19 +236,19 @@ class GitHubCopilotProvider(OpenAICompatProvider):
|
||||
token = await self._get_copilot_access_token()
|
||||
client = await self._ensure_client()
|
||||
self.api_key = token
|
||||
client.api_key = token
|
||||
cast(Any, client).api_key = token
|
||||
return token
|
||||
|
||||
async def chat(
|
||||
self,
|
||||
messages: list[dict[str, object]],
|
||||
tools: list[dict[str, object]] | None = None,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]] | None = None,
|
||||
model: str | None = None,
|
||||
max_tokens: int = 4096,
|
||||
temperature: float = 0.7,
|
||||
reasoning_effort: str | None = None,
|
||||
tool_choice: str | dict[str, object] | None = None,
|
||||
):
|
||||
tool_choice: str | dict[str, Any] | None = None,
|
||||
) -> LLMResponse:
|
||||
await self._refresh_client_api_key()
|
||||
return await super().chat(
|
||||
messages=messages,
|
||||
@@ -258,17 +262,17 @@ class GitHubCopilotProvider(OpenAICompatProvider):
|
||||
|
||||
async def chat_stream(
|
||||
self,
|
||||
messages: list[dict[str, object]],
|
||||
tools: list[dict[str, object]] | None = None,
|
||||
messages: list[dict[str, Any]],
|
||||
tools: list[dict[str, Any]] | None = None,
|
||||
model: str | None = None,
|
||||
max_tokens: int = 4096,
|
||||
temperature: float = 0.7,
|
||||
reasoning_effort: str | None = None,
|
||||
tool_choice: str | dict[str, object] | None = None,
|
||||
on_content_delta: Callable[[str], None] | None = None,
|
||||
tool_choice: str | dict[str, Any] | None = None,
|
||||
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||
on_tool_call_delta: Callable[[dict[str, object]], Awaitable[None]] | None = None,
|
||||
):
|
||||
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||
) -> LLMResponse:
|
||||
await self._refresh_client_api_key()
|
||||
return await super().chat_stream(
|
||||
messages=messages,
|
||||
|
||||
@@ -9,12 +9,13 @@ import re
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
from urllib.parse import urljoin
|
||||
|
||||
import httpx
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.config.schema import Config, ProviderConfig
|
||||
from nanobot.providers.registry import find_by_name
|
||||
from nanobot.security.network import (
|
||||
PinnedDNSAsyncTransport,
|
||||
@@ -81,6 +82,18 @@ class GeneratedImageResponse:
|
||||
raw: dict[str, Any]
|
||||
|
||||
|
||||
def _as_json_object(value: object) -> dict[str, Any] | None:
|
||||
"""Narrow an untrusted provider response value to a JSON object."""
|
||||
return cast(dict[str, Any], value) if isinstance(value, dict) else None
|
||||
|
||||
|
||||
def _as_json_objects(value: object) -> list[dict[str, Any]]:
|
||||
"""Return object entries from an untrusted provider response array."""
|
||||
if not isinstance(value, list):
|
||||
return []
|
||||
return [cast(dict[str, Any], item) for item in cast(list[object], value) if isinstance(item, dict)]
|
||||
|
||||
|
||||
def _read_image_b64(path: str | Path) -> tuple[str, str]:
|
||||
"""Return ``(mime, base64)`` for the image at ``path``."""
|
||||
p = Path(path).expanduser()
|
||||
@@ -249,7 +262,7 @@ def image_gen_provider_names() -> tuple[str, ...]:
|
||||
return tuple(_IMAGE_GEN_PROVIDERS)
|
||||
|
||||
|
||||
def image_gen_provider_configs(config: Any) -> dict[str, Any]:
|
||||
def image_gen_provider_configs(config: Config) -> dict[str, ProviderConfig]:
|
||||
providers_cfg = config.providers
|
||||
return {
|
||||
name: pc
|
||||
@@ -315,7 +328,7 @@ class ImageGenerationProvider(ABC):
|
||||
def _require_images(self, images: list[str], data: dict[str, Any]) -> None:
|
||||
if images:
|
||||
return
|
||||
provider_error = data.get("error") if isinstance(data, dict) else None
|
||||
provider_error = data.get("error")
|
||||
label = self.provider_name
|
||||
if provider_error:
|
||||
raise ImageGenerationError(f"{label} returned no images: {provider_error}")
|
||||
@@ -410,20 +423,17 @@ class OpenRouterImageGenerationClient(ImageGenerationProvider):
|
||||
detail = response.text[:500]
|
||||
raise ImageGenerationError(f"OpenRouter image generation failed: {detail}") from exc
|
||||
|
||||
data = response.json()
|
||||
data = _as_json_object(response.json()) or {}
|
||||
images: list[str] = []
|
||||
text_parts: list[str] = []
|
||||
for choice in data.get("choices") or []:
|
||||
if not isinstance(choice, dict):
|
||||
continue
|
||||
message = choice.get("message") or {}
|
||||
if isinstance(message.get("content"), str):
|
||||
text_parts.append(message["content"])
|
||||
for image in message.get("images") or []:
|
||||
if not isinstance(image, dict):
|
||||
continue
|
||||
image_url = image.get("image_url") or image.get("imageUrl") or {}
|
||||
url_value = image_url.get("url") if isinstance(image_url, dict) else None
|
||||
for choice in _as_json_objects(data.get("choices")):
|
||||
message = _as_json_object(choice.get("message")) or {}
|
||||
message_content = message.get("content")
|
||||
if isinstance(message_content, str):
|
||||
text_parts.append(message_content)
|
||||
for image in _as_json_objects(message.get("images")):
|
||||
image_url = _as_json_object(image.get("image_url") or image.get("imageUrl"))
|
||||
url_value = image_url.get("url") if image_url is not None else None
|
||||
if isinstance(url_value, str) and url_value.startswith("data:image/"):
|
||||
images.append(url_value)
|
||||
|
||||
@@ -527,7 +537,7 @@ class AIHubMixImageGenerationClient(ImageGenerationProvider):
|
||||
detail = response.text[:500]
|
||||
raise ImageGenerationError(f"AIHubMix image generation failed: {detail}") from exc
|
||||
|
||||
payload = response.json()
|
||||
payload = _as_json_object(response.json()) or {}
|
||||
images = await _aihubmix_images_from_payload(payload, proxy=self.proxy)
|
||||
|
||||
self._require_images(images, payload)
|
||||
@@ -538,11 +548,12 @@ class AIHubMixImageGenerationClient(ImageGenerationProvider):
|
||||
def _http_error_detail(response: httpx.Response) -> str:
|
||||
"""Extract a readable error message from an HTTP error response."""
|
||||
try:
|
||||
data = response.json()
|
||||
if isinstance(data, dict):
|
||||
err = data.get("error")
|
||||
if isinstance(err, dict):
|
||||
return err.get("message") or str(err)
|
||||
data = _as_json_object(response.json())
|
||||
if data is not None:
|
||||
err = _as_json_object(data.get("error"))
|
||||
if err is not None:
|
||||
message = err.get("message")
|
||||
return message if isinstance(message, str) else str(err)
|
||||
if err:
|
||||
return str(err)
|
||||
except Exception:
|
||||
@@ -595,11 +606,11 @@ def _ollama_image_data_url(value: str) -> str:
|
||||
def _ollama_images_from_payload(payload: dict[str, Any]) -> list[str]:
|
||||
images: list[str] = []
|
||||
|
||||
def collect(value: Any) -> None:
|
||||
def collect(value: object) -> None:
|
||||
if isinstance(value, str) and value:
|
||||
images.append(_ollama_image_data_url(value))
|
||||
elif isinstance(value, list):
|
||||
for item in value:
|
||||
for item in cast(list[object], value):
|
||||
collect(item)
|
||||
|
||||
collect(payload.get("image"))
|
||||
@@ -768,14 +779,12 @@ class GeminiImageGenerationClient(ImageGenerationProvider):
|
||||
f"Gemini Imagen generation failed (HTTP {response.status_code}): {detail}"
|
||||
) from exc
|
||||
|
||||
data = response.json()
|
||||
data = _as_json_object(response.json()) or {}
|
||||
images: list[str] = []
|
||||
for prediction in data.get("predictions") or []:
|
||||
if not isinstance(prediction, dict):
|
||||
continue
|
||||
for prediction in _as_json_objects(data.get("predictions")):
|
||||
b64 = prediction.get("bytesBase64Encoded")
|
||||
mime = prediction.get("mimeType", "image/png")
|
||||
if isinstance(b64, str) and b64:
|
||||
if isinstance(b64, str) and b64 and isinstance(mime, str):
|
||||
images.append(f"data:{mime};base64,{b64}")
|
||||
|
||||
self._require_images(images, data)
|
||||
@@ -824,23 +833,21 @@ class GeminiImageGenerationClient(ImageGenerationProvider):
|
||||
f"Gemini image generation failed (HTTP {response.status_code}): {detail}"
|
||||
) from exc
|
||||
|
||||
data = response.json()
|
||||
data = _as_json_object(response.json()) or {}
|
||||
images: list[str] = []
|
||||
text_parts: list[str] = []
|
||||
for candidate in data.get("candidates") or []:
|
||||
if not isinstance(candidate, dict):
|
||||
continue
|
||||
content = candidate.get("content") or {}
|
||||
for part in content.get("parts") or []:
|
||||
if not isinstance(part, dict):
|
||||
continue
|
||||
for candidate in _as_json_objects(data.get("candidates")):
|
||||
content = _as_json_object(candidate.get("content")) or {}
|
||||
for part in _as_json_objects(content.get("parts")):
|
||||
if "text" in part:
|
||||
text_parts.append(part["text"])
|
||||
inline = part.get("inlineData")
|
||||
if isinstance(inline, dict):
|
||||
text = part["text"]
|
||||
if isinstance(text, str):
|
||||
text_parts.append(text)
|
||||
inline = _as_json_object(part.get("inlineData"))
|
||||
if inline is not None:
|
||||
mime = inline.get("mimeType", "image/png")
|
||||
b64 = inline.get("data", "")
|
||||
if b64:
|
||||
if isinstance(mime, str) and isinstance(b64, str) and b64:
|
||||
images.append(f"data:{mime};base64,{b64}")
|
||||
|
||||
self._require_images(images, data)
|
||||
@@ -914,9 +921,9 @@ async def _aihubmix_images_from_payload(
|
||||
if "output" in payload:
|
||||
candidates.append(payload["output"])
|
||||
|
||||
async def collect(value: Any) -> None:
|
||||
async def collect(value: object) -> None:
|
||||
if isinstance(value, list):
|
||||
for item in value:
|
||||
for item in cast(list[object], value):
|
||||
await collect(item)
|
||||
return
|
||||
if isinstance(value, str):
|
||||
@@ -925,32 +932,38 @@ async def _aihubmix_images_from_payload(
|
||||
elif value.startswith(("http://", "https://")):
|
||||
images.append(await _download_image_data_url(value, proxy=proxy))
|
||||
return
|
||||
if not isinstance(value, dict):
|
||||
value_object = _as_json_object(value)
|
||||
if value_object is None:
|
||||
return
|
||||
|
||||
b64_json = value.get("b64_json")
|
||||
b64_json = value_object.get("b64_json")
|
||||
if isinstance(b64_json, str) and b64_json:
|
||||
images.append(_b64_image_data_url(b64_json))
|
||||
elif b64_json is not None:
|
||||
await collect(b64_json)
|
||||
|
||||
bytes_base64 = value.get("bytesBase64") or value.get("bytes_base64") or value.get("base64")
|
||||
bytes_base64 = (
|
||||
value_object.get("bytesBase64")
|
||||
or value_object.get("bytes_base64")
|
||||
or value_object.get("base64")
|
||||
)
|
||||
if isinstance(bytes_base64, str) and bytes_base64:
|
||||
images.append(_b64_image_data_url(bytes_base64))
|
||||
|
||||
image_url = value.get("image_url") or value.get("imageUrl")
|
||||
if isinstance(image_url, dict):
|
||||
await collect(image_url.get("url"))
|
||||
image_url = value_object.get("image_url") or value_object.get("imageUrl")
|
||||
image_url_object = _as_json_object(image_url)
|
||||
if image_url_object is not None:
|
||||
await collect(image_url_object.get("url"))
|
||||
elif image_url is not None:
|
||||
await collect(image_url)
|
||||
|
||||
url_value = value.get("url")
|
||||
url_value = value_object.get("url")
|
||||
if url_value is not None:
|
||||
await collect(url_value)
|
||||
|
||||
for key in ("images", "image", "output"):
|
||||
if key in value:
|
||||
await collect(value[key])
|
||||
if key in value_object:
|
||||
await collect(value_object[key])
|
||||
|
||||
for candidate in candidates:
|
||||
await collect(candidate)
|
||||
@@ -1061,9 +1074,10 @@ def _minimax_images_from_payload(payload: dict[str, Any]) -> list[str]:
|
||||
"""
|
||||
images: list[str] = []
|
||||
data = payload.get("data")
|
||||
if not isinstance(data, dict):
|
||||
data_object = _as_json_object(data)
|
||||
if data_object is None:
|
||||
return images
|
||||
for b64 in data.get("image_base64") or []:
|
||||
for b64 in cast(list[object], data_object.get("image_base64") or []):
|
||||
if isinstance(b64, str) and b64:
|
||||
images.append(_b64_image_data_url(b64))
|
||||
return images
|
||||
@@ -1381,11 +1395,14 @@ class CodexImageGenerationClient(ImageGenerationProvider):
|
||||
image_size: str | None = None,
|
||||
) -> GeneratedImageResponse:
|
||||
try:
|
||||
from oauth_cli_kit import get_token as get_codex_token
|
||||
from oauth_cli_kit import ( # pyright: ignore[reportMissingTypeStubs]
|
||||
get_token as _get_codex_token,
|
||||
)
|
||||
except ImportError:
|
||||
raise ImageGenerationError(self.missing_key_message)
|
||||
|
||||
try:
|
||||
get_codex_token = cast(Any, _get_codex_token)
|
||||
token_kwargs = {"proxy": self.proxy} if self.proxy else {}
|
||||
token = await asyncio.to_thread(get_codex_token, **token_kwargs)
|
||||
except Exception as exc:
|
||||
@@ -1405,9 +1422,9 @@ class CodexImageGenerationClient(ImageGenerationProvider):
|
||||
len(reference_images),
|
||||
)
|
||||
|
||||
headers = {
|
||||
headers: dict[str, str] = {
|
||||
"Authorization": f"Bearer {token.access}",
|
||||
"chatgpt-account-id": token.account_id,
|
||||
"chatgpt-account-id": str(token.account_id),
|
||||
"OpenAI-Beta": "responses=experimental",
|
||||
"originator": "nanobot",
|
||||
"User-Agent": "nanobot (python)",
|
||||
@@ -1537,9 +1554,7 @@ async def _openai_images_from_payload(
|
||||
Handles both ``b64_json`` (preferred) and ``url`` (downloaded) formats.
|
||||
"""
|
||||
images: list[str] = []
|
||||
for item in payload.get("data") or []:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
for item in _as_json_objects(payload.get("data")):
|
||||
b64 = item.get("b64_json")
|
||||
if isinstance(b64, str) and b64:
|
||||
images.append(_b64_image_data_url(b64))
|
||||
@@ -1567,7 +1582,7 @@ async def _parse_codex_sse_images(
|
||||
line = line_bytes.strip()
|
||||
if line == "":
|
||||
if buffer:
|
||||
data_lines = []
|
||||
data_lines: list[str] = []
|
||||
for bl in buffer:
|
||||
if bl.startswith("data:"):
|
||||
data_lines.append(bl[5:].strip())
|
||||
@@ -1577,9 +1592,11 @@ async def _parse_codex_sse_images(
|
||||
if raw == "[DONE]":
|
||||
break
|
||||
try:
|
||||
event = _json.loads(raw)
|
||||
event = _as_json_object(_json.loads(raw))
|
||||
except Exception:
|
||||
continue
|
||||
if event is None:
|
||||
continue
|
||||
ev_type = event.get("type", "")
|
||||
if ev_type in ("error", "response.failed"):
|
||||
logger.error("Codex SSE failure: {}", raw[:2000])
|
||||
@@ -1596,12 +1613,13 @@ async def _parse_codex_sse_images(
|
||||
raw = "".join(data_lines)
|
||||
if raw and raw != "[DONE]":
|
||||
try:
|
||||
event = _json.loads(raw)
|
||||
event = _as_json_object(_json.loads(raw))
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
_collect_images_from_sse_event(event, images)
|
||||
_collect_text_from_sse_event(event, text_parts)
|
||||
if event is not None:
|
||||
_collect_images_from_sse_event(event, images)
|
||||
_collect_text_from_sse_event(event, text_parts)
|
||||
|
||||
return images, "".join(text_parts).strip()
|
||||
|
||||
@@ -1609,7 +1627,7 @@ async def _parse_codex_sse_images(
|
||||
def _collect_images_from_sse_event(event: dict[str, Any], images: list[str]) -> None:
|
||||
if event.get("type") != "response.output_item.done":
|
||||
return
|
||||
item = event.get("item") or {}
|
||||
item = _as_json_object(event.get("item")) or {}
|
||||
if item.get("type") != "image_generation_call":
|
||||
return
|
||||
result = item.get("result")
|
||||
@@ -1618,8 +1636,8 @@ def _collect_images_from_sse_event(event: dict[str, Any], images: list[str]) ->
|
||||
images.append(result)
|
||||
else:
|
||||
images.append(_b64_image_data_url(result))
|
||||
elif isinstance(result, dict):
|
||||
image_url = result.get("image_url") or result.get("image") or ""
|
||||
elif (result_object := _as_json_object(result)) is not None:
|
||||
image_url = result_object.get("image_url") or result_object.get("image") or ""
|
||||
if isinstance(image_url, str):
|
||||
if image_url.startswith("data:image/"):
|
||||
images.append(image_url)
|
||||
@@ -1749,9 +1767,7 @@ def _stepfun_images_from_payload(payload: dict[str, Any]) -> list[str]:
|
||||
StepFun returns images in ``data[].b64_json`` (base64 strings).
|
||||
"""
|
||||
images: list[str] = []
|
||||
for item in payload.get("data") or []:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
for item in _as_json_objects(payload.get("data")):
|
||||
b64 = item.get("b64_json")
|
||||
if isinstance(b64, str) and b64:
|
||||
images.append(_b64_image_data_url(b64))
|
||||
@@ -1894,9 +1910,7 @@ async def _zhipu_images_from_payload(
|
||||
We download and re-encode as base64 data URLs.
|
||||
"""
|
||||
images: list[str] = []
|
||||
for item in payload.get("data") or []:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
for item in _as_json_objects(payload.get("data")):
|
||||
url = item.get("url")
|
||||
if isinstance(url, str) and url:
|
||||
images.append(await _download_image_data_url(url, proxy=proxy))
|
||||
@@ -2080,7 +2094,7 @@ class ModelScopeImageGenerationClient(ImageGenerationProvider):
|
||||
data: dict[str, Any],
|
||||
) -> list[str]:
|
||||
images: list[str] = []
|
||||
for url in data.get("output_images") or []:
|
||||
for url in cast(list[object], data.get("output_images") or []):
|
||||
if isinstance(url, str) and url:
|
||||
if url.startswith("data:image/"):
|
||||
images.append(url)
|
||||
|
||||
@@ -1,12 +1,14 @@
|
||||
"""OpenAI Codex Responses Provider."""
|
||||
|
||||
# pyright: reportMissingTypeStubs=false, reportPrivateUsage=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import json
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
import httpx
|
||||
from loguru import logger
|
||||
@@ -83,7 +85,7 @@ class OpenAICodexProvider(LLMProvider):
|
||||
stage = "oauth_token"
|
||||
try:
|
||||
token = await asyncio.to_thread(get_codex_token, proxy=self.proxy)
|
||||
headers = _build_headers(token.account_id, token.access)
|
||||
headers = _build_headers(cast(str, token.account_id), token.access)
|
||||
|
||||
stage = "codex_request"
|
||||
try:
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
"""OpenAI-compatible provider for all non-Anthropic LLM APIs."""
|
||||
|
||||
# pyright: reportPrivateImportUsage=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
@@ -13,9 +15,9 @@ import string
|
||||
import time
|
||||
import uuid
|
||||
from collections import deque
|
||||
from collections.abc import Awaitable, Callable
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable, Iterable
|
||||
from ipaddress import ip_address
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from loguru import logger
|
||||
@@ -91,12 +93,18 @@ _OPENAI_COMPAT_REQUEST_TIMEOUT_S = 120.0
|
||||
# Maps ProviderSpec.thinking_style → extra_body builder.
|
||||
# Each builder takes a bool (thinking_enabled) and returns the dict to
|
||||
# merge into extra_body, keeping the style→wire-format mapping in one place.
|
||||
_THINKING_STYLE_MAP: dict[str, Any] = {
|
||||
_THINKING_STYLE_MAP: dict[
|
||||
str,
|
||||
Callable[[bool], dict[str, Any]],
|
||||
] = {
|
||||
"thinking_type": lambda on: {"thinking": {"type": "enabled" if on else "disabled"}},
|
||||
"enable_thinking": lambda on: {"enable_thinking": on},
|
||||
"reasoning_split": lambda on: {"reasoning_split": on},
|
||||
}
|
||||
_GATEWAY_REASONING_STYLE_MAP: dict[str, Any] = {
|
||||
_GATEWAY_REASONING_STYLE_MAP: dict[
|
||||
str,
|
||||
Callable[[str], dict[str, Any]],
|
||||
] = {
|
||||
"reasoning_effort": lambda effort: {"reasoning": {"effort": effort}},
|
||||
}
|
||||
_QWEN_THINKING_MODELS: frozenset[str] = frozenset({
|
||||
@@ -202,23 +210,30 @@ def _extract_text_tool_calls(content: str | None) -> tuple[str | None, list[Tool
|
||||
spans: list[tuple[int, int]] = []
|
||||
for match in _TEXT_TOOL_CALL_RE.finditer(content):
|
||||
try:
|
||||
payload = json.loads(_strip_json_fence(match.group(1)))
|
||||
raw_payload: object = json.loads(
|
||||
_strip_json_fence(match.group(1))
|
||||
)
|
||||
except Exception:
|
||||
continue
|
||||
if not isinstance(payload, dict):
|
||||
if not isinstance(raw_payload, dict):
|
||||
continue
|
||||
payload = cast(dict[str, Any], raw_payload)
|
||||
|
||||
nested = payload.get("tool_call")
|
||||
nested = cast(object, payload.get("tool_call"))
|
||||
if isinstance(nested, dict):
|
||||
payload = nested
|
||||
function = payload.get("function")
|
||||
payload = cast(dict[str, Any], nested)
|
||||
function = cast(object, payload.get("function"))
|
||||
if not isinstance(function, dict):
|
||||
function = payload
|
||||
name = function.get("name")
|
||||
function_data = cast(dict[str, Any], function)
|
||||
name = cast(object, function_data.get("name"))
|
||||
if not isinstance(name, str) or not name:
|
||||
continue
|
||||
|
||||
arguments = function.get("arguments", payload.get("arguments", {}))
|
||||
arguments = function_data.get(
|
||||
"arguments",
|
||||
payload.get("arguments", {}),
|
||||
)
|
||||
tool_calls.append(ToolCallRequest(
|
||||
id=str(payload.get("id") or _short_tool_id()),
|
||||
name=name,
|
||||
@@ -239,24 +254,24 @@ def _extract_text_tool_calls(content: str | None) -> tuple[str | None, list[Tool
|
||||
return visible_content, tool_calls
|
||||
|
||||
|
||||
def _get(obj: Any, key: str) -> Any:
|
||||
def _get(obj: object, key: str) -> Any:
|
||||
"""Get a value from dict or object attribute, returning None if absent."""
|
||||
if isinstance(obj, dict):
|
||||
return obj.get(key)
|
||||
return cast(dict[str, Any], obj).get(key)
|
||||
return getattr(obj, key, None)
|
||||
|
||||
|
||||
def _coerce_dict(value: Any) -> dict[str, Any] | None:
|
||||
def _coerce_dict(value: object) -> dict[str, Any] | None:
|
||||
"""Try to coerce *value* to a dict; return None if not possible or empty."""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, dict):
|
||||
return value if value else None
|
||||
return cast(dict[str, Any], value) if value else None
|
||||
model_dump = getattr(value, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
dumped = model_dump()
|
||||
dumped: object = model_dump()
|
||||
if isinstance(dumped, dict) and dumped:
|
||||
return dumped
|
||||
return cast(dict[str, Any], dumped)
|
||||
return None
|
||||
|
||||
|
||||
@@ -368,19 +383,25 @@ def _deep_merge(base: dict[str, Any], override: dict[str, Any]) -> dict[str, Any
|
||||
and isinstance(merged[key], dict)
|
||||
and isinstance(value, dict)
|
||||
):
|
||||
merged[key] = _deep_merge(merged[key], value)
|
||||
merged[key] = _deep_merge(
|
||||
cast(dict[str, Any], merged[key]),
|
||||
cast(dict[str, Any], value),
|
||||
)
|
||||
else:
|
||||
merged[key] = value
|
||||
return merged
|
||||
|
||||
|
||||
def _merge_unique_list(base: Any, override: Any) -> Any:
|
||||
def _merge_unique_list(base: object, override: object) -> object:
|
||||
"""Append list values while preserving order and removing duplicates."""
|
||||
if not isinstance(base, list) or not isinstance(override, list):
|
||||
return override
|
||||
result: list[Any] = []
|
||||
result: list[object] = []
|
||||
seen: set[str] = set()
|
||||
for value in [*base, *override]:
|
||||
for value in [
|
||||
*cast(list[object], base),
|
||||
*cast(list[object], override),
|
||||
]:
|
||||
try:
|
||||
key = json.dumps(value, sort_keys=True, ensure_ascii=False)
|
||||
except Exception:
|
||||
@@ -513,7 +534,7 @@ class OpenAICompatProvider(LLMProvider):
|
||||
http_client=http_client,
|
||||
)
|
||||
|
||||
async def _ensure_client(self):
|
||||
async def _ensure_client(self) -> AsyncOpenAIType:
|
||||
"""Return the shared OpenAI client, creating it on first call."""
|
||||
if self._client is not None:
|
||||
return self._client
|
||||
@@ -534,6 +555,8 @@ class OpenAICompatProvider(LLMProvider):
|
||||
AsyncOpenAI = _AsyncOpenAI
|
||||
|
||||
self._build_client()
|
||||
if self._client is None:
|
||||
raise RuntimeError("OpenAI client initialization did not produce a client")
|
||||
return self._client
|
||||
|
||||
def _setup_env(self, api_key: str, api_base: str | None) -> None:
|
||||
@@ -567,7 +590,7 @@ class OpenAICompatProvider(LLMProvider):
|
||||
{"type": "text", "text": content, "cache_control": cache_marker},
|
||||
]}
|
||||
if isinstance(content, list) and content:
|
||||
nc = list(content)
|
||||
nc = list(cast(list[dict[str, Any]], content))
|
||||
nc[-1] = {**nc[-1], "cache_control": cache_marker}
|
||||
return {**msg, "content": nc}
|
||||
return msg
|
||||
@@ -662,23 +685,24 @@ class OpenAICompatProvider(LLMProvider):
|
||||
return map_id(value)
|
||||
|
||||
for clean in sanitized:
|
||||
if isinstance(clean.get("tool_calls"), list):
|
||||
normalized = []
|
||||
tool_calls_value = cast(object, clean.get("tool_calls"))
|
||||
if isinstance(tool_calls_value, list):
|
||||
normalized: list[Any] = []
|
||||
used_ids: set[str] = set()
|
||||
for idx, tc in enumerate(clean["tool_calls"]):
|
||||
for idx, tc in enumerate(cast(list[object], tool_calls_value)):
|
||||
if not isinstance(tc, dict):
|
||||
normalized.append(tc)
|
||||
continue
|
||||
tc_clean = dict(tc)
|
||||
tc_clean = dict(cast(dict[str, Any], tc))
|
||||
raw_id = tc_clean.get("id")
|
||||
mapped_id = unique_tool_id(raw_id, used_ids, idx)
|
||||
tc_clean["id"] = mapped_id
|
||||
used_ids.add(mapped_id)
|
||||
if isinstance(raw_id, str) and raw_id:
|
||||
pending_tool_ids.setdefault(raw_id, deque()).append(mapped_id)
|
||||
function = tc_clean.get("function")
|
||||
function = cast(object, tc_clean.get("function"))
|
||||
if isinstance(function, dict):
|
||||
function_clean = dict(function)
|
||||
function_clean = dict(cast(dict[str, Any], function))
|
||||
if "arguments" in function_clean:
|
||||
function_clean["arguments"] = tool_arguments_json_for_replay(
|
||||
function_clean.get("arguments")
|
||||
@@ -715,9 +739,13 @@ class OpenAICompatProvider(LLMProvider):
|
||||
route_prefixes = getattr(spec, "strip_model_prefixes", ())
|
||||
if not isinstance(route_prefixes, tuple) or not route_prefixes:
|
||||
return model_name
|
||||
typed_route_prefixes = cast(tuple[str, ...], route_prefixes)
|
||||
model_prefix, routed_model = model_name.split("/", 1)
|
||||
model_prefix_key = _provider_prefix_key(model_prefix)
|
||||
if any(_provider_prefix_key(prefix) == model_prefix_key for prefix in route_prefixes):
|
||||
if any(
|
||||
_provider_prefix_key(prefix) == model_prefix_key
|
||||
for prefix in typed_route_prefixes
|
||||
):
|
||||
return routed_model
|
||||
return model_name
|
||||
|
||||
@@ -1050,25 +1078,25 @@ class OpenAICompatProvider(LLMProvider):
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _maybe_mapping(value: Any) -> dict[str, Any] | None:
|
||||
def _maybe_mapping(value: object) -> dict[str, Any] | None:
|
||||
if isinstance(value, dict):
|
||||
return value
|
||||
return cast(dict[str, Any], value)
|
||||
model_dump = getattr(value, "model_dump", None)
|
||||
if callable(model_dump):
|
||||
dumped = model_dump()
|
||||
dumped: object = model_dump()
|
||||
if isinstance(dumped, dict):
|
||||
return dumped
|
||||
return cast(dict[str, Any], dumped)
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def _extract_text_content(cls, value: Any) -> str | None:
|
||||
def _extract_text_content(cls, value: object) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
if isinstance(value, list):
|
||||
parts: list[str] = []
|
||||
for item in value:
|
||||
for item in cast(list[object], value):
|
||||
item_map = cls._maybe_mapping(item)
|
||||
if item_map:
|
||||
# Skip Mistral-style {"type":"thinking","thinking":[...]}
|
||||
@@ -1089,7 +1117,7 @@ class OpenAICompatProvider(LLMProvider):
|
||||
return str(value)
|
||||
|
||||
@classmethod
|
||||
def _extract_thinking_content(cls, value: Any) -> str | None:
|
||||
def _extract_thinking_content(cls, value: object) -> str | None:
|
||||
"""Extract reasoning text from Mistral-style thinking blocks.
|
||||
|
||||
Mistral returns content as a list mixing
|
||||
@@ -1101,7 +1129,7 @@ class OpenAICompatProvider(LLMProvider):
|
||||
if not isinstance(value, list):
|
||||
return None
|
||||
parts: list[str] = []
|
||||
for item in value:
|
||||
for item in cast(list[object], value):
|
||||
item_map = cls._maybe_mapping(item)
|
||||
if not item_map:
|
||||
continue
|
||||
@@ -1163,21 +1191,21 @@ class OpenAICompatProvider(LLMProvider):
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _get_nested_int(obj: Any, path: tuple[str, ...]) -> int:
|
||||
def _get_nested_int(obj: object, path: tuple[str, ...]) -> int:
|
||||
"""Drill into *obj* by *path* segments and return an ``int`` value.
|
||||
|
||||
Supports both dict-key access and attribute access so it works
|
||||
uniformly with raw JSON dicts **and** SDK Pydantic models.
|
||||
"""
|
||||
current = obj
|
||||
current: object = obj
|
||||
for segment in path:
|
||||
if current is None:
|
||||
return 0
|
||||
if isinstance(current, dict):
|
||||
current = current.get(segment)
|
||||
current = cast(dict[str, Any], current).get(segment)
|
||||
else:
|
||||
current = getattr(current, segment, None)
|
||||
return int(current or 0) if current is not None else 0
|
||||
return int(cast(Any, current) or 0) if current is not None else 0
|
||||
|
||||
def _parse(self, response: Any) -> LLMResponse:
|
||||
if isinstance(response, str):
|
||||
@@ -1185,7 +1213,10 @@ class OpenAICompatProvider(LLMProvider):
|
||||
|
||||
response_map = self._maybe_mapping(response)
|
||||
if response_map is not None:
|
||||
choices = response_map.get("choices") or []
|
||||
choices = cast(
|
||||
list[object],
|
||||
response_map.get("choices") or [],
|
||||
)
|
||||
if not choices:
|
||||
content = self._extract_text_content(
|
||||
response_map.get("content") or response_map.get("output_text")
|
||||
@@ -1211,7 +1242,7 @@ class OpenAICompatProvider(LLMProvider):
|
||||
content = self._extract_text_content(msg0.get("content"))
|
||||
finish_reason = str(choice0.get("finish_reason") or "stop")
|
||||
|
||||
raw_tool_calls: list[Any] = []
|
||||
raw_tool_calls: list[object] = []
|
||||
# StepFun: fallback to reasoning field when content is empty
|
||||
if not content and msg0.get("reasoning") and self._spec and self._spec.reasoning_as_content:
|
||||
content = self._extract_text_content(msg0.get("reasoning"))
|
||||
@@ -1227,9 +1258,11 @@ class OpenAICompatProvider(LLMProvider):
|
||||
for ch in choices:
|
||||
ch_map = self._maybe_mapping(ch) or {}
|
||||
m = self._maybe_mapping(ch_map.get("message")) or {}
|
||||
tool_calls = m.get("tool_calls")
|
||||
if isinstance(tool_calls, list) and tool_calls:
|
||||
raw_tool_calls.extend(tool_calls)
|
||||
message_tool_calls = cast(object, m.get("tool_calls"))
|
||||
if isinstance(message_tool_calls, list) and message_tool_calls:
|
||||
raw_tool_calls.extend(
|
||||
cast(list[object], message_tool_calls)
|
||||
)
|
||||
if ch_map.get("finish_reason") in ("tool_calls", "stop"):
|
||||
finish_reason = str(ch_map["finish_reason"])
|
||||
if not content:
|
||||
@@ -1240,7 +1273,7 @@ class OpenAICompatProvider(LLMProvider):
|
||||
# Deduplicate tool call IDs (same pattern as streaming path)
|
||||
# Some providers reuse the same ID for parallel tool calls.
|
||||
_seen_tc_ids: set[str] = set()
|
||||
parsed_tool_calls = []
|
||||
parsed_tool_calls: list[ToolCallRequest] = []
|
||||
for tc in raw_tool_calls:
|
||||
tc_map = self._maybe_mapping(tc) or {}
|
||||
fn = self._maybe_mapping(tc_map.get("function")) or {}
|
||||
@@ -1281,11 +1314,11 @@ class OpenAICompatProvider(LLMProvider):
|
||||
content = msg.content
|
||||
finish_reason = choice.finish_reason
|
||||
|
||||
raw_tool_calls: list[Any] = []
|
||||
raw_sdk_tool_calls: list[Any] = []
|
||||
for ch in response.choices:
|
||||
m = ch.message
|
||||
if hasattr(m, "tool_calls") and m.tool_calls:
|
||||
raw_tool_calls.extend(m.tool_calls)
|
||||
raw_sdk_tool_calls.extend(m.tool_calls)
|
||||
if ch.finish_reason in ("tool_calls", "stop"):
|
||||
finish_reason = ch.finish_reason
|
||||
if not content and m.content:
|
||||
@@ -1293,8 +1326,8 @@ class OpenAICompatProvider(LLMProvider):
|
||||
if not content and getattr(m, "reasoning", None) and self._spec and self._spec.reasoning_as_content:
|
||||
content = m.reasoning
|
||||
|
||||
tool_calls = []
|
||||
for tc in raw_tool_calls:
|
||||
tool_calls: list[ToolCallRequest] = []
|
||||
for tc in raw_sdk_tool_calls:
|
||||
args = parse_tool_arguments(tc.function.arguments)
|
||||
ec, prov, fn_prov = _extract_tc_extras(tc)
|
||||
tool_calls.append(ToolCallRequest(
|
||||
@@ -1376,7 +1409,10 @@ class OpenAICompatProvider(LLMProvider):
|
||||
|
||||
chunk_map = cls._maybe_mapping(chunk)
|
||||
if chunk_map is not None:
|
||||
choices = chunk_map.get("choices") or []
|
||||
choices = cast(
|
||||
list[object],
|
||||
chunk_map.get("choices") or [],
|
||||
)
|
||||
if not choices:
|
||||
usage = cls._extract_usage(chunk_map) or usage
|
||||
text = cls._extract_text_content(
|
||||
@@ -1402,7 +1438,12 @@ class OpenAICompatProvider(LLMProvider):
|
||||
text = cls._extract_thinking_content(raw_delta_content)
|
||||
if text:
|
||||
reasoning_parts.append(text)
|
||||
for idx, tc in enumerate(delta.get("tool_calls") or []):
|
||||
for idx, tc in enumerate(
|
||||
cast(
|
||||
Iterable[object],
|
||||
delta.get("tool_calls") or [],
|
||||
)
|
||||
):
|
||||
_accum_tc(tc, idx)
|
||||
_accum_legacy_function_call(delta.get("function_call"))
|
||||
usage = cls._extract_usage(chunk_map) or usage
|
||||
@@ -1430,7 +1471,12 @@ class OpenAICompatProvider(LLMProvider):
|
||||
text = cls._extract_text_content(reasoning)
|
||||
if text:
|
||||
reasoning_parts.append(text)
|
||||
for tc in (getattr(delta, "tool_calls", None) or []) if delta else []:
|
||||
delta_tool_calls = (
|
||||
cast(Iterable[object], getattr(delta, "tool_calls", None) or [])
|
||||
if delta
|
||||
else ()
|
||||
)
|
||||
for tc in delta_tool_calls:
|
||||
_accum_tc(tc, getattr(tc, "index", 0))
|
||||
if delta:
|
||||
_accum_legacy_function_call(getattr(delta, "function_call", None))
|
||||
@@ -1563,7 +1609,7 @@ class OpenAICompatProvider(LLMProvider):
|
||||
reasoning_effort: str | None = None,
|
||||
tool_choice: str | dict[str, Any] | None = None,
|
||||
) -> LLMResponse:
|
||||
await self._ensure_client()
|
||||
client = await self._ensure_client()
|
||||
try:
|
||||
if self._should_use_responses_api(model, reasoning_effort):
|
||||
try:
|
||||
@@ -1571,7 +1617,11 @@ class OpenAICompatProvider(LLMProvider):
|
||||
messages, tools, model, max_tokens, temperature,
|
||||
reasoning_effort, tool_choice,
|
||||
)
|
||||
result = parse_response_output(await self._client.responses.create(**body))
|
||||
responses_raw = cast(
|
||||
Any,
|
||||
await client.responses.create(**body),
|
||||
)
|
||||
result = parse_response_output(responses_raw)
|
||||
self._record_responses_success(model, reasoning_effort)
|
||||
return result
|
||||
except Exception as responses_error:
|
||||
@@ -1590,7 +1640,11 @@ class OpenAICompatProvider(LLMProvider):
|
||||
messages, tools, model, max_tokens, temperature,
|
||||
reasoning_effort, tool_choice,
|
||||
)
|
||||
return self._parse(await self._client.chat.completions.create(**kwargs))
|
||||
chat_raw = cast(
|
||||
Any,
|
||||
await client.chat.completions.create(**kwargs),
|
||||
)
|
||||
return self._parse(chat_raw)
|
||||
except Exception as e:
|
||||
return self._handle_error(e, spec=self._spec, api_base=self.api_base)
|
||||
|
||||
@@ -1607,7 +1661,7 @@ class OpenAICompatProvider(LLMProvider):
|
||||
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
|
||||
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
|
||||
) -> LLMResponse:
|
||||
await self._ensure_client()
|
||||
client = await self._ensure_client()
|
||||
idle_timeout_s = resolve_stream_idle_timeout_s()
|
||||
try:
|
||||
if self._should_use_responses_api(model, reasoning_effort):
|
||||
@@ -1617,10 +1671,13 @@ class OpenAICompatProvider(LLMProvider):
|
||||
reasoning_effort, tool_choice,
|
||||
)
|
||||
body["stream"] = True
|
||||
stream = await self._client.responses.create(**body)
|
||||
responses_stream = cast(
|
||||
Any,
|
||||
await client.responses.create(**body),
|
||||
)
|
||||
|
||||
async def _timed_stream():
|
||||
stream_iter = stream.__aiter__()
|
||||
async def _timed_stream() -> AsyncIterator[Any]:
|
||||
stream_iter: AsyncIterator[Any] = responses_stream.__aiter__()
|
||||
while True:
|
||||
try:
|
||||
yield await asyncio.wait_for(
|
||||
@@ -1673,12 +1730,15 @@ class OpenAICompatProvider(LLMProvider):
|
||||
kwargs.setdefault("extra_body", {})["tool_stream"] = True
|
||||
kwargs["stream"] = True
|
||||
kwargs["stream_options"] = {"include_usage": True}
|
||||
stream = await self._client.chat.completions.create(**kwargs)
|
||||
chat_stream = cast(
|
||||
Any,
|
||||
await client.chat.completions.create(**kwargs),
|
||||
)
|
||||
chunks: list[Any] = []
|
||||
stream_iter = stream.__aiter__()
|
||||
stream_iter: AsyncIterator[Any] = chat_stream.__aiter__()
|
||||
while True:
|
||||
try:
|
||||
chunk = await asyncio.wait_for(
|
||||
chunk: Any = await asyncio.wait_for(
|
||||
stream_iter.__anext__(),
|
||||
timeout=idle_timeout_s,
|
||||
)
|
||||
|
||||
@@ -3,11 +3,15 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
from nanobot.providers.base import tool_arguments_json_for_replay
|
||||
|
||||
|
||||
def _as_json_object(value: object) -> dict[str, Any] | None:
|
||||
return cast(dict[str, Any], value) if isinstance(value, dict) else None
|
||||
|
||||
|
||||
def convert_messages(messages: list[dict[str, Any]]) -> tuple[str, list[dict[str, Any]]]:
|
||||
"""Convert Chat Completions messages to Responses API input items.
|
||||
|
||||
@@ -39,8 +43,11 @@ def convert_messages(messages: list[dict[str, Any]]) -> tuple[str, list[dict[str
|
||||
"content": [{"type": "output_text", "text": content}],
|
||||
"status": "completed", "id": message_id,
|
||||
})
|
||||
for tool_call in msg.get("tool_calls", []) or []:
|
||||
fn = tool_call.get("function") or {}
|
||||
for raw_tool_call in cast(list[object], msg.get("tool_calls", []) or []):
|
||||
tool_call = _as_json_object(raw_tool_call)
|
||||
if tool_call is None:
|
||||
continue
|
||||
fn = _as_json_object(tool_call.get("function")) or {}
|
||||
call_id, item_id = split_tool_call_id(tool_call.get("id"))
|
||||
response_item_id = _unique_item_id(item_id or f"fc_{idx}", used_item_ids)
|
||||
input_items.append({
|
||||
@@ -70,13 +77,15 @@ def convert_user_message(content: Any) -> dict[str, Any]:
|
||||
return {"role": "user", "content": [{"type": "input_text", "text": content}]}
|
||||
if isinstance(content, list):
|
||||
converted: list[dict[str, Any]] = []
|
||||
for item in content:
|
||||
if not isinstance(item, dict):
|
||||
for raw_item in cast(list[object], content):
|
||||
item = _as_json_object(raw_item)
|
||||
if item is None:
|
||||
continue
|
||||
if item.get("type") == "text":
|
||||
converted.append({"type": "input_text", "text": item.get("text", "")})
|
||||
elif item.get("type") == "image_url":
|
||||
url = (item.get("image_url") or {}).get("url")
|
||||
image = _as_json_object(item.get("image_url")) or {}
|
||||
url = image.get("url")
|
||||
if url:
|
||||
converted.append({"type": "input_image", "image_url": url, "detail": "auto"})
|
||||
if converted:
|
||||
@@ -97,8 +106,9 @@ def convert_tool_output(content: Any) -> str | list[dict[str, Any]]:
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
converted: list[dict[str, Any]] = []
|
||||
for item in content:
|
||||
if not isinstance(item, dict):
|
||||
for raw_item in cast(list[object], content):
|
||||
item = _as_json_object(raw_item)
|
||||
if item is None:
|
||||
break
|
||||
item_type = item.get("type")
|
||||
if item_type in {"text", "input_text"}:
|
||||
@@ -110,15 +120,16 @@ def convert_tool_output(content: Any) -> str | list[dict[str, Any]]:
|
||||
converted.append({"type": "input_text", "text": text})
|
||||
elif item_type in {"image_url", "input_image"}:
|
||||
image = item.get("image_url")
|
||||
if isinstance(image, dict) and set(image) - {"url", "detail"}:
|
||||
image_object = _as_json_object(image)
|
||||
if image_object is not None and set(image_object) - {"url", "detail"}:
|
||||
break
|
||||
if set(item) - {"type", "image_url", "file_id", "detail", "_meta"}:
|
||||
break
|
||||
url = image.get("url") if isinstance(image, dict) else image
|
||||
url = image_object.get("url") if image_object is not None else image
|
||||
file_id = item.get("file_id")
|
||||
detail = item.get(
|
||||
"detail",
|
||||
image.get("detail", "auto") if isinstance(image, dict) else "auto",
|
||||
image_object.get("detail", "auto") if image_object is not None else "auto",
|
||||
)
|
||||
if detail not in {"low", "high", "auto", "original"}:
|
||||
break
|
||||
@@ -160,11 +171,11 @@ def convert_tools(tools: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
"""Convert OpenAI function-calling tool schema to Responses API flat format."""
|
||||
converted: list[dict[str, Any]] = []
|
||||
for tool in tools:
|
||||
fn = (tool.get("function") or {}) if tool.get("type") == "function" else tool
|
||||
fn = _as_json_object(tool.get("function")) or {} if tool.get("type") == "function" else tool
|
||||
name = fn.get("name")
|
||||
if not name:
|
||||
continue
|
||||
params = fn.get("parameters") or {}
|
||||
params: object = fn.get("parameters") or {}
|
||||
converted.append({
|
||||
"type": "function",
|
||||
"name": name,
|
||||
|
||||
@@ -4,7 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any, AsyncGenerator
|
||||
from typing import Any, AsyncGenerator, cast
|
||||
|
||||
import httpx
|
||||
from loguru import logger
|
||||
@@ -19,23 +19,58 @@ FINISH_REASON_MAP = {
|
||||
}
|
||||
|
||||
|
||||
def _as_json_object(value: object) -> dict[str, Any] | None:
|
||||
"""Narrow untyped Responses API JSON payloads at the wire boundary."""
|
||||
return cast(dict[str, Any], value) if isinstance(value, dict) else None
|
||||
|
||||
|
||||
def _response_object(value: object) -> dict[str, Any] | None:
|
||||
"""Convert a Responses SDK model or JSON object to a dictionary."""
|
||||
object_value = _as_json_object(value)
|
||||
if object_value is not None:
|
||||
return object_value
|
||||
dump = getattr(value, "model_dump", None)
|
||||
if callable(dump):
|
||||
return _as_json_object(dump())
|
||||
try:
|
||||
return _as_json_object(vars(value))
|
||||
except TypeError:
|
||||
return None
|
||||
|
||||
|
||||
def _response_object_list(value: object) -> list[dict[str, Any]]:
|
||||
"""Normalize a Responses API array that may contain SDK model objects."""
|
||||
if not isinstance(value, list):
|
||||
return []
|
||||
return [
|
||||
item
|
||||
for raw in cast(list[object], value)
|
||||
if (item := _response_object(raw)) is not None
|
||||
]
|
||||
|
||||
|
||||
def map_finish_reason(status: str | None) -> str:
|
||||
"""Map a Responses API status string to a Chat-Completions-style finish_reason."""
|
||||
return FINISH_REASON_MAP.get(status or "completed", "stop")
|
||||
|
||||
|
||||
def _usage_from_response_obj(response: Any) -> dict[str, int]:
|
||||
usage_raw = response.get("usage") if isinstance(response, dict) else getattr(response, "usage", None)
|
||||
def _usage_from_response_obj(response: object) -> dict[str, int]:
|
||||
response_object = _response_object(response)
|
||||
usage_raw: object = (
|
||||
response_object.get("usage")
|
||||
if response_object is not None
|
||||
else getattr(response, "usage", None)
|
||||
)
|
||||
if not usage_raw:
|
||||
return {}
|
||||
if not isinstance(usage_raw, dict):
|
||||
dump = getattr(usage_raw, "model_dump", None)
|
||||
usage_raw = dump() if callable(dump) else vars(usage_raw)
|
||||
prompt_tokens = int(usage_raw.get("input_tokens") or usage_raw.get("prompt_tokens") or 0)
|
||||
usage = _response_object(usage_raw)
|
||||
if usage is None:
|
||||
return {}
|
||||
prompt_tokens = int(usage.get("input_tokens") or usage.get("prompt_tokens") or 0)
|
||||
completion_tokens = int(
|
||||
usage_raw.get("output_tokens") or usage_raw.get("completion_tokens") or 0
|
||||
usage.get("output_tokens") or usage.get("completion_tokens") or 0
|
||||
)
|
||||
total_tokens = int(usage_raw.get("total_tokens") or prompt_tokens + completion_tokens)
|
||||
total_tokens = int(usage.get("total_tokens") or prompt_tokens + completion_tokens)
|
||||
return {
|
||||
"prompt_tokens": prompt_tokens,
|
||||
"completion_tokens": completion_tokens,
|
||||
@@ -77,7 +112,7 @@ async def iter_sse(response: httpx.Response) -> AsyncGenerator[dict[str, Any], N
|
||||
if not data or data == "[DONE]":
|
||||
return None
|
||||
try:
|
||||
return json.loads(data)
|
||||
return _as_json_object(json.loads(data))
|
||||
except Exception:
|
||||
logger.warning("Failed to parse SSE event JSON: {}", data[:200])
|
||||
return None
|
||||
@@ -134,7 +169,7 @@ async def consume_sse_with_reasoning(
|
||||
await on_response_event(event)
|
||||
event_type = event.get("type")
|
||||
if event_type == "response.output_item.added":
|
||||
item = event.get("item") or {}
|
||||
item = _as_json_object(event.get("item")) or {}
|
||||
if item.get("type") == "function_call":
|
||||
call_id = item.get("call_id")
|
||||
if not call_id:
|
||||
@@ -170,7 +205,7 @@ async def consume_sse_with_reasoning(
|
||||
if on_reasoning_delta:
|
||||
await on_reasoning_delta(text)
|
||||
elif event_type == "response.reasoning_summary_part.done":
|
||||
part = event.get("part") or {}
|
||||
part = _as_json_object(event.get("part")) or {}
|
||||
text = part.get("text") if part.get("type") == "summary_text" else None
|
||||
if text and not streamed_reasoning and not reasoning_content:
|
||||
reasoning_content = text
|
||||
@@ -203,7 +238,7 @@ async def consume_sse_with_reasoning(
|
||||
"arguments": "" if arguments is None else str(arguments),
|
||||
})
|
||||
elif event_type == "response.output_item.done":
|
||||
item = event.get("item") or {}
|
||||
item = _as_json_object(event.get("item")) or {}
|
||||
if item.get("type") == "function_call":
|
||||
call_id = item.get("call_id")
|
||||
if not call_id:
|
||||
@@ -235,12 +270,12 @@ async def consume_sse_with_reasoning(
|
||||
if on_reasoning_delta:
|
||||
await on_reasoning_delta(summary)
|
||||
elif event_type == "response.completed":
|
||||
response_obj = event.get("response") or {}
|
||||
response_obj = _response_object(event.get("response")) or {}
|
||||
status = response_obj.get("status")
|
||||
finish_reason = map_finish_reason(status)
|
||||
usage = _usage_from_response_obj(response_obj) or usage
|
||||
if not reasoning_content:
|
||||
summary = _extract_reasoning_summary_from_output(response_obj.get("output") or [])
|
||||
summary = _extract_reasoning_summary_from_output(response_obj.get("output"))
|
||||
if summary:
|
||||
reasoning_content = summary
|
||||
if on_reasoning_delta:
|
||||
@@ -252,54 +287,42 @@ async def consume_sse_with_reasoning(
|
||||
return content, tool_calls, finish_reason, usage, reasoning_content
|
||||
|
||||
|
||||
def _extract_reasoning_summary_from_output(output: Any) -> str | None:
|
||||
def _extract_reasoning_summary_from_output(output: object) -> str | None:
|
||||
parts: list[str] = []
|
||||
for item in output or []:
|
||||
if not isinstance(item, dict):
|
||||
dump = getattr(item, "model_dump", None)
|
||||
item = dump() if callable(dump) else vars(item)
|
||||
for item in _response_object_list(output):
|
||||
if item.get("type") != "reasoning":
|
||||
continue
|
||||
for summary in item.get("summary") or []:
|
||||
if not isinstance(summary, dict):
|
||||
dump = getattr(summary, "model_dump", None)
|
||||
summary = dump() if callable(dump) else vars(summary)
|
||||
for summary in _response_object_list(item.get("summary")):
|
||||
if summary.get("type") == "summary_text" and summary.get("text"):
|
||||
parts.append(summary["text"])
|
||||
text = summary.get("text")
|
||||
if isinstance(text, str):
|
||||
parts.append(text)
|
||||
return "".join(parts) or None
|
||||
|
||||
|
||||
def parse_response_output(response: Any) -> LLMResponse:
|
||||
def parse_response_output(response: object) -> LLMResponse:
|
||||
"""Parse an SDK ``Response`` object into an ``LLMResponse``."""
|
||||
if not isinstance(response, dict):
|
||||
dump = getattr(response, "model_dump", None)
|
||||
response = dump() if callable(dump) else vars(response)
|
||||
response_object = _response_object(response) or {}
|
||||
|
||||
output = response.get("output") or []
|
||||
output = _response_object_list(response_object.get("output"))
|
||||
content_parts: list[str] = []
|
||||
tool_calls: list[ToolCallRequest] = []
|
||||
reasoning_content: str | None = None
|
||||
|
||||
for item in output:
|
||||
if not isinstance(item, dict):
|
||||
dump = getattr(item, "model_dump", None)
|
||||
item = dump() if callable(dump) else vars(item)
|
||||
|
||||
item_type = item.get("type")
|
||||
if item_type == "message":
|
||||
for block in item.get("content") or []:
|
||||
if not isinstance(block, dict):
|
||||
dump = getattr(block, "model_dump", None)
|
||||
block = dump() if callable(dump) else vars(block)
|
||||
for block in _response_object_list(item.get("content")):
|
||||
if block.get("type") == "output_text":
|
||||
content_parts.append(block.get("text") or "")
|
||||
text = block.get("text")
|
||||
if isinstance(text, str):
|
||||
content_parts.append(text)
|
||||
elif item_type == "reasoning":
|
||||
for s in item.get("summary") or []:
|
||||
if not isinstance(s, dict):
|
||||
dump = getattr(s, "model_dump", None)
|
||||
s = dump() if callable(dump) else vars(s)
|
||||
for s in _response_object_list(item.get("summary")):
|
||||
if s.get("type") == "summary_text" and s.get("text"):
|
||||
reasoning_content = (reasoning_content or "") + s["text"]
|
||||
text = s.get("text")
|
||||
if isinstance(text, str):
|
||||
reasoning_content = (reasoning_content or "") + text
|
||||
elif item_type == "function_call":
|
||||
call_id = item.get("call_id") or ""
|
||||
item_id = item.get("id") or "fc_0"
|
||||
@@ -311,10 +334,10 @@ def parse_response_output(response: Any) -> LLMResponse:
|
||||
arguments=args,
|
||||
))
|
||||
|
||||
usage = _usage_from_response_obj(response)
|
||||
usage = _usage_from_response_obj(response_object)
|
||||
|
||||
status = response.get("status")
|
||||
finish_reason = map_finish_reason(status)
|
||||
status = response_object.get("status")
|
||||
finish_reason = map_finish_reason(status if isinstance(status, str) else None)
|
||||
|
||||
return LLMResponse(
|
||||
content="".join(content_parts) or None,
|
||||
@@ -339,7 +362,8 @@ async def consume_sdk_stream(
|
||||
usage: dict[str, int] = {}
|
||||
reasoning_content: str | None = None
|
||||
|
||||
async for event in stream:
|
||||
async for raw_event in stream:
|
||||
event: Any = raw_event
|
||||
event_type = getattr(event, "type", None)
|
||||
if event_type == "response.output_item.added":
|
||||
item = getattr(event, "item", None)
|
||||
@@ -431,9 +455,9 @@ async def consume_sdk_stream(
|
||||
"completion_tokens": int(getattr(usage_obj, "output_tokens", 0) or 0),
|
||||
"total_tokens": int(getattr(usage_obj, "total_tokens", 0) or 0),
|
||||
}
|
||||
for out_item in getattr(resp, "output", None) or []:
|
||||
for out_item in cast(list[Any], getattr(resp, "output", None) or []):
|
||||
if getattr(out_item, "type", None) == "reasoning":
|
||||
for s in getattr(out_item, "summary", None) or []:
|
||||
for s in cast(list[Any], getattr(out_item, "summary", None) or []):
|
||||
if getattr(s, "type", None) == "summary_text":
|
||||
text = getattr(s, "text", None)
|
||||
if text:
|
||||
|
||||
@@ -13,7 +13,7 @@ import mimetypes
|
||||
import os
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, cast
|
||||
|
||||
import httpx
|
||||
from loguru import logger
|
||||
@@ -116,7 +116,7 @@ async def _request_json_with_retry(
|
||||
url: str,
|
||||
*,
|
||||
provider_label: str,
|
||||
**kwargs: object,
|
||||
**kwargs: Any,
|
||||
) -> dict[str, Any] | None:
|
||||
for attempt in range(_MAX_RETRIES + 1):
|
||||
try:
|
||||
@@ -190,7 +190,7 @@ async def _request_json_with_retry(
|
||||
type(payload).__name__,
|
||||
)
|
||||
return None
|
||||
return payload
|
||||
return cast(dict[str, Any], payload)
|
||||
return None
|
||||
|
||||
|
||||
@@ -383,6 +383,7 @@ async def _post_stepfun_asr_with_retry(
|
||||
payload = json.loads(payload_str)
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
continue
|
||||
payload = cast(dict[str, Any], payload)
|
||||
event_type = payload.get("type", "")
|
||||
if event_type == "error":
|
||||
msg = payload.get("message", "unknown error")
|
||||
@@ -503,7 +504,7 @@ async def _post_with_retry(
|
||||
type(payload).__name__,
|
||||
)
|
||||
return ""
|
||||
return extract_text(payload)
|
||||
return extract_text(cast(dict[str, Any], payload))
|
||||
return ""
|
||||
|
||||
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user