Compare commits

..
Author SHA1 Message Date
chengyongru 40e95ba0da refactor(webui): split settings backend by domain 2026-08-10 18:52:59 +08:00
chengyongruandGitHub c281e090d0 refactor(webui): make gateway own settings services (#5321) 2026-08-10 18:10:55 +08:00
chengyongruandGitHub 85a452e5c7 refactor(agent): replace reflective runtime state access (#5319) 2026-08-10 16:44:26 +08:00
chengyongruandchengyongru 05d73803e7 refactor(webui): extract event projection helpers 2026-08-10 16:24:07 +08:00
chengyongruandchengyongru 5d733b1c7c fix(webui): move mutations to authenticated websocket requests 2026-08-10 16:23:47 +08:00
chengyongruandGitHub 71a99b0780 fix(webui): improve UX recovery and empty states (#5315) 2026-08-10 15:22:25 +08:00
chengyongruandchengyongru 43511decc9 fix(weixin): install QR code dependency 2026-08-10 13:48:49 +08:00
chengyongruandchengyongru 8dd2059be3 fix(weixin): require fresh credentials for forced login 2026-08-10 13:48:49 +08:00
KDBandchengyongru 7b1646f58c fix(weixin): honor forced QR login 2026-08-10 13:48:49 +08:00
chengyongruandchengyongru e620944150 fix(mcp): clean up failed HTTP connections 2026-08-10 13:15:04 +08:00
chengyongruandGitHub 66316f21da docs: refresh WebUI user guidance (#5312) 2026-08-10 11:47:58 +08:00
chengyongruandGitHub 55ecda275d test: strengthen user-path coverage and CI gates (#5308) 2026-08-09 21:37:06 +08:00
chengyongruandGitHub 411d6061ae fix(webui): explain HTTPS requirement for voice input (#5304) 2026-08-09 21:17:57 +08:00
Xubin Ren af52fbcbc4 fix(webui): emphasize temporary chat expiry 2026-08-08 23:20:59 +08:00
Xubin Ren 92eb91338a fix(webui): label temporary chat guidance 2026-08-08 23:20:59 +08:00
Xubin Ren c410ea444c fix(agent): stop session-owned exec processes 2026-08-08 23:20:59 +08:00
chengyongruandXubin Ren 8e04f12720 fix(webui): name temporary chats from first message 2026-08-08 23:20:59 +08:00
chengyongruandXubin Ren 516ae11c33 fix(webui): simplify temporary chat closing 2026-08-08 23:20:59 +08:00
chengyongruandXubin Ren 656e0d606b fix(webui): tighten temporary chat close action 2026-08-08 23:20:59 +08:00
chengyongruandXubin Ren 75e333a3c5 fix(webui): derive temporary chats from session policy 2026-08-08 23:20:59 +08:00
chengyongruandXubin Ren a5bc3bfbb9 fix(webui): complete temporary chat mode 2026-08-08 23:20:59 +08:00
Xubin Ren c9a6145878 feat(webui): add temporary chat mode 2026-08-08 23:20:59 +08:00
chengyongruandchengyongru 113e8d67ad refactor: remove verified dead code 2026-08-08 21:10:34 +08:00
chengyongruandGitHub 4e063f5695 fix(webui): prevent image hover clipping (#5294) 2026-08-08 18:05:00 +08:00
chengyongruandchengyongru bd8d3ad5b6 fix(channels): preserve global progress defaults 2026-08-07 17:19:14 +08:00
chengyongruandGitHub 332c159b93 fix(weixin): harden protocol delivery, streaming, and login (#5263) 2026-08-07 16:53:31 +08:00
chengyongruandchengyongru edb3b7e446 fix(webui): preserve newly created topic route 2026-08-07 16:10:01 +08:00
chengyongruandchengyongru cdb2a474f9 refactor(webui): remove legacy session messages route 2026-08-07 15:16:57 +08:00
chengyongruandGitHub ff6deda178 fix: modernize dependency recovery guidance (#5282) 2026-08-07 13:58:13 +08:00
chengyongruandchengyongru 02a002a0e6 fix(webui): preserve activity text rendering 2026-08-07 13:04:42 +08:00
Xubin Ren 3836c32874 fix(webui): scope preset editor to one row 2026-08-07 12:42:37 +08:00
Xubin Ren 3fc69b2922 style(webui): inset expanded preset editor 2026-08-07 12:42:37 +08:00
Xubin Ren eb5d7e1a32 style(webui): distinguish expanded preset editor 2026-08-07 12:42:37 +08:00
Xubin Ren b77e1133cb fix(webui): preserve preset deletion workflow 2026-08-07 12:42:37 +08:00
Xubin Ren 1b12fbae39 fix(webui): explain disabled preset deletion 2026-08-07 12:42:37 +08:00
Xubin Ren 6f2512ce9a style(webui): retain model preset colors 2026-08-07 12:42:37 +08:00
Xubin Ren c8bc4d8510 fix(webui): make active model presets deletable 2026-08-07 12:42:37 +08:00
Xubin Ren e971e81b6c refactor(webui): expand model preset editor inline 2026-08-07 12:42:37 +08:00
Xubin Ren ada07aa799 feat(webui): add responsive model preset detail pane 2026-08-07 12:42:37 +08:00
chengyongruandchengyongru 2c7943a133 fix(memory): archive short idle sessions for Dream 2026-08-07 11:45:49 +08:00
chengyongruandchengyongru 8dfce4c162 fix(session): require user anchor for delivery retention 2026-08-07 10:53:55 +08:00
ziuusandchengyongru 60282d1588 fix(session): preserve proactive channel delivery during session retention trimming 2026-08-07 10:53:55 +08:00
Xubin Ren c2fd41b44d fix(webui): persist large sidebar ordering state 2026-08-06 19:11:05 +08:00
Xubin Ren 1d290614c9 fix(webui): align composer mention metrics 2026-08-06 19:11:05 +08:00
Xubin Ren 9af6bb91c7 fix(webui): preserve session drag contracts 2026-08-06 19:11:05 +08:00
Xubin Ren f44a766f98 feat(webui): preview dragged session mentions 2026-08-06 19:11:05 +08:00
Xubin Ren 9cf6cf0639 feat(webui): persist manual session ordering 2026-08-06 19:11:05 +08:00
Xubin Ren 2c8e63446f feat(webui): drag sessions into composer mentions 2026-08-06 19:11:05 +08:00
Orrin WittandGitHub 5c4c2cb819 fix(matrix): send non-empty POST body on room join for Continuwuity compatibility (#5248) 2026-08-06 18:29:57 +08:00
chengyongruandchengyongru 223b911e7e fix(webui): tighten interactive motion 2026-08-06 18:28:45 +08:00
200 changed files with 18588 additions and 8016 deletions
+1 -1
View File
@@ -173,7 +173,7 @@ jobs:
- name: Test WebUI
working-directory: webui
run: bun run test
run: bun run test:coverage
- name: Build WebUI
working-directory: webui
+3 -2
View File
@@ -241,7 +241,7 @@ Prefer your own infrastructure? Follow the [deployment guide](./docs/deployment.
## 🌐 WebUI
The WebUI ships **inside the published wheel** with no separate frontend build. It is the browser workbench for persistent topics, visible agent activity, workspace controls, Apps, Skills, Automations, and settings.
The WebUI ships **inside the published wheel** with no separate frontend build. It is the browser workbench for persistent topics, temporary chats, visible agent activity, workspace controls, Apps, Skills, Automations, and settings.
<p align="center">
<img src="images/nanobot_webui.png" alt="nanobot webui preview" width="900">
@@ -250,9 +250,10 @@ The WebUI ships **inside the published wheel** with no separate frontend build.
Use it to:
- keep separate topics for different tasks and projects;
- use temporary chats when a conversation should not be saved to history or memory;
- inspect reasoning, tool calls, file edits, diffs, command output, and generated artifacts;
- switch models and workspaces without leaving the conversation;
- configure providers, chat channels, Apps, Skills, and Automations from one place.
- configure providers and chat channels, connect Apps, discover Skills, and manage Automations from one place.
See the [WebUI guide](./docs/webui.md) for LAN access, background operation, workspace controls, and the full feature tour. Working on the frontend itself? Use [`webui/README.md`](./webui/README.md).
@@ -27,7 +27,7 @@ nanobot agent -m "Hello!"
Install Langfuse:
```bash
python -m pip install langfuse
nanobot plugins enable langfuse
```
## Minimal working example
+1 -1
View File
@@ -549,7 +549,7 @@ This recipe applies after the agent works and you want observability for OpenAI-
Install the optional package in the same Python environment that runs nanobot:
```bash
python -m pip install langfuse
nanobot plugins enable langfuse
```
Set the environment variables before starting nanobot:
+6
View File
@@ -270,6 +270,12 @@ http://127.0.0.1:8765
If accessing from another device, bind the WebSocket channel to `0.0.0.0` and set `token` or `tokenIssueSecret`. The WebSocket channel refuses public binds without a token or token issue secret.
| Symptom | Check |
|---|---|
| A temporary chat disappeared after a reload or reconnect | This is expected. Temporary chats exist only for the current WebUI connection and are not saved to history or memory. Use a regular topic for anything you need to retain. |
| A skills.sh install says that `npx` is required | Install Node.js with `npx` on the gateway machine, or choose a SkillHub skill that does not require `npx`. |
| A remote browser says skill installation is disabled | Install from a same-machine WebUI. For a private deployment where every authenticated user is trusted to install third-party skill instructions or scripts, explicitly enable `tools.webuiAllowRemotePackageInstall`. |
See [`webui.md#lan-access`](./webui.md#lan-access) for LAN setup and [`../webui/README.md`](../webui/README.md) for frontend development.
## Chat App Problems
+68 -19
View File
@@ -1,10 +1,10 @@
# Nanobot WebUI: Browser Workbench for Self-Hosted AI Agents
<!-- Meta description: Run nanobot from a browser WebUI with persistent topics, visible tool activity, workspace controls, Apps, MCP presets, Skills, settings, and Automations. -->
<!-- Meta description: Run nanobot from a browser WebUI with persistent and temporary chats, visible tool activity, workspace controls, Apps, skill discovery, settings, and Automations. -->
The WebUI is nanobot's browser workbench for persistent topics, visible
agent activity, workspace controls, Apps, Skills, settings, and Automations in
one place.
The WebUI is nanobot's browser workbench for persistent topics, temporary
chats, visible agent activity, workspace controls, Apps, skill discovery,
settings, and Automations in one place.
The published `nanobot-ai` wheel already includes the WebUI bundle. You only need
the `webui/` source directory when you are changing the frontend itself.
@@ -72,14 +72,14 @@ This path avoids hand-editing `config.json` for normal setup. Use the reference
| Area | Use it for |
|---|---|
| Topics | Start, switch, search, fork, and delete browser topics |
| Topics | Start persistent topics or temporary chats; switch, search, reorder, fork, or delete persistent topics |
| Agent activity | See thinking, tool calls, file edits with diffs, command output, and generated artifacts in context |
| Workspace | Pick the project workspace before asking for file or shell work |
| Access | Choose the access mode for local capabilities allowed by your gateway configuration |
| Composer | Send text, images, voice input, slash commands, and `@` mentions for topics, Apps, or MCP presets |
| Channels | Connect and validate chat platforms, install their optional support, and manage saved channel setup |
| Apps | Install, test, update, and use local CLI App adapters and MCP presets |
| Skills | Inspect available built-in and workspace skills before relying on them |
| Skills | Inspect and manage installed skills, or discover skills from supported marketplaces |
| Automations | Review, search, run, pause, edit, and delete scheduled and local-trigger agent turns |
| Settings | Adjust models, providers, image generation, voice, web tools, runtime, and safety options |
@@ -90,6 +90,10 @@ workspace selection, and linked automations. Use a new topic when you want a
separate context; use fork when you want to continue from an existing point
without changing the original thread.
Drag a topic within its current sidebar group to keep frequently used work in
your preferred order. Drag a topic from the sidebar into the composer when you
want to reference it in the next message instead of switching to it.
The message timeline shows both user-visible replies and agent activity. Long
tool or reasoning sections can be expanded when you need the details.
@@ -103,6 +107,28 @@ File previews follow the active session access mode. Restricted workspace access
previews only files under the selected workspace. Full Access can preview files
outside the workspace when that access mode is allowed by the gateway.
## Temporary Chats
Use a temporary chat for a conversation that should not be added to nanobot's
topic history or long-term memory:
1. Select **New topic**.
2. Select the **Temporary chat** control in the page header.
3. Send the first message.
You can keep more than one temporary chat open and switch between them under
**Temporary chats** in the sidebar while the current WebUI connection remains
open. Reloading or closing the page, restarting the gateway, or losing the
WebSocket connection ends all of them. They cannot be recovered afterward.
Temporary does not mean consequence-free. Requests still go to the configured
model provider, and tools can still change files, run commands, or affect
external services. Temporary chats always use the default workspace in
Restricted mode; the project picker and Full Access are unavailable. Commands
and tools that create durable goals, automations, or subagent work are also
unavailable. Use a regular topic when you need reusable context, scheduled work,
or a result you must retain.
## Workspace and Access
Use the workspace picker before starting project-specific work. This gives the
@@ -145,7 +171,8 @@ clients.
The composer supports plain messages, image attachments, voice input when
transcription is configured, slash commands, and `@` mentions for installed Apps
or MCP presets. Select another topic from the `@` menu to attach a stable
reference; plain text that happens to start with `@` does not attach history.
reference, or drag that topic from the sidebar into the composer. Plain text
that happens to start with `@` does not attach history.
Restricted chats offer topics from the same project, while Full Access chats can
reference any WebUI topic. Nanobot reads a referenced topic only when its history
is relevant and can link it in the response. The model badge shows the current
@@ -204,10 +231,20 @@ After an App or integration is available, mention it from the composer with
## Skills
The Skills view shows the skill instructions available to the agent, including
built-in skills and workspace-provided skills. Check this view when you want to
know whether nanobot already has a focused workflow for a task before you ask it
to perform that task.
Open **Skills → Installed** to review built-in and workspace-provided skills.
You can search and filter them, inspect their instructions and setup
requirements, enable or disable them, and delete workspace skills you no longer
want.
Open **Skills → Discover** to browse or search skills from skills.sh and
SkillHub. A marketplace skill is copied into the active agent workspace after
you confirm the installation. skills.sh installation requires Node.js with
`npx`; SkillHub installation does not.
Marketplace skills are third-party instructions and may include executable
scripts. Review the source and instructions before installing one, and enable
only skills you trust with the same files, tools, and credentials available to
your agent.
## Automations
@@ -288,10 +325,17 @@ The gateway refuses to start with `host` set to `"0.0.0.0"` unless `token` or
`http://<your-ip>:8765` from the other device and enter the secret in the login
form.
Remote WebUI clients with a valid token can view and use Apps. Actions that
install missing nanobot support packages, such as adding a channel dependency,
are blocked by default. To let trusted remote administrators change the Python
environment through the WebUI, opt in explicitly:
Plain HTTP is enough for basic WebUI access, but browsers expose microphone
capture only in secure contexts. Voice input works on same-machine localhost;
from another device, serve the WebUI over HTTPS with a certificate that device
trusts. Configure [`sslCertfile` and `sslKeyfile`](./websocket.md#tlsssl) on the
WebSocket channel and open `https://<your-host>:8765`, or terminate HTTPS at a
reverse proxy and use that proxy's HTTPS URL.
Remote WebUI clients with a valid token can view and use Apps and installed
skills. Actions that install missing nanobot support packages or third-party
marketplace skills are blocked by default. To let trusted remote administrators
perform those installations through the WebUI, opt in explicitly:
```json
{
@@ -302,12 +346,13 @@ environment through the WebUI, opt in explicitly:
```
Use this only for a private deployment where every authenticated WebUI user is
trusted to change the Python environment that nanobot runs in. If you publish
the WebUI through Nginx, Caddy, Cloudflare Tunnel, or a similar service, treat it
as remote access and leave package installs disabled unless that is intentional.
trusted to change nanobot's Python environment and install workspace skill
instructions or scripts. If you publish the WebUI through Nginx, Caddy,
Cloudflare Tunnel, or a similar service, treat it as remote access and leave
package and skill installs disabled unless that is intentional.
Optional feature installs use pip's configured package index, including
`PIP_INDEX_URL`.
`PIP_INDEX_URL`. skills.sh marketplace installs use `npx` instead.
Leave remote package installs disabled when the WebUI is exposed beyond a
private, trusted network.
@@ -322,6 +367,10 @@ If the page does not open, check these in order:
4. You are opening port `8765`, not the gateway health port.
5. LAN access uses `host: "0.0.0.0"` and a token or token issue secret.
If voice input asks for a secure connection, use HTTPS with a certificate the
device trusts. Browsers do not expose microphone capture to
`http://<your-ip>` origins.
For detailed diagnostics, see
[`troubleshooting.md#webui-problems`](./troubleshooting.md#webui-problems).
For frontend development, see [`../webui/README.md`](../webui/README.md).
+5 -21
View File
@@ -8,7 +8,7 @@ from typing import TYPE_CHECKING, Any, Callable, Coroutine, cast
from loguru import logger
from nanobot.session.manager import Session, SessionManager
from nanobot.session.manager import MIN_COMPACTED_REPLAY_MESSAGES, Session, SessionManager
if TYPE_CHECKING:
from nanobot.agent.memory import Consolidator
@@ -16,7 +16,7 @@ if TYPE_CHECKING:
class AutoCompact:
_RECENT_SUFFIX_MESSAGES = 8
_RECENT_SUFFIX_MESSAGES = MIN_COMPACTED_REPLAY_MESSAGES
_INTERNAL_SESSION_PREFIXES = ("dream:",)
def __init__(self, sessions: SessionManager, consolidator: Consolidator,
@@ -45,25 +45,9 @@ class AutoCompact:
return False
return idle_seconds >= self._ttl * 60
def _has_compactable_idle_tail(self, key: str) -> bool:
def _has_unarchived_messages(self, key: str) -> bool:
session = self.sessions.get_or_create(key)
tail = list(session.messages[session.last_consolidated:])
if not tail:
return False
probe = Session(
key=session.key,
messages=tail,
created_at=session.created_at,
updated_at=session.updated_at,
metadata={},
last_consolidated=0,
)
result = probe.retain_recent_legal_suffix(
self._RECENT_SUFFIX_MESSAGES,
extend_to_user=True,
)
messages_to_remove = result.dropped[result.already_consolidated_count:]
return bool(messages_to_remove)
return session.last_consolidated < len(session.messages)
@staticmethod
def _format_summary(text: str, last_active: datetime) -> str:
@@ -88,7 +72,7 @@ class AutoCompact:
if key in active_session_keys:
continue
updated_at = info.get("updated_at")
if self._is_expired(updated_at, now) and self._has_compactable_idle_tail(key):
if self._is_expired(updated_at, now) and self._has_unarchived_messages(key):
session = self.sessions.get_or_create(key)
try:
runtime = resolve_runtime(session)
-7
View File
@@ -140,10 +140,3 @@ class AutomationTurnCoordinator:
if pending_id:
pending_ids.add(pending_id)
return pending_ids
async def publish_next_deferred(self, session_key: str) -> bool:
return await publish_next_deferred_turn(
deferred_queues=self.deferred_queues,
publish_inbound=self._publish_inbound,
session_key=session_key,
)
+15 -4
View File
@@ -13,7 +13,11 @@ from nanobot.agent.tools import mcp as mcp_tools
from nanobot.agent.tools import sessions as session_tools
from nanobot.agent.tools.registry import ToolRegistry
from nanobot.apps.cli import utils as cli_app_utils
from nanobot.bus.events import InboundMessage
from nanobot.bus.events import (
INBOUND_META_RUNTIME_CONTROL,
RUNTIME_CONTROL_SESSION_DISCARD,
InboundMessage,
)
from nanobot.runtime_context import (
RUNTIME_CONTEXT_END,
RUNTIME_CONTEXT_MESSAGE_META,
@@ -47,6 +51,9 @@ async def close_mcp(state: Any) -> None:
async def handle_runtime_control(state: Any, msg: InboundMessage, tools: ToolRegistry) -> bool:
if msg.metadata.get(INBOUND_META_RUNTIME_CONTROL) == RUNTIME_CONTROL_SESSION_DISCARD:
await state.discard_session(msg.session_key)
return True
for handler in (
image_generation_tools.handle_runtime_control,
mcp_tools.handle_runtime_control,
@@ -79,6 +86,7 @@ class ContextBuilder:
channel: str | None = None,
session_summary: str | None = None,
workspace: Path | None = None,
include_memory: bool = True,
include_memory_recent_history: bool = True,
session_key: str | None = None,
unified_session: bool = False,
@@ -93,9 +101,10 @@ class ContextBuilder:
parts.append(render_template("agent/tool_contract.md"))
memory = self.memory.read_memory()
if memory and not self._is_template_content(memory, "memory/MEMORY.md"):
parts.append(f"# Memory\n\n## Long-term Memory\n{memory}")
if include_memory:
memory = self.memory.read_memory()
if memory and not self._is_template_content(memory, "memory/MEMORY.md"):
parts.append(f"# Memory\n\n## Long-term Memory\n{memory}")
active_skills = self.skills.get_always_skills()
active_skills.extend(
@@ -219,6 +228,7 @@ class ContextBuilder:
session_summary: str | None = None,
runtime_context_blocks: Sequence[RuntimeContextBlock] | None = None,
workspace: Path | None = None,
include_memory: bool = True,
include_memory_recent_history: bool = True,
session_key: str | None = None,
unified_session: bool = False,
@@ -238,6 +248,7 @@ class ContextBuilder:
channel=channel,
session_summary=session_summary,
workspace=root,
include_memory=include_memory,
include_memory_recent_history=include_memory_recent_history,
session_key=session_key,
unified_session=unified_session,
+63 -15
View File
@@ -36,6 +36,7 @@ from nanobot.agent.tools.exec_session import ExecSessionManager
from nanobot.agent.tools.file_state import FileStateStore, bind_file_states, reset_file_states
from nanobot.agent.tools.message import MessageTool
from nanobot.agent.tools.registry import ToolRegistry
from nanobot.agent.tools.runtime_control import AgentRuntimeControl
from nanobot.agent.tools.self import MyTool
from nanobot.agent.turn_delivery import (
TurnDelivery,
@@ -197,6 +198,11 @@ class AgentLoop:
def tool_names(self) -> list[str]:
return self.tools.tool_names
@property
def last_usage(self) -> Mapping[str, int]:
"""Latest aggregate usage exposed through the runtime-control snapshot."""
return self._last_usage
@property
def provider(self) -> LLMProvider:
"""Provider selected for future turn admissions."""
@@ -398,6 +404,7 @@ class AgentLoop:
self._mcp_connecting = False
self._runtime_context_providers: list[RuntimeContextProvider] = []
self._active_tasks: dict[str, set[asyncio.Task[Any]]] = {}
self._discarding_sessions: set[str] = set()
self._background_tasks: set[asyncio.Task[Any]] = set()
self._close_mcp_lock = asyncio.Lock()
self._session_locks: weakref.WeakValueDictionary[str, asyncio.Lock] = (
@@ -447,7 +454,6 @@ class AgentLoop:
if model_preset:
self.set_model_preset(model_preset, publish_update=False)
self._register_default_tools(provider_snapshot_loader=provider_snapshot_loader)
self._runtime_vars: dict[str, Any] = {}
self._current_iteration: int = 0
self.commands = CommandRouter()
register_builtin_commands(self.commands)
@@ -622,10 +628,13 @@ class AgentLoop:
loader = ToolLoader()
registered = loader.load(ctx, self.tools)
# MyTool needs runtime state reference — manual registration
# MyTool receives only the explicit runtime-control capability.
if self.tools_config.my.enable:
self.tools.register(
MyTool(runtime_state=self, modify_allowed=self.tools_config.my.allow_set)
MyTool(
runtime_control=AgentRuntimeControl(self),
modify_allowed=self.tools_config.my.allow_set,
)
)
registered.append("my")
@@ -721,6 +730,7 @@ class AgentLoop:
session_summary=ctx.pending_summary,
workspace=scope.project_path,
runtime_context_blocks=ctx.runtime_context_blocks,
include_memory=ctx.session.policy.persist,
include_memory_recent_history=not ctx.ephemeral,
session_key=ctx.session.key,
unified_session=self._unified_session,
@@ -786,9 +796,9 @@ class AgentLoop:
logger.warning("Command '{}' matched but dispatch returned None", raw)
async def _cancel_active_tasks(self, key: str) -> int:
"""Cancel and await all active tasks and subagents for *key*.
"""Cancel and await all active work for *key*.
Returns the total number of cancelled tasks + subagents.
Returns the total number of cancelled tasks, subagents, and exec sessions.
"""
tasks = tuple(self._active_tasks.pop(key, set()))
cancelled = sum(1 for t in tasks if not t.done() and t.cancel())
@@ -796,7 +806,17 @@ class AgentLoop:
with suppress(asyncio.CancelledError, Exception):
await t
sub_cancelled = await self.subagents.cancel_by_session(key)
return cancelled + sub_cancelled
exec_cancelled = await self._exec_session_manager.terminate_by_owner(key)
return cancelled + sub_cancelled + exec_cancelled
async def discard_session(self, key: str) -> None:
"""Stop active work for *key* and forget its cached session."""
self._discarding_sessions.add(key)
try:
self.sessions.invalidate(key)
await self._cancel_active_tasks(key)
finally:
self._discarding_sessions.discard(key)
def _effective_session_key(self, msg: InboundMessage) -> str:
"""Return the session key used for task routing and mid-turn injections."""
@@ -1161,6 +1181,11 @@ class AgentLoop:
effective_key = self._effective_session_key(msg)
if await agent_context.handle_runtime_control(self, msg, self.tools):
continue
if (
msg.require_existing_session
and self.sessions.get_cached(effective_key) is None
):
continue
if self.commands.is_priority(raw):
await self._dispatch_command_inline(
msg, effective_key, raw,
@@ -1279,6 +1304,8 @@ class AgentLoop:
# _emit_checkpoint during tool execution; materializing
# it into session history now makes it visible in the
# next conversation turn.
if session_key in self._discarding_sessions:
raise
try:
key = self._effective_session_key(msg)
session = self.sessions.get_or_create(key)
@@ -1556,6 +1583,7 @@ class AgentLoop:
had_injections: bool,
streamed_content: bool,
*,
log_content: bool = True,
turn_latency_ms: int | None = None,
) -> OutboundMessage | None:
"""Assemble the final outbound message from turn results."""
@@ -1564,8 +1592,11 @@ class AgentLoop:
if not had_injections or stop_reason == "empty_final_response":
return None
preview = final_content[:120] + "..." if len(final_content) > 120 else final_content
logger.info("Response to {}:{}: {}", msg.channel, msg.sender_id, preview)
if log_content:
preview = final_content[:120] + "..." if len(final_content) > 120 else final_content
logger.info("Response to {}:{}: {}", msg.channel, msg.sender_id, preview)
else:
logger.info("Response to {}:{}: [content hidden]", msg.channel, msg.sender_id)
event = None
meta = dict(msg.metadata or {})
@@ -1594,17 +1625,33 @@ class AgentLoop:
ctx.msg = dataclasses.replace(msg, content=new_content, media=image_paths)
msg = ctx.msg
preview = msg.content[:80] + "..." if len(msg.content) > 80 else msg.content
if ctx.session is None:
if msg.require_existing_session:
ctx.session = self.sessions.get_cached(ctx.session_key)
if ctx.session is None:
raise RuntimeError("required session is not active")
else:
ctx.session = self.sessions.get_or_create(ctx.session_key)
session = ctx.session
ctx.ephemeral = ctx.ephemeral or not session.policy.persist
tools = ctx.tools or self.tools
if session.policy.disabled_tools:
restricted = ToolRegistry()
for name in tools.tool_names:
tool = tools.get(name)
if name not in session.policy.disabled_tools and tool:
restricted.register(tool)
tools = restricted
ctx.tools = tools
if ctx.kind is TurnKind.SYSTEM:
logger.info("Processing system message from {}", msg.sender_id)
else:
elif session.policy.log_content:
preview = msg.content[:80] + "..." if len(msg.content) > 80 else msg.content
logger.info("Processing message from {}:{}: {}", msg.channel, msg.sender_id, preview)
else:
logger.info("Processing message from {}:{}: [content hidden]", msg.channel, msg.sender_id)
# Session is already fetched by the caller (_process_message) but
# 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(
session,
msg,
@@ -1907,6 +1954,7 @@ class AgentLoop:
ctx.stop_reason,
ctx.had_injections,
ctx.streamed_content,
log_content=ctx.require_session().policy.log_content,
turn_latency_ms=ctx.turn_latency_ms,
)
if ctx.ephemeral and ctx.outbound is not None:
+36 -36
View File
@@ -21,7 +21,7 @@ from typing import TYPE_CHECKING, Any, Callable, Iterator, cast
from loguru import logger
from nanobot.runtime_context import public_history_messages
from nanobot.session.manager import Session, SessionManager
from nanobot.session.manager import MIN_COMPACTED_REPLAY_MESSAGES, Session, SessionManager
from nanobot.utils.gitstore import GitStore
from nanobot.utils.helpers import (
content_with_media_breadcrumbs,
@@ -858,14 +858,13 @@ class Consolidator:
return last_boundary
@staticmethod
def _full_unconsolidated_history(
def _full_replay_history(
session: Session,
) -> list[dict[str, Any]]:
"""Return the whole unconsolidated tail for consolidation decisions."""
unconsolidated_count = len(session.messages) - session.last_consolidated
if unconsolidated_count <= 0:
"""Return all messages that can reach the next model prompt."""
if not session.messages:
return []
return session.get_history(max_messages=unconsolidated_count)
return session.get_history(max_messages=len(session.messages))
@staticmethod
def _replay_overflow_boundary(
@@ -948,8 +947,8 @@ class Consolidator:
*,
runtime: LLMRuntime,
) -> tuple[int, str]:
"""Estimate prompt size from the full unconsolidated session tail."""
history = self._full_unconsolidated_history(session)
"""Estimate prompt size from the full replayable session history."""
history = self._full_replay_history(session)
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")
@@ -1160,42 +1159,37 @@ class Consolidator:
session_key: str,
*,
runtime: LLMRuntime,
max_suffix: int = 8,
max_suffix: int = MIN_COMPACTED_REPLAY_MESSAGES,
) -> str | None:
"""Archive an idle prefix and hide it from replay without deleting it."""
"""Archive the full idle tail while keeping recent messages replayable.
``max_suffix`` remains accepted for SDK compatibility. Replay retention
is now derived independently from archive progress using the project-wide
compacted-session window.
"""
if max_suffix != MIN_COMPACTED_REPLAY_MESSAGES:
logger.debug(
"Idle-session compact for {} uses the fixed replay window ({}, requested {})",
session_key,
MIN_COMPACTED_REPLAY_MESSAGES,
max_suffix,
)
lock = self.get_lock(session_key)
async with lock:
self.sessions.invalidate(session_key)
session = self.sessions.get_or_create(session_key)
messages_to_summarize = list(session.messages[session.last_consolidated:])
if not messages_to_summarize:
self.sessions.save(session)
return ""
probe = Session(
key=session.key,
messages=messages_to_summarize.copy(),
created_at=session.created_at,
updated_at=session.updated_at,
metadata={},
last_consolidated=0,
)
result = probe.retain_recent_legal_suffix(max_suffix, extend_to_user=True)
visible_suffix = probe.messages
messages_to_remove = result.dropped
if not messages_to_remove:
self.sessions.save(session)
archive_start = session.last_consolidated
messages_to_archive = list(session.messages[archive_start:])
if not messages_to_archive:
return ""
last_active = session.updated_at
# The visible suffix informs the summary but stays out of raw fallback.
archive_end = archive_start + len(messages_to_archive)
summary = await self.archive(
messages_to_remove,
messages_to_archive,
runtime=runtime,
session_key=session_key,
summary_messages=messages_to_summarize,
)
if summary and summary != "(nothing)":
@@ -1204,16 +1198,22 @@ class Consolidator:
"last_active": last_active.isoformat(),
}
# Preserve history and advance only the replay boundary.
session.last_consolidated = len(session.messages) - len(visible_suffix)
# A turn can append while the provider call is in flight. Advance only
# through the captured batch so new messages remain eligible next time.
session.last_consolidated = archive_end
session.provider_state = None
self.sessions.save(session)
visible = session.get_history(
max_messages=MIN_COMPACTED_REPLAY_MESSAGES,
extend_to_user=True,
)
logger.info(
"Idle-session compact for {}: archived={}, visible={}, retained={}, summary={}",
session_key,
len(messages_to_remove),
len(visible_suffix),
len(messages_to_archive),
len(visible),
len(session.messages),
bool(summary),
)
+5
View File
@@ -5,6 +5,7 @@ import json
import time
import uuid
import warnings
from collections.abc import Mapping
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Callable, TypedDict
@@ -157,6 +158,10 @@ class SubagentManager:
self._task_statuses: dict[str, SubagentStatus] = {}
self._session_tasks: dict[str, set[str]] = {} # session_key -> {task_id, ...}
def runtime_statuses(self) -> Mapping[str, SubagentStatus]:
"""Return the observable task statuses used by runtime-control snapshots."""
return self._task_statuses
def set_provider(self, provider: LLMProvider, model: str) -> None:
"""Update the deprecated runtime source used by legacy ``spawn`` calls."""
warnings.warn(
-16
View File
@@ -785,22 +785,6 @@ def _best_window(old_text: str, content: str) -> tuple[float, int, list[str], li
return best_ratio, best_start, best_window_lines, hints
def _find_match(content: str, old_text: str) -> tuple[str | None, int]:
"""Locate old_text in content with a multi-level fallback chain:
1. Exact substring match
2. Line-trimmed sliding window (handles indentation differences)
3. Smart quote normalization (curly ↔ straight quotes)
Both inputs should use LF line endings (caller normalises CRLF).
Returns (matched_fragment, count) or (None, 0).
"""
matches = _find_matches(content, old_text)
if not matches:
return None, 0
return matches[0].text, len(matches)
@tool_parameters(
tool_parameters_schema(
path=StringSchema("The file path to edit"),
+1 -1
View File
@@ -19,7 +19,7 @@ if TYPE_CHECKING:
_SKIP_MODULES = frozenset({
"base", "schema", "registry", "context", "loader", "config",
"file_state", "sandbox", "mcp", "__init__", "runtime_state",
"file_state", "sandbox", "mcp", "__init__", "runtime_control",
})
+19 -29
View File
@@ -975,11 +975,8 @@ async def connect_mcp_servers(
from mcp.client.streamable_http import streamable_http_client
async def open_single_server(
name: str, cfg: "MCPServerConfig"
) -> tuple[str, AsyncExitStack | None]:
server_stack = AsyncExitStack()
await server_stack.__aenter__()
name: str, cfg: "MCPServerConfig", server_stack: AsyncExitStack
) -> bool:
try:
transport_type = cfg.type
if not transport_type:
@@ -991,8 +988,7 @@ async def connect_mcp_servers(
)
else:
logger.warning("MCP server '{}': no command or url configured, skipping", name)
await server_stack.aclose()
return name, None
return False
if transport_type in {"sse", "streamableHttp"}:
ok, error = validate_url_target(cfg.url)
@@ -1003,8 +999,7 @@ async def connect_mcp_servers(
_redact_url(cfg.url),
error,
)
await server_stack.aclose()
return name, None
return False
if transport_type == "stdio":
command, args, env = _normalize_windows_stdio_command(
@@ -1022,8 +1017,7 @@ async def connect_mcp_servers(
elif transport_type == "sse":
if not await _probe_http_url(cfg.url):
logger.warning("MCP server '{}': {} unreachable, skipping", name, _redact_url(cfg.url))
await server_stack.aclose()
return name, None
return False
def httpx_client_factory(
headers: dict[str, str] | None = None,
@@ -1050,8 +1044,7 @@ async def connect_mcp_servers(
elif transport_type == "streamableHttp":
if not await _probe_http_url(cfg.url):
logger.warning("MCP server '{}': {} unreachable, skipping", name, _redact_url(cfg.url))
await server_stack.aclose()
return name, None
return False
http_client = await server_stack.enter_async_context(
httpx.AsyncClient(
@@ -1067,8 +1060,7 @@ async def connect_mcp_servers(
)
else:
logger.warning("MCP server '{}': unknown transport type '{}'", name, transport_type)
await server_stack.aclose()
return name, None
return False
read = _filter_malformed_mcp_progress_notifications(read, name)
session = await server_stack.enter_async_context(ClientSession(read, write))
@@ -1171,7 +1163,7 @@ async def connect_mcp_servers(
logger.info(
"MCP server '{}': connected, {} capabilities registered", name, registered_count
)
return name, server_stack
return True
except Exception as e:
hint = ""
@@ -1191,9 +1183,7 @@ async def connect_mcp_servers(
"only JSON-RPC to stdout and sends logs/debug output to stderr instead."
)
logger.exception("MCP server '{}': failed to connect: {}", name, hint)
with suppress(Exception):
await server_stack.aclose()
return name, None
return False
async def connect_single_server(
name: str, cfg: "MCPServerConfig"
@@ -1203,30 +1193,30 @@ async def connect_mcp_servers(
close_requested = asyncio.Event()
async def own_connection() -> None:
stack: AsyncExitStack | None = None
try:
_, stack = await open_single_server(name, cfg)
if not ready.done():
ready.set_result(stack is not None)
if stack is not None:
await close_requested.wait()
async with AsyncExitStack() as stack:
connected = await open_single_server(name, cfg, stack)
if not ready.done():
ready.set_result(connected)
if connected:
await close_requested.wait()
except BaseException as exc:
if not ready.done():
ready.set_exception(exc)
raise
finally:
if stack is not None:
await stack.aclose()
owner = asyncio.create_task(own_connection(), name=f"mcp:{name}")
connection = _OwnedMCPConnection(owner, close_requested)
try:
connected = await ready
except BaseException:
except BaseException as exc:
close_requested.set()
owner.cancel()
with suppress(BaseException):
await asyncio.shield(owner)
if isinstance(exc, asyncio.CancelledError) and not task_is_cancelling():
logger.warning("MCP server '{}': connection cancelled by server/SDK", name)
return name, None
raise
if not connected:
await connection.aclose()
+1 -9
View File
@@ -3,15 +3,7 @@
from pathlib import Path
from nanobot.config.paths import get_media_dir
from nanobot.security.workspace_policy import (
is_path_within,
resolve_allowed_path,
)
def is_under(path: Path, directory: Path) -> bool:
"""Return True when path resolves under directory."""
return is_path_within(path, directory)
from nanobot.security.workspace_policy import resolve_allowed_path
def resolve_workspace_path(
+319
View File
@@ -0,0 +1,319 @@
"""Explicit runtime state boundary used by :class:`MyTool`."""
from __future__ import annotations
from collections.abc import Mapping
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING, Protocol, TypeAlias, runtime_checkable
if TYPE_CHECKING:
from nanobot.agent.subagent import SubagentManager, SubagentStatus
from nanobot.agent.tools.shell import ExecToolConfig
from nanobot.agent.tools.web import WebToolsConfig
from nanobot.config.schema import ModelPresetConfig
from nanobot.utils.llm_runtime import LLMRuntime
JsonScalar: TypeAlias = str | int | float | bool | None
JsonValue: TypeAlias = JsonScalar | list["JsonValue"] | dict[str, "JsonValue"]
RUNTIME_SNAPSHOT_KEYS = frozenset({
"model",
"model_preset",
"model_presets",
"max_iterations",
"context_window_tokens",
"workspace",
"provider_retry_mode",
"max_tool_result_chars",
"current_iteration",
"_current_iteration",
"tool_names",
"web_config",
"exec_config",
"subagents",
"_last_usage",
})
RUNTIME_COMMAND_KEYS = frozenset({
"model",
"model_preset",
"max_iterations",
"context_window_tokens",
"provider_retry_mode",
"max_tool_result_chars",
"workspace",
})
@dataclass(frozen=True, slots=True)
class RuntimeSnapshot:
"""Detached, allowlisted values available to self-inspection."""
model: str
model_preset: str | None
model_presets: dict[str, dict[str, object]]
max_iterations: int
context_window_tokens: int
workspace: Path | str
provider_retry_mode: str
max_tool_result_chars: int
current_iteration: int
tool_names: list[str]
web_config: dict[str, object]
exec_config: dict[str, object]
subagent_statuses: dict[str, dict[str, object]]
last_usage: dict[str, int]
scratchpad: dict[str, JsonValue]
def as_mapping(self) -> Mapping[str, object]:
"""Return the fixed public names understood by ``MyTool``."""
values: dict[str, object] = {
"model": self.model,
"model_preset": self.model_preset,
"model_presets": self.model_presets,
"max_iterations": self.max_iterations,
"context_window_tokens": self.context_window_tokens,
"workspace": self.workspace,
"provider_retry_mode": self.provider_retry_mode,
"max_tool_result_chars": self.max_tool_result_chars,
"current_iteration": self.current_iteration,
"_current_iteration": self.current_iteration,
"tool_names": self.tool_names,
"web_config": self.web_config,
"exec_config": self.exec_config,
"subagents": {"_task_statuses": self.subagent_statuses},
"_last_usage": self.last_usage,
}
assert values.keys() == RUNTIME_SNAPSHOT_KEYS
return values
@runtime_checkable
class RuntimeControl(Protocol):
"""The complete runtime capability exposed to ``MyTool``."""
def snapshot(self) -> RuntimeSnapshot: ...
def set_model(self, model: str) -> LLMRuntime: ...
def set_model_preset(
self,
name: str,
*,
session_key: str | None,
) -> LLMRuntime: ...
def set_max_iterations(self, value: int) -> None: ...
def set_context_window_tokens(self, value: int) -> LLMRuntime: ...
def set_provider_retry_mode(self, value: str) -> None: ...
def set_max_tool_result_chars(self, value: int) -> None: ...
def set_workspace_display(self, value: str) -> None: ...
def set_scratchpad(self, key: str, value: JsonValue, *, max_keys: int) -> None: ...
class _RuntimeControlTarget(Protocol):
"""Narrow structural dependency required by ``AgentRuntimeControl``."""
max_iterations: int
provider_retry_mode: str
max_tool_result_chars: int
web_config: WebToolsConfig
exec_config: ExecToolConfig
subagents: SubagentManager
@property
def model(self) -> str: ...
@property
def model_preset(self) -> str | None: ...
@property
def model_presets(self) -> Mapping[str, ModelPresetConfig]: ...
@property
def context_window_tokens(self) -> int: ...
@property
def workspace(self) -> Path: ...
@property
def current_iteration(self) -> int: ...
@property
def tool_names(self) -> list[str]: ...
@property
def last_usage(self) -> Mapping[str, int]: ...
def set_runtime_model(self, model: str) -> LLMRuntime: ...
def set_runtime_context_window(self, context_window_tokens: int) -> LLMRuntime: ...
def set_model_preset(self, name: str | None) -> LLMRuntime: ...
def set_session_model_preset(self, session_key: str, name: str) -> LLMRuntime: ...
class AgentRuntimeControl:
"""Allowlisted adapter from agent-loop state to ``RuntimeControl``."""
def __init__(self, target: _RuntimeControlTarget) -> None:
self.__target = target
self.__scratchpad: dict[str, JsonValue] = {}
self.__workspace_display: str | None = None
def snapshot(self) -> RuntimeSnapshot:
target = self.__target
return RuntimeSnapshot(
model=target.model,
model_preset=target.model_preset,
model_presets=_snapshot_model_presets(target.model_presets),
max_iterations=target.max_iterations,
context_window_tokens=target.context_window_tokens,
workspace=(
self.__workspace_display
if self.__workspace_display is not None
else target.workspace
),
provider_retry_mode=target.provider_retry_mode,
max_tool_result_chars=target.max_tool_result_chars,
current_iteration=target.current_iteration,
tool_names=list(target.tool_names),
web_config=_snapshot_web_config(target.web_config),
exec_config=_snapshot_exec_config(target.exec_config),
subagent_statuses=_snapshot_subagent_statuses(target.subagents),
last_usage=dict(target.last_usage),
scratchpad=_snapshot_json_mapping(self.__scratchpad),
)
def set_model(self, model: str) -> LLMRuntime:
return self.__target.set_runtime_model(model)
def set_model_preset(
self,
name: str,
*,
session_key: str | None,
) -> LLMRuntime:
if session_key is not None:
return self.__target.set_session_model_preset(session_key, name)
return self.__target.set_model_preset(name)
def set_max_iterations(self, value: int) -> None:
self.__target.max_iterations = value
self.__target.subagents.max_iterations = value
def set_context_window_tokens(self, value: int) -> LLMRuntime:
return self.__target.set_runtime_context_window(value)
def set_provider_retry_mode(self, value: str) -> None:
self.__target.provider_retry_mode = value
def set_max_tool_result_chars(self, value: int) -> None:
self.__target.max_tool_result_chars = value
def set_workspace_display(self, value: str) -> None:
"""Preserve MyTool display compatibility without changing path enforcement."""
self.__workspace_display = value
def set_scratchpad(self, key: str, value: JsonValue, *, max_keys: int) -> None:
if key not in self.__scratchpad and len(self.__scratchpad) >= max_keys:
raise ValueError(f"scratchpad is full (max {max_keys} keys)")
self.__scratchpad[key] = value
def _snapshot_model_presets(
presets: Mapping[str, ModelPresetConfig],
) -> dict[str, dict[str, object]]:
return {
name: {
"label": preset.label,
"model": preset.model,
"provider": preset.provider,
"max_tokens": preset.max_tokens,
"context_window_tokens": preset.context_window_tokens,
"temperature": preset.temperature,
"reasoning_effort": preset.reasoning_effort,
}
for name, preset in presets.items()
}
def _snapshot_web_config(config: WebToolsConfig) -> dict[str, object]:
return {
"enable": config.enable,
# Proxy URLs may embed credentials. Presence is enough for diagnosis.
"proxy": "<configured>" if config.proxy else config.proxy,
"user_agent": config.user_agent,
"search": {
"provider": config.search.provider,
"base_url": config.search.base_url,
"max_results": config.search.max_results,
"timeout": config.search.timeout,
},
"fetch": {
"use_jina_reader": config.fetch.use_jina_reader,
},
}
def _snapshot_exec_config(config: ExecToolConfig) -> dict[str, object]:
return {
"enable": config.enable,
"timeout": config.timeout,
"path_prepend": config.path_prepend,
"path_append": config.path_append,
"sandbox": config.sandbox,
"sandbox_ro_binds": list(config.sandbox_ro_binds),
"sandbox_rw_binds": list(config.sandbox_rw_binds),
"allowed_env_keys": list(config.allowed_env_keys),
"allow_patterns": list(config.allow_patterns),
"deny_patterns": list(config.deny_patterns),
}
def _snapshot_subagent_statuses(
manager: SubagentManager,
) -> dict[str, dict[str, object]]:
return {
task_id: _snapshot_subagent_status(status)
for task_id, status in manager.runtime_statuses().items()
}
def _snapshot_subagent_status(status: SubagentStatus) -> dict[str, object]:
return {
"task_id": status.task_id,
"label": status.label,
"task_description": status.task_description,
"started_at": status.started_at,
"phase": status.phase,
"iteration": status.iteration,
"tool_events": [dict(event) for event in status.tool_events],
"usage": dict(status.usage),
"stop_reason": status.stop_reason,
"error": status.error,
}
def _snapshot_json_mapping(values: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
return {key: _snapshot_json_value(value) for key, value in values.items()}
def _snapshot_json_value(value: JsonValue) -> JsonValue:
if isinstance(value, list):
return [_snapshot_json_value(item) for item in value]
if isinstance(value, dict):
return {
key: _snapshot_json_value(item)
for key, item in value.items()
}
return value
-76
View File
@@ -1,76 +0,0 @@
"""RuntimeState protocol: agent loop state exposed to MyTool."""
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):
"""Minimum contract that MyTool requires from its runtime state provider.
In practice, this is always satisfied by ``AgentLoop``. MyTool also
accesses arbitrary attributes dynamically (via ``getattr`` / ``setattr``)
for dot-path inspection and modification; those paths are validated at
runtime rather than by this protocol.
"""
@property
def model(self) -> str: ...
@property
def max_iterations(self) -> int: ...
@property
def current_iteration(self) -> int: ...
@property
def tool_names(self) -> list[str]: ...
@property
def workspace(self) -> Path: ...
@property
def provider_retry_mode(self) -> str: ...
@property
def max_tool_result_chars(self) -> int: ...
@property
def context_window_tokens(self) -> int: ...
@property
def web_config(self) -> WebToolsConfig: ...
@property
def exec_config(self) -> ExecToolConfig: ...
@property
def subagents(self) -> SubagentManager: ...
@property
def _runtime_vars(self) -> dict[str, Any]: ...
@property
def _last_usage(self) -> dict[str, int]: ...
def _sync_subagent_runtime_limits(self) -> None: ...
def set_runtime_model(self, model: str) -> LLMRuntime: ...
def set_runtime_context_window(self, context_window_tokens: int) -> LLMRuntime: ...
def set_session_model_preset(
self,
session_key: str,
name: str,
) -> LLMRuntime: ...
@property
def model_preset(self) -> str | None: ...
+213 -182
View File
@@ -1,8 +1,7 @@
"""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
# Tool.execute accepts heterogeneous schemas.
# pyright: reportIncompatibleMethodOverride=false
from __future__ import annotations
@@ -14,7 +13,13 @@ from loguru import logger
from nanobot.agent.tools.base import Tool, ToolResult
from nanobot.agent.tools.context import current_request_context, current_request_session_key
from nanobot.agent.tools.runtime_state import RuntimeState
from nanobot.agent.tools.runtime_control import (
RUNTIME_COMMAND_KEYS,
RUNTIME_SNAPSHOT_KEYS,
JsonValue,
RuntimeControl,
RuntimeSnapshot,
)
from nanobot.config_base import Base
if TYPE_CHECKING:
@@ -28,25 +33,28 @@ class MyToolConfig(Base):
allow_set: bool = False
def _has_real_attr(obj: Any, key: str) -> bool:
"""Check if obj has a real (explicitly set) attribute, not auto-generated by mock."""
if isinstance(obj, dict):
return key in obj
d = getattr(obj, "__dict__", None)
if d is not None and key in d:
return True
for cls in type(obj).__mro__:
if key in cls.__dict__:
return True
return False
def _is_subagent_status(value: object) -> TypeGuard[SubagentStatus]:
from nanobot.agent.subagent import SubagentStatus
return isinstance(value, SubagentStatus)
def _is_subagent_status_snapshot(value: object) -> TypeGuard[Mapping[str, object]]:
if not isinstance(value, Mapping):
return False
return all(
field in value
for field in ("task_id", "label", "task_description", "started_at", "phase")
)
def _is_string_mapping(value: object) -> TypeGuard[Mapping[str, object]]:
if not isinstance(value, Mapping):
return False
mapping = cast(Mapping[object, object], value)
return all(isinstance(key, str) for key in mapping)
class MyTool(Tool):
"""Check and set the agent loop's runtime configuration."""
@@ -79,7 +87,10 @@ class MyTool(Tool):
READ_ONLY = frozenset({
"subagents", # observable but replacing it would break the system
"tool_names",
"current_iteration",
"_current_iteration", # updated by runner only
"_last_usage",
"exec_config", # inspect allowed (e.g. check sandbox), modify blocked
"web_config", # inspect allowed (e.g. check enable), modify blocked
"model_presets", # config-derived catalog; changes require config reload
@@ -103,13 +114,6 @@ class MyTool(Tool):
"private_key", "access_token", "refresh_token", "auth",
})
@classmethod
def _is_sensitive_field_name(cls, name: str) -> bool:
lowered = name.lower()
return lowered in cls._SENSITIVE_NAMES or any(
part in cls._SENSITIVE_NAMES for part in lowered.split("_")
)
RESTRICTED: dict[str, dict[str, Any]] = {
"max_iterations": {"type": int, "min": 1, "max": 100},
"context_window_tokens": {"type": int, "min": 4096, "max": 1_000_000},
@@ -123,15 +127,15 @@ class MyTool(Tool):
"context_window_tokens",
})
def __init__(self, runtime_state: RuntimeState, modify_allowed: bool = True) -> None:
self._runtime_state = runtime_state
def __init__(self, runtime_control: RuntimeControl, modify_allowed: bool = True) -> None:
self._runtime_control = runtime_control
self._modify_allowed = modify_allowed
def __deepcopy__(self, memo: dict[int, Any]) -> MyTool:
cls = self.__class__
result = cls.__new__(cls)
memo[id(self)] = result
result._runtime_state = self._runtime_state
result._runtime_control = self._runtime_control
result._modify_allowed = self._modify_allowed
return result
@@ -208,9 +212,12 @@ class MyTool(Tool):
# Path resolution
# ------------------------------------------------------------------
def _resolve_path(self, path: str) -> tuple[Any, str | None]:
def _resolve_path(
self,
snapshot: RuntimeSnapshot,
path: str,
) -> tuple[object | None, str | None]:
parts = path.split(".")
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"
@@ -218,17 +225,13 @@ class MyTool(Tool):
return None, f"'{part}' is not accessible"
if part.lower() in self._SENSITIVE_NAMES:
return None, f"'{part}' is not accessible"
try:
if isinstance(obj, Mapping):
mapping = cast(Mapping[str, Any], obj)
if part in mapping:
obj = mapping[part]
else:
return None, f"'{part}' not found in mapping"
else:
obj = getattr(obj, part)
except (KeyError, AttributeError) as e:
return None, f"'{part}' not found: {e}"
obj: object = snapshot.as_mapping()
for part in parts:
if not _is_string_mapping(obj):
return None, f"'{part}' not found"
if part not in obj:
return None, f"'{part}' not found in mapping"
obj = obj[part]
return obj, None
@staticmethod
@@ -242,20 +245,48 @@ class MyTool(Tool):
# ------------------------------------------------------------------
@staticmethod
def _format_status(st: "SubagentStatus", indent: str = " ") -> str:
elapsed = time.monotonic() - st.started_at
tool_summary = ", ".join(
f"{e.get('name', '?')}({e.get('status', '?')})" for e in st.tool_events[-5:]
) or "none"
def _format_status(
st: "SubagentStatus | Mapping[str, object]",
indent: str = " ",
) -> str:
if isinstance(st, Mapping):
started_at = st.get("started_at", time.monotonic())
raw_events = st.get("tool_events", [])
phase = st.get("phase", "unknown")
iteration = st.get("iteration", 0)
usage = st.get("usage", {})
error = st.get("error")
stop_reason = st.get("stop_reason")
else:
started_at = st.started_at
raw_events = st.tool_events
phase = st.phase
iteration = st.iteration
usage = st.usage
error = st.error
stop_reason = st.stop_reason
elapsed = time.monotonic() - (
float(started_at) if isinstance(started_at, (int, float)) else time.monotonic()
)
tool_events = cast(list[object], raw_events) if isinstance(raw_events, list) else []
tool_summaries: list[str] = []
for raw_event in tool_events[-5:]:
if not isinstance(raw_event, Mapping):
continue
event = cast(Mapping[str, object], raw_event)
tool_summaries.append(
f"{event.get('name', '?')}({event.get('status', '?')})"
)
tool_summary = ", ".join(tool_summaries) or "none"
lines = [
f"{indent}phase: {st.phase}, iteration: {st.iteration}, elapsed: {elapsed:.1f}s",
f"{indent}phase: {phase}, iteration: {iteration}, elapsed: {elapsed:.1f}s",
f"{indent}tools: {tool_summary}",
f"{indent}usage: {st.usage or 'n/a'}",
f"{indent}usage: {usage or 'n/a'}",
]
if st.error:
lines.append(f"{indent}error: {st.error}")
if st.stop_reason:
lines.append(f"{indent}stop_reason: {st.stop_reason}")
if error:
lines.append(f"{indent}error: {error}")
if stop_reason:
lines.append(f"{indent}stop_reason: {stop_reason}")
return "\n".join(lines)
@staticmethod
@@ -264,29 +295,38 @@ class MyTool(Tool):
header = f"Subagent [{val.task_id}] '{val.label}'"
detail = MyTool._format_status(val, " ")
return f"{header}\n task: {val.task_description}\n{detail}"
# SubagentManager: delegate to its _task_statuses dict
task_statuses = getattr(val, "_task_statuses", None)
if isinstance(task_statuses, dict):
return MyTool._format_value(task_statuses, key)
if _is_subagent_status_snapshot(val):
header = f"Subagent [{val['task_id']}] '{val['label']}'"
detail = MyTool._format_status(val, " ")
return f"{header}\n task: {val['task_description']}\n{detail}"
if isinstance(val, Mapping):
mapping = cast(Mapping[object, object], val)
else:
mapping = None
if mapping and set(mapping) == {"_task_statuses"}:
task_statuses = mapping["_task_statuses"]
if isinstance(task_statuses, Mapping):
return MyTool._format_value(task_statuses, key)
if (
mapping
and _is_subagent_status(next(iter(mapping.values())))
and (
_is_subagent_status(next(iter(mapping.values())))
or _is_subagent_status_snapshot(next(iter(mapping.values())))
)
):
status_mapping: Mapping[object, SubagentStatus] = cast(Any, mapping)
prefix = f"{key}: " if key else ""
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}")
lines = [f"{prefix}{len(mapping)} subagent(s):"]
for tid, st in mapping.items():
if _is_subagent_status(st):
detail = MyTool._format_status(st, " ")
label = st.label
elif _is_subagent_status_snapshot(st):
detail = MyTool._format_status(st, " ")
label = st.get("label", "?")
else:
continue
lines.append(f" [{tid}] '{label}'\n{detail}")
return "\n".join(lines)
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)
@@ -311,32 +351,6 @@ class MyTool(Tool):
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
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: list[str] = []
for f in fields:
fv = getattr(val, f, "?")
if MyTool._is_sensitive_field_name(f):
continue
if isinstance(fv, (str, int, float, bool, type(None))):
pairs.append(f"{f}={fv!r}")
else:
pairs.append(f"{f}=<{type(fv).__name__}>")
preview = ", ".join(pairs)
return f"{key}: {preview}" if key else preview
else:
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 ""
return f"{key}: <{cls_name}> [{preview}{suffix}]" if key else f"<{cls_name}> [{preview}{suffix}]"
r = repr(val)
return f"{key}: {r}" if key else r
@@ -366,7 +380,12 @@ class MyTool(Tool):
runtime = request_ctx.runtime if request_ctx is not None else None
if runtime is None or key not in self._MODEL_RUNTIME_FIELDS:
return False, None
return True, getattr(runtime, key)
values: dict[str, object] = {
"model": runtime.model,
"model_preset": runtime.model_preset,
"context_window_tokens": runtime.context_window_tokens,
}
return True, values[key]
def _inspect(self, key: str | None) -> str:
if not key:
@@ -375,62 +394,64 @@ class MyTool(Tool):
request_ctx = current_request_context()
if request_ctx is None:
return ToolResult.error("Error: current request context is unavailable")
request_values: dict[str, str | None] = {
"channel": request_ctx.channel,
"chat_id": request_ctx.chat_id,
"sender_id": request_ctx.sender_id,
}
if key == "request":
return self._format_value(
{field: getattr(request_ctx, field) for field in self._REQUEST_FIELDS},
key,
)
return self._format_value(request_values, key)
field = key.removeprefix("request.")
if field not in self._REQUEST_FIELDS:
return ToolResult.error(f"Error: '{key}' not found")
return self._format_value(getattr(request_ctx, field), key)
return self._format_value(request_values[field], key)
if "." not in key:
found, value = self._current_runtime_value(key)
if found:
return self._format_value(value, key)
snapshot = self._runtime_control.snapshot()
top = key.split(".")[0]
if top in self._DENIED_ATTRS or top.startswith("__"):
return ToolResult.error(f"Error: '{top}' is not accessible")
obj, err = self._resolve_path(key)
obj, err = self._resolve_path(snapshot, key)
if err:
# "scratchpad" alias for _runtime_vars
if key == "scratchpad":
rv = self._runtime_state._runtime_vars
return self._format_value(rv, "scratchpad") if rv else "scratchpad is empty"
# Fallback: check _runtime_vars for simple keys stored by modify
if "." not in key and key in self._runtime_state._runtime_vars:
return self._format_value(self._runtime_state._runtime_vars[key], key)
return (
self._format_value(snapshot.scratchpad, "scratchpad")
if snapshot.scratchpad
else "scratchpad is empty"
)
if "." not in key and key in snapshot.scratchpad:
return self._format_value(snapshot.scratchpad[key], key)
return ToolResult.error(f"Error: {err}")
# Guard against mock auto-generated attributes
if "." not in key and not _has_real_attr(self._runtime_state, key):
if key in self._runtime_state._runtime_vars:
return self._format_value(self._runtime_state._runtime_vars[key], key)
return ToolResult.error(f"Error: '{key}' not found")
return self._format_value(obj, key)
def _inspect_all(self) -> str:
state = self._runtime_state
snapshot = self._runtime_control.snapshot()
values = snapshot.as_mapping()
parts: list[str] = []
# RESTRICTED keys
for k in self.RESTRICTED:
found, value = self._current_runtime_value(k)
parts.append(self._format_value(value if found else getattr(state, k, None), k))
parts.append(self._format_value(value if found else values[k], k))
found, value = self._current_runtime_value("model_preset")
parts.append(self._format_value(
value if found else state.model_preset,
value if found else snapshot.model_preset,
"model_preset",
))
# Other useful top-level keys shown in description
for k in ("workspace", "provider_retry_mode", "max_tool_result_chars", "_current_iteration", "web_config", "exec_config", "workspace_sandbox", "subagents"):
if _has_real_attr(state, k):
parts.append(self._format_value(getattr(state, k, None), k))
# Token usage
usage = state._last_usage
if usage:
parts.append(self._format_value(usage, "_last_usage"))
rv = state._runtime_vars
if rv:
parts.append(self._format_value(rv, "scratchpad"))
for k in (
"workspace",
"provider_retry_mode",
"max_tool_result_chars",
"_current_iteration",
"web_config",
"exec_config",
"subagents",
):
parts.append(self._format_value(values[k], k))
if snapshot.last_usage:
parts.append(self._format_value(snapshot.last_usage, "_last_usage"))
if snapshot.scratchpad:
parts.append(self._format_value(snapshot.scratchpad, "scratchpad"))
return "\n".join(parts)
# -- modify --
@@ -454,48 +475,49 @@ class MyTool(Tool):
if leaf.lower() in self._SENSITIVE_NAMES:
self._audit("modify", f"BLOCKED sensitive leaf '{leaf}'")
return ToolResult.error(f"Error: '{leaf}' is not accessible")
parent, err = self._resolve_path(parent_path)
snapshot = self._runtime_control.snapshot()
_parent, err = self._resolve_path(snapshot, parent_path)
if err:
return ToolResult.error(f"Error: {err}")
if isinstance(parent, dict):
parent[leaf] = value
else:
setattr(parent, leaf, value)
self._audit("modify", f"{key} = {value!r}")
return f"Set {key} = {value!r}"
self._audit("modify", f"READ_ONLY {key}")
return ToolResult.error(f"Error: '{key}' is read-only and cannot be modified")
if key == "model_preset":
return self._modify_model_preset(value)
if key in self.RESTRICTED:
return self._modify_restricted(key, value)
return self._modify_free(key, value)
if key in RUNTIME_COMMAND_KEYS:
return self._modify_runtime_setting(key, value)
if key in RUNTIME_SNAPSHOT_KEYS:
self._audit("modify", f"READ_ONLY {key}")
return ToolResult.error(f"Error: '{key}' is read-only and cannot be modified")
return self._modify_scratchpad(key, value)
def _modify_model_preset(self, value: Any) -> str:
if not isinstance(value, str) or not value.strip():
return ToolResult.error("Error: 'model_preset' must be a non-empty string")
name = value.strip()
session_key = current_request_session_key()
old = self._runtime_control.snapshot().model_preset
try:
runtime = self._runtime_control.set_model_preset(
name,
session_key=session_key,
)
except (KeyError, ValueError) as exc:
message = str(exc.args[0]) if exc.args else str(exc)
punctuation = "" if message.endswith((".", "!", "?")) else "."
return ToolResult.error(f"Error: {message}{punctuation}")
if session_key:
try:
runtime = self._runtime_state.set_session_model_preset(
session_key,
name,
)
except (KeyError, ValueError) as exc:
message = str(exc.args[0]) if exc.args else str(exc)
punctuation = "" if message.endswith((".", "!", "?")) else "."
return ToolResult.error(f"Error: {message}{punctuation}")
self._audit("modify", f"model_preset = {name!r}")
return (
f"Set model_preset = {name!r} for the next turn; "
f"model will be {runtime.model!r}; "
f"context_window_tokens will be {runtime.context_window_tokens!r}"
)
result = self._modify_free("model_preset", name)
if isinstance(result, ToolResult) and result.is_error:
return result if result.endswith((".", "!", "?")) else ToolResult.error(f"{result}.")
self._audit("modify", f"model_preset: {old!r} -> {name!r}")
return (
f"{result}; model is now {self._runtime_state.model!r}; "
f"context_window_tokens is now {self._runtime_state.context_window_tokens!r}"
f"Set model_preset = {name!r} (was {old!r}); model is now {runtime.model!r}; "
f"context_window_tokens is now {runtime.context_window_tokens!r}"
)
def _modify_restricted(self, key: str, value: Any) -> str:
@@ -508,7 +530,7 @@ class MyTool(Tool):
value = expected(value)
except (ValueError, TypeError):
return ToolResult.error(f"Error: '{key}' must be {expected.__name__}, got {type(value).__name__}")
old = getattr(self._runtime_state, key)
old = self._runtime_control.snapshot().as_mapping()[key]
if "min" in spec and value < spec["min"]:
return ToolResult.error(f"Error: '{key}' must be >= {spec['min']}")
if "max" in spec and value > spec["max"]:
@@ -521,41 +543,46 @@ class MyTool(Tool):
"during an active session; use a configured model_preset"
)
if key == "model":
self._runtime_state.set_runtime_model(cast(str, value))
self._runtime_control.set_model(cast(str, value))
elif key == "context_window_tokens":
self._runtime_state.set_runtime_context_window(cast(int, value))
self._runtime_control.set_context_window_tokens(cast(int, value))
else:
setattr(self._runtime_state, key, value)
if key == "max_iterations" and hasattr(
self._runtime_state,
"_sync_subagent_runtime_limits",
):
self._runtime_state._sync_subagent_runtime_limits()
self._runtime_control.set_max_iterations(cast(int, value))
self._audit("modify", f"{key}: {old!r} -> {value!r}")
return f"Set {key} = {value!r} (was {old!r})"
def _modify_free(self, key: str, value: Any) -> str:
if _has_real_attr(self._runtime_state, key):
old = getattr(self._runtime_state, key)
if isinstance(old, (str, int, float, bool)):
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:
self._audit(
"modify",
f"REJECTED type mismatch {key}: expects {old_t.__name__}, got {new_t.__name__}",
)
return ToolResult.error(f"Error: '{key}' expects {old_t.__name__}, got {new_t.__name__}")
try:
setattr(self._runtime_state, key, value)
except (ValueError, KeyError) as e:
message = str(e.args[0] if isinstance(e, KeyError) and e.args else e).strip('"')
self._audit("modify", f"REJECTED {key}: {message}")
return ToolResult.error(f"Error: {message}")
self._audit("modify", f"{key}: {old!r} -> {value!r}")
return f"Set {key} = {value!r} (was {old!r})"
def _modify_runtime_setting(self, key: str, value: Any) -> str:
old = self._runtime_control.snapshot().as_mapping()[key]
if key == "workspace":
if not isinstance(value, str):
return ToolResult.error(
f"Error: 'workspace' expects str, got {type(value).__name__}"
)
self._runtime_control.set_workspace_display(value)
self._audit("modify", f"workspace: {old!r} -> {value!r}")
return f"Set workspace = {value!r} (was {old!r})"
old_t = type(old)
new_t = cast(type[Any], type(value))
if old_t is float and new_t is int:
pass
elif old_t is not new_t:
self._audit(
"modify",
f"REJECTED type mismatch {key}: expects {old_t.__name__}, got {new_t.__name__}",
)
return ToolResult.error(
f"Error: '{key}' expects {old_t.__name__}, got {new_t.__name__}"
)
if key == "provider_retry_mode":
self._runtime_control.set_provider_retry_mode(cast(str, value))
elif key == "max_tool_result_chars":
self._runtime_control.set_max_tool_result_chars(cast(int, value))
else:
raise AssertionError(f"Unhandled runtime command: {key}")
self._audit("modify", f"{key}: {old!r} -> {value!r}")
return f"Set {key} = {value!r} (was {old!r})"
def _modify_scratchpad(self, key: str, value: Any) -> str:
if callable(value):
self._audit("modify", f"REJECTED callable {key}")
return ToolResult.error("Error: cannot store callable values")
@@ -563,12 +590,16 @@ class MyTool(Tool):
if err:
self._audit("modify", f"REJECTED {key}: {err}")
return ToolResult.error(f"Error: {err}")
if key not in self._runtime_state._runtime_vars and len(self._runtime_state._runtime_vars) >= self._MAX_RUNTIME_KEYS:
try:
self._runtime_control.set_scratchpad(
key,
cast(JsonValue, value),
max_keys=self._MAX_RUNTIME_KEYS,
)
except ValueError as exc:
self._audit("modify", f"REJECTED {key}: max keys ({self._MAX_RUNTIME_KEYS}) reached")
return ToolResult.error(f"Error: scratchpad is full (max {self._MAX_RUNTIME_KEYS} keys). Remove unused keys first.")
old = self._runtime_state._runtime_vars.get(key)
self._runtime_state._runtime_vars[key] = value
self._audit("modify", f"scratchpad.{key}: {old!r} -> {value!r}")
return ToolResult.error(f"Error: {exc}. Remove unused keys first.")
self._audit("modify", f"scratchpad.{key} = {value!r}")
return f"Set scratchpad.{key} = {value!r}"
@classmethod
+6 -3
View File
@@ -453,12 +453,15 @@ class WebSearchTool(Tool):
async def _search_olostep(self, query: str, n: int) -> str:
try:
from olostep import ( # pyright: ignore[reportMissingImports]
from olostep import ( # pyright: ignore[reportMissingImports, reportMissingTypeStubs]
AsyncOlostep, # pyright: ignore[reportUnknownVariableType]
Olostep_BaseError, # pyright: ignore[reportUnknownVariableType]
Olostep_BaseError, # pyright: ignore[reportAttributeAccessIssue, reportUnknownVariableType]
)
except ImportError:
return ToolResult.error("Error: olostep package not installed. Run: pip install olostep")
return ToolResult.error(
"Error: Olostep support is not installed. "
"Run `nanobot plugins enable 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", "")
-9
View File
@@ -12,15 +12,6 @@ def session_extra(metadata: Mapping[str, Any] | None) -> dict[str, Any]:
return {"cli_apps": cli_apps} if isinstance(cli_apps, list) and cli_apps else {}
def runtime_lines(message: Any, workspace: Path, *, skip: bool = False) -> list[str]:
"""Return model-visible CLI app annotations for the current turn."""
if skip:
return []
text = message.content if isinstance(getattr(message, "content", None), str) else ""
metadata = message.metadata if isinstance(getattr(message, "metadata", None), Mapping) else None
return runtime_lines_for_request(text, metadata, workspace)
def runtime_lines_for_request(
text: str,
metadata: Mapping[str, Any] | None,
+2
View File
@@ -18,6 +18,7 @@ INBOUND_META_RUNTIME_CONTROL = "_runtime_control"
RUNTIME_CONTROL_ACK = "_ack"
RUNTIME_CONTROL_MCP_RELOAD = "mcp_reload"
RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD = "image_generation_reload"
RUNTIME_CONTROL_SESSION_DISCARD = "session_discard"
@dataclass
@@ -32,6 +33,7 @@ class InboundMessage:
media: list[str] = field(default_factory=list) # Media URLs
metadata: dict[str, Any] = field(default_factory=dict) # Channel-specific data
session_key_override: str | None = None # Optional override for thread-scoped sessions
require_existing_session: bool = False
@property
def session_key(self) -> str:
+27
View File
@@ -101,6 +101,31 @@ class BaseChannel(ABC):
"""
pass
def progress_transport_defaults(self) -> tuple[bool, bool] | None:
"""Return channel-owned defaults for progress and tool-hint messages.
``None`` keeps the global channel policy. Channels should override this
only when their transport requires different defaults.
"""
return None
def should_retry_send_error(self, error: Exception) -> bool:
"""Return whether the channel manager may retry a failed delivery.
Channels with protocol-level business errors can override this hook to
prevent retries that cannot succeed until external state changes.
Transport and unexpected errors remain retryable by default.
"""
return True
def start_error_message(self, error: Exception) -> str | None:
"""Return an actionable public message for a channel startup failure.
Channel-specific exception handling stays in the owning channel. Returning
``None`` keeps the manager's generic fallback.
"""
return None
async def send_delta(
self,
chat_id: str,
@@ -237,6 +262,7 @@ class BaseChannel(ABC):
session_key: str | None = None,
is_dm: bool = False,
authorization_id: str | None = None,
require_existing_session: bool = False,
) -> None:
"""Handle a message after checking its authorization subject.
@@ -289,6 +315,7 @@ class BaseChannel(ABC):
media=media or [],
metadata=meta,
session_key_override=session_key,
require_existing_session=require_existing_session,
)
await self.bus.publish_inbound(msg)
-9
View File
@@ -470,15 +470,6 @@ def _extract_post_content(content_json: dict[str, Any]) -> tuple[str, list[str]]
return "", []
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.
"""
text, _ = _extract_post_content(content_json)
return text
# =============================================================================
# QR scan-to-create onboarding
#
@@ -238,20 +238,6 @@ class TestStreamEndReactionCleanup:
ch._remove_reaction.assert_not_called()
@pytest.mark.asyncio
async def test_no_removal_when_both_ids_missing(self):
ch = _make_channel()
ch._stream_bufs["oc_chat1"] = _FeishuStreamBuf(
text="Done", card_id="card_1", sequence=3, last_edit=0.0,
)
ch._client.cardkit.v1.card_element.content.return_value = MagicMock(success=MagicMock(return_value=True))
ch._client.cardkit.v1.card.settings.return_value = MagicMock(success=MagicMock(return_value=True))
ch._remove_reaction = AsyncMock()
await ch.send_delta("oc_chat1", "", stream_end=True)
ch._remove_reaction.assert_not_called()
@pytest.mark.asyncio
async def test_no_removal_when_not_stream_end(self):
ch = _make_channel()
@@ -15,6 +15,7 @@ import type {
NanobotFeatureInfo,
NanobotFeaturesPayload,
} from "@/lib/types";
import { useClient } from "@/providers/ClientProvider";
import { FeishuConnectFlow } from "./FeishuConnectFlow";
@@ -33,7 +34,6 @@ export function FeishuAssistantsPanel({
return (
<ChannelInstancesPanel
token={token}
feature={feature}
showBrandLogos={showBrandLogos}
chatAppsDocsUrl={chatAppsDocsUrl}
@@ -92,6 +92,7 @@ function FeishuInstanceAction({
instance: NanobotChannelInstanceInfo;
onFeaturesUpdate: (payload: NanobotFeaturesPayload) => void;
}) {
const { client } = useClient();
const { t } = useTranslation();
const tx = channelTranslator(t, "feishu");
const [busy, setBusy] = useState(false);
@@ -114,7 +115,7 @@ function FeishuInstanceAction({
setError(null);
try {
onFeaturesUpdate(
await enableNanobotFeature(token, "feishu", { instanceId: instance.id }),
await enableNanobotFeature(client, "feishu", { instanceId: instance.id }),
);
} catch (err) {
setError((err as Error).message);
+28 -5
View File
@@ -101,8 +101,14 @@ class ChannelManager:
webui_runtime_surface: str = "browser",
webui_runtime_capabilities: dict[str, Any] | None = None,
webui_skill_state_action: Callable[[set[str]], None] | None = None,
config_path: Path | None = None,
):
if config_path is None:
from nanobot.config.loader import get_config_path
config_path = get_config_path()
self.config = config
self._config_path = config_path.expanduser().resolve(strict=False)
self.bus = bus
self._session_manager = session_manager
self._cron_service = cron_service
@@ -170,6 +176,7 @@ class ChannelManager:
static_dist_path=static_path,
workspace_path=workspace,
default_restrict_to_workspace=self.config.tools.restrict_to_workspace,
config_path=self._config_path,
disabled_skills=set(self.config.agents.defaults.disabled_skills),
runtime_model_name=self._webui_runtime_model_name,
runtime_surface=self._webui_runtime_surface,
@@ -187,11 +194,15 @@ class ChannelManager:
channel = cls(section, self.bus, **kwargs)
if runtime_name and runtime_name != channel.name:
channel.name = runtime_name
progress_default, tool_hints_default = channel.progress_transport_defaults() or (
self.config.channels.send_progress,
self.config.channels.send_tool_hints,
)
channel.send_progress = self._resolve_bool_override(
section, "send_progress", self.config.channels.send_progress,
section, "send_progress", progress_default,
)
channel.send_tool_hints = self._resolve_bool_override(
section, "send_tool_hints", self.config.channels.send_tool_hints,
section, "send_tool_hints", tool_hints_default,
)
channel.show_reasoning = self._resolve_bool_override(
section, "show_reasoning", self.config.channels.show_reasoning,
@@ -347,9 +358,13 @@ class ChannelManager:
await channel.start()
except asyncio.CancelledError:
raise
except Exception:
errors[name] = "Channel failed to start. Check gateway logs."
logger.exception("Failed to start channel {}", name)
except Exception as exc:
public_error = channel.start_error_message(exc)
errors[name] = public_error or "Channel failed to start. Check gateway logs."
if public_error:
logger.error("Failed to start channel {}: {}", name, public_error)
else:
logger.exception("Failed to start channel {}", name)
def _start_channel_task(self, name: str, channel: BaseChannel) -> asyncio.Task[None]:
logger.info("Starting {} channel...", name)
@@ -912,6 +927,14 @@ class ChannelManager:
except asyncio.CancelledError:
raise # Propagate cancellation for graceful shutdown
except Exception as e:
if not channel.should_retry_send_error(e):
logger.error(
"Send to {} failed with a non-retryable {}: {}",
msg.channel,
type(e).__name__,
e,
)
return
loop = asyncio.get_running_loop()
exhausted = (
attempt >= max_attempts
+48 -2
View File
@@ -24,10 +24,12 @@ try:
import nh3
from mistune import HTMLRenderer, create_markdown
from nio import (
Api,
AsyncClient,
AsyncClientConfig,
InviteEvent,
JoinError,
JoinResponse,
KeyVerificationCancel,
KeyVerificationEvent,
KeyVerificationKey,
@@ -43,6 +45,7 @@ try:
RoomSendResponse,
RoomTypingError,
SyncError,
SyncResponse,
ToDeviceError,
UploadError,
)
@@ -701,6 +704,7 @@ class MatrixChannel(BaseChannel):
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)
client.add_response_callback(self._on_sync_invite_fallback, SyncResponse)
def _is_sas_sender_allowed(self, sender: str) -> bool:
return bool(sender and self.is_allowed(sender))
@@ -782,6 +786,49 @@ class MatrixChannel(BaseChannel):
with suppress(Exception):
self.client.stop_sync_forever()
async def _join_room_safe(self, room_id: str) -> bool:
"""Join a room, sending a non-empty POST body.
nio's ``Api.join()`` produces a POST with no body. Some homeservers
(notably Continuwuity) reject empty bodies with ``M_BAD_JSON``.
Sending ``"{}"`` satisfies both strict and lenient servers.
"""
client = self._require_client()
method, path = Api.join(client.access_token, room_id)
try:
resp = cast(
JoinResponse | JoinError,
await client._send( # type: ignore[reportPrivateUsage, reportUnknownMemberType]
JoinResponse, method, path, data="{}"
),
)
except Exception:
self.logger.error("Matrix join request exception for room={}", room_id, exc_info=True)
return False
if isinstance(resp, JoinError):
self.logger.error("Matrix auto-join failed for room={}: {}", room_id, resp)
return False
self.logger.info("Matrix auto-join succeeded: {}", room_id)
return True
async def _on_sync_invite_fallback(self, response: SyncResponse) -> None:
"""Safety net: join pending invites that the event callback may have missed.
Some homeservers (e.g. Continuwuity) deliver each invite only once.
If ``_on_room_invite`` fires but the join fails, the sync token
advances and the invite is never re-delivered. This callback inspects
the same ``SyncResponse`` for pending invites and joins them, acting
as a fallback alongside the event-based callback.
"""
if not response.rooms or not response.rooms.invite:
return
for room_id, invite_info in response.rooms.invite.items():
for event in cast(list[Any], invite_info.invite_state):
sender = getattr(event, "sender", None)
if sender and self.is_allowed(cast(str, sender)):
await self._join_room_safe(room_id)
break
async def _on_join_error(self, response: JoinError) -> None:
self._log_response_error("join", response)
@@ -838,8 +885,7 @@ class MatrixChannel(BaseChannel):
async def _on_room_invite(self, room: MatrixRoom, event: InviteEvent) -> None:
if self.is_allowed(event.sender):
client = self._require_client()
await client.join(room.room_id)
await self._join_room_safe(room.room_id)
def _is_direct_room(self, room: MatrixRoom) -> bool:
count = getattr(room, "member_count", None)
@@ -4,13 +4,14 @@ import asyncio
import sys
from pathlib import Path
from types import SimpleNamespace
from urllib.parse import unquote
import pytest
pytest.importorskip("nio")
pytest.importorskip("nh3")
pytest.importorskip("mistune")
from nio import RoomSendResponse, SyncError
from nio import JoinResponse, RoomSendResponse, SyncError
import nanobot.channels.matrix.runtime as matrix_module
from nanobot.bus.events import OutboundMessage
@@ -104,6 +105,15 @@ class _FakeAsyncClient:
async def join(self, room_id: str) -> None:
self.join_calls.append(room_id)
async def _send(self, response_class, method, path, data=None, **kwargs):
"""Minimal mock for nio's ``_send`` used by ``_join_room_safe``."""
if response_class is JoinResponse and method == "POST" and "/join/" in path:
encoded = path.split("/join/")[1].split("?")[0]
room_id = unquote(encoded)
self.join_calls.append(room_id)
return JoinResponse(room_id=room_id)
return response_class()
async def accept_key_verification(self, transaction_id: str):
self.operation_calls.append(f"accept:{transaction_id}")
self.accept_key_verification_calls.append(transaction_id)
@@ -308,7 +318,7 @@ async def test_start_skips_load_store_when_device_id_missing(
assert clients[0].load_store_called is False
assert len(clients[0].callbacks) == 3
assert clients[0].to_device_callbacks == []
assert len(clients[0].response_callbacks) == 3
assert len(clients[0].response_callbacks) == 4
await channel.stop()
@@ -590,6 +600,7 @@ async def test_room_invite_joins_when_sender_allowed() -> None:
assert client.join_calls == ["!room:matrix.org"]
@pytest.mark.asyncio
async def test_room_invite_respects_allow_list_when_configured() -> None:
channel = MatrixChannel(_make_config(allow_from=["@bob:matrix.org"]), MessageBus())
@@ -604,6 +615,61 @@ async def test_room_invite_respects_allow_list_when_configured() -> None:
assert client.join_calls == []
@pytest.mark.asyncio
async def test_on_sync_invite_fallback_joins_pending_invites() -> None:
"""_on_sync_invite_fallback joins rooms from sync invite_state for allowed senders."""
channel = MatrixChannel(
_make_config(allow_from=["@alice:matrix.org"]), MessageBus()
)
client = _FakeAsyncClient("", "", "", None)
channel.client = client
invite_event = SimpleNamespace(sender="@alice:matrix.org")
invite_info = SimpleNamespace(invite_state=[invite_event])
rooms = SimpleNamespace(invite={"!room:matrix.org": invite_info})
response = SimpleNamespace(rooms=rooms)
await channel._on_sync_invite_fallback(response)
assert client.join_calls == ["!room:matrix.org"]
@pytest.mark.asyncio
async def test_on_sync_invite_fallback_skips_when_no_invites() -> None:
"""_on_sync_invite_fallback is a no-op when sync has no invites."""
channel = MatrixChannel(
_make_config(allow_from=["@alice:matrix.org"]), MessageBus()
)
client = _FakeAsyncClient("", "", "", None)
channel.client = client
rooms = SimpleNamespace(invite={})
response = SimpleNamespace(rooms=rooms)
await channel._on_sync_invite_fallback(response)
assert client.join_calls == []
@pytest.mark.asyncio
async def test_on_sync_invite_fallback_skips_denied_sender() -> None:
"""_on_sync_invite_fallback respects the allow list."""
channel = MatrixChannel(
_make_config(allow_from=["@bob:matrix.org"]), MessageBus()
)
client = _FakeAsyncClient("", "", "", None)
channel.client = client
invite_event = SimpleNamespace(sender="@alice:matrix.org")
invite_info = SimpleNamespace(invite_state=[invite_event])
rooms = SimpleNamespace(invite={"!room:matrix.org": invite_info})
response = SimpleNamespace(rooms=rooms)
await channel._on_sync_invite_fallback(response)
assert client.join_calls == []
@pytest.mark.asyncio
async def test_on_message_sets_typing_for_allowed_sender() -> None:
channel = MatrixChannel(_make_config(), MessageBus())
-8
View File
@@ -658,11 +658,6 @@ class MattermostChannel(BaseChannel):
resp.raise_for_status()
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._require_http_client().put(path, json=json_data)
resp.raise_for_status()
return cast(dict[str, Any], resp.json())
async def _create_post(
self,
channel_id: str,
@@ -681,9 +676,6 @@ class MattermostChannel(BaseChannel):
body["file_ids"] = file_ids
return await self._api_post("/api/v4/posts", body)
async def _edit_post(self, post_id: str, message: str) -> dict[str, Any]:
return await self._api_put(f"/api/v4/posts/{post_id}", {"id": post_id, "message": message})
async def _upload_file(self, channel_id: str, file_path: str) -> str | None:
path = Path(file_path)
if not path.exists():
-5
View File
@@ -811,11 +811,6 @@ class MSTeamsChannel(BaseChannel):
except Exception as e:
self.logger.warning("Failed to save conversation refs: {}", e)
def _save_refs(self, *, prune: bool = True) -> None:
"""Persist conversation references."""
with self._refs_guard:
self._save_refs_locked(prune=prune)
async def _get_access_token(self) -> str:
"""Fetch an access token for Bot Framework / Azure Bot auth."""
@@ -228,7 +228,8 @@ def test_save_prunes_unsupported_conversation_refs(make_channel, tmp_path, monke
),
}
ch._save_refs()
with ch._refs_guard:
ch._save_refs_locked()
assert set(ch._conversation_refs.keys()) == {"conv-valid"}
@@ -378,7 +379,8 @@ def test_save_uses_atomic_replace_and_keeps_existing_file_on_replace_error(make_
raise OSError("replace failed")
monkeypatch.setattr(msteams_module.os, "replace", _raise_replace)
ch._save_refs()
with ch._refs_guard:
ch._save_refs_locked()
persisted = json.loads(refs_path.read_text(encoding="utf-8"))
assert set(persisted.keys()) == {"conv-old"}
@@ -934,7 +936,8 @@ def test_save_refs_prunes_webchat_and_stale_refs(make_channel):
),
}
ch._save_refs()
with ch._refs_guard:
ch._save_refs_locked()
assert set(ch._conversation_refs) == {"teams-good"}
saved = json.loads(ch._refs_path.read_text(encoding="utf-8"))
+2
View File
@@ -431,6 +431,7 @@ class SignalChannel(BaseChannel):
session_key: str | None = None,
is_dm: bool = False,
authorization_id: str | None = None,
require_existing_session: bool = False,
) -> None:
"""Handle an inbound message whose policy has already been checked.
@@ -453,6 +454,7 @@ class SignalChannel(BaseChannel):
media=media or [],
metadata=meta,
session_key_override=session_key,
require_existing_session=require_existing_session,
)
)
+317 -35
View File
@@ -20,7 +20,10 @@ from websockets.asyncio.server import ServerConnection, serve, unix_serve
from websockets.exceptions import ConnectionClosed
from websockets.http11 import Request as WsRequest
from nanobot.bus.events import OUTBOUND_META_AGENT_UI, OutboundMessage
from nanobot.bus.events import (
OUTBOUND_META_AGENT_UI,
OutboundMessage,
)
from nanobot.bus.outbound_events import (
GoalStateSyncEvent,
GoalStatusEvent,
@@ -30,7 +33,6 @@ from nanobot.bus.outbound_events import (
TurnEndEvent,
TurnModelUpdatedEvent,
outbound_event_from_message,
outbound_message_for_event,
)
from nanobot.bus.queue import MessageBus
from nanobot.channels.base import BaseChannel
@@ -49,6 +51,7 @@ from nanobot.security.workspace_access import (
from nanobot.session.goal_state import goal_state_ws_blob
from nanobot.session.webui_turns import (
clear_websocket_turn_if_current,
clear_websocket_turns,
mark_websocket_turn_transcript_persistence_failed,
register_queued_websocket_turn_if_idle,
websocket_turn_id,
@@ -81,6 +84,8 @@ from nanobot.webui.session_access import (
WebuiSessionAccess,
session_mentions_runtime_context,
)
from nanobot.webui.sidebar_state import write_webui_sidebar_state
from nanobot.webui.temporary_chats import TemporaryChatError
from nanobot.webui.transcript import WEBUI_TRANSCRIPT_INCOMPLETE_KEY
from nanobot.webui.transcription_ws import webui_transcription_event
from nanobot.webui.websocket_logging import websockets_server_logger
@@ -279,21 +284,6 @@ class WebSocketConfig(Base):
)
def publish_runtime_model_update(
bus: MessageBus,
model: str,
model_preset: str | None,
) -> None:
"""Enqueue a runtime model snapshot for websocket subscribers (fan-out in-channel)."""
bus.outbound.put_nowait(
outbound_message_for_event(
channel="websocket",
chat_id="*",
event=RuntimeModelUpdatedEvent(model=model, model_preset=model_preset),
)
)
def _parse_inbound_payload(raw: str) -> str | None:
"""Parse a client frame into text; return None for empty or unrecognized content."""
text = raw.strip()
@@ -383,6 +373,13 @@ class WebSocketChannel(BaseChannel):
self._conn_default: dict[ServerConnection, str] = {}
# Connections authenticated with a one-time token from /webui/bootstrap.
self._webui_connections: set[ServerConnection] = set()
# Request/reply mutations aren't replayed across reconnects. Tasks may
# finish after a client-side deadline so an already-started mutation
# isn't ambiguously cancelled halfway through.
self._webui_request_tasks: dict[
tuple[ServerConnection, str],
asyncio.Task[None],
] = {}
self._stop_event: asyncio.Event | None = None
self._server_task: asyncio.Task[None] | None = None
@@ -393,6 +390,7 @@ class WebSocketChannel(BaseChannel):
self._ingress = gateway.ingress
self._transcripts = gateway.transcripts
self._workspaces = gateway.workspaces
self._temporary_chats = gateway.temporary_chats
self._session_access = (
WebuiSessionAccess(gateway.session_manager)
if gateway.session_manager is not None
@@ -411,6 +409,33 @@ class WebSocketChannel(BaseChannel):
self._subs.setdefault(chat_id, set()).add(connection)
self._conn_chats.setdefault(connection, set()).add(chat_id)
def _detach(self, connection: ServerConnection, chat_id: str) -> None:
chats = self._conn_chats.get(connection)
if chats is not None:
chats.discard(chat_id)
if not chats:
self._conn_chats.pop(connection, None)
subscribers = self._subs.get(chat_id)
if subscribers is not None:
subscribers.discard(connection)
if not subscribers:
self._subs.pop(chat_id, None)
def _clear_stream_buffers(self, chat_id: str) -> None:
for key in tuple(self._stream_text_buffers):
if key[0] == chat_id:
self._stream_text_buffers.pop(key, None)
async def _discard_connection_owned_chat(
self,
connection: ServerConnection,
chat_id: str,
) -> None:
await self._temporary_chats.discard(connection, chat_id)
self._detach(connection, chat_id)
clear_websocket_turns(chat_id)
self._clear_stream_buffers(chat_id)
async def send_webui_protocol_error(
self,
connection: ServerConnection,
@@ -439,16 +464,16 @@ class WebSocketChannel(BaseChannel):
)
await self._hydrate_after_subscribe(fork_id)
def _cleanup_connection(self, connection: ServerConnection) -> None:
async 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())
chat_ids = tuple(self._conn_chats.get(connection, ()))
for cid in chat_ids:
subs = self._subs.get(cid)
if subs is None:
continue
subs.discard(connection)
if not subs:
self._subs.pop(cid, None)
if self._temporary_chats.owns(connection, cid):
await self._discard_connection_owned_chat(connection, cid)
else:
self._detach(connection, cid)
for cid in self._temporary_chats.chat_ids_for_owner(connection):
await self._discard_connection_owned_chat(connection, cid)
self._conn_default.pop(connection, None)
self._webui_connections.discard(connection)
@@ -501,7 +526,7 @@ class WebSocketChannel(BaseChannel):
try:
await connection.send(raw)
except ConnectionClosed:
self._cleanup_connection(connection)
await self._cleanup_connection(connection)
except Exception as e:
self.logger.warning("failed to send {} event: {}", event, e)
@@ -728,7 +753,7 @@ class WebSocketChannel(BaseChannel):
except Exception as e:
self.logger.debug("connection ended: {}", e)
finally:
self._cleanup_connection(connection)
await self._cleanup_connection(connection)
# -- Inbound WebSocket envelopes ---------------------------------------
@@ -740,6 +765,9 @@ class WebSocketChannel(BaseChannel):
) -> None:
"""Route one typed inbound envelope (``new_chat`` / ``attach`` / ``message``)."""
t = envelope.get("type")
if t == "webui_request":
await self._start_webui_request(connection, envelope)
return
if t == "new_chat":
new_id = str(uuid.uuid4())
scope = await self._workspace_scope_or_error(
@@ -763,23 +791,84 @@ class WebSocketChannel(BaseChannel):
)
await self._hydrate_after_subscribe(new_id)
return
if t == "new_temporary_chat":
try:
new_id = self._temporary_chats.create(
connection,
trusted_webui=connection in self._webui_connections,
)
except TemporaryChatError as exc:
await self._send_event(connection, "error", detail=exc.detail)
return
self._attach(connection, new_id)
await self._send_event(
connection,
"attached",
chat_id=new_id,
temporary=True,
)
return
if t == "fork_chat":
await handle_webui_fork_chat(self, connection, envelope)
return
if t == "discard_temporary_chat":
cid = envelope.get("chat_id")
if not _is_valid_chat_id(cid):
await self._send_event(connection, "error", detail="invalid temporary chat_id")
return
try:
await self._discard_connection_owned_chat(connection, cid)
except TemporaryChatError as exc:
await self._send_event(connection, "error", detail=exc.detail, chat_id=cid)
return
if t == "attach":
cid = envelope.get("chat_id")
if not _is_valid_chat_id(cid):
await self._send_event(connection, "error", detail="invalid chat_id")
return
try:
self._temporary_chats.validate_attach(cid)
except TemporaryChatError as exc:
await self._send_event(connection, "error", detail=exc.detail, chat_id=cid)
return
self._attach(connection, cid)
await self._send_event(connection, "attached", chat_id=cid)
await self._hydrate_after_subscribe(cid)
return
if t == "set_sidebar_state":
if connection not in self._webui_connections:
await self._send_event(connection, "error", detail="access_denied")
return
state = envelope.get("state")
if not isinstance(state, dict):
await self._send_event(
connection,
"error",
detail="invalid_sidebar_state",
)
return
try:
await asyncio.to_thread(
write_webui_sidebar_state,
cast(dict[str, Any], state),
)
except (OSError, ValueError):
await self._send_event(
connection,
"error",
detail="invalid_sidebar_state",
)
return
if t == "set_workspace_scope":
cid = envelope.get("chat_id")
if not _is_valid_chat_id(cid):
await self._send_event(connection, "error", detail="invalid chat_id")
return
try:
self._temporary_chats.validate_workspace_update(cid)
except TemporaryChatError as exc:
await self._send_event(connection, "error", detail=exc.detail, chat_id=cid)
return
scope = await self._workspace_scope_or_error(
connection,
lambda: self._workspaces.scope_for_set_request(
@@ -848,6 +937,21 @@ class WebSocketChannel(BaseChannel):
)
return
try:
temporary_policy = self._temporary_chats.message_policy(
connection,
cid,
content,
)
except TemporaryChatError as exc:
await self._send_event(
connection,
"error",
detail=exc.detail,
**rejection_fields,
)
return
raw_media = envelope.get("media")
media_paths: list[str] = []
if raw_media is not None:
@@ -870,6 +974,8 @@ class WebSocketChannel(BaseChannel):
**rejection_fields,
)
return
if temporary_policy is not None:
self._temporary_chats.register_media(connection, cid, media_paths)
# Allow media-only turns (content may be empty when attachments are present).
if not content.strip() and not media_paths:
@@ -882,16 +988,21 @@ class WebSocketChannel(BaseChannel):
return
# Auto-attach on first use so clients can one-shot without a separate attach.
self._attach(connection, cid)
await self._hydrate_after_subscribe(cid)
if temporary_policy is None or temporary_policy.hydrate_transcript:
await self._hydrate_after_subscribe(cid)
# Resolve after hydration so a concurrent downgrade cannot be overwritten.
scope = await self._workspace_scope_or_error(
connection,
lambda: self._workspaces.scope_for_message(
envelope,
chat_id=cid,
chat_running=websocket_turn_wall_started_at(cid) is not None,
controls_available=self._workspace_controls_available(connection),
lambda: (
temporary_policy.workspace_scope
if temporary_policy is not None
else self._workspaces.scope_for_message(
envelope,
chat_id=cid,
chat_running=websocket_turn_wall_started_at(cid) is not None,
controls_available=self._workspace_controls_available(connection),
)
),
chat_id=cid,
turn_id=turn_id,
@@ -944,7 +1055,13 @@ class WebSocketChannel(BaseChannel):
metadata[WEBSOCKET_TURN_OWNER_METADATA_KEY] = queued_owner
accepted = False
try:
if is_webui:
if (
is_webui
and (
temporary_policy is None
or temporary_policy.persist_transcript
)
):
self._transcripts.append_user_message(
cid,
content,
@@ -973,6 +1090,16 @@ class WebSocketChannel(BaseChannel):
media=media_paths or None,
metadata=metadata,
is_dm=False,
session_key=(
temporary_policy.session_key
if temporary_policy is not None
else None
),
require_existing_session=(
temporary_policy.require_existing_session
if temporary_policy is not None
else False
),
)
accepted = True
finally:
@@ -988,6 +1115,152 @@ class WebSocketChannel(BaseChannel):
return
await self._send_event(connection, "error", detail=f"unknown type: {t!r}")
async def _start_webui_request(
self,
connection: ServerConnection,
envelope: dict[str, Any],
) -> None:
request_id = envelope.get("request_id")
if not isinstance(request_id, str) or re.fullmatch(
r"[A-Za-z0-9._:-]{1,128}",
request_id,
) is None:
await self._send_event(
connection,
"error",
detail="invalid webui request_id",
)
return
if connection not in self._webui_connections:
await self._send_webui_response(
connection,
request_id,
status=403,
message="access_denied",
)
return
action = envelope.get("action")
payload = envelope.get("payload")
if not isinstance(action, str) or re.fullmatch(
r"[a-z][a-z0-9_.]{0,127}",
action,
) is None:
await self._send_webui_response(
connection,
request_id,
status=400,
message="invalid WebUI mutation action",
)
return
if not isinstance(payload, dict):
await self._send_webui_response(
connection,
request_id,
status=400,
message="WebUI mutation payload must be an object",
)
return
key = (connection, request_id)
if key in self._webui_request_tasks:
await self._send_webui_response(
connection,
request_id,
status=409,
message="duplicate WebUI request_id",
)
return
task = asyncio.create_task(
self._complete_webui_request(
connection,
request_id,
action,
cast(dict[str, Any], payload),
)
)
self._webui_request_tasks[key] = task
async def _complete_webui_request(
self,
connection: ServerConnection,
request_id: str,
action: str,
payload: dict[str, Any],
) -> None:
try:
response = await self._http_router.dispatch_webui_mutation(
connection,
action,
payload,
)
status = response.status_code
body = bytes(response.body).decode("utf-8", errors="replace").strip()
if 200 <= status < 300:
try:
result = json.loads(body)
except json.JSONDecodeError:
await self._send_webui_response(
connection,
request_id,
status=502,
message="WebUI mutation returned an invalid response",
)
return
await self._send_webui_response(
connection,
request_id,
result=result,
)
return
await self._send_webui_response(
connection,
request_id,
status=status,
message=body or response.reason_phrase,
)
except asyncio.CancelledError:
raise
except Exception:
self.logger.exception("WebUI mutation '{}' failed", action)
await self._send_webui_response(
connection,
request_id,
status=500,
message="WebUI mutation failed",
)
finally:
self._webui_request_tasks.pop((connection, request_id), None)
async def _send_webui_response(
self,
connection: ServerConnection,
request_id: str,
*,
result: Any = None,
status: int | None = None,
message: str | None = None,
) -> None:
if status is None:
await self._send_event(
connection,
"webui_response",
request_id=request_id,
ok=True,
result=result,
)
return
await self._send_event(
connection,
"webui_response",
request_id=request_id,
ok=False,
error={
"status": status,
"message": message or "WebUI mutation failed",
},
)
async def _workspace_scope_or_error(
self,
connection: ServerConnection,
@@ -1028,11 +1301,18 @@ class WebSocketChannel(BaseChannel):
except Exception as e:
self.logger.warning("server task error during shutdown: {}", e)
self._server_task = None
mutation_tasks = tuple(self._webui_request_tasks.values())
for task in mutation_tasks:
task.cancel()
if mutation_tasks:
await asyncio.gather(*mutation_tasks, return_exceptions=True)
self._webui_request_tasks.clear()
self._subs.clear()
self._conn_chats.clear()
self._conn_default.clear()
self._webui_connections.clear()
self._tokens.clear()
self._temporary_chats.close()
async def _safe_send_to(
self,
@@ -1045,7 +1325,7 @@ class WebSocketChannel(BaseChannel):
try:
await connection.send(raw)
except ConnectionClosed:
self._cleanup_connection(connection)
await self._cleanup_connection(connection)
self.logger.warning("connection gone{}", label)
except Exception:
self.logger.exception("send failed{}", label)
@@ -1062,6 +1342,8 @@ class WebSocketChannel(BaseChannel):
transcript_overrides: dict[str, Any] | None = None,
) -> bool:
"""Persist one canonical turn event and retain unsafe owners on failure."""
if not self._temporary_chats.should_persist_transcript(chat_id):
return True
persisted = self._transcripts.prepare_and_append(
chat_id,
event,
@@ -3,17 +3,23 @@
import asyncio
import json
import time
import uuid
from pathlib import Path
from types import SimpleNamespace
from typing import Any
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
import websockets
from websockets.datastructures import Headers
from websockets.exceptions import ConnectionClosed
from websockets.frames import Close
from nanobot.bus.events import (
INBOUND_META_RUNTIME_CONTROL,
OUTBOUND_META_AGENT_UI,
RUNTIME_CONTROL_SESSION_DISCARD,
OutboundMessage,
)
from nanobot.bus.outbound_events import (
@@ -32,14 +38,20 @@ from nanobot.channels.websocket.runtime import (
_is_valid_chat_id,
_parse_envelope,
_parse_inbound_payload,
publish_runtime_model_update,
)
from nanobot.config.loader import load_config, save_config
from nanobot.config.schema import Config, ModelPresetConfig
from nanobot.runtime_context import RUNTIME_CONTEXT_INPUT_META, WEBUI_QUOTE_SOURCE
from nanobot.security.workspace_access import WORKSPACE_SCOPE_METADATA_KEY
from nanobot.session import webui_turns as wth
from nanobot.session.manager import SessionManager
from nanobot.webui.gateway_services import GatewayServices, build_gateway_services
from nanobot.webui.http_utils import (
http_error as _http_error,
)
from nanobot.webui.http_utils import (
http_json_response as _http_json_response,
)
from nanobot.webui.http_utils import (
issue_route_secret_matches as _issue_route_secret_matches,
)
@@ -117,6 +129,38 @@ def _basic_handler(bus: Any, **kw: Any) -> GatewayServices:
)
async def _webui_mutate(
client: Any,
action: str,
payload: dict[str, Any] | None = None,
) -> httpx.Response:
request_id = f"test-{uuid.uuid4().hex}"
await client.send(json.dumps({
"type": "webui_request",
"request_id": request_id,
"action": action,
"payload": payload or {},
}))
while True:
envelope = json.loads(await asyncio.wait_for(client.recv(), timeout=5))
if envelope.get("event") != "webui_response":
continue
if envelope.get("request_id") != request_id:
continue
if envelope.get("ok") is True:
status = 200
body = envelope.get("result")
else:
error = envelope.get("error") or {}
status = int(error.get("status") or 500)
body = {"error": str(error.get("message") or "WebUI mutation failed")}
return httpx.Response(
status,
json=body,
request=httpx.Request("WS", "http://nanobot.local/webui-mutation"),
)
@pytest.mark.asyncio
async def test_stop_treats_cancelled_server_task_as_shutdown() -> None:
channel = _ch(MessageBus())
@@ -193,6 +237,302 @@ def isolate_webui_workspace_state(tmp_path, monkeypatch) -> None:
wth._WEBSOCKET_TURN_OWNERS.clear()
async def _new_temporary_chat(
channel: WebSocketChannel,
connection: AsyncMock,
) -> str:
channel._webui_connections.add(connection)
await channel._dispatch_envelope(
connection,
"webui-client",
{"type": "new_temporary_chat"},
)
payload = json.loads(connection.send.await_args.args[0])
assert payload["event"] == "attached"
assert payload["temporary"] is True
connection.send.reset_mock()
return payload["chat_id"]
@pytest.mark.asyncio
async def test_temporary_chat_is_transient_and_discarded(bus, tmp_path) -> None:
sessions = SessionManager(tmp_path)
selected_project = tmp_path / "selected-project"
selected_project.mkdir()
channel = WebSocketChannel(
{"enabled": True, "allowFrom": ["*"]},
bus,
gateway=_basic_handler(
bus,
session_manager=sessions,
workspace_path=tmp_path,
),
)
connection = AsyncMock()
connection.remote_address = ("127.0.0.1", 5000)
chat_id = await _new_temporary_chat(channel, connection)
upload = tmp_path / "temporary-upload.txt"
upload.write_text("private attachment", encoding="utf-8")
channel.gateway.media.store_inbound_attachments = MagicMock(
return_value=([str(upload)], None),
)
await channel._dispatch_envelope(
connection,
"webui-client",
{
"type": "message",
"chat_id": chat_id,
"content": "read this",
"media": [{"data_url": "data:text/plain;base64,cHJpdmF0ZQ=="}],
"cli_apps": [{"name": "drawio"}],
"workspace_scope": {
"project_path": str(selected_project),
"access_mode": "full",
},
"turn_id": "turn-1",
"webui": True,
},
)
inbound = bus.publish_inbound.await_args_list[0].args[0]
assert inbound.session_key == f"websocket:{chat_id}"
assert inbound.session_key_override == f"websocket:{chat_id}"
assert inbound.require_existing_session is True
assert inbound.metadata["cli_apps"] == [{"name": "drawio"}]
assert inbound.metadata[WORKSPACE_SCOPE_METADATA_KEY] == {
"project_path": str(tmp_path.resolve()),
"access_mode": "restricted",
}
session = sessions.get_cached(inbound.session_key)
assert session is not None
assert session.policy.persist is False
assert upload.exists()
assert read_transcript_lines(inbound.session_key) == []
assert [payload["event"] for payload in _sent_ws_payloads(connection)] == [
"message_accepted",
]
await channel._dispatch_envelope(
connection,
"webui-client",
{"type": "discard_temporary_chat", "chat_id": chat_id},
)
control = bus.publish_inbound.await_args_list[1].args[0]
assert bus.publish_inbound.await_count == 2
assert control.session_key == inbound.session_key
assert control.metadata[INBOUND_META_RUNTIME_CONTROL] == (
RUNTIME_CONTROL_SESSION_DISCARD
)
assert sessions.get_cached(inbound.session_key) is None
assert chat_id not in channel._subs
assert chat_id not in channel._conn_chats.get(connection, set())
assert not upload.exists()
assert read_transcript_lines(inbound.session_key) == []
@pytest.mark.asyncio
@pytest.mark.parametrize("content", ["/goal private", "/trigger later", "/dream"])
async def test_temporary_chat_rejects_persistent_commands(bus, tmp_path, content) -> None:
sessions = SessionManager(tmp_path)
channel = WebSocketChannel(
{"enabled": True, "allowFrom": ["*"]},
bus,
gateway=_basic_handler(bus, session_manager=sessions, workspace_path=tmp_path),
)
connection = AsyncMock()
connection.remote_address = ("127.0.0.1", 5000)
chat_id = await _new_temporary_chat(channel, connection)
await channel._dispatch_envelope(connection, "webui-client", {
"type": "message",
"chat_id": chat_id,
"content": content,
"webui": True,
})
assert bus.publish_inbound.await_count == 0
assert sessions.get_cached(f"websocket:{chat_id}") is not None
assert json.loads(connection.send.await_args.args[0])["detail"] == (
"temporary_chat_command_rejected"
)
@pytest.mark.asyncio
async def test_disconnect_discards_temporary_chat(bus, tmp_path) -> None:
sessions = SessionManager(tmp_path)
channel = WebSocketChannel(
{"enabled": True, "allowFrom": ["*"]},
bus,
gateway=_basic_handler(
bus,
session_manager=sessions,
workspace_path=tmp_path,
),
)
connection = AsyncMock()
chat_id = await _new_temporary_chat(channel, connection)
await channel._dispatch_envelope(
connection,
"webui-client",
{
"type": "message",
"chat_id": chat_id,
"content": "hello",
"webui": True,
},
)
await channel._cleanup_connection(connection)
session_key = f"websocket:{chat_id}"
control = bus.publish_inbound.await_args_list[-1].args[0]
assert control.session_key == session_key
assert control.metadata[INBOUND_META_RUNTIME_CONTROL] == (
RUNTIME_CONTROL_SESSION_DISCARD
)
assert sessions.get_cached(session_key) is None
assert chat_id not in channel._subs
@pytest.mark.asyncio
async def test_temporary_chat_creation_requires_authenticated_webui_connection(bus, tmp_path) -> None:
sessions = SessionManager(tmp_path)
channel = WebSocketChannel(
{"enabled": True, "allowFrom": ["*"]},
bus,
gateway=_basic_handler(bus, session_manager=sessions, workspace_path=tmp_path),
)
connection = AsyncMock()
await channel._dispatch_envelope(
connection,
"generic-websocket-client",
{"type": "new_temporary_chat"},
)
assert json.loads(connection.send.await_args.args[0])["detail"] == "access_denied"
assert sessions.list_sessions() == []
@pytest.mark.asyncio
async def test_temporary_chat_cannot_be_claimed_by_another_connection(bus, tmp_path) -> None:
sessions = SessionManager(tmp_path)
channel = WebSocketChannel(
{"enabled": True, "allowFrom": ["*"]},
bus,
gateway=_basic_handler(bus, session_manager=sessions, workspace_path=tmp_path),
)
owner = AsyncMock()
other = AsyncMock()
channel._webui_connections.add(other)
chat_id = await _new_temporary_chat(channel, owner)
await channel._dispatch_envelope(
other,
"other-webui-client",
{
"type": "message",
"chat_id": chat_id,
"content": "claim it",
"webui": True,
},
)
assert json.loads(other.send.await_args.args[0])["detail"] == (
"temporary_chat_unavailable"
)
assert bus.publish_inbound.await_count == 0
assert sessions.get_cached(f"websocket:{chat_id}") is not None
@pytest.mark.asyncio
async def test_temporary_chat_cannot_persist_workspace_scope(bus, tmp_path) -> None:
sessions = SessionManager(tmp_path)
channel = WebSocketChannel(
{"enabled": True, "allowFrom": ["*"]},
bus,
gateway=_basic_handler(bus, session_manager=sessions, workspace_path=tmp_path),
)
connection = AsyncMock()
chat_id = await _new_temporary_chat(channel, connection)
await channel._dispatch_envelope(
connection,
"webui-client",
{
"type": "set_workspace_scope",
"chat_id": chat_id,
"workspace_scope": {
"project_path": str(tmp_path),
"access_mode": "full",
},
},
)
payload = json.loads(connection.send.await_args.args[0])
assert payload["detail"] == "temporary_chat_workspace_rejected"
session = sessions.get_cached(f"websocket:{chat_id}")
assert session is not None
assert WORKSPACE_SCOPE_METADATA_KEY not in session.metadata
assert sessions.list_sessions() == []
@pytest.mark.asyncio
async def test_temporary_looking_id_does_not_define_session_policy(bus, tmp_path) -> None:
sessions = SessionManager(tmp_path)
channel = WebSocketChannel(
{"enabled": True, "allowFrom": ["*"]},
bus,
gateway=_basic_handler(bus, session_manager=sessions, workspace_path=tmp_path),
)
connection = AsyncMock()
channel._webui_connections.add(connection)
await channel._dispatch_envelope(
connection,
"webui-client",
{
"type": "message",
"chat_id": "temporary-looking-but-persistent",
"content": "/goal ordinary chat",
"webui": True,
},
)
inbound = bus.publish_inbound.await_args.args[0]
assert inbound.require_existing_session is False
assert inbound.session_key_override is None
session = sessions.get_cached("websocket:temporary-looking-but-persistent")
assert session is not None
assert session.policy.persist is True
@pytest.mark.asyncio
async def test_discard_temporary_chat_does_not_detach_persistent_chat(bus, tmp_path) -> None:
sessions = SessionManager(tmp_path)
channel = WebSocketChannel(
{"enabled": True, "allowFrom": ["*"]},
bus,
gateway=_basic_handler(bus, session_manager=sessions, workspace_path=tmp_path),
)
connection = AsyncMock()
channel._attach(connection, "ordinary-chat")
await channel._dispatch_envelope(
connection,
"webui-client",
{"type": "discard_temporary_chat", "chat_id": "ordinary-chat"},
)
assert json.loads(connection.send.await_args.args[0])["detail"] == (
"temporary_chat_unavailable"
)
assert connection in channel._subs["ordinary-chat"]
assert "ordinary-chat" in channel._conn_chats[connection]
@pytest.mark.asyncio
async def test_send_session_updated_broadcasts_to_other_webui_connections(bus) -> None:
class Conn:
@@ -559,6 +899,136 @@ def test_only_bootstrap_tokens_mark_webui_connections(bus: MagicMock) -> None:
assert client_connection not in channel._webui_connections
@pytest.mark.asyncio
async def test_authenticated_webui_request_returns_correlated_success(bus: MagicMock) -> None:
channel = _ch(bus)
conn = AsyncMock()
channel._webui_connections.add(conn)
channel.gateway.http.dispatch_webui_mutation = AsyncMock(
return_value=_http_json_response({"saved": True})
)
await channel._dispatch_envelope(
conn,
"webui-client",
{
"type": "webui_request",
"request_id": "request-1",
"action": "settings.provider.update",
"payload": {"provider": "openrouter", "apiKey": "secret"},
},
)
await asyncio.gather(*tuple(channel._webui_request_tasks.values()))
channel.gateway.http.dispatch_webui_mutation.assert_awaited_once_with(
conn,
"settings.provider.update",
{"provider": "openrouter", "apiKey": "secret"},
)
assert json.loads(conn.send.await_args.args[0]) == {
"event": "webui_response",
"request_id": "request-1",
"ok": True,
"result": {"saved": True},
}
@pytest.mark.asyncio
async def test_webui_request_returns_correlated_route_error(bus: MagicMock) -> None:
channel = _ch(bus)
conn = AsyncMock()
channel._webui_connections.add(conn)
channel.gateway.http.dispatch_webui_mutation = AsyncMock(
return_value=_http_error(400, "invalid settings payload")
)
await channel._dispatch_envelope(
conn,
"webui-client",
{
"type": "webui_request",
"request_id": "request-2",
"action": "settings.agent.update",
"payload": {},
},
)
await asyncio.gather(*tuple(channel._webui_request_tasks.values()))
assert json.loads(conn.send.await_args.args[0]) == {
"event": "webui_response",
"request_id": "request-2",
"ok": False,
"error": {"status": 400, "message": "invalid settings payload"},
}
@pytest.mark.asyncio
async def test_webui_request_requires_bootstrap_authenticated_connection(
bus: MagicMock,
) -> None:
channel = _ch(bus)
conn = AsyncMock()
channel.gateway.http.dispatch_webui_mutation = AsyncMock()
await channel._dispatch_envelope(
conn,
"static-token-client",
{
"type": "webui_request",
"request_id": "request-3",
"action": "settings.agent.update",
"payload": {},
},
)
channel.gateway.http.dispatch_webui_mutation.assert_not_awaited()
assert json.loads(conn.send.await_args.args[0]) == {
"event": "webui_response",
"request_id": "request-3",
"ok": False,
"error": {"status": 403, "message": "access_denied"},
}
@pytest.mark.asyncio
async def test_webui_persists_sidebar_state_larger_than_http_request_line(
bus: MagicMock,
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
channel = _ch(bus)
conn = AsyncMock()
conn.request = SimpleNamespace(headers=Headers())
channel._webui_connections.add(conn)
session_order = [f"websocket:{index:04d}-{'x' * 48}" for index in range(160)]
request_id = "sidebar-large-state"
envelope = {
"type": "webui_request",
"request_id": request_id,
"action": "sidebar.update",
"payload": {"state": {
"session_order": session_order,
"view": {"sort": "manual"},
}},
}
assert len(json.dumps(envelope).encode()) > 8_192
await channel._dispatch_envelope(conn, "webui-client", envelope)
await asyncio.gather(*tuple(channel._webui_request_tasks.values()))
saved = json.loads((tmp_path / "webui" / "sidebar-state.json").read_text(encoding="utf-8"))
assert saved["session_order"] == session_order
assert saved["view"]["sort"] == "manual"
assert json.loads(conn.send.await_args.args[0]) == {
"event": "webui_response",
"request_id": request_id,
"ok": True,
"result": saved,
}
@pytest.mark.asyncio
async def test_client_cannot_self_assert_webui_quote_context(bus: MagicMock) -> None:
channel = _ch(bus)
@@ -1077,8 +1547,14 @@ async def test_send_broadcasts_runtime_model_updates() -> None:
mock_ws = AsyncMock()
channel._attach(mock_ws, "chat-1")
publish_runtime_model_update(bus, "openai/gpt-4.1", "fast")
await channel.send(bus.outbound.get_nowait())
await channel.send(
OutboundMessage(
channel="websocket",
chat_id="*",
content="",
event=RuntimeModelUpdatedEvent(model="openai/gpt-4.1", model_preset="fast"),
)
)
payload = json.loads(mock_ws.send.call_args[0][0])
assert payload["event"] == "runtime_model_updated"
@@ -1113,26 +1589,6 @@ async def test_send_scopes_turn_model_updates_to_the_subscribed_chat() -> None:
chat_two.send.assert_not_awaited()
@pytest.mark.asyncio
async def test_runtime_model_update_publisher_uses_websocket_outbound_event() -> None:
bus = MessageBus()
publish_runtime_model_update(
bus,
"openai/gpt-4.1",
"fast",
)
event = bus.outbound.get_nowait()
assert event.channel == "websocket"
assert event.chat_id == "*"
assert event.content == ""
assert event.metadata == {}
assert isinstance(event.event, RuntimeModelUpdatedEvent)
assert event.event.model == "openai/gpt-4.1"
assert event.event.model_preset == "fast"
@pytest.mark.asyncio
async def test_send_stages_external_media_as_signed_url(monkeypatch, tmp_path) -> None:
bus = MagicMock()
@@ -2575,7 +3031,15 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
server_task = asyncio.create_task(channel.start())
await asyncio.sleep(0.3)
webui_client = None
try:
webui_token = channel.gateway.tokens.issue_token(300, audience="webui")
webui_client = await websockets.connect(
f"ws://127.0.0.1:{port}/ws?token={webui_token}&client_id=settings-test"
)
ready = json.loads(await asyncio.wait_for(webui_client.recv(), timeout=5))
assert ready["event"] == "ready"
settings = await _http_get(
f"http://127.0.0.1:{port}/api/settings",
headers={"Authorization": "Bearer tok"},
@@ -2659,11 +3123,14 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
assert unknown_api.status_code == 404
assert "<!doctype html>" not in unknown_api.text.lower()
provider_updated = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/provider/update?provider=openrouter"
"&api_key=sk-or-test&api_base=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1",
headers={"Authorization": "Bearer tok"},
provider_updated = await _webui_mutate(
webui_client,
"settings.provider.update",
{
"provider": "openrouter",
"apiKey": "sk-or-test",
"apiBase": "https://openrouter.ai/api/v1",
},
)
assert provider_updated.status_code == 200
provider_body = provider_updated.json()
@@ -2673,22 +3140,18 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
assert provider_body["image_generation"]["provider_configured"] is True
assert "sk-or-test" not in provider_updated.text
custom_provider_created = await _http_get(
f"http://127.0.0.1:{port}/api/settings/provider/create",
headers={
"Authorization": "Bearer tok",
"X-Nanobot-Provider-Values": json.dumps(
{
"name": "Company Gateway",
"apiBase": "https://gateway.example/v1",
"apiKey": "sk-company",
"extraHeaders": json.dumps({"X-Tenant": "engineering"}),
"extraBody": json.dumps({"service_tier": "priority"}),
"extraQuery": json.dumps({"api-version": "2026-01-01"}),
"proxy": "http://127.0.0.1:7890",
"thinkingStyle": "enable_thinking",
}
),
custom_provider_created = await _webui_mutate(
webui_client,
"settings.provider.create",
{
"name": "Company Gateway",
"apiBase": "https://gateway.example/v1",
"apiKey": "sk-company",
"extraHeaders": json.dumps({"X-Tenant": "engineering"}),
"extraBody": json.dumps({"service_tier": "priority"}),
"extraQuery": json.dumps({"api-version": "2026-01-01"}),
"proxy": "http://127.0.0.1:7890",
"thinkingStyle": "enable_thinking",
},
)
assert custom_provider_created.status_code == 200
@@ -2703,11 +3166,10 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
}
assert "sk-company" not in custom_provider_created.text
local_provider_updated = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/provider/update?provider=atomic_chat"
"&api_base=http%3A%2F%2Flocalhost%3A1337%2Fv1",
headers={"Authorization": "Bearer tok"},
local_provider_updated = await _webui_mutate(
webui_client,
"settings.provider.update",
{"provider": "atomic_chat", "apiBase": "http://localhost:1337/v1"},
)
assert local_provider_updated.status_code == 200
local_provider_body = local_provider_updated.json()
@@ -2717,38 +3179,44 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
assert local_provider_rows["atomic_chat"]["configured"] is True
assert "localhost:1337" in local_provider_updated.text
updated = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/update?model=atomic_chat/test"
"&provider=atomic_chat&timezone=Asia%2FShanghai"
"&bot_name=Nano&bot_icon=N&tool_hint_max_length=120",
headers={"Authorization": "Bearer tok"},
updated = await _webui_mutate(
webui_client,
"settings.agent.update",
{
"model": "atomic_chat/test",
"provider": "atomic_chat",
"timezone": "Asia/Shanghai",
"tool_hint_max_length": 120,
},
)
assert updated.status_code == 200
updated_body = updated.json()
assert updated_body["requires_restart"] is True
assert updated_body["restart_required_sections"] == ["runtime"]
preset_updated = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/update?model_preset=deep",
headers={"Authorization": "Bearer tok"},
preset_updated = await _webui_mutate(
webui_client,
"settings.agent.update",
{"model_preset": "deep"},
)
assert preset_updated.status_code == 200
assert preset_updated.json()["agent"]["model"] == "anthropic/claude-opus-4-5"
bad_preset = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/update?model_preset=missing",
headers={"Authorization": "Bearer tok"},
bad_preset = await _webui_mutate(
webui_client,
"settings.agent.update",
{"model_preset": "missing"},
)
assert bad_preset.status_code == 400
created_preset = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/model-configurations/create"
"?label=Fast%20writing&provider=openai&model=openai%2Fgpt-4.1-mini",
headers={"Authorization": "Bearer tok"},
created_preset = await _webui_mutate(
webui_client,
"settings.model_configuration.create",
{
"label": "Fast writing",
"provider": "openai",
"model": "openai/gpt-4.1-mini",
},
)
assert created_preset.status_code == 200
created_body = created_preset.json()
@@ -2762,11 +3230,15 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
assert created_presets["fast-writing"]["label"] == "Fast writing"
assert created_presets["fast-writing"]["provider"] == "openai"
updated_preset = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/model-configurations/update"
"?name=fast-writing&label=Codex&provider=openai&model=openai%2Fgpt-5.5",
headers={"Authorization": "Bearer tok"},
updated_preset = await _webui_mutate(
webui_client,
"settings.model_configuration.update",
{
"name": "fast-writing",
"label": "Codex",
"provider": "openai",
"model": "openai/gpt-5.5",
},
)
assert updated_preset.status_code == 200
updated_preset_body = updated_preset.json()
@@ -2777,11 +3249,10 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
}
assert updated_presets["fast-writing"]["label"] == "Codex"
call_order_updated = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/model-call-order/update"
"?order=%5B%22fast-writing%22%2C%22deep%22%5D",
headers={"Authorization": "Bearer tok"},
call_order_updated = await _webui_mutate(
webui_client,
"settings.model_call_order.update",
{"order": ["fast-writing", "deep"]},
)
assert call_order_updated.status_code == 200
call_order_body = call_order_updated.json()
@@ -2789,20 +3260,27 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
assert call_order_body["agent"]["model"] == "openai/gpt-5.5"
assert call_order_body["model_call_order"] == ["fast-writing", "deep"]
duplicate_preset = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/model-configurations/create"
"?label=Fast%20writing&provider=openai&model=openai%2Fgpt-4.1-mini",
headers={"Authorization": "Bearer tok"},
duplicate_preset = await _webui_mutate(
webui_client,
"settings.model_configuration.create",
{
"label": "Fast writing",
"provider": "openai",
"model": "openai/gpt-4.1-mini",
},
)
assert duplicate_preset.status_code == 409
search_updated = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/web-search/update?provider=searxng"
"&base_url=https%3A%2F%2Fsearch.example.com"
"&max_results=8&timeout=45&use_jina_reader=false",
headers={"Authorization": "Bearer tok"},
search_updated = await _webui_mutate(
webui_client,
"settings.web_search.update",
{
"provider": "searxng",
"base_url": "https://search.example.com",
"max_results": 8,
"timeout": 45,
"use_jina_reader": False,
},
)
assert search_updated.status_code == 200
search_body = search_updated.json()
@@ -2814,10 +3292,13 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
assert search_body["web_search"]["max_results"] == 8
assert search_body["web"]["fetch"]["use_jina_reader"] is False
network_safety_updated = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/network-safety/update?webui_allow_local_service_access=false&webui_default_access_mode=full",
headers={"Authorization": "Bearer tok"},
network_safety_updated = await _webui_mutate(
webui_client,
"settings.network_safety.update",
{
"webui_allow_local_service_access": False,
"webui_default_access_mode": "full",
},
)
assert network_safety_updated.status_code == 200
network_safety_body = network_safety_updated.json()
@@ -2827,13 +3308,17 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
assert network_safety_body["advanced"]["webui_default_access_mode"] == "full"
assert network_safety_body["advanced"]["private_service_protection_enabled"] is True
image_updated = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/image-generation/update?enabled=true"
"&provider=openrouter&model=openai%2Fgpt-image-1"
"&default_aspect_ratio=16%3A9&default_image_size=2K"
"&max_images_per_turn=3",
headers={"Authorization": "Bearer tok"},
image_updated = await _webui_mutate(
webui_client,
"settings.image_generation.update",
{
"enabled": True,
"provider": "openrouter",
"model": "openai/gpt-image-1",
"default_aspect_ratio": "16:9",
"default_image_size": "2K",
"max_images_per_turn": 3,
},
)
assert image_updated.status_code == 200
image_body = image_updated.json()
@@ -2845,11 +3330,14 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
assert image_body["image_generation"]["default_image_size"] == "2K"
assert image_body["image_generation"]["max_images_per_turn"] == 3
image_provider_updated = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/provider/update?provider=openrouter"
"&api_key=sk-or-next&api_base=https%3A%2F%2Fopenrouter.ai%2Fapi%2Fv1",
headers={"Authorization": "Bearer tok"},
image_provider_updated = await _webui_mutate(
webui_client,
"settings.provider.update",
{
"provider": "openrouter",
"apiKey": "sk-or-next",
"apiBase": "https://openrouter.ai/api/v1",
},
)
assert image_provider_updated.status_code == 200
assert image_provider_updated.json()["requires_restart"] is True
@@ -2857,17 +3345,17 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
assert "sk-or-next" not in image_provider_updated.text
assert image_reload.await_count == 2
bad_web = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/web-search/update?provider=duckduckgo&max_results=99",
headers={"Authorization": "Bearer tok"},
bad_web = await _webui_mutate(
webui_client,
"settings.web_search.update",
{"provider": "duckduckgo", "max_results": 99},
)
assert bad_web.status_code == 400
bad_image = await _http_get(
"http://127.0.0.1:"
f"{port}/api/settings/image-generation/update?provider=missing",
headers={"Authorization": "Bearer tok"},
bad_image = await _webui_mutate(
webui_client,
"settings.image_generation.update",
{"provider": "missing"},
)
assert bad_image.status_code == 400
@@ -2904,6 +3392,8 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
assert saved.tools.image_generation.default_image_size == "2K"
assert saved.tools.image_generation.max_images_per_turn == 3
finally:
if webui_client is not None:
await webui_client.close()
await channel.stop()
await server_task
@@ -2936,11 +3426,17 @@ async def test_image_settings_hot_reload_without_restart(
channel.gateway.tokens.api_tokens["tok"] = time.monotonic() + 300
server_task = asyncio.create_task(channel.start())
await asyncio.sleep(0.3)
webui_client = None
try:
response = await _http_get(
f"http://127.0.0.1:{port}/api/settings/image-generation/update"
"?enabled=true&provider=openrouter&model=openai%2Fgpt-image-1",
headers={"Authorization": "Bearer tok"},
webui_token = channel.gateway.tokens.issue_token(300, audience="webui")
webui_client = await websockets.connect(
f"ws://127.0.0.1:{port}/ws?token={webui_token}&client_id=image-reload-test"
)
assert json.loads(await webui_client.recv())["event"] == "ready"
response = await _webui_mutate(
webui_client,
"settings.image_generation.update",
{"enabled": True, "provider": "openrouter", "model": "openai/gpt-image-1"},
)
assert response.status_code == 200
@@ -2948,6 +3444,8 @@ async def test_image_settings_hot_reload_without_restart(
assert response.json()["restart_required_sections"] == []
image_reload.assert_awaited_once_with(bus)
finally:
if webui_client is not None:
await webui_client.close()
await channel.stop()
await server_task
@@ -2979,17 +3477,25 @@ async def test_image_settings_fall_back_to_restart_when_hot_reload_fails(
channel.gateway.tokens.api_tokens["tok"] = time.monotonic() + 300
server_task = asyncio.create_task(channel.start())
await asyncio.sleep(0.3)
webui_client = None
try:
response = await _http_get(
f"http://127.0.0.1:{port}/api/settings/image-generation/update"
"?enabled=true&provider=openrouter&model=openai%2Fgpt-image-1",
headers={"Authorization": "Bearer tok"},
webui_token = channel.gateway.tokens.issue_token(300, audience="webui")
webui_client = await websockets.connect(
f"ws://127.0.0.1:{port}/ws?token={webui_token}&client_id=image-fallback-test"
)
assert json.loads(await webui_client.recv())["event"] == "ready"
response = await _webui_mutate(
webui_client,
"settings.image_generation.update",
{"enabled": True, "provider": "openrouter", "model": "openai/gpt-image-1"},
)
assert response.status_code == 200
assert response.json()["requires_restart"] is True
assert response.json()["restart_required_sections"] == ["image"]
finally:
if webui_client is not None:
await webui_client.close()
await channel.stop()
await server_task
File diff suppressed because it is too large Load Diff
@@ -1,11 +1,8 @@
"""Tests for the signed ``/api/media/<sig>/<payload>`` route and its replay
integration on ``/api/sessions/<key>/messages``.
"""Tests for the signed ``/api/media/<sig>/<payload>`` route and WebUI replay.
The route is the return path for images attached to persisted user turns:
:meth:`WebSocketChannel.gateway.media.sign_media_path` mints URLs during session reads,
and :meth:`GatewayHTTPHandler._handle_media_fetch` serves the bytes back.
These tests cover the two halves end-to-end plus the adversarial edges
(bad signatures, ``..`` traversal, non-existent files, non-image types).
The route is the return path for local media rendered by the WebUI. These tests
cover URL signing and serving end-to-end plus the adversarial edges (bad
signatures, ``..`` traversal, non-existent files, non-image types).
"""
from __future__ import annotations
@@ -20,11 +17,12 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from nanobot.channels.websocket.runtime import WebSocketChannel, WebSocketConfig
from nanobot.session.manager import Session, SessionManager
from nanobot.session.manager import SessionManager
from nanobot.webui.gateway_services import build_gateway_services
from nanobot.webui.media_api import (
b64url_decode,
b64url_encode,
sign_media_path,
)
from .ws_test_client import InProcessHttpChannel
@@ -87,8 +85,16 @@ def _fake_media_dir(root: Path):
return inner
def _sign_media_path(channel: WebSocketChannel, path: Path) -> str | None:
return sign_media_path(
path,
secret=channel.gateway.media.secret,
media_dir=channel.gateway.media._media_dir,
)
# ---------------------------------------------------------------------------
# gateway.media.sign_media_path: the URL minter
# media_api.sign_media_path: the URL minter
# ---------------------------------------------------------------------------
@@ -108,10 +114,10 @@ def test_sign_media_path_rejects_paths_outside_media_root(
media.mkdir()
channel = _ch(bus, port=0)
with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media):
assert channel.gateway.media.sign_media_path(outside) is None
assert _sign_media_path(channel, outside) is None
# Traversal via the media root is also rejected — the resolve() step
# normalises ``..`` out before the relative_to check.
assert channel.gateway.media.sign_media_path(media / ".." / "secrets" / "cred.txt") is None
assert _sign_media_path(channel, media / ".." / "secrets" / "cred.txt") is None
def test_sign_media_path_round_trips_via_hmac(
@@ -123,7 +129,7 @@ def test_sign_media_path_round_trips_via_hmac(
(media / "a.png").write_bytes(_PNG_BYTES)
channel = _ch(bus, port=0)
with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media):
url = channel.gateway.media.sign_media_path(media / "a.png")
url = _sign_media_path(channel, media / "a.png")
assert url is not None
assert url.startswith("/api/media/")
sig, payload = url[len("/api/media/"):].split("/", 1)
@@ -238,7 +244,7 @@ async def test_media_route_serves_signed_file(
channel = _ch(bus, port=29920)
with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media):
url_path = channel.gateway.media.sign_media_path(target)
url_path = _sign_media_path(channel, target)
assert url_path is not None
server_task = asyncio.create_task(channel.start())
try:
@@ -270,7 +276,7 @@ async def test_media_route_serves_video_byte_ranges(
channel = _ch(bus, port=29927)
with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media):
url_path = channel.gateway.media.sign_media_path(target)
url_path = _sign_media_path(channel, target)
assert url_path is not None
server_task = asyncio.create_task(channel.start())
try:
@@ -301,7 +307,7 @@ async def test_media_route_serves_suffix_video_byte_ranges(
channel = _ch(bus, port=29928)
with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media):
url_path = channel.gateway.media.sign_media_path(target)
url_path = _sign_media_path(channel, target)
assert url_path is not None
server_task = asyncio.create_task(channel.start())
try:
@@ -329,7 +335,7 @@ async def test_media_route_rejects_unsatisfiable_byte_range(
channel = _ch(bus, port=29929)
with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media):
url_path = channel.gateway.media.sign_media_path(target)
url_path = _sign_media_path(channel, target)
assert url_path is not None
server_task = asyncio.create_task(channel.start())
try:
@@ -361,7 +367,7 @@ async def test_media_route_rejects_bad_signature(
channel = _ch(bus, port=29921)
with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media):
good = channel.gateway.media.sign_media_path(media / "f.png")
good = _sign_media_path(channel, media / "f.png")
assert good is not None
_, payload = good[len("/api/media/"):].split("/", 1)
# Forge a sig with a *different* secret.
@@ -426,7 +432,7 @@ async def test_media_route_404s_missing_file(
channel = _ch(bus, port=29923)
with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media):
url_path = channel.gateway.media.sign_media_path(target)
url_path = _sign_media_path(channel, target)
assert url_path is not None
target.unlink() # the file vanishes between signing and fetching
server_task = asyncio.create_task(channel.start())
@@ -483,7 +489,7 @@ async def test_media_route_serves_svg_with_strict_csp(
channel = _ch(bus, port=29928)
with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media):
url_path = channel.gateway.media.sign_media_path(target)
url_path = _sign_media_path(channel, target)
assert url_path is not None
server_task = asyncio.create_task(channel.start())
try:
@@ -497,91 +503,3 @@ async def test_media_route_serves_svg_with_strict_csp(
assert resp.headers.get("x-content-type-options") == "nosniff"
assert "default-src 'none'" in resp.headers.get("content-security-policy", "")
assert "sandbox" in resp.headers.get("content-security-policy", "")
# ---------------------------------------------------------------------------
# /api/sessions/<key>/messages: media_urls hydration on session read
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_session_messages_exposes_signed_media_urls(
bus: MagicMock, tmp_path: Path
) -> None:
"""The read path must map persisted ``media`` paths onto signed URLs
and strip the raw path the client never learns the server's layout."""
media = tmp_path / "media"
media.mkdir()
img = media / "u.png"
img.write_bytes(_PNG_BYTES)
sm = SessionManager(tmp_path / "ws_state")
sess = Session(key="websocket:media-hydrate")
sess.add_message("user", "look at this", media=[str(img)])
sess.add_message("assistant", "nice")
sm.save(sess)
channel = _ch(bus, session_manager=sm, port=29925)
with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media):
server_task = asyncio.create_task(channel.start())
try:
token = channel.gateway.tokens.issue_api_token(300)
auth = {"Authorization": f"Bearer {token}"}
resp = await _http_get(
"http://127.0.0.1:29925/api/sessions/websocket:media-hydrate/messages",
headers=auth,
)
body = resp.json()
# The signed URL round-trips end-to-end: fetching it yields the same bytes.
user_msg = next(m for m in body["messages"] if m["role"] == "user")
urls = user_msg["media_urls"]
assert isinstance(urls, list) and len(urls) == 1
assert urls[0]["name"] == "u.png"
assert urls[0]["url"].startswith("/api/media/")
# Raw paths must not leak to the wire.
assert "media" not in user_msg
# And the URL actually works.
fetched = await _http_get(f"http://127.0.0.1:29925{urls[0]['url']}")
assert fetched.status_code == 200
assert fetched.content == _PNG_BYTES
finally:
await channel.stop()
await server_task
@pytest.mark.asyncio
async def test_session_messages_skips_vanished_media(
bus: MagicMock, tmp_path: Path
) -> None:
"""Paths that no longer resolve inside the media root produce no URL —
the message is still delivered, just without the preview."""
media = tmp_path / "media"
media.mkdir()
sm = SessionManager(tmp_path / "ws_state")
sess = Session(key="websocket:vanished")
sess.add_message("user", "missing pic", media=[str(media / "absent.png")])
sm.save(sess)
channel = _ch(bus, session_manager=sm, port=29926)
with patch("nanobot.webui.media_gateway.get_media_dir", return_value=media):
server_task = asyncio.create_task(channel.start())
try:
token = channel.gateway.tokens.issue_api_token(300)
resp = await _http_get(
"http://127.0.0.1:29926/api/sessions/websocket:vanished/messages",
headers={"Authorization": f"Bearer {token}"},
)
user_msg = next(m for m in resp.json()["messages"] if m["role"] == "user")
# absent.png lives inside the media root so it *does* get a signed
# URL (we don't stat the file at signing time — that would slow
# the listing). Fetching the URL is where the 404 surfaces.
urls = user_msg.get("media_urls") or []
assert len(urls) == 1
fetched = await _http_get(f"http://127.0.0.1:29926{urls[0]['url']}")
assert fetched.status_code == 404
assert "media" not in user_msg
finally:
await channel.stop()
await server_task
+95 -9
View File
@@ -22,6 +22,7 @@ class WeixinConnectSession:
channel: WeixinChannel
current_poll_base_url: str
refresh_count: int
force: bool
created_wall: float
deadline: float
last_error: str | None = None
@@ -47,7 +48,10 @@ class WeixinConnectStore:
if not session_id:
raise ChannelConnectError("missing WeChat connect session")
if action == "poll":
return await self.poll(session_id)
return await self.poll(
session_id,
verify_code=(query_first(query, "verify_code") or "").strip(),
)
if action == "cancel":
return await self.cancel(session_id)
raise ChannelConnectError(f"unsupported WeChat connect action: {action}", status=404)
@@ -69,7 +73,7 @@ class WeixinConnectStore:
channel.connect_open_client()
try:
qrcode_id, qr_url = await channel.connect_fetch_qr_code()
qrcode_id, qr_url = await channel.connect_fetch_qr_code(force=force)
except Exception as exc:
await self._close_channel(channel)
raise ChannelConnectError(
@@ -86,12 +90,13 @@ class WeixinConnectStore:
channel=channel,
current_poll_base_url=channel.connect_base_url,
refresh_count=0,
force=force,
created_wall=now_wall,
deadline=time.monotonic() + 600,
)
return self._start_payload(self._sessions[session_id])
async def poll(self, session_id: str) -> dict[str, Any]:
async def poll(self, session_id: str, *, verify_code: str = "") -> dict[str, Any]:
await self._cleanup()
session = self._sessions.get(session_id)
if session is None:
@@ -105,6 +110,7 @@ class WeixinConnectStore:
status_data = await session.channel.connect_poll_qr_code(
base_url=session.current_poll_base_url,
qrcode_id=session.qrcode_id,
verify_code=verify_code,
)
except Exception as exc:
if session.channel.connect_poll_error_is_retryable(exc):
@@ -120,6 +126,8 @@ class WeixinConnectStore:
status_payload = status_data
status = status_payload.get("status", "")
from nanobot.channels.weixin.runtime import MAX_QR_REFRESH_COUNT
if status == "confirmed":
if self._sessions.get(session_id) is not session:
return {
@@ -157,9 +165,77 @@ class WeixinConnectStore:
)
return self._pending_payload(session)
if status == "expired":
from nanobot.channels.weixin.runtime import MAX_QR_REFRESH_COUNT
if status == "need_verifycode":
return self._pending_payload(
session,
challenge="verify_code",
message=(
"That verification code did not match. Enter the new number shown in WeChat."
if verify_code
else "Enter the number shown in WeChat to continue."
),
verification_failed=bool(verify_code),
)
if status == "verify_code_blocked":
session.refresh_count += 1
if session.refresh_count > MAX_QR_REFRESH_COUNT:
self._sessions.pop(session_id, None)
await self._close_channel(session.channel)
return {
"session_id": session_id,
"status": "failed",
"message": "Too many incorrect verification attempts. Try again later.",
}
try:
session.qrcode_id, session.qr_url = (
await session.channel.connect_fetch_qr_code(force=session.force)
)
except Exception as exc:
self._sessions.pop(session_id, None)
await self._close_channel(session.channel)
return {
"session_id": session_id,
"status": "failed",
"message": f"Could not refresh WeChat QR code: {exc}",
}
session.current_poll_base_url = session.channel.connect_base_url
return self._pending_payload(
session,
message="Verification was blocked. Scan the refreshed QR code to try again.",
)
if status == "binded_redirect":
if session.force:
self._sessions.pop(session_id, None)
await self._close_channel(session.channel)
return {
"session_id": session_id,
"status": "failed",
"message": (
"Unable to complete a new WeChat login. "
"Start again and scan with the account you want to connect."
),
}
if not session.channel.connect_load_state():
self._sessions.pop(session_id, None)
await self._close_channel(session.channel)
return {
"session_id": session_id,
"status": "failed",
"message": (
"WeChat reports an existing binding, but no local credentials were found."
),
}
self._sessions.pop(session_id, None)
await self._close_channel(session.channel)
return {
"session_id": session_id,
"status": "succeeded",
"message": "WeChat is already connected to this nanobot instance.",
}
if status == "expired":
session.refresh_count += 1
if session.refresh_count > MAX_QR_REFRESH_COUNT:
self._sessions.pop(session_id, None)
@@ -171,7 +247,7 @@ class WeixinConnectStore:
}
try:
session.qrcode_id, session.qr_url = (
await session.channel.connect_fetch_qr_code()
await session.channel.connect_fetch_qr_code(force=session.force)
)
except Exception as exc:
self._sessions.pop(session_id, None)
@@ -238,15 +314,25 @@ class WeixinConnectStore:
}
@staticmethod
def _pending_payload(session: WeixinConnectSession) -> dict[str, Any]:
return {
def _pending_payload(
session: WeixinConnectSession,
*,
challenge: str = "",
message: str = "Waiting for WeChat scan.",
verification_failed: bool = False,
) -> dict[str, Any]:
payload: dict[str, Any] = {
"session_id": session.id,
"status": "pending",
"qr_url": session.qr_url,
"interval_ms": 2000,
"expires_at_ms": int((session.created_wall + 600) * 1000),
"message": "Waiting for WeChat scan.",
"message": message,
}
if challenge:
payload["challenge"] = challenge
payload["verification_failed"] = verification_failed
return payload
__all__ = ["WeixinConnectStore"]
+14
View File
@@ -10,6 +10,20 @@ SETUP_SPEC = ChannelSetupSpec(
fields={
"token": field("secret"),
"allowFrom": field("list"),
"baseUrl": field(default="https://ilinkai.weixin.qq.com"),
"cdnBaseUrl": field(default="https://novac2c.cdn.weixin.qq.com/c2c"),
"routeTag": field(),
"stateDir": field(),
"pollTimeout": field("int", default=35),
"sendProgress": field("bool", default=False),
"sendToolHints": field("bool", default=False),
"replyProgressMessages": field("bool", default=False),
"replyProgressMaxMessages": field("int", default=2),
"contextMessageBudget": field("int", default=8),
"streaming": field("bool", default=True),
"blockStreaming": field("bool", default=False),
"blockStreamingMinChars": field("int", default=1200),
"blockStreamingMaxMessages": field("int", default=3),
},
required=(required("token"),),
official_url="https://weixin.qq.com/",
File diff suppressed because it is too large Load Diff
+160 -4
View File
@@ -25,7 +25,9 @@ async def test_weixin_connect_store_saves_confirmed_qr_login(
)
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
async def fake_fetch_qr_code(self: WeixinChannel) -> tuple[str, str]:
async def fake_fetch_qr_code(
self: WeixinChannel, **_kwargs: Any
) -> tuple[str, str]:
return "qr-1", "https://qr.example/1"
async def fake_api_get_with_base(
@@ -86,14 +88,31 @@ async def test_weixin_reconnect_keeps_existing_account_until_scan_succeeds(
)
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
async def fake_fetch_qr_code(self: WeixinChannel) -> tuple[str, str]:
return "qr-reconnect", "https://qr.example/reconnect"
observed_force: list[bool] = []
async def fake_fetch_qr_code(
self: WeixinChannel,
*,
force: bool = False,
) -> tuple[str, str]:
observed_force.append(force)
return f"qr-reconnect-{len(observed_force)}", "https://qr.example/reconnect"
async def fake_api_get_with_base(
self: WeixinChannel,
**_kwargs: Any,
) -> dict[str, str]:
return {"status": "expired"}
monkeypatch.setattr(WeixinChannel, "_fetch_qr_code", fake_fetch_qr_code)
monkeypatch.setattr(WeixinChannel, "_api_get_with_base", fake_api_get_with_base)
store = WeixinConnectStore()
started = await store.start(force=True)
refreshed = await store.poll(started["session_id"])
assert refreshed["status"] == "pending"
assert observed_force == [True, True]
assert json.loads(state_file.read_text(encoding="utf-8")) == existing
cancelled = await store.cancel(started["session_id"])
assert cancelled["status"] == "cancelled"
@@ -116,7 +135,9 @@ async def test_weixin_cancel_wins_over_inflight_confirmation(
poll_started = asyncio.Event()
release_poll = asyncio.Event()
async def fake_fetch_qr_code(self: WeixinChannel) -> tuple[str, str]:
async def fake_fetch_qr_code(
self: WeixinChannel, **_kwargs: Any
) -> tuple[str, str]:
return "qr-cancel", "https://qr.example/cancel"
async def fake_api_get_with_base(
@@ -147,3 +168,138 @@ async def test_weixin_cancel_wins_over_inflight_confirmation(
assert cancelled["status"] == "cancelled"
assert completed["status"] == "cancelled"
assert not (state_dir / "account.json").exists()
@pytest.mark.asyncio
async def test_weixin_connect_store_handles_verification_code(
tmp_path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
state_dir = tmp_path / "weixin-state"
config_path = tmp_path / "config.json"
save_config(
Config.model_validate({"channels": {"weixin": {"stateDir": str(state_dir)}}}),
config_path,
)
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
async def fake_fetch_qr_code(
self: WeixinChannel, **_kwargs: Any
) -> tuple[str, str]:
return "qr-verify", "https://qr.example/verify"
responses = [
{"status": "need_verifycode"},
{
"status": "confirmed",
"bot_token": "verified-token",
"ilink_user_id": "wx-user",
},
]
async def fake_api_get_with_base(
self: WeixinChannel,
*,
params: dict[str, Any],
**_kwargs: Any,
) -> dict[str, str]:
if len(responses) == 1:
assert params == {"qrcode": "qr-verify", "verify_code": "1234"}
return responses.pop(0)
monkeypatch.setattr(WeixinChannel, "_fetch_qr_code", fake_fetch_qr_code)
monkeypatch.setattr(WeixinChannel, "_api_get_with_base", fake_api_get_with_base)
store = WeixinConnectStore()
started = await store.start()
challenged = await store.poll(started["session_id"])
completed = await store.handle(
"poll",
{
"session_id": [started["session_id"]],
"verify_code": ["1234"],
},
)
assert challenged["status"] == "pending"
assert challenged["challenge"] == "verify_code"
assert completed["status"] == "succeeded"
@pytest.mark.asyncio
async def test_weixin_connect_store_rejects_existing_binding_during_forced_login(
tmp_path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
state_dir = tmp_path / "weixin-state"
state_dir.mkdir()
(state_dir / "account.json").write_text(
json.dumps({"token": "working-token"}),
encoding="utf-8",
)
config_path = tmp_path / "config.json"
save_config(
Config.model_validate({"channels": {"weixin": {"stateDir": str(state_dir)}}}),
config_path,
)
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
async def fake_fetch_qr_code(
self: WeixinChannel,
*,
force: bool = False,
) -> tuple[str, str]:
assert force is True
return "qr-existing", "https://qr.example/existing"
async def fake_api_get_with_base(
self: WeixinChannel,
**_kwargs: Any,
) -> dict[str, str]:
return {"status": "binded_redirect"}
monkeypatch.setattr(WeixinChannel, "_fetch_qr_code", fake_fetch_qr_code)
monkeypatch.setattr(WeixinChannel, "_api_get_with_base", fake_api_get_with_base)
store = WeixinConnectStore()
started = await store.start(force=True)
completed = await store.poll(started["session_id"])
assert completed["status"] == "failed"
assert "new WeChat login" in completed["message"]
assert json.loads((state_dir / "account.json").read_text())["token"] == "working-token"
@pytest.mark.asyncio
async def test_weixin_connect_store_rejects_existing_binding_without_local_credentials(
tmp_path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
state_dir = tmp_path / "weixin-state"
config_path = tmp_path / "config.json"
save_config(
Config.model_validate({"channels": {"weixin": {"stateDir": str(state_dir)}}}),
config_path,
)
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
async def fake_fetch_qr_code(
self: WeixinChannel, **_kwargs: Any
) -> tuple[str, str]:
return "qr-missing", "https://qr.example/missing"
async def fake_api_get_with_base(
self: WeixinChannel,
**_kwargs: Any,
) -> dict[str, str]:
return {"status": "binded_redirect"}
monkeypatch.setattr(WeixinChannel, "_fetch_qr_code", fake_fetch_qr_code)
monkeypatch.setattr(WeixinChannel, "_api_get_with_base", fake_api_get_with_base)
store = WeixinConnectStore()
started = await store.start(force=False)
completed = await store.poll(started["session_id"])
assert completed["status"] == "failed"
assert "no local credentials" in completed["message"]
@@ -17,6 +17,7 @@ from nanobot.channels.weixin.runtime import (
ITEM_TEXT,
MESSAGE_TYPE_BOT,
WEIXIN_CHANNEL_VERSION,
WeixinAuthError,
WeixinChannel,
WeixinConfig,
_decrypt_aes_ecb,
@@ -67,11 +68,11 @@ def test_make_headers_includes_route_tag_when_configured() -> None:
assert headers["Authorization"] == "Bearer token"
assert headers["SKRouteTag"] == "123"
assert headers["iLink-App-Id"] == "bot"
assert headers["iLink-App-ClientVersion"] == str((2 << 16) | (1 << 8) | 1)
assert headers["iLink-App-ClientVersion"] == str((2 << 16) | (4 << 8) | 6)
def test_channel_version_matches_reference_plugin_version() -> None:
assert WEIXIN_CHANNEL_VERSION == "2.1.1"
assert WEIXIN_CHANNEL_VERSION == "2.4.6"
def test_save_and_load_state_persists_context_tokens(tmp_path) -> None:
@@ -159,6 +160,29 @@ def test_save_state_persists_explicit_config_token_over_stale_state(tmp_path) ->
assert saved["get_updates_buf"] == "current-cursor"
def test_save_state_preserves_qr_replacement_of_configured_token(tmp_path) -> None:
config = WeixinConfig(
enabled=True,
allow_from=["*"],
token="configured-token",
state_dir=str(tmp_path),
)
old_runtime = WeixinChannel(config, MessageBus())
old_runtime._token = "configured-token"
replacement = WeixinChannel(config, MessageBus())
replacement.connect_commit_account(
token="replacement-token",
base_url="https://new.example",
)
old_runtime._save_state()
saved = json.loads((tmp_path / "account.json").read_text())
assert saved["token"] == "replacement-token"
assert saved["base_url"] == "https://new.example"
def test_save_state_with_empty_runtime_token_preserves_persisted_account(tmp_path) -> None:
channel = WeixinChannel(
WeixinConfig(enabled=True, allow_from=["*"], state_dir=str(tmp_path)),
@@ -172,6 +196,86 @@ def test_save_state_with_empty_runtime_token_preserves_persisted_account(tmp_pat
assert json.loads((tmp_path / "account.json").read_text()) == persisted
@pytest.mark.asyncio
async def test_login_force_ignores_persisted_account_through_qr_flow(tmp_path) -> None:
persisted = {
"token": "persisted-token",
"get_updates_buf": "persisted-cursor",
"context_tokens": {"wx-user": "ctx-persisted"},
"typing_tickets": {"wx-user": {"ticket": "ticket-persisted"}},
"base_url": "https://persisted.example",
}
channel = WeixinChannel(
WeixinConfig(
enabled=True,
allow_from=["*"],
token="configured-token",
state_dir=str(tmp_path),
),
MessageBus(),
)
(tmp_path / "account.json").write_text(
json.dumps(persisted),
encoding="utf-8",
)
channel._print_qr_code = lambda _url: None
channel._api_post = AsyncMock(
side_effect=[
{"qrcode": "qr-1", "qrcode_img_content": "url-1"},
{"qrcode": "qr-2", "qrcode_img_content": "url-2"},
]
)
channel._api_get_with_base = AsyncMock(
side_effect=[
{"status": "expired"},
{"status": "binded_redirect"},
]
)
ok = await channel.login(force=True)
assert ok is False
assert [call.args[1]["local_token_list"] for call in channel._api_post.await_args_list] == [
[],
[],
]
assert channel._token == ""
assert channel._get_updates_buf == ""
assert channel._context_tokens == {}
assert channel._typing_tickets == {}
assert channel.config.base_url == "https://ilinkai.weixin.qq.com"
assert json.loads((tmp_path / "account.json").read_text()) == persisted
@pytest.mark.asyncio
async def test_login_without_force_reuses_persisted_account(tmp_path) -> None:
channel = WeixinChannel(
WeixinConfig(enabled=True, allow_from=["*"], state_dir=str(tmp_path)),
MessageBus(),
)
(tmp_path / "account.json").write_text(
json.dumps(
{
"token": "persisted-token",
"get_updates_buf": "persisted-cursor",
"context_tokens": {"wx-user": "ctx-persisted"},
"base_url": "https://persisted.example",
}
),
encoding="utf-8",
)
channel._qr_login = AsyncMock(return_value=False)
ok = await channel.login(force=False)
assert ok is True
channel._qr_login.assert_not_awaited()
assert channel._token == "persisted-token"
assert channel._get_updates_buf == "persisted-cursor"
assert channel._context_tokens == {"wx-user": "ctx-persisted"}
assert channel.config.base_url == "https://persisted.example"
@pytest.mark.asyncio
async def test_process_message_deduplicates_inbound_ids() -> None:
channel, bus = _make_channel()
@@ -442,15 +546,15 @@ async def test_send_without_context_token_raises() -> None:
@pytest.mark.asyncio
async def test_send_raises_when_session_is_paused() -> None:
async def test_send_raises_when_authentication_is_required() -> None:
channel, _bus = _make_channel()
channel._client = object()
channel._token = "token"
channel._context_tokens["wx-user"] = "ctx-2"
channel._pause_session(60)
channel._auth_required = True
channel._send_text = AsyncMock()
with pytest.raises(RuntimeError, match="session paused"):
with pytest.raises(WeixinAuthError, match="bot token is stale"):
await channel.send(
type("Msg", (), {"chat_id": "wx-user", "content": "pong", "media": [], "metadata": {}})()
)
@@ -525,20 +629,21 @@ async def test_send_still_sends_text_when_typing_ticket_missing() -> None:
@pytest.mark.asyncio
async def test_poll_once_pauses_session_on_expired_errcode() -> None:
async def test_poll_once_requires_login_on_stale_token() -> None:
channel, _bus = _make_channel()
channel._client = SimpleNamespace(timeout=None)
channel._token = "token"
channel._api_post = AsyncMock(return_value={"ret": 0, "errcode": -14, "errmsg": "expired"})
await channel._poll_once()
with pytest.raises(WeixinAuthError, match="no replacement credentials"):
await channel._poll_once()
assert channel._session_pause_remaining_s() > 0
assert channel._auth_required is True
@pytest.mark.asyncio
async def test_poll_once_reloads_refreshed_state_after_session_pause(
tmp_path, monkeypatch: pytest.MonkeyPatch
async def test_poll_once_reloads_refreshed_state_after_stale_token(
tmp_path,
) -> None:
channel = WeixinChannel(
WeixinConfig(enabled=True, allow_from=["*"], state_dir=str(tmp_path)),
@@ -550,8 +655,13 @@ async def test_poll_once_reloads_refreshed_state_after_session_pause(
json.dumps({"token": "new-token", "base_url": "https://new.example"}),
encoding="utf-8",
)
channel._session_pause_until = time.time() + 10
monkeypatch.setattr(weixin_mod.asyncio, "sleep", AsyncMock())
channel._client = object()
channel._api_post = AsyncMock(
side_effect=[
{"ret": 0, "errcode": -14, "errmsg": "stale"},
{"ret": 0},
]
)
await channel._poll_once()
@@ -560,8 +670,8 @@ async def test_poll_once_reloads_refreshed_state_after_session_pause(
@pytest.mark.asyncio
async def test_poll_once_keeps_explicit_token_after_session_pause(
tmp_path, monkeypatch: pytest.MonkeyPatch
async def test_poll_once_keeps_explicit_token_and_requires_login(
tmp_path,
) -> None:
channel = WeixinChannel(
WeixinConfig(
@@ -577,24 +687,132 @@ async def test_poll_once_keeps_explicit_token_after_session_pause(
json.dumps({"token": "stale-token", "base_url": "https://stale.example"}),
encoding="utf-8",
)
channel._session_pause_until = time.time() + 10
monkeypatch.setattr(weixin_mod.asyncio, "sleep", AsyncMock())
channel._client = object()
channel._api_post = AsyncMock(
return_value={"ret": 0, "errcode": -14, "errmsg": "stale"}
)
await channel._poll_once()
with pytest.raises(WeixinAuthError, match="no replacement credentials"):
await channel._poll_once()
assert channel._token == "configured-token"
assert channel.config.base_url == "https://ilinkai.weixin.qq.com"
@pytest.mark.asyncio
async def test_poll_once_loads_qr_replacement_for_configured_token(tmp_path) -> None:
config = WeixinConfig(
enabled=True,
allow_from=["*"],
token="configured-token",
state_dir=str(tmp_path),
)
replacement = WeixinChannel(config, MessageBus())
replacement.connect_commit_account(
token="replacement-token",
base_url="https://new.example",
)
channel = WeixinChannel(config, MessageBus())
channel._token = "configured-token"
channel._client = object()
channel._api_post = AsyncMock(
side_effect=[
{"ret": 0, "errcode": -14, "errmsg": "stale"},
{"ret": 0},
]
)
await channel._poll_once()
assert channel._token == "replacement-token"
assert channel.config.base_url == "https://new.example"
@pytest.mark.asyncio
async def test_start_uses_qr_replacement_for_configured_token(tmp_path) -> None:
config = WeixinConfig(
enabled=True,
allow_from=["*"],
token="configured-token",
state_dir=str(tmp_path),
)
connector = WeixinChannel(config, MessageBus())
connector.connect_commit_account(
token="replacement-token",
base_url="https://new.example",
)
channel = WeixinChannel(config, MessageBus())
observed_tokens: list[str] = []
async def stop_after_first_poll() -> None:
observed_tokens.append(channel._token)
channel._running = False
channel._notify_lifecycle = AsyncMock() # type: ignore[method-assign]
channel._poll_once = stop_after_first_poll # type: ignore[method-assign]
await channel.start()
await channel.stop()
assert observed_tokens == ["replacement-token"]
assert channel.config.base_url == "https://new.example"
@pytest.mark.asyncio
async def test_manager_surfaces_actionable_weixin_auth_error_without_traceback(
tmp_path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
from nanobot.channels import manager as manager_mod
channel = WeixinChannel(
WeixinConfig(enabled=True, allow_from=["*"], state_dir=str(tmp_path)),
MessageBus(),
)
channel.start = AsyncMock( # type: ignore[method-assign]
side_effect=WeixinAuthError(
"getupdates",
errcode=-14,
errmsg="stale",
)
)
errors: list[str] = []
tracebacks: list[str] = []
monkeypatch.setattr(
manager_mod.logger,
"error",
lambda message, *args: errors.append(message.format(*args)),
)
monkeypatch.setattr(
manager_mod.logger,
"exception",
lambda message, *args: tracebacks.append(message.format(*args)),
)
manager = manager_mod.ChannelManager.__new__(manager_mod.ChannelManager)
manager._channel_errors = {}
await manager._start_channel("weixin", channel)
assert manager._channel_errors["weixin"] == (
"WeChat login expired. Scan again to reconnect."
)
assert errors == [
"Failed to start channel weixin: WeChat login expired. Scan again to reconnect."
]
assert tracebacks == []
@pytest.mark.asyncio
async def test_qr_login_refreshes_expired_qr_and_then_succeeds(
no_qr_poll_delay,
) -> None:
channel, _bus = _make_channel()
channel._running = True
channel._save_state = lambda: None
channel._save_state = lambda **_kwargs: None
channel._print_qr_code = lambda url: None
channel._api_get = AsyncMock(
channel._api_post = AsyncMock(
side_effect=[
{"qrcode": "qr-1", "qrcode_img_content": "url-1"},
{"qrcode": "qr-2", "qrcode_img_content": "url-2"},
@@ -627,7 +845,7 @@ async def test_qr_login_returns_false_after_too_many_expired_qr_codes(
channel, _bus = _make_channel()
channel._running = True
channel._print_qr_code = lambda url: None
channel._api_get = AsyncMock(
channel._api_post = AsyncMock(
side_effect=[
{"qrcode": "qr-1", "qrcode_img_content": "url-1"},
{"qrcode": "qr-2", "qrcode_img_content": "url-2"},
@@ -655,7 +873,7 @@ async def test_qr_login_switches_polling_base_url_on_redirect_status(
) -> None:
channel, _bus = _make_channel()
channel._running = True
channel._save_state = lambda: None
channel._save_state = lambda **_kwargs: None
channel._print_qr_code = lambda url: None
channel._fetch_qr_code = AsyncMock(return_value=("qr-1", "url-1"))
@@ -689,7 +907,7 @@ async def test_qr_login_redirect_without_host_keeps_current_polling_base_url(
) -> None:
channel, _bus = _make_channel()
channel._running = True
channel._save_state = lambda: None
channel._save_state = lambda **_kwargs: None
channel._print_qr_code = lambda url: None
channel._fetch_qr_code = AsyncMock(return_value=("qr-1", "url-1"))
@@ -723,7 +941,7 @@ async def test_qr_login_resets_redirect_base_url_after_qr_refresh(
) -> None:
channel, _bus = _make_channel()
channel._running = True
channel._save_state = lambda: None
channel._save_state = lambda **_kwargs: None
channel._print_qr_code = lambda url: None
channel._fetch_qr_code = AsyncMock(side_effect=[("qr-1", "url-1"), ("qr-2", "url-2")])
@@ -1015,7 +1233,7 @@ async def test_qr_login_treats_temporary_connect_error_as_wait_and_recovers(
) -> None:
channel, _bus = _make_channel()
channel._running = True
channel._save_state = lambda: None
channel._save_state = lambda **_kwargs: None
channel._print_qr_code = lambda url: None
channel._fetch_qr_code = AsyncMock(return_value=("qr-1", "url-1"))
@@ -1045,7 +1263,7 @@ async def test_qr_login_treats_5xx_gateway_response_error_as_wait_and_recovers(
) -> None:
channel, _bus = _make_channel()
channel._running = True
channel._save_state = lambda: None
channel._save_state = lambda **_kwargs: None
channel._print_qr_code = lambda url: None
channel._fetch_qr_code = AsyncMock(return_value=("qr-1", "url-1"))
@@ -1080,6 +1298,32 @@ def test_decrypt_aes_ecb_strips_valid_pkcs7_padding() -> None:
assert decrypted == plaintext
def test_missing_aes_dependency_recommends_weixin_plugin(monkeypatch) -> None:
real_import = __import__
def fake_import(name, *args, **kwargs):
if name.startswith(("Crypto", "cryptography")):
raise ImportError("missing AES dependency")
return real_import(name, *args, **kwargs)
warnings: list[str] = []
monkeypatch.setattr("builtins.__import__", fake_import)
monkeypatch.setattr(
weixin_mod.logger,
"warning",
lambda message, *args: warnings.append(message.format(*args)),
)
key_b64 = "MDEyMzQ1Njc4OWFiY2RlZg=="
data = b"unencrypted media"
assert _encrypt_aes_ecb(data, key_b64) == data
assert _decrypt_aes_ecb(data, key_b64) == data
assert warnings == [
"Cannot encrypt media. Run `nanobot plugins enable weixin` to install WeChat support.",
"Cannot decrypt media. Run `nanobot plugins enable weixin` to install WeChat support.",
]
class _DummyDownloadResponse:
def __init__(self, content: bytes, status_code: int = 200) -> None:
self.content = content
@@ -1412,7 +1656,7 @@ async def test_send_text_raises_on_api_error() -> None:
return_value={"errcode": -14, "errmsg": "session expired"}
)
with pytest.raises(RuntimeError, match="WeChat send text error.*-14"):
with pytest.raises(WeixinAuthError, match="WeChat sendmessage failed.*errcode=-14"):
await channel._send_text("wx-user", "hello", "ctx-expired")
channel._api_post.assert_awaited_once()
@@ -1445,7 +1689,7 @@ async def test_send_text_raises_on_nonzero_ret_even_when_errcode_zero() -> None:
return_value={"ret": -100, "errcode": 0, "errmsg": "internal error"}
)
with pytest.raises(RuntimeError, match="WeChat send text error.*ret=-100.*errcode=0"):
with pytest.raises(RuntimeError, match="WeChat sendmessage failed.*ret=-100.*errcode=0"):
await channel._send_text("wx-user", "hello", "ctx-ok")
channel._api_post.assert_awaited_once()
@@ -0,0 +1,441 @@
from __future__ import annotations
import asyncio
import json
import time
from unittest.mock import AsyncMock
import httpx
import pytest
from nanobot.bus.events import OutboundMessage
from nanobot.bus.outbound_events import ProgressEvent
from nanobot.bus.queue import MessageBus
from nanobot.channels.manager import ChannelManager
from nanobot.channels.weixin.manifest import SETUP_SPEC
from nanobot.channels.weixin.runtime import (
ITEM_TOOL_CALL_RESULT,
ITEM_TOOL_CALL_START,
WEIXIN_MAX_MESSAGE_LEN,
WeixinAPIError,
WeixinAuthError,
WeixinChannel,
WeixinConfig,
WeixinQuotaError,
sanitize_weixin_markdown,
split_weixin_message,
)
from nanobot.config.schema import Config
def _channel(**config: object) -> WeixinChannel:
return WeixinChannel(
WeixinConfig.model_validate(
{"enabled": True, "allowFrom": ["*"], **config}
),
MessageBus(),
)
def _ready_channel(**config: object) -> WeixinChannel:
channel = _channel(**config)
channel._client = object()
channel._token = "bot-token"
channel._context_tokens["wx-user"] = "ctx-1"
channel._context_token_at["wx-user"] = time.time()
channel._typing_tickets["wx-user"] = {
"ticket": "",
"next_fetch_at": time.time() + 3600,
}
return channel
def test_weixin_defaults_protect_context_quota() -> None:
config = WeixinConfig()
assert WEIXIN_MAX_MESSAGE_LEN == 1800
assert config.send_progress is False
assert config.send_tool_hints is False
assert config.reply_progress_messages is False
assert config.context_message_budget == 8
assert config.block_streaming is False
def test_weixin_webui_manifest_covers_runtime_configuration() -> None:
runtime_fields = set(WeixinConfig().model_dump(mode="json", by_alias=True))
assert set(SETUP_SPEC.fields) == runtime_fields - {"enabled"}
def test_reply_progress_opt_in_enables_progress_transport() -> None:
config = WeixinConfig(reply_progress_messages=True)
assert config.send_progress is True
assert config.send_tool_hints is True
@pytest.mark.parametrize(
("section", "send_progress", "send_tool_hints"),
[
({"enabled": True}, False, False),
({"enabled": True, "replyProgressMessages": True}, True, True),
({"enabled": True, "sendProgress": True, "sendToolHints": False}, True, False),
],
)
def test_channel_manager_preserves_weixin_quota_defaults(
section: dict[str, object],
send_progress: bool,
send_tool_hints: bool,
) -> None:
manager = ChannelManager.__new__(ChannelManager)
manager.config = Config.model_validate({"channels": {"weixin": section}})
manager.bus = MessageBus()
channel = manager._build_channel("weixin", WeixinChannel, section)
assert channel.send_progress is send_progress
assert channel.send_tool_hints is send_tool_hints
@pytest.mark.asyncio
async def test_channel_manager_does_not_retry_permanent_weixin_error(monkeypatch) -> None:
manager = ChannelManager.__new__(ChannelManager)
manager.config = Config.model_validate({"channels": {"sendMaxRetries": 3}})
manager.bus = MessageBus()
channel = _channel()
channel.send = AsyncMock(
side_effect=WeixinAPIError(
"sendmessage",
errcode=-1,
errmsg="business rejection",
retryable=False,
)
)
sleep = AsyncMock()
monkeypatch.setattr("nanobot.channels.manager.asyncio.sleep", sleep)
await manager._send_with_retry(
channel,
OutboundMessage(channel="weixin", chat_id="wx-user", content="test"),
)
channel.send.assert_awaited_once()
sleep.assert_not_awaited()
@pytest.mark.asyncio
async def test_weixin_http_clients_ignore_system_proxy(tmp_path, monkeypatch) -> None:
captured: list[dict[str, object]] = []
class FakeClient:
async def aclose(self) -> None:
return None
def make_client(**kwargs: object) -> FakeClient:
captured.append(kwargs)
return FakeClient()
monkeypatch.setattr("nanobot.channels.weixin.runtime.httpx.AsyncClient", make_client)
connect_channel = _channel(stateDir=str(tmp_path / "connect"))
connect_channel.connect_open_client()
await connect_channel.connect_close_client()
login_channel = _channel(stateDir=str(tmp_path / "login"))
login_channel._qr_login = AsyncMock(return_value=True)
assert await login_channel.login() is True
start_channel = _channel(token="configured-token", stateDir=str(tmp_path / "start"))
async def stop_after_poll() -> None:
start_channel._running = False
start_channel._notify_lifecycle = AsyncMock()
start_channel._poll_once = AsyncMock(side_effect=stop_after_poll)
await start_channel.start()
await start_channel.stop()
assert len(captured) == 3
assert all(kwargs["trust_env"] is False for kwargs in captured)
def test_markdown_sanitizer_preserves_code_and_escapes_bare_angles() -> None:
content = "before <tag> `x<y>`\n```python\na<b\n```\n![drop](https://x.test/a.png)"
sanitized = sanitize_weixin_markdown(content)
assert "before tag" in sanitized
assert "`x<y>`" in sanitized
assert "a<b" in sanitized
assert "![drop]" not in sanitized
def test_markdown_split_balances_fences_and_stays_within_limit() -> None:
chunks = split_weixin_message("```python\n" + ("x" * 4000) + "\n```")
assert len(chunks) >= 3
assert all(len(chunk) <= WEIXIN_MAX_MESSAGE_LEN for chunk in chunks)
assert all(chunk.count("```") % 2 == 0 for chunk in chunks)
@pytest.mark.asyncio
async def test_qr_fetch_posts_known_local_tokens(tmp_path) -> None:
state_dir = tmp_path / "weixin"
state_dir.mkdir()
(state_dir / "account.json").write_text(
json.dumps({"token": "persisted-token"}),
encoding="utf-8",
)
channel = _channel(stateDir=str(state_dir))
channel._api_post = AsyncMock(
return_value={"qrcode": "qr-1", "qrcode_img_content": "https://qr.test/1"}
)
assert await channel._fetch_qr_code() == ("qr-1", "https://qr.test/1")
channel._api_post.assert_awaited_once_with(
"ilink/bot/get_bot_qrcode?bot_type=3",
{"local_token_list": ["persisted-token"]},
auth=False,
include_base_info=False,
)
@pytest.mark.asyncio
async def test_qr_fetch_retries_without_rejected_local_tokens(tmp_path) -> None:
state_dir = tmp_path / "weixin"
state_dir.mkdir()
(state_dir / "account.json").write_text(
json.dumps({"token": "invalid-token"}),
encoding="utf-8",
)
channel = _channel(stateDir=str(state_dir))
channel._api_post = AsyncMock(
side_effect=[
{"ret": -3},
{"ret": 0, "qrcode": "qr-1", "qrcode_img_content": "https://qr.test/1"},
]
)
assert await channel._fetch_qr_code() == ("qr-1", "https://qr.test/1")
assert [call.args[1] for call in channel._api_post.await_args_list] == [
{"local_token_list": ["invalid-token"]},
{"local_token_list": []},
]
@pytest.mark.asyncio
async def test_qr_fetch_does_not_retry_invalid_request_without_local_tokens(tmp_path) -> None:
channel = _channel(stateDir=str(tmp_path / "weixin"))
channel._api_post = AsyncMock(return_value={"ret": -3})
with pytest.raises(WeixinAPIError, match="get_bot_qrcode failed.*ret=-3"):
await channel._fetch_qr_code()
channel._api_post.assert_awaited_once()
@pytest.mark.asyncio
async def test_lifecycle_notifications_are_best_effort() -> None:
channel = _ready_channel()
channel._api_post = AsyncMock(return_value={"ret": 0})
await channel._notify_lifecycle("start")
await channel._notify_lifecycle("stop")
assert [call.args[0] for call in channel._api_post.await_args_list] == [
"ilink/bot/msg/notifystart",
"ilink/bot/msg/notifystop",
]
def test_business_errors_have_explicit_retry_contracts() -> None:
channel = _channel()
with pytest.raises(WeixinQuotaError) as quota:
channel._raise_for_api_error("sendmessage", {"ret": -2})
with pytest.raises(WeixinAuthError) as auth:
channel._raise_for_api_error("getupdates", {"errcode": -14})
with pytest.raises(WeixinAPIError) as rejected:
channel._raise_for_api_error("sendmessage", {"ret": -100})
assert channel.should_retry_send_error(quota.value) is False
assert channel.should_retry_send_error(auth.value) is False
assert channel.should_retry_send_error(rejected.value) is False
assert channel.should_retry_send_error(httpx.ReadTimeout("slow")) is True
request = httpx.Request("POST", "https://ilinkai.weixin.qq.com/send")
for status_code in (408, 425, 429, 503):
response = httpx.Response(status_code, request=request)
error = httpx.HTTPStatusError(
"retryable response",
request=request,
response=response,
)
assert channel.should_retry_send_error(error) is True
rejected_response = httpx.Response(400, request=request)
rejected_http = httpx.HTTPStatusError(
"bad request",
request=request,
response=rejected_response,
)
assert channel.should_retry_send_error(rejected_http) is False
def test_error_classification_checks_ret_and_errcode_independently() -> None:
channel = _channel()
with pytest.raises(WeixinQuotaError):
channel._raise_for_api_error(
"sendmessage",
{"ret": -2, "errcode": -100},
)
with pytest.raises(WeixinAuthError):
channel._raise_for_api_error(
"getupdates",
{"ret": -14, "errcode": -100},
)
@pytest.mark.asyncio
async def test_stop_cancels_inflight_long_poll() -> None:
channel = _channel(token="configured-token")
poll_started = asyncio.Event()
poll_cancelled = asyncio.Event()
class FakeClient:
async def aclose(self) -> None:
return None
async def blocking_poll() -> None:
poll_started.set()
try:
await asyncio.Event().wait()
except asyncio.CancelledError:
poll_cancelled.set()
raise
channel._new_http_client = lambda _timeout: FakeClient() # type: ignore[method-assign]
channel._notify_lifecycle = AsyncMock()
channel._poll_once = blocking_poll # type: ignore[method-assign]
start_task = asyncio.create_task(channel.start())
await asyncio.wait_for(poll_started.wait(), timeout=1)
await asyncio.wait_for(channel.stop(), timeout=1)
await asyncio.wait_for(start_task, timeout=1)
assert poll_cancelled.is_set()
assert channel._poll_task is None
@pytest.mark.asyncio
async def test_retry_reuses_client_id_and_skips_completed_chunks() -> None:
channel = _ready_channel()
request = httpx.Request("POST", "https://ilinkai.weixin.qq.com/ilink/bot/sendmessage")
channel._api_post = AsyncMock(
side_effect=[
{"ret": 0},
httpx.ReadTimeout("ambiguous timeout", request=request),
{"ret": 0},
]
)
msg = OutboundMessage(
channel="weixin",
chat_id="wx-user",
content="x" * (WEIXIN_MAX_MESSAGE_LEN + 200),
)
with pytest.raises(httpx.ReadTimeout):
await channel.send(msg)
await channel.send(msg)
bodies = [call.args[1] for call in channel._api_post.await_args_list]
client_ids = [body["msg"]["client_id"] for body in bodies]
assert client_ids[0] != client_ids[1]
assert client_ids[1] == client_ids[2]
assert channel._context_send_counts["ctx-1"] == 2
@pytest.mark.asyncio
async def test_quota_rejection_defers_final_until_fresh_context() -> None:
channel = _ready_channel()
channel._api_post = AsyncMock(side_effect=[{"ret": -2}, {"ret": 0}])
msg = OutboundMessage(
channel="weixin",
chat_id="wx-user",
content="deferred answer",
)
with pytest.raises(WeixinQuotaError):
await channel.send(msg)
first_client_id = channel._api_post.await_args_list[0].args[1]["msg"]["client_id"]
assert "wx-user" in channel._deferred_outbound
channel._context_tokens["wx-user"] = "ctx-2"
channel._context_token_at["wx-user"] = time.time()
await channel._retry_deferred_messages("wx-user")
second_client_id = channel._api_post.await_args_list[1].args[1]["msg"]["client_id"]
assert second_client_id == first_client_id
assert "wx-user" not in channel._deferred_outbound
@pytest.mark.asyncio
async def test_local_context_budget_stops_before_extra_api_call() -> None:
channel = _ready_channel(contextMessageBudget=1)
channel._api_post = AsyncMock(return_value={"ret": 0})
await channel._send_text("wx-user", "one", "ctx-1")
with pytest.raises(WeixinQuotaError, match="local safety budget"):
await channel._send_text("wx-user", "two", "ctx-1")
channel._api_post.assert_awaited_once()
@pytest.mark.asyncio
async def test_bounded_block_streaming_reserves_one_final_message() -> None:
channel = _ready_channel(
blockStreaming=True,
blockStreamingMinChars=200,
blockStreamingMaxMessages=3,
)
channel._send_text = AsyncMock()
await channel.send_delta("wx-user", "a" * 250, stream_id="stream-1")
await channel.send_delta("wx-user", "b" * 250, stream_id="stream-1")
await channel.send_delta("wx-user", "c" * 250, stream_id="stream-1")
await channel.send_delta("wx-user", "done", stream_id="stream-1", stream_end=True)
assert channel._send_text.await_count == 3
assert "stream-1" not in channel._stream_buffers
assert "stream-1" not in channel._stream_sent_counts
@pytest.mark.asyncio
async def test_structured_progress_is_capped_and_uses_one_run_id() -> None:
channel = _ready_channel(
replyProgressMessages=True,
replyProgressMaxMessages=2,
)
channel._send_message_item = AsyncMock()
events = [
{"phase": "start", "call_id": "call-1", "name": "read_file"},
{"phase": "end", "call_id": "call-1", "name": "read_file"},
{"phase": "start", "call_id": "call-2", "name": "exec"},
]
await channel.send(
OutboundMessage(
channel="weixin",
chat_id="wx-user",
content="read_file",
event=ProgressEvent(content="read_file", tool_hint=True, tool_events=events),
)
)
assert channel._send_message_item.await_count == 2
first = channel._send_message_item.await_args_list[0]
second = channel._send_message_item.await_args_list[1]
assert first.args[1]["type"] == ITEM_TOOL_CALL_START
assert second.args[1]["type"] == ITEM_TOOL_CALL_RESULT
assert first.kwargs["run_id"] == second.kwargs["run_id"]
@@ -1,25 +1,148 @@
import { useState } from "react";
import { useTranslation } from "react-i18next";
import { channelTranslator } from "@/channel-plugins/i18n";
import {
channelTranslator,
type ChannelTranslator,
} from "@/channel-plugins/i18n";
import type { ChannelPluginConnectFlowProps } from "@/channel-plugins/types";
import { ChannelQrConnectFlow } from "@/components/settings/channels/ChannelQrConnectFlow";
import {
ChannelQrConnectFlow,
type ChannelQrConnectPendingContext,
} from "@/components/settings/channels/ChannelQrConnectFlow";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import type { ChannelConnectPayload } from "@/lib/types";
type WeixinVerificationPayload = ChannelConnectPayload & {
challenge: "verify_code";
verification_failed?: boolean;
};
export const WEIXIN_AUTH_EXPIRED_MESSAGE =
"WeChat login expired. Scan again to reconnect.";
function isVerificationChallenge(
payload: ChannelConnectPayload,
): payload is WeixinVerificationPayload {
return (
"challenge" in payload
&& payload.challenge === "verify_code"
&& (
!("verification_failed" in payload)
|| typeof payload.verification_failed === "boolean"
)
);
}
function weixinConnectMessage(
payload: ChannelConnectPayload,
tx: ChannelTranslator,
): string {
if (payload.status === "succeeded") {
return tx("custom.connected", "WeChat is connected.");
}
if (payload.status === "expired") {
return tx("custom.expired", WEIXIN_AUTH_EXPIRED_MESSAGE);
}
if (payload.status === "failed") {
return payload.message
?? tx("custom.failed", "Unable to connect WeChat. Try again.");
}
if (payload.status === "cancelled") {
return tx("custom.stopped", "WeChat login stopped.");
}
if (isVerificationChallenge(payload)) {
return payload.verification_failed
? tx(
"custom.verifyMismatch",
"That code did not match. Enter the new number shown in WeChat.",
)
: tx(
"custom.verifyDescription",
"Enter the number shown in WeChat to continue.",
);
}
return tx("custom.waiting", "Waiting for WeChat scan...");
}
export function WeixinConnectFlow({
token,
feature,
idleLabel,
connectRequestId,
onFeaturesUpdate,
}: ChannelPluginConnectFlowProps) {
const { t } = useTranslation();
const tx = channelTranslator(t, "weixin");
const [verificationCode, setVerificationCode] = useState("");
const authExpired = feature.runtime_error === WEIXIN_AUTH_EXPIRED_MESSAGE;
const scanAgainLabel = t("settings.channels.scanAgain", {
defaultValue: "Scan again",
});
const renderVerification = ({
connect,
busy,
poll,
}: ChannelQrConnectPendingContext) => {
if (!isVerificationChallenge(connect)) return null;
return (
<form
className="mt-3 space-y-2"
onSubmit={(event) => {
event.preventDefault();
const code = verificationCode.trim();
if (!code) return;
void poll({ verify_code: code }).then((payload) => {
if (payload && !isVerificationChallenge(payload)) {
setVerificationCode("");
}
});
}}
>
<div className="text-[12px] font-semibold text-foreground">
{tx("custom.verifyTitle", "Verification required")}
</div>
<p className="text-[12px] leading-5 text-muted-foreground">
{weixinConnectMessage(connect, tx)}
</p>
<div className="flex gap-2">
<Input
value={verificationCode}
onChange={(event) => setVerificationCode(event.target.value)}
inputMode="numeric"
autoComplete="one-time-code"
placeholder={tx("custom.verifyPlaceholder", "Code")}
className="h-8 max-w-40"
aria-invalid={connect.verification_failed || undefined}
/>
<Button
type="submit"
size="sm"
className="h-8 rounded-full px-3 text-[12px] font-semibold"
disabled={busy || !verificationCode.trim()}
>
{tx("custom.verifySubmit", "Verify")}
</Button>
</div>
</form>
);
};
return (
<ChannelQrConnectFlow
token={token}
channelName="weixin"
idleLabel={idleLabel}
startOptions={{ force: authExpired }}
idleLabel={authExpired ? scanAgainLabel : idleLabel}
connectRequestId={connectRequestId}
forceOnRepeat
onFeaturesUpdate={onFeaturesUpdate}
pausePolling={isVerificationChallenge}
suppressSucceeded={feature.runtime_status === "failed"}
renderPending={renderVerification}
resolveMessage={(payload) => weixinConnectMessage(payload, tx)}
labels={{
qrAlt: tx("custom.qrAlt", "WeChat login QR code"),
scanTitle: tx("custom.scanTitle", "Scan with WeChat"),
@@ -31,7 +154,7 @@ export function WeixinConnectFlow({
connected: tx("custom.connected", "WeChat is connected."),
stopped: tx("custom.stopped", "WeChat login stopped."),
connecting: tx("custom.connecting", "Connecting..."),
scanAgain: t("settings.channels.scanAgain", { defaultValue: "Scan again" }),
scanAgain: scanAgainLabel,
connect: t("settings.channels.connect", { defaultValue: "Connect" }),
}}
/>
@@ -0,0 +1,555 @@
import { useCallback, useEffect, useMemo, useRef, useState, type ReactNode } from "react";
import { Check, ChevronDown, ExternalLink, Loader2, Plus } from "lucide-react";
import { useTranslation } from "react-i18next";
import { channelFieldMessageKey, channelTranslator } from "@/channel-plugins/i18n";
import { channelLocaleMessages } from "@/channel-plugins/locale-registry";
import type { ChannelPluginPanelProps } from "@/channel-plugins/types";
import { ToggleButton } from "@/components/settings/ToggleButton";
import {
chatAppGuideUrl,
docsUrlWithBase,
type ChannelConfigField,
} from "@/components/settings/channels/catalog";
import {
CredentialForm,
channelValuesForSave,
defaultChannelFieldValues,
} from "@/components/settings/channels/CredentialForm";
import { Button } from "@/components/ui/button";
import { useLogoFallback } from "@/hooks/useLogoFallback";
import { normalizeLocale } from "@/i18n/config";
import { configureChannel } from "@/lib/api";
import { logoFallbackUrls } from "@/lib/provider-brand";
import type {
ChannelRuntimeStatus,
ChannelSetupContractField,
NanobotFeatureInfo,
} from "@/lib/types";
import { cn } from "@/lib/utils";
import { useClient } from "@/providers/ClientProvider";
import {
WEIXIN_AUTH_EXPIRED_MESSAGE,
WeixinConnectFlow,
} from "./WeixinConnectFlow";
export const WEIXIN_PRIMARY_FIELD_KEYS = [
"channels.weixin.sendProgress",
"channels.weixin.sendToolHints",
"channels.weixin.streaming",
] as const;
export const WEIXIN_ADVANCED_FIELD_KEYS = [
"channels.weixin.allowFrom",
"channels.weixin.token",
"channels.weixin.replyProgressMessages",
"channels.weixin.replyProgressMaxMessages",
"channels.weixin.contextMessageBudget",
"channels.weixin.blockStreaming",
"channels.weixin.blockStreamingMinChars",
"channels.weixin.blockStreamingMaxMessages",
"channels.weixin.baseUrl",
"channels.weixin.cdnBaseUrl",
"channels.weixin.routeTag",
"channels.weixin.stateDir",
"channels.weixin.pollTimeout",
] as const;
export function WeixinPanel({
token,
feature,
actionKey,
chatAppsDocsUrl,
showBrandLogos,
onAction,
onFeaturesUpdate,
}: ChannelPluginPanelProps) {
const { client } = useClient();
const { t, i18n } = useTranslation();
const tx = (key: string, fallback: string) => t(key, { defaultValue: fallback });
const channelTx = channelTranslator(t, "weixin");
const runtimeError = weixinRuntimeError(feature.runtime_error, channelTx);
const displayName = channelTx("displayName", "WeChat");
const enabledBusy = actionKey === `enable:${feature.name}`;
const disabledBusy = actionKey === `disable:${feature.name}`;
const channelBusy = enabledBusy || disabledBusy;
const channelChecked =
feature.runtime_status === "running" || feature.runtime_status === "starting";
const missingSupport = feature.enabled && !feature.installed;
const alwaysEnabled = feature.capabilities?.includes("always_enabled") ?? false;
const toggleChecked = alwaysEnabled || channelChecked;
const channelToggleDisabled =
alwaysEnabled
|| channelBusy
|| (!feature.install_supported && !feature.installed && !feature.enabled);
const [connectRequestId, setConnectRequestId] = useState(0);
const [visibleSecrets, setVisibleSecrets] = useState<Record<string, boolean>>({});
const [touchedFields, setTouchedFields] = useState<Set<string>>(() => new Set());
const [saving, setSaving] = useState(false);
const [saveRevision, setSaveRevision] = useState(0);
const [attemptedRevision, setAttemptedRevision] = useState(0);
const [saveState, setSaveState] = useState<"idle" | "saved">("idle");
const [saveError, setSaveError] = useState<string | null>(null);
const configValuesKey = JSON.stringify(feature.config_values ?? {});
const setupFieldsKey = JSON.stringify(feature.setup?.fields ?? []);
const configuredFields = useMemo(
() => new Set(feature.configured_fields ?? []),
[feature.configured_fields],
);
const onLabel = tx("settings.values.on", "On");
const offLabel = tx("settings.values.off", "Off");
const setupFields = weixinSetupFields(
feature,
i18n.resolvedLanguage ?? i18n.language,
);
const primaryFields = localizeBooleanFields(setupFields.primary, onLabel, offLabel);
const advancedFields = localizeBooleanFields(setupFields.advanced, onLabel, offLabel);
const editableFields = [...primaryFields, ...advancedFields];
const docsUrl = docsUrlWithBase(chatAppGuideUrl("wechat"), chatAppsDocsUrl)
?? chatAppGuideUrl("wechat");
const [fieldValues, setFieldValues] = useState<Record<string, string>>(() =>
defaultChannelFieldValues(editableFields, feature.config_values),
);
const fieldValuesRef = useRef(fieldValues);
const touchedFieldsRef = useRef(touchedFields);
const editableFieldsRef = useRef(editableFields);
const saveContextRef = useRef({
token,
enabled: feature.enabled,
onFeaturesUpdate,
});
editableFieldsRef.current = editableFields;
saveContextRef.current = {
token,
enabled: feature.enabled,
onFeaturesUpdate,
};
useEffect(() => {
const nextValues = defaultChannelFieldValues(editableFields, feature.config_values);
for (const key of touchedFieldsRef.current) {
nextValues[key] = fieldValuesRef.current[key] ?? "";
}
fieldValuesRef.current = nextValues;
setFieldValues(nextValues);
setVisibleSecrets({});
}, [configValuesKey, setupFieldsKey]);
useEffect(() => {
if (saveState !== "saved") return;
const timeout = window.setTimeout(() => setSaveState("idle"), 1500);
return () => window.clearTimeout(timeout);
}, [saveState]);
const saveSettings = useCallback(async (
values: Record<string, string>,
savedFields: Set<string>,
) => {
const context = saveContextRef.current;
setSaving(true);
setSaveError(null);
setSaveState("idle");
try {
const payload = await configureChannel(
client,
"weixin",
channelValuesForSave(editableFieldsRef.current, values),
{ enable: context.enabled },
);
const remainingFields = new Set(touchedFieldsRef.current);
for (const key of savedFields) {
if (fieldValuesRef.current[key] === values[key]) remainingFields.delete(key);
}
touchedFieldsRef.current = remainingFields;
setTouchedFields(remainingFields);
setSaveState(remainingFields.size ? "idle" : "saved");
if (payload.nanobot_features) context.onFeaturesUpdate(payload.nanobot_features);
} catch (err) {
setSaveError((err as Error).message);
} finally {
setSaving(false);
}
}, [client]);
useEffect(() => {
if (
!editableFields.length
|| !touchedFields.size
|| saving
|| saveRevision <= attemptedRevision
) return;
const timeout = window.setTimeout(() => {
setAttemptedRevision(saveRevision);
void saveSettings(
{ ...fieldValuesRef.current },
new Set(touchedFieldsRef.current),
);
}, 500);
return () => window.clearTimeout(timeout);
}, [
attemptedRevision,
editableFields.length,
saveRevision,
saveSettings,
saving,
touchedFields.size,
]);
const setFieldValue = (key: string, value: string) => {
if (fieldValuesRef.current[key] === value) return;
const nextValues = { ...fieldValuesRef.current, [key]: value };
const nextTouchedFields = new Set(touchedFieldsRef.current).add(key);
fieldValuesRef.current = nextValues;
touchedFieldsRef.current = nextTouchedFields;
setFieldValues(nextValues);
setTouchedFields(nextTouchedFields);
setSaveError(null);
setSaveState("idle");
setSaveRevision((current) => current + 1);
};
const toggleAriaLabel = t("settings.channels.toggleChannel", {
name: displayName,
defaultValue: "{{name}} channel",
});
return (
<aside className="min-h-full rounded-[20px] bg-settings-surface p-5">
<div className="flex items-start justify-between gap-4">
<div className="flex min-w-0 items-start gap-3">
<WeixinLogo showBrandLogos={showBrandLogos} />
<div className="min-w-0 flex-1">
<h3 className="truncate text-[18px] font-semibold leading-6 text-foreground">
{displayName}
</h3>
<p className="mt-1 text-[13px] leading-5 text-muted-foreground">
{channelTx("description", "Use nanobot from WeChat conversations.")}
</p>
{missingSupport && feature.install_supported ? (
<Button
type="button"
size="sm"
variant="secondary"
disabled={enabledBusy}
onClick={() => onAction("enable", feature.name)}
className="mt-2 h-8 rounded-full px-3 text-[12px] font-semibold"
>
{enabledBusy ? (
<Loader2 className="mr-1.5 h-3.5 w-3.5 animate-spin" aria-hidden />
) : (
<Plus className="mr-1.5 h-3.5 w-3.5" aria-hidden />
)}
{tx("settings.nanobotFeatures.installSupport", "Install support")}
</Button>
) : null}
</div>
</div>
<div className="flex shrink-0 items-center gap-2 pt-1">
<WeixinStatusBadge status={feature.runtime_status}>
{weixinStatusLabel(feature, tx)}
</WeixinStatusBadge>
{channelBusy ? (
<Loader2 className="h-3.5 w-3.5 animate-spin text-muted-foreground" aria-hidden />
) : null}
<ToggleButton
checked={toggleChecked}
disabled={channelToggleDisabled}
ariaLabel={toggleAriaLabel}
label={toggleChecked ? onLabel : offLabel}
onChange={(checked) => {
if (checked && !channelChecked && feature.configured === false) {
setConnectRequestId((current) => current + 1);
return;
}
onAction(checked ? "enable" : "disable", feature.name);
}}
/>
</div>
</div>
{runtimeError ? (
<div className="mt-4 rounded-[12px] border border-destructive/20 bg-destructive/5 px-3 py-2 text-[12px] leading-5 text-destructive">
{runtimeError}
</div>
) : null}
<div className="mt-4 space-y-4">
<WeixinConnectFlow
token={token}
feature={feature}
idleLabel={channelTx("setup.primaryAction", "Connect WeChat")}
connectRequestId={connectRequestId}
onFeaturesUpdate={onFeaturesUpdate}
/>
{primaryFields.length ? (
<CredentialForm
fields={primaryFields}
values={fieldValues}
configuredFields={configuredFields}
visibleSecrets={visibleSecrets}
onChange={setFieldValue}
onToggleSecret={(key) => {
setVisibleSecrets((current) => ({ ...current, [key]: !current[key] }));
}}
compact
/>
) : null}
<div
role="status"
aria-live="polite"
aria-atomic="true"
className={cn(
"flex items-center justify-end gap-1.5 text-[11px] leading-4 text-muted-foreground",
!saving && saveState !== "saved" && "sr-only",
)}
>
{saving ? (
<>
<Loader2 className="h-3 w-3 animate-spin" aria-hidden />
{tx("settings.actions.saving", "Saving")}
</>
) : saveState === "saved" ? (
<>
<Check className="h-3 w-3" aria-hidden />
{tx("settings.channels.savedSettings", "Saved settings.")}
</>
) : null}
</div>
{saveError ? (
<div
role="alert"
className="rounded-[12px] border border-destructive/20 bg-destructive/5 px-3 py-2 text-[12px] leading-5 text-destructive"
>
{saveError}
</div>
) : null}
{advancedFields.length ? (
<details className="group text-[12px] leading-5 text-muted-foreground">
<summary className="cursor-pointer list-none text-[12px] font-semibold text-foreground">
<span className="inline-flex items-center gap-1.5">
{tx("settings.channels.advanced", "Advanced")}
<ChevronDown
className="h-3.5 w-3.5 transition-transform group-open:rotate-180"
aria-hidden
/>
</span>
</summary>
<div className="mt-3">
<CredentialForm
fields={advancedFields}
values={fieldValues}
configuredFields={configuredFields}
visibleSecrets={visibleSecrets}
onChange={setFieldValue}
onToggleSecret={(key) => {
setVisibleSecrets((current) => ({ ...current, [key]: !current[key] }));
}}
compact
/>
</div>
</details>
) : null}
<div className="flex justify-end">
<WeixinGuideLink
url={docsUrl}
label={channelTx("setup.docsLabel", "Open WeChat setup")}
/>
</div>
</div>
</aside>
);
}
function weixinSetupFields(
feature: NanobotFeatureInfo,
locale: string,
): { primary: ChannelConfigField[]; advanced: ChannelConfigField[] } {
const fields = feature.setup?.fields ?? [];
const fieldsByKey = new Map(fields.map((field) => [field.key, field]));
const messages = channelLocaleMessages("weixin", normalizeLocale(locale))?.setup;
const knownKeys = new Set<string>([
...WEIXIN_PRIMARY_FIELD_KEYS,
...WEIXIN_ADVANCED_FIELD_KEYS,
]);
const extraKeys = fields
.map((field) => field.key)
.filter((key) => !knownKeys.has(key));
const hydrate = (keys: readonly string[]) => keys.flatMap((key) => {
const field = fieldsByKey.get(key);
if (!field) return [];
const copy = messages?.fields?.[channelFieldMessageKey("weixin", key)];
return [weixinConfigField(field, copy)];
});
return {
primary: hydrate(WEIXIN_PRIMARY_FIELD_KEYS),
advanced: hydrate([...WEIXIN_ADVANCED_FIELD_KEYS, ...extraKeys]),
};
}
function weixinConfigField(
field: ChannelSetupContractField,
copy: { label: string; placeholder?: string; help?: string; choices?: Record<string, string> }
| undefined,
): ChannelConfigField {
const choices = field.kind === "bool" ? ["true", "false"] : field.choices;
return {
key: field.key,
label: copy?.label ?? fieldLabel(field.field),
placeholder: copy?.placeholder,
help: copy?.help,
secret: field.kind === "secret",
optional: !field.required,
inputType: field.kind === "int" ? "number" : undefined,
defaultValue: field.default_value,
options:
field.kind === "enum" || field.kind === "bool"
? choices.map((choice) => ({
value: choice,
label: copy?.choices?.[choice] ?? fieldLabel(choice),
}))
: undefined,
};
}
function fieldLabel(value: string): string {
const spaced = value
.replace(/([a-z0-9])([A-Z])/g, "$1 $2")
.replace(/[_-]+/g, " ")
.trim();
return spaced ? spaced[0].toUpperCase() + spaced.slice(1) : value;
}
function WeixinLogo({ showBrandLogos }: { showBrandLogos: boolean }) {
const logoUrls = useMemo(() => logoFallbackUrls("https://weixin.qq.com/favicon.ico"), []);
const { logoUrl, onLogoError, onLogoLoad } = useLogoFallback(logoUrls);
if (showBrandLogos && logoUrl) {
return (
<span className="grid h-10 w-10 shrink-0 place-items-center rounded-[12px] bg-background">
<img
src={logoUrl}
alt=""
decoding="async"
loading="lazy"
className="h-5.5 w-5.5 max-h-6 max-w-6 object-contain"
onLoad={onLogoLoad}
onError={onLogoError}
/>
</span>
);
}
return (
<span
className="flex h-10 w-10 shrink-0 items-center justify-center rounded-[12px] bg-background text-[11px] font-bold"
style={{ color: "#07C160" }}
aria-hidden
>
WX
</span>
);
}
function WeixinGuideLink({ url, label }: { url: string; label: string }) {
const logoUrls = useMemo(() => logoFallbackUrls("https://weixin.qq.com/favicon.ico"), []);
const { logoUrl, onLogoError, onLogoLoad } = useLogoFallback(logoUrls);
return (
<a
href={url}
target="_blank"
rel="noreferrer"
className="inline-flex max-w-full items-center gap-2 rounded-full bg-background/80 py-1 pl-1 pr-2.5 text-[11.5px] font-semibold text-foreground transition-colors hover:bg-background"
>
<span
className="grid h-5 w-5 shrink-0 place-items-center overflow-hidden rounded-full bg-muted/70 text-[9px] font-bold"
style={{ color: "#07C160" }}
aria-hidden
>
{logoUrl ? (
<img
src={logoUrl}
alt=""
decoding="async"
loading="lazy"
className="h-3.5 w-3.5 object-contain"
onLoad={onLogoLoad}
onError={onLogoError}
/>
) : (
"WX"
)}
</span>
<span className="truncate">{label}</span>
<ExternalLink className="h-3.5 w-3.5 shrink-0 text-muted-foreground" aria-hidden />
</a>
);
}
function WeixinStatusBadge({
children,
status,
}: {
children: ReactNode;
status?: ChannelRuntimeStatus;
}) {
return (
<span className={cn(
"shrink-0 rounded-full px-2 py-0.5 text-[11px] font-medium leading-4",
status === "failed"
? "bg-destructive/10 text-destructive"
: status === "running"
? "bg-emerald-500/10 text-emerald-700 dark:text-emerald-200"
: "bg-muted/75 text-muted-foreground",
)}>
{children}
</span>
);
}
function weixinStatusLabel(
feature: NanobotFeatureInfo,
tx: (key: string, fallback: string) => string,
): string {
if (feature.runtime_status === "failed") {
return tx("settings.channels.runtimeFailed", "Failed");
}
if (feature.runtime_status === "starting") {
return tx("settings.channels.runtimeStarting", "Starting");
}
if (feature.runtime_status === "running") return tx("settings.values.on", "On");
if (feature.enabled) return tx("settings.channels.runtimeStopped", "Not running");
return tx("settings.values.off", "Off");
}
function weixinRuntimeError(
error: string | undefined,
tx: (key: string, fallback: string) => string,
): string | undefined {
if (error === WEIXIN_AUTH_EXPIRED_MESSAGE) {
return tx("custom.expired", error);
}
return error;
}
function localizeBooleanFields(
fields: ChannelConfigField[],
onLabel: string,
offLabel: string,
): ChannelConfigField[] {
return fields.map((field) => {
const values = new Set(field.options?.map((option) => option.value));
if (values.size !== 2 || !values.has("true") || !values.has("false")) return field;
return {
...field,
options: field.options?.map((option) => ({
...option,
label: option.value === "true" ? onLabel : offLabel,
})),
};
});
}
+8 -4
View File
@@ -2,8 +2,14 @@ import type { ChannelUiContribution } from "@/channel-plugins/types";
import { chatAppGuideUrl } from "@/components/settings/channels/catalog";
import { WeixinConnectFlow } from "./WeixinConnectFlow";
import {
WEIXIN_ADVANCED_FIELD_KEYS,
WEIXIN_PRIMARY_FIELD_KEYS,
WeixinPanel,
} from "./WeixinPanel";
export default {
Panel: WeixinPanel,
ConnectFlow: WeixinConnectFlow,
canConnectBeforeConfigured: true,
aliases: {
@@ -18,10 +24,8 @@ export default {
mode: "connect",
command: "nanobot channels login weixin",
docsUrl: chatAppGuideUrl("wechat"),
manualFields: [
{ key: "channels.weixin.allowFrom" },
{ key: "channels.weixin.token" },
],
fields: WEIXIN_PRIMARY_FIELD_KEYS.map((key) => ({ key })),
manualFields: WEIXIN_ADVANCED_FIELD_KEYS.map((key) => ({ key })),
},
},
} satisfies ChannelUiContribution;
+23 -2
View File
@@ -20,7 +20,21 @@
"token": {
"label": "Token",
"placeholder": "Saved by QR login"
}
},
"sendProgress": { "label": "Send progress" },
"sendToolHints": { "label": "Send tool hints" },
"streaming": { "label": "Use streaming API" },
"replyProgressMessages": { "label": "Send structured progress" },
"replyProgressMaxMessages": { "label": "Structured progress limit" },
"contextMessageBudget": { "label": "Context message budget" },
"blockStreaming": { "label": "Send response blocks" },
"blockStreamingMinChars": { "label": "Minimum block size" },
"blockStreamingMaxMessages": { "label": "Block message limit" },
"baseUrl": { "label": "API URL" },
"cdnBaseUrl": { "label": "CDN URL" },
"routeTag": { "label": "Route tag" },
"stateDir": { "label": "State directory" },
"pollTimeout": { "label": "Poll timeout" }
}
},
"custom": {
@@ -30,6 +44,13 @@
"waiting": "Waiting for WeChat scan...",
"connected": "WeChat is connected.",
"stopped": "WeChat login stopped.",
"connecting": "Connecting..."
"connecting": "Connecting...",
"verifyTitle": "Verification required",
"verifyDescription": "Enter the number shown in WeChat to continue.",
"verifyMismatch": "That code did not match. Enter the new number shown in WeChat.",
"expired": "WeChat login expired. Scan again to reconnect.",
"failed": "Unable to connect WeChat. Try again.",
"verifyPlaceholder": "Code",
"verifySubmit": "Verify"
}
}
+23 -2
View File
@@ -20,7 +20,21 @@
"token": {
"label": "Token",
"placeholder": "Guardado al iniciar sesión por QR"
}
},
"sendProgress": { "label": "Enviar progreso" },
"sendToolHints": { "label": "Enviar indicaciones de herramientas" },
"streaming": { "label": "Usar API de streaming" },
"replyProgressMessages": { "label": "Enviar progreso estructurado" },
"replyProgressMaxMessages": { "label": "Límite de progreso estructurado" },
"contextMessageBudget": { "label": "Presupuesto de mensajes por contexto" },
"blockStreaming": { "label": "Enviar respuestas por bloques" },
"blockStreamingMinChars": { "label": "Tamaño mínimo del bloque" },
"blockStreamingMaxMessages": { "label": "Límite de mensajes por bloques" },
"baseUrl": { "label": "URL de la API" },
"cdnBaseUrl": { "label": "URL de la CDN" },
"routeTag": { "label": "Etiqueta de ruta" },
"stateDir": { "label": "Directorio de estado" },
"pollTimeout": { "label": "Tiempo de espera de consulta" }
}
},
"custom": {
@@ -30,6 +44,13 @@
"waiting": "Esperando el escaneo de WeChat...",
"connected": "WeChat está conectado.",
"stopped": "Inicio de WeChat detenido.",
"connecting": "Conectando..."
"connecting": "Conectando...",
"verifyTitle": "Se requiere verificación",
"verifyDescription": "Introduce el número que aparece en WeChat para continuar.",
"verifyMismatch": "El código no coincide. Introduce el nuevo número que aparece en WeChat.",
"expired": "El inicio de sesión de WeChat caducó. Escanea de nuevo para volver a conectarte.",
"failed": "No se pudo conectar WeChat. Inténtalo de nuevo.",
"verifyPlaceholder": "Código",
"verifySubmit": "Verificar"
}
}
+23 -2
View File
@@ -20,7 +20,21 @@
"token": {
"label": "Jeton",
"placeholder": "Enregistré après la connexion QR"
}
},
"sendProgress": { "label": "Envoyer la progression" },
"sendToolHints": { "label": "Envoyer les indications doutils" },
"streaming": { "label": "Utiliser lAPI de streaming" },
"replyProgressMessages": { "label": "Envoyer la progression structurée" },
"replyProgressMaxMessages": { "label": "Limite de progression structurée" },
"contextMessageBudget": { "label": "Budget de messages du contexte" },
"blockStreaming": { "label": "Envoyer la réponse par blocs" },
"blockStreamingMinChars": { "label": "Taille minimale dun bloc" },
"blockStreamingMaxMessages": { "label": "Limite de messages par blocs" },
"baseUrl": { "label": "URL de lAPI" },
"cdnBaseUrl": { "label": "URL du CDN" },
"routeTag": { "label": "Étiquette de routage" },
"stateDir": { "label": "Répertoire d’état" },
"pollTimeout": { "label": "Délai dinterrogation" }
}
},
"custom": {
@@ -30,6 +44,13 @@
"waiting": "En attente du scan WeChat...",
"connected": "WeChat est connecté.",
"stopped": "Connexion WeChat arrêtée.",
"connecting": "Connexion..."
"connecting": "Connexion...",
"verifyTitle": "Vérification requise",
"verifyDescription": "Saisissez le nombre affiché dans WeChat pour continuer.",
"verifyMismatch": "Le code ne correspond pas. Saisissez le nouveau nombre affiché dans WeChat.",
"expired": "La connexion WeChat a expiré. Scannez à nouveau pour vous reconnecter.",
"failed": "Impossible de connecter WeChat. Réessayez.",
"verifyPlaceholder": "Code",
"verifySubmit": "Vérifier"
}
}
+23 -2
View File
@@ -20,7 +20,21 @@
"token": {
"label": "Token",
"placeholder": "Disimpan saat login QR"
}
},
"sendProgress": { "label": "Kirim progres" },
"sendToolHints": { "label": "Kirim petunjuk alat" },
"streaming": { "label": "Gunakan API streaming" },
"replyProgressMessages": { "label": "Kirim progres terstruktur" },
"replyProgressMaxMessages": { "label": "Batas progres terstruktur" },
"contextMessageBudget": { "label": "Anggaran pesan konteks" },
"blockStreaming": { "label": "Kirim respons per blok" },
"blockStreamingMinChars": { "label": "Ukuran blok minimum" },
"blockStreamingMaxMessages": { "label": "Batas pesan blok" },
"baseUrl": { "label": "URL API" },
"cdnBaseUrl": { "label": "URL CDN" },
"routeTag": { "label": "Tag rute" },
"stateDir": { "label": "Direktori status" },
"pollTimeout": { "label": "Batas waktu polling" }
}
},
"custom": {
@@ -30,6 +44,13 @@
"waiting": "Menunggu pemindaian WeChat...",
"connected": "WeChat sudah terhubung.",
"stopped": "Login WeChat dihentikan.",
"connecting": "Menghubungkan..."
"connecting": "Menghubungkan...",
"verifyTitle": "Verifikasi diperlukan",
"verifyDescription": "Masukkan angka yang ditampilkan di WeChat untuk melanjutkan.",
"verifyMismatch": "Kode tidak cocok. Masukkan angka baru yang ditampilkan di WeChat.",
"expired": "Login WeChat telah kedaluwarsa. Pindai lagi untuk menghubungkan kembali.",
"failed": "Tidak dapat menghubungkan WeChat. Coba lagi.",
"verifyPlaceholder": "Kode",
"verifySubmit": "Verifikasi"
}
}
+23 -2
View File
@@ -20,7 +20,21 @@
"token": {
"label": "トークン",
"placeholder": "QR ログインで保存"
}
},
"sendProgress": { "label": "進捗を送信" },
"sendToolHints": { "label": "ツールのヒントを送信" },
"streaming": { "label": "ストリーミング API を使用" },
"replyProgressMessages": { "label": "構造化された進捗を送信" },
"replyProgressMaxMessages": { "label": "構造化進捗の上限" },
"contextMessageBudget": { "label": "コンテキストのメッセージ予算" },
"blockStreaming": { "label": "応答をブロック単位で送信" },
"blockStreamingMinChars": { "label": "最小ブロックサイズ" },
"blockStreamingMaxMessages": { "label": "ブロックメッセージの上限" },
"baseUrl": { "label": "API URL" },
"cdnBaseUrl": { "label": "CDN URL" },
"routeTag": { "label": "ルートタグ" },
"stateDir": { "label": "状態ディレクトリ" },
"pollTimeout": { "label": "ポーリングタイムアウト" }
}
},
"custom": {
@@ -30,6 +44,13 @@
"waiting": "WeChat のスキャンを待っています...",
"connected": "WeChat に接続しました。",
"stopped": "WeChat ログインを停止しました。",
"connecting": "接続中..."
"connecting": "接続中...",
"verifyTitle": "確認が必要です",
"verifyDescription": "WeChat に表示された数字を入力してください。",
"verifyMismatch": "コードが一致しません。WeChat に表示された新しい数字を入力してください。",
"expired": "WeChat のログイン期限が切れました。再接続するにはもう一度スキャンしてください。",
"failed": "WeChat に接続できません。もう一度お試しください。",
"verifyPlaceholder": "コード",
"verifySubmit": "確認"
}
}
+23 -2
View File
@@ -20,7 +20,21 @@
"token": {
"label": "토큰",
"placeholder": "QR 로그인으로 저장됨"
}
},
"sendProgress": { "label": "진행 상황 보내기" },
"sendToolHints": { "label": "도구 힌트 보내기" },
"streaming": { "label": "스트리밍 API 사용" },
"replyProgressMessages": { "label": "구조화된 진행 상황 보내기" },
"replyProgressMaxMessages": { "label": "구조화된 진행 메시지 한도" },
"contextMessageBudget": { "label": "컨텍스트 메시지 예산" },
"blockStreaming": { "label": "응답을 블록으로 보내기" },
"blockStreamingMinChars": { "label": "최소 블록 크기" },
"blockStreamingMaxMessages": { "label": "블록 메시지 한도" },
"baseUrl": { "label": "API URL" },
"cdnBaseUrl": { "label": "CDN URL" },
"routeTag": { "label": "경로 태그" },
"stateDir": { "label": "상태 디렉터리" },
"pollTimeout": { "label": "폴링 제한 시간" }
}
},
"custom": {
@@ -30,6 +44,13 @@
"waiting": "WeChat 스캔을 기다리는 중...",
"connected": "WeChat이 연결되었습니다.",
"stopped": "WeChat 로그인이 중지되었습니다.",
"connecting": "연결 중..."
"connecting": "연결 중...",
"verifyTitle": "인증 필요",
"verifyDescription": "계속하려면 WeChat에 표시된 숫자를 입력하세요.",
"verifyMismatch": "코드가 일치하지 않습니다. WeChat에 표시된 새 숫자를 입력하세요.",
"expired": "WeChat 로그인이 만료되었습니다. 다시 연결하려면 다시 스캔하세요.",
"failed": "WeChat에 연결할 수 없습니다. 다시 시도하세요.",
"verifyPlaceholder": "코드",
"verifySubmit": "인증"
}
}
@@ -20,7 +20,21 @@
"token": {
"label": "Token",
"placeholder": "Salvo pelo login via QR"
}
},
"sendProgress": { "label": "Enviar progresso" },
"sendToolHints": { "label": "Enviar dicas de ferramentas" },
"streaming": { "label": "Usar API de streaming" },
"replyProgressMessages": { "label": "Enviar progresso estruturado" },
"replyProgressMaxMessages": { "label": "Limite de progresso estruturado" },
"contextMessageBudget": { "label": "Orçamento de mensagens do contexto" },
"blockStreaming": { "label": "Enviar resposta em blocos" },
"blockStreamingMinChars": { "label": "Tamanho mínimo do bloco" },
"blockStreamingMaxMessages": { "label": "Limite de mensagens em blocos" },
"baseUrl": { "label": "URL da API" },
"cdnBaseUrl": { "label": "URL da CDN" },
"routeTag": { "label": "Etiqueta de rota" },
"stateDir": { "label": "Diretório de estado" },
"pollTimeout": { "label": "Tempo limite da consulta" }
}
},
"custom": {
@@ -30,6 +44,13 @@
"waiting": "Aguardando leitura do WeChat...",
"connected": "WeChat está conectado.",
"stopped": "Login do WeChat interrompido.",
"connecting": "Conectando..."
"connecting": "Conectando...",
"verifyTitle": "Verificação necessária",
"verifyDescription": "Digite o número exibido no WeChat para continuar.",
"verifyMismatch": "O código não corresponde. Digite o novo número exibido no WeChat.",
"expired": "O login do WeChat expirou. Escaneie novamente para reconectar.",
"failed": "Não foi possível conectar o WeChat. Tente novamente.",
"verifyPlaceholder": "Código",
"verifySubmit": "Verificar"
}
}
+23 -2
View File
@@ -20,7 +20,21 @@
"token": {
"label": "Token",
"placeholder": "Được lưu khi đăng nhập QR"
}
},
"sendProgress": { "label": "Gửi tiến trình" },
"sendToolHints": { "label": "Gửi gợi ý công cụ" },
"streaming": { "label": "Sử dụng API phát trực tiếp" },
"replyProgressMessages": { "label": "Gửi tiến trình có cấu trúc" },
"replyProgressMaxMessages": { "label": "Giới hạn tiến trình có cấu trúc" },
"contextMessageBudget": { "label": "Ngân sách tin nhắn ngữ cảnh" },
"blockStreaming": { "label": "Gửi phản hồi theo khối" },
"blockStreamingMinChars": { "label": "Kích thước khối tối thiểu" },
"blockStreamingMaxMessages": { "label": "Giới hạn tin nhắn theo khối" },
"baseUrl": { "label": "URL API" },
"cdnBaseUrl": { "label": "URL CDN" },
"routeTag": { "label": "Thẻ định tuyến" },
"stateDir": { "label": "Thư mục trạng thái" },
"pollTimeout": { "label": "Thời gian chờ thăm dò" }
}
},
"custom": {
@@ -30,6 +44,13 @@
"waiting": "Đang chờ quét WeChat...",
"connected": "WeChat đã kết nối.",
"stopped": "Đăng nhập WeChat đã dừng.",
"connecting": "Đang kết nối..."
"connecting": "Đang kết nối...",
"verifyTitle": "Cần xác minh",
"verifyDescription": "Nhập số hiển thị trong WeChat để tiếp tục.",
"verifyMismatch": "Mã không khớp. Nhập số mới hiển thị trong WeChat.",
"expired": "Đăng nhập WeChat đã hết hạn. Hãy quét lại để kết nối lại.",
"failed": "Không thể kết nối WeChat. Hãy thử lại.",
"verifyPlaceholder": "Mã",
"verifySubmit": "Xác minh"
}
}
@@ -21,7 +21,21 @@
"token": {
"label": "令牌",
"placeholder": "二维码登录后自动保存"
}
},
"sendProgress": { "label": "发送进度消息" },
"sendToolHints": { "label": "发送工具提示" },
"streaming": { "label": "使用流式 API" },
"replyProgressMessages": { "label": "发送结构化进度" },
"replyProgressMaxMessages": { "label": "结构化进度消息上限" },
"contextMessageBudget": { "label": "上下文消息预算" },
"blockStreaming": { "label": "分块发送回复" },
"blockStreamingMinChars": { "label": "最小分块字符数" },
"blockStreamingMaxMessages": { "label": "分块消息上限" },
"baseUrl": { "label": "API 地址" },
"cdnBaseUrl": { "label": "CDN 地址" },
"routeTag": { "label": "路由标签" },
"stateDir": { "label": "状态目录" },
"pollTimeout": { "label": "轮询超时" }
}
},
"custom": {
@@ -31,6 +45,13 @@
"waiting": "正在等待微信扫码...",
"connected": "微信已连接。",
"stopped": "微信登录已停止。",
"connecting": "正在连接..."
"connecting": "正在连接...",
"verifyTitle": "需要验证",
"verifyDescription": "输入手机微信中显示的数字以继续。",
"verifyMismatch": "验证码不匹配,请输入微信中显示的新数字。",
"expired": "微信登录已过期,请重新扫码连接。",
"failed": "无法连接微信,请重试。",
"verifyPlaceholder": "验证码",
"verifySubmit": "验证"
}
}
@@ -21,7 +21,21 @@
"token": {
"label": "權杖",
"placeholder": "二維碼登入後自動儲存"
}
},
"sendProgress": { "label": "傳送進度訊息" },
"sendToolHints": { "label": "傳送工具提示" },
"streaming": { "label": "使用串流 API" },
"replyProgressMessages": { "label": "傳送結構化進度" },
"replyProgressMaxMessages": { "label": "結構化進度訊息上限" },
"contextMessageBudget": { "label": "上下文訊息預算" },
"blockStreaming": { "label": "分塊傳送回覆" },
"blockStreamingMinChars": { "label": "最小分塊字元數" },
"blockStreamingMaxMessages": { "label": "分塊訊息上限" },
"baseUrl": { "label": "API 位址" },
"cdnBaseUrl": { "label": "CDN 位址" },
"routeTag": { "label": "路由標籤" },
"stateDir": { "label": "狀態目錄" },
"pollTimeout": { "label": "輪詢逾時" }
}
},
"custom": {
@@ -31,6 +45,13 @@
"waiting": "正在等待微信掃碼...",
"connected": "微信已連接。",
"stopped": "微信登入已停止。",
"connecting": "正在連接..."
"connecting": "正在連接...",
"verifyTitle": "需要驗證",
"verifyDescription": "輸入手機微信中顯示的數字以繼續。",
"verifyMismatch": "驗證碼不符,請輸入微信中顯示的新數字。",
"expired": "微信登入已過期,請重新掃碼連線。",
"failed": "無法連接微信,請重試。",
"verifyPlaceholder": "驗證碼",
"verifySubmit": "驗證"
}
}
+1
View File
@@ -669,6 +669,7 @@ def _run_gateway(
webui_runtime_surface=webui_runtime_surface,
webui_runtime_capabilities=webui_runtime_capabilities,
webui_skill_state_action=_webui_skill_state_action,
config_path=Path(config_path),
)
def _pick_heartbeat_target() -> tuple[str, str]:
+2 -1
View File
@@ -32,6 +32,7 @@ from nanobot.cli.models import (
)
from nanobot.config.loader import get_config_path, load_config, resolve_config_env_vars
from nanobot.config.schema import Config, ModelPresetConfig
from nanobot.providers.oauth_guidance import OAUTH_CLI_KIT_MISSING_MESSAGE
console = Console()
@@ -1674,7 +1675,7 @@ def _quick_start_oauth_login(config: Config, provider_name: str) -> bool:
login_oauth_interactive,
)
except ImportError:
console.print("[red]oauth_cli_kit not installed. Run: pip install oauth-cli-kit[/red]")
console.print(f"[red]{OAUTH_CLI_KIT_MISSING_MESSAGE}[/red]")
return False
try:
+6 -5
View File
@@ -12,6 +12,7 @@ import typer
from rich.console import Console
from nanobot import __logo__
from nanobot.providers.oauth_guidance import OAUTH_CLI_KIT_MISSING_MESSAGE
if TYPE_CHECKING:
from nanobot.providers.registry import ProviderSpec
@@ -74,7 +75,7 @@ def _required_module_attribute(module_name: str, attribute: str) -> object:
def _load_openai_oauth_client() -> tuple[_GetOAuthToken, _LoginOAuthInteractive]:
"""Load the optional untyped OAuth client behind a typed boundary."""
"""Load the untyped OAuth client behind a typed boundary."""
return (
cast(_GetOAuthToken, _required_module_attribute("oauth_cli_kit", "get_token")),
cast(
@@ -85,7 +86,7 @@ def _load_openai_oauth_client() -> tuple[_GetOAuthToken, _LoginOAuthInteractive]
def _load_openai_oauth_storage() -> tuple[_OAuthProviderConfig, _FileTokenStorageFactory]:
"""Load the optional untyped OAuth storage API behind a typed boundary."""
"""Load the untyped OAuth storage API behind a typed boundary."""
return (
cast(
_OAuthProviderConfig,
@@ -241,7 +242,7 @@ def _login_openai_codex() -> None:
f"[green]✓ Authenticated with OpenAI Codex[/green] [dim]{token.account_id}[/dim]"
)
except ImportError:
console.print("[red]oauth_cli_kit not installed. Run: pip install oauth-cli-kit[/red]")
console.print(f"[red]{OAUTH_CLI_KIT_MISSING_MESSAGE}[/red]")
raise typer.Exit(1)
@@ -250,7 +251,7 @@ def _logout_openai_codex() -> None:
try:
provider_config, storage_factory = _load_openai_oauth_storage()
except ImportError:
console.print("[red]oauth_cli_kit not installed. Run: pip install oauth-cli-kit[/red]")
console.print(f"[red]{OAUTH_CLI_KIT_MISSING_MESSAGE}[/red]")
raise typer.Exit(1)
storage = storage_factory(token_filename=provider_config.token_filename)
@@ -309,7 +310,7 @@ def _logout_github_copilot() -> None:
try:
from nanobot.providers.github_copilot_provider import get_storage
except ImportError:
console.print("[red]oauth_cli_kit not installed. Run: pip install oauth-cli-kit[/red]")
console.print(f"[red]{OAUTH_CLI_KIT_MISSING_MESSAGE}[/red]")
raise typer.Exit(1)
storage = get_storage()
-7
View File
@@ -25,9 +25,6 @@ from nanobot.cron.types import (
CronSchedule,
CronStore,
)
from nanobot.utils.run_records import (
safe_run_record_name,
)
from nanobot.utils.run_records import (
write_run_record as write_automation_run_record,
)
@@ -440,10 +437,6 @@ class CronService:
tmp_path.unlink(missing_ok=True)
raise
@staticmethod
def _safe_run_record_name(run_id: str) -> str:
return safe_run_record_name(run_id)
def write_run_record(self, run_id: str, record: dict[str, Any]) -> None:
"""Write an internal audit record for one cron execution."""
write_automation_run_record(self._run_records_dir, run_id, record)
-6
View File
@@ -7,7 +7,6 @@ from typing import Any, Mapping
from nanobot.cron.types import CronJob
from nanobot.session.automation_turns import (
AutomationTurnSpec,
automation_history_overrides_for_spec,
automation_trigger,
)
@@ -63,11 +62,6 @@ def cron_run_id(metadata: Mapping[str, Any] | None) -> str | None:
return value if isinstance(value, str) and value else None
def cron_history_overrides(metadata: Mapping[str, Any] | None) -> tuple[str | None, dict[str, Any]]:
"""Return session-history text/metadata overrides for a cron turn."""
return automation_history_overrides_for_spec(metadata, CRON_AUTOMATION_SPEC)
def is_bound_cron_job(job: CronJob) -> bool:
"""True for session-bound cron jobs with complete delivery context."""
payload = job.payload
+6
View File
@@ -0,0 +1,6 @@
"""Shared recovery guidance for OAuth dependency failures."""
OAUTH_CLI_KIT_MISSING_MESSAGE = (
"This nanobot installation is missing the required oauth-cli-kit package. "
"Reinstall or upgrade nanobot-ai using the same installation method."
)
+1 -1
View File
@@ -586,7 +586,7 @@ class OpenAICompatProvider(LLMProvider):
if os.environ.get("LANGFUSE_SECRET_KEY"):
logger.warning(
"LANGFUSE_SECRET_KEY is set but langfuse is not installed; "
"install with `pip install langfuse` to enable tracing"
"run `nanobot plugins enable langfuse` to enable tracing"
)
from openai import AsyncOpenAI as _AsyncOpenAI
AsyncOpenAI = _AsyncOpenAI
+70 -17
View File
@@ -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, Protocol, TypedDict, cast
from typing import Any, Callable, Collection, Protocol, TypedDict, cast
from weakref import WeakValueDictionary
from loguru import logger
@@ -36,6 +36,7 @@ from nanobot.utils.subagent_channel_display import scrub_subagent_announce_body
FILE_MAX_MESSAGES = 2000
SESSION_CACHE_MAX_SIZE = 128
MIN_REPLAY_MAX_MESSAGES = 120
MIN_COMPACTED_REPLAY_MESSAGES = 8
REPLAY_TOKENS_PER_MESSAGE = 100
_MESSAGE_TIME_PREFIX_RE = re.compile(r"^\[Message Time: [^\]]+\]\n?")
_LOCAL_IMAGE_BREADCRUMB_RE = re.compile(r"^\[image: (?:/|~)[^\]]+\]\s*$")
@@ -146,6 +147,15 @@ class RetentionResult:
already_consolidated_count: int
@dataclass(frozen=True)
class SessionPolicy:
"""Runtime rules that do not belong in durable session data."""
persist: bool = True
log_content: bool = True
disabled_tools: frozenset[str] = frozenset()
@dataclass
class Session:
"""A conversation session."""
@@ -157,6 +167,7 @@ class Session:
metadata: dict[str, Any] = field(default_factory=dict)
last_consolidated: int = 0 # Number of messages already consolidated to files
provider_state: ProviderConversationState | None = field(default=None, repr=False)
policy: SessionPolicy = field(default_factory=SessionPolicy, repr=False, compare=False)
def __post_init__(self) -> None:
if not isinstance(cast(object, self.metadata), dict):
@@ -191,19 +202,37 @@ class Session:
extend_to_user: bool = False,
include_runtime_context: bool = True,
) -> list[dict[str, Any]]:
"""Return unconsolidated messages for LLM input.
"""Return recent replayable messages for LLM input.
History is sliced by message count first (``max_messages``), then by
token budget from the tail (``max_tokens``) when provided.
"""
unconsolidated = self.messages[self.last_consolidated:]
replay_start = self.last_consolidated
if replay_start:
# ``last_consolidated`` is archive progress, not a replay boundary.
# Keep a small raw suffix for continuity, extending back to the user
# that started an assistant/tool sequence when necessary.
recent_start = recent_message_start_index(
self.messages,
MIN_COMPACTED_REPLAY_MESSAGES,
extend_to_user=True,
)
replay_start = min(replay_start, recent_start)
replayable = self.messages[replay_start:]
max_messages = max_messages if max_messages > 0 else FILE_MAX_MESSAGES
start_idx = recent_message_start_index(
unconsolidated,
max_messages,
extend_to_user=extend_to_user,
)
sliced = unconsolidated[start_idx:]
unarchived_count = len(self.messages) - self.last_consolidated
if replay_start < self.last_consolidated and unarchived_count < max_messages:
# The archived replay suffix can exceed the nominal count when one
# tool-heavy turn spans the boundary. Preserve that complete turn.
start_idx = 0
else:
start_idx = recent_message_start_index(
replayable,
max_messages,
extend_to_user=extend_to_user,
)
sliced = replayable[start_idx:]
# Avoid starting mid-turn when possible, except for proactive
# assistant deliveries that the user may be replying to.
@@ -352,17 +381,24 @@ class Session:
start_idx = max(0, len(self.messages) - max_messages)
if extend_to_user:
start_idx = next(
recovered_user = next(
(i for i in range(start_idx, -1, -1) if self.messages[i].get("role") == "user"),
start_idx,
None,
)
if recovered_user is not None:
start_idx = recovered_user
if start_idx > 0 and self.messages[start_idx - 1].get("_channel_delivery"):
start_idx -= 1
retained = self.messages[start_idx:]
# Prefer starting at a user turn when one exists within the retained window.
# Prefer starting at a user turn (or its preceding _channel_delivery) when one exists within the retained window.
first_user = next((i for i, m in enumerate(retained) if m.get("role") == "user"), None)
if first_user is not None:
retained = retained[first_user:]
if first_user > 0 and retained[first_user - 1].get("_channel_delivery"):
retained = retained[first_user - 1:]
else:
retained = retained[first_user:]
elif not extend_to_user:
# If the hard-capped tail is assistant/tool-only, anchor to the
# latest user in the full session and take a capped forward window.
@@ -1053,6 +1089,24 @@ class SessionManager:
self._remember(session)
return session
def get_or_create_transient(
self,
key: str,
*,
disabled_tools: Collection[str] = (),
) -> Session:
"""Return a fresh, non-persistent session without loading history."""
policy = SessionPolicy(
persist=False,
log_content=False,
disabled_tools=frozenset(disabled_tools),
)
session = self.get_cached(key)
if session is None or session.policy != policy:
session = Session(key=key, policy=policy)
self._remember(session)
return session
def _load(self, key: str) -> Session | None:
return self._store.load(key)
@@ -1060,12 +1114,11 @@ class SessionManager:
"""Attempt to recover a session from a corrupt JSONL file."""
return self._jsonl_store.repair(key, path=path)
@staticmethod
def _session_payload(session: Session) -> SessionPayload:
return JsonlSessionStore.session_payload(session)
def save(self, session: Session, *, fsync: bool = False) -> None:
"""Persist a session and retain it in the cache."""
if not session.policy.persist:
return
archiver = self._file_cap_archiver
if archiver is not None:
session.enforce_file_cap(
+6
View File
@@ -334,6 +334,12 @@ def clear_websocket_turn_if_current(
return False
def clear_websocket_turns(chat_id: str) -> None:
"""Forget every in-process turn projection for a discarded chat."""
_WEBSOCKET_ACTIVE_TURNS.pop(chat_id, None)
_sync_websocket_turn_projection(chat_id)
def build_bus_progress_callback(
bus: MessageBus,
msg: InboundMessage,
-11
View File
@@ -6,7 +6,6 @@ from typing import Any, Mapping
from nanobot.session.automation_turns import (
AutomationTurnSpec,
automation_history_overrides_for_spec,
automation_trigger,
)
@@ -50,13 +49,3 @@ def local_trigger_delivery_id(metadata: Mapping[str, Any] | None) -> str | None:
return None
value = trigger.get("delivery_id")
return value if isinstance(value, str) and value else None
def local_trigger_history_overrides(
metadata: Mapping[str, Any] | None,
) -> tuple[str | None, dict[str, Any]]:
"""Return session-history text/metadata overrides for a local trigger turn."""
return automation_history_overrides_for_spec(
metadata,
LOCAL_TRIGGER_AUTOMATION_SPEC,
)
-29
View File
@@ -11,35 +11,6 @@ from loguru import logger
from nanobot.utils.helpers import detect_image_mime
# Supported file extensions for text extraction
SUPPORTED_EXTENSIONS: set[str] = {
# Document formats
".pdf",
".docx",
".xlsx",
".pptx",
# Text formats
".txt",
".md",
".csv",
".json",
".xml",
".html",
".htm",
".log",
".yaml",
".yml",
".toml",
".ini",
".cfg",
# Image formats (for future OCR support)
".png",
".jpg",
".jpeg",
".gif",
".webp",
}
_MAX_TEXT_LENGTH = 200_000
_MAX_EXTRACT_FILE_SIZE = 50 * 1024 * 1024 # 50 MB
_MAX_OFFICE_ARCHIVE_MEMBERS = 10_000
-18
View File
@@ -274,24 +274,6 @@ def _text_line_count(text: str) -> int:
return line_count if last_was_newline else line_count + 1
def prepare_file_edit_tracker(
*,
call_id: str,
tool_name: str,
tool: Any,
workspace: Path | None,
params: dict[str, Any] | None,
) -> FileEditTracker | None:
trackers = prepare_file_edit_trackers(
call_id=call_id,
tool_name=tool_name,
tool=tool,
workspace=workspace,
params=params,
)
return trackers[0] if trackers else None
def prepare_file_edit_trackers(
*,
call_id: str,
+2 -56
View File
@@ -5,14 +5,13 @@ from __future__ import annotations
import io
import time
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import TYPE_CHECKING, Iterable, cast
from typing import TYPE_CHECKING, cast
from loguru import logger
if TYPE_CHECKING:
from dulwich.objects import Blob, Commit, ObjectID, Tree, TreeEntry
from dulwich.objects import Blob, Commit, ObjectID, Tree
from dulwich.refs import Ref
from dulwich.repo import Repo
@@ -45,25 +44,6 @@ class CommitInfo:
return f"{header}\n(no file changes)"
@dataclass
class LineAge:
"""Age of a single line based on git blame."""
age_days: int # days since last modification
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] = []
for (commit, _tree_entry), _line_bytes in annotated:
dt = datetime.fromtimestamp(commit.commit_time, tz=timezone.utc).date()
ages.append(LineAge(age_days=(now - dt).days))
return ages
class GitStore:
"""Git-backed version control for memory files."""
@@ -293,33 +273,6 @@ class GitStore:
except Exception as exc:
raise GitStoreError("Git log failed") from exc
def line_ages(self, file_path: str) -> list[LineAge]:
"""Compute the age of each line in a tracked file via git blame.
Returns one LineAge per line, in order.
Returns an empty list if the repo is not initialized or the file is
empty. Annotation failures raise :class:`GitStoreError`.
"""
if not self.is_initialized():
return []
target = self._workspace / file_path
if not target.exists() or target.stat().st_size == 0:
return []
try:
from dulwich import porcelain
annotated = porcelain.annotate(str(self._workspace), file_path)
except Exception as exc:
raise GitStoreError(f"Git line annotation failed for {file_path}") from exc
if not annotated:
return []
return _compute_line_ages(annotated)
def diff_commits(self, sha1: str, sha2: str) -> str:
"""Show diff between two commits."""
if not self.is_initialized():
@@ -461,13 +414,6 @@ class GitStore:
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."""
for c in self.log(max_entries=max_entries):
if c.sha.startswith(short_sha):
return c
return None
def show_commit_diff(
self,
short_sha: str,
-12
View File
@@ -351,18 +351,6 @@ def timestamp() -> str:
return datetime.now().isoformat()
def current_time_str(timezone: str | None = None) -> str:
"""Return the current time string."""
from zoneinfo import ZoneInfo
tz = ZoneInfo(timezone) if timezone else None
now = datetime.now(tz=tz) if tz else datetime.now().astimezone()
offset = now.strftime("%z")
offset_fmt = f"{offset[:3]}:{offset[3:]}" if len(offset) == 5 else offset
tz_name = timezone or (time.strftime("%Z") or "UTC")
return f"{now.strftime('%Y-%m-%d %H:%M (%A)')} ({tz_name}, UTC{offset_fmt})"
_UNSAFE_CHARS = re.compile(r'[<>:"/\\|?*]')
_TOOL_RESULT_PREVIEW_CHARS = 1200
_TOOL_RESULTS_DIR = ".nanobot/tool-results"
+4 -18
View File
@@ -3,14 +3,13 @@
Persisted subagent announcements mirror ``agent/subagent_announce.md``: header,
full ``Task:`` assignment (model context), ``Result:``, and a trailing model-only
``Summarize`` instruction. External channels (embedded WebUI, session previews)
should show only the header plus a truncated result body."""
should show only the header plus a truncated result body.
"""
from __future__ import annotations
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).
# Cap the Result section so session previews stay readable; full text remains on
# disk for LLM replay.
_SUBAGENT_CHANNEL_RESULT_MAX_CHARS = 800
@@ -44,16 +43,3 @@ def scrub_subagent_announce_body(content: str) -> str:
if header and body:
return f"{header}\n\n{body}"
return header or body or stripped
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(cast(object, msg), dict):
continue
if msg.get("injected_event") != "subagent_result":
continue
raw = msg.get("content")
if not isinstance(raw, str) or not raw.strip():
continue
msg["content"] = scrub_subagent_announce_body(raw)
+16 -6
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import asyncio
import re
import time
from pathlib import Path
from typing import Any, cast
from nanobot.apps.cli import CliAppError, CliAppManager, CliAppsRuntimeConfig
@@ -89,8 +90,8 @@ def _query_first(query: QueryParams, key: str) -> str | None:
return values[0] if values else None
def _manager() -> CliAppManager:
config = load_config()
def _manager(config_path: Path | None = None) -> CliAppManager:
config = load_config(config_path) if config_path is not None else load_config()
cli_cfg = config.tools.cli_apps
return CliAppManager(
workspace=config.workspace_path,
@@ -102,8 +103,12 @@ def _manager() -> CliAppManager:
)
async def cli_apps_payload(*, installed_only: bool = False) -> dict[str, Any]:
manager = _manager()
async def cli_apps_payload(
*,
installed_only: bool = False,
config_path: Path | None = None,
) -> dict[str, Any]:
manager = _manager(config_path) if config_path is not None else _manager()
if installed_only:
return manager.installed_payload()
payload = manager.payload(cache_only=True)
@@ -118,11 +123,16 @@ async def cli_apps_payload(*, installed_only: bool = False) -> dict[str, Any]:
return payload
def cli_apps_action(action: str, query: QueryParams) -> dict[str, Any]:
def cli_apps_action(
action: str,
query: QueryParams,
*,
config_path: Path | None = None,
) -> dict[str, Any]:
name = (_query_first(query, "name") or "").strip()
if not name:
raise CliAppError("missing CLI app name")
manager = _manager()
manager = _manager(config_path) if config_path is not None else _manager()
if action == "install":
return manager.install(name)
if action == "update":
+16
View File
@@ -8,9 +8,12 @@ from typing import TYPE_CHECKING, Any, Callable
from loguru import logger as default_logger
from nanobot.config.loader import get_config_path
from nanobot.webui.gateway_tokens import GatewayTokenStore
from nanobot.webui.ingress_policy import DEFAULT_WEBUI_INGRESS_POLICY, WebUIIngressPolicy
from nanobot.webui.media_gateway import WebUIMediaGateway
from nanobot.webui.settings_services import WebUISettingsServices
from nanobot.webui.temporary_chats import WebUITemporaryChats
from nanobot.webui.transcript import WebUITranscriptRecorder
from nanobot.webui.workspaces import WebUIWorkspaceController
from nanobot.webui.ws_http import GatewayHTTPHandler
@@ -28,11 +31,13 @@ class GatewayServices:
"""Explicit dependencies shared by WebSocket transport and HTTP routes."""
http: GatewayHTTPHandler
settings: WebUISettingsServices
tokens: GatewayTokenStore
media: WebUIMediaGateway
ingress: WebUIIngressPolicy
transcripts: WebUITranscriptRecorder
workspaces: WebUIWorkspaceController
temporary_chats: WebUITemporaryChats
session_manager: SessionManager | None
cron_service: CronService | None
local_trigger_store: LocalTriggerStore | None
@@ -48,6 +53,7 @@ def build_gateway_services(
static_dist_path: Path | None,
workspace_path: Path,
default_restrict_to_workspace: bool,
config_path: Path | None = None,
runtime_model_name: Callable[[], str | None] | None,
runtime_surface: str,
runtime_capabilities_overrides: dict[str, Any] | None,
@@ -61,6 +67,7 @@ def build_gateway_services(
skill_state_action: Callable[[set[str]], None] | None = None,
logger: Any = default_logger,
) -> GatewayServices:
settings = WebUISettingsServices.create(config_path or get_config_path())
tokens = GatewayTokenStore()
ingress = DEFAULT_WEBUI_INGRESS_POLICY
minimum_frame_bytes = ingress.minimum_full_policy_frame_bytes()
@@ -82,6 +89,12 @@ def build_gateway_services(
default_workspace=workspace_path,
default_restrict_to_workspace=default_restrict_to_workspace,
)
temporary_chats = WebUITemporaryChats(
bus=bus,
session_manager=session_manager,
workspaces=workspaces,
logger=logger,
)
http = GatewayHTTPHandler(
config=config,
session_manager=session_manager,
@@ -94,6 +107,7 @@ def build_gateway_services(
media=media,
ingress=ingress,
workspaces=workspaces,
settings=settings,
skills_workspace_path=workspace_path,
disabled_skills=disabled_skills,
cron_service=cron_service,
@@ -107,11 +121,13 @@ def build_gateway_services(
)
return GatewayServices(
http=http,
settings=settings,
tokens=tokens,
media=media,
ingress=ingress,
transcripts=transcripts,
workspaces=workspaces,
temporary_chats=temporary_chats,
session_manager=session_manager,
cron_service=cron_service,
local_trigger_store=local_trigger_store,
+86 -35
View File
@@ -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, cast
from typing import TYPE_CHECKING, Any, Literal, Mapping, cast
from nanobot.agent.tools.registry import ToolRegistry
from nanobot.apps.protocol import app_manifest, compact_dict
@@ -25,6 +25,9 @@ from nanobot.utils.helpers import ensure_dir
QueryParams = dict[str, list[str]]
if TYPE_CHECKING:
from nanobot.webui.settings_services import WebUISettingsConfig
_MCP_PRESET_NAME_RE = re.compile(r"^[a-z0-9][a-z0-9_-]{0,63}$", re.IGNORECASE)
_SECRET_QUERY_RE = re.compile(
r"([?&](?:[^=&]*(?:api[_-]?key|token|secret|password|bearer)[^=&]*)=)[^&#\s]+",
@@ -841,8 +844,9 @@ def mcp_presets_payload(
*,
last_action: dict[str, Any] | None = None,
tool_preview: Mapping[str, list[str]] | None = None,
config_path: Path | None = None,
) -> dict[str, Any]:
config = load_config()
config = load_config(config_path) if config_path is not None else load_config()
known = _known_preset_names()
preset_rows = [
_preset_payload(preset, config.tools.mcp_servers)
@@ -928,7 +932,11 @@ async def _close_mcp_stacks(stacks: Mapping[str, Any]) -> None:
await stack.aclose()
async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]:
async def mcp_presets_test_action(
query: QueryParams,
*,
config_path: Path | None = None,
) -> dict[str, Any]:
"""Connect to an enabled MCP preset and report its tool surface."""
from nanobot.agent.tools.mcp import connect_mcp_servers
@@ -941,16 +949,22 @@ async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]:
display_name = _display_name_for(name, preset)
try:
config = resolve_config_env_vars(load_config())
config = resolve_config_env_vars(
load_config(config_path) if config_path is not None else load_config(),
config_path=config_path,
)
except ValueError as exc:
return mcp_presets_payload(last_action={
"ok": False,
"message": _scrub_test_error(str(exc)),
"error": _scrub_test_error(str(exc)),
"tool_count": 0,
"tool_names": [],
"checked_at": _checked_at(),
})
return mcp_presets_payload(
last_action={
"ok": False,
"message": _scrub_test_error(str(exc)),
"error": _scrub_test_error(str(exc)),
"tool_count": 0,
"tool_names": [],
"checked_at": _checked_at(),
},
config_path=config_path,
)
cfg = config.tools.mcp_servers.get(name)
if cfg is None:
@@ -968,7 +982,7 @@ async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]:
"tool_names": [],
"checked_at": _checked_at(),
}
return mcp_presets_payload(last_action=last_action)
return mcp_presets_payload(last_action=last_action, config_path=config_path)
if cfg.command and not _command_available(cfg.command):
last_action = {
@@ -979,7 +993,7 @@ async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]:
"tool_names": [],
"checked_at": _checked_at(),
}
return mcp_presets_payload(last_action=last_action)
return mcp_presets_payload(last_action=last_action, config_path=config_path)
registry = ToolRegistry()
stacks: dict[str, Any] = {}
@@ -1040,7 +1054,11 @@ async def mcp_presets_test_action(query: QueryParams) -> dict[str, Any]:
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)
return mcp_presets_payload(
last_action=last_action,
tool_preview=preview,
config_path=config_path,
)
def _parse_json_value(raw: str | None, *, fallback: Any) -> Any:
@@ -1221,24 +1239,35 @@ def _import_mcp_servers(raw_json: str | None) -> dict[str, MCPServerConfig]:
return out
def custom_mcp_action(action: str, query: QueryParams) -> dict[str, Any]:
config = load_config()
def custom_mcp_action(
action: str,
query: QueryParams,
*,
config_path: Path | None = None,
) -> dict[str, Any]:
config = load_config(config_path) if config_path is not None else load_config()
if action == "custom":
name, cfg = _custom_server_from_query(query)
config.tools.mcp_servers[name] = cfg
save_config(config)
payload = mcp_presets_payload(last_action=_server_action_message(action, name))
save_config(config, config_path)
payload = mcp_presets_payload(
last_action=_server_action_message(action, name),
config_path=config_path,
)
payload["requires_restart"] = True
return payload
if action in {"import", "import-cursor"}:
servers = _import_mcp_servers(_query_first(query, "config"))
config.tools.mcp_servers.update(servers)
save_config(config)
payload = mcp_presets_payload(last_action={
"ok": True,
"message": f"Imported {len(servers)} MCP server(s).",
})
save_config(config, config_path)
payload = mcp_presets_payload(
last_action={
"ok": True,
"message": f"Imported {len(servers)} MCP server(s).",
},
config_path=config_path,
)
payload["requires_restart"] = True
return payload
@@ -1249,29 +1278,40 @@ def custom_mcp_action(action: str, query: QueryParams) -> dict[str, Any]:
raise McpPresetError("unknown MCP server", status=404)
cfg.enabled_tools = _parse_enabled_tools(_query_first(query, "enabled_tools"))
config.tools.mcp_servers[name] = cfg
save_config(config)
payload = mcp_presets_payload(last_action=_server_action_message(action, name))
save_config(config, config_path)
payload = mcp_presets_payload(
last_action=_server_action_message(action, name),
config_path=config_path,
)
payload["requires_restart"] = True
return payload
raise McpPresetError(f"unknown MCP action '{action}'", status=404)
def mcp_presets_action(action: str, query: QueryParams) -> dict[str, Any]:
def mcp_presets_action(
action: str,
query: QueryParams,
*,
config_path: Path | None = None,
) -> dict[str, Any]:
name = (_query_first(query, "name") or "").strip()
if not name:
raise McpPresetError("missing MCP preset name")
preset = _preset_by_name_optional(name)
config = load_config()
config = load_config(config_path) if config_path is not None else load_config()
existing = config.tools.mcp_servers.get(name)
if action == "enable":
if preset is None:
raise McpPresetError("unknown MCP preset", status=404)
config.tools.mcp_servers[preset.name] = _materialize_server(preset, query, existing)
save_config(config)
payload = mcp_presets_payload(last_action=_action_message(action, preset))
save_config(config, config_path)
payload = mcp_presets_payload(
last_action=_action_message(action, preset),
config_path=config_path,
)
payload["requires_restart"] = True
return payload
@@ -1287,7 +1327,7 @@ def mcp_presets_action(action: str, query: QueryParams) -> dict[str, Any]:
except OSError as exc:
cleanup_error = str(exc)
del config.tools.mcp_servers[name]
save_config(config)
save_config(config, config_path)
last_action = (
_action_message(action, preset)
if preset is not None
@@ -1303,7 +1343,10 @@ def mcp_presets_action(action: str, query: QueryParams) -> dict[str, Any]:
f"{last_action['message']} Could not remove managed runtime files: {cleanup_error}"
)
last_action["verification_failed"] = ["managed_paths_absent"]
payload = mcp_presets_payload(last_action=last_action)
payload = mcp_presets_payload(
last_action=last_action,
config_path=config_path,
)
payload["requires_restart"] = True
return payload
@@ -1339,13 +1382,21 @@ async def mcp_presets_settings_action(
query: QueryParams,
*,
reload_mcp: McpReload | None = None,
config: WebUISettingsConfig | None = None,
) -> dict[str, Any]:
"""Run a WebUI MCP preset action and hot-reload the agent when config changes."""
config_path = config.path if config is not None else None
if action is None:
return mcp_presets_payload()
return mcp_presets_payload(config_path=config_path)
if action == "test":
return await mcp_presets_test_action(query)
if action in _CUSTOM_ACTIONS:
return await mcp_presets_test_action(query, config_path=config_path)
if config is not None:
operation = custom_mcp_action if action in _CUSTOM_ACTIONS else mcp_presets_action
payload = await asyncio.to_thread(
config.run_serialized,
lambda path: operation(action, query, config_path=path),
)
elif action in _CUSTOM_ACTIONS:
payload = await asyncio.to_thread(custom_mcp_action, action, query)
else:
payload = await asyncio.to_thread(mcp_presets_action, action, query)
+1 -33
View File
@@ -13,7 +13,7 @@ import shutil
import uuid
from collections.abc import Callable
from pathlib import Path
from typing import Any, cast
from typing import Any
from websockets.http11 import Request as WsRequest
from websockets.http11 import Response
@@ -32,7 +32,6 @@ from nanobot.webui.http_utils import (
MediaDirProvider = Callable[[str | None], Path]
SignedMediaPath = Callable[[Path], dict[str, str] | None]
SignedMediaUrl = Callable[[Path], str | None]
def b64url_encode(data: bytes) -> str:
@@ -190,37 +189,6 @@ def signed_media_attachments(
return out
def attach_signed_media_urls(
payload: dict[str, Any],
*,
sign_path: SignedMediaUrl,
) -> None:
"""Replace raw media path lists in a WebUI session payload with signed URLs."""
messages = payload.get("messages")
if not isinstance(messages, list):
return
raw_messages = cast(list[Any], messages)
for msg in raw_messages:
if not isinstance(msg, dict):
continue
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_entries:
if not isinstance(entry, str) or not entry:
continue
signed = sign_path(Path(entry))
if signed is None:
continue
urls.append({"url": signed, "name": Path(entry).name})
if urls:
message["media_urls"] = urls
message.pop("media", None)
def serve_signed_media(
sig: str,
payload: str,
-12
View File
@@ -17,9 +17,7 @@ from nanobot.webui.attachment_ingress import (
)
from nanobot.webui.ingress_policy import AttachmentIngressLimits
from nanobot.webui.media_api import (
attach_signed_media_urls,
serve_signed_media,
sign_media_path,
sign_or_stage_media_path,
signed_media_attachments,
)
@@ -72,13 +70,6 @@ class WebUIMediaGateway:
media_dir=self._media_dir,
)
def sign_media_path(self, abs_path: Path) -> str | None:
return sign_media_path(
abs_path,
secret=self.secret,
media_dir=self._media_dir,
)
def sign_or_stage_media_path(self, path: Path) -> dict[str, str] | None:
return sign_or_stage_media_path(
path,
@@ -99,9 +90,6 @@ class WebUIMediaGateway:
sign_path=self.sign_or_stage_media_path,
)
def augment_media_urls(self, payload: dict[str, Any]) -> None:
attach_signed_media_urls(payload, sign_path=self.sign_media_path)
def augment_transcript_media(self, paths: list[str]) -> list[dict[str, Any]]:
return signed_media_attachments(
paths,
+20 -4
View File
@@ -1,6 +1,7 @@
"""Nanobot optional feature helpers for WebUI Settings."""
from __future__ import annotations
from pathlib import Path
from typing import Any
from nanobot.channels.registry import load_channel_plugin
@@ -15,8 +16,13 @@ from nanobot.webui.http_utils import query_first
QueryParams = dict[str, list[str]]
def nanobot_features_payload() -> dict[str, Any]:
return optional_features_payload()
def nanobot_features_payload(*, config_path: Path | None = None) -> dict[str, Any]:
if config_path is None:
return optional_features_payload()
from nanobot.config.loader import load_config
return optional_features_payload(config=load_config(config_path))
def nanobot_feature_instance_target(query: QueryParams) -> str | None:
@@ -32,13 +38,19 @@ def nanobot_features_action(
query: QueryParams,
*,
allow_install: bool = True,
config_path: Path | None = None,
) -> dict[str, Any]:
name = (query_first(query, "name") or "").strip()
instance_id = nanobot_feature_instance_target(query)
if not name:
raise OptionalFeatureError("missing feature name")
if action == "enable":
return enable_optional_feature(name, allow_install=allow_install, instance_id=instance_id)
return enable_optional_feature(
name,
config_path=config_path,
allow_install=allow_install,
instance_id=instance_id,
)
if action == "disable":
try:
plugin = load_channel_plugin(name)
@@ -50,5 +62,9 @@ def nanobot_features_action(
f"Use `nanobot plugins disable {name}` from a terminal if you need to disable it.",
status=400,
)
return disable_optional_feature(name, instance_id=instance_id)
return disable_optional_feature(
name,
config_path=config_path,
instance_id=instance_id,
)
raise OptionalFeatureError(f"unknown feature action '{action}'", status=404)
File diff suppressed because it is too large Load Diff
+804
View File
@@ -0,0 +1,804 @@
"""Capability settings domain logic for Web, media, network, and API features."""
from __future__ import annotations
import asyncio
import os
import re
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Awaitable, Callable, TypedDict
from nanobot.agent.tools.web import SEARCH_PROVIDER_OPTIONS
from nanobot.api.runtime import ApiRuntime, ApiStartOptions
from nanobot.audio.transcription import resolve_transcription_config
from nanobot.audio.transcription_registry import (
resolve_transcription_provider,
transcription_provider_names,
)
from nanobot.config.schema import Config
from nanobot.optional_features import (
OptionalFeatureError,
extra_installed,
optional_dependency_groups,
)
from nanobot.providers.image_generation import (
get_image_gen_provider,
image_gen_provider_names,
)
from nanobot.providers.registry import find_by_name
from nanobot.security.network import is_loopback_host
from nanobot.webui.settings_contracts import (
QueryParams,
SettingsRequest,
SettingsRouteResult,
WebUISettingsError,
parse_bool,
query_first,
query_first_alias,
)
from nanobot.webui.settings_models import (
OAuthStatusReader,
mask_secret_hint,
provider_configured_for_settings,
)
from nanobot.webui.workspaces import (
read_webui_default_access_mode,
)
if TYPE_CHECKING:
from nanobot.webui.settings_services import WebUISettingsServices
SettingsOperation = Callable[..., dict[str, Any]]
@dataclass(frozen=True)
class CapabilitySettingsOperations:
update_web_search: SettingsOperation
update_api: SettingsOperation
update_image: SettingsOperation
update_transcription: SettingsOperation
update_network: SettingsOperation
nanobot_features_action: SettingsOperation
api_runtime: Callable[[], ApiRuntime]
reload_image: Callable[[], Awaitable[dict[str, Any]]]
class CapabilitySettingsPayload(TypedDict):
web_search: dict[str, Any]
web: dict[str, Any]
api: dict[str, Any]
observability: dict[str, Any]
image_generation: dict[str, Any]
transcription: dict[str, Any]
_WEB_SEARCH_PROVIDER_OPTIONS = SEARCH_PROVIDER_OPTIONS
_WEB_SEARCH_PROVIDER_BY_NAME = {
provider["name"]: provider for provider in _WEB_SEARCH_PROVIDER_OPTIONS
}
_IMAGE_GENERATION_ASPECT_RATIOS = {
"1:1",
"3:4",
"9:16",
"4:3",
"16:9",
"3:2",
"2:3",
"21:9",
}
def _image_generation_provider_rows(
config: Config,
*,
oauth_status: OAuthStatusReader,
) -> list[dict[str, Any]]:
rows: list[dict[str, Any]] = []
for name in image_gen_provider_names():
image_provider = get_image_gen_provider(name)
spec = find_by_name(name)
provider_config = getattr(config.providers, name, None)
configured = (
provider_configured_for_settings(spec, provider_config, oauth_status)
if spec is not None and provider_config is not None
else bool(getattr(provider_config, "api_key", None))
)
rows.append(
{
"name": name,
"label": spec.label if spec is not None else name,
"configured": configured,
"auth_type": "oauth" if spec is not None and spec.is_oauth else "api_key",
"api_key_hint": mask_secret_hint(getattr(provider_config, "api_key", None)),
"api_base": getattr(provider_config, "api_base", None),
"default_api_base": (
spec.default_api_base if spec and spec.default_api_base else None
),
"models": list(image_provider.model_options) if image_provider else [],
"default_model": (
image_provider.model_options[0]
if image_provider and image_provider.model_options
else None
),
}
)
return rows
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)
provider_config = getattr(config.providers, name, None)
rows.append(
{
"name": name,
"label": spec.label if spec is not None else name,
"configured": bool(getattr(provider_config, "api_key", None)),
"api_key_hint": mask_secret_hint(getattr(provider_config, "api_key", None)),
"api_base": getattr(provider_config, "api_base", None),
"default_api_base": (
spec.default_api_base if spec and spec.default_api_base else None
),
}
)
return rows
def capability_settings_payload(
config: Config,
*,
oauth_status: OAuthStatusReader,
) -> CapabilitySettingsPayload:
search_config = config.tools.web.search
image_config = config.tools.image_generation
transcription = resolve_transcription_config(config)
search_provider = (
search_config.provider
if search_config.provider in _WEB_SEARCH_PROVIDER_BY_NAME
else "duckduckgo"
)
image_providers = _image_generation_provider_rows(config, oauth_status=oauth_status)
selected_image_provider = next(
(
provider
for provider in image_providers
if provider["name"] == image_config.provider
),
None,
)
return {
"web_search": {
"provider": search_provider,
"api_key_hint": mask_secret_hint(search_config.api_key),
"base_url": search_config.base_url or None,
"max_results": search_config.max_results,
"timeout": search_config.timeout,
"providers": list(_WEB_SEARCH_PROVIDER_OPTIONS),
},
"web": {
"enable": config.tools.web.enable,
"proxy": config.tools.web.proxy,
"user_agent": config.tools.web.user_agent,
"search": {
"max_results": search_config.max_results,
"timeout": search_config.timeout,
},
"fetch": {
"use_jina_reader": config.tools.web.fetch.use_jina_reader,
},
},
"api": {
"host": config.api.host,
"port": config.api.port,
"timeout": config.api.timeout,
"api_key_hint": mask_secret_hint(config.api.api_key),
},
"observability": {
"provider": "langfuse",
"configured": bool(
os.environ.get("LANGFUSE_SECRET_KEY")
and os.environ.get("LANGFUSE_PUBLIC_KEY")
),
"base_url": os.environ.get("LANGFUSE_BASE_URL")
or "https://cloud.langfuse.com",
},
"image_generation": {
"enabled": image_config.enabled,
"provider": image_config.provider,
"provider_configured": bool(
selected_image_provider and selected_image_provider["configured"]
),
"model": image_config.model,
"default_aspect_ratio": image_config.default_aspect_ratio,
"default_image_size": image_config.default_image_size,
"max_images_per_turn": image_config.max_images_per_turn,
"save_dir": image_config.save_dir,
"providers": image_providers,
},
"transcription": {
"enabled": transcription.enabled,
"provider": transcription.provider,
"provider_configured": transcription.configured,
"model": transcription.model,
"language": transcription.language,
"max_duration_sec": transcription.max_duration_sec,
"max_upload_mb": transcription.max_upload_mb,
"providers": _transcription_provider_rows(config),
},
}
def update_network_safety_settings(
config: Config,
query: QueryParams,
) -> tuple[bool, str | None]:
raw_allow = (
query_first_alias(
query,
"webui_allow_local_service_access",
"webuiAllowLocalServiceAccess",
)
or query_first_alias(
query,
"allow_local_preview_access",
"allowLocalPreviewAccess",
)
)
raw_default_access_mode = query_first_alias(
query,
"webui_default_access_mode",
"webuiDefaultAccessMode",
)
if raw_allow is None and raw_default_access_mode is None:
raise WebUISettingsError(
"webui_allow_local_service_access or webui_default_access_mode is required"
)
changed = False
if raw_allow is not None:
allow_local = parse_bool(raw_allow, "webui_allow_local_service_access")
if config.tools.webui_allow_local_service_access != allow_local:
config.tools.webui_allow_local_service_access = allow_local
changed = True
default_access_mode: str | None = None
if raw_default_access_mode is not None:
default_access_mode = raw_default_access_mode.strip().lower()
if default_access_mode == "restricted":
default_access_mode = "default"
if default_access_mode not in {"default", "full"}:
raise WebUISettingsError(
"webui_default_access_mode must be default or full"
)
return changed, default_access_mode
def update_web_search_settings(config: Config, query: QueryParams) -> tuple[bool, bool]:
provider_name = (query_first(query, "provider") or "").strip().lower()
provider_option = _WEB_SEARCH_PROVIDER_BY_NAME.get(provider_name)
if provider_option is None:
raise WebUISettingsError("unknown web search provider")
search_config = config.tools.web.search
web_config = config.tools.web
previous_provider = search_config.provider
changed = False
restart_required = False
def set_search_value(attr: str, value: object) -> None:
nonlocal changed
if getattr(search_config, attr) != value:
setattr(search_config, attr, value)
changed = True
def set_fetch_value(attr: str, value: object) -> None:
nonlocal changed
if getattr(web_config.fetch, attr) != value:
setattr(web_config.fetch, attr, value)
changed = True
if search_config.provider != provider_name:
search_config.provider = provider_name
changed = True
credential = provider_option["credential"]
if credential == "none":
set_search_value("api_key", "")
set_search_value("base_url", "")
elif credential == "base_url":
base_url = query_first_alias(query, "base_url", "baseUrl")
base_url = base_url.strip() if base_url is not None else None
if not base_url and previous_provider == provider_name and search_config.base_url:
base_url = search_config.base_url
if not base_url:
raise WebUISettingsError("base_url is required")
set_search_value("base_url", base_url)
set_search_value("api_key", "")
elif credential in {"api_key", "optional_api_key"}:
raw_api_key = query_first_alias(query, "api_key", "apiKey")
api_key = raw_api_key.strip() if raw_api_key is not None else None
if api_key is None and previous_provider == provider_name and search_config.api_key:
api_key = search_config.api_key
if credential == "api_key" and not api_key:
raise WebUISettingsError("api_key is required")
set_search_value("api_key", api_key or "")
set_search_value("base_url", "")
else:
raise WebUISettingsError("unknown web search credential type")
max_results = query_first_alias(query, "max_results", "maxResults")
if max_results is not None:
try:
parsed = int(max_results)
except ValueError:
raise WebUISettingsError("max_results must be an integer") from None
if parsed < 1 or parsed > 10:
raise WebUISettingsError("max_results must be between 1 and 10")
set_search_value("max_results", parsed)
timeout = query_first(query, "timeout")
if timeout is not None:
try:
parsed_timeout = int(timeout)
except ValueError:
raise WebUISettingsError("timeout must be an integer") from None
if parsed_timeout < 1 or parsed_timeout > 120:
raise WebUISettingsError("timeout must be between 1 and 120")
set_search_value("timeout", parsed_timeout)
use_jina_reader = query_first_alias(query, "use_jina_reader", "useJinaReader")
if use_jina_reader is not None:
previous_jina_reader = web_config.fetch.use_jina_reader
set_fetch_value("use_jina_reader", parse_bool(use_jina_reader, "use_jina_reader"))
if web_config.fetch.use_jina_reader != previous_jina_reader:
restart_required = True
return changed, restart_required
def update_api_settings(config: Config, query: QueryParams) -> None:
"""Update the managed OpenAI-compatible API configuration."""
api = config.api
host = query_first(query, "host")
if host is not None:
host = host.strip()
if not host:
raise WebUISettingsError("host is required")
api.host = host
port = query_first(query, "port")
if port is not None:
try:
parsed_port = int(port)
except ValueError:
raise WebUISettingsError("port must be an integer") from None
if parsed_port < 1 or parsed_port > 65535:
raise WebUISettingsError("port must be between 1 and 65535")
api.port = parsed_port
timeout = query_first(query, "timeout")
if timeout is not None:
try:
parsed_timeout = float(timeout)
except ValueError:
raise WebUISettingsError("timeout must be a number") from None
if parsed_timeout < 1 or parsed_timeout > 3600:
raise WebUISettingsError("timeout must be between 1 and 3600")
api.timeout = parsed_timeout
api_key = query_first_alias(query, "api_key", "apiKey")
if api_key is not None:
api.api_key = api_key.strip()
if not is_loopback_host(api.host) and not api.api_key.strip():
raise WebUISettingsError(
"an API key is required when the API is available on the network"
)
def update_image_generation_settings(
config: Config,
query: QueryParams,
*,
oauth_status: OAuthStatusReader,
) -> bool:
image_config = config.tools.image_generation
changed = False
provider_name = query_first(query, "provider")
if provider_name is not None:
provider_name = provider_name.strip().lower()
if not provider_name:
raise WebUISettingsError("image generation provider is required")
if get_image_gen_provider(provider_name) is None:
raise WebUISettingsError("unknown image generation provider")
if image_config.provider != provider_name:
image_config.provider = provider_name
changed = True
enabled = query_first(query, "enabled")
if enabled is not None:
parsed_enabled = parse_bool(enabled, "enabled")
if image_config.enabled != parsed_enabled:
image_config.enabled = parsed_enabled
changed = True
model = query_first(query, "model")
if model is not None:
model = model.strip()
if not model:
raise WebUISettingsError("image generation model is required")
if len(model) > 200:
raise WebUISettingsError("image generation model is too long")
if image_config.model != model:
image_config.model = model
changed = True
default_aspect_ratio = query_first_alias(
query,
"default_aspect_ratio",
"defaultAspectRatio",
)
if default_aspect_ratio is not None:
default_aspect_ratio = default_aspect_ratio.strip()
if default_aspect_ratio not in _IMAGE_GENERATION_ASPECT_RATIOS:
raise WebUISettingsError("unsupported image generation aspect ratio")
if image_config.default_aspect_ratio != default_aspect_ratio:
image_config.default_aspect_ratio = default_aspect_ratio
changed = True
default_image_size = query_first_alias(
query,
"default_image_size",
"defaultImageSize",
)
if default_image_size is not None:
default_image_size = default_image_size.strip()
if not default_image_size:
raise WebUISettingsError("default image size is required")
if len(default_image_size) > 32 or not all(
char.isascii() and (char.isalnum() or char in {"x", "X", ":", "-", "_"})
for char in default_image_size
):
raise WebUISettingsError("unsupported image generation size")
if image_config.default_image_size != default_image_size:
image_config.default_image_size = default_image_size
changed = True
max_images_per_turn = query_first_alias(
query,
"max_images_per_turn",
"maxImagesPerTurn",
)
if max_images_per_turn is not None:
try:
parsed_max = int(max_images_per_turn)
except ValueError:
raise WebUISettingsError("max_images_per_turn must be an integer") from None
if parsed_max < 1 or parsed_max > 8:
raise WebUISettingsError("max_images_per_turn must be between 1 and 8")
if image_config.max_images_per_turn != parsed_max:
image_config.max_images_per_turn = parsed_max
changed = True
if image_config.enabled:
selected_provider = next(
(
provider
for provider in _image_generation_provider_rows(
config,
oauth_status=oauth_status,
)
if provider["name"] == image_config.provider
),
None,
)
if not selected_provider or not selected_provider["configured"]:
raise WebUISettingsError("image generation provider is not configured")
return changed
def update_transcription_settings(config: Config, query: QueryParams) -> bool:
transcription = config.transcription
changed = False
enabled = query_first(query, "enabled")
if enabled is not None:
parsed_enabled = parse_bool(enabled, "enabled")
if transcription.enabled != parsed_enabled:
transcription.enabled = parsed_enabled
changed = True
provider = query_first(query, "provider")
if provider is not None:
provider = provider.strip().lower()
provider_spec = resolve_transcription_provider(provider)
if provider_spec is None:
raise WebUISettingsError("unknown transcription provider")
provider = provider_spec.name
if transcription.provider != provider:
transcription.provider = provider
changed = True
model = query_first(query, "model")
if model is not None:
model = model.strip() or None
if model is not None and len(model) > 200:
raise WebUISettingsError("transcription model is too long")
if transcription.model != model:
transcription.model = model
changed = True
language = query_first(query, "language")
if language is not None:
language = language.strip().lower() or None
if language is not None and not re.fullmatch(r"[a-z]{2,3}", language):
raise WebUISettingsError(
"transcription language must be 2-3 lowercase letters"
)
if transcription.language != language:
transcription.language = language
changed = True
max_duration_sec = query_first_alias(query, "max_duration_sec", "maxDurationSec")
if max_duration_sec is not None:
try:
parsed_duration = int(max_duration_sec)
except ValueError:
raise WebUISettingsError("max_duration_sec must be an integer") from None
if parsed_duration < 1 or parsed_duration > 600:
raise WebUISettingsError("max_duration_sec must be between 1 and 600")
if transcription.max_duration_sec != parsed_duration:
transcription.max_duration_sec = parsed_duration
changed = True
max_upload_mb = query_first_alias(query, "max_upload_mb", "maxUploadMb")
if max_upload_mb is not None:
try:
parsed_upload = int(max_upload_mb)
except ValueError:
raise WebUISettingsError("max_upload_mb must be an integer") from None
if parsed_upload < 1 or parsed_upload > 100:
raise WebUISettingsError("max_upload_mb must be between 1 and 100")
if transcription.max_upload_mb != parsed_upload:
transcription.max_upload_mb = parsed_upload
changed = True
return changed
def network_safety_payload(config: Config) -> dict[str, Any]:
"""Return the network-related fields embedded in the advanced DTO."""
return {
"webui_allow_local_service_access": config.tools.webui_allow_local_service_access,
"allow_local_preview_access": config.tools.webui_allow_local_service_access,
"webui_default_access_mode": read_webui_default_access_mode(),
"private_service_protection_enabled": True,
"ssrf_whitelist_count": len(config.tools.ssrf_whitelist),
}
def masked_api_secret(value: str) -> str | None:
value = value.strip()
if not value:
return None
return f"{value[:3]}...{value[-4:]}" if len(value) > 8 else "configured"
def api_runtime_message(message: str) -> str:
known = {
"api_exited_during_startup": "API server exited during startup. Check its log for details.",
"api_stop_timeout": "API server did not stop in time.",
"api_state_stale": "API server state was stale; try starting it again.",
}
if message in known:
return known[message]
if message.startswith("api_"):
return f"API server {message.removeprefix('api_').replace('_', ' ')}"
return message.replace("_", " ")
def api_service_payload(
settings: WebUISettingsServices,
runtime: ApiRuntime,
*,
last_action: str | None = None,
) -> dict[str, Any]:
config = settings.config.load()
status = runtime.status()
extras = optional_dependency_groups()
connect_host = (
"127.0.0.1" if config.api.host in {"0.0.0.0", "::"} else config.api.host
)
payload = {
"installed": extra_installed("api", extras.get("api")),
"running": status.running,
"managed": status.running,
"host": config.api.host,
"port": config.api.port,
"timeout": config.api.timeout,
"api_key_hint": masked_api_secret(config.api.api_key),
"endpoint": f"http://{connect_host}:{config.api.port}/v1",
"command": "nanobot serve",
"log_path": str(status.log_path),
}
if last_action:
payload["last_action"] = last_action
return payload
class CapabilitySettingsHandler:
"""Handle capability commands after transport authentication and decoding."""
def __init__(self, settings: WebUISettingsServices, logger: Any) -> None:
self.settings = settings
self.logger = logger
async def handle(
self,
action: str,
request: SettingsRequest,
operations: CapabilitySettingsOperations,
) -> SettingsRouteResult:
if action == "api-status":
return SettingsRouteResult.success(
api_service_payload(self.settings, operations.api_runtime())
)
if action == "api-start":
return await self._start_api(request, operations)
if action == "api-stop":
return await self._stop_api(operations)
mutation = {
"web-search-update": (
operations.update_web_search,
"browser",
False,
),
"transcription-update": (
operations.update_transcription,
None,
False,
),
"network-update": (
operations.update_network,
"runtime",
False,
),
"image-update": (
operations.update_image,
"image",
True,
),
}.get(action)
if mutation is None:
return SettingsRouteResult.failure(404, "unknown settings action")
operation, section, apply_image_reload = mutation
try:
payload = self.settings.mutate(operation, request.query)
except WebUISettingsError as exc:
return SettingsRouteResult.failure(exc.status, exc.message)
if apply_image_reload:
payload, image_restart_cleared = await self.apply_image_runtime_change(
payload,
operations.reload_image,
)
else:
image_restart_cleared = False
return SettingsRouteResult.success(
payload,
decorate_restart=True,
restart_section=section,
clear_restart_section=("image" if image_restart_cleared else None),
)
async def apply_image_runtime_change(
self,
payload: dict[str, Any],
reload_image: Callable[[], Awaitable[dict[str, Any]]],
) -> tuple[dict[str, Any], bool]:
"""Hot-apply image settings, preserving restart fallback on failure."""
if not payload.get("requires_restart"):
return payload, False
try:
result = await reload_image()
except Exception:
self.logger.exception("failed to hot-reload image generation settings")
return payload, False
applied = bool(result.get("ok")) and not result.get("requires_restart")
updated = dict(payload)
updated["requires_restart"] = not applied
if not applied:
self.logger.warning(
"image generation settings were saved but require restart: {}",
result.get("message") or "hot reload failed",
)
return updated, applied
async def _start_api(
self,
request: SettingsRequest,
operations: CapabilitySettingsOperations,
) -> SettingsRouteResult:
api_key = (request.payload or {}).get("api_key")
if api_key is not None and not isinstance(api_key, str):
return SettingsRouteResult.failure(
400,
"API service API key must be a string",
)
try:
await asyncio.to_thread(
self.settings.mutate,
operations.nanobot_features_action,
"enable",
{"name": ["api"]},
allow_install=self._allow_feature_package_install(request),
)
self.settings.mutate(operations.update_api, request.query)
config = self.settings.config.load()
runtime = operations.api_runtime()
options = ApiStartOptions(
host=config.api.host,
port=config.api.port,
workspace=str(config.workspace_path),
config_path=str(self.settings.config.path),
)
current = runtime.status()
result = await asyncio.to_thread(
runtime.restart if current.running else runtime.start_background,
options,
)
if not result.ok:
return SettingsRouteResult.failure(
500,
api_runtime_message(result.message),
)
except (WebUISettingsError, OptionalFeatureError) as exc:
return SettingsRouteResult.failure(
getattr(exc, "status", 400),
getattr(exc, "message", str(exc)),
)
except Exception as exc:
self.logger.exception("failed to start managed API service")
return SettingsRouteResult.failure(500, str(exc))
return SettingsRouteResult.success(
api_service_payload(
self.settings,
operations.api_runtime(),
last_action="started",
)
)
async def _stop_api(
self,
operations: CapabilitySettingsOperations,
) -> SettingsRouteResult:
runtime = operations.api_runtime()
try:
result = await asyncio.to_thread(runtime.stop)
except Exception as exc:
self.logger.exception("failed to stop managed API service")
return SettingsRouteResult.failure(500, str(exc))
if not result.ok and result.message != "api_not_running":
return SettingsRouteResult.failure(
500,
api_runtime_message(result.message),
)
return SettingsRouteResult.success(
api_service_payload(
self.settings,
operations.api_runtime(),
last_action="stopped",
)
)
def _allow_feature_package_install(self, request: SettingsRequest) -> bool:
if request.local_browser:
return True
try:
return bool(
self.settings.config.load().tools.webui_allow_remote_package_install
)
except Exception:
self.logger.exception("failed to load remote package install policy")
return False
+82
View File
@@ -0,0 +1,82 @@
"""Stable request and error contracts shared by WebUI settings domains."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
QueryParams = dict[str, list[str]]
@dataclass(frozen=True)
class SettingsRequest:
"""Transport-neutral input decoded by the settings route facade."""
query: QueryParams
payload: dict[str, Any] | None = None
local_browser: bool = False
@dataclass(frozen=True)
class SettingsRouteResult:
"""Transport-neutral result returned by a settings domain handler."""
payload: dict[str, Any] | None = None
status: int = 200
error: str | None = None
decorate_restart: bool = False
restart_section: str | None = None
clear_restart_section: str | None = None
restart_payload_key: str | None = None
@classmethod
def success(
cls,
payload: dict[str, Any],
*,
decorate_restart: bool = False,
restart_section: str | None = None,
clear_restart_section: str | None = None,
restart_payload_key: str | None = None,
) -> SettingsRouteResult:
return cls(
payload=payload,
decorate_restart=decorate_restart,
restart_section=restart_section,
clear_restart_section=clear_restart_section,
restart_payload_key=restart_payload_key,
)
@classmethod
def failure(cls, status: int, error: str) -> SettingsRouteResult:
return cls(status=status, error=error)
class WebUISettingsError(ValueError):
"""User-facing settings validation failure."""
def __init__(self, message: str, *, status: int = 400) -> None:
super().__init__(message)
self.message = message
self.status = status
def query_first(query: QueryParams, key: str) -> str | None:
values = query.get(key)
return values[0] if values else None
def query_first_alias(query: QueryParams, snake: str, camel: str) -> str | None:
value = query_first(query, snake)
return query_first(query, camel) if value is None else value
def query_has_alias(query: QueryParams, snake: str, camel: str) -> bool:
return snake in query or camel in query
def parse_bool(value: str, field: str) -> bool:
normalized = value.strip().lower()
if normalized not in {"1", "0", "true", "false", "yes", "no"}:
raise WebUISettingsError(f"{field} must be boolean")
return normalized in {"1", "true", "yes"}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+148
View File
@@ -0,0 +1,148 @@
"""Gateway-owned state for the WebUI settings surface."""
from __future__ import annotations
import threading
from collections.abc import Callable
from dataclasses import dataclass
from pathlib import Path
from typing import Any, TypeVar
from nanobot.config.loader import load_config, save_config
from nanobot.config.schema import Config
_T = TypeVar("_T")
_WEBUI_OAUTH_MAX_FLOWS = 8
class WebUISettingsConfig:
"""Instance-scoped config access with serialized read-modify-write operations."""
def __init__(self, config_path: Path) -> None:
self.path = config_path.expanduser().resolve(strict=False)
self._lock = threading.RLock()
def load(self) -> Config:
"""Load this gateway's config without consulting the process-global path."""
with self._lock:
return load_config(self.path)
def update(self, mutation: Callable[[Config], _T]) -> _T:
"""Apply and atomically persist one in-process read-modify-write operation."""
with self._lock:
config = load_config(self.path)
result = mutation(config)
save_config(config, self.path)
return result
def run_serialized(self, operation: Callable[[Path], _T]) -> _T:
"""Run a path-aware read-modify-write operation under the instance lock."""
with self._lock:
return operation(self.path)
class WebUIOAuthFlowRegistry:
"""Bounded, thread-safe OAuth flows owned by one gateway instance."""
def __init__(self, *, max_flows: int = _WEBUI_OAUTH_MAX_FLOWS) -> None:
if max_flows < 1:
raise ValueError("max_flows must be at least one")
self._max_flows = max_flows
self._flows: dict[str, tuple[str, Any]] = {}
self._lock = threading.Lock()
def register(self, provider_name: str, flow_id: str, flow: Any) -> None:
discarded: list[Any] = []
with self._lock:
for existing_id, (_provider_name, existing) in list(self._flows.items()):
if existing.expired:
discarded.append(self._flows.pop(existing_id)[1])
while len(self._flows) >= self._max_flows:
oldest_id = next(iter(self._flows))
discarded.append(self._flows.pop(oldest_id)[1])
self._flows[flow_id] = (provider_name, flow)
for existing in discarded:
existing.cancel()
def get(self, provider_name: str, flow_id: str) -> Any | None:
with self._lock:
registered = self._flows.get(flow_id)
if registered is None or registered[0] != provider_name:
return None
flow = registered[1]
if not flow.expired:
return flow
self._flows.pop(flow_id, None)
flow.cancel()
return None
def remove(
self,
provider_name: str,
flow_id: str,
flow: Any,
*,
cancel: bool = True,
) -> None:
with self._lock:
registered = self._flows.get(flow_id)
if (
registered is not None
and registered[0] == provider_name
and registered[1] is flow
):
self._flows.pop(flow_id)
if cancel:
flow.cancel()
def clear(self, provider_name: str) -> None:
with self._lock:
flow_ids = [
flow_id
for flow_id, (registered_provider, _flow) in self._flows.items()
if registered_provider == provider_name
]
flows = [self._flows.pop(flow_id)[1] for flow_id in flow_ids]
for flow in flows:
flow.cancel()
@dataclass(frozen=True)
class WebUISettingsServices:
"""Settings dependencies composed once for a gateway instance."""
config: WebUISettingsConfig
oauth_flows: WebUIOAuthFlowRegistry
@classmethod
def create(cls, config_path: Path) -> WebUISettingsServices:
return cls(
config=WebUISettingsConfig(config_path),
oauth_flows=WebUIOAuthFlowRegistry(),
)
def read(
self,
operation: Callable[..., _T],
/,
*args: Any,
**kwargs: Any,
) -> _T:
"""Run a settings read against this gateway's explicit config path."""
return operation(*args, config_path=self.config.path, **kwargs)
def mutate(
self,
operation: Callable[..., _T],
/,
*args: Any,
**kwargs: Any,
) -> _T:
"""Serialize a path-aware settings read-modify-write operation."""
return self.config.run_serialized(
lambda config_path: operation(
*args,
config_path=config_path,
**kwargs,
)
)
+957
View File
@@ -0,0 +1,957 @@
"""System and channel settings domain logic."""
from __future__ import annotations
import asyncio
import inspect
import re
import time
from collections.abc import Callable, Iterable
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING, Any, TypedDict, cast
from zoneinfo import ZoneInfo
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,
)
from nanobot.config.schema import Config
from nanobot.optional_features import OptionalFeatureError, with_channel_runtime_status
from nanobot.security.workspace_access import workspace_sandbox_status
from nanobot.webui.settings_capabilities import network_safety_payload
from nanobot.webui.settings_contracts import (
QueryParams,
SettingsRequest,
SettingsRouteResult,
WebUISettingsError,
query_first,
query_first_alias,
)
from nanobot.webui.token_usage import token_usage_payload
if TYPE_CHECKING:
from nanobot.webui.settings_services import WebUISettingsServices
LoadChannelPlugin = Callable[[str], Any]
ListPendingPairings = Callable[[], Iterable[dict[str, Any]]]
SettingsOperation = Callable[..., Any]
@dataclass(frozen=True)
class SystemSettingsOperations:
cli_apps_payload: SettingsOperation
cli_apps_action: SettingsOperation
nanobot_features_payload: SettingsOperation
nanobot_features_action: SettingsOperation
nanobot_feature_instance_target: SettingsOperation
validate_channel_config: SettingsOperation
load_channel_plugin: LoadChannelPlugin
list_pending: ListPendingPairings
approve_code: SettingsOperation
deny_code: SettingsOperation
mcp_presets_action: SettingsOperation
reload_mcp: SettingsOperation
check_for_update: SettingsOperation
channel_feature_action: SettingsOperation | None = None
channel_runtime_status: Callable[[], dict[str, Any]] | None = None
class SystemSettingsPayload(TypedDict):
runtime: dict[str, Any]
usage: dict[str, Any]
advanced: dict[str, Any]
version: dict[str, Any]
docs: dict[str, Any]
_DOCS_STABLE_VERSION_RE = re.compile(r"^\d+\.\d+\.\d+(?:\.post\d+)?$")
_DOCS_LATEST_URL = "https://nanobot.wiki/docs/latest"
_SKIP_FIELD = object()
def docs_version(version: str) -> str:
"""Map package versions to the matching public docs path."""
normalized = version.strip()
if _DOCS_STABLE_VERSION_RE.fullmatch(normalized):
return normalized
return "latest"
def docs_payload(version: str) -> dict[str, Any]:
selected_version = docs_version(version)
base_url = f"https://nanobot.wiki/docs/{selected_version}"
return {
"version": selected_version,
"base_url": base_url,
"chat_apps_url": f"{base_url}/getting-started/chat-apps",
"latest_url": _DOCS_LATEST_URL,
}
def system_settings_payload(
config: Config,
*,
config_path: Path,
version: str,
) -> SystemSettingsPayload:
defaults = config.agents.defaults
exec_config = config.tools.exec
sandbox_status = workspace_sandbox_status(
restrict_to_workspace=config.tools.restrict_to_workspace,
workspace=config.workspace_path,
)
return {
"runtime": {
"config_path": str(config_path.expanduser()),
"workspace_path": str(config.workspace_path),
"gateway_host": config.gateway.host,
"gateway_port": config.gateway.port,
"heartbeat": {
"enabled": config.gateway.heartbeat.enabled,
"interval_s": config.gateway.heartbeat.interval_s,
"keep_recent_messages": config.gateway.heartbeat.keep_recent_messages,
},
"dream": {
"schedule": defaults.dream.describe_schedule(),
},
"unified_session": defaults.unified_session,
},
"usage": token_usage_payload(timezone_name=defaults.timezone),
"advanced": {
"restrict_to_workspace": config.tools.restrict_to_workspace,
"workspace_sandbox": sandbox_status.as_dict(),
**network_safety_payload(config),
"mcp_server_count": len(config.tools.mcp_servers),
"exec_enabled": exec_config.enable,
"exec_sandbox": exec_config.sandbox or None,
"exec_path_prepend_set": bool(exec_config.path_prepend),
"exec_path_append_set": bool(exec_config.path_append),
},
"version": {"current": version},
"docs": docs_payload(version),
}
def settings_usage_payload(config: Config) -> dict[str, Any]:
"""Return the lightweight token usage slice for Overview refreshes."""
return token_usage_payload(timezone_name=config.agents.defaults.timezone)
def update_agent_system_settings(config: Config, query: QueryParams) -> tuple[bool, bool]:
defaults = config.agents.defaults
changed = False
restart_required = False
timezone = query_first(query, "timezone")
if timezone is not None:
timezone = timezone.strip()
if not timezone:
raise WebUISettingsError("timezone is required")
try:
ZoneInfo(timezone)
except Exception:
raise WebUISettingsError("invalid timezone") from None
timezone_changed = defaults.timezone != timezone
if timezone_changed or defaults.timezone_mode != "manual":
defaults.timezone = timezone
defaults.timezone_mode = "manual"
changed = True
restart_required = timezone_changed
tool_hint_max_length = query_first_alias(
query,
"tool_hint_max_length",
"toolHintMaxLength",
)
if tool_hint_max_length is not None:
try:
parsed = int(tool_hint_max_length)
except ValueError:
raise WebUISettingsError(
"tool_hint_max_length must be an integer"
) from None
if parsed < 20 or parsed > 500:
raise WebUISettingsError(
"tool_hint_max_length must be between 20 and 500"
)
if defaults.tool_hint_max_length != parsed:
defaults.tool_hint_max_length = parsed
changed = True
restart_required = True
return changed, restart_required
def save_channel_config_values(
config: Config,
name: str,
raw_values: dict[str, Any],
instance_id: str = "default",
*,
load_channel_plugin: LoadChannelPlugin,
) -> list[str]:
if not name:
raise WebUISettingsError("missing channel name")
try:
plugin = load_channel_plugin(name)
except ImportError:
raise WebUISettingsError(f"unknown channel '{name}'", status=404) from None
setup_spec = channel_setup_spec(name, plugin=plugin)
if setup_spec is None:
raise WebUISettingsError(
f"channel '{name}' cannot be configured from WebUI",
status=404,
)
field_types = setup_spec.route_field_types
if not raw_values:
return []
section = getattr(config.channels, name, None)
channel_config = channel_instance_config(
plugin,
section,
instance_id=instance_id,
)
saved: list[str] = []
prefix = f"channels.{name}."
for raw_key, raw_value in raw_values.items():
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)
if value_type is None:
raise WebUISettingsError(f"'{raw_key}' cannot be configured from WebUI")
value = coerce_channel_value(raw_key, raw_value, value_type)
if value is _SKIP_FIELD:
continue
assign_channel_config_value(channel_config, field, value)
saved.append(raw_key)
try:
updated_section = channel_update_instance_config(
plugin,
section,
channel_config,
instance_id=instance_id,
)
except ValueError as exc:
raise WebUISettingsError(
f"Invalid {name} configuration: {exc}",
status=400,
) from exc
setattr(config.channels, name, updated_section)
return saved
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]
else:
kind = value_type
allowed = None
if kind in {"string", "secret"}:
value = raw_value.strip() if isinstance(raw_value, str) else str(raw_value)
if kind == "secret" and not value:
return _SKIP_FIELD
return value
if kind == "list":
if raw_value is None:
return []
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 cast(list[Any], raw_value)
if str(item).strip()
]
raise WebUISettingsError(f"'{raw_key}' must be a comma-separated list")
if kind == "int":
if raw_value in (None, ""):
return _SKIP_FIELD
try:
return int(raw_value)
except (TypeError, ValueError) as exc:
raise WebUISettingsError(f"'{raw_key}' must be a number") from exc
if kind == "bool":
if isinstance(raw_value, bool):
return raw_value
value = str(raw_value).strip().lower()
if value in {"true", "1", "yes", "on"}:
return True
if value in {"false", "0", "no", "off"}:
return False
raise WebUISettingsError(f"'{raw_key}' must be true or false")
if kind == "enum":
value = raw_value.strip() if isinstance(raw_value, str) else str(raw_value)
if not value:
return _SKIP_FIELD
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
raise WebUISettingsError(f"'{raw_key}' has an unsupported field type")
def assign_channel_config_value(
channel_config: dict[str, Any],
field: str,
value: Any,
) -> None:
target = channel_config
parts = field.split(".")
for part in parts[:-1]:
current: object = target.get(part)
if not isinstance(current, dict):
current = {}
target[part] = current
target = cast(dict[str, Any], current)
target[parts[-1]] = value
def pairing_payload(
list_pending: ListPendingPairings,
last_action: dict[str, Any] | None = None,
*,
now: float | None = None,
) -> dict[str, Any]:
current_time = time.time() if now is None else now
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)
requests.append(
{
"code": str(item.get("code", "")),
"channel": str(item.get("channel", "")),
"sender_id": str(item.get("sender_id", "")),
"created_at_ms": int(created_at * 1000) if created_at else None,
"expires_at_ms": int(expires_at * 1000) if expires_at else None,
"expires_in_seconds": (
max(0, int(expires_at - current_time)) if expires_at else None
),
}
)
payload: dict[str, Any] = {"requests": requests}
if last_action is not None:
payload["last_action"] = last_action
return payload
class SystemSettingsHandler:
"""Handle channel and system commands behind a transport-neutral request DTO."""
def __init__(self, settings: WebUISettingsServices, logger: Any) -> None:
self.settings = settings
self.logger = logger
self._channel_connectors: dict[str, Any] = {}
async def handle(
self,
action: str,
request: SettingsRequest,
operations: SystemSettingsOperations,
*,
channel_name: str | None = None,
connect_action: str | None = None,
) -> SettingsRouteResult:
if action == "cli-list":
return await self._cli_apps(request, operations)
if action.startswith("cli-"):
return await self._cli_apps_action(
request,
action.removeprefix("cli-"),
operations,
)
if action == "features-list":
return await self._features(operations)
if action in {"features-enable", "features-disable"}:
return await self._features_action(
request,
action.removeprefix("features-"),
operations,
)
if action == "channel-validate":
return await self._channel_validate(request, operations)
if action == "channel-configure":
return await self._channel_configure(request, operations)
if action == "channel-connect" and channel_name and connect_action:
return await self._channel_connect(
request,
channel_name,
connect_action,
operations,
)
if action == "pairing-list":
return SettingsRouteResult.success(pairing_payload(operations.list_pending))
if action in {"pairing-approve", "pairing-deny"}:
return self._pairing_action(
request,
action.removeprefix("pairing-"),
operations,
)
if action == "mcp-list":
return await self._mcp_presets(request, None, operations)
if action.startswith("mcp-"):
return await self._mcp_presets(
request,
action.removeprefix("mcp-"),
operations,
)
if action == "version-check":
return await self._version_check(operations)
return SettingsRouteResult.failure(404, "unknown settings action")
async def _cli_apps(
self,
request: SettingsRequest,
operations: SystemSettingsOperations,
) -> SettingsRouteResult:
installed_only = (query_first(request.query, "installed_only") or "").lower() in {
"1",
"true",
"yes",
}
try:
payload = await operations.cli_apps_payload(
installed_only=installed_only,
config_path=self.settings.config.path,
)
except Exception:
self.logger.exception("failed to load CLI Apps payload")
return SettingsRouteResult.failure(500, "failed to load CLI Apps")
return SettingsRouteResult.success(payload)
async def _cli_apps_action(
self,
request: SettingsRequest,
action: str,
operations: SystemSettingsOperations,
) -> SettingsRouteResult:
try:
payload = await asyncio.to_thread(
operations.cli_apps_action,
action,
request.query,
config_path=self.settings.config.path,
)
except WebUISettingsError as exc:
return SettingsRouteResult.failure(exc.status, exc.message)
except Exception as exc:
status = getattr(exc, "status", 500)
message = getattr(exc, "message", str(exc))
if status >= 500:
self.logger.exception("CLI Apps action '{}' failed", action)
return SettingsRouteResult.failure(status, message)
return SettingsRouteResult.success(payload)
async def _features(
self,
operations: SystemSettingsOperations,
) -> SettingsRouteResult:
try:
payload = await asyncio.to_thread(
operations.nanobot_features_payload,
config_path=self.settings.config.path,
)
except Exception:
self.logger.exception("failed to load nanobot features")
return SettingsRouteResult.failure(500, "failed to load nanobot features")
return SettingsRouteResult.success(
self._with_channel_runtime_status(payload, operations)
)
def _nanobot_features_payload(
self,
operations: SystemSettingsOperations,
) -> dict[str, Any]:
return operations.nanobot_features_payload(config_path=self.settings.config.path)
def _nanobot_features_action(
self,
action: str,
query: QueryParams,
operations: SystemSettingsOperations,
*,
allow_install: bool = True,
) -> dict[str, Any]:
return self.settings.mutate(
operations.nanobot_features_action,
action,
query,
allow_install=allow_install,
)
async def _features_action(
self,
request: SettingsRequest,
action: str,
operations: SystemSettingsOperations,
) -> SettingsRouteResult:
try:
payload = await asyncio.to_thread(
self._nanobot_features_action,
action,
request.query,
operations,
allow_install=(
action != "enable"
or self.allow_feature_package_install(request)
),
)
except OptionalFeatureError as exc:
return SettingsRouteResult.failure(exc.status, exc.message)
except Exception as exc:
status = getattr(exc, "status", 500)
message = getattr(exc, "message", str(exc))
if status >= 500:
self.logger.exception(
"nanobot feature action '{}' failed",
action,
)
return SettingsRouteResult.failure(status, message)
payload = await self._apply_feature_runtime_change(
action,
request.query,
payload,
operations,
)
payload = self._with_channel_runtime_status(payload, operations)
return SettingsRouteResult.success(
payload,
decorate_restart=True,
restart_section="runtime",
)
def _with_channel_runtime_status(
self,
payload: dict[str, Any],
operations: SystemSettingsOperations,
) -> dict[str, Any]:
if operations.channel_runtime_status is None:
return payload
try:
return with_channel_runtime_status(
payload,
operations.channel_runtime_status(),
)
except Exception:
self.logger.exception("failed to load channel runtime status")
return payload
async def _apply_feature_runtime_change(
self,
action: str,
query: QueryParams,
payload: dict[str, Any],
operations: SystemSettingsOperations,
) -> dict[str, Any]:
if operations.channel_feature_action is None:
return payload
name = (query_first(query, "name") or "").strip()
if not name:
return payload
try:
instance_id = operations.nanobot_feature_instance_target(query)
result = operations.channel_feature_action(action, name, instance_id)
if inspect.isawaitable(result):
result = await result
except Exception as exc:
self.logger.exception("failed to apply channel '{}' without restart", name)
return self.feature_runtime_fallback(
payload,
message=(
f"{name} channel config was saved, but hot reload failed: {exc}"
),
)
if not isinstance(result, dict):
return payload
result = cast(dict[str, Any], result)
if not result.get("handled"):
return payload
updated = dict(payload)
updated["requires_restart"] = bool(result.get("requires_restart"))
message = result.get("message")
if isinstance(message, str) and message:
last_action = dict(updated.get("last_action") or {})
previous = last_action.get("message")
last_action["message"] = (
f"{previous}. {message}"
if isinstance(previous, str) and previous
else message
)
last_action["hot_reload"] = not updated["requires_restart"]
if "ok" in result:
last_action["ok"] = bool(result["ok"])
updated["last_action"] = last_action
return updated
@staticmethod
def feature_runtime_fallback(
payload: dict[str, Any],
*,
message: str,
) -> dict[str, Any]:
updated = dict(payload)
updated["requires_restart"] = True
last_action = dict(updated.get("last_action") or {})
previous = last_action.get("message")
last_action["message"] = (
f"{previous}. {message}"
if isinstance(previous, str) and previous
else message
)
last_action["hot_reload"] = False
updated["last_action"] = last_action
return updated
async def _channel_configure(
self,
request: SettingsRequest,
operations: SystemSettingsOperations,
) -> SettingsRouteResult:
name = (query_first(request.query, "name") or "").strip()
instance_id = (
query_first(request.query, "instance_id") or "default"
).strip()
enable = (query_first(request.query, "enable") or "").strip().lower() in {
"1",
"true",
"yes",
}
try:
saved = await asyncio.to_thread(
self._save_channel_config_values,
name,
self.parse_channel_values(request),
instance_id,
operations,
)
except WebUISettingsError as exc:
return SettingsRouteResult.failure(exc.status, exc.message)
except Exception:
self.logger.exception("failed to save channel '{}' settings", name)
return SettingsRouteResult.failure(500, "failed to save channel settings")
payload: dict[str, Any] = {
"name": name,
"saved": True,
"saved_keys": saved,
}
if not enable:
features = await asyncio.to_thread(
self._nanobot_features_payload,
operations,
)
payload["nanobot_features"] = self._with_channel_runtime_status(
features,
operations,
)
return SettingsRouteResult.success(
payload,
decorate_restart=True,
restart_section="runtime",
restart_payload_key="nanobot_features",
)
feature_query = {"name": [name]}
if instance_id:
feature_query["instance_id"] = [instance_id]
try:
features = await asyncio.to_thread(
self._nanobot_features_action,
"enable",
feature_query,
operations,
allow_install=self.allow_feature_package_install(request),
)
except OptionalFeatureError as exc:
return SettingsRouteResult.failure(
exc.status,
f"Settings saved, but {exc.message}",
)
except Exception as exc:
self.logger.exception(
"failed to enable channel '{}' after settings save",
name,
)
return SettingsRouteResult.failure(
500,
f"Settings saved, but enabling {name} failed: {exc}",
)
features = await self._apply_feature_runtime_change(
"enable",
feature_query,
features,
operations,
)
payload["nanobot_features"] = self._with_channel_runtime_status(
features,
operations,
)
return SettingsRouteResult.success(
payload,
decorate_restart=True,
restart_section="runtime",
restart_payload_key="nanobot_features",
)
async def _channel_validate(
self,
request: SettingsRequest,
operations: SystemSettingsOperations,
) -> SettingsRouteResult:
name = (query_first(request.query, "name") or "").strip()
instance_id = (
query_first(request.query, "instance_id") or "default"
).strip()
try:
payload = await asyncio.to_thread(
operations.validate_channel_config,
name,
self.parse_channel_values(request),
instance_id=instance_id,
)
except WebUISettingsError as exc:
return SettingsRouteResult.failure(exc.status, exc.message)
except Exception:
self.logger.exception("failed to validate channel '{}' settings", name)
return SettingsRouteResult.failure(
500,
"failed to validate channel settings",
)
return SettingsRouteResult.success(payload)
@staticmethod
def parse_channel_values(request: SettingsRequest) -> dict[str, Any]:
if request.payload is None or "values" not in request.payload:
return {}
values = request.payload.get("values")
if not isinstance(values, dict):
raise WebUISettingsError(
"channel settings payload must be a JSON object"
)
return cast(dict[str, Any], values)
def _save_channel_config_values(
self,
name: str,
raw_values: dict[str, Any],
instance_id: str,
operations: SystemSettingsOperations,
) -> list[str]:
return self.settings.config.update(
lambda config: save_channel_config_values(
config,
name,
raw_values,
instance_id,
load_channel_plugin=operations.load_channel_plugin,
)
)
async def _channel_connect(
self,
request: SettingsRequest,
channel_name: str,
action: str,
operations: SystemSettingsOperations,
) -> SettingsRouteResult:
try:
connector = self._channel_connectors.get(channel_name)
if connector is None:
plugin = operations.load_channel_plugin(channel_name)
connector = plugin.load_connector()
self._channel_connectors[channel_name] = connector
except ImportError:
return SettingsRouteResult.failure(
404,
f"channel '{channel_name}' does not support connect",
)
try:
payload = await connector.handle(action, request.query)
except ChannelConnectError as exc:
return SettingsRouteResult.failure(exc.status, exc.message)
except Exception:
self.logger.exception(
"failed to run {} WebUI connect action for {}",
action,
channel_name,
)
return SettingsRouteResult.failure(
500,
f"failed to {action} {channel_name} connection",
)
if payload.get("status") != "succeeded":
return SettingsRouteResult.success(payload)
payload = await self._with_channel_connect_success(
request,
channel_name,
payload,
operations,
)
return SettingsRouteResult.success(
payload,
decorate_restart=True,
restart_section="runtime",
restart_payload_key="nanobot_features",
)
async def _with_channel_connect_success(
self,
request: SettingsRequest,
channel_name: str,
payload: dict[str, Any],
operations: SystemSettingsOperations,
) -> dict[str, Any]:
target = {"name": [channel_name]}
if payload.get("instance_id"):
target["instance_id"] = [str(payload["instance_id"])]
try:
features = await asyncio.to_thread(
self._nanobot_features_action,
"enable",
target,
operations,
allow_install=self.allow_feature_package_install(request),
)
except OptionalFeatureError as exc:
features = self.feature_runtime_fallback(
self._nanobot_features_payload(operations),
message=(
f"{channel_name} connected, but enabling channel support failed: "
f"{exc.message}"
),
)
else:
features = await self._apply_feature_runtime_change(
"enable",
target,
features,
operations,
)
updated = dict(payload)
updated["nanobot_features"] = self._with_channel_runtime_status(
features,
operations,
)
return updated
def allow_feature_package_install(self, request: SettingsRequest) -> bool:
if request.local_browser:
return True
try:
return bool(
self.settings.config.load().tools.webui_allow_remote_package_install
)
except Exception:
self.logger.exception("failed to load remote package install policy")
return False
def _pairing_action(
self,
request: SettingsRequest,
action: str,
operations: SystemSettingsOperations,
) -> SettingsRouteResult:
code = (query_first(request.query, "code") or "").strip()
if not code:
return SettingsRouteResult.failure(400, "Missing pairing code")
if action == "approve":
result = operations.approve_code(code)
if result is None:
return SettingsRouteResult.failure(
404,
"Pairing code not found or expired",
)
channel, sender_id = result
return SettingsRouteResult.success(
pairing_payload(
operations.list_pending,
{
"ok": True,
"action": "approve",
"message": f"Approved {sender_id} for {channel}",
"channel": channel,
"sender_id": sender_id,
"code": code,
},
)
)
if not operations.deny_code(code):
return SettingsRouteResult.failure(
404,
"Pairing code not found or expired",
)
return SettingsRouteResult.success(
pairing_payload(
operations.list_pending,
{
"ok": True,
"action": "deny",
"message": f"Denied pairing code {code}",
"code": code,
},
)
)
async def _mcp_presets(
self,
request: SettingsRequest,
action: str | None,
operations: SystemSettingsOperations,
) -> SettingsRouteResult:
try:
payload = await operations.mcp_presets_action(
action,
request.query,
reload_mcp=operations.reload_mcp,
config=self.settings.config,
)
except Exception as exc:
status = getattr(exc, "status", 500)
message = getattr(exc, "message", str(exc))
if status >= 500:
self.logger.exception(
"MCP preset action '{}' failed",
action or "list",
)
return SettingsRouteResult.failure(status, message)
return SettingsRouteResult.success(
payload,
decorate_restart=action is not None,
restart_section="runtime" if action is not None else None,
)
async def _version_check(
self,
operations: SystemSettingsOperations,
) -> SettingsRouteResult:
try:
update_info = await asyncio.to_thread(operations.check_for_update)
except Exception:
self.logger.exception("version check failed")
return SettingsRouteResult.failure(500, "version check failed")
return SettingsRouteResult.success({"updateAvailable": update_info})
+3 -1
View File
@@ -25,7 +25,7 @@ _MAX_KEY_LEN = 512
_MAX_TITLE_LEN = 160
_MAX_TAG_LEN = 40
_ALLOWED_DENSITIES = {"comfortable", "compact"}
_ALLOWED_SORTS = {"updated_desc", "created_desc", "title_asc"}
_ALLOWED_SORTS = {"updated_desc", "created_desc", "title_asc", "manual"}
def webui_sidebar_state_path() -> Path:
@@ -37,6 +37,7 @@ def default_webui_sidebar_state() -> dict[str, Any]:
"schema_version": WEBUI_SIDEBAR_STATE_SCHEMA_VERSION,
"pinned_keys": [],
"archived_keys": [],
"session_order": [],
"title_overrides": {},
"project_name_overrides": {},
"tags_by_key": {},
@@ -138,6 +139,7 @@ def normalize_webui_sidebar_state(raw: Any) -> dict[str, Any]:
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"))
state["session_order"] = _clean_string_list(raw.get("session_order"))
state["title_overrides"] = _clean_title_overrides(raw.get("title_overrides"))
state["project_name_overrides"] = _clean_title_overrides(
raw.get("project_name_overrides")
+218
View File
@@ -0,0 +1,218 @@
"""Connection-owned Temporary Chat behavior for the WebUI."""
from __future__ import annotations
import uuid
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from nanobot.bus.events import (
INBOUND_META_RUNTIME_CONTROL,
RUNTIME_CONTROL_SESSION_DISCARD,
InboundMessage,
)
from nanobot.bus.queue import MessageBus
from nanobot.security.workspace_access import WorkspaceScope
from nanobot.session.manager import Session, SessionManager
from nanobot.webui.workspaces import WebUIWorkspaceController
_TEMPORARY_CHAT_DISABLED_TOOLS = frozenset({
"create_goal",
"update_goal",
"spawn",
"cron",
})
_TEMPORARY_CHAT_COMMANDS = frozenset({"/model", "/stop"})
class TemporaryChatError(ValueError):
"""A stable WebUI protocol error for a Temporary Chat operation."""
def __init__(self, detail: str) -> None:
super().__init__(detail)
self.detail = detail
@dataclass(frozen=True)
class TemporaryChatMessagePolicy:
"""Server-owned message rules for one active Temporary Chat."""
session_key: str
workspace_scope: WorkspaceScope
require_existing_session: bool = True
hydrate_transcript: bool = False
persist_transcript: bool = False
class WebUITemporaryChats:
"""Own Temporary Chat creation, policy, attachments, and disposal."""
def __init__(
self,
*,
bus: MessageBus,
session_manager: SessionManager | None,
workspaces: WebUIWorkspaceController,
logger: Any,
channel_name: str = "websocket",
) -> None:
self._bus = bus
self._sessions = session_manager
self._workspaces = workspaces
self._logger = logger
self._channel_name = channel_name
self._owners: dict[str, object] = {}
self._owner_chat_ids: dict[object, set[str]] = {}
# Keep active sessions alive if the bounded manager cache evicts them
# between WebUI turns. SessionPolicy remains the authority below.
self._active_sessions: dict[str, Session] = {}
# Retain policy-derived tombstones until shutdown so late outbound
# events cannot create a durable transcript after a chat is discarded.
self._known_transient_chat_ids: set[str] = set()
self._media_paths: dict[str, set[str]] = {}
def _session_key(self, chat_id: str) -> str:
return f"{self._channel_name}:{chat_id}"
def _cached_session_is_transient(self, chat_id: str) -> bool:
if self._sessions is None:
return False
session = self._sessions.get_cached(self._session_key(chat_id))
return session is not None and not session.policy.persist
def create(self, owner: object, *, trusted_webui: bool) -> str:
"""Create a server-identified chat owned by one authenticated WebUI connection."""
if not trusted_webui:
raise TemporaryChatError("access_denied")
if self._sessions is None:
raise TemporaryChatError("temporary_chat_unavailable")
chat_id = str(uuid.uuid4())
session = self._sessions.get_or_create_transient(
self._session_key(chat_id),
disabled_tools=_TEMPORARY_CHAT_DISABLED_TOOLS,
)
if session.policy.persist:
raise RuntimeError("Temporary Chat must use a non-persistent session policy")
self._owners[chat_id] = owner
self._owner_chat_ids.setdefault(owner, set()).add(chat_id)
self._active_sessions[chat_id] = session
self._known_transient_chat_ids.add(chat_id)
return chat_id
def message_policy(
self,
owner: object,
chat_id: str,
content: str,
) -> TemporaryChatMessagePolicy | None:
"""Return Temporary Chat rules, or ``None`` for an ordinary chat."""
if not self._cached_session_is_transient(chat_id):
if chat_id in self._known_transient_chat_ids:
raise TemporaryChatError("temporary_chat_unavailable")
return None
if self._owners.get(chat_id) is not owner or self._sessions is None:
raise TemporaryChatError("temporary_chat_unavailable")
session = self._sessions.get_cached(self._session_key(chat_id))
if session is None:
raise TemporaryChatError("temporary_chat_unavailable")
command = content.strip().split(maxsplit=1)[0].lower() if content.strip() else ""
if command.startswith("/") and command not in _TEMPORARY_CHAT_COMMANDS:
raise TemporaryChatError("temporary_chat_command_rejected")
return TemporaryChatMessagePolicy(
session_key=self._session_key(chat_id),
workspace_scope=self._workspaces.restricted_default_scope(),
)
def validate_attach(self, chat_id: str) -> None:
"""Reject attempts to recover a non-persistent session."""
if not self._cached_session_is_transient(chat_id):
if chat_id in self._known_transient_chat_ids:
raise TemporaryChatError("temporary_chat_unavailable")
return
raise TemporaryChatError("temporary_chat_unavailable")
def validate_workspace_update(self, chat_id: str) -> None:
"""Prevent non-persistent sessions from acquiring durable workspace state."""
if self._cached_session_is_transient(chat_id):
raise TemporaryChatError("temporary_chat_workspace_rejected")
if chat_id in self._known_transient_chat_ids:
raise TemporaryChatError("temporary_chat_unavailable")
def register_media(self, owner: object, chat_id: str, paths: list[str]) -> None:
if not paths:
return
if self._owners.get(chat_id) is not owner:
raise TemporaryChatError("temporary_chat_unavailable")
self._media_paths.setdefault(chat_id, set()).update(paths)
def chat_ids_for_owner(self, owner: object) -> tuple[str, ...]:
return tuple(self._owner_chat_ids.get(owner, ()))
def owns(self, owner: object, chat_id: str) -> bool:
return self._owners.get(chat_id) is owner
def should_persist_transcript(self, chat_id: str) -> bool:
"""Apply the session policy and retain it for late events after disposal."""
return (
not self._cached_session_is_transient(chat_id)
and chat_id not in self._known_transient_chat_ids
)
def _discard_media(self, chat_id: str) -> None:
for raw_path in self._media_paths.pop(chat_id, set()):
try:
Path(raw_path).unlink(missing_ok=True)
except OSError:
self._logger.warning("failed to remove a temporary WebUI attachment")
def _forget_owner(self, owner: object, chat_id: str) -> None:
self._owners.pop(chat_id, None)
chat_ids = self._owner_chat_ids.get(owner)
if chat_ids is None:
return
chat_ids.discard(chat_id)
if not chat_ids:
self._owner_chat_ids.pop(owner, None)
async def discard(self, owner: object, chat_id: str) -> None:
"""Forget one owned chat and cancel any active work through the message bus."""
if (
not self._cached_session_is_transient(chat_id)
or self._owners.get(chat_id) is not owner
):
raise TemporaryChatError("temporary_chat_unavailable")
session_key = self._session_key(chat_id)
self._forget_owner(owner, chat_id)
self._active_sessions.pop(chat_id, None)
self._discard_media(chat_id)
if self._sessions is not None:
self._sessions.invalidate(session_key)
await self._bus.publish_inbound(
InboundMessage(
channel=self._channel_name,
sender_id="webui",
chat_id=chat_id,
content="",
metadata={
INBOUND_META_RUNTIME_CONTROL: RUNTIME_CONTROL_SESSION_DISCARD,
},
session_key_override=session_key,
)
)
def close(self) -> None:
"""Release process-local resources during gateway shutdown."""
for chat_id in tuple(self._owners):
self._discard_media(chat_id)
if self._sessions is not None:
self._sessions.invalidate(self._session_key(chat_id))
self._owners.clear()
self._owner_chat_ids.clear()
self._active_sessions.clear()
self._known_transient_chat_ids.clear()
-29
View File
@@ -1313,21 +1313,6 @@ def _recover_incomplete_turns(
return recovered
def recover_incomplete_turns_from_session(
lines: list[dict[str, Any]],
session_messages: list[dict[str, Any]] | None,
*,
session_key: str,
) -> list[dict[str, Any]]:
"""Recover marked transcript answers only when one durable session turn matches."""
if not lines or not session_messages or not _needs_incomplete_turn_recovery(lines):
return lines
session_turns = _session_backfill_turns(session_key, session_messages)
if not session_turns:
return lines
return _recover_incomplete_turns(lines, session_turns)
def _with_backfilled_user(
records: list[dict[str, Any]],
user_event: dict[str, Any],
@@ -1365,20 +1350,6 @@ def _inject_missing_user_events(
return out
def inject_missing_user_events_from_session(
session_key: str,
lines: list[dict[str, Any]],
session_messages: list[dict[str, Any]] | None,
) -> list[dict[str, Any]]:
"""Backfill user rows for legacy WebUI transcripts that only stored assistant streams."""
if not lines or not session_messages or not _needs_user_event_backfill(lines):
return lines
session_turns = _session_backfill_turns(session_key, session_messages)
if not session_turns:
return lines
return _inject_missing_user_events(lines, session_turns)
def _format_tool_call_trace(call: Any) -> str | None:
if not call or not isinstance(call, dict):
return None
+8 -1
View File
@@ -10,6 +10,7 @@ import time
from typing import Any
import httpx
from packaging.version import InvalidVersion, Version
from nanobot import __version__
@@ -42,7 +43,13 @@ def check_for_update() -> dict[str, Any] | None:
return None
_cache = (now, latest)
if not latest or latest == __version__:
if not isinstance(latest, str) or not latest:
return None
try:
if Version(latest) <= Version(__version__):
return None
except InvalidVersion:
logger.debug("PyPI returned an invalid nanobot version: %r", latest)
return None
return {
"currentVersion": __version__,
+8
View File
@@ -191,6 +191,14 @@ class WebUIWorkspaceController:
self._default_restrict_to_workspace,
)
def restricted_default_scope(self) -> WorkspaceScope:
"""Return the default workspace with access restricted for this request."""
return build_workspace_scope(
self._default_workspace,
"restricted",
source_channel=_WEBUI_SCOPE_CHANNEL,
)
def _scope_from_metadata_value(
self,
raw_scope: object,
+175 -59
View File
@@ -17,19 +17,18 @@ import time
from collections.abc import Callable
from pathlib import Path
from typing import TYPE_CHECKING, Any, cast
from urllib.parse import unquote
from urllib.parse import quote, unquote
from loguru import logger
from websockets.datastructures import Headers
from websockets.http11 import Request as WsRequest
from websockets.http11 import Response
from nanobot.command.builtin import builtin_command_palette
from nanobot.cron.session_turns import is_bound_cron_job
from nanobot.cron.types import CronJob, CronSchedule
from nanobot.runtime_context import public_history_messages
from nanobot.security.workspace_access import WorkspaceScope
from nanobot.triggers.local_types import LocalTrigger
from nanobot.utils.subagent_channel_display import scrub_subagent_messages_for_channel
from nanobot.webui.file_preview import (
WebUIFilePreviewError,
file_preview_availability_payload,
@@ -120,7 +119,60 @@ from nanobot.webui.transcript import build_webui_thread_response
from nanobot.webui.workspaces import WebUIWorkspaceController
_SLOW_WEBUI_HTTP_LOG_MS = 1_000
_AUTOMATION_VALUES_HEADER = "X-Nanobot-Automation-Values"
_WEBUI_MUTATION_PAYLOAD_ATTR = "_nanobot_webui_mutation_payload"
_WEBUI_MUTATION_REQUEST_ATTR = "_nanobot_webui_mutation_request"
_WEBUI_MUTATION_PATHS = {
"automation.enable": "/api/webui/automations/enable",
"automation.disable": "/api/webui/automations/disable",
"automation.delete": "/api/webui/automations/delete",
"automation.run": "/api/webui/automations/run",
"automation.update": "/api/webui/automations/update",
"skill.install": "/api/webui/skills/install",
"skill.update": "/api/webui/skills/update",
"skill.delete": "/api/webui/skills/delete",
"sidebar.update": "/api/webui/sidebar-state/update",
"settings.agent.update": "/api/settings/update",
"settings.model_configuration.create": "/api/settings/model-configurations/create",
"settings.model_configuration.update": "/api/settings/model-configurations/update",
"settings.model_configuration.delete": "/api/settings/model-configurations/delete",
"settings.model_configuration.migrate": "/api/settings/model-configurations/migrate",
"settings.model_call_order.update": "/api/settings/model-call-order/update",
"settings.provider.update": "/api/settings/provider/update",
"settings.provider.create": "/api/settings/provider/create",
"settings.provider.oauth_login": "/api/settings/provider/oauth-login",
"settings.provider.oauth_complete": "/api/settings/provider/oauth-login/complete",
"settings.provider.oauth_logout": "/api/settings/provider/oauth-logout",
"settings.web_search.update": "/api/settings/web-search/update",
"settings.api_service.start": "/api/settings/api-service/start",
"settings.api_service.stop": "/api/settings/api-service/stop",
"settings.image_generation.update": "/api/settings/image-generation/update",
"settings.transcription.update": "/api/settings/transcription/update",
"settings.network_safety.update": "/api/settings/network-safety/update",
"settings.cli_app.install": "/api/settings/cli-apps/install",
"settings.cli_app.update": "/api/settings/cli-apps/update",
"settings.cli_app.uninstall": "/api/settings/cli-apps/uninstall",
"settings.cli_app.test": "/api/settings/cli-apps/test",
"settings.feature.enable": "/api/settings/nanobot-features/enable",
"settings.feature.disable": "/api/settings/nanobot-features/disable",
"settings.channel.validate": "/api/settings/channels/validate",
"settings.channel.configure": "/api/settings/channels/configure",
"settings.pairing.approve": "/api/settings/pairing/approve",
"settings.pairing.deny": "/api/settings/pairing/deny",
"settings.mcp.enable": "/api/settings/mcp-presets/enable",
"settings.mcp.remove": "/api/settings/mcp-presets/remove",
"settings.mcp.test": "/api/settings/mcp-presets/test",
"settings.mcp.custom": "/api/settings/mcp-presets/custom",
"settings.mcp.import": "/api/settings/mcp-presets/import",
"settings.mcp.import_cursor": "/api/settings/mcp-presets/import-cursor",
"settings.mcp.tools": "/api/settings/mcp-presets/tools",
}
_WEBUI_CHANNEL_CONNECT_ACTIONS = {
"settings.channel.connect.start": "start",
"settings.channel.connect.poll": "poll",
"settings.channel.connect.cancel": "cancel",
}
# Fix for #5190: On Windows, mimetypes.guess_type() reads the registry key
# HKEY_CLASSES_ROOT\.js\Content Type, which is commonly set to 'text/plain'
@@ -152,6 +204,7 @@ if TYPE_CHECKING:
from nanobot.cron.service import CronService
from nanobot.session.manager import SessionManager
from nanobot.triggers.local_store import LocalTriggerStore
from nanobot.webui.settings_services import WebUISettingsServices
def _decode_api_key(raw_key: str) -> str | None:
key = unquote(raw_key)
@@ -161,6 +214,33 @@ def _decode_api_key(raw_key: str) -> str | None:
return key
def _mutation_payload(request: WsRequest) -> dict[str, Any] | None:
payload = getattr(request, _WEBUI_MUTATION_PAYLOAD_ATTR, None)
if not isinstance(payload, dict):
return None
return cast(dict[str, Any], payload)
def _request_query(request: WsRequest) -> dict[str, list[str]]:
payload = _mutation_payload(request)
if payload is None:
return _parse_query(request.path)
query: dict[str, list[str]] = {}
for key, value in payload.items():
if not key:
continue
if isinstance(value, bool):
text = "true" if value else "false"
elif value is None:
text = ""
elif isinstance(value, (dict, list)):
text = json.dumps(value, ensure_ascii=False, separators=(",", ":"))
else:
text = str(value)
query[key] = [text]
return query
def _default_model_name_from_config() -> str | None:
try:
from nanobot.config.loader import load_config
@@ -213,6 +293,7 @@ class GatewayHTTPHandler:
media: WebUIMediaGateway,
ingress: WebUIIngressPolicy,
workspaces: WebUIWorkspaceController,
settings: WebUISettingsServices,
skills_workspace_path: Path,
disabled_skills: set[str] | None = None,
cron_service: CronService | None = None,
@@ -233,6 +314,7 @@ class GatewayHTTPHandler:
self.media = media
self.ingress = ingress
self.workspaces = workspaces
self.settings = settings
self.skills_workspace_path = skills_workspace_path
self.disabled_skills: set[str] = (
disabled_skills if disabled_skills is not None else set()
@@ -251,6 +333,7 @@ class GatewayHTTPHandler:
self._capabilities = _rc(runtime_surface, runtime_capabilities_overrides or {})
self.settings_routes = WebUISettingsRouter(
settings=settings,
bus=bus,
logger=self._log,
check_api_token=self.check_api_token,
@@ -287,11 +370,86 @@ class GatewayHTTPHandler:
)
try:
if self._is_webui_mutation_path(got):
return _http_error(
405,
"WebUI mutations require an authenticated WebSocket",
)
response = await self._dispatch_resolved(connection, request, got)
return response
finally:
self._log_slow_http(got, response, started)
async def dispatch_webui_mutation(
self,
connection: Any,
action: str,
payload: dict[str, Any],
) -> Response:
"""Run one explicitly allowlisted mutation for an authenticated WebUI socket."""
path = self._webui_mutation_path(action, payload)
if isinstance(path, Response):
return path
source_request = getattr(connection, "request", None)
source_headers = getattr(source_request, "headers", None)
if source_headers is None:
headers = Headers()
else:
try:
headers = Headers(source_headers.raw_items())
except (AttributeError, TypeError):
try:
headers = Headers(source_headers)
except TypeError:
headers = Headers()
request = WsRequest(path, headers)
setattr(request, "_nanobot_trusted_proxy_authenticated", True)
setattr(request, _WEBUI_MUTATION_REQUEST_ATTR, True)
setattr(request, _WEBUI_MUTATION_PAYLOAD_ATTR, dict(payload))
response = await self._dispatch_resolved(connection, request, path)
if isinstance(response, Response):
return response
return _http_error(404, "WebUI mutation action not found")
def _is_webui_mutation_path(self, path: str) -> bool:
if self.settings_routes.is_mutation_path(path):
return True
if re.match(r"^/api/sessions/[^/]+/delete$", path):
return True
if re.match(r"^/api/webui/automations/(enable|disable|delete|run|update)$", path):
return True
return path in {
"/api/webui/skills/install",
"/api/webui/skills/update",
"/api/webui/skills/delete",
"/api/webui/sidebar-state/update",
}
@staticmethod
def _webui_mutation_path(
action: str,
payload: dict[str, Any],
) -> str | Response:
path = _WEBUI_MUTATION_PATHS.get(action)
if path is not None:
return path
if action == "session.delete":
key = payload.get("key")
if not isinstance(key, str) or not key.strip():
return _http_error(400, "missing session key")
return f"/api/sessions/{quote(key, safe='')}/delete"
connect_action = _WEBUI_CHANNEL_CONNECT_ACTIONS.get(action)
if connect_action is not None:
channel = payload.get("channel")
if not isinstance(channel, str) or re.fullmatch(
r"[A-Za-z0-9_-]{1,64}",
channel,
) is None:
return _http_error(400, "invalid channel name")
return f"/api/settings/channels/{channel}/connect/{connect_action}"
return _http_error(404, "unknown WebUI mutation action")
async def _dispatch_resolved(
self,
connection: Any,
@@ -462,10 +620,6 @@ class GatewayHTTPHandler:
# -- Session routes -----------------------------------------------------
async def _dispatch_session_routes(self, request: WsRequest, got: str) -> Response | None:
m = re.match(r"^/api/sessions/([^/]+)/messages$", got)
if m:
return self._handle_session_messages(request, m.group(1))
m = re.match(r"^/api/sessions/([^/]+)/webui-thread$", got)
if m:
return self._handle_webui_thread_get(request, m.group(1))
@@ -527,34 +681,6 @@ class GatewayHTTPHandler:
cleaned.append(row)
return {"sessions": cleaned}
def _handle_session_messages(self, request: WsRequest, key: str) -> Response:
if not self.check_api_token(request):
return _http_error(401, "Unauthorized")
if self.session_manager is None:
return _http_error(503, "session manager unavailable")
decoded_key = _decode_api_key(key)
if decoded_key is None:
return _http_error(400, "invalid session key")
if not _is_websocket_channel_session_key(decoded_key):
return _http_error(404, "session not found")
data = self.session_manager.read_session_file(decoded_key)
if data is None:
return _http_error(404, "session not found")
messages = data.get("messages")
if isinstance(messages, list):
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(
[
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)
def _handle_webui_thread_get(self, request: WsRequest, key: str) -> Response:
if not self.check_api_token(request):
return _http_error(401, "Unauthorized")
@@ -680,7 +806,7 @@ class GatewayHTTPHandler:
return _http_error(400, "invalid session key")
if not _is_websocket_channel_session_key(decoded_key):
return _http_error(404, "session not found")
query = _parse_query(request.path)
query = _request_query(request)
delete_automations = (_query_first(query, "delete_automations") or "").lower()
automation_jobs = session_automation_jobs(
self.cron_service,
@@ -776,7 +902,7 @@ class GatewayHTTPHandler:
if self.cron_service is None and self.local_trigger_store is None:
return _http_error(503, "automation service unavailable")
query = _parse_query(request.path)
query = _request_query(request)
job_id = (_query_first(query, "id") or _query_first(query, "job_id") or "").strip()
if not job_id:
return _http_error(400, "missing automation id")
@@ -1008,7 +1134,7 @@ class GatewayHTTPHandler:
if self._skill_install_lock.locked():
return _http_error(409, "another skill installation is already in progress")
query = _parse_query(request.path)
query = _request_query(request)
provider = _query_first(query, "provider") or "skills_sh"
source = _query_first(query, "source") or ""
skill_id = _query_first(query, "skill") or ""
@@ -1049,7 +1175,7 @@ class GatewayHTTPHandler:
def _handle_webui_skill_update(self, request: WsRequest) -> Response:
if not self.check_api_token(request):
return _http_error(401, "Unauthorized")
query = _parse_query(request.path)
query = _request_query(request)
name = _query_first(query, "name") or ""
raw_enabled = (_query_first(query, "enabled") or "").lower()
if raw_enabled not in {"true", "false"}:
@@ -1081,7 +1207,7 @@ class GatewayHTTPHandler:
return _http_error(401, "Unauthorized")
if not _is_local_browser_request(connection, request.headers):
return _http_error(403, "remote skill deletion is disabled")
name = _query_first(_parse_query(request.path), "name") or ""
name = _query_first(_request_query(request), "name") or ""
try:
action = delete_webui_skill(
self.skills_workspace_path,
@@ -1128,18 +1254,14 @@ class GatewayHTTPHandler:
def _handle_webui_sidebar_state_update(self, request: WsRequest) -> Response:
if not self.check_api_token(request):
return _http_error(401, "Unauthorized")
query = _parse_query(request.path)
raw_state = _query_first(query, "state")
if raw_state is None:
payload = _mutation_payload(request)
state_value = payload.get("state") if payload is not None else None
if state_value is None:
return _http_error(400, "missing state")
try:
decoded = json.loads(raw_state)
except json.JSONDecodeError:
return _http_error(400, "state must be JSON")
if not isinstance(decoded, dict):
if not isinstance(state_value, dict):
return _http_error(400, "state must be an object")
try:
state = write_webui_sidebar_state(cast(dict[str, Any], decoded))
state = write_webui_sidebar_state(cast(dict[str, Any], state_value))
except ValueError as e:
return _http_error(400, str(e))
except OSError:
@@ -1208,16 +1330,10 @@ class GatewayHTTPHandler:
def _automation_values_from_request(request: WsRequest) -> dict[str, Any] | None:
raw = _case_insensitive_header(request.headers, _AUTOMATION_VALUES_HEADER)
if not raw:
payload = _mutation_payload(request)
if payload is None or "values" not in payload:
return {}
try:
values = json.loads(raw)
except Exception:
try:
values = json.loads(unquote(raw))
except Exception:
return None
values = payload.get("values")
return cast(dict[str, Any], values) if isinstance(values, dict) else None
+1
View File
@@ -37,6 +37,7 @@ dependencies = [
"readability-lxml>=0.8.4,<1.0.0",
"lxml-html-clean>=0.4.0,<1.0.0",
"rich>=14.0.0,<15.0.0",
"qrcode[pil]>=8.0",
"croniter>=6.0.0,<7.0.0",
"prompt-toolkit>=3.0.50,<4.0.0",
"questionary>=2.0.0,<3.0.0",
+9 -28
View File
@@ -80,8 +80,6 @@ def _make_fake_compact(
track_archived: list | None = None,
track_count: bool = False,
):
from nanobot.session.manager import Session as _Session
state = {"count": 0}
async def _fake_compact(key: str, *, runtime, max_suffix: int = 8) -> str:
@@ -92,25 +90,8 @@ def _make_fake_compact(
if not tail:
loop.sessions.save(session)
return ""
probe = _Session(
key=session.key,
messages=tail.copy(),
created_at=session.created_at,
updated_at=session.updated_at,
metadata={},
last_consolidated=0,
)
result = probe.retain_recent_legal_suffix(
max_suffix,
extend_to_user=True,
)
visible_suffix = probe.messages
archive_msgs = result.dropped
if not archive_msgs:
loop.sessions.save(session)
return ""
archive_end = session.last_consolidated + len(tail)
archive_msgs = tail
last_active = session.updated_at
s = summary
@@ -126,7 +107,7 @@ def _make_fake_compact(
"last_active": last_active.isoformat(),
}
session.last_consolidated = len(session.messages) - len(visible_suffix)
session.last_consolidated = archive_end
loop.sessions.save(session)
return s
@@ -365,7 +346,7 @@ class TestAutoCompact:
await loop.close_mcp()
@pytest.mark.asyncio
async def test_auto_compact_archives_prefix_without_deleting_history(self, tmp_path):
async def test_auto_compact_archives_full_tail_without_deleting_history(self, tmp_path):
loop = _make_loop(tmp_path, session_ttl_minutes=15)
session = loop.sessions.get_or_create("cli:test")
_add_turns(session, 6)
@@ -378,7 +359,7 @@ class TestAutoCompact:
await loop.auto_compact._archive("cli:test", runtime=loop.llm_runtime())
assert len(archived_messages) == 4
assert len(archived_messages) == 12
session_after = loop.sessions.get_or_create("cli:test")
assert len(session_after.messages) == 12
assert session_after.messages[0]["content"] == "msg user 0"
@@ -473,7 +454,7 @@ class TestAutoCompact:
await loop.auto_compact._archive("cli:test", runtime=loop.llm_runtime())
assert len(archived_messages) == 2
assert len(archived_messages) == 10
await loop.close_mcp()
@@ -515,7 +496,7 @@ class TestAutoCompactIdleDetection:
await loop._process_message(msg)
session_after = loop.sessions.get_or_create("cli:test")
assert len(archived_messages) == 4
assert len(archived_messages) == 12
assert any(m["content"] == "old user 0" for m in session_after.messages)
assert not any(
m["content"] == "old user 0"
@@ -724,7 +705,7 @@ class TestAutoCompactEdgeCases:
await loop._process_message(msg)
session_after = loop.sessions.get_or_create("cli:test")
assert archived_messages == []
assert [message["content"] for message in archived_messages] == ["previous message"]
assert any(m["content"] == "previous message" for m in session_after.messages)
assert any(m["content"] == "interrupted response" for m in session_after.messages)
@@ -912,7 +893,7 @@ class TestProactiveAutoCompact:
assert len(session_after.get_history(max_messages=10)) == (
loop.auto_compact._RECENT_SUFFIX_MESSAGES
)
assert len(archived_messages) == 2
assert len(archived_messages) == 10
entry = loop.auto_compact._summaries.get("cli:test")
assert entry is not None
assert entry[0] == "User chatted about old things."
+26 -2
View File
@@ -405,13 +405,37 @@ class TestCheckExpired:
scheduler.assert_not_called()
assert "dream:20260602-155256" not in ac._archiving
def test_already_trimmed_session_skips(self):
"""Expired session with no removable tail should not be re-scheduled."""
def test_short_unarchived_session_schedules(self):
"""A short idle session still needs an archive entry for Dream."""
ac = _make_autocompact(ttl=15)
mock_sm = MagicMock(spec=SessionManager)
last_active = datetime(2026, 1, 1, 10, 0, 0)
session = _make_session("cli:short", updated_at=last_active)
_add_turns(session, 2)
mock_sm.list_sessions.return_value = [
{"key": "cli:short", "updated_at": last_active.isoformat()},
]
mock_sm.get_or_create.return_value = session
ac.sessions = mock_sm
scheduled = []
def scheduler(coro):
scheduled.append(coro)
coro.close()
ac.check_expired(scheduler, _runtime)
assert len(scheduled) == 1
assert ac._archiving == {"cli:short"}
def test_fully_archived_session_skips(self):
ac = _make_autocompact(ttl=15)
mock_sm = MagicMock(spec=SessionManager)
last_active = datetime(2026, 1, 1, 10, 0, 0)
session = _make_session("cli:done", updated_at=last_active)
_add_turns(session, 2)
session.last_consolidated = len(session.messages)
mock_sm.list_sessions.return_value = [
{"key": "cli:done", "updated_at": last_active.isoformat()},
]
+110 -11
View File
@@ -391,6 +391,25 @@ class TestConsolidatorTokenBudget:
assert len(captured["history"]) == 160
assert captured["history"][0]["content"].endswith("msg-0")
async def test_estimate_includes_recent_archived_replay(self, consolidator, runtime):
session = Session(key="test:archived-replay")
for i in range(10):
session.add_message("user", f"msg-{i}")
session.last_consolidated = len(session.messages)
captured: dict[str, list[dict]] = {}
def build_messages(**kwargs):
captured["history"] = kwargs["history"]
return kwargs["history"]
consolidator._build_messages = build_messages
consolidator.estimate_session_prompt_tokens(session, runtime=runtime)
assert len(captured["history"]) == 8
assert captured["history"][0]["content"] == "msg-2"
async def test_replay_window_overflow_is_archived_even_under_token_budget(
self,
consolidator,
@@ -620,7 +639,7 @@ class TestCompactIdleSession:
)
@pytest.mark.asyncio
async def test_archives_prefix_preserves_messages_and_hides_prefix(
async def test_archives_full_tail_preserves_messages_and_replays_recent_suffix(
self, real_consolidator, mock_provider, runtime
):
mock_provider.chat_with_retry.return_value = MagicMock(
@@ -645,7 +664,7 @@ class TestCompactIdleSession:
reloaded = sessions.get_or_create("cli:test")
assert len(reloaded.messages) == 40
assert reloaded.messages[0]["content"] == "user msg 0"
assert reloaded.last_consolidated == 32
assert reloaded.last_consolidated == 40
assert reloaded.provider_state is None
visible = reloaded.get_history(max_messages=40)
assert len(visible) == 8
@@ -657,6 +676,82 @@ class TestCompactIdleSession:
assert "last_active" in meta
assert reloaded.updated_at == old_ts
@pytest.mark.asyncio
async def test_short_idle_session_archives_once(
self, real_consolidator, mock_provider, store, runtime
):
mock_provider.chat_with_retry.return_value = MagicMock(
content="Short summary.", finish_reason="stop"
)
sessions = real_consolidator.sessions
session = sessions.get_or_create("cli:short")
session.add_message("user", "hello")
session.add_message("assistant", "hi")
sessions.save(session)
first = await real_consolidator.compact_idle_session("cli:short", runtime=runtime)
second = await real_consolidator.compact_idle_session("cli:short", runtime=runtime)
assert first == "Short summary."
assert second == ""
mock_provider.chat_with_retry.assert_awaited_once()
assert len(store.read_unprocessed_history(since_cursor=0)) == 1
reloaded = sessions.get_or_create("cli:short")
assert reloaded.last_consolidated == 2
assert [message["content"] for message in reloaded.get_history()] == ["hello", "hi"]
@pytest.mark.asyncio
async def test_new_messages_advance_existing_archive_progress(
self, real_consolidator, mock_provider, runtime
):
mock_provider.chat_with_retry.return_value = MagicMock(
content="Summary.", finish_reason="stop"
)
sessions = real_consolidator.sessions
session = sessions.get_or_create("cli:incremental")
session.add_message("user", "first user")
session.add_message("assistant", "first assistant")
sessions.save(session)
await real_consolidator.compact_idle_session("cli:incremental", runtime=runtime)
current = sessions.get_or_create("cli:incremental")
current.add_message("user", "second user")
current.add_message("assistant", "second assistant")
sessions.save(current)
await real_consolidator.compact_idle_session("cli:incremental", runtime=runtime)
assert mock_provider.chat_with_retry.await_count == 2
latest_prompt = mock_provider.chat_with_retry.await_args_list[-1].kwargs["messages"][1][
"content"
]
assert "second user" in latest_prompt
assert "first user" not in latest_prompt
assert sessions.get_or_create("cli:incremental").last_consolidated == 4
@pytest.mark.asyncio
async def test_concurrent_append_remains_unarchived(
self, real_consolidator, mock_provider, runtime
):
sessions = real_consolidator.sessions
session = sessions.get_or_create("cli:concurrent")
session.add_message("user", "captured user")
session.add_message("assistant", "captured assistant")
sessions.save(session)
async def append_during_archive(**_kwargs):
current = sessions.get_or_create("cli:concurrent")
current.add_message("user", "late user")
current.add_message("assistant", "late assistant")
return LLMResponse(content="Summary.", finish_reason="stop")
mock_provider.chat_with_retry.side_effect = append_during_archive
await real_consolidator.compact_idle_session("cli:concurrent", runtime=runtime)
reloaded = sessions.get_or_create("cli:concurrent")
assert len(reloaded.messages) == 4
assert reloaded.last_consolidated == 2
@pytest.mark.asyncio
async def test_summarizes_retained_suffix_not_just_dropped_prefix(
self, real_consolidator, mock_provider, runtime
@@ -686,10 +781,10 @@ class TestCompactIdleSession:
assert "CORRECTED_FINAL_RESULT_alpha" in summarized
@pytest.mark.asyncio
async def test_raw_dumps_only_dropped_messages_on_llm_failure(
async def test_raw_dumps_full_archive_batch_on_llm_failure(
self, real_consolidator, mock_provider, store, runtime
):
"""Extra summary context must not enter raw fallback. Regression for #4264."""
"""The fallback covers the same full range as successful idle archival."""
mock_provider.chat_with_retry.side_effect = RuntimeError("LLM unavailable")
sessions = real_consolidator.sessions
session = sessions.get_or_create("cli:rawdrop")
@@ -707,7 +802,7 @@ class TestCompactIdleSession:
raw = "\n".join(e["content"] for e in store.read_unprocessed_history(since_cursor=0))
assert "[RAW]" in raw
assert "user msg 0" in raw
assert "RETAINED_SUFFIX_marker" not in raw
assert "RETAINED_SUFFIX_marker" in raw
reloaded = sessions.get_or_create("cli:rawdrop")
assert len(reloaded.messages) == 38
assert reloaded.messages[-1]["content"] == "RETAINED_SUFFIX_marker"
@@ -805,8 +900,12 @@ class TestCompactIdleSession:
reloaded = sessions.get_or_create("cli:fail")
assert len(reloaded.messages) == 20
assert reloaded.messages[0]["content"] == "u0"
assert reloaded.last_consolidated == 16
assert reloaded.last_consolidated == 20
assert [m["content"] for m in reloaded.get_history(max_messages=20)] == [
"u6",
"a6",
"u7",
"a7",
"u8",
"a8",
"u9",
@@ -835,10 +934,10 @@ class TestCompactIdleSession:
assert result == "Tail summary."
reloaded = sessions.get_or_create("cli:offset")
assert len(reloaded.messages) == 60
assert reloaded.last_consolidated == 56
assert reloaded.last_consolidated == 60
# Verify only the unconsolidated tail was processed:
# 10 unconsolidated messages (50-59), keep suffix of 4 → archive 6
# All 10 unconsolidated messages (50-59) are archived exactly once.
archived_call = mock_provider.chat_with_retry.call_args
user_content = archived_call.kwargs["messages"][1]["content"]
# Should contain only tail messages, not early ones
@@ -846,7 +945,7 @@ class TestCompactIdleSession:
assert "u25" in user_content or "a25" in user_content
@pytest.mark.asyncio
async def test_extended_suffix_archives_only_hidden_prefix(
async def test_full_archive_keeps_extended_legal_replay_suffix(
self,
real_consolidator,
mock_provider,
@@ -870,7 +969,7 @@ class TestCompactIdleSession:
reloaded = sessions.get_or_create("cli:noncontiguous")
assert len(reloaded.messages) == 25
assert reloaded.last_consolidated == 14
assert reloaded.last_consolidated == 25
assert [m["content"] for m in reloaded.get_history(max_messages=25)] == [
"user-14",
"assistant-00",
@@ -1034,7 +1133,7 @@ class TestConsolidatorSessionRefresh:
session_after = sessions.get_or_create("cli:test")
assert len(session_after.messages) == 40
assert session_after.last_consolidated == 32
assert session_after.last_consolidated == 40
assert len(session_after.get_history(max_messages=40)) == 8
-13
View File
@@ -169,19 +169,6 @@ class TestDiffCommits:
assert git_ready.diff_commits("deadbeef", "cafebabe") == ""
class TestFindCommit:
def test_finds_by_prefix(self, git_ready):
ws = git_ready._workspace
(ws / "SOUL.md").write_text("v2", encoding="utf-8")
sha = git_ready.auto_commit("v2")
found = git_ready.find_commit(sha[:4])
assert found is not None
assert found.sha == sha
def test_returns_none_for_unknown(self, git_ready):
assert git_ready.find_commit("deadbeef") is None
class TestShowCommitDiff:
def test_returns_commit_with_diff(self, git_ready):
ws = git_ready._workspace
+2 -3
View File
@@ -1074,9 +1074,8 @@ async def test_process_message_persists_media_paths_on_user_turn(tmp_path: Path)
"""User turns that attach images must record the media paths alongside
the text so the webui can rehydrate previews on session replay.
This is the producer half of the signed-media-URL round-trip: paths are
stored here, then :meth:`WebSocketChannel._augment_media_urls` maps them
onto signed URLs on the way out.
The WebUI transcript replay can use these paths to restore attachment
previews when it backfills from canonical session history.
"""
img_a = tmp_path / "uuid-1.png"
img_a.write_bytes(_PNG_1X1)
+165
View File
@@ -0,0 +1,165 @@
import asyncio
from unittest.mock import AsyncMock, MagicMock
import pytest
from nanobot.agent.loop import AgentLoop
from nanobot.bus.events import (
INBOUND_META_RUNTIME_CONTROL,
RUNTIME_CONTROL_SESSION_DISCARD,
InboundMessage,
)
from nanobot.bus.queue import MessageBus
from nanobot.providers.base import GenerationSettings, LLMResponse
from nanobot.session.keys import UNIFIED_SESSION_KEY
def _message(key: str, content: str) -> InboundMessage:
return InboundMessage(
channel="websocket",
sender_id="user",
chat_id=key.removeprefix("websocket:"),
content=content,
session_key_override=key,
require_existing_session=True,
)
def _loop(tmp_path, responses: list[str], **kwargs) -> AgentLoop:
provider = MagicMock()
provider.get_default_model.return_value = "test-model"
provider.generation = GenerationSettings()
provider.chat_with_retry = AsyncMock(
side_effect=[LLMResponse(content=response, usage={}) for response in responses]
)
return AgentLoop(
bus=MessageBus(),
provider=provider,
workspace=tmp_path,
model="test-model",
cron_service=MagicMock(),
**kwargs,
)
@pytest.mark.asyncio
async def test_transient_session_keeps_history_without_persisting_or_durable_tools(tmp_path) -> None:
loop = _loop(tmp_path, ["first answer", "second answer"])
loop.context.memory.write_memory("private durable memory")
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock()
key = "websocket:transient-test"
loop.sessions.get_or_create_transient(
key,
disabled_tools={"create_goal", "update_goal", "spawn", "cron"},
)
await loop._process_message(_message(key, "first question"))
await loop._process_message(_message(key, "second question"))
calls = loop.provider.chat_with_retry.await_args_list
assert "private durable memory" not in str(calls[0].kwargs["messages"])
tool_names = {item["function"]["name"] for item in calls[0].kwargs["tools"]}
assert "read_session" in tool_names
assert {"create_goal", "update_goal", "spawn", "cron"}.isdisjoint(tool_names)
assert "first answer" in str(calls[1].kwargs["messages"])
session = loop.sessions.get_cached(key)
assert session is not None
assert [message["role"] for message in session.messages] == [
"user",
"assistant",
"user",
"assistant",
]
assert loop.sessions.read_session_file(key) is None
loop.consolidator.maybe_consolidate_by_tokens.assert_not_awaited()
@pytest.mark.asyncio
async def test_transient_session_stays_outside_unified_session(tmp_path) -> None:
loop = _loop(tmp_path, ["private answer"], unified_session=True)
durable = loop.sessions.get_or_create(UNIFIED_SESSION_KEY)
durable.add_message("user", "durable question")
loop.sessions.save(durable)
key = "websocket:transient-unified"
transient = loop.sessions.get_or_create_transient(key)
await loop._dispatch(_message(key, "private question"))
assert [message["content"] for message in transient.messages] == [
"private question",
"private answer",
]
assert [message["content"] for message in durable.messages] == ["durable question"]
assert loop.sessions.read_session_file(key) is None
@pytest.mark.asyncio
async def test_missing_required_session_cannot_fall_back_to_disk(tmp_path) -> None:
loop = _loop(tmp_path, [])
key = "websocket:transient-stale"
loop.sessions.get_or_create_transient(key)
loop.sessions.invalidate(key)
with pytest.raises(RuntimeError, match="required session is not active"):
await loop._process_message(_message(key, "stale private message"))
loop.provider.chat_with_retry.assert_not_awaited()
assert loop.sessions.read_session_file(key) is None
@pytest.mark.asyncio
async def test_session_discard_control_cancels_active_turn(tmp_path, monkeypatch) -> None:
provider_started = asyncio.Event()
async def block_provider(**_kwargs: object) -> LLMResponse:
provider_started.set()
await asyncio.Event().wait()
raise AssertionError("provider blocker unexpectedly released")
loop = _loop(tmp_path, [])
async def wait_for_discard(key: str) -> None:
while loop.sessions.get_cached(key) is not None or key in loop._discarding_sessions:
await asyncio.sleep(0)
loop.provider.chat_with_retry = AsyncMock(side_effect=block_provider)
monkeypatch.setattr(loop, "_connect_mcp", AsyncMock())
monkeypatch.setattr(loop, "close_mcp", AsyncMock())
terminate_exec_sessions = AsyncMock(return_value=1)
monkeypatch.setattr(
loop._exec_session_manager,
"terminate_by_owner",
terminate_exec_sessions,
)
key = "websocket:transient-cancelled"
loop.sessions.get_or_create_transient(
key,
disabled_tools={"create_goal", "update_goal", "spawn", "cron"},
)
run_task = asyncio.create_task(loop.run())
await loop.bus.publish_inbound(_message(key, "private"))
await asyncio.wait_for(provider_started.wait(), timeout=2)
active_task = next(iter(loop._active_tasks[key]))
await loop.bus.publish_inbound(
InboundMessage(
channel="websocket",
sender_id="webui",
chat_id="transient-cancelled",
content="",
metadata={
INBOUND_META_RUNTIME_CONTROL: RUNTIME_CONTROL_SESSION_DISCARD,
},
session_key_override=key,
)
)
with pytest.raises(asyncio.CancelledError):
await asyncio.wait_for(active_task, timeout=2)
await asyncio.wait_for(wait_for_discard(key), timeout=2)
assert loop.sessions.get_cached(key) is None
terminate_exec_sessions.assert_awaited_once_with(key)
loop.stop()
await loop.bus.publish_inbound(_message(key, "wake"))
await asyncio.wait_for(run_task, timeout=2)
+17
View File
@@ -1092,6 +1092,23 @@ class TestMainMenuUpdate:
assert config.providers.openai.api_key == "${UNRELATED_MISSING_KEY}"
assert config.providers.openai_codex.proxy == "${CODEX_PROXY}"
def test_quick_start_openai_codex_reports_incomplete_installation(self, monkeypatch):
import oauth_cli_kit
messages: list[str] = []
monkeypatch.delattr(oauth_cli_kit, "get_token")
monkeypatch.setattr(
onboard_wizard.console,
"print",
lambda message, *args, **kwargs: messages.append(str(message)),
)
assert onboard_wizard._quick_start_oauth_login(Config(), "openai_codex") is False
assert messages == [
"[red]This nanobot installation is missing the required oauth-cli-kit package. "
"Reinstall or upgrade nanobot-ai using the same installation method.[/red]"
]
def test_quick_start_openai_codex_runs_interactive_login_for_bad_cached_token(
self, monkeypatch
):
+11 -2
View File
@@ -9,6 +9,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from agent.runner_helpers import make_run_spec
from nanobot.agent.automation_turns import publish_next_deferred_turn
from nanobot.config.schema import AgentDefaults
from nanobot.providers.base import LLMResponse, ToolCallRequest
@@ -1047,7 +1048,11 @@ async def test_cron_turn_deferred_while_session_active(tmp_path):
assert loop._cron_turns.deferred_queues[session_key] == [msg]
assert loop.pending_cron_job_ids_for_session(session_key) == {"job-1"}
await loop._cron_turns.publish_next_deferred(session_key)
await publish_next_deferred_turn(
deferred_queues=loop._cron_turns.deferred_queues,
publish_inbound=loop.bus.publish_inbound,
session_key=session_key,
)
queued = await asyncio.wait_for(loop.bus.consume_inbound(), timeout=0.5)
assert queued is msg
assert session_key not in loop._cron_turns.deferred_queues
@@ -1097,7 +1102,11 @@ async def test_local_trigger_turn_deferred_while_session_active(tmp_path):
assert loop._local_trigger_turns.deferred_queues[session_key] == [msg]
assert loop.pending_local_trigger_ids_for_session(session_key) == {"trg_123"}
assert await loop._local_trigger_turns.publish_next_deferred(session_key) is True
assert await publish_next_deferred_turn(
deferred_queues=loop._local_trigger_turns.deferred_queues,
publish_inbound=loop.bus.publish_inbound,
session_key=session_key,
) is True
queued = await asyncio.wait_for(loop.bus.consume_inbound(), timeout=0.5)
assert queued is msg
assert session_key not in loop._local_trigger_turns.deferred_queues
+16 -8
View File
@@ -5,6 +5,7 @@ import pytest
from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools.context import RequestContext, request_context
from nanobot.agent.tools.runtime_control import AgentRuntimeControl
from nanobot.agent.tools.self import MyTool
from nanobot.bus.queue import MessageBus
from nanobot.config.schema import ModelPresetConfig
@@ -34,6 +35,13 @@ def _make_loop(tmp_path, presets=None, active_preset=None):
)
def _my_tool(loop: AgentLoop) -> MyTool:
return MyTool(
runtime_control=AgentRuntimeControl(loop),
modify_allowed=True,
)
def test_model_preset_getter_none_when_not_set(tmp_path) -> None:
loop = _make_loop(tmp_path)
assert loop.model_preset is None
@@ -240,7 +248,7 @@ def test_self_tool_inspect_shows_model_preset(tmp_path) -> None:
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
}
loop = _make_loop(tmp_path, presets=presets, active_preset="fast")
tool = MyTool(runtime_state=loop, modify_allowed=True)
tool = _my_tool(loop)
output = tool._inspect_all()
assert "model_preset: 'fast'" in output
@@ -250,7 +258,7 @@ def test_self_tool_set_model_preset_via_modify(tmp_path) -> None:
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
}
loop = _make_loop(tmp_path, presets=presets)
tool = MyTool(runtime_state=loop, modify_allowed=True)
tool = _my_tool(loop)
result = tool._modify("model_preset", "fast")
assert "Error" not in result
assert loop.model_preset == "fast"
@@ -263,7 +271,7 @@ def test_self_tool_set_model_preset_switches_back_to_default(tmp_path) -> None:
"fast": ModelPresetConfig(model="openai/gpt-4.1", context_window_tokens=32_768),
}
loop = _make_loop(tmp_path, presets=presets, active_preset="fast")
tool = MyTool(runtime_state=loop, modify_allowed=True)
tool = _my_tool(loop)
result = tool._modify("model_preset", "default")
@@ -280,7 +288,7 @@ def test_self_tool_set_model_preset_unknown_lists_available(tmp_path) -> None:
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
}
loop = _make_loop(tmp_path, presets=presets)
tool = MyTool(runtime_state=loop, modify_allowed=True)
tool = _my_tool(loop)
result = tool._modify("model_preset", "missing")
@@ -295,7 +303,7 @@ def test_self_tool_sets_model_preset_for_current_session(tmp_path) -> None:
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
}
loop = _make_loop(tmp_path, presets=presets)
tool = MyTool(runtime_state=loop, modify_allowed=True)
tool = _my_tool(loop)
with request_context(RequestContext(
channel="cli",
@@ -318,7 +326,7 @@ def test_self_tool_reports_session_preset_provider_configuration_error(tmp_path)
loop.set_session_model_preset = MagicMock(
side_effect=ValueError("No API key configured for provider 'openai'.")
)
tool = MyTool(runtime_state=loop, modify_allowed=True)
tool = _my_tool(loop)
with request_context(RequestContext(
channel="cli",
@@ -343,7 +351,7 @@ def test_self_tool_rejects_instance_runtime_changes_in_session(
value: object,
) -> None:
loop = _make_loop(tmp_path)
tool = MyTool(runtime_state=loop, modify_allowed=True)
tool = _my_tool(loop)
session = loop.sessions.get_or_create("cli:one")
with request_context(RequestContext(
@@ -366,7 +374,7 @@ def test_self_tool_set_model_clears_active_preset(tmp_path) -> None:
"fast": ModelPresetConfig(model="openai/gpt-4.1"),
}
loop = _make_loop(tmp_path, presets=presets, active_preset="fast")
tool = MyTool(runtime_state=loop, modify_allowed=True)
tool = _my_tool(loop)
result = tool._modify("model", "anthropic/claude-opus-4-5")
assert "Error" not in result
assert loop.model_preset is None

Some files were not shown because too many files have changed in this diff Show More