From 757ad9c76416e36d92bd83ec0b8acf983432841c Mon Sep 17 00:00:00 2001 From: chengyongru <61816729+chengyongru@users.noreply.github.com> Date: Wed, 29 Jul 2026 21:37:11 +0800 Subject: [PATCH] refactor: enforce BasedPyright strict type checking (#5158) --- .agent/design.md | 8 + .github/workflows/ci.yml | 4 + AGENTS.md | 5 + CONTRIBUTING.md | 14 + nanobot/__init__.py | 28 +- nanobot/agent/autocompact.py | 13 +- nanobot/agent/context.py | 17 +- nanobot/agent/context_governance.py | 36 +- nanobot/agent/hooks/file_edit_activity.py | 10 +- nanobot/agent/loop.py | 246 ++++++++---- nanobot/agent/memory.py | 64 ++-- nanobot/agent/model_presets.py | 7 +- nanobot/agent/model_runtime.py | 8 +- nanobot/agent/progress_hook.py | 6 +- nanobot/agent/runner.py | 105 ++++-- nanobot/agent/skills.py | 43 ++- nanobot/agent/subagent.py | 32 +- nanobot/agent/tools/apply_patch.py | 14 +- nanobot/agent/tools/base.py | 51 ++- nanobot/agent/tools/cli_apps.py | 9 +- nanobot/agent/tools/context.py | 32 +- nanobot/agent/tools/cron.py | 21 +- nanobot/agent/tools/exec_session.py | 44 +-- nanobot/agent/tools/file_state.py | 6 +- nanobot/agent/tools/filesystem.py | 10 +- nanobot/agent/tools/image_generation.py | 44 ++- nanobot/agent/tools/loader.py | 12 +- nanobot/agent/tools/long_task.py | 34 +- nanobot/agent/tools/mcp.py | 145 +++++--- nanobot/agent/tools/message.py | 36 +- nanobot/agent/tools/registry.py | 15 +- nanobot/agent/tools/runtime_state.py | 30 +- nanobot/agent/tools/search.py | 2 + nanobot/agent/tools/self.py | 83 +++-- nanobot/agent/tools/shell.py | 30 +- nanobot/agent/tools/spawn.py | 10 +- nanobot/agent/tools/web.py | 155 +++++--- nanobot/agent/turn_delivery.py | 9 +- nanobot/api/runtime.py | 2 +- nanobot/api/server.py | 82 +++- nanobot/apps/cli/service.py | 58 +-- nanobot/apps/cli/utils.py | 13 +- nanobot/audio/transcription.py | 20 +- nanobot/bus/outbound_events.py | 18 +- nanobot/bus/runtime_events.py | 11 +- nanobot/channels/base.py | 14 +- nanobot/channels/contracts.py | 95 +++-- nanobot/channels/dingtalk/runtime.py | 102 +++-- nanobot/channels/discord/runtime.py | 35 +- nanobot/channels/email/runtime.py | 20 +- nanobot/channels/feishu/connect.py | 2 + nanobot/channels/feishu/instances.py | 25 +- nanobot/channels/feishu/runtime.py | 350 +++++++++++------- nanobot/channels/feishu/websocket.py | 13 +- nanobot/channels/manager.py | 31 +- nanobot/channels/matrix/runtime.py | 154 ++++++-- nanobot/channels/mattermost/runtime.py | 106 ++++-- nanobot/channels/mochat/runtime.py | 161 +++++--- nanobot/channels/msteams/runtime.py | 99 +++-- nanobot/channels/napcat/runtime.py | 45 ++- nanobot/channels/plugin.py | 10 +- nanobot/channels/qq/runtime.py | 94 +++-- nanobot/channels/signal/runtime.py | 127 ++++--- nanobot/channels/slack/runtime.py | 198 +++++++--- nanobot/channels/telegram/runtime.py | 153 +++++--- .../telegram/tests/test_telegram_channel.py | 65 +++- nanobot/channels/telegram/validation.py | 4 +- nanobot/channels/validation.py | 18 +- nanobot/channels/websocket/runtime.py | 99 +++-- nanobot/channels/wecom/runtime.py | 95 +++-- nanobot/channels/weixin/connect.py | 73 ++-- nanobot/channels/weixin/runtime.py | 194 +++++++--- nanobot/channels/whatsapp/runtime.py | 14 +- nanobot/cli/commands.py | 241 ++++++------ nanobot/cli/gateway.py | 2 + nanobot/cli/onboard.py | 185 +++++---- nanobot/cli/stream.py | 7 +- nanobot/command/builtin.py | 87 +++-- nanobot/command/router.py | 3 +- nanobot/config/loader.py | 82 ++-- nanobot/config/paths.py | 2 +- nanobot/config/schema.py | 7 +- nanobot/config/watcher.py | 2 +- nanobot/cron/__init__.py | 7 +- nanobot/cron/bound_runner.py | 7 +- nanobot/cron/service.py | 38 +- nanobot/cron/types.py | 35 +- nanobot/gateway/runtime.py | 2 +- nanobot/optional_features.py | 37 +- nanobot/pairing/store.py | 67 ++-- nanobot/process_runtime.py | 20 +- nanobot/providers/anthropic_provider.py | 89 +++-- nanobot/providers/azure_openai_provider.py | 6 +- nanobot/providers/base.py | 64 ++-- nanobot/providers/bedrock_provider.py | 155 +++++--- nanobot/providers/factory.py | 2 + nanobot/providers/fallback_provider.py | 8 +- nanobot/providers/github_copilot_provider.py | 26 +- nanobot/providers/image_generation.py | 162 ++++---- nanobot/providers/openai_codex_provider.py | 6 +- nanobot/providers/openai_compat_provider.py | 190 ++++++---- .../providers/openai_responses/converters.py | 37 +- nanobot/providers/openai_responses/parsing.py | 124 ++++--- nanobot/providers/transcription.py | 9 +- nanobot/providers/unconfigured_provider.py | 8 +- nanobot/providers/xai_grok_provider.py | 46 ++- nanobot/providers/xai_oauth.py | 30 +- nanobot/runtime_context.py | 46 ++- nanobot/sdk/clients.py | 4 +- nanobot/sdk/streaming.py | 2 + nanobot/sdk/types.py | 13 +- nanobot/security/network.py | 16 +- nanobot/security/workspace_access.py | 18 +- nanobot/security/workspace_policy.py | 2 +- nanobot/session/automation_turns.py | 6 +- nanobot/session/goal_state.py | 6 +- nanobot/session/manager.py | 197 ++++++---- nanobot/session/model_selection.py | 6 +- nanobot/session/turn_continuation.py | 11 +- nanobot/session/webui_turns.py | 8 +- .../skill-creator/scripts/init_skill.py | 27 +- .../skill-creator/scripts/package_skill.py | 6 +- .../skill-creator/scripts/quick_validate.py | 12 +- nanobot/triggers/local_store.py | 16 +- nanobot/triggers/local_types.py | 13 +- nanobot/utils/document.py | 21 +- nanobot/utils/file_edit_events.py | 13 +- nanobot/utils/gitstore.py | 75 ++-- nanobot/utils/helpers.py | 82 ++-- nanobot/utils/progress_events.py | 13 +- nanobot/utils/restart.py | 4 +- nanobot/utils/runtime.py | 8 +- nanobot/utils/searchusage.py | 11 +- nanobot/utils/subagent_channel_display.py | 4 +- nanobot/utils/tool_hints.py | 36 +- nanobot/webui/attachment_ingress.py | 14 +- nanobot/webui/build.py | 8 +- nanobot/webui/cli_apps_api.py | 13 +- nanobot/webui/forking.py | 37 +- nanobot/webui/gateway_services.py | 27 +- nanobot/webui/http_utils.py | 4 +- nanobot/webui/mcp_presets_api.py | 83 +++-- nanobot/webui/media_api.py | 15 +- nanobot/webui/media_gateway.py | 6 +- nanobot/webui/session_automations.py | 13 +- nanobot/webui/session_list_index.py | 30 +- nanobot/webui/settings_api.py | 109 ++++-- nanobot/webui/settings_routes.py | 35 +- nanobot/webui/sidebar_state.py | 24 +- nanobot/webui/token_usage.py | 40 +- nanobot/webui/transcript.py | 199 ++++++---- nanobot/webui/workspaces.py | 18 +- nanobot/webui/ws_http.py | 32 +- pyproject.toml | 7 + tests/agent/test_runner_governance.py | 5 +- tests/channels/test_channel_plugins.py | 2 + tests/cli/test_commands.py | 35 +- tests/command/test_builtin_dream.py | 2 + tests/command/test_router_dispatchable.py | 8 + tests/config/test_env_interpolation.py | 7 + tests/cron/test_cron_service.py | 11 + tests/pairing/test_store.py | 19 + tests/providers/test_transcription.py | 2 +- tests/test_api_attachment.py | 22 ++ tests/test_document_parsing.py | 7 + tests/tools/test_mcp_tool.py | 8 + 166 files changed, 4728 insertions(+), 2621 deletions(-) diff --git a/.agent/design.md b/.agent/design.md index 75ea7607b..e598d99be 100644 --- a/.agent/design.md +++ b/.agent/design.md @@ -24,6 +24,14 @@ Fix bugs by changing only what is necessary. Do not bundle unrelated refactors o A bugfix should make the protected invariant clear, change the smallest surface that enforces it, and add only the closest regression test. If a diff starts changing ownership boundaries or mixing behavior changes with clean-up, split it before it becomes hard to review. +## Type dynamic boundaries at the edge + +Wire payloads, persisted records, and third-party SDK objects are untrusted dynamic boundaries. Prefer a parser or small normalizer at the owning edge, and use `TypedDict` for stable dictionary shapes, so validation happens once and internal code receives a concrete type. Do not spread raw dynamic dictionaries or SDK objects through the core. + +Stable first-party dependencies must be typed where they are stored or passed. Do not declare an internal service, context field, or callback result as `Any` and then recover its real type with consumer-side casts. Use the concrete type or a narrow `Protocol`; reserve `Any` for genuinely dynamic boundaries. + +`typing.cast` performs no runtime validation. Every new cast must be supported by a runtime check on the same path or by an explicit invariant that is clear from construction and control flow (and documented locally when it is not obvious). If input can violate the claimed type, handle that invalid case before casting; never use `cast` only to silence BasedPyright. + ## Explicit over magical Configuration must be declared explicitly in `config/schema.py` Pydantic models. Error handling should raise clear exceptions rather than silently correcting bad input. Provider auto-detection exists, but every resolution path must be traceable from the factory to the concrete provider class. diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 515cbb2eb..dacf6acb7 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -129,6 +129,10 @@ jobs: if: matrix.coverage run: uv run --no-sync ruff check nanobot tests conftest.py + - name: Type check with BasedPyright (strict) + if: matrix.coverage + run: uv run --no-sync basedpyright + - name: Run tests with coverage if: matrix.coverage run: >- diff --git a/AGENTS.md b/AGENTS.md index 58ea83854..8217f0ad5 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -11,6 +11,11 @@ nanobot is a lightweight, open-source AI agent framework written in Python with pytest tests/test_openai_api.py::test_function -v ruff check nanobot/ +# Strict type checking (matches CI) +uv sync --all-extras --dev +uv run --no-sync python -m scripts.install_channel_dependencies --all-channels +uv run --no-sync basedpyright + # WebUI: dev server (proxies API/WS to gateway :8765), build, test # Build outputs to ../nanobot/web/dist (bundled into the Python wheel) cd webui && bun run dev # or NANOBOT_API_URL=... bun run dev diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index c897514fc..7c7183418 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -78,6 +78,20 @@ ruff check nanobot/ ruff format ``` +### Strict Type Checking + +Strict type checking covers optional providers and channels. Reproduce the CI environment +with the same dependency sources and commands: + +```bash +uv sync --all-extras --dev +uv run --no-sync python -m scripts.install_channel_dependencies --all-channels +uv run --no-sync basedpyright +``` + +Keep `--no-sync` on the final commands: channel dependencies come from their package +manifests and are installed explicitly by the setup step. + ## Contribution License By submitting a contribution, you confirm that you have the right to submit it diff --git a/nanobot/__init__.py b/nanobot/__init__.py index cf3c19ec8..20986b2bb 100644 --- a/nanobot/__init__.py +++ b/nanobot/__init__.py @@ -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}") diff --git a/nanobot/agent/autocompact.py b/nanobot/agent/autocompact.py index d73bf9446..ba1f629d2 100644 --- a/nanobot/agent/autocompact.py +++ b/nanobot/agent/autocompact.py @@ -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 diff --git a/nanobot/agent/context.py b/nanobot/agent/context.py index f96200921..7713917ea 100644 --- a/nanobot/agent/context.py +++ b/nanobot/agent/context.py @@ -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(): diff --git a/nanobot/agent/context_governance.py b/nanobot/agent/context_governance.py index 9a1a18776..98b1291dc 100644 --- a/nanobot/agent/context_governance.py +++ b/nanobot/agent/context_governance.py @@ -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: diff --git a/nanobot/agent/hooks/file_edit_activity.py b/nanobot/agent/hooks/file_edit_activity.py index 8de68f051..f45c56d4a 100644 --- a/nanobot/agent/hooks/file_edit_activity.py +++ b/nanobot/agent/hooks/file_edit_activity.py @@ -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, diff --git a/nanobot/agent/loop.py b/nanobot/agent/loop.py index 61de18fb2..e31e66c03 100644 --- a/nanobot/agent/loop.py +++ b/nanobot/agent/loop.py @@ -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", diff --git a/nanobot/agent/memory.py b/nanobot/agent/memory.py index e3dcc9a7f..753f9c954 100644 --- a/nanobot/agent/memory.py +++ b/nanobot/agent/memory.py @@ -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. diff --git a/nanobot/agent/model_presets.py b/nanobot/agent/model_presets.py index eb2a9a643..79a7d4e80 100644 --- a/nanobot/agent/model_presets.py +++ b/nanobot/agent/model_presets.py @@ -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: diff --git a/nanobot/agent/model_runtime.py b/nanobot/agent/model_runtime.py index 0a6ece829..d2d7b469f 100644 --- a/nanobot/agent/model_runtime.py +++ b/nanobot/agent/model_runtime.py @@ -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") diff --git a/nanobot/agent/progress_hook.py b/nanobot/agent/progress_hook.py index 826093d9a..82b493cd4 100644 --- a/nanobot/agent/progress_hook.py +++ b/nanobot/agent/progress_hook.py @@ -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, ) diff --git a/nanobot/agent/runner.py b/nanobot/agent/runner.py index 3e3811880..3c5f88c36 100644 --- a/nanobot/agent/runner.py +++ b/nanobot/agent/runner.py @@ -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: diff --git a/nanobot/agent/skills.py b/nanobot/agent/skills.py index 1e748f04b..31a81c086 100644 --- a/nanobot/agent/skills.py +++ b/nanobot/agent/skills.py @@ -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 diff --git a/nanobot/agent/subagent.py b/nanobot/agent/subagent.py index b4f4cece8..c4d151f74 100644 --- a/nanobot/agent/subagent.py +++ b/nanobot/agent/subagent.py @@ -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] = [] diff --git a/nanobot/agent/tools/apply_patch.py b/nanobot/agent/tools/apply_patch.py index 43c526f50..9a14d1f82 100644 --- a/nanobot/agent/tools/apply_patch.py +++ b/nanobot/agent/tools/apply_patch.py @@ -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: diff --git a/nanobot/agent/tools/base.py b/nanobot/agent/tools/base.py index bcea7e5ab..46cc4b472 100644 --- a/nanobot/agent/tools/base.py +++ b/nanobot/agent/tools/base.py @@ -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": diff --git a/nanobot/agent/tools/cli_apps.py b/nanobot/agent/tools/cli_apps.py index 7642e390f..76758a386 100644 --- a/nanobot/agent/tools/cli_apps.py +++ b/nanobot/agent/tools/cli_apps.py @@ -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), diff --git a/nanobot/agent/tools/context.py b/nanobot/agent/tools/context.py index 7baa71a66..1db6a3dd3 100644 --- a/nanobot/agent/tools/context.py +++ b/nanobot/agent/tools/context.py @@ -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 diff --git a/nanobot/agent/tools/cron.py b/nanobot/agent/tools/cron.py index 89f389f11..6a4fd653a 100644 --- a/nanobot/agent/tools/cron.py +++ b/nanobot/agent/tools/cron.py @@ -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})"] diff --git a/nanobot/agent/tools/exec_session.py b/nanobot/agent/tools/exec_session.py index 9390bc02e..1245edfdc 100644 --- a/nanobot/agent/tools/exec_session.py +++ b/nanobot/agent/tools/exec_session.py @@ -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: diff --git a/nanobot/agent/tools/file_state.py b/nanobot/agent/tools/file_state.py index 33673b3ef..3dd4667d5 100644 --- a/nanobot/agent/tools/file_state.py +++ b/nanobot/agent/tools/file_state.py @@ -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) diff --git a/nanobot/agent/tools/filesystem.py b/nanobot/agent/tools/filesystem.py index 18da2c145..d1406604f 100644 --- a/nanobot/agent/tools/filesystem.py +++ b/nanobot/agent/tools/filesystem.py @@ -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: diff --git a/nanobot/agent/tools/image_generation.py b/nanobot/agent/tools/image_generation.py index b1e448e70..de164f116 100644 --- a/nanobot/agent/tools/image_generation.py +++ b/nanobot/agent/tools/image_generation.py @@ -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 diff --git a/nanobot/agent/tools/loader.py b/nanobot/agent/tools/loader.py index fb420562e..27760cec6 100644 --- a/nanobot/agent/tools/loader.py +++ b/nanobot/agent/tools/loader.py @@ -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) diff --git a/nanobot/agent/tools/long_task.py b/nanobot/agent/tools/long_task.py index 5aa36fc10..359b60c6c 100644 --- a/nanobot/agent/tools/long_task.py +++ b/nanobot/agent/tools/long_task.py @@ -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: diff --git a/nanobot/agent/tools/mcp.py b/nanobot/agent/tools/mcp.py index 478032839..27a54de33 100644 --- a/nanobot/agent/tools/mcp.py +++ b/nanobot/agent/tools/mcp.py @@ -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 "" -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 diff --git a/nanobot/agent/tools/message.py b/nanobot/agent/tools/message.py index d8a660090..12e008f10 100644 --- a/nanobot/agent/tools/message.py +++ b/nanobot/agent/tools/message.py @@ -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)}") diff --git a/nanobot/agent/tools/registry.py b/nanobot/agent/tools/registry.py index e21222840..f5f94f654 100644 --- a/nanobot/agent/tools/registry.py +++ b/nanobot/agent/tools/registry.py @@ -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.""" diff --git a/nanobot/agent/tools/runtime_state.py b/nanobot/agent/tools/runtime_state.py index 288988699..3efe8e870 100644 --- a/nanobot/agent/tools/runtime_state.py +++ b/nanobot/agent/tools/runtime_state.py @@ -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: ... diff --git a/nanobot/agent/tools/search.py b/nanobot/agent/tools/search.py index 775311fcd..feba01a56 100644 --- a/nanobot/agent/tools/search.py +++ b/nanobot/agent/tools/search.py @@ -1,5 +1,7 @@ """Search tools: file discovery and grep.""" +# pyright: reportIncompatibleMethodOverride=false, reportPrivateUsage=false + from __future__ import annotations import fnmatch diff --git a/nanobot/agent/tools/self.py b/nanobot/agent/tools/self.py index ae60f96e2..3cad85242 100644 --- a/nanobot/agent/tools/self.py +++ b/nanobot/agent/tools/self.py @@ -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): diff --git a/nanobot/agent/tools/shell.py b/nanobot/agent/tools/shell.py index 6650a2af2..be3d5c04b 100644 --- a/nanobot/agent/tools/shell.py +++ b/nanobot/agent/tools/shell.py @@ -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 diff --git a/nanobot/agent/tools/spawn.py b/nanobot/agent/tools/spawn.py index 8c64076df..1936434e1 100644 --- a/nanobot/agent/tools/spawn.py +++ b/nanobot/agent/tools/spawn.py @@ -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: diff --git a/nanobot/agent/tools/web.py b/nanobot/agent/tools/web.py index 834988cf6..c2d8e011c 100644 --- a/nanobot/agent/tools/web.py +++ b/nanobot/agent/tools/web.py @@ -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 diff --git a/nanobot/agent/turn_delivery.py b/nanobot/agent/turn_delivery.py index 9f2aba0a4..5b3746b2d 100644 --- a/nanobot/agent/turn_delivery.py +++ b/nanobot/agent/turn_delivery.py @@ -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: diff --git a/nanobot/api/runtime.py b/nanobot/api/runtime.py index ef062156d..97fa3af90 100644 --- a/nanobot/api/runtime.py +++ b/nanobot/api/runtime.py @@ -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" diff --git a/nanobot/api/server.py b/nanobot/api/server.py index bc2f8a7c2..9ad57deef 100644 --- a/nanobot/api/server.py +++ b/nanobot/api/server.py @@ -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) diff --git a/nanobot/apps/cli/service.py b/nanobot/apps/cli/service.py index 8c5d63916..413b449bd 100644 --- a/nanobot/apps/cli/service.py +++ b/nanobot/apps/cli/service.py @@ -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() diff --git a/nanobot/apps/cli/utils.py b/nanobot/apps/cli/utils.py index 850dc598f..5668a486d 100644 --- a/nanobot/apps/cli/utils.py +++ b/nanobot/apps/cli/utils.py @@ -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 [ diff --git a/nanobot/audio/transcription.py b/nanobot/audio/transcription.py index 539f90b95..c336cc88e 100644 --- a/nanobot/audio/transcription.py +++ b/nanobot/audio/transcription.py @@ -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) diff --git a/nanobot/bus/outbound_events.py b/nanobot/bus/outbound_events.py index 1a5b8d551..f750b2c74 100644 --- a/nanobot/bus/outbound_events.py +++ b/nanobot/bus/outbound_events.py @@ -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 diff --git a/nanobot/bus/runtime_events.py b/nanobot/bus/runtime_events.py index 0d9a83027..30be6f402 100644 --- a/nanobot/bus/runtime_events.py +++ b/nanobot/bus/runtime_events.py @@ -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: diff --git a/nanobot/channels/base.py b/nanobot/channels/base.py index 01a794a44..aed1407ec 100644 --- a/nanobot/channels/base.py +++ b/nanobot/channels/base.py @@ -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: diff --git a/nanobot/channels/contracts.py b/nanobot/channels/contracts.py index 560f17d4e..75bb4a0db 100644 --- a/nanobot/channels/contracts.py +++ b/nanobot/channels/contracts.py @@ -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: diff --git a/nanobot/channels/dingtalk/runtime.py b/nanobot/channels/dingtalk/runtime.py index dd3989153..f00d75e49 100644 --- a/nanobot/channels/dingtalk/runtime.py +++ b/nanobot/channels/dingtalk/runtime.py @@ -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") diff --git a/nanobot/channels/discord/runtime.py b/nanobot/channels/discord/runtime.py index bab06fe24..9b5afc373 100644 --- a/nanobot/channels/discord/runtime.py +++ b/nanobot/channels/discord/runtime.py @@ -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) diff --git a/nanobot/channels/email/runtime.py b/nanobot/channels/email/runtime.py index 5f025e0fc..f3aaefa5f 100644 --- a/nanobot/channels/email/runtime.py +++ b/nanobot/channels/email/runtime.py @@ -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): diff --git a/nanobot/channels/feishu/connect.py b/nanobot/channels/feishu/connect.py index 3258d1550..b41e57fae 100644 --- a/nanobot/channels/feishu/connect.py +++ b/nanobot/channels/feishu/connect.py @@ -1,5 +1,7 @@ """Short-lived WebUI channel connection sessions.""" +# pyright: reportPrivateUsage=false + from __future__ import annotations import asyncio diff --git a/nanobot/channels/feishu/instances.py b/nanobot/channels/feishu/instances.py index 44de6f9bd..adff0e636 100644 --- a/nanobot/channels/feishu/instances.py +++ b/nanobot/channels/feishu/instances.py @@ -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) diff --git a/nanobot/channels/feishu/runtime.py b/nanobot/channels/feishu/runtime.py index a9753e6a1..96ad3a328 100644 --- a/nanobot/channels/feishu/runtime.py +++ b/nanobot/channels/feishu/runtime.py @@ -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: diff --git a/nanobot/channels/feishu/websocket.py b/nanobot/channels/feishu/websocket.py index d7d005605..e378c5c1d 100644 --- a/nanobot/channels/feishu/websocket.py +++ b/nanobot/channels/feishu/websocket.py @@ -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 diff --git a/nanobot/channels/manager.py b/nanobot/channels/manager.py index 3fbe32ede..7c8d392b2 100644 --- a/nanobot/channels/manager.py +++ b/nanobot/channels/manager.py @@ -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)) diff --git a/nanobot/channels/matrix/runtime.py b/nanobot/channels/matrix/runtime.py index 992aba7a9..04f8b254f 100644 --- a/nanobot/channels/matrix/runtime.py +++ b/nanobot/channels/matrix/runtime.py @@ -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) diff --git a/nanobot/channels/mattermost/runtime.py b/nanobot/channels/mattermost/runtime.py index 473dac316..cbd575423 100644 --- a/nanobot/channels/mattermost/runtime.py +++ b/nanobot/channels/mattermost/runtime.py @@ -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"(? 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: diff --git a/nanobot/channels/mochat/runtime.py b/nanobot/channels/mochat/runtime.py index de99711fe..e5e44e863 100644 --- a/nanobot/channels/mochat/runtime.py +++ b/nanobot/channels/mochat/runtime.py @@ -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 diff --git a/nanobot/channels/msteams/runtime.py b/nanobot/channels/msteams/runtime.py index addb4164f..e040092e5 100644 --- a/nanobot/channels/msteams/runtime.py +++ b/nanobot/channels/msteams/runtime.py @@ -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 diff --git a/nanobot/channels/napcat/runtime.py b/nanobot/channels/napcat/runtime.py index 3bfdae3c8..b431f2359 100644 --- a/nanobot/channels/napcat/runtime.py +++ b/nanobot/channels/napcat/runtime.py @@ -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)) diff --git a/nanobot/channels/plugin.py b/nanobot/channels/plugin.py index beee1c20b..3bdd8d49a 100644 --- a/nanobot/channels/plugin.py +++ b/nanobot/channels/plugin.py @@ -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") diff --git a/nanobot/channels/qq/runtime.py b/nanobot/channels/qq/runtime.py index 0fb30de10..b744fec14 100644 --- a/nanobot/channels/qq/runtime.py +++ b/nanobot/channels/qq/runtime.py @@ -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) diff --git a/nanobot/channels/signal/runtime.py b/nanobot/channels/signal/runtime.py index 3a282fb76..8afffeabe 100644 --- a/nanobot/channels/signal/runtime.py +++ b/nanobot/channels/signal/runtime.py @@ -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)}} diff --git a/nanobot/channels/slack/runtime.py b/nanobot/channels/slack/runtime.py index 6b7b37a41..56512165c 100644 --- a/nanobot/channels/slack/runtime.py +++ b/nanobot/channels/slack/runtime.py @@ -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: diff --git a/nanobot/channels/telegram/runtime.py b/nanobot/channels/telegram/runtime.py index 9e42b2df1..7fac96b0a 100644 --- a/nanobot/channels/telegram/runtime.py +++ b/nanobot/channels/telegram/runtime.py @@ -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( diff --git a/nanobot/channels/telegram/tests/test_telegram_channel.py b/nanobot/channels/telegram/tests/test_telegram_channel.py index 498aa892e..9e5ef17ec 100644 --- a/nanobot/channels/telegram/tests/test_telegram_channel.py +++ b/nanobot/channels/telegram/tests/test_telegram_channel.py @@ -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" diff --git a/nanobot/channels/telegram/validation.py b/nanobot/channels/telegram/validation.py index 7244b1c6e..5931a5de6 100644 --- a/nanobot/channels/telegram/validation.py +++ b/nanobot/channels/telegram/validation.py @@ -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]: diff --git a/nanobot/channels/validation.py b/nanobot/channels/validation.py index 91b3779e3..661664ca1 100644 --- a/nanobot/channels/validation.py +++ b/nanobot/channels/validation.py @@ -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: diff --git a/nanobot/channels/websocket/runtime.py b/nanobot/channels/websocket/runtime.py index 5dd1bcfb5..fc2471ba0 100644 --- a/nanobot/channels/websocket/runtime.py +++ b/nanobot/channels/websocket/runtime.py @@ -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) diff --git a/nanobot/channels/wecom/runtime.py b/nanobot/channels/wecom/runtime.py index e850504c0..066de3b3a 100644 --- a/nanobot/channels/wecom/runtime.py +++ b/nanobot/channels/wecom/runtime.py @@ -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, diff --git a/nanobot/channels/weixin/connect.py b/nanobot/channels/weixin/connect.py index a254b64d3..36a14d143 100644 --- a/nanobot/channels/weixin/connect.py +++ b/nanobot/channels/weixin/connect.py @@ -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]: diff --git a/nanobot/channels/weixin/runtime.py b/nanobot/channels/weixin/runtime.py index ea5d88b59..ac9869e63 100644 --- a/nanobot/channels/weixin/runtime.py +++ b/nanobot/channels/weixin/runtime.py @@ -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: diff --git a/nanobot/channels/whatsapp/runtime.py b/nanobot/channels/whatsapp/runtime.py index b70c05dab..e576bfe99 100644 --- a/nanobot/channels/whatsapp/runtime.py +++ b/nanobot/channels/whatsapp/runtime.py @@ -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: diff --git a/nanobot/cli/commands.py b/nanobot/cli/commands.py index 0df030eac..98eb0f34b 100644 --- a/nanobot/cli/commands.py +++ b/nanobot/cli/commands.py @@ -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 diff --git a/nanobot/cli/gateway.py b/nanobot/cli/gateway.py index 40011d0a7..1485fd34d 100644 --- a/nanobot/cli/gateway.py +++ b/nanobot/cli/gateway.py @@ -1,5 +1,7 @@ """Typer commands for foreground and background gateway control.""" +# pyright: reportUnusedFunction=false + from __future__ import annotations import subprocess diff --git a/nanobot/cli/onboard.py b/nanobot/cli/onboard.py index 226cc3fcf..892376a3b 100644 --- a/nanobot/cli/onboard.py +++ b/nanobot/cli/onboard.py @@ -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() diff --git a/nanobot/cli/stream.py b/nanobot/cli/stream.py index 24a141cdd..90bf064bb 100644 --- a/nanobot/cli/stream.py +++ b/nanobot/cli/stream.py @@ -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) diff --git a/nanobot/command/builtin.py b/nanobot/command/builtin.py index 4da0b01de..6a5ac1a5f 100644 --- a/nanobot/command/builtin.py +++ b/nanobot/command/builtin.py @@ -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 diff --git a/nanobot/command/router.py b/nanobot/command/router.py index 2a6a9c6f0..eb2939847 100644 --- a/nanobot/command/router.py +++ b/nanobot/command/router.py @@ -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) diff --git a/nanobot/config/loader.py b/nanobot/config/loader.py index 2b7664941..db6522df4 100644 --- a/nanobot/config/loader.py +++ b/nanobot/config/loader.py @@ -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: diff --git a/nanobot/config/paths.py b/nanobot/config/paths.py index 82796038b..bed717b1f 100644 --- a/nanobot/config/paths.py +++ b/nanobot/config/paths.py @@ -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) diff --git a/nanobot/config/schema.py b/nanobot/config/schema.py index 7505dde18..deb312bcb 100644 --- a/nanobot/config/schema.py +++ b/nanobot/config/schema.py @@ -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: diff --git a/nanobot/config/watcher.py b/nanobot/config/watcher.py index 76fa0e6d2..32f10fac4 100644 --- a/nanobot/config/watcher.py +++ b/nanobot/config/watcher.py @@ -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: diff --git a/nanobot/cron/__init__.py b/nanobot/cron/__init__.py index a85f44d1f..70c377f05 100644 --- a/nanobot/cron/__init__.py +++ b/nanobot/cron/__init__.py @@ -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}") diff --git a/nanobot/cron/bound_runner.py b/nanobot/cron/bound_runner.py index 0dfd901ae..eff121a13 100644 --- a/nanobot/cron/bound_runner.py +++ b/nanobot/cron/bound_runner.py @@ -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: ... diff --git a/nanobot/cron/service.py b/nanobot/cron/service.py index 336e7a6f6..f3b04eae9 100644 --- a/nanobot/cron/service.py +++ b/nanobot/cron/service.py @@ -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 { diff --git a/nanobot/cron/types.py b/nanobot/cron/types.py index 89a2d5417..77273f547 100644 --- a/nanobot/cron/types.py +++ b/nanobot/cron/types.py @@ -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: diff --git a/nanobot/gateway/runtime.py b/nanobot/gateway/runtime.py index 60d562f46..6c9740d02 100644 --- a/nanobot/gateway/runtime.py +++ b/nanobot/gateway/runtime.py @@ -69,7 +69,7 @@ class GatewayRuntimePaths(ProcessRuntimePaths): ) -class GatewayRuntime(ManagedProcessRuntime): +class GatewayRuntime(ManagedProcessRuntime[ProcessStartOptions]): """Manage a background ``nanobot gateway`` process.""" service_name = "gateway" diff --git a/nanobot/optional_features.py b/nanobot/optional_features.py index 05b36f03a..6f427c100 100644 --- a/nanobot/optional_features.py +++ b/nanobot/optional_features.py @@ -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: diff --git a/nanobot/pairing/store.py b/nanobot/pairing/store.py index 490369cbc..38253368a 100644 --- a/nanobot/pairing/store.py +++ b/nanobot/pairing/store.py @@ -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 ` or `/pairing revoke `" diff --git a/nanobot/process_runtime.py b/nanobot/process_runtime.py index 2cc17cfc7..23f3d4f80 100644 --- a/nanobot/process_runtime.py +++ b/nanobot/process_runtime.py @@ -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) diff --git a/nanobot/providers/anthropic_provider.py b/nanobot/providers/anthropic_provider.py index 68bb8316b..943a46038 100644 --- a/nanobot/providers/anthropic_provider.py +++ b/nanobot/providers/anthropic_provider.py @@ -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): diff --git a/nanobot/providers/azure_openai_provider.py b/nanobot/providers/azure_openai_provider.py index 5100344e9..8ce72c0d0 100644 --- a/nanobot/providers/azure_openai_provider.py +++ b/nanobot/providers/azure_openai_provider.py @@ -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) ) diff --git a/nanobot/providers/base.py b/nanobot/providers/base.py index 2b0ef3320..642f24177 100644 --- a/nanobot/providers/base.py +++ b/nanobot/providers/base.py @@ -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: diff --git a/nanobot/providers/bedrock_provider.py b/nanobot/providers/bedrock_provider.py index b704195a8..7a4728302 100644 --- a/nanobot/providers/bedrock_provider.py +++ b/nanobot/providers/bedrock_provider.py @@ -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), diff --git a/nanobot/providers/factory.py b/nanobot/providers/factory.py index 3f1a966bd..6e3e0706d 100644 --- a/nanobot/providers/factory.py +++ b/nanobot/providers/factory.py @@ -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, diff --git a/nanobot/providers/fallback_provider.py b/nanobot/providers/fallback_provider.py index 8b19890a5..bf14bdd0a 100644 --- a/nanobot/providers/fallback_provider.py +++ b/nanobot/providers/fallback_provider.py @@ -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: diff --git a/nanobot/providers/github_copilot_provider.py b/nanobot/providers/github_copilot_provider.py index 45c9f3d65..bac99ccaa 100644 --- a/nanobot/providers/github_copilot_provider.py +++ b/nanobot/providers/github_copilot_provider.py @@ -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, diff --git a/nanobot/providers/image_generation.py b/nanobot/providers/image_generation.py index 06170bac0..2a28a0fc3 100644 --- a/nanobot/providers/image_generation.py +++ b/nanobot/providers/image_generation.py @@ -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) diff --git a/nanobot/providers/openai_codex_provider.py b/nanobot/providers/openai_codex_provider.py index 0d7eac0bf..4b9a1d311 100644 --- a/nanobot/providers/openai_codex_provider.py +++ b/nanobot/providers/openai_codex_provider.py @@ -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: diff --git a/nanobot/providers/openai_compat_provider.py b/nanobot/providers/openai_compat_provider.py index c8bad3c86..605d8e717 100644 --- a/nanobot/providers/openai_compat_provider.py +++ b/nanobot/providers/openai_compat_provider.py @@ -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, ) diff --git a/nanobot/providers/openai_responses/converters.py b/nanobot/providers/openai_responses/converters.py index d023781fe..903340b8b 100644 --- a/nanobot/providers/openai_responses/converters.py +++ b/nanobot/providers/openai_responses/converters.py @@ -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, diff --git a/nanobot/providers/openai_responses/parsing.py b/nanobot/providers/openai_responses/parsing.py index bb9ecdc70..d999910c6 100644 --- a/nanobot/providers/openai_responses/parsing.py +++ b/nanobot/providers/openai_responses/parsing.py @@ -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: diff --git a/nanobot/providers/transcription.py b/nanobot/providers/transcription.py index 426f0088e..555f433bf 100644 --- a/nanobot/providers/transcription.py +++ b/nanobot/providers/transcription.py @@ -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 "" diff --git a/nanobot/providers/unconfigured_provider.py b/nanobot/providers/unconfigured_provider.py index d7a15f68c..98d7b69ec 100644 --- a/nanobot/providers/unconfigured_provider.py +++ b/nanobot/providers/unconfigured_provider.py @@ -2,6 +2,8 @@ from __future__ import annotations +from typing import Any + from nanobot.providers.base import LLMProvider, LLMResponse @@ -14,13 +16,13 @@ class UnconfiguredProvider(LLMProvider): async def chat( self, - messages: list[dict], - tools: list[dict] | 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 | None = None, + tool_choice: str | dict[str, Any] | None = None, ) -> LLMResponse: return LLMResponse( content=( diff --git a/nanobot/providers/xai_grok_provider.py b/nanobot/providers/xai_grok_provider.py index 50c10b47b..b99f95696 100644 --- a/nanobot/providers/xai_grok_provider.py +++ b/nanobot/providers/xai_grok_provider.py @@ -9,7 +9,7 @@ import re import time import uuid from collections.abc import Awaitable, Callable -from typing import Any +from typing import Any, cast import httpx from loguru import logger @@ -309,7 +309,7 @@ def _decode_access_token_claims(token: str) -> dict[str, Any]: claims = json.loads(decoded) except (ValueError, TypeError): return {} - return claims if isinstance(claims, dict) else {} + return cast(dict[str, Any], claims) if isinstance(claims, dict) else {} class _XAIHTTPError(RuntimeError): @@ -356,7 +356,8 @@ async def _fetch_xai_model_capabilities( def _parse_xai_model_capabilities(payload: Any) -> dict[str, bool]: if isinstance(payload, dict): - rows = payload.get("data") + payload = cast(dict[str, Any], payload) + rows: object = payload.get("data") if not isinstance(rows, list): rows = payload.get("models") else: @@ -365,10 +366,12 @@ def _parse_xai_model_capabilities(payload: Any) -> dict[str, bool]: return {} capabilities: dict[str, bool] = {} - for row in rows: - if not isinstance(row, dict): + for row_value in cast(list[object], rows): + if not isinstance(row_value, dict): continue - meta = row.get("_meta") if isinstance(row.get("_meta"), dict) else {} + row = cast(dict[str, Any], row_value) + meta_value = row.get("_meta") + meta = cast(dict[str, Any], meta_value) if isinstance(meta_value, dict) else {} support_value = row.get("supportsBackendSearch") if not isinstance(support_value, bool): support_value = row.get("supports_backend_search") @@ -444,7 +447,10 @@ def _xai_hosted_tool_event(event: dict[str, Any]) -> dict[str, Any] | None: if event_type != "response.output_item.done": return None item = event.get("item") - if not isinstance(item, dict) or item.get("type") != "custom_tool_call": + if not isinstance(item, dict): + return None + item = cast(dict[str, Any], item) + if item.get("type") != "custom_tool_call": return None tool_name = item.get("name") if not isinstance(tool_name, str) or not tool_name.startswith("x_"): @@ -468,14 +474,14 @@ def _xai_hosted_tool_event(event: dict[str, Any]) -> dict[str, Any] | None: def _xai_hosted_tool_arguments(value: Any) -> dict[str, Any]: if isinstance(value, dict): - return dict(value) + return cast(dict[str, Any], value) if not isinstance(value, str) or not value.strip(): return {} try: parsed = json.loads(value) except (TypeError, ValueError): return {} - return parsed if isinstance(parsed, dict) else {} + return cast(dict[str, Any], parsed) if isinstance(parsed, dict) else {} def _build_xai_http_error( @@ -483,8 +489,8 @@ def _build_xai_http_error( headers: httpx.Headers, raw: str, ) -> _XAIHTTPError: - retry_after = LLMProvider._extract_retry_after_from_headers(headers) - error_type, error_code = LLMProvider._extract_error_type_code(raw) + retry_after = LLMProvider._extract_retry_after_from_headers(headers) # pyright: ignore[reportPrivateUsage] + error_type, error_code = LLMProvider._extract_error_type_code(raw) # pyright: ignore[reportPrivateUsage] response_body = _bounded_error_body(raw) return _XAIHTTPError( _friendly_error(status_code, response_body), @@ -522,12 +528,16 @@ def _bounded_error_body(raw: str) -> str | None: def _redact_error_payload(payload: Any) -> Any: if isinstance(payload, dict): - return { - key: "[REDACTED]" if _is_sensitive_error_key(key) else _redact_error_payload(value) - for key, value in payload.items() - } + redacted: dict[str, Any] = {} + payload_mapping: dict[str, Any] = cast(dict[str, Any], payload) + for key in payload_mapping: + value = payload_mapping[key] + redacted[key] = ( + "[REDACTED]" if _is_sensitive_error_key(key) else _redact_error_payload(value) + ) + return redacted if isinstance(payload, list): - return [_redact_error_payload(value) for value in payload] + return [_redact_error_payload(value) for value in cast(list[Any], payload)] return payload @@ -593,7 +603,7 @@ def _should_retry_status( content: str | None, ) -> bool: if status_code == 429: - return LLMProvider._is_retryable_429_response( + return LLMProvider._is_retryable_429_response( # pyright: ignore[reportPrivateUsage] LLMResponse( content=content or "", finish_reason="error", @@ -602,4 +612,4 @@ def _should_retry_status( error_code=error_code, ) ) - return status_code in LLMProvider._RETRYABLE_STATUS_CODES or status_code >= 500 + return status_code in LLMProvider._RETRYABLE_STATUS_CODES or status_code >= 500 # pyright: ignore[reportPrivateUsage] diff --git a/nanobot/providers/xai_oauth.py b/nanobot/providers/xai_oauth.py index ca5e0ddae..e77d676fb 100644 --- a/nanobot/providers/xai_oauth.py +++ b/nanobot/providers/xai_oauth.py @@ -23,7 +23,7 @@ from dataclasses import asdict, dataclass from http import HTTPStatus from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path -from typing import Any +from typing import Any, cast from urllib.parse import parse_qs, urlencode, urlsplit import httpx @@ -31,7 +31,7 @@ from filelock import FileLock 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] XAI_OAUTH_ISSUER = "https://auth.x.ai" XAI_CLIENT_ID = "b1a00492-073a-47ea-816f-4c329264a828" @@ -73,17 +73,18 @@ class XAIToken: def from_dict(cls, value: Any) -> XAIToken | None: if not isinstance(value, dict): return None - access = value.get("access") + token_data = cast(dict[str, Any], value) + access = token_data.get("access") if not isinstance(access, str) or not access: return None - refresh = value.get("refresh") + refresh = token_data.get("refresh") if not isinstance(refresh, str) or not refresh: refresh = None try: - expires = int(value.get("expires") or 0) + expires = int(token_data.get("expires") or 0) except (TypeError, ValueError): expires = 0 - account_id = value.get("account_id") + account_id = token_data.get("account_id") if not isinstance(account_id, str) or not account_id: account_id = None return cls(access=access, refresh=refresh, expires=expires, account_id=account_id) @@ -521,7 +522,7 @@ def _make_callback_server( self.send_header("Vary", "Origin") self.send_header("Access-Control-Allow-Private-Network", "true") - def log_message(self, _format: str, *_args: Any) -> None: + def log_message(self, format: str, *_args: Any) -> None: # noqa: A002 # Callback query strings contain an authorization code. return @@ -641,9 +642,12 @@ def _token_payload(response: httpx.Response) -> dict[str, Any]: payload = response.json() except ValueError as exc: raise XAIOAuthError("xAI sign-in returned an invalid token response.") from exc - if not isinstance(payload, dict) or not isinstance(payload.get("access_token"), str): + if not isinstance(payload, dict): raise XAIOAuthError("xAI sign-in returned no access token.") - return payload + token_payload = cast(dict[str, Any], payload) + if not isinstance(token_payload.get("access_token"), str): + raise XAIOAuthError("xAI sign-in returned no access token.") + return token_payload def _token_from_response( @@ -680,8 +684,9 @@ def _fetch_account(endpoint: str | None, access_token: str, proxy: str | None) - return None if not isinstance(payload, dict): return None + account_payload = cast(dict[str, Any], payload) for key in ("email", "preferred_username", "name", "sub"): - value = payload.get(key) + value = account_payload.get(key) if isinstance(value, str) and value: return value return None @@ -693,8 +698,9 @@ def _oauth_http_error(response: httpx.Response, action: str) -> XAIOAuthError: with suppress(ValueError): payload = response.json() if isinstance(payload, dict): - raw_code = payload.get("error") - raw_description = payload.get("error_description") or payload.get("message") + error_payload = cast(dict[str, Any], payload) + raw_code = error_payload.get("error") + raw_description = error_payload.get("error_description") or error_payload.get("message") code = raw_code[:80] if isinstance(raw_code, str) else None description = raw_description[:200] if isinstance(raw_description, str) else None detail = ": ".join(value for value in (code, description) if value) diff --git a/nanobot/runtime_context.py b/nanobot/runtime_context.py index e489fd349..052572cae 100644 --- a/nanobot/runtime_context.py +++ b/nanobot/runtime_context.py @@ -6,7 +6,7 @@ import json from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence from copy import deepcopy from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, TypeAlias +from typing import TYPE_CHECKING, Any, TypeAlias, cast if TYPE_CHECKING: from nanobot.agent.tools.context import RequestContext @@ -79,7 +79,10 @@ def normalize_runtime_context_blocks(result: RuntimeContextResult) -> list[Runti """Return validated, non-empty blocks while preserving provider order.""" if result is None: return [] - values = [result] if isinstance(result, RuntimeContextBlock) else list(result) + if isinstance(cast(object, result), RuntimeContextBlock): + values: list[object] = [result] + else: + values = list(cast(Sequence[object], result)) blocks: list[RuntimeContextBlock] = [] for block in values: if not isinstance(block, RuntimeContextBlock): @@ -147,16 +150,17 @@ def detach_runtime_context( marker: Mapping[str, Any], ) -> tuple[Any, list[str], list[dict[str, Any]]] | None: """Detach one validated runtime-context suffix for safe message merging.""" - if marker.get("version") != 1: + marker_data = marker + if marker_data.get("version") != 1: return None - raw_sources = marker.get("sources") - sources = [ + raw_sources = marker_data.get("sources") + sources: list[str] = [ source - for source in raw_sources + for source in cast(list[Any], raw_sources) if isinstance(source, str) and source ] if isinstance(raw_sources, list) else [] - suffix = marker.get("suffix") + suffix = marker_data.get("suffix") if isinstance(content, str) and isinstance(suffix, str) and suffix: if content == suffix: clean_content = "" @@ -166,12 +170,14 @@ def detach_runtime_context( return None return clean_content, sources, [{"type": "text", "text": suffix}] - expected = marker.get("blocks") + expected = marker_data.get("blocks") if isinstance(content, list) and isinstance(expected, list) and expected: - count = len(expected) - if content[-count:] != expected: + content_blocks = cast(list[Any], content) + expected_blocks = cast(list[dict[str, Any]], expected) + count = len(expected_blocks) + if content_blocks[-count:] != expected_blocks: return None - return content[:-count], sources, deepcopy(expected) + return content_blocks[:-count], sources, deepcopy(expected_blocks) return None @@ -194,8 +200,8 @@ def reattach_runtime_context( "suffix": suffix, } - visible_blocks = ( - [*content] + visible_blocks: list[Any] = ( + [*cast(list[Any], content)] if isinstance(content, list) else ([] if content is None else [{"type": "text", "text": str(content)}]) ) @@ -210,11 +216,14 @@ def public_history_message(message: Mapping[str, Any]) -> dict[str, Any]: """Return a user-visible copy with trusted runtime context removed exactly.""" cleaned = deepcopy(dict(message)) marker = cleaned.pop(RUNTIME_CONTEXT_HISTORY_META, None) - if not isinstance(marker, Mapping) or marker.get("version") != 1: + if not isinstance(marker, Mapping): + return cleaned + marker_data = cast(Mapping[str, Any], marker) + if marker_data.get("version") != 1: return cleaned content = cleaned.get("content") - suffix = marker.get("suffix") + suffix = marker_data.get("suffix") if isinstance(content, str) and isinstance(suffix, str) and suffix: if content == suffix: cleaned["content"] = "" @@ -222,10 +231,11 @@ def public_history_message(message: Mapping[str, Any]) -> dict[str, Any]: cleaned["content"] = content[: -(len(suffix) + 2)] return cleaned - expected = marker.get("blocks") + expected = marker_data.get("blocks") if isinstance(content, list) and isinstance(expected, list) and expected: - count = len(expected) - if content[-count:] == expected: + expected_blocks = cast(list[Any], expected) + count = len(expected_blocks) + if content[-count:] == expected_blocks: cleaned["content"] = content[:-count] return cleaned diff --git a/nanobot/sdk/clients.py b/nanobot/sdk/clients.py index c9f132290..a1d189af3 100644 --- a/nanobot/sdk/clients.py +++ b/nanobot/sdk/clients.py @@ -67,7 +67,7 @@ class SessionClient: def get(self, session_key: str) -> SessionSnapshot | None: """Return a display-safe snapshot without creating a new session on disk.""" - cached = self._loop.sessions._cached(session_key) + cached = self._loop.sessions.get_cached(session_key) if cached is not None: return snapshot_from_session(cached) payload = self._loop.sessions.read_session_file(session_key) @@ -91,7 +91,7 @@ class SessionClient: def export(self, session_key: str) -> SessionSnapshot | None: """Return a trusted full snapshot, including model-only runtime context.""" - cached = self._loop.sessions._cached(session_key) + cached = self._loop.sessions.get_cached(session_key) if cached is not None: return snapshot_from_session(cached, include_runtime_context=True) payload = self._loop.sessions.read_session_file(session_key) diff --git a/nanobot/sdk/streaming.py b/nanobot/sdk/streaming.py index b53bf53d0..1abfd64fc 100644 --- a/nanobot/sdk/streaming.py +++ b/nanobot/sdk/streaming.py @@ -59,6 +59,8 @@ class RunStream: if item is _STREAM_SENTINEL: self._events_done = True break + if not isinstance(item, StreamEvent): + raise TypeError("SDK event queue contained an invalid item") yield item finally: self._stream_active = False diff --git a/nanobot/sdk/types.py b/nanobot/sdk/types.py index 019b44f38..e28feef99 100644 --- a/nanobot/sdk/types.py +++ b/nanobot/sdk/types.py @@ -4,7 +4,7 @@ from __future__ import annotations from copy import deepcopy from dataclasses import dataclass, field -from typing import Any, Literal, Mapping, TypeAlias +from typing import Any, Literal, Mapping, TypeAlias, cast from nanobot.runtime_context import public_history_messages @@ -126,7 +126,7 @@ def snapshot_from_session( *, include_runtime_context: bool = False, ) -> SessionSnapshot: - messages = deepcopy(session.messages) + messages = cast(list[dict[str, Any]], deepcopy(session.messages)) if not include_runtime_context: messages = public_history_messages(messages) return SessionSnapshot( @@ -143,9 +143,10 @@ def snapshot_from_payload( *, include_runtime_context: bool = False, ) -> SessionSnapshot: - messages = [ - deepcopy(dict(message)) - for message in list(payload.get("messages") or []) + raw_messages: list[Any] = list(payload.get("messages") or []) + messages: list[dict[str, Any]] = [ + deepcopy(dict(cast(Mapping[str, Any], message))) + for message in raw_messages if isinstance(message, Mapping) ] if not include_runtime_context: @@ -154,7 +155,7 @@ def snapshot_from_payload( key=str(payload.get("key") or ""), created_at=payload.get("created_at"), updated_at=payload.get("updated_at"), - metadata=deepcopy(dict(payload.get("metadata") or {})), + metadata=deepcopy(dict(cast(Mapping[str, Any], payload.get("metadata") or {}))), messages=messages, ) diff --git a/nanobot/security/network.py b/nanobot/security/network.py index dba5e14ad..a8ffb9fa6 100644 --- a/nanobot/security/network.py +++ b/nanobot/security/network.py @@ -7,6 +7,7 @@ import ipaddress import re import socket from contextlib import contextmanager, suppress +from typing import Any, cast from urllib.parse import urlparse from urllib.request import getproxies, proxy_bypass @@ -45,7 +46,7 @@ def is_loopback_host(host: str) -> bool: def configure_ssrf_whitelist(cidrs: list[str]) -> None: """Allow specific CIDR ranges to bypass SSRF blocking (e.g. Tailscale's 100.64.0.0/10).""" global _allowed_networks - nets = [] + nets: list[ipaddress.IPv4Network | ipaddress.IPv6Network] = [] for cidr in cidrs: with suppress(ValueError): nets.append(ipaddress.ip_network(cidr, strict=False)) @@ -229,10 +230,17 @@ def pin_resolved_url_dns(url: str, resolved_ips: tuple[str, ...]): pinned_host = hostname.rstrip(".").lower() original_getaddrinfo = socket.getaddrinfo - def _getaddrinfo(host, port, family=0, type=0, proto=0, flags=0): # noqa: A002 + def _getaddrinfo( + host: Any, + port: Any, + family: int = 0, + type: int = 0, # noqa: A002 + proto: int = 0, + flags: int = 0, + ) -> list[Any]: if str(host).rstrip(".").lower() != pinned_host: return original_getaddrinfo(host, port, family, type, proto, flags) - infos = [] + infos: list[Any] = [] for ip in resolved_ips: addr = ipaddress.ip_address(ip) addr_family = socket.AF_INET6 if addr.version == 6 else socket.AF_INET @@ -242,7 +250,7 @@ def pin_resolved_url_dns(url: str, resolved_ips: tuple[str, ...]): infos.append((addr_family, type or socket.SOCK_STREAM, proto, "", sockaddr)) return infos - socket.getaddrinfo = _getaddrinfo + socket.getaddrinfo = cast(Any, _getaddrinfo) try: yield finally: diff --git a/nanobot/security/workspace_access.py b/nanobot/security/workspace_access.py index 59c54559d..72d0ae4b6 100644 --- a/nanobot/security/workspace_access.py +++ b/nanobot/security/workspace_access.py @@ -6,7 +6,7 @@ import os from contextvars import ContextVar, Token from dataclasses import dataclass from pathlib import Path -from typing import Any, Literal +from typing import Any, Literal, cast WorkspaceAccessMode = Literal["restricted", "full"] WORKSPACE_SCOPE_METADATA_KEY = "workspace_scope" @@ -158,9 +158,10 @@ class WorkspaceScopeResolver: metadata = getattr(msg, "metadata", None) if not isinstance(metadata, dict): return - raw = metadata.get(WORKSPACE_SCOPE_METADATA_KEY) + metadata_data = cast(dict[str, Any], metadata) + raw = metadata_data.get(WORKSPACE_SCOPE_METADATA_KEY) if isinstance(raw, dict): - session.metadata[WORKSPACE_SCOPE_METADATA_KEY] = dict(raw) + session.metadata[WORKSPACE_SCOPE_METADATA_KEY] = dict(cast(dict[str, Any], raw)) def workspace_sandbox_status( @@ -261,8 +262,9 @@ def validate_workspace_scope_payload( ) if not isinstance(raw, dict): raise WorkspaceScopeError("workspace_scope must be an object") + scope_data = cast(dict[str, Any], raw) - raw_path = raw.get("project_path") or raw.get("path") + raw_path = scope_data.get("project_path") or scope_data.get("path") if raw_path is None or raw_path == "": raw_path = str(Path(default_workspace).expanduser().resolve(strict=False)) if not isinstance(raw_path, str): @@ -277,7 +279,7 @@ def validate_workspace_scope_payload( if not project.is_dir(): raise WorkspaceScopeError("project_path must be an existing directory") - raw_mode = raw.get("access_mode") + raw_mode = scope_data.get("access_mode") if raw_mode is None: raw_mode = default_access_mode(default_restrict_to_workspace) if not isinstance(raw_mode, str): @@ -300,8 +302,9 @@ def workspace_scope_from_metadata( source_channel=source_channel, ) try: + metadata_data = cast(dict[str, Any], metadata) return validate_workspace_scope_payload( - metadata.get(WORKSPACE_SCOPE_METADATA_KEY), + metadata_data.get(WORKSPACE_SCOPE_METADATA_KEY), default_workspace=default_workspace, default_restrict_to_workspace=default_restrict_to_workspace, source_channel=source_channel, @@ -323,8 +326,9 @@ def resolve_effective_workspace_scope( source_channel: str | None = None, ) -> WorkspaceScope: if isinstance(message_metadata, dict) and WORKSPACE_SCOPE_METADATA_KEY in message_metadata: + message_metadata_data = cast(dict[str, Any], message_metadata) return workspace_scope_from_metadata( - message_metadata, + message_metadata_data, default_workspace=default_workspace, default_restrict_to_workspace=default_restrict_to_workspace, source_channel=source_channel, diff --git a/nanobot/security/workspace_policy.py b/nanobot/security/workspace_policy.py index a91cd8809..44758a6a9 100644 --- a/nanobot/security/workspace_policy.py +++ b/nanobot/security/workspace_policy.py @@ -108,7 +108,7 @@ def resolve_allowed_path( if allowed_root is None and not files: return resolve_path(path, workspace, strict=strict) if strict else resolved - roots = [] + roots: list[str | Path] = [] if allowed_root is not None: roots.append(allowed_root) roots.extend(extra_allowed_roots or []) diff --git a/nanobot/session/automation_turns.py b/nanobot/session/automation_turns.py index ebd73c579..336f9f491 100644 --- a/nanobot/session/automation_turns.py +++ b/nanobot/session/automation_turns.py @@ -5,7 +5,7 @@ from __future__ import annotations from collections.abc import Callable, Mapping from dataclasses import dataclass, field from functools import lru_cache -from typing import Any +from typing import Any, cast AUTOMATION_HISTORY_META = "_automation_turn" @@ -17,7 +17,7 @@ class AutomationTurnSpec: kind: str trigger_meta_key: str legacy_history_meta_key: str | None = None - history_fields: Mapping[str, str] = field(default_factory=dict) + history_fields: Mapping[str, str] = field(default_factory=dict[str, str]) text_builder: Callable[[Mapping[str, Any]], str | None] | None = None @@ -27,7 +27,7 @@ def automation_trigger( ) -> dict[str, Any] | None: """Return source trigger metadata for *spec* when present.""" raw = (metadata or {}).get(spec.trigger_meta_key) - return raw if isinstance(raw, dict) else None + return cast(dict[str, Any], raw) if isinstance(raw, dict) else None def automation_history_overrides_for_spec( diff --git a/nanobot/session/goal_state.py b/nanobot/session/goal_state.py index 184374060..4c2193fcf 100644 --- a/nanobot/session/goal_state.py +++ b/nanobot/session/goal_state.py @@ -8,7 +8,7 @@ for older sessions. Callers use ``goal_state_runtime_lines``, ``goal_state_ws_bl from __future__ import annotations import json -from typing import Any, Mapping, MutableMapping +from typing import Any, Mapping, MutableMapping, cast from nanobot.session.manager import SessionManager @@ -66,13 +66,13 @@ def parse_goal_state(blob: Any) -> dict[str, Any] | None: if blob is None: return None if isinstance(blob, dict): - return blob + return cast(dict[str, Any], blob) if isinstance(blob, str): try: parsed = json.loads(blob) except json.JSONDecodeError: return None - return parsed if isinstance(parsed, dict) else None + return cast(dict[str, Any], parsed) if isinstance(parsed, dict) else None return None diff --git a/nanobot/session/manager.py b/nanobot/session/manager.py index 98ccce75e..1d8c57dac 100644 --- a/nanobot/session/manager.py +++ b/nanobot/session/manager.py @@ -11,7 +11,7 @@ from copy import deepcopy from dataclasses import dataclass, field from datetime import datetime from pathlib import Path -from typing import Any, Callable +from typing import Any, Callable, cast from weakref import WeakValueDictionary from loguru import logger @@ -53,6 +53,13 @@ _FORK_VOLATILE_METADATA_KEYS = { } +def _json_object(value: object) -> dict[str, Any]: + """Narrow a decoded JSON object while preserving its original values.""" + if not isinstance(value, dict): + raise ValueError("session records must be JSON objects") + return cast(dict[str, Any], value) + + def replay_max_messages_for_context(context_window_tokens: int | None) -> int: if not context_window_tokens or context_window_tokens <= 0: return FILE_MAX_MESSAGES @@ -78,15 +85,18 @@ def _sanitize_assistant_replay_text(content: str) -> str: return "\n".join(lines).strip() -def _text_preview(content: Any) -> str: +def _text_preview(content: object) -> str: """Return compact display text for session lists.""" if isinstance(content, str): text = content elif isinstance(content, list): parts: list[str] = [] - for block in content: - if isinstance(block, dict) and block.get("type") == "text": - value = block.get("text") + for block in cast(list[object], content): + if isinstance(block, dict): + block_data = cast(dict[object, object], block) + if block_data.get("type") != "text": + continue + value = block_data.get("text") if isinstance(value, str): parts.append(value) text = " ".join(parts) @@ -102,26 +112,27 @@ def _text_preview(content: Any) -> str: def _message_preview_text(message: dict[str, Any]) -> str: """Session list preview text; subagent inject blobs are shortened for display.""" message = public_history_message(message) - content: Any = message.get("content") + content = cast(object, message.get("content")) if message.get("injected_event") == "subagent_result" and isinstance(content, str): content = scrub_subagent_announce_body(content) return _text_preview(content) -def _metadata_title(metadata: Any) -> str: +def _metadata_title(metadata: object) -> str: if not isinstance(metadata, dict): return "" - title = metadata.get("title") + metadata_data = cast(dict[object, object], metadata) + title = metadata_data.get("title") if not isinstance(title, str): return "" - if metadata.get("title_user_edited") is True: + if metadata_data.get("title_user_edited") is True: return title return strip_think(title) @dataclass class RetentionResult: - dropped: list[dict] + dropped: list[dict[str, Any]] already_consolidated_count: int @@ -137,13 +148,14 @@ class Session: last_consolidated: int = 0 # Number of messages already consolidated to files def __post_init__(self) -> None: - if not isinstance(self.metadata, dict): + if not isinstance(cast(object, self.metadata), dict): self.metadata = {} # An out-of-range offset (corrupt metadata) would hide all history; reset it. + last_consolidated = cast(object, self.last_consolidated) if ( - isinstance(self.last_consolidated, bool) - or not isinstance(self.last_consolidated, int) - or not 0 <= self.last_consolidated <= len(self.messages) + isinstance(last_consolidated, bool) + or not isinstance(last_consolidated, int) + or not 0 <= last_consolidated <= len(self.messages) ): self.last_consolidated = 0 @@ -219,7 +231,7 @@ class Session: content, message.get("media"), ) - cli_apps = message.get("cli_apps") + cli_apps = cast(object, message.get("cli_apps")) if ( include_runtime_context and not has_persisted_runtime_context @@ -229,15 +241,18 @@ class Session: and isinstance(content, str) ): cli_lines: list[str] = [] - for item in cli_apps[:8]: + for item in cast(list[object], cli_apps[:8]): if not isinstance(item, dict): continue - name = str(item.get("name") or "").strip().lower() + item_data = cast(dict[object, object], item) + name = str(item_data.get("name") or "").strip().lower() if not name: continue - entry = str(item.get("entry_point") or "unknown").strip() or "unknown" + entry_point = ( + str(item_data.get("entry_point") or "unknown").strip() or "unknown" + ) cli_lines.append( - f"[CLI App Attachment: @{name}; tool=run_cli_app; entry_point={entry}; " + f"[CLI App Attachment: @{name}; tool=run_cli_app; entry_point={entry_point}; " f"skill=skills/cli-app-{name}/SKILL.md]" ) if cli_lines: @@ -389,7 +404,7 @@ class Session: def enforce_file_cap( self, - on_archive: Any = None, + on_archive: Callable[[list[dict[str, Any]]], None] | None = None, limit: int = FILE_MAX_MESSAGES, ) -> None: """Bound session message growth by archiving and trimming old prefixes.""" @@ -449,6 +464,10 @@ class SessionManager: self._remember(session) return session + def get_cached(self, key: str) -> Session | None: + """Return a cached session without creating or loading one from disk.""" + return self._cached(key) + def set_file_cap_archiver(self, archiver: Callable[..., None]) -> None: """Archive unconsolidated overflow whenever a session is persisted.""" self._file_cap_archiver = archiver @@ -475,6 +494,11 @@ class SessionManager: except _SESSION_DATA_ERRORS: return None + @staticmethod + def decode_storage_key(stem: str) -> str | None: + """Public decoder for components that inspect canonical session filenames.""" + return SessionManager._decode_storage_key(stem) + @classmethod def _session_key_from_path(cls, path: Path) -> str | None: """Decode a session key only from a canonical collision-resistant filename.""" @@ -523,11 +547,11 @@ class SessionManager: return None try: - messages = [] - metadata = {} - created_at = None - updated_at = None - last_consolidated = 0 + messages: list[dict[str, Any]] = [] + metadata: object = {} + created_at: datetime | None = None + updated_at: datetime | None = None + last_consolidated: object = 0 with open(path, encoding="utf-8") as f: for line in f: @@ -535,15 +559,27 @@ class SessionManager: if not line: continue - data = json.loads(line) - if not isinstance(data, dict): - raise ValueError("session records must be JSON objects") + raw_data: object = json.loads(line) + data = _json_object(raw_data) if data.get("_type") == "metadata": - metadata = data.get("metadata", {}) - created_at = datetime.fromisoformat(data["created_at"]) if data.get("created_at") else None - updated_at = datetime.fromisoformat(data["updated_at"]) if data.get("updated_at") else None - last_consolidated = data.get("last_consolidated", 0) + metadata = cast(object, data.get("metadata", {})) + created_at_value = cast(object, data.get("created_at")) + updated_at_value = cast(object, data.get("updated_at")) + created_at = ( + datetime.fromisoformat(cast(str, created_at_value)) + if created_at_value + else None + ) + updated_at = ( + datetime.fromisoformat(cast(str, updated_at_value)) + if updated_at_value + else None + ) + last_consolidated = cast( + object, + data.get("last_consolidated", 0), + ) else: messages.append(data) @@ -552,8 +588,8 @@ class SessionManager: messages=messages, created_at=created_at or datetime.now(), updated_at=updated_at or datetime.now(), - metadata=metadata, - last_consolidated=last_consolidated + metadata=cast(dict[str, Any], metadata), + last_consolidated=cast(int, last_consolidated), ) except _SESSION_DATA_ERRORS as e: logger.warning("Failed to load session {}: {}", key, e) @@ -571,10 +607,10 @@ class SessionManager: try: messages: list[dict[str, Any]] = [] - metadata: dict[str, Any] = {} + metadata: object = {} created_at: datetime | None = None updated_at: datetime | None = None - last_consolidated = 0 + last_consolidated: object = 0 skipped = 0 with open(path, encoding="utf-8") as f: @@ -583,23 +619,33 @@ class SessionManager: if not line: continue try: - data = json.loads(line) + raw_data: object = json.loads(line) except json.JSONDecodeError: skipped += 1 continue - if not isinstance(data, dict): + if not isinstance(raw_data, dict): skipped += 1 continue + data = cast(dict[str, Any], raw_data) if data.get("_type") == "metadata": - metadata = data.get("metadata", {}) - if data.get("created_at"): + metadata = cast(object, data.get("metadata", {})) + created_at_value = cast(object, data.get("created_at")) + if created_at_value: with suppress(ValueError, TypeError): - created_at = datetime.fromisoformat(data["created_at"]) - if data.get("updated_at"): + created_at = datetime.fromisoformat( + cast(str, created_at_value) + ) + updated_at_value = cast(object, data.get("updated_at")) + if updated_at_value: with suppress(ValueError, TypeError): - updated_at = datetime.fromisoformat(data["updated_at"]) - last_consolidated = data.get("last_consolidated", 0) + updated_at = datetime.fromisoformat( + cast(str, updated_at_value) + ) + last_consolidated = cast( + object, + data.get("last_consolidated", 0), + ) else: messages.append(data) @@ -614,8 +660,8 @@ class SessionManager: messages=messages, created_at=created_at or datetime.now(), updated_at=updated_at or datetime.now(), - metadata=metadata, - last_consolidated=last_consolidated + metadata=cast(dict[str, Any], metadata), + last_consolidated=cast(int, last_consolidated), ) except _SESSION_DATA_ERRORS as e: logger.warning("Repair failed for session {}: {}", key, e) @@ -641,9 +687,10 @@ class SessionManager: write-back caching (e.g. rclone VFS, NFS, FUSE mounts) do not lose the most recent writes. """ - if self._file_cap_archiver is not None: + archiver = self._file_cap_archiver + if archiver is not None: session.enforce_file_cap( - on_archive=lambda messages: self._file_cap_archiver( + on_archive=lambda messages: archiver( messages, session_key=session.key, ) @@ -803,21 +850,22 @@ class SessionManager: return None try: messages: list[dict[str, Any]] = [] - metadata: dict[str, Any] = {} - created_at: str | None = None - updated_at: str | None = None - stored_key: str | None = None + metadata: object = {} + created_at: object = None + updated_at: object = None + stored_key: object = None with open(path, encoding="utf-8") as f: for line in f: line = line.strip() if not line: continue - data = json.loads(line) + raw_data: object = json.loads(line) + data = _json_object(raw_data) if data.get("_type") == "metadata": - metadata = data.get("metadata", {}) - created_at = data.get("created_at") - updated_at = data.get("updated_at") - stored_key = data.get("key") + metadata = cast(object, data.get("metadata", {})) + created_at = cast(object, data.get("created_at")) + updated_at = cast(object, data.get("updated_at")) + stored_key = cast(object, data.get("key")) else: messages.append(data) return { @@ -850,17 +898,20 @@ class SessionManager: line = line.strip() if not line: continue - data = json.loads(line) - if not isinstance(data, dict): - raise ValueError("session records must be JSON objects") + raw_data: object = json.loads(line) + data = _json_object(raw_data) if data.get("_type") != "metadata": return None - metadata = data.get("metadata", {}) + metadata = cast(object, data.get("metadata", {})) return { "key": data.get("key") or key, "created_at": data.get("created_at"), "updated_at": data.get("updated_at"), - "metadata": metadata if isinstance(metadata, dict) else {}, + "metadata": ( + cast(dict[str, Any], metadata) + if isinstance(metadata, dict) + else {} + ), } return None except _SESSION_DATA_ERRORS as e: @@ -883,7 +934,7 @@ class SessionManager: Returns: List of session info dicts. """ - sessions = [] + sessions: list[dict[str, Any]] = [] for path in self.sessions_dir.glob("*.jsonl"): storage_key = self._session_key_from_path(path) @@ -894,12 +945,11 @@ class SessionManager: with open(path, encoding="utf-8") as f: first_line = f.readline().strip() if first_line: - data = json.loads(first_line) - if not isinstance(data, dict): - raise ValueError("session records must be JSON objects") + raw_data: object = json.loads(first_line) + data = _json_object(raw_data) if data.get("_type") == "metadata": - key = data.get("key") or storage_key - metadata = data.get("metadata", {}) + key = cast(object, data.get("key")) or storage_key + metadata = cast(object, data.get("metadata", {})) title = _metadata_title(metadata) preview = "" fallback_preview = "" @@ -915,9 +965,8 @@ class SessionManager: or scanned_chars > _SESSION_LIST_PREVIEW_MAX_CHARS ): break - item = json.loads(line) - if not isinstance(item, dict): - raise ValueError("session records must be JSON objects") + raw_item: object = json.loads(line) + item = _json_object(raw_item) if item.get("_type") == "metadata": continue text = _message_preview_text(item) @@ -963,4 +1012,8 @@ class SessionManager: } ) continue - return sorted(sessions, key=lambda x: x.get("updated_at", ""), reverse=True) + return sorted( + sessions, + key=lambda item: cast(str, item.get("updated_at", "")), + reverse=True, + ) diff --git a/nanobot/session/model_selection.py b/nanobot/session/model_selection.py index bf5be146b..225376c6c 100644 --- a/nanobot/session/model_selection.py +++ b/nanobot/session/model_selection.py @@ -3,6 +3,7 @@ from __future__ import annotations from collections.abc import Mapping +from typing import cast # Session.metadata is public SDK data, so internal selectors use a reserved namespace. SESSION_MODEL_PRESET_METADATA_KEY = "_nanobot_model_preset" @@ -12,9 +13,10 @@ def model_preset_from_metadata(metadata: object) -> str | None: """Read the canonical session preset name from persisted metadata.""" if not isinstance(metadata, Mapping): return None - if SESSION_MODEL_PRESET_METADATA_KEY not in metadata: + typed_metadata = cast(Mapping[object, object], metadata) + if SESSION_MODEL_PRESET_METADATA_KEY not in typed_metadata: return None - value = metadata[SESSION_MODEL_PRESET_METADATA_KEY] + value = typed_metadata[SESSION_MODEL_PRESET_METADATA_KEY] if not isinstance(value, str) or not value.strip(): raise ValueError("session model preset must be a non-empty string") return value.strip() diff --git a/nanobot/session/turn_continuation.py b/nanobot/session/turn_continuation.py index 97d9c941f..cc98ea08e 100644 --- a/nanobot/session/turn_continuation.py +++ b/nanobot/session/turn_continuation.py @@ -8,7 +8,7 @@ continuation is allowed and, when it is, queue the next turn directly. from __future__ import annotations import dataclasses -from typing import Any, Mapping, MutableMapping +from typing import TYPE_CHECKING, Any, Mapping, MutableMapping from loguru import logger @@ -18,6 +18,9 @@ from nanobot.session.goal_state import ( sustained_goal_turn, ) +if TYPE_CHECKING: + from nanobot.agent.loop import TurnContext + INTERNAL_CONTINUATION_META = "_internal_continuation" INTERNAL_CONTINUATION_KIND_META = "_internal_continuation_kind" INTERNAL_CONTINUATION_PENDING_META = "_internal_continuation_pending" @@ -101,7 +104,7 @@ def should_finalize_on_max_iterations( ) -async def maybe_continue_turn(ctx: Any) -> bool: +async def maybe_continue_turn(ctx: TurnContext) -> bool: """Queue an internal continuation for *ctx* when policy allows it.""" if ctx.session is None or ctx.pending_queue is None: return False @@ -115,7 +118,7 @@ async def maybe_continue_turn(ctx: Any) -> bool: metadata = _internal_continuation_metadata( ctx.msg.metadata, - run_started_at=getattr(ctx, "visible_run_started_at", None), + run_started_at=ctx.visible_run_started_at, ) content = _goal_continuation_prompt(ctx.session.metadata) messages = _strip_terminal_assistant(ctx.all_messages, ctx.final_content) @@ -139,7 +142,7 @@ async def maybe_continue_turn(ctx: Any) -> bool: return True -def prepare_save_boundary(ctx: Any) -> None: +def prepare_save_boundary(ctx: TurnContext) -> None: """Prepare continuation bookkeeping and the history append boundary.""" if ctx.session is not None: clear_internal_continuation_state(ctx.session.metadata) diff --git a/nanobot/session/webui_turns.py b/nanobot/session/webui_turns.py index 37bfc731e..0e526891e 100644 --- a/nanobot/session/webui_turns.py +++ b/nanobot/session/webui_turns.py @@ -73,6 +73,11 @@ class _WebsocketTurn: _WEBSOCKET_ACTIVE_TURNS: dict[str, dict[str, _WebsocketTurn]] = {} +def _validated_llm_runtime(value: object) -> LLMRuntime | None: + """Keep runtime-event consumers defensive if an external publisher violates the contract.""" + return value if isinstance(value, LLMRuntime) else None + + def _sync_websocket_turn_projection(chat_id: str) -> None: turns = _WEBSOCKET_ACTIVE_TURNS.get(chat_id) if not turns: @@ -599,11 +604,10 @@ class WebuiTurnCoordinator: ) def _schedule_title_update_from_event(self, event: TurnCompleted) -> None: - title_context = event.runtime + title_context = _validated_llm_runtime(event.runtime) if ( event.context.metadata.get("webui") is not True or title_context is None - or not isinstance(title_context, LLMRuntime) ): return diff --git a/nanobot/skills/skill-creator/scripts/init_skill.py b/nanobot/skills/skill-creator/scripts/init_skill.py index 8633fe9e3..e44addbf7 100755 --- a/nanobot/skills/skill-creator/scripts/init_skill.py +++ b/nanobot/skills/skill-creator/scripts/init_skill.py @@ -191,7 +191,7 @@ Note: This is a text placeholder. Actual assets can be any file type. """ -def normalize_skill_name(skill_name): +def normalize_skill_name(skill_name: str) -> str: """Normalize a skill name to lowercase hyphen-case.""" normalized = skill_name.strip().lower() normalized = re.sub(r"[^a-z0-9]+", "-", normalized) @@ -200,12 +200,12 @@ def normalize_skill_name(skill_name): return normalized -def title_case_skill_name(skill_name): +def title_case_skill_name(skill_name: str) -> str: """Convert hyphenated skill name to Title Case for display.""" return " ".join(word.capitalize() for word in skill_name.split("-")) -def parse_resources(raw_resources): +def parse_resources(raw_resources: str) -> list[str]: if not raw_resources: return [] resources = [item.strip() for item in raw_resources.split(",") if item.strip()] @@ -215,8 +215,8 @@ def parse_resources(raw_resources): print(f"[ERROR] Unknown resource type(s): {', '.join(invalid)}") print(f" Allowed: {allowed}") sys.exit(1) - deduped = [] - seen = set() + deduped: list[str] = [] + seen: set[str] = set() for resource in resources: if resource not in seen: deduped.append(resource) @@ -224,7 +224,13 @@ def parse_resources(raw_resources): return deduped -def create_resource_dirs(skill_dir, skill_name, skill_title, resources, include_examples): +def create_resource_dirs( + skill_dir: Path, + skill_name: str, + skill_title: str, + resources: list[str], + include_examples: bool, +) -> None: for resource in resources: resource_dir = skill_dir / resource resource_dir.mkdir(exist_ok=True) @@ -252,7 +258,12 @@ def create_resource_dirs(skill_dir, skill_name, skill_title, resources, include_ print("[OK] Created assets/") -def init_skill(skill_name, path, resources, include_examples): +def init_skill( + skill_name: str, + path: str | Path, + resources: list[str], + include_examples: bool, +) -> Path | None: """ Initialize a new skill directory with template SKILL.md. @@ -317,7 +328,7 @@ def init_skill(skill_name, path, resources, include_examples): return skill_dir -def main(): +def main() -> None: parser = argparse.ArgumentParser( description="Create a new skill directory with a SKILL.md template.", ) diff --git a/nanobot/skills/skill-creator/scripts/package_skill.py b/nanobot/skills/skill-creator/scripts/package_skill.py index 1494b10e4..67b898362 100755 --- a/nanobot/skills/skill-creator/scripts/package_skill.py +++ b/nanobot/skills/skill-creator/scripts/package_skill.py @@ -31,7 +31,7 @@ def _cleanup_partial_archive(skill_filename: Path) -> None: skill_filename.unlink() -def package_skill(skill_path, output_dir=None): +def package_skill(skill_path: str | Path, output_dir: str | Path | None = None) -> Path | None: """ Package a skill folder into a .skill file. @@ -80,7 +80,7 @@ def package_skill(skill_path, output_dir=None): excluded_dirs = {".git", ".svn", ".hg", "__pycache__", "node_modules"} - files_to_package = [] + files_to_package: list[Path] = [] resolved_archive = skill_filename.resolve() for file_path in skill_path.rglob("*"): @@ -124,7 +124,7 @@ def package_skill(skill_path, output_dir=None): return None -def main(): +def main() -> None: if len(sys.argv) < 2: print("Usage: python package_skill.py [output-directory]") print("\nExample:") diff --git a/nanobot/skills/skill-creator/scripts/quick_validate.py b/nanobot/skills/skill-creator/scripts/quick_validate.py index 03d246d6e..e7953762f 100644 --- a/nanobot/skills/skill-creator/scripts/quick_validate.py +++ b/nanobot/skills/skill-creator/scripts/quick_validate.py @@ -6,7 +6,7 @@ Minimal validator for nanobot skill folders. import re import sys from pathlib import Path -from typing import Optional +from typing import Any, Optional, cast try: import yaml @@ -83,7 +83,7 @@ def _parse_simple_frontmatter(frontmatter_text: str) -> Optional[dict[str, str]] return parsed -def _load_frontmatter(frontmatter_text: str) -> tuple[Optional[dict], Optional[str]]: +def _load_frontmatter(frontmatter_text: str) -> tuple[dict[str, Any] | None, str | None]: if yaml is not None: try: frontmatter = yaml.safe_load(frontmatter_text) @@ -91,7 +91,7 @@ def _load_frontmatter(frontmatter_text: str) -> tuple[Optional[dict], Optional[s return None, f"Invalid YAML in frontmatter: {exc}" if not isinstance(frontmatter, dict): return None, "Frontmatter must be a YAML dictionary" - return frontmatter, None + return cast(dict[str, Any], frontmatter), None frontmatter = _parse_simple_frontmatter(frontmatter_text) if frontmatter is None: @@ -129,7 +129,7 @@ def _validate_description(description: str) -> Optional[str]: return None -def validate_skill(skill_path): +def validate_skill(skill_path: str | Path) -> tuple[bool, str]: """Validate a skill folder structure and required frontmatter.""" skill_path = Path(skill_path).resolve() @@ -152,8 +152,8 @@ def validate_skill(skill_path): return False, "Invalid frontmatter format" frontmatter, error = _load_frontmatter(frontmatter_text) - if error: - return False, error + if error or frontmatter is None: + return False, error or "Invalid frontmatter" unexpected_keys = sorted(set(frontmatter.keys()) - ALLOWED_FRONTMATTER_KEYS) if unexpected_keys: diff --git a/nanobot/triggers/local_store.py b/nanobot/triggers/local_store.py index 977143fa5..83e1094ff 100644 --- a/nanobot/triggers/local_store.py +++ b/nanobot/triggers/local_store.py @@ -10,7 +10,7 @@ import time import uuid from contextlib import suppress from pathlib import Path -from typing import Any +from typing import Any, cast from filelock import FileLock from loguru import logger @@ -204,7 +204,7 @@ class LocalTriggerStore: logger.exception("Trigger: failed to parse delivery {}", path) self._move_bad_delivery_unlocked(path) continue - os.replace(path, delivery.path) + os.replace(path, cast(Path, delivery.path)) claimed.append(delivery) return claimed @@ -311,9 +311,11 @@ class LocalTriggerStore: return [] try: data = json.loads(self.store_path.read_text(encoding="utf-8")) + store_data = cast(dict[str, Any], data) + raw_triggers = cast(list[Any], store_data.get("triggers", [])) return [ - LocalTrigger.from_dict(raw) - for raw in data.get("triggers", []) + LocalTrigger.from_dict(cast(dict[str, Any], raw)) + for raw in raw_triggers if isinstance(raw, dict) ] except Exception as exc: @@ -386,10 +388,12 @@ class LocalTriggerStore: data = json.loads(path.read_text(encoding="utf-8")) except Exception: return None - raw = data.get("delivery", data) if isinstance(data, dict) else None + payload = cast(dict[str, Any], data) if isinstance(data, dict) else None + raw = payload.get("delivery", payload) if payload is not None else None if not isinstance(raw, dict): return None - trigger_id = raw.get("triggerId", raw.get("trigger_id", "")) + delivery_data = cast(dict[str, Any], raw) + trigger_id = delivery_data.get("triggerId", delivery_data.get("trigger_id", "")) return str(trigger_id) if trigger_id else None @staticmethod diff --git a/nanobot/triggers/local_types.py b/nanobot/triggers/local_types.py index 9dc890ad6..4dda5e9cb 100644 --- a/nanobot/triggers/local_types.py +++ b/nanobot/triggers/local_types.py @@ -4,7 +4,7 @@ from __future__ import annotations from dataclasses import dataclass, field from pathlib import Path -from typing import Any, Literal +from typing import Any, Literal, cast from nanobot.utils.dict_keys import get_camel_snake as _get @@ -68,9 +68,14 @@ class LocalTrigger: @classmethod def from_dict(cls, data: dict[str, Any]) -> "LocalTrigger": - raw_history = data.get("runHistory", data.get("run_history", [])) or [] - history = [ - record if isinstance(record, TriggerRunRecord) else TriggerRunRecord.from_dict(record) + raw_history = cast( + list[Any], + data.get("runHistory", data.get("run_history", [])) or [], + ) + history: list[TriggerRunRecord] = [ + record + if isinstance(record, TriggerRunRecord) + else TriggerRunRecord.from_dict(cast(dict[str, Any], record)) for record in raw_history if isinstance(record, (dict, TriggerRunRecord)) ] diff --git a/nanobot/utils/document.py b/nanobot/utils/document.py index 632271da5..6d27fadff 100644 --- a/nanobot/utils/document.py +++ b/nanobot/utils/document.py @@ -4,6 +4,7 @@ import mimetypes from collections.abc import Iterator from dataclasses import dataclass from pathlib import Path +from typing import Any from zipfile import BadZipFile, ZipFile from loguru import logger @@ -102,7 +103,7 @@ class PdfExtraction: end_page: int -def extract_text(path: Path) -> str | None: +def extract_text(path: str | Path) -> str | None: """Extract text from a file. Args: @@ -112,9 +113,7 @@ def extract_text(path: Path) -> str | None: Extracted text as string, None for unsupported types, or error string for failures. """ - if not isinstance(path, Path): - path = Path(path) - + path = Path(path) if not path.exists(): return f"[error: file not found: {path}]" try: @@ -217,14 +216,14 @@ def _extract_docx(path: Path) -> str: """Extract text from DOCX using python-docx.""" try: from docx import Document as DocxDocument - from docx.table import Table, _Cell + from docx.table import Table, _Cell # pyright: ignore[reportPrivateUsage] from docx.text.paragraph import Paragraph except ImportError: return "[error: python-docx not installed]" try: if error := _office_archive_error(path): return error - doc = DocxDocument(path) + doc = DocxDocument(str(path)) collector = _TextCollector(_MAX_TEXT_LENGTH) table_cell_count = 0 @@ -235,7 +234,7 @@ def _extract_docx(path: Path) -> str: text = " ".join(block.text.split()) if text: parts.append(text) - elif isinstance(block, Table): + elif isinstance(block, Table): # pyright: ignore[reportUnnecessaryIsInstance] parts.extend(row.replace("\t", " | ") for row in table_rows(block, depth + 1)) return " ".join(parts) @@ -249,7 +248,7 @@ def _extract_docx(path: Path) -> str: cells: list[str] = [] # row.cells expands w:gridSpan before callers can apply a bound. # Physical w:tc elements keep malformed documents proportional to XML size. - for tc in row._tr.tc_lst: + for tc in row._tr.tc_lst: # pyright: ignore[reportPrivateUsage] table_cell_count += 1 if table_cell_count > _MAX_DOCX_TABLE_CELLS: raise DocxSafetyError( @@ -265,7 +264,7 @@ def _extract_docx(path: Path) -> str: if text and not collector.add(text, separator="\n\n"): break continue - if not isinstance(block, Table): + if not isinstance(block, Table): # pyright: ignore[reportUnnecessaryIsInstance] continue first_row = True for row_text in table_rows(block, 1): @@ -325,7 +324,7 @@ def _extract_pptx(path: Path) -> str: try: if error := _office_archive_error(path): return error - prs = PptxPresentation(path) + prs = PptxPresentation(str(path)) collector = _TextCollector(_MAX_TEXT_LENGTH) for i, slide in enumerate(prs.slides, 1): slide_text: list[str] = [] @@ -343,7 +342,7 @@ def _extract_pptx(path: Path) -> str: return f"[error: failed to extract PPTX: {e!s}]" -def _collect_pptx_shape_text(shape, out: list[str]) -> None: +def _collect_pptx_shape_text(shape: Any, out: list[str]) -> None: """Collect text from a PPTX shape, recursing into groups and tables. Groups have ``has_text_frame=False`` and must be walked via ``.shapes``; diff --git a/nanobot/utils/file_edit_events.py b/nanobot/utils/file_edit_events.py index 8ec6b3fdb..e6740df8c 100644 --- a/nanobot/utils/file_edit_events.py +++ b/nanobot/utils/file_edit_events.py @@ -6,7 +6,7 @@ import difflib import re from dataclasses import dataclass from pathlib import Path -from typing import Any +from typing import Any, cast TRACKED_FILE_EDIT_TOOLS = frozenset({"write_file", "edit_file", "apply_patch"}) _MAX_SNAPSHOT_BYTES = 2 * 1024 * 1024 @@ -351,12 +351,13 @@ def _resolve_apply_patch_paths( edits = params.get("edits") if not isinstance(edits, list): return [] + patch_edits = cast(list[Any], edits) paths: list[Path] = [] seen: set[Path] = set() - for edit in edits: + for edit in patch_edits: if not isinstance(edit, dict): continue - raw_path = edit.get("path") + raw_path = cast(dict[str, Any], edit).get("path") if not isinstance(raw_path, str): continue raw_path = raw_path.strip() @@ -379,7 +380,7 @@ def _resolve_single_path(tool: Any, workspace: Path | None, raw_path: Any) -> Pa if isinstance(resolved, Path): return resolved if resolved: - return Path(resolved) + return Path(cast(str, resolved)) except Exception: return None resolver = getattr(tool, "_resolve", None) @@ -389,7 +390,7 @@ def _resolve_single_path(tool: Any, workspace: Path | None, raw_path: Any) -> Pa if isinstance(resolved, Path): return resolved if resolved: - return Path(resolved) + return Path(cast(str, resolved)) except Exception: return None if workspace is None: @@ -407,7 +408,7 @@ def _display_workspace(tool: Any, fallback: Path | None) -> Path | None: if isinstance(value, Path): return value if value: - return Path(value) + return Path(cast(str, value)) return fallback diff --git a/nanobot/utils/gitstore.py b/nanobot/utils/gitstore.py index 51f3c9b69..01ab1fb9a 100644 --- a/nanobot/utils/gitstore.py +++ b/nanobot/utils/gitstore.py @@ -7,9 +7,15 @@ import time from dataclasses import dataclass from datetime import datetime, timezone from pathlib import Path +from typing import TYPE_CHECKING, Iterable, cast from loguru import logger +if TYPE_CHECKING: + from dulwich.objects import Blob, Commit, ObjectID, Tree, TreeEntry + from dulwich.refs import Ref + from dulwich.repo import Repo + # Cap on the unified-diff block embedded in Dream commit messages. Memory files # are tiny in practice, but a pathological rewrite must not blow up the audit # record. The structured per-file summary is always emitted in full regardless. @@ -46,7 +52,9 @@ class LineAge: age_days: int # days since last modification -def _compute_line_ages(annotated) -> list[LineAge]: +def _compute_line_ages( + annotated: Iterable[tuple[tuple["Commit", "TreeEntry"], bytes]], +) -> list[LineAge]: """Convert annotate results to per-line ages.""" now = datetime.now(tz=timezone.utc).date() ages: list[LineAge] = [] @@ -148,10 +156,17 @@ class GitStore: # .gitignore excludes everything except tracked files, # so any staged/unstaged change must be in our files. st = porcelain.status(str(self._workspace)) - if not st.unstaged and not any(st.staged.values()): + unstaged = cast(list[object], st.unstaged) + staged = cast(dict[object, list[object]], st.staged) + if not unstaged and not any(staged.values()): return None - msg_bytes = message.encode("utf-8") if isinstance(message, str) else message + message_value = cast(object, message) + msg_bytes = ( + message_value.encode("utf-8") + if isinstance(message_value, str) + else cast(bytes, message_value) + ) porcelain.add(str(self._workspace), paths=self._staging_paths(*self._tracked_files)) sha_bytes = porcelain.commit( str(self._workspace), @@ -159,7 +174,7 @@ class GitStore: author=b"nanobot ", committer=b"nanobot ", ) - if sha_bytes is None: + if cast(object, sha_bytes) is None: return None sha = sha_bytes.hex()[:8] logger.debug("Git auto-commit: {} ({})", sha, message) @@ -180,16 +195,17 @@ class GitStore: with Repo(str(self._workspace)) as repo: try: - sha = repo.refs[b"HEAD"] + sha: ObjectID | None = repo.refs[cast("Ref", b"HEAD")] except KeyError: return None while sha: if sha.hex().startswith(short_sha): return sha - commit = repo[sha] - if commit.type_name != b"commit": + commit_obj = repo[sha] + if commit_obj.type_name != b"commit": break + commit = cast("Commit", commit_obj) sha = commit.parents[0] if commit.parents else None return None except Exception as exc: @@ -247,15 +263,16 @@ class GitStore: entries: list[CommitInfo] = [] with Repo(str(self._workspace)) as repo: try: - head = repo.refs[b"HEAD"] + head = repo.refs[cast("Ref", b"HEAD")] except KeyError: return [] - sha = head + sha: ObjectID | None = head while sha and len(entries) < max_entries: - commit = repo[sha] - if commit.type_name != b"commit": + commit_obj = repo[sha] + if commit_obj.type_name != b"commit": break + commit = cast("Commit", commit_obj) ts = time.strftime( "%Y-%m-%d %H:%M", time.localtime(commit.commit_time), @@ -429,16 +446,17 @@ class GitStore: return body @staticmethod - def _head_tree(repo) -> object | None: + def _head_tree(repo: "Repo") -> "Tree | None": """Return the tree object at HEAD, or None if there are no commits.""" try: - head = repo.refs[b"HEAD"] + head = repo.refs[cast("Ref", b"HEAD")] except KeyError: return None - commit = repo[head] - if commit.type_name != b"commit": + commit_obj = repo[head] + if commit_obj.type_name != b"commit": return None - return repo[commit.tree] + commit = cast("Commit", commit_obj) + return cast("Tree", repo[commit.tree]) def find_commit(self, short_sha: str, max_entries: int = 20) -> CommitInfo | None: """Find a commit by short SHA prefix match.""" @@ -464,7 +482,7 @@ class GitStore: if not full_sha: return None with Repo(str(self._workspace)) as repo: - commit = repo[full_sha] + commit = cast("Commit", repo[full_sha]) parent = commit.parents[0] if commit.parents else None diff = self.diff_commits(parent.hex()[:8], c.sha) if parent else "" return c, diff @@ -500,8 +518,12 @@ class GitStore: commit_obj = repo[full_sha] if commit_obj.type_name != b"commit": return None + typed_commit = cast("Commit", commit_obj) - commit_message = commit_obj.message.decode("utf-8", errors="replace").strip() + commit_message = typed_commit.message.decode( + "utf-8", + errors="replace", + ).strip() if message_prefix is not None and not commit_message.startswith(message_prefix): logger.warning( "Git revert: commit {} does not match message prefix {!r}", @@ -510,13 +532,13 @@ class GitStore: ) return None - if not commit_obj.parents: + if not typed_commit.parents: logger.warning("Git revert: cannot revert root commit {}", commit) return None # Use the parent's tree — this undoes the commit's changes - parent_obj = repo[commit_obj.parents[0]] - tree = repo[parent_obj.tree] + parent_obj = cast("Commit", repo[typed_commit.parents[0]]) + tree = cast("Tree", repo[parent_obj.tree]) restored: list[str] = [] for filepath in self._tracked_files: @@ -536,7 +558,11 @@ class GitStore: raise GitStoreError(f"Git revert failed for {commit}") from exc @staticmethod - def _read_blob_from_tree(repo, tree, filepath: str) -> str | None: + def _read_blob_from_tree( + repo: "Repo", + tree: "Tree", + filepath: str, + ) -> str | None: """Read a blob's content from a tree object by walking path parts.""" parts = Path(filepath).parts current = tree @@ -547,9 +573,10 @@ class GitStore: return None obj = repo[entry[1]] if obj.type_name == b"blob": - return obj.data.decode("utf-8", errors="replace") + blob = cast("Blob", obj) + return blob.data.decode("utf-8", errors="replace") if obj.type_name == b"tree": - current = obj + current = cast("Tree", obj) else: return None return None diff --git a/nanobot/utils/helpers.py b/nanobot/utils/helpers.py index b2ab7b495..8ab5f32bb 100644 --- a/nanobot/utils/helpers.py +++ b/nanobot/utils/helpers.py @@ -12,16 +12,25 @@ from contextlib import suppress from datetime import datetime from functools import lru_cache from pathlib import Path -from typing import Any +from typing import Any, TypeVar, cast, overload import tiktoken from loguru import logger _TOOLS_TOKEN_CACHE_MAX_ENTRIES = 64 _TOOLS_TOKEN_CACHE: dict[int, tuple[tuple[int, ...], dict[bool, int]]] = {} +_T = TypeVar("_T") -def sanitize_surrogates(text: str) -> str: +@overload +def sanitize_surrogates(text: str) -> str: ... + + +@overload +def sanitize_surrogates(text: _T) -> _T: ... + + +def sanitize_surrogates(text: Any) -> Any: """Reconstruct surrogate pairs and replace unpaired surrogates. Lone UTF-16 surrogate code points (``U+D800``..``U+DFFF``) cannot be @@ -62,24 +71,29 @@ def sanitize_surrogates_deep(value: Any) -> Any: if isinstance(value, list): result_list: list[Any] = [] mutated = False - for item in value: + for item in cast(list[Any], value): new_item = sanitize_surrogates_deep(item) if new_item is not item: mutated = True result_list.append(new_item) - return result_list if mutated else value + return result_list if mutated else cast(Any, value) if isinstance(value, dict): result_dict: dict[Any, Any] = {} mutated = False - for key, item in value.items(): + for key, item in cast(dict[Any, Any], value).items(): new_item = sanitize_surrogates_deep(item) if new_item is not item: mutated = True result_dict[key] = new_item - return result_dict if mutated else value + return result_dict if mutated else cast(Any, value) if isinstance(value, tuple): - result_tuple = tuple(sanitize_surrogates_deep(item) for item in value) - return result_tuple if any(a is not b for a, b in zip(result_tuple, value)) else value + tuple_value = cast(tuple[Any, ...], value) + result_tuple = tuple(sanitize_surrogates_deep(item) for item in tuple_value) + return ( + result_tuple + if any(a is not b for a, b in zip(result_tuple, tuple_value)) + else cast(Any, value) + ) return value @@ -289,7 +303,7 @@ def extract_reasoning( parts = [ strip_reasoning_tags(tb.get("thinking", "")) for tb in thinking_blocks - if isinstance(tb, dict) and tb.get("type") == "thinking" + if tb.get("type") == "thinking" ] joined = "\n\n".join(p for p in parts if p) return (joined or None), strip_think(content) if content else content @@ -377,7 +391,7 @@ def content_with_media_breadcrumbs( return content breadcrumbs = "\n".join( image_placeholder_text(path) - for path in media + for path in cast(list[object], media) if isinstance(path, str) and path ) if not breadcrumbs: @@ -468,9 +482,10 @@ def find_legal_message_start(messages: list[dict[str, Any]]) -> int: for i, 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 raw_call in cast(list[object], msg.get("tool_calls") or []): + tool_call = cast(dict[str, Any], raw_call) if isinstance(raw_call, dict) else None + if tool_call is not None and tool_call.get("id"): + declared.add(str(tool_call["id"])) elif role == "tool": tid = msg.get("tool_call_id") if tid and str(tid) not in declared: @@ -479,11 +494,12 @@ def find_legal_message_start(messages: list[dict[str, Any]]) -> int: return start -def stringify_text_blocks(content: list[dict[str, Any]]) -> str | None: +def stringify_text_blocks(content: list[object]) -> str | None: parts: list[str] = [] - for block in content: - if not isinstance(block, dict): + for raw_block in content: + if not isinstance(raw_block, dict): return None + block = cast(dict[str, Any], raw_block) if block.get("type") != "text": return None text = block.get("text") @@ -574,15 +590,15 @@ def maybe_persist_tool_result( if isinstance(content, str): text_payload = content elif isinstance(content, list): - text_payload = stringify_text_blocks(content) + text_payload = stringify_text_blocks(cast(list[object], content)) if text_payload is None: - return content + return cast(Any, content) suffix = "json" else: return content if len(text_payload) <= max_chars: - return content + return cast(Any, content) root = ensure_dir(workspace / _TOOL_RESULTS_DIR) bucket = ensure_dir(root / safe_filename(session_key or "default")) @@ -645,7 +661,7 @@ def build_assistant_message( content: str | None, tool_calls: list[dict[str, Any]] | None = None, reasoning_content: str | None = None, - thinking_blocks: list[dict] | None = None, + thinking_blocks: list[dict[str, Any]] | None = None, ) -> dict[str, Any]: """Build a provider-safe assistant message with optional reasoning fields.""" msg: dict[str, Any] = {"role": "assistant", "content": content or ""} @@ -677,11 +693,12 @@ def _estimate_prompt_tokens_with_source( if isinstance(content, str): parts.append(content) elif isinstance(content, list): - for part in content: - if isinstance(part, dict) and part.get("type") == "text": - txt = part.get("text", "") - if txt: - parts.append(txt) + for raw_part in cast(list[object], content): + part = cast(dict[str, Any], raw_part) if isinstance(raw_part, dict) else None + if part is not None and part.get("type") == "text": + text = part.get("text", "") + if isinstance(text, str) and text: + parts.append(text) tc = msg.get("tool_calls") if tc: @@ -732,13 +749,14 @@ def estimate_message_tokens(message: dict[str, Any]) -> int: if isinstance(content, str): parts.append(content) elif isinstance(content, list): - for part in content: - if isinstance(part, dict) and part.get("type") == "text": + for raw_part in cast(list[object], content): + part = cast(dict[str, Any], raw_part) if isinstance(raw_part, dict) else None + if part is not None and part.get("type") == "text": text = part.get("text", "") - if text: + if isinstance(text, str) and text: parts.append(text) else: - parts.append(json.dumps(part, ensure_ascii=False)) + parts.append(json.dumps(raw_part, ensure_ascii=False)) elif content is not None: parts.append(json.dumps(content, ensure_ascii=False)) @@ -764,7 +782,7 @@ def estimate_message_tokens(message: dict[str, Any]) -> int: def estimate_prompt_tokens_chain( - provider: Any, + provider: object, model: str | None, messages: list[dict[str, Any]], tools: list[dict[str, Any]] | None = None, @@ -773,7 +791,7 @@ def estimate_prompt_tokens_chain( provider_counter = getattr(provider, "estimate_prompt_tokens", None) if callable(provider_counter): with suppress(Exception): - tokens, source = provider_counter(messages, tools, model) + tokens, source = cast(tuple[object, object], provider_counter(messages, tools, model)) if isinstance(tokens, (int, float)) and tokens > 0: return int(tokens), str(source or "provider_counter") estimated, source = _estimate_prompt_tokens_with_source(messages, tools) @@ -851,7 +869,7 @@ def sync_workspace_templates(workspace: Path, silent: bool = False) -> list[str] added: list[str] = [] - def _write(src, dest: Path): + def _write(src: Any, dest: Path) -> None: content = src.read_text(encoding="utf-8") if src else "" if dest.exists(): return diff --git a/nanobot/utils/progress_events.py b/nanobot/utils/progress_events.py index 645a351d6..e3fd72406 100644 --- a/nanobot/utils/progress_events.py +++ b/nanobot/utils/progress_events.py @@ -4,7 +4,7 @@ from __future__ import annotations import inspect from collections.abc import Awaitable, Callable -from typing import Any +from typing import Any, cast from nanobot.agent.hook import AgentHookContext @@ -51,7 +51,7 @@ async def invoke_file_edit_progress( def _tool_event_arguments(tool_call: Any) -> dict[str, Any]: arguments = getattr(tool_call, "arguments", {}) or {} - return arguments if isinstance(arguments, dict) else {} + return cast(dict[str, Any], arguments) if isinstance(arguments, dict) else {} def build_tool_event_start_payload(tool_call: Any) -> dict[str, Any]: @@ -71,8 +71,11 @@ def build_tool_event_start_payload(tool_call: Any) -> dict[str, Any]: def tool_event_result_extras(result: Any) -> tuple[list[Any], list[Any]]: if not isinstance(result, dict): return [], [] - files = result.get("files") if isinstance(result.get("files"), list) else [] - embeds = result.get("embeds") if isinstance(result.get("embeds"), list) else [] + result_data = cast(dict[str, Any], result) + raw_files = result_data.get("files") + raw_embeds = result_data.get("embeds") + files: list[Any] = cast(list[Any], raw_files) if isinstance(raw_files, list) else [] + embeds: list[Any] = cast(list[Any], raw_embeds) if isinstance(raw_embeds, list) else [] return files, embeds @@ -82,7 +85,7 @@ def build_tool_event_finish_payloads(context: AgentHookContext) -> list[dict[str for idx in range(count): tool_call = context.tool_calls[idx] result = context.tool_results[idx] - event = context.tool_events[idx] if isinstance(context.tool_events[idx], dict) else {} + event = context.tool_events[idx] status = event.get("status") phase = "end" if status == "ok" else "error" files, embeds = tool_event_result_extras(result) diff --git a/nanobot/utils/restart.py b/nanobot/utils/restart.py index 2f9fa3544..03cbf8b1f 100644 --- a/nanobot/utils/restart.py +++ b/nanobot/utils/restart.py @@ -7,7 +7,7 @@ import os import time from contextlib import suppress from dataclasses import dataclass, field -from typing import Any +from typing import Any, cast from nanobot.webui.metadata import WEBUI_TURN_METADATA_KEY @@ -71,7 +71,7 @@ def consume_restart_notice_from_env() -> RestartNotice | None: except (TypeError, ValueError): parsed = None if isinstance(parsed, dict): - metadata = parsed + metadata = cast(dict[str, Any], parsed) return RestartNotice( channel=channel, chat_id=chat_id, diff --git a/nanobot/utils/runtime.py b/nanobot/utils/runtime.py index e755fa5be..fc850648c 100644 --- a/nanobot/utils/runtime.py +++ b/nanobot/utils/runtime.py @@ -4,7 +4,7 @@ from __future__ import annotations import re from pathlib import Path -from typing import Any +from typing import Any, cast from loguru import logger @@ -61,10 +61,10 @@ def ensure_nonempty_tool_result(tool_name: str, content: Any) -> Any: if isinstance(content, list): if not content: return empty_tool_result_message(tool_name) - text_payload = stringify_text_blocks(content) + text_payload = stringify_text_blocks(cast(list[Any], content)) if text_payload is not None and not text_payload.strip(): return empty_tool_result_message(tool_name) - return content + return cast(Any, content) def is_blank_text(content: str | None) -> bool: @@ -106,6 +106,7 @@ def external_lookup_signature(tool_name: str, arguments: Any) -> str | None: """Stable signature for repeated external lookups we want to throttle.""" if not isinstance(arguments, dict): return None + arguments = cast(dict[str, Any], arguments) if tool_name == "web_fetch": url = str(arguments.get("url") or "").strip() if url: @@ -153,6 +154,7 @@ def workspace_violation_signature( """Return a stable cross-tool signature for the outside-workspace target.""" if not isinstance(arguments, dict): return None + arguments = cast(dict[str, Any], arguments) for key in ("path", "file_path", "target", "source", "destination"): val = arguments.get(key) if isinstance(val, str) and val.strip(): diff --git a/nanobot/utils/searchusage.py b/nanobot/utils/searchusage.py index 94f76775d..af22973dd 100644 --- a/nanobot/utils/searchusage.py +++ b/nanobot/utils/searchusage.py @@ -4,7 +4,7 @@ from __future__ import annotations import os from dataclasses import dataclass -from typing import Any +from typing import Any, cast @dataclass @@ -28,7 +28,7 @@ class SearchUsageInfo: def format(self) -> str: """Return a human-readable multi-line string for /status output.""" - lines = [f"🔍 Web Search: {self.provider}"] + lines: list[str] = [f"🔍 Web Search: {self.provider}"] if not self.supported: lines.append(" Usage tracking: not available for this provider") @@ -44,7 +44,7 @@ class SearchUsageInfo: lines.append(f" Usage: {self.used} requests") # Tavily breakdown - breakdown_parts = [] + breakdown_parts: list[str] = [] if self.search_used is not None: breakdown_parts.append(f"Search: {self.search_used}") if self.extract_used is not None: @@ -109,7 +109,7 @@ async def _fetch_tavily_usage(api_key: str | None) -> SearchUsageInfo: headers={"Authorization": f"Bearer {key}"}, ) r.raise_for_status() - data: dict[str, Any] = r.json() + data = cast(dict[str, Any], r.json()) return _parse_tavily_usage(data) except httpx.HTTPStatusError as e: return SearchUsageInfo( @@ -145,7 +145,8 @@ def _parse_tavily_usage(data: dict[str, Any]) -> SearchUsageInfo: } } """ - account = data.get("account") or {} + raw_account = data.get("account") + account = cast(dict[str, Any], raw_account) if isinstance(raw_account, dict) else {} used = _optional_int(account.get("plan_usage")) limit = _optional_int(account.get("plan_limit")) diff --git a/nanobot/utils/subagent_channel_display.py b/nanobot/utils/subagent_channel_display.py index 3a939dd8e..88eee37a7 100644 --- a/nanobot/utils/subagent_channel_display.py +++ b/nanobot/utils/subagent_channel_display.py @@ -7,7 +7,7 @@ should show only the header plus a truncated result body.""" from __future__ import annotations -from typing import Any +from typing import Any, cast # Cap Result section length so WebSocket session replay stays readable; full text # remains on disk for LLM replay (we only mutate outgoing API copies in websocket). @@ -49,7 +49,7 @@ def scrub_subagent_announce_body(content: str) -> str: def scrub_subagent_messages_for_channel(messages: list[dict[str, Any]]) -> None: """Mutate message dicts in place when they carry ``subagent_result`` inject.""" for msg in messages: - if not isinstance(msg, dict): + if not isinstance(cast(object, msg), dict): continue if msg.get("injected_event") != "subagent_result": continue diff --git a/nanobot/utils/tool_hints.py b/nanobot/utils/tool_hints.py index a1212aeb0..15542e6a6 100644 --- a/nanobot/utils/tool_hints.py +++ b/nanobot/utils/tool_hints.py @@ -3,7 +3,9 @@ from __future__ import annotations import re +from typing import cast +from nanobot.providers.base import ToolCallRequest from nanobot.utils.path import abbreviate_path # Registry: tool_name -> (key_args, template, is_path, is_command) @@ -29,12 +31,15 @@ _PATH_IN_CMD_RE = re.compile( ) -def format_tool_hints(tool_calls: list, max_length: int = 40) -> str: +ToolFormat = tuple[list[str], str, bool, bool] + + +def format_tool_hints(tool_calls: list[ToolCallRequest], max_length: int = 40) -> str: """Format tool calls as concise hints with smart abbreviation.""" if not tool_calls: return "" - formatted = [] + formatted: list[str] = [] for tc in tool_calls: name = getattr(tc, "name", None) if not isinstance(name, str) or not name: @@ -49,7 +54,7 @@ def format_tool_hints(tool_calls: list, max_length: int = 40) -> str: else: formatted.append(_fmt_fallback(tc, max_length)) - hints = [] + hints: list[tuple[str, int]] = [] for hint in formatted: if hints and hints[-1][0] == hint: hints[-1] = (hint, hints[-1][1] + 1) @@ -61,22 +66,23 @@ def format_tool_hints(tool_calls: list, max_length: int = 40) -> str: ) -def _get_args(tc) -> dict: +def _get_args(tc: ToolCallRequest) -> dict[str, object]: """Extract args dict from tc.arguments, handling list/dict/None/empty.""" if tc.arguments is None: return {} - if isinstance(tc.arguments, list): - return tc.arguments[0] if tc.arguments else {} - if isinstance(tc.arguments, dict): - return tc.arguments + arguments = tc.arguments + if isinstance(arguments, list): + argument_list = cast(list[object], arguments) + first_argument = argument_list[0] if argument_list else None + return cast(dict[str, object], first_argument) if isinstance(first_argument, dict) else {} + if isinstance(arguments, dict): + return cast(dict[str, object], arguments) return {} -def _extract_arg(tc, key_args: list[str]) -> str | None: +def _extract_arg(tc: ToolCallRequest, key_args: list[str]) -> str | None: """Extract the first available value from preferred key names.""" args = _get_args(tc) - if not isinstance(args, dict): - return None for key in key_args: val = args.get(key) if isinstance(val, str) and val: @@ -87,7 +93,7 @@ def _extract_arg(tc, key_args: list[str]) -> str | None: return None -def _fmt_known(tc, fmt: tuple, max_length: int = 40) -> str: +def _fmt_known(tc: ToolCallRequest, fmt: ToolFormat, max_length: int = 40) -> str: """Format a registered tool using its template.""" if not fmt[0] and "{}" not in fmt[1]: return fmt[1] @@ -118,7 +124,7 @@ def _abbreviate_command(cmd: str, max_len: int = 40) -> str: return abbreviated[:max_len - 1] + "\u2026" -def _fmt_mcp(tc, max_length: int = 40) -> str: +def _fmt_mcp(tc: ToolCallRequest, max_length: int = 40) -> str: """Format MCP tool as server::tool.""" name = tc.name if "__" in name: @@ -139,10 +145,10 @@ def _fmt_mcp(tc, max_length: int = 40) -> str: return f'{server}::{tool}("{abbreviate_path(val, max_length)}")' -def _fmt_fallback(tc, max_length: int = 40) -> str: +def _fmt_fallback(tc: ToolCallRequest, max_length: int = 40) -> str: """Original formatting logic for unregistered tools.""" args = _get_args(tc) - val = next(iter(args.values()), None) if isinstance(args, dict) else None + val = next(iter(args.values()), None) if not isinstance(val, str): return tc.name return f'{tc.name}("{abbreviate_path(val, max_length)}")' if len(val) > max_length else f'{tc.name}("{val}")' diff --git a/nanobot/webui/attachment_ingress.py b/nanobot/webui/attachment_ingress.py index a9e19bd28..8e4eeb5f9 100644 --- a/nanobot/webui/attachment_ingress.py +++ b/nanobot/webui/attachment_ingress.py @@ -4,7 +4,7 @@ from __future__ import annotations import re from pathlib import Path -from typing import Any, Literal +from typing import Any, Literal, cast from nanobot.utils.media_decode import FileSizeExceeded, save_base64_data_url from nanobot.webui.ingress_policy import ( @@ -93,9 +93,10 @@ def store_inbound_attachments( video_count = 0 document_count = 0 for item in media: + attachment = cast(dict[str, Any], item) if isinstance(item, dict) else None mime = ( - extract_data_url_mime(item.get("data_url", "")) - if isinstance(item, dict) + extract_data_url_mime(attachment.get("data_url", "")) + if attachment is not None else None ) if mime in _VIDEO_MIME_ALLOWED: @@ -125,7 +126,8 @@ def store_inbound_attachments( for item in media: if not isinstance(item, dict): return abort("malformed") - data_url = item.get("data_url") + attachment = cast(dict[str, Any], item) + data_url = attachment.get("data_url") if not isinstance(data_url, str) or not data_url: return abort("malformed") mime = extract_data_url_mime(data_url) @@ -140,8 +142,8 @@ def store_inbound_attachments( else limits.max_file_bytes ) name = ( - item.get("name") - if is_document and isinstance(item.get("name"), str) + attachment.get("name") + if is_document and isinstance(attachment.get("name"), str) else None ) try: diff --git a/nanobot/webui/build.py b/nanobot/webui/build.py index 6c53afc6c..ef8a247e6 100644 --- a/nanobot/webui/build.py +++ b/nanobot/webui/build.py @@ -9,7 +9,7 @@ from collections.abc import Callable, Mapping from contextlib import suppress from dataclasses import dataclass from pathlib import Path -from typing import Literal +from typing import Any, Literal BuildMode = Literal["auto", "prompt", "warn", "skip"] @@ -186,7 +186,7 @@ def build_webui_bundle( source_dir: Path | None = None, dist_dir: Path | None = None, runner: str | None = None, - subprocess_run: Callable[..., subprocess.CompletedProcess] = subprocess.run, + subprocess_run: Callable[..., subprocess.CompletedProcess[Any]] = subprocess.run, output: Callable[[str], None] | None = None, ) -> WebUIBundleStatus: """Install frontend dependencies and build the WebUI bundle.""" @@ -221,7 +221,7 @@ def ensure_webui_bundle( output: Callable[[str], None] | None = None, runner: str | None = None, environ: Mapping[str, str] | None = None, - subprocess_run: Callable[..., subprocess.CompletedProcess] = subprocess.run, + subprocess_run: Callable[..., subprocess.CompletedProcess[Any]] = subprocess.run, ) -> WebUIBundleStatus: """Ensure or warn about a stale WebUI bundle according to the selected mode.""" env = environ or os.environ @@ -275,7 +275,7 @@ def _run_frontend_command( command: list[str], *, cwd: Path, - subprocess_run: Callable[..., subprocess.CompletedProcess], + subprocess_run: Callable[..., subprocess.CompletedProcess[Any]], ) -> None: try: subprocess_run(command, cwd=cwd, check=True) diff --git a/nanobot/webui/cli_apps_api.py b/nanobot/webui/cli_apps_api.py index d9eeb8cc0..b0d1640e5 100644 --- a/nanobot/webui/cli_apps_api.py +++ b/nanobot/webui/cli_apps_api.py @@ -5,7 +5,7 @@ from __future__ import annotations import asyncio import re import time -from typing import Any +from typing import Any, cast from nanobot.apps.cli import CliAppError, CliAppManager, CliAppsRuntimeConfig from nanobot.config.loader import load_config @@ -58,12 +58,14 @@ def normalize_cli_app_mentions(raw: Any) -> list[dict[str, str]]: """Sanitize structured CLI app mentions sent by the WebUI.""" if not isinstance(raw, list): return [] + raw_items = cast(list[Any], raw) out: list[dict[str, str]] = [] seen: set[str] = set() - for item in raw[:8]: + for item in raw_items[:8]: if not isinstance(item, dict): continue - name = _clip_ws_string(item.get("name"), 64) + app_data = cast(dict[str, Any], item) + name = _clip_ws_string(app_data.get("name"), 64) if not name or _CLI_APP_NAME_RE.match(name) is None: continue key = name.lower() @@ -72,7 +74,10 @@ def normalize_cli_app_mentions(raw: Any) -> list[dict[str, str]]: seen.add(key) row: dict[str, str] = {"name": key} for field in _CLI_APP_ATTACHMENT_KEYS[1:]: - value = _clip_ws_string(item.get(field), 512 if field == "logo_url" else 160) + value = _clip_ws_string( + app_data.get(field), + 512 if field == "logo_url" else 160, + ) if value: row[field] = value out.append(row) diff --git a/nanobot/webui/forking.py b/nanobot/webui/forking.py index c67f559a6..169777556 100644 --- a/nanobot/webui/forking.py +++ b/nanobot/webui/forking.py @@ -5,7 +5,7 @@ from __future__ import annotations import re import uuid from collections.abc import Mapping -from typing import Any +from typing import TYPE_CHECKING, Any, TypeGuard from nanobot.session.manager import SessionManager from nanobot.session.webui_turns import WEBUI_TITLE_METADATA_KEY, clean_generated_title @@ -16,10 +16,15 @@ from nanobot.webui.transcript import ( write_session_messages_as_transcript, ) +if TYPE_CHECKING: + from websockets.asyncio.server import ServerConnection + + from nanobot.channels.websocket.runtime import WebSocketChannel + _WEBUI_CHAT_ID_RE = re.compile(r"^[A-Za-z0-9_:-]{1,64}$") -def _valid_webui_chat_id(value: Any) -> bool: +def _valid_webui_chat_id(value: Any) -> TypeGuard[str]: return isinstance(value, str) and _WEBUI_CHAT_ID_RE.match(value) is not None @@ -63,7 +68,11 @@ def create_webui_chat_fork( return new_id, target_key -async def handle_webui_fork_chat(channel: Any, connection: Any, envelope: Mapping[str, Any]) -> None: +async def handle_webui_fork_chat( + channel: WebSocketChannel, + connection: ServerConnection, + envelope: Mapping[str, Any], +) -> None: """Handle the WebUI ``fork_chat`` websocket command. ``websocket.py`` owns the transport. This module owns WebUI fork semantics: @@ -73,15 +82,15 @@ async def handle_webui_fork_chat(channel: Any, connection: Any, envelope: Mappin source_chat_id = envelope.get("source_chat_id") raw_index = envelope.get("before_user_index") if not _valid_webui_chat_id(source_chat_id): - await channel._send_event(connection, "error", detail="invalid source_chat_id") + await channel.send_webui_protocol_error(connection, "invalid source_chat_id") return if isinstance(raw_index, bool) or not isinstance(raw_index, int) or raw_index < 0: - await channel._send_event(connection, "error", detail="invalid before_user_index") + await channel.send_webui_protocol_error(connection, "invalid before_user_index") return session_manager = channel.gateway.session_manager if session_manager is None: - await channel._send_event(connection, "error", detail="session_manager_unavailable") + await channel.send_webui_protocol_error(connection, "session_manager_unavailable") return try: @@ -92,22 +101,16 @@ async def handle_webui_fork_chat(channel: Any, connection: Any, envelope: Mappin title=envelope.get("title") if isinstance(envelope.get("title"), str) else None, ) if forked is None: - await channel._send_event(connection, "error", detail="invalid fork source or index") + await channel.send_webui_protocol_error(connection, "invalid fork source or index") return fork_id, fork_key = forked except Exception as exc: channel.logger.warning("fork_chat failed: {}", exc) - await channel._send_event(connection, "error", detail="fork_chat_failed") + await channel.send_webui_protocol_error(connection, "fork_chat_failed") return - scope = channel._workspaces.scope_for_session_key(fork_key) - channel._attach(connection, fork_id) - await channel._send_event(connection, "attached", chat_id=fork_id) - await channel._send_event( + await channel.attach_webui_fork( connection, - "session_updated", - chat_id=fork_id, - scope="metadata", - workspace_scope=scope.payload(), + fork_id=fork_id, + fork_key=fork_key, ) - await channel._hydrate_after_subscribe(fork_id) diff --git a/nanobot/webui/gateway_services.py b/nanobot/webui/gateway_services.py index 6bf438f39..eb944d766 100644 --- a/nanobot/webui/gateway_services.py +++ b/nanobot/webui/gateway_services.py @@ -4,7 +4,7 @@ from __future__ import annotations from dataclasses import dataclass from pathlib import Path -from typing import Any, Callable +from typing import TYPE_CHECKING, Any, Callable from loguru import logger as default_logger @@ -15,6 +15,13 @@ from nanobot.webui.transcript import WebUITranscriptRecorder from nanobot.webui.workspaces import WebUIWorkspaceController from nanobot.webui.ws_http import GatewayHTTPHandler +if TYPE_CHECKING: + from nanobot.bus.queue import MessageBus + from nanobot.channels.websocket.runtime import WebSocketConfig + from nanobot.cron.service import CronService + from nanobot.session.manager import SessionManager + from nanobot.triggers.local_store import LocalTriggerStore + @dataclass(frozen=True) class GatewayServices: @@ -26,27 +33,27 @@ class GatewayServices: ingress: WebUIIngressPolicy transcripts: WebUITranscriptRecorder workspaces: WebUIWorkspaceController - session_manager: Any | None - cron_service: Any | None - local_trigger_store: Any | None + session_manager: SessionManager | None + cron_service: CronService | None + local_trigger_store: LocalTriggerStore | None cron_pending_job_ids: Callable[[str], set[str]] | None local_trigger_pending_ids: Callable[[str], set[str]] | None def build_gateway_services( *, - config: Any, - bus: Any, - session_manager: Any | None, + config: WebSocketConfig, + bus: MessageBus, + session_manager: SessionManager | None, static_dist_path: Path | None, workspace_path: Path, default_restrict_to_workspace: bool, - runtime_model_name: Any | None, + runtime_model_name: Callable[[], str | None] | None, runtime_surface: str, runtime_capabilities_overrides: dict[str, Any] | None, disabled_skills: set[str] | None = None, - cron_service: Any | None = None, - local_trigger_store: Any | None = None, + cron_service: CronService | None = None, + local_trigger_store: LocalTriggerStore | None = None, cron_pending_job_ids: Callable[[str], set[str]] | None = None, local_trigger_pending_ids: Callable[[str], set[str]] | None = None, channel_feature_action: Callable[..., Any] | None = None, diff --git a/nanobot/webui/http_utils.py b/nanobot/webui/http_utils.py index 22a36d2e5..e261b6f5b 100644 --- a/nanobot/webui/http_utils.py +++ b/nanobot/webui/http_utils.py @@ -8,7 +8,7 @@ import http import ipaddress import json import re -from typing import Any +from typing import Any, cast from urllib.parse import parse_qs, urlparse from websockets.datastructures import Headers @@ -120,7 +120,7 @@ def is_localhost(connection: Any) -> bool: addr = getattr(connection, "remote_address", None) if not addr: return False - host = addr[0] if isinstance(addr, tuple) else addr + host = cast(Any, addr[0] if isinstance(addr, tuple) else addr) if not isinstance(host, str): return False if host.startswith("::ffff:"): diff --git a/nanobot/webui/mcp_presets_api.py b/nanobot/webui/mcp_presets_api.py index 0521848f2..ca4147a33 100644 --- a/nanobot/webui/mcp_presets_api.py +++ b/nanobot/webui/mcp_presets_api.py @@ -14,7 +14,7 @@ from contextlib import suppress from dataclasses import dataclass from datetime import datetime, timezone from pathlib import Path -from typing import Any, Literal, Mapping +from typing import Any, Literal, Mapping, cast from nanobot.agent.tools.registry import ToolRegistry from nanobot.apps.protocol import app_manifest, compact_dict @@ -455,9 +455,10 @@ def normalize_mcp_preset_mentions(raw: Any) -> list[dict[str, Any]]: known = _known_mcp_names() out: list[dict[str, Any]] = [] seen: set[str] = set() - for item in raw[:8]: - if not isinstance(item, dict): + for item_value in cast(list[object], raw)[:8]: + if not isinstance(item_value, dict): continue + item = cast(dict[str, Any], item_value) name = _clip_ws_string(item.get("name"), 64) if not name or _MCP_PRESET_NAME_RE.match(name) is None: continue @@ -959,7 +960,7 @@ async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]: "missing_dependency" if cfg.command and not _command_available(cfg.command) else "configured" ) if status == "missing_credentials": - last_action = { + last_action: dict[str, Any] = { "ok": False, "message": f"{display_name} is missing required credentials.", "error": "missing credentials", @@ -988,7 +989,11 @@ async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]: timeout=_test_timeout(cfg), ) tool_prefix = f"mcp_{name}_" - tool_names = sorted(name for name in registry.tool_names if name.startswith(tool_prefix)) + tool_names = sorted( + tool_name + for tool_name in registry.tool_names + if tool_name.startswith(tool_prefix) + ) ok = name in stacks if ok: last_action = { @@ -1033,7 +1038,8 @@ async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]: finally: await _close_mcp_stacks(stacks) - preview = {name: last_action.get("tool_names", [])} if last_action.get("tool_names") else None + tool_names = last_action.get("tool_names", []) + preview = {name: tool_names} if tool_names else None return mcp_presets_payload(last_action=last_action, tool_preview=preview) @@ -1050,8 +1056,10 @@ def _parse_string_list(raw: str | None) -> list[str]: if raw is None or not raw.strip(): return [] parsed = _parse_json_value(raw, fallback=None) - if isinstance(parsed, list) and all(isinstance(item, str) for item in parsed): - return [item for item in parsed if item.strip()] + if isinstance(parsed, list): + items = cast(list[object], parsed) + if all(isinstance(item, str) for item in items): + return [item for item in cast(list[str], items) if item.strip()] if isinstance(parsed, str): return shlex.split(parsed) raise McpPresetError("expected a JSON string array") @@ -1062,7 +1070,7 @@ def _parse_string_map(raw: str | None) -> dict[str, str]: if not isinstance(parsed, dict): raise McpPresetError("expected a JSON object") out: dict[str, str] = {} - for key, value in parsed.items(): + for key, value in cast(dict[object, object], parsed).items(): if not isinstance(key, str) or not isinstance(value, str): raise McpPresetError("JSON object values must be strings") if key.strip(): @@ -1141,42 +1149,56 @@ def _mcp_server_config(name: str, raw: Any) -> tuple[str, MCPServerConfig]: server_name = _validated_server_name(name) if not isinstance(raw, Mapping): raise McpPresetError(f"MCP server '{server_name}' must be an object") - command = str(raw.get("command") or "").strip() - url = str(raw.get("url") or "").strip() - transport_value = str(raw.get("type", raw.get("transport", "")) or "") + server = cast(Mapping[str, Any], raw) + command = str(server.get("command") or "").strip() + url = str(server.get("url") or "").strip() + transport_value = str(server.get("type", server.get("transport", "")) or "") transport = _normalize_transport(transport_value, command=command, url=url) if transport == "stdio" and not command: raise McpPresetError(f"MCP server '{server_name}' stdio transport requires a command") if transport in {"sse", "streamableHttp"} and not url: raise McpPresetError(f"MCP server '{server_name}' remote transport requires a URL") - args = raw.get("args") or [] - env = raw.get("env") or {} - headers = raw.get("headers") or {} - cwd = str(raw.get("cwd") or "").strip() - enabled_tools = raw.get("enabledTools", raw.get("enabled_tools", ["*"])) - tool_timeout = raw.get("toolTimeout", raw.get("tool_timeout", _DEFAULT_CUSTOM_TIMEOUT)) + args_value: object = server.get("args") or [] + env_value: object = server.get("env") or {} + headers_value: object = server.get("headers") or {} + cwd = str(server.get("cwd") or "").strip() + enabled_tools_value: object = server.get("enabledTools", server.get("enabled_tools", ["*"])) + tool_timeout: object = server.get("toolTimeout", server.get("tool_timeout", _DEFAULT_CUSTOM_TIMEOUT)) try: - timeout_int = max(5, min(int(tool_timeout), 600)) + timeout_int = max(5, min(int(cast(Any, tool_timeout)), 600)) except (TypeError, ValueError): timeout_int = _DEFAULT_CUSTOM_TIMEOUT - if not isinstance(args, list) or not all(isinstance(item, str) for item in args): + if not isinstance(args_value, list): raise McpPresetError(f"MCP server '{server_name}' args must be a string array") - if not isinstance(env, dict) or not all(isinstance(k, str) and isinstance(v, str) for k, v in env.items()): + args = cast(list[object], args_value) + if not all(isinstance(item, str) for item in args): + raise McpPresetError(f"MCP server '{server_name}' args must be a string array") + if not isinstance(env_value, dict): raise McpPresetError(f"MCP server '{server_name}' env must be a string object") - if not isinstance(headers, dict) or not all(isinstance(k, str) and isinstance(v, str) for k, v in headers.items()): + env = cast(dict[object, object], env_value) + if not all(isinstance(k, str) and isinstance(v, str) for k, v in env.items()): + raise McpPresetError(f"MCP server '{server_name}' env must be a string object") + if not isinstance(headers_value, dict): raise McpPresetError(f"MCP server '{server_name}' headers must be a string object") - if not isinstance(enabled_tools, list) or not all(isinstance(item, str) for item in enabled_tools): - enabled_tools = ["*"] + headers = cast(dict[object, object], headers_value) + if not all(isinstance(k, str) and isinstance(v, str) for k, v in headers.items()): + raise McpPresetError(f"MCP server '{server_name}' headers must be a string object") + if not isinstance(enabled_tools_value, list): + enabled_tools_value = ["*"] + else: + enabled_tools = cast(list[object], enabled_tools_value) + if not all(isinstance(item, str) for item in enabled_tools): + enabled_tools_value = ["*"] return server_name, MCPServerConfig( type=transport, command=command if transport == "stdio" else "", - args=args, - env=dict(env), + args=cast(list[str], args), + env=cast(dict[str, str], env), cwd=cwd if transport == "stdio" else "", url=url if transport in {"sse", "streamableHttp"} else "", - headers=dict(headers), + headers=cast(dict[str, str], headers), tool_timeout=timeout_int, - enabled_tools=list(enabled_tools), + enabled_tools=cast(list[str], enabled_tools_value), ) @@ -1184,11 +1206,12 @@ def _import_mcp_servers(raw_json: str | None) -> dict[str, MCPServerConfig]: parsed = _parse_json_value(raw_json, fallback=None) if not isinstance(parsed, Mapping): raise McpPresetError("MCP config must be a JSON object") - servers = parsed.get("mcpServers", parsed) + parsed_mapping = cast(Mapping[str, Any], parsed) + servers: object = parsed_mapping.get("mcpServers", parsed_mapping) if not isinstance(servers, Mapping): raise McpPresetError("MCP config must contain mcpServers") out: dict[str, MCPServerConfig] = {} - for name, raw_server in servers.items(): + for name, raw_server in cast(Mapping[object, object], servers).items(): if not isinstance(name, str): raise McpPresetError("MCP server names must be strings") server_name, cfg = _mcp_server_config(name, raw_server) diff --git a/nanobot/webui/media_api.py b/nanobot/webui/media_api.py index f8292d40d..76db933dc 100644 --- a/nanobot/webui/media_api.py +++ b/nanobot/webui/media_api.py @@ -12,7 +12,7 @@ import shutil import uuid from collections.abc import Callable from pathlib import Path -from typing import Any +from typing import Any, cast from websockets.http11 import Request as WsRequest from websockets.http11 import Response @@ -182,14 +182,17 @@ def attach_signed_media_urls( messages = payload.get("messages") if not isinstance(messages, list): return - for msg in messages: + raw_messages = cast(list[Any], messages) + for msg in raw_messages: if not isinstance(msg, dict): continue - media = msg.get("media") + message = cast(dict[str, Any], msg) + media = message.get("media") if not isinstance(media, list) or not media: continue + media_entries = cast(list[Any], media) urls: list[dict[str, str]] = [] - for entry in media: + for entry in media_entries: if not isinstance(entry, str) or not entry: continue signed = sign_path(Path(entry)) @@ -197,8 +200,8 @@ def attach_signed_media_urls( continue urls.append({"url": signed, "name": Path(entry).name}) if urls: - msg["media_urls"] = urls - msg.pop("media", None) + message["media_urls"] = urls + message.pop("media", None) def serve_signed_media( diff --git a/nanobot/webui/media_gateway.py b/nanobot/webui/media_gateway.py index 4483e3344..b1accb566 100644 --- a/nanobot/webui/media_gateway.py +++ b/nanobot/webui/media_gateway.py @@ -26,6 +26,10 @@ from nanobot.webui.media_api import ( from nanobot.webui.transcript import rewrite_local_markdown_images +def _default_media_dir(channel: str | None) -> Path: + return get_media_dir(channel) + + class WebUIMediaGateway: """Own media URL signing and WebUI markdown/media augmentation.""" @@ -40,7 +44,7 @@ class WebUIMediaGateway: ) -> None: self.workspace_path = workspace_path self.logger = logger - self._media_dir = media_dir or (lambda channel=None: get_media_dir(channel)) + self._media_dir: Callable[[str | None], Path] = media_dir or _default_media_dir self.secret = secret or secrets.token_bytes(32) self.attachment_limits = attachment_limits or AttachmentIngressLimits() diff --git a/nanobot/webui/session_automations.py b/nanobot/webui/session_automations.py index 37cec5ddb..be54c330f 100644 --- a/nanobot/webui/session_automations.py +++ b/nanobot/webui/session_automations.py @@ -3,11 +3,11 @@ from __future__ import annotations from collections.abc import Collection -from typing import Any, Protocol +from typing import Any, Protocol, cast from nanobot.cron.types import CronJob from nanobot.session.history_visibility import is_hidden_history_message -from nanobot.session.manager import _message_preview_text +from nanobot.session.manager import _message_preview_text # pyright: ignore[reportPrivateUsage] from nanobot.triggers.local_types import LocalTrigger AutomationJob = CronJob | LocalTrigger @@ -140,7 +140,7 @@ def _serialize_job( session_manager=session_manager, ) - payload = { + payload: dict[str, Any] = { "id": job.id, "name": job.name, "enabled": job.enabled, @@ -195,7 +195,7 @@ def _serialize_trigger( session_manager: _SessionManagerLike | None = None, ) -> dict[str, Any]: command = f'nanobot trigger {trigger.id} "message"' - payload = { + payload: dict[str, Any] = { "id": trigger.id, "name": trigger.name, "enabled": trigger.enabled, @@ -325,9 +325,10 @@ def _session_preview(messages: Any) -> str: if not isinstance(messages, list): return "" fallback_preview = "" - for message in messages: - if not isinstance(message, dict): + for message_value in cast(list[object], messages): + if not isinstance(message_value, dict): continue + message = cast(dict[str, Any], message_value) if is_hidden_history_message(message): continue text = _message_preview_text(message) diff --git a/nanobot/webui/session_list_index.py b/nanobot/webui/session_list_index.py index 104a82067..31dbb79c2 100644 --- a/nanobot/webui/session_list_index.py +++ b/nanobot/webui/session_list_index.py @@ -11,19 +11,19 @@ import json import os from datetime import datetime from pathlib import Path -from typing import Any +from typing import Any, cast from loguru import logger from nanobot.config.paths import get_webui_dir from nanobot.session.history_visibility import is_hidden_history_message from nanobot.session.manager import ( - _SESSION_LIST_PREVIEW_MAX_CHARS, - _SESSION_LIST_PREVIEW_MAX_RECORDS, + _SESSION_LIST_PREVIEW_MAX_CHARS, # pyright: ignore[reportPrivateUsage] + _SESSION_LIST_PREVIEW_MAX_RECORDS, # pyright: ignore[reportPrivateUsage] Session, SessionManager, - _message_preview_text, - _metadata_title, + _message_preview_text, # pyright: ignore[reportPrivateUsage] + _metadata_title, # pyright: ignore[reportPrivateUsage] ) from nanobot.session.model_selection import model_preset_from_metadata @@ -57,7 +57,7 @@ def _reconcile_index(session_manager: SessionManager) -> tuple[list[dict[str, An paths = sorted( path for path in session_manager.sessions_dir.glob("*.jsonl") - if SessionManager._session_key_from_path(path) is not None + if SessionManager._session_key_from_path(path) is not None # pyright: ignore[reportPrivateUsage] ) rows: list[dict[str, Any]] = [] changed = existing_rows is None @@ -92,12 +92,18 @@ def _read_index_rows(sessions_dir: Path) -> list[dict[str, Any]] | None: data = json.loads(path.read_text(encoding="utf-8")) except (OSError, json.JSONDecodeError): return None - if not isinstance(data, dict) or data.get("version") != _INDEX_VERSION: + if not isinstance(data, dict): return None - rows = data.get("sessions") - if not isinstance(rows, list) or not all(isinstance(row, dict) for row in rows): + index_data = cast(dict[str, Any], data) + if index_data.get("version") != _INDEX_VERSION: return None - return rows + rows = index_data.get("sessions") + if not isinstance(rows, list): + return None + index_rows = cast(list[Any], rows) + if not all(isinstance(row, dict) for row in index_rows): + return None + return [cast(dict[str, Any], row) for row in index_rows] def _write_index_rows(sessions_dir: Path, rows: list[dict[str, Any]]) -> None: @@ -272,7 +278,7 @@ def _indexed_row_for_session(session: Session, path: Path) -> dict[str, Any]: def _scan_session_row(session_manager: SessionManager, path: Path) -> dict[str, Any] | None: - storage_key = SessionManager._session_key_from_path(path) + storage_key = SessionManager._session_key_from_path(path) # pyright: ignore[reportPrivateUsage] if storage_key is None: return None try: @@ -345,7 +351,7 @@ def _scan_session_row(session_manager: SessionManager, path: Path) -> dict[str, **activity_signature, } except Exception: - repaired = session_manager._repair(storage_key) + repaired = session_manager._repair(storage_key) # pyright: ignore[reportPrivateUsage] if repaired is None: return None return _indexed_row_for_session(repaired, path) diff --git a/nanobot/webui/settings_api.py b/nanobot/webui/settings_api.py index 85c52f6c0..038938913 100644 --- a/nanobot/webui/settings_api.py +++ b/nanobot/webui/settings_api.py @@ -4,6 +4,9 @@ The WebSocket channel owns transport/authentication. This module owns the settings payload shape and the allowlisted config mutations exposed to WebUI. """ +# oauth-cli-kit is an optional dependency and does not publish type stubs. +# pyright: reportMissingTypeStubs=false + from __future__ import annotations import json @@ -13,8 +16,9 @@ import re import secrets import threading import time +from collections.abc import Iterable from contextlib import suppress -from typing import Any, Literal +from typing import Any, Literal, cast from zoneinfo import ZoneInfo import httpx @@ -27,7 +31,7 @@ from nanobot.audio.transcription_registry import ( transcription_provider_names, ) from nanobot.config.loader import get_config_path, load_config, resolve_config_env_vars, save_config -from nanobot.config.schema import ModelPresetConfig, ProviderConfig +from nanobot.config.schema import Config, FallbackCandidate, ModelPresetConfig, ProviderConfig from nanobot.providers.image_generation import ( get_image_gen_provider, image_gen_provider_names, @@ -184,8 +188,12 @@ def decorate_settings_payload( surface_value = _normalize_surface(surface) sections = restart_required_sections if sections is None: - raw_sections = payload.get("restart_required_sections") or [] - sections = [str(section) for section in raw_sections if isinstance(section, str)] + raw_sections: object = payload.get("restart_required_sections") or [] + sections = [ + section + for section in cast(Iterable[object], raw_sections) + if isinstance(section, str) + ] sections = sorted(dict.fromkeys(sections)) result = dict(payload) result["surface"] = surface_value @@ -230,12 +238,12 @@ def _provider_json_setting( if not raw: return None try: - value = json.loads(raw) + value: object = json.loads(raw) except json.JSONDecodeError as exc: raise WebUISettingsError(f"{snake} must be a JSON object") from exc if not isinstance(value, dict): raise WebUISettingsError(f"{snake} must be a JSON object") - return value or None + return cast(dict[str, Any], value) or None _REDACTED_PROVIDER_SECRET = "••••••••" @@ -280,15 +288,19 @@ def _redact_provider_secret_values(value: Any, *, secret: bool = False) -> Any: if secret and value not in (None, ""): return _REDACTED_PROVIDER_SECRET if isinstance(value, dict): + value_mapping = cast(dict[str, Any], value) return { key: _redact_provider_secret_values( item, secret=_provider_setting_key_is_secret(key), ) - for key, item in value.items() + for key, item in value_mapping.items() } if isinstance(value, list): - return [_redact_provider_secret_values(item) for item in value] + return [ + _redact_provider_secret_values(item) + for item in cast(list[Any], value) + ] return value @@ -301,23 +313,25 @@ def _restore_redacted_provider_secret_values( if secret and submitted == _REDACTED_PROVIDER_SECRET: return current if isinstance(submitted, dict): - current_mapping = current if isinstance(current, dict) else {} + submitted_mapping = cast(dict[str, Any], submitted) + current_mapping = cast(dict[str, Any], current) if isinstance(current, dict) else {} return { key: _restore_redacted_provider_secret_values( item, current_mapping.get(key), secret=_provider_setting_key_is_secret(key), ) - for key, item in submitted.items() + for key, item in submitted_mapping.items() } if isinstance(submitted, list): - current_items = current if isinstance(current, list) else [] + submitted_items = cast(list[Any], submitted) + current_items = cast(list[Any], current) if isinstance(current, list) else [] return [ _restore_redacted_provider_secret_values( item, current_items[index] if index < len(current_items) else None, ) - for index, item in enumerate(submitted) + for index, item in enumerate(submitted_items) ] return submitted @@ -366,7 +380,12 @@ def _validated_provider_config( try: return config_type.model_validate(values) except ValueError as exc: - errors = getattr(exc, "errors", lambda: [])() + errors_callback = getattr(exc, "errors", None) + errors: list[dict[str, Any]] = ( + cast(Any, errors_callback)() + if callable(errors_callback) + else [] + ) if errors: error = errors[0] field = ".".join(str(part) for part in error.get("loc", ())) @@ -515,16 +534,17 @@ def _provider_configured_for_settings(spec: Any, provider_config: Any) -> bool: ) -def _dynamic_provider_items(config: Any) -> list[tuple[str, ProviderConfig]]: +def _dynamic_provider_items(config: Config) -> list[tuple[str, ProviderConfig]]: + model_extra = config.providers.model_extra or {} return [ (name, provider_config) - for name, provider_config in (config.providers.model_extra or {}).items() + for name, provider_config in model_extra.items() if isinstance(provider_config, ProviderConfig) ] def _resolve_settings_provider( - config: Any, + config: Config, provider_name: str, ) -> tuple[Any, str, ProviderConfig] | None: spec = find_by_name(provider_name) @@ -610,7 +630,7 @@ def _provider_settings_row( return row -def _provider_settings_rows(config: Any, selected_provider: str | None) -> list[dict[str, Any]]: +def _provider_settings_rows(config: Config, selected_provider: str | None) -> list[dict[str, Any]]: """Return one Settings row per provider family while preserving legacy configs.""" aliases: dict[str, list[Any]] = {} for spec in PROVIDERS: @@ -664,8 +684,9 @@ def _model_id_from_row(row: Any) -> str | None: return row.strip() or None if not isinstance(row, dict): return None + row_mapping = cast(dict[str, Any], row) for key in ("id", "name", "model"): - value = row.get(key) + value = row_mapping.get(key) if isinstance(value, str) and value.strip(): return value.strip() return None @@ -674,6 +695,7 @@ def _model_id_from_row(row: Any) -> str | None: def _model_context_window(row: Any) -> int | None: if not isinstance(row, dict): return None + row_mapping = cast(dict[str, Any], row) for key in ( "context_window", "context_length", @@ -681,7 +703,7 @@ def _model_context_window(row: Any) -> int | None: "max_model_len", "max_input_tokens", ): - value = row.get(key) + value = row_mapping.get(key) if isinstance(value, int) and value > 0: return value if isinstance(value, float) and value > 0: @@ -697,13 +719,22 @@ def _model_row_payload(row: Any) -> dict[str, Any] | None: description: str | None = None owned_by: str | None = None if isinstance(row, dict): - raw_label = row.get("display_name") or row.get("label") or row.get("name") + row_mapping = cast(dict[str, Any], row) + raw_label = ( + row_mapping.get("display_name") + or row_mapping.get("label") + or row_mapping.get("name") + ) if isinstance(raw_label, str) and raw_label.strip() and raw_label.strip() != model_id: label = raw_label.strip() - raw_description = row.get("description") + raw_description = row_mapping.get("description") if isinstance(raw_description, str) and raw_description.strip(): description = raw_description.strip() - raw_owner = row.get("owned_by") or row.get("owner") or row.get("organization") + raw_owner = ( + row_mapping.get("owned_by") + or row_mapping.get("owner") + or row_mapping.get("organization") + ) if isinstance(raw_owner, str) and raw_owner.strip(): owned_by = raw_owner.strip() payload = { @@ -718,12 +749,12 @@ def _model_row_payload(row: Any) -> dict[str, Any] | None: def _extract_model_rows(body: Any) -> list[dict[str, Any]]: - raw_rows = body.get("data") if isinstance(body, dict) else body + raw_rows = cast(dict[str, Any], body).get("data") if isinstance(body, dict) else body if not isinstance(raw_rows, list): return [] rows: list[dict[str, Any]] = [] seen: set[str] = set() - for raw_row in raw_rows: + for raw_row in cast(list[object], raw_rows): row = _model_row_payload(raw_row) if row is None or row["id"] in seen: continue @@ -907,7 +938,7 @@ def _model_configuration_slug(label: str) -> str: return normalized -def _custom_provider_key(config: Any, display_name: str) -> str: +def _custom_provider_key(config: Config, display_name: str) -> str: slug = _MODEL_CONFIGURATION_SLUG_RE.sub("-", display_name.strip().lower()).strip("-_") base = f"custom-{slug or 'provider'}" if len(base) > 56: @@ -925,7 +956,7 @@ def _custom_provider_key(config: Any, display_name: str) -> str: def _provider_display_name_exists( - config: Any, + config: Config, display_name: str, *, exclude_key: str | None = None, @@ -945,7 +976,7 @@ def _provider_display_name_exists( return False -def _unique_model_configuration_name(config: Any, label: str) -> str: +def _unique_model_configuration_name(config: Config, label: str) -> str: """Return a stable, unused preset name for a migrated model configuration.""" try: base = _model_configuration_slug(label) @@ -963,7 +994,7 @@ def _model_configuration_label(model: str) -> str: return model.rsplit("/", 1)[-1] or model -def _model_call_order_state(config: Any) -> tuple[list[str], bool]: +def _model_call_order_state(config: Config) -> tuple[list[str], bool]: defaults = config.agents.defaults primary = defaults.model_preset if not primary or primary == "default" or primary not in config.model_presets: @@ -976,7 +1007,7 @@ def _model_call_order_state(config: Any) -> tuple[list[str], bool]: return order, True -def _validate_configured_provider(config: Any, provider: str) -> None: +def _validate_configured_provider(config: Config, provider: str) -> None: if provider == "auto": return resolved_provider = _resolve_settings_provider(config, provider) @@ -989,7 +1020,7 @@ def _validate_configured_provider(config: Any, provider: str) -> None: raise WebUISettingsError("provider is not configured") -def _image_generation_provider_rows(config: Any) -> list[dict[str, Any]]: +def _image_generation_provider_rows(config: Config) -> list[dict[str, Any]]: rows: list[dict[str, Any]] = [] for name in image_gen_provider_names(): image_provider = get_image_gen_provider(name) @@ -1062,7 +1093,7 @@ def _reasoning_effort_values_for(provider_name: str, model: str) -> list[str]: return list(_DEFAULT_REASONING_EFFORT_VALUES) -def _transcription_provider_rows(config: Any) -> list[dict[str, Any]]: +def _transcription_provider_rows(config: Config) -> list[dict[str, Any]]: rows: list[dict[str, Any]] = [] for name in transcription_provider_names(): spec = find_by_name(name) @@ -1549,17 +1580,23 @@ def update_model_call_order(query: QueryParams) -> dict[str, Any]: if raw_order is None: raise WebUISettingsError("model call order is required") try: - order = json.loads(raw_order) + order: object = json.loads(raw_order) except json.JSONDecodeError: raise WebUISettingsError("model call order must be a JSON array") from None if ( not isinstance(order, list) or not order - or any(not isinstance(name, str) or not name.strip() for name in order) + or any( + not isinstance(name, str) or not name.strip() + for name in cast(list[object], order) + ) ): raise WebUISettingsError("model call order must contain at least one preset") - normalized_order = [name.strip() for name in order] + normalized_order = [ + cast(str, name).strip() + for name in cast(list[object], order) + ] config = load_config() _, editable = _model_call_order_state(config) if not editable: @@ -1572,7 +1609,7 @@ def update_model_call_order(query: QueryParams) -> dict[str, Any]: raise WebUISettingsError(f"unknown model preset: {unknown[0]}") defaults = config.agents.defaults - fallback_models = normalized_order[1:] + fallback_models: list[FallbackCandidate] = list(normalized_order[1:]) if ( defaults.model_preset != normalized_order[0] or defaults.fallback_models != fallback_models @@ -1605,7 +1642,7 @@ def migrate_model_configurations(_query: QueryParams | None = None) -> dict[str, defaults.model_preset = name created.append(name) - fallback_models: list[str] = [] + fallback_models: list[FallbackCandidate] = [] for fallback in defaults.fallback_models: if isinstance(fallback, str): fallback_models.append(fallback) diff --git a/nanobot/webui/settings_routes.py b/nanobot/webui/settings_routes.py index 4ff4f4fb0..aa451bc96 100644 --- a/nanobot/webui/settings_routes.py +++ b/nanobot/webui/settings_routes.py @@ -12,7 +12,7 @@ import inspect import json import time from collections.abc import Callable -from typing import Any +from typing import Any, cast from urllib.parse import unquote from websockets.http11 import Request as WsRequest @@ -25,6 +25,7 @@ from nanobot.bus.queue import MessageBus from nanobot.channels._setup import channel_setup_spec from nanobot.channels.connect import ChannelConnectError from nanobot.channels.contracts import ( + RouteFieldType, channel_instance_config, channel_update_instance_config, ) @@ -270,6 +271,7 @@ class WebUISettingsRouter: raise WebUISettingsError("invalid MCP settings payload") from exc if not isinstance(payload, dict): raise WebUISettingsError("MCP settings payload must be a JSON object") + payload = cast(dict[object, Any], payload) merged = {key: list(values) for key, values in query.items()} for key, value in payload.items(): if not isinstance(key, str) or not key: @@ -300,6 +302,7 @@ class WebUISettingsRouter: raise WebUISettingsError("invalid provider settings payload") from exc if not isinstance(payload, dict): raise WebUISettingsError("provider settings payload must be a JSON object") + payload = cast(dict[object, Any], payload) merged = {key: list(values) for key, values in query.items()} for key, value in payload.items(): @@ -553,6 +556,7 @@ class WebUISettingsRouter: raise WebUISettingsError("invalid API service settings payload") from exc if not isinstance(payload, dict): raise WebUISettingsError("API service settings payload must be a JSON object") + payload = cast(dict[str, Any], payload) unknown = set(payload) - {"api_key"} if unknown: @@ -787,7 +791,10 @@ class WebUISettingsRouter: message=f"{name} channel config was saved, but hot reload failed: {exc}", ) - if not isinstance(result, dict) or not result.get("handled"): + if not isinstance(result, dict): + return payload + result = cast(dict[str, Any], result) + if not result.get("handled"): return payload payload = dict(payload) @@ -914,7 +921,7 @@ class WebUISettingsRouter: raise WebUISettingsError("invalid channel settings payload") from exc if not isinstance(payload, dict): raise WebUISettingsError("channel settings payload must be a JSON object") - return payload + return cast(dict[str, Any], payload) def _save_channel_config_values( self, @@ -946,7 +953,7 @@ class WebUISettingsRouter: saved: list[str] = [] prefix = f"channels.{name}." for raw_key, raw_value in raw_values.items(): - if not isinstance(raw_key, str) or not raw_key: + if not raw_key: raise WebUISettingsError("channel settings payload contains an invalid key") field = raw_key[len(prefix):] if raw_key.startswith(prefix) else raw_key value_type = field_types.get(field) @@ -975,7 +982,11 @@ class WebUISettingsRouter: return saved @staticmethod - def _coerce_channel_value(raw_key: str, raw_value: Any, value_type: Any) -> Any: + def _coerce_channel_value( + raw_key: str, + raw_value: Any, + value_type: RouteFieldType, + ) -> Any: if isinstance(value_type, tuple): kind = value_type[0] allowed = value_type[1] @@ -995,7 +1006,7 @@ class WebUISettingsRouter: if isinstance(raw_value, str): return [item.strip() for item in raw_value.split(",") if item.strip()] if isinstance(raw_value, list): - return [str(item).strip() for item in raw_value if str(item).strip()] + return [str(item).strip() for item in cast(list[Any], raw_value) if str(item).strip()] raise WebUISettingsError(f"'{raw_key}' must be a comma-separated list") if kind == "int": @@ -1020,8 +1031,8 @@ class WebUISettingsRouter: value = raw_value.strip() if isinstance(raw_value, str) else str(raw_value) if not value: return _SKIP_FIELD - if value not in allowed: - options = ", ".join(sorted(allowed)) + if allowed is None or value not in allowed: + options = ", ".join(sorted(allowed or ())) raise WebUISettingsError(f"'{raw_key}' must be one of: {options}") return value @@ -1029,14 +1040,14 @@ class WebUISettingsRouter: @staticmethod def _assign_channel_config_value(channel_config: dict[str, Any], field: str, value: Any) -> None: - target = channel_config + target: dict[str, Any] = channel_config parts = field.split(".") for part in parts[:-1]: - current = target.get(part) + current: object = target.get(part) if not isinstance(current, dict): current = {} target[part] = current - target = current + target = cast(dict[str, Any], current) target[parts[-1]] = value async def _handle_settings_channel_connect( @@ -1162,7 +1173,7 @@ class WebUISettingsRouter: def _pairing_payload(last_action: dict[str, Any] | None = None) -> dict[str, Any]: now = time.time() - requests = [] + requests: list[dict[str, Any]] = [] for item in list_pending(): expires_at = float(item.get("expires_at", 0) or 0) created_at = float(item.get("created_at", 0) or 0) diff --git a/nanobot/webui/sidebar_state.py b/nanobot/webui/sidebar_state.py index 0a2f4cfcc..08ab708f3 100644 --- a/nanobot/webui/sidebar_state.py +++ b/nanobot/webui/sidebar_state.py @@ -11,7 +11,7 @@ import json import os import time from pathlib import Path -from typing import Any +from typing import Any, cast from loguru import logger @@ -66,7 +66,7 @@ def _clean_string_list(value: Any, *, max_len: int = _MAX_KEY_LEN) -> list[str]: return [] out: list[str] = [] seen: set[str] = set() - for item in value[:_MAX_LIST_ITEMS]: + for item in cast(list[Any], value)[:_MAX_LIST_ITEMS]: cleaned = _clean_string(item, max_len=max_len) if cleaned is None or cleaned in seen: continue @@ -79,7 +79,7 @@ def _clean_bool_map(value: Any) -> dict[str, bool]: if not isinstance(value, dict): return {} out: dict[str, bool] = {} - for key, raw in list(value.items())[:_MAX_MAP_ITEMS]: + for key, raw in list(cast(dict[Any, Any], value).items())[:_MAX_MAP_ITEMS]: cleaned_key = _clean_string(key) if cleaned_key is None: continue @@ -91,7 +91,7 @@ def _clean_title_overrides(value: Any) -> dict[str, str]: if not isinstance(value, dict): return {} out: dict[str, str] = {} - for key, raw_title in list(value.items())[:_MAX_MAP_ITEMS]: + for key, raw_title in list(cast(dict[Any, Any], value).items())[:_MAX_MAP_ITEMS]: cleaned_key = _clean_string(key) cleaned_title = _clean_string(raw_title, max_len=_MAX_TITLE_LEN) if cleaned_key is None or cleaned_title is None: @@ -104,7 +104,7 @@ def _clean_tags_by_key(value: Any) -> dict[str, list[str]]: if not isinstance(value, dict): return {} out: dict[str, list[str]] = {} - for key, raw_tags in list(value.items())[:_MAX_MAP_ITEMS]: + for key, raw_tags in list(cast(dict[Any, Any], value).items())[:_MAX_MAP_ITEMS]: cleaned_key = _clean_string(key) if cleaned_key is None: continue @@ -115,16 +115,17 @@ def _clean_tags_by_key(value: Any) -> dict[str, list[str]]: def _clean_view(value: Any) -> dict[str, Any]: - default = default_webui_sidebar_state()["view"] + default: dict[str, Any] = default_webui_sidebar_state()["view"] if not isinstance(value, dict): return dict(default) - density = value.get("density") - sort = value.get("sort") + view = cast(dict[str, Any], value) + density = view.get("density") + sort = view.get("sort") return { "density": density if density in _ALLOWED_DENSITIES else default["density"], - "show_previews": bool(value.get("show_previews", default["show_previews"])), - "show_timestamps": bool(value.get("show_timestamps", default["show_timestamps"])), - "show_archived": bool(value.get("show_archived", default["show_archived"])), + "show_previews": bool(view.get("show_previews", default["show_previews"])), + "show_timestamps": bool(view.get("show_timestamps", default["show_timestamps"])), + "show_archived": bool(view.get("show_archived", default["show_archived"])), "sort": sort if sort in _ALLOWED_SORTS else default["sort"], } @@ -133,6 +134,7 @@ def normalize_webui_sidebar_state(raw: Any) -> dict[str, Any]: """Return a schema-v1 sidebar state from any older/partial input.""" if not isinstance(raw, dict): raw = {} + raw = cast(dict[str, Any], raw) state = default_webui_sidebar_state() state["pinned_keys"] = _clean_string_list(raw.get("pinned_keys")) state["archived_keys"] = _clean_string_list(raw.get("archived_keys")) diff --git a/nanobot/webui/token_usage.py b/nanobot/webui/token_usage.py index 761cb63f8..1e72b5e69 100644 --- a/nanobot/webui/token_usage.py +++ b/nanobot/webui/token_usage.py @@ -8,7 +8,7 @@ import threading import time from datetime import datetime, timedelta, timezone from pathlib import Path -from typing import Any +from typing import Any, Mapping, cast from zoneinfo import ZoneInfo, ZoneInfoNotFoundError from loguru import logger @@ -126,9 +126,10 @@ def _normalize_usage_row(row: dict[str, Any]) -> dict[str, int]: def _normalize_sources(raw: Any, fallback: dict[str, int]) -> dict[str, dict[str, int]]: sources: dict[str, dict[str, int]] = {} if isinstance(raw, dict): - for source, row in raw.items(): - if not isinstance(row, dict): + for source, row_value in cast(dict[Any, Any], raw).items(): + if not isinstance(row_value, dict): continue + row = cast(dict[str, Any], row_value) normalized = _normalize_usage_row(row) if normalized["total_tokens"] <= 0 and normalized["requests"] <= 0: continue @@ -148,14 +149,16 @@ def normalize_token_usage_state(raw: Any) -> dict[str, Any]: state = default_token_usage_state() if not isinstance(raw, dict): return state + raw = cast(dict[str, Any], raw) days_raw = raw.get("days") if not isinstance(days_raw, dict): return state days: dict[str, dict[str, Any]] = {} - for date, row in sorted(days_raw.items())[-_MAX_DAYS_RETAINED:]: - if not isinstance(date, str) or len(date) != 10 or not isinstance(row, dict): + for date, row_value in sorted(cast(dict[Any, Any], days_raw).items())[-_MAX_DAYS_RETAINED:]: + if not isinstance(date, str) or len(date) != 10 or not isinstance(row_value, dict): continue + row = cast(dict[str, Any], row_value) normalized = _normalize_usage_row(row) if normalized["total_tokens"] <= 0 and normalized["requests"] <= 0: continue @@ -232,8 +235,9 @@ def record_token_usage( with _WRITE_LOCK: state = read_token_usage_state() + days_by_date = cast(dict[str, dict[str, Any]], state["days"]) day = _local_day(now, timezone_name=timezone_name) - row = dict(state["days"].get(day) or {"date": day, "requests": 0}) + row: dict[str, Any] = dict(days_by_date.get(day) or {"date": day, "requests": 0}) for key in _USAGE_KEYS: row[key] = _clean_int(row.get(key)) + normalized.get(key, 0) row["requests"] = _clean_int(row.get("requests")) + 1 @@ -243,8 +247,10 @@ def record_token_usage( row["provider_requests"] = _clean_int(row.get("provider_requests")) + 1 source_key = _clean_source(source) - sources = dict(row.get("sources") or {}) - source_row = dict(sources.get(source_key) or {"requests": 0}) + sources: dict[str, dict[str, Any]] = dict( + cast(Mapping[str, dict[str, Any]], row.get("sources") or {}) + ) + source_row: dict[str, Any] = dict(sources.get(source_key) or {"requests": 0}) for key in _USAGE_KEYS: source_row[key] = _clean_int(source_row.get(key)) + normalized.get(key, 0) source_row["requests"] = _clean_int(source_row.get("requests")) + 1 @@ -255,10 +261,9 @@ def record_token_usage( sources[source_key] = source_row row["sources"] = sources - state["days"][day] = row - if len(state["days"]) > _MAX_DAYS_RETAINED: - kept = dict(sorted(state["days"].items())[-_MAX_DAYS_RETAINED:]) - state["days"] = kept + days_by_date[day] = row + if len(days_by_date) > _MAX_DAYS_RETAINED: + state["days"] = dict(sorted(days_by_date.items())[-_MAX_DAYS_RETAINED:]) return write_token_usage_state(state) @@ -285,28 +290,29 @@ def token_usage_payload( now: datetime | None = None, ) -> dict[str, Any]: state = read_token_usage_state() + days_by_date = cast(dict[str, dict[str, Any]], state["days"]) today = datetime.fromisoformat(_local_day(now, timezone_name=timezone_name)).date() start = today - timedelta(days=max(1, days) - 1) day_rows = [ row - for date, row in sorted(state["days"].items()) + for date, row in sorted(days_by_date.items()) if start.isoformat() <= date <= today.isoformat() ] last_30_start = today - timedelta(days=29) last_30 = [ row - for date, row in state["days"].items() + for date, row in days_by_date.items() if last_30_start.isoformat() <= date <= today.isoformat() ] last_365_start = today - timedelta(days=364) last_365 = [ row - for date, row in state["days"].items() + for date, row in days_by_date.items() if last_365_start.isoformat() <= date <= today.isoformat() ] active_dates = { datetime.fromisoformat(date).date() - for date, row in state["days"].items() + for date, row in days_by_date.items() if _clean_int(row.get("total_tokens")) > 0 } current_streak = 0 @@ -324,7 +330,7 @@ def token_usage_payload( running_streak = 1 longest_streak = max(longest_streak, running_streak) - all_rows = list(state["days"].values()) + all_rows = list(days_by_date.values()) return { "days": day_rows, "total_tokens": sum(_clean_int(row.get("total_tokens")) for row in all_rows), diff --git a/nanobot/webui/transcript.py b/nanobot/webui/transcript.py index 54c79d71e..7aec02918 100644 --- a/nanobot/webui/transcript.py +++ b/nanobot/webui/transcript.py @@ -11,7 +11,7 @@ import shutil import time import uuid from pathlib import Path -from typing import Any, Callable, Mapping, NamedTuple +from typing import Any, Callable, Mapping, NamedTuple, cast from urllib.parse import unquote, urlparse from loguru import logger @@ -176,7 +176,7 @@ def _read_transcript_file(path: Path) -> list[dict[str, Any]]: logger.warning("bad jsonl at {} line {}", path, line_no) continue if isinstance(obj, dict): - lines_out.append(obj) + lines_out.append(cast(dict[str, Any], obj)) except OSError as e: logger.warning("read transcript failed {}: {}", path, e) return [] @@ -247,12 +247,13 @@ def _non_negative_int(value: Any) -> int | None: def _normalize_manifest_entry(session_key: str, entry: Any) -> dict[str, Any] | None: if not isinstance(entry, dict): return None - segment_id = entry.get("id") + manifest_entry = cast(dict[str, Any], entry) + segment_id = manifest_entry.get("id") if not isinstance(segment_id, str) or not _TRANSCRIPT_SEGMENT_RE.fullmatch(f"{segment_id}.jsonl"): return None segment_path = _segment_file_path(session_key, segment_id) values = { - key: _non_negative_int(entry.get(key)) + key: _non_negative_int(manifest_entry.get(key)) for key in ("bytes", "turn_count", "user_count") } if not segment_path.is_file() or values["bytes"] != segment_path.stat().st_size: @@ -306,11 +307,16 @@ def _read_segment_manifest_entries(session_key: str) -> list[dict[str, Any]]: return _rebuilt_segment_manifest_entries(session_key) try: data = json.loads(path.read_text(encoding="utf-8")) - raw_segments = data.get("segments") if isinstance(data, dict) else None - if data.get("version") != _TRANSCRIPT_SEGMENT_MANIFEST_VERSION or not isinstance(raw_segments, list): + manifest = cast(dict[str, Any], data) if isinstance(data, dict) else None + raw_segments = manifest.get("segments") if manifest is not None else None + if ( + manifest is None + or manifest.get("version") != _TRANSCRIPT_SEGMENT_MANIFEST_VERSION + or not isinstance(raw_segments, list) + ): return _rebuilt_segment_manifest_entries(session_key) entries: list[dict[str, Any]] = [] - for entry in raw_segments: + for entry in cast(list[Any], raw_segments): normalized = _normalize_manifest_entry(session_key, entry) if normalized is None: return _rebuilt_segment_manifest_entries(session_key) @@ -423,7 +429,8 @@ def _decode_page_cursor(value: str | None) -> int | None: return None if not isinstance(data, dict): return None - before_turn = data.get("before_turn") + cursor_data = cast(dict[str, Any], data) + before_turn = cursor_data.get("before_turn") if ( isinstance(before_turn, bool) or not isinstance(before_turn, int) @@ -630,11 +637,12 @@ def webui_message_source(metadata: dict[str, Any] | None) -> dict[str, str] | No raw = (metadata or {}).get(WEBUI_MESSAGE_SOURCE_METADATA_KEY) if not isinstance(raw, dict): return None - kind = raw.get("kind") - if not is_automation_kind(kind): + source_metadata = cast(dict[str, Any], raw) + kind = source_metadata.get("kind") + if not isinstance(kind, str) or not is_automation_kind(kind): return None source: dict[str, str] = {"kind": kind} - label = raw.get("label") + label = source_metadata.get("label") if isinstance(label, str) and label.strip(): source["label"] = label.strip() return source @@ -824,7 +832,9 @@ def write_session_messages_as_transcript( row: dict[str, Any] = {"event": "user", "chat_id": target_chat_id, "text": text} media = msg.get("media") if isinstance(media, list) and media: - row["media_paths"] = [str(p) for p in media if isinstance(p, str) and p] + row["media_paths"] = [ + str(p) for p in cast(list[Any], media) if isinstance(p, str) and p + ] for key in ("cli_apps", "mcp_presets"): value = msg.get(key) if isinstance(value, list) and value: @@ -833,7 +843,9 @@ def write_session_messages_as_transcript( row = {"event": "message", "chat_id": target_chat_id, "text": text} media = msg.get("media") if isinstance(media, list) and media: - row["media"] = [str(p) for p in media if isinstance(p, str) and p] + row["media"] = [ + str(p) for p in cast(list[Any], media) if isinstance(p, str) and p + ] else: continue rows.append(row) @@ -878,10 +890,18 @@ def build_user_transcript_event( } if paths: event["media_paths"] = paths - apps = [dict(app) for app in (cli_apps or []) if isinstance(app, Mapping)] + apps = [ + dict(cast(Mapping[str, Any], app)) + for app in (cli_apps or []) + if isinstance(app, Mapping) + ] if apps: event["cli_apps"] = apps - presets = [dict(preset) for preset in (mcp_presets or []) if isinstance(preset, Mapping)] + presets = [ + dict(cast(Mapping[str, Any], preset)) + for preset in (mcp_presets or []) + if isinstance(preset, Mapping) + ] if presets: event["mcp_presets"] = presets return event @@ -920,9 +940,9 @@ def _session_user_event( return build_user_transcript_event( chat_id, text, - media_paths=media if isinstance(media, list) else None, - cli_apps=cli_apps if isinstance(cli_apps, list) else None, - mcp_presets=mcp_presets if isinstance(mcp_presets, list) else None, + media_paths=cast(list[Any], media) if isinstance(media, list) else None, + cli_apps=cast(list[Any], cli_apps) if isinstance(cli_apps, list) else None, + mcp_presets=cast(list[Any], mcp_presets) if isinstance(mcp_presets, list) else None, ) @@ -940,7 +960,7 @@ def _session_assistant_event( content = message.get("content") text = content if isinstance(content, str) else "" media = message.get("media") - media_paths = [str(path) for path in media] if isinstance(media, list) else [] + media_paths = [str(path) for path in cast(list[Any], media)] if isinstance(media, list) else [] media_paths = [path for path in media_paths if path] if not text.strip() and not media_paths: return None @@ -1172,14 +1192,18 @@ def inject_missing_user_events_from_session( def _format_tool_call_trace(call: Any) -> str | None: if not call or not isinstance(call, dict): return None - fn = call.get("function") - name = fn.get("name") if isinstance(fn, dict) else None + call_data = cast(dict[str, Any], call) + fn = call_data.get("function") + function_data = cast(dict[str, Any], fn) if isinstance(fn, dict) else None + name = function_data.get("name") if function_data is not None else None if not isinstance(name, str) or not name: - raw_name = call.get("name") + raw_name = call_data.get("name") name = raw_name if isinstance(raw_name, str) else "" if not name: return None - args = (fn.get("arguments") if isinstance(fn, dict) else None) or call.get("arguments") + args = ( + function_data.get("arguments") if function_data is not None else None + ) or call_data.get("arguments") if isinstance(args, str) and args.strip(): return f"{name}({args})" if args and isinstance(args, dict): @@ -1192,17 +1216,18 @@ def tool_trace_lines_from_events(events: Any) -> list[str]: return [] lines: list[str] = [] seen: set[str] = set() - for event in events: + for event in cast(list[Any], events): if not event or not isinstance(event, dict): continue - if event.get("phase") not in {"start", "end", "error"}: + tool_event = cast(dict[str, Any], event) + if tool_event.get("phase") not in {"start", "end", "error"}: continue - call_id = event.get("call_id") + call_id = tool_event.get("call_id") if isinstance(call_id, str) and call_id: if call_id in seen: continue seen.add(call_id) - t = _format_tool_call_trace(event) + t = _format_tool_call_trace(tool_event) if t: lines.append(t) return lines @@ -1215,16 +1240,18 @@ def _normalize_tool_events(events: Any) -> list[dict[str, Any]]: if not isinstance(events, list): return [] out: list[dict[str, Any]] = [] - for event in events: + for event in cast(list[Any], events): if not event or not isinstance(event, dict): continue - if event.get("phase") not in {"start", "end", "error"}: + tool_event = cast(dict[str, Any], event) + if tool_event.get("phase") not in {"start", "end", "error"}: continue - if not isinstance(event.get("name"), str): - fn = event.get("function") - if not (isinstance(fn, dict) and isinstance(fn.get("name"), str)): + if not isinstance(tool_event.get("name"), str): + fn = tool_event.get("function") + function = cast(dict[str, Any], fn) if isinstance(fn, dict) else None + if function is None or not isinstance(function.get("name"), str): continue - out.append(dict(event)) + out.append(tool_event) return out @@ -1242,7 +1269,8 @@ def _tool_event_file_edit_key(event: dict[str, Any]) -> str | None: name = event.get("name") if not isinstance(name, str) or not name: fn = event.get("function") - name = fn.get("name") if isinstance(fn, dict) else "" + function = cast(dict[str, Any], fn) if isinstance(fn, dict) else None + name = function.get("name") if function is not None else "" if not isinstance(name, str) or name not in _FILE_EDIT_TOOL_NAMES: return None return f"{call_id}|{name}" @@ -1252,8 +1280,16 @@ def _merge_tool_events(previous: Any, incoming: list[dict[str, Any]]) -> list[di if not isinstance(previous, list) or not previous: return incoming if not incoming: - return [dict(event) for event in previous if isinstance(event, dict)] - merged = [dict(event) for event in previous if isinstance(event, dict)] + return [ + cast(dict[str, Any], event) + for event in cast(list[Any], previous) + if isinstance(event, dict) + ] + merged = [ + cast(dict[str, Any], event) + for event in cast(list[Any], previous) + if isinstance(event, dict) + ] index_by_key = {_tool_event_key(event): idx for idx, event in enumerate(merged)} for event in incoming: key = _tool_event_key(event) @@ -1300,8 +1336,9 @@ def _message_has_file_edit_for_tool_event( if not isinstance(edits, list): return False return any( - isinstance(edit, dict) and _file_edit_tool_event_key(edit) == key - for edit in edits + _file_edit_tool_event_key(cast(dict[str, Any], edit)) == key + for edit in cast(list[Any], edits) + if isinstance(edit, dict) ) @@ -1325,7 +1362,6 @@ def _strip_covered_file_edit_tool_hints( incoming_keys = { _file_edit_tool_event_key(edit) for edit in edits - if isinstance(edit, dict) } events = message.get("toolEvents") if not incoming_keys or not isinstance(events, list): @@ -1334,21 +1370,24 @@ def _strip_covered_file_edit_tool_hints( kept_events: list[dict[str, Any]] = [] removed_trace_lines: set[str] = set() changed = False - for event in events: + for event in cast(list[Any], events): if not isinstance(event, dict): continue - key = _tool_event_file_edit_key(event) + tool_event = cast(dict[str, Any], event) + key = _tool_event_file_edit_key(tool_event) if key and key in incoming_keys: changed = True - removed_trace_lines.update(tool_trace_lines_from_events([event])) + removed_trace_lines.update(tool_trace_lines_from_events([tool_event])) continue - kept_events.append(event) + kept_events.append(tool_event) if not changed: return message raw_traces = message.get("traces") if isinstance(raw_traces, list): - previous_traces = [trace for trace in raw_traces if isinstance(trace, str)] + previous_traces = [ + trace for trace in cast(list[Any], raw_traces) if isinstance(trace, str) + ] else: content = message.get("content") previous_traces = [content] if isinstance(content, str) and content else [] @@ -1383,14 +1422,17 @@ def _merge_unique_tool_trace_lines( def _media_from_signed_urls(value: Any) -> list[dict[str, Any]]: media: list[dict[str, Any]] = [] - urls = value if isinstance(value, list) else [] + urls = cast(list[Any], value) if isinstance(value, list) else [] for m in urls: - if isinstance(m, dict) and m.get("url"): - name = str(m.get("name") or "") + if isinstance(m, dict): + media_item = cast(dict[str, Any], m) + if not media_item.get("url"): + continue + name = str(media_item.get("name") or "") media.append( { "kind": _media_kind_from_name(name), - "url": str(m["url"]), + "url": str(media_item["url"]), "name": name, }, ) @@ -1465,11 +1507,12 @@ def replay_transcript_to_ui_messages( source = rec.get("source") if not isinstance(source, dict): return {} - kind = source.get("kind") - if not is_automation_kind(kind): + source_data = cast(dict[str, Any], source) + kind = source_data.get("kind") + if not isinstance(kind, str) or not is_automation_kind(kind): return {} out: dict[str, Any] = {"source": {"kind": kind}} - label = source.get("label") + label = source_data.get("label") if isinstance(label, str) and label.strip(): out["source"]["label"] = label.strip() return out @@ -1669,11 +1712,10 @@ def replay_transcript_to_ui_messages( segment: str | None, edits: list[dict[str, Any]], ) -> int | None: - incoming_keys = {_file_edit_key(edit) for edit in edits if isinstance(edit, dict)} + incoming_keys = {_file_edit_key(edit) for edit in edits} incoming_tool_event_keys = { _file_edit_tool_event_key(edit) for edit in edits - if isinstance(edit, dict) } for i in range(len(messages) - 1, -1, -1): candidate = messages[i] @@ -1685,15 +1727,16 @@ def replay_transcript_to_ui_messages( return i existing_edits = candidate.get("fileEdits") if isinstance(existing_edits, list): - for existing in existing_edits: + for existing in cast(list[Any], existing_edits): if not isinstance(existing, dict): continue + existing_edit = cast(dict[str, Any], existing) if ( - _file_edit_key(existing) in incoming_keys + _file_edit_key(existing_edit) in incoming_keys or ( - not existing.get("path") - and existing.get("pending") - and _file_edit_tool_event_key(existing) in incoming_tool_event_keys + not existing_edit.get("path") + and existing_edit.get("pending") + and _file_edit_tool_event_key(existing_edit) in incoming_tool_event_keys ) ): return i @@ -1702,7 +1745,10 @@ def replay_transcript_to_ui_messages( def trace_message_is_empty(message: dict[str, Any]) -> bool: traces = message.get("traces") if isinstance(traces, list): - has_trace = any(isinstance(trace, str) and trace.strip() for trace in traces) + has_trace = any( + isinstance(trace, str) and trace.strip() + for trace in cast(list[Any], traces) + ) else: has_trace = bool(str(message.get("content") or "").strip()) return ( @@ -1784,15 +1830,14 @@ def replay_transcript_to_ui_messages( if not segment: segment = _new_activity_segment(activate=False) active_file_edit_segment_id = segment - existing = list(last.get("fileEdits") or []) + raw_existing: Any = last.get("fileEdits") or [] + existing: list[Any] = list(cast(list[Any], raw_existing)) if isinstance(raw_existing, list) else [] index_by_key = { - _file_edit_key(edit): pos + _file_edit_key(cast(dict[str, Any], edit)): pos for pos, edit in enumerate(existing) if isinstance(edit, dict) } for edit in edits: - if not isinstance(edit, dict): - continue key = _file_edit_key(edit) pos = index_by_key.get(key) if pos is None and edit.get("path"): @@ -1800,9 +1845,9 @@ def replay_transcript_to_ui_messages( for existing_pos, existing_edit in enumerate(existing): if ( isinstance(existing_edit, dict) - and not existing_edit.get("path") - and existing_edit.get("pending") - and _file_edit_tool_event_key(existing_edit) == event_key + and not cast(dict[str, Any], existing_edit).get("path") + and cast(dict[str, Any], existing_edit).get("pending") + and _file_edit_tool_event_key(cast(dict[str, Any], existing_edit)) == event_key ): pos = existing_pos break @@ -1832,7 +1877,7 @@ def replay_transcript_to_ui_messages( media_paths = rec.get("media_paths") paths: list[str] = [] if isinstance(media_paths, list): - paths = [str(p) for p in media_paths if p] + paths = [str(p) for p in cast(list[Any], media_paths) if p] media_att: list[dict[str, Any]] | None = None if paths and augment_user_media is not None: media_att = augment_user_media(paths) @@ -1849,11 +1894,15 @@ def replay_transcript_to_ui_messages( row["images"] = [{"url": m.get("url"), "name": m.get("name")} for m in media_att] cli_apps = rec.get("cli_apps") if isinstance(cli_apps, list) and cli_apps: - row["cliApps"] = [dict(app) for app in cli_apps if isinstance(app, dict)] + row["cliApps"] = [ + dict(cast(dict[str, Any], app)) for app in cast(list[Any], cli_apps) if isinstance(app, dict) + ] mcp_presets = rec.get("mcp_presets") if isinstance(mcp_presets, list) and mcp_presets: row["mcpPresets"] = [ - dict(preset) for preset in mcp_presets if isinstance(preset, dict) + dict(cast(dict[str, Any], preset)) + for preset in cast(list[Any], mcp_presets) + if isinstance(preset, dict) ] messages.append(row) continue @@ -1862,7 +1911,7 @@ def replay_transcript_to_ui_messages( raw_edits = rec.get("edits") if isinstance(raw_edits, list): upsert_file_edits( - [e for e in raw_edits if isinstance(e, dict)], + [cast(dict[str, Any], e) for e in cast(list[Any], raw_edits) if isinstance(e, dict)], idx, _turn_fields(rec, "activity"), _created_at_ms(rec, idx), @@ -2011,7 +2060,11 @@ def replay_transcript_to_ui_messages( and not last.get("isStreaming") and (last.get("activitySegmentId") in (None, segment)) ): - prev_traces = list(last.get("traces") or [last.get("content")]) + prev_traces = [ + trace + for trace in cast(list[Any], last.get("traces") or [last.get("content")]) + if isinstance(trace, str) + ] if structured: merged_traces, added = _merge_unique_tool_trace_lines(prev_traces, structured) if not added and not visible_structured_events: @@ -2051,7 +2104,7 @@ def replay_transcript_to_ui_messages( content_s = text if isinstance(text, str) else "" media: list[dict[str, Any]] = [] raw_media = rec.get("media") - raw_media_list = raw_media if isinstance(raw_media, list) else [] + raw_media_list = cast(list[Any], raw_media) if isinstance(raw_media, list) else [] media_paths = [path for path in raw_media_list if isinstance(path, str) and path] if media_paths and augment_assistant_media is not None: media = augment_assistant_media(media_paths) @@ -2218,7 +2271,7 @@ def build_webui_thread_response( augment_assistant_media=augment_assistant_media, augment_assistant_text=augment_assistant_text, ) - payload = { + payload: dict[str, Any] = { "schemaVersion": WEBUI_TRANSCRIPT_SCHEMA_VERSION, "sessionKey": session_key, "messages": msgs, diff --git a/nanobot/webui/workspaces.py b/nanobot/webui/workspaces.py index 6076f6d77..56dc205f7 100644 --- a/nanobot/webui/workspaces.py +++ b/nanobot/webui/workspaces.py @@ -6,7 +6,7 @@ import json import os import time from pathlib import Path -from typing import Any +from typing import TYPE_CHECKING, Any, cast from loguru import logger @@ -20,6 +20,9 @@ from nanobot.security.workspace_access import ( validate_workspace_scope_payload, ) +if TYPE_CHECKING: + from nanobot.session.manager import SessionManager + WEBUI_WORKSPACE_STATE_SCHEMA_VERSION = 1 _MAX_STATE_FILE_BYTES = 128 * 1024 _DEFAULT_ACCESS_MODES = {"default", "full"} @@ -50,6 +53,7 @@ def default_webui_workspace_state() -> dict[str, Any]: def normalize_webui_workspace_state(raw: Any) -> dict[str, Any]: if not isinstance(raw, dict): raw = {} + raw = cast(dict[str, Any], raw) state = default_webui_workspace_state() updated_at = raw.get("updated_at") state["updated_at"] = updated_at if isinstance(updated_at, str) else None @@ -173,7 +177,7 @@ class WebUIWorkspaceController: def __init__( self, *, - session_manager: Any | None, + session_manager: SessionManager | None, default_workspace: Path, default_restrict_to_workspace: bool, ) -> None: @@ -190,14 +194,12 @@ class WebUIWorkspaceController: def scope_for_session_key(self, session_key: str) -> WorkspaceScope: if self._sessions is None: return self.default_scope() - metadata_reader = getattr(self._sessions, "read_session_metadata", None) - if callable(metadata_reader): - data = metadata_reader(session_key) - else: - data = self._sessions.read_session_file(session_key) - metadata = data.get("metadata", {}) if isinstance(data, dict) else {} + data = self._sessions.read_session_metadata(session_key) + session_data = data if data is not None else {} + metadata = session_data.get("metadata", {}) if not isinstance(metadata, dict) or WORKSPACE_SCOPE_METADATA_KEY not in metadata: return self.default_scope() + metadata = cast(dict[str, Any], metadata) try: return validate_workspace_scope_payload( metadata.get(WORKSPACE_SCOPE_METADATA_KEY), diff --git a/nanobot/webui/ws_http.py b/nanobot/webui/ws_http.py index fb3a373b6..e82e53b2a 100644 --- a/nanobot/webui/ws_http.py +++ b/nanobot/webui/ws_http.py @@ -16,7 +16,7 @@ import re import time from collections.abc import Callable from pathlib import Path -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, cast from urllib.parse import unquote from loguru import logger @@ -97,6 +97,7 @@ _AUTOMATION_VALUES_HEADER = "X-Nanobot-Automation-Values" if TYPE_CHECKING: from nanobot.bus.queue import MessageBus + from nanobot.channels.websocket.runtime import WebSocketConfig from nanobot.cron.service import CronService from nanobot.session.manager import SessionManager from nanobot.triggers.local_store import LocalTriggerStore @@ -151,7 +152,7 @@ class GatewayHTTPHandler: def __init__( self, *, - config: Any, # WebSocketConfig + config: WebSocketConfig, session_manager: SessionManager | None, static_dist_path: Path | None, runtime_model_name: Callable[[], str | None] | None, @@ -410,7 +411,7 @@ class GatewayHTTPHandler: sessions = list_webui_sessions(self.session_manager) from nanobot.session.webui_turns import websocket_turn_wall_started_at - cleaned = [] + cleaned: list[dict[str, Any]] = [] for s in sessions: key = s.get("key") if not (isinstance(key, str) and key.startswith("websocket:")): @@ -440,9 +441,15 @@ class GatewayHTTPHandler: return _http_error(404, "session not found") messages = data.get("messages") if isinstance(messages, list): - scrub_subagent_messages_for_channel(messages) + session_messages = cast(list[dict[str, Any]], messages) + scrub_subagent_messages_for_channel(session_messages) + raw_session_messages = cast(list[Any], messages) data["messages"] = public_history_messages( - message for message in messages if isinstance(message, dict) + [ + cast(dict[str, Any], message) + for message in raw_session_messages + if isinstance(message, dict) + ] ) self.media.augment_media_urls(data) return _http_json_response(data) @@ -461,7 +468,12 @@ class GatewayHTTPHandler: session_data = self.session_manager.read_session_file(decoded_key) raw_messages = session_data.get("messages") if isinstance(session_data, dict) else None if isinstance(raw_messages, list): - session_messages = [m for m in raw_messages if isinstance(m, dict)] + raw_session_messages = cast(list[Any], raw_messages) + session_messages = [ + cast(dict[str, Any], raw_message) + for raw_message in raw_session_messages + if isinstance(raw_message, dict) + ] query = _parse_query(request.path) raw_limit = _query_first(query, "limit") limit: int | None = None @@ -854,7 +866,7 @@ class GatewayHTTPHandler: if not isinstance(decoded, dict): return _http_error(400, "state must be an object") try: - state = write_webui_sidebar_state(decoded) + state = write_webui_sidebar_state(cast(dict[str, Any], decoded)) except ValueError as e: return _http_error(400, str(e)) except OSError: @@ -915,7 +927,7 @@ def _automation_values_from_request(request: WsRequest) -> dict[str, Any] | None values = json.loads(unquote(raw)) except Exception: return None - return values if isinstance(values, dict) else None + return cast(dict[str, Any], values) if isinstance(values, dict) else None def _parse_automation_update( @@ -944,7 +956,7 @@ def _parse_automation_update( raw_schedule = values.get("schedule") if not isinstance(raw_schedule, dict): return "schedule must be an object" - parsed_schedule = _parse_automation_schedule(raw_schedule) + parsed_schedule = _parse_automation_schedule(cast(dict[str, Any], raw_schedule)) if isinstance(parsed_schedule, str): return parsed_schedule if current_job is not None and _schedule_matches_job(parsed_schedule, current_job): @@ -1034,7 +1046,7 @@ def _validate_automation_schedule(schedule: CronSchedule) -> str | None: tz = ZoneInfo(schedule.tz) if schedule.tz else datetime.now().astimezone().tzinfo base = datetime.now(tz=tz) - croniter(schedule.expr, base).get_next(datetime) + croniter(cast(str, schedule.expr), base).get_next(datetime) except Exception: return "cron schedule is invalid" return None diff --git a/pyproject.toml b/pyproject.toml index 129fbe33b..a815f66af 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -94,6 +94,7 @@ dev = [ "pytest-cov>=6.0.0,<7.0.0", "pytest-xdist>=3.8.0,<4.0.0", "ruff>=0.1.0", + "basedpyright>=1.39.0,<2.0.0", "pymupdf>=1.25.0", "pypdf>=5.0.0,<6.0.0", "python-docx>=1.1.0,<2.0.0", @@ -164,6 +165,12 @@ target-version = "py311" select = ["E", "F", "I", "N", "W"] ignore = ["E501"] +[tool.basedpyright] +include = ["nanobot"] +exclude = ["**/tests"] +typeCheckingMode = "strict" +pythonVersion = "3.11" + [tool.pytest.ini_options] asyncio_mode = "auto" testpaths = ["tests", "nanobot/channels"] diff --git a/tests/agent/test_runner_governance.py b/tests/agent/test_runner_governance.py index 71a3c48cb..1c2a30279 100644 --- a/tests/agent/test_runner_governance.py +++ b/tests/agent/test_runner_governance.py @@ -891,7 +891,8 @@ def test_drop_malformed_tool_calls_trims_response(): tool_calls=[ ToolCallRequest(id="1", name=None, arguments={}), ToolCallRequest(id="2", name="", arguments={}), - ToolCallRequest(id="3", name="read_file", arguments={}), + ToolCallRequest(id="3", name={"unexpected": "object"}, arguments={}), + ToolCallRequest(id="4", name="read_file", arguments={}), ], finish_reason="tool_calls", ) @@ -899,7 +900,7 @@ def test_drop_malformed_tool_calls_trims_response(): assert [tc.name for tc in response.tool_calls] == ["read_file"] assert response.finish_reason == "tool_calls" assert response.should_execute_tools is True - assert dropped == 2 + assert dropped == 3 assert all_dropped is False assert orig == "tool_calls" diff --git a/tests/channels/test_channel_plugins.py b/tests/channels/test_channel_plugins.py index a57160ea2..fb904be35 100644 --- a/tests/channels/test_channel_plugins.py +++ b/tests/channels/test_channel_plugins.py @@ -1217,6 +1217,7 @@ def test_channels_login_uses_discovered_plugin_class(monkeypatch): async def login(self, force: bool = False) -> bool: seen["force"] = force seen["config"] = self.config + seen["bus"] = self.bus return True monkeypatch.setattr("nanobot.config.loader.load_config", lambda config_path=None: Config()) @@ -1229,6 +1230,7 @@ def test_channels_login_uses_discovered_plugin_class(monkeypatch): assert result.exit_code == 0 assert seen["force"] is True + assert isinstance(seen["bus"], MessageBus) def test_channels_login_sets_custom_config_path(monkeypatch, tmp_path): diff --git a/tests/cli/test_commands.py b/tests/cli/test_commands.py index 898b35e6c..e29a2c2bb 100644 --- a/tests/cli/test_commands.py +++ b/tests/cli/test_commands.py @@ -13,6 +13,7 @@ import pytest from typer.testing import CliRunner from nanobot.agent.memory import MemoryStore +from nanobot.agent.tools.registry import ToolRegistry from nanobot.agent.turn_delivery import TurnDeliveryFactory from nanobot.bus.events import InboundMessage, OutboundMessage from nanobot.cli import commands as cli_commands @@ -71,6 +72,26 @@ class _StopGatewayError(RuntimeError): pass +class _GatewayAgentContractStub: + """Minimal stable AgentLoop surface required by gateway assembly tests.""" + + tools = ToolRegistry() + + @staticmethod + def pending_cron_job_ids_for_session(_session_key: str) -> set[str]: + return set() + + @staticmethod + def pending_local_trigger_ids_for_session(_session_key: str) -> set[str]: + return set() + + async def submit_local_trigger_turn( + self, + _msg: InboundMessage, + ) -> OutboundMessage | None: + return None + + def test_gateway_signal_handler_first_signal_stops_and_second_forces() -> None: class _FakeLoop: def __init__(self) -> None: @@ -1949,7 +1970,7 @@ def test_heartbeat_empty_response_still_retains_recent_messages( def register_system_job(self, _job: CronJob) -> None: raise _StopGatewayError("stop") - class _FakeAgentLoop: + class _FakeAgentLoop(_GatewayAgentContractStub): @classmethod def from_config(cls, config, bus=None, **extra): return cls(**extra) @@ -2615,7 +2636,7 @@ def test_gateway_unbound_agent_cron_is_skipped( self.on_job = None seen["cron"] = self - class _FakeAgentLoop: + class _FakeAgentLoop(_GatewayAgentContractStub): @classmethod def from_config(cls, config, bus=None, **extra): return cls(**extra) @@ -2731,7 +2752,7 @@ def test_gateway_bound_cron_runs_as_session_turn( def write_run_record(self, run_id: str, record: dict[str, object]) -> None: seen["run_records"].append((run_id, record)) - class _FakeAgentLoop: + class _FakeAgentLoop(_GatewayAgentContractStub): @classmethod def from_config(cls, config, bus=None, **extra): return cls(**extra) @@ -2947,7 +2968,7 @@ def test_gateway_local_trigger_queue_submits_agent_turns( def register_system_job(self, _job) -> None: return None - class _FakeAgentLoop: + class _FakeAgentLoop(_GatewayAgentContractStub): @classmethod def from_config(cls, config, bus=None, **extra): seen["agent_from_config_kwargs"] = extra @@ -3200,7 +3221,7 @@ def test_gateway_health_endpoint_binds_and_serves_expected_responses( def flush_all(self) -> int: return 0 - class _FakeAgentLoop: + class _FakeAgentLoop(_GatewayAgentContractStub): @classmethod def from_config(cls, config, bus=None, **extra): return cls(**extra) @@ -3393,7 +3414,7 @@ def test_gateway_shutdown_lets_agent_task_own_mcp_cleanup( def flush_all(self) -> int: return 0 - class _FakeAgentLoop: + class _FakeAgentLoop(_GatewayAgentContractStub): @classmethod def from_config(cls, config, bus=None, **extra): return cls(**extra) @@ -3492,7 +3513,7 @@ def test_gateway_shutdown_event_exits_forever_runtime_tasks( def flush_all(self) -> int: return 0 - class _FakeAgentLoop: + class _FakeAgentLoop(_GatewayAgentContractStub): @classmethod def from_config(cls, config, bus=None, **extra): return cls(**extra) diff --git a/tests/command/test_builtin_dream.py b/tests/command/test_builtin_dream.py index b7c9aa5d1..1da8109c5 100644 --- a/tests/command/test_builtin_dream.py +++ b/tests/command/test_builtin_dream.py @@ -231,6 +231,7 @@ def _build_runnable_dream( context=SimpleNamespace(memory=store, timezone="UTC"), sessions=SimpleNamespace(sessions_dir=sessions_dir), process_direct=process_direct, + dream_runtime=lambda: None, ) ctx = CommandContext(msg=msg, session=None, key=msg.session_key, raw="/dream", args="", loop=loop) return ctx, store @@ -317,6 +318,7 @@ async def test_dream_noop_batch_unlocks_following_history(tmp_path) -> None: context=SimpleNamespace(memory=store, timezone="UTC"), sessions=SimpleNamespace(sessions_dir=sessions_dir), process_direct=process_direct, + dream_runtime=lambda: None, ) ctx = CommandContext(msg=msg, session=None, key=msg.session_key, raw="/dream", args="", loop=loop) diff --git a/tests/command/test_router_dispatchable.py b/tests/command/test_router_dispatchable.py index 6c3fbadf5..59ecec602 100644 --- a/tests/command/test_router_dispatchable.py +++ b/tests/command/test_router_dispatchable.py @@ -2,6 +2,7 @@ from __future__ import annotations +from inspect import Parameter, signature from unittest.mock import AsyncMock, MagicMock import pytest @@ -13,6 +14,13 @@ from nanobot.command.builtin import ( from nanobot.command.router import CommandContext, CommandRouter +def test_command_context_requires_loop_as_keyword_dependency() -> None: + loop_parameter = signature(CommandContext).parameters["loop"] + + assert loop_parameter.kind is Parameter.KEYWORD_ONLY + assert loop_parameter.default is Parameter.empty + + class TestIsDispatchableCommand: """Unit tests for the is_dispatchable_command() predicate.""" diff --git a/tests/config/test_env_interpolation.py b/tests/config/test_env_interpolation.py index 9119e226c..21ce71490 100644 --- a/tests/config/test_env_interpolation.py +++ b/tests/config/test_env_interpolation.py @@ -7,6 +7,7 @@ from nanobot.config.loader import ( _resolve_env_vars, load_config, resolve_config_env_vars, + resolve_env_refs, save_config, ) from nanobot.config.schema import Config @@ -50,6 +51,12 @@ class TestResolveEnvVars: _resolve_env_vars("${DOES_NOT_EXIST}") +class TestResolveSingleEnvRefs: + @pytest.mark.parametrize("value", [None, 42, True, {"key": "value"}]) + def test_non_string_values_pass_through_unchanged(self, value): + assert resolve_env_refs(value) is value + + class TestResolveConfig: def test_resolves_env_vars_in_config(self, tmp_path, monkeypatch): monkeypatch.setenv("TEST_API_KEY", "resolved-key") diff --git a/tests/cron/test_cron_service.py b/tests/cron/test_cron_service.py index 0d91d6a9d..36079e0ae 100644 --- a/tests/cron/test_cron_service.py +++ b/tests/cron/test_cron_service.py @@ -65,6 +65,17 @@ def test_load_jobs_accepts_snake_case_schedule_and_run_history(tmp_path) -> None assert jobs[0].state.run_history[0].duration_ms == 12 +def test_cron_job_from_dict_rejects_malformed_run_history() -> None: + with pytest.raises(TypeError): + CronJob.from_dict( + { + "id": "j1", + "name": "t", + "state": {"run_history": [None]}, + } + ) + + def test_load_jobs_coerces_string_schedule_and_state_ms(tmp_path) -> None: store_path = tmp_path / "cron" / "jobs.json" store_path.parent.mkdir(parents=True) diff --git a/tests/pairing/test_store.py b/tests/pairing/test_store.py index c4b4758af..d16f0705c 100644 --- a/tests/pairing/test_store.py +++ b/tests/pairing/test_store.py @@ -1,3 +1,5 @@ +import json + import pytest from nanobot.pairing import __all__ as pairing_all @@ -272,6 +274,23 @@ def test_load_treats_null_approved_and_pending_maps_as_empty(tmp_path, monkeypat assert store.get_approved("telegram") == [] +@pytest.mark.parametrize( + ("field", "value"), + [("approved", "corrupt"), ("pending", ["corrupt"])], +) +def test_load_treats_non_object_approved_and_pending_maps_as_empty( + tmp_path, monkeypatch, field, value +): + path = tmp_path / "pairing.json" + payload = {"approved": {}, "pending": {}} + payload[field] = value + path.write_text(json.dumps(payload), encoding="utf-8") + monkeypatch.setattr(store, "_store_path", lambda: path) + + assert store.is_approved("telegram", "123") is False + assert store.list_pending() == [] + + @pytest.mark.parametrize("payload", ["null", "[]", "true"]) def test_load_treats_non_object_store_as_empty(tmp_path, monkeypatch, payload): path = tmp_path / "pairing.json" diff --git a/tests/providers/test_transcription.py b/tests/providers/test_transcription.py index 8ba08d232..7c4cffd99 100644 --- a/tests/providers/test_transcription.py +++ b/tests/providers/test_transcription.py @@ -181,7 +181,7 @@ def test_resolver_env_ref_missing_var_degrades_to_not_configured() -> None: # Unresolved reference degrades to a falsy key rather than the literal # "${...}" string, so the config reports itself as not configured. - assert not resolved.api_key + assert resolved.api_key == "" assert resolved.configured is False diff --git a/tests/test_api_attachment.py b/tests/test_api_attachment.py index 694461096..a12c77ba0 100644 --- a/tests/test_api_attachment.py +++ b/tests/test_api_attachment.py @@ -156,6 +156,28 @@ def test_parse_json_content_validates_user_role() -> None: _parse_json_content(body) +@pytest.mark.parametrize( + ("part", "field"), + [ + ({"type": "text", "text": 1}, r"content\[\]\.text"), + ( + {"type": "image_url", "image_url": "not-an-object"}, + r"content\[\]\.image_url", + ), + ( + {"type": "image_url", "image_url": {"url": 1}}, + r"image_url\.url", + ), + ], +) +def test_parse_json_content_validates_typed_block_fields(part, field) -> None: + """Dynamic content blocks are checked before their values reach typed code.""" + body = {"messages": [{"role": "user", "content": [part]}]} + + with pytest.raises(TypeError, match=field): + _parse_json_content(body) + + def test_parse_json_content_rejects_oversized_base64_file(tmp_path) -> None: """Oversized JSON data URLs should fail before writing to disk.""" large_payload = base64.b64encode(b"x" * (11 * 1024 * 1024)).decode() diff --git a/tests/test_document_parsing.py b/tests/test_document_parsing.py index 980edf6b5..b2ae4e7d7 100644 --- a/tests/test_document_parsing.py +++ b/tests/test_document_parsing.py @@ -67,6 +67,13 @@ class TestExtractText: result = extract_text(txt_file) assert result == content + def test_extract_text_accepts_string_path(self, tmp_path: Path): + """String paths retain the compatibility behavior of Path inputs.""" + txt_file = tmp_path / "string-path.txt" + txt_file.write_text("string path", encoding="utf-8") + + assert extract_text(str(txt_file)) == "string path" + def test_extract_text_txt_file_with_truncation(self, tmp_path: Path): """Test that large text files are truncated.""" txt_file = tmp_path / "large.txt" diff --git a/tests/tools/test_mcp_tool.py b/tests/tools/test_mcp_tool.py index 229be0149..8ece82788 100644 --- a/tests/tools/test_mcp_tool.py +++ b/tests/tools/test_mcp_tool.py @@ -26,6 +26,14 @@ from nanobot.config.schema import MCPServerConfig _PROXY_ENV_VARS = ("HTTP_PROXY", "HTTPS_PROXY", "ALL_PROXY", "http_proxy", "https_proxy", "all_proxy") +def test_type_checking_only_mcp_annotations_are_deferred() -> None: + assert mcp_mod._MCPWrapperBase.__annotations__["_session"] == "ClientSession" + assert MCPToolWrapper.__init__.__annotations__["session"] == "ClientSession" + assert MCPResourceWrapper.__init__.__annotations__["resource_def"] == "Resource" + assert MCPPromptWrapper.__init__.__annotations__["prompt_def"] == "Prompt" + assert connect_mcp_servers.__annotations__["mcp_servers"] == "dict[str, MCPServerConfig]" + + class _FakeTextContent: def __init__(self, text: str) -> None: self.text = text