mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-08 21:38:40 +03:00
Compare commits
32
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0ccb7234d4 | ||
|
|
34535b4e7c | ||
|
|
5abe06f808 | ||
|
|
64c7ff5fdc | ||
|
|
54bcdb5a62 | ||
|
|
ffdf05a603 | ||
|
|
fd9e57703c | ||
|
|
661ab00656 | ||
|
|
b941233138 | ||
|
|
acb0e853ff | ||
|
|
afef27dd6c | ||
|
|
f32007c83f | ||
|
|
09bde468eb | ||
|
|
2ebf5c4972 | ||
|
|
1ed2c9a213 | ||
|
|
55b550ee01 | ||
|
|
b690a48336 | ||
|
|
7178ea3f13 | ||
|
|
2a0cd19a74 | ||
|
|
c78421cf16 | ||
|
|
03be51ade5 | ||
|
|
c757c5466c | ||
|
|
5f4cfbcb16 | ||
|
|
f6d1dba32a | ||
|
|
2ec4044217 | ||
|
|
a6d5e4f3b5 | ||
|
|
ed48325346 | ||
|
|
56443ac6e2 | ||
|
|
21aa900d64 | ||
|
|
b0258e8b20 | ||
|
|
8493560976 | ||
|
|
8d2c31eb6a |
@@ -107,6 +107,7 @@ File operations have path traversal protection, but:
|
||||
**API Calls:**
|
||||
- All external API calls use HTTPS by default
|
||||
- Timeouts are configured to prevent hanging requests
|
||||
- The OpenAI-compatible API server must set `api.api_key` when binding to `0.0.0.0` or `::`; otherwise startup fails to prevent unauthenticated network access
|
||||
- Consider using a firewall to restrict outbound connections if needed
|
||||
|
||||
**WhatsApp:**
|
||||
|
||||
+1
-1
@@ -51,7 +51,7 @@ If a local `nanobot agent` session can already answer normally, you can also ask
|
||||
|---|---|---|
|
||||
| Open the bundled browser UI | [`webui.md`](./webui.md) | WebUI on port `8765`, chat workspace, Apps, Skills, Automations, and settings |
|
||||
| Connect Telegram, Discord, WeChat, Slack, and other apps | [`chat-apps.md`](./chat-apps.md) | A gateway-backed chat channel with access control |
|
||||
| Use slash commands and periodic tasks | [`chat-commands.md`](./chat-commands.md) | Pairing, model presets, heartbeat tasks, and chat-side controls |
|
||||
| Use slash commands and automations | [`chat-commands.md`](./chat-commands.md) | Pairing, model presets, local triggers, heartbeat tasks, and chat-side controls |
|
||||
| Generate images | [`image-generation.md`](./image-generation.md) | Image provider config, WebUI image mode, and artifact behavior |
|
||||
| Run several isolated bots | [`multiple-instances.md`](./multiple-instances.md) | Separate configs, workspaces, ports, and sessions |
|
||||
| Deploy outside a terminal | [`deployment.md`](./deployment.md) | Docker, systemd user services, and macOS LaunchAgent setup |
|
||||
|
||||
@@ -103,7 +103,8 @@ class WebhookChannel(BaseChannel):
|
||||
msg.content — markdown text (convert to platform format as needed)
|
||||
msg.media — list of local file paths to attach
|
||||
msg.chat_id — the recipient (same chat_id you passed to _handle_message)
|
||||
msg.metadata — may contain "_progress": True for streaming chunks
|
||||
msg.metadata — channel routing context such as message/thread ids
|
||||
msg.event — typed runtime event for progress/status messages
|
||||
"""
|
||||
logger.info("[webhook] -> {}: {}", msg.chat_id, msg.content[:80])
|
||||
# In a real plugin: POST to a callback URL, send via SDK, etc.
|
||||
@@ -238,15 +239,15 @@ nanobot channels login <channel_name> --force # re-authenticate
|
||||
| `supports_streaming` (property) | `True` when config has `"streaming": true` **and** subclass overrides `send_delta()`. |
|
||||
| `is_running` | Returns `self._running`. |
|
||||
| `login(force=False)` | Perform interactive login (e.g. QR code scan). Returns `True` if already authenticated or login succeeds. Override in subclasses that support interactive login. |
|
||||
| `send_reasoning_delta(chat_id, delta, metadata?)` | Optional hook for streamed model reasoning/thinking content. Default is no-op. |
|
||||
| `send_reasoning_end(chat_id, metadata?)` | Optional hook marking the end of a reasoning block. Default is no-op. |
|
||||
| `send_reasoning_delta(chat_id, delta, metadata?, *, stream_id?)` | Optional hook for streamed model reasoning/thinking content. Default is no-op. |
|
||||
| `send_reasoning_end(chat_id, metadata?, *, stream_id?)` | Optional hook marking the end of a reasoning block. Default is no-op. |
|
||||
| `send_reasoning(msg)` | Optional one-shot reasoning fallback. Default translates to `send_reasoning_delta()` + `send_reasoning_end()`. |
|
||||
|
||||
### Optional (streaming)
|
||||
|
||||
| Method | Description |
|
||||
|--------|-------------|
|
||||
| `async send_delta(chat_id, delta, metadata?)` | Override to receive streaming chunks. See [Streaming Support](#streaming-support) for details. |
|
||||
| `async send_delta(chat_id, delta, metadata?, *, stream_id?, stream_end=False, resuming=False)` | Override to receive streaming chunks. See [Streaming Support](#streaming-support) for details. |
|
||||
|
||||
### Message Types
|
||||
|
||||
@@ -257,10 +258,12 @@ class OutboundMessage:
|
||||
chat_id: str # recipient (same value you passed to _handle_message)
|
||||
content: str # markdown text — convert to platform format as needed
|
||||
media: list[str] # local file paths to attach (images, audio, docs)
|
||||
metadata: dict # may contain: "_progress" (bool) for streaming chunks,
|
||||
# "message_id" for reply threading
|
||||
metadata: dict # channel routing context, e.g. "message_id" for threading
|
||||
event: object | None # typed runtime/UI event; usually inspect with isinstance()
|
||||
```
|
||||
|
||||
Runtime/UI semantics live on `msg.event`. Plugin-authored outbound messages should use typed events instead of legacy metadata flags such as `_progress`, `_stream_delta`, `_stream_end`, `_reasoning_delta`, `_turn_end`, or `_goal_status`. nanobot still accepts those old flags as a compatibility bridge for existing in-process extensions, but new plugin code should not add fresh dependencies on them.
|
||||
|
||||
## Streaming Support
|
||||
|
||||
Channels can opt into real-time streaming — the agent sends content token-by-token instead of one final message. This is entirely optional; channels work fine without it.
|
||||
@@ -279,10 +282,18 @@ If either is missing, the agent falls back to the normal one-shot `send()` path.
|
||||
Override `send_delta` to handle two types of calls:
|
||||
|
||||
```python
|
||||
async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None) -> None:
|
||||
meta = metadata or {}
|
||||
|
||||
if meta.get("_stream_end"):
|
||||
async def send_delta(
|
||||
self,
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
stream_end: bool = False,
|
||||
resuming: bool = False,
|
||||
) -> None:
|
||||
buffer_key = stream_id or chat_id
|
||||
if stream_end:
|
||||
# Streaming finished — do final formatting, cleanup, etc.
|
||||
return
|
||||
|
||||
@@ -290,12 +301,7 @@ async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] |
|
||||
# delta contains a small chunk of text (a few tokens)
|
||||
```
|
||||
|
||||
**Metadata flags:**
|
||||
|
||||
| Flag | Meaning |
|
||||
|------|---------|
|
||||
| `_stream_delta: True` | A content chunk (delta contains the new text) |
|
||||
| `_stream_end: True` | Streaming finished (delta is empty) |
|
||||
Streaming state is passed through keyword-only arguments, not `_stream_delta` or `_stream_end` metadata flags. Use `stream_id` to key any per-stream buffers; fall back to `chat_id` when it is missing.
|
||||
|
||||
### Example: Webhook with Streaming
|
||||
|
||||
@@ -310,18 +316,27 @@ class WebhookChannel(BaseChannel):
|
||||
super().__init__(config, bus)
|
||||
self._buffers: dict[str, str] = {}
|
||||
|
||||
async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None) -> None:
|
||||
meta = metadata or {}
|
||||
if meta.get("_stream_end"):
|
||||
text = self._buffers.pop(chat_id, "")
|
||||
async def send_delta(
|
||||
self,
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
stream_end: bool = False,
|
||||
resuming: bool = False,
|
||||
) -> None:
|
||||
buffer_key = stream_id or chat_id
|
||||
if stream_end:
|
||||
text = self._buffers.pop(buffer_key, "")
|
||||
# Final delivery — format and send the complete message
|
||||
await self._deliver(chat_id, text, final=True)
|
||||
return
|
||||
|
||||
self._buffers.setdefault(chat_id, "")
|
||||
self._buffers[chat_id] += delta
|
||||
self._buffers.setdefault(buffer_key, "")
|
||||
self._buffers[buffer_key] += delta
|
||||
# Incremental update — push partial text to the client
|
||||
await self._deliver(chat_id, self._buffers[chat_id], final=False)
|
||||
await self._deliver(chat_id, self._buffers[buffer_key], final=False)
|
||||
|
||||
async def send(self, msg: OutboundMessage) -> None:
|
||||
# Non-streaming path — unchanged
|
||||
@@ -350,7 +365,7 @@ When `streaming` is `false` (default) or omitted, only `send()` is called — no
|
||||
|
||||
| Method / Property | Description |
|
||||
|-------------------|-------------|
|
||||
| `async send_delta(chat_id, delta, metadata?)` | Override to handle streaming chunks. No-op by default. |
|
||||
| `async send_delta(chat_id, delta, metadata?, *, stream_id?, stream_end=False, resuming=False)` | Override to handle streaming chunks. No-op by default. |
|
||||
| `supports_streaming` (property) | Returns `True` when config has `streaming: true` **and** subclass overrides `send_delta`. |
|
||||
|
||||
## Progress, Tool Hints, and Reasoning
|
||||
@@ -359,18 +374,20 @@ Besides normal assistant text, nanobot can emit low-emphasis trace blocks. These
|
||||
|
||||
### Progress and Tool Hints
|
||||
|
||||
Progress and tool hints arrive through the normal `send(msg)` path. Check `msg.metadata` before rendering:
|
||||
Progress and tool hints arrive through the normal `send(msg)` path. Check `msg.event` before rendering:
|
||||
|
||||
```python
|
||||
async def send(self, msg: OutboundMessage) -> None:
|
||||
meta = msg.metadata or {}
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
|
||||
if meta.get("_tool_hint"):
|
||||
async def send(self, msg: OutboundMessage) -> None:
|
||||
event = msg.event
|
||||
|
||||
if isinstance(event, ProgressEvent) and event.tool_hint:
|
||||
# A short tool breadcrumb, e.g. read_file("config.json")
|
||||
await self._send_trace(msg.chat_id, msg.content, kind="tool")
|
||||
return
|
||||
|
||||
if meta.get("_progress"):
|
||||
if isinstance(event, ProgressEvent):
|
||||
# Generic non-final status, e.g. "Thinking..." or "Running command..."
|
||||
await self._send_trace(msg.chat_id, msg.content, kind="progress")
|
||||
return
|
||||
@@ -412,32 +429,33 @@ class WebhookChannel(BaseChannel):
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
) -> None:
|
||||
meta = metadata or {}
|
||||
stream_id = str(meta.get("_stream_id") or chat_id)
|
||||
self._reasoning_buffers[stream_id] = self._reasoning_buffers.get(stream_id, "") + delta
|
||||
await self._update_reasoning_block(chat_id, self._reasoning_buffers[stream_id], final=False)
|
||||
buffer_key = stream_id or chat_id
|
||||
self._reasoning_buffers[buffer_key] = self._reasoning_buffers.get(buffer_key, "") + delta
|
||||
await self._update_reasoning_block(chat_id, self._reasoning_buffers[buffer_key], final=False)
|
||||
|
||||
async def send_reasoning_end(
|
||||
self,
|
||||
chat_id: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
) -> None:
|
||||
meta = metadata or {}
|
||||
stream_id = str(meta.get("_stream_id") or chat_id)
|
||||
text = self._reasoning_buffers.pop(stream_id, "")
|
||||
buffer_key = stream_id or chat_id
|
||||
text = self._reasoning_buffers.pop(buffer_key, "")
|
||||
if text:
|
||||
await self._update_reasoning_block(chat_id, text, final=True)
|
||||
```
|
||||
|
||||
**Reasoning metadata flags:**
|
||||
**Reasoning arguments:**
|
||||
|
||||
| Flag | Meaning |
|
||||
| Argument | Meaning |
|
||||
|------|---------|
|
||||
| `_reasoning_delta: True` | A reasoning/thinking chunk; `delta` contains the new text. |
|
||||
| `_reasoning_end: True` | The current reasoning block is complete; `delta` is empty. |
|
||||
| `_reasoning: True` | Legacy one-shot reasoning. `BaseChannel.send_reasoning()` converts it to delta + end. |
|
||||
| `_stream_id` | Stable id for this assistant turn/segment. Use it to key buffers instead of only `chat_id`. |
|
||||
| `delta` | A reasoning/thinking chunk for `send_reasoning_delta()`. |
|
||||
| `stream_id` | Stable id for this assistant turn/segment. Use it to key buffers instead of only `chat_id`. |
|
||||
| `send_reasoning_end()` | The current reasoning block is complete. |
|
||||
|
||||
Reasoning visibility is controlled by `showReasoning` globally or per channel:
|
||||
|
||||
|
||||
@@ -16,6 +16,8 @@ These commands work inside chat channels and interactive agent sessions:
|
||||
| `/dream-restore` | List recent Dream memory versions |
|
||||
| `/dream-restore <sha>` | Restore memory to the state before a specific change |
|
||||
| `/skill` | List enabled skills and their descriptions |
|
||||
| `/trigger` | Show local trigger usage |
|
||||
| `/trigger <name>` | Create a named local trigger for the current chat/session |
|
||||
| `/pairing` | List pending pairing requests |
|
||||
| `/pairing approve <code>` | Approve a pairing code |
|
||||
| `/pairing deny <code>` | Deny a pending pairing request |
|
||||
@@ -55,6 +57,65 @@ To switch presets for future turns:
|
||||
|
||||
Preset names come from the top-level `modelPresets` config. Switching is runtime-only: it does not rewrite `config.json`, and an in-progress turn keeps using the model it started with. See [Configuration: Model presets](./configuration.md#model-presets) for setup details.
|
||||
|
||||
## Local triggers
|
||||
|
||||
Use `/trigger <name>` when a local script or another service should be able to
|
||||
send a message into the current chat/session later. A name is required; plain
|
||||
`/trigger` only shows the usage hint.
|
||||
|
||||
Create the trigger from the chat where future messages should arrive:
|
||||
|
||||
```text
|
||||
/trigger PR review
|
||||
```
|
||||
|
||||
nanobot replies with a trigger ID and a command shaped like:
|
||||
|
||||
```bash
|
||||
nanobot trigger trg_8K4P2Q9X "Review PR #4502"
|
||||
```
|
||||
|
||||
Replace `"Review PR #4502"` with the message you want nanobot to receive. The
|
||||
trigger is bound to the session where it was created, so the message goes back
|
||||
to that same chat. Keep `nanobot gateway` running so trigger messages can be
|
||||
delivered. The trigger message starts an automation turn recorded in that
|
||||
session with the message you passed to the CLI; it is not treated as a normal
|
||||
user message. If that session is already running a turn, the trigger waits
|
||||
until the session is idle instead of being injected into the active turn.
|
||||
|
||||
Trigger deliveries are stored in the workspace until their linked agent turn
|
||||
finishes successfully. If the gateway exits after claiming a delivery but before
|
||||
the turn completes, the next gateway start requeues that delivery. This is an
|
||||
at-least-once local queue: a delivery may run more than once if the process
|
||||
exits at the wrong time, so external scripts should make repeated trigger
|
||||
messages safe. If the delivery reaches the agent and the agent turn fails, the
|
||||
delivery is marked failed in Automations instead of retrying forever.
|
||||
|
||||
For longer or generated content, omit the message argument and pipe stdin:
|
||||
|
||||
```bash
|
||||
printf '%s\n' "Review the latest failed CI job" | nanobot trigger trg_8K4P2Q9X
|
||||
```
|
||||
|
||||
If an external webhook should wake nanobot up, run your own small webhook
|
||||
service and have it call the trigger command after it builds the final message:
|
||||
|
||||
```bash
|
||||
nanobot trigger <trigger-id> "<message>"
|
||||
```
|
||||
|
||||
If you run multiple nanobot instances, pass the same config or workspace
|
||||
selector used by the gateway:
|
||||
|
||||
```bash
|
||||
nanobot trigger --config ./bot-a/config.json trg_8K4P2Q9X "Nightly report"
|
||||
nanobot trigger --workspace ./bot-a/workspace trg_8K4P2Q9X "Nightly report"
|
||||
```
|
||||
|
||||
Manage triggers from the WebUI Automations view. You can search, pause/resume,
|
||||
rename, delete, and copy the trigger command there. A session may have multiple
|
||||
triggers, just like it may have multiple scheduled automations.
|
||||
|
||||
## Periodic Tasks
|
||||
|
||||
Periodic background checks are driven by `HEARTBEAT.md` in your workspace (`~/.nanobot/workspace/HEARTBEAT.md`). When `nanobot gateway` starts, it registers a protected heartbeat cron job by default. Every 30 minutes, that job checks the file; if it finds tasks under `## Active Tasks`, the agent executes them and delivers only results that pass the notification gate to your most recently active chat channel. If there are no active tasks, or the result is routine with nothing useful to report, the heartbeat is skipped silently.
|
||||
|
||||
@@ -13,6 +13,7 @@ Use this page when you know what you want to run and need the command shape. For
|
||||
| Send one test message | `nanobot agent -m "Hello!"` | First proof that install, config, provider, model, and workspace all work |
|
||||
| Chat in the terminal | `nanobot agent` | Interactive local chat; exit with `exit`, `/exit`, `:q`, or `Ctrl+D` |
|
||||
| Use WebUI or chat apps | `nanobot gateway` | Keep this terminal running, or use `nanobot gateway --background` |
|
||||
| Deliver a local trigger | `nanobot trigger <id> "message"` | Created first with `/trigger <name>` in the target chat/session |
|
||||
| Serve an OpenAI-compatible API | `nanobot serve` | Starts `/v1/chat/completions`, `/v1/models`, and `/health` |
|
||||
| Check chat channel setup | `nanobot channels status` | Useful before starting `nanobot gateway` |
|
||||
| Log in to QR/OAuth-style channels | `nanobot channels login <channel>` | Used by channels such as WhatsApp and WeChat |
|
||||
@@ -122,6 +123,52 @@ http://127.0.0.1:18790/health
|
||||
|
||||
The bundled WebUI is served by the WebSocket channel, usually on port `8765`, not by the gateway health endpoint.
|
||||
|
||||
## Local Triggers
|
||||
|
||||
`nanobot trigger` delivers one local message to a trigger that was created from
|
||||
a chat/session with `/trigger <name>`.
|
||||
|
||||
```bash
|
||||
nanobot trigger trg_8K4P2Q9X "Review PR #4502"
|
||||
```
|
||||
|
||||
Keep `nanobot gateway` running so the message can be delivered to the linked
|
||||
chat/session. The message is recorded as an automation turn in that session,
|
||||
not as a normal chat message typed by the user.
|
||||
|
||||
The command writes to a workspace-local durable queue. If `nanobot gateway` is
|
||||
not running yet, the message waits in that workspace. If the target session is
|
||||
already running a turn, the trigger waits for that session to become idle. If the
|
||||
gateway exits after claiming a delivery but before the linked turn completes,
|
||||
the next gateway start requeues that delivery. The queue is at-least-once, not
|
||||
exactly-once, so the same message can be delivered again after an interrupted
|
||||
process. If the agent receives the delivery and the turn fails, the delivery is
|
||||
marked failed instead of retried indefinitely. Each delivery also writes an
|
||||
audit record under `<workspace>/triggers/runs`. Run one gateway consumer per
|
||||
workspace; this local queue is not a distributed multi-consumer queue.
|
||||
|
||||
Use stdin when another local process generates the message:
|
||||
|
||||
```bash
|
||||
generate-report | nanobot trigger trg_8K4P2Q9X
|
||||
```
|
||||
|
||||
Options:
|
||||
|
||||
| Command | Description |
|
||||
|---|---|
|
||||
| `nanobot trigger <id> "message"` | Deliver one message through a trigger |
|
||||
| `nanobot trigger <id>` | Read the message from stdin |
|
||||
| `nanobot trigger --config <path> <id> "message"` | Use the workspace from a specific config |
|
||||
| `nanobot trigger --workspace <path> <id> "message"` | Use a specific workspace |
|
||||
|
||||
Triggers are managed in the WebUI Automations view instead of through separate
|
||||
`list`, `revoke`, or `delete` CLI subcommands. From there you can pause/resume,
|
||||
rename, delete, search, and copy the command for each trigger.
|
||||
|
||||
For webhooks or other external systems, run your own small service and have it
|
||||
call this CLI after it decides what message nanobot should receive.
|
||||
|
||||
## OpenAI-Compatible API
|
||||
|
||||
| Command | Description |
|
||||
|
||||
+18
-3
@@ -123,7 +123,7 @@ Tools are discovered automatically from built-in modules and plugin entry points
|
||||
- shell execution with configurable sandboxing;
|
||||
- web search and web fetch with SSRF checks;
|
||||
- MCP servers;
|
||||
- cron reminders and heartbeat tasks;
|
||||
- cron reminders, local triggers, and heartbeat tasks;
|
||||
- image generation;
|
||||
- subagents and runtime self-inspection.
|
||||
|
||||
@@ -131,14 +131,29 @@ Security-sensitive controls live in [`configuration.md#security`](./configuratio
|
||||
|
||||
## Background Jobs
|
||||
|
||||
When `nanobot gateway` starts, it creates workspace-scoped cron storage at `<workspace>/cron/jobs.json` and registers system jobs:
|
||||
When `nanobot gateway` starts, it runs workspace-scoped automations and
|
||||
registers system jobs:
|
||||
|
||||
- `dream`, when `agents.defaults.dream.enabled` is true;
|
||||
- `heartbeat`, when `gateway.heartbeat.enabled` is true.
|
||||
|
||||
Heartbeat reads `<workspace>/HEARTBEAT.md`. If the file has tasks under `## Active Tasks`, nanobot executes them and sends only useful/actionable results to the most recently active chat target. Routine "nothing changed" results are suppressed.
|
||||
|
||||
User-created reminders use the same cron service but are not the same as the protected heartbeat system job. They run as scheduled turns in their origin chat/session and normally deliver the result back to that channel.
|
||||
User-created reminders use the same cron service but are not the same as the
|
||||
protected heartbeat system job. They run as scheduled turns in their origin
|
||||
chat/session and normally deliver the result back to that channel.
|
||||
|
||||
Local triggers are also session-bound, but they do not have their own
|
||||
schedule. Create one from the target chat with `/trigger <name>`, then call
|
||||
`nanobot trigger <id> "<message>"` when a local script or external service wants
|
||||
nanobot to respond in that session. Webhook servers, third-party auth, and
|
||||
event-to-message formatting stay outside nanobot. Trigger deliveries are stored
|
||||
in the workspace until the linked agent turn finishes successfully. If the
|
||||
target session is busy, the trigger waits until that session is idle instead of
|
||||
being injected into the active turn. The message is recorded as an automation
|
||||
turn in that session. Delivery is at-least-once, so external systems should
|
||||
tolerate repeated trigger messages; a delivery that reaches the agent but fails
|
||||
is marked failed rather than retried forever.
|
||||
|
||||
## Where to Go Next
|
||||
|
||||
|
||||
@@ -12,6 +12,32 @@ Run the CLI check first. If `nanobot agent -m "Hello!"` fails, fix provider or c
|
||||
|
||||
For setup help, see [`quick-start.md`](./quick-start.md), [`providers.md`](./providers.md), and [`troubleshooting.md`](./troubleshooting.md).
|
||||
|
||||
## Authentication
|
||||
|
||||
Local-only `127.0.0.1` usage does not require an API key. If you bind the API
|
||||
server to all interfaces with `api.host: "0.0.0.0"` or `"::"`, nanobot requires
|
||||
`api.apiKey`; otherwise startup fails to avoid exposing an unauthenticated agent
|
||||
endpoint on the network.
|
||||
|
||||
```json
|
||||
{
|
||||
"api": {
|
||||
"host": "0.0.0.0",
|
||||
"port": 8900,
|
||||
"apiKey": "${NANOBOT_API_KEY}"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
When `api.apiKey` is set, send it as a Bearer token on API routes. The health
|
||||
endpoint remains unauthenticated so local probes and load balancers can still
|
||||
check process health.
|
||||
|
||||
```bash
|
||||
curl http://127.0.0.1:8900/v1/models \
|
||||
-H "Authorization: Bearer $NANOBOT_API_KEY"
|
||||
```
|
||||
|
||||
## Behavior
|
||||
|
||||
- Session isolation: pass `"session_id"` in the request body to isolate conversations; omit for a shared default session (`api:default`)
|
||||
|
||||
+35
-8
@@ -56,7 +56,7 @@ Enter `tokenIssueSecret` when the WebUI asks for a password.
|
||||
| Composer | Send text, images, voice input, slash commands, and `@` mentions for Apps or MCP presets |
|
||||
| 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 |
|
||||
| Automations | Review, search, run, pause, edit, and delete scheduled agent turns |
|
||||
| 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 |
|
||||
|
||||
## Chat Workspace
|
||||
@@ -116,10 +116,30 @@ to perform that task.
|
||||
|
||||
## Automations
|
||||
|
||||
Automations are scheduled agent turns. They should be created from the chat,
|
||||
channel, or session where they are supposed to run so nanobot keeps the correct
|
||||
target context. When an automation runs, it normally delivers the result back to
|
||||
that linked chat.
|
||||
Automations are agent turns that run later in a linked chat/session. They should
|
||||
be created from the chat, channel, or session where they are supposed to run so
|
||||
nanobot keeps the correct target context. When an automation runs, it normally
|
||||
delivers the result back to that linked chat.
|
||||
|
||||
There are two user-facing automation types:
|
||||
|
||||
- Scheduled automations, created by the agent's cron tool, run at a time,
|
||||
interval, or cron expression.
|
||||
- Local triggers, created with `/trigger <name>`, run when you call a local
|
||||
command such as `nanobot trigger trg_8K4P2Q9X "Review PR #4502"`.
|
||||
|
||||
If a GitHub webhook, CI system, or another service should wake nanobot up, keep
|
||||
that webhook/service outside nanobot and have it call the trigger command with
|
||||
the final message.
|
||||
|
||||
Trigger deliveries use the same workspace as the gateway. They survive gateway
|
||||
restarts and are requeued if the process exits before the linked turn completes.
|
||||
If the linked session is already running a turn, the local trigger waits until
|
||||
that session is idle instead of being injected into the active turn. This is an
|
||||
at-least-once local queue, so repeated delivery is possible after an interrupted
|
||||
process. A delivered trigger is recorded as an automation turn in the linked
|
||||
session; if the agent receives it but the turn fails, Automations marks the run
|
||||
failed instead of retrying indefinitely.
|
||||
|
||||
For recurring background checks that should stay quiet unless there is something
|
||||
useful to report, use the protected heartbeat job by editing `HEARTBEAT.md`
|
||||
@@ -128,18 +148,25 @@ instead of creating a chat automation.
|
||||
Use the Automations view to:
|
||||
|
||||
- Filter by all, active, paused, needs-attention, or system jobs.
|
||||
- Search by task name, message, linked chat, schedule, or status.
|
||||
- Search by task name, message, trigger command, linked chat, schedule, or status.
|
||||
- Sort by next run, last run, updated time, or name.
|
||||
- Run now, pause or resume, edit, or delete user-created automations.
|
||||
- Run scheduled automations now.
|
||||
- Pause or resume, rename, or delete user-created automations.
|
||||
- Copy the CLI command for local triggers.
|
||||
- Inspect protected system automations without changing them.
|
||||
|
||||
Search accepts plain text and field filters such as `name:backup`,
|
||||
`chat:WeChat`, `schedule:09:30`, `cron:"0 23 * * *"`, and `status:paused`.
|
||||
`chat:WeChat`, `schedule:09:30`, `cron:"0 23 * * *"`, `trigger`, and
|
||||
`status:paused`.
|
||||
|
||||
An automation without a linked chat cannot be enabled or run from the WebUI,
|
||||
because nanobot would not know where to deliver the scheduled turn. Recreate it
|
||||
from the target chat or channel so the automation has complete context.
|
||||
|
||||
Local triggers do not have a WebUI "Run now" action because each run needs a
|
||||
message. Use the copied `nanobot trigger ...` command and replace `"message"`
|
||||
with the content that should be delivered.
|
||||
|
||||
## Settings
|
||||
|
||||
Settings is the control surface for the browser session and gateway-backed
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
"""Shared coordination for session-bound automation turns."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import dataclasses
|
||||
from collections.abc import Awaitable, Callable, Iterable
|
||||
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
|
||||
|
||||
class AutomationTurnError(RuntimeError):
|
||||
"""Raised when an automation turn reaches the agent and finishes with an error."""
|
||||
|
||||
|
||||
async def publish_next_deferred_turn(
|
||||
*,
|
||||
deferred_queues: dict[str, list[InboundMessage]],
|
||||
publish_inbound: Callable[[InboundMessage], Awaitable[None]],
|
||||
session_key: str,
|
||||
) -> bool:
|
||||
"""Publish the next deferred automation turn for a session."""
|
||||
queue = deferred_queues.get(session_key)
|
||||
if not queue:
|
||||
return False
|
||||
msg = queue.pop(0)
|
||||
if not queue:
|
||||
deferred_queues.pop(session_key, None)
|
||||
await publish_inbound(msg)
|
||||
return True
|
||||
|
||||
|
||||
class AutomationTurnCoordinator:
|
||||
"""Manage automation turns without mixing them into live injections."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
publish_inbound: Callable[[InboundMessage], Awaitable[None]],
|
||||
dispatch: Callable[[InboundMessage], Awaitable[object]],
|
||||
is_running: Callable[[], bool],
|
||||
turn_id: Callable[[InboundMessage], str | None],
|
||||
pending_id: Callable[[InboundMessage], str | None],
|
||||
should_defer_turn: Callable[[InboundMessage, str, Iterable[str]], bool],
|
||||
missing_id_error: str,
|
||||
duplicate_id_error: Callable[[str], str],
|
||||
deferred_queues: dict[str, list[InboundMessage]] | None = None,
|
||||
) -> None:
|
||||
self._publish_inbound = publish_inbound
|
||||
self._dispatch = dispatch
|
||||
self._is_running = is_running
|
||||
self._turn_id = turn_id
|
||||
self._pending_id = pending_id
|
||||
self._should_defer_turn = should_defer_turn
|
||||
self._missing_id_error = missing_id_error
|
||||
self._duplicate_id_error = duplicate_id_error
|
||||
self.deferred_queues = deferred_queues if deferred_queues is not None else {}
|
||||
self._waiters: dict[str, asyncio.Future[OutboundMessage | None]] = {}
|
||||
self._pending_messages_by_turn_id: dict[str, InboundMessage] = {}
|
||||
|
||||
async def submit(self, msg: InboundMessage) -> OutboundMessage | None:
|
||||
"""Submit an automation turn and wait for its session response."""
|
||||
turn_id = self._turn_id(msg)
|
||||
if not turn_id:
|
||||
raise ValueError(self._missing_id_error)
|
||||
if turn_id in self._waiters:
|
||||
raise RuntimeError(self._duplicate_id_error(turn_id))
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
future: asyncio.Future[OutboundMessage | None] = loop.create_future()
|
||||
self._waiters[turn_id] = future
|
||||
self._pending_messages_by_turn_id[turn_id] = msg
|
||||
try:
|
||||
if self._is_running():
|
||||
await self._publish_inbound(msg)
|
||||
else:
|
||||
await self._dispatch(msg)
|
||||
try:
|
||||
return await future
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise AutomationTurnError(str(exc) or exc.__class__.__name__) from exc
|
||||
finally:
|
||||
self._waiters.pop(turn_id, None)
|
||||
self._pending_messages_by_turn_id.pop(turn_id, None)
|
||||
|
||||
def defer_if_active(
|
||||
self,
|
||||
msg: InboundMessage,
|
||||
*,
|
||||
session_key: str,
|
||||
active_session_keys: Iterable[str],
|
||||
) -> bool:
|
||||
"""Defer an automation turn when its target session is already active."""
|
||||
if not self._should_defer_turn(msg, session_key, active_session_keys):
|
||||
return False
|
||||
pending_msg = msg
|
||||
if session_key != msg.session_key:
|
||||
pending_msg = dataclasses.replace(
|
||||
msg,
|
||||
session_key_override=session_key,
|
||||
)
|
||||
self.deferred_queues.setdefault(session_key, []).append(pending_msg)
|
||||
return True
|
||||
|
||||
def complete(
|
||||
self,
|
||||
msg: InboundMessage,
|
||||
*,
|
||||
response: OutboundMessage | None = None,
|
||||
error: BaseException | None = None,
|
||||
) -> None:
|
||||
turn_id = self._turn_id(msg)
|
||||
if not turn_id:
|
||||
return
|
||||
future = self._waiters.get(turn_id)
|
||||
if future is None or future.done():
|
||||
return
|
||||
if error is not None:
|
||||
future.set_exception(error)
|
||||
else:
|
||||
future.set_result(response)
|
||||
|
||||
def pending_ids_for_session(self, session_key: str) -> set[str]:
|
||||
"""Return automation IDs that are waiting for or running in *session_key*."""
|
||||
pending_ids: set[str] = set()
|
||||
for msg in self.deferred_queues.get(session_key, []):
|
||||
pending_id = self._pending_id(msg)
|
||||
if pending_id:
|
||||
pending_ids.add(pending_id)
|
||||
for msg in self._pending_messages_by_turn_id.values():
|
||||
if msg.session_key != session_key:
|
||||
continue
|
||||
pending_id = self._pending_id(msg)
|
||||
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,
|
||||
)
|
||||
+22
-107
@@ -2,11 +2,10 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import dataclasses
|
||||
from collections.abc import Awaitable, Callable, Iterable
|
||||
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
from nanobot.agent.automation_turns import AutomationTurnCoordinator
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.cron.session_turns import (
|
||||
cron_run_id,
|
||||
cron_trigger,
|
||||
@@ -14,7 +13,7 @@ from nanobot.cron.session_turns import (
|
||||
)
|
||||
|
||||
|
||||
class CronTurnCoordinator:
|
||||
class CronTurnCoordinator(AutomationTurnCoordinator):
|
||||
"""Manage scheduled cron turns without mixing them into live injections."""
|
||||
|
||||
def __init__(
|
||||
@@ -23,115 +22,31 @@ class CronTurnCoordinator:
|
||||
publish_inbound: Callable[[InboundMessage], Awaitable[None]],
|
||||
dispatch: Callable[[InboundMessage], Awaitable[object]],
|
||||
is_running: Callable[[], bool],
|
||||
deferred_queues: dict[str, list[InboundMessage]] | None = None,
|
||||
) -> None:
|
||||
self._publish_inbound = publish_inbound
|
||||
self._dispatch = dispatch
|
||||
self._is_running = is_running
|
||||
self.deferred_queues: dict[str, list[InboundMessage]] = {}
|
||||
self._waiters: dict[str, asyncio.Future[OutboundMessage | None]] = {}
|
||||
self._pending_messages_by_run_id: dict[str, InboundMessage] = {}
|
||||
|
||||
async def submit(self, msg: InboundMessage) -> OutboundMessage | None:
|
||||
"""Submit a scheduled cron turn and wait for its session response."""
|
||||
run_id = cron_run_id(msg.metadata)
|
||||
if not run_id:
|
||||
raise ValueError("cron turn metadata must include a run_id")
|
||||
if run_id in self._waiters:
|
||||
raise RuntimeError(f"cron run {run_id!r} is already pending")
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
future: asyncio.Future[OutboundMessage | None] = loop.create_future()
|
||||
self._waiters[run_id] = future
|
||||
self._pending_messages_by_run_id[run_id] = msg
|
||||
try:
|
||||
if self._is_running():
|
||||
await self._publish_inbound(msg)
|
||||
else:
|
||||
await self._dispatch(msg)
|
||||
return await future
|
||||
finally:
|
||||
self._waiters.pop(run_id, None)
|
||||
self._pending_messages_by_run_id.pop(run_id, None)
|
||||
|
||||
def should_defer(
|
||||
self,
|
||||
msg: InboundMessage,
|
||||
*,
|
||||
session_key: str,
|
||||
active_session_keys: Iterable[str],
|
||||
) -> bool:
|
||||
return (
|
||||
defer_cron_until_session_idle(msg.metadata)
|
||||
and session_key in active_session_keys
|
||||
super().__init__(
|
||||
publish_inbound=publish_inbound,
|
||||
dispatch=dispatch,
|
||||
is_running=is_running,
|
||||
turn_id=lambda msg: cron_run_id(msg.metadata),
|
||||
pending_id=_cron_job_id,
|
||||
should_defer_turn=_should_defer_cron_turn,
|
||||
missing_id_error="cron turn metadata must include a run_id",
|
||||
duplicate_id_error=lambda run_id: f"cron run {run_id!r} is already pending",
|
||||
deferred_queues=deferred_queues,
|
||||
)
|
||||
|
||||
def defer_if_active(
|
||||
self,
|
||||
msg: InboundMessage,
|
||||
*,
|
||||
session_key: str,
|
||||
active_session_keys: Iterable[str],
|
||||
) -> bool:
|
||||
"""Defer a cron turn when its target session is already active."""
|
||||
if not self.should_defer(
|
||||
msg,
|
||||
session_key=session_key,
|
||||
active_session_keys=active_session_keys,
|
||||
):
|
||||
return False
|
||||
pending_msg = msg
|
||||
if session_key != msg.session_key:
|
||||
pending_msg = dataclasses.replace(
|
||||
msg,
|
||||
session_key_override=session_key,
|
||||
)
|
||||
self.defer(session_key, pending_msg)
|
||||
return True
|
||||
|
||||
def complete(
|
||||
self,
|
||||
msg: InboundMessage,
|
||||
*,
|
||||
response: OutboundMessage | None = None,
|
||||
error: BaseException | None = None,
|
||||
) -> None:
|
||||
run_id = cron_run_id(msg.metadata)
|
||||
if not run_id:
|
||||
return
|
||||
future = self._waiters.get(run_id)
|
||||
if future is None or future.done():
|
||||
return
|
||||
if error is not None:
|
||||
future.set_exception(error)
|
||||
else:
|
||||
future.set_result(response)
|
||||
|
||||
def defer(self, session_key: str, msg: InboundMessage) -> None:
|
||||
self.deferred_queues.setdefault(session_key, []).append(msg)
|
||||
|
||||
def pending_job_ids_for_session(self, session_key: str) -> set[str]:
|
||||
"""Return cron jobs that are waiting for or running in *session_key*."""
|
||||
job_ids: set[str] = set()
|
||||
for msg in self.deferred_queues.get(session_key, []):
|
||||
job_id = _cron_job_id(msg)
|
||||
if job_id:
|
||||
job_ids.add(job_id)
|
||||
for msg in self._pending_messages_by_run_id.values():
|
||||
if msg.session_key != session_key:
|
||||
continue
|
||||
job_id = _cron_job_id(msg)
|
||||
if job_id:
|
||||
job_ids.add(job_id)
|
||||
return job_ids
|
||||
return self.pending_ids_for_session(session_key)
|
||||
|
||||
async def publish_next_deferred(self, session_key: str) -> None:
|
||||
queue = self.deferred_queues.get(session_key)
|
||||
if not queue:
|
||||
return
|
||||
msg = queue.pop(0)
|
||||
if not queue:
|
||||
self.deferred_queues.pop(session_key, None)
|
||||
await self._publish_inbound(msg)
|
||||
|
||||
def _should_defer_cron_turn(
|
||||
msg: InboundMessage,
|
||||
session_key: str,
|
||||
active_session_keys: Iterable[str],
|
||||
) -> bool:
|
||||
return defer_cron_until_session_idle(msg.metadata) and session_key in active_session_keys
|
||||
|
||||
|
||||
def _cron_job_id(msg: InboundMessage) -> str | None:
|
||||
|
||||
+131
-268
@@ -18,6 +18,7 @@ from loguru import logger
|
||||
from nanobot.agent import context as agent_context
|
||||
from nanobot.agent import model_presets as preset_helpers
|
||||
from nanobot.agent.autocompact import AutoCompact
|
||||
from nanobot.agent.automation_turns import publish_next_deferred_turn
|
||||
from nanobot.agent.context import ContextBuilder
|
||||
from nanobot.agent.cron_turns import CronTurnCoordinator
|
||||
from nanobot.agent.hook import AgentHook, CompositeHook
|
||||
@@ -31,6 +32,13 @@ from nanobot.agent.tools.message import MessageTool
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
from nanobot.agent.tools.self import MyTool
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
from nanobot.bus.outbound_events import (
|
||||
RetryWaitEvent,
|
||||
StreamDeltaEvent,
|
||||
StreamedResponseEvent,
|
||||
StreamEndEvent,
|
||||
outbound_message_for_event,
|
||||
)
|
||||
from nanobot.bus.progress import build_bus_progress_callback
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.bus.runtime_events import (
|
||||
@@ -40,9 +48,6 @@ from nanobot.bus.runtime_events import (
|
||||
)
|
||||
from nanobot.command import CommandContext, CommandRouter, register_builtin_commands
|
||||
from nanobot.config.schema import AgentDefaults, ModelPresetConfig
|
||||
from nanobot.cron.session_turns import (
|
||||
cron_history_overrides,
|
||||
)
|
||||
from nanobot.providers.base import LLMProvider
|
||||
from nanobot.providers.factory import ProviderSnapshot
|
||||
from nanobot.security.workspace_access import (
|
||||
@@ -50,21 +55,22 @@ from nanobot.security.workspace_access import (
|
||||
bind_workspace_scope,
|
||||
reset_workspace_scope,
|
||||
)
|
||||
from nanobot.session import turn_continuation
|
||||
from nanobot.session import turn_continuation, turn_history
|
||||
from nanobot.session.automation_turns import automation_history_overrides
|
||||
from nanobot.session.goal_state import (
|
||||
goal_state_runtime_lines,
|
||||
runner_wall_llm_timeout_s,
|
||||
sustained_goal_active,
|
||||
)
|
||||
from nanobot.session.history_visibility import HIDDEN_HISTORY_META
|
||||
from nanobot.session.keys import UNIFIED_SESSION_KEY, session_key_for_channel
|
||||
from nanobot.session.manager import (
|
||||
Session,
|
||||
SessionManager,
|
||||
replay_max_messages_for_context,
|
||||
)
|
||||
from nanobot.triggers.local_turns import LocalTriggerTurnCoordinator
|
||||
from nanobot.utils.document import extract_documents, reference_non_image_attachments
|
||||
from nanobot.utils.helpers import image_placeholder_text
|
||||
from nanobot.utils.helpers import truncate_text as truncate_text_fn
|
||||
from nanobot.utils.image_generation_intent import image_generation_prompt
|
||||
from nanobot.utils.llm_runtime import LLMRuntime
|
||||
from nanobot.utils.runtime import (
|
||||
@@ -167,8 +173,8 @@ class AgentLoop:
|
||||
self._refresh_provider_snapshot()
|
||||
return LLMRuntime(self.provider, self.model)
|
||||
|
||||
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
|
||||
_PENDING_USER_TURN_KEY = "pending_user_turn"
|
||||
_RUNTIME_CHECKPOINT_KEY = turn_history.RUNTIME_CHECKPOINT_KEY
|
||||
_PENDING_USER_TURN_KEY = turn_history.PENDING_USER_TURN_KEY
|
||||
|
||||
# Event-driven state transition table.
|
||||
# Handlers return an event string; the driver looks up the next state here.
|
||||
@@ -219,6 +225,7 @@ class AgentLoop:
|
||||
runtime_events: RuntimeEventBus | None = None,
|
||||
runtime_model_publisher: Callable[[str, str | None], None] | None = None,
|
||||
restart_mode: str = "auto",
|
||||
local_trigger_store: Any | None = None,
|
||||
):
|
||||
from nanobot.config.schema import ToolsConfig
|
||||
|
||||
@@ -266,6 +273,7 @@ class AgentLoop:
|
||||
):
|
||||
self._image_generation_provider_configs["openrouter"] = image_generation_provider_config
|
||||
self.cron_service = cron_service
|
||||
self.local_trigger_store = local_trigger_store
|
||||
self.restrict_to_workspace = restrict_to_workspace
|
||||
self.workspace_scopes = WorkspaceScopeResolver(
|
||||
default_workspace=workspace,
|
||||
@@ -310,10 +318,22 @@ class AgentLoop:
|
||||
# When a session has an active task, new messages for that session
|
||||
# are routed here instead of creating a new task.
|
||||
self._pending_queues: dict[str, asyncio.Queue] = {}
|
||||
self._deferred_automation_turns: dict[str, list[InboundMessage]] = {}
|
||||
self._cron_turns = CronTurnCoordinator(
|
||||
publish_inbound=self.bus.publish_inbound,
|
||||
dispatch=self._dispatch,
|
||||
is_running=lambda: self._running,
|
||||
deferred_queues=self._deferred_automation_turns,
|
||||
)
|
||||
self._local_trigger_turns = LocalTriggerTurnCoordinator(
|
||||
publish_inbound=self.bus.publish_inbound,
|
||||
dispatch=self._dispatch,
|
||||
is_running=lambda: self._running,
|
||||
deferred_queues=self._deferred_automation_turns,
|
||||
)
|
||||
self._automation_turn_coordinators = (
|
||||
("cron", self._cron_turns),
|
||||
("local trigger", self._local_trigger_turns),
|
||||
)
|
||||
# NANOBOT_MAX_CONCURRENT_REQUESTS: <=0 means unlimited; default 3.
|
||||
_max = int(os.environ.get("NANOBOT_MAX_CONCURRENT_REQUESTS", "3"))
|
||||
@@ -567,14 +587,12 @@ class AgentLoop:
|
||||
"""Build a retry-wait callback that publishes to the message bus."""
|
||||
|
||||
async def _on_retry_wait(content: str) -> None:
|
||||
meta = dict(msg.metadata or {})
|
||||
meta["_retry_wait"] = True
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
outbound_message_for_event(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content=content,
|
||||
metadata=meta,
|
||||
event=RetryWaitEvent(content=content),
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -586,9 +604,22 @@ class AgentLoop:
|
||||
async def submit_cron_turn(self, msg: InboundMessage) -> OutboundMessage | None:
|
||||
return await self._cron_turns.submit(msg)
|
||||
|
||||
async def submit_local_trigger_turn(self, msg: InboundMessage) -> OutboundMessage | None:
|
||||
return await self._local_trigger_turns.submit(msg)
|
||||
|
||||
def pending_cron_job_ids_for_session(self, session_key: str) -> set[str]:
|
||||
return self._cron_turns.pending_job_ids_for_session(session_key)
|
||||
|
||||
def pending_local_trigger_ids_for_session(self, session_key: str) -> set[str]:
|
||||
return self._local_trigger_turns.pending_trigger_ids_for_session(session_key)
|
||||
|
||||
async def _publish_next_deferred_automation_turn(self, session_key: str) -> None:
|
||||
await publish_next_deferred_turn(
|
||||
deferred_queues=self._deferred_automation_turns,
|
||||
publish_inbound=self.bus.publish_inbound,
|
||||
session_key=session_key,
|
||||
)
|
||||
|
||||
def _persist_user_message_early(
|
||||
self,
|
||||
msg: InboundMessage,
|
||||
@@ -607,10 +638,10 @@ class AgentLoop:
|
||||
extra: dict[str, Any] = ({"media": list(media_paths)} if media_paths else {}) | agent_context.session_extra(msg.metadata)
|
||||
extra.update(kwargs)
|
||||
text = msg.content if isinstance(msg.content, str) else ""
|
||||
text_override, cron_extra = cron_history_overrides(msg.metadata)
|
||||
text_override, automation_extra = automation_history_overrides(msg.metadata)
|
||||
if text_override is not None:
|
||||
text = text_override
|
||||
extra.update(cron_extra)
|
||||
extra.update(automation_extra)
|
||||
session.add_message("user", text, **extra)
|
||||
self._mark_pending_user_turn(session)
|
||||
self.sessions.save(session)
|
||||
@@ -763,7 +794,20 @@ class AgentLoop:
|
||||
content, media = self._prepare_message_media(content, media)
|
||||
media = media or None
|
||||
user_content = self.context._build_user_content(content, media)
|
||||
return {"role": "user", "content": user_content}
|
||||
row: dict[str, Any] = {"role": "user", "content": user_content}
|
||||
metadata = pending_msg.metadata if isinstance(pending_msg.metadata, dict) else {}
|
||||
if (
|
||||
pending_msg.sender_id == "subagent"
|
||||
and metadata.get("injected_event") == "subagent_result"
|
||||
):
|
||||
marker: dict[str, Any] = {"kind": "subagent_result"}
|
||||
task_id = metadata.get("subagent_task_id")
|
||||
if isinstance(task_id, str) and task_id:
|
||||
marker["subagent_task_id"] = task_id
|
||||
row["subagent_task_id"] = task_id
|
||||
row[HIDDEN_HISTORY_META] = marker
|
||||
row["injected_event"] = "subagent_result"
|
||||
return row
|
||||
|
||||
items: list[dict[str, Any]] = []
|
||||
while len(items) < limit:
|
||||
@@ -918,15 +962,21 @@ class AgentLoop:
|
||||
self.commands.dispatch_priority,
|
||||
)
|
||||
continue
|
||||
if self._cron_turns.defer_if_active(
|
||||
msg,
|
||||
session_key=effective_key,
|
||||
active_session_keys=self._pending_queues.keys(),
|
||||
):
|
||||
logger.info(
|
||||
"Deferred cron turn for active session {}",
|
||||
effective_key,
|
||||
)
|
||||
deferred = False
|
||||
for label, coordinator in self._automation_turn_coordinators:
|
||||
if coordinator.defer_if_active(
|
||||
msg,
|
||||
session_key=effective_key,
|
||||
active_session_keys=self._pending_queues.keys(),
|
||||
):
|
||||
logger.info(
|
||||
"Deferred {} turn for active session {}",
|
||||
label,
|
||||
effective_key,
|
||||
)
|
||||
deferred = True
|
||||
break
|
||||
if deferred:
|
||||
continue
|
||||
# If this session already has an active pending queue (i.e. a task
|
||||
# is processing this session), route the message there for mid-turn
|
||||
@@ -999,26 +1049,31 @@ class AgentLoop:
|
||||
return f"{stream_base_id}:{stream_segment}"
|
||||
|
||||
async def on_stream(delta: str) -> None:
|
||||
meta = dict(msg.metadata or {})
|
||||
meta["_stream_delta"] = True
|
||||
meta["_stream_id"] = _current_stream_id()
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel, chat_id=msg.chat_id,
|
||||
content=delta,
|
||||
metadata=meta,
|
||||
))
|
||||
await self.bus.publish_outbound(
|
||||
outbound_message_for_event(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
event=StreamDeltaEvent(
|
||||
content=delta,
|
||||
stream_id=_current_stream_id(),
|
||||
),
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
|
||||
async def on_stream_end(*, resuming: bool = False) -> None:
|
||||
nonlocal stream_segment
|
||||
meta = dict(msg.metadata or {})
|
||||
meta["_stream_end"] = True
|
||||
meta["_resuming"] = resuming
|
||||
meta["_stream_id"] = _current_stream_id()
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel, chat_id=msg.chat_id,
|
||||
content="",
|
||||
metadata=meta,
|
||||
))
|
||||
await self.bus.publish_outbound(
|
||||
outbound_message_for_event(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
event=StreamEndEvent(
|
||||
stream_id=_current_stream_id(),
|
||||
resuming=resuming,
|
||||
),
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
stream_segment += 1
|
||||
|
||||
response = await self._process_message(
|
||||
@@ -1044,12 +1099,11 @@ class AgentLoop:
|
||||
session_key=session_key,
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
self._cron_turns.complete(msg, response=response)
|
||||
for _, coordinator in self._automation_turn_coordinators:
|
||||
coordinator.complete(msg, response=response)
|
||||
except asyncio.CancelledError:
|
||||
self._cron_turns.complete(
|
||||
msg,
|
||||
error=asyncio.CancelledError(),
|
||||
)
|
||||
for _, coordinator in self._automation_turn_coordinators:
|
||||
coordinator.complete(msg, error=asyncio.CancelledError())
|
||||
logger.info("Task cancelled for session {}", session_key)
|
||||
# Preserve partial context from the interrupted turn so
|
||||
# the user does not lose tool results and assistant
|
||||
@@ -1088,7 +1142,8 @@ class AgentLoop:
|
||||
session_key=session_key,
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
self._cron_turns.complete(msg, error=exc)
|
||||
for _, coordinator in self._automation_turn_coordinators:
|
||||
coordinator.complete(msg, error=exc)
|
||||
finally:
|
||||
# Drain any messages still in the pending queue and re-publish
|
||||
# them to the bus so they are processed as fresh inbound messages
|
||||
@@ -1119,14 +1174,14 @@ class AgentLoop:
|
||||
msg, session_key, "idle"
|
||||
)
|
||||
self._runtime_events().clear_turn(session_key)
|
||||
await self._cron_turns.publish_next_deferred(session_key)
|
||||
await self._publish_next_deferred_automation_turn(session_key)
|
||||
finally:
|
||||
if pending is None:
|
||||
await self._runtime_events().run_status_changed(
|
||||
msg, session_key, "idle"
|
||||
)
|
||||
self._runtime_events().clear_turn(session_key)
|
||||
await self._cron_turns.publish_next_deferred(session_key)
|
||||
await self._publish_next_deferred_automation_turn(session_key)
|
||||
|
||||
async def close_mcp(self) -> None:
|
||||
"""Drain pending background archives, then close MCP connections."""
|
||||
@@ -1371,9 +1426,10 @@ class AgentLoop:
|
||||
preview = final_content[:120] + "..." if len(final_content) > 120 else final_content
|
||||
logger.info("Response to {}:{}: {}", msg.channel, msg.sender_id, preview)
|
||||
|
||||
event = None
|
||||
meta = dict(msg.metadata or {})
|
||||
if on_stream is not None and stop_reason not in {"error", "tool_error"}:
|
||||
meta["_streamed"] = True
|
||||
event = StreamedResponseEvent()
|
||||
if turn_latency_ms is not None:
|
||||
meta["latency_ms"] = int(turn_latency_ms)
|
||||
|
||||
@@ -1381,6 +1437,7 @@ class AgentLoop:
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content=final_content,
|
||||
event=event,
|
||||
metadata=meta,
|
||||
)
|
||||
|
||||
@@ -1438,7 +1495,7 @@ class AgentLoop:
|
||||
# message. Mark messages with _command so get_history can filter
|
||||
# them out of LLM context. /new is excluded because it
|
||||
# intentionally clears the session.
|
||||
if raw.lower() != "/new":
|
||||
if cmd_ctx.raw.lower() != "/new":
|
||||
ctx.user_persisted_early = self._persist_user_message_early(
|
||||
ctx.msg, ctx.session, _command=True
|
||||
)
|
||||
@@ -1595,38 +1652,13 @@ class AgentLoop:
|
||||
should_truncate_text: bool = False,
|
||||
drop_runtime: bool = False,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Strip volatile multimodal payloads before writing session history."""
|
||||
filtered: list[dict[str, Any]] = []
|
||||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
filtered.append(block)
|
||||
continue
|
||||
|
||||
if (
|
||||
drop_runtime
|
||||
and block.get("type") == "text"
|
||||
and isinstance(block.get("text"), str)
|
||||
and block["text"].startswith(ContextBuilder._RUNTIME_CONTEXT_TAG)
|
||||
):
|
||||
continue
|
||||
|
||||
if block.get("type") == "image_url" and block.get("image_url", {}).get(
|
||||
"url", ""
|
||||
).startswith("data:image/"):
|
||||
path = (block.get("_meta") or {}).get("path", "")
|
||||
filtered.append({"type": "text", "text": image_placeholder_text(path)})
|
||||
continue
|
||||
|
||||
if block.get("type") == "text" and isinstance(block.get("text"), str):
|
||||
text = block["text"]
|
||||
if should_truncate_text and len(text) > self.max_tool_result_chars:
|
||||
text = truncate_text_fn(text, self.max_tool_result_chars)
|
||||
filtered.append({**block, "text": text})
|
||||
continue
|
||||
|
||||
filtered.append(block)
|
||||
|
||||
return filtered
|
||||
return turn_history.sanitize_persisted_blocks(
|
||||
content,
|
||||
max_tool_result_chars=self.max_tool_result_chars,
|
||||
runtime_context_tag=ContextBuilder._RUNTIME_CONTEXT_TAG,
|
||||
should_truncate_text=should_truncate_text,
|
||||
drop_runtime=drop_runtime,
|
||||
)
|
||||
|
||||
def _save_turn(
|
||||
self,
|
||||
@@ -1636,193 +1668,36 @@ class AgentLoop:
|
||||
*,
|
||||
turn_latency_ms: int | None = None,
|
||||
) -> None:
|
||||
"""Save new-turn messages into session, truncating large tool results."""
|
||||
from datetime import datetime
|
||||
|
||||
declared_tool_call_ids = {
|
||||
str(tc["id"])
|
||||
for m in session.messages
|
||||
if m.get("role") == "assistant"
|
||||
for tc in m.get("tool_calls") or []
|
||||
if isinstance(tc, dict) and tc.get("id")
|
||||
}
|
||||
last_assistant_idx: int | None = None
|
||||
for m in messages[skip:]:
|
||||
entry = dict(m)
|
||||
role, content = entry.get("role"), entry.get("content")
|
||||
if role == "assistant" and not content and not entry.get("tool_calls"):
|
||||
continue # skip empty assistant messages — they poison session context
|
||||
if role == "tool":
|
||||
tool_call_id = entry.get("tool_call_id")
|
||||
if not tool_call_id or str(tool_call_id) not in declared_tool_call_ids:
|
||||
# Undeclared tool results corrupt future provider requests.
|
||||
logger.warning(
|
||||
"Dropping orphaned tool result {} from session {} during persistence",
|
||||
tool_call_id or "(missing id)",
|
||||
session.key,
|
||||
)
|
||||
continue
|
||||
if isinstance(content, str) and len(content) > self.max_tool_result_chars:
|
||||
entry["content"] = truncate_text_fn(content, self.max_tool_result_chars)
|
||||
elif isinstance(content, list):
|
||||
filtered = self._sanitize_persisted_blocks(content, should_truncate_text=True)
|
||||
if not filtered:
|
||||
# Preserve the tool_call/result pair after block filtering.
|
||||
filtered = [
|
||||
{"type": "text", "text": "[tool result omitted during persistence]"}
|
||||
]
|
||||
entry["content"] = filtered
|
||||
elif role == "user":
|
||||
if isinstance(content, str) and ContextBuilder._RUNTIME_CONTEXT_TAG in content:
|
||||
# Strip the runtime-context block appended at the end.
|
||||
tag_pos = content.find(ContextBuilder._RUNTIME_CONTEXT_TAG)
|
||||
before = content[:tag_pos].rstrip("\n ")
|
||||
if before:
|
||||
entry["content"] = before
|
||||
else:
|
||||
continue
|
||||
if isinstance(content, list):
|
||||
filtered = self._sanitize_persisted_blocks(content, drop_runtime=True)
|
||||
if not filtered:
|
||||
continue
|
||||
entry["content"] = filtered
|
||||
entry.setdefault("timestamp", datetime.now().isoformat())
|
||||
session.messages.append(entry)
|
||||
if role == "assistant":
|
||||
last_assistant_idx = len(session.messages) - 1
|
||||
declared_tool_call_ids.update(
|
||||
str(tc["id"])
|
||||
for tc in entry.get("tool_calls") or []
|
||||
if isinstance(tc, dict) and tc.get("id")
|
||||
)
|
||||
if turn_latency_ms is not None and last_assistant_idx is not None:
|
||||
session.messages[last_assistant_idx]["latency_ms"] = int(turn_latency_ms)
|
||||
session.updated_at = datetime.now()
|
||||
turn_history.save_turn(
|
||||
session,
|
||||
messages,
|
||||
skip,
|
||||
max_tool_result_chars=self.max_tool_result_chars,
|
||||
runtime_context_tag=ContextBuilder._RUNTIME_CONTEXT_TAG,
|
||||
turn_latency_ms=turn_latency_ms,
|
||||
)
|
||||
|
||||
def _persist_subagent_followup(self, session: Session, msg: InboundMessage) -> bool:
|
||||
"""Persist subagent follow-ups before prompt assembly so history stays durable.
|
||||
|
||||
Returns True if a new entry was appended; False if the follow-up was
|
||||
deduped (same ``subagent_task_id`` already in session) or carries no
|
||||
content worth persisting.
|
||||
"""
|
||||
if not msg.content:
|
||||
return False
|
||||
task_id = msg.metadata.get("subagent_task_id") if isinstance(msg.metadata, dict) else None
|
||||
if task_id and any(
|
||||
m.get("injected_event") == "subagent_result" and m.get("subagent_task_id") == task_id
|
||||
for m in session.messages
|
||||
):
|
||||
return False
|
||||
session.add_message(
|
||||
"assistant",
|
||||
msg.content,
|
||||
sender_id=msg.sender_id,
|
||||
injected_event="subagent_result",
|
||||
subagent_task_id=task_id,
|
||||
)
|
||||
return True
|
||||
return turn_history.persist_subagent_followup(session, msg)
|
||||
|
||||
def _set_runtime_checkpoint(self, session: Session, payload: dict[str, Any]) -> None:
|
||||
"""Persist the latest in-flight turn state into session metadata."""
|
||||
session.metadata[self._RUNTIME_CHECKPOINT_KEY] = payload
|
||||
turn_history.set_runtime_checkpoint(session, payload)
|
||||
self.sessions.save(session)
|
||||
|
||||
def _mark_pending_user_turn(self, session: Session) -> None:
|
||||
session.metadata[self._PENDING_USER_TURN_KEY] = True
|
||||
turn_history.mark_pending_user_turn(session)
|
||||
|
||||
def _clear_pending_user_turn(self, session: Session) -> None:
|
||||
session.metadata.pop(self._PENDING_USER_TURN_KEY, None)
|
||||
turn_history.clear_pending_user_turn(session)
|
||||
|
||||
def _clear_runtime_checkpoint(self, session: Session) -> None:
|
||||
if self._RUNTIME_CHECKPOINT_KEY in session.metadata:
|
||||
session.metadata.pop(self._RUNTIME_CHECKPOINT_KEY, None)
|
||||
|
||||
@staticmethod
|
||||
def _checkpoint_message_key(message: dict[str, Any]) -> tuple[Any, ...]:
|
||||
return (
|
||||
message.get("role"),
|
||||
message.get("content"),
|
||||
message.get("tool_call_id"),
|
||||
message.get("name"),
|
||||
message.get("tool_calls"),
|
||||
message.get("reasoning_content"),
|
||||
message.get("thinking_blocks"),
|
||||
)
|
||||
turn_history.clear_runtime_checkpoint(session)
|
||||
|
||||
def _restore_runtime_checkpoint(self, session: Session) -> bool:
|
||||
"""Materialize an unfinished turn into session history before a new request."""
|
||||
from datetime import datetime
|
||||
|
||||
checkpoint = session.metadata.get(self._RUNTIME_CHECKPOINT_KEY)
|
||||
if not isinstance(checkpoint, dict):
|
||||
return False
|
||||
|
||||
assistant_message = checkpoint.get("assistant_message")
|
||||
completed_tool_results = checkpoint.get("completed_tool_results") or []
|
||||
pending_tool_calls = checkpoint.get("pending_tool_calls") or []
|
||||
|
||||
restored_messages: list[dict[str, Any]] = []
|
||||
if isinstance(assistant_message, dict):
|
||||
restored = dict(assistant_message)
|
||||
restored.setdefault("timestamp", datetime.now().isoformat())
|
||||
restored_messages.append(restored)
|
||||
for message in completed_tool_results:
|
||||
if isinstance(message, dict):
|
||||
restored = dict(message)
|
||||
restored.setdefault("timestamp", datetime.now().isoformat())
|
||||
restored_messages.append(restored)
|
||||
for tool_call in pending_tool_calls:
|
||||
if not isinstance(tool_call, dict):
|
||||
continue
|
||||
tool_id = tool_call.get("id")
|
||||
name = ((tool_call.get("function") or {}).get("name")) or "tool"
|
||||
restored_messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": tool_id,
|
||||
"name": name,
|
||||
"content": "Error: Task interrupted before this tool finished.",
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
}
|
||||
)
|
||||
|
||||
overlap = 0
|
||||
max_overlap = min(len(session.messages), len(restored_messages))
|
||||
for size in range(max_overlap, 0, -1):
|
||||
existing = session.messages[-size:]
|
||||
restored = restored_messages[:size]
|
||||
if all(
|
||||
self._checkpoint_message_key(left) == self._checkpoint_message_key(right)
|
||||
for left, right in zip(existing, restored)
|
||||
):
|
||||
overlap = size
|
||||
break
|
||||
session.messages.extend(restored_messages[overlap:])
|
||||
|
||||
self._clear_pending_user_turn(session)
|
||||
self._clear_runtime_checkpoint(session)
|
||||
return True
|
||||
return turn_history.restore_runtime_checkpoint(session)
|
||||
|
||||
def _restore_pending_user_turn(self, session: Session) -> bool:
|
||||
"""Close a turn that only persisted the user message before crashing."""
|
||||
from datetime import datetime
|
||||
|
||||
if not session.metadata.get(self._PENDING_USER_TURN_KEY):
|
||||
return False
|
||||
|
||||
if session.messages and session.messages[-1].get("role") == "user":
|
||||
session.messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Error: Task interrupted before a response was generated.",
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
}
|
||||
)
|
||||
session.updated_at = datetime.now()
|
||||
|
||||
self._clear_pending_user_turn(session)
|
||||
return True
|
||||
return turn_history.restore_pending_user_turn(session)
|
||||
|
||||
async def process_direct(
|
||||
self,
|
||||
@@ -1852,17 +1727,13 @@ class AgentLoop:
|
||||
)
|
||||
# Share the dispatch lock so direct calls serialize with bus turns.
|
||||
lock = self._session_locks.setdefault(session_key, asyncio.Lock())
|
||||
pending: asyncio.Queue[InboundMessage] = asyncio.Queue(maxsize=20)
|
||||
try:
|
||||
async with lock:
|
||||
self._pending_queues[session_key] = pending
|
||||
self.subagents.set_direct_result_queue(session_key, pending)
|
||||
kwargs: dict[str, Any] = {
|
||||
"session_key": session_key,
|
||||
"on_progress": on_progress,
|
||||
"on_stream": on_stream,
|
||||
"on_stream_end": on_stream_end,
|
||||
"pending_queue": pending,
|
||||
"ephemeral": ephemeral,
|
||||
}
|
||||
if _run_extra_hooks_for_ephemeral:
|
||||
@@ -1876,13 +1747,5 @@ class AgentLoop:
|
||||
**kwargs,
|
||||
)
|
||||
finally:
|
||||
self.subagents.clear_direct_result_queue(session_key, pending)
|
||||
if self._pending_queues.get(session_key) is pending:
|
||||
self._pending_queues.pop(session_key, None)
|
||||
while True:
|
||||
try:
|
||||
await self.bus.publish_inbound(pending.get_nowait())
|
||||
except asyncio.QueueEmpty:
|
||||
break
|
||||
await self._runtime_events().run_status_changed(msg, session_key, "idle")
|
||||
self._runtime_events().clear_turn(session_key)
|
||||
|
||||
@@ -18,8 +18,9 @@ from nanobot.agent.context_governance import (
|
||||
ContextGovernor,
|
||||
)
|
||||
from nanobot.agent.hook import AgentHook, AgentHookContext, AgentRunHookContext
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
from nanobot.agent.tools.registry import ToolRegistry, is_tool_error_result
|
||||
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
|
||||
from nanobot.session.history_visibility import is_hidden_history_message
|
||||
from nanobot.utils.file_edit_events import (
|
||||
StreamingFileEditTracker,
|
||||
build_file_edit_end_event,
|
||||
@@ -155,6 +156,8 @@ class AgentRunner:
|
||||
messages
|
||||
and injection.get("role") == "user"
|
||||
and messages[-1].get("role") == "user"
|
||||
and not is_hidden_history_message(injection)
|
||||
and not is_hidden_history_message(messages[-1])
|
||||
):
|
||||
merged = dict(messages[-1])
|
||||
merged["content"] = cls._merge_message_content(
|
||||
@@ -1266,7 +1269,7 @@ class AgentRunner:
|
||||
return payload, event, exc
|
||||
return payload, event, None
|
||||
|
||||
if isinstance(result, str) and result.startswith("Error"):
|
||||
if is_tool_error_result(tool_call.name, result):
|
||||
if file_edit_trackers and progress_callback is not None:
|
||||
await invoke_file_edit_progress(
|
||||
progress_callback,
|
||||
|
||||
@@ -118,22 +118,6 @@ class SubagentManager:
|
||||
self._running_tasks: dict[str, asyncio.Task[None]] = {}
|
||||
self._task_statuses: dict[str, SubagentStatus] = {}
|
||||
self._session_tasks: dict[str, set[str]] = {} # session_key -> {task_id, ...}
|
||||
self._direct_result_queues: dict[str, asyncio.Queue[InboundMessage]] = {}
|
||||
|
||||
def set_direct_result_queue(
|
||||
self,
|
||||
session_key: str,
|
||||
queue: asyncio.Queue[InboundMessage],
|
||||
) -> None:
|
||||
self._direct_result_queues[session_key] = queue
|
||||
|
||||
def clear_direct_result_queue(
|
||||
self,
|
||||
session_key: str,
|
||||
queue: asyncio.Queue[InboundMessage],
|
||||
) -> None:
|
||||
if self._direct_result_queues.get(session_key) is queue:
|
||||
self._direct_result_queues.pop(session_key, None)
|
||||
|
||||
def _subagent_tools_config(self) -> ToolsConfig:
|
||||
"""Build a ToolsConfig scoped for subagent use."""
|
||||
@@ -351,10 +335,6 @@ class SubagentManager:
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
if queue := self._direct_result_queues.get(override):
|
||||
await queue.put(msg)
|
||||
logger.debug("Subagent [{}] queued result directly for {}", task_id, override)
|
||||
return
|
||||
await self.bus.publish_inbound(msg)
|
||||
logger.debug("Subagent [{}] announced result to {}:{}", task_id, origin['channel'], origin['chat_id'])
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
"""Agent tools module."""
|
||||
|
||||
from nanobot.agent.tools.base import Schema, Tool, tool_parameters
|
||||
from nanobot.agent.tools.base import Schema, Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import ToolContext
|
||||
from nanobot.agent.tools.loader import ToolLoader
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
@@ -25,6 +25,7 @@ __all__ = [
|
||||
"Tool",
|
||||
"ToolContext",
|
||||
"ToolLoader",
|
||||
"ToolResult",
|
||||
"ToolRegistry",
|
||||
"tool_parameters",
|
||||
"tool_parameters_schema",
|
||||
|
||||
@@ -7,7 +7,7 @@ from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from nanobot.agent.tools.base import tool_parameters
|
||||
from nanobot.agent.tools.base import ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.filesystem import _FsTool
|
||||
from nanobot.agent.tools.schema import (
|
||||
ArraySchema,
|
||||
@@ -289,8 +289,8 @@ class ApplyPatchTool(_FsTool):
|
||||
_format_summary(summary) for summary in summaries
|
||||
)
|
||||
except PermissionError as exc:
|
||||
return f"Error: {exc}"
|
||||
return ToolResult.error(f"Error: {exc}")
|
||||
except _PatchError as exc:
|
||||
return f"Error applying patch: {exc}"
|
||||
return ToolResult.error(f"Error applying patch: {exc}")
|
||||
except Exception as exc:
|
||||
return f"Error applying patch: {exc}"
|
||||
return ToolResult.error(f"Error applying patch: {exc}")
|
||||
|
||||
@@ -128,6 +128,21 @@ class Schema(ABC):
|
||||
return Schema.validate_json_schema_value(value, self.to_json_schema(), path)
|
||||
|
||||
|
||||
class ToolResult(str):
|
||||
"""String-compatible tool output with structured status."""
|
||||
|
||||
is_error: bool
|
||||
|
||||
def __new__(cls, content: str, *, is_error: bool = False) -> ToolResult:
|
||||
obj = str.__new__(cls, content)
|
||||
obj.is_error = is_error
|
||||
return obj
|
||||
|
||||
@classmethod
|
||||
def error(cls, content: str) -> ToolResult:
|
||||
return cls(content, is_error=True)
|
||||
|
||||
|
||||
class Tool(ABC):
|
||||
"""Agent capability: read files, run commands, etc."""
|
||||
|
||||
@@ -193,9 +208,13 @@ class Tool(ABC):
|
||||
|
||||
@abstractmethod
|
||||
async def execute(self, **kwargs: Any) -> Any:
|
||||
"""Run the tool; returns a string or list of content blocks."""
|
||||
"""Run the tool; return content, or ``ToolResult.error(...)`` for failures."""
|
||||
...
|
||||
|
||||
@staticmethod
|
||||
def error(content: str) -> ToolResult:
|
||||
return ToolResult.error(content)
|
||||
|
||||
def _cast_object(self, obj: Any, schema: dict[str, Any]) -> dict[str, Any]:
|
||||
if not isinstance(obj, dict):
|
||||
return obj
|
||||
|
||||
@@ -7,7 +7,7 @@ from typing import Any
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
from nanobot.agent.tools.base import Tool, tool_parameters
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.schema import (
|
||||
ArraySchema,
|
||||
BooleanSchema,
|
||||
@@ -136,4 +136,4 @@ class CliAppsTool(Tool):
|
||||
restrict_to_workspace=access.restrict_to_workspace,
|
||||
)
|
||||
except CliAppError as exc:
|
||||
return f"Error: {exc.message}"
|
||||
return ToolResult.error(f"Error: {exc.message}")
|
||||
|
||||
+10
-10
@@ -6,7 +6,7 @@ from contextvars import ContextVar
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from nanobot.agent.tools.base import Tool, tool_parameters
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import ContextAware, RequestContext
|
||||
from nanobot.agent.tools.schema import (
|
||||
IntegerSchema,
|
||||
@@ -99,7 +99,7 @@ class CronTool(Tool, ContextAware):
|
||||
try:
|
||||
ZoneInfo(tz)
|
||||
except (KeyError, Exception):
|
||||
return f"Error: unknown timezone '{tz}'"
|
||||
return ToolResult.error(f"Error: unknown timezone '{tz}'")
|
||||
return None
|
||||
|
||||
def _display_timezone(self, schedule: CronSchedule) -> str:
|
||||
@@ -148,7 +148,7 @@ class CronTool(Tool, ContextAware):
|
||||
) -> str:
|
||||
if action == "add":
|
||||
if self._in_cron_context.get():
|
||||
return "Error: cannot schedule new jobs from within a cron job execution"
|
||||
return ToolResult.error("Error: cannot schedule new jobs from within a cron job execution")
|
||||
return self._add_job(name, message, every_seconds, cron_expr, tz, at)
|
||||
elif action == "list":
|
||||
return self._list_jobs()
|
||||
@@ -166,20 +166,20 @@ class CronTool(Tool, ContextAware):
|
||||
at: str | None,
|
||||
) -> str:
|
||||
if not message:
|
||||
return (
|
||||
return ToolResult.error(
|
||||
"Error: cron action='add' requires a non-empty 'message' parameter "
|
||||
"describing what to do when the job triggers "
|
||||
"(e.g. the reminder text). Retry including message=\"...\"."
|
||||
)
|
||||
session_key = self._session_key.get()
|
||||
if not session_key:
|
||||
return "Error: scheduled cron jobs must be created from a chat session"
|
||||
return ToolResult.error("Error: scheduled cron jobs must be created from a chat session")
|
||||
origin_channel = self._origin_channel.get()
|
||||
origin_chat_id = self._origin_chat_id.get()
|
||||
if not origin_channel or not origin_chat_id:
|
||||
return "Error: scheduled cron jobs must be created from a chat session"
|
||||
return ToolResult.error("Error: scheduled cron jobs must be created from a chat session")
|
||||
if tz and not cron_expr:
|
||||
return "Error: tz can only be used with cron_expr"
|
||||
return ToolResult.error("Error: tz can only be used with cron_expr")
|
||||
if tz:
|
||||
if err := self._validate_timezone(tz):
|
||||
return err
|
||||
@@ -199,7 +199,7 @@ class CronTool(Tool, ContextAware):
|
||||
try:
|
||||
dt = datetime.fromisoformat(at)
|
||||
except ValueError:
|
||||
return f"Error: invalid ISO datetime format '{at}'. Expected format: YYYY-MM-DDTHH:MM:SS"
|
||||
return ToolResult.error(f"Error: invalid ISO datetime format '{at}'. Expected format: YYYY-MM-DDTHH:MM:SS")
|
||||
if dt.tzinfo is None:
|
||||
if err := self._validate_timezone(self._default_timezone):
|
||||
return err
|
||||
@@ -208,7 +208,7 @@ class CronTool(Tool, ContextAware):
|
||||
schedule = CronSchedule(kind="at", at_ms=at_ms)
|
||||
delete_after = True
|
||||
else:
|
||||
return "Error: either every_seconds, cron_expr, or at is required"
|
||||
return ToolResult.error("Error: either every_seconds, cron_expr, or at is required")
|
||||
|
||||
job = self._cron.add_job(
|
||||
name=name or message[:30],
|
||||
@@ -279,7 +279,7 @@ class CronTool(Tool, ContextAware):
|
||||
|
||||
def _remove_job(self, job_id: str | None) -> str:
|
||||
if not job_id:
|
||||
return "Error: job_id is required for remove"
|
||||
return ToolResult.error("Error: job_id is required for remove")
|
||||
result = self._cron.remove_job(job_id)
|
||||
if result == "removed":
|
||||
return f"Removed job {job_id}"
|
||||
|
||||
@@ -9,7 +9,7 @@ from contextlib import suppress
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from nanobot.agent.tools.base import Tool, tool_parameters
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import current_request_session_key
|
||||
from nanobot.agent.tools.schema import (
|
||||
BooleanSchema,
|
||||
@@ -128,7 +128,15 @@ class _ExecSession:
|
||||
) -> _SessionPoll:
|
||||
self.last_access = time.monotonic()
|
||||
if yield_time_ms > 0 and self.process.returncode is None:
|
||||
await asyncio.sleep(min(yield_time_ms, MAX_YIELD_MS) / 1000)
|
||||
wait_s = min(yield_time_ms, MAX_YIELD_MS) / 1000
|
||||
remaining_s = self.deadline - time.monotonic()
|
||||
if remaining_s <= 0:
|
||||
wait_s = 0
|
||||
else:
|
||||
wait_s = min(wait_s, remaining_s)
|
||||
if wait_s > 0:
|
||||
with suppress(asyncio.TimeoutError):
|
||||
await asyncio.wait_for(self.process.wait(), timeout=wait_s)
|
||||
|
||||
if self.process.returncode is None and time.monotonic() >= self.deadline:
|
||||
self._timed_out = True
|
||||
@@ -492,11 +500,12 @@ class WriteStdinTool(Tool):
|
||||
max_output_chars=output_limit,
|
||||
owner_session_key=current_request_session_key(),
|
||||
)
|
||||
return format_session_poll(session_id, poll)
|
||||
result = format_session_poll(session_id, poll)
|
||||
return ToolResult.error(result) if poll.timed_out else result
|
||||
except KeyError:
|
||||
return f"Error: exec session not found: {session_id}"
|
||||
return ToolResult.error(f"Error: exec session not found: {session_id!r}")
|
||||
except Exception as exc:
|
||||
return f"Error writing to exec session: {exc}"
|
||||
return ToolResult.error(f"Error writing to exec session: {exc}")
|
||||
|
||||
async def _wait_for_output(
|
||||
self,
|
||||
@@ -532,13 +541,14 @@ class WriteStdinTool(Tool):
|
||||
joined = "".join(aggregate)
|
||||
if wait_for in joined:
|
||||
poll.output = joined
|
||||
return format_session_poll(session_id, poll)
|
||||
result = format_session_poll(session_id, poll)
|
||||
return ToolResult.error(result) if poll.timed_out else result
|
||||
if poll.done or remaining_ms <= 0:
|
||||
poll.output = "".join(aggregate)
|
||||
result = format_session_poll(session_id, poll)
|
||||
if wait_for not in poll.output:
|
||||
result += f"\nWait target not observed: {wait_for!r}"
|
||||
return result
|
||||
return ToolResult.error(result) if poll.timed_out else result
|
||||
|
||||
|
||||
@tool_parameters(tool_parameters_schema())
|
||||
@@ -606,4 +616,4 @@ class ListExecSessionsTool(Tool):
|
||||
)
|
||||
return "\n".join(lines)
|
||||
except Exception as exc:
|
||||
return f"Error listing exec sessions: {exc}"
|
||||
return ToolResult.error(f"Error listing exec sessions: {exc}")
|
||||
|
||||
@@ -7,7 +7,7 @@ from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from nanobot.agent.tools.base import Tool, tool_parameters
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.file_state import FileStates, _hash_file, current_file_states
|
||||
from nanobot.agent.tools.path_utils import resolve_workspace_path
|
||||
from nanobot.agent.tools.schema import (
|
||||
@@ -268,19 +268,19 @@ class ReadFileTool(_FsTool):
|
||||
) -> Any:
|
||||
try:
|
||||
if not path:
|
||||
return "Error reading file: Unknown path"
|
||||
return ToolResult.error("Error reading file: Unknown path")
|
||||
|
||||
# Device path blacklist
|
||||
if _is_blocked_device(path):
|
||||
return f"Error: Reading {path} is blocked (device path that could hang or produce infinite output)."
|
||||
return ToolResult.error(f"Error: Reading {path} is blocked (device path that could hang or produce infinite output).")
|
||||
|
||||
fp = self._resolve_read(path)
|
||||
if _is_blocked_device(fp):
|
||||
return f"Error: Reading {fp} is blocked (device path that could hang or produce infinite output)."
|
||||
return ToolResult.error(f"Error: Reading {fp} is blocked (device path that could hang or produce infinite output).")
|
||||
if not fp.exists():
|
||||
return f"Error: File not found: {path}"
|
||||
return ToolResult.error(f"Error: File not found: {path}")
|
||||
if not fp.is_file():
|
||||
return f"Error: Not a file: {path}"
|
||||
return ToolResult.error(f"Error: Not a file: {path}")
|
||||
|
||||
# PDF support
|
||||
if fp.suffix.lower() == ".pdf":
|
||||
@@ -343,7 +343,7 @@ class ReadFileTool(_FsTool):
|
||||
mime = detect_image_mime(raw) or mimetypes.guess_type(path)[0]
|
||||
if mime and mime.startswith("image/"):
|
||||
return build_image_content_blocks(raw, mime, str(fp), f"(Image file: {path})")
|
||||
return f"Error: Cannot read binary file {path} (MIME: {mime or 'unknown'}). Only UTF-8 text and images are supported."
|
||||
return ToolResult.error(f"Error: Cannot read binary file {path} (MIME: {mime or 'unknown'}). Only UTF-8 text and images are supported.")
|
||||
|
||||
# Normalize CRLF -> LF before line-splitting. Primarily a Windows
|
||||
# concern (git checkouts with autocrlf, editors saving CRLF) but
|
||||
@@ -357,7 +357,7 @@ class ReadFileTool(_FsTool):
|
||||
if offset < 1:
|
||||
offset = 1
|
||||
if offset > total:
|
||||
return f"Error: offset {offset} is beyond end of file ({total} lines)"
|
||||
return ToolResult.error(f"Error: offset {offset} is beyond end of file ({total} lines)")
|
||||
|
||||
start = offset - 1
|
||||
end = min(start + (limit or self._DEFAULT_LIMIT), total)
|
||||
@@ -381,20 +381,20 @@ class ReadFileTool(_FsTool):
|
||||
self._file_states.record_read(fp, offset=offset, limit=limit)
|
||||
return result
|
||||
except PermissionError as e:
|
||||
return f"Error: {e}"
|
||||
return ToolResult.error(f"Error: {e}")
|
||||
except Exception as e:
|
||||
return f"Error reading file: {e}"
|
||||
return ToolResult.error(f"Error reading file: {e}")
|
||||
|
||||
def _read_pdf(self, fp: Path, pages: str | None) -> str:
|
||||
try:
|
||||
import fitz # pymupdf
|
||||
except ImportError:
|
||||
return "Error: PDF reading requires pymupdf. Install with: pip install pymupdf"
|
||||
return ToolResult.error("Error: PDF reading requires pymupdf. Install with: pip install pymupdf")
|
||||
|
||||
try:
|
||||
doc = fitz.open(str(fp))
|
||||
except Exception as e:
|
||||
return f"Error reading PDF: {e}"
|
||||
return ToolResult.error(f"Error reading PDF: {e}")
|
||||
|
||||
total_pages = len(doc)
|
||||
if pages:
|
||||
@@ -402,10 +402,10 @@ class ReadFileTool(_FsTool):
|
||||
start, end = _parse_page_range(pages, total_pages)
|
||||
except (ValueError, IndexError):
|
||||
doc.close()
|
||||
return f"Error: Invalid page range '{pages}'. Use format like '1-5'."
|
||||
return ToolResult.error(f"Error: Invalid page range '{pages}'. Use format like '1-5'.")
|
||||
if start > end or start >= total_pages:
|
||||
doc.close()
|
||||
return f"Error: Page range '{pages}' is out of bounds (document has {total_pages} pages)."
|
||||
return ToolResult.error(f"Error: Page range '{pages}' is out of bounds (document has {total_pages} pages).")
|
||||
else:
|
||||
start = 0
|
||||
end = min(total_pages - 1, self._MAX_PDF_PAGES - 1)
|
||||
@@ -437,10 +437,10 @@ class ReadFileTool(_FsTool):
|
||||
result = extract_text(fp)
|
||||
|
||||
if result is None:
|
||||
return f"Error: Unsupported file format: {fp.suffix}"
|
||||
return ToolResult.error(f"Error: Unsupported file format: {fp.suffix}")
|
||||
|
||||
if result.startswith("[error:"):
|
||||
return f"Error reading {fp.suffix.upper()} file: {result}"
|
||||
return ToolResult.error(f"Error reading {fp.suffix.upper()} file: {result}")
|
||||
|
||||
if not result:
|
||||
return f"({fp.suffix.upper().lstrip('.')} has no extractable text: {fp})"
|
||||
@@ -492,9 +492,9 @@ class WriteFileTool(_FsTool):
|
||||
self._file_states.record_write(fp)
|
||||
return f"Successfully wrote {len(content)} characters to {fp}"
|
||||
except PermissionError as e:
|
||||
return f"Error: {e}"
|
||||
return ToolResult.error(f"Error: {e}")
|
||||
except Exception as e:
|
||||
return f"Error writing file: {e}"
|
||||
return ToolResult.error(f"Error writing file: {e}")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -830,11 +830,11 @@ class EditFileTool(_FsTool):
|
||||
if new_text is None:
|
||||
raise ValueError("Unknown new_text")
|
||||
if occurrence is not None and occurrence < 1:
|
||||
return "Error: occurrence must be >= 1."
|
||||
return ToolResult.error("Error: occurrence must be >= 1.")
|
||||
if line_hint is not None and line_hint < 1:
|
||||
return "Error: line_hint must be >= 1."
|
||||
return ToolResult.error("Error: line_hint must be >= 1.")
|
||||
if expected_replacements is not None and expected_replacements < 1:
|
||||
return "Error: expected_replacements must be >= 1."
|
||||
return ToolResult.error("Error: expected_replacements must be >= 1.")
|
||||
|
||||
fp = self._resolve_write(path)
|
||||
|
||||
@@ -853,14 +853,14 @@ class EditFileTool(_FsTool):
|
||||
except OSError:
|
||||
fsize = 0
|
||||
if fsize > self._MAX_EDIT_FILE_SIZE:
|
||||
return f"Error: File too large to edit ({fsize / (1024**3):.1f} GiB). Maximum is 1 GiB."
|
||||
return ToolResult.error(f"Error: File too large to edit ({fsize / (1024**3):.1f} GiB). Maximum is 1 GiB.")
|
||||
|
||||
# Create-file: old_text='' but file exists and not empty → reject
|
||||
if old_text == "":
|
||||
raw = fp.read_bytes()
|
||||
content = raw.decode("utf-8")
|
||||
if content.strip():
|
||||
return f"Error: Cannot create file — {path} already exists and is not empty."
|
||||
return ToolResult.error(f"Error: Cannot create file — {path} already exists and is not empty.")
|
||||
fp.write_text(new_text, encoding="utf-8")
|
||||
self._file_states.record_write(fp)
|
||||
return f"Successfully edited {fp}"
|
||||
@@ -878,15 +878,15 @@ class EditFileTool(_FsTool):
|
||||
return self._not_found_msg(old_text, content, path)
|
||||
count = len(matches)
|
||||
if replace_all and occurrence is not None:
|
||||
return "Error: occurrence cannot be used with replace_all=true."
|
||||
return ToolResult.error("Error: occurrence cannot be used with replace_all=true.")
|
||||
if replace_all and line_hint is not None:
|
||||
return "Error: line_hint cannot be used with replace_all=true."
|
||||
return ToolResult.error("Error: line_hint cannot be used with replace_all=true.")
|
||||
if occurrence is not None and line_hint is not None:
|
||||
return "Error: line_hint cannot be used with occurrence."
|
||||
return ToolResult.error("Error: line_hint cannot be used with occurrence.")
|
||||
if count > 1 and not replace_all:
|
||||
if occurrence is not None:
|
||||
if occurrence > count:
|
||||
return (
|
||||
return ToolResult.error(
|
||||
f"Error: occurrence {occurrence} is out of range; "
|
||||
f"old_text appears {count} times."
|
||||
)
|
||||
@@ -894,7 +894,7 @@ class EditFileTool(_FsTool):
|
||||
nearest = min(matches, key=lambda match: abs(match.line - line_hint))
|
||||
distance = abs(nearest.line - line_hint)
|
||||
if sum(1 for match in matches if abs(match.line - line_hint) == distance) > 1:
|
||||
return (
|
||||
return ToolResult.error(
|
||||
f"Error: line_hint {line_hint} is ambiguous; "
|
||||
f"old_text appears {count} times."
|
||||
)
|
||||
@@ -910,7 +910,7 @@ class EditFileTool(_FsTool):
|
||||
"or set replace_all=true."
|
||||
)
|
||||
elif occurrence is not None and occurrence > count:
|
||||
return (
|
||||
return ToolResult.error(
|
||||
f"Error: occurrence {occurrence} is out of range; "
|
||||
f"old_text appears {count} time."
|
||||
)
|
||||
@@ -928,7 +928,7 @@ class EditFileTool(_FsTool):
|
||||
else:
|
||||
selected = [matches[occurrence - 1 if occurrence else 0]]
|
||||
if expected_replacements is not None and len(selected) != expected_replacements:
|
||||
return (
|
||||
return ToolResult.error(
|
||||
f"Error: expected {expected_replacements} replacements but "
|
||||
f"would make {len(selected)}."
|
||||
)
|
||||
@@ -954,9 +954,9 @@ class EditFileTool(_FsTool):
|
||||
msg = f"{warning}\n{msg}"
|
||||
return msg
|
||||
except PermissionError as e:
|
||||
return f"Error: {e}"
|
||||
return ToolResult.error(f"Error: {e}")
|
||||
except Exception as e:
|
||||
return f"Error editing file: {e}"
|
||||
return ToolResult.error(f"Error editing file: {e}")
|
||||
|
||||
def _file_not_found_msg(self, path: str, fp: Path) -> str:
|
||||
"""Build an error message with 'Did you mean ...?' suggestions."""
|
||||
@@ -969,7 +969,7 @@ class EditFileTool(_FsTool):
|
||||
parts = [f"Error: File not found: {path}"]
|
||||
if suggestions:
|
||||
parts.append("Did you mean: " + ", ".join(suggestions) + "?")
|
||||
return "\n".join(parts)
|
||||
return ToolResult.error("\n".join(parts))
|
||||
|
||||
@staticmethod
|
||||
def _not_found_msg(old_text: str, content: str, path: str) -> str:
|
||||
@@ -985,18 +985,18 @@ class EditFileTool(_FsTool):
|
||||
hint_text = ""
|
||||
if hints:
|
||||
hint_text = "\nPossible cause: " + ", ".join(hints) + "."
|
||||
return (
|
||||
return ToolResult.error(
|
||||
f"Error: old_text not found in {path}."
|
||||
f"{hint_text}\nBest match ({best_ratio:.0%} similar) at line {best_start + 1}:\n{diff}"
|
||||
)
|
||||
|
||||
if hints:
|
||||
return (
|
||||
return ToolResult.error(
|
||||
f"Error: old_text not found in {path}. "
|
||||
f"Possible cause: {', '.join(hints)}. "
|
||||
"Copy the exact text from read_file and try again."
|
||||
)
|
||||
return f"Error: old_text not found in {path}. No similar text found. Verify the file content."
|
||||
return ToolResult.error(f"Error: old_text not found in {path}. No similar text found. Verify the file content.")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -1051,9 +1051,9 @@ class ListDirTool(_FsTool):
|
||||
raise ValueError("Unknown path")
|
||||
dp = self._resolve(path)
|
||||
if not dp.exists():
|
||||
return f"Error: Directory not found: {path}"
|
||||
return ToolResult.error(f"Error: Directory not found: {path}")
|
||||
if not dp.is_dir():
|
||||
return f"Error: Not a directory: {path}"
|
||||
return ToolResult.error(f"Error: Not a directory: {path}")
|
||||
|
||||
cap = max_entries or self._DEFAULT_MAX
|
||||
items: list[str] = []
|
||||
@@ -1084,6 +1084,6 @@ class ListDirTool(_FsTool):
|
||||
result += f"\n\n(truncated, showing first {cap} of {total} entries)"
|
||||
return result
|
||||
except PermissionError as e:
|
||||
return f"Error: {e}"
|
||||
return ToolResult.error(f"Error: {e}")
|
||||
except Exception as e:
|
||||
return f"Error listing directory: {e}"
|
||||
return ToolResult.error(f"Error listing directory: {e}")
|
||||
|
||||
@@ -7,7 +7,7 @@ from typing import TYPE_CHECKING, Any
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
from nanobot.agent.tools.base import Tool, tool_parameters
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.schema import (
|
||||
ArraySchema,
|
||||
IntegerSchema,
|
||||
@@ -172,11 +172,11 @@ class ImageGenerationTool(Tool):
|
||||
) -> str:
|
||||
client = self._provider_client()
|
||||
if client is None:
|
||||
return f"Error: unsupported image generation provider '{self.config.provider}'"
|
||||
return ToolResult.error(f"Error: unsupported image generation provider '{self.config.provider}'")
|
||||
|
||||
requested = count or 1
|
||||
if requested > self.config.max_images_per_turn:
|
||||
return (
|
||||
return ToolResult.error(
|
||||
"Error: count exceeds tools.imageGeneration.maxImagesPerTurn "
|
||||
f"({self.config.max_images_per_turn})"
|
||||
)
|
||||
@@ -206,4 +206,4 @@ class ImageGenerationTool(Tool):
|
||||
break
|
||||
return generated_image_tool_result(artifacts)
|
||||
except (ArtifactError, ImageGenerationError, OSError) as exc:
|
||||
return f"Error: {exc}"
|
||||
return ToolResult.error(f"Error: {exc}")
|
||||
|
||||
@@ -8,7 +8,7 @@ from typing import Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.agent.tools.base import Tool
|
||||
from nanobot.agent.tools.base import Tool, ToolResult
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
|
||||
_SKIP_MODULES = frozenset({
|
||||
@@ -96,6 +96,8 @@ class ToolLoader:
|
||||
if not tool_cls.enabled(ctx):
|
||||
continue
|
||||
tool = tool_cls.create(ctx)
|
||||
if is_plugin_source:
|
||||
tool = _LegacyErrorPrefixTool(tool)
|
||||
if registry.has(tool.name):
|
||||
if is_plugin_source and tool.name in builtin_names:
|
||||
logger.warning(
|
||||
@@ -114,3 +116,67 @@ class ToolLoader:
|
||||
except Exception:
|
||||
logger.exception("Failed to register tool: %s", cls_label)
|
||||
return registered
|
||||
|
||||
|
||||
class _LegacyErrorPrefixTool(Tool):
|
||||
"""Compatibility wrapper for external tools using the old error-string contract."""
|
||||
|
||||
_plugin_discoverable = False
|
||||
|
||||
def __init__(self, wrapped: Tool) -> None:
|
||||
self._wrapped = wrapped
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return self._wrapped.name
|
||||
|
||||
@property
|
||||
def description(self) -> str:
|
||||
return self._wrapped.description
|
||||
|
||||
@property
|
||||
def parameters(self) -> dict[str, Any]:
|
||||
return self._wrapped.parameters
|
||||
|
||||
@property
|
||||
def read_only(self) -> bool:
|
||||
return self._wrapped.read_only
|
||||
|
||||
@property
|
||||
def exclusive(self) -> bool:
|
||||
return self._wrapped.exclusive
|
||||
|
||||
@property
|
||||
def concurrency_safe(self) -> bool:
|
||||
return self._wrapped.concurrency_safe
|
||||
|
||||
@property
|
||||
def config_key(self) -> str:
|
||||
return getattr(self._wrapped, "config_key", "")
|
||||
|
||||
def set_context(self, ctx: Any) -> None:
|
||||
set_context = getattr(self._wrapped, "set_context", None)
|
||||
if callable(set_context):
|
||||
set_context(ctx)
|
||||
|
||||
def cast_params(self, params: dict[str, Any]) -> dict[str, Any]:
|
||||
return self._wrapped.cast_params(params)
|
||||
|
||||
def validate_params(self, params: dict[str, Any]) -> list[str]:
|
||||
return self._wrapped.validate_params(params)
|
||||
|
||||
def to_schema(self) -> dict[str, Any]:
|
||||
return self._wrapped.to_schema()
|
||||
|
||||
async def execute(self, **kwargs: Any) -> Any:
|
||||
result = await self._wrapped.execute(**kwargs)
|
||||
if (
|
||||
isinstance(result, str)
|
||||
and not isinstance(result, ToolResult)
|
||||
and result.startswith("Error:")
|
||||
):
|
||||
return ToolResult.error(result)
|
||||
return result
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
return getattr(self._wrapped, name)
|
||||
|
||||
@@ -20,7 +20,7 @@ from contextvars import ContextVar
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from nanobot.agent.tools.base import Tool, tool_parameters
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import ContextAware, RequestContext
|
||||
from nanobot.agent.tools.schema import StringSchema, tool_parameters_schema
|
||||
from nanobot.bus.runtime_events import GoalStateChanged, RuntimeEventBus, RuntimeEventContext
|
||||
@@ -97,8 +97,7 @@ class _GoalToolsMixin(ContextAware):
|
||||
"Sustained objective for this chat thread. First read the built-in **long-goal** skill, "
|
||||
"especially its Start fast section, then call this promptly once the user's intent is clear. "
|
||||
"The goal must still be idempotent, self-contained, bounded, and explicit about done-ness; "
|
||||
"do not delay this tool call to over-plan, research, or decide execution details. "
|
||||
"Do not use this for a single current-turn answer, including one that uses spawn subagents.",
|
||||
"do not delay this tool call to over-plan, research, or decide execution details.",
|
||||
max_length=12_000,
|
||||
),
|
||||
ui_summary=StringSchema(
|
||||
@@ -140,8 +139,6 @@ class LongTaskTool(Tool, _GoalToolsMixin):
|
||||
def description(self) -> str:
|
||||
return (
|
||||
"Mark this thread as a sustained long-running task. "
|
||||
"Use only when the user wants work to persist across future turns or background check-ins; "
|
||||
"do not use for a single current-turn answer, including one that uses spawn subagents. "
|
||||
"First read the built-in **long-goal** skill, especially its Start fast section; then call this "
|
||||
"as soon as the user's intent is clear. Write a good idempotent goal, but do not delay the tool "
|
||||
"call with long planning, research, or execution-detail thinking. "
|
||||
@@ -153,12 +150,12 @@ class LongTaskTool(Tool, _GoalToolsMixin):
|
||||
async def execute(self, goal: str, ui_summary: str | None = None, **kwargs: Any) -> str:
|
||||
sess = self._session()
|
||||
if sess is None:
|
||||
return (
|
||||
return ToolResult.error(
|
||||
"Error: long_task requires an active chat session (missing routing context)."
|
||||
)
|
||||
prior = parse_goal_state(goal_state_raw(sess.metadata))
|
||||
if isinstance(prior, dict) and prior.get("status") == "active":
|
||||
return (
|
||||
return ToolResult.error(
|
||||
"Error: a sustained goal is already active. "
|
||||
"Use complete_goal when finished, or ask the user before replacing it."
|
||||
)
|
||||
@@ -233,7 +230,7 @@ class CompleteGoalTool(Tool, _GoalToolsMixin):
|
||||
async def execute(self, recap: str | None = None, **kwargs: Any) -> str:
|
||||
sess = self._session()
|
||||
if sess is None:
|
||||
return "Error: complete_goal requires an active chat session."
|
||||
return ToolResult.error("Error: complete_goal requires an active chat session.")
|
||||
prior = parse_goal_state(goal_state_raw(sess.metadata))
|
||||
if not isinstance(prior, dict) or prior.get("status") != "active":
|
||||
return "No active goal to complete."
|
||||
|
||||
@@ -14,7 +14,7 @@ from weakref import WeakKeyDictionary
|
||||
import httpx
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.agent.tools.base import Tool
|
||||
from nanobot.agent.tools.base import Tool, ToolResult
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
from nanobot.bus.events import (
|
||||
INBOUND_META_RUNTIME_CONTROL,
|
||||
@@ -461,7 +461,10 @@ class MCPToolWrapper(_MCPWrapperBase):
|
||||
return f"(MCP tool call failed: {type(exc).__name__})"
|
||||
else:
|
||||
# Success — extract text and persist any image content as artifacts.
|
||||
return self._render_call_result(result.content, kwargs)
|
||||
rendered = self._render_call_result(result.content, kwargs)
|
||||
if getattr(result, "isError", False):
|
||||
return ToolResult.error(rendered)
|
||||
return rendered
|
||||
|
||||
return "(MCP tool call failed)" # Unreachable, but satisfies type checkers
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ from typing import Any, Awaitable, Callable
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.agent.tools.base import Tool, tool_parameters
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import ContextAware, RequestContext
|
||||
from nanobot.agent.tools.path_utils import resolve_workspace_path
|
||||
from nanobot.agent.tools.schema import ArraySchema, StringSchema, tool_parameters_schema
|
||||
@@ -198,7 +198,7 @@ class MessageTool(Tool, ContextAware):
|
||||
not isinstance(row, list) or any(not isinstance(label, str) for label in row)
|
||||
for row in buttons
|
||||
):
|
||||
return "Error: buttons must be a list of list of strings"
|
||||
return ToolResult.error("Error: buttons must be a list of list of strings")
|
||||
default_channel = self._default_channel.get()
|
||||
default_chat_id = self._default_chat_id.get()
|
||||
channel = channel or default_channel
|
||||
@@ -210,7 +210,7 @@ class MessageTool(Tool, ContextAware):
|
||||
and str(explicit_chat_id).strip() != ""
|
||||
and str(explicit_chat_id).strip() != str(default_chat_id).strip()
|
||||
):
|
||||
return (
|
||||
return ToolResult.error(
|
||||
"Error: chat_id does not match the active WebSocket conversation. "
|
||||
"Omit chat_id (and usually channel) so delivery uses the current "
|
||||
"conversation id from context — WebSocket client_id strings "
|
||||
@@ -229,16 +229,16 @@ class MessageTool(Tool, ContextAware):
|
||||
message_id = None
|
||||
|
||||
if not channel or not chat_id:
|
||||
return "Error: No target channel/chat specified"
|
||||
return ToolResult.error("Error: No target channel/chat specified")
|
||||
|
||||
if not self._send_callback:
|
||||
return "Error: Message sending not configured"
|
||||
return ToolResult.error("Error: Message sending not configured")
|
||||
|
||||
if media:
|
||||
try:
|
||||
media = self._resolve_media(media)
|
||||
except (OSError, PermissionError, ValueError) as e:
|
||||
return f"Error: media path is not allowed: {str(e)}"
|
||||
return ToolResult.error(f"Error: media path is not allowed: {str(e)}")
|
||||
|
||||
metadata = dict(self._default_metadata.get()) if same_target else {}
|
||||
if message_id:
|
||||
@@ -270,4 +270,4 @@ class MessageTool(Tool, ContextAware):
|
||||
button_info = f" with {sum(len(row) for row in buttons)} button(s)" if buttons else ""
|
||||
return f"Message sent to {channel}:{chat_id}{media_info}{button_info}"
|
||||
except Exception as e:
|
||||
return f"Error sending message: {str(e)}"
|
||||
return ToolResult.error(f"Error sending message: {str(e)}")
|
||||
|
||||
@@ -3,7 +3,11 @@
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from nanobot.agent.tools.base import Tool
|
||||
from nanobot.agent.tools.base import Tool, ToolResult
|
||||
|
||||
|
||||
def is_tool_error_result(name: str, result: Any) -> bool:
|
||||
return isinstance(result, ToolResult) and result.is_error
|
||||
|
||||
|
||||
class ToolRegistry:
|
||||
@@ -100,22 +104,26 @@ class ToolRegistry:
|
||||
suggestion = self._suggest_name(str(name))
|
||||
hint = f" Did you mean '{suggestion}'? Tool names must match exactly." if suggestion else ""
|
||||
return None, params, (
|
||||
f"Error: Tool '{name}' not found.{hint} Available: {', '.join(self.tool_names)}"
|
||||
ToolResult.error(
|
||||
f"Error: Tool '{name}' not found.{hint} Available: {', '.join(self.tool_names)}"
|
||||
)
|
||||
)
|
||||
|
||||
params = self._coerce_params(tool, params)
|
||||
if not isinstance(params, dict):
|
||||
return tool, params, (
|
||||
f"Error: Tool '{name}' parameters must be a JSON object, got "
|
||||
f"{type(params).__name__}. Use named parameters like "
|
||||
'tool_name(param1="value1", param2="value2") matching the tool schema.'
|
||||
ToolResult.error(
|
||||
f"Error: Tool '{name}' parameters must be a JSON object, got "
|
||||
f"{type(params).__name__}. Use named parameters like "
|
||||
'tool_name(param1="value1", param2="value2") matching the tool schema.'
|
||||
)
|
||||
)
|
||||
|
||||
cast_params = tool.cast_params(params)
|
||||
errors = tool.validate_params(cast_params)
|
||||
if errors:
|
||||
return tool, cast_params, (
|
||||
f"Error: Invalid parameters for tool '{name}': " + "; ".join(errors)
|
||||
ToolResult.error(f"Error: Invalid parameters for tool '{name}': " + "; ".join(errors))
|
||||
)
|
||||
return tool, cast_params, None
|
||||
|
||||
@@ -159,16 +167,16 @@ class ToolRegistry:
|
||||
hint = "\n\n[Analyze the error above and try a different approach.]"
|
||||
tool, params, error = self.prepare_call(name, params)
|
||||
if error:
|
||||
return error + hint
|
||||
return ToolResult.error(str(error) + hint)
|
||||
|
||||
try:
|
||||
assert tool is not None # guarded by prepare_call()
|
||||
result = await tool.execute(**params)
|
||||
if isinstance(result, str) and result.startswith("Error"):
|
||||
return result + hint
|
||||
if is_tool_error_result(name, result):
|
||||
return ToolResult.error(str(result) + hint)
|
||||
return result
|
||||
except Exception as e:
|
||||
return f"Error executing {name}: {str(e)}" + hint
|
||||
return ToolResult.error(f"Error executing {name}: {str(e)}" + hint)
|
||||
|
||||
@property
|
||||
def tool_names(self) -> list[str]:
|
||||
|
||||
@@ -9,6 +9,7 @@ from contextlib import suppress
|
||||
from pathlib import Path, PurePosixPath
|
||||
from typing import Any, Iterable, TypeVar
|
||||
|
||||
from nanobot.agent.tools.base import ToolResult
|
||||
from nanobot.agent.tools.filesystem import ListDirTool, _FsTool
|
||||
|
||||
_DEFAULT_HEAD_LIMIT = 250
|
||||
@@ -218,12 +219,12 @@ class FindFilesTool(_SearchTool):
|
||||
try:
|
||||
target = self._resolve(path or ".")
|
||||
if not target.exists():
|
||||
return f"Error: Path not found: {path}"
|
||||
return ToolResult.error(f"Error: Path not found: {path}")
|
||||
if not (target.is_dir() or target.is_file()):
|
||||
return f"Error: Unsupported path: {path}"
|
||||
return ToolResult.error(f"Error: Unsupported path: {path}")
|
||||
|
||||
if sort not in {"path", "modified"}:
|
||||
return "Error: sort must be 'path' or 'modified'"
|
||||
return ToolResult.error("Error: sort must be 'path' or 'modified'")
|
||||
|
||||
limit = (
|
||||
_DEFAULT_FILE_HEAD_LIMIT
|
||||
@@ -271,9 +272,9 @@ class FindFilesTool(_SearchTool):
|
||||
result += "\n\n" + note
|
||||
return result
|
||||
except PermissionError as e:
|
||||
return f"Error: {e}"
|
||||
return ToolResult.error(f"Error: {e}")
|
||||
except Exception as e:
|
||||
return f"Error finding files: {e}"
|
||||
return ToolResult.error(f"Error finding files: {e}")
|
||||
|
||||
|
||||
class GrepTool(_SearchTool):
|
||||
@@ -425,16 +426,16 @@ class GrepTool(_SearchTool):
|
||||
try:
|
||||
target = self._resolve(path or ".")
|
||||
if not target.exists():
|
||||
return f"Error: Path not found: {path}"
|
||||
return ToolResult.error(f"Error: Path not found: {path}")
|
||||
if not (target.is_dir() or target.is_file()):
|
||||
return f"Error: Unsupported path: {path}"
|
||||
return ToolResult.error(f"Error: Unsupported path: {path}")
|
||||
|
||||
flags = re.IGNORECASE if case_insensitive else 0
|
||||
try:
|
||||
needle = re.escape(pattern) if fixed_strings else pattern
|
||||
regex = re.compile(needle, flags)
|
||||
except re.error as e:
|
||||
return f"Error: invalid regex pattern: {e}"
|
||||
return ToolResult.error(f"Error: invalid regex pattern: {e}")
|
||||
|
||||
if head_limit is not None:
|
||||
limit = None if head_limit == 0 else head_limit
|
||||
@@ -579,6 +580,6 @@ class GrepTool(_SearchTool):
|
||||
result += "\n\n" + "\n".join(notes)
|
||||
return result
|
||||
except PermissionError as e:
|
||||
return f"Error: {e}"
|
||||
return ToolResult.error(f"Error: {e}")
|
||||
except Exception as e:
|
||||
return f"Error searching files: {e}"
|
||||
return ToolResult.error(f"Error searching files: {e}")
|
||||
|
||||
+24
-24
@@ -7,7 +7,7 @@ from typing import TYPE_CHECKING, Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.agent.tools.base import Tool
|
||||
from nanobot.agent.tools.base import Tool, ToolResult
|
||||
from nanobot.agent.tools.context import ContextAware, RequestContext
|
||||
from nanobot.agent.tools.runtime_state import RuntimeState
|
||||
from nanobot.config_base import Base
|
||||
@@ -216,7 +216,7 @@ class MyTool(Tool, ContextAware):
|
||||
@staticmethod
|
||||
def _validate_key(key: str | None, label: str = "key") -> str | None:
|
||||
if not key or not key.strip():
|
||||
return f"Error: '{label}' cannot be empty or whitespace"
|
||||
return ToolResult.error(f"Error: '{label}' cannot be empty or whitespace")
|
||||
return None
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
@@ -321,7 +321,7 @@ class MyTool(Tool, ContextAware):
|
||||
if action in ("inspect", "check"):
|
||||
return self._inspect(key)
|
||||
if not self._modify_allowed:
|
||||
return "Error: set is disabled (tools.my.allow_set is false)"
|
||||
return ToolResult.error("Error: set is disabled (tools.my.allow_set is false)")
|
||||
if action in ("modify", "set"):
|
||||
return self._modify(key, value)
|
||||
return f"Unknown action: {action}"
|
||||
@@ -333,7 +333,7 @@ class MyTool(Tool, ContextAware):
|
||||
return self._inspect_all()
|
||||
top = key.split(".")[0]
|
||||
if top in self._DENIED_ATTRS or top.startswith("__"):
|
||||
return f"Error: '{top}' is not accessible"
|
||||
return ToolResult.error(f"Error: '{top}' is not accessible")
|
||||
obj, err = self._resolve_path(key)
|
||||
if err:
|
||||
# "scratchpad" alias for _runtime_vars
|
||||
@@ -343,12 +343,12 @@ class MyTool(Tool, ContextAware):
|
||||
# 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 f"Error: {err}"
|
||||
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 f"Error: '{key}' not found"
|
||||
return ToolResult.error(f"Error: '{key}' not found")
|
||||
return self._format_value(obj, key)
|
||||
|
||||
def _inspect_all(self) -> str:
|
||||
@@ -379,21 +379,21 @@ class MyTool(Tool, ContextAware):
|
||||
top = key.split(".")[0]
|
||||
if top in self.BLOCKED or top in self._DENIED_ATTRS or top.startswith("__") or top.lower() in self._SENSITIVE_NAMES:
|
||||
self._audit("modify", f"BLOCKED {key}")
|
||||
return f"Error: '{key}' is protected and cannot be modified"
|
||||
return ToolResult.error(f"Error: '{key}' is protected and cannot be modified")
|
||||
if top in self.READ_ONLY:
|
||||
self._audit("modify", f"READ_ONLY {key}")
|
||||
return f"Error: '{key}' is read-only and cannot be modified"
|
||||
return ToolResult.error(f"Error: '{key}' is read-only and cannot be modified")
|
||||
if "." in key:
|
||||
parent_path, leaf = key.rsplit(".", 1)
|
||||
if leaf in self._DENIED_ATTRS or leaf.startswith("__"):
|
||||
self._audit("modify", f"BLOCKED leaf '{leaf}'")
|
||||
return f"Error: '{leaf}' is not accessible"
|
||||
return ToolResult.error(f"Error: '{leaf}' is not accessible")
|
||||
if leaf.lower() in self._SENSITIVE_NAMES:
|
||||
self._audit("modify", f"BLOCKED sensitive leaf '{leaf}'")
|
||||
return f"Error: '{leaf}' is not accessible"
|
||||
return ToolResult.error(f"Error: '{leaf}' is not accessible")
|
||||
parent, err = self._resolve_path(parent_path)
|
||||
if err:
|
||||
return f"Error: {err}"
|
||||
return ToolResult.error(f"Error: {err}")
|
||||
if isinstance(parent, dict):
|
||||
parent[leaf] = value
|
||||
else:
|
||||
@@ -408,11 +408,11 @@ class MyTool(Tool, ContextAware):
|
||||
|
||||
def _modify_model_preset(self, value: Any) -> str:
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
return "Error: 'model_preset' must be a non-empty string"
|
||||
return ToolResult.error("Error: 'model_preset' must be a non-empty string")
|
||||
name = value.strip()
|
||||
result = self._modify_free("model_preset", name)
|
||||
if result.startswith("Error:"):
|
||||
return result if result.endswith((".", "!", "?")) else f"{result}."
|
||||
if isinstance(result, ToolResult) and result.is_error:
|
||||
return result if result.endswith((".", "!", "?")) else ToolResult.error(f"{result}.")
|
||||
return (
|
||||
f"{result}; model is now {self._runtime_state.model!r}; "
|
||||
f"context_window_tokens is now {self._runtime_state.context_window_tokens!r}"
|
||||
@@ -422,19 +422,19 @@ class MyTool(Tool, ContextAware):
|
||||
spec = self.RESTRICTED[key]
|
||||
expected = spec["type"]
|
||||
if expected is int and isinstance(value, bool):
|
||||
return f"Error: '{key}' must be {expected.__name__}, got bool"
|
||||
return ToolResult.error(f"Error: '{key}' must be {expected.__name__}, got bool")
|
||||
if not isinstance(value, expected):
|
||||
try:
|
||||
value = expected(value)
|
||||
except (ValueError, TypeError):
|
||||
return f"Error: '{key}' must be {expected.__name__}, got {type(value).__name__}"
|
||||
return ToolResult.error(f"Error: '{key}' must be {expected.__name__}, got {type(value).__name__}")
|
||||
old = getattr(self._runtime_state, key)
|
||||
if "min" in spec and value < spec["min"]:
|
||||
return f"Error: '{key}' must be >= {spec['min']}"
|
||||
return ToolResult.error(f"Error: '{key}' must be >= {spec['min']}")
|
||||
if "max" in spec and value > spec["max"]:
|
||||
return f"Error: '{key}' must be <= {spec['max']}"
|
||||
return ToolResult.error(f"Error: '{key}' must be <= {spec['max']}")
|
||||
if "min_len" in spec and len(str(value)) < spec["min_len"]:
|
||||
return f"Error: '{key}' must be at least {spec['min_len']} characters"
|
||||
return ToolResult.error(f"Error: '{key}' must be at least {spec['min_len']} characters")
|
||||
setattr(self._runtime_state, key, value)
|
||||
if key == "model":
|
||||
self._runtime_state._active_preset = None
|
||||
@@ -458,25 +458,25 @@ class MyTool(Tool, ContextAware):
|
||||
"modify",
|
||||
f"REJECTED type mismatch {key}: expects {old_t.__name__}, got {new_t.__name__}",
|
||||
)
|
||||
return f"Error: '{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 f"Error: {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})"
|
||||
if callable(value):
|
||||
self._audit("modify", f"REJECTED callable {key}")
|
||||
return "Error: cannot store callable values"
|
||||
return ToolResult.error("Error: cannot store callable values")
|
||||
err = self._validate_json_safe(value)
|
||||
if err:
|
||||
self._audit("modify", f"REJECTED {key}: {err}")
|
||||
return f"Error: {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:
|
||||
self._audit("modify", f"REJECTED {key}: max keys ({self._MAX_RUNTIME_KEYS}) reached")
|
||||
return f"Error: scratchpad is full (max {self._MAX_RUNTIME_KEYS} keys). Remove unused keys first."
|
||||
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}")
|
||||
|
||||
@@ -15,7 +15,7 @@ from typing import Any
|
||||
from loguru import logger
|
||||
from pydantic import Field
|
||||
|
||||
from nanobot.agent.tools.base import Tool, tool_parameters
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.context import current_request_session_key
|
||||
from nanobot.agent.tools.exec_session import (
|
||||
DEFAULT_EXEC_SESSION_MANAGER,
|
||||
@@ -256,7 +256,7 @@ class ExecTool(Tool):
|
||||
command = command or cmd
|
||||
working_dir = working_dir or workdir
|
||||
if not command:
|
||||
return "Error: Missing command. Provide command or cmd."
|
||||
return ToolResult.error("Error: Missing command. Provide command or cmd.")
|
||||
if max_output_chars is None:
|
||||
max_output_chars = max_output_tokens
|
||||
|
||||
@@ -283,7 +283,7 @@ class ExecTool(Tool):
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
await self._kill_process(process)
|
||||
return f"Error: Command timed out after {prepared.timeout} seconds"
|
||||
return ToolResult.error(f"Error: Command timed out after {prepared.timeout} seconds")
|
||||
except asyncio.CancelledError:
|
||||
await self._kill_process(process)
|
||||
raise
|
||||
@@ -314,7 +314,7 @@ class ExecTool(Tool):
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
return f"Error executing command: {str(e)}"
|
||||
return ToolResult.error(f"Error executing command: {str(e)}")
|
||||
|
||||
async def _execute_session(
|
||||
self,
|
||||
@@ -339,9 +339,10 @@ class ExecTool(Tool):
|
||||
MAX_OUTPUT_CHARS,
|
||||
),
|
||||
)
|
||||
return format_session_poll(session_id, poll)
|
||||
result = format_session_poll(session_id, poll)
|
||||
return ToolResult.error(result) if poll.timed_out else result
|
||||
except Exception as exc:
|
||||
return f"Error executing command: {exc}"
|
||||
return ToolResult.error(f"Error executing command: {exc}")
|
||||
|
||||
def _resolve_timeout(self, timeout: int | None) -> int | None:
|
||||
"""Resolve the effective hard timeout in seconds (None = no limit).
|
||||
@@ -383,12 +384,12 @@ class ExecTool(Tool):
|
||||
requested = Path(cwd).expanduser().resolve()
|
||||
resolved_root = Path(workspace_root).expanduser().resolve()
|
||||
except Exception:
|
||||
return (
|
||||
return ToolResult.error(
|
||||
"Error: working_dir could not be resolved"
|
||||
+ _WORKSPACE_BOUNDARY_NOTE
|
||||
)
|
||||
if not is_path_within(requested, resolved_root):
|
||||
return (
|
||||
return ToolResult.error(
|
||||
"Error: working_dir is outside the configured workspace"
|
||||
+ _WORKSPACE_BOUNDARY_NOTE
|
||||
)
|
||||
@@ -504,24 +505,24 @@ class ExecTool(Tool):
|
||||
if not shell:
|
||||
return None, None
|
||||
if _IS_WINDOWS:
|
||||
return None, "Error: shell parameter is not supported on Windows"
|
||||
return None, ToolResult.error("Error: shell parameter is not supported on Windows")
|
||||
if "\0" in shell or "\n" in shell or "\r" in shell:
|
||||
return None, "Error: shell contains invalid characters"
|
||||
return None, ToolResult.error("Error: shell contains invalid characters")
|
||||
allowed = {"sh", "bash", "zsh"}
|
||||
path = Path(shell).expanduser()
|
||||
if path.is_absolute():
|
||||
if path.name not in allowed:
|
||||
return None, f"Error: unsupported shell {shell!r}. Allowed: bash, sh, zsh"
|
||||
return None, ToolResult.error(f"Error: unsupported shell {shell!r}. Allowed: bash, sh, zsh")
|
||||
if not path.is_file() or not os.access(path, os.X_OK):
|
||||
return None, f"Error: shell is not executable: {shell}"
|
||||
return None, ToolResult.error(f"Error: shell is not executable: {shell}")
|
||||
return str(path), None
|
||||
if "/" in shell or "\\" in shell:
|
||||
return None, "Error: shell must be a shell name or absolute path"
|
||||
return None, ToolResult.error("Error: shell must be a shell name or absolute path")
|
||||
if shell not in allowed:
|
||||
return None, f"Error: unsupported shell {shell!r}. Allowed: bash, sh, zsh"
|
||||
return None, ToolResult.error(f"Error: unsupported shell {shell!r}. Allowed: bash, sh, zsh")
|
||||
resolved = shutil.which(shell)
|
||||
if not resolved:
|
||||
return None, f"Error: shell not found: {shell}"
|
||||
return None, ToolResult.error(f"Error: shell not found: {shell}")
|
||||
return resolved, None
|
||||
|
||||
@staticmethod
|
||||
@@ -608,10 +609,10 @@ class ExecTool(Tool):
|
||||
if not explicitly_allowed:
|
||||
for pattern in self.deny_patterns:
|
||||
if re.search(pattern, lower):
|
||||
return "Error: Command blocked by deny pattern filter"
|
||||
return ToolResult.error("Error: Command blocked by deny pattern filter")
|
||||
|
||||
if self.allow_patterns:
|
||||
return "Error: Command blocked by allowlist filter (not in allowlist)"
|
||||
return ToolResult.error("Error: Command blocked by allowlist filter (not in allowlist)")
|
||||
|
||||
from nanobot.security.network import contains_internal_url
|
||||
if contains_internal_url(
|
||||
@@ -621,12 +622,12 @@ class ExecTool(Tool):
|
||||
),
|
||||
):
|
||||
# The runner turns this marker into a non-retryable security hint.
|
||||
return "Error: Command blocked by safety guard (internal/private URL detected)"
|
||||
return ToolResult.error("Error: Command blocked by safety guard (internal/private URL detected)")
|
||||
|
||||
should_restrict = self.restrict_to_workspace if restrict_to_workspace is None else restrict_to_workspace
|
||||
if should_restrict:
|
||||
if "..\\" in cmd or "../" in cmd:
|
||||
return (
|
||||
return ToolResult.error(
|
||||
"Error: Command blocked by safety guard (path traversal detected)"
|
||||
+ _WORKSPACE_BOUNDARY_NOTE
|
||||
)
|
||||
@@ -661,7 +662,7 @@ class ExecTool(Tool):
|
||||
if not allowed and resolved_workspace is not None:
|
||||
allowed = is_path_within(p, resolved_workspace)
|
||||
if p.is_absolute() and not allowed:
|
||||
return (
|
||||
return ToolResult.error(
|
||||
"Error: Command blocked by safety guard (path outside working dir)"
|
||||
+ _WORKSPACE_BOUNDARY_NOTE
|
||||
)
|
||||
|
||||
@@ -63,9 +63,6 @@ class SpawnTool(Tool, ContextAware):
|
||||
return (
|
||||
"Spawn a subagent to handle a task in the background. "
|
||||
"Use this for complex or time-consuming tasks that can run independently. "
|
||||
"For MapReduce-style work, spawn only independent map slices with clear "
|
||||
"boundaries; keep reduction, conflict resolution, and final user-facing "
|
||||
"synthesis in the main agent. "
|
||||
"The subagent will complete the task and report back when done. "
|
||||
"For deliverables or existing projects, inspect the workspace first "
|
||||
"and use a dedicated subdirectory when helpful."
|
||||
|
||||
+28
-28
@@ -14,7 +14,7 @@ import httpx
|
||||
from loguru import logger
|
||||
from pydantic import Field
|
||||
|
||||
from nanobot.agent.tools.base import Tool, tool_parameters
|
||||
from nanobot.agent.tools.base import Tool, ToolResult, tool_parameters
|
||||
from nanobot.agent.tools.schema import (
|
||||
BooleanSchema,
|
||||
IntegerSchema,
|
||||
@@ -395,13 +395,13 @@ class WebSearchTool(Tool):
|
||||
elif provider == "keenable":
|
||||
return await self._search_keenable(query, n)
|
||||
else:
|
||||
return f"Error: unknown search provider '{provider}'"
|
||||
return ToolResult.error(f"Error: unknown search provider '{provider}'")
|
||||
|
||||
async def _search_olostep(self, query: str, n: int) -> str:
|
||||
try:
|
||||
from olostep import AsyncOlostep, Olostep_BaseError
|
||||
except ImportError:
|
||||
return "Error: olostep package not installed. Run: pip install olostep"
|
||||
return ToolResult.error("Error: olostep package not installed. Run: pip install olostep")
|
||||
api_key = self.config.api_key or os.environ.get("OLOSTEP_API_KEY", "")
|
||||
if not api_key:
|
||||
logger.warning("OLOSTEP_API_KEY not set, falling back to DuckDuckGo")
|
||||
@@ -445,9 +445,9 @@ class WebSearchTool(Tool):
|
||||
items = [{"title": answer_text or "Olostep answer", "url": "", "content": "\n".join(source_lines)}]
|
||||
return _format_results(query, items, n)
|
||||
except Olostep_BaseError as e:
|
||||
return f"Olostep search error: {type(e).__name__}: {e}"
|
||||
return ToolResult.error(f"Error: Olostep search error: {type(e).__name__}: {e}")
|
||||
except Exception as e:
|
||||
return f"Olostep search error: {type(e).__name__}: {e}"
|
||||
return ToolResult.error(f"Error: Olostep search error: {type(e).__name__}: {e}")
|
||||
|
||||
async def _search_brave(self, query: str, n: int) -> str:
|
||||
api_key = self.config.api_key or os.environ.get("BRAVE_API_KEY", "")
|
||||
@@ -481,13 +481,13 @@ class WebSearchTool(Tool):
|
||||
return _format_results(query, items, n)
|
||||
except httpx.HTTPStatusError as e:
|
||||
if e.response.status_code == 429:
|
||||
return (
|
||||
return ToolResult.error(
|
||||
"Error: Brave search rate limited after retry. "
|
||||
"Retry later or reduce consecutive web_search calls."
|
||||
)
|
||||
return f"Error: {e}"
|
||||
return ToolResult.error(f"Error: {e}")
|
||||
except Exception as e:
|
||||
return f"Error: {e}"
|
||||
return ToolResult.error(f"Error: {e}")
|
||||
|
||||
async def _search_tavily(self, query: str, n: int) -> str:
|
||||
api_key = self.config.api_key or os.environ.get("TAVILY_API_KEY", "")
|
||||
@@ -505,7 +505,7 @@ class WebSearchTool(Tool):
|
||||
r.raise_for_status()
|
||||
return _format_results(query, r.json().get("results", []), n)
|
||||
except Exception as e:
|
||||
return f"Error: {e}"
|
||||
return ToolResult.error(f"Error: {e}")
|
||||
|
||||
async def _search_keenable(self, query: str, n: int) -> str:
|
||||
api_key = self.config.api_key or os.environ.get("KEENABLE_API_KEY", "")
|
||||
@@ -540,10 +540,10 @@ class WebSearchTool(Tool):
|
||||
return _format_results(query, items, n)
|
||||
except httpx.HTTPStatusError as e:
|
||||
if e.response.status_code == 429:
|
||||
return "Error: Keenable search rate limited. Try again later or reduce search frequency."
|
||||
return f"Error: Keenable search failed ({e.response.status_code}): {e}"
|
||||
return ToolResult.error("Error: Keenable search rate limited. Try again later or reduce search frequency.")
|
||||
return ToolResult.error(f"Error: Keenable search failed ({e.response.status_code}): {e}")
|
||||
except Exception as e:
|
||||
return f"Error: Keenable search failed: {e}"
|
||||
return ToolResult.error(f"Error: Keenable search failed: {e}")
|
||||
|
||||
async def _search_searxng(self, query: str, n: int) -> str:
|
||||
base_url = (self.config.base_url or os.environ.get("SEARXNG_BASE_URL", "")).strip()
|
||||
@@ -553,7 +553,7 @@ class WebSearchTool(Tool):
|
||||
endpoint = f"{base_url.rstrip('/')}/search"
|
||||
is_valid, error_msg = _validate_url(endpoint)
|
||||
if not is_valid:
|
||||
return f"Error: invalid SearXNG URL: {error_msg}"
|
||||
return ToolResult.error(f"Error: invalid SearXNG URL: {error_msg}")
|
||||
try:
|
||||
async with httpx.AsyncClient(proxy=self.proxy) as client:
|
||||
r = await client.get(
|
||||
@@ -565,7 +565,7 @@ class WebSearchTool(Tool):
|
||||
r.raise_for_status()
|
||||
return _format_results(query, r.json().get("results", []), n)
|
||||
except Exception as e:
|
||||
return f"Error: {e}"
|
||||
return ToolResult.error(f"Error: {e}")
|
||||
|
||||
async def _search_jina(self, query: str, n: int) -> str:
|
||||
api_key = self.config.api_key or os.environ.get("JINA_API_KEY", "")
|
||||
@@ -616,7 +616,7 @@ class WebSearchTool(Tool):
|
||||
]
|
||||
return _format_results(query, items, n)
|
||||
except Exception as e:
|
||||
return f"Error: {e}"
|
||||
return ToolResult.error(f"Error: {e}")
|
||||
|
||||
async def _search_exa(self, query: str, n: int) -> str:
|
||||
api_key = self.config.api_key or os.environ.get("EXA_API_KEY", "")
|
||||
@@ -663,10 +663,10 @@ class WebSearchTool(Tool):
|
||||
return _format_results(query, items, n)
|
||||
except httpx.HTTPStatusError as e:
|
||||
if e.response.status_code == 429:
|
||||
return "Error: Exa search rate limited. Try again later or reduce search frequency."
|
||||
return f"Error: Exa search failed ({e.response.status_code}): {e}"
|
||||
return ToolResult.error("Error: Exa search rate limited. Try again later or reduce search frequency.")
|
||||
return ToolResult.error(f"Error: Exa search failed ({e.response.status_code}): {e}")
|
||||
except Exception as e:
|
||||
return f"Error: Exa search failed: {e}"
|
||||
return ToolResult.error(f"Error: Exa search failed: {e}")
|
||||
|
||||
async def _search_volcengine(
|
||||
self,
|
||||
@@ -690,7 +690,7 @@ class WebSearchTool(Tool):
|
||||
normalized_time_range = _normalize_volcengine_time_range(time_range) if time_range else None
|
||||
normalized_auth_level = _normalize_volcengine_auth_level(auth_level) if auth_level is not None else None
|
||||
except ValueError as e:
|
||||
return f"Error: {e}"
|
||||
return ToolResult.error(f"Error: {e}")
|
||||
|
||||
body: dict[str, Any] = {
|
||||
"Query": query,
|
||||
@@ -723,18 +723,18 @@ class WebSearchTool(Tool):
|
||||
data = r.json()
|
||||
except httpx.HTTPStatusError as e:
|
||||
if e.response.status_code == 429:
|
||||
return "Error: Volcengine search rate limited. Try again later or reduce search frequency."
|
||||
return f"Error: Volcengine search failed ({e.response.status_code}): {e}"
|
||||
return ToolResult.error("Error: Volcengine search rate limited. Try again later or reduce search frequency.")
|
||||
return ToolResult.error(f"Error: Volcengine search failed ({e.response.status_code}): {e}")
|
||||
except Exception as e:
|
||||
return f"Error: Volcengine search failed: {e}"
|
||||
return ToolResult.error(f"Error: Volcengine search failed: {e}")
|
||||
|
||||
error = (data.get("ResponseMetadata") or {}).get("Error") or data.get("Error") or data.get("error")
|
||||
if error:
|
||||
if isinstance(error, dict):
|
||||
code = error.get("Code") or error.get("code") or "unknown"
|
||||
message = error.get("Message") or error.get("message") or error
|
||||
return f"Error: Volcengine search error {code}: {message}"
|
||||
return f"Error: Volcengine search error: {error}"
|
||||
return ToolResult.error(f"Error: Volcengine search error {code}: {message}")
|
||||
return ToolResult.error(f"Error: Volcengine search error: {error}")
|
||||
|
||||
result = data.get("Result") or data
|
||||
web_results = result.get("WebResults") or result.get("webResults") or result.get("results") or []
|
||||
@@ -791,7 +791,7 @@ class WebSearchTool(Tool):
|
||||
return _format_results(query, items, n)
|
||||
except Exception as e:
|
||||
logger.warning("DuckDuckGo search failed: {}", e)
|
||||
return f"Error: DuckDuckGo search failed ({e})"
|
||||
return ToolResult.error(f"Error: DuckDuckGo search failed ({e})")
|
||||
|
||||
async def _search_bocha(self, query: str, n: int, freshness: str = "noLimit") -> str:
|
||||
api_key = self.config.api_key or os.environ.get("BOCHA_API_KEY", "")
|
||||
@@ -819,7 +819,7 @@ class WebSearchTool(Tool):
|
||||
timeout=self.config.timeout,
|
||||
)
|
||||
if r.status_code == 429:
|
||||
return "Error: Bocha search rate-limited (HTTP 429). Wait and retry."
|
||||
return ToolResult.error("Error: Bocha search rate-limited (HTTP 429). Wait and retry.")
|
||||
r.raise_for_status()
|
||||
data = r.json()
|
||||
wrapped_data = data.get("data") if isinstance(data, dict) else None
|
||||
@@ -839,9 +839,9 @@ class WebSearchTool(Tool):
|
||||
]
|
||||
return _format_results(query, items, n)
|
||||
except httpx.HTTPStatusError as e:
|
||||
return f"Error: Bocha search HTTP {e.response.status_code}: {e.response.text[:200]}"
|
||||
return ToolResult.error(f"Error: Bocha search HTTP {e.response.status_code}: {e.response.text[:200]}")
|
||||
except Exception as e:
|
||||
return f"Error: {e}"
|
||||
return ToolResult.error(f"Error: {e}")
|
||||
|
||||
|
||||
@tool_parameters(
|
||||
|
||||
+22
-1
@@ -8,6 +8,7 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import hmac
|
||||
import json as _json
|
||||
import time
|
||||
import uuid
|
||||
@@ -392,7 +393,10 @@ async def handle_health(request: web.Request) -> web.Response:
|
||||
|
||||
|
||||
def create_app(
|
||||
agent_loop, model_name: str = "nanobot", request_timeout: float = 120.0
|
||||
agent_loop,
|
||||
model_name: str = "nanobot",
|
||||
request_timeout: float = 120.0,
|
||||
api_key: str = "",
|
||||
) -> web.Application:
|
||||
"""Create the aiohttp application.
|
||||
|
||||
@@ -400,6 +404,7 @@ def create_app(
|
||||
agent_loop: An initialized AgentLoop instance.
|
||||
model_name: Model name reported in responses.
|
||||
request_timeout: Per-request timeout in seconds.
|
||||
api_key: Optional API key for Bearer-token authentication.
|
||||
"""
|
||||
app = web.Application(client_max_size=20 * 1024 * 1024) # 20MB for base64 images
|
||||
app["agent_loop"] = agent_loop
|
||||
@@ -407,6 +412,22 @@ def create_app(
|
||||
app["request_timeout"] = request_timeout
|
||||
app["session_locks"] = {} # per-user locks, keyed by session_key
|
||||
|
||||
@web.middleware
|
||||
async def auth_middleware(request: web.Request, handler) -> web.StreamResponse:
|
||||
if not api_key:
|
||||
return await handler(request)
|
||||
# Allow unauthenticated health checks.
|
||||
if request.path == "/health":
|
||||
return await handler(request)
|
||||
auth = request.headers.get("Authorization", "")
|
||||
if not auth.startswith("Bearer "):
|
||||
return _error_json(401, "Missing Authorization header. Use: Bearer <api_key>")
|
||||
if not hmac.compare_digest(auth[len("Bearer "):], api_key):
|
||||
return _error_json(401, "Invalid API key")
|
||||
return await handler(request)
|
||||
|
||||
app.middlewares.append(auth_middleware)
|
||||
|
||||
app.router.add_post("/v1/chat/completions", handle_chat_completions)
|
||||
app.router.add_get("/v1/models", handle_models)
|
||||
app.router.add_get("/health", handle_health)
|
||||
|
||||
@@ -2,7 +2,10 @@
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.bus.outbound_events import OutboundEvent
|
||||
|
||||
# Optional ``OutboundMessage.metadata`` key for structured, channel-agnostic UI
|
||||
# payloads. Value is JSON-serializable with at least ``kind``; rich clients may
|
||||
@@ -39,9 +42,9 @@ class InboundMessage:
|
||||
class OutboundMessage:
|
||||
"""Message to send to a chat channel.
|
||||
|
||||
``metadata`` can carry routing (``message_id``, …), trace flags (``_progress``),
|
||||
and optional ``OUTBOUND_META_AGENT_UI`` blobs for rich clients; non-WebUI
|
||||
channels may ignore unknown keys.
|
||||
``event`` carries internal runtime/UI semantics. ``metadata`` is reserved
|
||||
for channel routing context (``message_id``, thread ids, etc.) and optional
|
||||
``OUTBOUND_META_AGENT_UI`` blobs for rich clients.
|
||||
"""
|
||||
|
||||
channel: str
|
||||
@@ -51,3 +54,4 @@ class OutboundMessage:
|
||||
media: list[str] = field(default_factory=list)
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
buttons: list[list[str]] = field(default_factory=list)
|
||||
event: "OutboundEvent | None" = None
|
||||
|
||||
@@ -0,0 +1,226 @@
|
||||
"""Typed outbound events carried by :class:`OutboundMessage`.
|
||||
|
||||
The message bus still transports :class:`nanobot.bus.events.OutboundMessage`
|
||||
because channels need chat routing fields. Runtime/UI semantics live on the
|
||||
message's explicit ``event`` field rather than in reserved metadata flags.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import Any
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
|
||||
|
||||
class OutboundEvent:
|
||||
"""Marker base for internal outbound runtime events."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ProgressEvent(OutboundEvent):
|
||||
content: str = ""
|
||||
tool_hint: bool = False
|
||||
reasoning: bool = False
|
||||
reasoning_delta: bool = False
|
||||
reasoning_end: bool = False
|
||||
stream_id: str | None = None
|
||||
tool_events: list[dict[str, Any]] | None = None
|
||||
file_edit_events: list[dict[str, Any]] | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RetryWaitEvent(OutboundEvent):
|
||||
content: str = ""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StreamDeltaEvent(OutboundEvent):
|
||||
content: str = ""
|
||||
stream_id: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StreamEndEvent(OutboundEvent):
|
||||
content: str = ""
|
||||
stream_id: str | None = None
|
||||
resuming: bool = False
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StreamedResponseEvent(OutboundEvent):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TurnEndEvent(OutboundEvent):
|
||||
latency_ms: int | None = None
|
||||
goal_state: dict[str, Any] | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GoalStatusEvent(OutboundEvent):
|
||||
status: str
|
||||
started_at: float | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GoalStateSyncEvent(OutboundEvent):
|
||||
goal_state: dict[str, Any]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SessionUpdatedEvent(OutboundEvent):
|
||||
scope: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RuntimeModelUpdatedEvent(OutboundEvent):
|
||||
model: str | None
|
||||
model_preset: str | None = None
|
||||
|
||||
|
||||
def outbound_message_for_event(
|
||||
*,
|
||||
channel: str,
|
||||
chat_id: str,
|
||||
event: OutboundEvent,
|
||||
content: str | None = None,
|
||||
metadata: Mapping[str, Any] | None = None,
|
||||
) -> OutboundMessage:
|
||||
"""Build an :class:`OutboundMessage` for a typed event."""
|
||||
|
||||
return OutboundMessage(
|
||||
channel=channel,
|
||||
chat_id=chat_id,
|
||||
content=_event_content(event) if content is None else content,
|
||||
event=event,
|
||||
metadata=dict(metadata or {}),
|
||||
)
|
||||
|
||||
|
||||
def outbound_event_from_message(msg: OutboundMessage) -> OutboundEvent | None:
|
||||
"""Return the typed outbound event carried by *msg*, if any."""
|
||||
|
||||
if msg.event is not None:
|
||||
return msg.event
|
||||
return _legacy_event_from_metadata(msg)
|
||||
|
||||
|
||||
def replace_outbound_event(
|
||||
msg: OutboundMessage,
|
||||
event: OutboundEvent,
|
||||
*,
|
||||
content: str | None = None,
|
||||
) -> OutboundMessage:
|
||||
"""Return *msg* with a new event and optional content."""
|
||||
|
||||
return replace(
|
||||
msg,
|
||||
content=_event_content(event) if content is None else content,
|
||||
event=event,
|
||||
)
|
||||
|
||||
|
||||
def _event_content(event: OutboundEvent) -> str:
|
||||
if isinstance(event, ProgressEvent | RetryWaitEvent | StreamDeltaEvent | StreamEndEvent):
|
||||
return event.content
|
||||
return ""
|
||||
|
||||
|
||||
def _legacy_event_from_metadata(msg: OutboundMessage) -> OutboundEvent | None:
|
||||
"""Bridge pre-typed outbound metadata flags into typed events.
|
||||
|
||||
New code should set ``OutboundMessage.event`` directly. The fallback keeps
|
||||
older in-process extensions and channel plugins from losing runtime events
|
||||
while they migrate off reserved metadata flags.
|
||||
"""
|
||||
|
||||
meta = msg.metadata or {}
|
||||
if meta.get("_runtime_model_updated"):
|
||||
return RuntimeModelUpdatedEvent(
|
||||
model=_metadata_str(meta, "model"),
|
||||
model_preset=_metadata_str(meta, "model_preset"),
|
||||
)
|
||||
if meta.get("_goal_state_sync"):
|
||||
goal_state = meta.get("goal_state")
|
||||
return GoalStateSyncEvent(goal_state if isinstance(goal_state, dict) else {"active": False})
|
||||
if meta.get("_goal_status"):
|
||||
status = meta.get("goal_status")
|
||||
if not isinstance(status, str) or not status:
|
||||
return None
|
||||
return GoalStatusEvent(
|
||||
status=status,
|
||||
started_at=_metadata_float(meta, "started_at", "goal_started_at"),
|
||||
)
|
||||
if meta.get("_turn_end"):
|
||||
goal_state = meta.get("goal_state")
|
||||
return TurnEndEvent(
|
||||
latency_ms=_metadata_int(meta, "latency_ms"),
|
||||
goal_state=goal_state if isinstance(goal_state, dict) else None,
|
||||
)
|
||||
if meta.get("_session_updated"):
|
||||
return SessionUpdatedEvent(scope=_metadata_str(meta, "_session_update_scope"))
|
||||
if meta.get("_retry_wait"):
|
||||
return RetryWaitEvent(content=msg.content)
|
||||
if meta.get("_stream_end"):
|
||||
return StreamEndEvent(
|
||||
content=msg.content,
|
||||
stream_id=_metadata_str(meta, "_stream_id"),
|
||||
resuming=bool(meta.get("_resuming")),
|
||||
)
|
||||
if meta.get("_stream_delta"):
|
||||
return StreamDeltaEvent(
|
||||
content=msg.content,
|
||||
stream_id=_metadata_str(meta, "_stream_id"),
|
||||
)
|
||||
if meta.get("_streamed"):
|
||||
return StreamedResponseEvent()
|
||||
if (
|
||||
meta.get("_progress")
|
||||
or meta.get("_reasoning_delta")
|
||||
or meta.get("_reasoning_end")
|
||||
or meta.get("_reasoning")
|
||||
or meta.get("_file_edit_events")
|
||||
or meta.get("_tool_events")
|
||||
):
|
||||
tool_events = meta.get("_tool_events")
|
||||
file_edit_events = meta.get("_file_edit_events")
|
||||
return ProgressEvent(
|
||||
content=msg.content,
|
||||
tool_hint=bool(meta.get("_tool_hint")),
|
||||
reasoning=bool(meta.get("_reasoning")),
|
||||
reasoning_delta=bool(meta.get("_reasoning_delta")),
|
||||
reasoning_end=bool(meta.get("_reasoning_end")),
|
||||
stream_id=_metadata_str(meta, "_stream_id"),
|
||||
tool_events=tool_events if isinstance(tool_events, list) else None,
|
||||
file_edit_events=file_edit_events if isinstance(file_edit_events, list) else None,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def _metadata_str(meta: Mapping[str, Any], key: str) -> str | None:
|
||||
value = meta.get(key)
|
||||
return value if isinstance(value, str) and value else None
|
||||
|
||||
|
||||
def _metadata_int(meta: Mapping[str, Any], key: str) -> int | None:
|
||||
value = meta.get(key)
|
||||
if isinstance(value, bool):
|
||||
return None
|
||||
if isinstance(value, int):
|
||||
return value
|
||||
if isinstance(value, float) and value.is_integer():
|
||||
return int(value)
|
||||
return None
|
||||
|
||||
|
||||
def _metadata_float(meta: Mapping[str, Any], *keys: str) -> float | None:
|
||||
for key in keys:
|
||||
value = meta.get(key)
|
||||
if isinstance(value, bool):
|
||||
continue
|
||||
if isinstance(value, int | float):
|
||||
return float(value)
|
||||
return None
|
||||
+12
-15
@@ -10,7 +10,8 @@ from __future__ import annotations
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent, outbound_message_for_event
|
||||
from nanobot.bus.queue import MessageBus
|
||||
|
||||
|
||||
@@ -29,23 +30,19 @@ def build_bus_progress_callback(
|
||||
reasoning: bool = False,
|
||||
reasoning_end: bool = False,
|
||||
) -> None:
|
||||
meta = dict(msg.metadata or {})
|
||||
meta["_progress"] = True
|
||||
meta["_tool_hint"] = tool_hint
|
||||
if reasoning:
|
||||
meta["_reasoning_delta"] = True
|
||||
if reasoning_end:
|
||||
meta["_reasoning_end"] = True
|
||||
if tool_events:
|
||||
meta["_tool_events"] = tool_events
|
||||
if file_edit_events:
|
||||
meta["_file_edit_events"] = file_edit_events
|
||||
await bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
outbound_message_for_event(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content=content,
|
||||
metadata=meta,
|
||||
event=ProgressEvent(
|
||||
content=content,
|
||||
tool_hint=tool_hint,
|
||||
reasoning_delta=reasoning,
|
||||
reasoning_end=reasoning_end,
|
||||
tool_events=tool_events,
|
||||
file_edit_events=file_edit_events,
|
||||
),
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
+37
-17
@@ -101,20 +101,33 @@ class BaseChannel(ABC):
|
||||
"""
|
||||
pass
|
||||
|
||||
async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None) -> None:
|
||||
async def send_delta(
|
||||
self,
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
stream_end: bool = False,
|
||||
resuming: bool = False,
|
||||
) -> None:
|
||||
"""Deliver a streaming text chunk.
|
||||
|
||||
Override in subclasses to enable streaming. Implementations should
|
||||
raise on delivery failure so the channel manager can retry.
|
||||
|
||||
Streaming contract: ``_stream_delta`` is a chunk, ``_stream_end`` ends
|
||||
the current segment, and stateful implementations must key buffers by
|
||||
``_stream_id`` rather than only by ``chat_id``.
|
||||
Stateful implementations should key buffers by ``stream_id`` rather
|
||||
than only by ``chat_id`` when it is provided.
|
||||
"""
|
||||
pass
|
||||
|
||||
async def send_reasoning_delta(
|
||||
self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None
|
||||
self,
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
) -> None:
|
||||
"""Stream a chunk of model reasoning/thinking content.
|
||||
|
||||
@@ -123,15 +136,17 @@ class BaseChannel(ABC):
|
||||
subtext, WebUI italic bubble, ...) override to render reasoning
|
||||
as a subordinate trace that updates in place as the model thinks.
|
||||
|
||||
Streaming contract mirrors :meth:`send_delta`: ``_reasoning_delta``
|
||||
is a chunk, ``_reasoning_end`` ends the current reasoning segment,
|
||||
and stateful implementations should key buffers by ``_stream_id``
|
||||
rather than only by ``chat_id``.
|
||||
Streaming contract mirrors :meth:`send_delta`: stateful implementations
|
||||
should key buffers by ``stream_id`` rather than only by ``chat_id``.
|
||||
"""
|
||||
return
|
||||
|
||||
async def send_reasoning_end(
|
||||
self, chat_id: str, metadata: dict[str, Any] | None = None
|
||||
self,
|
||||
chat_id: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
) -> None:
|
||||
"""Mark the end of a reasoning stream segment.
|
||||
|
||||
@@ -165,13 +180,18 @@ class BaseChannel(ABC):
|
||||
"""
|
||||
if not msg.content:
|
||||
return
|
||||
meta = dict(msg.metadata or {})
|
||||
meta.setdefault("_reasoning_delta", True)
|
||||
await self.send_reasoning_delta(msg.chat_id, msg.content, meta)
|
||||
end_meta = dict(meta)
|
||||
end_meta.pop("_reasoning_delta", None)
|
||||
end_meta["_reasoning_end"] = True
|
||||
await self.send_reasoning_end(msg.chat_id, end_meta)
|
||||
stream_id = getattr(msg.event, "stream_id", None)
|
||||
await self.send_reasoning_delta(
|
||||
msg.chat_id,
|
||||
msg.content,
|
||||
msg.metadata,
|
||||
stream_id=stream_id,
|
||||
)
|
||||
await self.send_reasoning_end(
|
||||
msg.chat_id,
|
||||
msg.metadata,
|
||||
stream_id=stream_id,
|
||||
)
|
||||
|
||||
@property
|
||||
def supports_streaming(self) -> bool:
|
||||
|
||||
@@ -13,6 +13,7 @@ from typing import TYPE_CHECKING, Any, Literal
|
||||
from pydantic import Field
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.command.builtin import build_help_text
|
||||
@@ -217,6 +218,16 @@ if DISCORD_AVAILABLE:
|
||||
command_text = f"/model {preset}" if preset else "/model"
|
||||
await self._forward_slash_command(interaction, command_text)
|
||||
|
||||
@self.tree.command(name="trigger", description="Create a named local trigger for this chat")
|
||||
@app_commands.describe(name="Trigger name")
|
||||
async def trigger_command(
|
||||
interaction: discord.Interaction,
|
||||
name: str,
|
||||
) -> None:
|
||||
name = name.strip()
|
||||
command_text = f"/trigger {name}" if name else "/trigger"
|
||||
await self._forward_slash_command(interaction, command_text)
|
||||
|
||||
@self.tree.command(name="help", description="Show available commands")
|
||||
async def help_command(interaction: discord.Interaction) -> None:
|
||||
sender_id = str(interaction.user.id)
|
||||
@@ -458,7 +469,7 @@ class DiscordChannel(BaseChannel):
|
||||
self.logger.warning("client not ready; dropping outbound message")
|
||||
return
|
||||
|
||||
is_progress = bool((msg.metadata or {}).get("_progress"))
|
||||
is_progress = isinstance(msg.event, ProgressEvent)
|
||||
|
||||
try:
|
||||
await client.send_outbound(msg)
|
||||
@@ -471,7 +482,14 @@ class DiscordChannel(BaseChannel):
|
||||
await self._clear_reactions(msg.chat_id)
|
||||
|
||||
async def send_delta(
|
||||
self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None
|
||||
self,
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
stream_end: bool = False,
|
||||
resuming: bool = False,
|
||||
) -> None:
|
||||
"""Progressive Discord delivery: send once, then edit until the stream ends."""
|
||||
client = self._client
|
||||
@@ -479,10 +497,7 @@ class DiscordChannel(BaseChannel):
|
||||
self.logger.warning("client not ready; dropping stream delta")
|
||||
return
|
||||
|
||||
meta = metadata or {}
|
||||
stream_id = meta.get("_stream_id")
|
||||
|
||||
if meta.get("_stream_end"):
|
||||
if stream_end:
|
||||
buf = self._stream_bufs.get(chat_id)
|
||||
if not buf or buf.message is None or not buf.text:
|
||||
return
|
||||
|
||||
@@ -23,6 +23,7 @@ from loguru import logger
|
||||
from pydantic import Field
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.config.paths import get_media_dir
|
||||
@@ -218,7 +219,7 @@ class EmailChannel(BaseChannel):
|
||||
return
|
||||
|
||||
# Skip progress messages to prevent sending an empty email after each tool call
|
||||
if (msg.metadata or {}).get("_progress"):
|
||||
if isinstance(msg.event, ProgressEvent):
|
||||
self.logger.debug("Skip progress message to {}", msg.chat_id)
|
||||
return
|
||||
|
||||
|
||||
@@ -22,6 +22,7 @@ from rich.panel import Panel
|
||||
from rich.text import Text
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.config.paths import get_media_dir
|
||||
@@ -1797,14 +1798,19 @@ class FeishuChannel(BaseChannel):
|
||||
return self._stream_update_text_sync(card_id, content, sequence), sequence
|
||||
|
||||
async def send_delta(
|
||||
self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None
|
||||
self,
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
stream_end: bool = False,
|
||||
resuming: bool = False,
|
||||
) -> None:
|
||||
"""Progressive streaming via CardKit: create card on first delta, stream-update on subsequent.
|
||||
|
||||
Supported metadata keys:
|
||||
_stream_end: Finalize the streaming card.
|
||||
_tool_hint: Delta is a formatted tool hint (for display only).
|
||||
message_id: Original message id (used with _stream_end for reaction cleanup).
|
||||
message_id: Original message id (used with stream end for reaction cleanup).
|
||||
chat_type: "group" or "p2p" — controls reply-in-thread for streaming cards.
|
||||
"""
|
||||
if not self._client:
|
||||
@@ -1815,14 +1821,14 @@ class FeishuChannel(BaseChannel):
|
||||
rid_type = "chat_id" if chat_id.startswith("oc_") else "open_id"
|
||||
|
||||
# --- stream end: final update or fallback ---
|
||||
if meta.get("_stream_end"):
|
||||
if stream_end:
|
||||
message_id = meta.get("message_id")
|
||||
# Only finalize the OnIt -> DONE reaction transition on the truly
|
||||
# final stream end. _resuming=True means the agent will keep
|
||||
# final stream end. resuming=True means the agent will keep
|
||||
# working (more tool-call rounds), so leave the reaction state
|
||||
# in place — otherwise the OnIt indicator disappears prematurely
|
||||
# and the DONE reaction fires after every tool call.
|
||||
if message_id and not meta.get("_resuming"):
|
||||
if message_id and not resuming:
|
||||
reaction_id = self._reaction_ids.pop(message_id, None)
|
||||
if reaction_id:
|
||||
await self._remove_reaction(message_id, reaction_id)
|
||||
@@ -1965,7 +1971,9 @@ class FeishuChannel(BaseChannel):
|
||||
# Handle tool hint messages. When a streaming card is active for
|
||||
# this chat, inline the hint into the card instead of sending a
|
||||
# separate message so the user experience stays cohesive.
|
||||
if msg.metadata.get("_tool_hint"):
|
||||
progress_event = msg.event if isinstance(msg.event, ProgressEvent) else None
|
||||
|
||||
if progress_event and progress_event.tool_hint:
|
||||
hint = (msg.content or "").strip()
|
||||
if not hint:
|
||||
return
|
||||
@@ -1976,6 +1984,7 @@ class FeishuChannel(BaseChannel):
|
||||
await self.send_delta(
|
||||
msg.chat_id,
|
||||
"\n\n" + self._format_tool_hint_delta(hint) + "\n\n",
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
return
|
||||
# No active streaming card — send as a regular interactive card
|
||||
@@ -2009,7 +2018,7 @@ class FeishuChannel(BaseChannel):
|
||||
reply_message_id: str | None = None
|
||||
_msg_id = msg.metadata.get("message_id")
|
||||
has_thread_id = msg.metadata.get("thread_id")
|
||||
if self.config.reply_to_message and not msg.metadata.get("_progress", False):
|
||||
if self.config.reply_to_message and progress_event is None:
|
||||
reply_message_id = _msg_id
|
||||
# For topic group messages, always reply to keep context in thread
|
||||
elif has_thread_id:
|
||||
|
||||
+159
-50
@@ -4,6 +4,7 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import inspect
|
||||
from collections.abc import Callable
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
@@ -12,6 +13,16 @@ from typing import TYPE_CHECKING, Any
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import (
|
||||
ProgressEvent,
|
||||
RetryWaitEvent,
|
||||
RuntimeModelUpdatedEvent,
|
||||
StreamDeltaEvent,
|
||||
StreamedResponseEvent,
|
||||
StreamEndEvent,
|
||||
outbound_event_from_message,
|
||||
replace_outbound_event,
|
||||
)
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.config.schema import Config
|
||||
@@ -57,8 +68,10 @@ class ChannelManager:
|
||||
*,
|
||||
session_manager: "SessionManager | None" = None,
|
||||
cron_service: Any | None = None,
|
||||
local_trigger_store: Any | None = None,
|
||||
webui_runtime_model_name: Callable[[], str | None] | None = None,
|
||||
webui_cron_pending_job_ids: Callable[[str], set[str]] | None = None,
|
||||
webui_local_trigger_pending_ids: Callable[[str], set[str]] | None = None,
|
||||
webui_static_dist: bool = True,
|
||||
webui_runtime_surface: str = "browser",
|
||||
webui_runtime_capabilities: dict[str, Any] | None = None,
|
||||
@@ -67,8 +80,10 @@ class ChannelManager:
|
||||
self.bus = bus
|
||||
self._session_manager = session_manager
|
||||
self._cron_service = cron_service
|
||||
self._local_trigger_store = local_trigger_store
|
||||
self._webui_runtime_model_name = webui_runtime_model_name
|
||||
self._webui_cron_pending_job_ids = webui_cron_pending_job_ids
|
||||
self._webui_local_trigger_pending_ids = webui_local_trigger_pending_ids
|
||||
self._webui_static_dist = webui_static_dist
|
||||
self._webui_runtime_surface = webui_runtime_surface
|
||||
self._webui_runtime_capabilities = dict(webui_runtime_capabilities or {})
|
||||
@@ -128,7 +143,9 @@ class ChannelManager:
|
||||
runtime_surface=self._webui_runtime_surface,
|
||||
runtime_capabilities_overrides=self._webui_runtime_capabilities,
|
||||
cron_service=self._cron_service,
|
||||
local_trigger_store=self._local_trigger_store,
|
||||
cron_pending_job_ids=self._webui_cron_pending_job_ids,
|
||||
local_trigger_pending_ids=self._webui_local_trigger_pending_ids,
|
||||
logger=logger,
|
||||
)
|
||||
kwargs["gateway"] = gateway
|
||||
@@ -266,7 +283,7 @@ class ChannelManager:
|
||||
|
||||
def _should_suppress_outbound(self, msg: OutboundMessage) -> bool:
|
||||
metadata = msg.metadata or {}
|
||||
if metadata.get("_progress"):
|
||||
if isinstance(outbound_event_from_message(msg), ProgressEvent):
|
||||
return False
|
||||
fingerprint = self._fingerprint_content(msg.content)
|
||||
if not fingerprint:
|
||||
@@ -305,57 +322,59 @@ class ChannelManager:
|
||||
timeout=1.0
|
||||
)
|
||||
|
||||
if (
|
||||
msg.metadata.get("_reasoning_delta")
|
||||
or msg.metadata.get("_reasoning_end")
|
||||
or msg.metadata.get("_reasoning")
|
||||
event = outbound_event_from_message(msg)
|
||||
progress_event = event if isinstance(event, ProgressEvent) else None
|
||||
if progress_event and (
|
||||
progress_event.reasoning_delta
|
||||
or progress_event.reasoning_end
|
||||
or progress_event.reasoning
|
||||
):
|
||||
# Reasoning rides its own plugin channel: only delivered
|
||||
# when the destination channel opts in via ``show_reasoning``
|
||||
# and overrides the streaming primitives. Channels without
|
||||
# a low-emphasis UI affordance keep the base no-op and the
|
||||
# content silently drops here. ``_reasoning`` (one-shot)
|
||||
# is accepted for backward compatibility with hooks that
|
||||
# haven't migrated to delta/end yet.
|
||||
# content silently drops here.
|
||||
channel = self.channels.get(msg.channel)
|
||||
if channel is not None and channel.show_reasoning:
|
||||
await self._send_with_retry(channel, msg)
|
||||
continue
|
||||
|
||||
if msg.metadata.get("_progress"):
|
||||
if msg.metadata.get("_tool_hint") and not self._should_send_progress(
|
||||
if progress_event:
|
||||
if progress_event.tool_hint and not self._should_send_progress(
|
||||
msg.channel, tool_hint=True,
|
||||
):
|
||||
continue
|
||||
if not msg.metadata.get("_tool_hint") and not self._should_send_progress(
|
||||
if not progress_event.tool_hint and not self._should_send_progress(
|
||||
msg.channel, tool_hint=False,
|
||||
):
|
||||
continue
|
||||
|
||||
if msg.metadata.get("_retry_wait"):
|
||||
if isinstance(event, RetryWaitEvent):
|
||||
continue
|
||||
|
||||
if (
|
||||
msg.metadata.get("_runtime_model_updated")
|
||||
isinstance(event, RuntimeModelUpdatedEvent)
|
||||
and msg.channel == "websocket"
|
||||
and "websocket" not in self.channels
|
||||
):
|
||||
continue
|
||||
|
||||
# Coalesce consecutive _stream_delta messages for the same (channel, chat_id)
|
||||
# Coalesce consecutive stream delta messages for the same (channel, chat_id)
|
||||
# to reduce API calls and improve streaming latency
|
||||
if msg.metadata.get("_stream_delta") and not msg.metadata.get("_stream_end"):
|
||||
if isinstance(event, StreamDeltaEvent):
|
||||
msg, extra_pending = self._coalesce_stream_deltas(msg)
|
||||
pending.extend(extra_pending)
|
||||
event = outbound_event_from_message(msg)
|
||||
|
||||
channel = self.channels.get(msg.channel)
|
||||
if channel:
|
||||
# Duplicate suppression is scoped to a known source message
|
||||
# so repeated content from separate turns is still delivered.
|
||||
if (
|
||||
not msg.metadata.get("_stream_delta")
|
||||
and not msg.metadata.get("_stream_end")
|
||||
and not msg.metadata.get("_streamed")
|
||||
not isinstance(
|
||||
event,
|
||||
StreamDeltaEvent | StreamEndEvent | StreamedResponseEvent,
|
||||
)
|
||||
):
|
||||
if self._should_suppress_outbound(msg):
|
||||
logger.info("Suppressing duplicate outbound message to {}:{}", msg.channel, msg.chat_id)
|
||||
@@ -369,34 +388,116 @@ class ChannelManager:
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
|
||||
@staticmethod
|
||||
def _accepts_keyword(callable_obj: Callable[..., Any], name: str) -> bool:
|
||||
try:
|
||||
signature = inspect.signature(callable_obj)
|
||||
except (TypeError, ValueError):
|
||||
return True
|
||||
return any(
|
||||
parameter.kind is inspect.Parameter.VAR_KEYWORD or parameter.name == name
|
||||
for parameter in signature.parameters.values()
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def _send_reasoning_delta(cls, channel: BaseChannel, msg: OutboundMessage, event: ProgressEvent) -> None:
|
||||
metadata = msg.metadata
|
||||
kwargs: dict[str, Any] = {}
|
||||
if cls._accepts_keyword(channel.send_reasoning_delta, "stream_id"):
|
||||
kwargs["stream_id"] = event.stream_id
|
||||
else:
|
||||
metadata = dict(metadata or {})
|
||||
metadata["_reasoning_delta"] = True
|
||||
if event.stream_id is not None:
|
||||
metadata["_stream_id"] = event.stream_id
|
||||
await channel.send_reasoning_delta(
|
||||
msg.chat_id,
|
||||
msg.content,
|
||||
metadata,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def _send_reasoning_end(cls, channel: BaseChannel, msg: OutboundMessage, event: ProgressEvent) -> None:
|
||||
metadata = msg.metadata
|
||||
kwargs: dict[str, Any] = {}
|
||||
if cls._accepts_keyword(channel.send_reasoning_end, "stream_id"):
|
||||
kwargs["stream_id"] = event.stream_id
|
||||
else:
|
||||
metadata = dict(metadata or {})
|
||||
metadata["_reasoning_end"] = True
|
||||
if event.stream_id is not None:
|
||||
metadata["_stream_id"] = event.stream_id
|
||||
await channel.send_reasoning_end(
|
||||
msg.chat_id,
|
||||
metadata,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def _send_stream_event(
|
||||
cls,
|
||||
channel: BaseChannel,
|
||||
msg: OutboundMessage,
|
||||
event: StreamDeltaEvent | StreamEndEvent,
|
||||
) -> None:
|
||||
metadata = msg.metadata
|
||||
kwargs: dict[str, Any] = {}
|
||||
if cls._accepts_keyword(channel.send_delta, "stream_id"):
|
||||
kwargs["stream_id"] = event.stream_id
|
||||
else:
|
||||
metadata = dict(metadata or {})
|
||||
if event.stream_id is not None:
|
||||
metadata["_stream_id"] = event.stream_id
|
||||
|
||||
if isinstance(event, StreamEndEvent):
|
||||
if cls._accepts_keyword(channel.send_delta, "stream_end"):
|
||||
kwargs["stream_end"] = True
|
||||
else:
|
||||
metadata = dict(metadata or {})
|
||||
metadata["_stream_end"] = True
|
||||
if cls._accepts_keyword(channel.send_delta, "resuming"):
|
||||
kwargs["resuming"] = event.resuming
|
||||
elif not kwargs:
|
||||
metadata = dict(metadata or {})
|
||||
metadata["_stream_delta"] = True
|
||||
|
||||
await channel.send_delta(
|
||||
msg.chat_id,
|
||||
msg.content,
|
||||
metadata,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _send_once(channel: BaseChannel, msg: OutboundMessage) -> None:
|
||||
"""Send one outbound message without retry policy."""
|
||||
if msg.metadata.get("_reasoning_end"):
|
||||
await channel.send_reasoning_end(msg.chat_id, msg.metadata)
|
||||
elif msg.metadata.get("_reasoning_delta"):
|
||||
await channel.send_reasoning_delta(msg.chat_id, msg.content, msg.metadata)
|
||||
elif msg.metadata.get("_reasoning"):
|
||||
# Back-compat: one-shot reasoning. BaseChannel translates this
|
||||
# to a single delta + end pair so plugins only implement the
|
||||
# streaming primitives.
|
||||
event = outbound_event_from_message(msg)
|
||||
if isinstance(event, ProgressEvent) and event.reasoning_end:
|
||||
await ChannelManager._send_reasoning_end(channel, msg, event)
|
||||
elif isinstance(event, ProgressEvent) and event.reasoning_delta:
|
||||
await ChannelManager._send_reasoning_delta(channel, msg, event)
|
||||
elif isinstance(event, ProgressEvent) and event.reasoning:
|
||||
# BaseChannel translates one-shot reasoning to a single delta +
|
||||
# end pair so plugins only implement the streaming primitives.
|
||||
await channel.send_reasoning(msg)
|
||||
elif msg.metadata.get("_file_edit_events"):
|
||||
edits = msg.metadata.get("_file_edit_events")
|
||||
elif isinstance(event, ProgressEvent) and event.file_edit_events:
|
||||
await channel.send_file_edit_events(
|
||||
msg.chat_id,
|
||||
edits if isinstance(edits, list) else [],
|
||||
event.file_edit_events,
|
||||
msg.metadata,
|
||||
)
|
||||
elif msg.metadata.get("_stream_delta") or msg.metadata.get("_stream_end"):
|
||||
await channel.send_delta(msg.chat_id, msg.content, msg.metadata)
|
||||
elif not msg.metadata.get("_streamed"):
|
||||
elif isinstance(event, StreamDeltaEvent):
|
||||
await ChannelManager._send_stream_event(channel, msg, event)
|
||||
elif isinstance(event, StreamEndEvent):
|
||||
await ChannelManager._send_stream_event(channel, msg, event)
|
||||
elif not isinstance(event, StreamedResponseEvent):
|
||||
await channel.send(msg)
|
||||
|
||||
def _coalesce_stream_deltas(
|
||||
self, first_msg: OutboundMessage
|
||||
) -> tuple[OutboundMessage, list[OutboundMessage]]:
|
||||
"""Merge consecutive _stream_delta messages for the same (channel, chat_id, _stream_id).
|
||||
"""Merge consecutive stream deltas for the same (channel, chat_id, stream_id).
|
||||
|
||||
This reduces the number of API calls when the queue has accumulated multiple
|
||||
deltas, which happens when LLM generates faster than the channel can process.
|
||||
@@ -404,10 +505,15 @@ class ChannelManager:
|
||||
Returns:
|
||||
tuple of (merged_message, list_of_non_matching_messages)
|
||||
"""
|
||||
first_metadata = first_msg.metadata or {}
|
||||
target_key = (first_msg.channel, first_msg.chat_id, first_metadata.get("_stream_id"))
|
||||
first_event = outbound_event_from_message(first_msg)
|
||||
first_stream_id = first_event.stream_id if isinstance(first_event, StreamDeltaEvent) else None
|
||||
target_key = (first_msg.channel, first_msg.chat_id, first_stream_id)
|
||||
combined_content = first_msg.content
|
||||
final_metadata = dict(first_msg.metadata or {})
|
||||
final_event: StreamDeltaEvent | StreamEndEvent = (
|
||||
first_event
|
||||
if isinstance(first_event, StreamDeltaEvent)
|
||||
else StreamDeltaEvent(stream_id=first_stream_id)
|
||||
)
|
||||
non_matching: list[OutboundMessage] = []
|
||||
|
||||
# Only merge consecutive deltas. As soon as we hit any other message,
|
||||
@@ -419,21 +525,29 @@ class ChannelManager:
|
||||
break
|
||||
|
||||
# Check if this message belongs to the same stream
|
||||
next_metadata = next_msg.metadata or {}
|
||||
next_event = outbound_event_from_message(next_msg)
|
||||
next_stream_id = (
|
||||
next_event.stream_id
|
||||
if isinstance(next_event, StreamDeltaEvent | StreamEndEvent)
|
||||
else None
|
||||
)
|
||||
same_target = (
|
||||
next_msg.channel,
|
||||
next_msg.chat_id,
|
||||
next_metadata.get("_stream_id"),
|
||||
next_stream_id,
|
||||
) == target_key
|
||||
is_delta = next_metadata.get("_stream_delta")
|
||||
is_end = next_metadata.get("_stream_end")
|
||||
is_delta = isinstance(next_event, StreamDeltaEvent)
|
||||
is_end = isinstance(next_event, StreamEndEvent)
|
||||
|
||||
if same_target and is_delta and not final_metadata.get("_stream_end"):
|
||||
if same_target and (is_delta or (is_end and next_msg.content)):
|
||||
# Accumulate content
|
||||
combined_content += next_msg.content
|
||||
# If we see _stream_end, remember it and stop coalescing this stream
|
||||
if is_end:
|
||||
final_metadata["_stream_end"] = True
|
||||
# If we see stream_end, remember it and stop coalescing this stream
|
||||
if isinstance(next_event, StreamEndEvent):
|
||||
final_event = StreamEndEvent(
|
||||
stream_id=next_stream_id,
|
||||
resuming=next_event.resuming,
|
||||
)
|
||||
# Stream ended - stop coalescing this stream
|
||||
break
|
||||
else:
|
||||
@@ -441,12 +555,7 @@ class ChannelManager:
|
||||
non_matching.append(next_msg)
|
||||
break
|
||||
|
||||
merged = OutboundMessage(
|
||||
channel=first_msg.channel,
|
||||
chat_id=first_msg.chat_id,
|
||||
content=combined_content,
|
||||
metadata=final_metadata,
|
||||
)
|
||||
merged = replace_outbound_event(first_msg, final_event, content=combined_content)
|
||||
return merged, non_matching
|
||||
|
||||
async def _send_with_retry(self, channel: BaseChannel, msg: OutboundMessage) -> None:
|
||||
|
||||
@@ -49,6 +49,7 @@ except ImportError as e:
|
||||
) from e
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.config.paths import get_data_dir, get_media_dir
|
||||
@@ -504,7 +505,7 @@ class MatrixChannel(BaseChannel):
|
||||
text = msg.content or ""
|
||||
candidates = self._collect_outbound_media_candidates(msg.media)
|
||||
relates_to = self._build_thread_relates_to(msg.metadata)
|
||||
is_progress = bool((msg.metadata or {}).get("_progress"))
|
||||
is_progress = isinstance(msg.event, ProgressEvent)
|
||||
try:
|
||||
failures: list[str] = []
|
||||
if candidates:
|
||||
@@ -528,11 +529,19 @@ class MatrixChannel(BaseChannel):
|
||||
if not is_progress:
|
||||
await self._stop_typing_keepalive(msg.chat_id, clear_typing=True)
|
||||
|
||||
async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None) -> None:
|
||||
meta = metadata or {}
|
||||
async def send_delta(
|
||||
self,
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
stream_end: bool = False,
|
||||
resuming: bool = False,
|
||||
) -> None:
|
||||
relates_to = self._build_thread_relates_to(metadata)
|
||||
|
||||
if meta.get("_stream_end"):
|
||||
if stream_end:
|
||||
buf = self._stream_bufs.pop(chat_id, None)
|
||||
if not buf or not buf.event_id or not buf.text:
|
||||
return
|
||||
|
||||
@@ -18,6 +18,7 @@ import httpx
|
||||
from pydantic import Field, computed_field, field_validator
|
||||
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.config.paths import get_media_dir
|
||||
@@ -539,7 +540,7 @@ class SignalChannel(BaseChannel):
|
||||
|
||||
async def send(self, msg: OutboundMessage) -> None:
|
||||
"""Send a message through Signal."""
|
||||
is_progress_message = bool(msg.metadata.get("_progress"))
|
||||
is_progress_message = isinstance(msg.event, ProgressEvent)
|
||||
try:
|
||||
plain_text, text_styles = _markdown_to_signal(msg.content)
|
||||
if not plain_text and not msg.media:
|
||||
|
||||
@@ -14,6 +14,7 @@ from slack_sdk.web.async_client import AsyncWebClient
|
||||
from slackify_markdown import slackify_markdown
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.config.paths import get_media_dir
|
||||
@@ -164,7 +165,7 @@ class SlackChannel(BaseChannel):
|
||||
# only makes sense within the originating conversation.
|
||||
thread_ts_param = thread_ts if thread_ts and target_chat_id == origin_chat_id else None
|
||||
|
||||
is_progress = (msg.metadata or {}).get("_progress", False)
|
||||
is_progress = isinstance(msg.event, ProgressEvent)
|
||||
if is_progress and not msg.content:
|
||||
pass # skip empty progress messages (e.g. tool-event-only updates)
|
||||
elif msg.content or not (msg.media or []):
|
||||
@@ -190,7 +191,7 @@ class SlackChannel(BaseChannel):
|
||||
self.logger.exception("Failed to upload file {}", media_path)
|
||||
|
||||
# Update reaction emoji when the final (non-progress) response is sent
|
||||
if not (msg.metadata or {}).get("_progress"):
|
||||
if not is_progress:
|
||||
event = slack_meta.get("event", {})
|
||||
await self._update_react_emoji(origin_chat_id, event.get("ts"))
|
||||
|
||||
|
||||
@@ -26,6 +26,7 @@ from telegram.ext import Application, CallbackQueryHandler, ContextTypes, Messag
|
||||
from telegram.request import HTTPXRequest
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.command.builtin import build_help_text
|
||||
@@ -36,7 +37,7 @@ from nanobot.utils.helpers import split_message
|
||||
|
||||
TELEGRAM_MAX_MESSAGE_LEN = 4000 # Telegram message character limit
|
||||
# Telegram's actual API limit is 4096; we split raw markdown at 4000 as a
|
||||
# safety margin for mid-stream edits (plain text). For _stream_end, we split
|
||||
# safety margin for mid-stream edits (plain text). On stream end, we split
|
||||
# raw markdown into chunks whose rendered HTML fits Telegram's true 4096-char
|
||||
# boundary so the final rendered message never overflows.
|
||||
TELEGRAM_HTML_MAX_LEN = 4096
|
||||
@@ -410,6 +411,7 @@ class TelegramChannel(BaseChannel):
|
||||
BotCommand("status", "Show bot status"),
|
||||
BotCommand("history", "Show recent conversation messages"),
|
||||
BotCommand("goal", "Start a sustained objective (long-running task)"),
|
||||
BotCommand("trigger", "Create a named local trigger"),
|
||||
BotCommand("pairing", "Manage DM pairing (approve/deny/list)"),
|
||||
BotCommand("model", "Switch runtime model preset"),
|
||||
BotCommand("skill", "List enabled skills"),
|
||||
@@ -422,7 +424,7 @@ class TelegramChannel(BaseChannel):
|
||||
# Regex for slash commands routed to AgentLoop via ``_forward_command``.
|
||||
# Hyphenated ``dream-*`` commands stay on a separate handler (below).
|
||||
TELEGRAM_BUS_SLASH_COMMAND_RE = re.compile(
|
||||
r"^/(?:new|stop|restart|status|dream|history|goal|pairing|model|skill)(?:@\w+)?(?:\s+.*)?$"
|
||||
r"^/(?:new|stop|restart|status|dream|history|goal|trigger|pairing|model|skill)(?:@\w+)?(?:\s+.*)?$"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -706,8 +708,10 @@ class TelegramChannel(BaseChannel):
|
||||
self.logger.warning("bot not running")
|
||||
return
|
||||
|
||||
progress_event = msg.event if isinstance(msg.event, ProgressEvent) else None
|
||||
|
||||
# Only stop typing indicator and remove reaction for final responses
|
||||
if not msg.metadata.get("_progress", False):
|
||||
if progress_event is None:
|
||||
self._stop_typing(msg.chat_id)
|
||||
if reply_to_message_id := msg.metadata.get("message_id"):
|
||||
with suppress(ValueError):
|
||||
@@ -792,7 +796,7 @@ class TelegramChannel(BaseChannel):
|
||||
|
||||
# Send text content
|
||||
if msg.content and msg.content != "[empty message]":
|
||||
render_as_blockquote = bool(msg.metadata.get("_tool_hint"))
|
||||
render_as_blockquote = bool(progress_event and progress_event.tool_hint)
|
||||
buttons = getattr(msg, "buttons", None) or []
|
||||
reply_markup = self._build_keyboard(buttons) if buttons else None
|
||||
text = msg.content
|
||||
@@ -887,15 +891,23 @@ class TelegramChannel(BaseChannel):
|
||||
def _is_not_modified_error(exc: Exception) -> bool:
|
||||
return isinstance(exc, BadRequest) and "message is not modified" in str(exc).lower()
|
||||
|
||||
async def send_delta(self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None) -> None:
|
||||
async def send_delta(
|
||||
self,
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
stream_end: bool = False,
|
||||
resuming: bool = False,
|
||||
) -> None:
|
||||
"""Progressive message editing: send on first delta, edit on subsequent ones."""
|
||||
if not self._app:
|
||||
return
|
||||
meta = metadata or {}
|
||||
int_chat_id = int(chat_id)
|
||||
stream_id = meta.get("_stream_id")
|
||||
|
||||
if meta.get("_stream_end"):
|
||||
if stream_end:
|
||||
buf = self._stream_bufs.get(chat_id)
|
||||
if not buf or not buf.message_id or not buf.text:
|
||||
return
|
||||
|
||||
@@ -19,6 +19,16 @@ 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.outbound_events import (
|
||||
GoalStateSyncEvent,
|
||||
GoalStatusEvent,
|
||||
ProgressEvent,
|
||||
RuntimeModelUpdatedEvent,
|
||||
SessionUpdatedEvent,
|
||||
TurnEndEvent,
|
||||
outbound_event_from_message,
|
||||
outbound_message_for_event,
|
||||
)
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.config.paths import get_media_dir
|
||||
@@ -148,16 +158,13 @@ def publish_runtime_model_update(
|
||||
model_preset: str | None,
|
||||
) -> None:
|
||||
"""Enqueue a runtime model snapshot for websocket subscribers (fan-out in-channel)."""
|
||||
bus.outbound.put_nowait(OutboundMessage(
|
||||
channel="websocket",
|
||||
chat_id="*",
|
||||
content="",
|
||||
metadata={
|
||||
"_runtime_model_updated": True,
|
||||
"model": model,
|
||||
"model_preset": model_preset,
|
||||
},
|
||||
))
|
||||
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:
|
||||
@@ -851,70 +858,63 @@ class WebSocketChannel(BaseChannel):
|
||||
raise
|
||||
|
||||
async def send(self, msg: OutboundMessage) -> None:
|
||||
if msg.metadata.get("_runtime_model_updated"):
|
||||
event = outbound_event_from_message(msg)
|
||||
progress_event = event if isinstance(event, ProgressEvent) else None
|
||||
if isinstance(event, RuntimeModelUpdatedEvent):
|
||||
await self.send_runtime_model_updated(
|
||||
model_name=msg.metadata.get("model"),
|
||||
model_preset=msg.metadata.get("model_preset"),
|
||||
model_name=event.model,
|
||||
model_preset=event.model_preset,
|
||||
)
|
||||
return
|
||||
|
||||
# Snapshot the subscriber set so ConnectionClosed cleanups mid-iteration are safe.
|
||||
conns = list(self._subs.get(msg.chat_id, ()))
|
||||
if not conns:
|
||||
if (
|
||||
msg.metadata.get("_progress")
|
||||
or msg.metadata.get("_file_edit_events")
|
||||
or msg.metadata.get("_turn_end")
|
||||
or msg.metadata.get("_session_updated")
|
||||
or msg.metadata.get("_goal_status")
|
||||
or msg.metadata.get("_goal_state_sync")
|
||||
if isinstance(
|
||||
event,
|
||||
ProgressEvent
|
||||
| TurnEndEvent
|
||||
| SessionUpdatedEvent
|
||||
| GoalStatusEvent
|
||||
| GoalStateSyncEvent,
|
||||
):
|
||||
self.logger.debug("no active subscribers for chat_id={}", msg.chat_id)
|
||||
else:
|
||||
self.logger.warning("no active subscribers for chat_id={}", msg.chat_id)
|
||||
if msg.metadata.get("_goal_state_sync"):
|
||||
if isinstance(event, GoalStateSyncEvent):
|
||||
if conns:
|
||||
blob = msg.metadata.get("goal_state")
|
||||
await self.send_goal_state(msg.chat_id, blob if isinstance(blob, dict) else {"active": False})
|
||||
await self.send_goal_state(msg.chat_id, event.goal_state or {"active": False})
|
||||
return
|
||||
if msg.metadata.get("_goal_status"):
|
||||
if isinstance(event, GoalStatusEvent):
|
||||
if conns:
|
||||
status = msg.metadata.get("goal_status")
|
||||
if status in ("running", "idle"):
|
||||
started_raw = msg.metadata.get("started_at", msg.metadata.get("goal_started_at"))
|
||||
if event.status in ("running", "idle"):
|
||||
await self.send_goal_status(
|
||||
msg.chat_id,
|
||||
status,
|
||||
started_at=float(started_raw) if isinstance(started_raw, int | float) else None,
|
||||
event.status,
|
||||
started_at=event.started_at,
|
||||
)
|
||||
return
|
||||
# Signal that the agent has fully finished processing the current turn.
|
||||
if msg.metadata.get("_turn_end"):
|
||||
lat = msg.metadata.get("latency_ms")
|
||||
lat_i = int(lat) if isinstance(lat, (int, float)) else None
|
||||
gs = msg.metadata.get("goal_state")
|
||||
gs_blob = gs if isinstance(gs, dict) else None
|
||||
if isinstance(event, TurnEndEvent):
|
||||
await self.send_turn_end(
|
||||
msg.chat_id,
|
||||
latency_ms=lat_i,
|
||||
goal_state=gs_blob,
|
||||
latency_ms=event.latency_ms,
|
||||
goal_state=event.goal_state,
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
await self.send_session_updated(msg.chat_id, scope="thread")
|
||||
return
|
||||
if msg.metadata.get("_session_updated"):
|
||||
if isinstance(event, SessionUpdatedEvent):
|
||||
if conns:
|
||||
scope = msg.metadata.get("_session_update_scope")
|
||||
await self.send_session_updated(
|
||||
msg.chat_id,
|
||||
scope=scope if isinstance(scope, str) else None,
|
||||
scope=event.scope,
|
||||
)
|
||||
return
|
||||
if msg.metadata.get("_file_edit_events"):
|
||||
edits = msg.metadata.get("_file_edit_events")
|
||||
if progress_event and progress_event.file_edit_events:
|
||||
await self.send_file_edit_events(
|
||||
msg.chat_id,
|
||||
edits if isinstance(edits, list) else [],
|
||||
progress_event.file_edit_events,
|
||||
msg.metadata,
|
||||
)
|
||||
return
|
||||
@@ -939,17 +939,17 @@ class WebSocketChannel(BaseChannel):
|
||||
lat = msg.metadata.get("latency_ms")
|
||||
if isinstance(lat, (int, float)):
|
||||
payload["latency_ms"] = int(lat)
|
||||
if msg.metadata.get("_tool_events"):
|
||||
payload["tool_events"] = msg.metadata["_tool_events"]
|
||||
if progress_event and progress_event.tool_events:
|
||||
payload["tool_events"] = progress_event.tool_events
|
||||
agent_ui = msg.metadata.get(OUTBOUND_META_AGENT_UI)
|
||||
if agent_ui is not None:
|
||||
payload["agent_ui"] = agent_ui
|
||||
# Mark intermediate agent breadcrumbs (tool-call hints, generic
|
||||
# progress strings) so WS clients can render them as subordinate
|
||||
# trace rows rather than conversational replies.
|
||||
if msg.metadata.get("_tool_hint"):
|
||||
if progress_event and progress_event.tool_hint:
|
||||
payload["kind"] = "tool_hint"
|
||||
elif msg.metadata.get("_progress"):
|
||||
elif progress_event:
|
||||
payload["kind"] = "progress"
|
||||
phase = "activity" if payload.get("kind") in ("tool_hint", "progress") else "answer"
|
||||
self._transcripts.prepare_and_append(
|
||||
@@ -971,6 +971,8 @@ class WebSocketChannel(BaseChannel):
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
) -> None:
|
||||
"""Push one chunk of model reasoning. Mirrors ``send_delta`` shape so
|
||||
clients receive a stream that opens, updates in place, and closes —
|
||||
@@ -986,7 +988,6 @@ class WebSocketChannel(BaseChannel):
|
||||
"chat_id": chat_id,
|
||||
"text": delta,
|
||||
}
|
||||
stream_id = meta.get("_stream_id")
|
||||
if stream_id is not None:
|
||||
body["stream_id"] = stream_id
|
||||
self._transcripts.prepare_and_append(
|
||||
@@ -1005,6 +1006,8 @@ class WebSocketChannel(BaseChannel):
|
||||
self,
|
||||
chat_id: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
) -> None:
|
||||
"""Close the current reasoning stream segment for in-place renderers."""
|
||||
conns = list(self._subs.get(chat_id, ()))
|
||||
@@ -1013,7 +1016,6 @@ class WebSocketChannel(BaseChannel):
|
||||
"event": "reasoning_end",
|
||||
"chat_id": chat_id,
|
||||
}
|
||||
stream_id = meta.get("_stream_id")
|
||||
if stream_id is not None:
|
||||
body["stream_id"] = stream_id
|
||||
self._transcripts.prepare_and_append(
|
||||
@@ -1057,11 +1059,15 @@ class WebSocketChannel(BaseChannel):
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
stream_end: bool = False,
|
||||
resuming: bool = False,
|
||||
) -> None:
|
||||
conns = list(self._subs.get(chat_id, ()))
|
||||
meta = metadata or {}
|
||||
stream_key = (chat_id, str(meta.get("_stream_id") or ""))
|
||||
if meta.get("_stream_end"):
|
||||
stream_key = (chat_id, str(stream_id or ""))
|
||||
if stream_end:
|
||||
body: dict[str, Any] = {"event": "stream_end", "chat_id": chat_id}
|
||||
buffered = self._stream_text_buffers.pop(stream_key, [])
|
||||
if delta:
|
||||
@@ -1077,8 +1083,8 @@ class WebSocketChannel(BaseChannel):
|
||||
"text": delta,
|
||||
}
|
||||
self._stream_text_buffers.setdefault(stream_key, []).append(delta)
|
||||
if meta.get("_stream_id") is not None:
|
||||
body["stream_id"] = meta["_stream_id"]
|
||||
if stream_id is not None:
|
||||
body["stream_id"] = stream_id
|
||||
self._transcripts.prepare_and_append(
|
||||
chat_id,
|
||||
body,
|
||||
|
||||
@@ -13,6 +13,7 @@ from typing import Any
|
||||
from pydantic import Field
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.config.paths import get_media_dir
|
||||
@@ -497,7 +498,7 @@ class WecomChannel(BaseChannel):
|
||||
|
||||
try:
|
||||
content = (msg.content or "").strip()
|
||||
is_progress = bool(msg.metadata.get("_progress"))
|
||||
is_progress = isinstance(msg.event, ProgressEvent)
|
||||
|
||||
# Get the stored frame for this chat
|
||||
frame = self._chat_frames.get(msg.chat_id)
|
||||
|
||||
+23
-14
@@ -29,6 +29,7 @@ from loguru import logger
|
||||
from pydantic import Field
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.config.paths import get_media_dir, get_runtime_subdir
|
||||
@@ -1101,11 +1102,13 @@ class WeixinChannel(BaseChannel):
|
||||
raise RuntimeError("WeChat client not initialized or not authenticated")
|
||||
self._assert_session_active()
|
||||
|
||||
is_progress = bool((msg.metadata or {}).get("_progress", False))
|
||||
event = getattr(msg, "event", None)
|
||||
progress_event = event if isinstance(event, ProgressEvent) else None
|
||||
is_progress = progress_event is not None
|
||||
|
||||
# Buffer tool hints to coalesce consecutive ones and avoid burning
|
||||
# WeChat iLink rate-limit quota (~7 msgs / 5 min).
|
||||
if is_progress and (msg.metadata or {}).get("_tool_hint"):
|
||||
if progress_event and progress_event.tool_hint:
|
||||
if not self.send_tool_hints:
|
||||
return
|
||||
self._pending_tool_hints.setdefault(msg.chat_id, []).append(msg.content)
|
||||
@@ -1118,7 +1121,7 @@ class WeixinChannel(BaseChannel):
|
||||
|
||||
# Reasoning deltas are invisible in WeChat (there is no reasoning
|
||||
# UI). Skip them entirely — do not send and do not flush buffer.
|
||||
if is_progress and (msg.metadata or {}).get("_reasoning_delta"):
|
||||
if progress_event and (progress_event.reasoning_delta or progress_event.reasoning):
|
||||
self.logger.debug(
|
||||
"Dropped invisible reasoning delta for {}", msg.chat_id
|
||||
)
|
||||
@@ -1232,40 +1235,46 @@ class WeixinChannel(BaseChannel):
|
||||
await self._send_typing(msg.chat_id, typing_ticket, TYPING_STATUS_CANCEL)
|
||||
|
||||
async def send_delta(
|
||||
self, chat_id: str, delta: str, metadata: dict[str, Any] | None = None
|
||||
self,
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict[str, Any] | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
stream_end: bool = False,
|
||||
resuming: bool = False,
|
||||
) -> None:
|
||||
"""Deliver a streamed reply to WeChat.
|
||||
|
||||
WeChat iLink has no native incremental delivery, and the manager
|
||||
bypasses :meth:`send` for the ``_streamed`` final answer. So we
|
||||
accumulate the content deltas here and flush the full reply as a
|
||||
single message at ``_stream_end`` — otherwise a streamed reply would
|
||||
never reach the user. Reasoning deltas are invisible in WeChat and are
|
||||
dropped.
|
||||
accumulate content deltas and flush the full reply as a single message
|
||||
at stream end. Reasoning deltas are invisible in WeChat and are dropped.
|
||||
"""
|
||||
meta = metadata or {}
|
||||
if meta.get("_reasoning_delta") or meta.get("_reasoning"):
|
||||
return
|
||||
is_end = meta.get("_stream_end")
|
||||
# Accumulate intermediate deltas. The _stream_end message's own content
|
||||
is_end = stream_end or bool(meta.get("_stream_end"))
|
||||
buffer_key = stream_id or chat_id
|
||||
# Accumulate intermediate deltas. The stream_end message's own content
|
||||
# (present when the manager coalesces deltas into the end message) is
|
||||
# folded into `full` below instead of appended here, so a send retry
|
||||
# recomputes the same `full` from an unchanged buffer rather than
|
||||
# double-counting that delta.
|
||||
if delta and not is_end:
|
||||
self._stream_buffers.setdefault(chat_id, []).append(delta)
|
||||
self._stream_buffers.setdefault(buffer_key, []).append(delta)
|
||||
if not is_end:
|
||||
return
|
||||
full = ("".join(self._stream_buffers.get(chat_id, [])) + (delta or "")).strip()
|
||||
full = ("".join(self._stream_buffers.get(buffer_key, [])) + (delta or "")).strip()
|
||||
await self._flush_tool_hints(chat_id)
|
||||
if full:
|
||||
# Send before clearing the buffer: if the send raises, the buffer is
|
||||
# left intact so ChannelManager._send_with_retry can re-deliver the
|
||||
# same _stream_end message instead of silently losing the reply.
|
||||
# same stream_end message instead of silently losing the reply.
|
||||
await self.send(
|
||||
OutboundMessage(channel=self.name, chat_id=chat_id, content=full)
|
||||
)
|
||||
self._stream_buffers.pop(chat_id, None)
|
||||
self._stream_buffers.pop(buffer_key, None)
|
||||
|
||||
async def _start_typing(self, chat_id: str, context_token: str = "") -> None:
|
||||
"""Start typing indicator immediately when a message is received."""
|
||||
|
||||
+99
-18
@@ -50,6 +50,14 @@ from rich.text import Text # noqa: E402
|
||||
|
||||
from nanobot import __logo__, __version__ # noqa: E402
|
||||
from nanobot.agent.loop import AgentLoop # noqa: E402
|
||||
from nanobot.bus.outbound_events import ( # noqa: E402
|
||||
ProgressEvent,
|
||||
RetryWaitEvent,
|
||||
StreamDeltaEvent,
|
||||
StreamedResponseEvent,
|
||||
StreamEndEvent,
|
||||
outbound_event_from_message,
|
||||
)
|
||||
from nanobot.cli.gateway import create_gateway_app # noqa: E402
|
||||
from nanobot.cli.stream import StreamRenderer, ThinkingSpinner # noqa: E402
|
||||
from nanobot.config.paths import get_workspace_path, is_default_workspace # noqa: E402
|
||||
@@ -461,25 +469,25 @@ async def _maybe_print_interactive_progress(
|
||||
renderer: StreamRenderer | None = None,
|
||||
reasoning_buffer: _ReasoningBuffer | None = None,
|
||||
) -> bool:
|
||||
metadata = msg.metadata or {}
|
||||
if metadata.get("_retry_wait"):
|
||||
event = outbound_event_from_message(msg)
|
||||
if isinstance(event, RetryWaitEvent):
|
||||
await _print_interactive_progress_line(msg.content, thinking, renderer)
|
||||
return True
|
||||
|
||||
if not metadata.get("_progress"):
|
||||
if not isinstance(event, ProgressEvent):
|
||||
return False
|
||||
|
||||
reasoning_buffer = reasoning_buffer or _ReasoningBuffer()
|
||||
|
||||
if metadata.get("_reasoning_end"):
|
||||
if event.reasoning_end:
|
||||
if channels_config and not channels_config.show_reasoning:
|
||||
reasoning_buffer.clear()
|
||||
else:
|
||||
_flush_cli_reasoning(reasoning_buffer, thinking, renderer)
|
||||
return True
|
||||
|
||||
is_tool_hint = metadata.get("_tool_hint", False)
|
||||
is_reasoning = metadata.get("_reasoning", False) or metadata.get("_reasoning_delta", False)
|
||||
is_tool_hint = event.tool_hint
|
||||
is_reasoning = event.reasoning or event.reasoning_delta
|
||||
if is_reasoning:
|
||||
if channels_config and not channels_config.show_reasoning:
|
||||
reasoning_buffer.clear()
|
||||
@@ -710,6 +718,21 @@ def _load_runtime_config(config: str | None = None, workspace: str | None = None
|
||||
return loaded
|
||||
|
||||
|
||||
def _read_trigger_cli_message(message: str | None) -> str:
|
||||
"""Read a trigger message from an argument or stdin."""
|
||||
if message and message.strip():
|
||||
return message
|
||||
try:
|
||||
if not sys.stdin.isatty():
|
||||
content = sys.stdin.read()
|
||||
if content.strip():
|
||||
return content
|
||||
except Exception:
|
||||
pass
|
||||
console.print("[red]Error: trigger message is required[/red]")
|
||||
raise typer.Exit(1)
|
||||
|
||||
|
||||
def _warn_deprecated_config_keys(config_path: Path | None) -> None:
|
||||
"""Hint users to remove obsolete keys from their config file."""
|
||||
import json
|
||||
@@ -741,6 +764,35 @@ def _migrate_cron_store(config: "Config") -> None:
|
||||
shutil.move(str(legacy_path), str(new_path))
|
||||
|
||||
|
||||
@app.command()
|
||||
def trigger(
|
||||
trigger_id: str = typer.Argument(..., help="Trigger ID returned by /trigger"),
|
||||
message: str | None = typer.Argument(None, help="Message to deliver; stdin is used when omitted"),
|
||||
workspace: str | None = typer.Option(None, "--workspace", "-w", help="Workspace directory"),
|
||||
config: str | None = typer.Option(None, "--config", "-c", help="Config file path"),
|
||||
):
|
||||
"""Deliver a local trigger message to its bound chat session."""
|
||||
from nanobot.triggers.local_store import (
|
||||
LocalTriggerStore,
|
||||
TriggerDisabledError,
|
||||
TriggerNotFoundError,
|
||||
TriggerStoreError,
|
||||
)
|
||||
|
||||
runtime_config = _load_runtime_config(config, workspace)
|
||||
content = _read_trigger_cli_message(message)
|
||||
store = LocalTriggerStore(runtime_config.workspace_path)
|
||||
try:
|
||||
delivery = store.enqueue(trigger_id, content)
|
||||
except (TriggerNotFoundError, TriggerDisabledError) as exc:
|
||||
console.print(f"[red]Error: {exc}[/red]")
|
||||
raise typer.Exit(1) from exc
|
||||
except (TriggerStoreError, ValueError) as exc:
|
||||
console.print(f"[red]Error: {exc}[/red]")
|
||||
raise typer.Exit(1) from exc
|
||||
console.print(f"[green]Queued[/green] {delivery.trigger_id} ({delivery.id})")
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# OpenAI-Compatible API Server
|
||||
# ============================================================================
|
||||
@@ -798,14 +850,24 @@ def serve(
|
||||
console.print(f" [cyan]Model[/cyan] : {model_name}{preset_tag}")
|
||||
console.print(" [cyan]Session[/cyan] : api:default")
|
||||
console.print(f" [cyan]Timeout[/cyan] : {timeout}s")
|
||||
api_key = api_cfg.api_key.strip() if api_cfg.api_key else ""
|
||||
if host in {"0.0.0.0", "::"}:
|
||||
if not api_key:
|
||||
console.print(
|
||||
"[red]Error: host is 0.0.0.0 (all interfaces) but api_key is not set. "
|
||||
"Set api.api_key in config to prevent unauthenticated access.[/red]"
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
console.print(
|
||||
"[yellow]Warning:[/yellow] API is bound to all interfaces. "
|
||||
"Only do this behind a trusted network boundary, firewall, or reverse proxy."
|
||||
"[yellow]API is bound to all interfaces "
|
||||
"(authentication required).[/yellow]"
|
||||
)
|
||||
console.print()
|
||||
|
||||
api_app = create_app(agent_loop, model_name=model_name, request_timeout=timeout)
|
||||
api_app = create_app(
|
||||
agent_loop, model_name=model_name, request_timeout=timeout,
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
async def on_startup(_app):
|
||||
await agent_loop._connect_mcp()
|
||||
@@ -847,6 +909,8 @@ def _run_gateway(
|
||||
from nanobot.providers.image_generation import image_gen_provider_configs
|
||||
from nanobot.session.manager import SessionManager
|
||||
from nanobot.session.webui_turns import WebuiTurnCoordinator
|
||||
from nanobot.triggers.local_runner import run_local_trigger_queue
|
||||
from nanobot.triggers.local_store import LocalTriggerStore
|
||||
from nanobot.webui.token_usage import TokenUsageHook
|
||||
|
||||
port = port if port is not None else config.gateway.port
|
||||
@@ -869,6 +933,7 @@ def _run_gateway(
|
||||
# Create cron service with workspace-scoped store
|
||||
cron_store_path = config.workspace_path / "cron" / "jobs.json"
|
||||
cron = CronService(cron_store_path)
|
||||
trigger_store = LocalTriggerStore(config.workspace_path)
|
||||
|
||||
# Create agent with cron service
|
||||
agent = AgentLoop.from_config(
|
||||
@@ -883,13 +948,13 @@ def _run_gateway(
|
||||
runtime_events=runtime_events,
|
||||
provider_signature=provider_snapshot.signature,
|
||||
hooks=[TokenUsageHook(timezone_name=config.agents.defaults.timezone)],
|
||||
local_trigger_store=trigger_store,
|
||||
)
|
||||
WebuiTurnCoordinator(
|
||||
bus=bus,
|
||||
sessions=session_manager,
|
||||
schedule_background=lambda coro: agent._schedule_background(coro),
|
||||
).subscribe(runtime_events)
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.session.keys import session_key_for_channel
|
||||
|
||||
@@ -1085,8 +1150,14 @@ def _run_gateway(
|
||||
bus,
|
||||
session_manager=session_manager,
|
||||
cron_service=cron,
|
||||
local_trigger_store=trigger_store,
|
||||
webui_runtime_model_name=_webui_runtime_model_name,
|
||||
webui_cron_pending_job_ids=getattr(agent, "pending_cron_job_ids_for_session", None),
|
||||
webui_local_trigger_pending_ids=getattr(
|
||||
agent,
|
||||
"pending_local_trigger_ids_for_session",
|
||||
None,
|
||||
),
|
||||
webui_static_dist=webui_static_dist,
|
||||
webui_runtime_surface=webui_runtime_surface,
|
||||
webui_runtime_capabilities=webui_runtime_capabilities,
|
||||
@@ -1227,6 +1298,13 @@ def _run_gateway(
|
||||
tasks = [
|
||||
asyncio.create_task(agent.run(), name="nanobot-agent-loop"),
|
||||
asyncio.create_task(channels.start_all(), name="nanobot-channels"),
|
||||
asyncio.create_task(
|
||||
run_local_trigger_queue(
|
||||
store=trigger_store,
|
||||
submit_turn=getattr(agent, "submit_local_trigger_turn", None),
|
||||
),
|
||||
name="nanobot-local-triggers",
|
||||
),
|
||||
]
|
||||
if health_server_enabled:
|
||||
tasks.append(asyncio.create_task(
|
||||
@@ -1446,7 +1524,7 @@ def agent(
|
||||
bus_task = asyncio.create_task(agent_loop.run())
|
||||
turn_done = asyncio.Event()
|
||||
turn_done.set()
|
||||
turn_response: list[tuple[str, dict]] = []
|
||||
turn_response: list[Any] = []
|
||||
renderer: StreamRenderer | None = None
|
||||
reasoning_buffer = _ReasoningBuffer()
|
||||
|
||||
@@ -1454,18 +1532,19 @@ def agent(
|
||||
while True:
|
||||
try:
|
||||
msg = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0)
|
||||
event = outbound_event_from_message(msg)
|
||||
|
||||
if msg.metadata.get("_stream_delta"):
|
||||
if isinstance(event, StreamDeltaEvent):
|
||||
if renderer:
|
||||
await renderer.on_delta(msg.content)
|
||||
continue
|
||||
if msg.metadata.get("_stream_end"):
|
||||
if isinstance(event, StreamEndEvent):
|
||||
if renderer:
|
||||
await renderer.on_end(
|
||||
resuming=msg.metadata.get("_resuming", False),
|
||||
resuming=event.resuming,
|
||||
)
|
||||
continue
|
||||
if msg.metadata.get("_streamed"):
|
||||
if isinstance(event, StreamedResponseEvent):
|
||||
turn_done.set()
|
||||
continue
|
||||
|
||||
@@ -1480,7 +1559,7 @@ def agent(
|
||||
|
||||
if not turn_done.is_set():
|
||||
if msg.content:
|
||||
turn_response.append((msg.content, dict(msg.metadata or {})))
|
||||
turn_response.append(msg)
|
||||
turn_done.set()
|
||||
elif msg.content:
|
||||
await _print_interactive_response(
|
||||
@@ -1533,8 +1612,10 @@ def agent(
|
||||
await turn_done.wait()
|
||||
|
||||
if turn_response:
|
||||
content, meta = turn_response[0]
|
||||
if content and not meta.get("_streamed"):
|
||||
response_msg = turn_response[0]
|
||||
content = response_msg.content
|
||||
meta = response_msg.metadata
|
||||
if content and not isinstance(response_msg.event, StreamedResponseEvent):
|
||||
if renderer:
|
||||
await renderer.close()
|
||||
print_kwargs: dict[str, Any] = {}
|
||||
|
||||
@@ -81,6 +81,13 @@ BUILTIN_COMMAND_SPECS: tuple[BuiltinCommandSpec, ...] = (
|
||||
"activity",
|
||||
"<goal>",
|
||||
),
|
||||
BuiltinCommandSpec(
|
||||
"/trigger",
|
||||
"Create named local trigger",
|
||||
"Create a named CLI trigger bound to this chat session.",
|
||||
"zap",
|
||||
"<name>",
|
||||
),
|
||||
BuiltinCommandSpec(
|
||||
"/dream",
|
||||
"Run Dream",
|
||||
@@ -718,6 +725,61 @@ async def cmd_skill(ctx: CommandContext) -> OutboundMessage:
|
||||
metadata=dict(ctx.msg.metadata or {}),
|
||||
)
|
||||
|
||||
|
||||
async def cmd_trigger(ctx: CommandContext) -> OutboundMessage:
|
||||
"""Create a local trigger bound to the current session."""
|
||||
name = ctx.args.strip()
|
||||
if not name:
|
||||
return OutboundMessage(
|
||||
channel=ctx.msg.channel,
|
||||
chat_id=ctx.msg.chat_id,
|
||||
content=(
|
||||
"Usage: /trigger <name>\n\n"
|
||||
"Create a named local trigger bound to this chat session."
|
||||
),
|
||||
metadata={**dict(ctx.msg.metadata or {}), "render_as": "text"},
|
||||
)
|
||||
|
||||
from nanobot.triggers.local_store import LocalTriggerStore
|
||||
|
||||
loop = ctx.loop
|
||||
workspace = getattr(loop, "workspace", None)
|
||||
if workspace is None:
|
||||
workspace = getattr(getattr(loop, "context", None), "workspace", None)
|
||||
if workspace is None:
|
||||
raise RuntimeError("workspace unavailable for trigger creation")
|
||||
|
||||
store = getattr(loop, "local_trigger_store", None)
|
||||
if store is None:
|
||||
store = LocalTriggerStore(workspace)
|
||||
|
||||
from nanobot.session.keys import UNIFIED_SESSION_KEY
|
||||
|
||||
session_key = (
|
||||
ctx.msg.session_key
|
||||
if ctx.key == UNIFIED_SESSION_KEY
|
||||
else ctx.key
|
||||
)
|
||||
trigger = store.create(
|
||||
name=name,
|
||||
channel=ctx.msg.channel,
|
||||
chat_id=ctx.msg.chat_id,
|
||||
session_key=session_key,
|
||||
sender_id="trigger",
|
||||
origin_metadata=dict(ctx.msg.metadata or {}),
|
||||
)
|
||||
command = f'nanobot trigger {trigger.id} "message"'
|
||||
return OutboundMessage(
|
||||
channel=ctx.msg.channel,
|
||||
chat_id=ctx.msg.chat_id,
|
||||
content=(
|
||||
f"Trigger created: {trigger.name}\n"
|
||||
f"ID: {trigger.id}\n\n"
|
||||
f"Command:\n{command}"
|
||||
),
|
||||
metadata={**dict(ctx.msg.metadata or {}), "render_as": "text"},
|
||||
)
|
||||
|
||||
async def cmd_help(ctx: CommandContext) -> OutboundMessage:
|
||||
"""Return available slash commands."""
|
||||
return OutboundMessage(
|
||||
@@ -752,6 +814,8 @@ def register_builtin_commands(router: CommandRouter) -> None:
|
||||
router.prefix("/history ", cmd_history)
|
||||
router.exact("/goal", cmd_goal)
|
||||
router.prefix("/goal ", cmd_goal)
|
||||
router.exact("/trigger", cmd_trigger)
|
||||
router.prefix("/trigger ", cmd_trigger)
|
||||
router.exact("/dream", cmd_dream)
|
||||
router.exact("/dream-log", cmd_dream_log)
|
||||
router.prefix("/dream-log ", cmd_dream_log)
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Any, Awaitable, Callable
|
||||
|
||||
@@ -10,6 +11,26 @@ if TYPE_CHECKING:
|
||||
from nanobot.session.manager import Session
|
||||
|
||||
Handler = Callable[["CommandContext"], Awaitable["OutboundMessage | None"]]
|
||||
_BOT_SUFFIX_RE = re.compile(r"^[A-Za-z0-9_]+$")
|
||||
|
||||
|
||||
def normalize_command_text(text: str) -> str:
|
||||
"""Normalize slash-command transport variants before routing.
|
||||
|
||||
Telegram and Discord-style command dispatch can produce ``/cmd@bot args``.
|
||||
The bot suffix belongs to the transport, not the command name, so strip it
|
||||
once at the router boundary while preserving user arguments verbatim.
|
||||
"""
|
||||
stripped = text.strip()
|
||||
if not stripped.startswith("/"):
|
||||
return stripped
|
||||
first, sep, rest = stripped.partition(" ")
|
||||
if "@" not in first:
|
||||
return stripped
|
||||
command, suffix = first.rsplit("@", 1)
|
||||
if command and suffix and _BOT_SUFFIX_RE.fullmatch(suffix):
|
||||
return f"{command}{sep}{rest}" if sep else command
|
||||
return stripped
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -50,7 +71,7 @@ class CommandRouter:
|
||||
self._prefix.sort(key=lambda p: len(p[0]), reverse=True)
|
||||
|
||||
def is_priority(self, text: str) -> bool:
|
||||
return text.strip().lower() in self._priority
|
||||
return normalize_command_text(text).lower() in self._priority
|
||||
|
||||
def is_dispatchable_command(self, text: str) -> bool:
|
||||
"""Check whether *text* matches any non-priority command tier (exact or prefix).
|
||||
@@ -58,7 +79,7 @@ class CommandRouter:
|
||||
Does NOT check priority tier.
|
||||
If this returns True, ``dispatch()`` is guaranteed to match a handler.
|
||||
"""
|
||||
cmd = text.strip().lower()
|
||||
cmd = normalize_command_text(text).lower()
|
||||
if cmd in self._exact:
|
||||
return True
|
||||
for pfx, _ in self._prefix:
|
||||
@@ -68,6 +89,7 @@ class CommandRouter:
|
||||
|
||||
async def dispatch_priority(self, ctx: CommandContext) -> OutboundMessage | None:
|
||||
"""Dispatch a priority command. Called from run() without the lock."""
|
||||
ctx.raw = normalize_command_text(ctx.raw)
|
||||
handler = self._priority.get(ctx.raw.lower())
|
||||
if handler:
|
||||
return await handler(ctx)
|
||||
@@ -75,6 +97,7 @@ class CommandRouter:
|
||||
|
||||
async def dispatch(self, ctx: CommandContext) -> OutboundMessage | None:
|
||||
"""Try exact, then prefix handlers. Returns None if unhandled."""
|
||||
ctx.raw = normalize_command_text(ctx.raw)
|
||||
cmd = ctx.raw.lower()
|
||||
|
||||
if handler := self._exact.get(cmd):
|
||||
|
||||
@@ -307,6 +307,18 @@ class ApiConfig(Base):
|
||||
host: str = "127.0.0.1" # Safer default: local-only bind.
|
||||
port: int = 8900
|
||||
timeout: float = 120.0 # Per-request timeout in seconds.
|
||||
api_key: str = Field(default="", repr=False)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def wildcard_host_requires_auth(self) -> "ApiConfig":
|
||||
if self.host not in ("0.0.0.0", "::"):
|
||||
return self
|
||||
if self.api_key.strip():
|
||||
return self
|
||||
raise ValueError(
|
||||
"host is 0.0.0.0 (all interfaces) but api_key is not set "
|
||||
"- set api.api_key to prevent unauthenticated access"
|
||||
)
|
||||
|
||||
|
||||
class GatewayConfig(Base):
|
||||
|
||||
+15
-13
@@ -1,6 +1,7 @@
|
||||
"""Cron service for scheduling agent tasks."""
|
||||
|
||||
import asyncio
|
||||
import errno
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
@@ -23,6 +24,12 @@ 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,
|
||||
)
|
||||
|
||||
|
||||
class CronJobSkippedError(Exception):
|
||||
@@ -456,11 +463,15 @@ class CronService:
|
||||
os.replace(tmp_path, path)
|
||||
# fsync the parent directory so the rename itself is durable.
|
||||
# Skip on Windows where opening a directory raises PermissionError;
|
||||
# NTFS journals metadata synchronously so this is a no-op there.
|
||||
# some shared filesystems reject directory fsync with EINVAL.
|
||||
with suppress(PermissionError):
|
||||
fd = os.open(str(path.parent), os.O_RDONLY)
|
||||
try:
|
||||
os.fsync(fd)
|
||||
try:
|
||||
os.fsync(fd)
|
||||
except OSError as exc:
|
||||
if exc.errno != errno.EINVAL:
|
||||
raise
|
||||
finally:
|
||||
os.close(fd)
|
||||
except BaseException:
|
||||
@@ -469,20 +480,11 @@ class CronService:
|
||||
|
||||
@staticmethod
|
||||
def _safe_run_record_name(run_id: str) -> str:
|
||||
return "".join(c if c.isalnum() or c in "._-" else "_" for c in run_id)
|
||||
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."""
|
||||
name = self._safe_run_record_name(run_id)
|
||||
if not name:
|
||||
name = str(uuid.uuid4())
|
||||
path = self._run_records_dir / f"{name}.json"
|
||||
payload = {
|
||||
**record,
|
||||
"run_id": run_id,
|
||||
"updated_at_ms": _now_ms(),
|
||||
}
|
||||
self._atomic_write(path, json.dumps(payload, indent=2, ensure_ascii=False))
|
||||
write_automation_run_record(self._run_records_dir, run_id, record)
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Start the cron service."""
|
||||
|
||||
@@ -5,16 +5,43 @@ from __future__ import annotations
|
||||
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,
|
||||
)
|
||||
|
||||
CRON_TRIGGER_META = "_cron_trigger"
|
||||
CRON_DEFER_UNTIL_IDLE_META = "_cron_defer_until_session_idle"
|
||||
CRON_HISTORY_META = "_cron_turn"
|
||||
|
||||
|
||||
def _cron_history_text(trigger: Mapping[str, Any]) -> str | None:
|
||||
persist_content = trigger.get("persist_content")
|
||||
return (
|
||||
persist_content
|
||||
if isinstance(persist_content, str) and persist_content.strip()
|
||||
else None
|
||||
)
|
||||
|
||||
|
||||
CRON_AUTOMATION_SPEC = AutomationTurnSpec(
|
||||
kind="cron",
|
||||
trigger_meta_key=CRON_TRIGGER_META,
|
||||
legacy_history_meta_key=CRON_HISTORY_META,
|
||||
history_fields={
|
||||
"cron_job_id": "job_id",
|
||||
"cron_job_name": "job_name",
|
||||
"cron_run_id": "run_id",
|
||||
"cron_prompt_ref": "prompt_ref",
|
||||
},
|
||||
text_builder=_cron_history_text,
|
||||
)
|
||||
|
||||
|
||||
def cron_trigger(metadata: Mapping[str, Any] | None) -> dict[str, Any] | None:
|
||||
"""Return structured cron trigger metadata when present."""
|
||||
raw = (metadata or {}).get(CRON_TRIGGER_META)
|
||||
return raw if isinstance(raw, dict) else None
|
||||
return automation_trigger(metadata, CRON_AUTOMATION_SPEC)
|
||||
|
||||
|
||||
def is_cron_turn(metadata: Mapping[str, Any] | None) -> bool:
|
||||
@@ -38,22 +65,7 @@ def cron_run_id(metadata: Mapping[str, Any] | None) -> str | 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."""
|
||||
trigger = cron_trigger(metadata)
|
||||
if not trigger:
|
||||
return None, {}
|
||||
persist_content = trigger.get("persist_content")
|
||||
text = (
|
||||
persist_content
|
||||
if isinstance(persist_content, str) and persist_content.strip()
|
||||
else None
|
||||
)
|
||||
return text, {
|
||||
CRON_HISTORY_META: True,
|
||||
"cron_job_id": trigger.get("job_id"),
|
||||
"cron_job_name": trigger.get("job_name"),
|
||||
"cron_run_id": trigger.get("run_id"),
|
||||
"cron_prompt_ref": trigger.get("prompt_ref"),
|
||||
}
|
||||
return automation_history_overrides_for_spec(metadata, CRON_AUTOMATION_SPEC)
|
||||
|
||||
|
||||
def is_bound_cron_job(job: CronJob) -> bool:
|
||||
|
||||
@@ -32,6 +32,7 @@ class ProviderSpec:
|
||||
keywords: tuple[str, ...] # model-name keywords for matching (lowercase)
|
||||
env_key: str # env var for API key, e.g. "DASHSCOPE_API_KEY"
|
||||
display_name: str = "" # shown in `nanobot status`
|
||||
model_catalog: str = "auto" # WebUI model-list source
|
||||
|
||||
# which provider implementation to use
|
||||
# "openai_compat" | "anthropic" | "azure_openai" | "openai_codex" | "github_copilot" | "bedrock"
|
||||
@@ -221,6 +222,7 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
|
||||
keywords=("skywork", "skyclaw", "apifree"),
|
||||
env_key="SKYWORK_API_KEY",
|
||||
display_name="Skywork",
|
||||
model_catalog="official",
|
||||
backend="openai_compat",
|
||||
env_extras=(("APIFREE_API_KEY", "{api_key}"),),
|
||||
is_gateway=True,
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
"""Shared handling for session-bound automation turns."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from functools import lru_cache
|
||||
from typing import Any
|
||||
|
||||
AUTOMATION_HISTORY_META = "_automation_turn"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AutomationTurnSpec:
|
||||
"""Source-specific wiring for one session-bound automation turn type."""
|
||||
|
||||
kind: str
|
||||
trigger_meta_key: str
|
||||
legacy_history_meta_key: str | None = None
|
||||
history_fields: Mapping[str, str] = field(default_factory=dict)
|
||||
text_builder: Callable[[Mapping[str, Any]], str | None] | None = None
|
||||
|
||||
|
||||
def automation_trigger(
|
||||
metadata: Mapping[str, Any] | None,
|
||||
spec: AutomationTurnSpec,
|
||||
) -> dict[str, Any] | None:
|
||||
"""Return source trigger metadata for *spec* when present."""
|
||||
raw = (metadata or {}).get(spec.trigger_meta_key)
|
||||
return raw if isinstance(raw, dict) else None
|
||||
|
||||
|
||||
def automation_history_overrides_for_spec(
|
||||
metadata: Mapping[str, Any] | None,
|
||||
spec: AutomationTurnSpec,
|
||||
) -> tuple[str | None, dict[str, Any]]:
|
||||
"""Return hidden session-history text/metadata overrides for *spec*."""
|
||||
trigger = automation_trigger(metadata, spec)
|
||||
if not trigger:
|
||||
return None, {}
|
||||
|
||||
details: dict[str, Any] = {"kind": spec.kind}
|
||||
extra: dict[str, Any] = {AUTOMATION_HISTORY_META: details}
|
||||
if spec.legacy_history_meta_key:
|
||||
extra[spec.legacy_history_meta_key] = True
|
||||
for history_key, trigger_key in spec.history_fields.items():
|
||||
value = trigger.get(trigger_key)
|
||||
extra[history_key] = value
|
||||
details[history_key] = value
|
||||
|
||||
text = spec.text_builder(trigger) if spec.text_builder else None
|
||||
return text, extra
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _automation_specs() -> tuple[AutomationTurnSpec, ...]:
|
||||
# Source modules import the generic helpers above, so keep spec loading lazy.
|
||||
from nanobot.cron.session_turns import CRON_AUTOMATION_SPEC
|
||||
from nanobot.triggers.local_session_turns import LOCAL_TRIGGER_AUTOMATION_SPEC
|
||||
|
||||
return (CRON_AUTOMATION_SPEC, LOCAL_TRIGGER_AUTOMATION_SPEC)
|
||||
|
||||
|
||||
def automation_history_overrides(
|
||||
metadata: Mapping[str, Any] | None,
|
||||
) -> tuple[str | None, dict[str, Any]]:
|
||||
"""Return session-history text/metadata overrides for supported automation turns."""
|
||||
for spec in _automation_specs():
|
||||
text, extra = automation_history_overrides_for_spec(metadata, spec)
|
||||
if extra:
|
||||
return text, extra
|
||||
return None, {}
|
||||
|
||||
|
||||
def is_automation_history_message(message: Mapping[str, Any] | None) -> bool:
|
||||
"""True for hidden automation trigger records in session history."""
|
||||
if not message:
|
||||
return False
|
||||
marker = message.get(AUTOMATION_HISTORY_META)
|
||||
if marker is True or isinstance(marker, Mapping):
|
||||
return True
|
||||
return any(
|
||||
spec.legacy_history_meta_key
|
||||
and message.get(spec.legacy_history_meta_key) is True
|
||||
for spec in _automation_specs()
|
||||
)
|
||||
|
||||
|
||||
def is_automation_kind(value: Any) -> bool:
|
||||
return isinstance(value, str) and (
|
||||
value == "trigger" or any(spec.kind == value for spec in _automation_specs())
|
||||
)
|
||||
@@ -0,0 +1,22 @@
|
||||
"""Visibility helpers for persisted session history messages."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
|
||||
from nanobot.session.automation_turns import is_automation_history_message
|
||||
|
||||
HIDDEN_HISTORY_META = "_hidden_history"
|
||||
|
||||
|
||||
def _has_hidden_history_marker(message: Mapping[str, Any] | None) -> bool:
|
||||
if not message:
|
||||
return False
|
||||
marker = message.get(HIDDEN_HISTORY_META)
|
||||
return marker is True or isinstance(marker, Mapping)
|
||||
|
||||
|
||||
def is_hidden_history_message(message: Mapping[str, Any] | None) -> bool:
|
||||
"""True for persisted messages that should not be shown as chat turns."""
|
||||
return _has_hidden_history_marker(message) or is_automation_history_message(message)
|
||||
@@ -15,6 +15,7 @@ from typing import Any
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.config.paths import get_legacy_sessions_dir
|
||||
from nanobot.session.turn_history import PENDING_USER_TURN_KEY, RUNTIME_CHECKPOINT_KEY
|
||||
from nanobot.utils.helpers import (
|
||||
ensure_dir,
|
||||
estimate_message_tokens,
|
||||
@@ -37,8 +38,8 @@ _SESSION_LIST_PREVIEW_MAX_RECORDS = 200
|
||||
_SESSION_LIST_PREVIEW_MAX_CHARS = 1_000_000
|
||||
_FORK_VOLATILE_METADATA_KEYS = {
|
||||
"goal_state",
|
||||
"pending_user_turn",
|
||||
"runtime_checkpoint",
|
||||
PENDING_USER_TURN_KEY,
|
||||
RUNTIME_CHECKPOINT_KEY,
|
||||
"thread_goal",
|
||||
"title",
|
||||
"title_user_edited",
|
||||
|
||||
@@ -29,10 +29,6 @@ _GOAL_CONTINUATION_SENDER = "system:continuation"
|
||||
_GOAL_CONTINUATION_ROUNDS_KEY = "_sustained_goal_continuation_rounds"
|
||||
_MAX_GOAL_CONTINUATION_ROUNDS = 12
|
||||
_STRIPPED_INBOUND_META_KEYS = {
|
||||
"_stream_id",
|
||||
"_stream_delta",
|
||||
"_stream_end",
|
||||
"_resuming",
|
||||
INTERNAL_CONTINUATION_PENDING_META,
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,266 @@
|
||||
"""Turn history persistence and interrupted-turn recovery."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.utils.helpers import image_placeholder_text
|
||||
from nanobot.utils.helpers import truncate_text as truncate_text_fn
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from nanobot.session.manager import Session
|
||||
|
||||
RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
|
||||
PENDING_USER_TURN_KEY = "pending_user_turn"
|
||||
|
||||
|
||||
def sanitize_persisted_blocks(
|
||||
content: list[dict[str, Any]],
|
||||
*,
|
||||
max_tool_result_chars: int,
|
||||
runtime_context_tag: str,
|
||||
should_truncate_text: bool = False,
|
||||
drop_runtime: bool = False,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Strip volatile multimodal payloads before writing session history."""
|
||||
filtered: list[dict[str, Any]] = []
|
||||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
filtered.append(block)
|
||||
continue
|
||||
|
||||
if (
|
||||
drop_runtime
|
||||
and block.get("type") == "text"
|
||||
and isinstance(block.get("text"), str)
|
||||
and block["text"].startswith(runtime_context_tag)
|
||||
):
|
||||
continue
|
||||
|
||||
if block.get("type") == "image_url" and block.get("image_url", {}).get(
|
||||
"url", ""
|
||||
).startswith("data:image/"):
|
||||
path = (block.get("_meta") or {}).get("path", "")
|
||||
filtered.append({"type": "text", "text": image_placeholder_text(path)})
|
||||
continue
|
||||
|
||||
if block.get("type") == "text" and isinstance(block.get("text"), str):
|
||||
text = block["text"]
|
||||
if should_truncate_text and len(text) > max_tool_result_chars:
|
||||
text = truncate_text_fn(text, max_tool_result_chars)
|
||||
filtered.append({**block, "text": text})
|
||||
continue
|
||||
|
||||
filtered.append(block)
|
||||
|
||||
return filtered
|
||||
|
||||
|
||||
def save_turn(
|
||||
session: Session,
|
||||
messages: list[dict],
|
||||
skip: int,
|
||||
*,
|
||||
max_tool_result_chars: int,
|
||||
runtime_context_tag: str,
|
||||
turn_latency_ms: int | None = None,
|
||||
) -> None:
|
||||
"""Save new-turn messages into session, truncating large tool results."""
|
||||
declared_tool_call_ids = {
|
||||
str(tc["id"])
|
||||
for m in session.messages
|
||||
if m.get("role") == "assistant"
|
||||
for tc in m.get("tool_calls") or []
|
||||
if isinstance(tc, dict) and tc.get("id")
|
||||
}
|
||||
last_assistant_idx: int | None = None
|
||||
for m in messages[skip:]:
|
||||
entry = dict(m)
|
||||
role, content = entry.get("role"), entry.get("content")
|
||||
if role == "assistant" and not content and not entry.get("tool_calls"):
|
||||
continue # skip empty assistant messages - they poison session context
|
||||
if role == "tool":
|
||||
tool_call_id = entry.get("tool_call_id")
|
||||
if not tool_call_id or str(tool_call_id) not in declared_tool_call_ids:
|
||||
# Undeclared tool results corrupt future provider requests.
|
||||
logger.warning(
|
||||
"Dropping orphaned tool result {} from session {} during persistence",
|
||||
tool_call_id or "(missing id)",
|
||||
session.key,
|
||||
)
|
||||
continue
|
||||
if isinstance(content, str) and len(content) > max_tool_result_chars:
|
||||
entry["content"] = truncate_text_fn(content, max_tool_result_chars)
|
||||
elif isinstance(content, list):
|
||||
filtered = sanitize_persisted_blocks(
|
||||
content,
|
||||
max_tool_result_chars=max_tool_result_chars,
|
||||
runtime_context_tag=runtime_context_tag,
|
||||
should_truncate_text=True,
|
||||
)
|
||||
if not filtered:
|
||||
# Preserve the tool_call/result pair after block filtering.
|
||||
filtered = [
|
||||
{"type": "text", "text": "[tool result omitted during persistence]"}
|
||||
]
|
||||
entry["content"] = filtered
|
||||
elif role == "user":
|
||||
if isinstance(content, str) and runtime_context_tag in content:
|
||||
# Strip the runtime-context block appended at the end.
|
||||
tag_pos = content.find(runtime_context_tag)
|
||||
before = content[:tag_pos].rstrip("\n ")
|
||||
if before:
|
||||
entry["content"] = before
|
||||
else:
|
||||
continue
|
||||
if isinstance(content, list):
|
||||
filtered = sanitize_persisted_blocks(
|
||||
content,
|
||||
max_tool_result_chars=max_tool_result_chars,
|
||||
runtime_context_tag=runtime_context_tag,
|
||||
drop_runtime=True,
|
||||
)
|
||||
if not filtered:
|
||||
continue
|
||||
entry["content"] = filtered
|
||||
entry.setdefault("timestamp", datetime.now().isoformat())
|
||||
session.messages.append(entry)
|
||||
if role == "assistant":
|
||||
last_assistant_idx = len(session.messages) - 1
|
||||
declared_tool_call_ids.update(
|
||||
str(tc["id"])
|
||||
for tc in entry.get("tool_calls") or []
|
||||
if isinstance(tc, dict) and tc.get("id")
|
||||
)
|
||||
if turn_latency_ms is not None and last_assistant_idx is not None:
|
||||
session.messages[last_assistant_idx]["latency_ms"] = int(turn_latency_ms)
|
||||
session.updated_at = datetime.now()
|
||||
|
||||
|
||||
def persist_subagent_followup(session: Session, msg: Any) -> bool:
|
||||
"""Persist subagent follow-ups before prompt assembly so history stays durable.
|
||||
|
||||
Returns True if a new entry was appended; False if the follow-up was
|
||||
deduped (same ``subagent_task_id`` already in session) or carries no
|
||||
content worth persisting.
|
||||
"""
|
||||
if not msg.content:
|
||||
return False
|
||||
task_id = msg.metadata.get("subagent_task_id") if isinstance(msg.metadata, dict) else None
|
||||
if task_id and any(
|
||||
m.get("injected_event") == "subagent_result" and m.get("subagent_task_id") == task_id
|
||||
for m in session.messages
|
||||
):
|
||||
return False
|
||||
session.add_message(
|
||||
"assistant",
|
||||
msg.content,
|
||||
sender_id=msg.sender_id,
|
||||
injected_event="subagent_result",
|
||||
subagent_task_id=task_id,
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
def set_runtime_checkpoint(session: Session, payload: dict[str, Any]) -> None:
|
||||
"""Persist the latest in-flight turn state into session metadata."""
|
||||
session.metadata[RUNTIME_CHECKPOINT_KEY] = payload
|
||||
|
||||
|
||||
def mark_pending_user_turn(session: Session) -> None:
|
||||
session.metadata[PENDING_USER_TURN_KEY] = True
|
||||
|
||||
|
||||
def clear_pending_user_turn(session: Session) -> None:
|
||||
session.metadata.pop(PENDING_USER_TURN_KEY, None)
|
||||
|
||||
|
||||
def clear_runtime_checkpoint(session: Session) -> None:
|
||||
session.metadata.pop(RUNTIME_CHECKPOINT_KEY, None)
|
||||
|
||||
|
||||
def checkpoint_message_key(message: dict[str, Any]) -> tuple[Any, ...]:
|
||||
return (
|
||||
message.get("role"),
|
||||
message.get("content"),
|
||||
message.get("tool_call_id"),
|
||||
message.get("name"),
|
||||
message.get("tool_calls"),
|
||||
message.get("reasoning_content"),
|
||||
message.get("thinking_blocks"),
|
||||
)
|
||||
|
||||
|
||||
def restore_runtime_checkpoint(session: Session) -> bool:
|
||||
"""Materialize an unfinished turn into session history before a new request."""
|
||||
checkpoint = session.metadata.get(RUNTIME_CHECKPOINT_KEY)
|
||||
if not isinstance(checkpoint, dict):
|
||||
return False
|
||||
|
||||
assistant_message = checkpoint.get("assistant_message")
|
||||
completed_tool_results = checkpoint.get("completed_tool_results") or []
|
||||
pending_tool_calls = checkpoint.get("pending_tool_calls") or []
|
||||
|
||||
restored_messages: list[dict[str, Any]] = []
|
||||
if isinstance(assistant_message, dict):
|
||||
restored = dict(assistant_message)
|
||||
restored.setdefault("timestamp", datetime.now().isoformat())
|
||||
restored_messages.append(restored)
|
||||
for message in completed_tool_results:
|
||||
if isinstance(message, dict):
|
||||
restored = dict(message)
|
||||
restored.setdefault("timestamp", datetime.now().isoformat())
|
||||
restored_messages.append(restored)
|
||||
for tool_call in pending_tool_calls:
|
||||
if not isinstance(tool_call, dict):
|
||||
continue
|
||||
tool_id = tool_call.get("id")
|
||||
name = ((tool_call.get("function") or {}).get("name")) or "tool"
|
||||
restored_messages.append(
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": tool_id,
|
||||
"name": name,
|
||||
"content": "Error: Task interrupted before this tool finished.",
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
}
|
||||
)
|
||||
|
||||
overlap = 0
|
||||
max_overlap = min(len(session.messages), len(restored_messages))
|
||||
for size in range(max_overlap, 0, -1):
|
||||
existing = session.messages[-size:]
|
||||
restored = restored_messages[:size]
|
||||
if all(
|
||||
checkpoint_message_key(left) == checkpoint_message_key(right)
|
||||
for left, right in zip(existing, restored)
|
||||
):
|
||||
overlap = size
|
||||
break
|
||||
session.messages.extend(restored_messages[overlap:])
|
||||
|
||||
clear_pending_user_turn(session)
|
||||
clear_runtime_checkpoint(session)
|
||||
return True
|
||||
|
||||
|
||||
def restore_pending_user_turn(session: Session) -> bool:
|
||||
"""Close a turn that only persisted the user message before crashing."""
|
||||
if not session.metadata.get(PENDING_USER_TURN_KEY):
|
||||
return False
|
||||
|
||||
if session.messages and session.messages[-1].get("role") == "user":
|
||||
session.messages.append(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": "Error: Task interrupted before a response was generated.",
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
}
|
||||
)
|
||||
session.updated_at = datetime.now()
|
||||
|
||||
clear_pending_user_turn(session)
|
||||
return True
|
||||
@@ -11,7 +11,15 @@ from typing import Any
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.bus import progress as bus_progress
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.outbound_events import (
|
||||
GoalStateSyncEvent,
|
||||
GoalStatusEvent,
|
||||
RuntimeModelUpdatedEvent,
|
||||
SessionUpdatedEvent,
|
||||
TurnEndEvent,
|
||||
outbound_message_for_event,
|
||||
)
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.bus.runtime_events import (
|
||||
GoalStateChanged,
|
||||
@@ -22,9 +30,9 @@ from nanobot.bus.runtime_events import (
|
||||
TurnCompleted,
|
||||
TurnRunStatusChanged,
|
||||
)
|
||||
from nanobot.cron.session_turns import CRON_HISTORY_META
|
||||
from nanobot.providers.base import LLMProvider
|
||||
from nanobot.session.goal_state import goal_state_ws_blob
|
||||
from nanobot.session.history_visibility import is_hidden_history_message
|
||||
from nanobot.session.manager import Session, SessionManager
|
||||
from nanobot.utils.helpers import strip_think, truncate_text
|
||||
from nanobot.utils.llm_runtime import LLMRuntime
|
||||
@@ -69,7 +77,7 @@ def _title_inputs(session: Session) -> tuple[str, str]:
|
||||
for message in session.messages:
|
||||
if message.get("_command") is True:
|
||||
continue
|
||||
if message.get(CRON_HISTORY_META) is True:
|
||||
if is_hidden_history_message(message):
|
||||
continue
|
||||
role = message.get("role")
|
||||
content = message.get("content")
|
||||
@@ -206,26 +214,22 @@ async def publish_turn_run_status(
|
||||
if msg.channel != "websocket":
|
||||
return
|
||||
cid = str(msg.chat_id)
|
||||
meta: dict[str, Any] = {
|
||||
**dict(msg.metadata or {}),
|
||||
"_goal_status": True,
|
||||
"goal_status": status,
|
||||
}
|
||||
started_at_event: float | None = None
|
||||
if status == "running":
|
||||
if isinstance(started_at, int | float) and started_at > 0:
|
||||
t0 = float(started_at)
|
||||
else:
|
||||
t0 = time.time()
|
||||
meta["started_at"] = t0
|
||||
started_at_event = t0
|
||||
_WEBSOCKET_TURN_WALL_STARTED_AT[cid] = t0
|
||||
else:
|
||||
_WEBSOCKET_TURN_WALL_STARTED_AT.pop(cid, None)
|
||||
await bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
outbound_message_for_event(
|
||||
channel=msg.channel,
|
||||
chat_id=cid,
|
||||
content="",
|
||||
metadata=meta,
|
||||
event=GoalStatusEvent(status=status, started_at=started_at_event),
|
||||
metadata=msg.metadata,
|
||||
),
|
||||
)
|
||||
|
||||
@@ -318,28 +322,25 @@ class WebuiTurnCoordinator:
|
||||
if not cid:
|
||||
return
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
outbound_message_for_event(
|
||||
channel=event.context.channel,
|
||||
chat_id=cid,
|
||||
content="",
|
||||
metadata={
|
||||
"_goal_state_sync": True,
|
||||
"goal_state": goal_state_ws_blob(event.session_metadata),
|
||||
},
|
||||
event=GoalStateSyncEvent(
|
||||
goal_state=goal_state_ws_blob(event.session_metadata),
|
||||
),
|
||||
metadata=event.context.metadata,
|
||||
),
|
||||
)
|
||||
|
||||
async def _handle_runtime_model_changed(self, event: RuntimeModelChanged) -> None:
|
||||
await self.bus.publish_outbound(
|
||||
OutboundMessage(
|
||||
outbound_message_for_event(
|
||||
channel="websocket",
|
||||
chat_id="*",
|
||||
content="",
|
||||
metadata={
|
||||
"_runtime_model_updated": True,
|
||||
"model": event.model,
|
||||
"model_preset": event.model_preset,
|
||||
},
|
||||
event=RuntimeModelUpdatedEvent(
|
||||
model=event.model,
|
||||
model_preset=event.model_preset,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -374,17 +375,18 @@ class WebuiTurnCoordinator:
|
||||
if msg.channel != "websocket":
|
||||
return
|
||||
|
||||
turn_metadata: dict[str, Any] = {**msg.metadata, "_turn_end": True}
|
||||
if latency_ms is not None:
|
||||
turn_metadata["latency_ms"] = int(latency_ms)
|
||||
session = self.sessions.get_or_create(session_key)
|
||||
turn_metadata["goal_state"] = goal_state_ws_blob(session.metadata)
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="",
|
||||
metadata=turn_metadata,
|
||||
))
|
||||
await self.bus.publish_outbound(
|
||||
outbound_message_for_event(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
event=TurnEndEvent(
|
||||
latency_ms=latency_ms,
|
||||
goal_state=goal_state_ws_blob(session.metadata),
|
||||
),
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
)
|
||||
self._schedule_title_update(msg, session_key=session_key)
|
||||
|
||||
def _schedule_title_update(self, msg: InboundMessage, *, session_key: str) -> None:
|
||||
@@ -404,16 +406,11 @@ class WebuiTurnCoordinator:
|
||||
model=title_llm.model,
|
||||
)
|
||||
if generated:
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
await self._publish_session_metadata_updated(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content="",
|
||||
metadata={
|
||||
**msg.metadata,
|
||||
"_session_updated": True,
|
||||
"_session_update_scope": "metadata",
|
||||
},
|
||||
))
|
||||
metadata=msg.metadata,
|
||||
)
|
||||
|
||||
self.schedule_background(_generate_title_and_notify())
|
||||
|
||||
@@ -438,15 +435,26 @@ class WebuiTurnCoordinator:
|
||||
model=title_llm.model,
|
||||
)
|
||||
if generated:
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
await self._publish_session_metadata_updated(
|
||||
channel=event.context.channel,
|
||||
chat_id=event.context.chat_id,
|
||||
content="",
|
||||
metadata={
|
||||
**event.context.metadata,
|
||||
"_session_updated": True,
|
||||
"_session_update_scope": "metadata",
|
||||
},
|
||||
))
|
||||
metadata=event.context.metadata,
|
||||
)
|
||||
|
||||
self.schedule_background(_generate_title_and_notify())
|
||||
|
||||
async def _publish_session_metadata_updated(
|
||||
self,
|
||||
*,
|
||||
channel: str,
|
||||
chat_id: str,
|
||||
metadata: dict[str, Any],
|
||||
) -> None:
|
||||
await self.bus.publish_outbound(
|
||||
outbound_message_for_event(
|
||||
channel=channel,
|
||||
chat_id=chat_id,
|
||||
event=SessionUpdatedEvent(scope="metadata"),
|
||||
metadata=metadata,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -5,7 +5,4 @@ Task: {{ task }}
|
||||
Result:
|
||||
{{ result }}
|
||||
|
||||
Use this result as evidence for the current turn. For MapReduce-style work,
|
||||
preserve any Summary / Evidence / Open issues structure when reducing multiple
|
||||
results. Mention gaps or failures if they affect the answer; avoid exposing
|
||||
internal task IDs unless they are needed for clarity.
|
||||
Summarize this naturally for the user. Keep it brief (1-2 sentences). Do not mention technical details like "subagent" or task IDs.
|
||||
|
||||
@@ -4,15 +4,6 @@
|
||||
|
||||
You are a subagent spawned by the main agent to complete a specific task.
|
||||
Stay focused on the assigned task. Your final response will be reported back to the main agent.
|
||||
If this task is one slice of a larger MapReduce-style effort, treat yourself as
|
||||
the map step: do only the assigned slice, avoid cross-slice coordination, and
|
||||
leave reduction or final synthesis to the main agent.
|
||||
|
||||
For MapReduce-style slices, end with a compact, mergeable result:
|
||||
|
||||
- Summary: what you found or changed
|
||||
- Evidence: relevant files, commands, URLs, or observations
|
||||
- Open issues: blockers, failures, or "none"
|
||||
|
||||
{% include 'agent/_snippets/untrusted_content.md' %}
|
||||
|
||||
|
||||
@@ -0,0 +1,19 @@
|
||||
"""Local trigger support."""
|
||||
|
||||
from nanobot.triggers.local_store import (
|
||||
LocalTriggerStore,
|
||||
TriggerDisabledError,
|
||||
TriggerNotFoundError,
|
||||
TriggerStoreError,
|
||||
)
|
||||
from nanobot.triggers.local_types import LocalTrigger, TriggerDelivery, TriggerRunRecord
|
||||
|
||||
__all__ = [
|
||||
"LocalTrigger",
|
||||
"LocalTriggerStore",
|
||||
"TriggerDelivery",
|
||||
"TriggerDisabledError",
|
||||
"TriggerNotFoundError",
|
||||
"TriggerRunRecord",
|
||||
"TriggerStoreError",
|
||||
]
|
||||
@@ -0,0 +1,209 @@
|
||||
"""Gateway delivery loop for local triggers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import uuid
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Any
|
||||
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.agent.automation_turns import AutomationTurnError
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
from nanobot.triggers.local_session_turns import LOCAL_TRIGGER_META
|
||||
from nanobot.triggers.local_store import LocalTriggerStore
|
||||
from nanobot.triggers.local_types import LocalTrigger, TriggerDelivery
|
||||
from nanobot.webui.metadata import WEBUI_MESSAGE_SOURCE_METADATA_KEY, WEBUI_TURN_METADATA_KEY
|
||||
|
||||
|
||||
async def run_local_trigger_queue(
|
||||
*,
|
||||
store: LocalTriggerStore,
|
||||
submit_turn: Callable[[InboundMessage], Awaitable[OutboundMessage | None]] | None = None,
|
||||
poll_interval_s: float = 0.5,
|
||||
batch_size: int = 20,
|
||||
) -> None:
|
||||
"""Poll local trigger deliveries and submit them as session turns."""
|
||||
if submit_turn is None:
|
||||
raise ValueError("run_local_trigger_queue requires submit_turn")
|
||||
logger.info("Local trigger queue started")
|
||||
recovered = store.recover_processing_deliveries()
|
||||
if recovered:
|
||||
logger.warning(
|
||||
"Trigger: recovered {} interrupted delivery file(s) from processing",
|
||||
recovered,
|
||||
)
|
||||
while True:
|
||||
deliveries = store.claim_deliveries(limit=batch_size)
|
||||
if not deliveries:
|
||||
await asyncio.sleep(poll_interval_s)
|
||||
continue
|
||||
|
||||
for delivery in deliveries:
|
||||
try:
|
||||
await _deliver_delivery(
|
||||
store,
|
||||
delivery,
|
||||
submit_turn=submit_turn,
|
||||
)
|
||||
store.complete_delivery(delivery)
|
||||
except asyncio.CancelledError as exc:
|
||||
store.retry_delivery(delivery, str(exc) or exc.__class__.__name__)
|
||||
_write_delivery_run_record(
|
||||
store,
|
||||
delivery,
|
||||
status="interrupted",
|
||||
error=str(exc) or exc.__class__.__name__,
|
||||
)
|
||||
raise
|
||||
except _TerminalDeliveryError as exc:
|
||||
store.record_delivery(
|
||||
delivery.trigger_id,
|
||||
status="error",
|
||||
error=str(exc),
|
||||
run_at_ms=delivery.created_at_ms,
|
||||
)
|
||||
_write_delivery_run_record(
|
||||
store,
|
||||
delivery,
|
||||
status="error",
|
||||
error=str(exc),
|
||||
)
|
||||
store.complete_delivery(delivery)
|
||||
logger.warning(
|
||||
"Trigger: dropped delivery {} for {}: {}",
|
||||
delivery.id,
|
||||
delivery.trigger_id,
|
||||
exc,
|
||||
)
|
||||
except AutomationTurnError as exc:
|
||||
error = str(exc) or exc.__class__.__name__
|
||||
store.record_delivery(
|
||||
delivery.trigger_id,
|
||||
status="error",
|
||||
error=error,
|
||||
run_at_ms=delivery.created_at_ms,
|
||||
)
|
||||
_write_delivery_run_record(
|
||||
store,
|
||||
delivery,
|
||||
status="error",
|
||||
error=error,
|
||||
)
|
||||
store.complete_delivery(delivery)
|
||||
logger.warning(
|
||||
"Trigger: delivery {} for {} reached the agent but failed: {}",
|
||||
delivery.id,
|
||||
delivery.trigger_id,
|
||||
error,
|
||||
)
|
||||
except Exception as exc:
|
||||
error = str(exc) or exc.__class__.__name__
|
||||
retried = store.retry_delivery(delivery, error)
|
||||
_write_delivery_run_record(
|
||||
store,
|
||||
delivery,
|
||||
status="retrying" if retried else "error",
|
||||
error=error,
|
||||
)
|
||||
store.record_delivery(
|
||||
delivery.trigger_id,
|
||||
status="error",
|
||||
error=error,
|
||||
run_at_ms=delivery.created_at_ms,
|
||||
)
|
||||
logger.exception(
|
||||
"Trigger: failed delivery {} for {}{}",
|
||||
delivery.id,
|
||||
delivery.trigger_id,
|
||||
"; queued retry" if retried else "; moved to failed queue",
|
||||
)
|
||||
|
||||
|
||||
class _TerminalDeliveryError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
async def _deliver_delivery(
|
||||
store: LocalTriggerStore,
|
||||
delivery: TriggerDelivery,
|
||||
*,
|
||||
submit_turn: Callable[[InboundMessage], Awaitable[OutboundMessage | None]],
|
||||
) -> None:
|
||||
trigger = store.get(delivery.trigger_id)
|
||||
if trigger is None:
|
||||
raise _TerminalDeliveryError("trigger not found")
|
||||
if not trigger.enabled:
|
||||
raise _TerminalDeliveryError("trigger is disabled")
|
||||
|
||||
store.write_delivery_run_record(delivery, trigger=trigger, status="processing")
|
||||
msg = InboundMessage(
|
||||
channel=trigger.channel,
|
||||
sender_id=trigger.sender_id,
|
||||
chat_id=trigger.chat_id,
|
||||
content=delivery.content,
|
||||
metadata=_delivery_metadata(trigger, delivery),
|
||||
session_key_override=trigger.session_key,
|
||||
)
|
||||
response = await submit_turn(msg)
|
||||
store.record_delivery(
|
||||
trigger.id,
|
||||
status="ok",
|
||||
run_at_ms=delivery.created_at_ms,
|
||||
)
|
||||
_write_delivery_run_record(
|
||||
store,
|
||||
delivery,
|
||||
trigger=trigger,
|
||||
status="ok",
|
||||
response=response.content if response else "",
|
||||
)
|
||||
|
||||
|
||||
def _write_delivery_run_record(
|
||||
store: LocalTriggerStore,
|
||||
delivery: TriggerDelivery,
|
||||
*,
|
||||
status: str,
|
||||
trigger: LocalTrigger | None = None,
|
||||
error: str | None = None,
|
||||
response: str | None = None,
|
||||
) -> None:
|
||||
try:
|
||||
store.write_delivery_run_record(
|
||||
delivery,
|
||||
trigger=trigger,
|
||||
status=status,
|
||||
error=error,
|
||||
response=response,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Trigger: failed to write run record for delivery {}",
|
||||
delivery.id,
|
||||
)
|
||||
|
||||
|
||||
def _delivery_metadata(trigger: LocalTrigger, delivery: TriggerDelivery) -> dict[str, Any]:
|
||||
metadata = dict(trigger.origin_metadata or {})
|
||||
metadata[LOCAL_TRIGGER_META] = {
|
||||
"trigger_id": trigger.id,
|
||||
"trigger_name": trigger.name,
|
||||
"delivery_id": delivery.id,
|
||||
"created_at_ms": delivery.created_at_ms,
|
||||
"persist_content": _history_content(trigger, delivery),
|
||||
}
|
||||
if trigger.channel == "websocket":
|
||||
metadata.pop(WEBUI_TURN_METADATA_KEY, None)
|
||||
metadata[WEBUI_TURN_METADATA_KEY] = f"trigger:{trigger.id}:{uuid.uuid4().hex}"
|
||||
source: dict[str, str] = {"kind": "local_trigger"}
|
||||
if trigger.name:
|
||||
source["label"] = trigger.name
|
||||
metadata[WEBUI_MESSAGE_SOURCE_METADATA_KEY] = source
|
||||
return metadata
|
||||
|
||||
|
||||
def _history_content(trigger: LocalTrigger, delivery: TriggerDelivery) -> str:
|
||||
label = trigger.name.strip() if trigger.name else trigger.id
|
||||
return f"Local trigger received: {label}\n\n{delivery.content}"
|
||||
@@ -0,0 +1,62 @@
|
||||
"""Shared metadata helpers for local trigger session turns."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Mapping
|
||||
|
||||
from nanobot.session.automation_turns import (
|
||||
AutomationTurnSpec,
|
||||
automation_history_overrides_for_spec,
|
||||
automation_trigger,
|
||||
)
|
||||
|
||||
LOCAL_TRIGGER_META = "_local_trigger"
|
||||
|
||||
|
||||
def _local_trigger_history_text(trigger: Mapping[str, Any]) -> str:
|
||||
persist_content = trigger.get("persist_content")
|
||||
if isinstance(persist_content, str) and persist_content.strip():
|
||||
return persist_content
|
||||
name = trigger.get("trigger_name")
|
||||
trigger_id = trigger.get("trigger_id")
|
||||
label = name if isinstance(name, str) and name.strip() else trigger_id
|
||||
return (
|
||||
f"Local trigger received: {label}"
|
||||
if isinstance(label, str) and label.strip()
|
||||
else "Local trigger received"
|
||||
)
|
||||
|
||||
|
||||
LOCAL_TRIGGER_AUTOMATION_SPEC = AutomationTurnSpec(
|
||||
kind="local_trigger",
|
||||
trigger_meta_key=LOCAL_TRIGGER_META,
|
||||
history_fields={
|
||||
"trigger_id": "trigger_id",
|
||||
"trigger_name": "trigger_name",
|
||||
"trigger_delivery_id": "delivery_id",
|
||||
},
|
||||
text_builder=_local_trigger_history_text,
|
||||
)
|
||||
|
||||
|
||||
def local_trigger(metadata: Mapping[str, Any] | None) -> dict[str, Any] | None:
|
||||
"""Return structured local trigger metadata when present."""
|
||||
return automation_trigger(metadata, LOCAL_TRIGGER_AUTOMATION_SPEC)
|
||||
|
||||
|
||||
def local_trigger_delivery_id(metadata: Mapping[str, Any] | None) -> str | None:
|
||||
trigger = local_trigger(metadata)
|
||||
if not trigger:
|
||||
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,
|
||||
)
|
||||
@@ -0,0 +1,474 @@
|
||||
"""Workspace-scoped local trigger store and delivery queue."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import errno
|
||||
import json
|
||||
import os
|
||||
import secrets
|
||||
import time
|
||||
import uuid
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from filelock import FileLock
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.triggers.local_types import LocalTrigger, TriggerDelivery, TriggerRunRecord
|
||||
from nanobot.utils.helpers import truncate_text
|
||||
from nanobot.utils.run_records import write_run_record as write_automation_run_record
|
||||
|
||||
_TRIGGER_ID_ALPHABET = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
|
||||
_MAX_RUN_HISTORY = 20
|
||||
_MAX_DELIVERY_ATTEMPTS = 10
|
||||
_RUN_RECORD_TEXT_MAX_CHARS = 4000
|
||||
_PROCESSING_RECOVERY_ERROR = "delivery was recovered from interrupted processing"
|
||||
|
||||
|
||||
class TriggerStoreError(RuntimeError):
|
||||
"""Base class for trigger store errors."""
|
||||
|
||||
|
||||
class TriggerNotFoundError(TriggerStoreError):
|
||||
"""Raised when a trigger ID does not exist."""
|
||||
|
||||
|
||||
class TriggerDisabledError(TriggerStoreError):
|
||||
"""Raised when a trigger is disabled."""
|
||||
|
||||
|
||||
class LocalTriggerStore:
|
||||
"""Persistent local triggers for one workspace."""
|
||||
|
||||
def __init__(self, workspace_path: Path):
|
||||
self.workspace_path = Path(workspace_path)
|
||||
self.root = self.workspace_path / "triggers"
|
||||
self.store_path = self.root / "triggers.json"
|
||||
self.inbox_dir = self.root / "inbox"
|
||||
self.processing_dir = self.root / "processing"
|
||||
self.failed_dir = self.root / "failed"
|
||||
self.runs_dir = self.root / "runs"
|
||||
self._lock = FileLock(str(self.root / ".lock"))
|
||||
|
||||
def create(
|
||||
self,
|
||||
*,
|
||||
name: str,
|
||||
channel: str,
|
||||
chat_id: str,
|
||||
session_key: str,
|
||||
sender_id: str = "trigger",
|
||||
origin_metadata: dict[str, Any] | None = None,
|
||||
) -> LocalTrigger:
|
||||
"""Create a new session-bound local trigger."""
|
||||
clean_name = _clean_name(name)
|
||||
channel = channel.strip()
|
||||
chat_id = chat_id.strip()
|
||||
session_key = session_key.strip()
|
||||
if not channel or not chat_id or not session_key:
|
||||
raise ValueError("channel, chat_id, and session_key are required")
|
||||
|
||||
now = _now_ms()
|
||||
self._ensure_dirs()
|
||||
with self._lock:
|
||||
triggers = self._load_triggers_unlocked()
|
||||
existing_ids = {trigger.id for trigger in triggers}
|
||||
trigger_id = _new_trigger_id(existing_ids)
|
||||
trigger = LocalTrigger(
|
||||
id=trigger_id,
|
||||
name=clean_name,
|
||||
enabled=True,
|
||||
channel=channel,
|
||||
chat_id=chat_id,
|
||||
session_key=session_key,
|
||||
sender_id=sender_id.strip() or "trigger",
|
||||
origin_metadata=dict(origin_metadata or {}),
|
||||
created_at_ms=now,
|
||||
updated_at_ms=now,
|
||||
)
|
||||
triggers.append(trigger)
|
||||
self._save_triggers_unlocked(triggers)
|
||||
return trigger
|
||||
|
||||
def list_triggers(self, *, include_disabled: bool = False) -> list[LocalTrigger]:
|
||||
"""List triggers in this workspace."""
|
||||
self._ensure_dirs()
|
||||
with self._lock:
|
||||
triggers = self._load_triggers_unlocked()
|
||||
if not include_disabled:
|
||||
triggers = [trigger for trigger in triggers if trigger.enabled]
|
||||
return sorted(triggers, key=lambda trigger: (trigger.updated_at_ms, trigger.id), reverse=True)
|
||||
|
||||
def list_for_session(
|
||||
self,
|
||||
session_key: str,
|
||||
*,
|
||||
include_disabled: bool = True,
|
||||
) -> list[LocalTrigger]:
|
||||
"""List triggers bound to one session key."""
|
||||
return [
|
||||
trigger
|
||||
for trigger in self.list_triggers(include_disabled=include_disabled)
|
||||
if trigger.session_key == session_key
|
||||
]
|
||||
|
||||
def get(self, trigger_id: str) -> LocalTrigger | None:
|
||||
"""Return one trigger by ID."""
|
||||
self._ensure_dirs()
|
||||
with self._lock:
|
||||
return self._find_unlocked(self._load_triggers_unlocked(), trigger_id)
|
||||
|
||||
def enable(self, trigger_id: str, *, enabled: bool) -> LocalTrigger | None:
|
||||
"""Enable or disable a trigger."""
|
||||
self._ensure_dirs()
|
||||
with self._lock:
|
||||
triggers = self._load_triggers_unlocked()
|
||||
trigger = self._find_unlocked(triggers, trigger_id)
|
||||
if trigger is None:
|
||||
return None
|
||||
trigger.enabled = enabled
|
||||
trigger.updated_at_ms = _now_ms()
|
||||
self._save_triggers_unlocked(triggers)
|
||||
return trigger
|
||||
|
||||
def update(self, trigger_id: str, *, name: str | None = None) -> LocalTrigger | None:
|
||||
"""Update mutable trigger fields."""
|
||||
self._ensure_dirs()
|
||||
with self._lock:
|
||||
triggers = self._load_triggers_unlocked()
|
||||
trigger = self._find_unlocked(triggers, trigger_id)
|
||||
if trigger is None:
|
||||
return None
|
||||
if name is not None:
|
||||
trigger.name = _clean_name(name)
|
||||
trigger.updated_at_ms = _now_ms()
|
||||
self._save_triggers_unlocked(triggers)
|
||||
return trigger
|
||||
|
||||
def delete(self, trigger_id: str) -> bool:
|
||||
"""Delete a trigger by ID."""
|
||||
trigger_id = trigger_id.strip()
|
||||
self._ensure_dirs()
|
||||
with self._lock:
|
||||
triggers = self._load_triggers_unlocked()
|
||||
remaining = [trigger for trigger in triggers if trigger.id != trigger_id]
|
||||
if len(remaining) == len(triggers):
|
||||
return False
|
||||
self._save_triggers_unlocked(remaining)
|
||||
self._delete_delivery_files_for_trigger_unlocked(trigger_id)
|
||||
return True
|
||||
|
||||
def enqueue(self, trigger_id: str, content: str) -> TriggerDelivery:
|
||||
"""Queue a delivery for the gateway process to consume."""
|
||||
trigger_id = trigger_id.strip()
|
||||
if not content.strip():
|
||||
raise ValueError("trigger message is required")
|
||||
self._ensure_dirs()
|
||||
with self._lock:
|
||||
trigger = self._find_unlocked(self._load_triggers_unlocked(), trigger_id)
|
||||
if trigger is None:
|
||||
raise TriggerNotFoundError(f"trigger not found: {trigger_id}")
|
||||
if not trigger.enabled:
|
||||
raise TriggerDisabledError(f"trigger is disabled: {trigger_id}")
|
||||
delivery = TriggerDelivery(
|
||||
id=f"tdl_{uuid.uuid4().hex[:12]}",
|
||||
trigger_id=trigger_id,
|
||||
content=content,
|
||||
created_at_ms=_now_ms(),
|
||||
)
|
||||
path = self.inbox_dir / f"{delivery.created_at_ms}-{delivery.id}.json"
|
||||
self._atomic_write(path, json.dumps(_delivery_payload(delivery), ensure_ascii=False))
|
||||
delivery.path = path
|
||||
try:
|
||||
self.write_delivery_run_record(delivery, trigger=trigger, status="queued")
|
||||
except BaseException:
|
||||
path.unlink(missing_ok=True)
|
||||
delivery.path = None
|
||||
raise
|
||||
return delivery
|
||||
|
||||
def claim_deliveries(self, *, limit: int = 20) -> list[TriggerDelivery]:
|
||||
"""Move pending deliveries into processing and return them."""
|
||||
self._ensure_dirs()
|
||||
claimed: list[TriggerDelivery] = []
|
||||
with self._lock:
|
||||
for path in sorted(self.inbox_dir.glob("*.json"))[: max(0, limit)]:
|
||||
try:
|
||||
data = json.loads(path.read_text(encoding="utf-8"))
|
||||
delivery = TriggerDelivery.from_dict(
|
||||
data.get("delivery", data),
|
||||
path=self.processing_dir / path.name,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Trigger: failed to parse delivery {}", path)
|
||||
self._move_bad_delivery_unlocked(path)
|
||||
continue
|
||||
os.replace(path, delivery.path)
|
||||
claimed.append(delivery)
|
||||
return claimed
|
||||
|
||||
def recover_processing_deliveries(self) -> int:
|
||||
"""Requeue deliveries left in processing by an interrupted gateway."""
|
||||
self._ensure_dirs()
|
||||
recovered = 0
|
||||
with self._lock:
|
||||
for path in sorted(self.processing_dir.glob("*.json")):
|
||||
try:
|
||||
data = json.loads(path.read_text(encoding="utf-8"))
|
||||
delivery = TriggerDelivery.from_dict(
|
||||
data.get("delivery", data),
|
||||
path=path,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Trigger: failed to parse processing delivery {}", path)
|
||||
self._move_bad_delivery_unlocked(path)
|
||||
continue
|
||||
if self._retry_delivery_unlocked(delivery, _PROCESSING_RECOVERY_ERROR):
|
||||
recovered += 1
|
||||
return recovered
|
||||
|
||||
def complete_delivery(self, delivery: TriggerDelivery) -> None:
|
||||
"""Delete a claimed delivery after it is handled."""
|
||||
if delivery.path is None:
|
||||
return
|
||||
self._ensure_dirs()
|
||||
with self._lock:
|
||||
delivery.path.unlink(missing_ok=True)
|
||||
|
||||
def retry_delivery(self, delivery: TriggerDelivery, error: str) -> bool:
|
||||
"""Retry a claimed delivery unless it exceeded the attempt limit."""
|
||||
if delivery.path is None:
|
||||
return False
|
||||
self._ensure_dirs()
|
||||
with self._lock:
|
||||
return self._retry_delivery_unlocked(delivery, error)
|
||||
|
||||
def record_delivery(
|
||||
self,
|
||||
trigger_id: str,
|
||||
*,
|
||||
status: str,
|
||||
error: str | None = None,
|
||||
run_at_ms: int | None = None,
|
||||
) -> None:
|
||||
"""Record the latest delivery status on a trigger."""
|
||||
self._ensure_dirs()
|
||||
run_at_ms = run_at_ms or _now_ms()
|
||||
with self._lock:
|
||||
triggers = self._load_triggers_unlocked()
|
||||
trigger = self._find_unlocked(triggers, trigger_id)
|
||||
if trigger is None:
|
||||
return
|
||||
trigger.last_run_at_ms = run_at_ms
|
||||
trigger.last_status = "ok" if status == "ok" else "error"
|
||||
trigger.last_error = None if status == "ok" else (error or "delivery failed")
|
||||
trigger.updated_at_ms = _now_ms()
|
||||
trigger.run_history.append(
|
||||
TriggerRunRecord(
|
||||
run_at_ms=run_at_ms,
|
||||
status=trigger.last_status,
|
||||
error=trigger.last_error,
|
||||
)
|
||||
)
|
||||
trigger.run_history = trigger.run_history[-_MAX_RUN_HISTORY:]
|
||||
self._save_triggers_unlocked(triggers)
|
||||
|
||||
def write_run_record(self, run_id: str, record: dict[str, Any]) -> Path:
|
||||
"""Write an internal audit record for one local trigger delivery."""
|
||||
self._ensure_dirs()
|
||||
return write_automation_run_record(self.runs_dir, run_id, record)
|
||||
|
||||
def write_delivery_run_record(
|
||||
self,
|
||||
delivery: TriggerDelivery,
|
||||
*,
|
||||
status: str,
|
||||
trigger: LocalTrigger | None = None,
|
||||
error: str | None = None,
|
||||
response: str | None = None,
|
||||
) -> Path:
|
||||
"""Write the durable audit record for one local trigger delivery."""
|
||||
if trigger is None:
|
||||
trigger = self.get(delivery.trigger_id)
|
||||
record = _delivery_run_record(delivery, trigger)
|
||||
record["status"] = status
|
||||
if error:
|
||||
record["error"] = _run_record_text(error)
|
||||
if response is not None:
|
||||
record["response"] = _run_record_text(response)
|
||||
return self.write_run_record(delivery.id, record)
|
||||
|
||||
def _ensure_dirs(self) -> None:
|
||||
self.root.mkdir(parents=True, exist_ok=True)
|
||||
self.inbox_dir.mkdir(parents=True, exist_ok=True)
|
||||
self.processing_dir.mkdir(parents=True, exist_ok=True)
|
||||
self.failed_dir.mkdir(parents=True, exist_ok=True)
|
||||
self.runs_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
def _load_triggers_unlocked(self) -> list[LocalTrigger]:
|
||||
if not self.store_path.exists():
|
||||
return []
|
||||
try:
|
||||
data = json.loads(self.store_path.read_text(encoding="utf-8"))
|
||||
return [
|
||||
LocalTrigger.from_dict(raw)
|
||||
for raw in data.get("triggers", [])
|
||||
if isinstance(raw, dict)
|
||||
]
|
||||
except Exception as exc:
|
||||
backup = self.store_path.with_suffix(
|
||||
self.store_path.suffix + f".corrupt-{int(time.time())}"
|
||||
)
|
||||
with suppress(OSError):
|
||||
os.replace(self.store_path, backup)
|
||||
raise TriggerStoreError(
|
||||
f"trigger store at {self.store_path} could not be loaded and was preserved "
|
||||
"as a .corrupt-<ts> backup"
|
||||
) from exc
|
||||
|
||||
def _save_triggers_unlocked(self, triggers: list[LocalTrigger]) -> None:
|
||||
payload = {
|
||||
"version": 1,
|
||||
"triggers": [trigger.to_dict() for trigger in triggers],
|
||||
}
|
||||
self._atomic_write(self.store_path, json.dumps(payload, indent=2, ensure_ascii=False))
|
||||
|
||||
@staticmethod
|
||||
def _find_unlocked(
|
||||
triggers: list[LocalTrigger],
|
||||
trigger_id: str,
|
||||
) -> LocalTrigger | None:
|
||||
return next((trigger for trigger in triggers if trigger.id == trigger_id), None)
|
||||
|
||||
def _move_bad_delivery_unlocked(self, path: Path) -> None:
|
||||
target = self.failed_dir / f"{path.name}.bad"
|
||||
with suppress(OSError):
|
||||
os.replace(path, target)
|
||||
|
||||
def _retry_delivery_unlocked(self, delivery: TriggerDelivery, error: str) -> bool:
|
||||
if delivery.path is None:
|
||||
return False
|
||||
if delivery.attempts + 1 >= _MAX_DELIVERY_ATTEMPTS:
|
||||
delivery.attempts += 1
|
||||
delivery.last_error = error
|
||||
failed = self.failed_dir / delivery.path.name
|
||||
self._atomic_write(failed, json.dumps(_delivery_payload(delivery), ensure_ascii=False))
|
||||
delivery.path.unlink(missing_ok=True)
|
||||
return False
|
||||
delivery.attempts += 1
|
||||
delivery.last_error = error
|
||||
target = self.inbox_dir / delivery.path.name
|
||||
self._atomic_write(target, json.dumps(_delivery_payload(delivery), ensure_ascii=False))
|
||||
delivery.path.unlink(missing_ok=True)
|
||||
return True
|
||||
|
||||
def _delete_delivery_files_for_trigger_unlocked(self, trigger_id: str) -> None:
|
||||
for directory in (self.inbox_dir, self.processing_dir, self.failed_dir):
|
||||
for path in directory.iterdir():
|
||||
if not path.is_file():
|
||||
continue
|
||||
if self._delivery_file_trigger_id(path) != trigger_id:
|
||||
continue
|
||||
try:
|
||||
path.unlink(missing_ok=True)
|
||||
except OSError as exc:
|
||||
logger.warning(
|
||||
"Trigger: failed to delete delivery file {} for deleted trigger {}: {}",
|
||||
path,
|
||||
trigger_id,
|
||||
exc,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _delivery_file_trigger_id(path: Path) -> str | None:
|
||||
try:
|
||||
data = json.loads(path.read_text(encoding="utf-8"))
|
||||
except Exception:
|
||||
return None
|
||||
raw = data.get("delivery", data) if isinstance(data, dict) else None
|
||||
if not isinstance(raw, dict):
|
||||
return None
|
||||
trigger_id = raw.get("triggerId", raw.get("trigger_id", ""))
|
||||
return str(trigger_id) if trigger_id else None
|
||||
|
||||
@staticmethod
|
||||
def _atomic_write(path: Path, content: str) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
tmp_path = path.with_name(f".{path.name}.{uuid.uuid4().hex}.tmp")
|
||||
try:
|
||||
with open(tmp_path, "w", encoding="utf-8") as f:
|
||||
f.write(content)
|
||||
f.flush()
|
||||
os.fsync(f.fileno())
|
||||
os.replace(tmp_path, path)
|
||||
with suppress(PermissionError):
|
||||
fd = os.open(str(path.parent), os.O_RDONLY)
|
||||
try:
|
||||
try:
|
||||
os.fsync(fd)
|
||||
except OSError as exc:
|
||||
if exc.errno != errno.EINVAL:
|
||||
raise
|
||||
finally:
|
||||
os.close(fd)
|
||||
except BaseException:
|
||||
tmp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
|
||||
|
||||
def _new_trigger_id(existing_ids: set[str]) -> str:
|
||||
for _ in range(100):
|
||||
suffix = "".join(secrets.choice(_TRIGGER_ID_ALPHABET) for _ in range(8))
|
||||
candidate = f"trg_{suffix}"
|
||||
if candidate not in existing_ids:
|
||||
return candidate
|
||||
raise TriggerStoreError("could not allocate a unique trigger id")
|
||||
|
||||
|
||||
def _clean_name(name: str) -> str:
|
||||
stripped = " ".join(name.strip().split())
|
||||
return (stripped or "Local trigger")[:120]
|
||||
|
||||
|
||||
def _now_ms() -> int:
|
||||
return int(time.time() * 1000)
|
||||
|
||||
|
||||
def _delivery_payload(delivery: TriggerDelivery) -> dict[str, Any]:
|
||||
return {
|
||||
"version": 1,
|
||||
"delivery": delivery.to_dict(),
|
||||
}
|
||||
|
||||
|
||||
def _delivery_run_record(
|
||||
delivery: TriggerDelivery,
|
||||
trigger: LocalTrigger | None,
|
||||
) -> dict[str, Any]:
|
||||
record: dict[str, Any] = {
|
||||
"kind": "local_trigger",
|
||||
"trigger_id": delivery.trigger_id,
|
||||
"delivery_id": delivery.id,
|
||||
"content": _run_record_text(delivery.content),
|
||||
"created_at_ms": delivery.created_at_ms,
|
||||
"attempts": delivery.attempts,
|
||||
}
|
||||
if delivery.last_error:
|
||||
record["last_error"] = _run_record_text(delivery.last_error)
|
||||
if trigger is not None:
|
||||
record.update(
|
||||
{
|
||||
"trigger_name": trigger.name,
|
||||
"session_key": trigger.session_key,
|
||||
"channel": trigger.channel,
|
||||
"chat_id": trigger.chat_id,
|
||||
"sender_id": trigger.sender_id,
|
||||
"origin_metadata": trigger.origin_metadata,
|
||||
}
|
||||
)
|
||||
return record
|
||||
|
||||
|
||||
def _run_record_text(value: str) -> str:
|
||||
return truncate_text(value, _RUN_RECORD_TEXT_MAX_CHARS)
|
||||
@@ -0,0 +1,55 @@
|
||||
"""Coordination for local trigger turns."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable, Iterable
|
||||
|
||||
from nanobot.agent.automation_turns import AutomationTurnCoordinator
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.triggers.local_session_turns import local_trigger, local_trigger_delivery_id
|
||||
|
||||
|
||||
class LocalTriggerTurnCoordinator(AutomationTurnCoordinator):
|
||||
"""Manage local trigger turns without mixing them into live injections."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
publish_inbound: Callable[[InboundMessage], Awaitable[None]],
|
||||
dispatch: Callable[[InboundMessage], Awaitable[object]],
|
||||
is_running: Callable[[], bool],
|
||||
deferred_queues: dict[str, list[InboundMessage]] | None = None,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
publish_inbound=publish_inbound,
|
||||
dispatch=dispatch,
|
||||
is_running=is_running,
|
||||
turn_id=lambda msg: local_trigger_delivery_id(msg.metadata),
|
||||
pending_id=_local_trigger_id,
|
||||
should_defer_turn=_should_defer_local_trigger_turn,
|
||||
missing_id_error="local trigger turn metadata must include a delivery_id",
|
||||
duplicate_id_error=lambda delivery_id: (
|
||||
f"local trigger delivery {delivery_id!r} is already pending"
|
||||
),
|
||||
deferred_queues=deferred_queues,
|
||||
)
|
||||
|
||||
def pending_trigger_ids_for_session(self, session_key: str) -> set[str]:
|
||||
"""Return local triggers waiting for or running in *session_key*."""
|
||||
return self.pending_ids_for_session(session_key)
|
||||
|
||||
|
||||
def _should_defer_local_trigger_turn(
|
||||
msg: InboundMessage,
|
||||
session_key: str,
|
||||
active_session_keys: Iterable[str],
|
||||
) -> bool:
|
||||
return local_trigger(msg.metadata) is not None and session_key in active_session_keys
|
||||
|
||||
|
||||
def _local_trigger_id(msg: InboundMessage) -> str | None:
|
||||
trigger = local_trigger(msg.metadata)
|
||||
if not trigger:
|
||||
return None
|
||||
value = trigger.get("trigger_id")
|
||||
return value if isinstance(value, str) and value else None
|
||||
@@ -0,0 +1,141 @@
|
||||
"""Persistent types for local triggers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
|
||||
TriggerStatus = Literal["ok", "error"]
|
||||
|
||||
|
||||
def _get(data: dict[str, Any], camel: str, snake: str, default: Any = None) -> Any:
|
||||
if camel in data:
|
||||
return data[camel]
|
||||
return data.get(snake, default)
|
||||
|
||||
|
||||
@dataclass
|
||||
class TriggerRunRecord:
|
||||
"""A single local trigger delivery record."""
|
||||
|
||||
run_at_ms: int
|
||||
status: TriggerStatus
|
||||
error: str | None = None
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any]) -> "TriggerRunRecord":
|
||||
return cls(
|
||||
run_at_ms=int(_get(data, "runAtMs", "run_at_ms", 0)),
|
||||
status=str(data.get("status") or "error"), # type: ignore[arg-type]
|
||||
error=data.get("error"),
|
||||
)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"runAtMs": self.run_at_ms,
|
||||
"status": self.status,
|
||||
"error": self.error,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class LocalTrigger:
|
||||
"""A session-bound local trigger."""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
enabled: bool
|
||||
channel: str
|
||||
chat_id: str
|
||||
session_key: str
|
||||
sender_id: str = "trigger"
|
||||
origin_metadata: dict[str, Any] = field(default_factory=dict)
|
||||
created_at_ms: int = 0
|
||||
updated_at_ms: int = 0
|
||||
last_run_at_ms: int | None = None
|
||||
last_status: TriggerStatus | None = None
|
||||
last_error: str | None = None
|
||||
run_history: list[TriggerRunRecord] = field(default_factory=list)
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any]) -> "LocalTrigger":
|
||||
history = [
|
||||
record if isinstance(record, TriggerRunRecord) else TriggerRunRecord.from_dict(record)
|
||||
for record in data.get("runHistory", data.get("run_history", []))
|
||||
if isinstance(record, (dict, TriggerRunRecord))
|
||||
]
|
||||
return cls(
|
||||
id=str(data["id"]),
|
||||
name=str(data.get("name") or data["id"]),
|
||||
enabled=bool(data.get("enabled", True)),
|
||||
channel=str(data.get("channel") or ""),
|
||||
chat_id=str(_get(data, "chatId", "chat_id", "")),
|
||||
session_key=str(_get(data, "sessionKey", "session_key", "")),
|
||||
sender_id=str(_get(data, "senderId", "sender_id", "trigger") or "trigger"),
|
||||
origin_metadata=dict(_get(data, "originMetadata", "origin_metadata", {}) or {}),
|
||||
created_at_ms=int(_get(data, "createdAtMs", "created_at_ms", 0)),
|
||||
updated_at_ms=int(_get(data, "updatedAtMs", "updated_at_ms", 0)),
|
||||
last_run_at_ms=_get(data, "lastRunAtMs", "last_run_at_ms"),
|
||||
last_status=_get(data, "lastStatus", "last_status"), # type: ignore[arg-type]
|
||||
last_error=_get(data, "lastError", "last_error"),
|
||||
run_history=history,
|
||||
)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"id": self.id,
|
||||
"name": self.name,
|
||||
"enabled": self.enabled,
|
||||
"channel": self.channel,
|
||||
"chatId": self.chat_id,
|
||||
"sessionKey": self.session_key,
|
||||
"senderId": self.sender_id,
|
||||
"originMetadata": self.origin_metadata,
|
||||
"createdAtMs": self.created_at_ms,
|
||||
"updatedAtMs": self.updated_at_ms,
|
||||
"lastRunAtMs": self.last_run_at_ms,
|
||||
"lastStatus": self.last_status,
|
||||
"lastError": self.last_error,
|
||||
"runHistory": [record.to_dict() for record in self.run_history],
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class TriggerDelivery:
|
||||
"""One pending local trigger delivery written by the CLI."""
|
||||
|
||||
id: str
|
||||
trigger_id: str
|
||||
content: str
|
||||
created_at_ms: int
|
||||
attempts: int = 0
|
||||
last_error: str | None = None
|
||||
path: Path | None = field(default=None, compare=False, repr=False)
|
||||
|
||||
@classmethod
|
||||
def from_dict(
|
||||
cls,
|
||||
data: dict[str, Any],
|
||||
*,
|
||||
path: Path | None = None,
|
||||
) -> "TriggerDelivery":
|
||||
return cls(
|
||||
id=str(data["id"]),
|
||||
trigger_id=str(_get(data, "triggerId", "trigger_id", "")),
|
||||
content=str(data.get("content") or ""),
|
||||
created_at_ms=int(_get(data, "createdAtMs", "created_at_ms", 0)),
|
||||
attempts=int(data.get("attempts", 0)),
|
||||
last_error=data.get("lastError") or data.get("last_error"),
|
||||
path=path,
|
||||
)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"id": self.id,
|
||||
"triggerId": self.trigger_id,
|
||||
"content": self.content,
|
||||
"createdAtMs": self.created_at_ms,
|
||||
"attempts": self.attempts,
|
||||
"lastError": self.last_error,
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
"""Durable JSON run records for automation executions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import errno
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from contextlib import suppress
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
|
||||
def safe_run_record_name(run_id: str) -> str:
|
||||
"""Return a filesystem-safe filename stem for a run ID."""
|
||||
return "".join(c if c.isalnum() or c in "._-" else "_" for c in run_id)
|
||||
|
||||
|
||||
def write_run_record(runs_dir: Path, run_id: str, record: dict[str, Any]) -> Path:
|
||||
"""Write or replace one durable automation run audit record."""
|
||||
name = safe_run_record_name(run_id) or str(uuid.uuid4())
|
||||
path = runs_dir / f"{name}.json"
|
||||
payload = {
|
||||
**record,
|
||||
"run_id": run_id,
|
||||
"updated_at_ms": _now_ms(),
|
||||
}
|
||||
_atomic_write(path, json.dumps(payload, indent=2, ensure_ascii=False))
|
||||
return path
|
||||
|
||||
|
||||
def _now_ms() -> int:
|
||||
return int(time.time() * 1000)
|
||||
|
||||
|
||||
def _atomic_write(path: Path, content: str) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
tmp_path = path.with_name(f".{path.name}.{uuid.uuid4().hex}.tmp")
|
||||
try:
|
||||
with open(tmp_path, "w", encoding="utf-8") as f:
|
||||
f.write(content)
|
||||
f.flush()
|
||||
os.fsync(f.fileno())
|
||||
os.replace(tmp_path, path)
|
||||
with suppress(PermissionError):
|
||||
fd = os.open(str(path.parent), os.O_RDONLY)
|
||||
try:
|
||||
try:
|
||||
os.fsync(fd)
|
||||
except OSError as exc:
|
||||
if exc.errno != errno.EINVAL:
|
||||
raise
|
||||
finally:
|
||||
os.close(fd)
|
||||
except BaseException:
|
||||
tmp_path.unlink(missing_ok=True)
|
||||
raise
|
||||
@@ -26,7 +26,9 @@ class GatewayServices:
|
||||
workspaces: WebUIWorkspaceController
|
||||
session_manager: Any | None
|
||||
cron_service: Any | None
|
||||
local_trigger_store: Any | None
|
||||
cron_pending_job_ids: Callable[[str], set[str]] | None
|
||||
local_trigger_pending_ids: Callable[[str], set[str]] | None
|
||||
|
||||
|
||||
def build_gateway_services(
|
||||
@@ -42,7 +44,9 @@ def build_gateway_services(
|
||||
runtime_capabilities_overrides: dict[str, Any] | None,
|
||||
disabled_skills: set[str] | None = None,
|
||||
cron_service: Any | None = None,
|
||||
local_trigger_store: Any | None = None,
|
||||
cron_pending_job_ids: Callable[[str], set[str]] | None = None,
|
||||
local_trigger_pending_ids: Callable[[str], set[str]] | None = None,
|
||||
logger: Any = default_logger,
|
||||
) -> GatewayServices:
|
||||
tokens = GatewayTokenStore()
|
||||
@@ -70,7 +74,9 @@ def build_gateway_services(
|
||||
skills_workspace_path=workspace_path,
|
||||
disabled_skills=disabled_skills,
|
||||
cron_service=cron_service,
|
||||
local_trigger_store=local_trigger_store,
|
||||
cron_pending_job_ids=cron_pending_job_ids,
|
||||
local_trigger_pending_ids=local_trigger_pending_ids,
|
||||
log=logger,
|
||||
)
|
||||
return GatewayServices(
|
||||
@@ -81,5 +87,7 @@ def build_gateway_services(
|
||||
workspaces=workspaces,
|
||||
session_manager=session_manager,
|
||||
cron_service=cron_service,
|
||||
local_trigger_store=local_trigger_store,
|
||||
cron_pending_job_ids=cron_pending_job_ids,
|
||||
local_trigger_pending_ids=local_trigger_pending_ids,
|
||||
)
|
||||
|
||||
@@ -5,9 +5,12 @@ from __future__ import annotations
|
||||
from collections.abc import Collection
|
||||
from typing import Any, Protocol
|
||||
|
||||
from nanobot.cron.session_turns import CRON_HISTORY_META
|
||||
from nanobot.cron.types import CronJob
|
||||
from nanobot.session.history_visibility import is_hidden_history_message
|
||||
from nanobot.session.manager import _message_preview_text
|
||||
from nanobot.triggers.local_types import LocalTrigger
|
||||
|
||||
AutomationJob = CronJob | LocalTrigger
|
||||
|
||||
|
||||
class _CronServiceLike(Protocol):
|
||||
@@ -21,6 +24,17 @@ class _CronServiceLike(Protocol):
|
||||
) -> list[CronJob]: ...
|
||||
|
||||
|
||||
class _LocalTriggerStoreLike(Protocol):
|
||||
def list_triggers(self, *, include_disabled: bool = False) -> list[LocalTrigger]: ...
|
||||
|
||||
def list_for_session(
|
||||
self,
|
||||
session_key: str,
|
||||
*,
|
||||
include_disabled: bool = True,
|
||||
) -> list[LocalTrigger]: ...
|
||||
|
||||
|
||||
class _SessionManagerLike(Protocol):
|
||||
def read_session_file(self, key: str) -> dict[str, Any] | None: ...
|
||||
|
||||
@@ -28,26 +42,43 @@ class _SessionManagerLike(Protocol):
|
||||
def session_automation_jobs(
|
||||
cron_service: _CronServiceLike | None,
|
||||
session_key: str,
|
||||
) -> list[CronJob]:
|
||||
*,
|
||||
local_trigger_store: _LocalTriggerStoreLike | None = None,
|
||||
) -> list[AutomationJob]:
|
||||
"""Return user automations attached to the WebUI session."""
|
||||
if cron_service is None:
|
||||
return []
|
||||
return cron_service.list_bound_cron_jobs_for_session(
|
||||
session_key,
|
||||
include_disabled=True,
|
||||
)
|
||||
jobs: list[AutomationJob] = []
|
||||
if cron_service is not None:
|
||||
jobs.extend(
|
||||
cron_service.list_bound_cron_jobs_for_session(
|
||||
session_key,
|
||||
include_disabled=True,
|
||||
)
|
||||
)
|
||||
if local_trigger_store is not None:
|
||||
jobs.extend(
|
||||
local_trigger_store.list_for_session(
|
||||
session_key,
|
||||
include_disabled=True,
|
||||
)
|
||||
)
|
||||
return jobs
|
||||
|
||||
|
||||
def session_automations_payload(
|
||||
cron_service: _CronServiceLike | None,
|
||||
session_key: str,
|
||||
*,
|
||||
local_trigger_store: _LocalTriggerStoreLike | None = None,
|
||||
pending_job_ids: Collection[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Return user-created automation jobs attached to a WebUI session."""
|
||||
return {
|
||||
"jobs": serialize_automation_jobs(
|
||||
session_automation_jobs(cron_service, session_key),
|
||||
session_automation_jobs(
|
||||
cron_service,
|
||||
session_key,
|
||||
local_trigger_store=local_trigger_store,
|
||||
),
|
||||
pending_job_ids=pending_job_ids,
|
||||
)
|
||||
}
|
||||
@@ -56,11 +87,16 @@ def session_automations_payload(
|
||||
def all_automations_payload(
|
||||
cron_service: _CronServiceLike | None,
|
||||
*,
|
||||
local_trigger_store: _LocalTriggerStoreLike | None = None,
|
||||
session_manager: _SessionManagerLike | None = None,
|
||||
pending_job_ids: Collection[str] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Return all cron jobs visible to the WebUI automation manager."""
|
||||
jobs = cron_service.list_jobs(include_disabled=True) if cron_service is not None else []
|
||||
jobs: list[AutomationJob] = []
|
||||
if cron_service is not None:
|
||||
jobs.extend(cron_service.list_jobs(include_disabled=True))
|
||||
if local_trigger_store is not None:
|
||||
jobs.extend(local_trigger_store.list_triggers(include_disabled=True))
|
||||
return {
|
||||
"jobs": serialize_automation_jobs(
|
||||
jobs,
|
||||
@@ -72,7 +108,7 @@ def all_automations_payload(
|
||||
|
||||
|
||||
def serialize_automation_jobs(
|
||||
jobs: list[CronJob],
|
||||
jobs: list[AutomationJob],
|
||||
*,
|
||||
pending_job_ids: Collection[str] | None = None,
|
||||
include_details: bool = False,
|
||||
@@ -90,12 +126,20 @@ def serialize_automation_jobs(
|
||||
|
||||
|
||||
def _serialize_job(
|
||||
job: CronJob,
|
||||
job: AutomationJob,
|
||||
*,
|
||||
pending: bool = False,
|
||||
include_details: bool = False,
|
||||
session_manager: _SessionManagerLike | None = None,
|
||||
) -> dict[str, Any]:
|
||||
if isinstance(job, LocalTrigger):
|
||||
return _serialize_trigger(
|
||||
job,
|
||||
pending=pending,
|
||||
include_details=include_details,
|
||||
session_manager=session_manager,
|
||||
)
|
||||
|
||||
payload = {
|
||||
"id": job.id,
|
||||
"name": job.name,
|
||||
@@ -143,6 +187,67 @@ def _serialize_job(
|
||||
return payload
|
||||
|
||||
|
||||
def _serialize_trigger(
|
||||
trigger: LocalTrigger,
|
||||
*,
|
||||
pending: bool = False,
|
||||
include_details: bool = False,
|
||||
session_manager: _SessionManagerLike | None = None,
|
||||
) -> dict[str, Any]:
|
||||
command = f'nanobot trigger {trigger.id} "message"'
|
||||
payload = {
|
||||
"id": trigger.id,
|
||||
"name": trigger.name,
|
||||
"enabled": trigger.enabled,
|
||||
"kind": "local_trigger",
|
||||
"schedule": {
|
||||
"kind": "local",
|
||||
"at_ms": None,
|
||||
"every_ms": None,
|
||||
"expr": None,
|
||||
"tz": None,
|
||||
},
|
||||
"payload": {
|
||||
"kind": "local_trigger",
|
||||
"message": command,
|
||||
"command": command,
|
||||
},
|
||||
"state": {
|
||||
"next_run_at_ms": None,
|
||||
"last_status": trigger.last_status,
|
||||
"pending": pending,
|
||||
},
|
||||
}
|
||||
if not include_details:
|
||||
return payload
|
||||
|
||||
payload["protected"] = False
|
||||
payload["delete_after_run"] = False
|
||||
payload["created_at_ms"] = trigger.created_at_ms
|
||||
payload["updated_at_ms"] = trigger.updated_at_ms
|
||||
payload["state"].update(
|
||||
{
|
||||
"last_run_at_ms": trigger.last_run_at_ms,
|
||||
"last_error": trigger.last_error,
|
||||
"run_history": [
|
||||
{
|
||||
"run_at_ms": record.run_at_ms,
|
||||
"status": record.status,
|
||||
"duration_ms": 0,
|
||||
"error": record.error,
|
||||
}
|
||||
for record in trigger.run_history[-5:]
|
||||
],
|
||||
}
|
||||
)
|
||||
payload["origin"] = _trigger_origin_payload(trigger, session_manager)
|
||||
payload["trigger"] = {
|
||||
"id": trigger.id,
|
||||
"command": command,
|
||||
}
|
||||
return payload
|
||||
|
||||
|
||||
def _origin_payload(
|
||||
job: CronJob,
|
||||
session_manager: _SessionManagerLike | None,
|
||||
@@ -161,6 +266,46 @@ def _origin_payload(
|
||||
}
|
||||
|
||||
session_key = f"{channel}:{chat_id}"
|
||||
return _websocket_origin_payload(
|
||||
session_key=session_key,
|
||||
channel=channel,
|
||||
chat_id=chat_id,
|
||||
session_manager=session_manager,
|
||||
)
|
||||
|
||||
|
||||
def _trigger_origin_payload(
|
||||
trigger: LocalTrigger,
|
||||
session_manager: _SessionManagerLike | None,
|
||||
) -> dict[str, Any] | None:
|
||||
channel = trigger.channel
|
||||
chat_id = trigger.chat_id
|
||||
if not channel or not chat_id:
|
||||
return None
|
||||
if channel != "websocket":
|
||||
return {
|
||||
"channel": channel,
|
||||
"title": "",
|
||||
"preview": "",
|
||||
}
|
||||
|
||||
return _websocket_origin_payload(
|
||||
session_key=trigger.session_key or f"{channel}:{chat_id}",
|
||||
channel=channel,
|
||||
chat_id=chat_id,
|
||||
session_manager=session_manager,
|
||||
)
|
||||
|
||||
|
||||
def _websocket_origin_payload(
|
||||
*,
|
||||
session_key: str,
|
||||
channel: str,
|
||||
chat_id: str,
|
||||
session_manager: _SessionManagerLike | None,
|
||||
) -> dict[str, Any]:
|
||||
title = ""
|
||||
preview = ""
|
||||
if session_manager is not None:
|
||||
data = session_manager.read_session_file(session_key)
|
||||
if isinstance(data, dict):
|
||||
@@ -183,7 +328,7 @@ def _session_preview(messages: Any) -> str:
|
||||
for message in messages:
|
||||
if not isinstance(message, dict):
|
||||
continue
|
||||
if message.get(CRON_HISTORY_META) is True:
|
||||
if is_hidden_history_message(message):
|
||||
continue
|
||||
text = _message_preview_text(message)
|
||||
if not text:
|
||||
|
||||
@@ -16,7 +16,7 @@ from typing import Any
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.config.paths import get_webui_dir
|
||||
from nanobot.cron.session_turns import CRON_HISTORY_META
|
||||
from nanobot.session.history_visibility import is_hidden_history_message
|
||||
from nanobot.session.manager import (
|
||||
_SESSION_LIST_PREVIEW_MAX_CHARS,
|
||||
_SESSION_LIST_PREVIEW_MAX_RECORDS,
|
||||
@@ -154,7 +154,7 @@ def _preview_from_messages(messages: list[dict[str, Any]]) -> str:
|
||||
or scanned_chars > _SESSION_LIST_PREVIEW_MAX_CHARS
|
||||
):
|
||||
break
|
||||
if item.get(CRON_HISTORY_META) is True:
|
||||
if is_hidden_history_message(item):
|
||||
continue
|
||||
text = _message_preview_text(item)
|
||||
if not text:
|
||||
@@ -216,7 +216,7 @@ def _latest_updated_at(stored: str | None, activity: str | None) -> str | None:
|
||||
|
||||
|
||||
def _visible_message_timestamp(item: dict[str, Any]) -> str | None:
|
||||
if item.get(CRON_HISTORY_META) is True:
|
||||
if is_hidden_history_message(item):
|
||||
return None
|
||||
if item.get("role") not in _VISIBLE_TRANSCRIPT_ROLES:
|
||||
return None
|
||||
@@ -296,7 +296,9 @@ def _scan_session_row(session_manager: SessionManager, path: Path) -> dict[str,
|
||||
):
|
||||
preview_done = True
|
||||
continue
|
||||
if item.get(CRON_HISTORY_META) is True:
|
||||
if item.get("_type") == "metadata":
|
||||
continue
|
||||
if is_hidden_history_message(item):
|
||||
continue
|
||||
text = _message_preview_text(item)
|
||||
if not text:
|
||||
|
||||
@@ -99,47 +99,6 @@ _CONTEXT_WINDOW_TOKEN_OPTIONS = {65_536, 200_000, 262_144}
|
||||
_MODEL_CONFIGURATION_SLUG_RE = re.compile(r"[^a-z0-9_-]+")
|
||||
_ENV_REF_RE = re.compile(r"\$\{([A-Za-z_][A-Za-z0-9_]*)\}")
|
||||
|
||||
_MODEL_LIST_UNSUPPORTED_BACKENDS = {
|
||||
"anthropic",
|
||||
"azure_openai",
|
||||
"bedrock",
|
||||
"github_copilot",
|
||||
"openai_codex",
|
||||
}
|
||||
|
||||
_MODEL_LIST_CATALOG_PROVIDERS = {
|
||||
"aihubmix",
|
||||
"byteplus",
|
||||
"byteplus_coding_plan",
|
||||
"huggingface",
|
||||
"novita",
|
||||
"openrouter",
|
||||
"siliconflow",
|
||||
"volcengine",
|
||||
"volcengine_coding_plan",
|
||||
}
|
||||
|
||||
_MODEL_LIST_OFFICIAL_PROVIDERS = {
|
||||
"ant_ling",
|
||||
"dashscope",
|
||||
"deepseek",
|
||||
"gemini",
|
||||
"groq",
|
||||
"longcat",
|
||||
"minimax",
|
||||
"minimax_anthropic",
|
||||
"mistral",
|
||||
"moonshot",
|
||||
"nvidia",
|
||||
"openai",
|
||||
"qianfan",
|
||||
"skywork",
|
||||
"stepfun",
|
||||
"xiaomi_mimo",
|
||||
"zhipu",
|
||||
}
|
||||
|
||||
|
||||
class WebUISettingsError(ValueError):
|
||||
"""User-facing settings validation failure."""
|
||||
|
||||
@@ -394,10 +353,13 @@ def _provider_settings_row(
|
||||
|
||||
|
||||
def _model_catalog_kind(spec: Any) -> str:
|
||||
if spec.name in _MODEL_LIST_CATALOG_PROVIDERS:
|
||||
return "catalog"
|
||||
if spec.name in _MODEL_LIST_OFFICIAL_PROVIDERS:
|
||||
return "official"
|
||||
catalog = getattr(spec, "model_catalog", "auto")
|
||||
if catalog != "auto":
|
||||
return catalog
|
||||
if spec.is_transcription_only or spec.is_oauth:
|
||||
return "unsupported"
|
||||
if spec.backend != "openai_compat" and spec.name != "minimax_anthropic":
|
||||
return "unsupported"
|
||||
if spec.is_local:
|
||||
return "local"
|
||||
if spec.is_direct:
|
||||
@@ -490,27 +452,20 @@ def provider_models_payload(query: QueryParams) -> dict[str, Any]:
|
||||
raise WebUISettingsError("unknown provider")
|
||||
spec, provider_key, provider_config = resolved_provider
|
||||
|
||||
catalog_kind = _model_catalog_kind(spec)
|
||||
base_payload: dict[str, Any] = {
|
||||
"provider": provider_key,
|
||||
"label": spec.label,
|
||||
"catalog_kind": _model_catalog_kind(spec),
|
||||
"catalog_kind": catalog_kind,
|
||||
"models": [],
|
||||
"model_count": 0,
|
||||
"message": None,
|
||||
"fetched_at": time.time(),
|
||||
}
|
||||
if (
|
||||
spec.is_transcription_only
|
||||
or (
|
||||
spec.backend in _MODEL_LIST_UNSUPPORTED_BACKENDS
|
||||
and spec.name != "minimax_anthropic"
|
||||
)
|
||||
or spec.is_oauth
|
||||
):
|
||||
if catalog_kind == "unsupported":
|
||||
return {
|
||||
**base_payload,
|
||||
"status": "unsupported",
|
||||
"catalog_kind": "unsupported",
|
||||
"message": "Model list is not available for this provider. Type a model ID manually.",
|
||||
}
|
||||
|
||||
|
||||
@@ -17,7 +17,8 @@ from urllib.parse import unquote, urlparse
|
||||
from loguru import logger
|
||||
|
||||
from nanobot.config.paths import get_webui_dir
|
||||
from nanobot.cron.session_turns import CRON_HISTORY_META
|
||||
from nanobot.session.automation_turns import is_automation_kind
|
||||
from nanobot.session.history_visibility import is_hidden_history_message
|
||||
from nanobot.session.manager import SessionManager
|
||||
from nanobot.webui.metadata import WEBUI_MESSAGE_SOURCE_METADATA_KEY, WEBUI_TURN_METADATA_KEY
|
||||
|
||||
@@ -598,9 +599,12 @@ def normalize_webui_turn_id(value: Any) -> str:
|
||||
|
||||
def webui_message_source(metadata: dict[str, Any] | None) -> dict[str, str] | None:
|
||||
raw = (metadata or {}).get(WEBUI_MESSAGE_SOURCE_METADATA_KEY)
|
||||
if not isinstance(raw, dict) or raw.get("kind") != "cron":
|
||||
if not isinstance(raw, dict):
|
||||
return None
|
||||
source: dict[str, str] = {"kind": "cron"}
|
||||
kind = raw.get("kind")
|
||||
if not is_automation_kind(kind):
|
||||
return None
|
||||
source: dict[str, str] = {"kind": kind}
|
||||
label = raw.get("label")
|
||||
if isinstance(label, str) and label.strip():
|
||||
source["label"] = label.strip()
|
||||
@@ -779,6 +783,8 @@ def write_session_messages_as_transcript(
|
||||
target_chat_id = _chat_id_from_session_key(target_key)
|
||||
rows: list[dict[str, Any]] = []
|
||||
for msg in messages:
|
||||
if is_hidden_history_message(msg):
|
||||
continue
|
||||
role = msg.get("role")
|
||||
content = msg.get("content")
|
||||
text = content if isinstance(content, str) else ""
|
||||
@@ -849,13 +855,28 @@ def build_user_transcript_event(
|
||||
return event
|
||||
|
||||
|
||||
def _is_legacy_raw_subagent_result(message: dict[str, Any]) -> bool:
|
||||
content = message.get("content")
|
||||
if not isinstance(content, str):
|
||||
return False
|
||||
text = content.replace("\r\n", "\n").strip()
|
||||
return (
|
||||
text.startswith("[Subagent '")
|
||||
and "\n\nTask:" in text
|
||||
and "\n\nResult:" in text
|
||||
and "Summarize this naturally" in text
|
||||
)
|
||||
|
||||
|
||||
def _session_user_event(
|
||||
session_key: str,
|
||||
message: dict[str, Any],
|
||||
) -> dict[str, Any] | None:
|
||||
if message.get("role") != "user":
|
||||
return None
|
||||
if message.get(CRON_HISTORY_META) is True:
|
||||
if is_hidden_history_message(message):
|
||||
return None
|
||||
if _is_legacy_raw_subagent_result(message):
|
||||
return None
|
||||
content = message.get("content")
|
||||
text = content if isinstance(content, str) else ""
|
||||
@@ -1271,9 +1292,12 @@ def replay_transcript_to_ui_messages(
|
||||
|
||||
def _source_fields(rec: dict[str, Any]) -> dict[str, Any]:
|
||||
source = rec.get("source")
|
||||
if not isinstance(source, dict) or source.get("kind") != "cron":
|
||||
if not isinstance(source, dict):
|
||||
return {}
|
||||
out: dict[str, Any] = {"source": {"kind": "cron"}}
|
||||
kind = source.get("kind")
|
||||
if not is_automation_kind(kind):
|
||||
return {}
|
||||
out: dict[str, Any] = {"source": {"kind": kind}}
|
||||
label = source.get("label")
|
||||
if isinstance(label, str) and label.strip():
|
||||
out["source"]["label"] = label.strip()
|
||||
|
||||
+101
-9
@@ -26,6 +26,7 @@ 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.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_payload
|
||||
from nanobot.webui.gateway_tokens import GatewayTokenStore, token_response_payload
|
||||
@@ -89,6 +90,7 @@ if TYPE_CHECKING:
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.cron.service import CronService
|
||||
from nanobot.session.manager import SessionManager
|
||||
from nanobot.triggers.local_store import LocalTriggerStore
|
||||
|
||||
|
||||
def _decode_api_key(raw_key: str) -> str | None:
|
||||
@@ -153,7 +155,9 @@ class GatewayHTTPHandler:
|
||||
skills_workspace_path: Path,
|
||||
disabled_skills: set[str] | None = None,
|
||||
cron_service: CronService | None = None,
|
||||
local_trigger_store: LocalTriggerStore | None = None,
|
||||
cron_pending_job_ids: Callable[[str], set[str]] | None = None,
|
||||
local_trigger_pending_ids: Callable[[str], set[str]] | None = None,
|
||||
log: Any = logger,
|
||||
) -> None:
|
||||
self.config = config
|
||||
@@ -167,7 +171,9 @@ class GatewayHTTPHandler:
|
||||
self.skills_workspace_path = skills_workspace_path
|
||||
self.disabled_skills = disabled_skills or set()
|
||||
self.cron_service = cron_service
|
||||
self.local_trigger_store = local_trigger_store
|
||||
self.cron_pending_job_ids = cron_pending_job_ids
|
||||
self.local_trigger_pending_ids = local_trigger_pending_ids
|
||||
self._log = log
|
||||
self._runtime_surface = runtime_surface
|
||||
|
||||
@@ -483,13 +489,12 @@ 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")
|
||||
pending_job_ids: set[str] = set()
|
||||
if self.cron_pending_job_ids is not None:
|
||||
pending_job_ids = self.cron_pending_job_ids(decoded_key)
|
||||
pending_job_ids = self._pending_automation_ids_for_session(decoded_key)
|
||||
return _http_json_response(
|
||||
session_automations_payload(
|
||||
self.cron_service,
|
||||
decoded_key,
|
||||
local_trigger_store=self.local_trigger_store,
|
||||
pending_job_ids=pending_job_ids,
|
||||
)
|
||||
)
|
||||
@@ -506,7 +511,11 @@ class GatewayHTTPHandler:
|
||||
return _http_error(404, "session not found")
|
||||
query = _parse_query(request.path)
|
||||
delete_automations = (_query_first(query, "delete_automations") or "").lower()
|
||||
automation_jobs = session_automation_jobs(self.cron_service, decoded_key)
|
||||
automation_jobs = session_automation_jobs(
|
||||
self.cron_service,
|
||||
decoded_key,
|
||||
local_trigger_store=self.local_trigger_store,
|
||||
)
|
||||
if automation_jobs and delete_automations not in {"1", "true", "yes"}:
|
||||
return _http_json_response(
|
||||
{
|
||||
@@ -515,9 +524,13 @@ class GatewayHTTPHandler:
|
||||
"automations": serialize_automation_jobs(automation_jobs),
|
||||
}
|
||||
)
|
||||
if automation_jobs and self.cron_service is not None:
|
||||
if automation_jobs:
|
||||
for job in automation_jobs:
|
||||
self.cron_service.remove_job(job.id)
|
||||
if isinstance(job, LocalTrigger):
|
||||
if self.local_trigger_store is not None:
|
||||
self.local_trigger_store.delete(job.id)
|
||||
elif self.cron_service is not None:
|
||||
self.cron_service.remove_job(job.id)
|
||||
deleted = self.session_manager.delete_session(decoded_key)
|
||||
delete_webui_thread(decoded_key)
|
||||
return _http_json_response({"deleted": bool(deleted)})
|
||||
@@ -548,14 +561,37 @@ class GatewayHTTPHandler:
|
||||
pending.update(self.cron_pending_job_ids(session_key))
|
||||
return pending
|
||||
|
||||
def _pending_local_trigger_ids_for_all(self) -> set[str]:
|
||||
if self.local_trigger_store is None or self.local_trigger_pending_ids is None:
|
||||
return set()
|
||||
pending: set[str] = set()
|
||||
for trigger in self.local_trigger_store.list_triggers(include_disabled=True):
|
||||
session_key = trigger.session_key
|
||||
if not session_key and trigger.channel and trigger.chat_id:
|
||||
session_key = f"{trigger.channel}:{trigger.chat_id}"
|
||||
if session_key:
|
||||
pending.update(self.local_trigger_pending_ids(session_key))
|
||||
return pending
|
||||
|
||||
def _pending_automation_ids_for_session(self, session_key: str) -> set[str]:
|
||||
pending: set[str] = set()
|
||||
if self.cron_pending_job_ids is not None:
|
||||
pending.update(self.cron_pending_job_ids(session_key))
|
||||
if self.local_trigger_pending_ids is not None:
|
||||
pending.update(self.local_trigger_pending_ids(session_key))
|
||||
return pending
|
||||
|
||||
def _handle_webui_automations(self, request: WsRequest) -> Response:
|
||||
if not self.check_api_token(request):
|
||||
return _http_error(401, "Unauthorized")
|
||||
pending_job_ids = self._pending_cron_job_ids_for_all()
|
||||
pending_job_ids.update(self._pending_local_trigger_ids_for_all())
|
||||
return _http_json_response(
|
||||
all_automations_payload(
|
||||
self.cron_service,
|
||||
local_trigger_store=self.local_trigger_store,
|
||||
session_manager=self.session_manager,
|
||||
pending_job_ids=self._pending_cron_job_ids_for_all(),
|
||||
pending_job_ids=pending_job_ids,
|
||||
)
|
||||
)
|
||||
|
||||
@@ -566,13 +602,19 @@ class GatewayHTTPHandler:
|
||||
) -> Response:
|
||||
if not self.check_api_token(request):
|
||||
return _http_error(401, "Unauthorized")
|
||||
if self.cron_service is None:
|
||||
return _http_error(503, "cron service unavailable")
|
||||
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)
|
||||
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")
|
||||
trigger = self.local_trigger_store.get(job_id) if self.local_trigger_store else None
|
||||
if trigger is not None:
|
||||
return self._handle_local_trigger_action(request, action, trigger)
|
||||
|
||||
if self.cron_service is None:
|
||||
return _http_error(404, "automation not found")
|
||||
job = self.cron_service.get_job(job_id)
|
||||
if job is None:
|
||||
return _http_error(404, "automation not found")
|
||||
@@ -618,6 +660,40 @@ class GatewayHTTPHandler:
|
||||
|
||||
return self._handle_webui_automations(request)
|
||||
|
||||
def _handle_local_trigger_action(
|
||||
self,
|
||||
request: WsRequest,
|
||||
action: str,
|
||||
trigger: LocalTrigger,
|
||||
) -> Response:
|
||||
if self.local_trigger_store is None:
|
||||
return _http_error(503, "trigger service unavailable")
|
||||
if action == "enable":
|
||||
if self.local_trigger_store.enable(trigger.id, enabled=True) is None:
|
||||
return _http_error(404, "automation not found")
|
||||
elif action == "disable":
|
||||
if self.local_trigger_store.enable(trigger.id, enabled=False) is None:
|
||||
return _http_error(404, "automation not found")
|
||||
elif action == "delete":
|
||||
if not self.local_trigger_store.delete(trigger.id):
|
||||
return _http_error(404, "automation not found")
|
||||
elif action == "run":
|
||||
return _http_error(409, "local trigger requires a CLI message")
|
||||
elif action == "update":
|
||||
values = _automation_values_from_request(request)
|
||||
if values is None:
|
||||
return _http_error(400, "invalid automation update payload")
|
||||
parsed = _parse_local_trigger_update(values)
|
||||
if isinstance(parsed, str):
|
||||
return _http_error(400, parsed)
|
||||
if parsed:
|
||||
if self.local_trigger_store.update(trigger.id, **parsed) is None:
|
||||
return _http_error(404, "automation not found")
|
||||
else:
|
||||
return _http_error(404, "unknown automation action")
|
||||
|
||||
return self._handle_webui_automations(request)
|
||||
|
||||
@staticmethod
|
||||
def _log_automation_run_result(task: asyncio.Task[bool]) -> None:
|
||||
try:
|
||||
@@ -830,6 +906,22 @@ def _parse_automation_update(
|
||||
return update
|
||||
|
||||
|
||||
def _parse_local_trigger_update(values: dict[str, Any]) -> dict[str, Any] | str:
|
||||
update: dict[str, Any] = {}
|
||||
if "name" in values:
|
||||
raw_name = values.get("name")
|
||||
if not isinstance(raw_name, str):
|
||||
return "name must be a string"
|
||||
name = raw_name.strip()
|
||||
if not name:
|
||||
return "name cannot be empty"
|
||||
update["name"] = name
|
||||
forbidden = [key for key in ("message", "schedule") if key in values]
|
||||
if forbidden:
|
||||
return "local trigger updates only support name"
|
||||
return update
|
||||
|
||||
|
||||
def _parse_automation_schedule(values: dict[str, Any]) -> CronSchedule | str:
|
||||
raw_kind = values.get("kind")
|
||||
if not isinstance(raw_kind, str):
|
||||
|
||||
@@ -5,6 +5,7 @@ import pytest
|
||||
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import GoalStatusEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import GenerationSettings, LLMResponse
|
||||
from nanobot.session.webui_turns import WebuiTurnCoordinator
|
||||
@@ -54,13 +55,13 @@ async def test_process_direct_websocket_clears_run_status(tmp_path) -> None:
|
||||
events.append(await loop.bus.consume_outbound())
|
||||
|
||||
statuses = [
|
||||
event.metadata
|
||||
event.event
|
||||
for event in events
|
||||
if event.metadata.get("_goal_status") is True
|
||||
if isinstance(event.event, GoalStatusEvent)
|
||||
]
|
||||
assert [status["goal_status"] for status in statuses] == ["running", "idle"]
|
||||
assert isinstance(statuses[0].get("started_at"), float)
|
||||
assert "started_at" not in statuses[1]
|
||||
assert [status.status for status in statuses] == ["running", "idle"]
|
||||
assert isinstance(statuses[0].started_at, float)
|
||||
assert statuses[1].started_at is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -9,6 +9,15 @@ import pytest
|
||||
import nanobot.agent.runner as runner_module
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.outbound_events import (
|
||||
GoalStatusEvent,
|
||||
ProgressEvent,
|
||||
SessionUpdatedEvent,
|
||||
StreamDeltaEvent,
|
||||
StreamedResponseEvent,
|
||||
StreamEndEvent,
|
||||
TurnEndEvent,
|
||||
)
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||
from nanobot.session.webui_turns import WebuiTurnCoordinator
|
||||
@@ -260,25 +269,45 @@ class TestToolEventProgress:
|
||||
)
|
||||
await loop._dispatch(msg)
|
||||
|
||||
# Drain all outbound messages and find the one carrying _tool_events
|
||||
# Drain all outbound messages and find the one carrying tool events.
|
||||
outbound = []
|
||||
while bus.outbound_size > 0:
|
||||
outbound.append(await bus.consume_outbound())
|
||||
|
||||
tool_event_msgs = [m for m in outbound if m.metadata and m.metadata.get("_tool_events")]
|
||||
assert tool_event_msgs, "expected at least one outbound message with _tool_events"
|
||||
tool_event_msgs = [
|
||||
m
|
||||
for m in outbound
|
||||
if isinstance(m.event, ProgressEvent) and m.event.tool_events
|
||||
]
|
||||
assert tool_event_msgs, "expected at least one outbound message with tool events"
|
||||
|
||||
start_msgs = [m for m in tool_event_msgs if m.metadata["_tool_events"][0]["phase"] == "start"]
|
||||
finish_msgs = [m for m in tool_event_msgs if m.metadata["_tool_events"][0]["phase"] in ("end", "error")]
|
||||
start_msgs = [
|
||||
m
|
||||
for m in tool_event_msgs
|
||||
if isinstance(m.event, ProgressEvent)
|
||||
and m.event.tool_events
|
||||
and m.event.tool_events[0]["phase"] == "start"
|
||||
]
|
||||
finish_msgs = [
|
||||
m
|
||||
for m in tool_event_msgs
|
||||
if isinstance(m.event, ProgressEvent)
|
||||
and m.event.tool_events
|
||||
and m.event.tool_events[0]["phase"] in ("end", "error")
|
||||
]
|
||||
assert start_msgs, "expected a start-phase tool event"
|
||||
assert finish_msgs, "expected a finish-phase tool event"
|
||||
|
||||
start = start_msgs[0].metadata["_tool_events"][0]
|
||||
assert isinstance(start_msgs[0].event, ProgressEvent)
|
||||
assert start_msgs[0].event.tool_events is not None
|
||||
start = start_msgs[0].event.tool_events[0]
|
||||
assert start["name"] == "exec"
|
||||
assert start["call_id"] == "tc1"
|
||||
assert start["result"] is None
|
||||
|
||||
finish = finish_msgs[0].metadata["_tool_events"][0]
|
||||
assert isinstance(finish_msgs[0].event, ProgressEvent)
|
||||
assert finish_msgs[0].event.tool_events is not None
|
||||
finish = finish_msgs[0].event.tool_events[0]
|
||||
assert finish["phase"] == "end"
|
||||
assert finish["result"] == "file.txt"
|
||||
|
||||
@@ -309,7 +338,8 @@ class TestToolEventProgress:
|
||||
await invoke_file_edit_progress(progress, edit_events)
|
||||
outbound = await bus.consume_outbound()
|
||||
assert outbound.channel == "telegram"
|
||||
assert outbound.metadata["_file_edit_events"] == edit_events
|
||||
assert isinstance(outbound.event, ProgressEvent)
|
||||
assert outbound.event.file_edit_events == edit_events
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_goal_turn_keeps_live_file_edit_progress_for_webui(self, tmp_path: Path) -> None:
|
||||
@@ -389,7 +419,8 @@ class TestToolEventProgress:
|
||||
edit_events = [
|
||||
event
|
||||
for msg in outbound
|
||||
for event in msg.metadata.get("_file_edit_events", [])
|
||||
if isinstance(msg.event, ProgressEvent)
|
||||
for event in msg.event.file_edit_events or []
|
||||
]
|
||||
assert any(
|
||||
event["status"] == "editing"
|
||||
@@ -433,8 +464,8 @@ class TestToolEventProgress:
|
||||
outbound.append(await bus.consume_outbound())
|
||||
|
||||
assert [m.content for m in outbound] == ["Hello"]
|
||||
assert not any(m.metadata.get("_progress") for m in outbound)
|
||||
assert not any(m.metadata.get("_streamed") for m in outbound)
|
||||
assert not any(isinstance(m.event, ProgressEvent) for m in outbound)
|
||||
assert not any(isinstance(m.event, StreamedResponseEvent) for m in outbound)
|
||||
provider.chat_stream_with_retry.assert_not_awaited()
|
||||
provider.chat_with_retry.assert_awaited_once()
|
||||
|
||||
@@ -443,7 +474,7 @@ class TestToolEventProgress:
|
||||
self,
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
"""Streaming channels still receive provider deltas through _stream_delta messages."""
|
||||
"""Streaming channels still receive provider deltas through stream events."""
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.supports_progress_deltas = True
|
||||
@@ -473,21 +504,19 @@ class TestToolEventProgress:
|
||||
while bus.outbound_size > 0:
|
||||
outbound.append(await bus.consume_outbound())
|
||||
|
||||
deltas = [m for m in outbound if m.metadata.get("_stream_delta")]
|
||||
stream_end = [m for m in outbound if m.metadata.get("_stream_end")]
|
||||
deltas = [m for m in outbound if isinstance(m.event, StreamDeltaEvent)]
|
||||
stream_end = [m for m in outbound if isinstance(m.event, StreamEndEvent)]
|
||||
final = [
|
||||
m for m in outbound
|
||||
if not m.metadata.get("_stream_delta")
|
||||
and not m.metadata.get("_stream_end")
|
||||
and not m.metadata.get("_turn_end")
|
||||
and not m.metadata.get("_goal_status")
|
||||
if not isinstance(m.event, StreamDeltaEvent | StreamEndEvent)
|
||||
and not isinstance(m.event, TurnEndEvent | GoalStatusEvent)
|
||||
]
|
||||
|
||||
assert [m.content for m in deltas] == ["Hel", "lo"]
|
||||
assert len(stream_end) == 1
|
||||
assert final[-1].content == "Hello"
|
||||
assert final[-1].metadata.get("_streamed") is True
|
||||
turn_end_msgs = [m for m in outbound if m.metadata.get("_turn_end")]
|
||||
assert isinstance(final[-1].event, StreamedResponseEvent)
|
||||
turn_end_msgs = [m for m in outbound if isinstance(m.event, TurnEndEvent)]
|
||||
assert len(turn_end_msgs) == 1
|
||||
assert turn_end_msgs[0].content == ""
|
||||
provider.chat_with_retry.assert_not_awaited()
|
||||
@@ -528,23 +557,28 @@ class TestToolEventProgress:
|
||||
while bus.outbound_size > 0:
|
||||
outbound.append(await bus.consume_outbound())
|
||||
|
||||
deltas = [m for m in outbound if m.metadata.get("_stream_delta")]
|
||||
stream_end = [m for m in outbound if m.metadata.get("_stream_end")]
|
||||
deltas = [m for m in outbound if isinstance(m.event, StreamDeltaEvent)]
|
||||
stream_end = [m for m in outbound if isinstance(m.event, StreamEndEvent)]
|
||||
final = [
|
||||
m for m in outbound
|
||||
if not m.metadata.get("_stream_delta")
|
||||
and not m.metadata.get("_stream_end")
|
||||
and not m.metadata.get("_turn_end")
|
||||
and not m.metadata.get("_goal_status")
|
||||
if not isinstance(m.event, StreamDeltaEvent | StreamEndEvent)
|
||||
and not isinstance(m.event, TurnEndEvent | GoalStatusEvent)
|
||||
]
|
||||
|
||||
assert [m.content for m in deltas] == ["partial", "full retry response"]
|
||||
assert [m.metadata.get("_resuming") for m in stream_end] == [True, False]
|
||||
assert deltas[0].metadata.get("_stream_id") == stream_end[0].metadata.get("_stream_id")
|
||||
assert deltas[1].metadata.get("_stream_id") == stream_end[1].metadata.get("_stream_id")
|
||||
assert deltas[0].metadata.get("_stream_id") != deltas[1].metadata.get("_stream_id")
|
||||
assert [m.event.resuming for m in stream_end if isinstance(m.event, StreamEndEvent)] == [
|
||||
True,
|
||||
False,
|
||||
]
|
||||
assert isinstance(deltas[0].event, StreamDeltaEvent)
|
||||
assert isinstance(deltas[1].event, StreamDeltaEvent)
|
||||
assert isinstance(stream_end[0].event, StreamEndEvent)
|
||||
assert isinstance(stream_end[1].event, StreamEndEvent)
|
||||
assert deltas[0].event.stream_id == stream_end[0].event.stream_id
|
||||
assert deltas[1].event.stream_id == stream_end[1].event.stream_id
|
||||
assert deltas[0].event.stream_id != deltas[1].event.stream_id
|
||||
assert final[-1].content == "full retry response"
|
||||
assert final[-1].metadata.get("_streamed") is True
|
||||
assert isinstance(final[-1].event, StreamedResponseEvent)
|
||||
provider.chat_with_retry.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -623,9 +657,9 @@ class TestToolEventProgress:
|
||||
|
||||
done_msgs = [m for m in outbound if m.content == "Done"]
|
||||
assert len(done_msgs) == 1
|
||||
assert not done_msgs[0].metadata.get("_turn_end")
|
||||
assert not isinstance(done_msgs[0].event, TurnEndEvent)
|
||||
|
||||
turn_end_msgs = [m for m in outbound if m.metadata.get("_turn_end")]
|
||||
turn_end_msgs = [m for m in outbound if isinstance(m.event, TurnEndEvent)]
|
||||
assert len(turn_end_msgs) == 1
|
||||
assert turn_end_msgs[0].content == ""
|
||||
assert turn_end_msgs[0].chat_id == "chat1"
|
||||
@@ -659,14 +693,14 @@ class TestToolEventProgress:
|
||||
outbound.append(await bus.consume_outbound())
|
||||
|
||||
error_msgs = [m for m in outbound if m.content == "Sorry, I encountered an error."]
|
||||
turn_end_msgs = [m for m in outbound if m.metadata.get("_turn_end")]
|
||||
statuses = [m for m in outbound if m.metadata.get("_goal_status")]
|
||||
turn_end_msgs = [m for m in outbound if isinstance(m.event, TurnEndEvent)]
|
||||
statuses = [m for m in outbound if isinstance(m.event, GoalStatusEvent)]
|
||||
|
||||
assert len(error_msgs) == 1
|
||||
assert len(turn_end_msgs) == 1
|
||||
assert turn_end_msgs[0].content == ""
|
||||
assert turn_end_msgs[0].chat_id == "chat1"
|
||||
assert [m.metadata["goal_status"] for m in statuses] == ["idle"]
|
||||
assert [m.event.status for m in statuses if isinstance(m.event, GoalStatusEvent)] == ["idle"]
|
||||
assert outbound.index(error_msgs[0]) < outbound.index(turn_end_msgs[0])
|
||||
assert outbound.index(turn_end_msgs[0]) < outbound.index(statuses[-1])
|
||||
|
||||
@@ -705,27 +739,27 @@ class TestToolEventProgress:
|
||||
outbound: list = []
|
||||
for _ in range(12):
|
||||
outbound.append(await asyncio.wait_for(bus.consume_outbound(), timeout=0.5))
|
||||
if outbound[-1].metadata.get("_turn_end"):
|
||||
if isinstance(outbound[-1].event, TurnEndEvent):
|
||||
break
|
||||
else:
|
||||
raise AssertionError("_turn_end message not found")
|
||||
raise AssertionError("turn-end event not found")
|
||||
|
||||
done_with_body = [m for m in outbound if m.content == "Done"]
|
||||
assert len(done_with_body) == 1
|
||||
assert outbound[-1].metadata.get("_turn_end") is True
|
||||
assert isinstance(outbound[-1].event, TurnEndEvent)
|
||||
|
||||
await asyncio.wait_for(title_started.wait(), timeout=0.5)
|
||||
release_title.set()
|
||||
session_updated = None
|
||||
for _ in range(10):
|
||||
candidate = await asyncio.wait_for(bus.consume_outbound(), timeout=0.5)
|
||||
if (candidate.metadata or {}).get("_session_updated"):
|
||||
if isinstance(candidate.event, SessionUpdatedEvent):
|
||||
session_updated = candidate
|
||||
break
|
||||
assert session_updated is not None
|
||||
|
||||
assert (session_updated.metadata or {}).get("_session_updated") is True
|
||||
assert (session_updated.metadata or {}).get("_session_update_scope") == "metadata"
|
||||
assert isinstance(session_updated.event, SessionUpdatedEvent)
|
||||
assert session_updated.event.scope == "metadata"
|
||||
assert provider.chat_with_retry.await_count == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -837,4 +871,4 @@ class TestToolEventProgress:
|
||||
|
||||
assert len(outbound) == 1
|
||||
assert outbound[0].content == "Done"
|
||||
assert (outbound[0].metadata or {}).get("_turn_end") is not True
|
||||
assert not isinstance(outbound[0].event, TurnEndEvent)
|
||||
|
||||
@@ -7,6 +7,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from nanobot.bus.outbound_events import StreamedResponseEvent
|
||||
from nanobot.config.schema import AgentDefaults
|
||||
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||
|
||||
@@ -23,8 +24,8 @@ def _make_loop(tmp_path):
|
||||
|
||||
with patch("nanobot.agent.loop.ContextBuilder"), \
|
||||
patch("nanobot.agent.loop.SessionManager"), \
|
||||
patch("nanobot.agent.loop.SubagentManager") as MockSubMgr:
|
||||
MockSubMgr.return_value.cancel_by_session = AsyncMock(return_value=0)
|
||||
patch("nanobot.agent.loop.SubagentManager") as mock_sub_mgr:
|
||||
mock_sub_mgr.return_value.cancel_by_session = AsyncMock(return_value=0)
|
||||
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path)
|
||||
return loop
|
||||
|
||||
@@ -193,8 +194,9 @@ async def test_streamed_flag_not_set_on_llm_error(tmp_path):
|
||||
|
||||
assert result is not None
|
||||
assert "503" in result.content
|
||||
assert not result.metadata.get("_streamed"), \
|
||||
"_streamed must not be set when stop_reason is error"
|
||||
assert not isinstance(result.event, StreamedResponseEvent), (
|
||||
"streamed response event must not be set when stop_reason is error"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -239,7 +241,7 @@ async def test_ssrf_soft_block_can_finalize_after_streamed_tool_call(tmp_path):
|
||||
|
||||
assert result is not None
|
||||
assert result.content == "I cannot access private URLs. Please share the local file."
|
||||
assert result.metadata.get("_streamed") is True
|
||||
assert isinstance(result.event, StreamedResponseEvent)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -8,9 +8,17 @@ import pytest
|
||||
from nanobot.agent.context import ContextBuilder
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.outbound_events import (
|
||||
GoalStatusEvent,
|
||||
StreamDeltaEvent,
|
||||
StreamedResponseEvent,
|
||||
StreamEndEvent,
|
||||
TurnEndEvent,
|
||||
)
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.cron.session_turns import CRON_HISTORY_META, CRON_TRIGGER_META
|
||||
from nanobot.providers.base import LLMResponse
|
||||
from nanobot.session.automation_turns import AUTOMATION_HISTORY_META
|
||||
from nanobot.session.goal_state import GOAL_STATE_KEY
|
||||
from nanobot.session.manager import Session, SessionManager
|
||||
from nanobot.session.turn_continuation import (
|
||||
@@ -26,6 +34,7 @@ from nanobot.session.webui_turns import (
|
||||
clean_generated_title,
|
||||
maybe_generate_webui_title,
|
||||
)
|
||||
from nanobot.triggers.local_session_turns import LOCAL_TRIGGER_META
|
||||
from nanobot.utils.llm_runtime import LLMRuntime
|
||||
|
||||
|
||||
@@ -94,6 +103,13 @@ def test_persist_cron_turn_uses_distinct_history_marker(tmp_path: Path) -> None:
|
||||
assert persisted is True
|
||||
message = session.messages[-1]
|
||||
assert message["content"] == "Scheduled cron job triggered: Daily check"
|
||||
assert message[AUTOMATION_HISTORY_META] == {
|
||||
"kind": "cron",
|
||||
"cron_job_id": "job-1",
|
||||
"cron_job_name": "Daily check",
|
||||
"cron_run_id": "job-1:1",
|
||||
"cron_prompt_ref": prompt_ref,
|
||||
}
|
||||
assert message[CRON_HISTORY_META] is True
|
||||
assert CRON_TRIGGER_META not in message
|
||||
assert message["cron_job_id"] == "job-1"
|
||||
@@ -102,6 +118,63 @@ def test_persist_cron_turn_uses_distinct_history_marker(tmp_path: Path) -> None:
|
||||
assert message["cron_prompt_ref"] == prompt_ref
|
||||
|
||||
|
||||
def test_persist_local_trigger_turn_uses_hidden_automation_marker(tmp_path: Path) -> None:
|
||||
loop = _make_full_loop(tmp_path)
|
||||
session = loop.sessions.get_or_create("websocket:auto")
|
||||
|
||||
persisted = loop._persist_user_message_early(
|
||||
InboundMessage(
|
||||
channel="websocket",
|
||||
sender_id="trigger",
|
||||
chat_id="auto",
|
||||
content="Review PR #4502",
|
||||
metadata={
|
||||
LOCAL_TRIGGER_META: {
|
||||
"trigger_id": "trg_123",
|
||||
"trigger_name": "PR review",
|
||||
"delivery_id": "tdel_456",
|
||||
"created_at_ms": 1_700_000_000_000,
|
||||
"persist_content": "Local trigger received: PR review\n\nReview PR #4502",
|
||||
}
|
||||
},
|
||||
),
|
||||
session,
|
||||
)
|
||||
|
||||
assert persisted is True
|
||||
message = session.messages[-1]
|
||||
assert message["content"] == "Local trigger received: PR review\n\nReview PR #4502"
|
||||
assert message[AUTOMATION_HISTORY_META] == {
|
||||
"kind": "local_trigger",
|
||||
"trigger_id": "trg_123",
|
||||
"trigger_name": "PR review",
|
||||
"trigger_delivery_id": "tdel_456",
|
||||
}
|
||||
assert LOCAL_TRIGGER_META not in message
|
||||
assert message["trigger_id"] == "trg_123"
|
||||
assert message["trigger_name"] == "PR review"
|
||||
assert message["trigger_delivery_id"] == "tdel_456"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_with_bot_suffix_does_not_persist_command(tmp_path: Path) -> None:
|
||||
loop = _make_full_loop(tmp_path)
|
||||
|
||||
response = await loop._process_message(
|
||||
InboundMessage(
|
||||
channel="websocket",
|
||||
sender_id="user",
|
||||
chat_id="chat-1",
|
||||
content="/new@nanobot_bot",
|
||||
)
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert response.content == "New session started."
|
||||
session = loop.sessions.get_or_create("websocket:chat-1")
|
||||
assert session.messages == []
|
||||
|
||||
|
||||
def test_clean_generated_title_strips_reasoning_tags() -> None:
|
||||
assert clean_generated_title("<think>reasoning</think> WebUI polish") == "WebUI polish"
|
||||
assert clean_generated_title("Title: <think> The user said hello") == ""
|
||||
@@ -765,7 +838,6 @@ async def test_internal_continuation_preserves_streaming_route_metadata(
|
||||
"_wants_stream": True,
|
||||
"message_id": "om_001",
|
||||
"origin_message_id": "root_001",
|
||||
"_stream_id": "old-stream",
|
||||
},
|
||||
))
|
||||
|
||||
@@ -775,23 +847,23 @@ async def test_internal_continuation_preserves_streaming_route_metadata(
|
||||
assert queued.metadata["_wants_stream"] is True
|
||||
assert queued.metadata["message_id"] == "om_001"
|
||||
assert queued.metadata["origin_message_id"] == "root_001"
|
||||
assert "_stream_id" not in queued.metadata
|
||||
|
||||
await loop._dispatch(queued)
|
||||
|
||||
outbound = []
|
||||
while loop.bus.outbound_size:
|
||||
outbound.append(await loop.bus.consume_outbound())
|
||||
deltas = [m for m in outbound if m.metadata.get("_stream_delta")]
|
||||
ends = [m for m in outbound if m.metadata.get("_stream_end")]
|
||||
streamed_markers = [m for m in outbound if m.metadata.get("_streamed")]
|
||||
deltas = [m for m in outbound if isinstance(m.event, StreamDeltaEvent)]
|
||||
ends = [m for m in outbound if isinstance(m.event, StreamEndEvent)]
|
||||
streamed_markers = [m for m in outbound if isinstance(m.event, StreamedResponseEvent)]
|
||||
|
||||
assert [m.content for m in deltas] == ["done"]
|
||||
assert len(ends) == 1
|
||||
assert ends[0].metadata["_resuming"] is False
|
||||
assert isinstance(ends[0].event, StreamEndEvent)
|
||||
assert ends[0].event.resuming is False
|
||||
assert ends[0].metadata["message_id"] == "om_001"
|
||||
assert ends[0].metadata["origin_message_id"] == "root_001"
|
||||
assert isinstance(ends[0].metadata.get("_stream_id"), str)
|
||||
assert isinstance(ends[0].event.stream_id, str)
|
||||
assert streamed_markers and streamed_markers[-1].content == "done"
|
||||
|
||||
|
||||
@@ -842,10 +914,10 @@ async def test_websocket_internal_continuation_keeps_single_visible_run(
|
||||
first_outbound = []
|
||||
while loop.bus.outbound_size:
|
||||
first_outbound.append(await loop.bus.consume_outbound())
|
||||
first_statuses = [m.metadata for m in first_outbound if m.metadata.get("_goal_status")]
|
||||
assert [m["goal_status"] for m in first_statuses] == ["running"]
|
||||
assert not [m for m in first_outbound if m.metadata.get("_turn_end")]
|
||||
started_at = first_statuses[0]["started_at"]
|
||||
first_statuses = [m.event for m in first_outbound if isinstance(m.event, GoalStatusEvent)]
|
||||
assert [m.status for m in first_statuses] == ["running"]
|
||||
assert not [m for m in first_outbound if isinstance(m.event, TurnEndEvent)]
|
||||
started_at = first_statuses[0].started_at
|
||||
|
||||
queued = await asyncio.wait_for(loop.bus.consume_inbound(), timeout=0.5)
|
||||
assert queued.metadata[INTERNAL_CONTINUATION_META] is True
|
||||
@@ -856,12 +928,13 @@ async def test_websocket_internal_continuation_keeps_single_visible_run(
|
||||
second_outbound = []
|
||||
while loop.bus.outbound_size:
|
||||
second_outbound.append(await loop.bus.consume_outbound())
|
||||
second_statuses = [m.metadata for m in second_outbound if m.metadata.get("_goal_status")]
|
||||
assert [m["goal_status"] for m in second_statuses] == ["running", "idle"]
|
||||
assert second_statuses[0]["started_at"] == started_at
|
||||
turn_end = [m for m in second_outbound if m.metadata.get("_turn_end")]
|
||||
second_statuses = [m.event for m in second_outbound if isinstance(m.event, GoalStatusEvent)]
|
||||
assert [m.status for m in second_statuses] == ["running", "idle"]
|
||||
assert second_statuses[0].started_at == started_at
|
||||
turn_end = [m for m in second_outbound if isinstance(m.event, TurnEndEvent)]
|
||||
assert len(turn_end) == 1
|
||||
assert isinstance(turn_end[0].metadata.get("latency_ms"), int)
|
||||
assert isinstance(turn_end[0].event, TurnEndEvent)
|
||||
assert isinstance(turn_end[0].event.latency_ms, int)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -853,10 +853,11 @@ class TestApiServerRegistration:
|
||||
config = Config()
|
||||
from nanobot.config.schema import ApiConfig
|
||||
|
||||
new_api = ApiConfig(host="0.0.0.0", port=9999)
|
||||
new_api = ApiConfig(host="0.0.0.0", port=9999, api_key="secret")
|
||||
_SETTINGS_SETTER["API Server"](config, new_api)
|
||||
assert config.api.host == "0.0.0.0"
|
||||
assert config.api.port == 9999
|
||||
assert config.api.api_key == "secret"
|
||||
|
||||
|
||||
class TestMainMenuUpdate:
|
||||
|
||||
@@ -103,6 +103,48 @@ async def test_llm_arrearage_error_surfaces_clear_message():
|
||||
assert result.final_content == _ARREARAGE_ERROR_MESSAGE
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("finish_reason", "expected_stop_reason"),
|
||||
[
|
||||
("refusal", "completed"),
|
||||
("content_filter", "completed"),
|
||||
("error", "error"),
|
||||
],
|
||||
)
|
||||
async def test_runner_ignores_tool_calls_when_finish_reason_blocks_execution(
|
||||
finish_reason: str,
|
||||
expected_stop_reason: str,
|
||||
):
|
||||
"""Provider/gateway-injected tool calls under terminal block reasons must not run."""
|
||||
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||
|
||||
provider = MagicMock(spec=LLMProvider)
|
||||
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
|
||||
content="Request blocked by provider policy.",
|
||||
finish_reason=finish_reason,
|
||||
tool_calls=[ToolCallRequest(id="call_1", name="exec", arguments={"command": "echo nope"})],
|
||||
usage={},
|
||||
))
|
||||
tools = MagicMock()
|
||||
tools.get_definitions.return_value = []
|
||||
tools.execute = AsyncMock(return_value="should not run")
|
||||
|
||||
result = await AgentRunner(provider).run(AgentRunSpec(
|
||||
initial_messages=[{"role": "user", "content": "run a command"}],
|
||||
tools=tools,
|
||||
model="test-model",
|
||||
max_iterations=2,
|
||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||
))
|
||||
|
||||
tools.execute.assert_not_awaited()
|
||||
assert result.stop_reason == expected_stop_reason
|
||||
assert result.tools_used == []
|
||||
assert result.final_content == "Request blocked by provider policy."
|
||||
assert not any(msg.get("role") == "tool" for msg in result.messages)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runner_tool_error_sets_final_content():
|
||||
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||
@@ -135,6 +177,46 @@ async def test_runner_tool_error_sets_final_content():
|
||||
assert result.stop_reason == "tool_error"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runner_preserves_successful_exec_output_that_starts_with_error():
|
||||
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||
|
||||
provider = MagicMock(spec=LLMProvider)
|
||||
|
||||
async def chat_with_retry(*, messages, **kwargs):
|
||||
if not any(msg.get("role") == "tool" for msg in messages):
|
||||
return LLMResponse(
|
||||
content="working",
|
||||
tool_calls=[
|
||||
ToolCallRequest(id="call_1", name="exec", arguments={"command": "report"})
|
||||
],
|
||||
usage={},
|
||||
)
|
||||
return LLMResponse(content="done", usage={})
|
||||
|
||||
provider.chat_with_retry = chat_with_retry
|
||||
output = "Error: generated report successfully\n\nExit code: 0"
|
||||
tools = MagicMock()
|
||||
tools.get_definitions.return_value = []
|
||||
tools.execute = AsyncMock(return_value=output)
|
||||
|
||||
runner = AgentRunner(provider)
|
||||
result = await runner.run(AgentRunSpec(
|
||||
initial_messages=[{"role": "user", "content": "run report"}],
|
||||
tools=tools,
|
||||
model="test-model",
|
||||
max_iterations=2,
|
||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||
fail_on_tool_error=True,
|
||||
))
|
||||
|
||||
assert result.final_content == "done"
|
||||
assert result.stop_reason == "completed"
|
||||
assert result.tool_events == [
|
||||
{"name": "exec", "status": "ok", "detail": "Error: generated report successfully Exit code: 0"}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runner_tool_error_preserves_tool_results_in_messages():
|
||||
"""When a tool raises a fatal error, its results must still be appended
|
||||
|
||||
@@ -465,6 +465,68 @@ async def test_loop_injected_followup_preserves_image_media(tmp_path):
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subagent_pending_injection_is_hidden_history_and_not_merged(tmp_path):
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.session.history_visibility import HIDDEN_HISTORY_META
|
||||
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
call_count = {"n": 0}
|
||||
|
||||
async def chat_with_retry(*, messages, **kwargs):
|
||||
call_count["n"] += 1
|
||||
if call_count["n"] == 1:
|
||||
return LLMResponse(content="first answer", tool_calls=[], usage={})
|
||||
return LLMResponse(content="second answer", tool_calls=[], usage={})
|
||||
|
||||
provider.chat_with_retry = chat_with_retry
|
||||
loop = AgentLoop(bus=bus, provider=provider, workspace=tmp_path, model="test-model")
|
||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||
|
||||
payload = (
|
||||
"[Subagent 'x' completed successfully]\n\n"
|
||||
"Task: t\n\n"
|
||||
"Result:\nr\n\n"
|
||||
"Summarize this naturally for the user."
|
||||
)
|
||||
pending_queue = asyncio.Queue()
|
||||
await pending_queue.put(InboundMessage(
|
||||
channel="cli",
|
||||
sender_id="user",
|
||||
chat_id="c",
|
||||
content="visible follow-up",
|
||||
))
|
||||
await pending_queue.put(InboundMessage(
|
||||
channel="system",
|
||||
sender_id="subagent",
|
||||
chat_id="cli:c",
|
||||
content=payload,
|
||||
metadata={"injected_event": "subagent_result", "subagent_task_id": "sub-1"},
|
||||
))
|
||||
|
||||
final_content, _, all_msgs, _, had_injections = await loop._run_agent_loop(
|
||||
[{"role": "user", "content": "hello"}],
|
||||
channel="cli",
|
||||
chat_id="c",
|
||||
pending_queue=pending_queue,
|
||||
)
|
||||
|
||||
assert final_content == "second answer"
|
||||
assert had_injections is True
|
||||
assert call_count["n"] == 2
|
||||
injected_users = [message for message in all_msgs if message.get("role") == "user"][-2:]
|
||||
assert [message["content"] for message in injected_users] == ["visible follow-up", payload]
|
||||
assert injected_users[1][HIDDEN_HISTORY_META] == {
|
||||
"kind": "subagent_result",
|
||||
"subagent_task_id": "sub-1",
|
||||
}
|
||||
assert injected_users[1]["injected_event"] == "subagent_result"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runner_merges_multiple_injected_user_messages_without_losing_media():
|
||||
"""Multiple injected follow-ups should not create lossy consecutive user messages."""
|
||||
@@ -730,6 +792,56 @@ async def test_cron_turn_deferred_while_session_active(tmp_path):
|
||||
assert loop.pending_cron_job_ids_for_session(session_key) == set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_local_trigger_turn_deferred_while_session_active(tmp_path):
|
||||
"""Local trigger turns wait for the active session instead of becoming injections."""
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.triggers.local_session_turns import LOCAL_TRIGGER_META
|
||||
|
||||
loop = _make_loop(tmp_path)
|
||||
loop._dispatch = AsyncMock() # type: ignore[method-assign]
|
||||
|
||||
session_key = "websocket:chat-1"
|
||||
pending = asyncio.Queue(maxsize=20)
|
||||
loop._pending_queues[session_key] = pending
|
||||
|
||||
run_task = asyncio.create_task(loop.run())
|
||||
msg = InboundMessage(
|
||||
channel="websocket",
|
||||
sender_id="trigger",
|
||||
chat_id="chat-1",
|
||||
content="review failed CI",
|
||||
metadata={
|
||||
LOCAL_TRIGGER_META: {
|
||||
"trigger_id": "trg_123",
|
||||
"trigger_name": "CI review",
|
||||
"delivery_id": "tdl_123",
|
||||
},
|
||||
},
|
||||
session_key_override=session_key,
|
||||
)
|
||||
await loop.bus.publish_inbound(msg)
|
||||
|
||||
for _ in range(20):
|
||||
if loop._local_trigger_turns.deferred_queues.get(session_key):
|
||||
break
|
||||
await asyncio.sleep(0.05)
|
||||
|
||||
loop.stop()
|
||||
await asyncio.wait_for(run_task, timeout=2)
|
||||
|
||||
assert pending.empty()
|
||||
assert loop._dispatch.await_count == 0
|
||||
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
|
||||
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
|
||||
assert loop.pending_local_trigger_ids_for_session(session_key) == set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_submitted_cron_turn_reports_pending_until_completed(tmp_path):
|
||||
"""Bound cron jobs remain marked pending while their session turn is in flight."""
|
||||
@@ -766,6 +878,48 @@ async def test_submitted_cron_turn_reports_pending_until_completed(tmp_path):
|
||||
assert loop.pending_cron_job_ids_for_session(session_key) == set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_submitted_local_trigger_turn_reports_pending_until_completed(tmp_path):
|
||||
"""Local triggers remain marked pending while their session turn is in flight."""
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
from nanobot.triggers.local_session_turns import LOCAL_TRIGGER_META
|
||||
|
||||
loop = _make_loop(tmp_path)
|
||||
loop._running = True
|
||||
|
||||
session_key = "websocket:chat-1"
|
||||
msg = InboundMessage(
|
||||
channel="websocket",
|
||||
sender_id="trigger",
|
||||
chat_id="chat-1",
|
||||
content="review failed CI",
|
||||
metadata={
|
||||
LOCAL_TRIGGER_META: {
|
||||
"trigger_id": "trg_123",
|
||||
"trigger_name": "CI review",
|
||||
"delivery_id": "tdl_123",
|
||||
},
|
||||
},
|
||||
session_key_override=session_key,
|
||||
)
|
||||
|
||||
submit_task = asyncio.create_task(loop.submit_local_trigger_turn(msg))
|
||||
queued = await asyncio.wait_for(loop.bus.consume_inbound(), timeout=0.5)
|
||||
|
||||
assert queued is msg
|
||||
assert loop.pending_local_trigger_ids_for_session(session_key) == {"trg_123"}
|
||||
|
||||
response = OutboundMessage(
|
||||
channel="websocket",
|
||||
chat_id="chat-1",
|
||||
content="done",
|
||||
)
|
||||
loop._local_trigger_turns.complete(msg, response=response)
|
||||
|
||||
assert await asyncio.wait_for(submit_task, timeout=0.5) is response
|
||||
assert loop.pending_local_trigger_ids_for_session(session_key) == set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pending_queue_preserves_overflow_for_next_injection_cycle(tmp_path):
|
||||
"""Pending queue should leave overflow messages queued for later drains."""
|
||||
|
||||
@@ -6,6 +6,8 @@ from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from nanobot.agent.runner import AgentRunner, AgentRunSpec
|
||||
from nanobot.agent.tools import ToolResult
|
||||
from nanobot.config.schema import AgentDefaults
|
||||
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||
|
||||
@@ -20,8 +22,6 @@ async def test_runner_does_not_abort_on_workspace_violation_anymore():
|
||||
we now hand the error back to the LLM as a recoverable tool result and
|
||||
rely on ``repeated_workspace_violation_error`` to throttle bypass loops.
|
||||
"""
|
||||
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||
|
||||
provider = MagicMock()
|
||||
provider.chat_with_retry = AsyncMock(side_effect=[
|
||||
LLMResponse(
|
||||
@@ -64,8 +64,6 @@ async def test_runner_does_not_abort_on_workspace_violation_anymore():
|
||||
|
||||
def test_is_ssrf_violation_recognizes_private_url_blocks():
|
||||
"""SSRF rejections are classified separately from workspace boundaries."""
|
||||
from nanobot.agent.runner import AgentRunner
|
||||
|
||||
ssrf_msg = "Error: Command blocked by safety guard (internal/private URL detected)"
|
||||
assert AgentRunner._is_ssrf_violation(ssrf_msg) is True
|
||||
assert AgentRunner._is_ssrf_violation(
|
||||
@@ -88,8 +86,6 @@ def test_is_ssrf_violation_recognizes_private_url_blocks():
|
||||
@pytest.mark.asyncio
|
||||
async def test_runner_returns_non_retryable_hint_on_ssrf_violation():
|
||||
"""SSRF stays blocked, but the runtime gives the LLM a final chance to recover."""
|
||||
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||
|
||||
provider = MagicMock()
|
||||
provider.chat_with_retry = AsyncMock(side_effect=[
|
||||
LLMResponse(
|
||||
@@ -107,7 +103,7 @@ async def test_runner_returns_non_retryable_hint_on_ssrf_violation():
|
||||
])
|
||||
tools = MagicMock()
|
||||
tools.get_definitions.return_value = []
|
||||
tools.execute = AsyncMock(return_value=(
|
||||
tools.execute = AsyncMock(return_value=ToolResult.error(
|
||||
"Error: Command blocked by safety guard (internal/private URL detected)"
|
||||
))
|
||||
|
||||
@@ -141,8 +137,6 @@ async def test_runner_lets_llm_recover_from_shell_guard_path_outside():
|
||||
turn (silent hang on Telegram per #3605); now the LLM gets the soft
|
||||
error back and can finalize on the next iteration.
|
||||
"""
|
||||
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||
|
||||
provider = MagicMock()
|
||||
captured_second_call: list[dict] = []
|
||||
|
||||
@@ -163,7 +157,9 @@ async def test_runner_lets_llm_recover_from_shell_guard_path_outside():
|
||||
tools = MagicMock()
|
||||
tools.get_definitions.return_value = []
|
||||
tools.execute = AsyncMock(
|
||||
return_value="Error: Command blocked by safety guard (path outside working dir)"
|
||||
return_value=ToolResult.error(
|
||||
"Error: Command blocked by safety guard (path outside working dir)"
|
||||
)
|
||||
)
|
||||
|
||||
runner = AgentRunner(provider)
|
||||
@@ -195,8 +191,6 @@ async def test_runner_throttles_repeated_workspace_bypass_attempts():
|
||||
the runner replaces the tool result with a hard "stop trying" message
|
||||
so the model finally gives up and surfaces the boundary to the user.
|
||||
"""
|
||||
from nanobot.agent.runner import AgentRunSpec, AgentRunner
|
||||
|
||||
bypass_attempts = [
|
||||
ToolCallRequest(
|
||||
id=f"a{i}", name="exec",
|
||||
@@ -215,7 +209,9 @@ async def test_runner_throttles_repeated_workspace_bypass_attempts():
|
||||
tools = MagicMock()
|
||||
tools.get_definitions.return_value = []
|
||||
tools.execute = AsyncMock(
|
||||
return_value="Error: Command blocked by safety guard (path outside working dir)"
|
||||
return_value=ToolResult.error(
|
||||
"Error: Command blocked by safety guard (path outside working dir)"
|
||||
)
|
||||
)
|
||||
|
||||
runner = AgentRunner(provider)
|
||||
|
||||
@@ -8,7 +8,9 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
import pytest
|
||||
|
||||
from nanobot.agent.runner import AgentRunner, AgentRunSpec
|
||||
from nanobot.agent.tools.base import Tool
|
||||
from nanobot.agent.tools.base import Tool, ToolResult
|
||||
from nanobot.agent.tools.context import ToolContext
|
||||
from nanobot.agent.tools.loader import ToolLoader
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
from nanobot.config.schema import AgentDefaults
|
||||
from nanobot.providers.base import LLMResponse, ToolCallRequest
|
||||
@@ -61,6 +63,40 @@ class _DelayTool(Tool):
|
||||
return self._name
|
||||
|
||||
|
||||
class _LegacyErrorPluginTool(Tool):
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return "legacy_plugin"
|
||||
|
||||
@property
|
||||
def description(self) -> str:
|
||||
return "legacy entry-point plugin"
|
||||
|
||||
@property
|
||||
def parameters(self) -> dict:
|
||||
return {"type": "object", "properties": {}, "required": []}
|
||||
|
||||
async def execute(self, **kwargs):
|
||||
return "Error: legacy plugin failed"
|
||||
|
||||
|
||||
class _StructuredSuccessPluginTool(Tool):
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return "structured_success_plugin"
|
||||
|
||||
@property
|
||||
def description(self) -> str:
|
||||
return "structured entry-point plugin"
|
||||
|
||||
@property
|
||||
def parameters(self) -> dict:
|
||||
return {"type": "object", "properties": {}, "required": []}
|
||||
|
||||
async def execute(self, **kwargs):
|
||||
return ToolResult("Error: generated report successfully")
|
||||
|
||||
|
||||
async def _run_optional_tool_response(response: LLMResponse):
|
||||
provider = MagicMock()
|
||||
calls = {"n": 0}
|
||||
@@ -91,6 +127,20 @@ async def _run_optional_tool_response(response: LLMResponse):
|
||||
return result, shared_events
|
||||
|
||||
|
||||
def _load_entry_point_plugin(tool_cls: type[Tool], tmp_path) -> ToolRegistry:
|
||||
mock_ep = MagicMock()
|
||||
mock_ep.name = tool_cls.__name__
|
||||
mock_ep.load.return_value = tool_cls
|
||||
|
||||
registry = ToolRegistry()
|
||||
with patch("nanobot.agent.tools.loader.entry_points", return_value=[mock_ep]):
|
||||
ToolLoader(test_classes=[]).load(
|
||||
ToolContext(config=None, workspace=str(tmp_path)),
|
||||
registry,
|
||||
)
|
||||
return registry
|
||||
|
||||
|
||||
def _tool_message(result, tool_call_id: str) -> dict:
|
||||
return [
|
||||
msg for msg in result.messages
|
||||
@@ -320,6 +370,63 @@ async def test_runner_rejects_openai_responses_array_arguments_without_executing
|
||||
assert "parameters must be a JSON object" in tool_message["content"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runner_treats_legacy_entry_point_error_prefix_as_tool_error(tmp_path):
|
||||
provider = MagicMock()
|
||||
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
|
||||
content="working",
|
||||
tool_calls=[ToolCallRequest(id="call_1", name="legacy_plugin", arguments={})],
|
||||
usage={},
|
||||
))
|
||||
|
||||
result = await AgentRunner(provider).run(AgentRunSpec(
|
||||
initial_messages=[{"role": "user", "content": "run plugin"}],
|
||||
tools=_load_entry_point_plugin(_LegacyErrorPluginTool, tmp_path),
|
||||
model="test-model",
|
||||
max_iterations=1,
|
||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||
fail_on_tool_error=True,
|
||||
))
|
||||
|
||||
assert result.stop_reason == "tool_error"
|
||||
assert result.tool_events == [
|
||||
{"name": "legacy_plugin", "status": "error", "detail": "Error: legacy plugin failed"}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runner_preserves_structured_plugin_success_that_starts_with_error(tmp_path):
|
||||
provider = MagicMock()
|
||||
provider.chat_with_retry = AsyncMock(side_effect=[
|
||||
LLMResponse(
|
||||
content="working",
|
||||
tool_calls=[
|
||||
ToolCallRequest(id="call_1", name="structured_success_plugin", arguments={})
|
||||
],
|
||||
usage={},
|
||||
),
|
||||
LLMResponse(content="done", tool_calls=[], usage={}),
|
||||
])
|
||||
|
||||
result = await AgentRunner(provider).run(AgentRunSpec(
|
||||
initial_messages=[{"role": "user", "content": "run plugin"}],
|
||||
tools=_load_entry_point_plugin(_StructuredSuccessPluginTool, tmp_path),
|
||||
model="test-model",
|
||||
max_iterations=2,
|
||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||
fail_on_tool_error=True,
|
||||
))
|
||||
|
||||
assert result.stop_reason == "completed"
|
||||
assert result.tool_events == [
|
||||
{
|
||||
"name": "structured_success_plugin",
|
||||
"status": "ok",
|
||||
"detail": "Error: generated report successfully",
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runner_blocks_repeated_external_fetches():
|
||||
provider = MagicMock()
|
||||
|
||||
@@ -127,6 +127,7 @@ class TestDispatch:
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_streaming_preserves_message_metadata(self):
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.outbound_events import StreamDeltaEvent, StreamEndEvent
|
||||
|
||||
loop, bus = _make_loop()
|
||||
msg = InboundMessage(
|
||||
@@ -156,10 +157,10 @@ class TestDispatch:
|
||||
|
||||
assert first.metadata["thread_root_event_id"] == "$root1"
|
||||
assert first.metadata["thread_reply_to_event_id"] == "$reply1"
|
||||
assert first.metadata["_stream_delta"] is True
|
||||
assert isinstance(first.event, StreamDeltaEvent)
|
||||
assert second.metadata["thread_root_event_id"] == "$root1"
|
||||
assert second.metadata["thread_reply_to_event_id"] == "$reply1"
|
||||
assert second.metadata["_stream_end"] is True
|
||||
assert isinstance(second.event, StreamEndEvent)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_processing_lock_serializes(self):
|
||||
|
||||
@@ -1,7 +1,11 @@
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from nanobot.agent.tools.base import Tool
|
||||
from nanobot.agent.tools.context import ToolContext
|
||||
from nanobot.agent.tools.loader import ToolLoader
|
||||
from nanobot.agent.tools.registry import ToolRegistry, is_tool_error_result
|
||||
|
||||
|
||||
def test_loader_discovers_entry_point_tools():
|
||||
@@ -74,3 +78,67 @@ def test_loader_skips_abstract_entry_point_tools():
|
||||
discovered = loader._discover_plugins()
|
||||
|
||||
assert "abstract_plugin" not in discovered
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_loader_entry_point_error_wrapper_preserves_tool_api(tmp_path):
|
||||
"""Only adapt legacy plugin error strings; keep the wrapped tool API intact."""
|
||||
mock_ep = MagicMock()
|
||||
mock_ep.name = "api_plugin"
|
||||
|
||||
class _ApiPluginTool(Tool):
|
||||
config_key = "api_plugin"
|
||||
|
||||
@property
|
||||
def name(self) -> str:
|
||||
return "api_plugin"
|
||||
|
||||
@property
|
||||
def description(self) -> str:
|
||||
return "Entry-point plugin with custom tool API methods."
|
||||
|
||||
@property
|
||||
def parameters(self) -> dict:
|
||||
return {"type": "object", "properties": {"value": {"type": "string"}}}
|
||||
|
||||
@property
|
||||
def read_only(self) -> bool:
|
||||
return True
|
||||
|
||||
@property
|
||||
def concurrency_safe(self) -> bool:
|
||||
return False
|
||||
|
||||
def cast_params(self, params: dict) -> dict:
|
||||
return {"value": str(params["value"])}
|
||||
|
||||
def validate_params(self, params: dict) -> list[str]:
|
||||
return [] if params == {"value": "1"} else ["bad value"]
|
||||
|
||||
def to_schema(self) -> dict:
|
||||
return {"name": self.name, "custom": True}
|
||||
|
||||
async def execute(self, **_):
|
||||
return "Error: plugin failed"
|
||||
|
||||
mock_ep.load.return_value = _ApiPluginTool
|
||||
|
||||
registry = ToolRegistry()
|
||||
with patch("nanobot.agent.tools.loader.entry_points", return_value=[mock_ep]):
|
||||
ToolLoader(test_classes=[]).load(
|
||||
ToolContext(config=None, workspace=str(tmp_path)),
|
||||
registry,
|
||||
)
|
||||
|
||||
tool = registry.get("api_plugin")
|
||||
assert tool is not None
|
||||
assert tool.config_key == "api_plugin"
|
||||
assert tool.read_only is True
|
||||
assert tool.concurrency_safe is False
|
||||
assert tool.cast_params({"value": 1}) == {"value": "1"}
|
||||
assert tool.validate_params({"value": "1"}) == []
|
||||
assert tool.to_schema() == {"name": "api_plugin", "custom": True}
|
||||
|
||||
result = await tool.execute(value="1")
|
||||
assert is_tool_error_result("api_plugin", result) is True
|
||||
assert str(result) == "Error: plugin failed"
|
||||
|
||||
@@ -13,6 +13,7 @@ from nanobot.agent.tools.long_task import (
|
||||
CompleteGoalTool,
|
||||
LongTaskTool,
|
||||
)
|
||||
from nanobot.bus.outbound_events import GoalStateSyncEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.bus.runtime_events import RuntimeEventBus
|
||||
from nanobot.session.goal_state import GOAL_STATE_KEY
|
||||
@@ -144,8 +145,8 @@ async def test_long_task_publishes_goal_state_ws_after_save(tmp_path):
|
||||
call = bus.publish_outbound.await_args.args[0]
|
||||
assert call.channel == "websocket"
|
||||
assert call.chat_id == "chat-99"
|
||||
assert call.metadata.get("_goal_state_sync") is True
|
||||
assert call.metadata["goal_state"] == {
|
||||
assert isinstance(call.event, GoalStateSyncEvent)
|
||||
assert call.event.goal_state == {
|
||||
"active": True,
|
||||
"ui_summary": "alpha",
|
||||
"objective": "Objective alpha",
|
||||
@@ -180,7 +181,8 @@ async def test_complete_goal_publishes_inactive_goal_state_ws(tmp_path):
|
||||
|
||||
bus.publish_outbound.assert_awaited_once()
|
||||
call = bus.publish_outbound.await_args.args[0]
|
||||
assert call.metadata["goal_state"] == {"active": False}
|
||||
assert isinstance(call.event, GoalStateSyncEvent)
|
||||
assert call.event.goal_state == {"active": False}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -7,7 +7,6 @@ from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.config.schema import AgentDefaults
|
||||
|
||||
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
|
||||
@@ -483,49 +482,3 @@ async def test_drain_pending_timeout(tmp_path):
|
||||
await hang_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_direct_routes_subagent_results_to_pending_queue(tmp_path):
|
||||
"""Single-message CLI mode should consume subagent announcements mid-turn."""
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.bus.events import InboundMessage
|
||||
from nanobot.bus.queue import MessageBus
|
||||
|
||||
loop = AgentLoop(
|
||||
bus=MessageBus(),
|
||||
provider=MagicMock(),
|
||||
workspace=tmp_path,
|
||||
model="test-model",
|
||||
)
|
||||
loop._connect_mcp = AsyncMock() # type: ignore[method-assign]
|
||||
|
||||
async def fake_process_message(msg, **kwargs):
|
||||
pending_queue = kwargs["pending_queue"]
|
||||
await loop.bus.publish_inbound(InboundMessage(
|
||||
channel="other",
|
||||
sender_id="u",
|
||||
chat_id="room",
|
||||
content="unrelated",
|
||||
))
|
||||
await loop.subagents._announce_result(
|
||||
"sub-1",
|
||||
"label",
|
||||
"task",
|
||||
"subagent result",
|
||||
{"channel": "cli", "chat_id": "direct", "session_key": "cli:direct"},
|
||||
"ok",
|
||||
)
|
||||
routed = await asyncio.wait_for(pending_queue.get(), timeout=1)
|
||||
assert "subagent result" in routed.content
|
||||
assert routed.metadata["subagent_task_id"] == "sub-1"
|
||||
return OutboundMessage(channel="cli", chat_id="direct", content="done")
|
||||
|
||||
loop._process_message = fake_process_message # type: ignore[method-assign]
|
||||
|
||||
response = await loop.process_direct("start", session_key="cli:direct")
|
||||
|
||||
assert response is not None
|
||||
assert response.content == "done"
|
||||
unrelated = await asyncio.wait_for(loop.bus.consume_inbound(), timeout=1)
|
||||
assert unrelated.content == "unrelated"
|
||||
|
||||
@@ -0,0 +1,244 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import (
|
||||
GoalStateSyncEvent,
|
||||
GoalStatusEvent,
|
||||
ProgressEvent,
|
||||
RetryWaitEvent,
|
||||
RuntimeModelUpdatedEvent,
|
||||
SessionUpdatedEvent,
|
||||
StreamDeltaEvent,
|
||||
StreamedResponseEvent,
|
||||
StreamEndEvent,
|
||||
TurnEndEvent,
|
||||
outbound_event_from_message,
|
||||
outbound_message_for_event,
|
||||
replace_outbound_event,
|
||||
)
|
||||
|
||||
|
||||
def test_progress_event_lives_on_outbound_message_event_field() -> None:
|
||||
tool_events = [{"phase": "start", "name": "read_file"}]
|
||||
file_edit_events = [{"phase": "end", "path": "app.py"}]
|
||||
|
||||
msg = outbound_message_for_event(
|
||||
channel="websocket",
|
||||
chat_id="chat-1",
|
||||
event=ProgressEvent(
|
||||
content="working",
|
||||
tool_hint=True,
|
||||
reasoning_delta=True,
|
||||
stream_id="r1",
|
||||
tool_events=tool_events,
|
||||
file_edit_events=file_edit_events,
|
||||
),
|
||||
metadata={"origin_message_id": "m1"},
|
||||
)
|
||||
|
||||
assert msg.content == "working"
|
||||
assert msg.metadata == {"origin_message_id": "m1"}
|
||||
|
||||
event = outbound_event_from_message(msg)
|
||||
assert isinstance(event, ProgressEvent)
|
||||
assert event.content == "working"
|
||||
assert event.tool_hint is True
|
||||
assert event.reasoning_delta is True
|
||||
assert event.stream_id == "r1"
|
||||
assert event.tool_events == tool_events
|
||||
assert event.file_edit_events == file_edit_events
|
||||
|
||||
|
||||
def test_normal_outbound_message_has_no_runtime_event() -> None:
|
||||
msg = OutboundMessage(channel="websocket", chat_id="chat-1", content="hello")
|
||||
|
||||
assert outbound_event_from_message(msg) is None
|
||||
|
||||
|
||||
def test_legacy_progress_metadata_flags_create_runtime_event() -> None:
|
||||
tool_events = [{"phase": "start", "name": "read_file"}]
|
||||
file_edit_events = [{"phase": "end", "path": "app.py"}]
|
||||
msg = OutboundMessage(
|
||||
channel="websocket",
|
||||
chat_id="chat-1",
|
||||
content="legacy progress",
|
||||
metadata={
|
||||
"_progress": True,
|
||||
"_tool_hint": True,
|
||||
"_reasoning_delta": True,
|
||||
"_stream_id": "r1",
|
||||
"_tool_events": tool_events,
|
||||
"_file_edit_events": file_edit_events,
|
||||
"message_id": "platform-routing-context",
|
||||
},
|
||||
)
|
||||
|
||||
event = outbound_event_from_message(msg)
|
||||
assert isinstance(event, ProgressEvent)
|
||||
assert event.content == "legacy progress"
|
||||
assert event.tool_hint is True
|
||||
assert event.reasoning_delta is True
|
||||
assert event.stream_id == "r1"
|
||||
assert event.tool_events == tool_events
|
||||
assert event.file_edit_events == file_edit_events
|
||||
|
||||
|
||||
def test_legacy_stream_metadata_flags_create_runtime_events() -> None:
|
||||
delta = OutboundMessage(
|
||||
channel="websocket",
|
||||
chat_id="chat-1",
|
||||
content="hello",
|
||||
metadata={"_stream_delta": True, "_stream_id": "s1"},
|
||||
)
|
||||
end = OutboundMessage(
|
||||
channel="websocket",
|
||||
chat_id="chat-1",
|
||||
content="",
|
||||
metadata={"_stream_end": True, "_stream_id": "s1", "_resuming": True},
|
||||
)
|
||||
|
||||
delta_event = outbound_event_from_message(delta)
|
||||
assert isinstance(delta_event, StreamDeltaEvent)
|
||||
assert delta_event.content == "hello"
|
||||
assert delta_event.stream_id == "s1"
|
||||
|
||||
end_event = outbound_event_from_message(end)
|
||||
assert isinstance(end_event, StreamEndEvent)
|
||||
assert end_event.stream_id == "s1"
|
||||
assert end_event.resuming is True
|
||||
|
||||
|
||||
def test_legacy_webui_runtime_metadata_flags_create_runtime_events() -> None:
|
||||
runtime = OutboundMessage(
|
||||
channel="websocket",
|
||||
chat_id="*",
|
||||
content="",
|
||||
metadata={
|
||||
"_runtime_model_updated": True,
|
||||
"model": "gpt-5.5",
|
||||
"model_preset": "high",
|
||||
},
|
||||
)
|
||||
goal_state = OutboundMessage(
|
||||
channel="websocket",
|
||||
chat_id="chat-1",
|
||||
content="",
|
||||
metadata={"_goal_state_sync": True, "goal_state": {"active": True}},
|
||||
)
|
||||
goal_status = OutboundMessage(
|
||||
channel="websocket",
|
||||
chat_id="chat-1",
|
||||
content="",
|
||||
metadata={"_goal_status": True, "goal_status": "running", "started_at": 1.25},
|
||||
)
|
||||
turn_end = OutboundMessage(
|
||||
channel="websocket",
|
||||
chat_id="chat-1",
|
||||
content="",
|
||||
metadata={"_turn_end": True, "latency_ms": 42.0, "goal_state": {"active": False}},
|
||||
)
|
||||
session_updated = OutboundMessage(
|
||||
channel="websocket",
|
||||
chat_id="chat-1",
|
||||
content="",
|
||||
metadata={"_session_updated": True, "_session_update_scope": "metadata"},
|
||||
)
|
||||
|
||||
runtime_event = outbound_event_from_message(runtime)
|
||||
assert isinstance(runtime_event, RuntimeModelUpdatedEvent)
|
||||
assert runtime_event.model == "gpt-5.5"
|
||||
assert runtime_event.model_preset == "high"
|
||||
|
||||
goal_state_event = outbound_event_from_message(goal_state)
|
||||
assert isinstance(goal_state_event, GoalStateSyncEvent)
|
||||
assert goal_state_event.goal_state == {"active": True}
|
||||
|
||||
goal_status_event = outbound_event_from_message(goal_status)
|
||||
assert isinstance(goal_status_event, GoalStatusEvent)
|
||||
assert goal_status_event.status == "running"
|
||||
assert goal_status_event.started_at == 1.25
|
||||
|
||||
turn_end_event = outbound_event_from_message(turn_end)
|
||||
assert isinstance(turn_end_event, TurnEndEvent)
|
||||
assert turn_end_event.latency_ms == 42
|
||||
assert turn_end_event.goal_state == {"active": False}
|
||||
|
||||
session_updated_event = outbound_event_from_message(session_updated)
|
||||
assert isinstance(session_updated_event, SessionUpdatedEvent)
|
||||
assert session_updated_event.scope == "metadata"
|
||||
|
||||
|
||||
def test_legacy_metadata_numbers_ignore_bool_values() -> None:
|
||||
goal_status = OutboundMessage(
|
||||
channel="websocket",
|
||||
chat_id="chat-1",
|
||||
content="",
|
||||
metadata={"_goal_status": True, "goal_status": "running", "started_at": True},
|
||||
)
|
||||
turn_end = OutboundMessage(
|
||||
channel="websocket",
|
||||
chat_id="chat-1",
|
||||
content="",
|
||||
metadata={"_turn_end": True, "latency_ms": True},
|
||||
)
|
||||
|
||||
goal_status_event = outbound_event_from_message(goal_status)
|
||||
assert isinstance(goal_status_event, GoalStatusEvent)
|
||||
assert goal_status_event.started_at is None
|
||||
|
||||
turn_end_event = outbound_event_from_message(turn_end)
|
||||
assert isinstance(turn_end_event, TurnEndEvent)
|
||||
assert turn_end_event.latency_ms is None
|
||||
|
||||
|
||||
def test_legacy_retry_wait_and_streamed_flags_create_runtime_events() -> None:
|
||||
retry = OutboundMessage(
|
||||
channel="cli",
|
||||
chat_id="direct",
|
||||
content="waiting",
|
||||
metadata={"_retry_wait": True},
|
||||
)
|
||||
streamed = OutboundMessage(
|
||||
channel="cli",
|
||||
chat_id="direct",
|
||||
content="final answer",
|
||||
metadata={"_streamed": True},
|
||||
)
|
||||
|
||||
retry_event = outbound_event_from_message(retry)
|
||||
assert isinstance(retry_event, RetryWaitEvent)
|
||||
assert retry_event.content == "waiting"
|
||||
assert isinstance(outbound_event_from_message(streamed), StreamedResponseEvent)
|
||||
|
||||
|
||||
def test_replace_outbound_event_keeps_routing_metadata() -> None:
|
||||
msg = outbound_message_for_event(
|
||||
channel="websocket",
|
||||
chat_id="chat-1",
|
||||
event=StreamDeltaEvent(content="hello", stream_id="s1"),
|
||||
metadata={"message_id": "m1"},
|
||||
)
|
||||
|
||||
updated = replace_outbound_event(
|
||||
msg,
|
||||
StreamEndEvent(stream_id="s1", resuming=True),
|
||||
content="hello world",
|
||||
)
|
||||
|
||||
assert updated.content == "hello world"
|
||||
assert updated.metadata == {"message_id": "m1"}
|
||||
assert isinstance(updated.event, StreamEndEvent)
|
||||
assert updated.event.stream_id == "s1"
|
||||
assert updated.event.resuming is True
|
||||
|
||||
|
||||
def test_streamed_response_event_keeps_final_content_outside_event_payload() -> None:
|
||||
msg = outbound_message_for_event(
|
||||
channel="cli",
|
||||
chat_id="direct",
|
||||
event=StreamedResponseEvent(),
|
||||
content="final answer",
|
||||
)
|
||||
|
||||
assert msg.content == "final answer"
|
||||
assert isinstance(outbound_event_from_message(msg), StreamedResponseEvent)
|
||||
@@ -1,10 +1,19 @@
|
||||
"""Tests for ChannelManager delta coalescing to reduce streaming latency."""
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import (
|
||||
ProgressEvent,
|
||||
RetryWaitEvent,
|
||||
StreamDeltaEvent,
|
||||
StreamEndEvent,
|
||||
outbound_event_from_message,
|
||||
outbound_message_for_event,
|
||||
)
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.channels.manager import ChannelManager
|
||||
@@ -29,221 +38,187 @@ class MockChannel(BaseChannel):
|
||||
pass
|
||||
|
||||
async def send(self, msg):
|
||||
"""Implement abstract method."""
|
||||
return await self._send_mock(msg)
|
||||
|
||||
async def send_delta(self, chat_id, delta, metadata=None):
|
||||
"""Override send_delta for testing."""
|
||||
return await self._send_delta_mock(chat_id, delta, metadata)
|
||||
async def send_delta(
|
||||
self,
|
||||
chat_id,
|
||||
delta,
|
||||
metadata=None,
|
||||
*,
|
||||
stream_id=None,
|
||||
stream_end=False,
|
||||
resuming=False,
|
||||
):
|
||||
return await self._send_delta_mock(
|
||||
chat_id,
|
||||
delta,
|
||||
metadata,
|
||||
stream_id=stream_id,
|
||||
stream_end=stream_end,
|
||||
resuming=resuming,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def config():
|
||||
"""Create a minimal config for testing."""
|
||||
return Config()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def bus():
|
||||
"""Create a message bus for testing."""
|
||||
return MessageBus()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def manager(config, bus):
|
||||
"""Create a channel manager with a mock channel."""
|
||||
manager = ChannelManager(config, bus)
|
||||
manager.channels["mock"] = MockChannel({}, bus)
|
||||
return manager
|
||||
|
||||
|
||||
def _delta(content: str, *, chat_id: str = "chat1", stream_id: str | None = None):
|
||||
return outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id=chat_id,
|
||||
event=StreamDeltaEvent(content=content, stream_id=stream_id),
|
||||
)
|
||||
|
||||
|
||||
def _end(
|
||||
content: str = "",
|
||||
*,
|
||||
chat_id: str = "chat1",
|
||||
stream_id: str | None = None,
|
||||
resuming: bool = False,
|
||||
):
|
||||
return outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id=chat_id,
|
||||
event=StreamEndEvent(content=content, stream_id=stream_id, resuming=resuming),
|
||||
)
|
||||
|
||||
|
||||
class TestDeltaCoalescing:
|
||||
"""Tests for _stream_delta message coalescing."""
|
||||
"""Tests for stream delta message coalescing."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_single_delta_not_coalesced(self, manager, bus):
|
||||
"""A single delta should be sent as-is."""
|
||||
msg = OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="Hello",
|
||||
metadata={"_stream_delta": True},
|
||||
)
|
||||
msg = _delta("Hello")
|
||||
await bus.publish_outbound(msg)
|
||||
|
||||
# Process one message
|
||||
async def process_one():
|
||||
try:
|
||||
m = await asyncio.wait_for(bus.consume_outbound(), timeout=0.1)
|
||||
if m.metadata.get("_stream_delta"):
|
||||
event = outbound_event_from_message(m)
|
||||
if isinstance(event, StreamDeltaEvent):
|
||||
m, pending = manager._coalesce_stream_deltas(m)
|
||||
# Put pending back (none expected)
|
||||
for p in pending:
|
||||
await bus.publish_outbound(p)
|
||||
channel = manager.channels.get(m.channel)
|
||||
if channel:
|
||||
await channel.send_delta(m.chat_id, m.content, m.metadata)
|
||||
event = outbound_event_from_message(m)
|
||||
if channel and isinstance(event, StreamDeltaEvent):
|
||||
await channel.send_delta(
|
||||
m.chat_id,
|
||||
m.content,
|
||||
m.metadata,
|
||||
stream_id=event.stream_id,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
pass
|
||||
|
||||
await process_one()
|
||||
|
||||
manager.channels["mock"]._send_delta_mock.assert_called_once_with(
|
||||
"chat1", "Hello", {"_stream_delta": True}
|
||||
"chat1",
|
||||
"Hello",
|
||||
{},
|
||||
stream_id=None,
|
||||
stream_end=False,
|
||||
resuming=False,
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multiple_deltas_coalesced(self, manager, bus):
|
||||
"""Multiple consecutive deltas for same chat should be merged."""
|
||||
# Put multiple deltas in queue
|
||||
for text in ["Hello", " ", "world", "!"]:
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content=text,
|
||||
metadata={"_stream_delta": True},
|
||||
))
|
||||
await bus.publish_outbound(_delta(text))
|
||||
|
||||
# Process using coalescing logic
|
||||
first_msg = await bus.consume_outbound()
|
||||
merged, pending = manager._coalesce_stream_deltas(first_msg)
|
||||
|
||||
# Should have merged all deltas
|
||||
assert merged.content == "Hello world!"
|
||||
assert merged.metadata.get("_stream_delta") is True
|
||||
# No pending messages (all were coalesced)
|
||||
assert isinstance(merged.event, StreamDeltaEvent)
|
||||
assert len(pending) == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deltas_different_chats_not_coalesced(self, manager, bus):
|
||||
"""Deltas for different chats should not be merged."""
|
||||
# Put deltas for different chats
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="Hello",
|
||||
metadata={"_stream_delta": True},
|
||||
))
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat2",
|
||||
content="World",
|
||||
metadata={"_stream_delta": True},
|
||||
))
|
||||
await bus.publish_outbound(_delta("Hello", chat_id="chat1"))
|
||||
await bus.publish_outbound(_delta("World", chat_id="chat2"))
|
||||
|
||||
first_msg = await bus.consume_outbound()
|
||||
merged, pending = manager._coalesce_stream_deltas(first_msg)
|
||||
|
||||
# First chat should not include second chat's content
|
||||
assert merged.content == "Hello"
|
||||
assert merged.chat_id == "chat1"
|
||||
# Second chat should be in pending
|
||||
assert len(pending) == 1
|
||||
assert pending[0].chat_id == "chat2"
|
||||
assert pending[0].content == "World"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deltas_different_stream_ids_not_coalesced(self, manager, bus):
|
||||
"""Deltas for the same chat but different streams should not be merged."""
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="A1",
|
||||
metadata={"_stream_delta": True, "_stream_id": "stream-a"},
|
||||
))
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="B1",
|
||||
metadata={"_stream_delta": True, "_stream_id": "stream-b"},
|
||||
))
|
||||
await bus.publish_outbound(_delta("A1", stream_id="stream-a"))
|
||||
await bus.publish_outbound(_delta("B1", stream_id="stream-b"))
|
||||
|
||||
first_msg = await bus.consume_outbound()
|
||||
merged, pending = manager._coalesce_stream_deltas(first_msg)
|
||||
|
||||
assert merged.content == "A1"
|
||||
assert merged.metadata.get("_stream_id") == "stream-a"
|
||||
assert isinstance(merged.event, StreamDeltaEvent)
|
||||
assert merged.event.stream_id == "stream-a"
|
||||
assert len(pending) == 1
|
||||
assert pending[0].content == "B1"
|
||||
assert pending[0].metadata.get("_stream_id") == "stream-b"
|
||||
assert isinstance(pending[0].event, StreamDeltaEvent)
|
||||
assert pending[0].event.stream_id == "stream-b"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_end_terminates_coalescing(self, manager, bus):
|
||||
"""_stream_end should stop coalescing and be included in final message."""
|
||||
# Put deltas with stream_end at the end
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="Hello",
|
||||
metadata={"_stream_delta": True},
|
||||
))
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content=" world",
|
||||
metadata={"_stream_delta": True, "_stream_end": True},
|
||||
))
|
||||
await bus.publish_outbound(_delta("Hello"))
|
||||
await bus.publish_outbound(_end(" world"))
|
||||
|
||||
first_msg = await bus.consume_outbound()
|
||||
merged, pending = manager._coalesce_stream_deltas(first_msg)
|
||||
|
||||
# Should have merged content
|
||||
assert merged.content == "Hello world"
|
||||
# Should have stream_end flag
|
||||
assert merged.metadata.get("_stream_end") is True
|
||||
# No pending
|
||||
assert isinstance(merged.event, StreamEndEvent)
|
||||
assert len(pending) == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_coalescing_stops_at_first_non_matching_boundary(self, manager, bus):
|
||||
"""Only consecutive deltas should be merged; later deltas stay queued."""
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="Hello",
|
||||
metadata={"_stream_delta": True, "_stream_id": "seg-1"},
|
||||
))
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="",
|
||||
metadata={"_stream_end": True, "_stream_id": "seg-1"},
|
||||
))
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="world",
|
||||
metadata={"_stream_delta": True, "_stream_id": "seg-2"},
|
||||
))
|
||||
await bus.publish_outbound(_delta("Hello", stream_id="seg-1"))
|
||||
await bus.publish_outbound(_end(stream_id="seg-1"))
|
||||
await bus.publish_outbound(_delta("world", stream_id="seg-2"))
|
||||
|
||||
first_msg = await bus.consume_outbound()
|
||||
merged, pending = manager._coalesce_stream_deltas(first_msg)
|
||||
|
||||
assert merged.content == "Hello"
|
||||
assert merged.metadata.get("_stream_end") is None
|
||||
assert isinstance(merged.event, StreamDeltaEvent)
|
||||
assert len(pending) == 1
|
||||
assert pending[0].metadata.get("_stream_end") is True
|
||||
assert pending[0].metadata.get("_stream_id") == "seg-1"
|
||||
assert isinstance(pending[0].event, StreamEndEvent)
|
||||
assert pending[0].event.stream_id == "seg-1"
|
||||
|
||||
# The next stream segment must remain in queue order for later dispatch.
|
||||
remaining = await bus.consume_outbound()
|
||||
assert remaining.content == "world"
|
||||
assert remaining.metadata.get("_stream_id") == "seg-2"
|
||||
assert isinstance(remaining.event, StreamDeltaEvent)
|
||||
assert remaining.event.stream_id == "seg-2"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_delta_message_preserved(self, manager, bus):
|
||||
"""Non-delta messages should be preserved in pending list."""
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="Delta",
|
||||
metadata={"_stream_delta": True},
|
||||
))
|
||||
await bus.publish_outbound(_delta("Delta"))
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="Final message",
|
||||
metadata={}, # Not a delta
|
||||
))
|
||||
|
||||
first_msg = await bus.consume_outbound()
|
||||
@@ -252,17 +227,11 @@ class TestDeltaCoalescing:
|
||||
assert merged.content == "Delta"
|
||||
assert len(pending) == 1
|
||||
assert pending[0].content == "Final message"
|
||||
assert pending[0].metadata.get("_stream_delta") is None
|
||||
assert pending[0].event is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_queue_stops_coalescing(self, manager, bus):
|
||||
"""Coalescing should stop when queue is empty."""
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="Only message",
|
||||
metadata={"_stream_delta": True},
|
||||
))
|
||||
await bus.publish_outbound(_delta("Only message"))
|
||||
|
||||
first_msg = await bus.consume_outbound()
|
||||
merged, pending = manager._coalesce_stream_deltas(first_msg)
|
||||
@@ -276,49 +245,35 @@ class TestDispatchOutboundWithCoalescing:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_coalesces_and_processes_pending(self, manager, bus):
|
||||
"""_dispatch_outbound should coalesce deltas and process pending messages."""
|
||||
# Put multiple deltas followed by a regular message
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="A",
|
||||
metadata={"_stream_delta": True},
|
||||
))
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="B",
|
||||
metadata={"_stream_delta": True},
|
||||
))
|
||||
await bus.publish_outbound(_delta("A"))
|
||||
await bus.publish_outbound(_delta("B"))
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="Final",
|
||||
metadata={}, # Regular message
|
||||
))
|
||||
|
||||
# Run one iteration of dispatch logic manually
|
||||
pending = []
|
||||
processed = []
|
||||
|
||||
# First iteration: should coalesce A+B
|
||||
if pending:
|
||||
msg = pending.pop(0)
|
||||
else:
|
||||
msg = await bus.consume_outbound()
|
||||
|
||||
if msg.metadata.get("_stream_delta") and not msg.metadata.get("_stream_end"):
|
||||
msg = pending.pop(0) if pending else await bus.consume_outbound()
|
||||
event = outbound_event_from_message(msg)
|
||||
if isinstance(event, StreamDeltaEvent):
|
||||
msg, extra_pending = manager._coalesce_stream_deltas(msg)
|
||||
pending.extend(extra_pending)
|
||||
|
||||
channel = manager.channels.get(msg.channel)
|
||||
if channel:
|
||||
await channel.send_delta(msg.chat_id, msg.content, msg.metadata)
|
||||
event = outbound_event_from_message(msg)
|
||||
if channel and isinstance(event, StreamDeltaEvent):
|
||||
await channel.send_delta(
|
||||
msg.chat_id,
|
||||
msg.content,
|
||||
msg.metadata,
|
||||
stream_id=event.stream_id,
|
||||
)
|
||||
processed.append(("delta", msg.content))
|
||||
|
||||
# Should have sent coalesced delta
|
||||
assert processed == [("delta", "AB")]
|
||||
# Should have pending regular message
|
||||
assert len(pending) == 1
|
||||
assert pending[0].content == "Final"
|
||||
|
||||
@@ -354,23 +309,20 @@ class TestProgressFiltering:
|
||||
|
||||
assert manager._resolve_bool_override(FakeSection(), "send_progress", True) is False
|
||||
assert manager._resolve_bool_override(FakeSection(), "send_tool_hints", False) is True
|
||||
# Missing attribute falls back to default
|
||||
assert manager._resolve_bool_override(FakeSection(), "unknown_key", True) is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_channel_override_can_drop_progress_message(self, manager, bus):
|
||||
manager.channels["mock"].send_progress = False
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
await bus.publish_outbound(outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="thinking",
|
||||
metadata={"_progress": True},
|
||||
event=ProgressEvent(content="thinking"),
|
||||
))
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="final answer",
|
||||
metadata={},
|
||||
))
|
||||
|
||||
task = asyncio.create_task(manager._dispatch_outbound())
|
||||
@@ -391,13 +343,37 @@ class TestProgressFiltering:
|
||||
assert send_mock.await_args_list[0].args[0].content == "final answer"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_channel_override_can_enable_tool_hints(self, manager, bus):
|
||||
manager.channels["mock"].send_tool_hints = True
|
||||
async def test_legacy_progress_flag_uses_runtime_progress_filter(self, manager, bus):
|
||||
manager.channels["mock"].send_progress = False
|
||||
await bus.publish_outbound(OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="read_file(foo.py)",
|
||||
metadata={"_progress": True, "_tool_hint": True},
|
||||
content="legacy progress-shaped message",
|
||||
metadata={"_progress": True},
|
||||
))
|
||||
|
||||
task = asyncio.create_task(manager._dispatch_outbound())
|
||||
try:
|
||||
for _ in range(30):
|
||||
if manager.channels["mock"]._send_mock.await_count >= 1:
|
||||
break
|
||||
await asyncio.sleep(0.05)
|
||||
finally:
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
assert manager.channels["mock"]._send_mock.await_count == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_channel_override_can_enable_tool_hints(self, manager, bus):
|
||||
manager.channels["mock"].send_tool_hints = True
|
||||
await bus.publish_outbound(outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
event=ProgressEvent(content="read_file(foo.py)", tool_hint=True),
|
||||
))
|
||||
|
||||
task = asyncio.create_task(manager._dispatch_outbound())
|
||||
@@ -423,24 +399,15 @@ class TestRetryWaitFiltering:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_retry_wait_message_dropped(self, manager, bus):
|
||||
"""A ``_retry_wait`` message must be filtered before channel dispatch.
|
||||
|
||||
Regression: provider retry diagnostics like
|
||||
``Model request failed, retry in 1s (attempt 1).`` were being
|
||||
delivered to end-user channels because the runner bound
|
||||
``on_retry_wait`` to the progress callback.
|
||||
"""
|
||||
retry_msg = OutboundMessage(
|
||||
retry_msg = outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="Model request failed, retry in 1s (attempt 1).",
|
||||
metadata={"_retry_wait": True},
|
||||
event=RetryWaitEvent(content="Model request failed, retry in 1s (attempt 1)."),
|
||||
)
|
||||
real_msg = OutboundMessage(
|
||||
channel="mock",
|
||||
chat_id="chat1",
|
||||
content="final answer",
|
||||
metadata={},
|
||||
)
|
||||
await bus.publish_outbound(retry_msg)
|
||||
await bus.publish_outbound(real_msg)
|
||||
@@ -462,4 +429,4 @@ class TestRetryWaitFiltering:
|
||||
assert send_mock.await_count == 1
|
||||
sent = send_mock.await_args_list[0].args[0]
|
||||
assert sent.content == "final answer"
|
||||
assert not sent.metadata.get("_retry_wait")
|
||||
assert sent.event is None
|
||||
|
||||
@@ -8,10 +8,9 @@ channels that opt in via ``channel.show_reasoning``; plugins without a
|
||||
low-emphasis UI primitive keep the base no-op and the content silently
|
||||
drops at dispatch.
|
||||
|
||||
One-shot ``_reasoning`` frames are accepted for back-compat with hooks
|
||||
that haven't migrated yet — ``BaseChannel.send_reasoning`` expands them
|
||||
to a single delta + end pair so plugins only implement the streaming
|
||||
primitives.
|
||||
One-shot reasoning frames are represented as typed progress events and
|
||||
``BaseChannel.send_reasoning`` expands them to a single delta + end pair so
|
||||
plugins only implement the streaming primitives.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -22,6 +21,7 @@ from unittest.mock import AsyncMock
|
||||
import pytest
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent, outbound_message_for_event
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.channels.manager import ChannelManager
|
||||
@@ -48,11 +48,11 @@ class _MockChannel(BaseChannel):
|
||||
async def send(self, msg):
|
||||
return await self._send_mock(msg)
|
||||
|
||||
async def send_reasoning_delta(self, chat_id, delta, metadata=None):
|
||||
return await self._delta_mock(chat_id, delta, metadata)
|
||||
async def send_reasoning_delta(self, chat_id, delta, metadata=None, *, stream_id=None):
|
||||
return await self._delta_mock(chat_id, delta, metadata, stream_id=stream_id)
|
||||
|
||||
async def send_reasoning_end(self, chat_id, metadata=None):
|
||||
return await self._end_mock(chat_id, metadata)
|
||||
async def send_reasoning_end(self, chat_id, metadata=None, *, stream_id=None):
|
||||
return await self._end_mock(chat_id, metadata, stream_id=stream_id)
|
||||
|
||||
async def send_file_edit_events(self, chat_id, edits, metadata=None):
|
||||
return await self._file_edit_mock(chat_id, edits, metadata)
|
||||
@@ -94,17 +94,17 @@ def test_websocket_gateway_uses_configured_workspace_restriction(tmp_path, monke
|
||||
@pytest.mark.asyncio
|
||||
async def test_reasoning_delta_routes_to_send_reasoning_delta(manager):
|
||||
channel = manager.channels["mock"]
|
||||
msg = OutboundMessage(
|
||||
msg = outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id="c1",
|
||||
content="step-by-step",
|
||||
metadata={"_progress": True, "_reasoning_delta": True, "_stream_id": "r1"},
|
||||
event=ProgressEvent(content="step-by-step", reasoning_delta=True, stream_id="r1"),
|
||||
)
|
||||
await manager._send_once(channel, msg)
|
||||
channel._delta_mock.assert_awaited_once()
|
||||
args = channel._delta_mock.await_args.args
|
||||
assert args[0] == "c1"
|
||||
assert args[1] == "step-by-step"
|
||||
assert channel._delta_mock.await_args.kwargs["stream_id"] == "r1"
|
||||
channel._send_mock.assert_not_awaited()
|
||||
channel._end_mock.assert_not_awaited()
|
||||
|
||||
@@ -112,11 +112,10 @@ async def test_reasoning_delta_routes_to_send_reasoning_delta(manager):
|
||||
@pytest.mark.asyncio
|
||||
async def test_reasoning_end_routes_to_send_reasoning_end(manager):
|
||||
channel = manager.channels["mock"]
|
||||
msg = OutboundMessage(
|
||||
msg = outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id="c1",
|
||||
content="",
|
||||
metadata={"_progress": True, "_reasoning_end": True, "_stream_id": "r1"},
|
||||
event=ProgressEvent(reasoning_end=True, stream_id="r1"),
|
||||
)
|
||||
await manager._send_once(channel, msg)
|
||||
channel._end_mock.assert_awaited_once()
|
||||
@@ -124,16 +123,13 @@ async def test_reasoning_end_routes_to_send_reasoning_end(manager):
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_legacy_one_shot_reasoning_expands_to_delta_plus_end(manager):
|
||||
"""`_reasoning` (no delta/end pair) falls back through `send_reasoning`
|
||||
which the base class expands to a single delta + end. Hooks that haven't
|
||||
migrated still surface in WebUI as a complete stream segment."""
|
||||
async def test_one_shot_reasoning_expands_to_delta_plus_end(manager):
|
||||
"""One-shot reasoning expands to a single delta + end."""
|
||||
channel = manager.channels["mock"]
|
||||
msg = OutboundMessage(
|
||||
msg = outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id="c1",
|
||||
content="one-shot reasoning",
|
||||
metadata={"_progress": True, "_reasoning": True},
|
||||
event=ProgressEvent(content="one-shot reasoning", reasoning=True),
|
||||
)
|
||||
await manager._send_once(channel, msg)
|
||||
channel._delta_mock.assert_awaited_once()
|
||||
@@ -144,11 +140,10 @@ async def test_legacy_one_shot_reasoning_expands_to_delta_plus_end(manager):
|
||||
async def test_dispatch_drops_reasoning_when_channel_opts_out(manager):
|
||||
channel = manager.channels["mock"]
|
||||
channel.show_reasoning = False
|
||||
msg = OutboundMessage(
|
||||
msg = outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id="c1",
|
||||
content="hidden thinking",
|
||||
metadata={"_progress": True, "_reasoning_delta": True},
|
||||
event=ProgressEvent(content="hidden thinking", reasoning_delta=True),
|
||||
)
|
||||
await manager.bus.publish_outbound(msg)
|
||||
|
||||
@@ -164,17 +159,15 @@ async def test_dispatch_delivers_reasoning_when_channel_opts_in(manager):
|
||||
channel = manager.channels["mock"]
|
||||
channel.show_reasoning = True
|
||||
for chunk in ("first ", "second"):
|
||||
await manager.bus.publish_outbound(OutboundMessage(
|
||||
await manager.bus.publish_outbound(outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id="c1",
|
||||
content=chunk,
|
||||
metadata={"_progress": True, "_reasoning_delta": True, "_stream_id": "r1"},
|
||||
event=ProgressEvent(content=chunk, reasoning_delta=True, stream_id="r1"),
|
||||
))
|
||||
await manager.bus.publish_outbound(OutboundMessage(
|
||||
await manager.bus.publish_outbound(outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id="c1",
|
||||
content="",
|
||||
metadata={"_progress": True, "_reasoning_end": True, "_stream_id": "r1"},
|
||||
event=ProgressEvent(reasoning_end=True, stream_id="r1"),
|
||||
))
|
||||
|
||||
await _pump_one(manager)
|
||||
@@ -185,11 +178,10 @@ async def test_dispatch_delivers_reasoning_when_channel_opts_in(manager):
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_silently_drops_reasoning_for_unknown_channel(manager):
|
||||
msg = OutboundMessage(
|
||||
msg = outbound_message_for_event(
|
||||
channel="ghost",
|
||||
chat_id="c1",
|
||||
content="nobody home",
|
||||
metadata={"_progress": True, "_reasoning_delta": True},
|
||||
event=ProgressEvent(content="nobody home", reasoning_delta=True),
|
||||
)
|
||||
await manager.bus.publish_outbound(msg)
|
||||
|
||||
@@ -229,17 +221,34 @@ async def test_base_channel_reasoning_primitives_are_noop_safe():
|
||||
async def test_file_edit_events_route_to_channel_capability(manager):
|
||||
channel = manager.channels["mock"]
|
||||
edits = [{"version": 1, "phase": "start", "path": "src/app.py"}]
|
||||
msg = OutboundMessage(
|
||||
msg = outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id="c1",
|
||||
content="",
|
||||
metadata={"_progress": True, "_file_edit_events": edits},
|
||||
event=ProgressEvent(file_edit_events=edits),
|
||||
)
|
||||
|
||||
await manager._send_once(channel, msg)
|
||||
|
||||
channel._file_edit_mock.assert_awaited_once_with(
|
||||
"c1", edits, {"_progress": True, "_file_edit_events": edits}
|
||||
"c1", edits, msg.metadata
|
||||
)
|
||||
channel._send_mock.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_typed_file_edit_event_routes_to_channel_capability(manager):
|
||||
channel = manager.channels["mock"]
|
||||
edits = [{"version": 1, "phase": "start", "path": "src/app.py"}]
|
||||
msg = outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id="c1",
|
||||
event=ProgressEvent(file_edit_events=edits),
|
||||
)
|
||||
|
||||
await manager._send_once(channel, msg)
|
||||
|
||||
channel._file_edit_mock.assert_awaited_once_with(
|
||||
"c1", edits, msg.metadata
|
||||
)
|
||||
channel._send_mock.assert_not_awaited()
|
||||
|
||||
@@ -270,11 +279,10 @@ async def test_reasoning_routing_does_not_consult_send_progress(manager):
|
||||
channel = manager.channels["mock"]
|
||||
channel.send_progress = False
|
||||
channel.show_reasoning = True
|
||||
await manager.bus.publish_outbound(OutboundMessage(
|
||||
await manager.bus.publish_outbound(outbound_message_for_event(
|
||||
channel="mock",
|
||||
chat_id="c1",
|
||||
content="still surfaces",
|
||||
metadata={"_progress": True, "_reasoning_delta": True},
|
||||
event=ProgressEvent(content="still surfaces", reasoning_delta=True),
|
||||
))
|
||||
|
||||
await _pump_one(manager)
|
||||
|
||||
@@ -9,6 +9,13 @@ from unittest.mock import AsyncMock, patch
|
||||
import pytest
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import (
|
||||
ProgressEvent,
|
||||
StreamDeltaEvent,
|
||||
StreamedResponseEvent,
|
||||
StreamEndEvent,
|
||||
outbound_message_for_event,
|
||||
)
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.base import BaseChannel
|
||||
from nanobot.channels.manager import ChannelManager
|
||||
@@ -718,7 +725,7 @@ async def test_send_with_retry_no_retry_when_max_is_zero():
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_with_retry_calls_send_delta():
|
||||
"""_send_with_retry should call send_delta when metadata has _stream_delta."""
|
||||
"""_send_with_retry should call send_delta for stream delta events."""
|
||||
send_delta_called = False
|
||||
|
||||
class _StreamingChannel(BaseChannel):
|
||||
@@ -734,7 +741,16 @@ async def test_send_with_retry_calls_send_delta():
|
||||
async def send(self, msg: OutboundMessage) -> None:
|
||||
pass # Should not be called
|
||||
|
||||
async def send_delta(self, chat_id: str, delta: str, metadata: dict | None = None) -> None:
|
||||
async def send_delta(
|
||||
self,
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
stream_end: bool = False,
|
||||
resuming: bool = False,
|
||||
) -> None:
|
||||
nonlocal send_delta_called
|
||||
send_delta_called = True
|
||||
|
||||
@@ -749,18 +765,147 @@ async def test_send_with_retry_calls_send_delta():
|
||||
mgr.channels = {"streaming": _StreamingChannel(fake_config, mgr.bus)}
|
||||
mgr._dispatch_task = None
|
||||
|
||||
msg = OutboundMessage(
|
||||
channel="streaming", chat_id="123", content="test delta",
|
||||
metadata={"_stream_delta": True}
|
||||
msg = outbound_message_for_event(
|
||||
channel="streaming",
|
||||
chat_id="123",
|
||||
event=StreamDeltaEvent(content="test delta"),
|
||||
)
|
||||
await mgr._send_with_retry(mgr.channels["streaming"], msg)
|
||||
|
||||
assert send_delta_called is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_with_retry_supports_legacy_stream_delta_signature():
|
||||
"""External plugins with the old send_delta signature should keep working."""
|
||||
calls: list[tuple[str, str, dict]] = []
|
||||
|
||||
class _LegacyStreamingChannel(BaseChannel):
|
||||
name = "legacy_streaming"
|
||||
display_name = "Legacy Streaming"
|
||||
|
||||
async def start(self) -> None:
|
||||
pass
|
||||
|
||||
async def stop(self) -> None:
|
||||
pass
|
||||
|
||||
async def send(self, msg: OutboundMessage) -> None:
|
||||
pass
|
||||
|
||||
async def send_delta(
|
||||
self,
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict | None = None,
|
||||
) -> None:
|
||||
calls.append((chat_id, delta, dict(metadata or {})))
|
||||
|
||||
fake_config = SimpleNamespace(
|
||||
channels=ChannelsConfig(send_max_retries=3),
|
||||
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||
)
|
||||
mgr = ChannelManager.__new__(ChannelManager)
|
||||
mgr.config = fake_config
|
||||
mgr.bus = MessageBus()
|
||||
mgr.channels = {"legacy_streaming": _LegacyStreamingChannel(fake_config, mgr.bus)}
|
||||
mgr._dispatch_task = None
|
||||
|
||||
await mgr._send_with_retry(
|
||||
mgr.channels["legacy_streaming"],
|
||||
outbound_message_for_event(
|
||||
channel="legacy_streaming",
|
||||
chat_id="123",
|
||||
event=StreamDeltaEvent(content="hello", stream_id="s1"),
|
||||
),
|
||||
)
|
||||
await mgr._send_with_retry(
|
||||
mgr.channels["legacy_streaming"],
|
||||
outbound_message_for_event(
|
||||
channel="legacy_streaming",
|
||||
chat_id="123",
|
||||
event=StreamEndEvent(content="", stream_id="s1", resuming=True),
|
||||
),
|
||||
)
|
||||
|
||||
assert calls == [
|
||||
("123", "hello", {"_stream_id": "s1", "_stream_delta": True}),
|
||||
("123", "", {"_stream_id": "s1", "_stream_end": True}),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_with_retry_supports_legacy_reasoning_signature():
|
||||
"""External plugins with the old reasoning hook signature should keep working."""
|
||||
deltas: list[tuple[str, str, dict]] = []
|
||||
ends: list[tuple[str, dict]] = []
|
||||
|
||||
class _LegacyReasoningChannel(BaseChannel):
|
||||
name = "legacy_reasoning"
|
||||
display_name = "Legacy Reasoning"
|
||||
|
||||
async def start(self) -> None:
|
||||
pass
|
||||
|
||||
async def stop(self) -> None:
|
||||
pass
|
||||
|
||||
async def send(self, msg: OutboundMessage) -> None:
|
||||
pass
|
||||
|
||||
async def send_reasoning_delta(
|
||||
self,
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict | None = None,
|
||||
) -> None:
|
||||
deltas.append((chat_id, delta, dict(metadata or {})))
|
||||
|
||||
async def send_reasoning_end(
|
||||
self,
|
||||
chat_id: str,
|
||||
metadata: dict | None = None,
|
||||
) -> None:
|
||||
ends.append((chat_id, dict(metadata or {})))
|
||||
|
||||
fake_config = SimpleNamespace(
|
||||
channels=ChannelsConfig(send_max_retries=3),
|
||||
providers=SimpleNamespace(groq=SimpleNamespace(api_key="")),
|
||||
)
|
||||
mgr = ChannelManager.__new__(ChannelManager)
|
||||
mgr.config = fake_config
|
||||
mgr.bus = MessageBus()
|
||||
mgr.channels = {"legacy_reasoning": _LegacyReasoningChannel(fake_config, mgr.bus)}
|
||||
mgr._dispatch_task = None
|
||||
|
||||
await mgr._send_with_retry(
|
||||
mgr.channels["legacy_reasoning"],
|
||||
outbound_message_for_event(
|
||||
channel="legacy_reasoning",
|
||||
chat_id="123",
|
||||
event=ProgressEvent(content="thinking", reasoning_delta=True, stream_id="r1"),
|
||||
),
|
||||
)
|
||||
await mgr._send_with_retry(
|
||||
mgr.channels["legacy_reasoning"],
|
||||
outbound_message_for_event(
|
||||
channel="legacy_reasoning",
|
||||
chat_id="123",
|
||||
event=ProgressEvent(reasoning_end=True, stream_id="r1"),
|
||||
),
|
||||
)
|
||||
|
||||
assert deltas == [
|
||||
("123", "thinking", {"_reasoning_delta": True, "_stream_id": "r1"}),
|
||||
]
|
||||
assert ends == [
|
||||
("123", {"_reasoning_end": True, "_stream_id": "r1"}),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_with_retry_skips_send_when_streamed():
|
||||
"""_send_with_retry should not call send when metadata has _streamed flag."""
|
||||
"""_send_with_retry should not call send for streamed response events."""
|
||||
send_called = False
|
||||
send_delta_called = False
|
||||
|
||||
@@ -778,7 +923,16 @@ async def test_send_with_retry_skips_send_when_streamed():
|
||||
nonlocal send_called
|
||||
send_called = True
|
||||
|
||||
async def send_delta(self, chat_id: str, delta: str, metadata: dict | None = None) -> None:
|
||||
async def send_delta(
|
||||
self,
|
||||
chat_id: str,
|
||||
delta: str,
|
||||
metadata: dict | None = None,
|
||||
*,
|
||||
stream_id: str | None = None,
|
||||
stream_end: bool = False,
|
||||
resuming: bool = False,
|
||||
) -> None:
|
||||
nonlocal send_delta_called
|
||||
send_delta_called = True
|
||||
|
||||
@@ -793,10 +947,11 @@ async def test_send_with_retry_skips_send_when_streamed():
|
||||
mgr.channels = {"streamed": _StreamedChannel(fake_config, mgr.bus)}
|
||||
mgr._dispatch_task = None
|
||||
|
||||
# _streamed means message was already sent via send_delta, so skip send
|
||||
msg = OutboundMessage(
|
||||
channel="streamed", chat_id="123", content="test",
|
||||
metadata={"_streamed": True}
|
||||
msg = outbound_message_for_event(
|
||||
channel="streamed",
|
||||
chat_id="123",
|
||||
event=StreamedResponseEvent(),
|
||||
content="test",
|
||||
)
|
||||
await mgr._send_with_retry(mgr.channels["streamed"], msg)
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ pytest.importorskip("discord")
|
||||
import discord
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.discord import (
|
||||
MAX_MESSAGE_LEN,
|
||||
@@ -718,9 +719,9 @@ async def test_send_delta_streams_by_editing_message(monkeypatch) -> None:
|
||||
times = iter([1.0, 3.0, 5.0])
|
||||
monkeypatch.setattr("nanobot.channels.discord.time.monotonic", lambda: next(times, 5.0))
|
||||
|
||||
await owner.send_delta("123", "hel", {"_stream_delta": True, "_stream_id": "s1"})
|
||||
await owner.send_delta("123", "lo", {"_stream_delta": True, "_stream_id": "s1"})
|
||||
await owner.send_delta("123", "", {"_stream_end": True, "_stream_id": "s1"})
|
||||
await owner.send_delta("123", "hel", stream_id="s1")
|
||||
await owner.send_delta("123", "lo", stream_id="s1")
|
||||
await owner.send_delta("123", "", stream_id="s1", stream_end=True)
|
||||
|
||||
assert target.sent_payloads[0] == {"content": "hel"}
|
||||
assert target.sent_messages[0].edits == [{"content": "hello"}, {"content": "hello"}]
|
||||
@@ -745,9 +746,9 @@ async def test_send_delta_stream_end_splits_oversized_reply(monkeypatch) -> None
|
||||
times = iter([1.0, 3.0])
|
||||
monkeypatch.setattr("nanobot.channels.discord.time.monotonic", lambda: next(times, 3.0))
|
||||
|
||||
await owner.send_delta("123", prefix, {"_stream_delta": True, "_stream_id": "s1"})
|
||||
await owner.send_delta("123", suffix, {"_stream_delta": True, "_stream_id": "s1"})
|
||||
await owner.send_delta("123", "", {"_stream_end": True, "_stream_id": "s1"})
|
||||
await owner.send_delta("123", prefix, stream_id="s1")
|
||||
await owner.send_delta("123", suffix, stream_id="s1")
|
||||
await owner.send_delta("123", "", stream_id="s1", stream_end=True)
|
||||
|
||||
assert target.sent_payloads == [{"content": prefix}, {"content": chunks[1]}]
|
||||
assert target.sent_messages[0].edits == [{"content": chunks[0]}, {"content": chunks[0]}]
|
||||
@@ -917,6 +918,31 @@ async def test_slash_model_forwards_optional_preset() -> None:
|
||||
assert handled[0]["metadata"]["is_slash_command"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_slash_trigger_forwards_required_name() -> None:
|
||||
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||
handled: list[dict] = []
|
||||
|
||||
async def capture_handle(**kwargs) -> None:
|
||||
handled.append(kwargs)
|
||||
|
||||
channel._handle_message = capture_handle # type: ignore[method-assign]
|
||||
client = DiscordBotClient(channel, intents=discord.Intents.none())
|
||||
interaction = _make_interaction()
|
||||
interaction.command.qualified_name = "trigger"
|
||||
|
||||
trigger_cmd = client.tree.get_command("trigger")
|
||||
assert trigger_cmd is not None
|
||||
await trigger_cmd.callback(interaction, name="PR review")
|
||||
|
||||
assert interaction.response.messages == [
|
||||
{"content": "Processing /trigger PR review...", "ephemeral": True}
|
||||
]
|
||||
assert len(handled) == 1
|
||||
assert handled[0]["content"] == "/trigger PR review"
|
||||
assert handled[0]["metadata"]["is_slash_command"] is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_slash_help_returns_ephemeral_help_text() -> None:
|
||||
channel = DiscordChannel(DiscordConfig(enabled=True, allow_from=["*"]), MessageBus())
|
||||
@@ -1073,7 +1099,7 @@ async def test_send_stops_typing_after_send() -> None:
|
||||
channel="discord",
|
||||
chat_id="123",
|
||||
content="progress",
|
||||
metadata={"_progress": True},
|
||||
event=ProgressEvent(content="progress"),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ from pathlib import Path
|
||||
import pytest
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.email import EmailChannel, EmailConfig
|
||||
|
||||
@@ -868,10 +869,7 @@ async def test_send_skips_progress_messages_before_smtp(monkeypatch) -> None:
|
||||
channel="email",
|
||||
chat_id="alice@example.com",
|
||||
content="",
|
||||
metadata={
|
||||
"_progress": True,
|
||||
"_tool_events": [{"phase": "end", "name": "exec"}],
|
||||
},
|
||||
event=ProgressEvent(tool_events=[{"phase": "end", "name": "exec"}]),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -193,7 +193,8 @@ class TestStreamEndReactionCleanup:
|
||||
|
||||
await ch.send_delta(
|
||||
"oc_chat1", "",
|
||||
metadata={"_stream_end": True, "message_id": "om_001"},
|
||||
metadata={"message_id": "om_001"},
|
||||
stream_end=True,
|
||||
)
|
||||
|
||||
ch._remove_reaction.assert_called_once_with("om_001", "rx_42")
|
||||
@@ -210,7 +211,7 @@ class TestStreamEndReactionCleanup:
|
||||
|
||||
await ch.send_delta(
|
||||
"oc_chat1", "",
|
||||
metadata={"_stream_end": True},
|
||||
stream_end=True,
|
||||
)
|
||||
|
||||
ch._remove_reaction.assert_not_called()
|
||||
@@ -227,7 +228,8 @@ class TestStreamEndReactionCleanup:
|
||||
|
||||
await ch.send_delta(
|
||||
"oc_chat1", "",
|
||||
metadata={"_stream_end": True, "message_id": "om_001"},
|
||||
metadata={"message_id": "om_001"},
|
||||
stream_end=True,
|
||||
)
|
||||
|
||||
ch._remove_reaction.assert_not_called()
|
||||
@@ -242,7 +244,7 @@ class TestStreamEndReactionCleanup:
|
||||
ch._client.cardkit.v1.card.settings.return_value = MagicMock(success=MagicMock(return_value=True))
|
||||
ch._remove_reaction = AsyncMock()
|
||||
|
||||
await ch.send_delta("oc_chat1", "", metadata={"_stream_end": True})
|
||||
await ch.send_delta("oc_chat1", "", stream_end=True)
|
||||
|
||||
ch._remove_reaction.assert_not_called()
|
||||
|
||||
@@ -260,7 +262,7 @@ class TestStreamEndReactionCleanup:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_removal_when_resuming(self):
|
||||
"""_resuming=True means more tool-call rounds follow; reaction must persist."""
|
||||
"""resuming=True means more tool-call rounds follow; reaction must persist."""
|
||||
ch = _make_channel()
|
||||
ch.config.done_emoji = "DONE"
|
||||
ch._stream_bufs["oc_chat1"] = _FeishuStreamBuf(
|
||||
@@ -274,7 +276,9 @@ class TestStreamEndReactionCleanup:
|
||||
|
||||
await ch.send_delta(
|
||||
"oc_chat1", "",
|
||||
metadata={"_stream_end": True, "_resuming": True, "message_id": "om_001"},
|
||||
metadata={"message_id": "om_001"},
|
||||
stream_end=True,
|
||||
resuming=True,
|
||||
)
|
||||
|
||||
ch._remove_reaction.assert_not_called()
|
||||
@@ -299,19 +303,23 @@ class TestStreamEndReactionCleanup:
|
||||
# Intermediate stream end (more tool calls coming).
|
||||
await ch.send_delta(
|
||||
"oc_chat1", "",
|
||||
metadata={"_stream_end": True, "_resuming": True, "message_id": "om_001"},
|
||||
metadata={"message_id": "om_001"},
|
||||
stream_end=True,
|
||||
resuming=True,
|
||||
)
|
||||
ch._remove_reaction.assert_not_called()
|
||||
ch._add_reaction.assert_not_called()
|
||||
|
||||
# Re-prime the stream buffer for the final round (the previous _stream_end popped it).
|
||||
# Re-prime the stream buffer for the final round (the previous stream end popped it).
|
||||
ch._stream_bufs["oc_chat1"] = _FeishuStreamBuf(
|
||||
text="t", card_id="card_1", sequence=5, last_edit=0.0,
|
||||
)
|
||||
# Final stream end (resuming=False): OnIt removed, done_emoji added.
|
||||
await ch.send_delta(
|
||||
"oc_chat1", "",
|
||||
metadata={"_stream_end": True, "_resuming": False, "message_id": "om_001"},
|
||||
metadata={"message_id": "om_001"},
|
||||
stream_end=True,
|
||||
resuming=False,
|
||||
)
|
||||
ch._remove_reaction.assert_called_once_with("om_001", "rx_42")
|
||||
ch._add_reaction.assert_called_once_with("om_001", "DONE")
|
||||
|
||||
@@ -18,6 +18,7 @@ if not FEISHU_AVAILABLE:
|
||||
pytest.skip("Feishu dependencies not installed (lark-oapi)", allow_module_level=True)
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.feishu import FeishuChannel, FeishuConfig
|
||||
|
||||
@@ -332,7 +333,8 @@ async def test_send_skips_reply_for_progress_messages() -> None:
|
||||
channel="feishu",
|
||||
chat_id="oc_abc",
|
||||
content="thinking...",
|
||||
metadata={"message_id": "om_001", "_progress": True},
|
||||
event=ProgressEvent(content="thinking..."),
|
||||
metadata={"message_id": "om_001"},
|
||||
))
|
||||
|
||||
channel._client.im.v1.message.create.assert_called_once()
|
||||
|
||||
@@ -6,6 +6,7 @@ from unittest.mock import MagicMock
|
||||
import pytest
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.feishu import FeishuChannel, FeishuConfig, _FeishuStreamBuf
|
||||
|
||||
@@ -272,7 +273,7 @@ class TestSendDelta:
|
||||
ch._client.cardkit.v1.card_element.content.return_value = _mock_content_response()
|
||||
ch._client.cardkit.v1.card.settings.return_value = _mock_content_response()
|
||||
|
||||
await ch.send_delta("oc_chat1", "", metadata={"_stream_end": True})
|
||||
await ch.send_delta("oc_chat1", "", stream_end=True)
|
||||
|
||||
assert "oc_chat1" not in ch._stream_bufs
|
||||
ch._client.cardkit.v1.card_element.content.assert_called_once()
|
||||
@@ -289,7 +290,7 @@ class TestSendDelta:
|
||||
)
|
||||
ch._client.im.v1.message.create.return_value = _mock_send_response("om_fb")
|
||||
|
||||
await ch.send_delta("oc_chat1", "", metadata={"_stream_end": True})
|
||||
await ch.send_delta("oc_chat1", "", stream_end=True)
|
||||
|
||||
assert "oc_chat1" not in ch._stream_bufs
|
||||
ch._client.cardkit.v1.card_element.content.assert_not_called()
|
||||
@@ -306,7 +307,8 @@ class TestSendDelta:
|
||||
await ch.send_delta(
|
||||
"oc_chat1",
|
||||
"",
|
||||
metadata={"_stream_end": True, "message_id": "om_001", "chat_type": "group"},
|
||||
metadata={"message_id": "om_001", "chat_type": "group"},
|
||||
stream_end=True,
|
||||
)
|
||||
|
||||
ch._client.im.v1.message.create.assert_called_once()
|
||||
@@ -326,11 +328,11 @@ class TestSendDelta:
|
||||
"oc_chat1",
|
||||
"",
|
||||
metadata={
|
||||
"_stream_end": True,
|
||||
"message_id": "om_001",
|
||||
"chat_type": "group",
|
||||
"thread_id": "ot_001",
|
||||
},
|
||||
stream_end=True,
|
||||
)
|
||||
|
||||
ch._client.im.v1.message.reply.assert_called_once()
|
||||
@@ -351,7 +353,8 @@ class TestSendDelta:
|
||||
await ch.send_delta(
|
||||
"oc_chat1",
|
||||
"",
|
||||
metadata={"_stream_end": True, "message_id": "om_001", "chat_type": "group"},
|
||||
metadata={"message_id": "om_001", "chat_type": "group"},
|
||||
stream_end=True,
|
||||
)
|
||||
|
||||
ch._client.im.v1.message.reply.assert_called_once()
|
||||
@@ -369,7 +372,7 @@ class TestSendDelta:
|
||||
ch._client.cardkit.v1.card_element.content.return_value = _mock_content_response(success=False)
|
||||
ch._client.im.v1.message.create.return_value = _mock_send_response("om_fb")
|
||||
|
||||
await ch.send_delta("oc_chat1", "", metadata={"_stream_end": True})
|
||||
await ch.send_delta("oc_chat1", "", stream_end=True)
|
||||
|
||||
assert "oc_chat1" not in ch._stream_bufs
|
||||
assert ch._client.cardkit.v1.card.settings.call_count == 2
|
||||
@@ -388,7 +391,7 @@ class TestSendDelta:
|
||||
]
|
||||
ch._client.cardkit.v1.card.settings.return_value = _mock_content_response(True)
|
||||
|
||||
await ch.send_delta("oc_chat1", "", metadata={"_stream_end": True})
|
||||
await ch.send_delta("oc_chat1", "", stream_end=True)
|
||||
|
||||
assert "oc_chat1" not in ch._stream_bufs
|
||||
assert ch._client.cardkit.v1.card_element.content.call_count == 2
|
||||
@@ -398,7 +401,7 @@ class TestSendDelta:
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_end_without_buf_is_noop(self):
|
||||
ch = _make_channel()
|
||||
await ch.send_delta("oc_chat1", "", metadata={"_stream_end": True})
|
||||
await ch.send_delta("oc_chat1", "", stream_end=True)
|
||||
ch._client.cardkit.v1.card_element.content.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -446,7 +449,7 @@ class TestToolHintInlineStreaming:
|
||||
msg = OutboundMessage(
|
||||
channel="feishu", chat_id="oc_chat1",
|
||||
content='web_fetch("https://example.com")',
|
||||
metadata={"_tool_hint": True},
|
||||
event=ProgressEvent(content='web_fetch("https://example.com")', tool_hint=True),
|
||||
)
|
||||
await ch.send(msg)
|
||||
|
||||
@@ -482,7 +485,7 @@ class TestToolHintInlineStreaming:
|
||||
msg = OutboundMessage(
|
||||
channel="feishu", chat_id="oc_chat1",
|
||||
content='read_file("path")',
|
||||
metadata={"_tool_hint": True},
|
||||
event=ProgressEvent(content='read_file("path")', tool_hint=True),
|
||||
)
|
||||
await ch.send(msg)
|
||||
|
||||
@@ -497,7 +500,8 @@ class TestToolHintInlineStreaming:
|
||||
msg = OutboundMessage(
|
||||
channel="feishu", chat_id="oc_chat1",
|
||||
content='read_file("path")',
|
||||
metadata={"_tool_hint": True, "message_id": "om_001", "chat_type": "group"},
|
||||
event=ProgressEvent(content='read_file("path")', tool_hint=True),
|
||||
metadata={"message_id": "om_001", "chat_type": "group"},
|
||||
)
|
||||
await ch.send(msg)
|
||||
|
||||
@@ -514,8 +518,8 @@ class TestToolHintInlineStreaming:
|
||||
msg = OutboundMessage(
|
||||
channel="feishu", chat_id="oc_chat1",
|
||||
content='read_file("path")',
|
||||
event=ProgressEvent(content='read_file("path")', tool_hint=True),
|
||||
metadata={
|
||||
"_tool_hint": True,
|
||||
"message_id": "om_001",
|
||||
"chat_type": "group",
|
||||
"thread_id": "ot_001",
|
||||
@@ -538,7 +542,8 @@ class TestToolHintInlineStreaming:
|
||||
msg = OutboundMessage(
|
||||
channel="feishu", chat_id="oc_chat1",
|
||||
content='read_file("path")',
|
||||
metadata={"_tool_hint": True, "message_id": "om_001", "chat_type": "group"},
|
||||
event=ProgressEvent(content='read_file("path")', tool_hint=True),
|
||||
metadata={"message_id": "om_001", "chat_type": "group"},
|
||||
)
|
||||
await ch.send(msg)
|
||||
|
||||
@@ -558,13 +563,15 @@ class TestToolHintInlineStreaming:
|
||||
|
||||
msg1 = OutboundMessage(
|
||||
channel="feishu", chat_id="oc_chat1",
|
||||
content='$ cd /project', metadata={"_tool_hint": True},
|
||||
content='$ cd /project',
|
||||
event=ProgressEvent(content='$ cd /project', tool_hint=True),
|
||||
)
|
||||
await ch.send(msg1)
|
||||
|
||||
msg2 = OutboundMessage(
|
||||
channel="feishu", chat_id="oc_chat1",
|
||||
content='$ git status', metadata={"_tool_hint": True},
|
||||
content='$ git status',
|
||||
event=ProgressEvent(content='$ git status', tool_hint=True),
|
||||
)
|
||||
await ch.send(msg2)
|
||||
|
||||
@@ -577,7 +584,7 @@ class TestToolHintInlineStreaming:
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_hint_preserved_on_final_stream_end(self):
|
||||
"""When final _stream_end closes the card, tool hint is kept in the final text."""
|
||||
"""When stream end closes the card, tool hint is kept in the final text."""
|
||||
ch = _make_channel()
|
||||
ch._stream_bufs["oc_chat1"] = _FeishuStreamBuf(
|
||||
text="Final content\n\n🔧 web_fetch(\"url\")\n\n",
|
||||
@@ -586,7 +593,7 @@ class TestToolHintInlineStreaming:
|
||||
ch._client.cardkit.v1.card_element.content.return_value = _mock_content_response()
|
||||
ch._client.cardkit.v1.card.settings.return_value = _mock_content_response()
|
||||
|
||||
await ch.send_delta("oc_chat1", "", metadata={"_stream_end": True})
|
||||
await ch.send_delta("oc_chat1", "", stream_end=True)
|
||||
|
||||
assert "oc_chat1" not in ch._stream_bufs
|
||||
update_call = ch._client.cardkit.v1.card_element.content.call_args[0][0]
|
||||
@@ -603,7 +610,8 @@ class TestToolHintInlineStreaming:
|
||||
for content in ("", " ", "\t\n"):
|
||||
msg = OutboundMessage(
|
||||
channel="feishu", chat_id="oc_chat1",
|
||||
content=content, metadata={"_tool_hint": True},
|
||||
content=content,
|
||||
event=ProgressEvent(content=content, tool_hint=True),
|
||||
)
|
||||
await ch.send(msg)
|
||||
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
"""Tests for FeishuChannel tool hint formatting."""
|
||||
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
@@ -18,6 +17,7 @@ if not FEISHU_AVAILABLE:
|
||||
pytest.skip("Feishu dependencies not installed (lark-oapi)", allow_module_level=True)
|
||||
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.channels.feishu import FeishuChannel
|
||||
|
||||
|
||||
@@ -51,7 +51,7 @@ async def test_tool_hint_sends_interactive_card(mock_feishu_channel):
|
||||
channel="feishu",
|
||||
chat_id="oc_123456",
|
||||
content='web_search("test query")',
|
||||
metadata={"_tool_hint": True}
|
||||
event=ProgressEvent(tool_hint=True),
|
||||
)
|
||||
|
||||
with patch.object(mock_feishu_channel, '_send_message_sync') as mock_send:
|
||||
@@ -72,7 +72,7 @@ async def test_tool_hint_empty_content_does_not_send(mock_feishu_channel):
|
||||
channel="feishu",
|
||||
chat_id="oc_123456",
|
||||
content=" ", # whitespace only
|
||||
metadata={"_tool_hint": True}
|
||||
event=ProgressEvent(tool_hint=True),
|
||||
)
|
||||
|
||||
with patch.object(mock_feishu_channel, '_send_message_sync') as mock_send:
|
||||
@@ -107,7 +107,7 @@ async def test_tool_hint_multiple_tools_in_one_message(mock_feishu_channel):
|
||||
channel="feishu",
|
||||
chat_id="oc_123456",
|
||||
content='web_search("query"), read_file("/path/to/file")',
|
||||
metadata={"_tool_hint": True}
|
||||
event=ProgressEvent(tool_hint=True),
|
||||
)
|
||||
|
||||
with patch.object(mock_feishu_channel, '_send_message_sync') as mock_send:
|
||||
@@ -127,7 +127,7 @@ async def test_tool_hint_new_format_basic(mock_feishu_channel):
|
||||
channel="feishu",
|
||||
chat_id="oc_123456",
|
||||
content='read src/main.py, grep "TODO"',
|
||||
metadata={"_tool_hint": True}
|
||||
event=ProgressEvent(tool_hint=True),
|
||||
)
|
||||
|
||||
with patch.object(mock_feishu_channel, '_send_message_sync') as mock_send:
|
||||
@@ -146,7 +146,7 @@ async def test_tool_hint_new_format_with_comma_in_quotes(mock_feishu_channel):
|
||||
channel="feishu",
|
||||
chat_id="oc_123456",
|
||||
content='grep "hello, world", $ echo test',
|
||||
metadata={"_tool_hint": True}
|
||||
event=ProgressEvent(tool_hint=True),
|
||||
)
|
||||
|
||||
with patch.object(mock_feishu_channel, '_send_message_sync') as mock_send:
|
||||
@@ -165,7 +165,7 @@ async def test_tool_hint_new_format_with_folding(mock_feishu_channel):
|
||||
channel="feishu",
|
||||
chat_id="oc_123456",
|
||||
content='read path × 3, grep "pattern"',
|
||||
metadata={"_tool_hint": True}
|
||||
event=ProgressEvent(tool_hint=True),
|
||||
)
|
||||
|
||||
with patch.object(mock_feishu_channel, '_send_message_sync') as mock_send:
|
||||
@@ -184,7 +184,7 @@ async def test_tool_hint_new_format_mcp(mock_feishu_channel):
|
||||
channel="feishu",
|
||||
chat_id="oc_123456",
|
||||
content='4_5v::analyze_image("photo.jpg")',
|
||||
metadata={"_tool_hint": True}
|
||||
event=ProgressEvent(tool_hint=True),
|
||||
)
|
||||
|
||||
with patch.object(mock_feishu_channel, '_send_message_sync') as mock_send:
|
||||
@@ -202,7 +202,7 @@ async def test_tool_hint_keeps_commas_inside_arguments(mock_feishu_channel):
|
||||
channel="feishu",
|
||||
chat_id="oc_123456",
|
||||
content='web_search("foo, bar"), read_file("/path/to/file")',
|
||||
metadata={"_tool_hint": True}
|
||||
event=ProgressEvent(tool_hint=True),
|
||||
)
|
||||
|
||||
with patch.object(mock_feishu_channel, '_send_message_sync') as mock_send:
|
||||
|
||||
@@ -11,6 +11,7 @@ from nio import RoomSendResponse, SyncError
|
||||
|
||||
import nanobot.channels.matrix as matrix_module
|
||||
from nanobot.bus.events import OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.matrix import (
|
||||
MATRIX_HTML_FORMAT,
|
||||
@@ -1522,7 +1523,7 @@ async def test_send_progress_keeps_typing_keepalive_running() -> None:
|
||||
channel="matrix",
|
||||
chat_id="!room:matrix.org",
|
||||
content="working...",
|
||||
metadata={"_progress": True, "_progress_kind": "reasoning"},
|
||||
event=ProgressEvent(content="working..."),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -1544,7 +1545,7 @@ async def test_send_empty_content_does_not_call_room_send() -> None:
|
||||
channel="matrix",
|
||||
chat_id="!room:matrix.org",
|
||||
content="",
|
||||
metadata={"_progress": True},
|
||||
event=ProgressEvent(),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -1563,7 +1564,7 @@ async def test_send_whitespace_only_content_does_not_call_room_send() -> None:
|
||||
channel="matrix",
|
||||
chat_id="!room:matrix.org",
|
||||
content=" \n\n ",
|
||||
metadata={"_progress": True},
|
||||
event=ProgressEvent(content=" \n\n "),
|
||||
)
|
||||
)
|
||||
|
||||
@@ -1883,7 +1884,7 @@ async def test_send_delta_stream_end_replaces_existing_message() -> None:
|
||||
last_edit=100.0,
|
||||
)
|
||||
|
||||
await channel.send_delta("!room:matrix.org", "", {"_stream_end": True})
|
||||
await channel.send_delta("!room:matrix.org", "", stream_end=True)
|
||||
|
||||
assert "!room:matrix.org" not in channel._stream_bufs
|
||||
assert client.typing_calls[-1] == ("!room:matrix.org", False, TYPING_NOTICE_TIMEOUT_MS)
|
||||
@@ -1933,7 +1934,7 @@ async def test_send_delta_threaded_edit_keeps_replace_and_thread_relation(monkey
|
||||
}
|
||||
await channel.send_delta("!room:matrix.org", "Hello", metadata)
|
||||
await channel.send_delta("!room:matrix.org", " world", metadata)
|
||||
await channel.send_delta("!room:matrix.org", "", {"_stream_end": True, **metadata})
|
||||
await channel.send_delta("!room:matrix.org", "", metadata, stream_end=True)
|
||||
|
||||
edit_content = client.room_send_calls[1]["content"]
|
||||
final_content = client.room_send_calls[2]["content"]
|
||||
@@ -1966,7 +1967,7 @@ async def test_send_delta_stream_end_noop_when_buffer_missing() -> None:
|
||||
client = _FakeAsyncClient("", "", "", None)
|
||||
channel.client = client
|
||||
|
||||
await channel.send_delta("!room:matrix.org", "", {"_stream_end": True})
|
||||
await channel.send_delta("!room:matrix.org", "", stream_end=True)
|
||||
|
||||
assert client.room_send_calls == []
|
||||
assert client.typing_calls == []
|
||||
|
||||
@@ -10,6 +10,7 @@ from unittest.mock import MagicMock
|
||||
import pytest
|
||||
|
||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||
from nanobot.bus.outbound_events import ProgressEvent
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.channels.signal import (
|
||||
SignalChannel,
|
||||
@@ -1341,7 +1342,7 @@ class TestSend:
|
||||
channel="signal",
|
||||
chat_id="+19995550001",
|
||||
content="working...",
|
||||
metadata={"_progress": True},
|
||||
event=ProgressEvent(content="working..."),
|
||||
)
|
||||
await ch.send(msg)
|
||||
# Progress messages should NOT stop the typing indicator
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user