Compare commits

..
Author SHA1 Message Date
chengyongru 9b7610709b chore(anthropic): raise SDK floor for native effort fields
Require Anthropic 0.100.0 so adaptive thinking, disabled thinking, and all advertised effort values use typed SDK parameters instead of extra_body compatibility.
2026-08-04 13:32:24 +08:00
chengyongru 356eeeb48c fix(anthropic): honor disabled thinking on Opus 5
Send an explicit disabled thinking mode for default-on Opus and Sonnet 5 models while preserving unset provider defaults and minimum-SDK compatibility.
2026-08-04 13:04:42 +08:00
chengyongru e971f6bb8f fix(anthropic): distinguish sampling restrictions
Reuse adaptive-only version thresholds where the capabilities align, while preserving Mythos Preview's supported manual thinking budgets.
2026-08-04 11:48:35 +08:00
chengyongru 5e0ef36cf2 fix(anthropic): preserve SDK and dated model compatibility 2026-08-04 10:50:26 +08:00
chengyongru 39b2294ecf fix(anthropic): support Opus 5 effort controls 2026-08-04 09:54:09 +08:00
chengyongruandchengyongru 44b7e1bf41 fix(providers): keep serde errors explicit 2026-08-03 18:06:45 +08:00
arcdrake22andchengyongru 6eda67b50c fix(providers): keep reasoning items wire-valid for DeepSeek Responses
convert_messages() emitted reasoning items with ``content`` as a plain
string whenever preserve_reasoning was enabled (the DeepSeek spec).
DeepSeek's Responses gateway rejects that shape with a serde error
("input: invalid type: string ..., expected a sequence"), which surfaced
only after token consolidation cleared provider_state and forced the
full-history conversion path; replayed server items already carry list
content, which is why normal multi-turn requests never failed. Serialize
reasoning content as a list of output_text parts, matching the OpenAI
Responses schema and DeepSeek's accepted wire shape (verified live against
api.deepseek.com/responses).

The serde fallback classifier introduced in the previous commit remains as
a last-resort safeguard for any remaining wire incompatibility.

Tests: extend test_preserves_deepseek_reasoning_content to the array shape;
add a full-history regression with the observed failing item, a
replay/consolidation regression covering both replayed and converted
reasoning items, and provider-level request fixtures for both paths.
Full suite: 5773 passed, 22 skipped (only the known local-only
channels/sms packaging failure remains).
2026-08-03 18:06:45 +08:00
arcdrake22andchengyongru fb2688fd37 fix(providers): fall back to chat completions on serde body rejections
DeepSeek's new Responses endpoint (deepseek-v4-flash) intermittently rejects valid request bodies with serde deserialization errors such as 'input: invalid type: string ..., expected a sequence'. These were not classified as compatibility errors, so affected conversations died instead of falling back to Chat Completions.

The wire format is correct (input serializes as a list), so this is a server-side Responses compatibility issue; Chat Completions is strictly more permissive, making fallback safe. Extend the fallback classifier to recognize serde body-parsing markers. Repeated failures still trip the existing circuit breaker.
2026-08-03 18:06:45 +08:00
chengyongruandchengyongru 2b63715282 fix(webui): complete i18n audit 2026-08-03 17:53:33 +08:00
Xubin Ren df11fd92a6 docs(providers): link ModelScope setup sources 2026-08-03 16:57:10 +08:00
Xubin Ren b29f9dcbcb docs(providers): align ModelScope setup with current config 2026-08-03 16:57:10 +08:00
Krislu1221andXubin Ren 02df20cd55 docs(providers): add ModelScope (魔搭) section
ModelScope is a fully implemented provider (nanobot/providers/registry.py,
image_generation.py, schema.py) with async image-generation task submission
and polling, but was previously undocumented in docs/providers.md.

This patch adds a ModelScope entry under 'Common Provider Patterns',
covering:

- Default base URL: https://api-inference.modelscope.cn/v1
- OpenAI-compatible chat/completions endpoint
- Async image-generation flow (task submit + status poll)
- Automatic 'modelscope/' prefix stripping when calling the API
- A minimal nanobot.yaml example

No code changes; docs-only.
2026-08-03 16:57:10 +08:00
chengyongruandGitHub f11710a578 fix(webui): show actual local trigger messages (#5228) 2026-08-03 16:43:01 +08:00
chengyongruandchengyongru eeecfac538 fix(webui): stabilize thread during IME input 2026-08-03 16:41:08 +08:00
Xubin Ren ac216c3e94 docs(providers): document Eden AI setup and WebUI parity 2026-08-03 16:40:13 +08:00
Xubin Ren e7ec981f79 test(providers): verify Eden AI gateway contract 2026-08-03 16:40:13 +08:00
Victor M. SMITHandXubin Ren f42a44817a feat(providers): add Eden AI as an OpenAI-compatible gateway provider
Eden AI (https://www.edenai.co) is an EU-hosted, OpenAI-compatible gateway exposing 100+ models from many providers through a single endpoint and API key. Models use the provider/model naming scheme (the full id is sent upstream, like OpenRouter).

Adds it following the registry's documented two-step recipe:
- a ProviderSpec in providers/registry.py (backend openai_compat, gateway, default_api_base https://api.edenai.run/v3, EDENAI_API_KEY, reasoning_effort)
- the matching field in ProvidersConfig (config/schema.py)

API key via EDENAI_API_KEY only; never hardcoded.

Signed-off-by: Victor M. SMITH <72023257+MVS-source@users.noreply.github.com>
2026-08-03 16:40:13 +08:00
Xubin Ren 84f98f5e92 test(cron): cover invalid schedule expressions 2026-08-03 16:20:22 +08:00
ferkans-amirandXubin Ren 73a0080484 fix(cron): validate expression syntax in _validate_schedule_for_add 2026-08-03 16:20:22 +08:00
arcdrake22andXubin Ren c6bd5f0075 test(gateway): align runtime-tasks gather tests with bounded retrieval
The helper never waits on the runtime-tasks gather after cancelling it
(its children are bounded individually), so the finished-gather test must
hand the helper an already-complete gather to exercise the bounded
retrieval path, and the cancelled-gather test must settle the gather
itself instead of expecting the helper to await a still-pending future.

Use a pre-completed child for the finished case and suppress(await) for
the cancelled case; both now assert done() and a single close.
2026-08-03 16:00:39 +08:00
Xubin Ren 39e1533c3b fix(gateway): make resource teardown cancellation-safe 2026-08-03 16:00:39 +08:00
arcdrake22andXubin Ren a91ce900ef test(gateway): add shutdown teardown regression coverage
Covers the lifecycle contract of _close_gateway_runtime: runtime tasks are
cancelled before shared resources close, pending background work is drained
before the close returns, cancellation-swallowing tasks and hanging cleanup
are bounded by their timeouts, a failing close is logged without blocking the
stop, duplicate cleanup is idempotent, and the runtime_tasks gather await path
is exercised for both completed and cancelled gathers.
2026-08-03 16:00:39 +08:00
arcdrake22andXubin Ren 8942c22d86 fix(gateway): close agent resources deterministically on shutdown
The gateway shutdown path never closed agent resources explicitly: it relied
on the agent loop task's own finally to run close_mcp() when that task is
cancelled. When the service stops with an in-flight exec session or MCP
subprocess, that path can be skipped or cut short, leaving asyncio subprocess
transports alive after the event loop closes. They are then finalized by
__del__ against a closed loop, producing "RuntimeError: Event loop is closed"
noise in the shutdown log, and in the worst case orphaned subprocesses with
the stop stalling until systemd's timeout kills the cgroup.

The teardown is now extracted into _close_gateway_runtime() with explicit
ordering and bounds:

- Runtime tasks (including the agent loop and any in-flight turn) are
  cancelled and awaited -- bounded -- before exec sessions, subagents, and MCP
  servers are closed, so no active turn is using a shared resource when it
  closes.
- Channel transports are closed before waiting for their runners to exit, since
  some SDKs swallow task cancellation while attempting to reconnect.
- agent.close_mcp() is invoked explicitly, bounded to 15s, and is idempotent:
  it is a no-op when the agent loop's own cleanup already ran, and the
  guaranteed final close otherwise.
- A coroutine that swallows cancellation (e.g. an SDK reconnect loop) can no
  longer hold the stop open until systemd's timeout kills the cgroup; cleanup
  failures are logged instead of blocking shutdown.
2026-08-03 16:00:39 +08:00
chengyongruandchengyongru a9bb39b833 fix(webui): dismiss mobile keyboard after send 2026-08-03 15:53:05 +08:00
KDBandXubin Ren 52bc79d3a0 fix(plugins): use uv when pip is unavailable 2026-08-03 15:42:26 +08:00
chengyongruandchengyongru 5c72fdcd88 fix(webui): remove unused bot identity settings 2026-08-03 14:06:44 +08:00
f7a6bc2d21 fix(webui): globally register correct MIME types for static assets (#5190)
On Windows, mimetypes.guess_type() reads the Content Type value from
HKEY_CLASSES_ROOT\.js (and other extensions) in the registry, which is
commonly set to text/plain because .js is associated with Windows Script
Host rather than web JavaScript. The registry value overrides Python's
built-in mapping and causes browsers to reject ES module scripts.

Fix by explicitly registering correct MIME types via mimetypes.add_type()
at module import time for common web static extensions (.js, .mjs, .css,
.html, .json, .svg, .wasm). Using strict=True ensures the values replace
the registry-backed standard mappings used by mimetypes.guess_type(). This
benefits all callers of mimetypes.guess_type() in the gateway process, not
just _serve_static.

Closes #5190

Co-authored-by: amkile <44280409+amkile@users.noreply.github.com>
2026-08-03 11:33:02 +08:00
arcdrake22andchengyongru 08fe9f7b3a fix(image): send Gemini Flash hints via generationConfig.imageConfig
The live v1beta API rejects the legacy responseFormat.image block
(enum-based aspectRatio/imageSize fields) for gemini-3.1-flash-lite-image
with INVALID_ARGUMENT, even for documented plain-string values. Gemini
Flash image models accept plain-string hints under
generationConfig.imageConfig instead (e.g. aspectRatio 16:9, imageSize
1K), which the API accepts. Switch the flash path to imageConfig and
update the provider tests accordingly. Other providers (aihubmix,
ollama, imagen) are untouched.
2026-08-03 11:22:09 +08:00
chengyongruandchengyongru 8fde956c64 fix(webui): show timestamps for replayed messages 2026-08-03 10:41:50 +08:00
chengyongruandGitHub 580824a15a perf(webui): accelerate JSONL session list and thread loading (#5194) 2026-08-03 09:51:35 +08:00
Xubin Ren db6c9effc3 fix(webui): position sidebar highlight on mount 2026-08-01 23:01:43 +08:00
Xubin Ren 0cb7dd5cc9 refactor(webui): reuse sidebar selection highlight 2026-08-01 23:01:43 +08:00
Xubin Ren e1894d6f0b fix(providers): respect explicit cloud namespaces 2026-08-01 20:25:58 +08:00
5eb818e800 fix(providers): require api_base before local provider wins on keyword match
Ollama's spec keeps "nemotron" as a keyword so bare `nemotron-3-nano`
auto-routes to a configured Ollama install (PR #1863). NVIDIA NIM was
later registered with the same "nemotron" keyword (commit 046d0831),
creating the only keyword collision in the registry.

In `_match_provider`, the keyword loop accepted any local provider on
`spec.is_local` alone — no api_base check. Models like
`nvidia/nemotron-3-super-120b-a12b` (intended for OpenRouter or NVIDIA
NIM) were therefore hijacked to http://localhost:11434/v1 even when the
user had never configured Ollama, causing silent connection errors at
runtime.

Add the same api_base gate the local-fallback loop already uses: a local
provider only wins by keyword when the user has actually set its
api_base. Preserves PR #1863's intent for users who configured Ollama;
fixes the silent hijack for everyone else.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-08-01 20:25:58 +08:00
santhrealandXubin Ren 4c387f6633 fix(memory): handle non-string timestamp and missing role in raw_archive 2026-08-01 20:14:28 +08:00
Xubin Ren e152e7bc0b test(cron): cover stop during manual execution 2026-08-01 20:03:19 +08:00
yu-xin-candXubin Ren e26e09c205 fix(cron): preserve manual run completion state 2026-08-01 20:03:19 +08:00
KDBandXubin Ren f3bbb543d0 refactor(cli): narrow Pyright suppressions 2026-08-01 19:52:08 +08:00
KDBandXubin Ren b1030ab131 fix(exec): preserve wait targets across response truncation 2026-08-01 19:40:36 +08:00
KDBandXubin Ren 39bb20c76b fix(session): tolerate malformed persisted session summary
AutoCompact.prepare_session runs on the turn hot path
(AgentLoop._compact_session) and read the persisted _last_summary metadata
with an unguarded meta['text'] and datetime.fromisoformat(meta['last_active']).
A _last_summary dict that was hand-edited or written by another version
(missing text/last_active, or a non-ISO last_active) raised KeyError/ValueError
out of the turn.

Sibling readers already tolerate the same data: estimate_session_prompt_tokens
uses .get('text') and _archive parses inside try/except. Mirror that tolerance:
skip when text is unusable, and fall back to the session's own updated_at (the
value the writer persists) when last_active is missing or unparseable, so the
archived summary is preserved instead of crashing the turn.
2026-08-01 19:29:16 +08:00
chengyongruandGitHub cdb75f8e7d feat(providers): support DeepSeek Responses API (#5197) 2026-08-01 11:53:51 +08:00
chengyongruandGitHub 971b977a84 fix(weixin): recover refreshed state after session expiry (#5196) 2026-08-01 00:28:21 +08:00
54650332fb fix(slack): scope channel thread openers to their own session
A top-level channel message that opens a thread fell back to the
channel-wide session, because the session key required `raw_thread_ts` —
which Slack only sets on messages that already arrived inside a thread.
Every new thread therefore began life in one shared channel session and
only became thread-scoped from its first reply onward, so unrelated
threads saw each other's opening turns.

Key off `thread_ts` instead. It is set both for messages arriving inside
a thread and for channel messages that `reply_in_thread` opens a thread
for. DM roots never get a `thread_ts`, so they keep the default per-chat
session and the DM routing from 82c5083 is preserved; with
`reply_in_thread` disabled no thread exists and the channel session is
still used.

This restores the per-thread isolation introduced in #1048.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
2026-07-31 23:47:26 +08:00
chengyongruandGitHub 172fe4f991 fix(webui): preserve user scroll ownership near tail (#5193) 2026-07-31 23:37:26 +08:00
shixi-liandchengyongru dda9b61b1e fix(config): install timezone data on all platforms 2026-07-31 19:55:22 +08:00
chengyongruandGitHub 6a1a45d07a feat: preserve Responses reasoning state and compact context (#5172) 2026-07-30 22:39:43 +08:00
Solaris-starandXubin Ren 511c764f45 fix(agent): route finish_reason='length' with blank content to length recovery
When an LLM response arrives with finish_reason='length' and has_tool_calls
but blank text content (e.g. the model spent its whole output budget on a
tool call whose closing tag was truncated), the runner dropped the tool
calls and then misrouted the blank response into the empty-response retry
branch. Retrying the same prompt cannot recover from output-budget
exhaustion, so every retry hit the same length ceiling and the turn ended
in the generic apology.

The length-recovery branch was gated on 'finish_reason == length and not
is_blank_text(clean)', so a blank-but-truncated turn could never reach it.

- The empty-response retry branch now excludes finish_reason == 'length'
  (in addition to 'error').
- The length-recovery branch no longer requires non-blank content, so a
  blank-but-truncated turn enters recovery and appends
  build_length_recovery_message (which handles a blank tail safely).

Adds a regression test asserting the length-recovery path is taken; it
fails on the unfixed code and passes with the fix.

Fixes #5133
2026-07-30 19:55:19 +08:00
Xubin Ren 0eac82984c test(mcp): stabilize idle reconnect timing 2026-07-30 19:44:09 +08:00
yu-xin-candXubin Ren 5e67fbf93e fix(exec): bound buffered session output 2026-07-30 19:44:09 +08:00
yu-xin-candXubin Ren 9ec4420104 fix(agent): release idle session locks 2026-07-30 19:17:37 +08:00
KDBandXubin Ren 52680dbe19 fix(pairing): keep approvals across transient store read failures
_load() treated any OSError like corruption and returned an empty store. When pairing.json was transiently unreadable, an unapproved DM could deny the sender, generate a pairing code from the empty view, and overwrite the store without its approved senders.

Keep the existing JSONDecodeError reset behavior, but propagate OSError so mutations cannot persist unreadable state. Read-only checks fail closed without writing; mutating /pairing subcommands report temporary unavailability; and the DM pairing path skips one reply instead of crashing the handler.

This mirrors the refuse-to-overwrite strategy used by the cron and trigger stores.
2026-07-30 19:02:37 +08:00
KDBandXubin Ren e633f867e8 fix(session): tolerate invalid idle-compaction timestamps 2026-07-30 18:52:16 +08:00
KDBandXubin Ren 07c2677eed fix(webui): drop malformed token-usage day keys
normalize_token_usage_state only length-checked persisted day keys, so a
hand-edited or foreign 10-char key (e.g. "not-a-dat3" or "2026-13-01") in
token-usage.json survived reads and atomic rewrites. token_usage_payload
then parsed every day key with an unguarded datetime.fromisoformat, so one
such key failed every /api/settings and /api/settings/usage request until
the file was repaired by hand.

Validate day keys in normalize_token_usage_state, the shared boundary that
every read, record, and rewrite already funnels through. Malformed keys are
dropped like other malformed rows and scrubbed from the file on the next
write; valid state is unchanged.
2026-07-30 18:41:44 +08:00
92361cbeac fix(gitstore): return real git object ids instead of hex-of-hex
`porcelain.commit()` and `repo.refs[...]` hand back object ids as a
40-character hex string that is already encoded to bytes. Calling `.hex()`
on that encodes the ASCII a second time, so every id GitStore produced or
displayed was double-encoded:

    auto_commit()          -> '62623234'
    git log --abbrev=8     -> 'bb244606'

The module is self-consistently wrong, so `/dream-log` and `/dream-restore`
work as long as the id came from nanobot itself. What does not work is
crossing the boundary: ids in logs and commit output match nothing in
`git log`, and an id copied from `git log` cannot be resolved:

    _resolve_sha(own id)      -> b'bb244606d780...'
    _resolve_sha(real git id) -> None

Use `.decode()` at the four sites that consume dulwich object ids. Nothing
persists an id — callers either display it or resolve it live — so there is
no stored state in the old format.

Adds two regression tests: the id returned by `auto_commit` must equal
`git log --abbrev=8`, and a real git id must resolve through `_resolve_sha`.

Co-Authored-By: Claude <noreply@anthropic.com>
2026-07-30 18:25:56 +08:00
chengyongruandchengyongru bb2f6cf324 fix(webui): preserve automation source on streamed replies 2026-07-30 17:57:31 +08:00
146 changed files with 10678 additions and 1925 deletions
+14
View File
@@ -268,6 +268,7 @@ Tracing covers the providers that go through nanobot's OpenAI-compatible client
|----------|---------|-------------|
| `custom` | Any OpenAI-compatible endpoint | — |
| `openrouter` | LLM gateway for hosted model families + Voice transcription (STT models) | [openrouter.ai](https://openrouter.ai) |
| `edenai` | LLM gateway for Eden AI's OpenAI-compatible model catalog | [app.edenai.run](https://app.edenai.run/) |
| `opencode` | LLM gateway (OpenCode Zen coding-agent models) | [opencode.ai/docs/zen](https://opencode.ai/docs/zen/) |
| `opencode_zen` | LLM gateway (legacy alias for OpenCode Zen) | [opencode.ai/docs/zen](https://opencode.ai/docs/zen/) |
| `opencode_go` | LLM gateway (OpenCode Go low-cost coding models) | [opencode.ai/docs/go](https://opencode.ai/docs/go/) |
@@ -348,6 +349,19 @@ Valid `apiType` values are exactly `auto`, `chat_completions`, and `responses`.
</details>
<a id="responses-state-and-compaction"></a>
### Responses conversation state and compaction
Providers that use the Responses API can keep reasoning context across a
conversation, which helps with multi-step tasks. Supported providers can also
compact long conversations automatically.
nanobot preserves Responses conversation state automatically for OpenAI Responses, OpenAI Codex, Azure OpenAI, DeepSeek V4 Flash, and compatible GitHub Copilot models.
Native compaction is also automatic when the provider supports it. The
threshold is derived from the active model's context window and reserved output
headroom; no provider configuration is required.
<details>
<summary><b>Azure OpenAI</b></summary>
+84 -2
View File
@@ -100,6 +100,39 @@ Gateway-style setup for model IDs served through OpenRouter.
Use the model ID exactly as OpenRouter lists it.
### Eden AI Gateway
Eden AI exposes an OpenAI-compatible chat-completions endpoint at
`https://api.edenai.run/v3`. Configure the built-in `edenai` provider and use
the full `provider/model` identifier listed by Eden AI:
```json
{
"providers": {
"edenai": {
"apiKey": "${EDENAI_API_KEY}"
}
},
"modelPresets": {
"primary": {
"provider": "edenai",
"model": "anthropic/claude-sonnet-4-5",
"maxTokens": 8192
}
},
"agents": {
"defaults": {
"modelPreset": "primary"
}
}
}
```
Nanobot sends the model ID unchanged, including its provider prefix. Use
Eden AI's [model listing](https://www.edenai.co/docs/v3/llms/listing-models)
to choose a currently available model. The WebUI can also load that catalog
after the Eden AI API key is saved under **Settings → Models**.
### OpenCode Zen and Go
OpenCode Zen and OpenCode Go are OpenCode-managed gateways for coding-agent models.
@@ -229,7 +262,9 @@ Arbitrary custom provider names are OpenAI-compatible only; they do not use the
}
```
`providers.openai.apiType` may be set when you need to force a specific OpenAI API surface. Other providers reject `apiType`; leave it unset outside `providers.openai`. Replace the model with a model ID available to your OpenAI account.
`providers.openai.apiType` may be set when you need to force a specific OpenAI API surface. Other providers reject `apiType`; leave it unset outside `providers.openai`. Replace the model with a model ID available to your OpenAI account. Direct OpenAI Responses, OpenAI Codex, Azure OpenAI Responses, and eligible GitHub Copilot models share [opaque Responses state retention](./configuration.md#responses-state-and-compaction); native compaction is enabled only where the backend supports it.
DeepSeek is the model-level exception in the OpenAI-compatible provider: `deepseek-v4-flash` automatically uses DeepSeek's native Responses API, while `deepseek-v4-pro` remains on Chat Completions.
### Custom OpenAI-Compatible Endpoint
@@ -302,6 +337,53 @@ If your custom endpoint documents a nonstandard thinking toggle, set `providers.
This named custom provider path is not for Anthropic-compatible endpoints. For Anthropic-compatible proxies, use `providers.anthropic.apiBase` and set the preset provider to `anthropic`.
### ModelScope
ModelScope (魔搭社区) exposes an OpenAI-compatible LLM endpoint plus a separate async image generation API. Both are covered by the built-in `modelscope` provider.
Create a ModelScope [access token](https://modelscope.cn/my/myaccesstoken), then choose a model whose page exposes API-Inference. The example below uses [`Qwen/Qwen3-32B`](https://modelscope.cn/models/Qwen/Qwen3-32B); hosted availability and quotas are controlled by ModelScope. See the official [API-Inference guide](https://modelscope.cn/docs/model-service/API-Inference/intro) for current service details.
```json
{
"providers": {
"modelscope": {
"apiKey": "${MODELSCOPE_API_KEY}"
}
},
"modelPresets": {
"primary": {
"provider": "modelscope",
"model": "Qwen/Qwen3-32B",
"maxTokens": 8192,
"contextWindowTokens": 65536
}
},
"agents": {
"defaults": {
"modelPreset": "primary"
}
}
}
```
Use an inference-enabled model ID exactly as ModelScope publishes it (usually `Namespace/model-name`). The default base URL is `https://api-inference.modelscope.cn/v1`; override `providers.modelscope.apiBase` only if your account routes through a different host. Chat model IDs may optionally be prefixed with `modelscope/`; nanobot strips that routing prefix before sending the request.
ModelScope image generation reuses the same provider key but is configured under `tools.imageGeneration`, not in a model preset:
```json
{
"tools": {
"imageGeneration": {
"enabled": true,
"provider": "modelscope",
"model": "Qwen/Qwen-Image-2512"
}
}
}
```
Use the image model's exact ModelScope ID without a leading `modelscope/`; the image client sends this value unchanged and handles ModelScope's async submit/poll flow. The example uses [`Qwen/Qwen-Image-2512`](https://modelscope.cn/models/Qwen/Qwen-Image-2512). See [Image Generation](./image-generation.md#modelscope) for supported sizes, aspect ratios, and the complete provider configuration.
### Ollama
Start Ollama separately, then point nanobot at the OpenAI-compatible endpoint.
@@ -458,7 +540,7 @@ For GitHub Copilot:
nanobot provider login github-copilot --set-main
```
Each command authenticates the selected provider and makes its current default model active. OAuth providers are not valid automatic fallbacks. See [`troubleshooting.md`](./troubleshooting.md#provider-and-model-problems) for proxy, headless-login, model-name, and config-key errors.
Each command authenticates the selected provider and makes its current default model active. OpenAI Codex and eligible GitHub Copilot models participate in [Responses state retention](./configuration.md#responses-state-and-compaction), while native compaction remains provider-capability-specific. OAuth providers are not valid automatic fallbacks. See [`troubleshooting.md`](./troubleshooting.md#provider-and-model-problems) for proxy, headless-login, model-name, and config-key errors.
## Provider Resolution
+28 -7
View File
@@ -31,9 +31,19 @@ class AutoCompact:
now: datetime | None = None) -> bool:
if self._ttl <= 0 or not ts:
return False
if isinstance(ts, str):
ts = datetime.fromisoformat(ts)
return ((now or datetime.now()) - ts).total_seconds() >= self._ttl * 60
try:
if isinstance(ts, str):
ts = datetime.fromisoformat(ts)
current = now or datetime.now()
if getattr(ts, "tzinfo", None) is not None or current.tzinfo is not None:
idle_seconds = current.timestamp() - ts.timestamp()
else:
idle_seconds = (current - ts).total_seconds()
except (OSError, OverflowError, TypeError, ValueError):
# list_sessions() forwards raw persisted metadata; an unusable value
# must not escape the idle scan and stop the agent loop.
return False
return idle_seconds >= self._ttl * 60
def _has_compactable_idle_tail(self, key: str) -> bool:
session = self.sessions.get_or_create(key)
@@ -124,10 +134,21 @@ class AutoCompact:
if entry:
return session, self._format_summary(entry[0], entry[1])
# Cold path: summary persisted in session metadata (process restarted).
# Persisted metadata may outlive schema changes; a malformed summary must
# not abort turn preparation.
meta = session.metadata.get("_last_summary")
if isinstance(meta, dict):
return session, self._format_summary(
cast(str, meta["text"]),
datetime.fromisoformat(cast(str, meta["last_active"])),
)
summary_meta = cast(dict[str, object], meta)
text = summary_meta.get("text")
if isinstance(text, str) and text:
raw_last_active = summary_meta.get("last_active")
try:
last_active = (
datetime.fromisoformat(raw_last_active)
if isinstance(raw_last_active, str)
else session.updated_at
)
except ValueError:
last_active = session.updated_at
return session, self._format_summary(text, last_active)
return session, None
+32 -9
View File
@@ -225,9 +225,6 @@ class ContextBuilder:
if current_role == "user"
else []
)
user_content = self.build_user_content(current_message, image_paths=media)
blocks = list(runtime_context_blocks or ()) if current_role == "user" else []
merged, runtime_context_meta = append_runtime_context(user_content, blocks)
messages: list[dict[str, Any]] = [
{
"role": "system",
@@ -243,21 +240,47 @@ class ContextBuilder:
},
*history,
]
current = self.build_current_message(
current_message,
media=media,
current_role=current_role,
runtime_context_blocks=runtime_context_blocks,
)
if messages[-1].get("role") == current_role:
last = dict(messages[-1])
last["content"] = self._merge_message_content(last.get("content"), merged)
if current_role == "user" and runtime_context_meta is not None:
last["content"] = self._merge_message_content(
last.get("content"),
current.get("content"),
)
current_meta = current.get("_meta")
if current_role == "user" and isinstance(current_meta, dict):
internal_meta = dict(last.get("_meta") or {})
internal_meta[RUNTIME_CONTEXT_MESSAGE_META] = runtime_context_meta
internal_meta.update(cast(dict[str, Any], current_meta))
last["_meta"] = internal_meta
messages[-1] = last
return messages
current: dict[str, Any] = {"role": current_role, "content": merged}
if current_role == "user" and runtime_context_meta is not None:
current["_meta"] = {RUNTIME_CONTEXT_MESSAGE_META: runtime_context_meta}
messages.append(current)
return messages
def build_current_message(
self,
current_message: str,
*,
media: list[str] | None = None,
current_role: str = "user",
runtime_context_blocks: Sequence[RuntimeContextBlock] | None = None,
) -> dict[str, Any]:
"""Build only the fresh turn message without merging it into history."""
content = self.build_user_content(current_message, image_paths=media)
blocks = list(runtime_context_blocks or ()) if current_role == "user" else []
merged, runtime_context_meta = append_runtime_context(content, blocks)
current: dict[str, Any] = {"role": current_role, "content": merged}
if current_role == "user" and runtime_context_meta is not None:
current["_meta"] = {
RUNTIME_CONTEXT_MESSAGE_META: runtime_context_meta,
}
return current
def build_user_content(
self,
text: str,
+167 -12
View File
@@ -9,6 +9,7 @@ import dataclasses
import inspect
import os
import time
import weakref
from collections.abc import Coroutine, Iterable, Mapping
from contextlib import AbstractContextManager, ExitStack, nullcontext, suppress
from dataclasses import dataclass, field
@@ -48,7 +49,7 @@ from nanobot.bus.queue import MessageBus
from nanobot.bus.runtime_events import RuntimeEventBus
from nanobot.command import CommandContext, CommandRouter, register_builtin_commands
from nanobot.config.schema import AgentDefaults, ModelPresetConfig
from nanobot.providers.base import LLMProvider
from nanobot.providers.base import LLMProvider, ProviderConversationState
from nanobot.providers.factory import ProviderSnapshot
from nanobot.runtime_context import (
RUNTIME_CONTEXT_HISTORY_META,
@@ -105,6 +106,7 @@ if TYPE_CHECKING:
from nanobot.triggers.local_store import LocalTriggerStore
_T = TypeVar("_T")
_SUBAGENT_PROVIDER_TASK_META = "subagent_provider_task_id"
class TurnKind(Enum):
@@ -125,6 +127,7 @@ class TurnContext:
history: list[dict[str, Any]] = field(default_factory=list)
initial_messages: list[dict[str, Any]] = field(default_factory=list)
provider_state: ProviderConversationState | None = field(default=None, repr=False)
request_context: RequestContext | None = None
runtime_context_blocks: list[RuntimeContextBlock] = field(default_factory=list)
attributes: dict[str, Any] = field(default_factory=dict)
@@ -242,6 +245,8 @@ class AgentLoop:
_RUNTIME_CHECKPOINT_KEY = "runtime_checkpoint"
_PENDING_USER_TURN_KEY = "pending_user_turn"
_PROVIDER_STATE_CHECKPOINT_VERSION_KEY = "provider_state_checkpoint_version"
_PROVIDER_STATE_CHECKPOINT_VERSION = "v1"
def __init__(
self,
@@ -394,7 +399,10 @@ class AgentLoop:
self._runtime_context_providers: list[RuntimeContextProvider] = []
self._active_tasks: dict[str, set[asyncio.Task[Any]]] = {}
self._background_tasks: set[asyncio.Task[Any]] = set()
self._session_locks: dict[str, asyncio.Lock] = {}
self._close_mcp_lock = asyncio.Lock()
self._session_locks: weakref.WeakValueDictionary[str, asyncio.Lock] = (
weakref.WeakValueDictionary()
)
# Per-session pending queues for mid-turn message injection.
# When a session has an active task, new messages for that session
# are routed here instead of creating a new task.
@@ -854,6 +862,7 @@ class AgentLoop:
turn_scopes: list[AbstractContextManager[Any]] | None = None,
tools: ToolRegistry | None = None,
request_context: RequestContext | None = None,
provider_state: ProviderConversationState | None = None,
) -> tuple[str | None, list[str], list[dict[str, Any]], str, bool]:
"""Run the agent iteration loop.
@@ -869,7 +878,18 @@ class AgentLoop:
async def _checkpoint(payload: dict[str, Any]) -> None:
if session is None:
return
self._set_runtime_checkpoint(session, payload)
public_payload = dict(payload)
private_state = public_payload.pop("provider_state", None)
public_payload.pop(self._PROVIDER_STATE_CHECKPOINT_VERSION_KEY, None)
if "provider_state" in payload and (
private_state is None
or isinstance(private_state, ProviderConversationState)
):
session.provider_state = private_state
public_payload[self._PROVIDER_STATE_CHECKPOINT_VERSION_KEY] = (
self._PROVIDER_STATE_CHECKPOINT_VERSION
)
self._set_runtime_checkpoint(session, public_payload)
async def _drain_pending(*, limit: int = _MAX_INJECTIONS_PER_TURN) -> list[dict[str, Any]]:
"""Drain follow-up messages from the pending queue.
@@ -1067,6 +1087,7 @@ class AgentLoop:
session_metadata=session_metadata,
message_metadata=metadata,
),
provider_state=provider_state,
))
finally:
turn_scope_stack.close()
@@ -1074,6 +1095,8 @@ class AgentLoop:
reset_request_context(request_token)
reset_file_states(file_state_token)
self._last_usage = result.usage
if session is not None and not ephemeral:
session.provider_state = result.provider_state
if result.stop_reason == "max_iterations":
logger.warning("Max iterations ({}) reached", self.max_iterations)
should_stream = turn_continuation.should_stream_budget_response(
@@ -1206,7 +1229,7 @@ class AgentLoop:
session_key = self._effective_session_key(msg)
if session_key != msg.session_key:
msg = dataclasses.replace(msg, session_key_override=session_key)
lock = self._session_locks.setdefault(session_key, asyncio.Lock())
lock = self._get_session_lock(session_key)
gate = self._concurrency_gate or nullcontext()
delivery = self.turn_delivery_factory.unrouted(msg, session_key)
@@ -1316,11 +1339,42 @@ class AgentLoop:
await self._publish_next_deferred_automation_turn(session_key)
async def close_mcp(self) -> None:
"""Drain background work, stop exec sessions, then close MCP connections."""
if self._background_tasks:
await asyncio.gather(*self._background_tasks, return_exceptions=True)
self._background_tasks.clear()
"""Stop active work, then close exec, subagent, and MCP resources.
Resource teardown must still run if cancellation interrupts task draining.
Gateway shutdown deliberately bounds this coroutine, so keeping the cleanup
phase in ``finally`` prevents a timed-out background task from leaving
subprocess transports alive after the event loop closes.
"""
# The agent loop closes itself from ``run()`` while gateway shutdown also
# performs a guaranteed final close. Serialize those owners so they cannot
# tear down the same subprocess transports concurrently.
close_lock = getattr(self, "_close_mcp_lock", None)
if close_lock is None:
close_lock = self._close_mcp_lock = asyncio.Lock()
async with close_lock:
await self._close_mcp_unlocked()
async def _close_mcp_unlocked(self) -> None:
errors: list[BaseException] = []
active_task_groups = getattr(self, "_active_tasks", {})
active_tasks = tuple({task for tasks in active_task_groups.values() for task in tasks})
active_task_groups.clear()
current_task = asyncio.current_task()
active_tasks = tuple(task for task in active_tasks if task is not current_task)
for task in active_tasks:
if not task.done():
task.cancel()
try:
if active_tasks:
await asyncio.gather(*active_tasks, return_exceptions=True)
if self._background_tasks:
await asyncio.gather(*self._background_tasks, return_exceptions=True)
except BaseException as exc:
errors.append(exc)
finally:
self._background_tasks.clear()
cleanup_steps = (
self.subagents.close,
self._exec_session_manager.close_all,
@@ -1657,14 +1711,24 @@ class AgentLoop:
"extend_to_user": is_subagent,
}
ctx.history = session.get_history(**_hist_kwargs)
stored_state = session.provider_state
subagent_followup_persisted = False
if is_subagent:
# Keep the durable internal delivery as an assistant record, but
# present this completion to the model as fresh follow-up input.
# Providers without assistant-prefill support drop trailing
# assistant messages, so using the persisted record as the current
# prompt would hide an independently dispatched subagent result.
if self._persist_subagent_followup(session, ctx.msg):
subagent_followup_persisted = self._persist_subagent_followup(
session,
ctx.msg,
)
if subagent_followup_persisted:
logger.debug("Subagent result persisted for session {}", ctx.session_key)
# Establish a durable, replay-safe baseline before any fallible
# provider compatibility or prompt assembly work. A compatible
# staged state replaces this in a second atomic save below.
session.provider_state = None
self.sessions.save(session)
ctx.input_persisted_early = True
ctx.delivery.record_runtime(runtime)
@@ -1672,13 +1736,65 @@ class AgentLoop:
ctx.request_context = self._request_context_for_turn(ctx)
if ctx.kind is TurnKind.USER:
ctx.runtime_context_blocks = await self._resolve_runtime_context_for_turn(ctx)
ctx.initial_messages = self._build_initial_messages(ctx)
staged_provider_state = False
if stored_state is not None and runtime.provider.can_resume_conversation_state(
stored_state,
runtime.model,
):
current_provider_message = self.context.build_current_message(
ctx.msg.content,
media=ctx.msg.media if ctx.kind is TurnKind.USER and ctx.msg.media else None,
runtime_context_blocks=ctx.runtime_context_blocks,
)
task_id = ctx.msg.metadata.get("subagent_task_id") if is_subagent else None
already_staged = False
if isinstance(task_id, str) and task_id:
internal_meta = current_provider_message.get("_meta")
current_provider_message["_meta"] = {
**(
cast(dict[str, Any], internal_meta)
if isinstance(internal_meta, dict)
else {}
),
_SUBAGENT_PROVIDER_TASK_META: task_id,
}
already_staged = any(
isinstance(message.get("_meta"), dict)
and cast(dict[str, Any], message["_meta"]).get(
_SUBAGENT_PROVIDER_TASK_META
)
== task_id
for message in stored_state.pending_messages
)
ctx.provider_state = (
stored_state
if already_staged
else stored_state.with_pending_messages([
*stored_state.pending_messages,
current_provider_message,
])
)
if (
not ctx.ephemeral
and (ctx.kind is TurnKind.USER or subagent_followup_persisted)
):
session.provider_state = ctx.provider_state
staged_provider_state = True
elif stored_state is not None:
session.provider_state = None
if ctx.kind is TurnKind.USER:
ctx.input_persisted_early = self._persist_user_message_early(
ctx.msg,
session,
runtime_context_blocks=ctx.runtime_context_blocks,
)
if staged_provider_state and not ctx.input_persisted_early:
session.provider_state = stored_state
elif subagent_followup_persisted and staged_provider_state:
# Upgrade the replay-safe baseline to the resumable state before
# prompt assembly and the first model checkpoint.
self.sessions.save(session)
ctx.initial_messages = self._build_initial_messages(ctx)
if ctx.on_progress is None:
ctx.on_progress = ctx.delivery.progress_callback()
@@ -1712,6 +1828,7 @@ class AgentLoop:
turn_scopes=ctx.turn_scopes,
tools=ctx.tools,
request_context=ctx.request_context,
provider_state=ctx.provider_state,
)
final_content, _, all_msgs, stop_reason, had_injections = result
ctx.final_content = final_content
@@ -2049,7 +2166,36 @@ class AgentLoop:
):
overlap = size
break
session.messages.extend(restored_messages[overlap:])
appended_messages = restored_messages[overlap:]
session.messages.extend(appended_messages)
assistant_message_data = (
cast(dict[str, Any], assistant_message)
if isinstance(assistant_message, dict)
else None
)
provider_state_is_synchronized = (
checkpoint_data.get(self._PROVIDER_STATE_CHECKPOINT_VERSION_KEY)
== self._PROVIDER_STATE_CHECKPOINT_VERSION
)
phase = checkpoint_data.get("phase")
exact_final_response = (
phase == "final_response"
and assistant_message_data is not None
and assistant_message_data.get("role") == "assistant"
and not bool(checkpoint_data.get("completed_tool_results"))
and not bool(checkpoint_data.get("pending_tool_calls"))
)
exact_completed_tools = (
phase == "tools_completed"
and assistant_message_data is not None
and assistant_message_data.get("role") == "assistant"
and not bool(checkpoint_data.get("pending_tool_calls"))
)
if not (
provider_state_is_synchronized
and (exact_final_response or exact_completed_tools)
):
session.provider_state = None
self._clear_pending_user_turn(session)
self._clear_runtime_checkpoint(session)
@@ -2070,6 +2216,7 @@ class AgentLoop:
"timestamp": datetime.now().isoformat(),
}
)
session.provider_state = None
session.updated_at = datetime.now()
self._clear_pending_user_turn(session)
@@ -2108,7 +2255,7 @@ class AgentLoop:
content=content, media=media or [], metadata=metadata,
)
# Share the dispatch lock so direct calls serialize with bus turns.
lock = self._session_locks.setdefault(session_key, asyncio.Lock())
lock = self._get_session_lock(session_key)
try:
async with lock:
kwargs: dict[str, Any] = {
@@ -2139,3 +2286,11 @@ class AgentLoop:
finally:
await self.runtime_event_publisher.run_status_changed(msg, session_key, "idle")
self.runtime_event_publisher.clear_turn(session_key)
def _get_session_lock(self, session_key: str) -> asyncio.Lock:
"""Return the shared lock while allowing idle session entries to expire."""
lock = self._session_locks.get(session_key)
if lock is None:
lock = asyncio.Lock()
self._session_locks[session_key] = lock
return lock
+7 -5
View File
@@ -713,11 +713,10 @@ class MemoryStore:
if tools_used
else ""
)
timestamp = cast(str, message.get("timestamp", "?"))
role = cast(str, message["role"])
lines.append(
f"[{timestamp[:16]}] {role.upper()}{tools}: {content}"
)
raw_timestamp = message.get("timestamp")
timestamp = str(raw_timestamp) if raw_timestamp is not None else "?"
role = str(message.get("role") or "unknown")
lines.append(f"[{timestamp[:16]}] {role.upper()}{tools}: {content}")
return "\n".join(lines)
def raw_archive(
@@ -931,6 +930,7 @@ class Consolidator:
session_key=session.key,
)
session.last_consolidated = end_idx
session.provider_state = None
self.sessions.save(session)
return summary
@@ -1136,6 +1136,7 @@ class Consolidator:
if summary:
last_summary = summary
session.last_consolidated = end_idx
session.provider_state = None
self.sessions.save(session)
if not summary:
# LLM is degraded — stop hammering it this call;
@@ -1205,6 +1206,7 @@ class Consolidator:
# Preserve history and advance only the replay boundary.
session.last_consolidated = len(session.messages) - len(visible_suffix)
session.provider_state = None
self.sessions.save(session)
logger.info(
+167 -29
View File
@@ -19,7 +19,17 @@ from nanobot.agent.context_governance import (
)
from nanobot.agent.hook import AgentHook, AgentHookContext, AgentRunHookContext
from nanobot.agent.tools.registry import ToolRegistry, is_tool_error_result
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
ToolCallRequest,
)
from nanobot.providers.conversation_state import (
ProviderConversationStateController,
allows_conversation_message_merge,
)
from nanobot.runtime_context import (
RUNTIME_CONTEXT_MESSAGE_META,
detach_runtime_context,
@@ -104,6 +114,7 @@ class AgentRunSpec:
goal_active_predicate: Callable[[], bool] | None = None
goal_continue_message: GoalContinueMessage | None = None
finalize_on_max_iterations: bool = True
provider_state: ProviderConversationState | None = None
@dataclass(slots=True)
@@ -120,6 +131,7 @@ class AgentRunResult:
had_injections: bool = False
# Terminal tail to emit when the preceding final-content prefix was already streamed.
pending_stream_content: str | None = None
provider_state: ProviderConversationState | None = field(default=None, repr=False)
class AgentRunner:
@@ -161,6 +173,7 @@ class AgentRunner:
and messages[-1].get("role") == "user"
and not is_hidden_history_message(injection)
and not is_hidden_history_message(messages[-1])
and allows_conversation_message_merge(messages[-1])
):
merged = dict(messages[-1])
left_meta = merged.get("_meta")
@@ -231,6 +244,7 @@ class AgentRunner:
assistant_message: dict[str, Any] | None,
injection_cycles: int,
*,
conversation_state: ProviderConversationStateController | None = None,
phase: str = "after error",
iteration: int | None = None,
allow_goal_continue: bool = False,
@@ -258,16 +272,21 @@ class AgentRunner:
if assistant_message is not None:
messages.append(assistant_message)
if iteration is not None:
checkpoint: dict[str, Any] = {
"phase": "final_response",
"iteration": iteration,
"model": spec.runtime.model,
"assistant_message": assistant_message,
"completed_tool_results": [],
"pending_tool_calls": [],
}
if conversation_state is not None:
checkpoint["provider_state"] = conversation_state.checkpoint(
messages
)
await self._emit_checkpoint(
spec,
{
"phase": "final_response",
"iteration": iteration,
"model": spec.runtime.model,
"assistant_message": assistant_message,
"completed_tool_results": [],
"pending_tool_calls": [],
},
checkpoint,
)
self._append_injected_messages(messages, injections)
if real_injection:
@@ -420,6 +439,12 @@ class AgentRunner:
injection_cycles = 0
compacted_tool_call_ids: set[str] = set()
pending_stream_content: str | None = None
conversation_state = ProviderConversationStateController(
provider=spec.runtime.provider,
model=spec.runtime.model,
messages=messages,
state=spec.provider_state,
)
governance_config = ContextGovernanceConfig(
provider=spec.runtime.provider,
model=spec.runtime.model,
@@ -450,7 +475,20 @@ class AgentRunner:
session_key=spec.session_key,
)
await hook.before_iteration(context)
response = await self._request_model(spec, messages_for_model, hook, context)
provider_context = conversation_state.prepare_request(
messages,
context_window_tokens=spec.runtime.context_window_tokens,
model_messages=messages_for_model,
)
response = await self._request_model(
spec,
messages_for_model,
hook,
context,
conversation_state=conversation_state,
provider_context=provider_context,
)
conversation_state.observe_response(response, messages)
context.response = response
context.tool_calls = list(response.tool_calls)
@@ -480,6 +518,10 @@ class AgentRunner:
reasoning_content=response.reasoning_content,
thinking_blocks=response.thinking_blocks,
)
assistant_message = conversation_state.project_response_message(
assistant_message,
response,
)
messages.append(assistant_message)
await self._emit_checkpoint(
spec,
@@ -544,6 +586,15 @@ class AgentRunner:
length_recovery_parts.clear()
continue
break
checkpoint_model_messages = (
self.context_governor.prepare_for_model(
governance_config,
messages,
compacted_tool_call_ids,
)
if response.provider_state is not None
else None
)
await self._emit_checkpoint(
spec,
{
@@ -553,6 +604,10 @@ class AgentRunner:
"assistant_message": assistant_message,
"completed_tool_results": completed_tool_results,
"pending_tool_calls": [],
"provider_state": conversation_state.checkpoint(
messages,
model_messages=checkpoint_model_messages,
),
},
)
empty_content_retries = 0
@@ -575,7 +630,11 @@ class AgentRunner:
)
clean = hook.finalize_content(context, response.content)
if response.finish_reason != "error" and is_blank_text(clean):
if (
response.finish_reason
not in {"error", "length", "refusal", "content_filter"}
and is_blank_text(clean)
):
empty_content_retries += 1
if empty_content_retries < _MAX_EMPTY_RETRIES:
logger.warning(
@@ -598,7 +657,12 @@ class AgentRunner:
if hook.wants_streaming():
await hook.on_stream_end(context, resuming=False)
retry_messages = self._finalization_retry_messages(messages_for_model)
response = await self._request_finalization_retry(spec, messages_for_model)
response = await self._request_finalization_retry(
spec,
messages_for_model,
transcript=messages,
conversation_state=conversation_state,
)
retry_usage = self._usage_or_estimate(spec, retry_messages, response)
self._accumulate_usage(usage, retry_usage)
raw_usage = self._merge_usage(raw_usage, retry_usage)
@@ -608,7 +672,7 @@ class AgentRunner:
original_content = response.content
clean = hook.finalize_content(context, response.content)
if response.finish_reason == "length" and not is_blank_text(clean):
if response.finish_reason == "length":
if len(length_recovery_parts) < _MAX_LENGTH_RECOVERIES:
length_recovery_parts.append(
_restore_outer_whitespace(clean or "", original_content)
@@ -623,10 +687,13 @@ class AgentRunner:
if hook.wants_streaming():
context.stream_continues_current_message = True
await hook.on_stream_end(context, resuming=True)
messages.append(build_assistant_message(
clean,
reasoning_content=response.reasoning_content,
thinking_blocks=response.thinking_blocks,
messages.append(conversation_state.project_response_message(
build_assistant_message(
clean,
reasoning_content=response.reasoning_content,
thinking_blocks=response.thinking_blocks,
),
response,
))
messages.append(build_length_recovery_message(clean or ""))
await hook.after_iteration(context)
@@ -656,15 +723,22 @@ class AgentRunner:
reasoning_content=response.reasoning_content,
thinking_blocks=response.thinking_blocks,
)
assistant_message = conversation_state.project_response_message(
assistant_message,
response,
)
# Check for mid-turn injections BEFORE signaling stream end.
# If injections are found we keep the stream alive (resuming=True)
# so streaming channels don't prematurely finalize the card.
should_continue, injection_cycles = await self._try_drain_injections(
spec, messages, assistant_message, injection_cycles,
conversation_state=conversation_state,
phase="after final response",
iteration=iteration,
allow_goal_continue=True,
allow_goal_continue=(
response.finish_reason not in {"refusal", "content_filter"}
),
)
if should_continue:
had_injections = True
@@ -717,11 +791,17 @@ class AgentRunner:
continue
break
messages.append(assistant_message or build_assistant_message(
clean,
reasoning_content=response.reasoning_content,
thinking_blocks=response.thinking_blocks,
))
messages.append(
assistant_message
or conversation_state.project_response_message(
build_assistant_message(
clean,
reasoning_content=response.reasoning_content,
thinking_blocks=response.thinking_blocks,
),
response,
)
)
await self._emit_checkpoint(
spec,
{
@@ -731,6 +811,7 @@ class AgentRunner:
"assistant_message": messages[-1],
"completed_tool_results": [],
"pending_tool_calls": [],
"provider_state": conversation_state.checkpoint(messages),
},
)
if length_recovery_parts:
@@ -764,6 +845,7 @@ class AgentRunner:
hook,
messages,
usage,
conversation_state,
)
if terminal_content is None:
terminal_content = self._max_iterations_fallback(spec)
@@ -787,6 +869,7 @@ class AgentRunner:
tool_events=tool_events,
had_injections=had_injections,
pending_stream_content=pending_stream_content,
provider_state=conversation_state.finish(messages),
)
def _build_request_kwargs(
@@ -817,6 +900,8 @@ class AgentRunner:
context: AgentHookContext,
*,
malformed_retry: bool = False,
conversation_state: ProviderConversationStateController,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
timeout_s: float | None = spec.llm_timeout_s
if timeout_s is None:
@@ -886,6 +971,7 @@ class AgentRunner:
coro = spec.runtime.provider.chat_stream_with_retry(
**kwargs,
provider_context=provider_context,
on_content_delta=_stream,
on_thinking_delta=_thinking,
on_tool_call_delta=_provider_tool_event,
@@ -920,11 +1006,15 @@ class AgentRunner:
coro = spec.runtime.provider.chat_stream_with_retry(
**kwargs,
provider_context=provider_context,
on_content_delta=_stream_progress,
on_tool_call_delta=_provider_tool_event,
)
else:
coro = spec.runtime.provider.chat_with_retry(**kwargs)
coro = spec.runtime.provider.chat_with_retry(
**kwargs,
provider_context=provider_context,
)
# Streaming requests also have provider-level idle timeouts
# (NANOBOT_STREAM_IDLE_TIMEOUT_S), but a stream that keeps producing
@@ -986,6 +1076,10 @@ class AgentRunner:
return await self._request_model(
spec, retry_messages, hook, context,
malformed_retry=True,
conversation_state=conversation_state,
provider_context=conversation_state.independent_request_context(
context_window_tokens=spec.runtime.context_window_tokens,
),
)
if (
all_dropped
@@ -998,7 +1092,13 @@ class AgentRunner:
fallback_messages = self._malformed_tool_call_retry_messages(
messages, response.content,
)
return await self._request_no_tools(spec, fallback_messages)
return await self._request_no_tools(
spec,
fallback_messages,
provider_context=conversation_state.independent_request_context(
context_window_tokens=spec.runtime.context_window_tokens,
),
)
return response
@staticmethod
@@ -1031,6 +1131,10 @@ class AgentRunner:
original_finish_reason,
)
response.tool_calls = valid
# The opaque candidate still contains every raw function_call item.
# Advancing it after dropping even one call would replay an unmatched
# call without a corresponding tool output on the next request.
response.provider_state = None
if not valid:
response.finish_reason = "stop"
return (dropped, not valid, original_finish_reason)
@@ -1060,9 +1164,27 @@ class AgentRunner:
self,
spec: AgentRunSpec,
messages: list[dict[str, Any]],
*,
transcript: list[dict[str, Any]],
conversation_state: ProviderConversationStateController,
) -> LLMResponse:
retry_messages = self._finalization_retry_messages(messages)
return await self._request_no_tools(spec, retry_messages)
provider_context = conversation_state.prepare_request(
transcript,
context_window_tokens=spec.runtime.context_window_tokens,
supplemental_messages=[retry_messages[-1]],
)
response = await self._request_no_tools(
spec,
retry_messages,
provider_context=provider_context,
)
conversation_state.observe_response(
response,
transcript,
adopt_candidate_state=False,
)
return response
@staticmethod
def _finalization_retry_messages(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
@@ -1076,10 +1198,17 @@ class AgentRunner:
hook: AgentHook,
messages: list[dict[str, Any]],
usage: dict[str, int],
conversation_state: ProviderConversationStateController,
) -> str | None:
retry_messages = self._budget_exhausted_finalization_messages(messages)
try:
response = await self._request_no_tools(spec, retry_messages)
response = await self._request_no_tools(
spec,
retry_messages,
provider_context=conversation_state.independent_request_context(
context_window_tokens=spec.runtime.context_window_tokens,
),
)
except Exception:
logger.exception(
"Budget-exhausted finalization failed for {}; using fallback",
@@ -1115,9 +1244,18 @@ class AgentRunner:
self,
spec: AgentRunSpec,
messages: list[dict[str, Any]],
*,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
kwargs = self._build_request_kwargs(spec, messages, tools=None)
return await spec.runtime.provider.chat_with_retry(**kwargs)
kwargs = self._build_request_kwargs(
spec,
messages,
tools=None,
)
return await spec.runtime.provider.chat_with_retry(
**kwargs,
provider_context=provider_context,
)
@staticmethod
def _budget_exhausted_finalization_messages(
+93 -28
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import asyncio
import time
import uuid
from collections import deque
from contextlib import suppress
from dataclasses import dataclass
from typing import Any
@@ -51,6 +52,66 @@ class ExecSessionInfo:
owner_session_key: str | None = None
class _BoundedOutputBuffer:
"""Keep the first and most recent characters within a fixed budget."""
def __init__(self, max_chars: int) -> None:
self.max_chars = max_chars
self._content = ""
self._tail: deque[str] = deque()
self._tail_chars = 0
self._total_chars = 0
self._truncated = False
@property
def has_output(self) -> bool:
return self._total_chars > 0
@property
def retained_chars(self) -> int:
return len(self._content) + self._tail_chars
def append(self, text: str) -> None:
if not text:
return
self._total_chars += len(text)
if not self._truncated:
combined = self._content + text
if len(combined) <= self.max_chars:
self._content = combined
return
head_chars = self.max_chars // 2
tail_chars = self.max_chars - head_chars
self._content = combined[:head_chars]
self._tail.append(combined[-tail_chars:])
self._tail_chars = tail_chars
self._truncated = True
return
tail_chars = self.max_chars - len(self._content)
self._tail.append(text)
self._tail_chars += len(text)
while self._tail_chars > tail_chars:
excess = self._tail_chars - tail_chars
first = self._tail[0]
if len(first) <= excess:
self._tail.popleft()
self._tail_chars -= len(first)
else:
self._tail[0] = first[excess:]
self._tail_chars -= excess
def drain(self) -> tuple[str, int]:
output = self._content + "".join(self._tail)
truncated_chars = self._total_chars - len(output)
self._content = ""
self._tail.clear()
self._tail_chars = 0
self._total_chars = 0
self._truncated = False
return output, truncated_chars
class _ExecSession:
def __init__(
self,
@@ -73,30 +134,27 @@ class _ExecSession:
# timeout None/0 means no limit; an infinite deadline is never reached.
self.deadline = time.monotonic() + timeout if timeout else float("inf")
self.last_access = time.monotonic()
self._chunks: list[str] = []
self._stdout = _BoundedOutputBuffer(MAX_OUTPUT_CHARS)
self._stderr = _BoundedOutputBuffer(MAX_OUTPUT_CHARS)
self._lock = asyncio.Lock()
self._timed_out = False
self._stdout_task = asyncio.create_task(self._read_stream(process.stdout, ""))
self._stderr_task = asyncio.create_task(self._read_stream(process.stderr, "STDERR:\n"))
self._stdout_task = asyncio.create_task(self._read_stream(process.stdout, self._stdout))
self._stderr_task = asyncio.create_task(self._read_stream(process.stderr, self._stderr))
async def _read_stream(
self,
stream: asyncio.StreamReader | None,
prefix: str,
buffer: _BoundedOutputBuffer,
) -> None:
if stream is None:
return
first = True
while True:
chunk = await stream.read(4096)
if not chunk:
break
text = chunk.decode("utf-8", errors="replace")
if prefix and first:
text = prefix + text
first = False
async with self._lock:
self._chunks.append(text)
buffer.append(text)
async def write(self, chars: str) -> str | None:
if self.process.returncode is not None:
@@ -157,10 +215,14 @@ class _ExecSession:
await self._wait_for_buffered_output()
async with self._lock:
output = "".join(self._chunks)
self._chunks.clear()
stdout, stdout_truncated = self._stdout.drain()
stderr, stderr_truncated = self._stderr.drain()
output, truncated = _truncate_output(output, max_output_chars)
output_parts = [stdout] if stdout else []
if stderr:
output_parts.append(f"STDERR:\n{stderr}")
output = "\n".join(output_parts)
output, response_truncated = _truncate_output(output, max_output_chars)
return _SessionPoll(
output=output,
done=self.process.returncode is not None,
@@ -169,7 +231,7 @@ class _ExecSession:
timed_out=self._timed_out,
terminated=terminated,
stdin_closed=stdin_closed,
truncated_chars=truncated,
truncated_chars=stdout_truncated + stderr_truncated + response_truncated,
)
async def kill(self) -> None:
@@ -195,7 +257,7 @@ class _ExecSession:
deadline = time.monotonic() + OUTPUT_DRAIN_GRACE_S
while time.monotonic() < deadline:
async with self._lock:
if self._chunks:
if self._stdout.has_output or self._stderr.has_output:
return
await asyncio.sleep(0.01)
@@ -403,20 +465,16 @@ def clamp_session_int(value: int | None, default: int, minimum: int, maximum: in
def _truncate_output(output: str, max_output_chars: int) -> tuple[str, int]:
if len(output) <= max_output_chars:
return output, 0
half = max_output_chars // 2
head_chars = max_output_chars // 2
tail_chars = max_output_chars - head_chars
omitted = len(output) - max_output_chars
return (
output[:half]
+ f"\n\n... ({omitted:,} chars truncated) ...\n\n"
+ output[-half:],
omitted,
)
return output[:head_chars] + output[-tail_chars:], omitted
def format_session_poll(session_id: str, poll: _SessionPoll) -> str:
parts = [poll.output] if poll.output else []
if poll.truncated_chars:
parts.append(f"(output truncated by {poll.truncated_chars:,} chars)")
parts.append(f"({poll.truncated_chars:,} chars truncated from output)")
if poll.timed_out:
parts.append("Error: Command timed out; session was terminated.")
if poll.terminated and not poll.timed_out:
@@ -587,7 +645,9 @@ class WriteStdinTool(Tool):
max_output_chars: int,
) -> str:
deadline = time.monotonic() + (wait_timeout_ms / 1000)
aggregate: list[str] = []
aggregate = _BoundedOutputBuffer(max_output_chars)
upstream_truncated = 0
search_overlap = ""
first = True
poll: _SessionPoll | None = None
@@ -600,19 +660,24 @@ class WriteStdinTool(Tool):
close_stdin=close_stdin if first else False,
terminate=terminate if first else False,
yield_time_ms=step_ms,
max_output_chars=max_output_chars,
max_output_chars=MAX_OUTPUT_CHARS,
owner_session_key=current_request_session_key(),
)
first = False
upstream_truncated += poll.truncated_chars
if poll.output:
aggregate.append(poll.output)
joined = "".join(aggregate)
if wait_for in joined:
poll.output = joined
searchable = search_overlap + poll.output
if wait_for in searchable:
poll.output, aggregate_truncated = aggregate.drain()
poll.truncated_chars = upstream_truncated + aggregate_truncated
result = format_session_poll(session_id, poll)
return ToolResult.error(result) if poll.timed_out else result
overlap_chars = max(0, len(wait_for) - 1)
search_overlap = searchable[-overlap_chars:] if overlap_chars else ""
if poll.done or remaining_ms <= 0:
poll.output = "".join(aggregate)
poll.output, aggregate_truncated = aggregate.drain()
poll.truncated_chars = upstream_truncated + aggregate_truncated
result = format_session_poll(session_id, poll)
if wait_for not in poll.output:
result += f"\nWait target not observed: {wait_for!r}"
+14 -13
View File
@@ -12,7 +12,7 @@ from contextlib import AsyncExitStack, suppress
from typing import TYPE_CHECKING, Any, Mapping, Protocol, cast
from weakref import WeakKeyDictionary
import httpx2 as httpx
import httpx
from loguru import logger
from nanobot.agent.tools.base import Tool, ToolResult
@@ -25,9 +25,9 @@ from nanobot.bus.events import (
)
from nanobot.bus.queue import MessageBus
from nanobot.security.network import (
Httpx2PinnedDNSAsyncTransport,
PinnedDNSAsyncTransport,
env_proxy_applies_to_url,
httpx2_env_proxy_mounts,
httpx_env_proxy_mounts,
resolve_url_target,
validate_url_target,
)
@@ -194,7 +194,7 @@ def _is_session_terminated(exc: BaseException) -> bool:
messages.append(str(getattr(error, "message", "")))
return any(
marker in message.lower()
for marker in ("session terminated", "session not found", "connection closed")
for marker in ("session terminated", "connection closed")
for message in messages
)
@@ -252,8 +252,8 @@ def _redact_url(url: str) -> str:
def _pinned_transport_kwargs() -> dict[str, Any]:
kwargs: dict[str, Any] = {"transport": Httpx2PinnedDNSAsyncTransport()}
mounts = httpx2_env_proxy_mounts()
kwargs: dict[str, Any] = {"transport": PinnedDNSAsyncTransport()}
mounts = httpx_env_proxy_mounts()
if mounts:
kwargs["mounts"] = mounts
return kwargs
@@ -518,7 +518,7 @@ def _image_block_data_url(block: Any, types: Any) -> str | None:
"""
image_cls = getattr(types, "ImageContent", None)
if image_cls is not None and isinstance(block, image_cls):
mime = getattr(block, "mime_type", None) or "image/png"
mime = getattr(block, "mimeType", None) or "image/png"
return f"data:{mime};base64,{block.data}"
embedded_cls = getattr(types, "EmbeddedResource", None)
@@ -527,7 +527,7 @@ def _image_block_data_url(block: Any, types: Any) -> str | None:
resource = getattr(block, "resource", None)
if blob_cls is not None and isinstance(resource, blob_cls):
blob_resource = cast(Any, resource)
mime = getattr(blob_resource, "mime_type", None) or ""
mime = getattr(blob_resource, "mimeType", None) or ""
if isinstance(mime, str) and mime.startswith("image/"):
return f"data:{mime};base64,{blob_resource.blob}"
return None
@@ -571,7 +571,7 @@ class MCPToolWrapper(_MCPWrapperBase):
self._original_name = tool_def.name
self._name = _sanitize_mcp_tool_name(f"mcp_{server_name}_{tool_def.name}")
self._description = tool_def.description or tool_def.name
raw_schema = tool_def.input_schema or {"type": "object", "properties": {}}
raw_schema = tool_def.inputSchema or {"type": "object", "properties": {}}
self._parameters = _normalize_schema_for_openai(raw_schema)
self._tool_timeout = tool_timeout
@@ -650,7 +650,7 @@ class MCPToolWrapper(_MCPWrapperBase):
# Success — extract text and persist any image content as artifacts.
try:
rendered = self._render_call_result(result.content, kwargs)
if getattr(result, "is_error", False):
if getattr(result, "isError", False):
return ToolResult.error(rendered)
return rendered
except Exception as exc:
@@ -876,7 +876,8 @@ class MCPPromptWrapper(_MCPWrapperBase):
return True
async def execute(self, **kwargs: Any) -> str:
from mcp import MCPError, types
from mcp import types
from mcp.shared.exceptions import McpError
retried_transient = False
refreshed_session = False
@@ -896,7 +897,7 @@ class MCPPromptWrapper(_MCPWrapperBase):
raise
logger.warning("MCP prompt '{}' was cancelled by server/SDK", self._name)
return "(MCP prompt call was cancelled)"
except MCPError as exc:
except McpError as exc:
if await self._refresh_session_after_termination(
exc,
refreshed_session,
@@ -1061,7 +1062,7 @@ async def connect_mcp_servers(
**_pinned_transport_kwargs(),
)
)
read, write = await server_stack.enter_async_context(
read, write, _ = await server_stack.enter_async_context(
streamable_http_client(cfg.url, http_client=http_client)
)
else:
+9 -1
View File
@@ -248,7 +248,15 @@ class BaseChannel(ABC):
permission_id = authorization_id if authorization_id is not None else sender_id
if not self.is_allowed(permission_id):
if is_dm:
code = generate_code(self.name, str(sender_id))
try:
code = generate_code(self.name, str(sender_id))
except OSError:
# Transient pairing-store I/O failure: skip the pairing
# reply for this message rather than crash the handler.
self.logger.warning(
"Pairing store unavailable; dropping DM from {}", sender_id
)
return
await self.send(
OutboundMessage(
channel=self.name,
+5 -6
View File
@@ -493,12 +493,11 @@ class SlackChannel(BaseChannel):
except Exception as e:
self.logger.debug("reactions_add failed: {}", e)
# Thread-scoped session key whenever the user is in a real thread
# (raw_thread_ts is set). DM threads get their own session, separate
# from the DM root, so context doesn't bleed across thread boundaries.
session_key = (
f"slack:{chat_id}:{thread_ts}" if thread_ts and raw_thread_ts else None
)
# Thread-scoped session key whenever the turn lives in a thread: either the
# message arrived inside one (raw_thread_ts) or reply_in_thread opens a new
# thread for this channel message. DM roots have no thread_ts and keep the
# default per-chat session, so context doesn't bleed across thread boundaries.
session_key = f"slack:{chat_id}:{thread_ts}" if thread_ts else None
media_paths: list[str] = []
file_markers: list[str] = []
for file_info in _as_json_list(event.get("files")) or []:
@@ -555,6 +555,113 @@ async def test_dm_thread_message_keeps_thread_ts_and_threaded_session() -> None:
assert kwargs["metadata"]["slack"]["thread_ts"] == "1700000000.000100"
def _channel_mention_request(envelope_id: str, ts: str) -> SimpleNamespace:
return SimpleNamespace(
type="events_api",
envelope_id=envelope_id,
payload={
"event": {
"type": "app_mention",
"user": "U1",
"channel": "C123",
"text": "<@UBOT> hello",
"ts": ts,
}
},
)
@pytest.mark.asyncio
async def test_channel_root_message_uses_thread_scoped_session() -> None:
"""A channel mention that opens a thread belongs to that thread's session."""
channel = SlackChannel(SlackConfig(enabled=True), MessageBus())
channel._bot_user_id = "UBOT"
channel._web_client = _FakeAsyncWebClient()
channel._handle_message = AsyncMock() # type: ignore[method-assign]
client = SimpleNamespace(send_socket_mode_response=AsyncMock())
req = _channel_mention_request("env-c1", "1700000000.000100")
await channel._on_socket_request(client, req)
channel._handle_message.assert_awaited_once()
kwargs = channel._handle_message.await_args.kwargs
assert kwargs["session_key"] == "slack:C123:1700000000.000100"
assert kwargs["metadata"]["slack"]["thread_ts"] == "1700000000.000100"
@pytest.mark.asyncio
async def test_channel_root_messages_do_not_share_one_session() -> None:
"""Two threads opened in the same channel must not collapse into one session."""
channel = SlackChannel(SlackConfig(enabled=True), MessageBus())
channel._bot_user_id = "UBOT"
channel._web_client = _FakeAsyncWebClient()
channel._handle_message = AsyncMock() # type: ignore[method-assign]
client = SimpleNamespace(send_socket_mode_response=AsyncMock())
first = _channel_mention_request("env-c1", "1700000000.000100")
second = _channel_mention_request("env-c2", "1700000000.000200")
await channel._on_socket_request(client, first)
await channel._on_socket_request(client, second)
session_keys = [call.kwargs["session_key"] for call in channel._handle_message.await_args_list]
assert session_keys == [
"slack:C123:1700000000.000100",
"slack:C123:1700000000.000200",
]
@pytest.mark.asyncio
async def test_channel_root_message_without_reply_in_thread_uses_channel_session() -> None:
"""With reply_in_thread disabled no thread is opened, so the channel session is used."""
channel = SlackChannel(SlackConfig(enabled=True, reply_in_thread=False), MessageBus())
channel._bot_user_id = "UBOT"
channel._web_client = _FakeAsyncWebClient()
channel._handle_message = AsyncMock() # type: ignore[method-assign]
client = SimpleNamespace(send_socket_mode_response=AsyncMock())
req = _channel_mention_request("env-c3", "1700000000.000300")
await channel._on_socket_request(client, req)
channel._handle_message.assert_awaited_once()
kwargs = channel._handle_message.await_args.kwargs
assert kwargs["session_key"] is None
assert kwargs["metadata"]["slack"]["thread_ts"] is None
@pytest.mark.asyncio
async def test_channel_thread_reply_keeps_thread_session() -> None:
"""A reply inside a channel thread stays in the session opened by the root message."""
channel = SlackChannel(SlackConfig(enabled=True), MessageBus())
channel._bot_user_id = "UBOT"
channel._web_client = _FakeAsyncWebClient()
channel._handle_message = AsyncMock() # type: ignore[method-assign]
channel._with_thread_context = AsyncMock(return_value="hello") # type: ignore[method-assign]
client = SimpleNamespace(send_socket_mode_response=AsyncMock())
req = SimpleNamespace(
type="events_api",
envelope_id="env-c4",
payload={
"event": {
"type": "app_mention",
"user": "U1",
"channel": "C123",
"text": "<@UBOT> follow up",
"ts": "1700000000.000400",
"thread_ts": "1700000000.000100",
}
},
)
await channel._on_socket_request(client, req)
channel._handle_message.assert_awaited_once()
kwargs = channel._handle_message.await_args.kwargs
assert kwargs["session_key"] == "slack:C123:1700000000.000100"
@pytest.mark.asyncio
async def test_slack_slash_command_skips_thread_context() -> None:
channel = SlackChannel(SlackConfig(enabled=True, allow_from=[]), MessageBus())
+1
View File
@@ -1216,6 +1216,7 @@ class WebSocketChannel(BaseChannel):
body,
metadata=meta,
phase="answer",
include_source=True,
)
raw = json.dumps(body, ensure_ascii=False)
if not conns:
@@ -51,6 +51,7 @@ from nanobot.webui.http_utils import (
)
from nanobot.webui.metadata import (
WEBSOCKET_TURN_OWNER_METADATA_KEY,
WEBUI_MESSAGE_SOURCE_METADATA_KEY,
WEBUI_SYSTEM_COMMAND_TURN_PREFIX,
WEBUI_TURN_METADATA_KEY,
)
@@ -1350,6 +1351,35 @@ async def test_send_delta_emits_delta_and_stream_end() -> None:
assert "text" not in second
@pytest.mark.asyncio
async def test_send_delta_preserves_webui_source_metadata() -> None:
bus = MagicMock()
channel = WebSocketChannel({"enabled": True, "allowFrom": ["*"], "streaming": True}, bus, gateway=_basic_handler(bus))
mock_ws = AsyncMock()
channel._attach(mock_ws, "chat-source-stream")
source = {"kind": "cron", "label": "Repo check"}
metadata = {WEBUI_MESSAGE_SOURCE_METADATA_KEY: source}
await channel.send_delta("chat-source-stream", "done", metadata=metadata, stream_id="sid")
await channel.send_delta(
"chat-source-stream",
"",
metadata=metadata,
stream_id="sid",
stream_end=True,
)
first = json.loads(mock_ws.send.call_args_list[0][0][0])
second = json.loads(mock_ws.send.call_args_list[1][0][0])
assert first["event"] == "delta"
assert first["source"] == source
assert second["event"] == "stream_end"
assert second["source"] == source
lines = read_transcript_lines("websocket:chat-source-stream")
assert lines[-2]["source"] == source
assert lines[-1]["source"] == source
@pytest.mark.asyncio
async def test_send_delta_marks_resuming_stream_end() -> None:
bus = MagicMock()
@@ -2553,6 +2583,8 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
assert body["agent"]["model_preset"] == "default"
assert body["agent"]["max_tokens"] == 8192
assert body["agent"]["timezone"] == "UTC"
assert "bot_name" not in body["agent"]
assert "bot_icon" not in body["agent"]
assert body["agent"]["tool_hint_max_length"] == 40
presets = {preset["name"]: preset for preset in body["model_presets"]}
assert presets["default"]["active"] is True
@@ -2844,8 +2876,8 @@ async def test_settings_api_returns_safe_subset_and_updates_whitelist(
assert saved.model_presets["fast-writing"].model == "openai/gpt-5.5"
assert saved.model_presets["fast-writing"].provider == "openai"
assert saved.agents.defaults.timezone == "Asia/Shanghai"
assert saved.agents.defaults.bot_name == "Nano"
assert saved.agents.defaults.bot_icon == "N"
assert saved.agents.defaults.bot_name == "nanobot"
assert saved.agents.defaults.bot_icon == "🐈"
assert saved.agents.defaults.tool_hint_max_length == 120
assert saved.providers.openrouter.api_key == "sk-or-next"
assert saved.providers.openrouter.api_base == "https://openrouter.ai/api/v1"
@@ -24,6 +24,7 @@ from nanobot.runtime_context import (
RuntimeContextBlock,
append_runtime_context,
)
from nanobot.security.workspace_access import WORKSPACE_SCOPE_METADATA_KEY
from nanobot.session.keys import UNIFIED_SESSION_KEY
from nanobot.session.manager import Session, SessionManager
from nanobot.triggers.local_store import LocalTriggerStore
@@ -427,6 +428,7 @@ async def test_session_automations_route_lists_local_triggers(
chat_id="abc",
session_key="websocket:abc",
)
trigger_store.enqueue(trigger.id, "Review PR #4591")
channel = _ch(
bus,
session_manager=_seed_session(tmp_path, key="websocket:abc"),
@@ -453,6 +455,7 @@ async def test_session_automations_route_lists_local_triggers(
assert job["kind"] == "local_trigger"
assert job["schedule"]["kind"] == "local"
assert job["payload"]["kind"] == "local_trigger"
assert job["payload"]["message"] == "Review PR #4591"
assert job["payload"]["command"] == f'nanobot trigger {trigger.id} "message"'
assert job["state"]["pending"] is True
finally:
@@ -2201,7 +2204,7 @@ async def test_mcp_presets_routes_require_token_and_return_payload(
@pytest.mark.asyncio
async def test_sessions_list_only_returns_websocket_sessions_by_default(
bus: MagicMock, tmp_path: Path
bus: MagicMock, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
# Seed a realistic multi-channel disk state: CLI, Slack, Lark and
# websocket sessions all live in the same ``sessions/`` directory.
@@ -2215,7 +2218,20 @@ async def test_sessions_list_only_returns_websocket_sessions_by_default(
"websocket:beta",
],
)
channel = _ch(bus, session_manager=sm, port=29906)
project = tmp_path / "project"
project.mkdir()
scoped = sm.get_or_create("websocket:beta")
scoped.metadata[WORKSPACE_SCOPE_METADATA_KEY] = {
"project_path": str(project),
"access_mode": "restricted",
}
sm.save(scoped)
def fail_metadata_read(_key: str) -> None:
raise AssertionError("the session list must use its own index metadata")
monkeypatch.setattr(sm, "read_session_metadata", fail_metadata_read)
channel = _ch(bus, session_manager=sm, workspace_path=tmp_path, port=29906)
server_task = asyncio.create_task(channel.start())
try:
token = channel.gateway.tokens.issue_api_token(300)
@@ -2225,10 +2241,17 @@ async def test_sessions_list_only_returns_websocket_sessions_by_default(
"http://127.0.0.1:29906/api/sessions", headers=auth
)
assert listing.status_code == 200
keys = {s["key"] for s in listing.json()["sessions"]}
sessions = listing.json()["sessions"]
keys = {s["key"] for s in sessions}
# Only websocket-channel sessions are part of the webui surface; CLI /
# Slack / Lark rows would be non-resumable from the browser.
assert keys == {"websocket:alpha", "websocket:beta"}
rows = {row["key"]: row for row in sessions}
assert rows["websocket:beta"]["workspace_scope"]["project_path"] == str(
project.resolve()
)
assert rows["websocket:beta"]["workspace_scope"]["access_mode"] == "restricted"
assert all(not any(key.startswith("_") for key in row) for row in sessions)
finally:
await channel.stop()
await server_task
@@ -2594,6 +2617,7 @@ async def test_webui_automations_route_manages_local_triggers(
by_id = {job["id"]: job for job in listed.json()["jobs"]}
assert by_id[trigger.id]["kind"] == "local_trigger"
assert by_id[trigger.id]["state"]["pending"] is True
assert by_id[trigger.id]["payload"]["message"] == "Review queued PR"
assert by_id[trigger.id]["trigger"]["command"] == f'nanobot trigger {trigger.id} "message"'
disabled = await _http_get(
@@ -2956,6 +2980,139 @@ async def test_webui_thread_resigns_assistant_media_urls(
await server_task
@pytest.mark.asyncio
async def test_sessions_list_negotiates_gzip_across_repeated_headers(
bus: MagicMock, tmp_path: Path
) -> None:
sm = _seed_many(tmp_path, [f"websocket:gzip-{index:03d}" for index in range(80)])
port = _free_port()
channel = _ch(bus, session_manager=sm, workspace_path=tmp_path, port=port)
server_task = asyncio.create_task(channel.start())
try:
token = channel.gateway.tokens.issue_api_token(300)
response = await _http_get(
f"http://127.0.0.1:{port}/api/sessions",
headers=[
("Authorization", f"Bearer {token}"),
("Accept-Encoding", "identity;q=0"),
("Accept-Encoding", "gzip"),
],
)
assert response.status_code == 200
assert response.headers["Content-Encoding"] == "gzip"
assert response.headers["Vary"] == "Accept-Encoding"
assert len(response.json()["sessions"]) == 80
finally:
await channel.stop()
await server_task
@pytest.mark.asyncio
async def test_webui_thread_complete_transcript_skips_session_history_read(
bus: MagicMock, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
from nanobot.webui.transcript import append_transcript_object
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
key = "websocket:fast-thread"
sm = _seed_session(tmp_path, key=key)
for event in (
{"event": "user", "chat_id": "fast-thread", "text": "hi"},
{"event": "message", "chat_id": "fast-thread", "text": "hello back"},
{"event": "turn_end", "chat_id": "fast-thread"},
):
append_transcript_object(key, event)
read_session_file = MagicMock(
side_effect=AssertionError("complete transcripts must not read canonical history")
)
monkeypatch.setattr(sm, "read_session_file", read_session_file)
port = _free_port()
channel = _ch(
bus,
session_manager=sm,
workspace_path=tmp_path,
port=port,
)
server_task = asyncio.create_task(channel.start())
try:
token = channel.gateway.tokens.issue_api_token(300)
response = await _http_get(
f"http://127.0.0.1:{port}/api/sessions/"
"websocket%3Afast-thread/webui-thread?limit=160&direction=latest",
headers={"Authorization": f"Bearer {token}"},
)
assert response.status_code == 200
assert [message["content"] for message in response.json()["messages"]] == [
"hi",
"hello back",
]
read_session_file.assert_not_called()
finally:
await channel.stop()
await server_task
@pytest.mark.asyncio
async def test_webui_thread_negotiates_gzip_for_large_payloads(
bus: MagicMock, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
from nanobot.webui.transcript import append_transcript_object
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
sm = SessionManager(tmp_path)
append_transcript_object(
"websocket:gzip-thread",
{
"event": "user",
"chat_id": "gzip-thread",
"text": "compress me " * 1_000,
},
)
port = _free_port()
channel = _ch(bus, session_manager=sm, workspace_path=tmp_path, port=port)
server_task = asyncio.create_task(channel.start())
try:
token = channel.gateway.tokens.issue_api_token(300)
url = (
f"http://127.0.0.1:{port}/api/sessions/"
"websocket%3Agzip-thread/webui-thread?limit=80&direction=latest"
)
compressed = await _http_get(
url,
headers={
"Authorization": f"Bearer {token}",
"Accept-Encoding": "br, gzip",
},
)
assert compressed.status_code == 200
assert compressed.headers["Content-Encoding"] == "gzip"
assert compressed.headers["Vary"] == "Accept-Encoding"
assert int(compressed.headers["Content-Length"]) < len(compressed.content)
assert compressed.json()["messages"][0]["content"].startswith("compress me")
identity = await _http_get(
url,
headers={
"Authorization": f"Bearer {token}",
"Accept-Encoding": "gzip;q=0, br",
},
)
assert identity.status_code == 200
assert "Content-Encoding" not in identity.headers
assert identity.json() == compressed.json()
unauthorized = await _http_get(url, headers={"Accept-Encoding": "gzip"})
assert unauthorized.status_code == 401
assert "Content-Encoding" not in unauthorized.headers
finally:
await channel.stop()
await server_task
@pytest.mark.asyncio
async def test_session_routes_reject_non_websocket_keys(
bus: MagicMock, tmp_path: Path
@@ -248,7 +248,7 @@ class WsTestClient:
async def http_get(
url: str,
headers: dict[str, str] | None = None,
headers: dict[str, str] | list[tuple[str, str]] | None = None,
) -> httpx.Response:
"""GET a local test server without loading an unused TLS trust store."""
request = httpx.Request("GET", url, headers=headers or {})
+25 -2
View File
@@ -230,9 +230,30 @@ class WeixinChannel(BaseChannel):
self.logger.error("Failed to load Weixin account state", exc_info=True)
return False
def _save_state(self) -> None:
def _save_state(self, *, force: bool = False) -> None:
state_file = self._get_state_dir() / "account.json"
with suppress(Exception):
if not force and state_file.exists():
persisted: object = None
try:
persisted = json.loads(state_file.read_text())
except Exception:
persisted = None
persisted_token = ""
if isinstance(persisted, dict):
persisted_mapping = cast(dict[str, object], persisted)
persisted_token = str(persisted_mapping.get("token", "") or "")
configured_token_is_authoritative: bool = bool(self.config.token) and (
self._token == self.config.token
)
if (
persisted_token
and persisted_token != self._token
and not configured_token_is_authoritative
):
# A concurrent QR login may have committed a newer token.
# Never let an older runtime snapshot overwrite it.
return
data = {
"token": self._token,
"get_updates_buf": self._get_updates_buf,
@@ -489,7 +510,7 @@ class WeixinChannel(BaseChannel):
self._token = token
if base_url:
self.config.base_url = base_url
self._save_state()
self._save_state(force=True)
async def connect_close_client(self) -> None:
self._running = False
@@ -613,6 +634,8 @@ class WeixinChannel(BaseChannel):
remaining = self._session_pause_remaining_s()
if remaining > 0:
await asyncio.sleep(remaining)
if not self.config.token:
self._load_state()
return
body: dict[str, Any] = {
@@ -98,6 +98,80 @@ def test_save_and_load_state_persists_context_tokens(tmp_path) -> None:
assert restored._context_tokens == {"wx-user": "ctx-1"}
def test_save_state_preserves_token_committed_by_another_instance(tmp_path) -> None:
channel = WeixinChannel(
WeixinConfig(enabled=True, allow_from=["*"], state_dir=str(tmp_path)),
MessageBus(),
)
channel._token = "old-token"
channel._save_state()
replacement = {
"token": "new-token",
"base_url": "https://new.example",
"get_updates_buf": "",
"context_tokens": {},
"typing_tickets": {},
}
(tmp_path / "account.json").write_text(json.dumps(replacement), encoding="utf-8")
channel._get_updates_buf = "stale-cursor"
channel._save_state()
assert json.loads((tmp_path / "account.json").read_text()) == replacement
def test_save_state_force_overwrites_replaced_token(tmp_path) -> None:
channel = WeixinChannel(
WeixinConfig(enabled=True, allow_from=["*"], state_dir=str(tmp_path)),
MessageBus(),
)
(tmp_path / "account.json").write_text(json.dumps({"token": "old-token"}), encoding="utf-8")
channel.connect_commit_account(token="new-token", base_url="https://new.example")
saved = json.loads((tmp_path / "account.json").read_text())
assert saved["token"] == "new-token"
assert saved["base_url"] == "https://new.example"
def test_save_state_persists_explicit_config_token_over_stale_state(tmp_path) -> None:
channel = WeixinChannel(
WeixinConfig(
enabled=True,
allow_from=["*"],
token="configured-token",
state_dir=str(tmp_path),
),
MessageBus(),
)
channel._token = "configured-token"
channel._get_updates_buf = "current-cursor"
(tmp_path / "account.json").write_text(
json.dumps({"token": "stale-token", "get_updates_buf": "stale-cursor"}),
encoding="utf-8",
)
channel._save_state()
saved = json.loads((tmp_path / "account.json").read_text())
assert saved["token"] == "configured-token"
assert saved["get_updates_buf"] == "current-cursor"
def test_save_state_with_empty_runtime_token_preserves_persisted_account(tmp_path) -> None:
channel = WeixinChannel(
WeixinConfig(enabled=True, allow_from=["*"], state_dir=str(tmp_path)),
MessageBus(),
)
persisted = {"token": "persisted-token", "get_updates_buf": "persisted-cursor"}
(tmp_path / "account.json").write_text(json.dumps(persisted), encoding="utf-8")
channel._save_state()
assert json.loads((tmp_path / "account.json").read_text()) == persisted
@pytest.mark.asyncio
async def test_process_message_deduplicates_inbound_ids() -> None:
channel, bus = _make_channel()
@@ -462,6 +536,56 @@ async def test_poll_once_pauses_session_on_expired_errcode() -> None:
assert channel._session_pause_remaining_s() > 0
@pytest.mark.asyncio
async def test_poll_once_reloads_refreshed_state_after_session_pause(
tmp_path, monkeypatch: pytest.MonkeyPatch
) -> None:
channel = WeixinChannel(
WeixinConfig(enabled=True, allow_from=["*"], state_dir=str(tmp_path)),
MessageBus(),
)
channel._token = "old-token"
channel._save_state()
(tmp_path / "account.json").write_text(
json.dumps({"token": "new-token", "base_url": "https://new.example"}),
encoding="utf-8",
)
channel._session_pause_until = time.time() + 10
monkeypatch.setattr(weixin_mod.asyncio, "sleep", AsyncMock())
await channel._poll_once()
assert channel._token == "new-token"
assert channel.config.base_url == "https://new.example"
@pytest.mark.asyncio
async def test_poll_once_keeps_explicit_token_after_session_pause(
tmp_path, monkeypatch: pytest.MonkeyPatch
) -> None:
channel = WeixinChannel(
WeixinConfig(
enabled=True,
allow_from=["*"],
token="configured-token",
state_dir=str(tmp_path),
),
MessageBus(),
)
channel._token = "configured-token"
(tmp_path / "account.json").write_text(
json.dumps({"token": "stale-token", "base_url": "https://stale.example"}),
encoding="utf-8",
)
channel._session_pause_until = time.time() + 10
monkeypatch.setattr(weixin_mod.asyncio, "sleep", AsyncMock())
await channel._poll_once()
assert channel._token == "configured-token"
assert channel.config.base_url == "https://ilinkai.weixin.qq.com"
@pytest.mark.asyncio
async def test_qr_login_refreshes_expired_qr_and_then_succeeds(
no_qr_poll_delay,
+8 -9
View File
@@ -1,7 +1,5 @@
"""Typer commands for foreground and background gateway control."""
# pyright: reportUnusedFunction=false
from __future__ import annotations
import subprocess
@@ -135,8 +133,9 @@ def create_gateway_app(
console.print()
console.print(result.content)
# Typer consumes these callbacks through decorator registration.
@gateway_app.callback(invoke_without_command=True)
def gateway(
def gateway( # pyright: ignore[reportUnusedFunction]
ctx: typer.Context,
port: int | None = typer.Option(None, "--port", "-p", help="Gateway port"),
workspace: str | None = typer.Option(None, "--workspace", "-w", help="Workspace directory"),
@@ -191,7 +190,7 @@ def create_gateway_app(
)
@gateway_app.command("status")
def gateway_status(
def gateway_status( # pyright: ignore[reportUnusedFunction]
workspace: str | None = typer.Option(None, "--workspace", "-w", help="Workspace directory"),
config: str | None = typer.Option(None, "--config", "-c", help="Path to config file"),
) -> None:
@@ -199,7 +198,7 @@ def create_gateway_app(
print_status(runtime_for_instance(workspace=workspace, config=config).status())
@gateway_app.command("logs")
def gateway_logs(
def gateway_logs( # pyright: ignore[reportUnusedFunction]
tail: int = typer.Option(200, "--tail", help="Number of recent lines to show"),
follow: bool = typer.Option(True, "--follow/--no-follow", help="Follow new log output"),
workspace: str | None = typer.Option(None, "--workspace", "-w", help="Workspace directory"),
@@ -217,7 +216,7 @@ def create_gateway_app(
console.print(line)
@gateway_app.command("stop")
def gateway_stop(
def gateway_stop( # pyright: ignore[reportUnusedFunction]
timeout: int = typer.Option(20, "--timeout", help="Stop timeout in seconds"),
workspace: str | None = typer.Option(None, "--workspace", "-w", help="Workspace directory"),
config: str | None = typer.Option(None, "--config", "-c", help="Path to config file"),
@@ -233,7 +232,7 @@ def create_gateway_app(
raise typer.Exit(1)
@gateway_app.command("restart")
def gateway_restart(
def gateway_restart( # pyright: ignore[reportUnusedFunction]
port: int | None = typer.Option(None, "--port", "-p", help="Gateway port"),
workspace: str | None = typer.Option(None, "--workspace", "-w", help="Workspace directory"),
verbose: bool = typer.Option(False, "--verbose", "-v", help="Verbose output"),
@@ -266,7 +265,7 @@ def create_gateway_app(
raise typer.Exit(1)
@gateway_app.command("install-service")
def gateway_install_service(
def gateway_install_service( # pyright: ignore[reportUnusedFunction]
port: int | None = typer.Option(None, "--port", "-p", help="Gateway port"),
workspace: str | None = typer.Option(None, "--workspace", "-w", help="Workspace directory"),
verbose: bool = typer.Option(False, "--verbose", "-v", help="Verbose output"),
@@ -302,7 +301,7 @@ def create_gateway_app(
raise typer.Exit(1)
@gateway_app.command("uninstall-service")
def gateway_uninstall_service(
def gateway_uninstall_service( # pyright: ignore[reportUnusedFunction]
name: str = typer.Option("nanobot-gateway", "--name", help="Service name"),
manager: ServiceManagerKind = typer.Option("auto", "--manager", help="auto, systemd, or launchd"),
dry_run: bool = typer.Option(False, "--dry-run", help="Print actions without uninstalling"),
+55 -13
View File
@@ -201,6 +201,58 @@ def _print_gateway_health_endpoint(host: str, port: int) -> None:
)
async def _close_gateway_runtime(
agent: AgentLoop,
channels: Any,
tasks: list[asyncio.Task[Any]],
runtime_tasks: asyncio.Future[list[Any]] | None,
*,
task_wait_timeout: float = 15.0,
close_timeout: float = 15.0,
) -> None:
"""Cancel runtime tasks, then deterministically close agent resources.
Order matters: runtime tasks (including the agent loop and any in-flight
turn) are cancelled and awaited -- bounded -- before exec sessions,
subagents, and MCP servers are torn down, so no active turn is using a
shared resource when it closes. The final close is bounded and idempotent:
the agent loop's own finally also calls ``close_mcp()``, so this runs again
as a no-op when that path already completed, and as the guaranteed final
close when it was skipped or cut short (which previously left asyncio
subprocess transports alive past ``loop.close()``, producing
"RuntimeError: Event loop is closed" noise and potentially orphaned
processes at interpreter exit).
"""
# Some SDKs swallow task cancellation while attempting to reconnect.
# Close channel transports before waiting for their runners to exit.
await channels.stop_all()
for task in tasks:
if not task.done():
task.cancel()
pending: set[asyncio.Task[Any]] = set()
if tasks:
# Bounded: a coroutine that swallows cancellation (e.g. an SDK reconnect
# loop) must not hold the stop open until systemd's timeout kills the
# cgroup. Anything still pending is abandoned and closed underneath.
_done, pending = await asyncio.wait(tasks, timeout=task_wait_timeout)
# A task can swallow the first cancellation while unwinding. Re-cancel
# timed-out tasks so an agent loop stuck draining background work reaches
# its resource-cleanup phase before the explicit final close below.
for task in pending:
task.cancel()
if runtime_tasks is not None and not runtime_tasks.done():
runtime_tasks.cancel()
try:
await asyncio.wait_for(agent.close_mcp(), timeout=close_timeout)
except BaseException as exc: # noqa: BLE001 - shutdown must proceed
logger.warning("Gateway shutdown: agent resource cleanup incomplete: {}", exc)
# Retrieving an already-finished gather prevents noisy unhandled exceptions,
# but never wait for it here: its children were bounded individually above.
if runtime_tasks is not None and runtime_tasks.done():
with suppress(asyncio.CancelledError, Exception):
await runtime_tasks
def _run_gateway(
config: Config,
*,
@@ -734,7 +786,6 @@ def _run_gateway(
tasks: list[asyncio.Task[Any]] = []
shutdown_task: asyncio.Task[Any] | None = None
runtime_tasks: asyncio.Future[list[Any]] | None = None
runtime_tasks_drained = False
shutdown_event = asyncio.Event()
cli_terminal._ensure_interactive_tty_mode()
restore_shutdown_handlers = _install_gateway_shutdown_handlers(
@@ -786,7 +837,6 @@ def _run_gateway(
return_when=asyncio.FIRST_COMPLETED,
)
if runtime_tasks in done:
runtime_tasks_drained = True
await runtime_tasks
else:
runtime_tasks.cancel()
@@ -805,17 +855,9 @@ def _run_gateway(
await shutdown_task
cron.stop()
agent.stop()
# Some SDKs swallow task cancellation while attempting to reconnect.
# Close channel transports before waiting for their runners to exit.
await channels.stop_all()
for task in tasks:
if not task.done():
task.cancel()
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
if runtime_tasks is not None and not runtime_tasks_drained:
with suppress(asyncio.CancelledError, Exception):
await runtime_tasks
# Cancel runtime tasks first, then deterministically close
# exec/MCP resources while the event loop is still alive.
await _close_gateway_runtime(agent, channels, tasks, runtime_tasks)
# Flush all cached sessions to durable storage before exit.
# This prevents data loss on filesystems with write-back
# caching (rclone VFS, NFS, FUSE mounts, etc.).
+16 -11
View File
@@ -1,7 +1,5 @@
"""Interactive onboarding questionnaire for nanobot."""
# pyright: reportMissingTypeStubs=false, reportUnusedFunction=false
import asyncio
import json
import types
@@ -206,35 +204,36 @@ def _select_with_back(
# Key bindings
bindings = KeyBindings()
# KeyBindings consumes these handlers through decorator registration.
@bindings.add(Keys.Up)
def _up(event: KeyPressEvent) -> None:
def _up(event: KeyPressEvent) -> None: # pyright: ignore[reportUnusedFunction]
nonlocal selected_index
selected_index = (selected_index - 1) % len(choices)
event.app.invalidate()
@bindings.add(Keys.Down)
def _down(event: KeyPressEvent) -> None:
def _down(event: KeyPressEvent) -> None: # pyright: ignore[reportUnusedFunction]
nonlocal selected_index
selected_index = (selected_index + 1) % len(choices)
event.app.invalidate()
@bindings.add(Keys.Enter)
def _enter(event: KeyPressEvent) -> None:
def _enter(event: KeyPressEvent) -> None: # pyright: ignore[reportUnusedFunction]
state["result"] = choices[selected_index]
event.app.exit()
@bindings.add("escape")
def _escape(event: KeyPressEvent) -> None:
def _escape(event: KeyPressEvent) -> None: # pyright: ignore[reportUnusedFunction]
state["result"] = _BACK_PRESSED
event.app.exit()
@bindings.add(Keys.Left)
def _left(event: KeyPressEvent) -> None:
def _left(event: KeyPressEvent) -> None: # pyright: ignore[reportUnusedFunction]
state["result"] = _BACK_PRESSED
event.app.exit()
@bindings.add(Keys.ControlC)
def _ctrl_c(event: KeyPressEvent) -> None:
def _ctrl_c(event: KeyPressEvent) -> None: # pyright: ignore[reportUnusedFunction]
state["result"] = None
event.app.exit()
@@ -532,8 +531,9 @@ def _input_back_key_bindings() -> KeyBindings:
"""Return key bindings that make Escape behave like a local back action."""
bindings = KeyBindings()
# KeyBindings consumes this handler through decorator registration.
@bindings.add("escape")
def _escape(event: KeyPressEvent) -> None:
def _escape(event: KeyPressEvent) -> None: # pyright: ignore[reportUnusedFunction]
event.app.exit(result=_BACK_PRESSED)
return bindings
@@ -1668,7 +1668,11 @@ def _quick_start_oauth_login(config: Config, provider_name: str) -> bool:
return False
try:
from oauth_cli_kit import get_token, login_oauth_interactive
# oauth-cli-kit does not publish type information.
from oauth_cli_kit import ( # pyright: ignore[reportMissingTypeStubs]
get_token,
login_oauth_interactive,
)
except ImportError:
console.print("[red]oauth_cli_kit not installed. Run: pip install oauth-cli-kit[/red]")
return False
@@ -1709,7 +1713,8 @@ def _quick_start_oauth_is_authenticated(config: Config, provider_name: str) -> b
if provider_name != "openai_codex":
return False
try:
from oauth_cli_kit import get_token
# oauth-cli-kit does not publish type information.
from oauth_cli_kit import get_token # pyright: ignore[reportMissingTypeStubs]
proxy = _quick_start_codex_proxy(config)
token = get_token(proxy=proxy)
+30 -11
View File
@@ -139,7 +139,7 @@ class AgentDefaults(Base):
validation_alias=AliasChoices("toolHintMaxLength"),
serialization_alias="toolHintMaxLength",
) # Max characters for tool hint display (e.g. "$ cd …/project && npm test")
reasoning_effort: str | None = None # low / medium / high / adaptive / none — LLM thinking effort; None preserves the provider default
reasoning_effort: str | None = None # low / medium / high / xhigh / max / adaptive / none — LLM thinking effort; None preserves the provider default
timezone: str = "UTC" # IANA timezone, e.g. "Asia/Shanghai", "America/New_York"
bot_name: str = "nanobot" # Display name shown in CLI prompts (e.g. "{name} is thinking...")
bot_icon: str = "🐈" # Short icon (emoji or text) shown next to the bot name in CLI; "" to omit
@@ -269,6 +269,7 @@ class ProvidersConfig(Base):
ant_ling: ProviderConfig = Field(default_factory=ProviderConfig) # Ant Ling
aihubmix: ProviderConfig = Field(default_factory=ProviderConfig) # AiHubMix API gateway
siliconflow: ProviderConfig = Field(default_factory=ProviderConfig) # SiliconFlow (硅基流动)
edenai: ProviderConfig = Field(default_factory=ProviderConfig) # Eden AI API gateway
novita: ProviderConfig = Field(default_factory=ProviderConfig) # Novita AI
volcengine: ProviderConfig = Field(default_factory=ProviderConfig) # VolcEngine (火山引擎)
volcengine_coding_plan: ProviderConfig = Field(default_factory=ProviderConfig) # VolcEngine Coding Plan
@@ -504,6 +505,7 @@ class Config(BaseSettings):
model_normalized = model_lower.replace("-", "_")
model_prefix = model_lower.split("/", 1)[0] if "/" in model_lower else ""
normalized_prefix = model_prefix.replace("-", "_")
prefixed_provider = find_by_name(model_prefix) if model_prefix else None
def _kw_matches(kw: str) -> bool:
kw = kw.lower()
@@ -533,6 +535,22 @@ class Config(BaseSettings):
continue
p = getattr(self.providers, spec.name, None)
if p and any(_kw_matches(kw) for kw in spec.keywords):
# Local providers (Ollama, vLLM, …) keep model-family keywords
# like "nemotron" or "llama" to enable bare-model auto-routing,
# but those keywords collide with cloud-hosted variants of the
# same family (e.g. `nvidia/nemotron-...` via OpenRouter). Only
# honor a local keyword match when the user has actually
# configured that local endpoint via `api_base` — mirrors the
# gate already used by the local-fallback loop below.
if spec.is_local:
# A qualified model belongs to its explicit provider or a
# gateway fallback, never to a different local provider
# whose model-family keyword happens to match.
foreign_prefix = bool(
prefixed_provider is not None and prefixed_provider.name != spec.name
)
if not p.api_base or foreign_prefix:
continue
if spec.is_oauth or spec.is_local or spec.is_direct or p.api_key:
return p, spec.name
@@ -541,16 +559,17 @@ class Config(BaseSettings):
# Prefer providers whose detect_by_base_keyword matches the configured api_base
# (e.g. Ollama's "11434" in "http://localhost:11434") over plain registry order.
local_fallback: tuple[ProviderConfig, str] | None = None
for spec in PROVIDERS:
if not spec.is_local:
continue
p = getattr(self.providers, spec.name, None)
if not (p and p.api_base):
continue
if spec.detect_by_base_keyword and spec.detect_by_base_keyword in p.api_base:
return p, spec.name
if local_fallback is None:
local_fallback = (p, spec.name)
if prefixed_provider is None:
for spec in PROVIDERS:
if not spec.is_local:
continue
p = getattr(self.providers, spec.name, None)
if not (p and p.api_base):
continue
if spec.detect_by_base_keyword and spec.detect_by_base_keyword in p.api_base:
return p, spec.name
if local_fallback is None:
local_fallback = (p, spec.name)
if local_fallback:
return local_fallback
+48 -33
View File
@@ -75,13 +75,22 @@ def _validate_schedule_for_add(schedule: CronSchedule) -> None:
if schedule.tz and schedule.kind != "cron":
raise ValueError("tz can only be used with cron schedules")
if schedule.kind == "cron" and schedule.tz:
if schedule.kind == "cron":
if not schedule.expr or not schedule.expr.strip():
raise ValueError("cron schedule requires a non-empty 'expr'")
try:
from zoneinfo import ZoneInfo
from croniter import croniter
ZoneInfo(schedule.tz)
except Exception:
raise ValueError(f"unknown timezone '{schedule.tz}'") from None
croniter(schedule.expr)
except Exception as exc:
raise ValueError(f"invalid cron expression '{schedule.expr}': {exc}") from None
if schedule.tz:
try:
from zoneinfo import ZoneInfo
ZoneInfo(schedule.tz)
except Exception:
raise ValueError(f"unknown timezone '{schedule.tz}'") from None
def _has_legacy_delivery_context(payload: CronPayload) -> bool:
@@ -163,9 +172,13 @@ class CronService:
self._store: CronStore | None = None
self._timer_task: asyncio.Task[None] | None = None
self._running = False
self._timer_active = False
self._active_executions = 0
self.max_sleep_ms = max_sleep_ms
def _should_persist_store(self) -> bool:
"""Return whether this instance currently owns the live store."""
return self._running or self._active_executions > 0
def _is_unbound_agent_job(self, job: CronJob) -> bool:
return job.payload.kind == "agent_turn" and not is_bound_cron_job(job)
@@ -278,23 +291,24 @@ class CronService:
logger.exception("load action line error")
continue
self._store.jobs = list(jobs_map.values()) # pyright: ignore[reportOptionalMemberAccess]
if self._running and changed:
if self._should_persist_store() and changed:
self._action_path.write_text("", encoding="utf-8")
self._save_store()
return
def _load_store(self) -> CronStore | None:
def _load_store(self, *, reload_during_execution: bool = False) -> CronStore | None:
"""Load jobs from disk. Reloads automatically if file was modified externally.
- Reload every time because it needs to merge operations on the jobs object from other instances.
- During _on_timer execution, return the existing store to prevent concurrent
- During job execution, return the existing store to prevent concurrent
_load_store calls (e.g. from list_jobs polling) from replacing it mid-execution.
The first execution explicitly reloads once when it takes ownership.
- When the on-disk store exists but is unreadable: keep using the
previous in-memory ``self._store`` if we already have one (so a
transient corruption does not drop live jobs); only the very first
load (during ``start``) can return ``None`` to signal an unrecoverable
state to the caller.
"""
if self._timer_active and self._store:
if self._active_executions > 0 and self._store and not reload_during_execution:
return self._store
loaded = self._load_jobs()
if loaded is None:
@@ -307,12 +321,12 @@ class CronService:
jobs, version = loaded
self._store = CronStore(version=version, jobs=jobs)
self._merge_action()
if self._enforce_store_agent_bindings() and self._running:
if self._enforce_store_agent_bindings() and self._should_persist_store():
self._save_store()
return self._store
def _require_store(self) -> CronStore:
def _require_store(self, *, reload_during_execution: bool = False) -> CronStore:
"""Return a usable store or raise a clear error.
``_load_store`` deliberately returns ``None`` when the first load sees
@@ -322,7 +336,7 @@ class CronService:
``AttributeError`` and, more importantly, prevents follow-up saves from
treating a corrupt store as an empty one.
"""
store = self._load_store()
store = self._load_store(reload_during_execution=reload_during_execution)
if store is None:
raise RuntimeError(
f"cron store at {self.store_path} could not be loaded and was preserved "
@@ -504,19 +518,20 @@ class CronService:
async def _on_timer(self) -> None:
"""Handle timer tick - run due jobs."""
self._load_store()
# If a hot reload found a corrupt store on disk, ``self._store`` may
# still hold the previous, known-good in-memory snapshot. Keep using
# it rather than crashing the timer or wiping live jobs.
if not self._store:
self._arm_timer()
return
self._timer_active = True
reload_store = self._active_executions == 0
self._active_executions += 1
try:
store = self._load_store(reload_during_execution=reload_store)
# If a hot reload found a corrupt store on disk, ``self._store`` may
# still hold the previous, known-good in-memory snapshot. Keep using
# it rather than crashing the timer or wiping live jobs.
if store is None:
self._arm_timer()
return
now = _now_ms()
due_jobs = [
j for j in self._store.jobs
j for j in store.jobs
if j.enabled and j.state.next_run_at_ms and now >= j.state.next_run_at_ms
]
@@ -525,7 +540,7 @@ class CronService:
self._save_store()
finally:
self._timer_active = False
self._active_executions -= 1
self._arm_timer()
async def _execute_job(self, job: CronJob) -> None:
@@ -657,7 +672,7 @@ class CronService:
)
_normalize_agent_turn_job(job)
self._enforce_agent_binding(job)
if self._running:
if self._should_persist_store():
store = self._require_store()
store.jobs.append(job)
self._save_store()
@@ -697,7 +712,7 @@ class CronService:
removed = len(store.jobs) < before
if removed:
if self._running:
if self._should_persist_store():
self._save_store()
self._arm_timer()
else:
@@ -719,7 +734,7 @@ class CronService:
job.state.next_run_at_ms = _compute_next_run(job.schedule, _now_ms())
else:
job.state.next_run_at_ms = None
if self._running:
if self._should_persist_store():
self._save_store()
self._arm_timer()
else:
@@ -775,7 +790,7 @@ class CronService:
else:
job.state.next_run_at_ms = None
if self._running:
if self._should_persist_store():
self._save_store()
self._arm_timer()
else:
@@ -786,10 +801,10 @@ class CronService:
async def run_job(self, job_id: str, force: bool = False) -> bool:
"""Manually run a job without disturbing the service's running state."""
was_running = self._running
self._running = True
reload_store = self._active_executions == 0
self._active_executions += 1
try:
store = self._require_store()
store = self._require_store(reload_during_execution=reload_store)
for job in store.jobs:
if job.id == job_id:
if self._is_unbound_agent_job(job):
@@ -803,8 +818,8 @@ class CronService:
return True
return False
finally:
self._running = was_running
if was_running:
self._active_executions -= 1
if self._running and self._active_executions == 0:
self._arm_timer()
def get_job(self, job_id: str) -> CronJob | None:
+22 -1
View File
@@ -2,6 +2,8 @@
from __future__ import annotations
import json
import os
import shutil
import subprocess
import sys
from dataclasses import dataclass
@@ -179,13 +181,18 @@ def extra_installed(extra: str, deps: list[str] | None) -> bool:
return all(requirement_installed(dep, extra) for dep in deps)
def run_install_command(argv: list[str]) -> subprocess.CompletedProcess[str]:
def run_install_command(
argv: list[str],
*,
env: dict[str, str] | None = None,
) -> subprocess.CompletedProcess[str]:
try:
return subprocess.run(
argv,
capture_output=True,
text=True,
timeout=_INSTALL_TIMEOUT_SECONDS,
env=env,
)
except subprocess.TimeoutExpired as exc:
stdout = exc.stdout.decode(errors="replace") if isinstance(exc.stdout, bytes) else exc.stdout
@@ -234,6 +241,20 @@ def install_extra(
failed_cmd = pip_cmd
failed_proc = proc
if missing_pip(proc):
if shutil.which("uv"):
uv_cmd = ["uv", "pip", "install", "--python", sys.executable, *install_args]
uv_env = os.environ.copy()
if index_url := os.environ.get("PIP_INDEX_URL", "").strip():
uv_env["UV_INDEX_URL"] = index_url
logger.info("pip missing while installing '{}'; running {}", extra, command_text(uv_cmd))
uv_proc = runner(uv_cmd, env=uv_env)
_log_completed_command(f"Optional feature '{extra}' uv install", uv_proc)
if uv_proc.returncode == 0:
importlib.invalidate_caches()
return InstallResult(True, label, pip_cmd)
output = (uv_proc.stderr or uv_proc.stdout or "").strip()
return InstallResult(False, label, pip_cmd, failed_cmd=uv_cmd, output=output)
ensure_cmd = [sys.executable, "-m", "ensurepip", "--upgrade"]
logger.info("pip missing while installing '{}'; running {}", extra, command_text(ensure_cmd))
ensure_proc = runner(ensure_cmd)
+29 -4
View File
@@ -40,9 +40,15 @@ def _load() -> dict[str, Any]:
data = json.load(f)
except FileNotFoundError:
return {"approved": {}, "pending": {}}
except (json.JSONDecodeError, OSError):
except json.JSONDecodeError:
logger.warning("Corrupted pairing store, resetting")
return {"approved": {}, "pending": {}}
except OSError:
# A transiently locked or busy file is not corruption. Propagate so
# mutating callers fail loudly instead of persisting an empty view
# that would erase every approved sender.
logger.warning("Pairing store temporarily unreadable: {}", path)
raise
if not isinstance(data, dict):
logger.warning("Corrupted pairing store, resetting")
return {"approved": {}, "pending": {}}
@@ -171,7 +177,11 @@ def deny_code(code: str) -> bool:
def is_approved(channel: str, sender_id: str) -> bool:
"""Check whether *sender_id* has been approved on *channel*."""
with _LOCK:
data = _load()
try:
data = _load()
except OSError:
# Fail closed for this check; the store itself stays untouched.
return False
approved: dict[str, set[str]] = data.get("approved", {})
return str(sender_id) in approved.get(channel, set())
@@ -179,7 +189,10 @@ def is_approved(channel: str, sender_id: str) -> bool:
def list_pending() -> list[dict[str, Any]]:
"""Return all non-expired pending pairing requests."""
with _LOCK:
data = _load()
try:
data = _load()
except OSError:
return []
_gc_pending(data)
return [
{"code": code, **info}
@@ -257,7 +270,10 @@ def clear_channel(channel: str) -> dict[str, int]:
def get_approved(channel: str) -> list[str]:
"""Return all approved sender IDs for *channel*."""
with _LOCK:
data = _load()
try:
data = _load()
except OSError:
return []
return sorted(data.get("approved", {}).get(channel, set()))
@@ -283,6 +299,15 @@ def handle_pairing_command(channel: str, subcommand_text: str) -> str:
This is a pure function (no side effects other than store mutations)
so it can be used from both the CLI and the agent CommandRouter.
"""
try:
return _handle_pairing_subcommand(channel, subcommand_text)
except OSError:
# Mutations fail loudly on a transient I/O error instead of lying
# ("invalid code") or silently rewriting the store from an empty view.
return "The pairing store is temporarily unavailable. Please try again."
def _handle_pairing_subcommand(channel: str, subcommand_text: str) -> str:
parts = subcommand_text.split()
sub = parts[0] if parts else "list"
arg = parts[1] if len(parts) > 1 else None
+50 -10
View File
@@ -31,6 +31,36 @@ def _gen_tool_id() -> str:
_VALID_TOOL_ID = re.compile(r"^[a-zA-Z0-9_-]+$")
_CLAUDE_MODEL_VERSION = re.compile(
r"claude-(?P<family>[a-z]+)-(?P<major>\d+)"
r"(?:-(?P<minor>\d{1,2})(?=-|$))?"
)
_ADAPTIVE_ONLY_MIN_VERSIONS = {
"opus": (4, 7),
"sonnet": (5, 0),
"fable": (5, 0),
"mythos": (5, 0),
}
_THINKING_DISABLE_MIN_VERSIONS = {
"opus": (5, 0),
"sonnet": (5, 0),
}
_SAMPLING_DEPRECATED_MODELS = {"claude-mythos-preview"}
def _model_version_at_least(
model_name: str,
minimum_versions: dict[str, tuple[int, int]],
) -> bool:
match = _CLAUDE_MODEL_VERSION.search(model_name.lower())
if match is None:
return False
minimum = minimum_versions.get(match.group("family"))
if minimum is None:
return False
version = (int(match.group("major")), int(match.group("minor") or 0))
return version >= minimum
def _sanitize_tool_id(tid: str) -> str:
"""Ensure tool_use/tool_result IDs match Anthropic's required pattern.
@@ -562,13 +592,13 @@ class AnthropicProvider(LLMProvider):
)
max_tokens = max(1, max_tokens)
thinking_enabled = bool(reasoning_effort) and reasoning_effort.lower() != "none"
# Several Anthropic models (opus-4-7, opus-4-8, sonnet-5, fable) deprecated the
# `temperature` parameter — the API returns 400 if it is present.
_model_lower = model_name.lower()
omit_temperature = any(
m in _model_lower for m in ("opus-4-7", "opus-4-8", "sonnet-5", "fable")
reasoning_effort_lower = reasoning_effort.lower() if reasoning_effort else None
thinking_enabled = reasoning_effort_lower not in (None, "", "none")
adaptive_only = _model_version_at_least(model_name, _ADAPTIVE_ONLY_MIN_VERSIONS)
# Mythos Preview rejects sampling parameters but still accepts manual
# thinking budgets, so it is not part of the adaptive-only capability.
omit_temperature = (
adaptive_only or model_name.lower() in _SAMPLING_DEPRECATED_MODELS
)
kwargs: dict[str, Any] = {
@@ -580,16 +610,26 @@ class AnthropicProvider(LLMProvider):
if system:
kwargs["system"] = system
if reasoning_effort == "adaptive":
if reasoning_effort_lower == "none" and _model_version_at_least(
model_name, _THINKING_DISABLE_MIN_VERSIONS
):
# These models think by default, so omission would not honor an
# explicit request to disable thinking.
kwargs["thinking"] = {"type": "disabled"}
elif reasoning_effort_lower == "adaptive":
# Adaptive thinking: model decides when and how much to think
# Supported on claude-sonnet-4-6 and claude-opus-4-6.
# Also auto-enables interleaved thinking between tool calls.
kwargs["thinking"] = {"type": "adaptive"}
if not omit_temperature:
kwargs["temperature"] = 1.0
elif thinking_enabled and adaptive_only:
# Newer Claude models removed manual token budgets. Their effort
# control is independent from the adaptive thinking mode.
kwargs["thinking"] = {"type": "adaptive"}
kwargs["output_config"] = {"effort": reasoning_effort_lower}
elif thinking_enabled:
budget_map = {"low": 1024, "medium": 4096, "high": max(8192, max_tokens)}
budget = budget_map.get(cast(str, reasoning_effort).lower(), 4096)
budget = budget_map.get(reasoning_effort_lower, 4096)
kwargs["thinking"] = {"type": "enabled", "budget_tokens": budget}
kwargs["max_tokens"] = max(max_tokens, budget + 4096)
if not omit_temperature:
+175 -11
View File
@@ -23,14 +23,26 @@ import uuid
from collections.abc import Awaitable, Callable
from typing import Any, cast
from loguru import logger
from openai import AsyncOpenAI
from nanobot.providers.base import LLMProvider, LLMResponse
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
)
from nanobot.providers.openai_responses import (
ResponsesStreamCapture,
build_responses_state,
consume_sdk_stream,
convert_messages,
convert_tools,
is_compaction_compatibility_error,
is_replayable_finish_reason,
parse_response_output,
prepare_responses_input,
resolve_compact_threshold,
responses_state_matches,
)
_AZURE_OPENAI_SCOPE = "https://cognitiveservices.azure.com/.default"
@@ -97,6 +109,7 @@ class AzureOpenAIProvider(LLMProvider):
):
super().__init__(api_key, api_base)
self.default_model = default_model
self._native_compaction_available = True
if not api_base:
raise ValueError("Azure OpenAI api_base is required")
@@ -142,6 +155,25 @@ class AzureOpenAIProvider(LLMProvider):
name = deployment_name.lower()
return not any(token in name for token in ("gpt-5", "o1", "o3", "o4"))
def _responses_state_provider(self) -> str:
return f"azure_openai:{str(self.api_base).rstrip('/')}"
def can_resume_conversation_state(
self,
state: ProviderConversationState,
model: str | None = None,
) -> bool:
return responses_state_matches(
state,
provider=self._responses_state_provider(),
model=model or self.default_model,
)
def supports_native_compaction(self, model: str | None = None) -> bool:
"""Azure's native Responses endpoint accepts context management."""
_ = model
return self._native_compaction_available
def _build_body(
self,
messages: list[dict[str, Any]],
@@ -151,10 +183,26 @@ class AzureOpenAIProvider(LLMProvider):
temperature: float,
reasoning_effort: str | None,
tool_choice: str | dict[str, Any] | None,
provider_context: ProviderCallContext | None = None,
) -> dict[str, Any]:
"""Build the Responses API request body from Chat-Completions-style args."""
deployment = model or self.default_model
instructions, input_items = convert_messages(self._sanitize_empty_content(messages))
sanitized_messages = self._sanitize_empty_content(messages)
sanitized_state = (
provider_context.conversation_state
if provider_context is not None
else None
)
if sanitized_state is not None:
sanitized_state = sanitized_state.with_pending_messages(
self._sanitize_empty_content(sanitized_state.pending_messages)
)
instructions, input_items, replayed = prepare_responses_input(
sanitized_messages,
state=sanitized_state,
provider=self._responses_state_provider(),
model=deployment,
)
body: dict[str, Any] = {
"model": deployment,
@@ -164,13 +212,29 @@ class AzureOpenAIProvider(LLMProvider):
"store": False,
"stream": False,
}
compact_threshold = resolve_compact_threshold(
(
provider_context.context_window_tokens
if provider_context is not None
else None
),
max_tokens,
)
if self.supports_native_compaction(deployment) and compact_threshold is not None:
body["context_management"] = [{
"type": "compaction",
"compact_threshold": compact_threshold,
}]
if self._supports_temperature(deployment, reasoning_effort):
body["temperature"] = temperature
if not self._supports_temperature(deployment, reasoning_effort):
body["include"] = ["reasoning.encrypted_content"]
if reasoning_effort and reasoning_effort.lower() != "none":
body["reasoning"] = {"effort": reasoning_effort}
body["include"] = ["reasoning.encrypted_content"]
if replayed and "gpt-5.6" in deployment.lower():
body.setdefault("reasoning", {})["context"] = "all_turns"
if tools:
body["tools"] = convert_tools(tools)
@@ -178,21 +242,97 @@ class AzureOpenAIProvider(LLMProvider):
return body
async def _create_response_with_compaction_fallback(
self,
body: dict[str, Any],
) -> Any:
"""Retry once without server compaction when Azure rejects the option."""
try:
return cast(Any, await self._client.responses.create(**body))
except Exception as exc:
if (
"context_management" not in body
or not is_compaction_compatibility_error(exc)
):
raise
self._native_compaction_available = False
body.pop("context_management", None)
logger.warning(
"Azure Responses server compaction unsupported; disabled for this provider "
"instance (status={})",
getattr(exc, "status_code", None),
)
return cast(Any, await self._client.responses.create(**body))
@staticmethod
def _handle_error(e: Exception) -> LLMResponse:
response = getattr(e, "response", None)
body = getattr(e, "body", None) or getattr(response, "text", None)
body_text = str(body).strip() if body is not None else ""
msg = f"Error: {body_text[:500]}" if body_text else f"Error calling Azure OpenAI: {e}"
retry_after = LLMProvider._extract_retry_after_from_headers(getattr(response, "headers", None))
headers = getattr(response, "headers", None)
retry_after = LLMProvider._extract_retry_after_from_headers(headers)
if retry_after is None:
retry_after = LLMProvider._extract_retry_after(msg)
return LLMResponse(content=msg, finish_reason="error", retry_after=retry_after)
status_code = getattr(e, "status_code", None)
if status_code is None and response is not None:
status_code = getattr(response, "status_code", None)
error_type, error_code = LLMProvider._extract_error_type_code(body)
should_retry: bool | None = None
if headers is not None:
raw_should_retry = headers.get("x-should-retry")
if isinstance(raw_should_retry, str):
lowered = raw_should_retry.strip().lower()
if lowered == "true":
should_retry = True
elif lowered == "false":
should_retry = False
error_name = type(e).__name__.lower()
error_kind = (
"timeout"
if "timeout" in error_name
else "connection"
if "connection" in error_name
else None
)
return LLMResponse(
content=msg,
finish_reason="error",
retry_after=retry_after,
error_status_code=int(status_code) if status_code is not None else None,
error_kind=error_kind,
error_type=error_type,
error_code=error_code,
error_retry_after_s=retry_after,
error_should_retry=should_retry,
)
# ------------------------------------------------------------------
# Public API
# ------------------------------------------------------------------
async def chat_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
return await self.chat(
**kwargs,
provider_context=provider_context,
)
async def chat_stream_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
return await self.chat_stream(
**kwargs,
provider_context=provider_context,
)
async def chat(
self,
messages: list[dict[str, Any]],
@@ -202,14 +342,21 @@ class AzureOpenAIProvider(LLMProvider):
temperature: float = 0.7,
reasoning_effort: str | None = None,
tool_choice: str | dict[str, Any] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
body = self._build_body(
messages, tools, model, max_tokens, temperature,
reasoning_effort, tool_choice,
provider_context,
)
try:
response = cast(Any, await self._client.responses.create(**body))
return parse_response_output(response)
response = await self._create_response_with_compaction_fallback(body)
return parse_response_output(
response,
state_provider=self._responses_state_provider(),
state_model=str(body["model"]),
state_input_items=cast(list[dict[str, Any]], body["input"]),
)
except Exception as e:
return self._handle_error(e)
@@ -225,26 +372,43 @@ class AzureOpenAIProvider(LLMProvider):
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
_ = on_thinking_delta
body = self._build_body(
messages, tools, model, max_tokens, temperature,
reasoning_effort, tool_choice,
provider_context,
)
body["stream"] = True
try:
stream = cast(Any, await self._client.responses.create(**body))
stream = await self._create_response_with_compaction_fallback(body)
capture = ResponsesStreamCapture()
content, tool_calls, finish_reason, usage, reasoning_content = (
await consume_sdk_stream(stream, on_content_delta, on_tool_call_delta)
await consume_sdk_stream(
stream,
on_content_delta,
on_tool_call_delta,
capture=capture,
)
)
return LLMResponse(
result = LLMResponse(
content=content or None,
tool_calls=tool_calls,
finish_reason=finish_reason,
usage=usage,
reasoning_content=reasoning_content,
)
if capture.completed and is_replayable_finish_reason(finish_reason):
result.provider_state = build_responses_state(
provider=self._responses_state_provider(),
model=str(body["model"]),
input_items=cast(list[dict[str, Any]], body["input"]),
output_items=capture.output_items,
usage=usage,
)
return result
except Exception as e:
return self._handle_error(e)
+201 -8
View File
@@ -1,5 +1,7 @@
"""Base LLM provider interface."""
from __future__ import annotations
import asyncio
import json
import os
@@ -7,6 +9,7 @@ import re
from abc import ABC, abstractmethod
from collections.abc import Awaitable, Callable
from contextlib import suppress
from copy import deepcopy
from dataclasses import dataclass, field
from datetime import datetime, timezone
from email.utils import parsedate_to_datetime
@@ -150,6 +153,104 @@ def tool_arguments_json_for_replay(arguments: Any) -> str:
return json.dumps(tool_arguments_object_for_replay(arguments), ensure_ascii=False)
@dataclass
class ProviderConversationState:
"""Opaque provider-owned continuation state.
``payload`` may contain encrypted reasoning or other provider-private
protocol items. Keep it out of normal logs and public chat history.
``pending_messages`` are Chat-style messages produced after the most
recent provider response and are materialized by the owning provider on
the next request.
"""
kind: str
provider: str
model: str
version: int
payload: dict[str, Any] = field(default_factory=dict, repr=False)
pending_messages: list[dict[str, Any]] = field(default_factory=list, repr=False)
def with_pending_messages(
self,
messages: list[dict[str, Any]],
) -> ProviderConversationState:
"""Return a state copy with an isolated pending-message list."""
return ProviderConversationState(
kind=self.kind,
provider=self.provider,
model=self.model,
version=self.version,
payload=self.payload,
pending_messages=deepcopy(messages),
)
def to_private_record(self) -> dict[str, Any]:
"""Serialize for the private session sidecar, never for public history."""
return {
"kind": self.kind,
"provider": self.provider,
"model": self.model,
"version": self.version,
"payload": deepcopy(self.payload),
"pending_messages": deepcopy(self.pending_messages),
}
@classmethod
def from_private_record(
cls,
value: object,
) -> ProviderConversationState | None:
"""Validate and deserialize a private session-sidecar value."""
if not isinstance(value, dict):
return None
data = cast(dict[str, Any], value)
kind = data.get("kind")
provider = data.get("provider")
model = data.get("model")
version = data.get("version")
payload = data.get("payload")
pending = data.get("pending_messages", [])
if (
not isinstance(kind, str)
or not kind
or not isinstance(provider, str)
or not provider
or not isinstance(model, str)
or not model
or isinstance(version, bool)
or not isinstance(version, int)
or not isinstance(payload, dict)
or not isinstance(pending, list)
or any(
not isinstance(message, dict)
for message in cast(list[object], pending)
)
):
return None
return cls(
kind=kind,
provider=provider,
model=model,
version=version,
payload=deepcopy(cast(dict[str, Any], payload)),
pending_messages=deepcopy(cast(list[dict[str, Any]], pending)),
)
@dataclass(frozen=True)
class ProviderCallContext:
"""Optional provider-owned continuation data for one model request.
The regular ``chat`` contract stays provider-agnostic. Responses-capable
providers consume this context through the opt-in ``chat_with_context``
hooks, while every other provider inherits the context-free delegation.
"""
conversation_state: ProviderConversationState | None = field(default=None, repr=False)
context_window_tokens: int | None = None
@dataclass
class LLMResponse:
"""Response from an LLM provider."""
@@ -160,6 +261,10 @@ class LLMResponse:
retry_after: float | None = None # Provider supplied retry wait in seconds.
reasoning_content: str | None = None # Kimi, DeepSeek-R1, MiMo etc.
thinking_blocks: list[dict[str, Any]] | None = None # Anthropic extended thinking
provider_state: ProviderConversationState | None = field(default=None, repr=False)
# Routing wrappers may preserve or discard an incoming provider-owned
# continuation independently of the final fallback error's retry policy.
preserve_provider_state_on_error: bool | None = field(default=None, repr=False)
# Structured error metadata used by retry policy when finish_reason == "error".
error_status_code: int | None = None
error_kind: str | None = None # e.g. "timeout", "connection"
@@ -274,6 +379,18 @@ class LLMProvider(ABC):
self.api_base = api_base
self.generation: GenerationSettings = GenerationSettings()
def can_resume_conversation_state(
self,
state: ProviderConversationState,
model: str | None = None,
) -> bool:
"""Whether this provider can safely consume an opaque saved state."""
return False
def supports_native_compaction(self, model: str | None = None) -> bool:
"""Whether requests may include provider-native context compaction."""
return False
@staticmethod
def _sanitize_empty_content(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""Sanitize message content: fix empty blocks, strip internal _meta fields.
@@ -416,7 +533,7 @@ class LLMProvider(ABC):
return any(marker in err for marker in cls._TRANSIENT_ERROR_MARKERS)
@classmethod
def _is_transient_response(cls, response: LLMResponse) -> bool:
def is_transient_response(cls, response: LLMResponse) -> bool:
"""Prefer structured error metadata, fallback to text markers for legacy providers."""
if response.error_should_retry is not None:
return bool(response.error_should_retry)
@@ -607,6 +724,21 @@ class LLMProvider(ABC):
result.append(msg)
return result if found else None
@staticmethod
def _contains_image_content(value: object) -> bool:
"""Return whether a JSON-like provider payload contains an input image."""
if isinstance(value, dict):
mapping = cast(dict[str, object], value)
if mapping.get("type") in {"image_url", "input_image"}:
return True
return any(LLMProvider._contains_image_content(item) for item in mapping.values())
if isinstance(value, list):
return any(
LLMProvider._contains_image_content(item)
for item in cast(list[object], value)
)
return False
@staticmethod
def _strip_image_content_inplace(messages: list[dict[str, Any]]) -> bool:
"""Replace image_url blocks with text placeholder *in-place*.
@@ -633,6 +765,12 @@ class LLMProvider(ABC):
async def _safe_chat(self, **kwargs: Any) -> LLMResponse:
"""Call chat() and convert unexpected exceptions to error responses."""
try:
provider_context = kwargs.pop("provider_context", None)
if isinstance(provider_context, ProviderCallContext):
return await self.chat_with_context(
provider_context=provider_context,
**kwargs,
)
return await self.chat(**kwargs)
except asyncio.CancelledError:
raise
@@ -666,17 +804,47 @@ class LLMProvider(ABC):
"""
_ = on_thinking_delta, on_tool_call_delta
response = await self.chat(
messages=messages, tools=tools, model=model,
max_tokens=max_tokens, temperature=temperature,
reasoning_effort=reasoning_effort, tool_choice=tool_choice,
messages=messages,
tools=tools,
model=model,
max_tokens=max_tokens,
temperature=temperature,
reasoning_effort=reasoning_effort,
tool_choice=tool_choice,
)
if on_content_delta and response.content:
await on_content_delta(response.content)
return response
async def chat_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
"""Opt-in continuation hook; ordinary providers delegate to ``chat``."""
_ = provider_context
return await self.chat(**kwargs)
async def chat_stream_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
"""Streaming continuation hook with a context-free default."""
_ = provider_context
return await self.chat_stream(**kwargs)
async def _safe_chat_stream(self, **kwargs: Any) -> LLMResponse:
"""Call chat_stream() and convert unexpected exceptions to error responses."""
try:
provider_context = kwargs.pop("provider_context", None)
if isinstance(provider_context, ProviderCallContext):
return await self.chat_stream_with_context(
provider_context=provider_context,
**kwargs,
)
return await self.chat_stream(**kwargs)
except asyncio.CancelledError:
raise
@@ -698,6 +866,7 @@ class LLMProvider(ABC):
on_stream_recover: Callable[[], Awaitable[None]] | None = None,
retry_mode: str = "standard",
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
"""Call chat_stream() with retry on transient provider failures."""
if max_tokens is self._SENTINEL or max_tokens is None:
@@ -730,6 +899,8 @@ class LLMProvider(ABC):
on_thinking_delta=on_thinking_delta,
on_tool_call_delta=on_tool_call_delta,
)
if provider_context is not None:
kw["provider_context"] = provider_context
if on_stream_recover and getattr(self, "supports_stream_recover_callback", False):
kw["on_stream_recover"] = _recover_stream
return await self._run_with_retry(
@@ -753,6 +924,7 @@ class LLMProvider(ABC):
tool_choice: str | dict[str, Any] | None = None,
retry_mode: str = "standard",
on_retry_wait: Callable[[str], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
"""Call chat() with retry on transient provider failures.
@@ -775,6 +947,8 @@ class LLMProvider(ABC):
max_tokens=max_tokens, temperature=temperature,
reasoning_effort=reasoning_effort, tool_choice=tool_choice,
)
if provider_context is not None:
kw["provider_context"] = provider_context
return await self._run_with_retry(
self._safe_chat,
kw,
@@ -932,14 +1106,33 @@ class LLMProvider(ABC):
last_error_key = error_key
identical_error_count = 1 if error_key else 0
if not self._is_transient_response(response):
stripped = self._strip_image_content(original_messages)
if stripped is not None and stripped != kw["messages"]:
if not self.is_transient_response(response):
stripped = self._strip_image_content(kw["messages"])
provider_context = kw.get("provider_context")
stripped_context: ProviderCallContext | None = None
if isinstance(provider_context, ProviderCallContext):
state = provider_context.conversation_state
if state is not None and (
stripped is not None
or self._strip_image_content(state.pending_messages) is not None
or self._contains_image_content(state.payload)
):
# Provider-owned payloads may retain earlier input_image items.
# Rebuild from the stripped public transcript for this retry.
stripped_context = ProviderCallContext(
context_window_tokens=(
provider_context.context_window_tokens
),
)
if stripped is not None or stripped_context is not None:
logger.warning(
"Non-transient LLM error with image content, retrying without images"
)
retry_kw = dict(kw)
retry_kw["messages"] = stripped
if stripped is not None:
retry_kw["messages"] = stripped
if stripped_context is not None:
retry_kw["provider_context"] = stripped_context
result = await call(**retry_kw)
# Permanently strip images from the original messages so
# subsequent iterations do not repeat the error-retry cycle.
+262
View File
@@ -0,0 +1,262 @@
"""Provider-owned conversation-state lifecycle coordination."""
from __future__ import annotations
from copy import deepcopy
from typing import Any, cast
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
)
_PROVIDER_STATE_OUTPUT_META = "provider_state_output"
_PROVIDER_STATE_BOUNDARY_META = "provider_state_boundary"
def allows_conversation_message_merge(message: dict[str, Any]) -> bool:
"""Return whether new same-role input may merge into *message*."""
internal_meta = cast(object, message.get("_meta"))
return not (
isinstance(internal_meta, dict)
and cast(dict[str, Any], internal_meta).get(
_PROVIDER_STATE_BOUNDARY_META
) is True
)
class ProviderConversationStateController:
"""Keep provider conversation-state semantics outside the agent runner.
The runner owns the tool loop and reports lifecycle events here. This
controller owns capability checks, transcript deltas, response projections,
retry transitions, and durable snapshots for provider-private state.
"""
def __init__(
self,
*,
provider: LLMProvider,
model: str | None,
messages: list[dict[str, Any]],
state: ProviderConversationState | None = None,
) -> None:
self._provider = provider
self._model = model
self._state = (
state
if state is not None
and provider.can_resume_conversation_state(state, model)
else None
)
self._boundary = len(messages)
self._request_messages: list[dict[str, Any]] = []
def independent_request_context(
self,
*,
context_window_tokens: int | None,
) -> ProviderCallContext | None:
"""Return typed provider context for a request that does not resume state."""
if context_window_tokens is None:
return None
return ProviderCallContext(context_window_tokens=context_window_tokens)
def prepare_request(
self,
messages: list[dict[str, Any]],
*,
context_window_tokens: int | None,
model_messages: list[dict[str, Any]] | None = None,
supplemental_messages: list[dict[str, Any]] | None = None,
) -> ProviderCallContext | None:
"""Build typed context for the next request and remember its durable delta."""
independent_context = self.independent_request_context(
context_window_tokens=context_window_tokens,
)
if self._state is None:
self._request_messages = []
return independent_context
if not self._provider.can_resume_conversation_state(
self._state,
self._model,
):
self._state = None
self._request_messages = []
return independent_context
durable_messages = self._messages_after_boundary(messages)
governed_messages = (
self._model_messages_after_boundary(model_messages)
if model_messages is not None and durable_messages
else None
)
request_messages = (
governed_messages
if governed_messages is not None
else durable_messages
)
supplemental = deepcopy(supplemental_messages or [])
self._request_messages = deepcopy(request_messages)
request_state = self._state.with_pending_messages([
*self._state.pending_messages,
*request_messages,
*supplemental,
])
return ProviderCallContext(
conversation_state=request_state,
context_window_tokens=(
independent_context.context_window_tokens
if independent_context is not None
else None
),
)
def observe_response(
self,
response: LLMResponse,
messages: list[dict[str, Any]],
*,
adopt_candidate_state: bool = True,
) -> None:
"""Advance, preserve, or discard state after one provider response."""
candidate = response.provider_state if adopt_candidate_state else None
candidate_is_replayable = response.finish_reason in {
"stop",
"tool_calls",
"function_call",
}
if (
candidate is not None
and candidate_is_replayable
and self._provider.can_resume_conversation_state(
candidate,
self._model,
)
):
self._state = candidate
self._boundary = len(messages)
self._seal_boundary(messages)
elif response.finish_reason == "error" and (
response.preserve_provider_state_on_error is True
or (
response.preserve_provider_state_on_error is None
and LLMProvider.is_transient_response(response)
)
):
if self._state is not None and self._request_messages:
self._state = self._state.with_pending_messages([
*self._state.pending_messages,
*self._request_messages,
])
self._boundary = len(messages)
else:
self._state = None
self._boundary = len(messages)
self._request_messages = []
@staticmethod
def project_response_message(
message: dict[str, Any],
response: LLMResponse,
) -> dict[str, Any]:
"""Mark a Chat projection already represented by provider output."""
if response.provider_state is None:
return message
internal_meta = dict(message.get("_meta") or {})
internal_meta[_PROVIDER_STATE_OUTPUT_META] = True
message["_meta"] = internal_meta
return message
def checkpoint(
self,
messages: list[dict[str, Any]],
*,
model_messages: list[dict[str, Any]] | None = None,
) -> ProviderConversationState | None:
"""Return a durable state snapshot without changing live state."""
if self._state is None:
return None
durable_messages = self._messages_after_boundary(messages)
governed_messages = (
self._model_messages_after_boundary(model_messages)
if model_messages is not None and durable_messages
else None
)
pending_messages = (
governed_messages
if governed_messages is not None
else durable_messages
)
return self._state.with_pending_messages([
*self._state.pending_messages,
*pending_messages,
])
def finish(
self,
messages: list[dict[str, Any]],
) -> ProviderConversationState | None:
"""Return the final durable state after all runner messages are known."""
self._state = self.checkpoint(messages)
return self._state
def _messages_after_boundary(
self,
messages: list[dict[str, Any]],
) -> list[dict[str, Any]]:
pending: list[dict[str, Any]] = []
for message in messages[self._boundary:]:
internal_meta = cast(object, message.get("_meta"))
if (
isinstance(internal_meta, dict)
and cast(dict[str, Any], internal_meta).get(
_PROVIDER_STATE_OUTPUT_META
) is True
):
continue
pending.append(deepcopy(message))
return pending
@staticmethod
def _model_messages_after_boundary(
messages: list[dict[str, Any]],
) -> list[dict[str, Any]] | None:
"""Return the governed delta after the latest provider-owned boundary."""
boundary = None
for idx in range(len(messages) - 1, -1, -1):
internal_meta = cast(object, messages[idx].get("_meta"))
if (
isinstance(internal_meta, dict)
and cast(dict[str, Any], internal_meta).get(
_PROVIDER_STATE_BOUNDARY_META
) is True
):
boundary = idx
break
if boundary is None:
return None
pending: list[dict[str, Any]] = []
for message in messages[boundary + 1:]:
internal_meta = cast(object, message.get("_meta"))
if (
isinstance(internal_meta, dict)
and cast(dict[str, Any], internal_meta).get(
_PROVIDER_STATE_OUTPUT_META
) is True
):
continue
pending.append(deepcopy(message))
return pending
@staticmethod
def _seal_boundary(messages: list[dict[str, Any]]) -> None:
"""Prevent later same-role injection merging across a state boundary."""
if not messages:
return
internal_meta = dict(messages[-1].get("_meta") or {})
internal_meta[_PROVIDER_STATE_BOUNDARY_META] = True
messages[-1]["_meta"] = internal_meta
+1
View File
@@ -261,6 +261,7 @@ def make_provider(
primary=provider,
fallback_presets=fallback_presets,
provider_factory=lambda fb: _make_provider_core(config, preset=fb),
primary_context_window_tokens=resolved.context_window_tokens,
)
return provider
+113 -2
View File
@@ -6,11 +6,18 @@ from __future__ import annotations
import time
from collections.abc import Awaitable, Callable
from dataclasses import replace
from typing import Any
from loguru import logger
from nanobot.providers.base import GenerationSettings, LLMProvider, LLMResponse
from nanobot.providers.base import (
GenerationSettings,
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
)
# Circuit breaker tuned to match OpenAICompatProvider's Responses API breaker.
_PRIMARY_FAILURE_THRESHOLD = 3
@@ -113,11 +120,13 @@ class FallbackProvider(LLMProvider):
fallback_presets: list[Any],
provider_factory: Callable[[Any], LLMProvider],
fallback_model_observer: FallbackModelObserver | None = None,
primary_context_window_tokens: int | None = None,
):
self._primary = primary
self._fallback_presets = list(fallback_presets)
self._provider_factory = provider_factory
self._fallback_model_observer = fallback_model_observer
self._primary_context_window_tokens = primary_context_window_tokens
self._has_fallbacks = bool(fallback_presets)
self._primary_failures = 0
self._primary_tripped_at: float | None = None
@@ -141,6 +150,33 @@ class FallbackProvider(LLMProvider):
def supports_progress_deltas(self) -> bool:
return bool(getattr(self._primary, "supports_progress_deltas", False))
def can_resume_conversation_state(
self,
state: ProviderConversationState,
model: str | None = None,
) -> bool:
return self._primary.can_resume_conversation_state(state, model)
def supports_native_compaction(self, model: str | None = None) -> bool:
return self._primary.supports_native_compaction(model)
def _primary_call_context(
self,
provider_context: ProviderCallContext,
model: str | None,
) -> ProviderCallContext:
context_window_tokens = (
self._primary_context_window_tokens
if self._primary_context_window_tokens is not None
else provider_context.context_window_tokens
)
if not self._primary.supports_native_compaction(model):
context_window_tokens = None
return ProviderCallContext(
conversation_state=provider_context.conversation_state,
context_window_tokens=context_window_tokens,
)
def _primary_available(self) -> bool:
"""Return True if the primary provider is not currently tripped."""
if self._primary_tripped_at is None:
@@ -157,6 +193,25 @@ class FallbackProvider(LLMProvider):
lambda p, kw: p.chat(**kw), kwargs, has_streamed=None
)
async def chat_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
call_kwargs: dict[str, Any] = dict(kwargs)
call_kwargs["provider_context"] = self._primary_call_context(
provider_context,
kwargs.get("model"),
)
if not self._has_fallbacks:
return await self._primary.chat_with_context(**call_kwargs)
return await self._try_with_fallback(
lambda p, kw: p.chat_with_context(**kw),
call_kwargs,
has_streamed=None,
)
async def chat_stream(self, **kwargs: Any) -> LLMResponse:
on_stream_recover = kwargs.pop("on_stream_recover", None)
if not self._has_fallbacks:
@@ -179,6 +234,38 @@ class FallbackProvider(LLMProvider):
on_stream_recover=on_stream_recover,
)
async def chat_stream_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
on_stream_recover = kwargs.pop("on_stream_recover", None)
call_kwargs: dict[str, Any] = dict(kwargs)
call_kwargs["provider_context"] = self._primary_call_context(
provider_context,
kwargs.get("model"),
)
if not self._has_fallbacks:
return await self._primary.chat_stream_with_context(**call_kwargs)
has_streamed: list[bool] = [False]
original_delta = call_kwargs.get("on_content_delta")
async def _tracking_delta(text: str) -> None:
if text:
has_streamed[0] = True
if original_delta:
await original_delta(text)
call_kwargs["on_content_delta"] = _tracking_delta
return await self._try_with_fallback(
lambda p, kw: p.chat_stream_with_context(**kw),
call_kwargs,
has_streamed=has_streamed,
on_stream_recover=on_stream_recover,
)
async def _try_with_fallback(
self,
call: Callable[[LLMProvider, dict[str, Any]], Awaitable[LLMResponse]],
@@ -189,6 +276,9 @@ class FallbackProvider(LLMProvider):
primary_model = kwargs.get("model") or self._primary.get_default_model()
primary_was_attempted = False
primary_error = "unknown error"
# A primary error eligible for failover did not return a replacement
# continuation, so the incoming primary state remains reusable.
preserve_primary_state = True
if self._primary_available():
primary_was_attempted = True
@@ -286,6 +376,23 @@ class FallbackProvider(LLMProvider):
"max_tokens": fallback.max_tokens,
"temperature": fallback.temperature,
}
provider_context = fallback_kwargs.get("provider_context")
if isinstance(provider_context, ProviderCallContext):
state = provider_context.conversation_state
if state is not None and not fallback_provider.can_resume_conversation_state(
state,
fallback_model,
):
state = None
context_window_tokens = (
fallback.context_window_tokens
if fallback_provider.supports_native_compaction(fallback_model)
else None
)
fallback_kwargs["provider_context"] = ProviderCallContext(
conversation_state=state,
context_window_tokens=context_window_tokens,
)
if fallback.reasoning_effort is None:
fallback_kwargs.pop("reasoning_effort", None)
else:
@@ -312,11 +419,15 @@ class FallbackProvider(LLMProvider):
)
# Return the last error response we saw (primary or last fallback).
if last_response is not None:
return last_response
return replace(
last_response,
preserve_provider_state_on_error=preserve_primary_state,
)
# Primary was tripped and we have no fallbacks — synthesize an error.
return LLMResponse(
content=f"Primary model '{primary_model}' circuit open and no fallbacks available",
finish_reason="error",
preserve_provider_state_on_error=preserve_primary_state,
)
async def _notify_fallback_model(self, model: str) -> None:
+5 -1
View File
@@ -16,7 +16,7 @@ import httpx
from oauth_cli_kit.models import OAuthToken
from oauth_cli_kit.storage import FileTokenStorage
from nanobot.providers.base import LLMResponse
from nanobot.providers.base import LLMResponse, ProviderCallContext
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
DEFAULT_GITHUB_DEVICE_CODE_URL = "https://github.com/login/device/code"
@@ -248,6 +248,7 @@ class GitHubCopilotProvider(OpenAICompatProvider):
temperature: float = 0.7,
reasoning_effort: str | None = None,
tool_choice: str | dict[str, Any] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
await self._refresh_client_api_key()
return await super().chat(
@@ -258,6 +259,7 @@ class GitHubCopilotProvider(OpenAICompatProvider):
temperature=temperature,
reasoning_effort=reasoning_effort,
tool_choice=tool_choice,
provider_context=provider_context,
)
async def chat_stream(
@@ -272,6 +274,7 @@ class GitHubCopilotProvider(OpenAICompatProvider):
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
await self._refresh_client_api_key()
return await super().chat_stream(
@@ -285,4 +288,5 @@ class GitHubCopilotProvider(OpenAICompatProvider):
on_content_delta=on_content_delta,
on_thinking_delta=on_thinking_delta,
on_tool_call_delta=on_tool_call_delta,
provider_context=provider_context,
)
+12 -5
View File
@@ -808,7 +808,12 @@ class GeminiImageGenerationClient(ImageGenerationProvider):
generation_config: dict[str, Any] = {"responseModalities": ["TEXT", "IMAGE"]}
image_config = _gemini_flash_image_config(model, aspect_ratio, image_size)
if image_config:
generation_config["responseFormat"] = {"image": image_config}
# Gemini Flash image models accept plain-string values under
# ``generationConfig.imageConfig``. The legacy
# ``responseFormat.image`` block is rejected with INVALID_ARGUMENT
# by gemini-3.1-flash-lite-image (enum-based fields), so it is not
# used here.
generation_config["imageConfig"] = image_config
body: dict[str, Any] = {
"contents": [{"role": "user", "parts": parts}],
@@ -864,11 +869,13 @@ def _gemini_flash_image_config(
aspect_ratio: str | None,
image_size: str | None,
) -> dict[str, str]:
"""Build the ``responseFormat.image`` config for Gemini Flash image models.
"""Build the ``generationConfig.imageConfig`` config for Gemini Flash image models.
Capabilities are model-specific: Gemini 3.1 Flash variants support four
additional extreme ratios, while configurable image sizes are limited to
the documented Gemini 3 image model families.
Values are the documented plain strings (e.g. ``16:9``, ``1K``) that the
live v1beta API accepts under ``imageConfig``. Capabilities are
model-specific: Gemini 3.1 Flash variants support four additional extreme
ratios, while configurable image sizes are limited to the documented
Gemini 3 image model families.
"""
config: dict[str, str] = {}
if aspect_ratio and aspect_ratio in _gemini_flash_supported_aspect_ratios(model):
+273 -41
View File
@@ -17,17 +17,27 @@ from oauth_cli_kit import get_token as get_codex_token
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ToolCallRequest,
ProviderCallContext,
ProviderConversationState,
resolve_stream_idle_timeout_s,
)
from nanobot.providers.openai_responses import (
ResponsesStreamCapture,
build_responses_state,
consume_sse_with_reasoning,
convert_messages,
convert_tools,
is_compaction_compatibility_error,
is_replayable_finish_reason,
prepare_responses_input,
resolve_compact_threshold,
responses_state_context_tokens,
responses_state_items,
responses_state_matches,
)
DEFAULT_CODEX_URL = "https://chatgpt.com/backend-api/codex/responses"
DEFAULT_ORIGINATOR = "nanobot"
_COMPACTION_RETAINED_CHAR_BUDGET = 256_000
class OpenAICodexProvider(LLMProvider):
@@ -45,21 +55,39 @@ class OpenAICodexProvider(LLMProvider):
self.default_model = default_model
self.proxy = proxy or None
self._extra_body = dict(extra_body or {})
self._native_compaction_available = True
async def _call_codex(
self,
messages: list[dict[str, Any]],
tools: list[dict[str, Any]] | None,
model: str | None,
max_tokens: int,
reasoning_effort: str | None,
tool_choice: str | dict[str, Any] | None,
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
"""Shared request logic for both chat() and chat_stream()."""
model = model or self.default_model
system_prompt, input_items = convert_messages(messages)
sanitized_messages = self._sanitize_empty_content(messages)
sanitized_state = (
provider_context.conversation_state
if provider_context is not None
else None
)
if sanitized_state is not None:
sanitized_state = sanitized_state.with_pending_messages(
self._sanitize_empty_content(sanitized_state.pending_messages)
)
system_prompt, input_items, replayed = prepare_responses_input(
sanitized_messages,
state=sanitized_state,
provider=self._responses_state_provider(),
model=_strip_model_prefix(model),
)
body: dict[str, Any] = {
"model": _strip_model_prefix(model),
@@ -68,12 +96,15 @@ class OpenAICodexProvider(LLMProvider):
"instructions": system_prompt,
"input": input_items,
"text": {"verbosity": "medium"},
"include": ["reasoning.encrypted_content"],
"prompt_cache_key": _prompt_cache_key(messages[:2]),
"tool_choice": tool_choice or "auto",
"parallel_tool_calls": True,
}
body["include"] = ["reasoning.encrypted_content"]
reasoning_options = _build_reasoning_options(reasoning_effort)
if replayed and "gpt-5.6" in _strip_model_prefix(model).lower():
reasoning_options = dict(reasoning_options or {})
reasoning_options["context"] = "all_turns"
if reasoning_options:
body["reasoning"] = reasoning_options
if tools:
@@ -87,33 +118,90 @@ class OpenAICodexProvider(LLMProvider):
token = await asyncio.to_thread(get_codex_token, proxy=self.proxy)
headers = _build_headers(cast(str, token.account_id), token.access)
stage = "codex_request"
try:
content, tool_calls, finish_reason, usage, reasoning_content = await _request_codex(
DEFAULT_CODEX_URL, headers, body, verify=True,
proxy=self.proxy,
on_content_delta=on_content_delta,
on_thinking_delta=on_thinking_delta,
on_tool_call_delta=on_tool_call_delta,
)
except Exception as e:
if "CERTIFICATE_VERIFY_FAILED" not in str(e):
raise
logger.warning("SSL verification failed for Codex API; retrying with verify=False")
content, tool_calls, finish_reason, usage, reasoning_content = await _request_codex(
DEFAULT_CODEX_URL, headers, body, verify=False,
proxy=self.proxy,
on_content_delta=on_content_delta,
on_thinking_delta=on_thinking_delta,
on_tool_call_delta=on_tool_call_delta,
)
return LLMResponse(
content=content,
tool_calls=tool_calls,
finish_reason=finish_reason,
usage=usage,
reasoning_content=reasoning_content,
async def _send(
request_body: dict[str, Any],
*,
emit_deltas: bool,
) -> LLMResponse:
wire_body = _without_response_item_ids(request_body)
try:
return await _request_codex(
DEFAULT_CODEX_URL,
headers,
wire_body,
verify=True,
proxy=self.proxy,
on_content_delta=on_content_delta if emit_deltas else None,
on_thinking_delta=on_thinking_delta if emit_deltas else None,
on_tool_call_delta=on_tool_call_delta if emit_deltas else None,
)
except Exception as exc:
if "CERTIFICATE_VERIFY_FAILED" not in str(exc):
raise
logger.warning(
"SSL verification failed for Codex API; retrying with verify=False"
)
return await _request_codex(
DEFAULT_CODEX_URL,
headers,
wire_body,
verify=False,
proxy=self.proxy,
on_content_delta=on_content_delta if emit_deltas else None,
on_thinking_delta=on_thinking_delta if emit_deltas else None,
on_tool_call_delta=on_tool_call_delta if emit_deltas else None,
)
compact_threshold = resolve_compact_threshold(
(
provider_context.context_window_tokens
if provider_context is not None
else None
),
max_tokens,
)
if (
self.supports_native_compaction(model)
and replayed
and sanitized_state is not None
and compact_threshold is not None
and responses_state_context_tokens(sanitized_state) >= compact_threshold
):
stage = "codex_compaction"
compact_body = {
**body,
"input": [*input_items, {"type": "compaction_trigger"}],
}
try:
compact_result = await _send(compact_body, emit_deltas=False)
compact_items = (
responses_state_items(compact_result.provider_state)
if compact_result.provider_state is not None
else None
)
if not compact_items or compact_items[-1].get("type") not in {
"compaction",
"compaction_summary",
"context_compaction",
}:
raise RuntimeError("Codex compaction returned no compaction item")
body["input"] = [
*_retained_compaction_messages(input_items),
*compact_items,
]
except Exception as compact_error:
if is_compaction_compatibility_error(compact_error):
self._native_compaction_available = False
logger.warning(
"Codex native compaction unavailable; continuing without it "
"(type={} status={} disabled={})",
type(compact_error).__name__,
getattr(compact_error, "status_code", None),
not self._native_compaction_available,
)
stage = "codex_request"
return await _send(body, emit_deltas=True)
except Exception as e:
response = _codex_error_response(e)
exc_type = "CodexHTTPError" if isinstance(e, _CodexHTTPError) else type(e).__name__
@@ -137,8 +225,28 @@ class OpenAICodexProvider(LLMProvider):
model: str | None = None, max_tokens: int = 4096, temperature: float = 0.7,
reasoning_effort: str | None = None,
tool_choice: str | dict[str, Any] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
return await self._call_codex(messages, tools, model, reasoning_effort, tool_choice)
return await self._call_codex(
messages,
tools,
model,
max_tokens,
reasoning_effort,
tool_choice,
provider_context=provider_context,
)
async def chat_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
return await self.chat(
**kwargs,
provider_context=provider_context,
)
async def chat_stream(
self, messages: list[dict[str, Any]], tools: list[dict[str, Any]] | None = None,
@@ -148,21 +256,55 @@ class OpenAICodexProvider(LLMProvider):
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
return await self._call_codex(
messages,
tools,
model,
reasoning_effort,
tool_choice,
on_content_delta,
on_thinking_delta,
on_tool_call_delta,
messages=messages,
tools=tools,
model=model,
max_tokens=max_tokens,
reasoning_effort=reasoning_effort,
tool_choice=tool_choice,
on_content_delta=on_content_delta,
on_thinking_delta=on_thinking_delta,
on_tool_call_delta=on_tool_call_delta,
provider_context=provider_context,
)
async def chat_stream_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
return await self.chat_stream(
**kwargs,
provider_context=provider_context,
)
def get_default_model(self) -> str:
return self.default_model
@staticmethod
def _responses_state_provider() -> str:
return f"openai_codex:{DEFAULT_CODEX_URL.rstrip('/')}"
def can_resume_conversation_state(
self,
state: ProviderConversationState,
model: str | None = None,
) -> bool:
return responses_state_matches(
state,
provider=self._responses_state_provider(),
model=_strip_model_prefix(model or self.default_model),
)
def supports_native_compaction(self, model: str | None = None) -> bool:
"""Use the Codex backend's inline compaction trigger when needed."""
_ = model
return self._native_compaction_available
def _strip_model_prefix(model: str) -> str:
if model.startswith("openai-codex/") or model.startswith("openai_codex/"):
@@ -170,6 +312,58 @@ def _strip_model_prefix(model: str) -> str:
return model
def _without_response_item_ids(
request_body: dict[str, Any],
) -> dict[str, Any]:
"""Match Codex's default ``store=false`` request-item contract."""
if request_body.get("store") is True:
return request_body
raw_input = request_body.get("input")
if not isinstance(raw_input, list):
return request_body
input_items: list[object] = cast(list[object], raw_input)
sanitized_input: list[object] = []
for raw_item in input_items:
if not isinstance(raw_item, dict):
sanitized_input.append(raw_item)
continue
item = cast(dict[str, Any], raw_item)
sanitized_input.append({
key: value
for key, value in item.items()
if key != "id"
})
body = dict(request_body)
body["input"] = sanitized_input
return body
def _retained_compaction_messages(
input_items: list[dict[str, Any]],
) -> list[dict[str, Any]]:
"""Mirror Codex's bounded retention of user/developer/system messages."""
retained_reversed: list[dict[str, Any]] = []
remaining = _COMPACTION_RETAINED_CHAR_BUDGET
for item in reversed(input_items):
if item.get("type") not in {None, "message"} or item.get("role") not in {
"user",
"developer",
"system",
}:
continue
size = len(json.dumps(item, ensure_ascii=False))
if size > remaining and retained_reversed:
continue
retained_reversed.append(item)
remaining = max(0, remaining - size)
if remaining == 0:
break
retained_reversed.reverse()
return retained_reversed
def _build_reasoning_options(reasoning_effort: str | None) -> dict[str, str] | None:
"""Opt in to visible summaries without changing provider-default effort."""
if reasoning_effort and reasoning_effort.lower() == "none":
@@ -202,6 +396,7 @@ class _CodexHTTPError(RuntimeError):
error_type: str | None = None,
error_code: str | None = None,
should_retry: bool | None = None,
compaction_unsupported: bool = False,
):
super().__init__(message)
self.status_code = status_code
@@ -209,6 +404,7 @@ class _CodexHTTPError(RuntimeError):
self.error_type = error_type
self.error_code = error_code
self.should_retry = should_retry
self.compaction_unsupported = compaction_unsupported
async def _request_codex(
@@ -220,7 +416,7 @@ async def _request_codex(
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
) -> tuple[str, list[ToolCallRequest], str, dict[str, int], str | None]:
) -> LLMResponse:
idle_timeout_s = resolve_stream_idle_timeout_s()
client_kwargs: dict[str, Any] = {"timeout": idle_timeout_s, "verify": verify}
if proxy:
@@ -233,6 +429,17 @@ async def _request_codex(
raw = text.decode("utf-8", "ignore")
retry_after = LLMProvider._extract_retry_after_from_headers(response.headers)
error_type, error_code = LLMProvider._extract_error_type_code(raw)
compaction_unsupported = (
response.status_code in {400, 404, 422}
and any(
marker in raw.lower()
for marker in (
"context_management",
"compact_threshold",
"compaction_trigger",
)
)
)
raise _CodexHTTPError(
_friendly_error(response.status_code, raw),
status_code=response.status_code,
@@ -240,13 +447,38 @@ async def _request_codex(
error_type=error_type,
error_code=error_code,
should_retry=_should_retry_status(response.status_code, error_type, error_code, raw),
compaction_unsupported=compaction_unsupported,
)
return await consume_sse_with_reasoning(
capture = ResponsesStreamCapture()
(
content,
tool_calls,
finish_reason,
usage,
reasoning_content,
) = await consume_sse_with_reasoning(
response,
on_content_delta=on_content_delta,
on_tool_call_delta=on_tool_call_delta,
on_reasoning_delta=on_thinking_delta,
capture=capture,
)
result = LLMResponse(
content=content,
tool_calls=tool_calls,
finish_reason=finish_reason,
usage=usage,
reasoning_content=reasoning_content,
)
if capture.completed and is_replayable_finish_reason(finish_reason):
result.provider_state = build_responses_state(
provider=f"openai_codex:{url.rstrip('/')}",
model=str(body.get("model") or ""),
input_items=cast(list[dict[str, Any]], body.get("input") or []),
output_items=capture.output_items,
usage=usage,
)
return result
def _prompt_cache_key(messages: list[dict[str, Any]]) -> str:
+171 -16
View File
@@ -26,16 +26,24 @@ from pydantic.alias_generators import to_snake
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
ToolCallRequest,
parse_tool_arguments,
resolve_stream_idle_timeout_s,
tool_arguments_json_for_replay,
)
from nanobot.providers.openai_responses import (
ResponsesStreamCapture,
build_responses_state,
consume_sdk_stream,
convert_messages,
convert_tools,
is_compaction_compatibility_error,
is_replayable_finish_reason,
parse_response_output,
prepare_responses_input,
resolve_compact_threshold,
responses_state_matches,
)
if TYPE_CHECKING:
@@ -443,6 +451,8 @@ class OpenAICompatProvider(LLMProvider):
registry lookups needed.
"""
_native_compaction_available = True
def __init__(
self,
api_key: str | None = None,
@@ -463,6 +473,7 @@ class OpenAICompatProvider(LLMProvider):
self._api_type = api_type if spec and spec.name == "openai" else "auto"
self._extra_query = extra_query or {}
self._proxy = proxy or None
self._native_compaction_available = True
if api_key and spec and spec.env_key:
self._setup_env(api_key, api_base)
@@ -947,22 +958,34 @@ class OpenAICompatProvider(LLMProvider):
model: str | None,
reasoning_effort: str | None,
) -> bool:
"""Use Responses API only for direct OpenAI requests that benefit from it."""
"""Choose Responses for providers/models that explicitly support it."""
if self._api_type == "chat_completions":
return False
if self._spec and self._spec.name not in ("openai", "github_copilot"):
spec_name = self._spec.name if self._spec is not None else None
model_name = self._request_model_name(model or self.default_model).lower()
supported_models = {
supported.lower()
for supported in getattr(self._spec, "responses_models", ())
}
model_responses = any(
model_name == supported or model_name.endswith(f"/{supported}")
for supported in supported_models
)
provider_responses = spec_name in ("openai", "github_copilot")
if not provider_responses and not model_responses:
return False
if self._api_type == "responses":
# Explicit configuration means Responses is mandatory; do not
# consult the circuit breaker or fall back to Chat Completions.
return True
if self._spec is None or self._spec.name != "github_copilot":
if provider_responses and (self._spec is None or self._spec.name != "github_copilot"):
if not _is_direct_openai_base(self._effective_base):
return False
model_name = (model or self.default_model).lower()
wants = False
if reasoning_effort and reasoning_effort.lower() != "none":
if model_responses:
wants = True
elif reasoning_effort and reasoning_effort.lower() != "none":
wants = True
elif any(token in model_name for token in ("gpt-5", "o1", "o3", "o4")):
wants = True
@@ -971,6 +994,37 @@ class OpenAICompatProvider(LLMProvider):
return self._responses_circuit_allows_probe(model, reasoning_effort)
def _responses_state_provider(self) -> str:
spec_name = self._spec.name if self._spec is not None else "custom"
effective_base = self._effective_base or "https://api.openai.com/v1"
return f"openai_compat:{spec_name}:{effective_base.rstrip('/')}"
def _responses_state_model(self, model: str | None) -> str:
return self._request_model_name(model or self.default_model)
def can_resume_conversation_state(
self,
state: ProviderConversationState,
model: str | None = None,
) -> bool:
return responses_state_matches(
state,
provider=self._responses_state_provider(),
model=self._responses_state_model(model),
)
def supports_native_compaction(self, model: str | None = None) -> bool:
"""Enable server compaction only on direct OpenAI Responses endpoints."""
_ = model
if (
not self._native_compaction_available
or self._api_type == "chat_completions"
):
return False
if self._spec is not None and self._spec.name != "openai":
return False
return _is_direct_openai_base(self._effective_base)
def _responses_circuit_allows_probe(
self,
model: str | None,
@@ -1040,12 +1094,31 @@ class OpenAICompatProvider(LLMProvider):
temperature: float,
reasoning_effort: str | None,
tool_choice: str | dict[str, Any] | None,
provider_context: ProviderCallContext | None = None,
) -> dict[str, Any]:
"""Build a Responses API body for direct OpenAI requests."""
model_name = model or self.default_model
model_name = self._request_model_name(model_name)
sanitized_messages = self._sanitize_messages(self._sanitize_empty_content(messages))
instructions, input_items = convert_messages(sanitized_messages)
sanitized_state = (
provider_context.conversation_state
if provider_context is not None
else None
)
if sanitized_state is not None:
sanitized_state = sanitized_state.with_pending_messages(
self._sanitize_messages(
self._sanitize_empty_content(sanitized_state.pending_messages)
)
)
preserve_reasoning = bool(self._spec and self._spec.name == "deepseek")
instructions, input_items, replayed = prepare_responses_input(
sanitized_messages,
state=sanitized_state,
provider=self._responses_state_provider(),
model=model_name,
preserve_reasoning=preserve_reasoning,
)
body: dict[str, Any] = {
"model": model_name,
@@ -1055,13 +1128,29 @@ class OpenAICompatProvider(LLMProvider):
"store": False,
"stream": False,
}
compact_threshold = resolve_compact_threshold(
(
provider_context.context_window_tokens
if provider_context is not None
else None
),
max_tokens,
)
if self.supports_native_compaction(model_name) and compact_threshold is not None:
body["context_management"] = [{
"type": "compaction",
"compact_threshold": compact_threshold,
}]
if self._supports_temperature(model_name, reasoning_effort):
body["temperature"] = temperature
if not self._supports_temperature(model_name, reasoning_effort) and not preserve_reasoning:
body["include"] = ["reasoning.encrypted_content"]
if reasoning_effort and reasoning_effort.lower() != "none":
body["reasoning"] = {"effort": reasoning_effort}
body["include"] = ["reasoning.encrypted_content"]
if replayed and "gpt-5.6" in model_name.lower():
body.setdefault("reasoning", {})["context"] = "all_turns"
if tools:
body["tools"] = convert_tools(tools)
@@ -1073,6 +1162,29 @@ class OpenAICompatProvider(LLMProvider):
return body
async def _create_response_with_compaction_fallback(
self,
client: Any,
body: dict[str, Any],
) -> Any:
"""Retry Responses once without server compaction on compatibility errors."""
try:
return await client.responses.create(**body)
except Exception as exc:
if (
"context_management" not in body
or not is_compaction_compatibility_error(exc)
):
raise
self._native_compaction_available = False
body.pop("context_management", None)
logger.warning(
"Responses server compaction unsupported; disabled for this provider instance "
"(status={})",
getattr(exc, "status_code", None),
)
return await client.responses.create(**body)
# ------------------------------------------------------------------
# Response parsing
# ------------------------------------------------------------------
@@ -1599,6 +1711,28 @@ class OpenAICompatProvider(LLMProvider):
# Public API
# ------------------------------------------------------------------
async def chat_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
return await self.chat(
**kwargs,
provider_context=provider_context,
)
async def chat_stream_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs: Any,
) -> LLMResponse:
return await self.chat_stream(
**kwargs,
provider_context=provider_context,
)
async def chat(
self,
messages: list[dict[str, Any]],
@@ -1608,6 +1742,7 @@ class OpenAICompatProvider(LLMProvider):
temperature: float = 0.7,
reasoning_effort: str | None = None,
tool_choice: str | dict[str, Any] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
client = await self._ensure_client()
try:
@@ -1616,12 +1751,18 @@ class OpenAICompatProvider(LLMProvider):
body = self._build_responses_body(
messages, tools, model, max_tokens, temperature,
reasoning_effort, tool_choice,
provider_context,
)
responses_raw = cast(
Any,
await client.responses.create(**body),
responses_raw = await self._create_response_with_compaction_fallback(
client,
body,
)
result = parse_response_output(
responses_raw,
state_provider=self._responses_state_provider(),
state_model=str(body["model"]),
state_input_items=cast(list[dict[str, Any]], body["input"]),
)
result = parse_response_output(responses_raw)
self._record_responses_success(model, reasoning_effort)
return result
except Exception as responses_error:
@@ -1660,6 +1801,7 @@ class OpenAICompatProvider(LLMProvider):
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_thinking_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
provider_context: ProviderCallContext | None = None,
) -> LLMResponse:
client = await self._ensure_client()
idle_timeout_s = resolve_stream_idle_timeout_s()
@@ -1669,11 +1811,12 @@ class OpenAICompatProvider(LLMProvider):
body = self._build_responses_body(
messages, tools, model, max_tokens, temperature,
reasoning_effort, tool_choice,
provider_context,
)
body["stream"] = True
responses_stream = cast(
Any,
await client.responses.create(**body),
responses_stream = await self._create_response_with_compaction_fallback(
client,
body,
)
async def _timed_stream() -> AsyncIterator[Any]:
@@ -1687,6 +1830,7 @@ class OpenAICompatProvider(LLMProvider):
except StopAsyncIteration:
break
capture = ResponsesStreamCapture()
(
content,
tool_calls,
@@ -1697,15 +1841,26 @@ class OpenAICompatProvider(LLMProvider):
_timed_stream(),
on_content_delta,
on_tool_call_delta=on_tool_call_delta,
on_reasoning_delta=on_thinking_delta,
capture=capture,
)
self._record_responses_success(model, reasoning_effort)
return LLMResponse(
result = LLMResponse(
content=content or None,
tool_calls=tool_calls,
finish_reason=finish_reason,
usage=usage,
reasoning_content=reasoning_content,
)
if capture.completed and is_replayable_finish_reason(finish_reason):
result.provider_state = build_responses_state(
provider=self._responses_state_provider(),
model=str(body["model"]),
input_items=cast(list[dict[str, Any]], body["input"]),
output_items=capture.output_items,
usage=usage,
)
return result
except Exception as responses_error:
if self._spec and self._spec.name == "github_copilot":
# Copilot gateway exposes GPT-5/o-series only via /responses;
+21 -1
View File
@@ -1,4 +1,4 @@
"""Shared helpers for OpenAI Responses API providers (Codex, Azure OpenAI)."""
"""Shared helpers for provider backends that implement the OpenAI Responses protocol."""
from nanobot.providers.openai_responses.converters import (
convert_messages,
@@ -8,13 +8,24 @@ from nanobot.providers.openai_responses.converters import (
)
from nanobot.providers.openai_responses.parsing import (
FINISH_REASON_MAP,
ResponsesStreamCapture,
consume_sdk_stream,
consume_sse,
consume_sse_with_reasoning,
is_replayable_finish_reason,
iter_sse,
map_finish_reason,
parse_response_output,
)
from nanobot.providers.openai_responses.state import (
build_responses_state,
is_compaction_compatibility_error,
prepare_responses_input,
resolve_compact_threshold,
responses_state_context_tokens,
responses_state_items,
responses_state_matches,
)
__all__ = [
"convert_messages",
@@ -25,7 +36,16 @@ __all__ = [
"consume_sse",
"consume_sse_with_reasoning",
"consume_sdk_stream",
"ResponsesStreamCapture",
"is_replayable_finish_reason",
"map_finish_reason",
"parse_response_output",
"build_responses_state",
"is_compaction_compatibility_error",
"prepare_responses_input",
"resolve_compact_threshold",
"responses_state_context_tokens",
"responses_state_items",
"responses_state_matches",
"FINISH_REASON_MAP",
]
@@ -12,7 +12,11 @@ def _as_json_object(value: object) -> dict[str, Any] | None:
return cast(dict[str, Any], value) if isinstance(value, dict) else None
def convert_messages(messages: list[dict[str, Any]]) -> tuple[str, list[dict[str, Any]]]:
def convert_messages(
messages: list[dict[str, Any]],
*,
preserve_reasoning: bool = False,
) -> tuple[str, list[dict[str, Any]]]:
"""Convert Chat Completions messages to Responses API input items.
Returns ``(system_prompt, input_items)`` where *system_prompt* is extracted
@@ -36,6 +40,13 @@ def convert_messages(messages: list[dict[str, Any]]) -> tuple[str, list[dict[str
continue
if role == "assistant":
if preserve_reasoning:
reasoning = msg.get("reasoning_content")
if isinstance(reasoning, str) and reasoning:
input_items.append({
"type": "reasoning",
"content": [{"type": "output_text", "text": reasoning}],
})
if isinstance(content, str) and content:
message_id = _unique_item_id(f"msg_{idx}", used_item_ids)
input_items.append({
+281 -23
View File
@@ -4,12 +4,14 @@ from __future__ import annotations
import json
from collections.abc import Awaitable, Callable
from dataclasses import dataclass, field
from typing import Any, AsyncGenerator, cast
import httpx
from loguru import logger
from nanobot.providers.base import LLMResponse, ToolCallRequest, parse_tool_arguments
from nanobot.providers.openai_responses.state import build_responses_state
FINISH_REASON_MAP = {
"completed": "stop",
@@ -17,6 +19,42 @@ FINISH_REASON_MAP = {
"failed": "error",
"cancelled": "error",
}
REPLAYABLE_FINISH_REASONS = frozenset({"stop", "tool_calls", "function_call"})
@dataclass(slots=True)
class ResponsesStreamCapture:
"""Losslessly capture terminal output items without changing stream results."""
completed: bool = False
response: dict[str, Any] | None = field(default=None, repr=False)
_items_by_index: dict[int, dict[str, Any]] = field(default_factory=dict, repr=False)
def record_output_item(self, index: object, item: object) -> None:
item_object = _response_object(item)
if item_object is None:
return
output_index = (
index
if isinstance(index, int) and not isinstance(index, bool)
else len(self._items_by_index)
)
self._items_by_index[output_index] = item_object
def record_completed(self, response: object) -> None:
response_object = _response_object(response)
if response_object is None:
return
self.completed = True
self.response = response_object
@property
def output_items(self) -> list[dict[str, Any]]:
if self.response is not None:
output = _response_object_list(self.response.get("output"))
if output:
return output
return [self._items_by_index[index] for index in sorted(self._items_by_index)]
def _as_json_object(value: object) -> dict[str, Any] | None:
@@ -31,7 +69,9 @@ def _response_object(value: object) -> dict[str, Any] | None:
return object_value
dump = getattr(value, "model_dump", None)
if callable(dump):
return _as_json_object(dump())
dumped = _as_json_object(dump())
if dumped is not None:
return dumped
try:
return _as_json_object(vars(value))
except TypeError:
@@ -54,6 +94,27 @@ def map_finish_reason(status: str | None) -> str:
return FINISH_REASON_MAP.get(status or "completed", "stop")
def is_replayable_finish_reason(finish_reason: str) -> bool:
"""Return whether a response can safely advance opaque conversation state."""
return finish_reason in REPLAYABLE_FINISH_REASONS
def _response_finish_reason(
response: object,
*,
fallback_status: str | None = None,
) -> str:
"""Map terminal response details without treating content filtering as truncation."""
response_object = _response_object(response) or {}
status = response_object.get("status")
terminal_status = status if isinstance(status, str) else fallback_status
if terminal_status == "incomplete":
details = _response_object(response_object.get("incomplete_details"))
if details is not None and details.get("reason") == "content_filter":
return "content_filter"
return map_finish_reason(terminal_status)
def _usage_from_response_obj(response: object) -> dict[str, int]:
response_object = _response_object(response)
usage_raw: object = (
@@ -99,6 +160,47 @@ def _tool_arguments_source(*values: Any) -> Any:
return "{}"
def _refusal_event_key(
item_id: object,
content_index: object,
) -> tuple[str | None, int | None]:
"""Identify one streamed refusal content part across delta/done events."""
return (
item_id if isinstance(item_id, str) else None,
(
content_index
if isinstance(content_index, int) and not isinstance(content_index, bool)
else None
),
)
def _remaining_refusal_text(streamed_text: str, refusal_text: str) -> str:
"""Return only text not already surfaced by refusal deltas."""
if not streamed_text:
return refusal_text
if refusal_text.startswith(streamed_text):
return refusal_text[len(streamed_text):]
return ""
def _extract_refusal_text_from_output(output: object) -> tuple[bool, str]:
"""Extract refusal content from terminal Responses output items."""
refusal_seen = False
parts: list[str] = []
for item in _response_object_list(output):
if item.get("type") != "message":
continue
for block in _response_object_list(item.get("content")):
if block.get("type") != "refusal":
continue
refusal_seen = True
refusal_text = block.get("refusal")
if isinstance(refusal_text, str):
parts.append(refusal_text)
return refusal_seen, "".join(parts)
async def iter_sse(response: httpx.Response) -> AsyncGenerator[dict[str, Any], None]:
"""Yield parsed JSON events from a Responses API SSE stream."""
buffer: list[str] = []
@@ -153,6 +255,7 @@ async def consume_sse_with_reasoning(
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
on_reasoning_delta: Callable[[str], Awaitable[None]] | None = None,
on_response_event: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
capture: ResponsesStreamCapture | None = None,
) -> tuple[str, list[ToolCallRequest], str, dict[str, int], str | None]:
"""Consume a Responses API SSE stream, including visible reasoning summaries."""
content = ""
@@ -163,6 +266,9 @@ async def consume_sse_with_reasoning(
usage: dict[str, int] = {}
reasoning_content: str | None = None
streamed_reasoning = False
refusal_seen = False
refusal_deltas: dict[tuple[str | None, int | None], str] = {}
emitted_refusal_text = ""
async for event in iter_sse(response):
if on_response_event:
@@ -191,6 +297,33 @@ async def consume_sse_with_reasoning(
content += delta_text
if on_content_delta and delta_text:
await on_content_delta(delta_text)
elif event_type == "response.refusal.delta":
refusal_seen = True
delta_text = event.get("delta")
if isinstance(delta_text, str) and delta_text:
key = _refusal_event_key(
event.get("item_id"),
event.get("content_index"),
)
refusal_deltas[key] = refusal_deltas.get(key, "") + delta_text
content += delta_text
emitted_refusal_text += delta_text
if on_content_delta:
await on_content_delta(delta_text)
elif event_type == "response.refusal.done":
refusal_seen = True
refusal_text = event.get("refusal")
key = _refusal_event_key(
event.get("item_id"),
event.get("content_index"),
)
streamed_text = refusal_deltas.pop(key, "")
if isinstance(refusal_text, str) and refusal_text:
remaining_text = _remaining_refusal_text(streamed_text, refusal_text)
content += remaining_text
emitted_refusal_text += remaining_text
if on_content_delta and remaining_text:
await on_content_delta(remaining_text)
elif event_type == "response.reasoning_summary_text.delta":
delta_text = event.get("delta") or ""
if delta_text:
@@ -239,6 +372,8 @@ async def consume_sse_with_reasoning(
})
elif event_type == "response.output_item.done":
item = _as_json_object(event.get("item")) or {}
if capture is not None:
capture.record_output_item(event.get("output_index"), item)
if item.get("type") == "function_call":
call_id = item.get("call_id")
if not call_id:
@@ -269,11 +404,28 @@ async def consume_sse_with_reasoning(
reasoning_content = summary
if on_reasoning_delta:
await on_reasoning_delta(summary)
elif event_type == "response.completed":
elif event_type in {"response.completed", "response.incomplete"}:
response_obj = _response_object(event.get("response")) or {}
status = response_obj.get("status")
finish_reason = map_finish_reason(status)
if capture is not None:
capture.record_completed(response_obj)
finish_reason = _response_finish_reason(
response_obj,
fallback_status=event_type.removeprefix("response."),
)
usage = _usage_from_response_obj(response_obj) or usage
terminal_refusal, terminal_refusal_text = _extract_refusal_text_from_output(
response_obj.get("output")
)
if terminal_refusal:
refusal_seen = True
remaining_text = _remaining_refusal_text(
emitted_refusal_text,
terminal_refusal_text,
)
content += remaining_text
emitted_refusal_text += remaining_text
if on_content_delta and remaining_text:
await on_content_delta(remaining_text)
if not reasoning_content:
summary = _extract_reasoning_summary_from_output(response_obj.get("output"))
if summary:
@@ -284,6 +436,8 @@ async def consume_sse_with_reasoning(
detail = event.get("error") or event.get("message") or event
raise RuntimeError(f"Response failed: {str(detail)[:500]}")
if refusal_seen:
finish_reason = "refusal"
return content, tool_calls, finish_reason, usage, reasoning_content
@@ -292,6 +446,14 @@ def _extract_reasoning_summary_from_output(output: object) -> str | None:
for item in _response_object_list(output):
if item.get("type") != "reasoning":
continue
content = item.get("content")
if isinstance(content, str) and content:
parts.append(content)
elif isinstance(content, list):
for block in _response_object_list(cast(list[object], content)):
text = block.get("text")
if isinstance(text, str) and text:
parts.append(text)
for summary in _response_object_list(item.get("summary")):
if summary.get("type") == "summary_text" and summary.get("text"):
text = summary.get("text")
@@ -300,7 +462,13 @@ def _extract_reasoning_summary_from_output(output: object) -> str | None:
return "".join(parts) or None
def parse_response_output(response: object) -> LLMResponse:
def parse_response_output(
response: object,
*,
state_provider: str | None = None,
state_model: str | None = None,
state_input_items: list[dict[str, Any]] | None = None,
) -> LLMResponse:
"""Parse an SDK ``Response`` object into an ``LLMResponse``."""
response_object = _response_object(response) or {}
@@ -308,21 +476,26 @@ def parse_response_output(response: object) -> LLMResponse:
content_parts: list[str] = []
tool_calls: list[ToolCallRequest] = []
reasoning_content: str | None = None
refusal_seen = False
for item in output:
item_type = item.get("type")
if item_type == "message":
for block in _response_object_list(item.get("content")):
if block.get("type") == "output_text":
block_type = block.get("type")
if block_type == "output_text":
text = block.get("text")
if isinstance(text, str):
content_parts.append(text)
elif block_type == "refusal":
refusal_seen = True
refusal = block.get("refusal")
if isinstance(refusal, str):
content_parts.append(refusal)
elif item_type == "reasoning":
for s in _response_object_list(item.get("summary")):
if s.get("type") == "summary_text" and s.get("text"):
text = s.get("text")
if isinstance(text, str):
reasoning_content = (reasoning_content or "") + text
text = _extract_reasoning_summary_from_output([item])
if text:
reasoning_content = (reasoning_content or "") + text
elif item_type == "function_call":
call_id = item.get("call_id") or ""
item_id = item.get("id") or "fc_0"
@@ -337,21 +510,38 @@ def parse_response_output(response: object) -> LLMResponse:
usage = _usage_from_response_obj(response_object)
status = response_object.get("status")
finish_reason = map_finish_reason(status if isinstance(status, str) else None)
finish_reason = "refusal" if refusal_seen else _response_finish_reason(response_object)
return LLMResponse(
result = LLMResponse(
content="".join(content_parts) or None,
tool_calls=tool_calls,
finish_reason=finish_reason,
usage=usage,
reasoning_content=reasoning_content if isinstance(reasoning_content, str) else None,
)
if (
state_provider is not None
and state_model is not None
and state_input_items is not None
and (status is None or status == "completed")
and is_replayable_finish_reason(finish_reason)
):
result.provider_state = build_responses_state(
provider=state_provider,
model=state_model,
input_items=state_input_items,
output_items=output,
usage=usage,
)
return result
async def consume_sdk_stream(
stream: Any,
on_content_delta: Callable[[str], Awaitable[None]] | None = None,
on_tool_call_delta: Callable[[dict[str, Any]], Awaitable[None]] | None = None,
on_reasoning_delta: Callable[[str], Awaitable[None]] | None = None,
capture: ResponsesStreamCapture | None = None,
) -> tuple[str, list[ToolCallRequest], str, dict[str, int], str | None]:
"""Consume an SDK async stream from ``client.responses.create(stream=True)``."""
content = ""
@@ -361,6 +551,10 @@ async def consume_sdk_stream(
finish_reason = "stop"
usage: dict[str, int] = {}
reasoning_content: str | None = None
streamed_reasoning = False
refusal_seen = False
refusal_deltas: dict[tuple[str | None, int | None], str] = {}
emitted_refusal_text = ""
async for raw_event in stream:
event: Any = raw_event
@@ -388,6 +582,46 @@ async def consume_sdk_stream(
content += delta_text
if on_content_delta and delta_text:
await on_content_delta(delta_text)
elif event_type == "response.reasoning_text.delta":
delta_text = getattr(event, "delta", "") or ""
if delta_text:
reasoning_content = (reasoning_content or "") + delta_text
streamed_reasoning = True
if on_reasoning_delta:
await on_reasoning_delta(delta_text)
elif event_type == "response.reasoning_text.done":
text = getattr(event, "text", "") or ""
if text and not streamed_reasoning and not reasoning_content:
reasoning_content = text
if on_reasoning_delta:
await on_reasoning_delta(text)
elif event_type == "response.refusal.delta":
refusal_seen = True
delta_text = getattr(event, "delta", None)
if isinstance(delta_text, str) and delta_text:
key = _refusal_event_key(
getattr(event, "item_id", None),
getattr(event, "content_index", None),
)
refusal_deltas[key] = refusal_deltas.get(key, "") + delta_text
content += delta_text
emitted_refusal_text += delta_text
if on_content_delta:
await on_content_delta(delta_text)
elif event_type == "response.refusal.done":
refusal_seen = True
refusal_text = getattr(event, "refusal", None)
key = _refusal_event_key(
getattr(event, "item_id", None),
getattr(event, "content_index", None),
)
streamed_text = refusal_deltas.pop(key, "")
if isinstance(refusal_text, str) and refusal_text:
remaining_text = _remaining_refusal_text(streamed_text, refusal_text)
content += remaining_text
emitted_refusal_text += remaining_text
if on_content_delta and remaining_text:
await on_content_delta(remaining_text)
elif event_type == "response.function_call_arguments.delta":
call_id = getattr(event, "call_id", None)
if call_id and call_id in tool_call_buffers:
@@ -416,6 +650,8 @@ async def consume_sdk_stream(
})
elif event_type == "response.output_item.done":
item = getattr(event, "item", None)
if capture is not None:
capture.record_output_item(getattr(event, "output_index", None), item)
if item and getattr(item, "type", None) == "function_call":
call_id = getattr(item, "call_id", None)
if not call_id:
@@ -443,10 +679,31 @@ async def consume_sdk_stream(
arguments=args,
)
)
elif event_type == "response.completed":
elif event_type in {"response.completed", "response.incomplete"}:
resp = getattr(event, "response", None)
status = getattr(resp, "status", None) if resp else None
finish_reason = map_finish_reason(status)
response_obj = _response_object(resp) or {}
if capture is not None:
capture.record_completed(resp)
finish_reason = _response_finish_reason(
resp,
fallback_status=event_type.removeprefix("response."),
)
terminal_output = response_obj.get("output")
if terminal_output is None:
terminal_output = getattr(resp, "output", None)
terminal_refusal, terminal_refusal_text = _extract_refusal_text_from_output(
terminal_output
)
if terminal_refusal:
refusal_seen = True
remaining_text = _remaining_refusal_text(
emitted_refusal_text,
terminal_refusal_text,
)
content += remaining_text
emitted_refusal_text += remaining_text
if on_content_delta and remaining_text:
await on_content_delta(remaining_text)
if resp:
usage_obj = getattr(resp, "usage", None)
if usage_obj:
@@ -455,15 +712,16 @@ async def consume_sdk_stream(
"completion_tokens": int(getattr(usage_obj, "output_tokens", 0) or 0),
"total_tokens": int(getattr(usage_obj, "total_tokens", 0) or 0),
}
for out_item in cast(list[Any], getattr(resp, "output", None) or []):
if getattr(out_item, "type", None) == "reasoning":
for s in cast(list[Any], getattr(out_item, "summary", None) or []):
if getattr(s, "type", None) == "summary_text":
text = getattr(s, "text", None)
if text:
reasoning_content = (reasoning_content or "") + text
if not reasoning_content:
reasoning_content = _extract_reasoning_summary_from_output(
getattr(resp, "output", None)
)
if reasoning_content and on_reasoning_delta:
await on_reasoning_delta(reasoning_content)
elif event_type in {"error", "response.failed"}:
detail = getattr(event, "error", None) or getattr(event, "message", None) or event
raise RuntimeError(f"Response failed: {str(detail)[:500]}")
if refusal_seen:
finish_reason = "refusal"
return content, tool_calls, finish_reason, usage, reasoning_content
+204
View File
@@ -0,0 +1,204 @@
"""Opaque conversation state for Responses API item replay."""
from __future__ import annotations
from copy import deepcopy
from typing import Any, cast
from loguru import logger
from nanobot.providers.base import ProviderConversationState
from nanobot.providers.openai_responses.converters import convert_messages
RESPONSES_STATE_KIND = "openai_responses"
RESPONSES_STATE_VERSION = 1
_ITEMS_KEY = "items"
_CONTEXT_TOKENS_KEY = "context_tokens"
_COMPACTION_ITEM_TYPES = frozenset({
"compaction",
"compaction_summary",
"context_compaction",
})
def responses_state_matches(
state: ProviderConversationState,
*,
provider: str,
model: str,
) -> bool:
"""Return whether *state* belongs to this exact Responses endpoint/model."""
return (
state.kind == RESPONSES_STATE_KIND
and state.version == RESPONSES_STATE_VERSION
and state.provider == provider
and state.model == model
and _state_items(state) is not None
)
def prepare_responses_input(
messages: list[dict[str, Any]],
*,
state: ProviderConversationState | None,
provider: str,
model: str,
preserve_reasoning: bool = False,
) -> tuple[str, list[dict[str, Any]], bool]:
"""Build a request from exact prior items plus only newly appended messages.
The full Chat transcript remains the source for the current instructions.
When no compatible state exists, it is converted normally as a safe
fallback.
"""
instructions, fallback_items = convert_messages(
messages,
preserve_reasoning=preserve_reasoning,
)
if state is None or not responses_state_matches(
state,
provider=provider,
model=model,
):
return instructions, fallback_items, False
prior_items = _state_items(state)
if prior_items is None:
return instructions, fallback_items, False
_, delta_items = convert_messages(
state.pending_messages,
preserve_reasoning=preserve_reasoning,
)
logger.debug(
"Replaying Responses state: prior_items={} pending_messages={}",
len(prior_items),
len(state.pending_messages),
)
return instructions, [*deepcopy(prior_items), *delta_items], True
def build_responses_state(
*,
provider: str,
model: str,
input_items: list[dict[str, Any]],
output_items: list[dict[str, Any]],
usage: dict[str, int] | None = None,
) -> ProviderConversationState:
"""Create the canonical next state from request input and every output item."""
unpruned_items = [*input_items, *output_items]
items = _prune_before_latest_output_compaction(input_items, output_items)
if len(items) < len(unpruned_items):
logger.info(
"Installed Responses compaction: dropped_items={} retained_items={}",
len(unpruned_items) - len(items),
len(items),
)
payload: dict[str, Any] = {_ITEMS_KEY: deepcopy(items)}
context_tokens = _context_tokens_from_usage(usage)
if context_tokens > 0:
payload[_CONTEXT_TOKENS_KEY] = context_tokens
return ProviderConversationState(
kind=RESPONSES_STATE_KIND,
provider=provider,
model=model,
version=RESPONSES_STATE_VERSION,
payload=payload,
)
def responses_state_items(
state: ProviderConversationState,
) -> list[dict[str, Any]] | None:
"""Return an isolated copy of canonical input items for tests/consumers."""
items = _state_items(state)
return deepcopy(items) if items is not None else None
def responses_state_context_tokens(state: ProviderConversationState) -> int:
"""Return the last server-reported active context size."""
value = state.payload.get(_CONTEXT_TOKENS_KEY)
if isinstance(value, bool) or not isinstance(value, int):
return 0
return max(0, value)
def resolve_compact_threshold(
context_window_tokens: int | None,
max_output_tokens: int,
) -> int | None:
"""Derive Codex-compatible 90% compaction headroom for a model window."""
if context_window_tokens is None or context_window_tokens <= 0:
return None
ninety_percent = max(1, context_window_tokens * 9 // 10)
output_headroom = max(1, context_window_tokens - max(1, max_output_tokens))
return min(ninety_percent, output_headroom)
def is_compaction_compatibility_error(exc: Exception) -> bool:
"""Recognize endpoints that reject native Responses compaction fields."""
if getattr(exc, "compaction_unsupported", False) is True:
return True
response = getattr(exc, "response", None)
status_code = getattr(exc, "status_code", None)
if status_code is None and response is not None:
status_code = getattr(response, "status_code", None)
body = (
getattr(exc, "body", None)
or getattr(exc, "doc", None)
or getattr(response, "text", None)
or str(exc)
)
text = str(body).lower()
has_compaction_marker = any(
marker in text
for marker in ("context_management", "compact_threshold", "compaction_trigger")
)
if not has_compaction_marker:
return False
return isinstance(exc, TypeError) or status_code in {400, 404, 422}
def _prune_before_latest_output_compaction(
input_items: list[dict[str, Any]],
output_items: list[dict[str, Any]],
) -> list[dict[str, Any]]:
"""Drop old input only when this response emits a new compaction item.
A canonical compacted input may intentionally retain messages before its
compaction item. Those messages must survive ordinary subsequent responses.
"""
latest = None
for index, item in enumerate(output_items):
if item.get("type") in _COMPACTION_ITEM_TYPES:
latest = index
if latest is None:
return [*input_items, *output_items]
return output_items[latest:]
def _context_tokens_from_usage(usage: dict[str, int] | None) -> int:
if not usage:
return 0
prompt_tokens = usage.get("prompt_tokens", 0)
completion_tokens = usage.get("completion_tokens", 0)
total_tokens = usage.get("total_tokens", 0)
values = (prompt_tokens, completion_tokens, total_tokens)
if any(isinstance(value, bool) for value in values):
return 0
return max(0, total_tokens or prompt_tokens + completion_tokens)
def _state_items(
state: ProviderConversationState,
) -> list[dict[str, Any]] | None:
raw_items = state.payload.get(_ITEMS_KEY)
if not isinstance(raw_items, list):
return None
items: list[dict[str, Any]] = []
for raw in cast(list[object], raw_items):
if not isinstance(raw, dict):
return None
items.append(cast(dict[str, Any], raw))
return items
+18
View File
@@ -111,6 +111,11 @@ class ProviderSpec:
# Substring match against the wire model name (lowercased).
implicit_reasoning_models: tuple[str, ...] = ()
# Models that expose the OpenAI Responses wire format. This is model-level
# because providers may add Responses support incrementally (DeepSeek V4
# Flash is supported before V4 Pro).
responses_models: tuple[str, ...] = ()
# When the model returns content as a list of {"type":"thinking",...} +
# {"type":"text",...} blocks, extract the thinking text into
# reasoning_content. Mistral's Magistral / reasoning-enabled responses use
@@ -191,6 +196,18 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
supports_prompt_caching=True,
gateway_reasoning_style="reasoning_effort",
),
# Eden AI: OpenAI-compatible gateway. Models use the "provider/model"
# naming scheme (e.g. "anthropic/claude-sonnet-4-5"); the full id is sent upstream.
ProviderSpec(
name="edenai",
keywords=("edenai",),
env_key="EDENAI_API_KEY",
display_name="Eden AI",
backend="openai_compat",
is_gateway=True,
detect_by_base_keyword="edenai",
default_api_base="https://api.edenai.run/v3",
),
# OpenCode Zen: OpenAI-compatible chat-completions gateway for coding models.
# models.dev/OpenCode use provider id "opencode" and model ids like
# "opencode/<model>"; send the bare model upstream.
@@ -461,6 +478,7 @@ PROVIDERS: tuple[ProviderSpec, ...] = (
backend="openai_compat",
default_api_base="https://api.deepseek.com",
thinking_style="thinking_type",
responses_models=("deepseek-v4-flash",),
),
# Gemini: Google's OpenAI-compatible endpoint
ProviderSpec(
+1 -58
View File
@@ -12,7 +12,6 @@ from urllib.parse import urlparse
from urllib.request import getproxies, proxy_bypass
import httpx
import httpx2
_BLOCKED_NETWORKS = [
ipaddress.ip_network("0.0.0.0/8"),
@@ -30,7 +29,6 @@ _BLOCKED_NETWORKS = [
_URL_RE = re.compile(r"https?://[^\s\"'`;|<>]+", re.IGNORECASE)
_allowed_networks: list[ipaddress.IPv4Network | ipaddress.IPv6Network] = []
_DNS_PIN_RESOLVER_LOCK = asyncio.Lock()
def is_loopback_host(host: str) -> bool:
@@ -197,30 +195,6 @@ def httpx_env_proxy_mounts() -> dict[str, httpx.AsyncBaseTransport | None]:
return mounts
def httpx2_env_proxy_mounts() -> dict[str, httpx2.AsyncBaseTransport | None]:
"""Build HTTPX2 proxy mounts while leaving direct routes to the base transport."""
proxies = getproxies()
mounts: dict[str, httpx2.AsyncBaseTransport | None] = {}
for scheme in ("http", "https", "all"):
proxy_url = proxies.get(scheme)
if proxy_url:
if "://" not in proxy_url:
proxy_url = f"http://{proxy_url}"
mounts[f"{scheme}://"] = httpx2.AsyncHTTPTransport(proxy=httpx2.Proxy(proxy_url))
if not mounts:
return {}
no_proxy = proxies.get("no", "")
if no_proxy == "*":
return {}
for entry in no_proxy.split(","):
pattern = _no_proxy_mount_pattern(entry.strip())
if pattern:
mounts[pattern] = None
return mounts
def _no_proxy_mount_pattern(hostname: str) -> str | None:
if not hostname:
return None
@@ -290,7 +264,7 @@ class UnsafeURLRequestError(httpx.RequestError):
class PinnedDNSAsyncTransport(httpx.AsyncBaseTransport):
"""HTTPX transport that pins each request to the IPs validated for its URL."""
_resolver_lock = _DNS_PIN_RESOLVER_LOCK
_resolver_lock = asyncio.Lock()
def __init__(
self,
@@ -314,37 +288,6 @@ class PinnedDNSAsyncTransport(httpx.AsyncBaseTransport):
await self._inner.aclose()
class Httpx2UnsafeURLRequestError(httpx2.RequestError):
"""Raised when an HTTPX2 request is rejected by URL safety validation."""
class Httpx2PinnedDNSAsyncTransport(httpx2.AsyncBaseTransport):
"""HTTPX2 transport that pins each request to the IPs validated for its URL."""
_resolver_lock = _DNS_PIN_RESOLVER_LOCK
def __init__(
self,
*,
allow_loopback: bool = False,
inner: httpx2.AsyncBaseTransport | None = None,
) -> None:
self._allow_loopback = allow_loopback
self._inner = inner or httpx2.AsyncHTTPTransport()
async def handle_async_request(self, request: httpx2.Request) -> httpx2.Response:
url = str(request.url)
ok, error, resolved_ips = resolve_url_target(url, allow_loopback=self._allow_loopback)
if not ok:
raise Httpx2UnsafeURLRequestError(error, request=request)
async with self._resolver_lock:
with pin_resolved_url_dns(url, resolved_ips):
return await self._inner.handle_async_request(request)
async def aclose(self) -> None:
await self._inner.aclose()
def validate_resolved_url(url: str) -> tuple[bool, str]:
"""Validate an already-fetched URL (e.g. after redirect). Only checks the IP, skips DNS."""
try:
+53 -5
View File
@@ -17,6 +17,7 @@ from weakref import WeakValueDictionary
from loguru import logger
from nanobot.config.paths import get_legacy_sessions_dir
from nanobot.providers.base import ProviderConversationState
from nanobot.runtime_context import (
RUNTIME_CONTEXT_HISTORY_META,
public_history_message,
@@ -43,6 +44,10 @@ _SESSION_PREVIEW_MAX_CHARS = 120
_SESSION_LIST_PREVIEW_MAX_RECORDS = 200
_SESSION_LIST_PREVIEW_MAX_CHARS = 1_000_000
_SESSION_DATA_ERRORS = (ValueError, TypeError, AttributeError, KeyError)
_PROVIDER_STATE_RECORD_TYPE = "provider_state"
_PROVIDER_STATE_RECORD_PREFIX_RE = re.compile(
r'^\s*\{\s*"_type"\s*:\s*"provider_state"\s*(?:,|\})'
)
_FORK_VOLATILE_METADATA_KEYS = {
"goal_state",
"pending_user_turn",
@@ -60,6 +65,11 @@ def _json_object(value: object) -> dict[str, Any]:
return cast(dict[str, Any], value)
def _is_provider_state_record_line(line: str) -> bool:
"""Recognize the canonical private record without decoding its opaque payload."""
return _PROVIDER_STATE_RECORD_PREFIX_RE.match(line) is not None
def replay_max_messages_for_context(context_window_tokens: int | None) -> int:
if not context_window_tokens or context_window_tokens <= 0:
return FILE_MAX_MESSAGES
@@ -146,10 +156,13 @@ class Session:
updated_at: datetime = field(default_factory=datetime.now)
metadata: dict[str, Any] = field(default_factory=dict)
last_consolidated: int = 0 # Number of messages already consolidated to files
provider_state: ProviderConversationState | None = field(default=None, repr=False)
def __post_init__(self) -> None:
if not isinstance(cast(object, self.metadata), dict):
self.metadata = {}
if not isinstance(cast(object, self.provider_state), ProviderConversationState):
self.provider_state = None
# An out-of-range offset (corrupt metadata) would hide all history; reset it.
last_consolidated = cast(object, self.last_consolidated)
if (
@@ -304,6 +317,7 @@ class Session:
"""Clear all messages and reset session to initial state."""
self.messages = []
self.last_consolidated = 0
self.provider_state = None
self.updated_at = datetime.now()
self.metadata.pop("_last_summary", None)
@@ -396,6 +410,8 @@ class Session:
self.messages = retained
self.last_consolidated = new_lc
if dropped:
self.provider_state = None
self.updated_at = datetime.now()
return RetentionResult(
dropped=dropped,
@@ -517,6 +533,7 @@ class JsonlSessionStore:
created_at: datetime | None = None
updated_at: datetime | None = None
last_consolidated = 0
provider_state: ProviderConversationState | None = None
with open(path, encoding="utf-8") as f:
for line in f:
@@ -527,7 +544,8 @@ class JsonlSessionStore:
raw_data: object = json.loads(line)
data = _json_object(raw_data)
if data.get("_type") == "metadata":
record_type = data.get("_type")
if record_type == "metadata":
metadata_value = cast(object, data.get("metadata", {}))
metadata = (
cast(dict[str, Any], metadata_value)
@@ -552,6 +570,10 @@ class JsonlSessionStore:
if isinstance(offset, int) and not isinstance(offset, bool)
else 0
)
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
provider_state = ProviderConversationState.from_private_record(
data.get("state")
)
else:
messages.append(data)
@@ -562,6 +584,7 @@ class JsonlSessionStore:
updated_at=updated_at or datetime.now(),
metadata=metadata,
last_consolidated=last_consolidated,
provider_state=provider_state,
)
except _SESSION_DATA_ERRORS as e:
logger.warning("Failed to load session {}: {}", key, e)
@@ -586,6 +609,7 @@ class JsonlSessionStore:
created_at: datetime | None = None
updated_at: datetime | None = None
last_consolidated = 0
provider_state: ProviderConversationState | None = None
skipped = 0
with open(path, encoding="utf-8") as f:
@@ -603,7 +627,8 @@ class JsonlSessionStore:
continue
data = cast(dict[str, Any], raw_data)
if data.get("_type") == "metadata":
record_type = data.get("_type")
if record_type == "metadata":
metadata_value = cast(object, data.get("metadata", {}))
metadata = (
cast(dict[str, Any], metadata_value)
@@ -624,13 +649,21 @@ class JsonlSessionStore:
if isinstance(offset, int) and not isinstance(offset, bool)
else 0
)
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
candidate = ProviderConversationState.from_private_record(
data.get("state")
)
if candidate is None:
skipped += 1
else:
provider_state = candidate
else:
messages.append(data)
if skipped:
logger.warning("Skipped {} corrupt lines in session {}", skipped, key)
if not messages and not metadata:
if not messages and not metadata and provider_state is None:
return None
return Session(
@@ -640,6 +673,7 @@ class JsonlSessionStore:
updated_at=updated_at or datetime.now(),
metadata=metadata,
last_consolidated=last_consolidated,
provider_state=provider_state,
)
except _SESSION_DATA_ERRORS as e:
logger.warning("Repair failed for session {}: {}", key, e)
@@ -670,6 +704,12 @@ class JsonlSessionStore:
"last_consolidated": session.last_consolidated,
}
f.write(json.dumps(metadata_line, ensure_ascii=False) + "\n")
if session.provider_state is not None:
provider_state_line = {
"_type": _PROVIDER_STATE_RECORD_TYPE,
"state": session.provider_state.to_private_record(),
}
f.write(json.dumps(provider_state_line, ensure_ascii=False) + "\n")
for msg in session.messages:
f.write(json.dumps(msg, ensure_ascii=False) + "\n")
if fsync:
@@ -726,7 +766,8 @@ class JsonlSessionStore:
continue
raw_data: object = json.loads(line)
data = _json_object(raw_data)
if data.get("_type") == "metadata":
record_type = data.get("_type")
if record_type == "metadata":
metadata_value = cast(object, data.get("metadata", {}))
metadata = (
cast(dict[str, Any], metadata_value)
@@ -745,6 +786,8 @@ class JsonlSessionStore:
stored_key = (
stored_key_value if isinstance(stored_key_value, str) else None
)
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
continue
else:
messages.append(data)
return {
@@ -837,6 +880,8 @@ class JsonlSessionStore:
for line in f:
if not line.strip():
continue
if _is_provider_state_record_line(line):
continue
scanned_records += 1
scanned_chars += len(line)
if (
@@ -846,7 +891,10 @@ class JsonlSessionStore:
break
raw_item: object = json.loads(line)
item = _json_object(raw_item)
if item.get("_type") == "metadata":
if item.get("_type") in {
"metadata",
_PROVIDER_STATE_RECORD_TYPE,
}:
continue
text = _message_preview_text(item)
if not text:
+13 -2
View File
@@ -166,7 +166,8 @@ class LocalTriggerStore:
raise ValueError("trigger message is required")
self._ensure_dirs()
with self._lock:
trigger = self._find_unlocked(self._load_triggers_unlocked(), trigger_id)
triggers = self._load_triggers_unlocked()
trigger = self._find_unlocked(triggers, trigger_id)
if trigger is None:
raise TriggerNotFoundError(f"trigger not found: {trigger_id}")
if not trigger.enabled:
@@ -180,10 +181,20 @@ class LocalTriggerStore:
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
run_record_path: Path | None = None
try:
self.write_delivery_run_record(delivery, trigger=trigger, status="queued")
run_record_path = self.write_delivery_run_record(
delivery,
trigger=trigger,
status="queued",
)
trigger.last_message = _run_record_text(content)
trigger.updated_at_ms = delivery.created_at_ms
self._save_triggers_unlocked(triggers)
except BaseException:
path.unlink(missing_ok=True)
if run_record_path is not None:
run_record_path.unlink(missing_ok=True)
delivery.path = None
raise
return delivery
+3
View File
@@ -61,6 +61,7 @@ class LocalTrigger:
origin_metadata: dict[str, Any] = field(default_factory=dict)
created_at_ms: int = 0
updated_at_ms: int = 0
last_message: str = ""
last_run_at_ms: int | None = None
last_status: TriggerStatus | None = None
last_error: str | None = None
@@ -90,6 +91,7 @@ class LocalTrigger:
origin_metadata=dict(_get(data, "originMetadata", "origin_metadata", {}) or {}),
created_at_ms=_int_or_zero(_get(data, "createdAtMs", "created_at_ms", 0)),
updated_at_ms=_int_or_zero(_get(data, "updatedAtMs", "updated_at_ms", 0)),
last_message=str(_get(data, "lastMessage", "last_message", "") or ""),
last_run_at_ms=_optional_int(_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"),
@@ -108,6 +110,7 @@ class LocalTrigger:
"originMetadata": self.origin_metadata,
"createdAtMs": self.created_at_ms,
"updatedAtMs": self.updated_at_ms,
"lastMessage": self.last_message,
"lastRunAtMs": self.last_run_at_ms,
"lastStatus": self.last_status,
"lastError": self.last_error,
+7 -4
View File
@@ -176,7 +176,10 @@ class GitStore:
)
if cast(object, sha_bytes) is None:
return None
sha = sha_bytes.hex()[:8]
# porcelain.commit returns the id as a 40-char hex string that is
# already encoded to bytes; .hex() would encode those ASCII bytes
# again and produce an id no git command can resolve.
sha = sha_bytes.decode()[:8]
logger.debug("Git auto-commit: {} ({})", sha, message)
return sha
except Exception as exc:
@@ -200,7 +203,7 @@ class GitStore:
return None
while sha:
if sha.hex().startswith(short_sha):
if sha.decode().startswith(short_sha):
return sha
commit_obj = repo[sha]
if commit_obj.type_name != b"commit":
@@ -280,7 +283,7 @@ class GitStore:
msg = commit.message.decode("utf-8", errors="replace").strip()
if message_prefix is None or msg.startswith(message_prefix):
entries.append(CommitInfo(
sha=sha.hex()[:8],
sha=sha.decode()[:8],
message=msg,
timestamp=ts,
))
@@ -484,7 +487,7 @@ class GitStore:
with Repo(str(self._workspace)) as repo:
commit = cast("Commit", repo[full_sha])
parent = commit.parents[0] if commit.parents else None
diff = self.diff_commits(parent.hex()[:8], c.sha) if parent else ""
diff = self.diff_commits(parent.decode()[:8], c.sha) if parent else ""
return c, diff
return None
except Exception as exc:
+51 -10
View File
@@ -3,6 +3,7 @@
from __future__ import annotations
import email.utils
import gzip
import hmac
import http
import ipaddress
@@ -16,6 +17,9 @@ from websockets.http11 import Response
QueryParams = dict[str, list[str]]
_JSON_GZIP_MIN_BYTES = 4 * 1024
_JSON_GZIP_LEVEL = 5
def strip_trailing_slash(path: str) -> str:
if len(path) > 1 and path.endswith("/"):
@@ -41,6 +45,15 @@ def case_insensitive_header(headers: Any, key: str) -> str:
return str(value or "").strip()
def combined_list_header(headers: Any, key: str) -> str:
"""Combine repeated values for a comma-separated HTTP list header."""
try:
values = headers.get_all(key)
except (AttributeError, KeyError):
return case_insensitive_header(headers, key)
return ", ".join(str(value).strip() for value in values if str(value).strip())
def safe_host_header(value: str) -> str:
"""Return a safe Host header value, or empty when it should not be echoed."""
value = value.strip()
@@ -62,18 +75,46 @@ def host_for_url(host: str, port: int) -> str:
return f"{host}:{port}"
def http_json_response(data: dict[str, Any], *, status: int = 200) -> Response:
def _accepts_gzip(value: str) -> bool:
wildcard_quality: float | None = None
for item in value.split(","):
name, *params = (part.strip() for part in item.split(";"))
quality = 1.0
for param in params:
key, separator, raw_value = param.partition("=")
if separator and key.strip().lower() == "q":
try:
quality = float(raw_value.strip())
except ValueError:
quality = 0.0
break
if name.lower() == "gzip":
return quality > 0
if name == "*":
wildcard_quality = quality
return wildcard_quality is not None and wildcard_quality > 0
def http_json_response(
data: dict[str, Any],
*,
status: int = 200,
accept_encoding: str | None = None,
) -> Response:
body = json.dumps(data, ensure_ascii=False).encode("utf-8")
headers = Headers(
[
("Date", email.utils.formatdate(usegmt=True)),
("Connection", "close"),
("Content-Length", str(len(body))),
("Content-Type", "application/json; charset=utf-8"),
]
)
headers = [
("Date", email.utils.formatdate(usegmt=True)),
("Connection", "close"),
("Content-Type", "application/json; charset=utf-8"),
]
if accept_encoding is not None:
headers.append(("Vary", "Accept-Encoding"))
if len(body) >= _JSON_GZIP_MIN_BYTES and _accepts_gzip(accept_encoding):
body = gzip.compress(body, compresslevel=_JSON_GZIP_LEVEL, mtime=0)
headers.append(("Content-Encoding", "gzip"))
headers.append(("Content-Length", str(len(body))))
reason = http.HTTPStatus(status).phrase
return Response(status, reason, headers, body)
return Response(status, reason, Headers(headers), body)
def http_response(
+1 -1
View File
@@ -209,7 +209,7 @@ def _serialize_trigger(
},
"payload": {
"kind": "local_trigger",
"message": command,
"message": trigger.last_message or command,
"command": command,
},
"state": {
+92 -19
View File
@@ -16,20 +16,30 @@ from typing import Any, cast
from loguru import logger
from nanobot.config.paths import get_webui_dir
from nanobot.security.workspace_access import WORKSPACE_SCOPE_METADATA_KEY
from nanobot.session.history_visibility import is_hidden_history_message
from nanobot.session.manager import (
_PROVIDER_STATE_RECORD_TYPE, # pyright: ignore[reportPrivateUsage]
_SESSION_LIST_PREVIEW_MAX_CHARS, # pyright: ignore[reportPrivateUsage]
_SESSION_LIST_PREVIEW_MAX_RECORDS, # pyright: ignore[reportPrivateUsage]
Session,
SessionManager,
_is_provider_state_record_line, # pyright: ignore[reportPrivateUsage]
_message_preview_text, # pyright: ignore[reportPrivateUsage]
_metadata_title, # pyright: ignore[reportPrivateUsage]
)
from nanobot.session.model_selection import model_preset_from_metadata
_INDEX_VERSION = 4
_INDEX_VERSION = 6
_INDEX_FILENAME = ".webui_session_index.json"
_MODEL_PRESET_FIELD = "model_preset"
_WORKSPACE_SCOPE_PRESENT_FIELD = "_workspace_scope_present"
_WORKSPACE_SCOPE_VALUE_FIELD = "_workspace_scope_value"
WEBUI_SESSION_INDEX_INTERNAL_FIELDS = frozenset(
{_WORKSPACE_SCOPE_PRESENT_FIELD, _WORKSPACE_SCOPE_VALUE_FIELD}
)
_INDEXED_WORKSPACE_SCOPE_KEYS = ("project_path", "path", "access_mode")
_MAX_INDEXED_WORKSPACE_SCOPE_BYTES = 4096
_WEBUI_ACTIVITY_MTIME_NS = "webui_activity_mtime_ns"
_WEBUI_ACTIVITY_SIZE = "webui_activity_size"
_VISIBLE_TRANSCRIPT_ROLES = {"user", "assistant"}
@@ -59,17 +69,21 @@ def _reconcile_index(session_manager: SessionManager) -> tuple[list[dict[str, An
for path in session_manager.sessions_dir.glob("*.jsonl")
if SessionManager._session_key_from_path(path) is not None # pyright: ignore[reportPrivateUsage]
)
if not paths:
return [], existing_rows != []
webui_dir = get_webui_dir()
rows: list[dict[str, Any]] = []
changed = existing_rows is None
for path in paths:
row = existing_by_file.get(path.name)
if row is not None and _indexed_row_matches_file(row, path):
if row is not None and _indexed_row_matches_file(row, path, webui_dir):
rows.append(row)
continue
changed = True
scanned = _scan_session_row(session_manager, path)
scanned = _scan_session_row(session_manager, path, webui_dir)
if scanned is not None:
rows.append(scanned)
@@ -123,18 +137,20 @@ def _file_signature(path: Path) -> dict[str, int]:
return {"mtime_ns": stat.st_mtime_ns, "size": stat.st_size}
def _indexed_row_matches_file(row: dict[str, Any], path: Path) -> bool:
def _indexed_row_matches_file(row: dict[str, Any], path: Path, webui_dir: Path) -> bool:
if not all(isinstance(row.get(key), str) for key in ("key", "created_at", "updated_at")):
return False
if not isinstance(row.get("title", ""), str) or not isinstance(row.get("preview", ""), str):
return False
if not isinstance(row.get(_WORKSPACE_SCOPE_PRESENT_FIELD), bool):
return False
if row.get("file") != path.name:
return False
try:
signature = _file_signature(path)
except OSError:
return False
activity_signature = _webui_activity_signature(str(row.get("key")))
activity_signature = _webui_activity_signature(str(row.get("key")), webui_dir)
return (
row.get("mtime_ns") == signature["mtime_ns"]
and row.get("size") == signature["size"]
@@ -151,10 +167,57 @@ def _public_row(sessions_dir: Path, row: dict[str, Any]) -> dict[str, Any]:
"title": row.get("title", ""),
"preview": row.get("preview", ""),
_MODEL_PRESET_FIELD: row.get(_MODEL_PRESET_FIELD),
_WORKSPACE_SCOPE_PRESENT_FIELD: row.get(_WORKSPACE_SCOPE_PRESENT_FIELD, False),
_WORKSPACE_SCOPE_VALUE_FIELD: row.get(_WORKSPACE_SCOPE_VALUE_FIELD),
"path": str(sessions_dir / str(row.get("file", ""))),
}
def indexed_workspace_scope(row: dict[str, Any]) -> tuple[bool, object]:
"""Return the cached sidebar scope value while preserving missing vs null."""
return (
row.get(_WORKSPACE_SCOPE_PRESENT_FIELD) is True,
cast(object, row.get(_WORKSPACE_SCOPE_VALUE_FIELD)),
)
def _indexed_workspace_scope_fields(metadata: object) -> dict[str, object]:
if not isinstance(metadata, dict):
return {
_WORKSPACE_SCOPE_PRESENT_FIELD: False,
_WORKSPACE_SCOPE_VALUE_FIELD: None,
}
metadata_data = cast(dict[str, Any], metadata)
if WORKSPACE_SCOPE_METADATA_KEY not in metadata_data:
return {
_WORKSPACE_SCOPE_PRESENT_FIELD: False,
_WORKSPACE_SCOPE_VALUE_FIELD: None,
}
raw_scope = metadata_data.get(WORKSPACE_SCOPE_METADATA_KEY)
indexed_scope: object = False
if raw_scope is None:
indexed_scope = None
elif isinstance(raw_scope, dict):
scope_data = cast(dict[object, object], raw_scope)
recognized = {
key: scope_data[key]
for key in _INDEXED_WORKSPACE_SCOPE_KEYS
if key in scope_data
}
try:
encoded = json.dumps(recognized, ensure_ascii=False)
except (TypeError, ValueError):
pass
else:
if len(encoded.encode("utf-8")) <= _MAX_INDEXED_WORKSPACE_SCOPE_BYTES:
indexed_scope = cast(object, json.loads(encoded))
return {
_WORKSPACE_SCOPE_PRESENT_FIELD: True,
_WORKSPACE_SCOPE_VALUE_FIELD: indexed_scope,
}
def _preview_from_messages(messages: list[dict[str, Any]]) -> str:
fallback_preview = ""
scanned_records = 0
@@ -179,19 +242,18 @@ def _preview_from_messages(messages: list[dict[str, Any]]) -> str:
return fallback_preview
def _webui_activity_paths(session_key: str) -> list[Path]:
def _webui_activity_paths(session_key: str, webui_dir: Path) -> list[Path]:
stem = SessionManager.safe_key(session_key)
webui_dir = get_webui_dir()
return [
webui_dir / f"{stem}.jsonl",
webui_dir / f"{stem}.json",
]
def _webui_activity_signature(session_key: str) -> dict[str, int]:
def _webui_activity_signature(session_key: str, webui_dir: Path) -> dict[str, int]:
latest_mtime_ns = 0
total_size = 0
for path in _webui_activity_paths(session_key):
for path in _webui_activity_paths(session_key, webui_dir):
try:
stat = path.stat()
except OSError:
@@ -229,10 +291,10 @@ def _latest_updated_at(stored: str | None, activity: str | None) -> str | None:
def _visible_message_timestamp(item: dict[str, Any]) -> str | None:
if is_hidden_history_message(item):
return None
if item.get("role") not in _VISIBLE_TRANSCRIPT_ROLES:
return None
if is_hidden_history_message(item):
return None
timestamp = item.get("timestamp")
return timestamp if isinstance(timestamp, str) else None
@@ -254,9 +316,9 @@ def _visible_activity_updated_at(
return _latest_updated_at(visible_message_at, webui_activity) or stored
def _indexed_row_for_session(session: Session, path: Path) -> dict[str, Any]:
def _indexed_row_for_session(session: Session, path: Path, webui_dir: Path) -> dict[str, Any]:
signature = _file_signature(path)
activity_signature = _webui_activity_signature(session.key)
activity_signature = _webui_activity_signature(session.key, webui_dir)
activity_updated_at = _webui_activity_updated_at(activity_signature)
visible_message_at = _last_visible_message_at(session.messages)
return {
@@ -270,6 +332,7 @@ def _indexed_row_for_session(session: Session, path: Path) -> dict[str, Any]:
"title": _metadata_title(session.metadata),
"preview": _preview_from_messages(session.messages),
_MODEL_PRESET_FIELD: model_preset_from_metadata(session.metadata),
**_indexed_workspace_scope_fields(session.metadata),
"file": path.name,
"mtime_ns": signature["mtime_ns"],
"size": signature["size"],
@@ -277,11 +340,16 @@ def _indexed_row_for_session(session: Session, path: Path) -> dict[str, Any]:
}
def _scan_session_row(session_manager: SessionManager, path: Path) -> dict[str, Any] | None:
def _scan_session_row(
session_manager: SessionManager,
path: Path,
webui_dir: Path,
) -> dict[str, Any] | None:
storage_key = SessionManager._session_key_from_path(path) # pyright: ignore[reportPrivateUsage]
if storage_key is None:
return None
try:
signature = _file_signature(path)
with open(path, encoding="utf-8") as f:
first_line = f.readline().strip()
if not first_line:
@@ -298,7 +366,11 @@ def _scan_session_row(session_manager: SessionManager, path: Path) -> dict[str,
for line in f:
if not line.strip():
continue
if _is_provider_state_record_line(line):
continue
item = json.loads(line)
if item.get("_type") == _PROVIDER_STATE_RECORD_TYPE:
continue
timestamp = _visible_message_timestamp(item)
if timestamp is not None:
visible_message_at = _latest_updated_at(visible_message_at, timestamp)
@@ -324,7 +396,6 @@ def _scan_session_row(session_manager: SessionManager, path: Path) -> dict[str,
continue
if not fallback_preview and item.get("role") == "assistant":
fallback_preview = text
signature = _file_signature(path)
created_at_s = data.get("created_at")
updated_at_s = data.get("updated_at")
if not created_at_s or not updated_at_s:
@@ -332,7 +403,8 @@ def _scan_session_row(session_manager: SessionManager, path: Path) -> dict[str,
created_at_s = created_at_s or fallback_time
updated_at_s = updated_at_s or fallback_time
key = data.get("key") or storage_key
activity_signature = _webui_activity_signature(key)
metadata = data.get("metadata", {})
activity_signature = _webui_activity_signature(key, webui_dir)
activity_updated_at = _webui_activity_updated_at(activity_signature)
return {
"key": key,
@@ -342,9 +414,10 @@ def _scan_session_row(session_manager: SessionManager, path: Path) -> dict[str,
visible_message_at,
activity_updated_at,
),
"title": _metadata_title(data.get("metadata", {})),
"title": _metadata_title(metadata),
"preview": preview or fallback_preview,
_MODEL_PRESET_FIELD: model_preset_from_metadata(data.get("metadata", {})),
_MODEL_PRESET_FIELD: model_preset_from_metadata(metadata),
**_indexed_workspace_scope_fields(metadata),
"file": path.name,
"mtime_ns": signature["mtime_ns"],
"size": signature["size"],
@@ -354,4 +427,4 @@ def _scan_session_row(session_manager: SessionManager, path: Path) -> dict[str,
repaired = session_manager._repair(storage_key) # pyright: ignore[reportPrivateUsage]
if repaired is None:
return None
return _indexed_row_for_session(repaired, path)
return _indexed_row_for_session(repaired, path, webui_dir)
-20
View File
@@ -1234,8 +1234,6 @@ def settings_payload(
"temperature": effective_preset.temperature,
"reasoning_effort": effective_preset.reasoning_effort,
"timezone": defaults.timezone,
"bot_name": defaults.bot_name,
"bot_icon": defaults.bot_icon,
"tool_hint_max_length": defaults.tool_hint_max_length,
},
"model_presets": model_presets,
@@ -1406,24 +1404,6 @@ def update_agent_settings(query: QueryParams) -> dict[str, Any]:
changed = True
restart_required = True
bot_name = _query_first_alias(query, "bot_name", "botName")
if bot_name is not None:
bot_name = bot_name.strip()
if not bot_name:
raise WebUISettingsError("bot_name is required")
if defaults.bot_name != bot_name:
defaults.bot_name = bot_name
changed = True
restart_required = True
bot_icon = _query_first_alias(query, "bot_icon", "botIcon")
if bot_icon is not None:
bot_icon = bot_icon.strip()
if defaults.bot_icon != bot_icon:
defaults.bot_icon = bot_icon
changed = True
restart_required = True
tool_hint_max_length = _query_first_alias(
query,
"tool_hint_max_length",
+7
View File
@@ -159,6 +159,13 @@ def normalize_token_usage_state(raw: Any) -> dict[str, Any]:
if not isinstance(date, str) or len(date) != 10 or not isinstance(row_value, dict):
continue
row = cast(dict[str, Any], row_value)
try:
datetime.fromisoformat(date)
except ValueError:
# A hand-edited or foreign day key that is not a real date would
# otherwise reach token_usage_payload's date parsing and fail every
# settings request; drop it like any other malformed row.
continue
normalized = _normalize_usage_row(row)
if normalized["total_tokens"] <= 0 and normalized["requests"] <= 0:
continue
+151 -65
View File
@@ -28,7 +28,8 @@ WEBUI_TRANSCRIPT_SCHEMA_VERSION = 3
WEBUI_FORK_MARKER_EVENT = "fork_marker"
WEBUI_TRANSCRIPT_INCOMPLETE_KEY = "transcript_incomplete"
_MAX_TRANSCRIPT_FILE_BYTES = 8 * 1024 * 1024
_TARGET_ACTIVE_TRANSCRIPT_BYTES = _MAX_TRANSCRIPT_FILE_BYTES // 2
_ACTIVE_TRANSCRIPT_ROTATE_BYTES = 2 * 1024 * 1024
_TARGET_ACTIVE_TRANSCRIPT_BYTES = _ACTIVE_TRANSCRIPT_ROTATE_BYTES // 2
_TRANSCRIPT_SEGMENT_MANIFEST_VERSION = 2
_TRANSCRIPT_ACTIVE_CHUNK_ID = "active"
_TRANSCRIPT_SEGMENT_RE = re.compile(r"^\d{6}\.jsonl$")
@@ -284,12 +285,12 @@ def _normalize_manifest_entry(session_key: str, entry: Any) -> dict[str, Any] |
}
def _write_segment_manifest(session_key: str, segment_ids: list[str]) -> None:
def _write_segment_manifest(session_key: str, entries: list[dict[str, Any]]) -> None:
directory = webui_transcript_segments_dir(session_key)
directory.mkdir(parents=True, exist_ok=True)
data = {
"version": _TRANSCRIPT_SEGMENT_MANIFEST_VERSION,
"segments": [_segment_manifest_entry(session_key, segment_id) for segment_id in segment_ids],
"segments": entries,
}
path = _webui_transcript_manifest_path(session_key)
tmp_path = path.with_suffix(".json.tmp")
@@ -301,17 +302,14 @@ def _write_segment_manifest(session_key: str, segment_ids: list[str]) -> None:
raise
def _rebuild_segment_manifest(session_key: str) -> list[str]:
def _rebuild_segment_manifest(session_key: str) -> list[dict[str, Any]]:
segment_ids = _segment_ids_on_disk(session_key)
if segment_ids:
_write_segment_manifest(session_key, segment_ids)
entries = [_segment_manifest_entry(session_key, segment_id) for segment_id in segment_ids]
if entries:
_write_segment_manifest(session_key, entries)
else:
_webui_transcript_manifest_path(session_key).unlink(missing_ok=True)
return segment_ids
def _rebuilt_segment_manifest_entries(session_key: str) -> list[dict[str, Any]]:
return [_segment_manifest_entry(session_key, segment_id) for segment_id in _rebuild_segment_manifest(session_key)]
return entries
def _read_segment_manifest_entries(session_key: str) -> list[dict[str, Any]]:
@@ -320,7 +318,7 @@ def _read_segment_manifest_entries(session_key: str) -> list[dict[str, Any]]:
return []
path = _webui_transcript_manifest_path(session_key)
if not path.is_file():
return _rebuilt_segment_manifest_entries(session_key)
return _rebuild_segment_manifest(session_key)
try:
data = json.loads(path.read_text(encoding="utf-8"))
manifest = cast(dict[str, Any], data) if isinstance(data, dict) else None
@@ -330,18 +328,18 @@ def _read_segment_manifest_entries(session_key: str) -> list[dict[str, Any]]:
or manifest.get("version") != _TRANSCRIPT_SEGMENT_MANIFEST_VERSION
or not isinstance(raw_segments, list)
):
return _rebuilt_segment_manifest_entries(session_key)
return _rebuild_segment_manifest(session_key)
entries: list[dict[str, Any]] = []
for entry in cast(list[Any], raw_segments):
normalized = _normalize_manifest_entry(session_key, entry)
if normalized is None:
return _rebuilt_segment_manifest_entries(session_key)
return _rebuild_segment_manifest(session_key)
entries.append(normalized)
if [entry["id"] for entry in entries] != _segment_ids_on_disk(session_key):
return _rebuilt_segment_manifest_entries(session_key)
return _rebuild_segment_manifest(session_key)
return entries
except (OSError, json.JSONDecodeError, TypeError, AttributeError):
return _rebuilt_segment_manifest_entries(session_key)
return _rebuild_segment_manifest(session_key)
def _read_segment_ids(session_key: str) -> list[str]:
@@ -351,26 +349,40 @@ def _read_segment_ids(session_key: str) -> list[str]:
def _append_segment_turns(session_key: str, turns: list[list[dict[str, Any]]]) -> None:
if not turns:
return
segment_ids = _read_segment_ids(session_key)
next_id = int(segment_ids[-1]) + 1 if segment_ids else 1
entries = _read_segment_manifest_entries(session_key)
next_id = int(entries[-1]["id"]) + 1 if entries else 1
batch: list[list[dict[str, Any]]] = []
batch_bytes = 0
def write_batch() -> None:
nonlocal next_id
segment_id = f"{next_id:06d}"
path = _segment_file_path(session_key, segment_id)
_write_records_to_path(path, _flatten_turns(batch))
entries.append({
"id": segment_id,
"bytes": path.stat().st_size,
"turn_count": len(batch),
"user_count": sum(
1
for turn in batch
for row in turn
if _is_user_transcript_row(row)
),
})
next_id += 1
for turn in turns:
turn_bytes = _records_bytes(turn)
if batch and batch_bytes + turn_bytes > _MAX_TRANSCRIPT_FILE_BYTES:
segment_id = f"{next_id:06d}"
_write_records_to_path(_segment_file_path(session_key, segment_id), _flatten_turns(batch))
segment_ids.append(segment_id)
next_id += 1
write_batch()
batch = []
batch_bytes = 0
batch.append(turn)
batch_bytes += turn_bytes
if batch:
segment_id = f"{next_id:06d}"
_write_records_to_path(_segment_file_path(session_key, segment_id), _flatten_turns(batch))
segment_ids.append(segment_id)
_write_segment_manifest(session_key, segment_ids)
write_batch()
_write_segment_manifest(session_key, entries)
def _rotate_active_transcript_if_needed(session_key: str) -> None:
@@ -378,7 +390,7 @@ def _rotate_active_transcript_if_needed(session_key: str) -> None:
if not path.is_file():
return
try:
if path.stat().st_size <= _MAX_TRANSCRIPT_FILE_BYTES:
if path.stat().st_size <= _ACTIVE_TRANSCRIPT_ROTATE_BYTES:
return
except OSError:
return
@@ -426,6 +438,16 @@ def _read_chunk_turns(session_key: str, chunk_id: str) -> list[list[dict[str, An
return _split_transcript_turns(_read_transcript_file(path))
def _cached_chunk_turns(
session_key: str,
chunk_id: str,
turn_cache: dict[str, list[list[dict[str, Any]]]],
) -> list[list[dict[str, Any]]]:
if chunk_id not in turn_cache:
turn_cache[chunk_id] = _read_chunk_turns(session_key, chunk_id)
return turn_cache[chunk_id]
def _encode_page_cursor(before_turn_ordinal: int) -> str:
raw = json.dumps(
{"before_turn": before_turn_ordinal},
@@ -462,7 +484,10 @@ def _coerce_page_limit(limit: int | None) -> int:
return max(1, min(_MAX_TRANSCRIPT_PAGE_LIMIT, int(limit)))
def _chunk_turn_refs(session_key: str) -> list[_TranscriptChunkRef]:
def _chunk_turn_refs(
session_key: str,
turn_cache: dict[str, list[list[dict[str, Any]]]],
) -> list[_TranscriptChunkRef]:
_rotate_active_transcript_if_needed(session_key)
refs: list[_TranscriptChunkRef] = []
ordinal = 0
@@ -474,7 +499,11 @@ def _chunk_turn_refs(session_key: str) -> list[_TranscriptChunkRef]:
refs.append(_TranscriptChunkRef(chunk_id, ordinal, turn_count, int(entry["user_count"])))
ordinal += turn_count
if webui_transcript_path(session_key).is_file():
active_turns = _read_chunk_turns(session_key, _TRANSCRIPT_ACTIVE_CHUNK_ID)
active_turns = _cached_chunk_turns(
session_key,
_TRANSCRIPT_ACTIVE_CHUNK_ID,
turn_cache,
)
active_turn_count = len(active_turns)
if active_turn_count > 0:
refs.append(
@@ -492,6 +521,7 @@ def _count_user_messages_before_ordinal(
session_key: str,
chunks: list[_TranscriptChunkRef],
before_ordinal: int,
turn_cache: dict[str, list[list[dict[str, Any]]]],
) -> int:
total = 0
for chunk in chunks:
@@ -503,7 +533,7 @@ def _count_user_messages_before_ordinal(
if local_end >= chunk.turn_count:
total += chunk.user_count
continue
turns = _read_chunk_turns(session_key, chunk.chunk_id)
turns = _cached_chunk_turns(session_key, chunk.chunk_id, turn_cache)
total += sum(
1
for turn in turns[:local_end]
@@ -521,7 +551,8 @@ def _select_transcript_page(
_manifest_rebuilt: bool = False,
) -> tuple[list[dict[str, Any]], dict[str, Any]]:
page_limit = _coerce_page_limit(limit)
chunks = _chunk_turn_refs(session_key)
turn_cache: dict[str, list[list[dict[str, Any]]]] = {}
chunks = _chunk_turn_refs(session_key, turn_cache)
total_turns = sum(chunk.turn_count for chunk in chunks)
before_ordinal = _decode_page_cursor(before)
upper_ordinal = total_turns if before_ordinal is None else min(before_ordinal, total_turns)
@@ -534,7 +565,7 @@ def _select_transcript_page(
local_upper = min(chunk.turn_count, upper_ordinal - chunk.start_ordinal)
if local_upper <= 0:
continue
turns = _read_chunk_turns(session_key, chunk.chunk_id)
turns = _cached_chunk_turns(session_key, chunk.chunk_id, turn_cache)
if (
chunk.chunk_id != _TRANSCRIPT_ACTIVE_CHUNK_ID
and len(turns) != chunk.turn_count
@@ -585,6 +616,7 @@ def _select_transcript_page(
session_key,
chunks,
first_ref.ordinal,
turn_cache,
),
}
return lines, page
@@ -1182,19 +1214,18 @@ def _is_recoverable_answer_record(record: dict[str, Any]) -> bool:
}
def recover_incomplete_turns_from_session(
lines: list[dict[str, Any]],
session_messages: list[dict[str, Any]] | None,
*,
session_key: str,
) -> list[dict[str, Any]]:
"""Recover marked transcript answers only when one durable session turn matches."""
if not lines or not session_messages:
return lines
session_turns = _session_backfill_turns(session_key, session_messages)
if not session_turns:
return lines
def _needs_incomplete_turn_recovery(lines: list[dict[str, Any]]) -> bool:
return any(
record.get("event") == "turn_end"
and record.get(WEBUI_TRANSCRIPT_INCOMPLETE_KEY) is True
for record in lines
)
def _recover_incomplete_turns(
lines: list[dict[str, Any]],
session_turns: list[_SessionBackfillTurn],
) -> list[dict[str, Any]]:
recovered: list[dict[str, Any]] = []
for turn in _split_transcript_turns(lines):
turn_end = turn[-1] if turn else None
@@ -1244,6 +1275,21 @@ def recover_incomplete_turns_from_session(
return recovered
def recover_incomplete_turns_from_session(
lines: list[dict[str, Any]],
session_messages: list[dict[str, Any]] | None,
*,
session_key: str,
) -> list[dict[str, Any]]:
"""Recover marked transcript answers only when one durable session turn matches."""
if not lines or not session_messages or not _needs_incomplete_turn_recovery(lines):
return lines
session_turns = _session_backfill_turns(session_key, session_messages)
if not session_turns:
return lines
return _recover_incomplete_turns(lines, session_turns)
def _with_backfilled_user(
records: list[dict[str, Any]],
user_event: dict[str, Any],
@@ -1254,18 +1300,19 @@ def _with_backfilled_user(
return records
def inject_missing_user_events_from_session(
session_key: str,
lines: list[dict[str, Any]],
session_messages: list[dict[str, Any]] | None,
) -> list[dict[str, Any]]:
"""Backfill user rows for legacy WebUI transcripts that only stored assistant streams."""
if not lines or not session_messages:
return lines
session_turns = _session_backfill_turns(session_key, session_messages)
if not session_turns:
return lines
def _needs_user_event_backfill(lines: list[dict[str, Any]]) -> bool:
for turn in _split_transcript_turns(lines):
if any(record.get("event") == "user" for record in turn):
continue
if _transcript_turn_signature(turn):
return True
return False
def _inject_missing_user_events(
lines: list[dict[str, Any]],
session_turns: list[_SessionBackfillTurn],
) -> list[dict[str, Any]]:
out: list[dict[str, Any]] = []
session_cursor = 0
for turn in _split_transcript_turns(lines):
@@ -1280,6 +1327,20 @@ def inject_missing_user_events_from_session(
return out
def inject_missing_user_events_from_session(
session_key: str,
lines: list[dict[str, Any]],
session_messages: list[dict[str, Any]] | None,
) -> list[dict[str, Any]]:
"""Backfill user rows for legacy WebUI transcripts that only stored assistant streams."""
if not lines or not session_messages or not _needs_user_event_backfill(lines):
return lines
session_turns = _session_backfill_turns(session_key, session_messages)
if not session_turns:
return lines
return _inject_missing_user_events(lines, session_turns)
def _format_tool_call_trace(call: Any) -> str | None:
if not call or not isinstance(call, dict):
return None
@@ -2026,6 +2087,7 @@ def replay_transcript_to_ui_messages(
continue
close_activity_for_answer()
turn_fields = _turn_fields(rec, "answer")
source_fields = _source_fields(rec)
adopted = find_active_placeholder(messages, turn_fields) if buffer_message_id is None else None
if buffer_message_id is None:
if adopted:
@@ -2038,7 +2100,8 @@ def replay_transcript_to_ui_messages(
"role": "assistant",
"content": "",
"isStreaming": True,
**_turn_fields(rec, "answer"),
**turn_fields,
**source_fields,
"createdAt": _created_at_ms(rec, idx),
},
)
@@ -2050,7 +2113,8 @@ def replay_transcript_to_ui_messages(
**m,
"content": combined,
"isStreaming": True,
**_turn_fields(rec, "answer"),
**turn_fields,
**source_fields,
}
break
continue
@@ -2062,6 +2126,8 @@ def replay_transcript_to_ui_messages(
continue
merge_next = rec.get("resuming") is True and rec.get("merge_next") is True
final_text = rec.get("text")
turn_fields = _turn_fields(rec, "answer")
source_fields = _source_fields(rec)
if isinstance(final_text, str):
if buffer_message_id is None:
buffer_message_id = _new_id("buf", idx)
@@ -2071,7 +2137,8 @@ def replay_transcript_to_ui_messages(
"role": "assistant",
"content": final_text,
"isStreaming": True,
**_turn_fields(rec, "answer"),
**turn_fields,
**source_fields,
"createdAt": _created_at_ms(rec, idx),
},
)
@@ -2082,11 +2149,21 @@ def replay_transcript_to_ui_messages(
**m,
"content": final_text,
"isStreaming": True,
**_turn_fields(rec, "answer"),
**turn_fields,
**source_fields,
}
break
if merge_next:
buffer_parts = [final_text]
elif source_fields and buffer_message_id is not None:
for i, m in enumerate(messages):
if m.get("id") == buffer_message_id:
messages[i] = {
**m,
**turn_fields,
**source_fields,
}
break
if not merge_next:
buffer_message_id = None
buffer_parts = []
@@ -2342,6 +2419,7 @@ def build_webui_thread_response(
augment_assistant_media: Callable[[list[str]], list[dict[str, Any]]] | None = None,
augment_assistant_text: Callable[[str], str] | None = None,
session_messages: list[dict[str, Any]] | None = None,
session_messages_loader: Callable[[], list[dict[str, Any]] | None] | None = None,
active_turn_started_at: float | None = None,
active_turn_id: str | None = None,
active_turn_transcript_persistence_failed: bool = False,
@@ -2358,12 +2436,20 @@ def build_webui_thread_response(
lines = _annotate_replay_identities(read_transcript_lines(session_key))
if not lines and active_turn_started_at is None:
return None
lines = inject_missing_user_events_from_session(session_key, lines, session_messages)
lines = recover_incomplete_turns_from_session(
lines,
session_messages,
session_key=session_key,
)
needs_user_backfill = _needs_user_event_backfill(lines)
needs_incomplete_recovery = _needs_incomplete_turn_recovery(lines)
if (
session_messages is None
and session_messages_loader is not None
and (needs_user_backfill or needs_incomplete_recovery)
):
session_messages = session_messages_loader()
if session_messages and (needs_user_backfill or needs_incomplete_recovery):
session_turns = _session_backfill_turns(session_key, session_messages)
if needs_user_backfill:
lines = _inject_missing_user_events(lines, session_turns)
if needs_incomplete_recovery:
lines = _recover_incomplete_turns(lines, session_turns)
lines = _ensure_replay_identities(lines)
fork_boundary = fork_boundary_message_count(lines)
msgs = replay_transcript_to_ui_messages(
+33 -10
View File
@@ -191,24 +191,47 @@ class WebUIWorkspaceController:
self._default_restrict_to_workspace,
)
def scope_for_session_key(self, session_key: str) -> WorkspaceScope:
if self._sessions is None:
return self.default_scope()
data = self._sessions.read_session_metadata(session_key)
session_data = data if data is not None else {}
metadata = session_data.get("metadata", {})
if not isinstance(metadata, dict) or WORKSPACE_SCOPE_METADATA_KEY not in metadata:
return self.default_scope()
metadata = cast(dict[str, Any], metadata)
def _scope_from_metadata_value(
self,
raw_scope: object,
*,
default_scope: WorkspaceScope | None = None,
) -> WorkspaceScope:
try:
return validate_workspace_scope_payload(
metadata.get(WORKSPACE_SCOPE_METADATA_KEY),
raw_scope,
default_workspace=self._default_workspace,
default_restrict_to_workspace=self._default_restrict_to_workspace,
source_channel=_WEBUI_SCOPE_CHANNEL,
)
except WorkspaceScopeError:
return default_scope if default_scope is not None else self.default_scope()
def scope_for_indexed_metadata(
self,
raw_scope: object,
*,
scope_present: bool,
default_scope: WorkspaceScope,
) -> WorkspaceScope:
"""Resolve a sidebar-only metadata snapshot without an authority-store read."""
if not scope_present:
return default_scope
return self._scope_from_metadata_value(raw_scope, default_scope=default_scope)
def scope_for_session_key(self, session_key: str) -> WorkspaceScope:
if self._sessions is None:
return self.default_scope()
data = self._sessions.read_session_metadata(session_key)
if not isinstance(data, dict):
return self.default_scope()
metadata = data.get("metadata", {})
if not isinstance(metadata, dict) or WORKSPACE_SCOPE_METADATA_KEY not in metadata:
return self.default_scope()
metadata_data = cast(dict[str, Any], metadata)
return self._scope_from_metadata_value(
cast(object, metadata_data.get(WORKSPACE_SCOPE_METADATA_KEY))
)
def payload(self, *, controls_available: bool) -> dict[str, Any]:
return workspaces_payload(
+69 -16
View File
@@ -27,6 +27,7 @@ from nanobot.command.builtin import builtin_command_palette
from nanobot.cron.session_turns import is_bound_cron_job
from nanobot.cron.types import CronJob, CronSchedule
from nanobot.runtime_context import public_history_messages
from nanobot.security.workspace_access import WorkspaceScope
from nanobot.triggers.local_types import LocalTrigger
from nanobot.utils.subagent_channel_display import scrub_subagent_messages_for_channel
from nanobot.webui.file_preview import (
@@ -38,6 +39,9 @@ from nanobot.webui.gateway_tokens import GatewayTokenStore, token_response_paylo
from nanobot.webui.http_utils import (
case_insensitive_header as _case_insensitive_header,
)
from nanobot.webui.http_utils import (
combined_list_header as _combined_list_header,
)
from nanobot.webui.http_utils import (
host_for_url as _host_for_url,
)
@@ -82,7 +86,11 @@ from nanobot.webui.session_automations import (
session_automation_jobs,
session_automations_payload,
)
from nanobot.webui.session_list_index import list_webui_sessions
from nanobot.webui.session_list_index import (
WEBUI_SESSION_INDEX_INTERNAL_FIELDS,
indexed_workspace_scope,
list_webui_sessions,
)
from nanobot.webui.sidebar_state import (
read_webui_sidebar_state,
write_webui_sidebar_state,
@@ -108,6 +116,30 @@ from nanobot.webui.workspaces import WebUIWorkspaceController
_SLOW_WEBUI_HTTP_LOG_MS = 1_000
_AUTOMATION_VALUES_HEADER = "X-Nanobot-Automation-Values"
# Fix for #5190: On Windows, mimetypes.guess_type() reads the registry key
# HKEY_CLASSES_ROOT\.js\Content Type, which is commonly set to 'text/plain'
# because .js is associated with Windows Script Host rather than web JavaScript.
# That registry value overrides Python's built-in mapping and causes browsers to
# reject ES module scripts with:
# Failed to load module script: Expected a JavaScript-or-Wasm module script
# but the server responded with a MIME type of "text/plain".
# We explicitly register correct MIME types for common web static assets here
# (module-import time) so all callers of mimetypes.guess_type() in this process
# benefit, regardless of host registry configuration.
_MIME_FIXES: dict[str, str] = {
".js": "application/javascript",
".mjs": "application/javascript",
".css": "text/css",
".html": "text/html",
".json": "application/json",
".svg": "image/svg+xml",
".wasm": "application/wasm",
}
for _ext, _ctype in _MIME_FIXES.items():
mimetypes.add_type(_ctype, _ext, strict=True)
if TYPE_CHECKING:
from nanobot.bus.queue import MessageBus
from nanobot.channels.websocket.runtime import WebSocketConfig
@@ -115,7 +147,6 @@ if TYPE_CHECKING:
from nanobot.session.manager import SessionManager
from nanobot.triggers.local_store import LocalTriggerStore
def _decode_api_key(raw_key: str) -> str | None:
key = unquote(raw_key)
_api_key_re = re.compile(r"^[A-Za-z0-9_:.-]{1,128}$")
@@ -422,7 +453,10 @@ class GatewayHTTPHandler:
if self.session_manager is None:
return _http_error(503, "session manager unavailable")
payload = await asyncio.to_thread(self._sessions_list_payload)
return _http_json_response(payload)
return _http_json_response(
payload,
accept_encoding=_combined_list_header(request.headers, "Accept-Encoding"),
)
def _sessions_list_payload(self) -> dict[str, Any]:
assert self.session_manager is not None
@@ -430,16 +464,28 @@ class GatewayHTTPHandler:
from nanobot.session.webui_turns import websocket_turn_wall_started_at
cleaned: list[dict[str, Any]] = []
default_scope: WorkspaceScope | None = None
for s in sessions:
key = s.get("key")
if not (isinstance(key, str) and key.startswith("websocket:")):
continue
row = {k: v for k, v in s.items() if k != "path"}
row = {
k: v
for k, v in s.items()
if k != "path" and k not in WEBUI_SESSION_INDEX_INTERNAL_FIELDS
}
chat_id = key.split(":", 1)[1]
started_at = websocket_turn_wall_started_at(chat_id)
if started_at is not None:
row["run_started_at"] = started_at
scope = self.workspaces.scope_for_session_key(key)
if default_scope is None:
default_scope = self.workspaces.default_scope()
scope_present, raw_scope = indexed_workspace_scope(s)
scope = self.workspaces.scope_for_indexed_metadata(
raw_scope,
scope_present=scope_present,
default_scope=default_scope,
)
row["workspace_scope"] = scope.payload()
cleaned.append(row)
return {"sessions": cleaned}
@@ -481,17 +527,21 @@ class GatewayHTTPHandler:
if not _is_websocket_channel_session_key(decoded_key):
return _http_error(404, "session not found")
scope = self.workspaces.scope_for_session_key(decoded_key)
session_messages: list[dict[str, Any]] | None = None
if self.session_manager is not None:
def load_session_messages() -> list[dict[str, Any]] | None:
if self.session_manager is None:
return None
session_data = self.session_manager.read_session_file(decoded_key)
raw_messages = session_data.get("messages") if isinstance(session_data, dict) else None
if isinstance(raw_messages, list):
raw_session_messages = cast(list[Any], raw_messages)
session_messages = [
cast(dict[str, Any], raw_message)
for raw_message in raw_session_messages
if isinstance(raw_message, dict)
]
if not isinstance(raw_messages, list):
return None
raw_session_messages = cast(list[Any], raw_messages)
return [
cast(dict[str, Any], raw_message)
for raw_message in raw_session_messages
if isinstance(raw_message, dict)
]
query = _parse_query(request.path)
raw_limit = _query_first(query, "limit")
limit: int | None = None
@@ -524,7 +574,7 @@ class GatewayHTTPHandler:
text,
workspace_path=scope.project_path,
),
session_messages=session_messages,
session_messages_loader=load_session_messages,
active_turn_started_at=active_turn_started_at,
active_turn_id=active_turn_id,
active_turn_transcript_persistence_failed=(
@@ -537,7 +587,10 @@ class GatewayHTTPHandler:
if data is None:
return _http_error(404, "webui thread not found")
data["workspace_scope"] = scope.payload()
return _http_json_response(data)
return _http_json_response(
data,
accept_encoding=_combined_list_header(request.headers, "Accept-Encoding"),
)
def _handle_file_preview(self, request: WsRequest, key: str) -> Response:
if not self.check_api_token(request):
+3 -5
View File
@@ -24,15 +24,13 @@ license-files = [
dependencies = [
"typer>=0.20.0,<1.0.0",
"anthropic>=0.45.0,<1.0.0",
"anthropic>=0.100.0,<1.0.0",
"pydantic>=2.12.0,<3.0.0",
"pydantic-settings>=2.12.0,<3.0.0",
# Feishu's lark-oapi currently requires websockets<16; core supports 15 and 16.
"websockets>=15.0,<17.0",
"websocket-client>=1.9.0,<2.0.0",
"httpx>=0.28.0,<1.0.0",
# MCP v2 uses the independently versioned httpx2 package for HTTP transports.
"httpx2>=2.5.0,<3.0.0",
"ddgs>=9.5.5,<10.0.0",
"oauth-cli-kit>=0.1.6,<1.0.0",
"loguru>=0.7.3,<1.0.0",
@@ -42,7 +40,7 @@ dependencies = [
"croniter>=6.0.0,<7.0.0",
"prompt-toolkit>=3.0.50,<4.0.0",
"questionary>=2.0.0,<3.0.0",
"mcp>=2.0.0,<3.0.0",
"mcp>=1.26.0,<2.0.0",
"json-repair>=0.57.0,<1.0.0",
"chardet>=3.0.2,<6.0.0",
"openai>=2.8.0",
@@ -53,7 +51,7 @@ dependencies = [
"filelock>=3.25.2",
"watchfiles>=1.1.1,<2.0.0",
"packaging>=24.0",
"tzdata>=2025.2; sys_platform == 'win32'",
"tzdata>=2025.2",
"defusedxml>=0.7.1,<1.0.0",
"pypdf>=5.0.0,<6.0.0",
"python-docx>=1.1.0,<2.0.0",
+102
View File
@@ -154,6 +154,26 @@ class TestIsExpired:
now_over = datetime(2026, 1, 1, 10, 10, 0)
assert ac._is_expired(ts, now=now_over) is True
def test_unparseable_string_timestamp_returns_false(self):
"""A persisted timestamp that no longer parses must not raise.
list_sessions() forwards the raw persisted updated_at string, and
SessionManager._load already tolerates a malformed value through its
recovery path. The idle scan must mirror that tolerance instead of crashing.
"""
ac = _make_autocompact(ttl=15)
assert ac._is_expired("not-a-timestamp") is False
def test_tz_aware_string_timestamp_is_compared_by_instant(self):
"""A valid timestamp with an offset remains eligible for expiry."""
ac = _make_autocompact(ttl=15)
now = datetime(2026, 1, 1, 12, 0, 0)
recent = (now - timedelta(minutes=10)).astimezone().isoformat()
expired = (now - timedelta(minutes=20)).astimezone().isoformat()
assert ac._is_expired(recent, now=now) is False
assert ac._is_expired(expired, now=now) is True
# ---------------------------------------------------------------------------
# _format_summary
@@ -221,6 +241,36 @@ class TestCheckExpired:
assert len(scheduled) == 1
assert "cli:old" in ac._archiving
def test_unparseable_updated_at_does_not_stop_scan(self):
"""A malformed timestamp is skipped without hiding later sessions.
The idle scan runs from the agent loop's inbound-timeout branch, so a
raised exception here would tear down the loop. list_sessions() forwards
the raw string, so check_expired must tolerate it like SessionManager
does when loading.
"""
ac = _make_autocompact(ttl=15)
mock_sm = MagicMock(spec=SessionManager)
old_dt = datetime.now() - timedelta(minutes=20)
session = _make_session("cli:old", updated_at=old_dt)
_add_turns(session, 5)
mock_sm.list_sessions.return_value = [
{"key": "cli:corrupt", "updated_at": "not-a-timestamp"},
{"key": "cli:old", "updated_at": old_dt.isoformat()},
]
mock_sm.get_or_create.return_value = session
ac.sessions = mock_sm
scheduled = []
def scheduler(coro):
scheduled.append(coro)
coro.close()
ac.check_expired(scheduler, _runtime)
assert len(scheduled) == 1
assert ac._archiving == {"cli:old"}
@pytest.mark.asyncio
async def test_runtime_is_captured_before_background_starts(self):
ac = _make_autocompact(ttl=15)
@@ -542,6 +592,58 @@ class TestPrepareSession:
assert summary is not None
assert "Cold summary." in summary
def test_cold_path_tolerates_malformed_last_active(self):
"""A malformed persisted last_active must not raise on the turn path.
prepare_session runs from _compact_session on every turn. Persisted
_last_summary can be hand-edited or written by another version, so a bad
last_active should degrade gracefully (mirror estimate_session_prompt_tokens
and _archive) instead of crashing the turn.
"""
ac = _make_autocompact(ttl=0)
fallback = datetime(2026, 1, 2, 3, 4, 5)
session = _make_session(
metadata={
"_last_summary": {"text": "Cold summary.", "last_active": "not-a-date"},
},
updated_at=fallback,
)
result_session, summary = ac.prepare_session(session, "cli:test")
assert result_session is session
assert summary is not None
assert "Cold summary." in summary
assert fallback.isoformat() in summary
def test_cold_path_tolerates_missing_last_active(self):
"""A _last_summary dict without last_active must not raise."""
ac = _make_autocompact(ttl=0)
fallback = datetime(2026, 1, 2, 3, 4, 5)
session = _make_session(
metadata={"_last_summary": {"text": "Cold summary."}},
updated_at=fallback,
)
result_session, summary = ac.prepare_session(session, "cli:test")
assert result_session is session
assert summary is not None
assert "Cold summary." in summary
assert fallback.isoformat() in summary
def test_cold_path_missing_text_returns_none(self):
"""A _last_summary without a non-empty string text yields no summary."""
ac = _make_autocompact()
session = _make_session(metadata={
"_last_summary": {"last_active": datetime(2026, 1, 1).isoformat()},
})
result_session, summary = ac.prepare_session(session, "cli:test")
assert result_session is session
assert summary is None
def test_no_summary_available_returns_none(self):
"""When no summary is available, should return (session, None)."""
ac = _make_autocompact()
+21 -1
View File
@@ -10,7 +10,11 @@ from nanobot.agent.memory import (
Consolidator,
MemoryStore,
)
from nanobot.providers.base import GenerationSettings, LLMResponse
from nanobot.providers.base import (
GenerationSettings,
LLMResponse,
ProviderConversationState,
)
from nanobot.runtime_context import (
RUNTIME_CONTEXT_HISTORY_META,
RuntimeContextBlock,
@@ -74,6 +78,16 @@ def _tool_round(call_id: str) -> list[dict]:
]
def _provider_state() -> ProviderConversationState:
return ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": []},
)
class TestConsolidatorSummarize:
async def test_archive_prompt_includes_media_breadcrumb(
self, consolidator, mock_provider, store, runtime
@@ -385,6 +399,7 @@ class TestConsolidatorTokenBudget:
"""Old messages that cannot be replayed should be materialized first."""
consolidator._SAFETY_BUFFER = 0
session = Session(key="test:replay-overflow")
session.provider_state = _provider_state()
for i in range(10):
session.add_message("user", f"u{i}")
session.add_message("assistant", f"a{i}")
@@ -404,6 +419,7 @@ class TestConsolidatorTokenBudget:
assert archived_chunk[-1]["content"] == "a6"
assert session.last_consolidated == 14
assert session.metadata["_last_summary"]["text"] == "old conversation summary"
assert session.provider_state is None
consolidator.sessions.save.assert_called()
async def test_replay_window_overflow_extends_to_long_recent_user_turn(
@@ -479,6 +495,7 @@ class TestConsolidatorTokenBudget:
session = MagicMock()
session.last_consolidated = 0
session.key = "test:key"
session.provider_state = _provider_state()
session.messages = [
{
"role": "user" if i in {0, 50, 61} else "assistant",
@@ -500,6 +517,7 @@ class TestConsolidatorTokenBudget:
# pick_consolidation_boundary returns (50, tokens) — user turn at idx 50
assert archived_chunk[0]["content"] == "m0"
assert session.last_consolidated > 0
assert session.provider_state is None
async def test_raw_archive_fallback_advances_last_consolidated(
self, consolidator, runtime
@@ -610,6 +628,7 @@ class TestCompactIdleSession:
)
sessions = real_consolidator.sessions
session = sessions.get_or_create("cli:test")
session.provider_state = _provider_state()
old_ts = session.updated_at
for i in range(20):
session.add_message("user", f"user msg {i}")
@@ -627,6 +646,7 @@ class TestCompactIdleSession:
assert len(reloaded.messages) == 40
assert reloaded.messages[0]["content"] == "user msg 0"
assert reloaded.last_consolidated == 32
assert reloaded.provider_state is None
visible = reloaded.get_history(max_messages=40)
assert len(visible) == 8
assert visible[0]["content"] == "user msg 16"
+14
View File
@@ -452,6 +452,20 @@ class TestBuildMessages:
assert "previous user message" in str(messages[1]["content"])
assert "new message" in str(messages[1]["content"])
def test_current_message_can_be_built_without_history_merge(self, tmp_path):
builder = _builder(tmp_path)
current = builder.build_current_message(
"new message",
runtime_context_blocks=[
RuntimeContextBlock(source="test", content="fresh context"),
],
)
assert current["role"] == "user"
assert "new message" in current["content"]
assert "fresh context" in current["content"]
assert current["_meta"]["runtime_context"]["sources"] == ["test"]
def test_different_role_appended(self, tmp_path):
builder = _builder(tmp_path)
history = [{"role": "assistant", "content": "previous response"}]
+308 -1
View File
@@ -1,4 +1,5 @@
import asyncio
import json
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
@@ -19,7 +20,7 @@ from nanobot.bus.outbound_events import (
)
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.providers.base import LLMProvider, LLMResponse, ProviderConversationState
from nanobot.providers.factory import ProviderSnapshot
from nanobot.runtime_context import (
RUNTIME_CONTEXT_HISTORY_META,
@@ -59,6 +60,16 @@ def _mk_loop() -> AgentLoop:
return loop
def _provider_state() -> ProviderConversationState:
return ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": []},
)
def _runtime_message(content, blocks: list[RuntimeContextBlock]) -> dict:
merged, marker = append_runtime_context(content, blocks)
assert marker is not None
@@ -494,6 +505,7 @@ def test_restore_runtime_checkpoint_rehydrates_completed_and_pending_tools() ->
loop = _mk_loop()
session = Session(
key="test:checkpoint",
provider_state=_provider_state(),
metadata={
AgentLoop._RUNTIME_CHECKPOINT_KEY: {
"assistant_message": {
@@ -539,6 +551,104 @@ def test_restore_runtime_checkpoint_rehydrates_completed_and_pending_tools() ->
assert session.messages[1]["tool_call_id"] == "call_done"
assert session.messages[2]["tool_call_id"] == "call_pending"
assert "interrupted before this tool finished" in session.messages[2]["content"].lower()
assert session.provider_state is None
def test_restore_final_response_checkpoint_preserves_matching_provider_state() -> None:
loop = _mk_loop()
state = _provider_state()
session = Session(
key="test:final-checkpoint",
provider_state=state,
metadata={
AgentLoop._RUNTIME_CHECKPOINT_KEY: {
"phase": "final_response",
AgentLoop._PROVIDER_STATE_CHECKPOINT_VERSION_KEY: (
AgentLoop._PROVIDER_STATE_CHECKPOINT_VERSION
),
"assistant_message": {
"role": "assistant",
"content": "finished",
},
"completed_tool_results": [],
"pending_tool_calls": [],
}
},
)
restored = loop._restore_runtime_checkpoint(session)
assert restored is True
assert session.messages[-1]["content"] == "finished"
assert session.provider_state is state
assert session.metadata.get(AgentLoop._RUNTIME_CHECKPOINT_KEY) is None
def test_restore_legacy_final_checkpoint_discards_unproven_provider_state() -> None:
loop = _mk_loop()
session = Session(
key="test:legacy-final-checkpoint",
provider_state=_provider_state(),
metadata={
AgentLoop._RUNTIME_CHECKPOINT_KEY: {
"phase": "final_response",
"assistant_message": {
"role": "assistant",
"content": "finished",
},
"completed_tool_results": [],
"pending_tool_calls": [],
}
},
)
restored = loop._restore_runtime_checkpoint(session)
assert restored is True
assert session.messages[-1]["content"] == "finished"
assert session.provider_state is None
def test_restore_completed_tools_checkpoint_preserves_matching_provider_state() -> None:
loop = _mk_loop()
tool_result = {
"role": "tool",
"tool_call_id": "call_done",
"name": "read_file",
"content": "compacted result",
}
state = _provider_state().with_pending_messages([tool_result])
session = Session(
key="test:completed-tools-checkpoint",
provider_state=state,
metadata={
AgentLoop._RUNTIME_CHECKPOINT_KEY: {
"phase": "tools_completed",
AgentLoop._PROVIDER_STATE_CHECKPOINT_VERSION_KEY: (
AgentLoop._PROVIDER_STATE_CHECKPOINT_VERSION
),
"assistant_message": {
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_done",
"type": "function",
"function": {"name": "read_file", "arguments": "{}"},
}
],
},
"completed_tool_results": [tool_result],
"pending_tool_calls": [],
}
},
)
restored = loop._restore_runtime_checkpoint(session)
assert restored is True
assert session.messages[-1]["content"] == "compacted result"
assert session.provider_state is state
def test_restore_runtime_checkpoint_dedupes_overlapping_tail() -> None:
@@ -616,6 +726,55 @@ def test_restore_runtime_checkpoint_dedupes_overlapping_tail() -> None:
assert session.messages[2]["tool_call_id"] == "call_pending"
@pytest.mark.asyncio
async def test_runtime_checkpoint_keeps_provider_state_out_of_public_metadata(
tmp_path: Path,
) -> None:
loop = _make_full_loop(tmp_path)
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={
"items": [
{
"type": "reasoning",
"encrypted_content": "private-checkpoint-blob",
}
]
},
)
loop.provider.can_resume_conversation_state.return_value = True
loop.provider.chat_with_retry = AsyncMock(
return_value=LLMResponse(content="done", provider_state=state)
)
session = loop.sessions.get_or_create("cli:private-checkpoint")
await loop._run_agent_loop(
[
{"role": "system", "content": "system"},
{"role": "user", "content": "question"},
],
runtime=loop.llm_runtime(),
session=session,
)
assert session.provider_state is not None
checkpoint = session.metadata[AgentLoop._RUNTIME_CHECKPOINT_KEY]
assert "provider_state" not in checkpoint
assert checkpoint[AgentLoop._PROVIDER_STATE_CHECKPOINT_VERSION_KEY] == (
AgentLoop._PROVIDER_STATE_CHECKPOINT_VERSION
)
assert "private-checkpoint-blob" not in json.dumps(session.metadata)
public_payload = loop.sessions.read_session_file(session.key)
assert public_payload is not None
assert "private-checkpoint-blob" not in json.dumps(public_payload)
raw = loop.sessions._get_session_path(session.key).read_text(encoding="utf-8")
assert "private-checkpoint-blob" in raw
@pytest.mark.asyncio
async def test_process_message_persists_user_message_before_turn_completes(tmp_path: Path) -> None:
loop = _make_full_loop(tmp_path)
@@ -634,6 +793,150 @@ async def test_process_message_persists_user_message_before_turn_completes(tmp_p
assert persisted.updated_at >= persisted.created_at
@pytest.mark.asyncio
async def test_subagent_followup_stages_provider_state_before_turn_runs(
tmp_path: Path,
) -> None:
loop = _make_full_loop(tmp_path)
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
loop._run_agent_loop = AsyncMock(side_effect=RuntimeError("boom")) # type: ignore[method-assign]
loop.provider.can_resume_conversation_state.return_value = True
session = loop.sessions.get_or_create("cli:subagent-crash")
session.provider_state = _provider_state()
loop.sessions.save(session)
msg = InboundMessage(
channel="system",
sender_id="subagent",
chat_id="cli:subagent-crash",
content="subagent result",
metadata={"subagent_task_id": "sub-1"},
)
with pytest.raises(RuntimeError, match="boom"):
await loop._process_message(msg)
loop.sessions.invalidate("cli:subagent-crash")
persisted = loop.sessions.get_or_create("cli:subagent-crash")
assert persisted.messages[-1]["content"] == "subagent result"
assert persisted.provider_state is not None
assert persisted.provider_state.pending_messages[-1]["role"] == "user"
assert persisted.provider_state.pending_messages[-1]["content"] == "subagent result"
@pytest.mark.asyncio
async def test_subagent_followup_state_is_durable_before_prompt_assembly(
tmp_path: Path,
) -> None:
loop = _make_full_loop(tmp_path)
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
loop.provider.can_resume_conversation_state.return_value = True
loop._build_initial_messages = MagicMock( # type: ignore[method-assign]
side_effect=RuntimeError("prompt boom"),
)
session = loop.sessions.get_or_create("cli:subagent-prompt-crash")
session.provider_state = _provider_state()
loop.sessions.save(session)
msg = InboundMessage(
channel="system",
sender_id="subagent",
chat_id="cli:subagent-prompt-crash",
content="subagent result",
metadata={"subagent_task_id": "sub-1"},
)
with pytest.raises(RuntimeError, match="prompt boom"):
await loop._process_message(msg)
loop.sessions.invalidate("cli:subagent-prompt-crash")
persisted = loop.sessions.get_or_create("cli:subagent-prompt-crash")
assert persisted.messages[-1]["content"] == "subagent result"
assert persisted.provider_state is not None
assert persisted.provider_state.pending_messages[-1]["content"] == (
"subagent result"
)
@pytest.mark.asyncio
async def test_subagent_redelivery_does_not_duplicate_staged_provider_input(
tmp_path: Path,
) -> None:
loop = _make_full_loop(tmp_path)
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
loop.provider.can_resume_conversation_state.return_value = True
build_initial_messages = loop._build_initial_messages
loop._build_initial_messages = MagicMock( # type: ignore[method-assign]
side_effect=RuntimeError("prompt boom"),
)
session = loop.sessions.get_or_create("cli:subagent-redelivery")
session.provider_state = _provider_state()
loop.sessions.save(session)
msg = InboundMessage(
channel="system",
sender_id="subagent",
chat_id="cli:subagent-redelivery",
content="subagent result",
metadata={"subagent_task_id": "sub-1"},
)
with pytest.raises(RuntimeError, match="prompt boom"):
await loop._process_message(msg)
loop.sessions.invalidate("cli:subagent-redelivery")
persisted = loop.sessions.get_or_create("cli:subagent-redelivery")
assert persisted.provider_state is not None
assert [
message.get("content")
for message in persisted.provider_state.pending_messages
].count("subagent result") == 1
loop._build_initial_messages = build_initial_messages # type: ignore[method-assign]
loop._run_agent_loop = AsyncMock( # type: ignore[method-assign]
side_effect=RuntimeError("provider boom"),
)
with pytest.raises(RuntimeError, match="provider boom"):
await loop._process_message(msg)
provider_state = loop._run_agent_loop.await_args.kwargs["provider_state"]
assert provider_state is not None
pending_results = [
message
for message in provider_state.pending_messages
if message.get("content") == "subagent result"
]
assert len(pending_results) == 1
assert LLMProvider._sanitize_empty_content(pending_results) == [
{"role": "user", "content": "subagent result"},
]
@pytest.mark.asyncio
async def test_subagent_followup_clears_state_before_compatibility_failure(
tmp_path: Path,
) -> None:
loop = _make_full_loop(tmp_path)
loop.consolidator.maybe_consolidate_by_tokens = AsyncMock(return_value=False) # type: ignore[method-assign]
loop.provider.can_resume_conversation_state.side_effect = RuntimeError(
"compatibility boom"
)
session = loop.sessions.get_or_create("cli:subagent-compat-crash")
session.provider_state = _provider_state()
loop.sessions.save(session)
msg = InboundMessage(
channel="system",
sender_id="subagent",
chat_id="cli:subagent-compat-crash",
content="subagent result",
metadata={"subagent_task_id": "sub-1"},
)
with pytest.raises(RuntimeError, match="compatibility boom"):
await loop._process_message(msg)
loop.sessions.invalidate("cli:subagent-compat-crash")
persisted = loop.sessions.get_or_create("cli:subagent-compat-crash")
assert persisted.messages[-1]["content"] == "subagent result"
assert persisted.provider_state is None
@pytest.mark.asyncio
async def test_process_message_persists_unified_session_delivery_route(tmp_path: Path) -> None:
loop = _make_full_loop(tmp_path)
@@ -1245,6 +1548,9 @@ async def test_next_turn_after_crash_closes_pending_user_turn_before_new_input(t
session = loop.sessions.get_or_create("feishu:c3")
session.add_message("user", "old question")
session.metadata[AgentLoop._PENDING_USER_TURN_KEY] = True
session.provider_state = _provider_state().with_pending_messages([
{"role": "user", "content": "old question"},
])
loop.sessions.save(session)
loop._run_agent_loop = AsyncMock(return_value=(
@@ -1278,6 +1584,7 @@ async def test_next_turn_after_crash_closes_pending_user_turn_before_new_input(t
{"role": "assistant", "content": "new answer"},
]
assert AgentLoop._PENDING_USER_TURN_KEY not in session.metadata
assert session.provider_state is None
@pytest.mark.asyncio
+13 -10
View File
@@ -10,9 +10,10 @@ from unittest.mock import MagicMock
import anyio
import pytest
from mcp import MCPError
from mcp import types as mcp_types
from mcp.shared.exceptions import McpError
from mcp.shared.message import SessionMessage
from mcp.types import ErrorData
from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools import mcp as mcp_runtime
@@ -25,10 +26,12 @@ from nanobot.config.schema import MCPServerConfig
def _mcp_notification(method: str, params: dict[str, Any] | None = None) -> SessionMessage:
return SessionMessage(
message=mcp_types.JSONRPCNotification(
jsonrpc="2.0",
method=method,
params=params,
message=mcp_types.JSONRPCMessage(
mcp_types.JSONRPCNotification(
jsonrpc="2.0",
method=method,
params=params,
)
)
)
@@ -425,7 +428,7 @@ async def test_mcp_tool_reconnects_after_session_terminated(
self.call_count += 1
assert arguments == {"symbol": "AAPL"}
if self.index == 1:
raise MCPError(-32000, "Session terminated")
raise McpError(ErrorData(code=-32000, message="Session terminated"))
return SimpleNamespace(
content=[mcp_types.TextContent(type="text", text="recovered")]
)
@@ -440,7 +443,7 @@ async def test_mcp_tool_reconnects_after_session_terminated(
tool_def = SimpleNamespace(
name="quote",
description="quote tool",
input_schema={"type": "object", "properties": {}},
inputSchema={"type": "object", "properties": {}},
)
registry.register(MCPToolWrapper(session, name, tool_def, tool_timeout=5))
stack = AsyncExitStack()
@@ -481,7 +484,7 @@ async def test_mcp_reconnect_handler_uses_sanitized_server_prefix(
async def call_tool(self, _name: str, arguments: dict[str, Any]) -> Any:
assert arguments == {}
if self.index == 1:
raise MCPError(-32000, "Session terminated")
raise McpError(ErrorData(code=-32000, message="Session terminated"))
return SimpleNamespace(
content=[mcp_types.TextContent(type="text", text="recovered")]
)
@@ -494,7 +497,7 @@ async def test_mcp_reconnect_handler_uses_sanitized_server_prefix(
tool_def = SimpleNamespace(
name="quote",
description="quote tool",
input_schema={"type": "object", "properties": {}},
inputSchema={"type": "object", "properties": {}},
)
registry.register(MCPToolWrapper(_FakeSession(connect_count), name, tool_def))
stack = AsyncExitStack()
@@ -529,7 +532,7 @@ async def test_concurrent_mcp_reconnect_reuses_fresh_session(
class _DeadSession:
async def read_resource(self, _uri: str) -> Any:
raise MCPError(-32000, "Session terminated")
raise McpError(ErrorData(code=-32000, message="Session terminated"))
class _LiveSession:
async def read_resource(self, uri: str) -> Any:
+19 -19
View File
@@ -17,7 +17,7 @@ import socket
import time
from unittest.mock import MagicMock
import httpx2 as httpx
import httpx
import pytest
from nanobot.agent.loop import AgentLoop
@@ -27,8 +27,10 @@ from nanobot.bus.queue import MessageBus
from nanobot.config.schema import MCPServerConfig
from nanobot.security import network as security_network
_IDLE_TIMEOUT_SECONDS = 0.25
_IDLE_EXPIRY_GRACE_SECONDS = 0.25
# Leave enough headroom for reconnect handshakes on slower CI hosts; each test
# still waits beyond this deadline explicitly before exercising recovery.
_IDLE_TIMEOUT_SECONDS = 1.0
_IDLE_EXPIRY_GRACE_SECONDS = 0.5
_TOOL_TIMEOUT_SECONDS = 10
@@ -39,29 +41,31 @@ def _free_port() -> int:
def _run_mcp_server(port: int, ready_event: multiprocessing.Event) -> None:
"""MCPServer target for ``multiprocessing.Process``.
"""FastMCP server target for ``multiprocessing.Process``.
The server exposes a single ``greet`` tool and terminates idle sessions
after ``_IDLE_TIMEOUT_SECONDS``.
"""
import uvicorn
from mcp.server import MCPServer
from mcp.server.fastmcp import FastMCP
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
mcp = MCPServer("IdleTimeoutDemo")
mcp = FastMCP("IdleTimeoutDemo", json_response=True, port=port)
@mcp.tool()
def greet(name: str = "World") -> str: # noqa: N802
"""Greet someone."""
return f"Hello, {name}!"
app = mcp.streamable_http_app(
json_response=True,
host="127.0.0.1",
mcp._session_manager = StreamableHTTPSessionManager(
app=mcp._mcp_server,
json_response=mcp.settings.json_response,
stateless=mcp.settings.stateless_http,
security_settings=mcp.settings.transport_security,
session_idle_timeout=_IDLE_TIMEOUT_SECONDS,
)
mcp.session_manager.session_idle_timeout = _IDLE_TIMEOUT_SECONDS
ready_event.set()
uvicorn.run(app, host="127.0.0.1", port=port, log_level="warning")
mcp.run(transport="streamable-http")
async def _wait_for_server(url: str, timeout: float = 10.0) -> bool:
@@ -126,14 +130,10 @@ def _make_loop(tmp_path, *, mcp_servers: dict) -> AgentLoop:
@pytest.fixture(autouse=True)
def allow_loopback_mcp_urls(monkeypatch: pytest.MonkeyPatch):
"""The repro server runs on 127.0.0.1; allow nanobot to talk to it."""
class TestPinnedDNSAsyncTransport(security_network.Httpx2PinnedDNSAsyncTransport):
class TestPinnedDNSAsyncTransport(security_network.PinnedDNSAsyncTransport):
_resolver_lock = asyncio.Lock()
monkeypatch.setattr(
mcp_module,
"Httpx2PinnedDNSAsyncTransport",
TestPinnedDNSAsyncTransport,
)
monkeypatch.setattr(mcp_module, "PinnedDNSAsyncTransport", TestPinnedDNSAsyncTransport)
monkeypatch.setattr(
mcp_module,
"validate_url_target",
@@ -156,7 +156,7 @@ def allow_loopback_mcp_urls(monkeypatch: pytest.MonkeyPatch):
)
monkeypatch.setattr(
mcp_module,
"httpx2_env_proxy_mounts",
"httpx_env_proxy_mounts",
lambda: {},
)
+10 -17
View File
@@ -5,8 +5,9 @@ from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
import pytest
from mcp import MCPError
from mcp import types as mcp_types
from mcp.shared.exceptions import McpError
from mcp.types import ErrorData
from nanobot.agent.tools.mcp import (
MCPPromptWrapper,
@@ -36,16 +37,12 @@ class _FakeEndOfStreamError(Exception):
_FakeEndOfStreamError.__name__ = "EndOfStream"
def _session_terminated_error() -> MCPError:
return MCPError(-32000, "Session terminated")
def _session_terminated_error() -> McpError:
return McpError(ErrorData(code=-32000, message="Session terminated"))
def _connection_closed_error() -> MCPError:
return MCPError(-32000, "Connection closed")
def _session_not_found_error() -> MCPError:
return MCPError(-32600, "Session not found")
def _connection_closed_error() -> McpError:
return McpError(ErrorData(code=-32000, message="Connection closed"))
def test_is_transient_recognizes_closed_resource():
@@ -88,10 +85,6 @@ def test_is_session_terminated_recognizes_connection_closed_mcp_error():
assert _is_session_terminated(_connection_closed_error())
def test_is_session_terminated_recognizes_v2_session_not_found_error():
assert _is_session_terminated(_session_not_found_error())
# ---------------------------------------------------------------------------
# MCPToolWrapper retry behaviour
# ---------------------------------------------------------------------------
@@ -101,7 +94,7 @@ def _make_tool_def(name="test_tool"):
return SimpleNamespace(
name=name,
description="A test tool",
input_schema={"type": "object", "properties": {}},
inputSchema={"type": "object", "properties": {}},
)
@@ -422,10 +415,10 @@ async def test_prompt_fails_after_retry_exhausted():
@pytest.mark.asyncio
async def test_prompt_no_retry_on_mcp_error():
"""MCPError (application-level) should NOT trigger retry."""
"""McpError (application-level) should NOT trigger retry."""
session = AsyncMock()
session.get_prompt = AsyncMock(
side_effect=MCPError(-1, "not found")
side_effect=McpError(ErrorData(code=-1, message="not found"))
)
wrapper = MCPPromptWrapper(session, "test_server", _make_prompt_def())
@@ -450,7 +443,7 @@ async def test_prompt_no_retry_on_non_transient():
@pytest.mark.asyncio
async def test_prompt_reconnects_on_session_terminated():
"""Prompt should reconnect once before falling back to MCPError handling."""
"""Prompt should reconnect once before falling back to McpError handling."""
old_session = AsyncMock()
old_session.get_prompt = AsyncMock(side_effect=_session_terminated_error())
new_session = AsyncMock()
+18
View File
@@ -579,3 +579,21 @@ def test_history_skips_non_dict_jsonl_lines(tmp_path: Path) -> None:
}]
next_cursor = memory.append_history("next", session_key="cli:t")
assert next_cursor == 2
def test_raw_archive_handles_none_timestamp_and_missing_role(tmp_path: Path) -> None:
"""raw_archive and _format_messages must safely format messages with None timestamp or missing role.
Prevents TypeError on NoneType[:16] slicing and KeyError on missing 'role'
when raw-dumping unconsolidated history entries without timestamps or role fields.
"""
memory = MemoryStore(tmp_path)
messages = [
{"content": "message with none timestamp", "timestamp": None, "role": "user"},
{"content": "message with int timestamp", "timestamp": 1720000000, "role": "assistant"},
{"content": "message with missing role", "timestamp": "2026-07-28T12:00:00"},
]
memory.raw_archive(messages, session_key="cli:test")
raw_history = memory.history_file.read_text(encoding="utf-8")
assert "[?] USER: message with none timestamp" in raw_history
assert "[1720000000] ASSISTANT: message with int timestamp" in raw_history
assert "[2026-07-28T12:00] UNKNOWN: message with missing role" in raw_history
+422 -1
View File
@@ -11,7 +11,13 @@ import pytest
from agent.runner_helpers import make_run_spec
from nanobot.config.schema import AgentDefaults
from nanobot.providers.base import LLMProvider, LLMResponse, ToolCallRequest
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
ToolCallRequest,
)
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
@@ -73,6 +79,311 @@ async def test_runner_preserves_reasoning_fields_and_tool_results():
)
@pytest.mark.asyncio
async def test_runner_replays_provider_state_without_chat_projection_duplicates():
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
provider.supports_native_compaction.return_value = False
captured_second_kwargs: dict = {}
checkpoints: list[dict] = []
calls = 0
async def checkpoint(payload: dict) -> None:
checkpoints.append(payload)
first_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "opaque"}]},
)
second_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "message", "role": "assistant"}]},
)
async def chat_with_retry(**kwargs):
nonlocal calls
calls += 1
if calls == 1:
provider_context = kwargs["provider_context"]
assert isinstance(provider_context, ProviderCallContext)
assert provider_context.conversation_state is None
return LLMResponse(
content=None,
tool_calls=[
ToolCallRequest(
id="call_1|fc_1",
name="list_dir",
arguments={"path": "."},
),
],
provider_state=first_state,
)
captured_second_kwargs.update(kwargs)
return LLMResponse(content="done", provider_state=second_state)
provider.chat_with_retry = chat_with_retry
tools = MagicMock()
tools.get_definitions.return_value = []
tools.execute = AsyncMock(return_value="tool result")
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[
{"role": "system", "content": "system"},
{"role": "user", "content": "do task"},
],
tools=tools,
model="gpt-5.6",
max_iterations=3,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
checkpoint_callback=checkpoint,
))
provider_context = captured_second_kwargs["provider_context"]
assert isinstance(provider_context, ProviderCallContext)
assert provider_context.conversation_state is not None
assert provider_context.conversation_state.payload == first_state.payload
assert provider_context.conversation_state.pending_messages == [{
"role": "tool",
"tool_call_id": "call_1|fc_1",
"name": "list_dir",
"content": "tool result",
}]
assert not any(
message.get("role") == "assistant"
for message in provider_context.conversation_state.pending_messages
)
assert result.provider_state is not None
assert result.provider_state.payload == second_state.payload
assert result.provider_state.pending_messages == []
assert checkpoints[0]["phase"] == "awaiting_tools"
assert "provider_state" not in checkpoints[0]
assert checkpoints[1]["phase"] == "tools_completed"
assert checkpoints[1]["provider_state"].pending_messages == [{
"role": "tool",
"tool_call_id": "call_1|fc_1",
"name": "list_dir",
"content": "tool result",
}]
assert checkpoints[2]["phase"] == "final_response"
assert checkpoints[2]["provider_state"].payload == second_state.payload
@pytest.mark.asyncio
async def test_runner_governs_tool_result_before_adding_it_to_provider_state():
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
provider.supports_native_compaction.return_value = False
calls = 0
captured_context: ProviderCallContext | None = None
checkpoints: list[dict] = []
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "opaque"}]},
)
async def chat_with_retry(**kwargs):
nonlocal calls, captured_context
calls += 1
if calls == 1:
return LLMResponse(
content=None,
tool_calls=[
ToolCallRequest(
id="call_1",
name="read_file",
arguments={"path": "large.txt"},
),
],
provider_state=state,
)
captured_context = kwargs["provider_context"]
return LLMResponse(content="done")
provider.chat_with_retry = chat_with_retry
tools = MagicMock()
tools.get_definitions.return_value = []
tools.execute = AsyncMock(return_value="x" * 5_000)
async def checkpoint(payload: dict) -> None:
checkpoints.append(payload)
await AgentRunner().run(make_run_spec(
provider,
initial_messages=[
{"role": "system", "content": "system"},
{"role": "user", "content": "read the file"},
],
tools=tools,
model="gpt-5.6",
context_window_tokens=3_000,
context_block_limit=200,
max_tokens=1_000,
max_iterations=3,
max_tool_result_chars=10_000,
checkpoint_callback=checkpoint,
))
assert captured_context is not None
assert captured_context.conversation_state is not None
pending = captured_context.conversation_state.pending_messages
assert len(pending) == 1
assert pending[0]["role"] == "tool"
assert "compacted to fit context" in pending[0]["content"]
assert pending[0]["content"] != "x" * 5_000
completed_checkpoint = next(
checkpoint
for checkpoint in checkpoints
if checkpoint["phase"] == "tools_completed"
)
checkpoint_pending = completed_checkpoint["provider_state"].pending_messages
assert "compacted to fit context" in checkpoint_pending[0]["content"]
assert checkpoint_pending[0]["content"] != "x" * 5_000
@pytest.mark.asyncio
async def test_injected_final_response_checkpoint_includes_provider_state():
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
provider.supports_native_compaction.return_value = False
first_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "message", "content": "first answer"}]},
)
second_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "message", "content": "second answer"}]},
)
provider.chat_with_retry = AsyncMock(side_effect=[
LLMResponse(content="first answer", provider_state=first_state),
LLMResponse(content="second answer", provider_state=second_state),
])
tools = MagicMock()
tools.get_definitions.return_value = []
checkpoints: list[dict] = []
injections = [[{"role": "user", "content": "follow up"}], []]
async def checkpoint(payload: dict) -> None:
checkpoints.append(payload)
async def inject() -> list[dict]:
return injections.pop(0)
await AgentRunner().run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "start"}],
tools=tools,
model="gpt-5.6",
max_iterations=3,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
checkpoint_callback=checkpoint,
injection_callback=inject,
))
assert checkpoints[0]["phase"] == "final_response"
assert checkpoints[0]["provider_state"].payload == first_state.payload
@pytest.mark.asyncio
async def test_runner_preserves_last_completed_provider_state_on_model_error():
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
content="temporary upstream failure",
finish_reason="error",
error_kind="timeout",
))
tools = MagicMock()
tools.get_definitions.return_value = []
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "opaque"}]},
)
unsaved_input = {"role": "user", "content": "ephemeral follow-up"}
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[
{"role": "system", "content": "system"},
unsaved_input,
],
tools=tools,
model="gpt-5.6",
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
provider_state=state.with_pending_messages([unsaved_input]),
))
assert result.stop_reason == "error"
assert result.provider_state is not None
assert result.provider_state.payload == state.payload
assert result.provider_state.pending_messages[0] == unsaved_input
assert result.provider_state.pending_messages[1]["role"] == "assistant"
assert "model error" in result.provider_state.pending_messages[1]["content"]
@pytest.mark.asyncio
async def test_runner_discards_provider_state_on_non_retryable_model_error():
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
content="context length exceeded",
finish_reason="error",
error_status_code=400,
error_should_retry=False,
))
tools = MagicMock()
tools.get_definitions.return_value = []
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "opaque"}]},
)
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "continue"}],
tools=tools,
model="gpt-5.6",
max_iterations=1,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
provider_state=state,
))
assert result.stop_reason == "error"
assert result.provider_state is None
@pytest.mark.asyncio
async def test_runner_returns_max_iterations_fallback():
from nanobot.agent.runner import AgentRunner
@@ -422,6 +733,66 @@ async def test_runner_retries_empty_final_response_with_summary_prompt():
assert result.usage["completion_tokens"] == 9
@pytest.mark.asyncio
@pytest.mark.parametrize("finish_reason", ["refusal", "content_filter"])
async def test_runner_does_not_retry_blank_policy_terminal(
finish_reason: str,
) -> None:
from nanobot.agent.runner import AgentRunner
from nanobot.utils.runtime import EMPTY_FINAL_RESPONSE_MESSAGE
provider = MagicMock(spec=LLMProvider)
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
content=None,
finish_reason=finish_reason,
))
tools = MagicMock()
tools.get_definitions.return_value = []
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "do task"}],
tools=tools,
model="test-model",
max_iterations=3,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
assert provider.chat_with_retry.await_count == 1
assert result.final_content == EMPTY_FINAL_RESPONSE_MESSAGE
assert result.stop_reason == "empty_final_response"
@pytest.mark.asyncio
@pytest.mark.parametrize("finish_reason", ["refusal", "content_filter"])
async def test_runner_does_not_auto_continue_goal_after_policy_terminal(
finish_reason: str,
) -> None:
from nanobot.agent.runner import AgentRunner
provider = MagicMock(spec=LLMProvider)
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(
content="Request blocked by provider policy.",
finish_reason=finish_reason,
))
tools = MagicMock()
tools.get_definitions.return_value = []
result = await AgentRunner().run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "do task"}],
tools=tools,
model="test-model",
max_iterations=3,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
goal_active_predicate=lambda: True,
))
assert provider.chat_with_retry.await_count == 1
assert result.final_content == "Request blocked by provider policy."
assert result.stop_reason == "completed"
@pytest.mark.asyncio
async def test_runner_uses_specific_message_after_empty_finalization_retry():
"""After silent retries + finalization all return empty, stop_reason is empty_final_response."""
@@ -450,6 +821,56 @@ async def test_runner_uses_specific_message_after_empty_finalization_retry():
assert result.stop_reason == "empty_final_response"
@pytest.mark.asyncio
async def test_empty_finalization_retry_discards_candidate_provider_state():
from nanobot.agent.runner import AgentRunner
candidate = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={
"items": [{
"type": "function_call",
"call_id": "call_1",
"name": "exec",
"arguments": "{}",
}],
},
)
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = True
provider.chat_with_retry = AsyncMock(side_effect=[
LLMResponse(content=None, tool_calls=[], usage={}),
LLMResponse(content=None, tool_calls=[], usage={}),
LLMResponse(
content="finalized without tools",
tool_calls=[ToolCallRequest(id="call_1", name="exec", arguments={})],
finish_reason="stop",
provider_state=candidate,
usage={},
),
])
tools = MagicMock()
tools.get_definitions.return_value = []
tools.execute = AsyncMock(return_value="must not run")
runner = AgentRunner()
result = await runner.run(make_run_spec(
provider,
initial_messages=[{"role": "user", "content": "do task"}],
tools=tools,
model="test-model",
max_iterations=3,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
tools.execute.assert_not_awaited()
assert result.final_content == "finalized without tools"
assert result.provider_state is None
@pytest.mark.asyncio
async def test_runner_length_recovery_returns_all_segments():
"""Recovered output segments are returned together instead of only the tail."""
+47
View File
@@ -310,3 +310,50 @@ async def test_runner_tool_error_preserves_tool_results_in_messages():
i for i, m in enumerate(result.messages) if m.get("role") == "tool"
]
assert all(ti > asst_tc_idx for ti in tool_indices)
@pytest.mark.asyncio
async def test_length_finish_with_blank_content_routes_to_length_recovery():
"""Regression test for #5133.
A response with finish_reason='length' and blank content (e.g. the model
spent its whole output budget on a tool call whose closing tag was
truncated) must take the length-recovery path, not the empty-response
retry path. Retrying the same prompt cannot recover from output-budget
exhaustion.
"""
from nanobot.agent.runner import AgentRunner
from nanobot.utils.runtime import LENGTH_RECOVERY_PROMPT
provider = MagicMock(spec=LLMProvider)
# First call: truncated (length) with blank content and a dropped tool call.
# Second call: normal completion so the loop can terminate.
provider.chat_with_retry = AsyncMock(side_effect=[
LLMResponse(
content="",
finish_reason="length",
tool_calls=[ToolCallRequest(id="call_1", name="exec", arguments={})],
usage={},
),
LLMResponse(content="done", finish_reason="stop", tool_calls=[], usage={}),
])
tools = MagicMock()
tools.get_definitions.return_value = []
runner = AgentRunner()
result = await runner.run(make_run_spec(provider,
initial_messages=[{"role": "user", "content": "do a long task"}],
tools=tools,
model="test-model",
max_iterations=5,
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
))
# The runner must have injected a length-recovery prompt and continued,
# rather than exhausting empty-response retries into a generic apology.
user_msgs = [m.get("content") or "" for m in result.messages if m.get("role") == "user"]
assert any(LENGTH_RECOVERY_PROMPT in c for c in user_msgs), (
"expected a length-recovery message to be appended for a "
"finish_reason='length' response with blank content"
)
assert result.final_content == "done"
+284 -1
View File
@@ -9,8 +9,15 @@ import pytest
from loguru import logger
from nanobot.config.schema import ModelPresetConfig
from nanobot.providers.base import LLMProvider, LLMResponse
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
)
from nanobot.providers.conversation_state import ProviderConversationStateController
from nanobot.providers.fallback_provider import FallbackProvider
from nanobot.providers.openai_responses import resolve_compact_threshold
def _make_response(
@@ -66,6 +73,9 @@ class _FakeProvider(LLMProvider):
self._response = response or _make_response()
self.chat_calls: list[dict[str, Any]] = []
self.chat_stream_calls: list[dict[str, Any]] = []
self.context_calls: list[ProviderCallContext | None] = []
self.resumable = False
self.compact = False
def get_default_model(self) -> str:
return f"{self.name}/model"
@@ -81,6 +91,26 @@ class _FakeProvider(LLMProvider):
await on_delta(self._response.content)
return self._response
async def chat_with_context(
self,
provider_context: ProviderCallContext | None = None,
**kwargs: Any,
) -> LLMResponse:
self.context_calls.append(provider_context)
return await self.chat(**kwargs)
def can_resume_conversation_state(
self,
state: ProviderConversationState,
model: str | None = None,
) -> bool:
_ = state, model
return self.resumable
def supports_native_compaction(self, model: str | None = None) -> bool:
_ = model
return self.compact
# -- config-level tests --
@@ -211,6 +241,8 @@ def test_provider_snapshot_uses_smallest_fallback_context_window() -> None:
snapshot = build_provider_snapshot(config)
assert snapshot.context_window_tokens == 64000
assert isinstance(snapshot.provider, FallbackProvider)
assert snapshot.provider._primary_context_window_tokens == 128000
def test_inline_fallback_reasoning_effort_does_not_inherit_primary() -> None:
@@ -285,6 +317,257 @@ class TestFallbackOnPrimaryError:
assert primary.chat_calls[0]["model"] == "primary-model"
assert fallback.chat_calls[0]["model"] == "fallback-a"
@pytest.mark.asyncio
async def test_primary_compaction_uses_primary_context_window(self) -> None:
primary = _FakeProvider("primary", _make_response("primary ok"))
primary.compact = True
fb = FallbackProvider(
primary=primary,
fallback_presets=[
_fallback("small-chat", context_window_tokens=50_000),
],
provider_factory=MagicMock(),
primary_context_window_tokens=200_000,
)
await fb.chat_with_context(
messages=[{"role": "user", "content": "hi"}],
model="gpt-5.6",
max_tokens=10_000,
provider_context=ProviderCallContext(context_window_tokens=50_000),
)
primary_context = primary.context_calls[0]
assert primary_context is not None
assert primary_context.context_window_tokens == 200_000
assert resolve_compact_threshold(
primary_context.context_window_tokens,
10_000,
) == 180_000
@pytest.mark.asyncio
async def test_native_fallback_compaction_uses_its_own_context_window(self) -> None:
primary = _FakeProvider("primary", _error_response())
primary.compact = True
fallback = _FakeProvider("fallback", _make_response("fallback ok"))
fallback.compact = True
fb = FallbackProvider(
primary=primary,
fallback_presets=[
_fallback("fallback-a", context_window_tokens=120_000),
],
provider_factory=MagicMock(return_value=fallback),
primary_context_window_tokens=200_000,
)
result = await fb.chat_with_context(
messages=[{"role": "user", "content": "hi"}],
model="gpt-5.6",
provider_context=ProviderCallContext(context_window_tokens=50_000),
)
assert result.content == "fallback ok"
assert primary.context_calls == [
ProviderCallContext(context_window_tokens=200_000)
]
assert fallback.context_calls == [
ProviderCallContext(context_window_tokens=120_000)
]
@pytest.mark.asyncio
async def test_native_fallback_gets_context_when_primary_does_not_use_it(self) -> None:
primary = _FakeProvider("primary", _error_response())
fallback = _FakeProvider("fallback", _make_response("fallback ok"))
fallback.compact = True
fb = FallbackProvider(
primary=primary,
fallback_presets=[
_fallback("fallback-a", context_window_tokens=120_000),
],
provider_factory=MagicMock(return_value=fallback),
primary_context_window_tokens=200_000,
)
messages = [{"role": "user", "content": "hi"}]
controller = ProviderConversationStateController(
provider=fb,
model="primary-model",
messages=messages,
)
assert fb.supports_native_compaction("primary-model") is False
provider_context = controller.prepare_request(
messages,
context_window_tokens=50_000,
)
assert provider_context == ProviderCallContext(
context_window_tokens=50_000
)
result = await fb.chat_with_context(
messages=messages,
model="primary-model",
provider_context=provider_context,
)
assert result.content == "fallback ok"
assert primary.context_calls == [ProviderCallContext()]
assert fallback.context_calls == [
ProviderCallContext(context_window_tokens=120_000)
]
@pytest.mark.asyncio
async def test_responses_chat_fallback_responses_rebuilds_state(self) -> None:
primary = _FakeProvider("primary", _error_response())
primary.resumable = True
primary.compact = True
fallback = _FakeProvider("fallback", _make_response("fallback ok"))
messages = [{"role": "user", "content": "hi"}]
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "opaque"}]},
pending_messages=list(messages),
)
fb = FallbackProvider(
primary=primary,
fallback_presets=[_fallback("fallback-a")],
provider_factory=MagicMock(return_value=fallback),
)
controller = ProviderConversationStateController(
provider=fb,
model="gpt-5.6",
messages=messages,
state=state,
)
provider_context = controller.prepare_request(
messages,
context_window_tokens=200_000,
)
assert provider_context is not None
result = await fb.chat_with_context(
messages=messages,
model="gpt-5.6",
provider_context=provider_context,
)
assert result.content == "fallback ok"
assert primary.context_calls == [provider_context]
assert fallback.context_calls == [ProviderCallContext()]
assert fallback.chat_calls[0]["messages"] == messages
controller.observe_response(result, messages)
messages.append({"role": "assistant", "content": result.content})
assert controller.finish(messages) is None
recovered_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "recovered"}]},
)
primary._response = LLMResponse(
content="primary recovered",
provider_state=recovered_state,
)
next_turn = ProviderConversationStateController(
provider=fb,
model="gpt-5.6",
messages=messages,
)
next_context = next_turn.prepare_request(
messages,
context_window_tokens=200_000,
)
assert next_context == ProviderCallContext(context_window_tokens=200_000)
recovered = await fb.chat_with_context(
messages=messages,
model="gpt-5.6",
provider_context=next_context,
)
assert recovered.provider_state is recovered_state
assert primary.context_calls[-1] == next_context
assert primary.chat_calls[-1]["messages"] == messages
@pytest.mark.asyncio
@pytest.mark.parametrize(
("primary_error_kind", "primary_status", "primary_should_retry"),
[
("server_error", 503, True),
("authentication", 401, False),
],
ids=["transient", "authentication"],
)
async def test_final_fallback_error_uses_primary_state_disposition(
self,
primary_error_kind: str,
primary_status: int,
primary_should_retry: bool,
) -> None:
primary = _FakeProvider(
"primary",
_make_response(
"primary unavailable",
finish_reason="error",
error_kind=primary_error_kind,
error_status_code=primary_status,
error_should_retry=primary_should_retry,
),
)
primary.resumable = True
fallback = _FakeProvider(
"fallback",
_make_response(
"fallback invalid request",
finish_reason="error",
error_kind="invalid_request",
error_status_code=400,
error_should_retry=False,
),
)
messages = [{"role": "user", "content": "continue"}]
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": "opaque"}]},
pending_messages=list(messages),
)
provider = FallbackProvider(
primary=primary,
fallback_presets=[_fallback("fallback-a")],
provider_factory=MagicMock(return_value=fallback),
)
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
state=state,
)
provider_context = controller.prepare_request(
messages,
context_window_tokens=200_000,
)
assert provider_context is not None
response = await provider.chat_with_context(
messages=messages,
model="gpt-5.6",
provider_context=provider_context,
)
controller.observe_response(response, messages)
assert response.content == "fallback invalid request"
assert response.preserve_provider_state_on_error is True
restored = controller.finish(messages)
assert restored is not None
assert restored.payload == state.payload
@pytest.mark.asyncio
async def test_reports_the_fallback_model_before_its_request(self) -> None:
primary = _FakeProvider("primary", _error_response())
+14 -1
View File
@@ -15,7 +15,11 @@ from nanobot.agent.context_governance import (
)
from nanobot.agent.runner import AgentRunSpec
from nanobot.config.schema import AgentDefaults
from nanobot.providers.base import LLMResponse, ToolCallRequest
from nanobot.providers.base import (
LLMResponse,
ProviderConversationState,
ToolCallRequest,
)
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
@@ -886,6 +890,13 @@ def test_drop_malformed_tool_calls_trims_response():
"""LLM response tool_calls with a missing/empty name are dropped in place."""
from nanobot.agent.runner import AgentRunner
candidate_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "function_call", "name": None}]},
)
response = LLMResponse(
content=None,
tool_calls=[
@@ -895,9 +906,11 @@ def test_drop_malformed_tool_calls_trims_response():
ToolCallRequest(id="4", name="read_file", arguments={}),
],
finish_reason="tool_calls",
provider_state=candidate_state,
)
dropped, all_dropped, orig = AgentRunner._drop_malformed_tool_calls(response)
assert [tc.name for tc in response.tool_calls] == ["read_file"]
assert response.provider_state is None
assert response.finish_reason == "tool_calls"
assert response.should_execute_tools is True
assert dropped == 3
+132
View File
@@ -4,6 +4,7 @@ import json
from datetime import datetime
from pathlib import Path
from nanobot.providers.base import ProviderConversationState
from nanobot.session.manager import Session, SessionManager
@@ -101,6 +102,137 @@ class TestAtomicSave:
for i in range(5):
assert loaded.messages[i]["content"] == f"msg{i}"
def test_provider_state_round_trips_in_private_record_only(self, tmp_path: Path):
mgr = SessionManager(tmp_path)
secret = "encrypted-reasoning-blob"
session = Session(
key="test:provider-state",
provider_state=ProviderConversationState(
kind="openai_responses",
provider="openai:https://api.openai.com/v1",
model="gpt-5.6",
version=1,
payload={
"items": [
{
"type": "reasoning",
"encrypted_content": secret,
}
]
},
pending_messages=[{"role": "user", "content": "continue"}],
),
)
session.add_message("user", "hello")
mgr.save(session)
records = [
json.loads(line)
for line in mgr._get_session_path(session.key)
.read_text(encoding="utf-8")
.splitlines()
]
assert [record.get("_type") for record in records] == [
"metadata",
"provider_state",
None,
]
assert secret in records[1]["state"]["payload"]["items"][0]["encrypted_content"]
mgr.invalidate(session.key)
loaded = mgr.get_or_create(session.key)
assert loaded.provider_state is not None
assert loaded.provider_state.to_private_record() == session.provider_state.to_private_record()
public_payload = mgr.read_session_file(session.key)
assert public_payload is not None
assert public_payload["messages"] == [session.messages[0]]
assert secret not in json.dumps(public_payload)
assert secret not in json.dumps(mgr.list_sessions())
def test_provider_state_does_not_consume_list_preview_budget(
self,
tmp_path: Path,
monkeypatch,
):
import nanobot.session.manager as session_manager
monkeypatch.setattr(session_manager, "_SESSION_LIST_PREVIEW_MAX_CHARS", 100)
mgr = SessionManager(tmp_path)
session = Session(
key="test:provider-state-preview",
provider_state=ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": [{"encrypted_content": "x" * 200}]},
),
)
session.add_message("user", "visible preview")
mgr.save(session)
assert mgr.list_sessions()[0]["preview"] == "visible preview"
def test_clear_and_fork_discard_provider_state(self, tmp_path: Path):
mgr = SessionManager(tmp_path)
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": []},
)
source = Session(key="test:state-source", provider_state=state)
source.add_message("user", "hello")
mgr.save(source)
fork = mgr.fork_session_before_user_index(
source.key,
"test:state-fork",
1,
)
assert fork is not None
assert fork.provider_state is None
source.clear()
assert source.provider_state is None
def test_invalid_provider_state_record_is_not_public_history(self, tmp_path: Path):
mgr = SessionManager(tmp_path)
path = mgr._get_session_path("test:bad-provider-state")
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(
"\n".join(
[
json.dumps(
{
"_type": "metadata",
"key": "test:bad-provider-state",
"created_at": datetime.now().isoformat(),
"updated_at": datetime.now().isoformat(),
"metadata": {},
"last_consolidated": 0,
}
),
json.dumps(
{
"_type": "provider_state",
"state": {"kind": "openai_responses"},
}
),
json.dumps({"role": "user", "content": "safe"}),
]
)
+ "\n",
encoding="utf-8",
)
loaded = mgr._load("test:bad-provider-state")
assert loaded is not None
assert loaded.provider_state is None
assert loaded.messages == [{"role": "user", "content": "safe"}]
class TestRepairCorruptFile:
def _write_corrupt_jsonl(self, path: Path, lines: list[str]) -> None:
@@ -0,0 +1,53 @@
from __future__ import annotations
import asyncio
import gc
from unittest.mock import MagicMock
import pytest
def _make_loop(loop_factory):
provider = MagicMock()
provider.get_default_model.return_value = "test-model"
return loop_factory(provider=provider)
def test_idle_agent_session_locks_are_released(loop_factory):
loop = _make_loop(loop_factory)
for index in range(1000):
lock = loop._get_session_lock(f"api:temporary-{index}")
del lock
gc.collect()
assert len(loop._session_locks) == 0
@pytest.mark.asyncio
async def test_waiter_keeps_agent_session_lock_alive(loop_factory):
loop = _make_loop(loop_factory)
owner_lock = loop._get_session_lock("api:shared")
await owner_lock.acquire()
waiter_started = asyncio.Event()
waiter_entered = asyncio.Event()
async def wait_for_lock() -> None:
lock = loop._get_session_lock("api:shared")
waiter_started.set()
async with lock:
waiter_entered.set()
waiter = asyncio.create_task(wait_for_lock())
await waiter_started.wait()
assert loop._get_session_lock("api:shared") is owner_lock
assert not waiter_entered.is_set()
owner_lock.release()
await waiter
del owner_lock
gc.collect()
assert "api:shared" not in loop._session_locks
+21 -2
View File
@@ -1,3 +1,4 @@
from nanobot.providers.base import ProviderConversationState
from nanobot.runtime_context import (
RUNTIME_CONTEXT_HISTORY_META,
RuntimeContextBlock,
@@ -769,7 +770,16 @@ def test_get_history_extend_to_user_keeps_newer_user_inside_window():
def test_retain_recent_legal_suffix_returns_dropped_messages():
"""retain_recent_legal_suffix returns the actually-dropped messages."""
session = Session(key="test:return-dropped")
session = Session(
key="test:return-dropped",
provider_state=ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": []},
),
)
for i in range(10):
session.messages.append({"role": "user", "content": f"msg{i}"})
@@ -779,11 +789,19 @@ def test_retain_recent_legal_suffix_returns_dropped_messages():
assert [m["content"] for m in result.dropped] == [f"msg{i}" for i in range(6)]
assert len(session.messages) == 4
assert result.already_consolidated_count == 0
assert session.provider_state is None
def test_retain_recent_legal_suffix_returns_empty_when_no_drop():
"""No messages dropped → empty list returned."""
session = Session(key="test:no-drop")
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": []},
)
session = Session(key="test:no-drop", provider_state=state)
for i in range(3):
session.messages.append({"role": "user", "content": f"msg{i}"})
@@ -792,6 +810,7 @@ def test_retain_recent_legal_suffix_returns_empty_when_no_drop():
assert result.dropped == []
assert result.already_consolidated_count == 0
assert len(session.messages) == 3
assert session.provider_state is state
def test_retain_recent_legal_suffix_returns_all_on_zero():
+56
View File
@@ -55,6 +55,62 @@ class TestHandleStop:
out = await cmd_stop(ctx)
assert "No active task" in out.content
@pytest.mark.asyncio
async def test_close_mcp_cancels_active_turn_before_resources(self):
loop, _bus = _make_loop()
events: list[str] = []
async def active_turn():
try:
await asyncio.sleep(60)
except asyncio.CancelledError:
events.append("turn_cancelled")
raise
task = asyncio.create_task(active_turn())
await asyncio.sleep(0)
loop._active_tasks["test:c1"] = {task}
async def close_subagents():
events.append("resources_closed")
loop.subagents.close = close_subagents
loop._exec_session_manager.close_all = AsyncMock()
with patch("nanobot.agent.loop.agent_context.close_mcp", AsyncMock()):
await loop.close_mcp()
assert events == ["turn_cancelled", "resources_closed"]
assert task.cancelled()
@pytest.mark.asyncio
async def test_close_mcp_serializes_duplicate_cleanup(self):
loop, _bus = _make_loop()
entered = asyncio.Event()
release = asyncio.Event()
concurrent = 0
max_concurrent = 0
async def close_subagents():
nonlocal concurrent, max_concurrent
concurrent += 1
max_concurrent = max(max_concurrent, concurrent)
entered.set()
await release.wait()
concurrent -= 1
loop.subagents.close = close_subagents
loop._exec_session_manager.close_all = AsyncMock()
with patch("nanobot.agent.loop.agent_context.close_mcp", AsyncMock()):
first = asyncio.create_task(loop.close_mcp())
await entered.wait()
second = asyncio.create_task(loop.close_mcp())
await asyncio.sleep(0)
assert not second.done()
release.set()
await asyncio.gather(first, second)
assert max_concurrent == 1
@pytest.mark.asyncio
async def test_stop_cancels_active_task(self):
from nanobot.bus.events import InboundMessage
+3
View File
@@ -504,6 +504,7 @@ async def test_drain_pending_blocks_while_subagents_running(tmp_path):
usage={},
had_injections=False,
tools_used=[],
provider_state=None,
)
loop.runner.run = AsyncMock(side_effect=fake_runner_run)
@@ -589,6 +590,7 @@ async def test_drain_pending_no_block_when_no_subagents(tmp_path):
usage={},
had_injections=False,
tools_used=[],
provider_state=None,
)
loop.runner.run = AsyncMock(side_effect=fake_runner_run)
@@ -638,6 +640,7 @@ async def test_drain_pending_timeout(tmp_path):
usage={},
had_injections=False,
tools_used=[],
provider_state=None,
)
loop.runner.run = AsyncMock(side_effect=fake_runner_run)
+43
View File
@@ -83,6 +83,49 @@ async def test_handle_message_dm_sends_pairing_code(monkeypatch) -> None:
assert msg.metadata.get("_pairing_code") == "ABCD-EFGH"
@pytest.mark.asyncio
async def test_dm_during_transient_store_failure_keeps_approvals(
tmp_path, monkeypatch
) -> None:
"""An unapproved DM while pairing.json is unreadable must not wipe approvals.
The pairing store treated a transient OSError like corruption and returned
an empty store; the DM pairing path then persisted that empty view,
erasing every approved sender.
"""
import builtins
from pathlib import Path
from nanobot.pairing import store
path = tmp_path / "pairing.json"
monkeypatch.setattr(store, "_store_path", lambda: path)
code = store.generate_code("dummy", "friend")
store.approve_code(code)
channel = _DummyChannel({"allowFrom": []}, MessageBus())
real_open = builtins.open
def flaky_open(file, mode="r", *args, **kwargs):
try:
same = Path(file) == path
except TypeError:
same = False
if same and "r" in mode and "+" not in mode:
raise PermissionError(13, "temporarily locked", str(path))
return real_open(file, mode, *args, **kwargs)
with monkeypatch.context() as m:
m.setattr(builtins, "open", flaky_open)
await channel._handle_message(
sender_id="stranger", chat_id="chat1", content="hello", is_dm=True
)
assert channel._sent == []
assert store.is_approved("dummy", "friend") is True
@pytest.mark.asyncio
async def test_handle_message_group_ignores_unknown() -> None:
channel = _DummyChannel({"allowFrom": []}, MessageBus())
+66 -1
View File
@@ -2479,6 +2479,69 @@ def test_optional_features_payload_preserves_legacy_flat_feishu_config(monkeypat
assert "instances" not in saved
@pytest.mark.parametrize(
"index_url",
[
"",
"https://mirror.example/simple",
],
)
def test_enable_uses_uv_when_tool_environment_has_no_pip(
monkeypatch,
index_url,
):
from nanobot import optional_features
calls: list[list[str]] = []
call_envs: list[dict[str, str] | None] = []
def _run(
argv: list[str],
*,
env: dict[str, str] | None = None,
) -> subprocess.CompletedProcess[str]:
calls.append(argv)
call_envs.append(env)
if len(calls) == 1:
return subprocess.CompletedProcess(argv, 1, stdout="", stderr="No module named pip")
if argv[0] == "uv":
return subprocess.CompletedProcess(argv, 0, stdout="", stderr="")
return subprocess.CompletedProcess(
argv,
1,
stdout="",
stderr="No module named ensurepip",
)
monkeypatch.setattr("shutil.which", lambda name: "uv" if name == "uv" else None)
monkeypatch.setenv("HTTPS_PROXY", "http://proxy.example:8080")
monkeypatch.delenv("UV_INDEX_URL", raising=False)
if index_url:
monkeypatch.setenv("PIP_INDEX_URL", index_url)
else:
monkeypatch.delenv("PIP_INDEX_URL", raising=False)
assert optional_features.install_extra("feishu", ["lark-oapi>=1.5.0"], runner=_run).ok is True
assert calls == [
[sys.executable, "-m", "pip", "install", "lark-oapi>=1.5.0"],
[
"uv",
"pip",
"install",
"--python",
sys.executable,
"lark-oapi>=1.5.0",
],
]
assert call_envs[0] is None
assert call_envs[1] is not None
assert call_envs[1]["HTTPS_PROXY"] == "http://proxy.example:8080"
if index_url:
assert call_envs[1]["UV_INDEX_URL"] == index_url
else:
assert "UV_INDEX_URL" not in call_envs[1]
def test_enable_bootstraps_pip_with_ensurepip(monkeypatch):
from nanobot import optional_features
@@ -2490,6 +2553,8 @@ def test_enable_bootstraps_pip_with_ensurepip(monkeypatch):
return subprocess.CompletedProcess(argv, 1, stdout="", stderr="No module named pip")
return subprocess.CompletedProcess(argv, 0, stdout="", stderr="")
monkeypatch.setattr("shutil.which", lambda _name: None)
assert optional_features.install_extra("bedrock", None, runner=_run).ok is True
assert calls == [
[sys.executable, "-m", "pip", "install", "nanobot-ai[bedrock]"],
@@ -2558,7 +2623,7 @@ def test_optional_dependency_metadata_for_enable():
):
assert not any(dep.startswith(dep_name) for dep in required)
for dependency in (
"tzdata>=2025.2; sys_platform == 'win32'",
"tzdata>=2025.2",
"defusedxml>=0.7.1,<1.0.0",
"pypdf>=5.0.0,<6.0.0",
"python-docx>=1.1.0,<2.0.0",
+57
View File
@@ -1160,6 +1160,63 @@ def test_config_falls_back_to_vllm_when_ollama_not_configured():
assert config.get_api_base() == "http://localhost:8000"
def test_config_cloud_nemotron_is_not_hijacked_by_unconfigured_ollama():
"""`nvidia/nemotron-*` via a gateway must not route to Ollama when no
Ollama endpoint is configured. Ollama keeps "nemotron" in its keywords
for bare-model auto-routing (PR #1863), which previously hijacked
cloud-hosted nemotron variants and silently sent traffic to
http://localhost:11434/v1."""
config = Config.model_validate(
{
"agents": {
"defaults": {
"provider": "auto",
"model": "nvidia/nemotron-3-super-120b-a12b",
}
},
"providers": {"openrouter": {"apiKey": "sk-or-test"}},
}
)
assert config.get_provider_name() == "openrouter"
assert config.get_api_base() == "https://openrouter.ai/api/v1"
def test_config_bare_nemotron_still_auto_routes_to_configured_ollama():
"""Preserves PR #1863 intent: when the user has actually configured an
Ollama endpoint, a bare nemotron model still auto-routes there."""
config = Config.model_validate(
{
"agents": {"defaults": {"provider": "auto", "model": "nemotron-3-nano"}},
"providers": {"ollama": {"apiBase": "http://localhost:11434/v1"}},
}
)
assert config.get_provider_name() == "ollama"
assert config.get_api_base() == "http://localhost:11434/v1"
def test_config_cloud_nemotron_is_not_hijacked_by_configured_ollama():
"""An explicit cloud namespace takes precedence over local keywords."""
config = Config.model_validate(
{
"agents": {
"defaults": {
"provider": "auto",
"model": "nvidia/nemotron-3-super-120b-a12b",
}
},
"providers": {
"ollama": {"apiBase": "http://localhost:11434/v1"},
"openrouter": {"apiKey": "sk-or-test"},
},
}
)
assert config.get_provider_name() == "openrouter"
assert config.get_api_base() == "https://openrouter.ai/api/v1"
def test_openai_compat_provider_passes_model_through():
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
+186
View File
@@ -0,0 +1,186 @@
"""Regression tests for gateway runtime resource teardown on stop.
Covers the lifecycle contract of ``_close_gateway_runtime``: runtime tasks
(including the agent loop and in-flight turns) are cancelled and awaited --
bounded -- before exec sessions, subagents, and MCP servers are closed, the
close is deterministic and idempotent, and a stuck or failing cleanup cannot
block the stop.
"""
import asyncio
import time
from contextlib import suppress
from nanobot.cli.gateway_runtime import _close_gateway_runtime
class _FakeAgent:
def __init__(self, events: list[str] | None = None) -> None:
self.close_calls = 0
self.events = events if events is not None else []
self.hang_on_close = False
self.raise_on_close = False
self.background: asyncio.Task[None] | None = None
async def close_mcp(self) -> None:
self.close_calls += 1
if self.hang_on_close:
await asyncio.sleep(3600)
if self.raise_on_close:
raise RuntimeError("cleanup exploded")
if self.background is not None:
await self.background
self.events.append("close_mcp")
class _FakeChannels:
def __init__(self) -> None:
self.stopped = 0
self.events: list[str] = []
async def stop_all(self) -> None:
self.stopped += 1
self.events.append("channels_stopped")
async def _cancellable_task(events: list[str]) -> None:
try:
await asyncio.sleep(3600)
except asyncio.CancelledError:
events.append("cancelled")
raise
async def _stubborn_task(events: list[str]) -> None:
"""Task that swallows cancellation and keeps running."""
try:
while True:
await asyncio.sleep(3600)
except asyncio.CancelledError:
events.append("swallowed")
await asyncio.sleep(3600)
async def test_runtime_tasks_cancelled_before_resources_closed() -> None:
events: list[str] = []
agent = _FakeAgent(events)
channels = _FakeChannels()
task = asyncio.create_task(_cancellable_task(events))
await asyncio.sleep(0) # let the task start (cancellation pre-start skips its body)
await _close_gateway_runtime(agent, channels, [task], None)
assert events == ["cancelled", "close_mcp"] # cancel happens before close
assert channels.stopped == 1
assert agent.close_calls == 1
assert task.cancelled()
async def test_pending_background_work_is_drained_before_close_returns() -> None:
agent = _FakeAgent()
channels = _FakeChannels()
done: dict[str, bool] = {"done": False}
async def background_work() -> None:
await asyncio.sleep(0.01)
done["done"] = True
agent.background = asyncio.create_task(background_work())
await _close_gateway_runtime(agent, channels, [], None)
assert done["done"] is True
assert agent.close_calls == 1
async def test_stubborn_task_does_not_block_past_wait_timeout() -> None:
agent = _FakeAgent()
channels = _FakeChannels()
events: list[str] = []
task = asyncio.create_task(_stubborn_task(events))
await asyncio.sleep(0) # let the task start (cancellation pre-start skips its body)
runtime_tasks = asyncio.gather(task)
start = time.monotonic()
await _close_gateway_runtime(
agent,
channels,
[task],
runtime_tasks,
task_wait_timeout=0.05,
)
elapsed = time.monotonic() - start
for _ in range(10):
await asyncio.sleep(0) # let the swallowed cancellation handler run
assert "swallowed" in events # task was cancelled, then refused to die
assert task.done() # the timed-out task received a second cancellation
assert runtime_tasks.done()
assert agent.close_calls == 1 # resources still closed underneath it
assert elapsed < 1.0 # bounded, not held open by the stubborn task
async def test_hanging_close_is_bounded_and_does_not_raise() -> None:
agent = _FakeAgent()
agent.hang_on_close = True
channels = _FakeChannels()
start = time.monotonic()
await _close_gateway_runtime(agent, channels, [], None, close_timeout=0.05)
elapsed = time.monotonic() - start
assert agent.close_calls == 1
assert channels.stopped == 1
assert elapsed < 1.0
async def test_failing_close_is_logged_but_shutdown_proceeds() -> None:
agent = _FakeAgent()
agent.raise_on_close = True
channels = _FakeChannels()
await _close_gateway_runtime(agent, channels, [], None)
assert agent.close_calls == 1
assert channels.stopped == 1 # teardown continued past the failure
async def test_duplicate_cleanup_is_idempotent() -> None:
agent = _FakeAgent()
channels = _FakeChannels()
task = asyncio.create_task(_cancellable_task([]))
await _close_gateway_runtime(agent, channels, [task], None)
await _close_gateway_runtime(agent, channels, [task], None)
assert agent.close_calls == 2 # second pass is a clean no-op
assert channels.stopped == 2
assert task.cancelled()
async def test_finished_runtime_tasks_gather_is_retrieved() -> None:
agent = _FakeAgent()
channels = _FakeChannels()
finished = asyncio.get_running_loop().create_future()
finished.set_result(None)
runtime_tasks = asyncio.gather(finished)
await asyncio.sleep(0) # let the gather observe the finished child
await _close_gateway_runtime(agent, channels, [], runtime_tasks)
assert runtime_tasks.done()
assert agent.close_calls == 1
async def test_cancelled_runtime_tasks_gather_does_not_raise() -> None:
agent = _FakeAgent()
channels = _FakeChannels()
runtime_tasks = asyncio.gather(asyncio.sleep(3600))
runtime_tasks.cancel()
await _close_gateway_runtime(agent, channels, [], runtime_tasks)
with suppress(asyncio.CancelledError):
await runtime_tasks # settle the cancelled gather without raising
assert runtime_tasks.done() # the cancelled gather was awaited without raising
assert agent.close_calls == 1
+30
View File
@@ -1,4 +1,8 @@
import json
import os
import subprocess
import sys
import textwrap
import warnings
import pytest
@@ -42,6 +46,32 @@ def test_agent_timezone_rejects_unknown_iana_name() -> None:
Config.model_validate({"agents": {"defaults": {"timezone": "Not/AZone"}}})
def test_agent_timezones_use_packaged_data_without_system_database() -> None:
script = textwrap.dedent(
"""\
from zoneinfo import TZPATH
from nanobot.config.schema import Config
assert not TZPATH
for name in ("UTC", "Asia/Shanghai"):
config = Config.model_validate({"agents": {"defaults": {"timezone": name}}})
serialized = config.model_dump(mode="json", by_alias=True)
restored = Config.model_validate(serialized)
assert restored.agents.defaults.timezone == name
"""
)
result = subprocess.run(
[sys.executable, "-c", script],
env=os.environ | {"PYTHONTZPATH": ""},
capture_output=True,
text=True,
check=False,
)
assert result.returncode == 0, result.stderr
def test_provider_api_type_accepts_exact_values_only() -> None:
config = Config.model_validate({
"providers": {
+138
View File
@@ -141,6 +141,33 @@ def test_add_job_accepts_valid_timezone(tmp_path) -> None:
assert job.state.next_run_at_ms is not None
@pytest.mark.parametrize("expr", [None, "", " "])
def test_add_job_rejects_missing_cron_expression(tmp_path, expr: str | None) -> None:
service = CronService(tmp_path / "cron" / "jobs.json")
with pytest.raises(ValueError, match="requires a non-empty 'expr'"):
service.add_job(
name="missing expression",
schedule=CronSchedule(kind="cron", expr=expr),
message="hello",
)
assert service.list_jobs(include_disabled=True) == []
def test_add_job_rejects_invalid_cron_expression_before_persisting(tmp_path) -> None:
service = CronService(tmp_path / "cron" / "jobs.json")
with pytest.raises(ValueError, match="invalid cron expression"):
service.add_job(
name="bad expression",
schedule=CronSchedule(kind="cron", expr="not a cron expression"),
message="hello",
)
assert service.list_jobs(include_disabled=True) == []
def test_write_run_record_uses_cron_runs_dir(tmp_path) -> None:
service = CronService(tmp_path / "cron" / "jobs.json")
@@ -600,6 +627,117 @@ async def test_run_job_preserves_running_service_state(tmp_path) -> None:
service.stop()
@pytest.mark.asyncio
async def test_manual_run_persists_completion_when_callback_lists_jobs(tmp_path) -> None:
store_path = tmp_path / "cron" / "jobs.json"
async def on_job(_job) -> None:
service.list_jobs(include_disabled=True)
await asyncio.sleep(0)
service = CronService(store_path, on_job=on_job)
job = service.add_job(
name="manual",
schedule=CronSchedule(kind="every", every_ms=60_000),
message="hello",
**_bound_chat(),
)
assert await service.run_job(job.id) is True
state = json.loads(store_path.read_text())["jobs"][0]["state"]
assert state["lastStatus"] == "ok"
assert state["lastError"] is None
assert len(state["runHistory"]) == 1
assert state["runHistory"][0]["status"] == "ok"
@pytest.mark.asyncio
async def test_overlapping_manual_runs_preserve_stopped_service_state(tmp_path) -> None:
store_path = tmp_path / "cron" / "jobs.json"
entered = [asyncio.Event(), asyncio.Event()]
release = [asyncio.Event(), asyncio.Event()]
call_count = 0
async def on_job(_job) -> None:
nonlocal call_count
call_index = call_count
call_count += 1
entered[call_index].set()
await release[call_index].wait()
service = CronService(store_path, on_job=on_job)
jobs = [
service.add_job(
name=f"manual-{index}",
schedule=CronSchedule(kind="every", every_ms=60_000),
message="hello",
**_bound_chat(str(index)),
)
for index in range(2)
]
first = asyncio.create_task(service.run_job(jobs[0].id))
await entered[0].wait()
second = asyncio.create_task(service.run_job(jobs[1].id))
try:
await entered[1].wait()
release[0].set()
assert await first is True
assert service._running is False
release[1].set()
assert await second is True
assert service._running is False
assert service._timer_task is None
states = {
item["name"]: item["state"]
for item in json.loads(store_path.read_text())["jobs"]
}
assert states["manual-0"]["lastStatus"] == "ok"
assert states["manual-1"]["lastStatus"] == "ok"
finally:
release[0].set()
release[1].set()
await asyncio.gather(first, second, return_exceptions=True)
service.stop()
@pytest.mark.asyncio
async def test_manual_run_does_not_restart_service_stopped_during_execution(tmp_path) -> None:
store_path = tmp_path / "cron" / "jobs.json"
entered = asyncio.Event()
release = asyncio.Event()
async def on_job(_job) -> None:
entered.set()
await release.wait()
service = CronService(store_path, on_job=on_job)
job = service.add_job(
name="manual-stop",
schedule=CronSchedule(kind="every", every_ms=60_000),
message="hello",
**_bound_chat(),
)
await service.start()
run = asyncio.create_task(service.run_job(job.id))
try:
await entered.wait()
service.stop()
release.set()
assert await run is True
assert service._running is False
assert service._timer_task is None
finally:
release.set()
await asyncio.gather(run, return_exceptions=True)
service.stop()
@pytest.mark.asyncio
async def test_running_service_honors_external_disable(tmp_path) -> None:
store_path = tmp_path / "cron" / "jobs.json"
+63
View File
@@ -323,3 +323,66 @@ def test_pending_gc_drops_malformed_entries(tmp_path, monkeypatch):
)
monkeypatch.setattr(store, "_store_path", lambda: path)
assert store.list_pending() == []
def _fail_reads_of(monkeypatch, path):
"""Make reads of *path* raise like a transiently locked/busy file."""
import builtins
from pathlib import Path
real_open = builtins.open
def flaky_open(file, mode="r", *args, **kwargs):
try:
same = Path(file) == path
except TypeError:
same = False
if same and "r" in mode and "+" not in mode:
raise PermissionError(13, "temporarily locked", str(path))
return real_open(file, mode, *args, **kwargs)
monkeypatch.setattr(builtins, "open", flaky_open)
class TestTransientReadFailure:
"""A transient I/O failure is not corruption and must never wipe the store."""
def test_generate_code_does_not_wipe_approvals(self, tmp_path, monkeypatch):
"""An unapproved DM during a read blip previously erased every approval.
_load treated OSError like corruption and returned an empty store;
generate_code then unconditionally saved it, overwriting pairing.json
with no approved senders.
"""
code = store.generate_code("telegram", "123")
store.approve_code(code)
with monkeypatch.context() as m:
_fail_reads_of(m, store._store_path())
with pytest.raises(OSError):
store.generate_code("telegram", "stranger")
assert store.is_approved("telegram", "123") is True
def test_reads_fail_closed_without_crashing(self, tmp_path, monkeypatch):
code = store.generate_code("telegram", "123")
store.approve_code(code)
with monkeypatch.context() as m:
_fail_reads_of(m, store._store_path())
assert store.is_approved("telegram", "123") is False
assert store.list_pending() == []
assert store.get_approved("telegram") == []
assert store.is_approved("telegram", "123") is True
def test_approve_command_reports_store_unavailable(self, tmp_path, monkeypatch):
"""/pairing approve must fail loudly instead of claiming the code is invalid."""
code = store.generate_code("telegram", "123")
with monkeypatch.context() as m:
_fail_reads_of(m, store._store_path())
reply = store.handle_pairing_command("telegram", f"approve {code}")
assert "unavailable" in reply.lower()
assert store.approve_code(code) == ("telegram", "123")
+72 -8
View File
@@ -4,6 +4,8 @@ from __future__ import annotations
from unittest.mock import patch
import pytest
from nanobot.providers.anthropic_provider import AnthropicProvider
@@ -65,17 +67,24 @@ def test_none_does_not_enable_thinking() -> None:
assert kw["temperature"] == 0.7
def test_empty_effort_does_not_enable_thinking() -> None:
kw = _build(_make_provider(), "")
assert "thinking" not in kw
assert kw["temperature"] == 0.7
def test_opus_4_7_omits_temperature_adaptive() -> None:
kw = _build(_make_provider("claude-opus-4-7"), "adaptive")
assert "temperature" not in kw
assert kw["thinking"] == {"type": "adaptive"}
def test_opus_4_7_omits_temperature_enabled() -> None:
"""Enabled thinking (high) must also omit temperature for opus-4-7."""
def test_opus_4_7_high_uses_adaptive_effort() -> None:
kw = _build(_make_provider("claude-opus-4-7"), "high", max_tokens=4096)
assert "temperature" not in kw
assert kw["thinking"]["type"] == "enabled"
assert kw["thinking"] == {"type": "adaptive"}
assert kw["output_config"] == {"effort": "high"}
assert kw["max_tokens"] == 4096
def test_opus_4_7_omits_temperature_none() -> None:
@@ -90,9 +99,11 @@ def test_opus_4_8_omits_temperature_adaptive() -> None:
assert "temperature" not in kw
def test_opus_4_8_omits_temperature_enabled() -> None:
def test_opus_4_8_high_uses_adaptive_effort() -> None:
kw = _build(_make_provider("claude-opus-4-8"), "high", max_tokens=4096)
assert "temperature" not in kw
assert kw["thinking"] == {"type": "adaptive"}
assert kw["output_config"] == {"effort": "high"}
def test_opus_4_8_omits_temperature_none() -> None:
@@ -105,9 +116,11 @@ def test_fable_omits_temperature_adaptive() -> None:
assert "temperature" not in kw
def test_fable_omits_temperature_enabled() -> None:
def test_fable_high_uses_adaptive_effort() -> None:
kw = _build(_make_provider("claude-fable-5"), "high", max_tokens=4096)
assert "temperature" not in kw
assert kw["thinking"] == {"type": "adaptive"}
assert kw["output_config"] == {"effort": "high"}
def test_fable_omits_temperature_none() -> None:
@@ -121,16 +134,67 @@ def test_sonnet_5_omits_temperature_adaptive() -> None:
assert kw["thinking"] == {"type": "adaptive"}
def test_sonnet_5_omits_temperature_enabled() -> None:
def test_sonnet_5_high_uses_adaptive_effort() -> None:
kw = _build(_make_provider("claude-sonnet-5"), "high", max_tokens=4096)
assert "temperature" not in kw
assert kw["thinking"]["type"] == "enabled"
assert kw["thinking"] == {"type": "adaptive"}
assert kw["output_config"] == {"effort": "high"}
def test_sonnet_5_omits_temperature_none() -> None:
kw = _build(_make_provider("anthropic/claude-sonnet-5"), None)
kw = _build(_make_provider("anthropic/claude-sonnet-5"), "none")
assert "temperature" not in kw
assert kw["thinking"] == {"type": "disabled"}
assert "output_config" not in kw
def test_mythos_preview_omits_temperature_but_keeps_manual_budget() -> None:
kw = _build(_make_provider("claude-mythos-preview"), "high", max_tokens=4096)
assert "temperature" not in kw
assert kw["thinking"] == {"type": "enabled", "budget_tokens": 8192}
assert "output_config" not in kw
@pytest.mark.parametrize(
"reasoning_effort", [None, "none", "adaptive", "low", "medium", "high", "xhigh", "max"]
)
def test_opus_5_omits_temperature(reasoning_effort: str | None) -> None:
kw = _build(_make_provider("claude-opus-5"), reasoning_effort)
assert "temperature" not in kw
def test_opus_5_none_disables_default_thinking() -> None:
kw = _build(_make_provider("claude-opus-5"), "none")
assert kw["thinking"] == {"type": "disabled"}
assert "output_config" not in kw
def test_opus_5_unset_preserves_provider_default() -> None:
kw = _build(_make_provider("claude-opus-5"), None)
assert "thinking" not in kw
assert "output_config" not in kw
@pytest.mark.parametrize("reasoning_effort", ["low", "medium", "high", "xhigh", "max"])
def test_opus_5_uses_adaptive_thinking_with_effort(reasoning_effort: str) -> None:
kw = _build(_make_provider("claude-opus-5"), reasoning_effort, max_tokens=4096)
assert kw["thinking"] == {"type": "adaptive"}
assert kw["output_config"] == {"effort": reasoning_effort}
assert kw["max_tokens"] == 4096
def test_dated_opus_5_model_uses_family_capabilities() -> None:
kw = _build(_make_provider("claude-opus-5-20260724"), "medium")
assert "temperature" not in kw
assert kw["thinking"] == {"type": "adaptive"}
assert kw["output_config"] == {"effort": "medium"}
def test_dated_opus_4_model_does_not_treat_date_as_minor_version() -> None:
kw = _build(_make_provider("claude-opus-4-20250514"), "high")
assert kw["temperature"] == 1.0
assert kw["thinking"] == {"type": "enabled", "budget_tokens": 8192}
assert "output_config" not in kw
def test_ordinary_model_sends_temperature() -> None:
+61 -1
View File
@@ -11,7 +11,7 @@ from nanobot.providers.azure_openai_provider import (
AzureOpenAIProvider,
_AzureTokenProvider,
)
from nanobot.providers.base import LLMResponse
from nanobot.providers.base import LLMResponse, ProviderCallContext
# ---------------------------------------------------------------------------
# Init & validation
@@ -234,6 +234,7 @@ def test_build_body_basic():
assert body["max_output_tokens"] == 4096
assert body["store"] is False
assert "reasoning" not in body
assert "include" not in body
# input should contain the converted user message only (system extracted)
assert any(
item.get("role") == "user"
@@ -241,6 +242,30 @@ def test_build_body_basic():
)
def test_build_body_enables_server_compaction():
provider = AzureOpenAIProvider(
api_key="k",
api_base="https://res.openai.azure.com",
default_model="gpt-5.6",
)
body = provider._build_body(
[{"role": "user", "content": "hello"}],
None,
None,
10_000,
0.1,
"high",
None,
provider_context=ProviderCallContext(context_window_tokens=200_000),
)
assert body["context_management"] == [{
"type": "compaction",
"compact_threshold": 180_000,
}]
def test_build_body_max_tokens_minimum():
"""max_output_tokens should never be less than 1."""
provider = AzureOpenAIProvider(api_key="k", api_base="https://r.com", default_model="gpt-4o")
@@ -358,6 +383,38 @@ async def test_chat_success():
assert result.usage["prompt_tokens"] == 10
@pytest.mark.asyncio
async def test_chat_retries_without_unsupported_server_compaction():
provider = AzureOpenAIProvider(
api_key="test-key",
api_base="https://test.openai.azure.com",
default_model="gpt-5.6",
)
class UnsupportedCompactionError(Exception):
status_code = 400
body = {"error": {"message": "Unknown parameter: context_management"}}
provider._client.responses = MagicMock()
provider._client.responses.create = AsyncMock(side_effect=[
UnsupportedCompactionError(),
_make_sdk_response(content="compaction fallback"),
])
result = await provider.chat(
[{"role": "user", "content": "Hi"}],
provider_context=ProviderCallContext(context_window_tokens=200_000),
)
create = provider._client.responses.create
assert result.content == "compaction fallback"
assert result.provider_state is not None
assert create.await_count == 2
assert "context_management" in create.call_args_list[0].kwargs
assert "context_management" not in create.call_args_list[1].kwargs
assert provider.supports_native_compaction() is False
@pytest.mark.asyncio
async def test_chat_uses_default_model():
provider = AzureOpenAIProvider(
@@ -411,6 +468,7 @@ async def test_chat_with_tool_calls():
assert len(result.tool_calls) == 1
assert result.tool_calls[0].name == "get_weather"
assert result.tool_calls[0].arguments == {"location": "SF"}
assert result.provider_state is not None
@pytest.mark.asyncio
@@ -510,6 +568,7 @@ async def test_chat_stream_with_tool_calls():
item_done.name = "get_weather"
ev_item_done = MagicMock(type="response.output_item.done", item=item_done)
resp_obj = MagicMock(status="completed")
resp_obj.model_dump.return_value = {"status": "completed", "output": []}
ev_completed = MagicMock(type="response.completed", response=resp_obj)
async def mock_stream():
@@ -527,6 +586,7 @@ async def test_chat_stream_with_tool_calls():
assert len(result.tool_calls) == 1
assert result.tool_calls[0].name == "get_weather"
assert result.tool_calls[0].arguments == {"location": "SF"}
assert result.provider_state is not None
@pytest.mark.asyncio
+291
View File
@@ -0,0 +1,291 @@
"""Tests for provider-owned conversation-state lifecycle coordination."""
from __future__ import annotations
from unittest.mock import MagicMock
import pytest
from nanobot.providers.base import (
LLMProvider,
LLMResponse,
ProviderConversationState,
ToolCallRequest,
)
from nanobot.providers.conversation_state import (
ProviderConversationStateController,
allows_conversation_message_merge,
)
def _provider(*, resumable: bool = True, compact: bool = False) -> MagicMock:
provider = MagicMock(spec=LLMProvider)
provider.can_resume_conversation_state.return_value = resumable
provider.supports_native_compaction.return_value = compact
return provider
def _state(label: str, *, pending: list[dict] | None = None) -> ProviderConversationState:
return ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={"items": [{"type": "reasoning", "encrypted_content": label}]},
pending_messages=pending or [],
)
def test_controller_replays_only_messages_after_provider_output() -> None:
provider = _provider()
messages = [
{"role": "system", "content": "system"},
{"role": "user", "content": "run a tool"},
]
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
)
state = _state("first")
controller.prepare_request(messages, context_window_tokens=200_000)
response = LLMResponse(content=None, provider_state=state)
controller.observe_response(response, messages)
assert allows_conversation_message_merge(messages[-1]) is False
messages.append(controller.project_response_message(
{
"role": "assistant",
"content": None,
"tool_calls": [{"id": "call_1", "type": "function"}],
},
response,
))
tool_message = {
"role": "tool",
"tool_call_id": "call_1",
"content": "tool result",
}
messages.append(tool_message)
provider_context = controller.prepare_request(
messages,
context_window_tokens=200_000,
)
assert provider_context is not None
assert provider_context.conversation_state is not None
assert provider_context.conversation_state.payload == state.payload
assert provider_context.conversation_state.pending_messages == [tool_message]
assert controller.checkpoint(messages).pending_messages == [tool_message]
def test_controller_uses_governed_messages_for_provider_state_delta() -> None:
provider = _provider()
messages = [
{"role": "user", "content": "run a tool"},
]
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
)
state = _state("first")
controller.prepare_request(messages, context_window_tokens=200_000)
response = LLMResponse(content=None, provider_state=state)
controller.observe_response(response, messages)
messages.extend([
controller.project_response_message(
{
"role": "assistant",
"content": None,
"tool_calls": [{"id": "call_1", "type": "function"}],
},
response,
),
{
"role": "tool",
"tool_call_id": "call_1",
"content": "raw oversized result",
},
])
governed_messages = [
messages[0],
messages[1],
{
"role": "tool",
"tool_call_id": "call_1",
"content": "compacted result",
},
]
provider_context = controller.prepare_request(
messages,
context_window_tokens=200_000,
model_messages=governed_messages,
)
assert provider_context is not None
assert provider_context.conversation_state is not None
assert provider_context.conversation_state.pending_messages == [{
"role": "tool",
"tool_call_id": "call_1",
"content": "compacted result",
}]
assert controller.checkpoint(messages).pending_messages[-1]["content"] == (
"raw oversized result"
)
governed_checkpoint = controller.checkpoint(
messages,
model_messages=governed_messages,
)
assert governed_checkpoint is not None
assert governed_checkpoint.pending_messages[-1]["content"] == "compacted result"
def test_transient_response_preserves_only_durable_request_messages() -> None:
provider = _provider()
current_message = {"role": "user", "content": "continue"}
supplemental = {"role": "user", "content": "internal finalization retry"}
messages = [{"role": "system", "content": "system"}, current_message]
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
state=_state("saved", pending=[
{"role": "tool", "content": "prior"},
current_message,
]),
)
provider_context = controller.prepare_request(
messages,
context_window_tokens=200_000,
supplemental_messages=[supplemental],
)
assert provider_context is not None
assert provider_context.conversation_state is not None
assert provider_context.conversation_state.pending_messages == [
{"role": "tool", "content": "prior"},
current_message,
supplemental,
]
controller.observe_response(
LLMResponse(
content="temporary failure",
finish_reason="error",
error_kind="timeout",
),
messages,
)
placeholder = {"role": "assistant", "content": "model error"}
messages.append(placeholder)
state = controller.finish(messages)
assert state is not None
assert state.pending_messages == [
{"role": "tool", "content": "prior"},
current_message,
placeholder,
]
def test_non_retryable_response_discards_saved_state() -> None:
provider = _provider()
messages = [{"role": "user", "content": "continue"}]
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
state=_state("saved"),
)
controller.prepare_request(messages, context_window_tokens=200_000)
controller.observe_response(
LLMResponse(
content="invalid request",
finish_reason="error",
error_status_code=400,
error_should_retry=False,
),
messages,
)
assert controller.finish(messages) is None
@pytest.mark.parametrize(
("finish_reason", "exposes_tool_call"),
[
("length", False),
("length", True),
("refusal", True),
("content_filter", True),
],
)
def test_terminal_response_discards_candidate_state(
finish_reason: str,
exposes_tool_call: bool,
) -> None:
provider = _provider()
messages = [{"role": "user", "content": "continue"}]
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
state=_state("saved"),
)
controller.prepare_request(messages, context_window_tokens=200_000)
candidate = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload={
"items": [{
"type": "function_call",
"call_id": "call_1",
"name": "exec",
"arguments": "{}",
}],
},
)
response = LLMResponse(
content="terminal response",
tool_calls=(
[ToolCallRequest(id="call_1", name="exec", arguments={})]
if exposes_tool_call
else []
),
finish_reason=finish_reason,
provider_state=candidate,
)
assert response.has_tool_calls is exposes_tool_call
assert response.should_execute_tools is False
controller.observe_response(response, messages)
assert controller.finish(messages) is None
def test_independent_request_exposes_context_without_capability_check() -> None:
provider = _provider(compact=False)
messages = [{"role": "user", "content": "hello"}]
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
state=_state("saved"),
)
provider_context = controller.independent_request_context(
context_window_tokens=200_000,
)
assert provider_context is not None
assert provider_context.conversation_state is None
assert provider_context.context_window_tokens == 200_000
provider.supports_native_compaction.assert_not_called()
+71
View File
@@ -0,0 +1,71 @@
"""Tests for the Eden AI provider registration."""
from unittest.mock import patch
from nanobot.config.schema import Config, ProvidersConfig
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
from nanobot.providers.registry import PROVIDERS, find_by_name
def test_edenai_config_field_exists() -> None:
assert hasattr(ProvidersConfig(), "edenai")
def test_edenai_registry_contract() -> None:
specs = {spec.name: spec for spec in PROVIDERS}
assert "edenai" in specs
edenai = specs["edenai"]
assert edenai.backend == "openai_compat"
assert edenai.env_key == "EDENAI_API_KEY"
assert edenai.display_name == "Eden AI"
assert edenai.is_gateway is True
assert edenai.detect_by_base_keyword == "edenai"
assert edenai.default_api_base == "https://api.edenai.run/v3"
assert edenai.strip_model_prefix is False
# Eden accepts OpenAI's top-level reasoning_effort parameter. Do not add
# OpenRouter's separate {"reasoning": {"effort": ...}} request shape.
assert edenai.gateway_reasoning_style == ""
def test_edenai_forced_provider_uses_default_api_base() -> None:
config = Config.model_validate(
{
"providers": {"edenai": {"apiKey": "eden-key"}},
"agents": {
"defaults": {
"provider": "edenai",
"model": "anthropic/claude-sonnet-4-5",
}
},
}
)
model = "anthropic/claude-sonnet-4-5"
assert config.get_provider_name(model) == "edenai"
assert config.get_api_key(model) == "eden-key"
assert config.get_api_base(model) == "https://api.edenai.run/v3"
def test_edenai_preserves_model_id_and_reasoning_effort() -> None:
spec = find_by_name("edenai")
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI"):
provider = OpenAICompatProvider(
api_key="eden-key",
default_model="anthropic/claude-sonnet-4-5",
spec=spec,
)
kwargs = provider._build_kwargs(
messages=[{"role": "user", "content": "hi"}],
tools=None,
model="anthropic/claude-sonnet-4-5",
max_tokens=1024,
temperature=0.7,
reasoning_effort="medium",
tool_choice=None,
)
assert kwargs["model"] == "anthropic/claude-sonnet-4-5"
assert kwargs["reasoning_effort"] == "medium"
assert "reasoning" not in kwargs.get("extra_body", {})
@@ -10,6 +10,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from nanobot.providers.base import ProviderCallContext
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
from nanobot.providers.registry import find_by_name
@@ -44,8 +45,10 @@ def test_build_responses_body_strips_github_copilot_prefix():
temperature=0.1,
reasoning_effort=None,
tool_choice=None,
provider_context=ProviderCallContext(context_window_tokens=128_000),
)
assert body["model"] == "gpt-5.4-mini"
assert "context_management" not in body
@pytest.mark.asyncio
+8 -8
View File
@@ -464,7 +464,7 @@ async def test_gemini_flash_forwards_aspect_ratio_and_image_size() -> None:
image_size="2K",
)
image_config = fake.calls[0]["json"]["generationConfig"]["responseFormat"]["image"]
image_config = fake.calls[0]["json"]["generationConfig"]["imageConfig"]
assert image_config == {"aspectRatio": "16:9", "imageSize": "2K"}
@@ -480,7 +480,7 @@ async def test_gemini_flash_2_5_drops_image_size() -> None:
image_size="1K",
)
image_config = fake.calls[0]["json"]["generationConfig"]["responseFormat"]["image"]
image_config = fake.calls[0]["json"]["generationConfig"]["imageConfig"]
assert image_config == {"aspectRatio": "4:3"}
@@ -496,7 +496,7 @@ async def test_gemini_flash_2_0_drops_image_size() -> None:
image_size="1K",
)
image_config = fake.calls[0]["json"]["generationConfig"]["responseFormat"]["image"]
image_config = fake.calls[0]["json"]["generationConfig"]["imageConfig"]
assert image_config == {"aspectRatio": "16:9"}
@@ -524,8 +524,8 @@ async def test_gemini_flash_scopes_extreme_aspect_ratios_by_model(
aspect_ratio=aspect_ratio,
)
response_format = fake.calls[0]["json"]["generationConfig"].get("responseFormat")
assert response_format == ({"image": expected} if expected else None)
image_config = fake.calls[0]["json"]["generationConfig"].get("imageConfig")
assert image_config == expected
@pytest.mark.parametrize(
@@ -553,8 +553,8 @@ async def test_gemini_flash_scopes_image_size_by_model(
image_size=image_size,
)
response_format = fake.calls[0]["json"]["generationConfig"].get("responseFormat")
assert response_format == ({"image": expected} if expected else None)
image_config = fake.calls[0]["json"]["generationConfig"].get("imageConfig")
assert image_config == expected
@pytest.mark.asyncio
@@ -571,7 +571,7 @@ async def test_gemini_flash_ignores_unsupported_hints() -> None:
image_size="1024x1024",
)
assert "responseFormat" not in fake.calls[0]["json"]["generationConfig"]
assert "imageConfig" not in fake.calls[0]["json"]["generationConfig"]
@pytest.mark.asyncio
+36
View File
@@ -14,6 +14,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from nanobot.providers.base import ProviderCallContext
from nanobot.providers.openai_compat_provider import OpenAICompatProvider
from nanobot.providers.registry import find_by_name
@@ -679,6 +680,7 @@ async def test_direct_openai_gpt5_uses_responses_api() -> None:
assert call_kwargs["max_output_tokens"] == 4096
assert "input" in call_kwargs
assert "messages" not in call_kwargs
assert call_kwargs["include"] == ["reasoning.encrypted_content"]
@pytest.mark.asyncio
@@ -710,6 +712,40 @@ async def test_direct_openai_reasoning_prefers_responses_api() -> None:
assert call_kwargs["include"] == ["reasoning.encrypted_content"]
@pytest.mark.asyncio
async def test_direct_openai_retries_without_unsupported_server_compaction() -> None:
mock_chat = AsyncMock(return_value=_fake_chat_response())
mock_responses = AsyncMock(side_effect=[
_FakeResponsesError(400, "Unknown parameter: context_management"),
_fake_responses_response("compaction fallback"),
])
spec = find_by_name("openai")
with patch("nanobot.providers.openai_compat_provider.AsyncOpenAI") as mock_client_class:
client_instance = mock_client_class.return_value
client_instance.chat.completions.create = mock_chat
client_instance.responses.create = mock_responses
provider = OpenAICompatProvider(
api_key="sk-test-key",
default_model="gpt-5.6",
spec=spec,
)
result = await provider.chat_with_context(
messages=[{"role": "user", "content": "hello"}],
model="gpt-5.6",
provider_context=ProviderCallContext(context_window_tokens=200_000),
)
assert result.content == "compaction fallback"
assert result.provider_state is not None
assert mock_responses.await_count == 2
assert "context_management" in mock_responses.call_args_list[0].kwargs
assert "context_management" not in mock_responses.call_args_list[1].kwargs
assert provider.supports_native_compaction("gpt-5.6") is False
mock_chat.assert_not_awaited()
@pytest.mark.asyncio
async def test_direct_openai_gpt4o_stays_on_chat_completions() -> None:
mock_chat = AsyncMock(return_value=_fake_chat_response())
+301 -5
View File
@@ -20,6 +20,7 @@ from nanobot.providers.openai_codex_provider import (
_request_codex,
_should_retry_status,
)
from nanobot.providers.openai_responses import build_responses_state
from nanobot.providers.registry import find_by_name
@@ -115,6 +116,48 @@ async def test_codex_request_non_200_populates_http_metadata(monkeypatch) -> Non
assert error.should_retry is True
@pytest.mark.asyncio
async def test_codex_request_marks_rejected_compaction_without_retaining_raw_body(
monkeypatch,
) -> None:
original_client = httpx.AsyncClient
secret = "PRIVATE PROMPT MUST NOT BE RETAINED"
def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(
400,
json={
"error": {
"message": f"Unknown input type compaction_trigger; {secret}",
},
},
request=request,
)
def fake_client(
*,
timeout: int,
verify: bool,
**_kwargs: object,
) -> httpx.AsyncClient:
return original_client(transport=httpx.MockTransport(handler), timeout=timeout)
monkeypatch.setattr("nanobot.providers.openai_codex_provider.httpx.AsyncClient", fake_client)
with pytest.raises(_CodexHTTPError) as caught:
await _request_codex(
"https://codex.example/responses",
{},
{"input": [{"type": "compaction_trigger"}]},
verify=True,
)
error = caught.value
assert error.compaction_unsupported is True
assert secret not in str(error)
assert not hasattr(error, "body")
@pytest.mark.asyncio
async def test_codex_request_honors_stream_idle_timeout_env(monkeypatch) -> None:
"""NANOBOT_STREAM_IDLE_TIMEOUT_S overrides the default Codex stream timeout."""
@@ -192,7 +235,7 @@ async def test_codex_prompt_cache_key_uses_stable_conversation_prefix(monkeypatc
):
_ = proxy, on_thinking_delta, on_tool_call_delta
bodies.append(body)
return "ok", [], "stop", {}, None
return provider_base.LLMResponse(content="ok")
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)
@@ -232,7 +275,7 @@ async def test_codex_provider_applies_extra_body_from_config(monkeypatch) -> Non
async def fake_request(_url, _headers, body, **_kwargs):
bodies.append(body)
return "ok", [], "stop", {}, None
return provider_base.LLMResponse(content="ok")
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)
config = Config.model_validate({
@@ -297,7 +340,7 @@ async def test_codex_provider_passes_proxy_to_oauth_and_response_request(monkeyp
):
_ = url, headers, body, verify, on_content_delta, on_thinking_delta, on_tool_call_delta
seen["request_proxy"] = proxy
return "ok", [], "stop", {}, None
return provider_base.LLMResponse(content="ok")
monkeypatch.setattr("nanobot.providers.openai_codex_provider.get_codex_token", fake_token)
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)
@@ -384,7 +427,7 @@ async def test_codex_retry_uses_structured_timeout_metadata(monkeypatch) -> None
calls += 1
if calls == 1:
raise httpx.ReadTimeout("")
return "ok", [], "stop", {}, None
return provider_base.LLMResponse(content="ok")
async def fake_sleep(delay: float) -> None:
delays.append(delay)
@@ -533,6 +576,254 @@ def test_codex_reasoning_options_request_summary_without_forcing_effort() -> Non
assert _build_reasoning_options("none") == {"effort": "none"}
@pytest.mark.asyncio
async def test_codex_replayed_tool_turn_omits_server_item_ids(monkeypatch) -> None:
_mock_codex_token(monkeypatch)
provider = OpenAICodexProvider(default_model="openai-codex/gpt-5.6-sol")
state = build_responses_state(
provider=provider._responses_state_provider(),
model="gpt-5.6-sol",
input_items=[{
"id": "msg_user",
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "Check the weather"}],
}],
output_items=[
{
"id": "rs_reasoning",
"type": "reasoning",
"encrypted_content": "opaque reasoning",
"summary": [],
},
{
"id": "fc_read",
"type": "function_call",
"call_id": "call_read",
"name": "read_file",
"arguments": '{"path":"weather/SKILL.md"}',
"status": "completed",
},
],
)
bodies: list[dict[str, Any]] = []
async def fake_request(
url,
headers,
body,
verify,
proxy=None,
on_content_delta=None,
on_thinking_delta=None,
on_tool_call_delta=None,
):
bodies.append(body)
return provider_base.LLMResponse(content="done")
monkeypatch.setattr(
"nanobot.providers.openai_codex_provider._request_codex",
fake_request,
)
response = await provider.chat(
[{"role": "user", "content": "Check the weather"}],
provider_context=provider_base.ProviderCallContext(
conversation_state=state.with_pending_messages([{
"role": "tool",
"tool_call_id": "call_read|fc_read",
"content": "weather skill contents",
}]),
),
)
assert response.content == "done"
assert len(bodies) == 1
input_items = bodies[0]["input"]
assert [item.get("type") for item in input_items] == [
"message",
"reasoning",
"function_call",
"function_call_output",
]
assert all("id" not in item for item in input_items)
assert input_items[1]["encrypted_content"] == "opaque reasoning"
assert input_items[2]["call_id"] == "call_read"
assert input_items[3]["call_id"] == "call_read"
@pytest.mark.asyncio
async def test_codex_compacts_state_at_ninety_percent_before_next_request(
monkeypatch,
) -> None:
_mock_codex_token(monkeypatch)
provider = OpenAICodexProvider(default_model="openai-codex/gpt-5.6-sol")
state_provider = provider._responses_state_provider()
state = build_responses_state(
provider=state_provider,
model="gpt-5.6-sol",
input_items=[{"type": "message", "role": "user", "content": "old question"}],
output_items=[
{"type": "reasoning", "encrypted_content": "old opaque reasoning"},
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "old answer"}],
},
],
usage={
"prompt_tokens": 90,
"completion_tokens": 5,
"total_tokens": 95,
},
)
bodies: list[dict[str, Any]] = []
async def fake_request(
url,
headers,
body,
verify,
proxy=None,
on_content_delta=None,
on_thinking_delta=None,
on_tool_call_delta=None,
):
_ = (
url,
headers,
verify,
proxy,
on_content_delta,
on_thinking_delta,
on_tool_call_delta,
)
bodies.append(body)
if body["input"][-1].get("type") == "compaction_trigger":
compact_item = {
"type": "compaction",
"encrypted_content": "compacted opaque state",
}
return provider_base.LLMResponse(
content=None,
provider_state=build_responses_state(
provider=state_provider,
model="gpt-5.6-sol",
input_items=body["input"],
output_items=[compact_item],
usage={
"prompt_tokens": 95,
"completion_tokens": 2,
"total_tokens": 97,
},
),
)
return provider_base.LLMResponse(content="done")
monkeypatch.setattr(
"nanobot.providers.openai_codex_provider._request_codex",
fake_request,
)
response = await provider.chat_with_retry(
[
{"role": "system", "content": "system"},
{"role": "user", "content": "new question"},
],
max_tokens=5,
provider_context=provider_base.ProviderCallContext(
conversation_state=state.with_pending_messages([
{"role": "user", "content": "new question"},
]),
context_window_tokens=100,
),
)
assert response.content == "done"
assert len(bodies) == 2
assert bodies[0]["input"][-1] == {"type": "compaction_trigger"}
assert bodies[1]["input"][-1] == {
"type": "compaction",
"encrypted_content": "compacted opaque state",
}
assert not any(
item.get("type") == "reasoning"
for item in bodies[1]["input"]
)
assert any(
item.get("role") == "user"
and "new question" in str(item.get("content"))
for item in bodies[1]["input"]
)
@pytest.mark.asyncio
async def test_codex_disables_unsupported_native_compaction_and_continues(
monkeypatch,
) -> None:
_mock_codex_token(monkeypatch)
provider = OpenAICodexProvider(default_model="openai-codex/gpt-5.6-sol")
state_provider = provider._responses_state_provider()
state = build_responses_state(
provider=state_provider,
model="gpt-5.6-sol",
input_items=[{"type": "message", "role": "user", "content": "old"}],
output_items=[{"type": "reasoning", "encrypted_content": "opaque"}],
usage={"prompt_tokens": 90, "completion_tokens": 5, "total_tokens": 95},
)
bodies: list[dict[str, Any]] = []
async def fake_request(
url,
headers,
body,
verify,
proxy=None,
on_content_delta=None,
on_thinking_delta=None,
on_tool_call_delta=None,
):
_ = (
url,
headers,
verify,
proxy,
on_content_delta,
on_thinking_delta,
on_tool_call_delta,
)
bodies.append(body)
if body["input"][-1].get("type") == "compaction_trigger":
raise _CodexHTTPError(
"HTTP 400: Codex API request failed",
status_code=400,
compaction_unsupported=True,
)
return provider_base.LLMResponse(content="done")
monkeypatch.setattr(
"nanobot.providers.openai_codex_provider._request_codex",
fake_request,
)
response = await provider.chat(
[{"role": "user", "content": "new"}],
max_tokens=5,
provider_context=provider_base.ProviderCallContext(
conversation_state=state.with_pending_messages([
{"role": "user", "content": "new"},
]),
context_window_tokens=100,
),
)
assert response.content == "done"
assert len(bodies) == 2
assert bodies[0]["input"][-1] == {"type": "compaction_trigger"}
assert bodies[1]["input"][-1] != {"type": "compaction_trigger"}
assert provider.supports_native_compaction() is False
@pytest.mark.asyncio
async def test_codex_stream_surfaces_reasoning_summary(monkeypatch) -> None:
def fake_token(**_kwargs):
@@ -559,7 +850,12 @@ async def test_codex_stream_surfaces_reasoning_summary(monkeypatch) -> None:
await on_content_delta("answer")
if on_thinking_delta:
await on_thinking_delta("summary")
return "answer", [], "stop", {"prompt_tokens": 10, "completion_tokens": 5}, "summary"
return provider_base.LLMResponse(
content="answer",
finish_reason="stop",
usage={"prompt_tokens": 10, "completion_tokens": 5},
reasoning_content="summary",
)
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)
+844 -5
View File
@@ -1,9 +1,11 @@
"""Tests for the shared openai_responses converters and parsers."""
import json
from io import StringIO
from unittest.mock import MagicMock, patch
import pytest
from loguru import logger
from nanobot.providers.openai_responses.converters import (
convert_messages,
@@ -12,12 +14,22 @@ from nanobot.providers.openai_responses.converters import (
split_tool_call_id,
)
from nanobot.providers.openai_responses.parsing import (
ResponsesStreamCapture,
consume_sdk_stream,
consume_sse,
consume_sse_with_reasoning,
is_replayable_finish_reason,
map_finish_reason,
parse_response_output,
)
from nanobot.providers.openai_responses.state import (
build_responses_state,
is_compaction_compatibility_error,
prepare_responses_input,
resolve_compact_threshold,
responses_state_context_tokens,
responses_state_items,
)
# ======================================================================
# converters - split_tool_call_id
@@ -138,6 +150,51 @@ class TestConvertMessages:
assert items[0]["content"][0]["type"] == "output_text"
assert items[0]["content"][0]["text"] == "I'll help"
def test_preserves_deepseek_reasoning_content(self):
_, items = convert_messages([
{"role": "assistant", "reasoning_content": "think first", "content": "answer"},
], preserve_reasoning=True)
assert items == [
{
"type": "reasoning",
"content": [{"type": "output_text", "text": "think first"}],
},
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "answer"}],
"status": "completed",
"id": "msg_0",
},
]
def test_reasoning_content_serialized_as_array_for_deepseek(self):
# Regression for PR #5214: DeepSeek's Responses gateway rejects
# reasoning items whose ``content`` is a plain string with
# "input: invalid type: string ..., expected a sequence" (observed
# after context consolidation cleared provider state and forced
# full-history conversion). ``content`` must be a list of parts,
# matching both the OpenAI Responses schema and DeepSeek's accepted
# wire shape.
_, items = convert_messages([
{
"role": "assistant",
"reasoning_content": "Michael topped up DeepSeek with $10.",
"content": "",
"tool_calls": [{
"id": "call_1|fc_1",
"function": {"name": "list_dir", "arguments": "{}"},
}],
},
], preserve_reasoning=True)
assert items[0]["type"] == "reasoning"
assert items[0]["content"] == [
{"type": "output_text", "text": "Michael topped up DeepSeek with $10."},
]
assert items[1]["type"] == "function_call"
def test_assistant_empty_content_skipped(self):
_, items = convert_messages([{"role": "assistant", "content": ""}])
assert len(items) == 0
@@ -398,6 +455,17 @@ class TestMapFinishReason:
def test_unknown_defaults_to_stop(self):
assert map_finish_reason("some_new_status") == "stop"
@pytest.mark.parametrize("finish_reason", ["stop", "tool_calls", "function_call"])
def test_replayable_finish_reasons(self, finish_reason):
assert is_replayable_finish_reason(finish_reason) is True
@pytest.mark.parametrize(
"finish_reason",
["length", "refusal", "content_filter", "error"],
)
def test_non_replayable_finish_reasons(self, finish_reason):
assert is_replayable_finish_reason(finish_reason) is False
# ======================================================================
# parsing - parse_response_output
@@ -418,6 +486,29 @@ class TestParseResponseOutput:
assert result.usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
assert result.tool_calls == []
def test_refusal_response_surfaces_text_without_advancing_state(self):
refusal = "I cant help with that request."
resp = {
"output": [{
"type": "message",
"role": "assistant",
"content": [{"type": "refusal", "refusal": refusal}],
}],
"status": "completed",
"usage": {},
}
result = parse_response_output(
resp,
state_provider="openai:test",
state_model="gpt-5.6",
state_input_items=[{"role": "user", "content": "request"}],
)
assert result.content == refusal
assert result.finish_reason == "refusal"
assert result.provider_state is None
def test_tool_call_response(self):
resp = {
"output": [{
@@ -429,12 +520,18 @@ class TestParseResponseOutput:
"status": "completed",
"usage": {},
}
result = parse_response_output(resp)
result = parse_response_output(
resp,
state_provider="openai:test",
state_model="gpt-5.6",
state_input_items=[{"role": "user", "content": "weather?"}],
)
assert result.content is None
assert len(result.tool_calls) == 1
assert result.tool_calls[0].name == "get_weather"
assert result.tool_calls[0].arguments == {"city": "SF"}
assert result.tool_calls[0].id == "call_1|fc_1"
assert result.provider_state is not None
def test_malformed_tool_arguments_logged(self):
"""Malformed JSON arguments should log a warning and remain non-object."""
@@ -487,16 +584,61 @@ class TestParseResponseOutput:
assert result.content == "42"
assert result.reasoning_content == "I think therefore I am."
def test_deepseek_reasoning_content_extracted(self):
resp = {
"output": [
{"type": "reasoning", "content": "think first"},
{"type": "message", "content": [
{"type": "output_text", "text": "answer"},
]},
],
"status": "completed", "usage": {},
}
result = parse_response_output(resp)
assert result.content == "answer"
assert result.reasoning_content == "think first"
def test_empty_output(self):
resp = {"output": [], "status": "completed", "usage": {}}
result = parse_response_output(resp)
assert result.content is None
assert result.tool_calls == []
def test_incomplete_status(self):
resp = {"output": [], "status": "incomplete", "usage": {}}
result = parse_response_output(resp)
assert result.finish_reason == "length"
@pytest.mark.parametrize(
("reason", "expected_finish_reason"),
[
("max_output_tokens", "length"),
("content_filter", "content_filter"),
],
)
def test_incomplete_status(self, reason, expected_finish_reason):
resp = {
"output": [],
"status": "incomplete",
"incomplete_details": {"reason": reason},
"usage": {},
}
result = parse_response_output(
resp,
state_provider="openai:test",
state_model="gpt-5.6",
state_input_items=[{"role": "user", "content": "prompt"}],
)
assert result.finish_reason == expected_finish_reason
assert result.provider_state is None
def test_unknown_status_does_not_advance_provider_state(self):
result = parse_response_output(
{"output": [], "status": "future_terminal_status", "usage": {}},
state_provider="openai:test",
state_model="gpt-5.6",
state_input_items=[{"role": "user", "content": "prompt"}],
)
assert result.finish_reason == "stop"
assert result.provider_state is None
def test_sdk_model_object(self):
"""parse_response_output should handle SDK objects with model_dump()."""
@@ -523,6 +665,247 @@ class TestParseResponseOutput:
assert result.usage["completion_tokens"] == 50
assert result.usage["total_tokens"] == 150
def test_preserves_every_output_item_as_opaque_state(self):
input_items = [{"role": "user", "content": "inspect the repo"}]
output = [
{
"id": "rs_1",
"type": "reasoning",
"encrypted_content": "opaque-secret",
"summary": [],
},
{
"id": "future_1",
"type": "future_item_type",
"provider_field": {"nested": True},
},
{
"id": "msg_1",
"type": "message",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": "done"}],
},
]
result = parse_response_output(
{"output": output, "status": "completed", "usage": {}},
state_provider="openai:test",
state_model="gpt-5.6",
state_input_items=input_items,
)
assert result.provider_state is not None
assert responses_state_items(result.provider_state) == [*input_items, *output]
class TestResponsesConversationState:
def test_server_compaction_prunes_superseded_prefix(self):
state = build_responses_state(
provider="openai:test",
model="gpt-5.6",
input_items=[
{"type": "message", "role": "user", "content": "old"},
{"type": "reasoning", "encrypted_content": "old-reasoning"},
],
output_items=[
{"type": "compaction", "encrypted_content": "compact"},
{"type": "message", "role": "assistant", "content": "new"},
],
usage={
"prompt_tokens": 90,
"completion_tokens": 10,
"total_tokens": 100,
},
)
assert responses_state_items(state) == [
{"type": "compaction", "encrypted_content": "compact"},
{"type": "message", "role": "assistant", "content": "new"},
]
assert responses_state_context_tokens(state) == 100
def test_existing_compaction_keeps_canonical_retained_prefix(self):
canonical_input = [
{"type": "message", "role": "user", "content": "retained"},
{"type": "compaction", "encrypted_content": "compact"},
]
output = [{"type": "message", "role": "assistant", "content": "new"}]
state = build_responses_state(
provider="openai:test",
model="gpt-5.6",
input_items=canonical_input,
output_items=output,
)
assert responses_state_items(state) == [*canonical_input, *output]
@pytest.mark.parametrize(
("context_window", "max_output", "expected"),
[
(200_000, 20_000, 180_000),
(100_000, 30_000, 70_000),
(0, 4_096, None),
],
)
def test_compact_threshold_reserves_codex_style_headroom(
self,
context_window,
max_output,
expected,
):
assert resolve_compact_threshold(context_window, max_output) == expected
def test_compaction_compatibility_recognizes_old_sdk_signature_error(self):
error = TypeError("create() got an unexpected keyword argument 'context_management'")
assert is_compaction_compatibility_error(error) is True
assert is_compaction_compatibility_error(TypeError("unrelated argument")) is False
def test_state_observability_logs_counts_without_opaque_content(self):
secret = "opaque-secret-that-must-not-be-logged"
state = build_responses_state(
provider=f"openai:https://example.test/?key={secret}",
model=f"secret-model-{secret}",
input_items=[{"role": "user", "content": secret}],
output_items=[{"type": "reasoning", "encrypted_content": secret}],
).with_pending_messages([{"role": "user", "content": secret}])
sink = StringIO()
sink_id = logger.add(sink, level="DEBUG", format="{message}")
try:
prepare_responses_input(
[{"role": "user", "content": secret}],
state=state,
provider=state.provider,
model=state.model,
)
build_responses_state(
provider=state.provider,
model=state.model,
input_items=[
{"role": "user", "content": secret},
{"type": "reasoning", "encrypted_content": secret},
],
output_items=[
{"type": "compaction", "encrypted_content": secret},
],
)
finally:
logger.remove(sink_id)
log_text = sink.getvalue()
assert "prior_items=2" in log_text
assert "pending_messages=1" in log_text
assert "dropped_items=2" in log_text
assert secret not in log_text
def test_replays_exact_items_then_only_pending_and_new_messages(self):
prior_items = [
{"role": "user", "content": "first"},
{
"type": "reasoning",
"id": "rs_1",
"encrypted_content": "opaque-secret",
},
{
"type": "function_call",
"id": "fc_1",
"call_id": "call_1",
"name": "read_file",
"arguments": '{"path":"a.py"}',
},
]
state = build_responses_state(
provider="openai:test",
model="gpt-5.6",
input_items=prior_items[:1],
output_items=prior_items[1:],
).with_pending_messages([
{
"role": "tool",
"tool_call_id": "call_1|fc_1",
"content": "file contents",
},
{"role": "user", "content": "continue"},
])
instructions, items, replayed = prepare_responses_input(
[
{"role": "system", "content": "current instructions"},
{"role": "user", "content": "a lossy public transcript"},
],
state=state,
provider="openai:test",
model="gpt-5.6",
)
assert instructions == "current instructions"
assert replayed is True
assert items[:3] == prior_items
assert items[3] == {
"type": "function_call_output",
"call_id": "call_1",
"output": "file contents",
}
assert items[4] == {
"role": "user",
"content": [{"type": "input_text", "text": "continue"}],
}
assert "lossy public transcript" not in str(items)
def test_replayed_and_delta_reasoning_items_keep_array_content(self):
# Regression for PR #5214: token consolidation clears
# ``provider_state``, so the next turn converts the full history
# (including assistant reasoning) instead of replaying server items.
# Both paths must keep reasoning ``content`` as a list - DeepSeek's
# Responses gateway rejects the string form with a serde error.
prior_items = [
{
"type": "reasoning",
"id": "rs_1",
"content": [{"type": "output_text", "text": "prior reasoning"}],
},
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "prior answer"}],
"status": "completed",
"id": "msg_0",
},
]
state = build_responses_state(
provider="openai:test",
model="deepseek-v4-flash",
input_items=prior_items,
output_items=[],
).with_pending_messages([
{
"role": "assistant",
"reasoning_content": "think before acting",
"content": "answer",
},
{"role": "user", "content": "audit the tools"},
])
instructions, items, replayed = prepare_responses_input(
[
{"role": "system", "content": "You are KITT."},
{"role": "user", "content": "audit the tools"},
],
state=state,
provider="openai:test",
model="deepseek-v4-flash",
preserve_reasoning=True,
)
assert instructions == "You are KITT."
assert replayed is True
reasoning_items = [item for item in items if item.get("type") == "reasoning"]
assert len(reasoning_items) == 2 # one replayed, one converted delta
for item in reasoning_items:
assert isinstance(item["content"], list)
assert item["content"][0]["type"] == "output_text"
# ======================================================================
# parsing - consume_sse
@@ -553,6 +936,122 @@ class TestConsumeSse:
assert tool_calls == []
assert finish_reason == "stop"
@pytest.mark.asyncio
async def test_refusal_events_reconcile_parts_and_terminal_output(self):
refusal = "First and second sentence. Done-only. Terminal suffix."
terminal_response = {
"status": "completed",
"output": [{
"type": "message",
"id": "msg_2",
"role": "assistant",
"content": [{"type": "refusal", "refusal": refusal}],
}],
}
response = _SseResponse([
{
"type": "response.refusal.delta",
"item_id": "msg_1",
"content_index": 0,
"delta": "First",
},
{
"type": "response.refusal.delta",
"item_id": "msg_1",
"content_index": 1,
"delta": " and second",
},
{
"type": "response.refusal.done",
"item_id": "msg_1",
"content_index": 0,
"refusal": "First",
},
{
"type": "response.refusal.done",
"item_id": "msg_1",
"content_index": 1,
"refusal": " and second sentence.",
},
{
"type": "response.refusal.done",
"item_id": "msg_2",
"content_index": 0,
"refusal": " Done-only.",
},
{
"type": "response.refusal.delta",
"item_id": "msg_2",
"content_index": 1,
"delta": " Terminal",
},
{"type": "response.completed", "response": terminal_response},
])
capture = ResponsesStreamCapture()
deltas: list[str] = []
async def on_content(delta: str) -> None:
deltas.append(delta)
content, _, finish_reason, _, _ = await consume_sse_with_reasoning(
response,
on_content_delta=on_content,
capture=capture,
)
assert content == refusal
assert deltas == [
"First",
" and second",
" sentence.",
" Done-only.",
" Terminal",
" suffix.",
]
assert finish_reason == "refusal"
assert capture.completed is True
assert is_replayable_finish_reason(finish_reason) is False
@pytest.mark.asyncio
@pytest.mark.parametrize("source", ["events", "terminal"])
async def test_refusal_without_deltas_has_non_replayable_finish(self, source: str):
refusal = "I cant help with that request."
terminal_response = {
"status": "completed",
"output": [{
"type": "message",
"id": "msg_1",
"role": "assistant",
"content": [{"type": "refusal", "refusal": refusal}],
}],
}
events = (
[
{"type": "response.refusal.done", "refusal": refusal},
{"type": "response.completed", "response": {"status": "completed"}},
]
if source == "events"
else [{"type": "response.completed", "response": terminal_response}]
)
response = _SseResponse(events)
capture = ResponsesStreamCapture()
deltas: list[str] = []
async def on_content(delta: str) -> None:
deltas.append(delta)
content, _, finish_reason, _, _ = await consume_sse_with_reasoning(
response,
on_content_delta=on_content,
capture=capture,
)
assert content == refusal
assert deltas == [refusal]
assert finish_reason == "refusal"
assert capture.completed is True
assert is_replayable_finish_reason(finish_reason) is False
@pytest.mark.asyncio
async def test_reasoning_summary_delta_extracted(self):
response = _SseResponse([
@@ -599,6 +1098,139 @@ class TestConsumeSse:
assert reasoning == "cached summary"
@pytest.mark.asyncio
async def test_capture_commits_exact_items_only_after_completed_event(self):
output = [
{
"type": "reasoning",
"id": "rs_1",
"encrypted_content": "opaque-secret",
},
{"type": "future_item_type", "id": "future_1", "value": 7},
]
capture = ResponsesStreamCapture()
response = _SseResponse([
{
"type": "response.output_item.done",
"output_index": 0,
"item": output[0],
},
{
"type": "response.output_item.done",
"output_index": 1,
"item": output[1],
},
{
"type": "response.completed",
"response": {"status": "completed", "output": output},
},
])
await consume_sse_with_reasoning(response, capture=capture)
assert capture.completed is True
assert capture.output_items == output
@pytest.mark.asyncio
async def test_capture_keeps_done_items_when_completed_output_is_empty(self):
output = [
{
"type": "reasoning",
"id": "rs_1",
"encrypted_content": "opaque-secret",
"summary": [],
},
{
"type": "function_call",
"id": "fc_1",
"call_id": "call_1",
"name": "read_file",
"arguments": '{"path":"weather/SKILL.md"}',
},
]
capture = ResponsesStreamCapture()
response = _SseResponse([
{
"type": "response.output_item.done",
"output_index": index,
"item": item,
}
for index, item in enumerate(output)
] + [{
"type": "response.completed",
"response": {"status": "completed", "output": []},
}])
await consume_sse_with_reasoning(response, capture=capture)
assert capture.completed is True
assert capture.output_items == output
@pytest.mark.asyncio
@pytest.mark.parametrize(
("reason", "expected_finish_reason"),
[
("max_output_tokens", "length"),
("content_filter", "content_filter"),
],
)
async def test_incomplete_event_commits_capture_usage(
self,
reason,
expected_finish_reason,
):
output = [
{
"type": "message",
"id": "msg_1",
"status": "incomplete",
"content": [{"type": "output_text", "text": "partial"}],
},
]
terminal_response = {
"id": "resp_1",
"status": "incomplete",
"incomplete_details": {"reason": reason},
"output": output,
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
}
capture = ResponsesStreamCapture()
response = _SseResponse([
{"type": "response.output_text.delta", "delta": "partial"},
{"type": "response.incomplete", "response": terminal_response},
])
content, _, finish_reason, usage, _ = await consume_sse_with_reasoning(
response,
capture=capture,
)
assert content == "partial"
assert finish_reason == expected_finish_reason
assert usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
assert capture.completed is True
assert capture.response == terminal_response
assert capture.output_items == output
@pytest.mark.asyncio
async def test_capture_does_not_commit_interrupted_stream(self):
capture = ResponsesStreamCapture()
response = _SseResponse([
{
"type": "response.output_item.done",
"output_index": 0,
"item": {
"type": "reasoning",
"id": "rs_1",
"encrypted_content": "opaque-secret",
},
},
])
await consume_sse_with_reasoning(response, capture=capture)
assert capture.completed is False
@pytest.mark.asyncio
async def test_reasoning_summary_from_done_item(self):
response = _SseResponse([
@@ -755,6 +1387,131 @@ class TestConsumeSdkStream:
assert tool_calls == []
assert finish_reason == "stop"
@pytest.mark.asyncio
async def test_refusal_events_reconcile_parts_and_terminal_output(self):
refusal = "First and second sentence. Done-only. Terminal suffix."
terminal_response = {
"status": "completed",
"output": [{
"type": "message",
"id": "msg_2",
"role": "assistant",
"content": [{"type": "refusal", "refusal": refusal}],
}],
}
resp_obj = MagicMock(status="completed", usage=None, output=[])
resp_obj.model_dump.return_value = terminal_response
events = [
MagicMock(
type="response.refusal.delta",
item_id="msg_1",
content_index=0,
delta="First",
),
MagicMock(
type="response.refusal.delta",
item_id="msg_1",
content_index=1,
delta=" and second",
),
MagicMock(
type="response.refusal.done",
item_id="msg_1",
content_index=0,
refusal="First",
),
MagicMock(
type="response.refusal.done",
item_id="msg_1",
content_index=1,
refusal=" and second sentence.",
),
MagicMock(
type="response.refusal.done",
item_id="msg_2",
content_index=0,
refusal=" Done-only.",
),
MagicMock(
type="response.refusal.delta",
item_id="msg_2",
content_index=1,
delta=" Terminal",
),
MagicMock(type="response.completed", response=resp_obj),
]
capture = ResponsesStreamCapture()
deltas: list[str] = []
async def on_content(delta: str) -> None:
deltas.append(delta)
async def stream():
for event in events:
yield event
content, _, finish_reason, _, _ = await consume_sdk_stream(
stream(),
on_content_delta=on_content,
capture=capture,
)
assert content == refusal
assert deltas == [
"First",
" and second",
" sentence.",
" Done-only.",
" Terminal",
" suffix.",
]
assert finish_reason == "refusal"
assert capture.completed is True
assert is_replayable_finish_reason(finish_reason) is False
@pytest.mark.asyncio
@pytest.mark.parametrize("source", ["events", "terminal"])
async def test_refusal_without_deltas_has_non_replayable_finish(self, source: str):
refusal = "I cant help with that request."
terminal_response = {
"status": "completed",
"output": [{
"type": "message",
"id": "msg_1",
"role": "assistant",
"content": [{"type": "refusal", "refusal": refusal}],
}],
}
resp_obj = MagicMock(status="completed", usage=None, output=[])
resp_obj.model_dump.return_value = terminal_response
capture = ResponsesStreamCapture()
deltas: list[str] = []
async def on_content(delta: str) -> None:
deltas.append(delta)
async def stream():
if source == "events":
yield MagicMock(type="response.refusal.done", refusal=refusal)
yield MagicMock(
type="response.completed",
response={"status": "completed"},
)
else:
yield MagicMock(type="response.completed", response=resp_obj)
content, _, finish_reason, _, _ = await consume_sdk_stream(
stream(),
on_content_delta=on_content,
capture=capture,
)
assert content == refusal
assert deltas == [refusal]
assert finish_reason == "refusal"
assert capture.completed is True
assert is_replayable_finish_reason(finish_reason) is False
@pytest.mark.asyncio
async def test_on_content_delta_called(self):
ev1 = MagicMock(type="response.output_text.delta", delta="hi")
@@ -919,6 +1676,64 @@ class TestConsumeSdkStream:
_, _, _, usage, _ = await consume_sdk_stream(stream())
assert usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
@pytest.mark.asyncio
@pytest.mark.parametrize(
("reason", "expected_finish_reason"),
[
("max_output_tokens", "length"),
("content_filter", "content_filter"),
],
)
async def test_incomplete_event_commits_capture_usage(
self,
reason,
expected_finish_reason,
):
output = [
{
"type": "message",
"id": "msg_1",
"status": "incomplete",
"content": [{"type": "output_text", "text": "partial"}],
},
]
usage_obj = MagicMock(input_tokens=10, output_tokens=5, total_tokens=15)
output_item = MagicMock(type="message")
terminal_response = {
"id": "resp_1",
"status": "incomplete",
"incomplete_details": {"reason": reason},
"output": output,
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
}
resp_obj = MagicMock(
status="incomplete",
usage=usage_obj,
output=[output_item],
)
resp_obj.model_dump.return_value = terminal_response
events = [
MagicMock(type="response.output_text.delta", delta="partial"),
MagicMock(type="response.incomplete", response=resp_obj),
]
capture = ResponsesStreamCapture()
async def stream():
for event in events:
yield event
content, _, finish_reason, usage, _ = await consume_sdk_stream(
stream(),
capture=capture,
)
assert content == "partial"
assert finish_reason == expected_finish_reason
assert usage == {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}
assert capture.completed is True
assert capture.response == terminal_response
assert capture.output_items == output
@pytest.mark.asyncio
async def test_reasoning_extracted(self):
summary_item = MagicMock(type="summary_text", text="thinking...")
@@ -932,6 +1747,30 @@ class TestConsumeSdkStream:
_, _, _, _, reasoning = await consume_sdk_stream(stream())
assert reasoning == "thinking..."
@pytest.mark.asyncio
async def test_deepseek_reasoning_text_streamed(self):
events = [
MagicMock(type="response.reasoning_text.delta", delta="step 1 "),
MagicMock(type="response.reasoning_text.delta", delta="step 2"),
MagicMock(type="response.reasoning_text.done", text="step 1 step 2"),
]
emitted: list[str] = []
async def stream():
for event in events:
yield event
async def on_reasoning_delta(delta: str) -> None:
emitted.append(delta)
_, _, _, _, reasoning = await consume_sdk_stream(
stream(),
on_reasoning_delta=on_reasoning_delta,
)
assert reasoning == "step 1 step 2"
assert emitted == ["step 1 ", "step 2"]
@pytest.mark.asyncio
async def test_error_event_raises(self):
ev = MagicMock(type="error", error="rate_limit_exceeded")
+81 -1
View File
@@ -3,7 +3,14 @@ import copy
import pytest
from nanobot.providers.base import RETRY_AFTER_BUFFER, GenerationSettings, LLMProvider, LLMResponse
from nanobot.providers.base import (
RETRY_AFTER_BUFFER,
GenerationSettings,
LLMProvider,
LLMResponse,
ProviderCallContext,
ProviderConversationState,
)
class ScriptedProvider(LLMProvider):
@@ -330,6 +337,79 @@ async def test_successful_image_retry_mutates_original_messages_in_place() -> No
assert any("not delivered" in (block.get("text") or "").lower() for block in content)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("messages", "payload", "pending_messages"),
[
(_IMAGE_MSG, {}, _IMAGE_MSG),
(
[{"role": "user", "content": "continue"}],
{
"items": [
{
"type": "message",
"role": "user",
"content": [
{
"type": "input_image",
"image_url": "data:image/png;base64,abc",
}
],
}
]
},
[],
),
],
ids=["pending-image", "opaque-payload-image"],
)
async def test_image_retry_discards_provider_state_with_images(
messages,
payload,
pending_messages,
) -> None:
class ContextScriptedProvider(ScriptedProvider):
def __init__(self, responses):
super().__init__(responses)
self.contexts: list[ProviderCallContext] = []
async def chat_with_context(
self,
*,
provider_context: ProviderCallContext,
**kwargs,
) -> LLMResponse:
self.contexts.append(provider_context)
return await self.chat(**kwargs)
provider = ContextScriptedProvider([
LLMResponse(content="model does not support images", finish_reason="error"),
LLMResponse(content="ok, no image"),
])
messages = copy.deepcopy(messages)
state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="gpt-5.6",
version=1,
payload=copy.deepcopy(payload),
pending_messages=copy.deepcopy(pending_messages),
)
response = await provider.chat_with_retry(
messages=messages,
provider_context=ProviderCallContext(conversation_state=state),
)
assert response.content == "ok, no image"
retry_context = provider.contexts[-1]
assert isinstance(retry_context, ProviderCallContext)
assert retry_context.conversation_state is None
public_content = messages[0]["content"]
if isinstance(public_content, list):
assert all(block.get("type") != "image_url" for block in public_content)
@pytest.mark.asyncio
async def test_non_transient_error_without_images_no_retry() -> None:
"""Non-transient errors without image content are returned immediately."""
@@ -4,11 +4,13 @@ import time
import pytest
from nanobot.providers.base import ProviderCallContext
from nanobot.providers.openai_compat_provider import (
_RESPONSES_FAILURE_THRESHOLD,
_RESPONSES_PROBE_INTERVAL_S,
OpenAICompatProvider,
)
from nanobot.providers.openai_responses.state import build_responses_state
@pytest.fixture()
@@ -28,6 +30,52 @@ def test_responses_api_available_by_default(provider):
assert provider._should_use_responses_api("gpt-5", None) is True
def test_deepseek_v4_flash_uses_responses_by_model(provider):
provider._spec = type("Spec", (), {
"name": "deepseek",
"responses_models": ("deepseek-v4-flash",),
"strip_model_prefix": False,
"strip_model_prefixes": (),
})()
provider._effective_base = "https://api.deepseek.com"
provider.default_model = "deepseek-v4-flash"
assert provider._should_use_responses_api("deepseek-v4-flash", None) is True
assert provider._should_use_responses_api("deepseek-v4-pro", None) is False
def test_deepseek_v4_flash_matches_provider_prefixed_model(provider):
provider._spec = type("Spec", (), {
"name": "deepseek",
"responses_models": ("deepseek-v4-flash",),
"strip_model_prefix": False,
"strip_model_prefixes": (),
})()
provider._effective_base = "https://api.deepseek.com"
assert provider._should_use_responses_api("deepseek/deepseek-v4-flash", None) is True
def test_direct_openai_enables_server_compaction(provider):
provider._extra_body = {}
body = provider._build_responses_body(
messages=[{"role": "user", "content": "hello"}],
tools=None,
model="gpt-5.6",
max_tokens=30_000,
temperature=0.1,
reasoning_effort="high",
tool_choice=None,
provider_context=ProviderCallContext(context_window_tokens=100_000),
)
assert body["context_management"] == [{
"type": "compaction",
"compact_threshold": 70_000,
}]
def test_api_type_chat_completions_disables_responses(provider):
provider._api_type = "chat_completions"
assert provider._should_use_responses_api("gpt-5", None) is False
@@ -103,3 +151,140 @@ def test_reasoning_effort_key_is_case_insensitive(provider):
for _ in range(_RESPONSES_FAILURE_THRESHOLD):
provider._record_responses_failure("o3", "High")
assert provider._should_use_responses_api("o3", "high") is False
# ======================================================================
# _should_fallback_from_responses_error
# ======================================================================
class _FakeAPIError(Exception):
def __init__(self, status_code, body):
super().__init__(str(body))
self.status_code = status_code
self.body = body
self.response = None
def test_serde_deserialize_error_does_not_trigger_fallback():
# Serde errors can also identify malformed user-provided request fields.
# The known DeepSeek wire-shape bug is fixed at serialization time instead.
err = _FakeAPIError(400, {
"message": (
"Failed to deserialize the JSON body into the target type: "
"input: invalid type: string \"Michael topped up DeepSeek ...\", "
"expected a sequence at line 1 column 268612"
),
"type": "invalid_request_error",
"param": None,
})
assert OpenAICompatProvider._should_fallback_from_responses_error(err) is False
def test_legacy_compatibility_markers_still_trigger_fallback():
err = _FakeAPIError(400, "parameter `instructions` is unsupported")
assert OpenAICompatProvider._should_fallback_from_responses_error(err) is True
# ======================================================================
# DeepSeek Responses wire shape (PR #5214 root cause)
# ======================================================================
def _deepseek_provider(provider):
provider._spec = type("Spec", (), {
"name": "deepseek",
"responses_models": ("deepseek-v4-flash",),
"strip_model_prefix": False,
"strip_model_prefixes": (),
})()
provider._effective_base = "https://api.deepseek.com"
provider.default_model = "deepseek-v4-flash"
provider._extra_body = {}
return provider
def test_deepseek_full_history_body_keeps_reasoning_content_as_array(provider):
# Full-history fixture: DeepSeek's Responses gateway rejects reasoning
# items whose ``content`` is a plain string ("input: invalid type: string
# ..., expected a sequence"); the wire body must keep it as a part list.
_deepseek_provider(provider)
body = provider._build_responses_body(
messages=[
{
"role": "assistant",
"reasoning_content": "Michael topped up DeepSeek with $10.",
"content": "All systems aligned now.",
},
{"role": "user", "content": "audit the custom tools"},
],
tools=None,
model="deepseek-v4-flash",
max_tokens=1000,
temperature=0.1,
reasoning_effort=None,
tool_choice=None,
)
reasoning_items = [item for item in body["input"] if item.get("type") == "reasoning"]
assert len(reasoning_items) == 1
assert reasoning_items[0]["content"] == [
{"type": "output_text", "text": "Michael topped up DeepSeek with $10."},
]
def test_deepseek_replay_body_keeps_reasoning_content_as_array(provider):
# Replay/consolidation fixture: after token consolidation clears
# provider_state the next turn converts full history on top of the
# replayed prior items. Both replayed and converted reasoning items must
# keep list content on the wire.
_deepseek_provider(provider)
prior_items = [
{
"type": "reasoning",
"id": "rs_1",
"content": [{"type": "output_text", "text": "prior reasoning"}],
},
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "prior answer"}],
"status": "completed",
"id": "msg_0",
},
]
state = build_responses_state(
provider=provider._responses_state_provider(),
model="deepseek-v4-flash",
input_items=prior_items,
output_items=[],
).with_pending_messages([
{
"role": "assistant",
"reasoning_content": "think first",
"content": "answer",
},
{"role": "user", "content": "audit the custom tools"},
])
body = provider._build_responses_body(
messages=[
{"role": "system", "content": "You are KITT."},
{"role": "user", "content": "audit the custom tools"},
],
tools=None,
model="deepseek-v4-flash",
max_tokens=1000,
temperature=0.1,
reasoning_effort=None,
tool_choice=None,
provider_context=ProviderCallContext(conversation_state=state),
)
reasoning_items = [item for item in body["input"] if item.get("type") == "reasoning"]
assert len(reasoning_items) == 2 # one replayed from state, one converted
for item in reasoning_items:
assert isinstance(item["content"], list)
assert item["content"][0]["type"] == "output_text"
-14
View File
@@ -9,12 +9,9 @@ from unittest.mock import patch
import pytest
from nanobot.security.network import (
Httpx2PinnedDNSAsyncTransport,
PinnedDNSAsyncTransport,
configure_ssrf_whitelist,
contains_internal_url,
env_proxy_applies_to_url,
httpx2_env_proxy_mounts,
httpx_env_proxy_mounts,
is_loopback_host,
pin_resolved_url_dns,
@@ -267,17 +264,6 @@ def test_env_proxy_helpers_respect_no_proxy(monkeypatch):
assert any(transport is None for transport in mounts.values())
assert any(transport is not None for transport in mounts.values())
httpx2_mounts = httpx2_env_proxy_mounts()
assert any(transport is None for transport in httpx2_mounts.values())
assert any(transport is not None for transport in httpx2_mounts.values())
def test_httpx_transports_share_global_dns_pin_lock():
assert (
Httpx2PinnedDNSAsyncTransport._resolver_lock
is PinnedDNSAsyncTransport._resolver_lock
)
# ---------------------------------------------------------------------------
# contains_internal_url — shell command scanning
+132
View File
@@ -16,9 +16,13 @@ from nanobot.agent import context as agent_context
from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools.context import RequestContext, bind_request_context, reset_request_context
from nanobot.agent.tools.exec_session import (
MAX_OUTPUT_CHARS,
ExecSessionManager,
ListExecSessionsTool,
WriteStdinTool,
_BoundedOutputBuffer,
_SessionPoll,
_truncate_output,
)
from nanobot.agent.tools.registry import is_tool_error_result
from nanobot.agent.tools.shell import ExecTool
@@ -143,6 +147,134 @@ def test_exec_session_accepts_max_output_tokens_alias(tmp_path):
assert "Exit code: 0" in result
def test_bounded_output_buffer_keeps_head_tail_and_exact_drop_count():
buffer = _BoundedOutputBuffer(10)
buffer.append("012345")
buffer.append("6789ABCDEF")
assert buffer.retained_chars == 10
assert buffer.drain() == ("01234BCDEF", 6)
assert buffer.retained_chars == 0
def test_exec_session_bounds_unpolled_stdout_and_stderr(tmp_path):
async def run() -> tuple[int, int, str, int]:
manager = ExecSessionManager()
tool = ExecTool(working_dir=str(tmp_path), timeout=5, session_manager=manager)
command = _python_command(
"import sys,time; time.sleep(0.05); "
"sys.stdout.write('OUT_HEAD' + 'o' * 200000 + 'OUT_TAIL'); "
"sys.stderr.write('ERR_HEAD' + 'e' * 200000 + 'ERR_TAIL')"
)
initial = await tool.execute(
command=command,
yield_time_ms=0,
max_output_chars=1000,
)
sid = _session_id(initial)
session = manager._sessions[sid]
await asyncio.wait_for(session.process.wait(), timeout=5)
await asyncio.wait_for(
asyncio.gather(session._stdout_task, session._stderr_task),
timeout=5,
)
retained_stdout = session._stdout.retained_chars
retained_stderr = session._stderr.retained_chars
poll = await manager.write(
session_id=sid,
chars=None,
close_stdin=False,
terminate=False,
yield_time_ms=0,
max_output_chars=1000,
)
return retained_stdout, retained_stderr, poll.output, poll.truncated_chars
retained_stdout, retained_stderr, output, truncated_chars = asyncio.run(run())
assert retained_stdout == 50000
assert retained_stderr == 50000
assert output.startswith("OUT_HEAD")
assert output.endswith("ERR_TAIL")
assert truncated_chars > 390000
def test_write_stdin_wait_for_keeps_aggregate_within_output_budget():
async def run() -> str:
manager = SimpleNamespace(
write=AsyncMock(side_effect=[
_SessionPoll(output="HEAD" + "a" * 596, done=False, exit_code=None),
_SessionPoll(output="b" * 600, done=False, exit_code=None),
_SessionPoll(output="c" * 590 + "TARGET", done=False, exit_code=None),
])
)
tool = WriteStdinTool(manager=manager)
return await tool._wait_for_output(
session_id="session",
chars=None,
close_stdin=False,
terminate=False,
wait_for="TARGET",
wait_timeout_ms=1000,
max_output_chars=1000,
)
result = asyncio.run(run())
assert result.startswith("HEAD")
assert "TARGET" in result
assert "(796 chars truncated from output)" in result
assert len(result) < 1100
def test_write_stdin_wait_for_searches_before_response_truncation():
async def run() -> tuple[str, list[int]]:
output = "A" * 1500 + "TARGET" + "B" * 1500
observed_limits: list[int] = []
async def write(
*,
session_id: str,
chars: str | None,
close_stdin: bool,
terminate: bool,
yield_time_ms: int,
max_output_chars: int,
owner_session_key: str | None,
) -> _SessionPoll:
del session_id, chars, close_stdin, terminate, yield_time_ms, owner_session_key
observed_limits.append(max_output_chars)
visible, truncated = _truncate_output(output, max_output_chars)
return _SessionPoll(
output=visible,
done=True,
exit_code=0,
truncated_chars=truncated,
)
manager = SimpleNamespace(write=AsyncMock(side_effect=write))
tool = WriteStdinTool(manager=manager)
result = await tool._wait_for_output(
session_id="session",
chars=None,
close_stdin=False,
terminate=False,
wait_for="TARGET",
wait_timeout_ms=1000,
max_output_chars=1000,
)
return result, observed_limits
result, observed_limits = asyncio.run(run())
assert observed_limits == [MAX_OUTPUT_CHARS]
assert "Wait target not observed" not in result
assert "(2,006 chars truncated from output)" in result
assert len(result) < 1100
def test_exec_one_shot_accepts_max_output_tokens_alias(tmp_path):
async def run() -> str:
tool = ExecTool(working_dir=str(tmp_path), timeout=5)
+31 -34
View File
@@ -7,7 +7,7 @@ from contextlib import asynccontextmanager
from pathlib import Path
from types import ModuleType, SimpleNamespace
import httpx2 as httpx
import httpx
import pytest
import nanobot.agent.tools.mcp as mcp_mod
@@ -52,7 +52,7 @@ class _FakeBlobResourceContents:
class _FakeImageContent:
def __init__(self, data: str, mime_type: str = "image/png") -> None:
self.data = data
self.mime_type = mime_type
self.mimeType = mime_type
@pytest.fixture
@@ -111,7 +111,7 @@ def _fake_mcp_module(
@asynccontextmanager
async def _fake_streamable_http_client(_url: str, http_client=None):
yield object(), object()
yield object(), object(), object()
mod.ClientSession = _FakeClientSession
mod.StdioServerParameters = _FakeStdioServerParameters
@@ -133,13 +133,12 @@ def _fake_mcp_module(
shared_mod = ModuleType("mcp.shared")
exc_mod = ModuleType("mcp.shared.exceptions")
class _FakeMCPError(Exception):
class _FakeMcpError(Exception):
def __init__(self, code: int = -1, message: str = "error"):
self.error = SimpleNamespace(code=code, message=message)
super().__init__(message)
mod.MCPError = _FakeMCPError
exc_mod.MCPError = _FakeMCPError
exc_mod.McpError = _FakeMcpError
monkeypatch.setitem(sys.modules, "mcp.shared", shared_mod)
monkeypatch.setitem(sys.modules, "mcp.shared.exceptions", exc_mod)
@@ -148,7 +147,7 @@ def _make_wrapper(session: object, *, timeout: float = 0.1) -> MCPToolWrapper:
tool_def = SimpleNamespace(
name="demo",
description="demo tool",
input_schema={"type": "object", "properties": {}},
inputSchema={"type": "object", "properties": {}},
)
return MCPToolWrapper(session, "test", tool_def, tool_timeout=timeout)
@@ -186,7 +185,7 @@ def test_wrapper_preserves_non_nullable_unions() -> None:
tool_def = SimpleNamespace(
name="demo",
description="demo tool",
input_schema={
inputSchema={
"type": "object",
"properties": {
"value": {
@@ -208,7 +207,7 @@ def test_wrapper_normalizes_nullable_property_type_union() -> None:
tool_def = SimpleNamespace(
name="demo",
description="demo tool",
input_schema={
inputSchema={
"type": "object",
"properties": {
"name": {"type": ["string", "null"]},
@@ -225,7 +224,7 @@ def test_wrapper_normalizes_nullable_property_anyof() -> None:
tool_def = SimpleNamespace(
name="demo",
description="demo tool",
input_schema={
inputSchema={
"type": "object",
"properties": {
"name": {
@@ -250,7 +249,7 @@ def test_wrapper_hoists_recursive_local_refs_into_defs() -> None:
tool_def = SimpleNamespace(
name="search_dataset",
description="search tool",
input_schema={
inputSchema={
"type": "object",
"properties": {
"filter": {
@@ -283,7 +282,7 @@ def test_wrapper_hoists_root_self_ref_into_defs() -> None:
tool_def = SimpleNamespace(
name="tree",
description="tree tool",
input_schema={
inputSchema={
"type": "object",
"properties": {
"children": {"type": "array", "items": {"$ref": "#"}},
@@ -305,7 +304,7 @@ def test_wrapper_preserves_existing_defs_refs() -> None:
tool_def = SimpleNamespace(
name="demo",
description="demo tool",
input_schema={
inputSchema={
"type": "object",
"$defs": {"value": {"type": "string"}},
"properties": {"value": {"$ref": "#/$defs/value"}},
@@ -322,7 +321,7 @@ def test_wrapper_resolves_uri_encoded_json_pointer() -> None:
tool_def = SimpleNamespace(
name="demo",
description="demo tool",
input_schema={
inputSchema={
"type": "object",
"properties": {
"space name/value": {"type": "string"},
@@ -450,7 +449,7 @@ async def test_execute_wraps_mcp_is_error_result() -> None:
async def call_tool(_name: str, arguments: dict) -> object:
return SimpleNamespace(
content=[_FakeTextContent("Error: server-side MCP failure")],
is_error=True,
isError=True,
)
wrapper = _make_wrapper(SimpleNamespace(call_tool=call_tool))
@@ -495,7 +494,7 @@ async def test_execute_preserves_success_text_that_starts_with_error() -> None:
async def call_tool(_name: str, arguments: dict) -> object:
return SimpleNamespace(
content=[_FakeTextContent("Error: generated report successfully")],
is_error=False,
isError=False,
)
wrapper = _make_wrapper(SimpleNamespace(call_tool=call_tool))
@@ -623,7 +622,7 @@ def _make_tool_def(name: str) -> SimpleNamespace:
return SimpleNamespace(
name=name,
description=f"{name} tool",
input_schema={"type": "object", "properties": {}},
inputSchema={"type": "object", "properties": {}},
)
@@ -937,7 +936,7 @@ async def test_connect_mcp_servers_env_proxy_adds_proxy_mounts_and_keeps_pinned_
@asynccontextmanager
async def _capturing_streamable_http_client(_url: str, http_client=None):
assert http_client is not None
yield object(), object()
yield object(), object(), object()
monkeypatch.setenv("HTTPS_PROXY", "http://proxy.example:8080")
monkeypatch.setenv("NO_PROXY", "localhost,127.0.0.1,::1")
@@ -945,11 +944,11 @@ async def test_connect_mcp_servers_env_proxy_adds_proxy_mounts_and_keeps_pinned_
monkeypatch.setattr(mcp_mod, "_probe_http_url", _reachable)
monkeypatch.setattr(
mcp_mod,
"Httpx2PinnedDNSAsyncTransport",
"PinnedDNSAsyncTransport",
lambda: httpx.MockTransport(lambda request: httpx.Response(200, request=request)),
)
monkeypatch.setattr(
"nanobot.security.network.httpx2.AsyncHTTPTransport",
"nanobot.security.network.httpx.AsyncHTTPTransport",
lambda **_kwargs: httpx.MockTransport(
lambda request: httpx.Response(200, request=request)
),
@@ -977,11 +976,11 @@ def test_mcp_http_clients_no_proxy_env_keeps_pinned_direct_route(monkeypatch):
monkeypatch.setenv("NO_PROXY", "mcp.example.com")
monkeypatch.setattr(
mcp_mod,
"Httpx2PinnedDNSAsyncTransport",
"PinnedDNSAsyncTransport",
lambda: httpx.MockTransport(lambda request: httpx.Response(200, request=request)),
)
monkeypatch.setattr(
"nanobot.security.network.httpx2.AsyncHTTPTransport",
"nanobot.security.network.httpx.AsyncHTTPTransport",
lambda **_kwargs: httpx.MockTransport(
lambda request: httpx.Response(200, request=request)
),
@@ -1051,15 +1050,13 @@ async def test_connect_mcp_servers_http_clients_reject_unsafe_redirect_targets(
assert http_client is not None
used_transports.append("streamableHttp")
await http_client.get("https://example.com/start")
yield object(), object()
yield object(), object(), object()
monkeypatch.setattr(mcp_mod, "validate_url_target", _validate)
monkeypatch.setattr(mcp_mod, "_probe_http_url", _reachable)
# Keep the redirect exercise isolated from host-level proxy settings.
monkeypatch.setattr(mcp_mod, "httpx2_env_proxy_mounts", lambda: {})
monkeypatch.setattr(
mcp_mod,
"Httpx2PinnedDNSAsyncTransport",
"PinnedDNSAsyncTransport",
lambda **_kwargs: httpx.MockTransport(_handler),
)
monkeypatch.setattr(mcp_mod.httpx, "AsyncClient", _async_client_with_mock_transport)
@@ -1141,13 +1138,13 @@ async def test_connect_mcp_servers_streamable_http_uses_finite_timeout(
@asynccontextmanager
async def _capturing_streamable_http_client(_url: str, http_client=None):
captured["timeout"] = http_client.timeout
yield object(), object()
yield object(), object(), object()
monkeypatch.setattr(mcp_mod, "validate_url_target", _validate)
monkeypatch.setattr(mcp_mod, "_probe_http_url", _reachable)
monkeypatch.setattr(
mcp_mod,
"Httpx2PinnedDNSAsyncTransport",
"PinnedDNSAsyncTransport",
lambda: httpx.MockTransport(lambda request: httpx.Response(200, request=request)),
)
monkeypatch.setattr(
@@ -1388,10 +1385,10 @@ async def test_prompt_wrapper_execute_handles_timeout() -> None:
@pytest.mark.asyncio
async def test_prompt_wrapper_execute_handles_mcp_error() -> None:
from mcp import MCPError
from mcp.shared.exceptions import McpError
async def get_prompt(name: str, arguments: dict | None = None) -> object:
raise MCPError(code=42, message="invalid argument")
raise McpError(code=42, message="invalid argument")
wrapper = _make_prompt_wrapper(SimpleNamespace(get_prompt=get_prompt))
result = await wrapper.execute()
@@ -1513,7 +1510,7 @@ def test_tool_wrapper_sanitizes_name() -> None:
tool_def = SimpleNamespace(
name="My Tool",
description="tool with spaces",
input_schema={"type": "object", "properties": {}},
inputSchema={"type": "object", "properties": {}},
)
wrapper = MCPToolWrapper(SimpleNamespace(call_tool=None), "srv", tool_def)
assert wrapper.name == "mcp_srv_My_Tool"
@@ -1544,7 +1541,7 @@ def test_tool_wrapper_preserves_original_name_for_mcp_call() -> None:
tool_def = SimpleNamespace(
name="My Tool",
description="tool with spaces",
input_schema={"type": "object", "properties": {}},
inputSchema={"type": "object", "properties": {}},
)
wrapper = MCPToolWrapper(SimpleNamespace(call_tool=None), "srv", tool_def)
# The sanitized API-facing name differs from the original MCP name
@@ -1622,12 +1619,12 @@ def test_long_server_name_tools_are_matched_by_server_name() -> None:
tool_def = SimpleNamespace(
name="search",
description="search tool",
input_schema={"type": "object", "properties": {}},
inputSchema={"type": "object", "properties": {}},
)
other_tool_def = SimpleNamespace(
name="search",
description="other search tool",
input_schema={"type": "object", "properties": {}},
inputSchema={"type": "object", "properties": {}},
)
wrapper = MCPToolWrapper(SimpleNamespace(call_tool=None), server_name, tool_def)
other_wrapper = MCPToolWrapper(SimpleNamespace(call_tool=None), "other", other_tool_def)
+33
View File
@@ -150,6 +150,36 @@ def test_enqueue_writes_trigger_run_record(tmp_path: Path) -> None:
assert record["content"] == "Review PR #4591"
assert record["origin_metadata"] == {"webui": True}
assert record["updated_at_ms"] > 0
stored = store.get(trigger.id)
assert stored is not None
assert stored.last_message == "Review PR #4591"
def test_enqueue_rolls_back_delivery_and_audit_when_trigger_save_fails(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
store = LocalTriggerStore(tmp_path)
trigger = store.create(
name="PR review",
channel="websocket",
chat_id="chat-1",
session_key="websocket:chat-1",
)
def fail_save(_triggers: list[LocalTrigger]) -> None:
raise OSError("store write failed")
monkeypatch.setattr(store, "_save_triggers_unlocked", fail_save)
with pytest.raises(OSError, match="store write failed"):
store.enqueue(trigger.id, "Review PR #4591")
assert list(store.inbox_dir.glob("*.json")) == []
assert list(store.runs_dir.glob("*.json")) == []
stored = LocalTriggerStore(tmp_path).get(trigger.id)
assert stored is not None
assert stored.last_message == ""
def test_delivery_run_record_truncates_large_content_and_response(tmp_path: Path) -> None:
@@ -168,6 +198,9 @@ def test_delivery_run_record_truncates_large_content_and_response(tmp_path: Path
assert queued_record["content"].startswith("content-")
assert queued_record["content"].endswith("\n... (truncated)")
assert len(queued_record["content"]) < len(large_content)
stored = store.get(trigger.id)
assert stored is not None
assert stored.last_message == queued_record["content"]
store.write_delivery_run_record(
delivery,
+22
View File
@@ -310,3 +310,25 @@ class TestNestedRepoProtection:
assert result is False
assert not (workspace / ".git").exists()
class TestCommitIdEncoding:
"""Commit ids must be usable with git, not hex-of-hex."""
def test_auto_commit_returns_the_real_short_sha(self, git, tmp_path):
(tmp_path / "MEMORY.md").write_text("- a fact\n", encoding="utf-8")
sha = git.auto_commit("memory update")
expected = subprocess.run(
["git", "-C", str(tmp_path), "log", "-1", "--format=%h", "--abbrev=8"],
capture_output=True, text=True, check=True,
).stdout.strip()
assert sha == expected
def test_a_real_git_sha_resolves(self, git, tmp_path):
(tmp_path / "MEMORY.md").write_text("- a fact\n", encoding="utf-8")
git.auto_commit("memory update")
real = subprocess.run(
["git", "-C", str(tmp_path), "log", "-1", "--format=%h", "--abbrev=8"],
capture_output=True, text=True, check=True,
).stdout.strip()
assert git._resolve_sha(real) is not None
+1
View File
@@ -13,6 +13,7 @@ from nanobot.webui.transcript import (
def test_delete_webui_thread_removes_legacy_json_and_transcript(tmp_path, monkeypatch) -> None:
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
monkeypatch.setattr("nanobot.webui.transcript._MAX_TRANSCRIPT_FILE_BYTES", 520)
monkeypatch.setattr("nanobot.webui.transcript._ACTIVE_TRANSCRIPT_ROTATE_BYTES", 520)
monkeypatch.setattr("nanobot.webui.transcript._TARGET_ACTIVE_TRANSCRIPT_BYTES", 260)
key = "websocket:k1"
json_path = webui_thread_file_path(key)
+176 -1
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
import nanobot.webui.transcript as transcript_module
from nanobot.session.history_visibility import HIDDEN_HISTORY_META
from nanobot.webui.transcript import (
WEBUI_TRANSCRIPT_SCHEMA_VERSION,
@@ -38,6 +39,7 @@ def test_append_stamps_created_at_ms(tmp_path, monkeypatch) -> None:
def _force_small_transcript_budget(monkeypatch, *, limit: int = 520, target: int = 260) -> None:
monkeypatch.setattr("nanobot.webui.transcript._MAX_TRANSCRIPT_FILE_BYTES", limit)
monkeypatch.setattr("nanobot.webui.transcript._ACTIVE_TRANSCRIPT_ROTATE_BYTES", limit)
monkeypatch.setattr("nanobot.webui.transcript._TARGET_ACTIVE_TRANSCRIPT_BYTES", target)
@@ -122,6 +124,28 @@ def test_segmented_transcript_paginates_latest_and_older_without_overlap(
]
def test_latest_page_reads_active_chunk_once(tmp_path, monkeypatch) -> None:
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
key = "websocket:single-active-read"
for idx in range(1, 7):
_append_numbered_turn(key, "single-active-read", idx)
original = transcript_module._read_chunk_turns
read_chunk_ids: list[str] = []
def track_read(session_key: str, chunk_id: str) -> list[list[dict]]:
read_chunk_ids.append(chunk_id)
return original(session_key, chunk_id)
monkeypatch.setattr(transcript_module, "_read_chunk_turns", track_read)
latest = build_webui_thread_response(key, limit=4, direction="latest")
assert latest is not None
assert _message_contents(latest) == _numbered_turn_texts(5, 6)
assert read_chunk_ids == ["active"]
def test_page_cursor_survives_active_rotation_after_latest_page(
tmp_path,
monkeypatch,
@@ -148,15 +172,53 @@ def test_segment_manifest_can_be_rebuilt_when_missing_or_corrupt(tmp_path, monke
key = "websocket:manifest"
_write_segmented_turns(tmp_path, monkeypatch, key, "manifest", 4)
manifest = webui_transcript_segments_dir(key) / "manifest.json"
segment_dir = webui_transcript_segments_dir(key)
segment_names = sorted(path.name for path in segment_dir.glob("*.jsonl"))
assert segment_names
original = transcript_module._read_transcript_file
segment_reads: list[str] = []
def track_read(path):
if path.parent == segment_dir and path.suffix == ".jsonl":
segment_reads.append(path.name)
return original(path)
monkeypatch.setattr(transcript_module, "_read_transcript_file", track_read)
manifest = segment_dir / "manifest.json"
manifest.write_text("{not json", encoding="utf-8")
entries = transcript_module._read_segment_manifest_entries(key)
assert [entry["id"] for entry in entries] == [path.removesuffix(".jsonl") for path in segment_names]
assert segment_reads == segment_names
lines = read_transcript_lines(key)
assert len([line for line in lines if line.get("event") == "user"]) == 4
assert manifest.read_text(encoding="utf-8").lstrip().startswith("{")
def test_rotation_does_not_reread_existing_segments(tmp_path, monkeypatch) -> None:
key = "websocket:manifest-append"
_write_segmented_turns(tmp_path, monkeypatch, key, "manifest-append", 4)
segment_dir = webui_transcript_segments_dir(key)
assert list(segment_dir.glob("*.jsonl"))
original = transcript_module._read_transcript_file
segment_reads: list[str] = []
def track_read(path):
if path.parent == segment_dir and path.suffix == ".jsonl":
segment_reads.append(path.name)
return original(path)
monkeypatch.setattr(transcript_module, "_read_transcript_file", track_read)
for idx in range(5, 9):
_append_numbered_turn(key, "manifest-append", idx)
assert segment_reads == []
def test_delete_webui_transcript_removes_segments(tmp_path, monkeypatch) -> None:
from nanobot.webui.thread_disk import webui_thread_file_path
from nanobot.webui.transcript import delete_webui_transcript, webui_transcript_path
@@ -696,6 +758,42 @@ def test_replay_preserves_local_trigger_source_metadata(tmp_path, monkeypatch) -
assert msgs[0]["source"] == {"kind": "local_trigger", "label": "PR review"}
def test_replay_preserves_automation_source_metadata_on_streamed_reply(
tmp_path,
monkeypatch,
) -> None:
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
key = "websocket:t-streamed-cron-source"
source = {"kind": "cron", "label": "Repo check"}
for record in (
{
"event": "delta",
"chat_id": "t-streamed-cron-source",
"text": "Repo ",
"source": source,
},
{
"event": "delta",
"chat_id": "t-streamed-cron-source",
"text": "clean.",
"source": source,
},
{
"event": "stream_end",
"chat_id": "t-streamed-cron-source",
"source": source,
},
{"event": "turn_end", "chat_id": "t-streamed-cron-source"},
):
append_transcript_object(key, record)
msgs = replay_transcript_to_ui_messages(read_transcript_lines(key))
assert msgs[0]["content"] == "Repo clean."
assert msgs[0]["source"] == source
def test_replay_preserves_legacy_trigger_source_metadata(tmp_path, monkeypatch) -> None:
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
key = "websocket:t-trigger-source"
@@ -750,6 +848,83 @@ def test_build_response_restores_session_users_for_legacy_transcript(
]
def test_complete_transcript_does_not_load_session_messages(tmp_path, monkeypatch) -> None:
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
key = "websocket:complete-fast-path"
for event in (
{"event": "user", "chat_id": "complete-fast-path", "text": "question"},
{"event": "message", "chat_id": "complete-fast-path", "text": "answer"},
{"event": "turn_end", "chat_id": "complete-fast-path"},
):
append_transcript_object(key, event)
def fail_if_loaded() -> list[dict]:
raise AssertionError("complete transcripts must not read canonical session history")
out = build_webui_thread_response(
key,
limit=4,
direction="latest",
session_messages_loader=fail_if_loaded,
)
assert out is not None
assert [(message["role"], message["content"]) for message in out["messages"]] == [
("user", "question"),
("assistant", "answer"),
]
def test_legacy_recovery_loads_session_and_builds_backfill_turns_once(
tmp_path,
monkeypatch,
) -> None:
monkeypatch.setattr("nanobot.config.paths.get_data_dir", lambda: tmp_path)
key = "websocket:lazy-legacy-recovery"
append_transcript_object(
key,
{"event": "message", "chat_id": "lazy-legacy-recovery", "text": "answer"},
)
append_transcript_object(
key,
{
"event": "turn_end",
"chat_id": "lazy-legacy-recovery",
"transcript_incomplete": True,
},
)
loader_calls = 0
backfill_calls = 0
original = transcript_module._session_backfill_turns
def load_session_messages() -> list[dict]:
nonlocal loader_calls
loader_calls += 1
return [
{"role": "user", "content": "question"},
{"role": "assistant", "content": "answer"},
]
def track_backfill(session_key: str, session_messages: list[dict]):
nonlocal backfill_calls
backfill_calls += 1
return original(session_key, session_messages)
monkeypatch.setattr(transcript_module, "_session_backfill_turns", track_backfill)
out = build_webui_thread_response(key, session_messages_loader=load_session_messages)
assert out is not None
assert loader_calls == 1
assert backfill_calls == 1
assert [(message["role"], message["content"]) for message in out["messages"]] == [
("user", "question"),
("assistant", "answer"),
]
assert out["has_pending_tool_calls"] is False
def test_build_response_restores_session_users_without_duplicating_new_transcript_users(
tmp_path,
monkeypatch,
+81 -2
View File
@@ -1,9 +1,14 @@
import json
from unittest.mock import MagicMock
import pytest
from nanobot.security.workspace_access import WorkspaceScopeError, default_workspace_scope
from nanobot.session.manager import SessionManager
from nanobot.security.workspace_access import (
WORKSPACE_SCOPE_METADATA_KEY,
WorkspaceScopeError,
default_workspace_scope,
)
from nanobot.session.manager import SessionManager, SessionStore
from nanobot.webui.workspaces import (
WebUIWorkspaceController,
read_webui_default_access_mode,
@@ -135,6 +140,33 @@ def test_webui_default_access_applies_to_unscoped_old_sessions(tmp_path, monkeyp
assert new_scope.access_mode == "full"
def test_indexed_scope_preserves_missing_and_explicit_null_semantics(tmp_path, monkeypatch) -> None:
monkeypatch.setattr("nanobot.webui.workspaces.get_webui_dir", lambda: tmp_path / "webui")
default = tmp_path / "default"
default.mkdir()
write_webui_default_access_mode("full")
controller = WebUIWorkspaceController(
session_manager=None,
default_workspace=default,
default_restrict_to_workspace=True,
)
webui_default = controller.default_scope()
missing = controller.scope_for_indexed_metadata(
None,
scope_present=False,
default_scope=webui_default,
)
explicit_null = controller.scope_for_indexed_metadata(
None,
scope_present=True,
default_scope=webui_default,
)
assert missing.access_mode == "full"
assert explicit_null.access_mode == "restricted"
def test_webui_default_access_does_not_override_explicit_session_scope(tmp_path, monkeypatch) -> None:
monkeypatch.setattr("nanobot.webui.workspaces.get_webui_dir", lambda: tmp_path / "webui")
default = tmp_path / "default"
@@ -185,6 +217,53 @@ def test_scope_for_session_key_reads_metadata_without_full_history(
assert scope.access_mode == "full"
def test_scope_for_session_key_always_reads_the_active_store(tmp_path, monkeypatch) -> None:
monkeypatch.setattr("nanobot.webui.workspaces.get_webui_dir", lambda: tmp_path / "webui")
default = tmp_path / "default"
project = tmp_path / "project"
default.mkdir()
project.mkdir()
workspace = tmp_path / "session-data"
full_scope = default_workspace_scope(project, restrict_to_workspace=False)
restricted_scope = default_workspace_scope(project, restrict_to_workspace=True)
residual_sessions = SessionManager(workspace)
residual = residual_sessions.get_or_create("websocket:cached")
residual.metadata[WORKSPACE_SCOPE_METADATA_KEY] = full_scope.metadata()
residual_sessions.save(residual)
store = MagicMock(spec=SessionStore)
store.read_metadata.side_effect = [
{
"key": "websocket:cached",
"created_at": None,
"updated_at": None,
"metadata": {WORKSPACE_SCOPE_METADATA_KEY: full_scope.metadata()},
},
{
"key": "websocket:cached",
"created_at": None,
"updated_at": None,
"metadata": {WORKSPACE_SCOPE_METADATA_KEY: restricted_scope.metadata()},
},
]
sessions = SessionManager(workspace, store=store)
controller = WebUIWorkspaceController(
session_manager=sessions,
default_workspace=default,
default_restrict_to_workspace=True,
)
first = controller.scope_for_session_key("websocket:cached")
second = controller.scope_for_session_key("websocket:cached")
assert first.project_path == project.resolve()
assert first.access_mode == "full"
assert second.project_path == project.resolve()
assert second.access_mode == "restricted"
assert store.read_metadata.call_count == 2
def test_remote_existing_chat_can_reduce_its_workspace_access(tmp_path, monkeypatch) -> None:
monkeypatch.setattr("nanobot.webui.workspaces.get_webui_dir", lambda: tmp_path / "webui")
default = tmp_path / "default"
+38
View File
@@ -0,0 +1,38 @@
"""Tests for shared embedded WebUI HTTP helpers."""
import gzip
import json
from nanobot.webui.http_utils import http_json_response
def test_http_json_response_compresses_large_payload_when_gzip_is_accepted() -> None:
payload = {"message": "响应内容" * 2_000}
response = http_json_response(payload, accept_encoding="br, gzip; q=0.5")
assert response.headers["Content-Encoding"] == "gzip"
assert response.headers["Vary"] == "Accept-Encoding"
assert int(response.headers["Content-Length"]) == len(response.body)
assert json.loads(gzip.decompress(response.body)) == payload
def test_http_json_response_preserves_identity_when_gzip_is_rejected() -> None:
payload = {"message": "x" * 8_000}
response = http_json_response(payload, accept_encoding="gzip;q=0, br")
assert "Content-Encoding" not in response.headers
assert response.headers["Vary"] == "Accept-Encoding"
assert int(response.headers["Content-Length"]) == len(response.body)
assert json.loads(response.body) == payload
def test_http_json_response_does_not_compress_small_payload() -> None:
payload = {"ok": True}
response = http_json_response(payload, accept_encoding="gzip")
assert "Content-Encoding" not in response.headers
assert response.headers["Vary"] == "Accept-Encoding"
assert json.loads(response.body) == payload
+119 -5
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
import io
import os
from datetime import datetime
from pathlib import Path
@@ -8,6 +9,8 @@ import pytest
import nanobot.webui.session_list_index as session_list_index
from nanobot.cron.session_turns import CRON_HISTORY_META
from nanobot.providers.base import ProviderConversationState
from nanobot.security.workspace_access import WORKSPACE_SCOPE_METADATA_KEY
from nanobot.session.automation_turns import AUTOMATION_HISTORY_META
from nanobot.session.history_visibility import HIDDEN_HISTORY_META
from nanobot.session.manager import SessionManager
@@ -27,7 +30,7 @@ def test_webui_session_list_reuses_valid_index_without_scanning_files(
assert list_webui_sessions(manager)[0]["preview"] == "indexed preview"
assert list_webui_sessions(manager)[0]["model_preset"] == "fast"
def fail_scan(session_manager: SessionManager, path: Path) -> None:
def fail_scan(session_manager: SessionManager, path: Path, webui_dir: Path) -> None:
raise AssertionError(f"unexpected session file scan: {path}")
monkeypatch.setattr(session_list_index, "_scan_session_row", fail_scan)
@@ -39,6 +42,89 @@ def test_webui_session_list_reuses_valid_index_without_scanning_files(
assert rows[0]["model_preset"] == "fast"
def test_webui_session_list_indexes_workspace_scope_and_preserves_null(
tmp_path: Path,
) -> None:
manager = SessionManager(tmp_path)
project = tmp_path / "project"
project.mkdir()
scoped = manager.get_or_create("websocket:scoped")
scoped.metadata[WORKSPACE_SCOPE_METADATA_KEY] = {
"project_path": str(project),
"access_mode": "full",
"future_extension": "x" * 5000,
}
manager.save(scoped)
explicit_null = manager.get_or_create("websocket:null")
explicit_null.metadata[WORKSPACE_SCOPE_METADATA_KEY] = None
manager.save(explicit_null)
manager.save(manager.get_or_create("websocket:missing"))
rows = {row["key"]: row for row in list_webui_sessions(manager)}
assert session_list_index.indexed_workspace_scope(rows["websocket:scoped"]) == (
True,
{"project_path": str(project), "access_mode": "full"},
)
assert session_list_index.indexed_workspace_scope(rows["websocket:null"]) == (True, None)
assert session_list_index.indexed_workspace_scope(rows["websocket:missing"]) == (False, None)
scoped.metadata[WORKSPACE_SCOPE_METADATA_KEY]["access_mode"] = "restricted"
manager.save(scoped)
refreshed = {row["key"]: row for row in list_webui_sessions(manager)}
assert session_list_index.indexed_workspace_scope(refreshed["websocket:scoped"])[1] == {
"project_path": str(project),
"access_mode": "restricted",
}
def test_webui_session_list_does_not_cache_old_snapshot_with_new_signature(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
manager = SessionManager(tmp_path)
session_key = "websocket:scope-race"
session = manager.get_or_create(session_key)
session.metadata[WORKSPACE_SCOPE_METADATA_KEY] = {
"project_path": str(tmp_path),
"access_mode": "full",
}
session.add_message("user", "hello")
manager.save(session)
session_path = manager._get_session_path(session_key)
original_open = open
scope_changed = False
class RacingReader(io.StringIO):
def __next__(self) -> str:
nonlocal scope_changed
if not scope_changed:
scope_changed = True
current = manager.get_or_create(session_key)
current.metadata[WORKSPACE_SCOPE_METADATA_KEY] = {
"project_path": str(tmp_path),
"access_mode": "restricted",
}
manager.save(current)
return super().__next__()
def racing_open(path, *args, **kwargs):
if Path(path) == session_path:
with original_open(path, *args, **kwargs) as source:
return RacingReader(source.read())
return original_open(path, *args, **kwargs)
monkeypatch.setattr(session_list_index, "open", racing_open, raising=False)
first = list_webui_sessions(manager)[0]
second = list_webui_sessions(manager)[0]
assert session_list_index.indexed_workspace_scope(first)[1]["access_mode"] == "full"
assert session_list_index.indexed_workspace_scope(second)[1]["access_mode"] == "restricted"
def test_webui_session_list_rejects_invalid_internal_model_preset_metadata(
tmp_path: Path,
) -> None:
@@ -73,9 +159,13 @@ def test_webui_session_list_rescans_only_changed_file(tmp_path: Path, monkeypatc
original_scan = session_list_index._scan_session_row
scanned: list[str] = []
def record_scan(session_manager: SessionManager, path: Path) -> dict | None:
def record_scan(
session_manager: SessionManager,
path: Path,
webui_dir: Path,
) -> dict | None:
scanned.append(path.name)
return original_scan(session_manager, path)
return original_scan(session_manager, path, webui_dir)
monkeypatch.setattr(session_list_index, "_scan_session_row", record_scan)
@@ -85,6 +175,26 @@ def test_webui_session_list_rescans_only_changed_file(tmp_path: Path, monkeypatc
assert {row["preview"] for row in rows} == {"first", "second after"}
def test_webui_session_list_skips_provider_state_before_preview_budget(
tmp_path: Path,
monkeypatch,
) -> None:
monkeypatch.setattr(session_list_index, "_SESSION_LIST_PREVIEW_MAX_CHARS", 100)
manager = SessionManager(tmp_path)
session = manager.get_or_create("websocket:private-state")
session.provider_state = ProviderConversationState(
kind="openai_responses",
provider="openai:test",
model="test-model",
version=1,
payload={"items": [{"encrypted_content": "x" * 200}]},
)
session.add_message("user", "visible preview")
manager.save(session)
assert list_webui_sessions(manager)[0]["preview"] == "visible preview"
def test_webui_session_list_drops_deleted_index_rows(tmp_path: Path) -> None:
manager = SessionManager(tmp_path)
session = manager.get_or_create("websocket:deleted")
@@ -226,9 +336,13 @@ def test_webui_session_list_rescans_when_transcript_changes(
original_scan = session_list_index._scan_session_row
scanned: list[str] = []
def record_scan(session_manager: SessionManager, path: Path) -> dict | None:
def record_scan(
session_manager: SessionManager,
path: Path,
webui_dir: Path,
) -> dict | None:
scanned.append(path.name)
return original_scan(session_manager, path)
return original_scan(session_manager, path, webui_dir)
monkeypatch.setattr(session_list_index, "_scan_session_row", record_scan)
+20
View File
@@ -86,6 +86,26 @@ def test_settings_payload_includes_versioned_docs(
}
def test_settings_payload_exposes_edenai_provider(
tmp_path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
config_path = tmp_path / "config.json"
config = Config()
config.providers.edenai.api_key = "eden-test-key"
save_config(config, config_path)
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
payload = settings_payload()
edenai = next(row for row in payload["providers"] if row["name"] == "edenai")
assert edenai["label"] == "Eden AI"
assert edenai["configured"] is True
assert edenai["default_api_base"] == "https://api.edenai.run/v3"
assert edenai["model_catalog"] == "catalog"
assert edenai["model_selectable"] is True
def test_settings_payload_includes_relocated_capabilities(
tmp_path,
monkeypatch: pytest.MonkeyPatch,

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