mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-25 14:53:34 +00:00
Compare commits
12
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d531adaf18 | ||
|
|
302fa8a38d | ||
|
|
1b9a8ad1b3 | ||
|
|
8c1553e2c9 | ||
|
|
3b57d5bec9 | ||
|
|
11c5e25422 | ||
|
|
12ba79ffc3 | ||
|
|
d3cf637d4a | ||
|
|
b6cf2a6b73 | ||
|
|
9851f2a6ae | ||
|
|
8dc4ca5a98 | ||
|
|
9f71bd0d44 |
@@ -6,6 +6,18 @@ DATABASE_POOL_SIZE=10
|
||||
# LLM_BACKEND=nearai # default
|
||||
# Possible values: nearai, ollama, openai_compatible, openai, anthropic, tinfoil
|
||||
|
||||
# === Anthropic Direct ===
|
||||
# Two auth modes:
|
||||
# 1. API key: Set ANTHROPIC_API_KEY (from console.anthropic.com/settings/keys)
|
||||
# 2. OAuth token: Set ANTHROPIC_OAUTH_TOKEN (from `claude login`)
|
||||
# OAuth tokens use Authorization: Bearer instead of x-api-key header.
|
||||
# ANTHROPIC_API_KEY=sk-ant-...
|
||||
# ANTHROPIC_OAUTH_TOKEN=sk-ant-oat01-... # from `claude login` credentials
|
||||
# ANTHROPIC_MODEL=claude-sonnet-4-20250514
|
||||
|
||||
# === OpenAI Direct ===
|
||||
# OPENAI_API_KEY=sk-...
|
||||
|
||||
# === NEAR AI (Chat Completions API) ===
|
||||
# Two auth modes:
|
||||
# 1. Session token (default): Uses browser OAuth (GitHub/Google) on first run.
|
||||
|
||||
@@ -1,3 +1,31 @@
|
||||
# Code Coverage Workflow
|
||||
#
|
||||
# This workflow runs test coverage analysis and uploads reports to Codecov.
|
||||
# Coverage reports help identify untested code paths and maintain code quality.
|
||||
#
|
||||
# What it does:
|
||||
# - Runs unit and integration tests with coverage instrumentation
|
||||
# - Runs E2E tests with coverage instrumentation
|
||||
# - Uploads coverage reports to Codecov (https://codecov.io/gh/nearai/ironclaw)
|
||||
#
|
||||
# Viewing coverage reports:
|
||||
# - PRs automatically get coverage comments showing changes in coverage
|
||||
# - Visit https://codecov.io/gh/nearai/ironclaw for detailed coverage reports
|
||||
# - Coverage reports are generated for three configurations:
|
||||
# 1. all-features: Full feature set
|
||||
# 2. default: Default features
|
||||
# 3. libsql-only: Minimal libSQL-only configuration
|
||||
# - E2E coverage tracks end-to-end test coverage separately
|
||||
#
|
||||
# Coverage files:
|
||||
# - Unit/integration: lcov.info (uploaded to Codecov with "unit" flag)
|
||||
# - E2E: e2e-coverage.info (uploaded to Codecov with "e2e" flag)
|
||||
#
|
||||
# Requirements:
|
||||
# - Uses cargo-llvm-cov for coverage instrumentation
|
||||
# - Requires PostgreSQL for integration tests (pgvector/pgvector:pg16)
|
||||
# - E2E tests require Python 3.12 and Playwright
|
||||
|
||||
name: Code Coverage
|
||||
on:
|
||||
push:
|
||||
|
||||
@@ -22,3 +22,4 @@ bench-results/
|
||||
# WASM build artifacts (loaded from disk, not bundled)
|
||||
*.wasm
|
||||
|
||||
trace_*.json
|
||||
|
||||
@@ -387,12 +387,27 @@ Dead code behind the wrong `#[cfg]` gate will only show up when building with a
|
||||
|
||||
**Zero clippy warnings policy:** Fix ALL clippy warnings before committing, including pre-existing ones in files you didn't change. Never leave warnings behind — treat `cargo clippy` output as a zero-tolerance gate.
|
||||
|
||||
**Transaction safety:** Multi-step database operations (INSERT+INSERT, UPDATE+DELETE, read-then-write) MUST be wrapped in a transaction. Never assume sequential calls are atomic. Before committing DB code, ask: "If this crashes between step N and N+1, is the database consistent?" If not, wrap in a transaction. This applies to both postgres and libsql backends.
|
||||
|
||||
**UTF-8 string safety:** Never use byte-index slicing (`&s[..n]`) on user-supplied or external strings — it panics on multi-byte characters. Use `is_char_boundary()` to walk backwards from the desired length, or iterate with `char_indices()`. Grep for `[..` in changed files to catch violations.
|
||||
|
||||
**Case-insensitive comparisons:** When comparing user-supplied strings (file paths, media types, extension names), always normalize to lowercase first with `.to_ascii_lowercase()`. On case-insensitive filesystems (macOS, Windows), path comparisons must be case-insensitive. File extension checks (`.png`, `.jpg`) and media type checks (`image/jpeg`) are common offenders.
|
||||
|
||||
**Decorator/wrapper trait delegation:** When adding a new method to `LlmProvider` (or any trait with decorator wrappers), you MUST update ALL wrapper types to delegate to their inner provider. Grep for `impl LlmProvider for` to find all implementations. Add a test that exercises the method through the full provider chain (`build_provider_chain()`), not just the base impl.
|
||||
|
||||
**Sensitive data in logs & events:** Tool parameters and outputs MUST be redacted before logging or broadcasting via SSE/WebSocket. Use `redact_params()` before any `tracing::info!`, `JobEvent`, or SSE emission that includes tool call data. Never log raw parameters from tool calls.
|
||||
|
||||
**Test temporary files:** Use the `tempfile` crate for test directories/files. Never hardcode `/tmp/...` paths — they collide in parallel test runs and break on non-Unix platforms.
|
||||
|
||||
**Trust boundaries in multi-process architecture:** Data from worker containers is untrusted. The orchestrator MUST validate: tool domain (never execute `Container`-domain tools on the host), nesting depth (server-side tracking, not client-supplied), and parameter sensitivity (redact before logging/broadcasting).
|
||||
|
||||
**Mechanical verification before committing:** Run these checks on changed files before committing:
|
||||
- `cargo clippy --all --benches --tests --examples --all-features` -- zero warnings
|
||||
- `grep -rnE '\.unwrap\(|\.expect\(' <files>` -- no panics in production
|
||||
- `grep -rn 'super::' <files>` -- use `crate::` imports
|
||||
- If you fixed a pattern bug, `grep` for other instances of that pattern across `src/`
|
||||
- Fix commits must include regression tests (enforced by `commit-msg` hook; bypass with `[skip-regression-check]`)
|
||||
- Run `scripts/pre-commit-safety.sh` to catch UTF-8, case-sensitivity, hardcoded /tmp, and logging issues
|
||||
|
||||
## Configuration
|
||||
|
||||
|
||||
+1
-1
@@ -56,7 +56,7 @@ rustls = { version = "0.23", optional = true, default-features = false }
|
||||
rustls-native-certs = { version = "0.8", optional = true }
|
||||
|
||||
# Database - libSQL/Turso (optional embedded database)
|
||||
libsql = { version = "0.6", optional = true, default-features = false, features = ["core", "replication"] }
|
||||
libsql = { version = "0.6", optional = true, default-features = false, features = ["core", "replication", "remote", "tls"] }
|
||||
|
||||
# Error handling
|
||||
thiserror = "2"
|
||||
|
||||
+7
-3
@@ -215,9 +215,13 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| NEAR AI | ✅ | ✅ | - | Primary provider |
|
||||
| Anthropic (Claude) | ✅ | 🚧 | - | Via NEAR AI proxy; Opus 4.5, Sonnet 4, Sonnet 4.6 |
|
||||
| OpenAI | ✅ | 🚧 | - | Via NEAR AI proxy |
|
||||
| AWS Bedrock | ✅ | ❌ | P3 | |
|
||||
| Google Gemini | ✅ | ❌ | P3 | |
|
||||
| NVIDIA API | ✅ | ❌ | P3 | New provider |
|
||||
| AWS Bedrock | ✅ | ✅ | P3 | Via `openai_compatible` adapter (e.g. LiteLLM) |
|
||||
| Google Gemini | ✅ | ✅ | P3 | Via `gemini` adapter |
|
||||
| io.net | ✅ | ✅ | P3 | Via `ionet` adapter |
|
||||
| Mistral | ✅ | ✅ | P3 | Via `mistral` adapter |
|
||||
| Yandex AI Studio | ✅ | ✅ | P3 | Via `yandex` adapter |
|
||||
| Cloudflare Workers AI | ✅ | ✅ | P3 | Via `cloudflare` adapter |
|
||||
| NVIDIA API | ✅ | ✅ | P3 | Via `nvidia` adapter and `providers.json` |
|
||||
| OpenRouter | ✅ | ✅ | - | Via OpenAI-compatible provider (RigAdapter) |
|
||||
| Tinfoil | ❌ | ✅ | - | Private inference provider (IronClaw-only) |
|
||||
| OpenAI-compatible | ❌ | ✅ | - | Generic OpenAI-compatible endpoint (RigAdapter) |
|
||||
|
||||
@@ -11,6 +11,12 @@ configurations.
|
||||
| NEAR AI | `nearai` | OAuth (browser) | Default; multi-model |
|
||||
| Anthropic | `anthropic` | `ANTHROPIC_API_KEY` | Claude models |
|
||||
| OpenAI | `openai` | `OPENAI_API_KEY` | GPT models |
|
||||
| Google Gemini | `gemini` | `GEMINI_API_KEY` | Gemini models |
|
||||
| AWS Bedrock | `bedrock` | `BEDROCK_ACCESS_KEY` | Requires OpenAI proxy (e.g. LiteLLM) |
|
||||
| io.net | `ionet` | `IONET_API_KEY` | Intelligence API |
|
||||
| Mistral | `mistral` | `MISTRAL_API_KEY` | Mistral models |
|
||||
| Yandex AI Studio | `yandex` | `YANDEX_API_KEY` | YandexGPT models |
|
||||
| Cloudflare Workers AI | `cloudflare` | `CLOUDFLARE_API_KEY` | Access to Workers AI |
|
||||
| Ollama | `ollama` | No | Local inference |
|
||||
| OpenRouter | `openai_compatible` | `LLM_API_KEY` | 300+ models |
|
||||
| Together AI | `openai_compatible` | `LLM_API_KEY` | Fast inference |
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
-- Partial unique indexes to prevent duplicate singleton conversations.
|
||||
-- These guard against TOCTOU races in get_or_create_routine_conversation
|
||||
-- and get_or_create_heartbeat_conversation.
|
||||
|
||||
-- One routine conversation per user per routine_id.
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS uq_conv_routine
|
||||
ON conversations (user_id, (metadata->>'routine_id'))
|
||||
WHERE metadata->>'routine_id' IS NOT NULL;
|
||||
|
||||
-- One heartbeat conversation per user.
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS uq_conv_heartbeat
|
||||
ON conversations (user_id)
|
||||
WHERE metadata->>'thread_type' = 'heartbeat';
|
||||
+161
-11
@@ -1,7 +1,9 @@
|
||||
[
|
||||
{
|
||||
"id": "openai",
|
||||
"aliases": ["open_ai"],
|
||||
"aliases": [
|
||||
"open_ai"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"api_key_env": "OPENAI_API_KEY",
|
||||
"api_key_required": true,
|
||||
@@ -19,7 +21,9 @@
|
||||
},
|
||||
{
|
||||
"id": "anthropic",
|
||||
"aliases": ["claude"],
|
||||
"aliases": [
|
||||
"claude"
|
||||
],
|
||||
"protocol": "anthropic",
|
||||
"api_key_env": "ANTHROPIC_API_KEY",
|
||||
"api_key_required": true,
|
||||
@@ -52,7 +56,10 @@
|
||||
},
|
||||
{
|
||||
"id": "openai_compatible",
|
||||
"aliases": ["openai-compatible", "compatible"],
|
||||
"aliases": [
|
||||
"openai-compatible",
|
||||
"compatible"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"base_url_env": "LLM_BASE_URL",
|
||||
"base_url_required": true,
|
||||
@@ -89,7 +96,9 @@
|
||||
},
|
||||
{
|
||||
"id": "openrouter",
|
||||
"aliases": ["open_router"],
|
||||
"aliases": [
|
||||
"open_router"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://openrouter.ai/api/v1",
|
||||
"api_key_env": "OPENROUTER_API_KEY",
|
||||
@@ -126,7 +135,10 @@
|
||||
},
|
||||
{
|
||||
"id": "nvidia",
|
||||
"aliases": ["nvidia_nim", "nim"],
|
||||
"aliases": [
|
||||
"nvidia_nim",
|
||||
"nim"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://integrate.api.nvidia.com/v1",
|
||||
"api_key_env": "NVIDIA_API_KEY",
|
||||
@@ -144,7 +156,10 @@
|
||||
},
|
||||
{
|
||||
"id": "venice",
|
||||
"aliases": ["venice_ai", "veniceai"],
|
||||
"aliases": [
|
||||
"venice_ai",
|
||||
"veniceai"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.venice.ai/api/v1",
|
||||
"api_key_env": "VENICE_API_KEY",
|
||||
@@ -162,7 +177,10 @@
|
||||
},
|
||||
{
|
||||
"id": "together",
|
||||
"aliases": ["together_ai", "togetherai"],
|
||||
"aliases": [
|
||||
"together_ai",
|
||||
"togetherai"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.together.xyz/v1",
|
||||
"api_key_env": "TOGETHER_API_KEY",
|
||||
@@ -180,7 +198,9 @@
|
||||
},
|
||||
{
|
||||
"id": "fireworks",
|
||||
"aliases": ["fireworks_ai"],
|
||||
"aliases": [
|
||||
"fireworks_ai"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.fireworks.ai/inference/v1",
|
||||
"api_key_env": "FIREWORKS_API_KEY",
|
||||
@@ -198,7 +218,9 @@
|
||||
},
|
||||
{
|
||||
"id": "deepseek",
|
||||
"aliases": ["deep_seek"],
|
||||
"aliases": [
|
||||
"deep_seek"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.deepseek.com/v1",
|
||||
"api_key_env": "DEEPSEEK_API_KEY",
|
||||
@@ -234,7 +256,9 @@
|
||||
},
|
||||
{
|
||||
"id": "sambanova",
|
||||
"aliases": ["samba_nova"],
|
||||
"aliases": [
|
||||
"samba_nova"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.sambanova.ai/v1",
|
||||
"api_key_env": "SAMBANOVA_API_KEY",
|
||||
@@ -249,5 +273,131 @@
|
||||
"display_name": "SambaNova",
|
||||
"can_list_models": false
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "gemini",
|
||||
"aliases": [
|
||||
"google_gemini",
|
||||
"google"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://generativelanguage.googleapis.com/v1beta/openai",
|
||||
"api_key_env": "GEMINI_API_KEY",
|
||||
"api_key_required": true,
|
||||
"model_env": "GEMINI_MODEL",
|
||||
"default_model": "gemini-2.5-flash",
|
||||
"description": "Google Gemini (via OpenAI-compatible endpoint)",
|
||||
"setup": {
|
||||
"kind": "api_key",
|
||||
"secret_name": "llm_gemini_api_key",
|
||||
"key_url": "https://aistudio.google.com/app/apikey",
|
||||
"display_name": "Google Gemini",
|
||||
"can_list_models": true
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "bedrock",
|
||||
"aliases": [
|
||||
"aws_bedrock",
|
||||
"aws"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"api_key_env": "BEDROCK_ACCESS_KEY",
|
||||
"api_key_required": false,
|
||||
"base_url_env": "BEDROCK_BASE_URL",
|
||||
"model_env": "BEDROCK_MODEL",
|
||||
"default_model": "anthropic.claude-3-5-sonnet-20241022-v2:0",
|
||||
"description": "AWS Bedrock (requires LiteLLM or OpenAI-compatible proxy)",
|
||||
"setup": {
|
||||
"kind": "open_ai_compatible",
|
||||
"secret_name": "llm_bedrock_api_key",
|
||||
"display_name": "AWS Bedrock",
|
||||
"can_list_models": false
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "ionet",
|
||||
"aliases": [
|
||||
"io_net",
|
||||
"io.net"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.intelligence.io.solutions/api/v1",
|
||||
"api_key_env": "IONET_API_KEY",
|
||||
"api_key_required": true,
|
||||
"model_env": "IONET_MODEL",
|
||||
"default_model": "deepseek-coder-v2-instruct",
|
||||
"description": "io.net Intelligence API",
|
||||
"setup": {
|
||||
"kind": "api_key",
|
||||
"secret_name": "llm_ionet_api_key",
|
||||
"key_url": "https://cloud.io.net/intelligence",
|
||||
"display_name": "io.net",
|
||||
"can_list_models": true
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "mistral",
|
||||
"aliases": [
|
||||
"mistral_ai",
|
||||
"mistralai"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://api.mistral.ai/v1",
|
||||
"api_key_env": "MISTRAL_API_KEY",
|
||||
"api_key_required": true,
|
||||
"model_env": "MISTRAL_MODEL",
|
||||
"default_model": "mistral-large-latest",
|
||||
"description": "Mistral AI API",
|
||||
"setup": {
|
||||
"kind": "api_key",
|
||||
"secret_name": "llm_mistral_api_key",
|
||||
"key_url": "https://console.mistral.ai/api-keys",
|
||||
"display_name": "Mistral",
|
||||
"can_list_models": true
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "yandex",
|
||||
"aliases": [
|
||||
"yandex_ai_studio",
|
||||
"yandexgpt",
|
||||
"yandex_gpt"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"default_base_url": "https://ai.api.cloud.yandex.net/v1",
|
||||
"api_key_env": "YANDEX_API_KEY",
|
||||
"api_key_required": true,
|
||||
"model_env": "YANDEX_MODEL",
|
||||
"extra_headers_env": "YANDEX_EXTRA_HEADERS",
|
||||
"default_model": "yandexgpt-lite",
|
||||
"description": "Yandex AI Studio (YandexGPT)",
|
||||
"setup": {
|
||||
"kind": "api_key",
|
||||
"secret_name": "llm_yandex_api_key",
|
||||
"key_url": "https://aistudio.yandex.ru/platform/folders/",
|
||||
"display_name": "Yandex AI Studio",
|
||||
"can_list_models": true
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": "cloudflare",
|
||||
"aliases": [
|
||||
"cloudflare_ai",
|
||||
"cf_ai"
|
||||
],
|
||||
"protocol": "open_ai_completions",
|
||||
"api_key_env": "CLOUDFLARE_API_KEY",
|
||||
"api_key_required": true,
|
||||
"base_url_env": "CLOUDFLARE_BASE_URL",
|
||||
"model_env": "CLOUDFLARE_MODEL",
|
||||
"default_model": "@cf/meta/llama-3.3-70b-instruct-fp8-fast",
|
||||
"description": "Cloudflare Workers AI",
|
||||
"setup": {
|
||||
"kind": "open_ai_compatible",
|
||||
"secret_name": "llm_cloudflare_api_key",
|
||||
"display_name": "Cloudflare Workers AI",
|
||||
"can_list_models": false
|
||||
}
|
||||
}
|
||||
]
|
||||
]
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/discord-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "030707431717bca3411a48f311c6ab5f92a45c747de26cafe4f6e3e23a8b3b2d"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "6ed36077b67ac70a041f06f760f93ba79b33269885413c3c3f2c8c87ee60807e"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "98c86895a9c4b0a1e19fe8a47f1ccbfe7e972e112b05e584bc897130dc32283a"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/whatsapp-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "bd35cad18d87292ea8d2f52db9b514ed9f814a414de910f59073d475c26c4c14"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -20,7 +20,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/github-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "6fcd32719a4ff15641a4b50fff8984686550f0c491dce60518f4126857d0c544"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/gmail-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "023da7000b17568bf0e64b2e5013c8a042b2f323c85f1632339231c73d500e39"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-calendar-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "fc42277b65881d6e9bcc5403dc54c7f5b3ddeaaaf04617fce2c5da05d76325f0"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-docs-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "385c04abd1e6b8011ccc330e1f4bd7ce58577e488959b51594aa04eb26cbe7cc"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-drive-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "1b107d575a5d52cc8c76d9a681802190f4373fb485f7f54f445533f097fa37c0"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-sheets-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "c4f6b1e8c5126ac2c8a4b98e4283a3afa32223d2488fc3c3a609758c0c9beb90"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -18,7 +18,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/google-slides-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "7110b8565340c888e51f99e9c013bf4de8f8a7f7b33bace00eb8fc47831ff20b"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -18,7 +18,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/slack-tool-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "6ed36077b67ac70a041f06f760f93ba79b33269885413c3c3f2c8c87ee60807e"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/telegram-mtproto-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "98c86895a9c4b0a1e19fe8a47f1ccbfe7e972e112b05e584bc897130dc32283a"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
"artifacts": {
|
||||
"wasm32-wasip2": {
|
||||
"url": "https://github.com/nearai/ironclaw/releases/latest/download/web-search-wasm32-wasip2.tar.gz",
|
||||
"sha256": null
|
||||
"sha256": "66cb2b9b00652385e9f30f17c74902b9222c17c53e9d3bd1ef42f5cab705bcf6"
|
||||
}
|
||||
},
|
||||
"auth_summary": {
|
||||
|
||||
@@ -51,9 +51,11 @@ echo "[6/6] Installing git hooks..."
|
||||
HOOKS_DIR=$(git rev-parse --git-path hooks 2>/dev/null) || true
|
||||
if [ -n "$HOOKS_DIR" ]; then
|
||||
mkdir -p "$HOOKS_DIR"
|
||||
SCRIPT_ABS="$(cd "$(dirname "$0")" && pwd)/commit-msg-regression.sh"
|
||||
ln -sf "$SCRIPT_ABS" "$HOOKS_DIR/commit-msg"
|
||||
SCRIPTS_ABS="$(cd "$(dirname "$0")" && pwd)"
|
||||
ln -sf "$SCRIPTS_ABS/commit-msg-regression.sh" "$HOOKS_DIR/commit-msg"
|
||||
echo " commit-msg hook installed (regression test enforcement)"
|
||||
ln -sf "$SCRIPTS_ABS/pre-commit-safety.sh" "$HOOKS_DIR/pre-commit"
|
||||
echo " pre-commit hook installed (UTF-8, case-sensitivity, /tmp, redaction checks)"
|
||||
else
|
||||
echo " Skipped: not a git repository"
|
||||
fi
|
||||
|
||||
Executable
+136
@@ -0,0 +1,136 @@
|
||||
#!/usr/bin/env bash
|
||||
# Pre-commit safety checks for common issues caught by AI code reviewers.
|
||||
#
|
||||
# Can be run standalone: bash scripts/pre-commit-safety.sh
|
||||
# Or installed as a git pre-commit hook via dev-setup.sh.
|
||||
#
|
||||
# Checks staged .rs files for:
|
||||
# 1. Unsafe UTF-8 byte slicing (panics on multi-byte chars)
|
||||
# 2. Case-sensitive file extension comparisons
|
||||
# 3. Hardcoded /tmp paths in tests (flaky in parallel runs)
|
||||
# 4. Tool parameters logged without redaction (secret leaks)
|
||||
# 5. Multi-step DB operations without transaction wrapping
|
||||
#
|
||||
# Suppress individual lines with an inline "// safety: <reason>" comment.
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
# Determine a suitable base ref for standalone diffs.
|
||||
resolve_base_ref() {
|
||||
local candidates=(
|
||||
"@{upstream}"
|
||||
"origin/HEAD"
|
||||
"origin/main"
|
||||
"origin/master"
|
||||
"main"
|
||||
"master"
|
||||
)
|
||||
|
||||
for ref in "${candidates[@]}"; do
|
||||
if git rev-parse --verify --quiet "$ref" >/dev/null 2>&1; then
|
||||
echo "$ref"
|
||||
return 0
|
||||
fi
|
||||
done
|
||||
|
||||
echo "pre-commit-safety: could not determine a base Git ref for diff (tried: ${candidates[*]})." >&2
|
||||
echo "pre-commit-safety: ensure your repository has an upstream or a local main/master branch." >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
# Support both pre-commit hook (staged files) and standalone (all changed vs base)
|
||||
if git diff --cached --quiet 2>/dev/null; then
|
||||
# No staged changes -- compare working tree against a resolved base ref
|
||||
BASE_REF="$(resolve_base_ref)"
|
||||
DIFF_OUTPUT=$(git diff "$BASE_REF" -- '*.rs' 2>/dev/null || true)
|
||||
else
|
||||
DIFF_OUTPUT=$(git diff --cached -U0 -- '*.rs' 2>/dev/null || true)
|
||||
fi
|
||||
|
||||
# Early exit if there are no relevant .rs changes
|
||||
if [ -z "$DIFF_OUTPUT" ]; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
WARNINGS=0
|
||||
|
||||
warn() {
|
||||
if [ "$WARNINGS" -eq 0 ]; then
|
||||
echo ""
|
||||
echo "=== Pre-commit Safety Checks ==="
|
||||
echo ""
|
||||
fi
|
||||
WARNINGS=$((WARNINGS + 1))
|
||||
echo " [$1] $2"
|
||||
}
|
||||
|
||||
# 1. Unsafe UTF-8 byte slicing: &s[..N] or &s[..some_var] on strings
|
||||
# Safe patterns: is_char_boundary, char_indices, // safety:
|
||||
if echo "$DIFF_OUTPUT" | grep -nE '^\+' | grep -E '\[\.\..*\]' | grep -vE 'is_char_boundary|char_indices|// safety:|as_bytes|Vec<|&\[u8\]|\[u8\]|bytes\(\)|&bytes' | head -3 | grep -q .; then
|
||||
warn "UTF8" "Possible unsafe byte-index string slicing. Use is_char_boundary() or char_indices()."
|
||||
echo "$DIFF_OUTPUT" | grep -nE '^\+' | grep -E '\[\.\..*\]' | grep -vE 'is_char_boundary|char_indices|// safety:|as_bytes|Vec<|&\[u8\]|\[u8\]|bytes\(\)|&bytes' | head -3 | sed 's/^/ /'
|
||||
fi
|
||||
|
||||
# 2. Case-sensitive file extension checks
|
||||
# Match: .ends_with(".png") without prior to_lowercase
|
||||
if echo "$DIFF_OUTPUT" | grep -nE '^\+.*ends_with\("\.([pP][nN][gG]|[jJ][pP][eE]?[gG]|[gG][iI][fF]|[wW][eE][bB][pP]|[mM][dD])"\)' | grep -vE 'to_lowercase|to_ascii_lowercase|// safety:' | head -3 | grep -q .; then
|
||||
warn "CASE" "Case-sensitive file extension comparison. Normalize to lowercase first."
|
||||
echo "$DIFF_OUTPUT" | grep -nE '^\+.*ends_with\("\.([pP][nN][gG]|[jJ][pP][eE]?[gG]|[gG][iI][fF]|[wW][eE][bB][pP]|[mM][dD])"\)' | grep -vE 'to_lowercase|to_ascii_lowercase|// safety:' | head -3 | sed 's/^/ /'
|
||||
fi
|
||||
|
||||
# 3. Hardcoded /tmp paths in test files
|
||||
if echo "$DIFF_OUTPUT" | grep -nE '^\+.*"/tmp/' | grep -vE 'tempfile|tempdir|// safety:' | head -3 | grep -q .; then
|
||||
warn "TMPDIR" "Hardcoded /tmp path. Use tempfile::tempdir() for parallel-safe tests."
|
||||
echo "$DIFF_OUTPUT" | grep -nE '^\+.*"/tmp/' | grep -vE 'tempfile|tempdir|// safety:' | head -3 | sed 's/^/ /'
|
||||
fi
|
||||
|
||||
# 4. Logging tool parameters without redaction
|
||||
if echo "$DIFF_OUTPUT" | grep -nE '^\+.*tracing::(info|debug|warn|error).*param' | grep -vE 'redact|// safety:' | head -3 | grep -q .; then
|
||||
warn "REDACT" "Logging tool parameters without redaction. Use redact_params() first."
|
||||
echo "$DIFF_OUTPUT" | grep -nE '^\+.*tracing::(info|debug|warn|error).*param' | grep -vE 'redact|// safety:' | head -3 | sed 's/^/ /'
|
||||
fi
|
||||
|
||||
# 5. Multi-step DB operations without transaction
|
||||
# Uses -W (function context) to reduce false positives from existing transactions.
|
||||
# Suppressible with "// safety:" in the hunk.
|
||||
DIFF_W_OUTPUT=$(git diff --cached -W -- '*.rs' 2>/dev/null || git diff "$(resolve_base_ref)" -W -- '*.rs' 2>/dev/null || true)
|
||||
if [ -n "$DIFF_W_OUTPUT" ]; then
|
||||
HUNK_COUNT=$(echo "$DIFF_W_OUTPUT" | awk '
|
||||
/^@@/ {
|
||||
if (count >= 2 && !has_tx && !has_safety) found++
|
||||
count=0; has_tx=0; has_safety=0
|
||||
}
|
||||
/^\+.*\.(execute|query)\(/ { count++ }
|
||||
/^\+.*(transaction|\.tx\.|\.begin\()/ { has_tx=1 }
|
||||
/ .*(transaction|\.tx\.|\.begin\()/ { has_tx=1 }
|
||||
/\/\/ safety:/ { has_safety=1 }
|
||||
END {
|
||||
if (count >= 2 && !has_tx && !has_safety) found++
|
||||
print found+0
|
||||
}
|
||||
')
|
||||
if [ "$HUNK_COUNT" -gt 0 ]; then
|
||||
warn "TX" "Multiple DB operations in same function without transaction. Wrap in a transaction for atomicity."
|
||||
echo "$DIFF_W_OUTPUT" | awk '
|
||||
/^@@/ {
|
||||
if (count >= 2 && !has_tx && !has_safety) { print buf }
|
||||
buf=""; count=0; has_tx=0; has_safety=0
|
||||
}
|
||||
/^\+.*\.(execute|query)\(/ { count++ }
|
||||
/^\+.*(transaction|\.tx\.|\.begin\()/ { has_tx=1 }
|
||||
/ .*(transaction|\.tx\.|\.begin\()/ { has_tx=1 }
|
||||
/\/\/ safety:/ { has_safety=1 }
|
||||
{ buf = buf "\n" $0 }
|
||||
END {
|
||||
if (count >= 2 && !has_tx && !has_safety) { print buf }
|
||||
}
|
||||
' | grep -E '^\+.*\.(execute|query)\(' | head -4 | sed 's/^/ /'
|
||||
fi
|
||||
fi
|
||||
|
||||
if [ "$WARNINGS" -gt 0 ]; then
|
||||
echo ""
|
||||
echo "Found $WARNINGS potential issue(s). Fix them or add '// safety: <reason>' to suppress."
|
||||
echo ""
|
||||
exit 1
|
||||
fi
|
||||
@@ -0,0 +1,54 @@
|
||||
---
|
||||
name: review-checklist
|
||||
version: 0.1.0
|
||||
description: Pre-merge review checklist based on recurring AI reviewer feedback patterns
|
||||
activation:
|
||||
patterns:
|
||||
- "review.*checklist"
|
||||
- "ready to merge"
|
||||
- "pre-merge check"
|
||||
- "check.*before.*merge"
|
||||
keywords:
|
||||
- review
|
||||
- checklist
|
||||
- merge
|
||||
- pre-merge
|
||||
max_context_tokens: 1500
|
||||
---
|
||||
|
||||
# Pre-Merge Review Checklist
|
||||
|
||||
Before merging, verify these items. They represent the most common issues caught by automated code reviewers (Copilot, Gemini) on IronClaw PRs.
|
||||
|
||||
## Database Operations
|
||||
- [ ] Multi-step DB operations are wrapped in transactions (INSERT+INSERT, UPDATE+DELETE, read-modify-write)
|
||||
- [ ] Both postgres AND libsql backends updated for any new Database trait methods
|
||||
- [ ] Migrations are atomic (SQL execution + version recording in same transaction)
|
||||
|
||||
## Security & Data Safety
|
||||
- [ ] Tool parameters are redacted via `redact_params()` before logging or SSE/WebSocket broadcast
|
||||
- [ ] URL validation resolves DNS before checking for private/loopback IPs (anti-SSRF via DNS rebinding)
|
||||
- [ ] Destructive tools have `requires_approval()` returning `Always` or `UnlessAutoApproved`
|
||||
- [ ] Data from worker containers is treated as untrusted (tool domain checks, server-side nesting depth)
|
||||
- [ ] No secrets or credentials in error messages, logs, or SSE events
|
||||
|
||||
## String Safety
|
||||
- [ ] No byte-index slicing (`&s[..n]`) on external/user strings -- use `is_char_boundary()` or `char_indices()`
|
||||
- [ ] File extension and media type comparisons are case-insensitive (`.to_ascii_lowercase()` before matching)
|
||||
- [ ] Path comparisons are case-insensitive where needed (macOS/Windows filesystems)
|
||||
|
||||
## Trait Wrappers & Decorator Chain
|
||||
- [ ] New `LlmProvider` trait methods are delegated in ALL wrapper types (grep `impl LlmProvider for`)
|
||||
- [ ] New trait methods are tested through the full decorator/provider chain, not just the base impl
|
||||
- [ ] Default trait method implementations are intentional -- wrappers that silently return defaults are bugs
|
||||
|
||||
## Tests
|
||||
- [ ] Temporary files/dirs use `tempfile` crate, no hardcoded `/tmp/` paths
|
||||
- [ ] Tests don't mutate global statics without synchronization (use per-test state or `serial_test`)
|
||||
- [ ] Tests don't make real network requests (use mocks, stubs, or RFC 5737 TEST-NET IPs like 192.0.2.1)
|
||||
- [ ] Test names and comments match actual test behavior and assertions
|
||||
|
||||
## Comments & Documentation
|
||||
- [ ] Code comments match actual behavior (especially route paths, tool names, function semantics)
|
||||
- [ ] Spec/README files updated if module behavior changed
|
||||
- [ ] Error messages are clear and non-redundant (don't nest tool name inside tool error that already contains it)
|
||||
+24
-1
@@ -96,6 +96,9 @@ pub struct Agent {
|
||||
pub(super) heartbeat_config: Option<HeartbeatConfig>,
|
||||
pub(super) hygiene_config: Option<crate::config::HygieneConfig>,
|
||||
pub(super) routine_config: Option<RoutineConfig>,
|
||||
/// Optional slot to expose the routine engine to the gateway for manual triggering.
|
||||
pub(super) routine_engine_slot:
|
||||
Option<Arc<tokio::sync::RwLock<Option<Arc<crate::agent::routine_engine::RoutineEngine>>>>>,
|
||||
}
|
||||
|
||||
impl Agent {
|
||||
@@ -148,9 +151,18 @@ impl Agent {
|
||||
heartbeat_config,
|
||||
hygiene_config,
|
||||
routine_config,
|
||||
routine_engine_slot: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Set the routine engine slot for exposing the engine to the gateway.
|
||||
pub fn set_routine_engine_slot(
|
||||
&mut self,
|
||||
slot: Arc<tokio::sync::RwLock<Option<Arc<crate::agent::routine_engine::RoutineEngine>>>>,
|
||||
) {
|
||||
self.routine_engine_slot = Some(slot);
|
||||
}
|
||||
|
||||
// Convenience accessors
|
||||
|
||||
/// Get the scheduler (for external wiring, e.g. CreateJobTool).
|
||||
@@ -342,8 +354,13 @@ impl Agent {
|
||||
let heartbeat_handle = if let Some(ref hb_config) = self.heartbeat_config {
|
||||
if hb_config.enabled {
|
||||
if let Some(workspace) = self.workspace() {
|
||||
let config = AgentHeartbeatConfig::default()
|
||||
let mut config = AgentHeartbeatConfig::default()
|
||||
.with_interval(std::time::Duration::from_secs(hb_config.interval_secs));
|
||||
if let (Some(user), Some(channel)) =
|
||||
(&hb_config.notify_user, &hb_config.notify_channel)
|
||||
{
|
||||
config = config.with_notify(user, channel);
|
||||
}
|
||||
|
||||
// Set up notification channel
|
||||
let (notify_tx, mut notify_rx) =
|
||||
@@ -396,6 +413,7 @@ impl Agent {
|
||||
self.cheap_llm().clone(),
|
||||
self.safety().clone(),
|
||||
Some(notify_tx),
|
||||
self.store().map(Arc::clone),
|
||||
))
|
||||
} else {
|
||||
tracing::warn!("Heartbeat enabled but no workspace available");
|
||||
@@ -486,6 +504,11 @@ impl Agent {
|
||||
// SAFETY: self is consumed by run(), we can smuggle the engine in
|
||||
// via a local to use in the message loop below.
|
||||
|
||||
// Expose engine to gateway for manual triggering
|
||||
if let Some(ref slot) = self.routine_engine_slot {
|
||||
*slot.write().await = Some(Arc::clone(&engine));
|
||||
}
|
||||
|
||||
tracing::info!(
|
||||
"Routines enabled: cron ticker every {}s, max {} concurrent",
|
||||
rt_config.cron_check_interval_secs,
|
||||
|
||||
+42
-7
@@ -131,10 +131,12 @@ impl CostGuard {
|
||||
// Check hourly rate
|
||||
if let Some(limit) = self.config.max_actions_per_hour {
|
||||
let mut window = self.action_window.lock().await;
|
||||
let cutoff = Instant::now() - std::time::Duration::from_secs(3600);
|
||||
// Drain expired entries
|
||||
while window.front().is_some_and(|t| *t < cutoff) {
|
||||
window.pop_front();
|
||||
// checked_sub avoids panic when system uptime < 1 hour (Windows)
|
||||
if let Some(cutoff) = Instant::now().checked_sub(std::time::Duration::from_secs(3600)) {
|
||||
// Drain expired entries
|
||||
while window.front().is_some_and(|t| *t < cutoff) {
|
||||
window.pop_front();
|
||||
}
|
||||
}
|
||||
let count = window.len() as u64;
|
||||
if count >= limit {
|
||||
@@ -260,9 +262,11 @@ impl CostGuard {
|
||||
/// Number of actions in the current hourly window.
|
||||
pub async fn actions_this_hour(&self) -> u64 {
|
||||
let mut window = self.action_window.lock().await;
|
||||
let cutoff = Instant::now() - std::time::Duration::from_secs(3600);
|
||||
while window.front().is_some_and(|t| *t < cutoff) {
|
||||
window.pop_front();
|
||||
// checked_sub avoids panic when system uptime < 1 hour (Windows)
|
||||
if let Some(cutoff) = Instant::now().checked_sub(std::time::Duration::from_secs(3600)) {
|
||||
while window.front().is_some_and(|t| *t < cutoff) {
|
||||
window.pop_front();
|
||||
}
|
||||
}
|
||||
window.len() as u64
|
||||
}
|
||||
@@ -621,4 +625,35 @@ mod tests {
|
||||
"surcharge should be 100% of input cost for 1h cache writes"
|
||||
);
|
||||
}
|
||||
|
||||
/// Regression test for #657: Instant::now() - Duration panics on Windows
|
||||
/// when system uptime is less than the subtracted duration.
|
||||
#[tokio::test]
|
||||
async fn test_checked_sub_no_panic_on_fresh_guard() {
|
||||
// A fresh CostGuard with rate limits should not panic even if
|
||||
// checked_sub returns None (simulating short uptime).
|
||||
let guard = CostGuard::new(CostGuardConfig {
|
||||
max_cost_per_day_cents: None,
|
||||
max_actions_per_hour: Some(100),
|
||||
});
|
||||
|
||||
// These must not panic regardless of system uptime
|
||||
assert!(guard.check_allowed().await.is_ok());
|
||||
assert_eq!(guard.actions_this_hour().await, 0);
|
||||
|
||||
// Record some actions and verify again
|
||||
guard
|
||||
.record_llm_call("gpt-4o", 10, 10, 0, 0, Decimal::ONE, Decimal::ONE, None)
|
||||
.await;
|
||||
assert!(guard.check_allowed().await.is_ok());
|
||||
assert_eq!(guard.actions_this_hour().await, 1);
|
||||
}
|
||||
|
||||
/// Verify that checked_sub itself behaves as expected for the pattern we use.
|
||||
#[test]
|
||||
fn test_instant_checked_sub_returns_none_for_overflow() {
|
||||
// Duration::MAX will always exceed uptime, so checked_sub must return None
|
||||
let result = Instant::now().checked_sub(std::time::Duration::MAX);
|
||||
assert!(result.is_none());
|
||||
}
|
||||
}
|
||||
|
||||
+55
-1
@@ -29,6 +29,7 @@ use std::time::Duration;
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use crate::channels::OutgoingResponse;
|
||||
use crate::db::Database;
|
||||
use crate::llm::{ChatMessage, CompletionRequest, LlmProvider, Reasoning};
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::workspace::Workspace;
|
||||
@@ -103,6 +104,7 @@ pub struct HeartbeatRunner {
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
safety: Arc<SafetyLayer>,
|
||||
response_tx: Option<mpsc::Sender<OutgoingResponse>>,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
consecutive_failures: u32,
|
||||
}
|
||||
|
||||
@@ -122,6 +124,7 @@ impl HeartbeatRunner {
|
||||
llm,
|
||||
safety,
|
||||
response_tx: None,
|
||||
store: None,
|
||||
consecutive_failures: 0,
|
||||
}
|
||||
}
|
||||
@@ -132,6 +135,12 @@ impl HeartbeatRunner {
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the database store for persistent heartbeat conversations.
|
||||
pub fn with_store(mut self, store: Arc<dyn Database>) -> Self {
|
||||
self.store = Some(store);
|
||||
self
|
||||
}
|
||||
|
||||
/// Run the heartbeat loop.
|
||||
///
|
||||
/// This runs forever, checking periodically based on the configured interval.
|
||||
@@ -292,9 +301,32 @@ impl HeartbeatRunner {
|
||||
return;
|
||||
};
|
||||
|
||||
let user_id = self.config.notify_user_id.as_deref().unwrap_or("default");
|
||||
|
||||
// Persist to heartbeat conversation and get thread_id
|
||||
let thread_id = if let Some(ref store) = self.store {
|
||||
match store.get_or_create_heartbeat_conversation(user_id).await {
|
||||
Ok(conv_id) => {
|
||||
if let Err(e) = store
|
||||
.add_conversation_message(conv_id, "assistant", message)
|
||||
.await
|
||||
{
|
||||
tracing::error!("Failed to persist heartbeat message: {}", e);
|
||||
}
|
||||
Some(conv_id.to_string())
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!("Failed to get heartbeat conversation: {}", e);
|
||||
None
|
||||
}
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let response = OutgoingResponse {
|
||||
content: format!("🔔 *Heartbeat Alert*\n\n{}", message),
|
||||
thread_id: None,
|
||||
thread_id,
|
||||
attachments: Vec::new(),
|
||||
metadata: serde_json::json!({
|
||||
"source": "heartbeat",
|
||||
@@ -356,11 +388,15 @@ pub fn spawn_heartbeat(
|
||||
llm: Arc<dyn LlmProvider>,
|
||||
safety: Arc<SafetyLayer>,
|
||||
response_tx: Option<mpsc::Sender<OutgoingResponse>>,
|
||||
store: Option<Arc<dyn Database>>,
|
||||
) -> tokio::task::JoinHandle<()> {
|
||||
let mut runner = HeartbeatRunner::new(config, hygiene_config, workspace, llm, safety);
|
||||
if let Some(tx) = response_tx {
|
||||
runner = runner.with_response_channel(tx);
|
||||
}
|
||||
if let Some(s) = store {
|
||||
runner = runner.with_store(s);
|
||||
}
|
||||
|
||||
tokio::spawn(async move {
|
||||
runner.run().await;
|
||||
@@ -495,4 +531,22 @@ mod tests {
|
||||
let content = "<!-- comment -->\nActual task here";
|
||||
assert!(!is_effectively_empty(content));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_spawn_heartbeat_accepts_store_param() {
|
||||
// Regression: spawn_heartbeat must accept an optional Database store
|
||||
// for persisting heartbeat notifications to a dedicated conversation.
|
||||
// Compile-time check: the 7th parameter is `Option<Arc<dyn Database>>`.
|
||||
#[allow(clippy::type_complexity)]
|
||||
let _fn_ptr: fn(
|
||||
HeartbeatConfig,
|
||||
HygieneConfig,
|
||||
Arc<crate::workspace::Workspace>,
|
||||
Arc<dyn crate::llm::LlmProvider>,
|
||||
Arc<crate::safety::SafetyLayer>,
|
||||
Option<tokio::sync::mpsc::Sender<crate::channels::OutgoingResponse>>,
|
||||
Option<Arc<dyn crate::db::Database>>,
|
||||
) -> tokio::task::JoinHandle<()> = spawn_heartbeat;
|
||||
let _ = _fn_ptr;
|
||||
}
|
||||
}
|
||||
|
||||
+257
-4
@@ -25,6 +25,7 @@ use crate::agent::routine::{
|
||||
};
|
||||
use crate::channels::{IncomingMessage, OutgoingResponse};
|
||||
use crate::config::RoutineConfig;
|
||||
use crate::context::JobState;
|
||||
use crate::db::Database;
|
||||
use crate::error::RoutineError;
|
||||
use crate::llm::{ChatMessage, CompletionRequest, FinishReason, LlmProvider};
|
||||
@@ -180,11 +181,139 @@ impl RoutineEngine {
|
||||
}
|
||||
}
|
||||
|
||||
/// Sync dispatched routine runs with their linked background job status.
|
||||
///
|
||||
/// Full-job routines are fire-and-forget: the routine run is created with
|
||||
/// `Running` status when the job is dispatched, but the run record is never
|
||||
/// updated when the background job completes or fails. This method checks
|
||||
/// all `Running` routine runs that have a linked job, queries the job's
|
||||
/// current state, and updates the routine run accordingly. It also sends
|
||||
/// failure/success notifications that would otherwise be lost.
|
||||
pub async fn sync_dispatched_runs(&self) {
|
||||
let runs = match self.store.list_dispatched_routine_runs().await {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
tracing::debug!("Failed to list dispatched routine runs: {}", e);
|
||||
return;
|
||||
}
|
||||
};
|
||||
|
||||
for run in runs {
|
||||
let Some(job_id) = run.job_id else {
|
||||
continue;
|
||||
};
|
||||
|
||||
// Check the linked job's current state
|
||||
let job = match self.store.get_job(job_id).await {
|
||||
Ok(Some(j)) => j,
|
||||
Ok(None) => {
|
||||
// Job was deleted — mark the routine run as failed
|
||||
tracing::warn!(
|
||||
run_id = %run.id,
|
||||
job_id = %job_id,
|
||||
"Linked job not found, marking routine run as failed"
|
||||
);
|
||||
self.complete_dispatched_run(
|
||||
&run,
|
||||
RunStatus::Failed,
|
||||
"Linked job not found (may have been deleted)",
|
||||
)
|
||||
.await;
|
||||
continue;
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::debug!(
|
||||
run_id = %run.id,
|
||||
job_id = %job_id,
|
||||
"Failed to query linked job: {}", e
|
||||
);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
// Extract the reason from the most recent state transition
|
||||
let last_reason = job.transitions.last().and_then(|t| t.reason.clone());
|
||||
|
||||
// Map job state to routine run status
|
||||
let (new_status, summary) = match job.state {
|
||||
JobState::Completed | JobState::Submitted | JobState::Accepted => {
|
||||
let summary =
|
||||
last_reason.unwrap_or_else(|| "Job completed successfully".to_string());
|
||||
(RunStatus::Ok, summary)
|
||||
}
|
||||
JobState::Failed => {
|
||||
let summary = last_reason
|
||||
.unwrap_or_else(|| "Job failed (no error message recorded)".to_string());
|
||||
(RunStatus::Failed, summary)
|
||||
}
|
||||
JobState::Cancelled => (RunStatus::Failed, "Job was cancelled".to_string()),
|
||||
// Still in progress — skip
|
||||
JobState::Pending | JobState::InProgress | JobState::Stuck => continue,
|
||||
};
|
||||
|
||||
tracing::info!(
|
||||
run_id = %run.id,
|
||||
job_id = %job_id,
|
||||
status = %new_status,
|
||||
"Syncing dispatched routine run with completed job"
|
||||
);
|
||||
|
||||
self.complete_dispatched_run(&run, new_status, &summary)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
|
||||
/// Complete a dispatched routine run and send the appropriate notification.
|
||||
async fn complete_dispatched_run(&self, run: &RoutineRun, status: RunStatus, summary: &str) {
|
||||
if let Err(e) = self
|
||||
.store
|
||||
.complete_routine_run(run.id, status, Some(summary), None)
|
||||
.await
|
||||
{
|
||||
tracing::error!(
|
||||
run_id = %run.id,
|
||||
"Failed to update dispatched routine run: {}", e
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
// Look up the routine to get its notify config and name
|
||||
match self.store.get_routine(run.routine_id).await {
|
||||
Ok(Some(routine)) => {
|
||||
send_notification(
|
||||
&self.notify_tx,
|
||||
&routine.notify,
|
||||
&routine.name,
|
||||
status,
|
||||
Some(summary),
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Ok(None) => {
|
||||
tracing::debug!(
|
||||
routine_id = %run.routine_id,
|
||||
"Routine not found for notification (may have been deleted)"
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::debug!(
|
||||
routine_id = %run.routine_id,
|
||||
"Failed to look up routine for notification: {}", e
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Fire a routine manually (from tool call or CLI).
|
||||
///
|
||||
/// Bypasses cooldown checks (those only apply to cron/event triggers).
|
||||
/// Still enforces enabled check and concurrent run limit.
|
||||
pub async fn fire_manual(&self, routine_id: Uuid) -> Result<Uuid, RoutineError> {
|
||||
pub async fn fire_manual(
|
||||
&self,
|
||||
routine_id: Uuid,
|
||||
user_id: Option<&str>,
|
||||
) -> Result<Uuid, RoutineError> {
|
||||
let routine = self
|
||||
.store
|
||||
.get_routine(routine_id)
|
||||
@@ -194,6 +323,13 @@ impl RoutineEngine {
|
||||
})?
|
||||
.ok_or(RoutineError::NotFound { id: routine_id })?;
|
||||
|
||||
// Enforce ownership when a user_id is provided (gateway calls).
|
||||
if let Some(uid) = user_id
|
||||
&& routine.user_id != uid
|
||||
{
|
||||
return Err(RoutineError::NotAuthorized { id: routine_id });
|
||||
}
|
||||
|
||||
if !routine.enabled {
|
||||
return Err(RoutineError::Disabled {
|
||||
name: routine.name.clone(),
|
||||
@@ -396,6 +532,39 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
|
||||
tracing::error!(routine = %routine.name, "Failed to update runtime state: {}", e);
|
||||
}
|
||||
|
||||
// Persist routine result to its dedicated conversation thread
|
||||
let thread_id = match ctx
|
||||
.store
|
||||
.get_or_create_routine_conversation(routine.id, &routine.name, &routine.user_id)
|
||||
.await
|
||||
{
|
||||
Ok(conv_id) => {
|
||||
tracing::debug!(
|
||||
routine = %routine.name,
|
||||
routine_id = %routine.id,
|
||||
conversation_id = %conv_id,
|
||||
"Resolved routine conversation thread"
|
||||
);
|
||||
// Record the run result as a conversation message
|
||||
let msg = match (&summary, status) {
|
||||
(Some(s), _) => format!("[{}] {}: {}", run.trigger_type, status, s),
|
||||
(None, _) => format!("[{}] {}", run.trigger_type, status),
|
||||
};
|
||||
if let Err(e) = ctx
|
||||
.store
|
||||
.add_conversation_message(conv_id, "assistant", &msg)
|
||||
.await
|
||||
{
|
||||
tracing::error!(routine = %routine.name, "Failed to persist routine message: {}", e);
|
||||
}
|
||||
Some(conv_id.to_string())
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::error!(routine = %routine.name, "Failed to get routine conversation: {}", e);
|
||||
None
|
||||
}
|
||||
};
|
||||
|
||||
// Send notifications based on config
|
||||
send_notification(
|
||||
&ctx.notify_tx,
|
||||
@@ -403,6 +572,7 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
|
||||
&routine.name,
|
||||
status,
|
||||
summary.as_deref(),
|
||||
thread_id.as_deref(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
@@ -489,9 +659,10 @@ async fn execute_full_job(
|
||||
);
|
||||
|
||||
let summary = format!(
|
||||
"Dispatched job {job_id} for full execution with tool access (max_iterations: {max_iterations})"
|
||||
"Dispatched job {job_id} for full execution with tool access (max_iterations: {max_iterations}). \
|
||||
Status will be updated when the job completes."
|
||||
);
|
||||
Ok((RunStatus::Ok, Some(summary), None))
|
||||
Ok((RunStatus::Running, Some(summary), None))
|
||||
}
|
||||
|
||||
/// Execute a lightweight routine (single LLM call).
|
||||
@@ -611,6 +782,7 @@ async fn send_notification(
|
||||
routine_name: &str,
|
||||
status: RunStatus,
|
||||
summary: Option<&str>,
|
||||
thread_id: Option<&str>,
|
||||
) {
|
||||
let should_notify = match status {
|
||||
RunStatus::Ok => notify.on_success,
|
||||
@@ -637,7 +809,7 @@ async fn send_notification(
|
||||
|
||||
let response = OutgoingResponse {
|
||||
content: message,
|
||||
thread_id: None,
|
||||
thread_id: thread_id.map(String::from),
|
||||
attachments: Vec::new(),
|
||||
metadata: serde_json::json!({
|
||||
"source": "routine",
|
||||
@@ -666,6 +838,7 @@ pub fn spawn_cron_ticker(
|
||||
loop {
|
||||
ticker.tick().await;
|
||||
engine.check_cron_triggers().await;
|
||||
engine.sync_dispatched_runs().await;
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -710,4 +883,84 @@ mod tests {
|
||||
let _ = status.to_string();
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_running_status_does_not_notify() {
|
||||
// Running status should not trigger notifications (job still in progress)
|
||||
let config = NotifyConfig {
|
||||
on_success: true,
|
||||
on_failure: true,
|
||||
on_attention: true,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
// RunStatus::Running maps to false in send_notification's match
|
||||
let should_notify = match RunStatus::Running {
|
||||
RunStatus::Ok => config.on_success,
|
||||
RunStatus::Attention => config.on_attention,
|
||||
RunStatus::Failed => config.on_failure,
|
||||
RunStatus::Running => false,
|
||||
};
|
||||
assert!(!should_notify);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_full_job_dispatch_returns_running_status() {
|
||||
// Verify the status text for Running is "running"
|
||||
assert_eq!(RunStatus::Running.to_string(), "running");
|
||||
}
|
||||
|
||||
/// Regression test for #697: full_job routines were immediately marked Ok
|
||||
/// on dispatch, so failures/completions were never synced back. The fix
|
||||
/// changed dispatch to return Running and added sync_dispatched_runs which
|
||||
/// maps terminal job states to routine run statuses.
|
||||
#[test]
|
||||
fn test_job_state_to_run_status_mapping() {
|
||||
use crate::context::JobState;
|
||||
|
||||
// Helper that replicates the mapping logic from sync_dispatched_runs
|
||||
let map_state = |state: JobState, reason: Option<&str>| -> Option<(RunStatus, String)> {
|
||||
let last_reason = reason.map(|s| s.to_string());
|
||||
match state {
|
||||
JobState::Completed | JobState::Submitted | JobState::Accepted => {
|
||||
let summary =
|
||||
last_reason.unwrap_or_else(|| "Job completed successfully".to_string());
|
||||
Some((RunStatus::Ok, summary))
|
||||
}
|
||||
JobState::Failed => {
|
||||
let summary = last_reason
|
||||
.unwrap_or_else(|| "Job failed (no error message recorded)".to_string());
|
||||
Some((RunStatus::Failed, summary))
|
||||
}
|
||||
JobState::Cancelled => Some((RunStatus::Failed, "Job was cancelled".to_string())),
|
||||
JobState::Pending | JobState::InProgress | JobState::Stuck => None,
|
||||
}
|
||||
};
|
||||
|
||||
// Terminal states produce a status update
|
||||
let (status, _) = map_state(JobState::Completed, None).unwrap();
|
||||
assert_eq!(status, RunStatus::Ok);
|
||||
|
||||
let (status, _) = map_state(JobState::Submitted, None).unwrap();
|
||||
assert_eq!(status, RunStatus::Ok);
|
||||
|
||||
let (status, _) = map_state(JobState::Accepted, None).unwrap();
|
||||
assert_eq!(status, RunStatus::Ok);
|
||||
|
||||
let (status, summary) = map_state(JobState::Failed, Some("OOM killed")).unwrap();
|
||||
assert_eq!(status, RunStatus::Failed);
|
||||
assert_eq!(summary, "OOM killed");
|
||||
|
||||
let (status, summary) = map_state(JobState::Failed, None).unwrap();
|
||||
assert_eq!(status, RunStatus::Failed);
|
||||
assert!(summary.contains("no error message"));
|
||||
|
||||
let (status, _) = map_state(JobState::Cancelled, None).unwrap();
|
||||
assert_eq!(status, RunStatus::Failed);
|
||||
|
||||
// In-progress states should NOT produce a status update (skip)
|
||||
assert!(map_state(JobState::Pending, None).is_none());
|
||||
assert!(map_state(JobState::InProgress, None).is_none());
|
||||
assert!(map_state(JobState::Stuck, None).is_none());
|
||||
}
|
||||
}
|
||||
|
||||
+64
-33
@@ -15,7 +15,8 @@ use crate::db::Database;
|
||||
use crate::error::Error;
|
||||
use crate::hooks::HookRegistry;
|
||||
use crate::llm::{
|
||||
ActionPlan, ChatMessage, LlmProvider, Reasoning, ReasoningContext, RespondResult, ToolSelection,
|
||||
ActionPlan, ChatMessage, LlmProvider, Reasoning, ReasoningContext, RespondResult, ToolCall,
|
||||
ToolSelection,
|
||||
};
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::tools::rate_limiter::RateLimitResult;
|
||||
@@ -576,37 +577,54 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
}
|
||||
}
|
||||
}
|
||||
} else if selections.len() == 1 {
|
||||
consecutive_tool_intent_nudges = 0;
|
||||
// Single tool: execute directly
|
||||
let selection = &selections[0];
|
||||
tracing::debug!(
|
||||
"Job {} selecting tool: {} - {}",
|
||||
self.job_id,
|
||||
selection.tool_name,
|
||||
selection.reasoning
|
||||
);
|
||||
|
||||
let result = self
|
||||
.execute_tool(&selection.tool_name, &selection.parameters)
|
||||
.await;
|
||||
|
||||
self.process_tool_result(reason_ctx, selection, result)
|
||||
.await?;
|
||||
} else {
|
||||
// Multiple tools: execute in parallel
|
||||
tracing::debug!(
|
||||
"Job {} executing {} tools in parallel",
|
||||
self.job_id,
|
||||
selections.len()
|
||||
);
|
||||
consecutive_tool_intent_nudges = 0;
|
||||
|
||||
let results = self.execute_tools_parallel(&selections).await;
|
||||
// Record the assistant tool_calls message so that tool_result
|
||||
// messages have a matching parent (prevents orphaned rewrites).
|
||||
let tool_calls: Vec<ToolCall> = selections
|
||||
.iter()
|
||||
.map(|s| ToolCall {
|
||||
id: s.tool_call_id.clone(),
|
||||
name: s.tool_name.clone(),
|
||||
arguments: s.parameters.clone(),
|
||||
})
|
||||
.collect();
|
||||
reason_ctx
|
||||
.messages
|
||||
.push(ChatMessage::assistant_with_tool_calls(None, tool_calls));
|
||||
|
||||
// Process all results
|
||||
for (selection, result) in selections.iter().zip(results) {
|
||||
self.process_tool_result(reason_ctx, selection, result.result)
|
||||
if selections.len() == 1 {
|
||||
// Single tool: execute directly
|
||||
let selection = &selections[0];
|
||||
tracing::debug!(
|
||||
"Job {} selecting tool: {} - {}",
|
||||
self.job_id,
|
||||
selection.tool_name,
|
||||
selection.reasoning
|
||||
);
|
||||
|
||||
let result = self
|
||||
.execute_tool(&selection.tool_name, &selection.parameters)
|
||||
.await;
|
||||
|
||||
self.process_tool_result(reason_ctx, selection, result)
|
||||
.await?;
|
||||
} else {
|
||||
// Multiple tools: execute in parallel
|
||||
tracing::debug!(
|
||||
"Job {} executing {} tools in parallel",
|
||||
self.job_id,
|
||||
selections.len()
|
||||
);
|
||||
|
||||
let results = self.execute_tools_parallel(&selections).await;
|
||||
|
||||
// Process all results
|
||||
for (selection, result) in selections.iter().zip(results) {
|
||||
self.process_tool_result(reason_ctx, selection, result.result)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1087,11 +1105,6 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
action.reasoning
|
||||
);
|
||||
|
||||
// Execute the planned tool
|
||||
let result = self
|
||||
.execute_tool(&action.tool_name, &action.parameters)
|
||||
.await;
|
||||
|
||||
// Create a synthetic ToolSelection for process_tool_result.
|
||||
// Plan actions don't originate from an LLM tool_call response so
|
||||
// there is no real tool_call_id; generate a unique one.
|
||||
@@ -1103,6 +1116,24 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
tool_call_id: format!("plan_{}_{}", self.job_id, i),
|
||||
};
|
||||
|
||||
// Record the assistant tool_calls message so that the tool_result
|
||||
// has a matching parent (prevents orphaned rewrites).
|
||||
reason_ctx
|
||||
.messages
|
||||
.push(ChatMessage::assistant_with_tool_calls(
|
||||
None,
|
||||
vec![ToolCall {
|
||||
id: selection.tool_call_id.clone(),
|
||||
name: selection.tool_name.clone(),
|
||||
arguments: selection.parameters.clone(),
|
||||
}],
|
||||
));
|
||||
|
||||
// Execute the planned tool
|
||||
let result = self
|
||||
.execute_tool(&action.tool_name, &action.parameters)
|
||||
.await;
|
||||
|
||||
// Process the result
|
||||
let completed = self
|
||||
.process_tool_result(reason_ctx, &selection, result)
|
||||
|
||||
+28
@@ -244,11 +244,28 @@ impl AppBuilder {
|
||||
let master_key = match self.config.secrets.master_key() {
|
||||
Some(k) => k,
|
||||
None => {
|
||||
// No secrets DB available, but we can still load tokens from
|
||||
// OS credential stores (e.g., Anthropic OAuth via Claude Code's
|
||||
// macOS Keychain / Linux ~/.claude/.credentials.json).
|
||||
crate::config::inject_os_credentials();
|
||||
|
||||
// Consume unused handles
|
||||
#[cfg(feature = "libsql")]
|
||||
{
|
||||
self.libsql_db.take();
|
||||
}
|
||||
|
||||
// Re-resolve config with OS credentials
|
||||
if let Some(ref db) = self.db {
|
||||
let toml_path = self.toml_path.as_deref();
|
||||
if let Ok(refreshed) =
|
||||
Config::from_db_with_toml(db.as_ref(), "default", toml_path).await
|
||||
{
|
||||
self.config = refreshed;
|
||||
tracing::debug!("LlmConfig re-resolved after OS credential injection");
|
||||
}
|
||||
}
|
||||
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
@@ -665,6 +682,17 @@ impl AppBuilder {
|
||||
self.init_database().await?;
|
||||
self.init_secrets().await?;
|
||||
|
||||
// Post-init validation: if a non-nearai backend was selected but
|
||||
// credentials were never resolved (deferred resolution found no keys),
|
||||
// fail early with a clear error instead of a confusing runtime failure.
|
||||
if self.config.llm.backend != "nearai" && self.config.llm.provider.is_none() {
|
||||
let backend = &self.config.llm.backend;
|
||||
anyhow::bail!(
|
||||
"LLM_BACKEND={backend} is configured but no credentials were found. \
|
||||
Set the appropriate API key environment variable or run the setup wizard."
|
||||
);
|
||||
}
|
||||
|
||||
let (llm, cheap_llm, recording_handle) = if let Some(llm) = self.llm_override.take() {
|
||||
(llm, None, None)
|
||||
} else {
|
||||
|
||||
@@ -426,7 +426,7 @@ pub async fn chat_threads_handler(
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
if let Ok(summaries) = store
|
||||
.list_conversations_with_preview(&state.user_id, "gateway", 50)
|
||||
.list_conversations_all_channels(&state.user_id, 50)
|
||||
.await
|
||||
{
|
||||
let mut assistant_thread = None;
|
||||
@@ -441,6 +441,7 @@ pub async fn chat_threads_handler(
|
||||
updated_at: s.last_activity.to_rfc3339(),
|
||||
title: s.title.clone(),
|
||||
thread_type: s.thread_type.clone(),
|
||||
channel: Some(s.channel.clone()),
|
||||
};
|
||||
|
||||
if s.id == assistant_id {
|
||||
@@ -460,6 +461,7 @@ pub async fn chat_threads_handler(
|
||||
updated_at: chrono::Utc::now().to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: Some("assistant".to_string()),
|
||||
channel: Some("gateway".to_string()),
|
||||
});
|
||||
}
|
||||
|
||||
@@ -472,9 +474,10 @@ pub async fn chat_threads_handler(
|
||||
}
|
||||
|
||||
// Fallback: in-memory only (no assistant thread without DB)
|
||||
let threads: Vec<ThreadInfo> = sess
|
||||
.threads
|
||||
.values()
|
||||
let mut sorted_threads: Vec<_> = sess.threads.values().collect();
|
||||
sorted_threads.sort_by(|a, b| b.updated_at.cmp(&a.updated_at));
|
||||
let threads: Vec<ThreadInfo> = sorted_threads
|
||||
.into_iter()
|
||||
.map(|t| ThreadInfo {
|
||||
id: t.id,
|
||||
state: format!("{:?}", t.state),
|
||||
@@ -483,6 +486,7 @@ pub async fn chat_threads_handler(
|
||||
updated_at: t.updated_at.to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: None,
|
||||
channel: Some("gateway".to_string()),
|
||||
})
|
||||
.collect();
|
||||
|
||||
@@ -502,38 +506,39 @@ pub async fn chat_new_thread_handler(
|
||||
))?;
|
||||
|
||||
let session = session_manager.get_or_create_session(&state.user_id).await;
|
||||
let mut sess = session.lock().await;
|
||||
let thread = sess.create_thread();
|
||||
let thread_id = thread.id;
|
||||
let info = ThreadInfo {
|
||||
id: thread.id,
|
||||
state: format!("{:?}", thread.state),
|
||||
turn_count: thread.turns.len(),
|
||||
created_at: thread.created_at.to_rfc3339(),
|
||||
updated_at: thread.updated_at.to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: Some("thread".to_string()),
|
||||
let (thread_id, info) = {
|
||||
let mut sess = session.lock().await;
|
||||
let thread = sess.create_thread();
|
||||
let id = thread.id;
|
||||
let info = ThreadInfo {
|
||||
id: thread.id,
|
||||
state: format!("{:?}", thread.state),
|
||||
turn_count: thread.turns.len(),
|
||||
created_at: thread.created_at.to_rfc3339(),
|
||||
updated_at: thread.updated_at.to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: Some("thread".to_string()),
|
||||
channel: Some("gateway".to_string()),
|
||||
};
|
||||
(id, info)
|
||||
};
|
||||
|
||||
// Persist the empty conversation row with thread_type metadata
|
||||
// Persist the empty conversation row with thread_type metadata synchronously
|
||||
// so that the subsequent loadThreads() call from the frontend sees it.
|
||||
if let Some(ref store) = state.store {
|
||||
let store = Arc::clone(store);
|
||||
let user_id = state.user_id.clone();
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = store
|
||||
.ensure_conversation(thread_id, "gateway", &user_id, None)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to persist new thread: {}", e);
|
||||
}
|
||||
let metadata_val = serde_json::json!("thread");
|
||||
if let Err(e) = store
|
||||
.update_conversation_metadata_field(thread_id, "thread_type", &metadata_val)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to set thread_type metadata: {}", e);
|
||||
}
|
||||
});
|
||||
if let Err(e) = store
|
||||
.ensure_conversation(thread_id, "gateway", &state.user_id, None)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to persist new thread: {}", e);
|
||||
}
|
||||
let metadata_val = serde_json::json!("thread");
|
||||
if let Err(e) = store
|
||||
.update_conversation_metadata_field(thread_id, "thread_type", &metadata_val)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to set thread_type metadata: {}", e);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Json(info))
|
||||
|
||||
@@ -10,9 +10,9 @@ use axum::{
|
||||
use serde::Deserialize;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::channels::IncomingMessage;
|
||||
use crate::channels::web::server::GatewayState;
|
||||
use crate::channels::web::types::*;
|
||||
use crate::error::RoutineError;
|
||||
|
||||
pub async fn routines_list_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
@@ -133,56 +133,27 @@ pub async fn routines_trigger_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
Path(id): Path<String>,
|
||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||
let store = state.store.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Database not available".to_string(),
|
||||
))?;
|
||||
// Clone the Arc out of the lock to avoid holding the RwLock across .await.
|
||||
let engine = {
|
||||
let guard = state.routine_engine.read().await;
|
||||
guard.as_ref().cloned().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Routine engine not available".to_string(),
|
||||
))?
|
||||
};
|
||||
|
||||
let routine_id = Uuid::parse_str(&id)
|
||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
|
||||
|
||||
let routine = store
|
||||
.get_routine(routine_id)
|
||||
let run_id = engine
|
||||
.fire_manual(routine_id, Some(&state.user_id))
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
||||
|
||||
if routine.user_id != state.user_id {
|
||||
return Err((StatusCode::FORBIDDEN, "Access denied".to_string()));
|
||||
}
|
||||
|
||||
// Send the routine prompt through the message pipeline as a manual trigger.
|
||||
let prompt = match &routine.action {
|
||||
crate::agent::routine::RoutineAction::Lightweight { prompt, .. } => prompt.clone(),
|
||||
crate::agent::routine::RoutineAction::FullJob {
|
||||
title, description, ..
|
||||
} => format!("{}: {}", title, description),
|
||||
};
|
||||
|
||||
let content = format!("[routine:{}] {}", routine.name, prompt);
|
||||
let thread_id = format!(
|
||||
"routine-{}-{}",
|
||||
routine_id,
|
||||
chrono::Utc::now().timestamp_millis()
|
||||
);
|
||||
let msg = IncomingMessage::new("gateway", &state.user_id, content).with_thread(thread_id);
|
||||
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
let tx = tx_guard.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Channel not started".to_string(),
|
||||
))?;
|
||||
|
||||
tx.send(msg).await.map_err(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
"Channel closed".to_string(),
|
||||
)
|
||||
})?;
|
||||
.map_err(|e| (routine_error_status(&e), e.to_string()))?;
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"status": "triggered",
|
||||
"routine_id": routine_id,
|
||||
"run_id": run_id,
|
||||
})))
|
||||
}
|
||||
|
||||
@@ -337,3 +308,13 @@ fn routine_to_info(r: &crate::agent::routine::Routine) -> RoutineInfo {
|
||||
status: status.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Map `RoutineError` variants to appropriate HTTP status codes.
|
||||
fn routine_error_status(err: &RoutineError) -> StatusCode {
|
||||
match err {
|
||||
RoutineError::NotFound { .. } => StatusCode::NOT_FOUND,
|
||||
RoutineError::NotAuthorized { .. } => StatusCode::FORBIDDEN,
|
||||
RoutineError::Disabled { .. } | RoutineError::MaxConcurrent { .. } => StatusCode::CONFLICT,
|
||||
_ => StatusCode::INTERNAL_SERVER_ERROR,
|
||||
}
|
||||
}
|
||||
|
||||
+21
-2
@@ -99,6 +99,7 @@ impl GatewayChannel {
|
||||
chat_rate_limiter: server::RateLimiter::new(30, 60),
|
||||
registry_entries: Vec::new(),
|
||||
cost_guard: None,
|
||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
||||
startup_time: std::time::Instant::now(),
|
||||
});
|
||||
|
||||
@@ -134,6 +135,7 @@ impl GatewayChannel {
|
||||
chat_rate_limiter: server::RateLimiter::new(30, 60),
|
||||
registry_entries: self.state.registry_entries.clone(),
|
||||
cost_guard: self.state.cost_guard.clone(),
|
||||
routine_engine: Arc::clone(&self.state.routine_engine),
|
||||
startup_time: self.state.startup_time,
|
||||
};
|
||||
mutate(&mut new_state);
|
||||
@@ -281,7 +283,15 @@ impl Channel for GatewayChannel {
|
||||
msg: &IncomingMessage,
|
||||
response: OutgoingResponse,
|
||||
) -> Result<(), ChannelError> {
|
||||
let thread_id = msg.thread_id.clone().unwrap_or_default();
|
||||
let thread_id = match &msg.thread_id {
|
||||
Some(tid) => tid.clone(),
|
||||
None => {
|
||||
tracing::warn!(
|
||||
"Gateway respond with no thread_id — skipping (clients would drop it)"
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
|
||||
self.state.sse.broadcast(SseEvent::Response {
|
||||
content: response.content,
|
||||
@@ -387,9 +397,18 @@ impl Channel for GatewayChannel {
|
||||
_user_id: &str,
|
||||
response: OutgoingResponse,
|
||||
) -> Result<(), ChannelError> {
|
||||
let thread_id = match response.thread_id {
|
||||
Some(tid) => tid,
|
||||
None => {
|
||||
tracing::warn!(
|
||||
"Gateway broadcast with no thread_id — skipping (clients would drop it)"
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
};
|
||||
self.state.sse.broadcast(SseEvent::Response {
|
||||
content: response.content,
|
||||
thread_id: String::new(),
|
||||
thread_id,
|
||||
});
|
||||
Ok(())
|
||||
}
|
||||
|
||||
+71
-67
@@ -57,6 +57,10 @@ pub type PromptQueue = Arc<
|
||||
>,
|
||||
>;
|
||||
|
||||
/// Slot for the routine engine, filled at runtime after the agent starts.
|
||||
pub type RoutineEngineSlot =
|
||||
Arc<tokio::sync::RwLock<Option<Arc<crate::agent::routine_engine::RoutineEngine>>>>;
|
||||
|
||||
/// Simple sliding-window rate limiter.
|
||||
///
|
||||
/// Tracks the number of requests in the current window. Resets when the window expires.
|
||||
@@ -165,6 +169,8 @@ pub struct GatewayState {
|
||||
pub registry_entries: Vec<crate::extensions::RegistryEntry>,
|
||||
/// Cost guard for token/cost tracking.
|
||||
pub cost_guard: Option<Arc<crate::agent::cost_guard::CostGuard>>,
|
||||
/// Routine engine slot for manual routine triggering (filled at runtime).
|
||||
pub routine_engine: RoutineEngineSlot,
|
||||
/// Server startup time for uptime calculation.
|
||||
pub startup_time: std::time::Instant,
|
||||
}
|
||||
@@ -1037,7 +1043,7 @@ async fn chat_threads_handler(
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
if let Ok(summaries) = store
|
||||
.list_conversations_with_preview(&state.user_id, "gateway", 50)
|
||||
.list_conversations_all_channels(&state.user_id, 50)
|
||||
.await
|
||||
{
|
||||
let mut assistant_thread = None;
|
||||
@@ -1052,6 +1058,7 @@ async fn chat_threads_handler(
|
||||
updated_at: s.last_activity.to_rfc3339(),
|
||||
title: s.title.clone(),
|
||||
thread_type: s.thread_type.clone(),
|
||||
channel: Some(s.channel.clone()),
|
||||
};
|
||||
|
||||
if s.id == assistant_id {
|
||||
@@ -1071,6 +1078,7 @@ async fn chat_threads_handler(
|
||||
updated_at: chrono::Utc::now().to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: Some("assistant".to_string()),
|
||||
channel: Some("gateway".to_string()),
|
||||
});
|
||||
}
|
||||
|
||||
@@ -1083,9 +1091,10 @@ async fn chat_threads_handler(
|
||||
}
|
||||
|
||||
// Fallback: in-memory only (no assistant thread without DB)
|
||||
let threads: Vec<ThreadInfo> = sess
|
||||
.threads
|
||||
.values()
|
||||
let mut sorted_threads: Vec<_> = sess.threads.values().collect();
|
||||
sorted_threads.sort_by(|a, b| b.updated_at.cmp(&a.updated_at));
|
||||
let threads: Vec<ThreadInfo> = sorted_threads
|
||||
.into_iter()
|
||||
.map(|t| ThreadInfo {
|
||||
id: t.id,
|
||||
state: format!("{:?}", t.state),
|
||||
@@ -1094,6 +1103,7 @@ async fn chat_threads_handler(
|
||||
updated_at: t.updated_at.to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: None,
|
||||
channel: Some("gateway".to_string()),
|
||||
})
|
||||
.collect();
|
||||
|
||||
@@ -1113,38 +1123,39 @@ async fn chat_new_thread_handler(
|
||||
))?;
|
||||
|
||||
let session = session_manager.get_or_create_session(&state.user_id).await;
|
||||
let mut sess = session.lock().await;
|
||||
let thread = sess.create_thread();
|
||||
let thread_id = thread.id;
|
||||
let info = ThreadInfo {
|
||||
id: thread.id,
|
||||
state: format!("{:?}", thread.state),
|
||||
turn_count: thread.turns.len(),
|
||||
created_at: thread.created_at.to_rfc3339(),
|
||||
updated_at: thread.updated_at.to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: Some("thread".to_string()),
|
||||
let (thread_id, info) = {
|
||||
let mut sess = session.lock().await;
|
||||
let thread = sess.create_thread();
|
||||
let id = thread.id;
|
||||
let info = ThreadInfo {
|
||||
id: thread.id,
|
||||
state: format!("{:?}", thread.state),
|
||||
turn_count: thread.turns.len(),
|
||||
created_at: thread.created_at.to_rfc3339(),
|
||||
updated_at: thread.updated_at.to_rfc3339(),
|
||||
title: None,
|
||||
thread_type: Some("thread".to_string()),
|
||||
channel: Some("gateway".to_string()),
|
||||
};
|
||||
(id, info)
|
||||
};
|
||||
|
||||
// Persist the empty conversation row with thread_type metadata
|
||||
// Persist the empty conversation row with thread_type metadata synchronously
|
||||
// so that the subsequent loadThreads() call from the frontend sees it.
|
||||
if let Some(ref store) = state.store {
|
||||
let store = Arc::clone(store);
|
||||
let user_id = state.user_id.clone();
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = store
|
||||
.ensure_conversation(thread_id, "gateway", &user_id, None)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to persist new thread: {}", e);
|
||||
}
|
||||
let metadata_val = serde_json::json!("thread");
|
||||
if let Err(e) = store
|
||||
.update_conversation_metadata_field(thread_id, "thread_type", &metadata_val)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to set thread_type metadata: {}", e);
|
||||
}
|
||||
});
|
||||
if let Err(e) = store
|
||||
.ensure_conversation(thread_id, "gateway", &state.user_id, None)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to persist new thread: {}", e);
|
||||
}
|
||||
let metadata_val = serde_json::json!("thread");
|
||||
if let Err(e) = store
|
||||
.update_conversation_metadata_field(thread_id, "thread_type", &metadata_val)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to set thread_type metadata: {}", e);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(Json(info))
|
||||
@@ -1965,47 +1976,35 @@ async fn routines_trigger_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
Path(id): Path<String>,
|
||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||
let store = state.store.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Database not available".to_string(),
|
||||
))?;
|
||||
let engine = {
|
||||
let guard = state.routine_engine.read().await;
|
||||
guard.as_ref().cloned().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Routine engine not available".to_string(),
|
||||
))?
|
||||
};
|
||||
|
||||
let routine_id = Uuid::parse_str(&id)
|
||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
|
||||
|
||||
let routine = store
|
||||
.get_routine(routine_id)
|
||||
let run_id = engine
|
||||
.fire_manual(routine_id, Some(&state.user_id))
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
||||
|
||||
// Send the routine prompt through the message pipeline as a manual trigger.
|
||||
let prompt = match &routine.action {
|
||||
crate::agent::routine::RoutineAction::Lightweight { prompt, .. } => prompt.clone(),
|
||||
crate::agent::routine::RoutineAction::FullJob {
|
||||
title, description, ..
|
||||
} => format!("{}: {}", title, description),
|
||||
};
|
||||
|
||||
let content = format!("[routine:{}] {}", routine.name, prompt);
|
||||
let msg = IncomingMessage::new("gateway", &state.user_id, content);
|
||||
|
||||
let tx_guard = state.msg_tx.read().await;
|
||||
let tx = tx_guard.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Channel not started".to_string(),
|
||||
))?;
|
||||
|
||||
tx.send(msg).await.map_err(|_| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
"Channel closed".to_string(),
|
||||
)
|
||||
})?;
|
||||
.map_err(|e| {
|
||||
let status = match &e {
|
||||
crate::error::RoutineError::NotFound { .. } => StatusCode::NOT_FOUND,
|
||||
crate::error::RoutineError::NotAuthorized { .. } => StatusCode::FORBIDDEN,
|
||||
crate::error::RoutineError::Disabled { .. }
|
||||
| crate::error::RoutineError::MaxConcurrent { .. } => StatusCode::CONFLICT,
|
||||
_ => StatusCode::INTERNAL_SERVER_ERROR,
|
||||
};
|
||||
(status, e.to_string())
|
||||
})?;
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"status": "triggered",
|
||||
"routine_id": routine_id,
|
||||
"run_id": run_id,
|
||||
})))
|
||||
}
|
||||
|
||||
@@ -2463,6 +2462,7 @@ mod tests {
|
||||
chat_rate_limiter: RateLimiter::new(30, 60),
|
||||
registry_entries: vec![],
|
||||
cost_guard: None,
|
||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
||||
startup_time: std::time::Instant::now(),
|
||||
})
|
||||
}
|
||||
@@ -2620,7 +2620,9 @@ mod tests {
|
||||
secrets,
|
||||
sse_sender: None,
|
||||
gateway_token: None,
|
||||
created_at: std::time::Instant::now() - std::time::Duration::from_secs(600),
|
||||
created_at: std::time::Instant::now()
|
||||
.checked_sub(std::time::Duration::from_secs(600))
|
||||
.expect("System uptime is too low to run expired flow test"),
|
||||
};
|
||||
|
||||
ext_mgr
|
||||
@@ -2727,7 +2729,9 @@ mod tests {
|
||||
sse_sender: None,
|
||||
gateway_token: None,
|
||||
// Expired — handler will reject after lookup (no network I/O)
|
||||
created_at: std::time::Instant::now() - std::time::Duration::from_secs(600),
|
||||
created_at: std::time::Instant::now()
|
||||
.checked_sub(std::time::Duration::from_secs(600))
|
||||
.expect("System uptime is too low to run expired flow test"),
|
||||
};
|
||||
|
||||
ext_mgr
|
||||
|
||||
+120
-17
@@ -5,6 +5,7 @@ let eventSource = null;
|
||||
let logEventSource = null;
|
||||
let currentTab = 'chat';
|
||||
let currentThreadId = null;
|
||||
let currentThreadIsReadOnly = false;
|
||||
let assistantThreadId = null;
|
||||
let hasMore = false;
|
||||
let oldestTimestamp = null;
|
||||
@@ -13,6 +14,8 @@ let sseHasConnectedBefore = false;
|
||||
let jobEvents = new Map(); // job_id -> Array of events
|
||||
let jobListRefreshTimer = null;
|
||||
let pairingPollInterval = null;
|
||||
let unreadThreads = new Map(); // thread_id -> unread count
|
||||
let _loadThreadsTimer = null;
|
||||
const JOB_EVENTS_CAP = 500;
|
||||
const MEMORY_SEARCH_QUERY_MAX_LENGTH = 100;
|
||||
|
||||
@@ -273,7 +276,13 @@ function connectSSE() {
|
||||
|
||||
eventSource.addEventListener('response', (e) => {
|
||||
const data = JSON.parse(e.data);
|
||||
if (!isCurrentThread(data.thread_id)) return;
|
||||
if (!isCurrentThread(data.thread_id)) {
|
||||
if (data.thread_id) {
|
||||
unreadThreads.set(data.thread_id, (unreadThreads.get(data.thread_id) || 0) + 1);
|
||||
debouncedLoadThreads();
|
||||
}
|
||||
return;
|
||||
}
|
||||
finalizeActivityGroup();
|
||||
addMessage('assistant', data.content);
|
||||
enableChatInput();
|
||||
@@ -288,7 +297,10 @@ function connectSSE() {
|
||||
|
||||
eventSource.addEventListener('thinking', (e) => {
|
||||
const data = JSON.parse(e.data);
|
||||
if (!isCurrentThread(data.thread_id)) return;
|
||||
if (!isCurrentThread(data.thread_id)) {
|
||||
if (data.thread_id) debouncedLoadThreads();
|
||||
return;
|
||||
}
|
||||
showActivityThinking(data.message);
|
||||
});
|
||||
|
||||
@@ -324,7 +336,10 @@ function connectSSE() {
|
||||
|
||||
eventSource.addEventListener('status', (e) => {
|
||||
const data = JSON.parse(e.data);
|
||||
if (!isCurrentThread(data.thread_id)) return;
|
||||
if (!isCurrentThread(data.thread_id)) {
|
||||
if (data.thread_id) debouncedLoadThreads();
|
||||
return;
|
||||
}
|
||||
// "Done" and "Awaiting approval" are terminal signals from the agent:
|
||||
// the agentic loop finished, so re-enable input as a safety net in case
|
||||
// the response SSE event is empty or lost.
|
||||
@@ -414,9 +429,9 @@ function connectSSE() {
|
||||
}
|
||||
|
||||
// Check if an SSE event belongs to the currently viewed thread.
|
||||
// Events without a thread_id (legacy) are always shown.
|
||||
// Events without a thread_id are dropped (prevents notification leaking).
|
||||
function isCurrentThread(threadId) {
|
||||
if (!threadId) return true;
|
||||
if (!threadId) return false;
|
||||
if (!currentThreadId) return true;
|
||||
return threadId === currentThreadId;
|
||||
}
|
||||
@@ -446,7 +461,14 @@ function sendMessage() {
|
||||
}
|
||||
|
||||
function enableChatInput() {
|
||||
// no-op: input and send button are always enabled
|
||||
if (currentThreadIsReadOnly) return;
|
||||
const input = document.getElementById('chat-input');
|
||||
const btn = document.getElementById('send-btn');
|
||||
if (input) {
|
||||
input.disabled = false;
|
||||
input.placeholder = 'Message or / for commands...';
|
||||
}
|
||||
if (btn) btn.disabled = false;
|
||||
}
|
||||
|
||||
// --- Slash Autocomplete ---
|
||||
@@ -1134,7 +1156,9 @@ function loadHistory(before) {
|
||||
// Fresh load: clear and render
|
||||
container.innerHTML = '';
|
||||
for (const turn of data.turns) {
|
||||
addMessage('user', turn.user_input);
|
||||
if (turn.user_input) {
|
||||
addMessage('user', turn.user_input);
|
||||
}
|
||||
if (turn.tool_calls && turn.tool_calls.length > 0) {
|
||||
addToolCallsSummary(turn.tool_calls);
|
||||
}
|
||||
@@ -1156,8 +1180,10 @@ function loadHistory(before) {
|
||||
const savedHeight = container.scrollHeight;
|
||||
const fragment = document.createDocumentFragment();
|
||||
for (const turn of data.turns) {
|
||||
const userDiv = createMessageElement('user', turn.user_input);
|
||||
fragment.appendChild(userDiv);
|
||||
if (turn.user_input) {
|
||||
const userDiv = createMessageElement('user', turn.user_input);
|
||||
fragment.appendChild(userDiv);
|
||||
}
|
||||
if (turn.tool_calls && turn.tool_calls.length > 0) {
|
||||
fragment.appendChild(createToolCallsSummaryElement(turn.tool_calls));
|
||||
}
|
||||
@@ -1256,6 +1282,37 @@ function removeScrollSpinner() {
|
||||
|
||||
// --- Threads ---
|
||||
|
||||
function threadTitle(thread) {
|
||||
if (thread.title) return thread.title;
|
||||
const ch = thread.channel || 'gateway';
|
||||
if (thread.thread_type === 'heartbeat') return 'Heartbeat Alerts';
|
||||
if (thread.thread_type === 'routine') return 'Routine';
|
||||
if (ch !== 'gateway') return ch.charAt(0).toUpperCase() + ch.slice(1);
|
||||
if (thread.turn_count === 0) return 'New chat';
|
||||
return thread.id.substring(0, 8);
|
||||
}
|
||||
|
||||
function relativeTime(isoStr) {
|
||||
if (!isoStr) return '';
|
||||
const diff = Date.now() - new Date(isoStr).getTime();
|
||||
const mins = Math.floor(diff / 60000);
|
||||
if (mins < 1) return 'now';
|
||||
if (mins < 60) return mins + 'm ago';
|
||||
const hrs = Math.floor(mins / 60);
|
||||
if (hrs < 24) return hrs + 'h ago';
|
||||
const days = Math.floor(hrs / 24);
|
||||
return days + 'd ago';
|
||||
}
|
||||
|
||||
function isReadOnlyChannel(channel) {
|
||||
return channel && channel !== 'gateway' && channel !== 'routine' && channel !== 'heartbeat';
|
||||
}
|
||||
|
||||
function debouncedLoadThreads() {
|
||||
if (_loadThreadsTimer) clearTimeout(_loadThreadsTimer);
|
||||
_loadThreadsTimer = setTimeout(() => { _loadThreadsTimer = null; loadThreads(); }, 500);
|
||||
}
|
||||
|
||||
function loadThreads() {
|
||||
apiFetch('/api/chat/threads').then((data) => {
|
||||
// Pinned assistant thread
|
||||
@@ -1264,9 +1321,13 @@ function loadThreads() {
|
||||
const el = document.getElementById('assistant-thread');
|
||||
const isActive = currentThreadId === assistantThreadId;
|
||||
el.className = 'assistant-item' + (isActive ? ' active' : '');
|
||||
const labelEl = document.getElementById('assistant-label');
|
||||
if (labelEl) {
|
||||
const at = data.assistant_thread;
|
||||
labelEl.textContent = 'Assistant';
|
||||
}
|
||||
const meta = document.getElementById('assistant-meta');
|
||||
const count = data.assistant_thread.turn_count || 0;
|
||||
meta.textContent = count > 0 ? count + ' turns' : '';
|
||||
meta.textContent = relativeTime(data.assistant_thread.updated_at);
|
||||
}
|
||||
|
||||
// Regular threads
|
||||
@@ -1275,16 +1336,38 @@ function loadThreads() {
|
||||
const threads = data.threads || [];
|
||||
for (const thread of threads) {
|
||||
const item = document.createElement('div');
|
||||
item.className = 'thread-item' + (thread.id === currentThreadId ? ' active' : '');
|
||||
const isActive = thread.id === currentThreadId;
|
||||
item.className = 'thread-item' + (isActive ? ' active' : '');
|
||||
|
||||
// Channel badge for non-gateway threads
|
||||
const ch = thread.channel || 'gateway';
|
||||
if (ch !== 'gateway') {
|
||||
const badge = document.createElement('span');
|
||||
badge.className = 'thread-badge thread-badge-' + ch;
|
||||
badge.textContent = ch;
|
||||
item.appendChild(badge);
|
||||
}
|
||||
|
||||
const label = document.createElement('span');
|
||||
label.className = 'thread-label';
|
||||
label.textContent = thread.title || thread.id.substring(0, 8);
|
||||
label.title = thread.title ? thread.title + ' (' + thread.id + ')' : thread.id;
|
||||
label.textContent = threadTitle(thread);
|
||||
label.title = (thread.title || '') + ' (' + thread.id + ')';
|
||||
item.appendChild(label);
|
||||
|
||||
const meta = document.createElement('span');
|
||||
meta.className = 'thread-meta';
|
||||
meta.textContent = (thread.turn_count || 0) + ' turns';
|
||||
meta.textContent = relativeTime(thread.updated_at);
|
||||
item.appendChild(meta);
|
||||
|
||||
// Unread dot
|
||||
const unread = unreadThreads.get(thread.id) || 0;
|
||||
if (unread > 0 && !isActive) {
|
||||
const dot = document.createElement('span');
|
||||
dot.className = 'thread-unread';
|
||||
dot.textContent = unread > 9 ? '9+' : String(unread);
|
||||
item.appendChild(dot);
|
||||
}
|
||||
|
||||
item.addEventListener('click', () => switchThread(thread.id));
|
||||
list.appendChild(item);
|
||||
}
|
||||
@@ -1294,17 +1377,36 @@ function loadThreads() {
|
||||
switchToAssistant();
|
||||
}
|
||||
|
||||
// Enable chat input once a thread is available
|
||||
// Enable/disable chat input based on channel type
|
||||
if (currentThreadId) {
|
||||
enableChatInput();
|
||||
const currentThread = threads.find(t => t.id === currentThreadId);
|
||||
const ch = currentThread ? currentThread.channel : 'gateway';
|
||||
currentThreadIsReadOnly = isReadOnlyChannel(ch);
|
||||
if (currentThreadIsReadOnly) {
|
||||
disableChatInputReadOnly();
|
||||
} else {
|
||||
enableChatInput();
|
||||
}
|
||||
}
|
||||
}).catch(() => {});
|
||||
}
|
||||
|
||||
function disableChatInputReadOnly() {
|
||||
const input = document.getElementById('chat-input');
|
||||
const btn = document.getElementById('send-btn');
|
||||
if (input) {
|
||||
input.disabled = true;
|
||||
input.placeholder = 'Read-only thread (external channel)';
|
||||
}
|
||||
if (btn) btn.disabled = true;
|
||||
}
|
||||
|
||||
function switchToAssistant() {
|
||||
if (!assistantThreadId) return;
|
||||
finalizeActivityGroup();
|
||||
currentThreadId = assistantThreadId;
|
||||
currentThreadIsReadOnly = false;
|
||||
unreadThreads.delete(assistantThreadId);
|
||||
hasMore = false;
|
||||
oldestTimestamp = null;
|
||||
loadHistory();
|
||||
@@ -1314,6 +1416,7 @@ function switchToAssistant() {
|
||||
function switchThread(threadId) {
|
||||
finalizeActivityGroup();
|
||||
currentThreadId = threadId;
|
||||
unreadThreads.delete(threadId);
|
||||
hasMore = false;
|
||||
oldestTimestamp = null;
|
||||
loadHistory();
|
||||
|
||||
@@ -113,12 +113,12 @@
|
||||
<div class="tab-panel active" id="tab-chat">
|
||||
<div class="thread-sidebar" id="thread-sidebar">
|
||||
<div class="thread-sidebar-header">
|
||||
<span>Threads</span>
|
||||
<button class="thread-new-btn" onclick="createNewThread()" title="New thread (Ctrl/Cmd+N)">+</button>
|
||||
<div class="spacer"></div>
|
||||
<button class="thread-toggle-btn" id="thread-toggle-btn" onclick="toggleThreadSidebar()" title="Toggle sidebar">«</button>
|
||||
</div>
|
||||
<div class="assistant-item" id="assistant-thread" onclick="switchToAssistant()">
|
||||
<span class="assistant-label">Assistant</span>
|
||||
<span class="assistant-label" id="assistant-label">Assistant</span>
|
||||
<span class="assistant-meta" id="assistant-meta"></span>
|
||||
</div>
|
||||
<div class="threads-section-header">
|
||||
|
||||
@@ -3074,7 +3074,7 @@ mark {
|
||||
}
|
||||
|
||||
.thread-sidebar {
|
||||
width: 200px;
|
||||
width: 240px;
|
||||
background: var(--bg-secondary);
|
||||
border-right: 1px solid var(--border);
|
||||
display: flex;
|
||||
@@ -3082,6 +3082,8 @@ mark {
|
||||
flex-shrink: 0;
|
||||
transition: width 0.2s ease;
|
||||
overflow: hidden;
|
||||
padding: 6px;
|
||||
gap: 2px;
|
||||
}
|
||||
|
||||
.thread-sidebar.collapsed {
|
||||
@@ -3099,8 +3101,7 @@ mark {
|
||||
.thread-sidebar-header {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
padding: 10px 12px;
|
||||
border-bottom: 1px solid var(--border);
|
||||
padding: 10px 10px;
|
||||
font-size: 13px;
|
||||
font-weight: 600;
|
||||
gap: 8px;
|
||||
@@ -3134,21 +3135,22 @@ mark {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
padding: 10px 12px;
|
||||
padding: 12px 14px;
|
||||
cursor: pointer;
|
||||
font-size: 13px;
|
||||
font-weight: 600;
|
||||
color: var(--text);
|
||||
border-bottom: 1px solid var(--border);
|
||||
background: var(--bg-secondary);
|
||||
background: var(--bg-tertiary);
|
||||
border-radius: var(--radius);
|
||||
margin-bottom: 2px;
|
||||
}
|
||||
|
||||
.assistant-item:hover {
|
||||
background: var(--bg-tertiary);
|
||||
background: rgba(255, 255, 255, 0.06);
|
||||
}
|
||||
|
||||
.assistant-item.active {
|
||||
background: rgba(52, 211, 153, 0.08);
|
||||
background: rgba(52, 211, 153, 0.1);
|
||||
color: var(--accent);
|
||||
border-left: 2px solid var(--accent);
|
||||
}
|
||||
@@ -3166,7 +3168,7 @@ mark {
|
||||
}
|
||||
|
||||
.threads-section-header {
|
||||
padding: 8px 12px 4px;
|
||||
padding: 10px 10px 4px;
|
||||
font-size: 11px;
|
||||
font-weight: 500;
|
||||
text-transform: uppercase;
|
||||
@@ -3196,11 +3198,11 @@ mark {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
padding: 8px 12px;
|
||||
padding: 10px 14px;
|
||||
cursor: pointer;
|
||||
font-size: 13px;
|
||||
color: var(--text-secondary);
|
||||
border-bottom: 1px solid rgba(255, 255, 255, 0.03);
|
||||
border-radius: var(--radius);
|
||||
}
|
||||
|
||||
.thread-item:hover {
|
||||
@@ -3222,6 +3224,43 @@ mark {
|
||||
.thread-meta {
|
||||
font-size: 11px;
|
||||
color: var(--text-secondary);
|
||||
flex-shrink: 0;
|
||||
}
|
||||
|
||||
.thread-badge {
|
||||
display: inline-block;
|
||||
font-size: 9px;
|
||||
font-weight: 600;
|
||||
text-transform: uppercase;
|
||||
letter-spacing: 0.5px;
|
||||
padding: 1px 5px;
|
||||
border-radius: 3px;
|
||||
background: rgba(255, 255, 255, 0.08);
|
||||
color: var(--text-secondary);
|
||||
margin-right: 6px;
|
||||
flex-shrink: 0;
|
||||
}
|
||||
|
||||
.thread-badge-routine { background: rgba(52, 211, 153, 0.15); color: var(--accent); }
|
||||
.thread-badge-heartbeat { background: rgba(245, 166, 35, 0.15); color: var(--warning); }
|
||||
.thread-badge-telegram { background: rgba(0, 136, 204, 0.15); color: #0088cc; }
|
||||
.thread-badge-signal { background: rgba(59, 118, 240, 0.15); color: #3b76f0; }
|
||||
.thread-badge-slack { background: rgba(74, 21, 75, 0.15); color: #e01e5a; }
|
||||
|
||||
.thread-unread {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
min-width: 16px;
|
||||
height: 16px;
|
||||
font-size: 10px;
|
||||
font-weight: 700;
|
||||
background: var(--accent);
|
||||
color: var(--bg);
|
||||
border-radius: 8px;
|
||||
padding: 0 4px;
|
||||
margin-left: auto;
|
||||
flex-shrink: 0;
|
||||
}
|
||||
|
||||
/* --- Memory editing --- */
|
||||
@@ -3620,7 +3659,7 @@ mark {
|
||||
left: 0;
|
||||
top: 0;
|
||||
bottom: 0;
|
||||
width: 200px;
|
||||
width: 240px;
|
||||
z-index: 50;
|
||||
}
|
||||
|
||||
|
||||
@@ -84,6 +84,7 @@ impl TestGatewayBuilder {
|
||||
chat_rate_limiter: RateLimiter::new(30, 60),
|
||||
registry_entries: Vec::new(),
|
||||
cost_guard: None,
|
||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
||||
startup_time: std::time::Instant::now(),
|
||||
})
|
||||
}
|
||||
|
||||
@@ -28,6 +28,8 @@ pub struct ThreadInfo {
|
||||
pub title: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub thread_type: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub channel: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
@@ -1063,4 +1065,40 @@ mod tests {
|
||||
let req: AuthCancelRequest = serde_json::from_str(json).unwrap();
|
||||
assert_eq!(req.extension_name, "telegram");
|
||||
}
|
||||
|
||||
// ---- ThreadInfo channel field tests ----
|
||||
|
||||
#[test]
|
||||
fn test_thread_info_channel_serialized() {
|
||||
let info = ThreadInfo {
|
||||
id: Uuid::nil(),
|
||||
state: "Idle".to_string(),
|
||||
turn_count: 0,
|
||||
created_at: "2026-01-01T00:00:00Z".to_string(),
|
||||
updated_at: "2026-01-01T00:00:00Z".to_string(),
|
||||
title: None,
|
||||
thread_type: None,
|
||||
channel: Some("telegram".to_string()),
|
||||
};
|
||||
let json = serde_json::to_string(&info).unwrap();
|
||||
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
|
||||
assert_eq!(parsed["channel"], "telegram");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_thread_info_channel_omitted_when_none() {
|
||||
let info = ThreadInfo {
|
||||
id: Uuid::nil(),
|
||||
state: "Idle".to_string(),
|
||||
turn_count: 0,
|
||||
created_at: "2026-01-01T00:00:00Z".to_string(),
|
||||
updated_at: "2026-01-01T00:00:00Z".to_string(),
|
||||
title: None,
|
||||
thread_type: None,
|
||||
channel: None,
|
||||
};
|
||||
let json = serde_json::to_string(&info).unwrap();
|
||||
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
|
||||
assert!(parsed.get("channel").is_none());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -83,6 +83,19 @@ pub fn build_turns_from_db_messages(
|
||||
|
||||
turns.push(turn);
|
||||
turn_number += 1;
|
||||
} else if msg.role == "assistant" {
|
||||
// Standalone assistant message (e.g. routine output, heartbeat)
|
||||
// with no preceding user message — render as a turn with empty input.
|
||||
turns.push(TurnInfo {
|
||||
turn_number,
|
||||
user_input: String::new(),
|
||||
response: Some(msg.content.clone()),
|
||||
state: "Completed".to_string(),
|
||||
started_at: msg.created_at.to_rfc3339(),
|
||||
completed_at: Some(msg.created_at.to_rfc3339()),
|
||||
tool_calls: Vec::new(),
|
||||
});
|
||||
turn_number += 1;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -220,6 +233,29 @@ mod tests {
|
||||
assert_eq!(turns[0].response.as_deref(), Some("Done"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_turns_standalone_assistant_messages() {
|
||||
// Routine conversations only have assistant messages (no user messages).
|
||||
let messages = vec![
|
||||
make_msg("assistant", "Routine executed: all checks passed", 0),
|
||||
make_msg("assistant", "Routine executed: found 2 issues", 5000),
|
||||
];
|
||||
let turns = build_turns_from_db_messages(&messages);
|
||||
assert_eq!(turns.len(), 2);
|
||||
// Standalone assistant messages should have empty user_input
|
||||
assert_eq!(turns[0].user_input, "");
|
||||
assert_eq!(
|
||||
turns[0].response.as_deref(),
|
||||
Some("Routine executed: all checks passed")
|
||||
);
|
||||
assert_eq!(turns[0].state, "Completed");
|
||||
assert_eq!(turns[1].user_input, "");
|
||||
assert_eq!(
|
||||
turns[1].response.as_deref(),
|
||||
Some("Routine executed: found 2 issues")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_turns_backward_compatible() {
|
||||
let messages = vec![
|
||||
|
||||
@@ -493,6 +493,7 @@ mod tests {
|
||||
chat_rate_limiter: crate::channels::web::server::RateLimiter::new(30, 60),
|
||||
registry_entries: Vec::new(),
|
||||
cost_guard: None,
|
||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
||||
startup_time: std::time::Instant::now(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -25,8 +25,13 @@ pub(crate) fn optional_env(key: &str) -> Result<Option<String>, ConfigError> {
|
||||
}
|
||||
|
||||
// Fall back to thread-safe overlay (secrets injected from DB)
|
||||
if let Some(val) = INJECTED_VARS.get().and_then(|map| map.get(key)) {
|
||||
return Ok(Some(val.clone()));
|
||||
if let Some(val) = INJECTED_VARS
|
||||
.lock()
|
||||
.unwrap_or_else(|p| p.into_inner())
|
||||
.get(key)
|
||||
.cloned()
|
||||
{
|
||||
return Ok(Some(val));
|
||||
}
|
||||
|
||||
Ok(None)
|
||||
|
||||
+141
-6
@@ -9,6 +9,13 @@ use crate::llm::registry::{ProviderProtocol, ProviderRegistry};
|
||||
use crate::llm::session::SessionConfig;
|
||||
use crate::settings::Settings;
|
||||
|
||||
/// Sentinel value used as `api_key` when only an OAuth token is present.
|
||||
///
|
||||
/// When we only have an OAuth token the provider factory in `llm/mod.rs`
|
||||
/// checks for this value and routes to `AnthropicOAuthProvider`, so this
|
||||
/// placeholder is never sent over the wire.
|
||||
pub const OAUTH_PLACEHOLDER: &str = "oauth-placeholder";
|
||||
|
||||
/// Prompt cache retention policy for Anthropic.
|
||||
///
|
||||
/// Controls Anthropic's automatic prompt caching via a top-level
|
||||
@@ -66,6 +73,7 @@ pub struct RegistryProviderConfig {
|
||||
/// Provider identifier (e.g., "groq", "openai", "tinfoil").
|
||||
pub provider_id: String,
|
||||
/// API key (optional for some providers like Ollama).
|
||||
/// For Anthropic OAuth, this is set to `OAUTH_PLACEHOLDER`.
|
||||
pub api_key: Option<SecretString>,
|
||||
/// Base URL for the API endpoint.
|
||||
pub base_url: String,
|
||||
@@ -73,6 +81,9 @@ pub struct RegistryProviderConfig {
|
||||
pub model: String,
|
||||
/// Extra HTTP headers injected into every request.
|
||||
pub extra_headers: Vec<(String, String)>,
|
||||
/// OAuth token for providers that support Bearer auth (e.g. Anthropic via `claude login`).
|
||||
/// When set, the provider factory routes to the OAuth-specific provider implementation.
|
||||
pub oauth_token: Option<SecretString>,
|
||||
}
|
||||
|
||||
/// LLM provider configuration.
|
||||
@@ -366,6 +377,22 @@ impl LlmConfig {
|
||||
Vec::new()
|
||||
};
|
||||
|
||||
// Resolve OAuth token (Anthropic-specific: `claude login` flow).
|
||||
// Only check for OAuth token when the provider is actually Anthropic.
|
||||
let oauth_token = if canonical_id == "anthropic" {
|
||||
optional_env("ANTHROPIC_OAUTH_TOKEN")?.map(SecretString::from)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let api_key = if api_key.is_none() && oauth_token.is_some() {
|
||||
// OAuth token present but no API key: use a placeholder so the
|
||||
// config block is populated. The provider factory will route to
|
||||
// the OAuth provider instead of rig-core's x-api-key client.
|
||||
Some(SecretString::from(OAUTH_PLACEHOLDER.to_string()))
|
||||
} else {
|
||||
api_key
|
||||
};
|
||||
|
||||
Ok(RegistryProviderConfig {
|
||||
protocol,
|
||||
provider_id: canonical_id.to_string(),
|
||||
@@ -373,6 +400,7 @@ impl LlmConfig {
|
||||
base_url,
|
||||
model,
|
||||
extra_headers,
|
||||
oauth_token,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -677,8 +705,6 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn backend_alias_normalized_to_canonical_id() {
|
||||
// When the user sets LLM_BACKEND to an alias (e.g., "open_ai"),
|
||||
// LlmConfig.backend should resolve to the canonical ID ("openai").
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_openai_compatible_env();
|
||||
// SAFETY: Under ENV_MUTEX.
|
||||
@@ -705,8 +731,6 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn unknown_backend_falls_back_to_openai_compatible() {
|
||||
// An unrecognized LLM_BACKEND should fall back to the openai_compatible
|
||||
// provider definition instead of erroring.
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_openai_compatible_env();
|
||||
// SAFETY: Under ENV_MUTEX.
|
||||
@@ -717,7 +741,6 @@ mod tests {
|
||||
|
||||
let settings = Settings::default();
|
||||
let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed");
|
||||
// Falls back to openai_compatible since "some_custom_provider" is unknown
|
||||
assert_eq!(cfg.backend, "openai_compatible");
|
||||
let provider = cfg.provider.expect("should have provider config");
|
||||
assert_eq!(provider.provider_id, "openai_compatible");
|
||||
@@ -759,7 +782,6 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn base_url_resolution_priority() {
|
||||
// Env var > settings > registry default
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_openai_compatible_env();
|
||||
|
||||
@@ -800,6 +822,119 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
// ── OAuth resolution tests ──────────────────────────────────────
|
||||
|
||||
/// Clear all Anthropic-related env vars.
|
||||
fn clear_anthropic_env() {
|
||||
// SAFETY: Only called under ENV_MUTEX in tests.
|
||||
unsafe {
|
||||
std::env::remove_var("LLM_BACKEND");
|
||||
std::env::remove_var("ANTHROPIC_API_KEY");
|
||||
std::env::remove_var("ANTHROPIC_OAUTH_TOKEN");
|
||||
std::env::remove_var("ANTHROPIC_MODEL");
|
||||
std::env::remove_var("ANTHROPIC_BASE_URL");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn anthropic_oauth_token_sets_placeholder_api_key() {
|
||||
use secrecy::ExposeSecret;
|
||||
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_anthropic_env();
|
||||
// SAFETY: Under ENV_MUTEX.
|
||||
unsafe {
|
||||
std::env::set_var("ANTHROPIC_OAUTH_TOKEN", "sk-ant-oat01-test-token");
|
||||
}
|
||||
|
||||
let settings = Settings {
|
||||
llm_backend: Some("anthropic".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed");
|
||||
let provider = cfg.provider.expect("provider config should be present");
|
||||
|
||||
assert_eq!(
|
||||
provider
|
||||
.api_key
|
||||
.as_ref()
|
||||
.map(|k| k.expose_secret().to_string()),
|
||||
Some(OAUTH_PLACEHOLDER.to_string()),
|
||||
"api_key should be the OAuth placeholder when only OAuth token is set"
|
||||
);
|
||||
assert!(
|
||||
provider.oauth_token.is_some(),
|
||||
"oauth_token should be populated"
|
||||
);
|
||||
assert_eq!(
|
||||
provider.oauth_token.as_ref().unwrap().expose_secret(),
|
||||
"sk-ant-oat01-test-token"
|
||||
);
|
||||
|
||||
clear_anthropic_env();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn anthropic_api_key_takes_priority_over_oauth() {
|
||||
use secrecy::ExposeSecret;
|
||||
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_anthropic_env();
|
||||
// SAFETY: Under ENV_MUTEX.
|
||||
unsafe {
|
||||
std::env::set_var("ANTHROPIC_API_KEY", "sk-ant-real-key");
|
||||
std::env::set_var("ANTHROPIC_OAUTH_TOKEN", "sk-ant-oat01-test-token");
|
||||
}
|
||||
|
||||
let settings = Settings {
|
||||
llm_backend: Some("anthropic".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed");
|
||||
let provider = cfg.provider.expect("provider config should be present");
|
||||
|
||||
assert_eq!(
|
||||
provider
|
||||
.api_key
|
||||
.as_ref()
|
||||
.map(|k| k.expose_secret().to_string()),
|
||||
Some("sk-ant-real-key".to_string()),
|
||||
"real API key should take priority over OAuth placeholder"
|
||||
);
|
||||
assert!(
|
||||
provider.oauth_token.is_some(),
|
||||
"oauth_token should still be populated"
|
||||
);
|
||||
|
||||
clear_anthropic_env();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn non_anthropic_provider_has_no_oauth_token() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
clear_anthropic_env();
|
||||
// SAFETY: Under ENV_MUTEX.
|
||||
unsafe {
|
||||
std::env::set_var("ANTHROPIC_OAUTH_TOKEN", "sk-ant-oat01-test-token");
|
||||
}
|
||||
|
||||
let settings = Settings {
|
||||
llm_backend: Some("openai".to_string()),
|
||||
..Default::default()
|
||||
};
|
||||
let cfg = LlmConfig::resolve(&settings).expect("resolve should succeed");
|
||||
let provider = cfg.provider.expect("provider config should be present");
|
||||
|
||||
assert!(
|
||||
provider.oauth_token.is_none(),
|
||||
"non-Anthropic providers should not pick up ANTHROPIC_OAUTH_TOKEN"
|
||||
);
|
||||
|
||||
clear_anthropic_env();
|
||||
}
|
||||
|
||||
// ── Cache retention tests ───────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn cache_retention_from_str_primary_values() {
|
||||
assert_eq!(
|
||||
|
||||
+73
-5
@@ -13,7 +13,7 @@ mod embeddings;
|
||||
mod heartbeat;
|
||||
pub(crate) mod helpers;
|
||||
mod hygiene;
|
||||
mod llm;
|
||||
pub(crate) mod llm;
|
||||
mod routines;
|
||||
mod safety;
|
||||
mod sandbox;
|
||||
@@ -24,7 +24,7 @@ mod tunnel;
|
||||
mod wasm;
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::OnceLock;
|
||||
use std::sync::{LazyLock, Mutex};
|
||||
|
||||
use crate::error::ConfigError;
|
||||
use crate::settings::Settings;
|
||||
@@ -53,7 +53,12 @@ pub use crate::llm::session::SessionConfig;
|
||||
/// Used by `inject_llm_keys_from_secrets()` to make API keys available to
|
||||
/// `optional_env()` without unsafe `set_var` calls. `optional_env()` checks
|
||||
/// real env vars first, then falls back to this overlay.
|
||||
static INJECTED_VARS: OnceLock<HashMap<String, String>> = OnceLock::new();
|
||||
///
|
||||
/// Uses `Mutex<HashMap>` instead of `OnceLock` so that both
|
||||
/// `inject_os_credentials()` and `inject_llm_keys_from_secrets()` can merge
|
||||
/// their data. Whichever runs first initialises the map; the second merges in.
|
||||
static INJECTED_VARS: LazyLock<Mutex<HashMap<String, String>>> =
|
||||
LazyLock::new(|| Mutex::new(HashMap::new()));
|
||||
|
||||
/// Main configuration for the agent.
|
||||
#[derive(Debug, Clone)]
|
||||
@@ -285,6 +290,9 @@ impl Config {
|
||||
/// env-var-first resolution in `LlmConfig::resolve()`. Keys in the overlay
|
||||
/// are read by `optional_env()` before falling back to `std::env::var()`,
|
||||
/// so explicit env vars always win.
|
||||
///
|
||||
/// Also loads tokens from OS credential stores (macOS Keychain, Linux
|
||||
/// credentials files) which don't require the secrets DB.
|
||||
pub async fn inject_llm_keys_from_secrets(
|
||||
secrets: &dyn crate::secrets::SecretsStore,
|
||||
user_id: &str,
|
||||
@@ -292,7 +300,10 @@ pub async fn inject_llm_keys_from_secrets(
|
||||
// Static mappings for well-known providers.
|
||||
// The registry's setup hints define secret_name -> env_var mappings,
|
||||
// so new providers added to providers.json get injection automatically.
|
||||
let mut mappings: Vec<(&str, &str)> = vec![("llm_nearai_api_key", "NEARAI_API_KEY")];
|
||||
let mut mappings: Vec<(&str, &str)> = vec![
|
||||
("llm_nearai_api_key", "NEARAI_API_KEY"),
|
||||
("llm_anthropic_oauth_token", "ANTHROPIC_OAUTH_TOKEN"),
|
||||
];
|
||||
|
||||
// Dynamically discover secret->env mappings from the provider registry.
|
||||
// Uses selectable() which deduplicates user overrides correctly.
|
||||
@@ -331,5 +342,62 @@ pub async fn inject_llm_keys_from_secrets(
|
||||
}
|
||||
}
|
||||
|
||||
let _ = INJECTED_VARS.set(injected);
|
||||
inject_os_credential_store_tokens(&mut injected);
|
||||
|
||||
merge_injected_vars(injected);
|
||||
}
|
||||
|
||||
/// Load tokens from OS credential stores (no DB required).
|
||||
///
|
||||
/// Called unconditionally during startup — even when the encrypted secrets DB
|
||||
/// is unavailable (no master key, no DB connection). This ensures OAuth tokens
|
||||
/// from `claude login` (macOS Keychain / Linux credentials.json)
|
||||
/// are available for config resolution.
|
||||
pub fn inject_os_credentials() {
|
||||
let mut injected = HashMap::new();
|
||||
inject_os_credential_store_tokens(&mut injected);
|
||||
merge_injected_vars(injected);
|
||||
}
|
||||
|
||||
/// Merge new entries into the global injected-vars overlay.
|
||||
///
|
||||
/// New keys are inserted; existing keys are overwritten (later callers win,
|
||||
/// e.g. fresh OS credential store tokens override stale DB copies).
|
||||
fn merge_injected_vars(new_entries: HashMap<String, String>) {
|
||||
if new_entries.is_empty() {
|
||||
return;
|
||||
}
|
||||
match INJECTED_VARS.lock() {
|
||||
Ok(mut map) => map.extend(new_entries),
|
||||
Err(poisoned) => poisoned.into_inner().extend(new_entries),
|
||||
}
|
||||
}
|
||||
|
||||
/// Inject a single key-value pair into the overlay.
|
||||
///
|
||||
/// Used by the setup wizard to make credentials available to `optional_env()`
|
||||
/// without calling `unsafe { std::env::set_var }`.
|
||||
pub fn inject_single_var(key: &str, value: &str) {
|
||||
match INJECTED_VARS.lock() {
|
||||
Ok(mut map) => {
|
||||
map.insert(key.to_string(), value.to_string());
|
||||
}
|
||||
Err(poisoned) => {
|
||||
poisoned
|
||||
.into_inner()
|
||||
.insert(key.to_string(), value.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Shared helper: extract tokens from OS credential stores into the overlay map.
|
||||
fn inject_os_credential_store_tokens(injected: &mut HashMap<String, String>) {
|
||||
// Try the OS credential store for a fresh Anthropic OAuth token.
|
||||
// Tokens from `claude login` expire in 8-12h, so the DB copy may be stale.
|
||||
// A fresh extraction from macOS Keychain / Linux credentials.json wins
|
||||
// over the (possibly expired) copy stored in the encrypted secrets DB.
|
||||
if let Some(fresh) = crate::config::ClaudeCodeConfig::extract_oauth_token() {
|
||||
injected.insert("ANTHROPIC_OAUTH_TOKEN".to_string(), fresh);
|
||||
tracing::debug!("Refreshed ANTHROPIC_OAUTH_TOKEN from OS credential store");
|
||||
}
|
||||
}
|
||||
|
||||
+16
-5
@@ -233,9 +233,14 @@ impl ClaudeCodeConfig {
|
||||
/// Expected shape: `{"claudeAiOauth": {"accessToken": "sk-ant-oat01-..."}}`
|
||||
fn parse_oauth_access_token(json: &str) -> Option<String> {
|
||||
let creds: serde_json::Value = serde_json::from_str(json).ok()?;
|
||||
creds["claudeAiOauth"]["accessToken"]
|
||||
.as_str()
|
||||
.map(String::from)
|
||||
let token = creds["claudeAiOauth"]["accessToken"].as_str()?;
|
||||
// Validate that the token looks like a real OAuth token before using it.
|
||||
// Claude CLI tokens start with "sk-ant-oat".
|
||||
if !token.starts_with("sk-ant-oat") {
|
||||
tracing::debug!("Ignoring credential store token with unexpected prefix");
|
||||
return None;
|
||||
}
|
||||
Some(token.to_string())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
@@ -401,14 +406,14 @@ mod tests {
|
||||
fn parse_oauth_token_nested_extra_fields() {
|
||||
let json = r#"{
|
||||
"claudeAiOauth": {
|
||||
"accessToken": "sk-ant-real-token",
|
||||
"accessToken": "sk-ant-oat01-real-token",
|
||||
"refreshToken": "rt-abc",
|
||||
"expiresAt": 1700000000
|
||||
}
|
||||
}"#;
|
||||
assert_eq!(
|
||||
parse_oauth_access_token(json),
|
||||
Some("sk-ant-real-token".to_string())
|
||||
Some("sk-ant-oat01-real-token".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
@@ -418,6 +423,12 @@ mod tests {
|
||||
assert_eq!(parse_oauth_access_token(json), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_oauth_token_rejects_invalid_prefix() {
|
||||
let json = r#"{"claudeAiOauth": {"accessToken": "not-an-oauth-token"}}"#;
|
||||
assert_eq!(parse_oauth_access_token(json), None);
|
||||
}
|
||||
|
||||
// ── default_claude_code_allowed_tools ───────────────────────────
|
||||
|
||||
#[test]
|
||||
|
||||
+346
-11
@@ -20,9 +20,10 @@ impl ConversationStore for LibSqlBackend {
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
let conn = self.connect().await?;
|
||||
let id = Uuid::new_v4();
|
||||
let now = fmt_ts(&Utc::now());
|
||||
conn.execute(
|
||||
"INSERT INTO conversations (id, channel, user_id, thread_id) VALUES (?1, ?2, ?3, ?4)",
|
||||
params![id.to_string(), channel, user_id, opt_text(thread_id)],
|
||||
"INSERT INTO conversations (id, channel, user_id, thread_id, started_at, last_activity) VALUES (?1, ?2, ?3, ?4, ?5, ?5)",
|
||||
params![id.to_string(), channel, user_id, opt_text(thread_id), now],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
@@ -71,8 +72,8 @@ impl ConversationStore for LibSqlBackend {
|
||||
let now = fmt_ts(&Utc::now());
|
||||
conn.execute(
|
||||
r#"
|
||||
INSERT INTO conversations (id, channel, user_id, thread_id)
|
||||
VALUES (?1, ?2, ?3, ?4)
|
||||
INSERT INTO conversations (id, channel, user_id, thread_id, started_at, last_activity)
|
||||
VALUES (?1, ?2, ?3, ?4, ?5, ?5)
|
||||
ON CONFLICT (id) DO UPDATE SET last_activity = ?5
|
||||
"#,
|
||||
params![id.to_string(), channel, user_id, opt_text(thread_id), now],
|
||||
@@ -97,6 +98,7 @@ impl ConversationStore for LibSqlBackend {
|
||||
c.started_at,
|
||||
c.last_activity,
|
||||
c.metadata,
|
||||
c.channel,
|
||||
(SELECT COUNT(*) FROM conversation_messages m WHERE m.conversation_id = c.id AND m.role = 'user') AS message_count,
|
||||
(SELECT substr(m2.content, 1, 100)
|
||||
FROM conversation_messages m2
|
||||
@@ -106,7 +108,7 @@ impl ConversationStore for LibSqlBackend {
|
||||
) AS title
|
||||
FROM conversations c
|
||||
WHERE c.user_id = ?1 AND c.channel = ?2
|
||||
ORDER BY c.last_activity DESC
|
||||
ORDER BY datetime(c.last_activity) DESC
|
||||
LIMIT ?3
|
||||
"#,
|
||||
params![user_id, channel, limit],
|
||||
@@ -125,6 +127,13 @@ impl ConversationStore for LibSqlBackend {
|
||||
.get("thread_type")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from);
|
||||
let sql_title = get_opt_text(&row, 6);
|
||||
let title = sql_title.or_else(|| {
|
||||
metadata
|
||||
.get("routine_name")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from)
|
||||
});
|
||||
results.push(ConversationSummary {
|
||||
id: row
|
||||
.get::<String>(0)
|
||||
@@ -133,14 +142,213 @@ impl ConversationStore for LibSqlBackend {
|
||||
.unwrap_or_default(),
|
||||
started_at: get_ts(&row, 1),
|
||||
last_activity: get_ts(&row, 2),
|
||||
message_count: get_i64(&row, 4),
|
||||
title: get_opt_text(&row, 5),
|
||||
message_count: get_i64(&row, 5),
|
||||
title,
|
||||
thread_type,
|
||||
channel: get_text(&row, 4),
|
||||
});
|
||||
}
|
||||
Ok(results)
|
||||
}
|
||||
|
||||
async fn list_conversations_all_channels(
|
||||
&self,
|
||||
user_id: &str,
|
||||
limit: i64,
|
||||
) -> Result<Vec<ConversationSummary>, DatabaseError> {
|
||||
let conn = self.connect().await?;
|
||||
let mut rows = conn
|
||||
.query(
|
||||
r#"
|
||||
SELECT
|
||||
c.id,
|
||||
c.started_at,
|
||||
c.last_activity,
|
||||
c.metadata,
|
||||
c.channel,
|
||||
(SELECT COUNT(*) FROM conversation_messages m WHERE m.conversation_id = c.id AND m.role = 'user') AS message_count,
|
||||
(SELECT substr(m2.content, 1, 100)
|
||||
FROM conversation_messages m2
|
||||
WHERE m2.conversation_id = c.id AND m2.role = 'user'
|
||||
ORDER BY m2.created_at ASC, m2.rowid ASC
|
||||
LIMIT 1
|
||||
) AS title
|
||||
FROM conversations c
|
||||
WHERE c.user_id = ?1
|
||||
ORDER BY datetime(c.last_activity) DESC
|
||||
LIMIT ?2
|
||||
"#,
|
||||
params![user_id, limit],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
|
||||
let mut results = Vec::new();
|
||||
while let Some(row) = rows
|
||||
.next()
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?
|
||||
{
|
||||
let metadata = get_json(&row, 3);
|
||||
let thread_type = metadata
|
||||
.get("thread_type")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from);
|
||||
let sql_title = get_opt_text(&row, 6);
|
||||
let title = sql_title.or_else(|| {
|
||||
metadata
|
||||
.get("routine_name")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from)
|
||||
});
|
||||
results.push(ConversationSummary {
|
||||
id: row
|
||||
.get::<String>(0)
|
||||
.unwrap_or_default()
|
||||
.parse()
|
||||
.unwrap_or_default(),
|
||||
started_at: get_ts(&row, 1),
|
||||
last_activity: get_ts(&row, 2),
|
||||
message_count: get_i64(&row, 5),
|
||||
title,
|
||||
thread_type,
|
||||
channel: get_text(&row, 4),
|
||||
});
|
||||
}
|
||||
Ok(results)
|
||||
}
|
||||
|
||||
/// Uses BEGIN IMMEDIATE to serialize concurrent writers and prevent
|
||||
/// duplicate routine conversations (TOCTOU race).
|
||||
async fn get_or_create_routine_conversation(
|
||||
&self,
|
||||
routine_id: Uuid,
|
||||
routine_name: &str,
|
||||
user_id: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
let conn = self.connect().await?;
|
||||
let rid = routine_id.to_string();
|
||||
|
||||
conn.execute("BEGIN IMMEDIATE", params![])
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
|
||||
let result: Result<Uuid, DatabaseError> = async {
|
||||
let mut rows = conn
|
||||
.query(
|
||||
r#"
|
||||
SELECT id FROM conversations
|
||||
WHERE user_id = ?1 AND json_extract(metadata, '$.routine_id') = ?2
|
||||
LIMIT 1
|
||||
"#,
|
||||
params![user_id, rid],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
|
||||
if let Some(row) = rows
|
||||
.next()
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?
|
||||
{
|
||||
let id_str: String = row.get(0).unwrap_or_default();
|
||||
return id_str
|
||||
.parse()
|
||||
.map_err(|_| DatabaseError::Serialization("Invalid UUID".to_string()));
|
||||
}
|
||||
|
||||
let id = Uuid::new_v4();
|
||||
let now = fmt_ts(&Utc::now());
|
||||
let metadata = serde_json::json!({
|
||||
"thread_type": "routine",
|
||||
"routine_id": routine_id.to_string(),
|
||||
"routine_name": routine_name,
|
||||
});
|
||||
conn.execute(
|
||||
"INSERT INTO conversations (id, channel, user_id, metadata, started_at, last_activity) VALUES (?1, ?2, ?3, ?4, ?5, ?5)",
|
||||
params![id.to_string(), "routine", user_id, metadata.to_string(), now],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
Ok(id)
|
||||
}
|
||||
.await;
|
||||
|
||||
match &result {
|
||||
Ok(_) => {
|
||||
conn.execute("COMMIT", params![])
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
}
|
||||
Err(_) => {
|
||||
let _ = conn.execute("ROLLBACK", params![]).await;
|
||||
}
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
/// Uses BEGIN IMMEDIATE to serialize concurrent writers and prevent
|
||||
/// duplicate heartbeat conversations (TOCTOU race).
|
||||
async fn get_or_create_heartbeat_conversation(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
let conn = self.connect().await?;
|
||||
|
||||
conn.execute("BEGIN IMMEDIATE", params![])
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
|
||||
let result: Result<Uuid, DatabaseError> = async {
|
||||
let mut rows = conn
|
||||
.query(
|
||||
r#"
|
||||
SELECT id FROM conversations
|
||||
WHERE user_id = ?1 AND json_extract(metadata, '$.thread_type') = 'heartbeat'
|
||||
LIMIT 1
|
||||
"#,
|
||||
params![user_id],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
|
||||
if let Some(row) = rows
|
||||
.next()
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?
|
||||
{
|
||||
let id_str: String = row.get(0).unwrap_or_default();
|
||||
return id_str
|
||||
.parse()
|
||||
.map_err(|_| DatabaseError::Serialization("Invalid UUID".to_string()));
|
||||
}
|
||||
|
||||
let id = Uuid::new_v4();
|
||||
let now = fmt_ts(&Utc::now());
|
||||
let metadata = serde_json::json!({ "thread_type": "heartbeat" });
|
||||
conn.execute(
|
||||
"INSERT INTO conversations (id, channel, user_id, metadata, started_at, last_activity) VALUES (?1, ?2, ?3, ?4, ?5, ?5)",
|
||||
params![id.to_string(), "heartbeat", user_id, metadata.to_string(), now],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
Ok(id)
|
||||
}
|
||||
.await;
|
||||
|
||||
match &result {
|
||||
Ok(_) => {
|
||||
conn.execute("COMMIT", params![])
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
}
|
||||
Err(_) => {
|
||||
let _ = conn.execute("ROLLBACK", params![]).await;
|
||||
}
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
async fn get_or_create_assistant_conversation(
|
||||
&self,
|
||||
user_id: &str,
|
||||
@@ -174,10 +382,11 @@ impl ConversationStore for LibSqlBackend {
|
||||
|
||||
// Create new
|
||||
let id = Uuid::new_v4();
|
||||
let now = fmt_ts(&Utc::now());
|
||||
let metadata = serde_json::json!({"thread_type": "assistant", "title": "Assistant"});
|
||||
conn.execute(
|
||||
"INSERT INTO conversations (id, channel, user_id, metadata) VALUES (?1, ?2, ?3, ?4)",
|
||||
params![id.to_string(), channel, user_id, metadata.to_string()],
|
||||
"INSERT INTO conversations (id, channel, user_id, metadata, started_at, last_activity) VALUES (?1, ?2, ?3, ?4, ?5, ?5)",
|
||||
params![id.to_string(), channel, user_id, metadata.to_string(), now],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
@@ -192,9 +401,10 @@ impl ConversationStore for LibSqlBackend {
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
let conn = self.connect().await?;
|
||||
let id = Uuid::new_v4();
|
||||
let now = fmt_ts(&Utc::now());
|
||||
conn.execute(
|
||||
"INSERT INTO conversations (id, channel, user_id, metadata) VALUES (?1, ?2, ?3, ?4)",
|
||||
params![id.to_string(), channel, user_id, metadata.to_string()],
|
||||
"INSERT INTO conversations (id, channel, user_id, metadata, started_at, last_activity) VALUES (?1, ?2, ?3, ?4, ?5, ?5)",
|
||||
params![id.to_string(), channel, user_id, metadata.to_string(), now],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
@@ -353,3 +563,128 @@ impl ConversationStore for LibSqlBackend {
|
||||
Ok(found.is_some())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::db::Database;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_get_or_create_routine_conversation_is_idempotent() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let db_path = dir.path().join("test_routine_conv.db");
|
||||
let backend = LibSqlBackend::new_local(&db_path).await.unwrap();
|
||||
backend.run_migrations().await.unwrap();
|
||||
|
||||
let routine_id = Uuid::new_v4();
|
||||
let user_id = "test_user";
|
||||
|
||||
// First call — creates the conversation
|
||||
let id1 = backend
|
||||
.get_or_create_routine_conversation(routine_id, "my-routine", user_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Second call — should return the SAME conversation
|
||||
let id2 = backend
|
||||
.get_or_create_routine_conversation(routine_id, "my-routine", user_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(id1, id2, "Expected same conversation ID on repeated calls");
|
||||
|
||||
// Third call — still the same
|
||||
let id3 = backend
|
||||
.get_or_create_routine_conversation(routine_id, "my-routine", user_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(id1, id3);
|
||||
|
||||
// Different routine_id should get a different conversation
|
||||
let other_routine_id = Uuid::new_v4();
|
||||
let id4 = backend
|
||||
.get_or_create_routine_conversation(other_routine_id, "other-routine", user_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_ne!(
|
||||
id1, id4,
|
||||
"Different routines should get different conversations"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_routine_conversation_persists_across_messages() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let db_path = dir.path().join("test_routine_persist.db");
|
||||
let backend = LibSqlBackend::new_local(&db_path).await.unwrap();
|
||||
backend.run_migrations().await.unwrap();
|
||||
|
||||
let routine_id = Uuid::new_v4();
|
||||
let user_id = "test_user";
|
||||
|
||||
// First invocation: create conversation and add a message
|
||||
let id1 = backend
|
||||
.get_or_create_routine_conversation(routine_id, "my-routine", user_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
backend
|
||||
.add_conversation_message(id1, "assistant", "[cron] Completed: all good")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Second invocation: should find existing conversation
|
||||
let id2 = backend
|
||||
.get_or_create_routine_conversation(routine_id, "my-routine", user_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(id1, id2, "Second invocation should reuse same conversation");
|
||||
|
||||
backend
|
||||
.add_conversation_message(id2, "assistant", "[cron] Completed: still good")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Verify only one routine conversation exists (not two)
|
||||
let convs = backend
|
||||
.list_conversations_all_channels(user_id, 50)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let routine_convs: Vec<_> = convs.iter().filter(|c| c.channel == "routine").collect();
|
||||
assert_eq!(
|
||||
routine_convs.len(),
|
||||
1,
|
||||
"Should have exactly 1 routine conversation, found {}",
|
||||
routine_convs.len()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_get_or_create_heartbeat_conversation_is_idempotent() {
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let db_path = dir.path().join("test_heartbeat_conv.db");
|
||||
let backend = LibSqlBackend::new_local(&db_path).await.unwrap();
|
||||
backend.run_migrations().await.unwrap();
|
||||
|
||||
let user_id = "test_user";
|
||||
|
||||
let id1 = backend
|
||||
.get_or_create_heartbeat_conversation(user_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let id2 = backend
|
||||
.get_or_create_heartbeat_conversation(user_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
id1, id2,
|
||||
"Expected same heartbeat conversation on repeated calls"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+59
-8
@@ -118,15 +118,37 @@ impl LibSqlBackend {
|
||||
/// Sets `PRAGMA busy_timeout = 5000` on every connection so concurrent
|
||||
/// writers wait up to 5 seconds instead of failing instantly with
|
||||
/// "database is locked".
|
||||
///
|
||||
/// Retries up to 3 times with exponential backoff to handle transient
|
||||
/// "unable to open database file" errors from concurrent connection
|
||||
/// creation (e.g. cron ticker vs main thread).
|
||||
pub async fn connect(&self) -> Result<Connection, DatabaseError> {
|
||||
let conn = self
|
||||
.db
|
||||
.connect()
|
||||
.map_err(|e| DatabaseError::Pool(format!("Failed to create connection: {}", e)))?;
|
||||
conn.query("PRAGMA busy_timeout = 5000", ())
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Pool(format!("Failed to set busy_timeout: {}", e)))?;
|
||||
Ok(conn)
|
||||
let mut last_err = None;
|
||||
for attempt in 0..3u32 {
|
||||
match self.db.connect() {
|
||||
Ok(conn) => {
|
||||
conn.query("PRAGMA busy_timeout = 5000", ())
|
||||
.await
|
||||
.map_err(|e| {
|
||||
DatabaseError::Pool(format!("Failed to set busy_timeout: {}", e))
|
||||
})?;
|
||||
return Ok(conn);
|
||||
}
|
||||
Err(e) => {
|
||||
last_err = Some(e);
|
||||
if attempt < 2 {
|
||||
tokio::time::sleep(std::time::Duration::from_millis(
|
||||
50 * 2u64.pow(attempt),
|
||||
))
|
||||
.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(DatabaseError::Pool(format!(
|
||||
"Failed to create connection after 3 attempts: {}",
|
||||
last_err.map(|e| e.to_string()).unwrap_or_default()
|
||||
)))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -459,4 +481,33 @@ mod tests {
|
||||
let count: i64 = row.get(0).unwrap();
|
||||
assert_eq!(count, 20);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_connect_retry_succeeds_on_valid_db() {
|
||||
// Verify connect() works with retry logic on a file-backed DB
|
||||
// (exercises the retry path even though transient failures are hard
|
||||
// to reproduce deterministically).
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let db_path = dir.path().join("test_retry.db");
|
||||
let backend = LibSqlBackend::new_local(&db_path).await.unwrap();
|
||||
backend.run_migrations().await.unwrap();
|
||||
|
||||
// Multiple concurrent connect() calls should all succeed
|
||||
let mut handles = Vec::new();
|
||||
for _ in 0..10 {
|
||||
let b = LibSqlBackend {
|
||||
db: backend.shared_db(),
|
||||
};
|
||||
handles.push(tokio::spawn(async move { b.connect().await }));
|
||||
}
|
||||
|
||||
for handle in handles {
|
||||
let result = handle.await.unwrap();
|
||||
assert!(
|
||||
result.is_ok(),
|
||||
"concurrent connect failed: {:?}",
|
||||
result.err()
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -423,4 +423,29 @@ impl RoutineStore for LibSqlBackend {
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn list_dispatched_routine_runs(&self) -> Result<Vec<RoutineRun>, DatabaseError> {
|
||||
let conn = self.connect().await?;
|
||||
let mut rows = conn
|
||||
.query(
|
||||
&format!(
|
||||
"SELECT {} FROM routine_runs \
|
||||
WHERE status = 'running' AND job_id IS NOT NULL",
|
||||
ROUTINE_RUN_COLUMNS
|
||||
),
|
||||
params![],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?;
|
||||
|
||||
let mut runs = Vec::new();
|
||||
while let Some(row) = rows
|
||||
.next()
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Query(e.to_string()))?
|
||||
{
|
||||
runs.push(row_to_routine_run_libsql(&row)?);
|
||||
}
|
||||
Ok(runs)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -45,6 +45,15 @@ CREATE INDEX IF NOT EXISTS idx_conversations_channel ON conversations(channel);
|
||||
CREATE INDEX IF NOT EXISTS idx_conversations_user ON conversations(user_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_conversations_last_activity ON conversations(last_activity);
|
||||
|
||||
-- Partial unique indexes to prevent duplicate singleton conversations.
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS uq_conv_routine
|
||||
ON conversations (user_id, json_extract(metadata, '$.routine_id'))
|
||||
WHERE json_extract(metadata, '$.routine_id') IS NOT NULL;
|
||||
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS uq_conv_heartbeat
|
||||
ON conversations (user_id)
|
||||
WHERE json_extract(metadata, '$.thread_type') = 'heartbeat';
|
||||
|
||||
CREATE TABLE IF NOT EXISTS conversation_messages (
|
||||
id TEXT PRIMARY KEY,
|
||||
conversation_id TEXT NOT NULL REFERENCES conversations(id) ON DELETE CASCADE,
|
||||
|
||||
@@ -125,6 +125,21 @@ pub trait ConversationStore: Send + Sync {
|
||||
channel: &str,
|
||||
limit: i64,
|
||||
) -> Result<Vec<ConversationSummary>, DatabaseError>;
|
||||
async fn list_conversations_all_channels(
|
||||
&self,
|
||||
user_id: &str,
|
||||
limit: i64,
|
||||
) -> Result<Vec<ConversationSummary>, DatabaseError>;
|
||||
async fn get_or_create_routine_conversation(
|
||||
&self,
|
||||
routine_id: Uuid,
|
||||
routine_name: &str,
|
||||
user_id: &str,
|
||||
) -> Result<Uuid, DatabaseError>;
|
||||
async fn get_or_create_heartbeat_conversation(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> Result<Uuid, DatabaseError>;
|
||||
async fn get_or_create_assistant_conversation(
|
||||
&self,
|
||||
user_id: &str,
|
||||
@@ -288,6 +303,10 @@ pub trait RoutineStore: Send + Sync {
|
||||
run_id: Uuid,
|
||||
job_id: Uuid,
|
||||
) -> Result<(), DatabaseError>;
|
||||
/// List routine runs that were dispatched as full_job (status = 'running'
|
||||
/// with a linked job_id). Used by the routine engine to sync completion
|
||||
/// status from the background job.
|
||||
async fn list_dispatched_routine_runs(&self) -> Result<Vec<RoutineRun>, DatabaseError>;
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
|
||||
@@ -116,6 +116,36 @@ impl ConversationStore for PgBackend {
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_conversations_all_channels(
|
||||
&self,
|
||||
user_id: &str,
|
||||
limit: i64,
|
||||
) -> Result<Vec<ConversationSummary>, DatabaseError> {
|
||||
self.store
|
||||
.list_conversations_all_channels(user_id, limit)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_or_create_routine_conversation(
|
||||
&self,
|
||||
routine_id: Uuid,
|
||||
routine_name: &str,
|
||||
user_id: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.store
|
||||
.get_or_create_routine_conversation(routine_id, routine_name, user_id)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_or_create_heartbeat_conversation(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.store
|
||||
.get_or_create_heartbeat_conversation(user_id)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_or_create_assistant_conversation(
|
||||
&self,
|
||||
user_id: &str,
|
||||
@@ -464,6 +494,10 @@ impl RoutineStore for PgBackend {
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.store.link_routine_run_to_job(run_id, job_id).await
|
||||
}
|
||||
|
||||
async fn list_dispatched_routine_runs(&self) -> Result<Vec<RoutineRun>, DatabaseError> {
|
||||
self.store.list_dispatched_routine_runs().await
|
||||
}
|
||||
}
|
||||
|
||||
// ==================== ToolFailureStore ====================
|
||||
|
||||
@@ -401,6 +401,9 @@ pub enum RoutineError {
|
||||
#[error("Routine not found: {id}")]
|
||||
NotFound { id: Uuid },
|
||||
|
||||
#[error("Not authorized to trigger routine {id}")]
|
||||
NotAuthorized { id: Uuid },
|
||||
|
||||
#[error("Routine {name} at max concurrent runs")]
|
||||
MaxConcurrent { name: String },
|
||||
|
||||
|
||||
+215
-1
@@ -1295,6 +1295,17 @@ impl Store {
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn list_dispatched_routine_runs(&self) -> Result<Vec<RoutineRun>, DatabaseError> {
|
||||
let conn = self.conn().await?;
|
||||
let rows = conn
|
||||
.query(
|
||||
"SELECT * FROM routine_runs WHERE status = 'running' AND job_id IS NOT NULL",
|
||||
&[],
|
||||
)
|
||||
.await?;
|
||||
rows.iter().map(row_to_routine_run).collect()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "postgres")]
|
||||
@@ -1377,6 +1388,8 @@ pub struct ConversationSummary {
|
||||
pub last_activity: DateTime<Utc>,
|
||||
/// Thread type extracted from metadata (e.g. "assistant", "thread").
|
||||
pub thread_type: Option<String>,
|
||||
/// Channel that owns this conversation (e.g. "gateway", "telegram", "routine").
|
||||
pub channel: String,
|
||||
}
|
||||
|
||||
/// A single message in a conversation.
|
||||
@@ -1429,6 +1442,7 @@ impl Store {
|
||||
c.started_at,
|
||||
c.last_activity,
|
||||
c.metadata,
|
||||
c.channel,
|
||||
(SELECT COUNT(*) FROM conversation_messages m WHERE m.conversation_id = c.id AND m.role = 'user') AS message_count,
|
||||
(SELECT LEFT(m2.content, 100)
|
||||
FROM conversation_messages m2
|
||||
@@ -1453,18 +1467,181 @@ impl Store {
|
||||
.get("thread_type")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from);
|
||||
let sql_title: Option<String> = r.get("title");
|
||||
let title = sql_title.or_else(|| {
|
||||
metadata
|
||||
.get("routine_name")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from)
|
||||
});
|
||||
ConversationSummary {
|
||||
id: r.get("id"),
|
||||
title: r.get("title"),
|
||||
title,
|
||||
message_count: r.get("message_count"),
|
||||
started_at: r.get("started_at"),
|
||||
last_activity: r.get("last_activity"),
|
||||
thread_type,
|
||||
channel: r.get("channel"),
|
||||
}
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
/// List conversations across all channels with a title derived from the first user message.
|
||||
pub async fn list_conversations_all_channels(
|
||||
&self,
|
||||
user_id: &str,
|
||||
limit: i64,
|
||||
) -> Result<Vec<ConversationSummary>, DatabaseError> {
|
||||
let conn = self.conn().await?;
|
||||
let rows = conn
|
||||
.query(
|
||||
r#"
|
||||
SELECT
|
||||
c.id,
|
||||
c.started_at,
|
||||
c.last_activity,
|
||||
c.metadata,
|
||||
c.channel,
|
||||
(SELECT COUNT(*) FROM conversation_messages m WHERE m.conversation_id = c.id AND m.role = 'user') AS message_count,
|
||||
(SELECT LEFT(m2.content, 100)
|
||||
FROM conversation_messages m2
|
||||
WHERE m2.conversation_id = c.id AND m2.role = 'user'
|
||||
ORDER BY m2.created_at ASC
|
||||
LIMIT 1
|
||||
) AS title
|
||||
FROM conversations c
|
||||
WHERE c.user_id = $1
|
||||
ORDER BY c.last_activity DESC
|
||||
LIMIT $2
|
||||
"#,
|
||||
&[&user_id, &limit],
|
||||
)
|
||||
.await?;
|
||||
|
||||
Ok(rows
|
||||
.iter()
|
||||
.map(|r| {
|
||||
let metadata: serde_json::Value = r.get("metadata");
|
||||
let thread_type = metadata
|
||||
.get("thread_type")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from);
|
||||
// For routine/heartbeat threads, derive title from metadata
|
||||
// since they may have no user messages.
|
||||
let sql_title: Option<String> = r.get("title");
|
||||
let title = sql_title.or_else(|| {
|
||||
metadata
|
||||
.get("routine_name")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(String::from)
|
||||
});
|
||||
ConversationSummary {
|
||||
id: r.get("id"),
|
||||
title,
|
||||
message_count: r.get("message_count"),
|
||||
started_at: r.get("started_at"),
|
||||
last_activity: r.get("last_activity"),
|
||||
thread_type,
|
||||
channel: r.get("channel"),
|
||||
}
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
/// Get or create a persistent conversation for a routine.
|
||||
///
|
||||
/// Looks for a conversation where `metadata->>'routine_id' = routine_id`.
|
||||
/// Creates one if it doesn't exist. Uses INSERT ON CONFLICT to avoid
|
||||
/// TOCTOU races under concurrent routine executions.
|
||||
pub async fn get_or_create_routine_conversation(
|
||||
&self,
|
||||
routine_id: Uuid,
|
||||
routine_name: &str,
|
||||
user_id: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
let conn = self.conn().await?;
|
||||
let rid = routine_id.to_string();
|
||||
|
||||
// Attempt insert first; the partial unique index
|
||||
// uq_conv_routine(user_id, (metadata->>'routine_id')) prevents duplicates.
|
||||
let new_id = Uuid::new_v4();
|
||||
let metadata = serde_json::json!({
|
||||
"thread_type": "routine",
|
||||
"routine_id": routine_id.to_string(),
|
||||
"routine_name": routine_name,
|
||||
});
|
||||
conn.execute(
|
||||
r#"
|
||||
INSERT INTO conversations (id, channel, user_id, metadata)
|
||||
VALUES ($1, 'routine', $2, $3)
|
||||
ON CONFLICT (user_id, (metadata->>'routine_id'))
|
||||
WHERE metadata->>'routine_id' IS NOT NULL
|
||||
DO NOTHING
|
||||
"#,
|
||||
&[&new_id, &user_id, &metadata],
|
||||
)
|
||||
.await?;
|
||||
|
||||
// Select back — always returns the winner.
|
||||
let row = conn
|
||||
.query_one(
|
||||
r#"
|
||||
SELECT id FROM conversations
|
||||
WHERE user_id = $1 AND metadata->>'routine_id' = $2
|
||||
LIMIT 1
|
||||
"#,
|
||||
&[&user_id, &rid],
|
||||
)
|
||||
.await?;
|
||||
|
||||
Ok(row.get("id"))
|
||||
}
|
||||
|
||||
/// Get or create the singleton heartbeat conversation for a user.
|
||||
///
|
||||
/// Looks for a conversation where `metadata->>'thread_type' = 'heartbeat'`.
|
||||
/// Creates one if it doesn't exist. Uses INSERT ON CONFLICT to avoid
|
||||
/// TOCTOU races under concurrent heartbeat sends.
|
||||
pub async fn get_or_create_heartbeat_conversation(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
let conn = self.conn().await?;
|
||||
|
||||
// Attempt insert; the partial unique index
|
||||
// uq_conv_heartbeat(user_id) prevents duplicates.
|
||||
let new_id = Uuid::new_v4();
|
||||
let metadata = serde_json::json!({
|
||||
"thread_type": "heartbeat",
|
||||
});
|
||||
conn.execute(
|
||||
r#"
|
||||
INSERT INTO conversations (id, channel, user_id, metadata)
|
||||
VALUES ($1, 'heartbeat', $2, $3)
|
||||
ON CONFLICT (user_id)
|
||||
WHERE metadata->>'thread_type' = 'heartbeat'
|
||||
DO NOTHING
|
||||
"#,
|
||||
&[&new_id, &user_id, &metadata],
|
||||
)
|
||||
.await?;
|
||||
|
||||
// Select back — always returns the winner.
|
||||
let row = conn
|
||||
.query_one(
|
||||
r#"
|
||||
SELECT id FROM conversations
|
||||
WHERE user_id = $1 AND metadata->>'thread_type' = 'heartbeat'
|
||||
LIMIT 1
|
||||
"#,
|
||||
&[&user_id],
|
||||
)
|
||||
.await?;
|
||||
|
||||
Ok(row.get("id"))
|
||||
}
|
||||
|
||||
/// Get or create the singleton "assistant" conversation for a user+channel.
|
||||
///
|
||||
/// Looks for a conversation where `metadata->>'thread_type' = 'assistant'`.
|
||||
@@ -1928,3 +2105,40 @@ impl Store {
|
||||
Ok(count > 0)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_conversation_summary_has_channel_field() {
|
||||
// Regression: ConversationSummary must include a `channel` field
|
||||
// so the gateway can distinguish thread origins.
|
||||
let summary = ConversationSummary {
|
||||
id: Uuid::nil(),
|
||||
title: Some("Hello".to_string()),
|
||||
message_count: 1,
|
||||
started_at: Utc::now(),
|
||||
last_activity: Utc::now(),
|
||||
thread_type: Some("thread".to_string()),
|
||||
channel: "telegram".to_string(),
|
||||
};
|
||||
assert_eq!(summary.channel, "telegram");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_conversation_summary_channel_various_values() {
|
||||
for ch in ["gateway", "routine", "heartbeat", "telegram", "signal"] {
|
||||
let summary = ConversationSummary {
|
||||
id: Uuid::nil(),
|
||||
title: None,
|
||||
message_count: 0,
|
||||
started_at: Utc::now(),
|
||||
last_activity: Utc::now(),
|
||||
thread_type: None,
|
||||
channel: ch.to_string(),
|
||||
};
|
||||
assert_eq!(summary.channel, ch);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,641 @@
|
||||
//! Anthropic OAuth provider (direct HTTP, `Authorization: Bearer`).
|
||||
//!
|
||||
//! This provider exists because the `rig-core` Anthropic client hardcodes the
|
||||
//! `x-api-key` header, which is rejected by Anthropic's OAuth tokens from
|
||||
//! `claude login`. OAuth tokens require `Authorization: Bearer <token>` instead.
|
||||
//!
|
||||
//! Pattern follows `nearai_chat.rs`: direct HTTP calls via `reqwest::Client`.
|
||||
|
||||
use async_trait::async_trait;
|
||||
use reqwest::Client;
|
||||
use rust_decimal::Decimal;
|
||||
use secrecy::{ExposeSecret, SecretString};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::config::RegistryProviderConfig;
|
||||
use crate::error::LlmError;
|
||||
use crate::llm::costs;
|
||||
use crate::llm::provider::{
|
||||
ChatMessage, CompletionRequest, CompletionResponse, FinishReason, LlmProvider, Role, ToolCall,
|
||||
ToolCompletionRequest, ToolCompletionResponse,
|
||||
};
|
||||
|
||||
const ANTHROPIC_API_URL: &str = "https://api.anthropic.com/v1/messages";
|
||||
/// OAuth beta requires 2023-06-01; the 2024-10-22 version is not valid with the beta flag.
|
||||
const ANTHROPIC_API_VERSION: &str = "2023-06-01";
|
||||
/// Required beta flag to enable OAuth Bearer auth on api.anthropic.com.
|
||||
/// Without this header, the API returns 401 "OAuth authentication is currently not supported."
|
||||
const ANTHROPIC_OAUTH_BETA: &str = "oauth-2025-04-20";
|
||||
const DEFAULT_MAX_TOKENS: u32 = 8192;
|
||||
|
||||
/// Anthropic provider using OAuth Bearer authentication.
|
||||
pub struct AnthropicOAuthProvider {
|
||||
client: Client,
|
||||
token: SecretString,
|
||||
model: String,
|
||||
base_url: Option<String>,
|
||||
active_model: std::sync::RwLock<String>,
|
||||
}
|
||||
|
||||
impl AnthropicOAuthProvider {
|
||||
pub fn new(config: &RegistryProviderConfig) -> Result<Self, LlmError> {
|
||||
let token = config
|
||||
.oauth_token
|
||||
.clone()
|
||||
.ok_or_else(|| LlmError::AuthFailed {
|
||||
provider: "anthropic_oauth".to_string(),
|
||||
})?;
|
||||
|
||||
let client = Client::builder()
|
||||
.timeout(std::time::Duration::from_secs(120))
|
||||
.build()
|
||||
.map_err(|e| LlmError::RequestFailed {
|
||||
provider: "anthropic_oauth".to_string(),
|
||||
reason: format!("Failed to build HTTP client: {}", e),
|
||||
})?;
|
||||
|
||||
let active_model = std::sync::RwLock::new(config.model.clone());
|
||||
let base_url = if config.base_url.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(config.base_url.clone())
|
||||
};
|
||||
|
||||
Ok(Self {
|
||||
client,
|
||||
token,
|
||||
model: config.model.clone(),
|
||||
base_url,
|
||||
active_model,
|
||||
})
|
||||
}
|
||||
|
||||
fn api_url(&self) -> String {
|
||||
if let Some(ref base) = self.base_url {
|
||||
let base = base.trim_end_matches('/');
|
||||
format!("{}/v1/messages", base)
|
||||
} else {
|
||||
ANTHROPIC_API_URL.to_string()
|
||||
}
|
||||
}
|
||||
|
||||
async fn send_request<R: for<'de> Deserialize<'de>>(
|
||||
&self,
|
||||
body: &AnthropicRequest,
|
||||
) -> Result<R, LlmError> {
|
||||
let url = self.api_url();
|
||||
|
||||
tracing::debug!("Sending request to Anthropic OAuth: {}", url);
|
||||
|
||||
let response = self
|
||||
.client
|
||||
.post(&url)
|
||||
.bearer_auth(self.token.expose_secret())
|
||||
.header("anthropic-version", ANTHROPIC_API_VERSION)
|
||||
.header("anthropic-beta", ANTHROPIC_OAUTH_BETA)
|
||||
.header("Content-Type", "application/json")
|
||||
.json(body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| LlmError::RequestFailed {
|
||||
provider: "anthropic_oauth".to_string(),
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
|
||||
let status = response.status();
|
||||
|
||||
if !status.is_success() {
|
||||
// Parse Retry-After header before consuming the body.
|
||||
let retry_after = response
|
||||
.headers()
|
||||
.get("retry-after")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|v| v.parse::<u64>().ok())
|
||||
.map(std::time::Duration::from_secs);
|
||||
|
||||
let response_text = response
|
||||
.text()
|
||||
.await
|
||||
.unwrap_or_else(|e| format!("(failed to read error body: {e})"));
|
||||
|
||||
if status.as_u16() == 401 {
|
||||
// OAuth tokens from `claude login` expire in ~8-12h. Attempt
|
||||
// to re-extract a fresh token from the OS credential store
|
||||
// (macOS Keychain / Linux credentials file) before giving up.
|
||||
if let Some(fresh) = crate::config::ClaudeCodeConfig::extract_oauth_token() {
|
||||
let fresh_token = SecretString::from(fresh);
|
||||
// Retry once with the refreshed token
|
||||
let retry = self
|
||||
.client
|
||||
.post(&url)
|
||||
.bearer_auth(fresh_token.expose_secret())
|
||||
.header("anthropic-version", ANTHROPIC_API_VERSION)
|
||||
.header("anthropic-beta", ANTHROPIC_OAUTH_BETA)
|
||||
.header("Content-Type", "application/json")
|
||||
.json(body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| LlmError::RequestFailed {
|
||||
provider: "anthropic_oauth".to_string(),
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
if retry.status().is_success() {
|
||||
let text = retry.text().await.map_err(|e| LlmError::RequestFailed {
|
||||
provider: "anthropic_oauth".to_string(),
|
||||
reason: format!("Failed to read response body: {}", e),
|
||||
})?;
|
||||
return serde_json::from_str(&text).map_err(|e| {
|
||||
let truncated = crate::agent::truncate_for_preview(&text, 512);
|
||||
LlmError::InvalidResponse {
|
||||
provider: "anthropic_oauth".to_string(),
|
||||
reason: format!("JSON parse error: {}. Raw: {}", e, truncated),
|
||||
}
|
||||
});
|
||||
}
|
||||
tracing::warn!(
|
||||
"Anthropic OAuth 401 retry with refreshed token also failed ({})",
|
||||
retry.status()
|
||||
);
|
||||
}
|
||||
return Err(LlmError::AuthFailed {
|
||||
provider: "anthropic_oauth".to_string(),
|
||||
});
|
||||
}
|
||||
if status.as_u16() == 429 {
|
||||
return Err(LlmError::RateLimited {
|
||||
provider: "anthropic_oauth".to_string(),
|
||||
retry_after,
|
||||
});
|
||||
}
|
||||
let truncated = crate::agent::truncate_for_preview(&response_text, 512);
|
||||
return Err(LlmError::RequestFailed {
|
||||
provider: "anthropic_oauth".to_string(),
|
||||
reason: format!("HTTP {}: {}", status, truncated),
|
||||
});
|
||||
}
|
||||
|
||||
let response_text = response.text().await.map_err(|e| LlmError::RequestFailed {
|
||||
provider: "anthropic_oauth".to_string(),
|
||||
reason: format!("Failed to read response body: {}", e),
|
||||
})?;
|
||||
|
||||
tracing::debug!(
|
||||
"Anthropic OAuth response: status={}, bytes={}",
|
||||
status,
|
||||
response_text.len()
|
||||
);
|
||||
|
||||
serde_json::from_str(&response_text).map_err(|e| {
|
||||
let truncated = crate::agent::truncate_for_preview(&response_text, 512);
|
||||
LlmError::InvalidResponse {
|
||||
provider: "anthropic_oauth".to_string(),
|
||||
reason: format!("JSON parse error: {}. Raw: {}", e, truncated),
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for AnthropicOAuthProvider {
|
||||
async fn complete(&self, req: CompletionRequest) -> Result<CompletionResponse, LlmError> {
|
||||
let model = req.model.unwrap_or_else(|| self.active_model_name());
|
||||
let (system, messages) = convert_messages(req.messages);
|
||||
|
||||
let request = AnthropicRequest {
|
||||
model,
|
||||
messages,
|
||||
system,
|
||||
max_tokens: req.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS),
|
||||
temperature: req.temperature,
|
||||
tools: None,
|
||||
tool_choice: None,
|
||||
};
|
||||
|
||||
let response: AnthropicResponse = self.send_request(&request).await?;
|
||||
let (content, _tool_calls) = extract_response_content(&response);
|
||||
|
||||
let finish_reason = match response.stop_reason.as_deref() {
|
||||
Some("end_turn") | Some("stop") => FinishReason::Stop,
|
||||
Some("max_tokens") => FinishReason::Length,
|
||||
Some("tool_use") => FinishReason::ToolUse,
|
||||
_ => FinishReason::Unknown,
|
||||
};
|
||||
|
||||
Ok(CompletionResponse {
|
||||
content: content.unwrap_or_default(),
|
||||
finish_reason,
|
||||
input_tokens: response.usage.input_tokens,
|
||||
output_tokens: response.usage.output_tokens,
|
||||
cache_creation_input_tokens: response.usage.cache_creation_input_tokens,
|
||||
cache_read_input_tokens: response.usage.cache_read_input_tokens,
|
||||
})
|
||||
}
|
||||
|
||||
async fn complete_with_tools(
|
||||
&self,
|
||||
req: ToolCompletionRequest,
|
||||
) -> Result<ToolCompletionResponse, LlmError> {
|
||||
let model = req.model.unwrap_or_else(|| self.active_model_name());
|
||||
let (system, messages) = convert_messages(req.messages);
|
||||
|
||||
let tools: Vec<AnthropicTool> = req
|
||||
.tools
|
||||
.into_iter()
|
||||
.map(|t| AnthropicTool {
|
||||
name: t.name,
|
||||
description: t.description,
|
||||
input_schema: t.parameters,
|
||||
})
|
||||
.collect();
|
||||
|
||||
// Map tool_choice from OpenAI format to Anthropic format
|
||||
let tool_choice = req.tool_choice.map(|tc| match tc.as_str() {
|
||||
"auto" => AnthropicToolChoice {
|
||||
choice_type: "auto".to_string(),
|
||||
name: None,
|
||||
},
|
||||
"required" => AnthropicToolChoice {
|
||||
choice_type: "any".to_string(),
|
||||
name: None,
|
||||
},
|
||||
"none" => AnthropicToolChoice {
|
||||
choice_type: "none".to_string(),
|
||||
name: None,
|
||||
},
|
||||
specific => AnthropicToolChoice {
|
||||
choice_type: "tool".to_string(),
|
||||
name: Some(specific.to_string()),
|
||||
},
|
||||
});
|
||||
|
||||
let request = AnthropicRequest {
|
||||
model,
|
||||
messages,
|
||||
system,
|
||||
max_tokens: req.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS),
|
||||
temperature: req.temperature,
|
||||
tools: if tools.is_empty() { None } else { Some(tools) },
|
||||
tool_choice,
|
||||
};
|
||||
|
||||
let response: AnthropicResponse = self.send_request(&request).await?;
|
||||
let (content, tool_calls) = extract_response_content(&response);
|
||||
|
||||
let finish_reason = match response.stop_reason.as_deref() {
|
||||
Some("end_turn") | Some("stop") => FinishReason::Stop,
|
||||
Some("max_tokens") => FinishReason::Length,
|
||||
Some("tool_use") => FinishReason::ToolUse,
|
||||
_ => {
|
||||
if !tool_calls.is_empty() {
|
||||
FinishReason::ToolUse
|
||||
} else {
|
||||
FinishReason::Unknown
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
Ok(ToolCompletionResponse {
|
||||
content,
|
||||
tool_calls,
|
||||
finish_reason,
|
||||
input_tokens: response.usage.input_tokens,
|
||||
output_tokens: response.usage.output_tokens,
|
||||
cache_creation_input_tokens: response.usage.cache_creation_input_tokens,
|
||||
cache_read_input_tokens: response.usage.cache_read_input_tokens,
|
||||
})
|
||||
}
|
||||
|
||||
fn model_name(&self) -> &str {
|
||||
&self.model
|
||||
}
|
||||
|
||||
fn cost_per_token(&self) -> (Decimal, Decimal) {
|
||||
let model = self.active_model_name();
|
||||
costs::model_cost(&model).unwrap_or_else(costs::default_cost)
|
||||
}
|
||||
|
||||
fn active_model_name(&self) -> String {
|
||||
match self.active_model.read() {
|
||||
Ok(guard) => guard.clone(),
|
||||
Err(poisoned) => poisoned.into_inner().clone(),
|
||||
}
|
||||
}
|
||||
|
||||
fn set_model(&self, model: &str) -> Result<(), LlmError> {
|
||||
match self.active_model.write() {
|
||||
Ok(mut guard) => {
|
||||
*guard = model.to_string();
|
||||
}
|
||||
Err(poisoned) => {
|
||||
*poisoned.into_inner() = model.to_string();
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
// --- Anthropic Messages API types ---
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct AnthropicRequest {
|
||||
model: String,
|
||||
messages: Vec<AnthropicMessage>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
system: Option<String>,
|
||||
max_tokens: u32,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
temperature: Option<f32>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
tools: Option<Vec<AnthropicTool>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
tool_choice: Option<AnthropicToolChoice>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct AnthropicMessage {
|
||||
role: String,
|
||||
content: AnthropicContent,
|
||||
}
|
||||
|
||||
/// Anthropic content can be a simple string or a list of content blocks.
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(untagged)]
|
||||
enum AnthropicContent {
|
||||
Text(String),
|
||||
Blocks(Vec<AnthropicContentBlock>),
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
#[serde(tag = "type")]
|
||||
enum AnthropicContentBlock {
|
||||
#[serde(rename = "text")]
|
||||
Text { text: String },
|
||||
#[serde(rename = "tool_use")]
|
||||
ToolUse {
|
||||
id: String,
|
||||
name: String,
|
||||
input: serde_json::Value,
|
||||
},
|
||||
#[serde(rename = "tool_result")]
|
||||
ToolResult {
|
||||
tool_use_id: String,
|
||||
content: String,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct AnthropicTool {
|
||||
name: String,
|
||||
description: String,
|
||||
input_schema: serde_json::Value,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
struct AnthropicToolChoice {
|
||||
#[serde(rename = "type")]
|
||||
choice_type: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
name: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct AnthropicResponse {
|
||||
content: Vec<AnthropicResponseBlock>,
|
||||
#[serde(default)]
|
||||
stop_reason: Option<String>,
|
||||
usage: AnthropicUsage,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(tag = "type")]
|
||||
enum AnthropicResponseBlock {
|
||||
#[serde(rename = "text")]
|
||||
Text { text: String },
|
||||
#[serde(rename = "tool_use")]
|
||||
ToolUse {
|
||||
id: String,
|
||||
name: String,
|
||||
input: serde_json::Value,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct AnthropicUsage {
|
||||
#[serde(default)]
|
||||
input_tokens: u32,
|
||||
#[serde(default)]
|
||||
output_tokens: u32,
|
||||
#[serde(default)]
|
||||
cache_creation_input_tokens: u32,
|
||||
#[serde(default)]
|
||||
cache_read_input_tokens: u32,
|
||||
}
|
||||
|
||||
/// Convert ChatMessage list to Anthropic format.
|
||||
///
|
||||
/// Extracts system messages to the top-level `system` parameter (Anthropic
|
||||
/// doesn't allow system messages in the `messages` array). Tool-call/tool-result
|
||||
/// pairs are converted to content blocks.
|
||||
fn convert_messages(messages: Vec<ChatMessage>) -> (Option<String>, Vec<AnthropicMessage>) {
|
||||
let mut system_parts: Vec<String> = Vec::new();
|
||||
let mut anthropic_msgs: Vec<AnthropicMessage> = Vec::new();
|
||||
|
||||
for msg in messages {
|
||||
match msg.role {
|
||||
Role::System => {
|
||||
if !msg.content.is_empty() {
|
||||
system_parts.push(msg.content);
|
||||
}
|
||||
}
|
||||
Role::User => {
|
||||
anthropic_msgs.push(AnthropicMessage {
|
||||
role: "user".to_string(),
|
||||
content: AnthropicContent::Text(msg.content),
|
||||
});
|
||||
}
|
||||
Role::Assistant => {
|
||||
if let Some(tool_calls) = msg.tool_calls {
|
||||
// Assistant message with tool calls → content blocks
|
||||
let mut blocks: Vec<AnthropicContentBlock> = Vec::new();
|
||||
if !msg.content.is_empty() {
|
||||
blocks.push(AnthropicContentBlock::Text { text: msg.content });
|
||||
}
|
||||
for tc in tool_calls {
|
||||
blocks.push(AnthropicContentBlock::ToolUse {
|
||||
id: tc.id,
|
||||
name: tc.name,
|
||||
input: tc.arguments,
|
||||
});
|
||||
}
|
||||
anthropic_msgs.push(AnthropicMessage {
|
||||
role: "assistant".to_string(),
|
||||
content: AnthropicContent::Blocks(blocks),
|
||||
});
|
||||
} else {
|
||||
anthropic_msgs.push(AnthropicMessage {
|
||||
role: "assistant".to_string(),
|
||||
content: AnthropicContent::Text(msg.content),
|
||||
});
|
||||
}
|
||||
}
|
||||
Role::Tool => {
|
||||
let Some(tool_call_id) = msg.tool_call_id else {
|
||||
tracing::warn!("Skipping Tool message without tool_call_id");
|
||||
continue;
|
||||
};
|
||||
// Tool results go into a user message with tool_result blocks
|
||||
let block = AnthropicContentBlock::ToolResult {
|
||||
tool_use_id: tool_call_id,
|
||||
content: msg.content,
|
||||
};
|
||||
// If the last message is already a user message with blocks,
|
||||
// append to it (Anthropic requires consecutive tool results
|
||||
// in one user message).
|
||||
if let Some(last) = anthropic_msgs.last_mut()
|
||||
&& last.role == "user"
|
||||
&& let AnthropicContent::Blocks(ref mut blocks) = last.content
|
||||
{
|
||||
blocks.push(block);
|
||||
continue;
|
||||
}
|
||||
anthropic_msgs.push(AnthropicMessage {
|
||||
role: "user".to_string(),
|
||||
content: AnthropicContent::Blocks(vec![block]),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let system = if system_parts.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(system_parts.join("\n\n"))
|
||||
};
|
||||
|
||||
(system, anthropic_msgs)
|
||||
}
|
||||
|
||||
/// Extract text content and tool calls from an Anthropic response.
|
||||
fn extract_response_content(response: &AnthropicResponse) -> (Option<String>, Vec<ToolCall>) {
|
||||
let mut text_parts: Vec<String> = Vec::new();
|
||||
let mut tool_calls: Vec<ToolCall> = Vec::new();
|
||||
|
||||
for block in &response.content {
|
||||
match block {
|
||||
AnthropicResponseBlock::Text { text } => {
|
||||
text_parts.push(text.clone());
|
||||
}
|
||||
AnthropicResponseBlock::ToolUse { id, name, input } => {
|
||||
tool_calls.push(ToolCall {
|
||||
id: id.clone(),
|
||||
name: name.clone(),
|
||||
arguments: input.clone(),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let content = if text_parts.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(text_parts.join(""))
|
||||
};
|
||||
|
||||
(content, tool_calls)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_convert_messages_extracts_system() {
|
||||
let messages = vec![
|
||||
ChatMessage::system("You are helpful."),
|
||||
ChatMessage::user("Hello"),
|
||||
];
|
||||
let (system, msgs) = convert_messages(messages);
|
||||
assert_eq!(system, Some("You are helpful.".to_string()));
|
||||
assert_eq!(msgs.len(), 1);
|
||||
assert_eq!(msgs[0].role, "user");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_convert_messages_multiple_systems() {
|
||||
let messages = vec![
|
||||
ChatMessage::system("System 1"),
|
||||
ChatMessage::system("System 2"),
|
||||
ChatMessage::user("Hello"),
|
||||
];
|
||||
let (system, msgs) = convert_messages(messages);
|
||||
assert_eq!(system, Some("System 1\n\nSystem 2".to_string()));
|
||||
assert_eq!(msgs.len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_convert_messages_tool_calls() {
|
||||
let tool_calls = vec![ToolCall {
|
||||
id: "call_1".to_string(),
|
||||
name: "search".to_string(),
|
||||
arguments: serde_json::json!({"q": "test"}),
|
||||
}];
|
||||
let messages = vec![
|
||||
ChatMessage::user("Search for test"),
|
||||
ChatMessage::assistant_with_tool_calls(Some("Let me search.".to_string()), tool_calls),
|
||||
ChatMessage::tool_result("call_1", "search", "found it"),
|
||||
];
|
||||
let (system, msgs) = convert_messages(messages);
|
||||
assert!(system.is_none());
|
||||
assert_eq!(msgs.len(), 3);
|
||||
assert_eq!(msgs[0].role, "user");
|
||||
assert_eq!(msgs[1].role, "assistant");
|
||||
// Tool result should be a user message
|
||||
assert_eq!(msgs[2].role, "user");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_response_text_only() {
|
||||
let response = AnthropicResponse {
|
||||
content: vec![AnthropicResponseBlock::Text {
|
||||
text: "Hello!".to_string(),
|
||||
}],
|
||||
stop_reason: Some("end_turn".to_string()),
|
||||
usage: AnthropicUsage {
|
||||
input_tokens: 10,
|
||||
output_tokens: 5,
|
||||
cache_creation_input_tokens: 0,
|
||||
cache_read_input_tokens: 0,
|
||||
},
|
||||
};
|
||||
let (content, tool_calls) = extract_response_content(&response);
|
||||
assert_eq!(content, Some("Hello!".to_string()));
|
||||
assert!(tool_calls.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_response_with_tool_use() {
|
||||
let response = AnthropicResponse {
|
||||
content: vec![
|
||||
AnthropicResponseBlock::Text {
|
||||
text: "Let me search.".to_string(),
|
||||
},
|
||||
AnthropicResponseBlock::ToolUse {
|
||||
id: "call_1".to_string(),
|
||||
name: "search".to_string(),
|
||||
input: serde_json::json!({"q": "test"}),
|
||||
},
|
||||
],
|
||||
stop_reason: Some("tool_use".to_string()),
|
||||
usage: AnthropicUsage {
|
||||
input_tokens: 20,
|
||||
output_tokens: 15,
|
||||
cache_creation_input_tokens: 0,
|
||||
cache_read_input_tokens: 0,
|
||||
},
|
||||
};
|
||||
let (content, tool_calls) = extract_response_content(&response);
|
||||
assert_eq!(content, Some("Let me search.".to_string()));
|
||||
assert_eq!(tool_calls.len(), 1);
|
||||
assert_eq!(tool_calls[0].name, "search");
|
||||
}
|
||||
}
|
||||
@@ -7,6 +7,7 @@
|
||||
//! - **Ollama**: Local model inference
|
||||
//! - **OpenAI-compatible**: Any endpoint that speaks the OpenAI API
|
||||
|
||||
mod anthropic_oauth;
|
||||
pub mod circuit_breaker;
|
||||
pub mod costs;
|
||||
pub mod failover;
|
||||
@@ -178,6 +179,24 @@ fn create_openai_compat_from_registry(
|
||||
fn create_anthropic_from_registry(
|
||||
config: &RegistryProviderConfig,
|
||||
) -> Result<Arc<dyn LlmProvider>, LlmError> {
|
||||
// Route to OAuth provider when an OAuth token is present and no real API
|
||||
// key was provided. When both are set, the API key takes priority (standard
|
||||
// x-api-key auth via rig-core).
|
||||
let api_key_is_placeholder = config
|
||||
.api_key
|
||||
.as_ref()
|
||||
.is_some_and(|k| k.expose_secret() == crate::config::llm::OAUTH_PLACEHOLDER);
|
||||
if config.oauth_token.is_some() && (config.api_key.is_none() || api_key_is_placeholder) {
|
||||
tracing::info!(
|
||||
provider = %config.provider_id,
|
||||
model = %config.model,
|
||||
base_url = if config.base_url.is_empty() { "default" } else { &config.base_url },
|
||||
"Using Anthropic OAuth API"
|
||||
);
|
||||
let provider = anthropic_oauth::AnthropicOAuthProvider::new(config)?;
|
||||
return Ok(Arc::new(provider));
|
||||
}
|
||||
|
||||
use crate::config::CacheRetention;
|
||||
use crate::config::helpers::optional_env;
|
||||
use rig::providers::anthropic;
|
||||
|
||||
@@ -522,4 +522,66 @@ mod tests {
|
||||
assert_eq!(messages[3].role, Role::User); // call_2 orphaned
|
||||
assert_eq!(messages[4].role, Role::User); // call_3 orphaned
|
||||
}
|
||||
|
||||
/// Regression: worker's select_tools/execute_plan now emit
|
||||
/// assistant_with_tool_calls before tool_result messages.
|
||||
/// Verify sanitize_tool_messages preserves all tool_results when
|
||||
/// each has a matching assistant tool_call.
|
||||
#[test]
|
||||
fn test_sanitize_preserves_tool_results_with_matching_assistant() {
|
||||
let tc1 = ToolCall {
|
||||
id: "call_sel_1".to_string(),
|
||||
name: "search".to_string(),
|
||||
arguments: serde_json::json!({"q": "test"}),
|
||||
};
|
||||
let tc2 = ToolCall {
|
||||
id: "call_sel_2".to_string(),
|
||||
name: "http".to_string(),
|
||||
arguments: serde_json::json!({"url": "https://example.com"}),
|
||||
};
|
||||
let mut messages = vec![
|
||||
ChatMessage::system("You are a helpful assistant."),
|
||||
ChatMessage::assistant_with_tool_calls(None, vec![tc1, tc2]),
|
||||
ChatMessage::tool_result("call_sel_1", "search", "found 3 results"),
|
||||
ChatMessage::tool_result("call_sel_2", "http", "200 OK"),
|
||||
];
|
||||
sanitize_tool_messages(&mut messages);
|
||||
|
||||
// All tool_results must keep Role::Tool -- none should be rewritten.
|
||||
assert_eq!(messages[2].role, Role::Tool);
|
||||
assert_eq!(messages[2].tool_call_id, Some("call_sel_1".to_string()));
|
||||
assert_eq!(messages[2].content, "found 3 results");
|
||||
|
||||
assert_eq!(messages[3].role, Role::Tool);
|
||||
assert_eq!(messages[3].tool_call_id, Some("call_sel_2".to_string()));
|
||||
assert_eq!(messages[3].content, "200 OK");
|
||||
}
|
||||
|
||||
/// Regression: the OLD buggy worker code pushed tool_result messages
|
||||
/// without a preceding assistant_with_tool_calls, causing
|
||||
/// sanitize_tool_messages to rewrite them as orphaned user messages.
|
||||
/// This test reproduces that buggy sequence and confirms the rewrite.
|
||||
#[test]
|
||||
fn test_sanitize_rewrites_orphaned_tool_results() {
|
||||
let mut messages = vec![
|
||||
ChatMessage::system("You are a helpful assistant."),
|
||||
// No assistant_with_tool_calls -- mimics the old bug.
|
||||
ChatMessage::tool_result("call_bug_1", "search", "found 3 results"),
|
||||
ChatMessage::tool_result("call_bug_2", "http", "200 OK"),
|
||||
];
|
||||
sanitize_tool_messages(&mut messages);
|
||||
|
||||
// Both tool_results must be rewritten to Role::User.
|
||||
assert_eq!(messages[1].role, Role::User);
|
||||
assert!(messages[1].content.contains("[Tool `search` returned:"));
|
||||
assert!(messages[1].content.contains("found 3 results"));
|
||||
assert!(messages[1].tool_call_id.is_none());
|
||||
assert!(messages[1].name.is_none());
|
||||
|
||||
assert_eq!(messages[2].role, Role::User);
|
||||
assert!(messages[2].content.contains("[Tool `http` returned:"));
|
||||
assert!(messages[2].content.contains("200 OK"));
|
||||
assert!(messages[2].tool_call_id.is_none());
|
||||
assert!(messages[2].name.is_none());
|
||||
}
|
||||
}
|
||||
|
||||
+86
-3
@@ -8,7 +8,8 @@ use serde::{Deserialize, Serialize};
|
||||
use crate::error::LlmError;
|
||||
|
||||
use crate::llm::{
|
||||
ChatMessage, CompletionRequest, LlmProvider, ToolCall, ToolCompletionRequest, ToolDefinition,
|
||||
ChatMessage, CompletionRequest, LlmProvider, Role, ToolCall, ToolCompletionRequest,
|
||||
ToolDefinition,
|
||||
};
|
||||
use crate::safety::SafetyLayer;
|
||||
|
||||
@@ -460,8 +461,15 @@ impl Reasoning {
|
||||
pub async fn plan(&self, context: &ReasoningContext) -> Result<ActionPlan, LlmError> {
|
||||
let system_prompt = self.build_planning_prompt(context);
|
||||
|
||||
let system_prompt = merge_system_messages(system_prompt, &context.messages);
|
||||
let mut messages = vec![ChatMessage::system(system_prompt)];
|
||||
messages.extend(context.messages.clone());
|
||||
messages.extend(
|
||||
context
|
||||
.messages
|
||||
.iter()
|
||||
.filter(|m| m.role != Role::System)
|
||||
.cloned(),
|
||||
);
|
||||
|
||||
if let Some(ref job) = context.job_description {
|
||||
messages.push(ChatMessage::user(format!(
|
||||
@@ -612,8 +620,15 @@ Respond in JSON format:
|
||||
None => self.build_system_prompt_with_tools(&context.available_tools),
|
||||
};
|
||||
|
||||
let system_prompt = merge_system_messages(system_prompt, &context.messages);
|
||||
let mut messages = vec![ChatMessage::system(system_prompt)];
|
||||
messages.extend(context.messages.clone());
|
||||
messages.extend(
|
||||
context
|
||||
.messages
|
||||
.iter()
|
||||
.filter(|m| m.role != Role::System)
|
||||
.cloned(),
|
||||
);
|
||||
|
||||
let effective_tools = if context.force_text {
|
||||
Vec::new()
|
||||
@@ -1026,6 +1041,22 @@ pub struct SuccessEvaluation {
|
||||
pub suggestions: Vec<String>,
|
||||
}
|
||||
|
||||
/// Merge the reasoning method's system prompt with any system messages already
|
||||
/// present in the conversation context. Strict LLM providers (e.g. Qwen)
|
||||
/// reject conversations with system messages that are not at the very
|
||||
/// beginning, so we concatenate all system content into a single prompt.
|
||||
fn merge_system_messages(primary: String, context_messages: &[ChatMessage]) -> String {
|
||||
let extra: Vec<&str> = context_messages
|
||||
.iter()
|
||||
.filter(|m| m.role == Role::System)
|
||||
.map(|m| m.content.as_str())
|
||||
.collect();
|
||||
if extra.is_empty() {
|
||||
return primary;
|
||||
}
|
||||
format!("{}\n\n---\n\n{}", primary, extra.join("\n\n"))
|
||||
}
|
||||
|
||||
/// Extract JSON from text that might contain other content.
|
||||
fn extract_json(text: &str) -> Option<&str> {
|
||||
// Find the first { and last } to extract JSON
|
||||
@@ -2198,6 +2229,58 @@ That's my plan."#;
|
||||
assert!(cleaned.contains("Here are the results."));
|
||||
}
|
||||
|
||||
// ---- merge_system_messages: duplicate system message regression (Bug #597) ----
|
||||
|
||||
#[test]
|
||||
fn test_merge_system_messages_no_system_in_context() {
|
||||
let messages = vec![
|
||||
ChatMessage::user("Hello"),
|
||||
ChatMessage::assistant("Hi there"),
|
||||
];
|
||||
let result = merge_system_messages("primary prompt".into(), &messages);
|
||||
assert_eq!(result, "primary prompt");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_merge_system_messages_merges_worker_system() {
|
||||
let messages = vec![
|
||||
ChatMessage::system("You are an autonomous agent working on a job.\n\nJob: Test Job"),
|
||||
ChatMessage::user("Do the thing"),
|
||||
];
|
||||
let result = merge_system_messages("planning prompt".into(), &messages);
|
||||
assert!(
|
||||
result.contains("planning prompt"),
|
||||
"must contain the primary prompt"
|
||||
);
|
||||
assert!(
|
||||
result.contains("autonomous agent"),
|
||||
"must contain worker system text"
|
||||
);
|
||||
assert!(
|
||||
result.contains("Test Job"),
|
||||
"must contain job description from worker system message"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_merge_system_messages_multiple_system() {
|
||||
let messages = vec![
|
||||
ChatMessage::system("First system instruction"),
|
||||
ChatMessage::system("Second system instruction"),
|
||||
ChatMessage::user("Hello"),
|
||||
];
|
||||
let result = merge_system_messages("primary".into(), &messages);
|
||||
assert!(result.contains("primary"), "must contain primary prompt");
|
||||
assert!(
|
||||
result.contains("First system instruction"),
|
||||
"must contain first system message"
|
||||
);
|
||||
assert!(
|
||||
result.contains("Second system instruction"),
|
||||
"must contain second system message"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_system_prompt_without_tools_omits_tools_section() {
|
||||
let reasoning = make_test_reasoning();
|
||||
|
||||
@@ -450,6 +450,8 @@ mod tests {
|
||||
if def.protocol == ProviderProtocol::OpenAiCompletions
|
||||
&& def.id != "openai"
|
||||
&& def.id != "openai_compatible"
|
||||
&& def.id != "bedrock"
|
||||
&& def.id != "cloudflare"
|
||||
{
|
||||
assert!(
|
||||
def.default_base_url.is_some(),
|
||||
|
||||
+13
-3
@@ -158,7 +158,10 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
wizard.run().await?;
|
||||
}
|
||||
|
||||
// Load initial config from env + disk + optional TOML (before DB is available)
|
||||
// Load initial config from env + disk + optional TOML (before DB is available).
|
||||
// Credentials may be missing at this point — that's fine. LlmConfig::resolve()
|
||||
// defers gracefully, and AppBuilder::build_all() re-resolves after loading
|
||||
// secrets from the encrypted DB.
|
||||
let toml_path = cli.config.as_deref();
|
||||
let config = match Config::from_env_with_toml(toml_path).await {
|
||||
Ok(c) => c,
|
||||
@@ -475,6 +478,7 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
let mut sse_sender: Option<
|
||||
tokio::sync::broadcast::Sender<ironclaw::channels::web::types::SseEvent>,
|
||||
> = None;
|
||||
let mut routine_engine_slot: Option<ironclaw::channels::web::server::RoutineEngineSlot> = None;
|
||||
if let Some(ref gw_config) = config.channels.gateway {
|
||||
let mut gw =
|
||||
GatewayChannel::new(gw_config.clone()).with_llm_provider(Arc::clone(&components.llm));
|
||||
@@ -528,10 +532,11 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
|
||||
tracing::info!("Web UI: http://{}:{}/", gw_config.host, gw_config.port);
|
||||
|
||||
// Capture SSE sender before moving gw into channels.
|
||||
// Capture SSE sender and routine engine slot before moving gw into channels.
|
||||
// IMPORTANT: This must come after all `with_*` calls since `rebuild_state`
|
||||
// creates a new SseManager, which would orphan this sender.
|
||||
sse_sender = Some(gw.state().sse.sender());
|
||||
routine_engine_slot = Some(Arc::clone(&gw.state().routine_engine));
|
||||
|
||||
channel_names.push("gateway".to_string());
|
||||
channels.add(Box::new(gw)).await;
|
||||
@@ -678,7 +683,7 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
)),
|
||||
};
|
||||
|
||||
let agent = Agent::new(
|
||||
let mut agent = Agent::new(
|
||||
config.agent.clone(),
|
||||
deps,
|
||||
channels,
|
||||
@@ -692,6 +697,11 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
// Fill the scheduler slot now that Agent (and its Scheduler) exist.
|
||||
*scheduler_slot.write().await = Some(agent.scheduler());
|
||||
|
||||
// Give the agent the routine engine slot so it can expose the engine to the gateway.
|
||||
if let Some(slot) = routine_engine_slot {
|
||||
agent.set_routine_engine_slot(slot);
|
||||
}
|
||||
|
||||
agent.run().await?;
|
||||
|
||||
// ── Shutdown ────────────────────────────────────────────────────────
|
||||
|
||||
@@ -498,6 +498,10 @@ pub struct SandboxSettings {
|
||||
/// Additional domains to allow through the network proxy.
|
||||
#[serde(default)]
|
||||
pub extra_allowed_domains: Vec<String>,
|
||||
|
||||
/// Whether Claude Code sandbox mode is enabled.
|
||||
#[serde(default)]
|
||||
pub claude_code_enabled: bool,
|
||||
}
|
||||
|
||||
fn default_sandbox_policy() -> String {
|
||||
@@ -531,6 +535,7 @@ impl Default for SandboxSettings {
|
||||
image: default_sandbox_image(),
|
||||
auto_pull_image: true,
|
||||
extra_allowed_domains: Vec::new(),
|
||||
claude_code_enabled: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+2
-6
@@ -152,12 +152,8 @@ This is OS-level behavior we cannot prevent. To minimize pain:
|
||||
rather than triggering system dialogs.
|
||||
|
||||
**Invariant:** After Step 2, `self.secrets_crypto` is `Some` if the user
|
||||
chose Keychain or env-var mode (both generate a key and initialize crypto
|
||||
immediately). It is `None` only if the user skipped secrets.
|
||||
|
||||
When env-var mode is chosen, the generated key is also stored in
|
||||
`self.secrets_master_key_hex` so that `write_bootstrap_env()` can persist
|
||||
it to `~/.ironclaw/.env` automatically.
|
||||
chose Keychain or generated a new key. It may be `None` if the user chose
|
||||
env-var mode or skipped secrets.
|
||||
|
||||
---
|
||||
|
||||
|
||||
+222
-69
@@ -22,6 +22,7 @@ use crate::bootstrap::ironclaw_base_dir;
|
||||
use crate::channels::wasm::{
|
||||
ChannelCapabilitiesFile, available_channel_names, install_bundled_channel,
|
||||
};
|
||||
use crate::config::llm::OAUTH_PLACEHOLDER;
|
||||
use crate::llm::{SessionConfig, SessionManager};
|
||||
use crate::secrets::{SecretsCrypto, SecretsStore};
|
||||
use crate::settings::{KeySource, Settings};
|
||||
@@ -90,8 +91,6 @@ pub struct SetupWizard {
|
||||
db_backend: Option<crate::db::libsql::LibSqlBackend>,
|
||||
/// Secrets crypto (created during setup).
|
||||
secrets_crypto: Option<Arc<SecretsCrypto>>,
|
||||
/// Generated master key hex (stored for writing to .env in env-var mode).
|
||||
secrets_master_key_hex: Option<String>,
|
||||
/// Cached API key from provider setup (used by model fetcher without env mutation).
|
||||
llm_api_key: Option<SecretString>,
|
||||
}
|
||||
@@ -108,7 +107,6 @@ impl SetupWizard {
|
||||
#[cfg(feature = "libsql")]
|
||||
db_backend: None,
|
||||
secrets_crypto: None,
|
||||
secrets_master_key_hex: None,
|
||||
llm_api_key: None,
|
||||
}
|
||||
}
|
||||
@@ -124,7 +122,6 @@ impl SetupWizard {
|
||||
#[cfg(feature = "libsql")]
|
||||
db_backend: None,
|
||||
secrets_crypto: None,
|
||||
secrets_master_key_hex: None,
|
||||
llm_api_key: None,
|
||||
}
|
||||
}
|
||||
@@ -772,25 +769,16 @@ impl SetupWizard {
|
||||
print_success("Master key generated and stored in OS keychain");
|
||||
}
|
||||
1 => {
|
||||
// Env var mode: generate key, initialize crypto, and persist to .env
|
||||
print_info("Generating master key...");
|
||||
// Env var mode
|
||||
print_info("Generate a key and add it to your environment:");
|
||||
let key_hex = crate::secrets::keychain::generate_master_key_hex();
|
||||
|
||||
// Initialize crypto so subsequent steps (API key storage) work
|
||||
self.secrets_crypto = Some(Arc::new(
|
||||
SecretsCrypto::new(SecretString::from(key_hex.clone()))
|
||||
.map_err(|e| SetupError::Config(e.to_string()))?,
|
||||
));
|
||||
|
||||
// Store for write_bootstrap_env to persist to ~/.ironclaw/.env
|
||||
self.secrets_master_key_hex = Some(key_hex.clone());
|
||||
|
||||
println!();
|
||||
print_info(&format!("Generated master key: {}", mask_api_key(&key_hex)));
|
||||
print_info("This key will be saved to ~/.ironclaw/.env automatically.");
|
||||
println!(" export SECRETS_MASTER_KEY={}", key_hex);
|
||||
println!();
|
||||
print_info("Add this to your shell profile or .env file.");
|
||||
|
||||
self.settings.secrets_master_key_source = KeySource::Env;
|
||||
print_success("Master key generated and configured for environment variable");
|
||||
print_success("Configured for environment variable");
|
||||
}
|
||||
_ => {
|
||||
self.settings.secrets_master_key_source = KeySource::None;
|
||||
@@ -899,6 +887,11 @@ impl SetupWizard {
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
// Anthropic has a custom flow: API key or OAuth token from `claude login`.
|
||||
if provider_id == "anthropic" {
|
||||
return self.setup_anthropic().await;
|
||||
}
|
||||
|
||||
match setup {
|
||||
crate::llm::registry::SetupHint::ApiKey {
|
||||
secret_name,
|
||||
@@ -1004,6 +997,112 @@ impl SetupWizard {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Anthropic provider setup: API key or OAuth token from `claude login`.
|
||||
async fn setup_anthropic(&mut self) -> Result<(), SetupError> {
|
||||
let options = &["Direct API Key", "OAuth Token (from `claude login`)"];
|
||||
let choice = select_one("How do you want to authenticate with Anthropic?", options)
|
||||
.map_err(SetupError::Io)?;
|
||||
|
||||
if choice == 0 {
|
||||
// Standard API key flow
|
||||
self.setup_api_key_provider(
|
||||
"anthropic",
|
||||
"ANTHROPIC_API_KEY",
|
||||
"llm_anthropic_api_key",
|
||||
"Anthropic API key",
|
||||
"https://console.anthropic.com/settings/keys",
|
||||
None,
|
||||
)
|
||||
.await
|
||||
} else {
|
||||
// OAuth token flow
|
||||
self.setup_anthropic_oauth().await
|
||||
}
|
||||
}
|
||||
|
||||
/// Anthropic OAuth setup: extract token from `claude login` credentials.
|
||||
async fn setup_anthropic_oauth(&mut self) -> Result<(), SetupError> {
|
||||
self.settings.llm_backend = Some("anthropic".to_string());
|
||||
if self.settings.selected_model.is_some() {
|
||||
self.settings.selected_model = None;
|
||||
}
|
||||
|
||||
// Try to extract existing OAuth token from Claude Code credentials
|
||||
if let Some(token) = crate::config::ClaudeCodeConfig::extract_oauth_token() {
|
||||
print_info(&format!("Found OAuth token: {}", mask_api_key(&token)));
|
||||
if confirm("Use this token?", true).map_err(SetupError::Io)? {
|
||||
return self.save_anthropic_oauth_token(&token).await;
|
||||
}
|
||||
} else {
|
||||
print_info("No OAuth token found from `claude login`.");
|
||||
print_info("Run `claude login` in a terminal to authenticate, then retry.");
|
||||
println!();
|
||||
|
||||
if confirm("Retry after running `claude login`?", true).map_err(SetupError::Io)? {
|
||||
// Block until the user has run `claude login` in another terminal
|
||||
input("Press Enter after running `claude login` in another terminal...")
|
||||
.map_err(SetupError::Io)?;
|
||||
if let Some(token) = crate::config::ClaudeCodeConfig::extract_oauth_token() {
|
||||
print_info(&format!("Found OAuth token: {}", mask_api_key(&token)));
|
||||
return self.save_anthropic_oauth_token(&token).await;
|
||||
}
|
||||
print_error("Still no OAuth token found.");
|
||||
}
|
||||
}
|
||||
|
||||
// Fallback: let user paste the token manually, or switch to API key
|
||||
print_info("You can paste your OAuth token directly (starts with sk-ant-oat01-).");
|
||||
print_info("Or press Enter with no input to switch to the API key flow.");
|
||||
let token = secret_input("Anthropic OAuth token").map_err(SetupError::Io)?;
|
||||
let token_str = token.expose_secret();
|
||||
if token_str.is_empty() {
|
||||
print_info("Switching to API key flow...");
|
||||
return self
|
||||
.setup_api_key_provider(
|
||||
"anthropic",
|
||||
"ANTHROPIC_API_KEY",
|
||||
"llm_anthropic_api_key",
|
||||
"Anthropic API key",
|
||||
"https://console.anthropic.com/settings/keys",
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
self.save_anthropic_oauth_token(token_str).await
|
||||
}
|
||||
|
||||
/// Save an Anthropic OAuth token to secrets and set env for immediate use.
|
||||
async fn save_anthropic_oauth_token(&mut self, token: &str) -> Result<(), SetupError> {
|
||||
// Validate token format to catch accidentally pasted API keys
|
||||
if !token.starts_with("sk-ant-oat") {
|
||||
print_error("Token doesn't look like an OAuth token (expected prefix: sk-ant-oat).");
|
||||
print_info("If you have an API key instead, use the 'Direct API Key' option.");
|
||||
return Err(SetupError::Config("Invalid OAuth token format".to_string()));
|
||||
}
|
||||
|
||||
// Store in secrets if available
|
||||
if let Ok(ctx) = self.init_secrets_context().await {
|
||||
let key = SecretString::from(token.to_string());
|
||||
ctx.save_secret("llm_anthropic_oauth_token", &key)
|
||||
.await
|
||||
.map_err(|e| SetupError::Config(format!("Failed to save OAuth token: {e}")))?;
|
||||
print_success("OAuth token encrypted and saved");
|
||||
} else {
|
||||
print_info("Secrets not available. Set ANTHROPIC_OAUTH_TOKEN in your environment.");
|
||||
}
|
||||
|
||||
// Make the token visible to `optional_env()` for subsequent config
|
||||
// resolution (model selection step). Uses the thread-safe overlay
|
||||
// instead of `std::env::set_var` to avoid UB on multi-threaded runtimes.
|
||||
crate::config::inject_single_var("ANTHROPIC_OAUTH_TOKEN", token);
|
||||
|
||||
// Cache for model fetching
|
||||
self.llm_api_key = Some(SecretString::from(token.to_string()));
|
||||
|
||||
print_success("Anthropic OAuth configured");
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Shared setup flow for API-key-based providers.
|
||||
async fn setup_api_key_provider(
|
||||
&mut self,
|
||||
@@ -1065,6 +1164,11 @@ impl SetupWizard {
|
||||
));
|
||||
}
|
||||
|
||||
// Make key visible to `optional_env()` for subsequent config resolution.
|
||||
// Uses the thread-safe overlay instead of `std::env::set_var` to avoid
|
||||
// UB on multi-threaded runtimes.
|
||||
crate::config::inject_single_var(env_var, key_str);
|
||||
|
||||
// Cache key in memory for model fetching later in the wizard
|
||||
self.llm_api_key = Some(SecretString::from(key_str.to_string()));
|
||||
|
||||
@@ -2001,6 +2105,67 @@ impl SetupWizard {
|
||||
}
|
||||
}
|
||||
|
||||
// Claude Code sandbox sub-step (only if Docker sandbox is enabled)
|
||||
if self.settings.sandbox.enabled {
|
||||
self.step_claude_code_sandbox().await?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Claude Code sandbox sub-step: enable Claude CLI inside Docker containers.
|
||||
async fn step_claude_code_sandbox(&mut self) -> Result<(), SetupError> {
|
||||
println!();
|
||||
print_info("Claude Code mode lets the agent delegate complex tasks to Claude CLI");
|
||||
print_info("running inside sandboxed Docker containers.");
|
||||
println!();
|
||||
|
||||
if !confirm("Enable Claude Code sandbox mode?", false).map_err(SetupError::Io)? {
|
||||
self.settings.sandbox.claude_code_enabled = false;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// Check for Anthropic credentials (API key or OAuth token).
|
||||
// Uses `optional_env()` which reads both real env vars and the
|
||||
// injected overlay (secrets DB, wizard-set values).
|
||||
let has_credentials = || {
|
||||
let has_api_key = crate::config::helpers::optional_env("ANTHROPIC_API_KEY")
|
||||
.ok()
|
||||
.flatten()
|
||||
.is_some_and(|v| !v.is_empty() && v != OAUTH_PLACEHOLDER);
|
||||
let has_oauth = crate::config::ClaudeCodeConfig::extract_oauth_token().is_some()
|
||||
|| crate::config::helpers::optional_env("ANTHROPIC_OAUTH_TOKEN")
|
||||
.ok()
|
||||
.flatten()
|
||||
.is_some_and(|v| !v.is_empty());
|
||||
has_api_key || has_oauth
|
||||
};
|
||||
|
||||
if has_credentials() {
|
||||
self.settings.sandbox.claude_code_enabled = true;
|
||||
print_success("Claude Code sandbox enabled");
|
||||
} else {
|
||||
print_error("No Anthropic credentials found.");
|
||||
print_info(
|
||||
"Claude Code needs ANTHROPIC_API_KEY or an OAuth token from `claude login`.",
|
||||
);
|
||||
println!();
|
||||
|
||||
if confirm("Retry after setting up credentials?", false).map_err(SetupError::Io)? {
|
||||
if has_credentials() {
|
||||
self.settings.sandbox.claude_code_enabled = true;
|
||||
print_success("Claude Code sandbox enabled");
|
||||
} else {
|
||||
self.settings.sandbox.claude_code_enabled = false;
|
||||
print_info("No credentials found. Claude Code disabled for now.");
|
||||
print_info("Set ANTHROPIC_API_KEY or run `claude login` and enable later.");
|
||||
}
|
||||
} else {
|
||||
self.settings.sandbox.claude_code_enabled = false;
|
||||
print_info("Claude Code disabled. Enable with CLAUDE_CODE_ENABLED=true later.");
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -2094,6 +2259,12 @@ impl SetupWizard {
|
||||
///
|
||||
/// These are the chicken-and-egg settings needed before the database is
|
||||
/// connected (DATABASE_BACKEND, DATABASE_URL, LLM_BACKEND, etc.).
|
||||
///
|
||||
/// **Credentials are NOT written here.** API keys and OAuth tokens live
|
||||
/// only in the encrypted secrets DB. `LlmConfig::resolve()` defers
|
||||
/// gracefully when credentials are missing during early startup, and the
|
||||
/// re-resolution in `AppBuilder::build_all()` fills them in after
|
||||
/// `inject_llm_keys_from_secrets()` loads from encrypted storage.
|
||||
fn write_bootstrap_env(&self) -> Result<(), SetupError> {
|
||||
let registry = crate::llm::ProviderRegistry::load();
|
||||
let mut env_vars: Vec<(String, String)> = Vec::new();
|
||||
@@ -2146,13 +2317,6 @@ impl SetupWizard {
|
||||
env_vars.push((base_url_env.clone(), base_url.clone()));
|
||||
}
|
||||
|
||||
// Persist SECRETS_MASTER_KEY when env-var mode was chosen in step 2
|
||||
if self.settings.secrets_master_key_source == KeySource::Env
|
||||
&& let Some(ref key_hex) = self.secrets_master_key_hex
|
||||
{
|
||||
env_vars.push(("SECRETS_MASTER_KEY".to_string(), key_hex.clone()));
|
||||
}
|
||||
|
||||
// Preserve NEARAI_API_KEY if present (set by API key auth flow)
|
||||
if let Ok(api_key) = std::env::var("NEARAI_API_KEY")
|
||||
&& !api_key.is_empty()
|
||||
@@ -2166,6 +2330,11 @@ impl SetupWizard {
|
||||
env_vars.push(("ONBOARD_COMPLETED".to_string(), "true".to_string()));
|
||||
}
|
||||
|
||||
// Claude Code sandbox mode
|
||||
if self.settings.sandbox.claude_code_enabled {
|
||||
env_vars.push(("CLAUDE_CODE_ENABLED".to_string(), "true".to_string()));
|
||||
}
|
||||
|
||||
// Signal channel env vars (chicken-and-egg: config resolves before DB).
|
||||
if let Some(ref url) = self.settings.channels.signal_http_url {
|
||||
env_vars.push(("SIGNAL_HTTP_URL".to_string(), url.clone()));
|
||||
@@ -2533,22 +2702,39 @@ async fn fetch_anthropic_models(cached_key: Option<&str>) -> Vec<(String, String
|
||||
let api_key = cached_key
|
||||
.map(String::from)
|
||||
.or_else(|| std::env::var("ANTHROPIC_API_KEY").ok())
|
||||
.filter(|k| !k.is_empty());
|
||||
.filter(|k| !k.is_empty() && k != crate::config::llm::OAUTH_PLACEHOLDER);
|
||||
|
||||
let api_key = match api_key {
|
||||
Some(k) => k,
|
||||
None => return static_defaults,
|
||||
// Fall back to OAuth token if no API key
|
||||
let oauth_token = if api_key.is_none() {
|
||||
crate::config::helpers::optional_env("ANTHROPIC_OAUTH_TOKEN")
|
||||
.ok()
|
||||
.flatten()
|
||||
.filter(|t| !t.is_empty())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let (key_or_token, is_oauth) = match (api_key, oauth_token) {
|
||||
(Some(k), _) => (k, false),
|
||||
(None, Some(t)) => (t, true),
|
||||
(None, None) => return static_defaults,
|
||||
};
|
||||
|
||||
let client = reqwest::Client::new();
|
||||
let resp = match client
|
||||
let mut request = client
|
||||
.get("https://api.anthropic.com/v1/models")
|
||||
.header("x-api-key", &api_key)
|
||||
.header("anthropic-version", "2023-06-01")
|
||||
.timeout(std::time::Duration::from_secs(5))
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
.timeout(std::time::Duration::from_secs(5));
|
||||
|
||||
if is_oauth {
|
||||
request = request
|
||||
.bearer_auth(&key_or_token)
|
||||
.header("anthropic-beta", "oauth-2025-04-20");
|
||||
} else {
|
||||
request = request.header("x-api-key", &key_or_token);
|
||||
}
|
||||
|
||||
let resp = match request.send().await {
|
||||
Ok(r) if r.status().is_success() => r,
|
||||
_ => return static_defaults,
|
||||
};
|
||||
@@ -3313,39 +3499,6 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
/// Regression test for #666: env var mode in step_security must initialize
|
||||
/// secrets_crypto (for immediate API key storage) and secrets_master_key_hex
|
||||
/// (for persisting to ~/.ironclaw/.env via write_bootstrap_env).
|
||||
#[test]
|
||||
fn test_env_var_mode_initializes_crypto_and_stores_key() {
|
||||
let mut wizard = SetupWizard::new();
|
||||
assert!(wizard.secrets_crypto.is_none());
|
||||
assert!(wizard.secrets_master_key_hex.is_none());
|
||||
|
||||
// Simulate the env-var branch of step_security
|
||||
let key_hex = crate::secrets::keychain::generate_master_key_hex();
|
||||
|
||||
// Verify it's a valid 64-char hex string (32 bytes = AES-256)
|
||||
assert_eq!(key_hex.len(), 64);
|
||||
assert!(key_hex.chars().all(|c| c.is_ascii_hexdigit()));
|
||||
|
||||
let crypto = SecretsCrypto::new(SecretString::from(key_hex.clone()))
|
||||
.expect("SecretsCrypto::new should succeed with generated hex key");
|
||||
|
||||
wizard.secrets_crypto = Some(Arc::new(crypto));
|
||||
wizard.secrets_master_key_hex = Some(key_hex.clone());
|
||||
wizard.settings.secrets_master_key_source = KeySource::Env;
|
||||
|
||||
// Verify crypto is usable for immediate secret encryption
|
||||
assert!(wizard.secrets_crypto.is_some());
|
||||
|
||||
// Verify the hex key is stored for write_bootstrap_env to persist
|
||||
assert_eq!(
|
||||
wizard.secrets_master_key_hex.as_deref(),
|
||||
Some(key_hex.as_str())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_run_provider_setup_no_setup_hint() {
|
||||
// A provider with setup: None should not error. It should set the
|
||||
|
||||
@@ -620,9 +620,13 @@ impl Tool for RoutineFireTool {
|
||||
.map_err(|e| ToolError::ExecutionFailed(format!("DB error: {e}")))?
|
||||
.ok_or_else(|| ToolError::ExecutionFailed(format!("routine '{}' not found", name)))?;
|
||||
|
||||
let run_id = self.engine.fire_manual(routine.id).await.map_err(|e| {
|
||||
ToolError::ExecutionFailed(format!("failed to fire routine '{}': {e}", name))
|
||||
})?;
|
||||
let run_id = self
|
||||
.engine
|
||||
.fire_manual(routine.id, None)
|
||||
.await
|
||||
.map_err(|e| {
|
||||
ToolError::ExecutionFailed(format!("failed to fire routine '{}': {e}", name))
|
||||
})?;
|
||||
|
||||
let result = serde_json::json!({
|
||||
"name": name,
|
||||
|
||||
@@ -204,7 +204,9 @@ mod tests {
|
||||
assert!(!session.is_stale(1800));
|
||||
|
||||
// Manually set last_activity to the past to simulate staleness
|
||||
session.last_activity = std::time::Instant::now() - std::time::Duration::from_secs(10);
|
||||
session.last_activity = std::time::Instant::now()
|
||||
.checked_sub(std::time::Duration::from_secs(10))
|
||||
.expect("System uptime is too low to run staleness test");
|
||||
assert!(session.is_stale(5));
|
||||
assert!(!session.is_stale(15));
|
||||
}
|
||||
|
||||
@@ -211,6 +211,7 @@ async fn start_test_server_with_provider(
|
||||
chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60),
|
||||
registry_entries: Vec::new(),
|
||||
cost_guard: None,
|
||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
||||
startup_time: std::time::Instant::now(),
|
||||
});
|
||||
|
||||
@@ -700,6 +701,7 @@ async fn test_no_llm_provider_returns_503() {
|
||||
chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60),
|
||||
registry_entries: Vec::new(),
|
||||
cost_guard: None,
|
||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
||||
startup_time: std::time::Instant::now(),
|
||||
});
|
||||
|
||||
|
||||
@@ -59,6 +59,7 @@ async fn start_test_server() -> (
|
||||
chat_rate_limiter: ironclaw::channels::web::server::RateLimiter::new(30, 60),
|
||||
registry_entries: Vec::new(),
|
||||
cost_guard: None,
|
||||
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
|
||||
startup_time: std::time::Instant::now(),
|
||||
});
|
||||
|
||||
|
||||
Reference in New Issue
Block a user