mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-28 00:20:16 +00:00
Compare commits
66
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3ca8b1bf68 | ||
|
|
d887309208 | ||
|
|
0c119b5c1e | ||
|
|
2784cef4d7 | ||
|
|
5c56032b88 | ||
|
|
4675e9618c | ||
|
|
d0cb5f0ac5 | ||
|
|
9065527761 | ||
|
|
c6128f4e41 | ||
|
|
ed0ed40dae | ||
|
|
1f209db0fa | ||
|
|
026beb00f2 | ||
|
|
e7ddd46039 | ||
|
|
fc18064be9 | ||
|
|
b50eddfe0a | ||
|
|
878a67cdb6 | ||
|
|
ea0fa7c2c5 | ||
|
|
f2587e1f44 | ||
|
|
218e8778b9 | ||
|
|
971b4c2ef4 | ||
|
|
4890e73a34 | ||
|
|
b8ddbeadb4 | ||
|
|
9aca6a1053 | ||
|
|
63a23550d6 | ||
|
|
4c7afdb0ca | ||
|
|
a580c1d75f | ||
|
|
d1c1bc79c5 | ||
|
|
4277a5a33a | ||
|
|
190c70cdbe | ||
|
|
aa3fac3edc | ||
|
|
ccdce69309 | ||
|
|
de214c23e0 | ||
|
|
946c040fff | ||
|
|
a357972908 | ||
|
|
0245c0f9e9 | ||
|
|
877f117096 | ||
|
|
0c31da46e7 | ||
|
|
596d17f04b | ||
|
|
9e41b8acea | ||
|
|
58a3eb1366 | ||
|
|
f618166ad8 | ||
|
|
3e0e35d1bc | ||
|
|
1b59eb6b39 | ||
|
|
81724cad93 | ||
|
|
c4e098d4e3 | ||
|
|
3debe41f71 | ||
|
|
f470f5db80 | ||
|
|
ca6d9f6ede | ||
|
|
a3c99f2801 | ||
|
|
2b8063a8cf | ||
|
|
3149c91116 | ||
|
|
a71a503870 | ||
|
|
8c2131db48 | ||
|
|
d7024f557f | ||
|
|
99dadcb0ea | ||
|
|
696d6a0bc8 | ||
|
|
ffbc0cd1d4 | ||
|
|
8391415bce | ||
|
|
edca67e8b1 | ||
|
|
7e8c0fbed6 | ||
|
|
6a1301bc5b | ||
|
|
6aae1f8a9e | ||
|
|
7a9396f081 | ||
|
|
6116c885e3 | ||
|
+7 |
a677b20701 | ||
|
|
8c094aec63 |
@@ -18,6 +18,11 @@ DATABASE_POOL_SIZE=10
|
||||
|
||||
# === OpenAI Direct ===
|
||||
# OPENAI_API_KEY=sk-...
|
||||
# Reuse Codex CLI auth.json instead of setting OPENAI_API_KEY manually.
|
||||
# Works with both OpenAI API-key mode and Codex ChatGPT OAuth mode.
|
||||
# In ChatGPT mode this uses the private `chatgpt.com/backend-api/codex` endpoint.
|
||||
# LLM_USE_CODEX_AUTH=true
|
||||
# CODEX_AUTH_PATH=~/.codex/auth.json
|
||||
|
||||
# === NEAR AI (Chat Completions API) ===
|
||||
# Two auth modes:
|
||||
|
||||
@@ -5,6 +5,8 @@ on:
|
||||
- cron: "0 6 * * 1" # Weekly Monday 6 AM UTC
|
||||
workflow_dispatch:
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "src/channels/web/**"
|
||||
- "tests/e2e/**"
|
||||
@@ -50,9 +52,11 @@ jobs:
|
||||
- group: core
|
||||
files: "tests/e2e/scenarios/test_connection.py tests/e2e/scenarios/test_chat.py tests/e2e/scenarios/test_sse_reconnect.py tests/e2e/scenarios/test_html_injection.py tests/e2e/scenarios/test_csp.py"
|
||||
- group: features
|
||||
files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py"
|
||||
files: "tests/e2e/scenarios/test_skills.py tests/e2e/scenarios/test_tool_approval.py tests/e2e/scenarios/test_webhook.py"
|
||||
- group: extensions
|
||||
files: "tests/e2e/scenarios/test_extensions.py tests/e2e/scenarios/test_extension_oauth.py tests/e2e/scenarios/test_wasm_lifecycle.py tests/e2e/scenarios/test_tool_execution.py tests/e2e/scenarios/test_pairing.py tests/e2e/scenarios/test_oauth_credential_fallback.py tests/e2e/scenarios/test_routine_oauth_credential_injection.py"
|
||||
files: "tests/e2e/scenarios/test_extensions.py tests/e2e/scenarios/test_extension_oauth.py tests/e2e/scenarios/test_telegram_token_validation.py tests/e2e/scenarios/test_telegram_hot_activation.py tests/e2e/scenarios/test_wasm_lifecycle.py tests/e2e/scenarios/test_tool_execution.py tests/e2e/scenarios/test_pairing.py tests/e2e/scenarios/test_mcp_auth_flow.py tests/e2e/scenarios/test_oauth_credential_fallback.py tests/e2e/scenarios/test_routine_oauth_credential_injection.py"
|
||||
- group: routines
|
||||
files: "tests/e2e/scenarios/test_owner_scope.py tests/e2e/scenarios/test_routine_event_batch.py"
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
|
||||
|
||||
@@ -17,7 +17,10 @@ jobs:
|
||||
matrix:
|
||||
include:
|
||||
- name: all-features
|
||||
flags: "--features postgres,libsql,html-to-markdown"
|
||||
# Keep product feature coverage broad without pulling in the
|
||||
# test-only `integration` feature, which is exercised separately
|
||||
# in the heavy integration job below.
|
||||
flags: "--no-default-features --features postgres,libsql,html-to-markdown,bedrock,import"
|
||||
- name: default
|
||||
flags: ""
|
||||
- name: libsql-only
|
||||
@@ -39,6 +42,26 @@ jobs:
|
||||
- name: Run Tests
|
||||
run: cargo test ${{ matrix.flags }} -- --nocapture
|
||||
|
||||
heavy-integration-tests:
|
||||
name: Heavy Integration Tests
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v6
|
||||
- name: Install Rust
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
with:
|
||||
targets: wasm32-wasip2
|
||||
- uses: Swatinem/rust-cache@v2
|
||||
with:
|
||||
key: heavy-integration
|
||||
- name: Build Telegram WASM channel
|
||||
run: cargo build --manifest-path channels-src/telegram/Cargo.toml --target wasm32-wasip2 --release
|
||||
- name: Run thread scheduling integration tests
|
||||
run: cargo test --no-default-features --features libsql,integration --test e2e_thread_scheduling -- --nocapture
|
||||
- name: Run Telegram thread-scope regression test
|
||||
run: cargo test --features integration --test telegram_auth_integration test_private_messages_use_chat_id_as_thread_scope -- --exact
|
||||
|
||||
telegram-tests:
|
||||
name: Telegram Channel Tests
|
||||
if: >
|
||||
@@ -65,7 +88,7 @@ jobs:
|
||||
matrix:
|
||||
include:
|
||||
- name: all-features
|
||||
flags: "--all-features"
|
||||
flags: "--no-default-features --features postgres,libsql,html-to-markdown,bedrock,import"
|
||||
- name: default
|
||||
flags: ""
|
||||
- name: libsql-only
|
||||
@@ -149,7 +172,7 @@ jobs:
|
||||
name: Run Tests
|
||||
runs-on: ubuntu-latest
|
||||
if: always()
|
||||
needs: [tests, telegram-tests, wasm-wit-compat, docker-build, windows-build, version-check, bench-compile]
|
||||
needs: [tests, heavy-integration-tests, telegram-tests, wasm-wit-compat, docker-build, windows-build, version-check, bench-compile]
|
||||
steps:
|
||||
- run: |
|
||||
# Unit tests must always pass
|
||||
@@ -157,6 +180,10 @@ jobs:
|
||||
echo "Unit tests failed"
|
||||
exit 1
|
||||
fi
|
||||
if [[ "${{ needs.heavy-integration-tests.result }}" != "success" ]]; then
|
||||
echo "Heavy integration tests failed"
|
||||
exit 1
|
||||
fi
|
||||
# Gated jobs: must pass on promotion PRs / push, skipped on developer PRs
|
||||
for job in telegram-tests wasm-wit-compat docker-build windows-build version-check bench-compile; do
|
||||
case "$job" in
|
||||
|
||||
@@ -33,3 +33,9 @@ trace_*.json
|
||||
# Local Claude Code settings (machine-specific, should not be committed)
|
||||
.claude/settings.local.json
|
||||
.worktrees/
|
||||
|
||||
# Python cache
|
||||
__pycache__/
|
||||
*.pyc
|
||||
*.pyo
|
||||
*.pyd
|
||||
|
||||
Generated
+5
-4
@@ -3461,6 +3461,7 @@ dependencies = [
|
||||
"dirs 6.0.0",
|
||||
"dotenvy",
|
||||
"ed25519-dalek",
|
||||
"eventsource-stream",
|
||||
"flate2",
|
||||
"fs4",
|
||||
"futures",
|
||||
@@ -4364,9 +4365,9 @@ dependencies = [
|
||||
|
||||
[[package]]
|
||||
name = "openssl"
|
||||
version = "0.10.75"
|
||||
version = "0.10.76"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "08838db121398ad17ab8531ce9de97b244589089e290a384c900cb9ff7434328"
|
||||
checksum = "951c002c75e16ea2c65b8c7e4d3d51d5530d8dfa7d060b4776828c88cfb18ecf"
|
||||
dependencies = [
|
||||
"bitflags 2.11.0",
|
||||
"cfg-if",
|
||||
@@ -4402,9 +4403,9 @@ checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe"
|
||||
|
||||
[[package]]
|
||||
name = "openssl-sys"
|
||||
version = "0.9.111"
|
||||
version = "0.9.112"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "82cab2d520aa75e3c58898289429321eb788c3106963d0dc886ec7a5f4adc321"
|
||||
checksum = "57d55af3b3e226502be1526dfdba67ab0e9c96fc293004e79576b2b9edb0dbdb"
|
||||
dependencies = [
|
||||
"cc",
|
||||
"libc",
|
||||
|
||||
@@ -40,6 +40,7 @@ eula = false
|
||||
tokio = { version = "1", features = ["full"] }
|
||||
tokio-stream = { version = "0.1", features = ["sync"] }
|
||||
futures = "0.3"
|
||||
eventsource-stream = "0.2"
|
||||
|
||||
# HTTP client
|
||||
reqwest = { version = "0.12", default-features = false, features = ["json", "multipart", "rustls-tls-native-roots", "stream"] }
|
||||
@@ -221,11 +222,17 @@ postgres = [
|
||||
"rust_decimal/db-tokio-postgres",
|
||||
]
|
||||
libsql = ["dep:libsql"]
|
||||
# Opt-in feature for especially heavy integration-test targets that run in a
|
||||
# dedicated CI job instead of the default Rust test matrix.
|
||||
integration = []
|
||||
html-to-markdown = ["dep:html-to-markdown-rs", "dep:readabilityrs"]
|
||||
bedrock = ["dep:aws-config", "dep:aws-sdk-bedrockruntime", "dep:aws-smithy-types"]
|
||||
import = ["dep:json5", "libsql"]
|
||||
|
||||
[[test]]
|
||||
name = "e2e_thread_scheduling"
|
||||
required-features = ["libsql", "integration"]
|
||||
|
||||
[[test]]
|
||||
name = "html_to_markdown"
|
||||
required-features = ["html-to-markdown"]
|
||||
|
||||
+4
-4
@@ -20,9 +20,9 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
|---------|----------|----------|-------|
|
||||
| Hub-and-spoke architecture | ✅ | ✅ | Web gateway as central hub |
|
||||
| WebSocket control plane | ✅ | ✅ | Gateway with WebSocket + SSE |
|
||||
| Single-user system | ✅ | ✅ | |
|
||||
| Single-user system | ✅ | ✅ | Explicit instance owner scope for persistent routines, secrets, jobs, settings, extensions, and workspace memory |
|
||||
| Multi-agent routing | ✅ | ❌ | Workspace isolation per-agent |
|
||||
| Session-based messaging | ✅ | ✅ | Per-sender sessions |
|
||||
| Session-based messaging | ✅ | ✅ | Owner scope is separate from sender identity and conversation scope |
|
||||
| Loopback-first networking | ✅ | ✅ | HTTP binds to 0.0.0.0 but can be configured |
|
||||
|
||||
### Owner: _Unassigned_
|
||||
@@ -66,9 +66,9 @@ This document tracks feature parity between IronClaw (Rust implementation) and O
|
||||
| CLI/TUI | ✅ | ✅ | - | Ratatui-based TUI |
|
||||
| HTTP webhook | ✅ | ✅ | - | axum with secret validation |
|
||||
| REPL (simple) | ✅ | ✅ | - | For testing |
|
||||
| WASM channels | ❌ | ✅ | - | IronClaw innovation |
|
||||
| WASM channels | ❌ | ✅ | - | IronClaw innovation; host resolves owner scope vs sender identity |
|
||||
| WhatsApp | ✅ | ❌ | P1 | Baileys (Web), same-phone mode with echo detection |
|
||||
| Telegram | ✅ | ✅ | - | WASM channel(MTProto), DM pairing, caption, /start, bot_username, DM topics |
|
||||
| Telegram | ✅ | ✅ | - | WASM channel(MTProto), DM pairing, caption, /start, bot_username, DM topics, setup-time owner auto-verification, owner-scoped persistence |
|
||||
| Discord | ✅ | ❌ | P2 | discord.js, thread parent binding inheritance |
|
||||
| Signal | ✅ | ✅ | P2 | signal-cli daemonPC, SSE listener HTTP/JSON-R, user/group allowlists, DM pairing |
|
||||
| Slack | ✅ | ✅ | - | WASM tool |
|
||||
|
||||
@@ -100,6 +100,14 @@ struct TelegramMessage {
|
||||
|
||||
/// Sticker.
|
||||
sticker: Option<TelegramSticker>,
|
||||
|
||||
/// Forum topic ID. Present when the message is sent inside a forum topic.
|
||||
#[serde(default)]
|
||||
message_thread_id: Option<i64>,
|
||||
|
||||
/// True when this message is sent inside a forum topic.
|
||||
#[serde(default)]
|
||||
is_topic_message: Option<bool>,
|
||||
}
|
||||
|
||||
/// Telegram PhotoSize object.
|
||||
@@ -290,6 +298,10 @@ struct TelegramMessageMetadata {
|
||||
|
||||
/// Whether this is a private (DM) chat.
|
||||
is_private: bool,
|
||||
|
||||
/// Forum topic thread ID (for routing replies back to the correct topic).
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
message_thread_id: Option<i64>,
|
||||
}
|
||||
|
||||
/// Channel configuration injected by host.
|
||||
@@ -491,8 +503,7 @@ impl Guest for TelegramChannel {
|
||||
|
||||
// Delete any existing webhook before polling. Telegram returns success
|
||||
// when no webhook exists, so any error here (e.g. 401) means a bad token.
|
||||
delete_webhook()
|
||||
.map_err(|e| format!("Bot token validation failed: {}", e))?;
|
||||
delete_webhook().map_err(|e| format!("Bot token validation failed: {}", e))?;
|
||||
}
|
||||
|
||||
// Configure polling only if not in webhook mode
|
||||
@@ -680,7 +691,12 @@ impl Guest for TelegramChannel {
|
||||
let metadata: TelegramMessageMetadata = serde_json::from_str(&response.metadata_json)
|
||||
.map_err(|e| format!("Failed to parse metadata: {}", e))?;
|
||||
|
||||
send_response(metadata.chat_id, &response, Some(metadata.message_id))
|
||||
send_response(
|
||||
metadata.chat_id,
|
||||
&response,
|
||||
Some(metadata.message_id),
|
||||
metadata.message_thread_id,
|
||||
)
|
||||
}
|
||||
|
||||
fn on_broadcast(user_id: String, response: AgentResponse) -> Result<(), String> {
|
||||
@@ -688,7 +704,7 @@ impl Guest for TelegramChannel {
|
||||
.parse()
|
||||
.map_err(|e| format!("Invalid chat_id '{}': {}", user_id, e))?;
|
||||
|
||||
send_response(chat_id, &response, None)
|
||||
send_response(chat_id, &response, None, None)
|
||||
}
|
||||
|
||||
fn on_status(update: StatusUpdate) {
|
||||
@@ -712,11 +728,15 @@ impl Guest for TelegramChannel {
|
||||
match action {
|
||||
TelegramStatusAction::Typing => {
|
||||
// POST /sendChatAction with action "typing"
|
||||
let payload = serde_json::json!({
|
||||
let mut payload = serde_json::json!({
|
||||
"chat_id": metadata.chat_id,
|
||||
"action": "typing"
|
||||
});
|
||||
|
||||
if let Some(thread_id) = metadata.message_thread_id {
|
||||
payload["message_thread_id"] = serde_json::Value::Number(thread_id.into());
|
||||
}
|
||||
|
||||
let payload_bytes = match serde_json::to_vec(&payload) {
|
||||
Ok(b) => b,
|
||||
Err(_) => return,
|
||||
@@ -743,9 +763,13 @@ impl Guest for TelegramChannel {
|
||||
}
|
||||
TelegramStatusAction::Notify(prompt) => {
|
||||
// Send user-visible status updates for actionable events.
|
||||
if let Err(first_err) =
|
||||
send_message(metadata.chat_id, &prompt, Some(metadata.message_id), None)
|
||||
{
|
||||
if let Err(first_err) = send_message(
|
||||
metadata.chat_id,
|
||||
&prompt,
|
||||
Some(metadata.message_id),
|
||||
None,
|
||||
metadata.message_thread_id,
|
||||
) {
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Warn,
|
||||
&format!(
|
||||
@@ -754,7 +778,13 @@ impl Guest for TelegramChannel {
|
||||
),
|
||||
);
|
||||
|
||||
if let Err(retry_err) = send_message(metadata.chat_id, &prompt, None, None) {
|
||||
if let Err(retry_err) = send_message(
|
||||
metadata.chat_id,
|
||||
&prompt,
|
||||
None,
|
||||
None,
|
||||
metadata.message_thread_id,
|
||||
) {
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Debug,
|
||||
&format!(
|
||||
@@ -797,6 +827,14 @@ impl std::fmt::Display for SendError {
|
||||
}
|
||||
}
|
||||
|
||||
/// Normalize `message_thread_id` for outbound API calls.
|
||||
///
|
||||
/// Telegram rejects `sendMessage` and file-send methods when
|
||||
/// `message_thread_id = 1` (the "General" topic), so omit it in that case.
|
||||
fn normalize_thread_id(thread_id: Option<i64>) -> Option<i64> {
|
||||
thread_id.filter(|&id| id != 1)
|
||||
}
|
||||
|
||||
/// Send a message via the Telegram Bot API.
|
||||
///
|
||||
/// Returns the sent message_id on success. When `parse_mode` is set and
|
||||
@@ -807,7 +845,10 @@ fn send_message(
|
||||
text: &str,
|
||||
reply_to_message_id: Option<i64>,
|
||||
parse_mode: Option<&str>,
|
||||
message_thread_id: Option<i64>,
|
||||
) -> Result<i64, SendError> {
|
||||
let message_thread_id = normalize_thread_id(message_thread_id);
|
||||
|
||||
let mut payload = serde_json::json!({
|
||||
"chat_id": chat_id,
|
||||
"text": text,
|
||||
@@ -821,6 +862,10 @@ fn send_message(
|
||||
payload["parse_mode"] = serde_json::Value::String(mode.to_string());
|
||||
}
|
||||
|
||||
if let Some(thread_id) = message_thread_id {
|
||||
payload["message_thread_id"] = serde_json::Value::Number(thread_id.into());
|
||||
}
|
||||
|
||||
let payload_bytes = serde_json::to_vec(&payload)
|
||||
.map_err(|e| SendError::Other(format!("Failed to serialize payload: {}", e)))?;
|
||||
|
||||
@@ -911,19 +956,20 @@ fn download_telegram_file(file_id: &str) -> Result<Vec<u8>, String> {
|
||||
);
|
||||
|
||||
let headers = serde_json::json!({});
|
||||
let result =
|
||||
channel_host::http_request("GET", &get_file_url, &headers.to_string(), None, None);
|
||||
let result = channel_host::http_request("GET", &get_file_url, &headers.to_string(), None, None);
|
||||
|
||||
let response = result.map_err(|e| format!("getFile request failed: {}", e))?;
|
||||
|
||||
if response.status != 200 {
|
||||
let body_str = String::from_utf8_lossy(&response.body);
|
||||
return Err(format!("getFile returned {}: {}", response.status, body_str));
|
||||
return Err(format!(
|
||||
"getFile returned {}: {}",
|
||||
response.status, body_str
|
||||
));
|
||||
}
|
||||
|
||||
let api_response: TelegramApiResponse<TelegramFile> =
|
||||
serde_json::from_slice(&response.body)
|
||||
.map_err(|e| format!("Failed to parse getFile response: {}", e))?;
|
||||
let api_response: TelegramApiResponse<TelegramFile> = serde_json::from_slice(&response.body)
|
||||
.map_err(|e| format!("Failed to parse getFile response: {}", e))?;
|
||||
|
||||
if !api_response.ok {
|
||||
return Err(format!(
|
||||
@@ -953,16 +999,12 @@ fn download_telegram_file(file_id: &str) -> Result<Vec<u8>, String> {
|
||||
file_path
|
||||
);
|
||||
|
||||
let result =
|
||||
channel_host::http_request("GET", &download_url, &headers.to_string(), None, None);
|
||||
let result = channel_host::http_request("GET", &download_url, &headers.to_string(), None, None);
|
||||
|
||||
let response = result.map_err(|e| format!("File download failed: {}", e))?;
|
||||
|
||||
if response.status != 200 {
|
||||
return Err(format!(
|
||||
"File download returned status {}",
|
||||
response.status
|
||||
));
|
||||
return Err(format!("File download returned status {}", response.status));
|
||||
}
|
||||
|
||||
// Post-download size guard: Telegram metadata file_size is optional,
|
||||
@@ -1036,7 +1078,10 @@ fn send_photo(
|
||||
mime_type: &str,
|
||||
data: &[u8],
|
||||
reply_to_message_id: Option<i64>,
|
||||
message_thread_id: Option<i64>,
|
||||
) -> Result<(), String> {
|
||||
let message_thread_id = normalize_thread_id(message_thread_id);
|
||||
|
||||
if data.len() > MAX_PHOTO_SIZE {
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Info,
|
||||
@@ -1046,7 +1091,14 @@ fn send_photo(
|
||||
data.len()
|
||||
),
|
||||
);
|
||||
return send_document(chat_id, filename, mime_type, data, reply_to_message_id);
|
||||
return send_document(
|
||||
chat_id,
|
||||
filename,
|
||||
mime_type,
|
||||
data,
|
||||
reply_to_message_id,
|
||||
message_thread_id,
|
||||
);
|
||||
}
|
||||
|
||||
let boundary = format!("ironclaw-{}", channel_host::now_millis());
|
||||
@@ -1054,7 +1106,20 @@ fn send_photo(
|
||||
|
||||
write_multipart_field(&mut body, &boundary, "chat_id", &chat_id.to_string());
|
||||
if let Some(msg_id) = reply_to_message_id {
|
||||
write_multipart_field(&mut body, &boundary, "reply_to_message_id", &msg_id.to_string());
|
||||
write_multipart_field(
|
||||
&mut body,
|
||||
&boundary,
|
||||
"reply_to_message_id",
|
||||
&msg_id.to_string(),
|
||||
);
|
||||
}
|
||||
if let Some(thread_id) = message_thread_id {
|
||||
write_multipart_field(
|
||||
&mut body,
|
||||
&boundary,
|
||||
"message_thread_id",
|
||||
&thread_id.to_string(),
|
||||
);
|
||||
}
|
||||
write_multipart_file(&mut body, &boundary, "photo", filename, mime_type, data);
|
||||
body.extend_from_slice(format!("--{}--\r\n", boundary).as_bytes());
|
||||
@@ -1097,13 +1162,29 @@ fn send_document(
|
||||
mime_type: &str,
|
||||
data: &[u8],
|
||||
reply_to_message_id: Option<i64>,
|
||||
message_thread_id: Option<i64>,
|
||||
) -> Result<(), String> {
|
||||
let message_thread_id = normalize_thread_id(message_thread_id);
|
||||
|
||||
let boundary = format!("ironclaw-{}", channel_host::now_millis());
|
||||
let mut body = Vec::new();
|
||||
|
||||
write_multipart_field(&mut body, &boundary, "chat_id", &chat_id.to_string());
|
||||
if let Some(msg_id) = reply_to_message_id {
|
||||
write_multipart_field(&mut body, &boundary, "reply_to_message_id", &msg_id.to_string());
|
||||
write_multipart_field(
|
||||
&mut body,
|
||||
&boundary,
|
||||
"reply_to_message_id",
|
||||
&msg_id.to_string(),
|
||||
);
|
||||
}
|
||||
if let Some(thread_id) = message_thread_id {
|
||||
write_multipart_field(
|
||||
&mut body,
|
||||
&boundary,
|
||||
"message_thread_id",
|
||||
&thread_id.to_string(),
|
||||
);
|
||||
}
|
||||
write_multipart_file(&mut body, &boundary, "document", filename, mime_type, data);
|
||||
body.extend_from_slice(format!("--{}--\r\n", boundary).as_bytes());
|
||||
@@ -1140,12 +1221,7 @@ fn send_document(
|
||||
}
|
||||
|
||||
/// Image MIME types that Telegram's sendPhoto API supports.
|
||||
const PHOTO_MIME_TYPES: &[&str] = &[
|
||||
"image/jpeg",
|
||||
"image/png",
|
||||
"image/gif",
|
||||
"image/webp",
|
||||
];
|
||||
const PHOTO_MIME_TYPES: &[&str] = &["image/jpeg", "image/png", "image/gif", "image/webp"];
|
||||
|
||||
/// Send a full agent response (attachments + text) to a chat.
|
||||
///
|
||||
@@ -1154,10 +1230,11 @@ fn send_response(
|
||||
chat_id: i64,
|
||||
response: &AgentResponse,
|
||||
reply_to_message_id: Option<i64>,
|
||||
message_thread_id: Option<i64>,
|
||||
) -> Result<(), String> {
|
||||
// Send attachments first (photos/documents)
|
||||
for attachment in &response.attachments {
|
||||
send_attachment(chat_id, attachment, reply_to_message_id)?;
|
||||
send_attachment(chat_id, attachment, reply_to_message_id, message_thread_id)?;
|
||||
}
|
||||
|
||||
// Skip text if empty and we already sent attachments
|
||||
@@ -1166,13 +1243,23 @@ fn send_response(
|
||||
}
|
||||
|
||||
// Try Markdown, fall back to plain text on parse errors
|
||||
match send_message(chat_id, &response.content, reply_to_message_id, Some("Markdown")) {
|
||||
match send_message(
|
||||
chat_id,
|
||||
&response.content,
|
||||
reply_to_message_id,
|
||||
Some("Markdown"),
|
||||
message_thread_id,
|
||||
) {
|
||||
Ok(_) => Ok(()),
|
||||
Err(SendError::ParseEntities(_)) => {
|
||||
send_message(chat_id, &response.content, reply_to_message_id, None)
|
||||
.map(|_| ())
|
||||
.map_err(|e| format!("Plain-text retry also failed: {}", e))
|
||||
}
|
||||
Err(SendError::ParseEntities(_)) => send_message(
|
||||
chat_id,
|
||||
&response.content,
|
||||
reply_to_message_id,
|
||||
None,
|
||||
message_thread_id,
|
||||
)
|
||||
.map(|_| ())
|
||||
.map_err(|e| format!("Plain-text retry also failed: {}", e)),
|
||||
Err(e) => Err(e.to_string()),
|
||||
}
|
||||
}
|
||||
@@ -1182,6 +1269,7 @@ fn send_attachment(
|
||||
chat_id: i64,
|
||||
attachment: &Attachment,
|
||||
reply_to_message_id: Option<i64>,
|
||||
message_thread_id: Option<i64>,
|
||||
) -> Result<(), String> {
|
||||
if PHOTO_MIME_TYPES.contains(&attachment.mime_type.as_str()) {
|
||||
send_photo(
|
||||
@@ -1190,6 +1278,7 @@ fn send_attachment(
|
||||
&attachment.mime_type,
|
||||
&attachment.data,
|
||||
reply_to_message_id,
|
||||
message_thread_id,
|
||||
)
|
||||
} else {
|
||||
send_document(
|
||||
@@ -1198,6 +1287,7 @@ fn send_attachment(
|
||||
&attachment.mime_type,
|
||||
&attachment.data,
|
||||
reply_to_message_id,
|
||||
message_thread_id,
|
||||
)
|
||||
}
|
||||
}
|
||||
@@ -1337,7 +1427,10 @@ fn register_webhook(tunnel_url: &str, webhook_secret: Option<&str>) -> Result<()
|
||||
let context = if retried { " (after retry)" } else { "" };
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Info,
|
||||
&format!("Webhook registered successfully{}: {}", context, webhook_url),
|
||||
&format!(
|
||||
"Webhook registered successfully{}: {}",
|
||||
context, webhook_url
|
||||
),
|
||||
);
|
||||
|
||||
Ok(())
|
||||
@@ -1357,6 +1450,7 @@ fn send_pairing_reply(chat_id: i64, code: &str) -> Result<(), String> {
|
||||
),
|
||||
None,
|
||||
Some("Markdown"),
|
||||
None,
|
||||
)
|
||||
.map(|_| ())
|
||||
.map_err(|e| e.to_string())
|
||||
@@ -1438,7 +1532,9 @@ fn extract_attachments(message: &TelegramMessage) -> Vec<InboundAttachment> {
|
||||
if let Some(ref doc) = message.document {
|
||||
attachments.push(make_inbound_attachment(
|
||||
doc.file_id.clone(),
|
||||
doc.mime_type.clone().unwrap_or_else(|| "application/octet-stream".to_string()),
|
||||
doc.mime_type
|
||||
.clone()
|
||||
.unwrap_or_else(|| "application/octet-stream".to_string()),
|
||||
doc.file_name.clone(),
|
||||
doc.file_size.map(|s| s as u64),
|
||||
Some(get_file_url(&doc.file_id)),
|
||||
@@ -1451,7 +1547,10 @@ fn extract_attachments(message: &TelegramMessage) -> Vec<InboundAttachment> {
|
||||
if let Some(ref audio) = message.audio {
|
||||
attachments.push(make_inbound_attachment(
|
||||
audio.file_id.clone(),
|
||||
audio.mime_type.clone().unwrap_or_else(|| "audio/mpeg".to_string()),
|
||||
audio
|
||||
.mime_type
|
||||
.clone()
|
||||
.unwrap_or_else(|| "audio/mpeg".to_string()),
|
||||
audio.file_name.clone(),
|
||||
audio.file_size.map(|s| s as u64),
|
||||
Some(get_file_url(&audio.file_id)),
|
||||
@@ -1464,7 +1563,10 @@ fn extract_attachments(message: &TelegramMessage) -> Vec<InboundAttachment> {
|
||||
if let Some(ref video) = message.video {
|
||||
attachments.push(make_inbound_attachment(
|
||||
video.file_id.clone(),
|
||||
video.mime_type.clone().unwrap_or_else(|| "video/mp4".to_string()),
|
||||
video
|
||||
.mime_type
|
||||
.clone()
|
||||
.unwrap_or_else(|| "video/mp4".to_string()),
|
||||
video.file_name.clone(),
|
||||
video.file_size.map(|s| s as u64),
|
||||
Some(get_file_url(&video.file_id)),
|
||||
@@ -1689,25 +1791,14 @@ fn handle_message(message: TelegramMessage) {
|
||||
|
||||
let is_private = message.chat.chat_type == "private";
|
||||
|
||||
// Owner validation: when owner_id is set, only that user can message
|
||||
let owner_id_str = channel_host::workspace_read(OWNER_ID_PATH).filter(|s| !s.is_empty());
|
||||
let owner_id = channel_host::workspace_read(OWNER_ID_PATH)
|
||||
.filter(|s| !s.is_empty())
|
||||
.and_then(|s| s.parse::<i64>().ok());
|
||||
let is_owner = owner_id == Some(from.id);
|
||||
|
||||
if let Some(ref id_str) = owner_id_str {
|
||||
if let Ok(owner_id) = id_str.parse::<i64>() {
|
||||
if from.id != owner_id {
|
||||
channel_host::log(
|
||||
channel_host::LogLevel::Debug,
|
||||
&format!(
|
||||
"Dropping message from non-owner user {} (owner: {})",
|
||||
from.id, owner_id
|
||||
),
|
||||
);
|
||||
return;
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// No owner_id: apply authorization based on dm_policy and allow_from
|
||||
// This applies to both private and group chats when owner_id is null
|
||||
if !is_owner {
|
||||
// Non-owner senders remain guests. Apply authorization based on
|
||||
// dm_policy / allow_from before letting them chat in their own scope.
|
||||
let dm_policy =
|
||||
channel_host::workspace_read(DM_POLICY_PATH).unwrap_or_else(|| "pairing".to_string());
|
||||
|
||||
@@ -1814,6 +1905,7 @@ fn handle_message(message: TelegramMessage) {
|
||||
message_id: message.message_id,
|
||||
user_id: from.id,
|
||||
is_private,
|
||||
message_thread_id: message.message_thread_id,
|
||||
};
|
||||
|
||||
let metadata_json = serde_json::to_string(&metadata).unwrap_or_else(|_| "{}".to_string());
|
||||
@@ -1838,7 +1930,7 @@ fn handle_message(message: TelegramMessage) {
|
||||
user_id: from.id.to_string(),
|
||||
user_name: Some(user_name),
|
||||
content: content_to_emit,
|
||||
thread_id: None, // Telegram doesn't have threads in the same way
|
||||
thread_id: Some(message.chat.id.to_string()),
|
||||
metadata_json,
|
||||
attachments,
|
||||
});
|
||||
@@ -2438,7 +2530,11 @@ mod tests {
|
||||
assert_eq!(attachments[0].id, "large_id"); // Largest photo
|
||||
assert_eq!(attachments[0].mime_type, "image/jpeg");
|
||||
assert_eq!(attachments[0].size_bytes, Some(54321));
|
||||
assert!(attachments[0].source_url.as_ref().unwrap().contains("large_id"));
|
||||
assert!(attachments[0]
|
||||
.source_url
|
||||
.as_ref()
|
||||
.unwrap()
|
||||
.contains("large_id"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -2490,9 +2586,7 @@ mod tests {
|
||||
attachments[0].filename.as_deref(),
|
||||
Some("voice_voice_xyz.ogg")
|
||||
);
|
||||
assert!(attachments[0]
|
||||
.extras_json
|
||||
.contains("\"duration_secs\":5"));
|
||||
assert!(attachments[0].extras_json.contains("\"duration_secs\":5"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -2638,18 +2732,33 @@ mod tests {
|
||||
};
|
||||
|
||||
// PDFs and Office docs should be downloaded
|
||||
assert!(is_downloadable_document(&make("application/pdf", Some("report.pdf"))));
|
||||
assert!(is_downloadable_document(&make(
|
||||
"application/pdf",
|
||||
Some("report.pdf")
|
||||
)));
|
||||
assert!(is_downloadable_document(&make(
|
||||
"application/vnd.openxmlformats-officedocument.wordprocessingml.document",
|
||||
Some("doc.docx"),
|
||||
)));
|
||||
assert!(is_downloadable_document(&make("text/plain", Some("notes.txt"))));
|
||||
assert!(is_downloadable_document(&make(
|
||||
"text/plain",
|
||||
Some("notes.txt")
|
||||
)));
|
||||
|
||||
// Voice, image, audio, video should NOT be downloaded
|
||||
assert!(!is_downloadable_document(&make("audio/ogg", Some("voice_123.ogg"))));
|
||||
assert!(!is_downloadable_document(&make(
|
||||
"audio/ogg",
|
||||
Some("voice_123.ogg")
|
||||
)));
|
||||
assert!(!is_downloadable_document(&make("image/jpeg", None)));
|
||||
assert!(!is_downloadable_document(&make("audio/mpeg", Some("song.mp3"))));
|
||||
assert!(!is_downloadable_document(&make("video/mp4", Some("clip.mp4"))));
|
||||
assert!(!is_downloadable_document(&make(
|
||||
"audio/mpeg",
|
||||
Some("song.mp3")
|
||||
)));
|
||||
assert!(!is_downloadable_document(&make(
|
||||
"video/mp4",
|
||||
Some("clip.mp4")
|
||||
)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -324,7 +324,7 @@ mod tests {
|
||||
let violations = policy.check(&payload);
|
||||
let elapsed = start.elapsed();
|
||||
assert!(
|
||||
elapsed.as_millis() < 100,
|
||||
elapsed.as_millis() < 500,
|
||||
"excessive_urls pattern took {}ms on 100KB near-miss",
|
||||
elapsed.as_millis()
|
||||
);
|
||||
@@ -349,7 +349,7 @@ mod tests {
|
||||
let violations = policy.check(&payload);
|
||||
let elapsed = start.elapsed();
|
||||
assert!(
|
||||
elapsed.as_millis() < 100,
|
||||
elapsed.as_millis() < 500,
|
||||
"obfuscated_string pattern took {}ms on 100KB near-miss",
|
||||
elapsed.as_millis()
|
||||
);
|
||||
@@ -370,7 +370,7 @@ mod tests {
|
||||
let _violations = policy.check(&payload);
|
||||
let elapsed = start.elapsed();
|
||||
assert!(
|
||||
elapsed.as_millis() < 100,
|
||||
elapsed.as_millis() < 500,
|
||||
"shell_injection pattern took {}ms on 100KB near-miss",
|
||||
elapsed.as_millis()
|
||||
);
|
||||
@@ -387,7 +387,7 @@ mod tests {
|
||||
let _violations = policy.check(&payload);
|
||||
let elapsed = start.elapsed();
|
||||
assert!(
|
||||
elapsed.as_millis() < 100,
|
||||
elapsed.as_millis() < 500,
|
||||
"sql_pattern took {}ms on 100KB near-miss",
|
||||
elapsed.as_millis()
|
||||
);
|
||||
@@ -405,7 +405,7 @@ mod tests {
|
||||
let _violations = policy.check(&payload);
|
||||
let elapsed = start.elapsed();
|
||||
assert!(
|
||||
elapsed.as_millis() < 100,
|
||||
elapsed.as_millis() < 500,
|
||||
"crypto_private_key pattern took {}ms on 100KB near-miss",
|
||||
elapsed.as_millis()
|
||||
);
|
||||
@@ -423,7 +423,7 @@ mod tests {
|
||||
let _violations = policy.check(&payload);
|
||||
let elapsed = start.elapsed();
|
||||
assert!(
|
||||
elapsed.as_millis() < 100,
|
||||
elapsed.as_millis() < 500,
|
||||
"system_file_access pattern took {}ms on 100KB near-miss",
|
||||
elapsed.as_millis()
|
||||
);
|
||||
@@ -441,7 +441,7 @@ mod tests {
|
||||
let _violations = policy.check(&payload);
|
||||
let elapsed = start.elapsed();
|
||||
assert!(
|
||||
elapsed.as_millis() < 100,
|
||||
elapsed.as_millis() < 500,
|
||||
"encoded_exploit pattern took {}ms on 100KB near-miss",
|
||||
elapsed.as_millis()
|
||||
);
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
-- Remove the legacy 'default' sentinel from routine notifications.
|
||||
-- A NULL notify_user now means "resolve the configured owner's last-seen
|
||||
-- channel target at send time."
|
||||
|
||||
ALTER TABLE routines
|
||||
ALTER COLUMN notify_user DROP NOT NULL,
|
||||
ALTER COLUMN notify_user DROP DEFAULT;
|
||||
|
||||
UPDATE routines
|
||||
SET notify_user = NULL
|
||||
WHERE notify_user = 'default';
|
||||
@@ -26,7 +26,7 @@ CREATE TABLE routines (
|
||||
|
||||
-- Notification preferences
|
||||
notify_channel TEXT, -- NULL = use default
|
||||
notify_user TEXT NOT NULL DEFAULT 'default',
|
||||
notify_user TEXT,
|
||||
notify_on_success BOOLEAN NOT NULL DEFAULT false,
|
||||
notify_on_failure BOOLEAN NOT NULL DEFAULT true,
|
||||
notify_on_attention BOOLEAN NOT NULL DEFAULT true,
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "discord",
|
||||
"display_name": "Discord Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Talk to your agent in Discord",
|
||||
"keywords": [
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "feishu",
|
||||
"display_name": "Feishu / Lark Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.1.0",
|
||||
"version": "0.1.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Talk to your agent through a Feishu or Lark bot",
|
||||
"keywords": [
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "telegram",
|
||||
"display_name": "Telegram Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.2.3",
|
||||
"version": "0.2.4",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Talk to your agent through a Telegram bot",
|
||||
"keywords": [
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "github",
|
||||
"display_name": "GitHub",
|
||||
"kind": "tool",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "GitHub integration for issues, PRs, repos, and code search",
|
||||
"keywords": [
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "web-search",
|
||||
"display_name": "Web Search",
|
||||
"kind": "tool",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Search the web using Brave Search API",
|
||||
"keywords": [
|
||||
|
||||
+224
-36
@@ -22,7 +22,7 @@ use crate::channels::{ChannelManager, IncomingMessage, OutgoingResponse};
|
||||
use crate::config::{AgentConfig, HeartbeatConfig, RoutineConfig, SkillsConfig};
|
||||
use crate::context::ContextManager;
|
||||
use crate::db::Database;
|
||||
use crate::error::Error;
|
||||
use crate::error::{ChannelError, Error};
|
||||
use crate::extensions::ExtensionManager;
|
||||
use crate::hooks::HookRegistry;
|
||||
use crate::llm::LlmProvider;
|
||||
@@ -54,10 +54,75 @@ pub(crate) fn truncate_for_preview(output: &str, max_chars: usize) -> String {
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
fn resolve_routine_notification_user(metadata: &serde_json::Value) -> Option<String> {
|
||||
resolve_owner_scope_notification_user(
|
||||
metadata.get("notify_user").and_then(|value| value.as_str()),
|
||||
metadata.get("owner_id").and_then(|value| value.as_str()),
|
||||
)
|
||||
}
|
||||
|
||||
fn trimmed_option(value: Option<&str>) -> Option<String> {
|
||||
value
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn resolve_owner_scope_notification_user(
|
||||
explicit_user: Option<&str>,
|
||||
owner_fallback: Option<&str>,
|
||||
) -> Option<String> {
|
||||
trimmed_option(explicit_user).or_else(|| trimmed_option(owner_fallback))
|
||||
}
|
||||
|
||||
async fn resolve_channel_notification_user(
|
||||
extension_manager: Option<&Arc<ExtensionManager>>,
|
||||
channel: Option<&str>,
|
||||
explicit_user: Option<&str>,
|
||||
owner_fallback: Option<&str>,
|
||||
) -> Option<String> {
|
||||
if let Some(user) = trimmed_option(explicit_user) {
|
||||
return Some(user);
|
||||
}
|
||||
|
||||
if let Some(channel_name) = trimmed_option(channel)
|
||||
&& let Some(extension_manager) = extension_manager
|
||||
&& let Some(target) = extension_manager
|
||||
.notification_target_for_channel(&channel_name)
|
||||
.await
|
||||
{
|
||||
return Some(target);
|
||||
}
|
||||
|
||||
resolve_owner_scope_notification_user(explicit_user, owner_fallback)
|
||||
}
|
||||
|
||||
async fn resolve_routine_notification_target(
|
||||
extension_manager: Option<&Arc<ExtensionManager>>,
|
||||
metadata: &serde_json::Value,
|
||||
) -> Option<String> {
|
||||
resolve_channel_notification_user(
|
||||
extension_manager,
|
||||
metadata
|
||||
.get("notify_channel")
|
||||
.and_then(|value| value.as_str()),
|
||||
metadata.get("notify_user").and_then(|value| value.as_str()),
|
||||
metadata.get("owner_id").and_then(|value| value.as_str()),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn should_fallback_routine_notification(error: &ChannelError) -> bool {
|
||||
!matches!(error, ChannelError::MissingRoutingTarget { .. })
|
||||
}
|
||||
|
||||
/// Core dependencies for the agent.
|
||||
///
|
||||
/// Bundles the shared components to reduce argument count.
|
||||
pub struct AgentDeps {
|
||||
/// Resolved durable owner scope for the instance.
|
||||
pub owner_id: String,
|
||||
pub store: Option<Arc<dyn Database>>,
|
||||
pub llm: Arc<dyn LlmProvider>,
|
||||
/// Cheap/fast LLM for lightweight tasks (heartbeat, routing, evaluation).
|
||||
@@ -102,6 +167,18 @@ pub struct Agent {
|
||||
}
|
||||
|
||||
impl Agent {
|
||||
pub(super) fn owner_id(&self) -> &str {
|
||||
if let Some(workspace) = self.deps.workspace.as_ref() {
|
||||
debug_assert_eq!(
|
||||
workspace.user_id(),
|
||||
self.deps.owner_id,
|
||||
"workspace.user_id() must stay aligned with deps.owner_id"
|
||||
);
|
||||
}
|
||||
|
||||
&self.deps.owner_id
|
||||
}
|
||||
|
||||
/// Create a new agent.
|
||||
///
|
||||
/// Optionally accepts pre-created `ContextManager` and `SessionManager` for sharing
|
||||
@@ -264,6 +341,7 @@ impl Agent {
|
||||
));
|
||||
let repair_interval = self.config.repair_check_interval;
|
||||
let repair_channels = self.channels.clone();
|
||||
let repair_owner_id = self.owner_id().to_string();
|
||||
let repair_handle = tokio::spawn(async move {
|
||||
loop {
|
||||
tokio::time::sleep(repair_interval).await;
|
||||
@@ -311,7 +389,9 @@ impl Agent {
|
||||
|
||||
if let Some(msg) = notification {
|
||||
let response = OutgoingResponse::text(format!("Self-Repair: {}", msg));
|
||||
let _ = repair_channels.broadcast_all("default", response).await;
|
||||
let _ = repair_channels
|
||||
.broadcast_all(&repair_owner_id, response)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -325,7 +405,9 @@ impl Agent {
|
||||
"Self-Repair: Tool '{}' repaired: {}",
|
||||
tool.name, message
|
||||
));
|
||||
let _ = repair_channels.broadcast_all("default", response).await;
|
||||
let _ = repair_channels
|
||||
.broadcast_all(&repair_owner_id, response)
|
||||
.await;
|
||||
}
|
||||
Ok(result) => {
|
||||
tracing::info!("Tool repair result: {:?}", result);
|
||||
@@ -362,8 +444,12 @@ impl Agent {
|
||||
.timezone
|
||||
.clone()
|
||||
.or_else(|| Some(self.config.default_timezone.clone()));
|
||||
if let (Some(user), Some(channel)) =
|
||||
(&hb_config.notify_user, &hb_config.notify_channel)
|
||||
let heartbeat_notify_user = resolve_owner_scope_notification_user(
|
||||
hb_config.notify_user.as_deref(),
|
||||
Some(self.owner_id()),
|
||||
);
|
||||
if let Some(channel) = &hb_config.notify_channel
|
||||
&& let Some(user) = heartbeat_notify_user.as_deref()
|
||||
{
|
||||
config = config.with_notify(user, channel);
|
||||
}
|
||||
@@ -374,15 +460,22 @@ impl Agent {
|
||||
|
||||
// Spawn notification forwarder that routes through channel manager
|
||||
let notify_channel = hb_config.notify_channel.clone();
|
||||
let notify_user = hb_config.notify_user.clone();
|
||||
let notify_target = resolve_channel_notification_user(
|
||||
self.deps.extension_manager.as_ref(),
|
||||
hb_config.notify_channel.as_deref(),
|
||||
hb_config.notify_user.as_deref(),
|
||||
Some(self.owner_id()),
|
||||
)
|
||||
.await;
|
||||
let notify_user = heartbeat_notify_user;
|
||||
let channels = self.channels.clone();
|
||||
tokio::spawn(async move {
|
||||
while let Some(response) = notify_rx.recv().await {
|
||||
let user = notify_user.as_deref().unwrap_or("default");
|
||||
|
||||
// Try the configured channel first, fall back to
|
||||
// broadcasting on all channels.
|
||||
let targeted_ok = if let Some(ref channel) = notify_channel {
|
||||
let targeted_ok = if let Some(ref channel) = notify_channel
|
||||
&& let Some(ref user) = notify_target
|
||||
{
|
||||
channels
|
||||
.broadcast(channel, user, response.clone())
|
||||
.await
|
||||
@@ -391,7 +484,7 @@ impl Agent {
|
||||
false
|
||||
};
|
||||
|
||||
if !targeted_ok {
|
||||
if !targeted_ok && let Some(ref user) = notify_user {
|
||||
let results = channels.broadcast_all(user, response).await;
|
||||
for (ch, result) in results {
|
||||
if let Err(e) = result {
|
||||
@@ -460,32 +553,60 @@ impl Agent {
|
||||
|
||||
// Spawn notification forwarder (mirrors heartbeat pattern)
|
||||
let channels = self.channels.clone();
|
||||
let extension_manager = self.deps.extension_manager.clone();
|
||||
tokio::spawn(async move {
|
||||
while let Some(response) = notify_rx.recv().await {
|
||||
let user = response
|
||||
.metadata
|
||||
.get("notify_user")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("default")
|
||||
.to_string();
|
||||
let notify_channel = response
|
||||
.metadata
|
||||
.get("notify_channel")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
let fallback_user = resolve_owner_scope_notification_user(
|
||||
response
|
||||
.metadata
|
||||
.get("notify_user")
|
||||
.and_then(|v| v.as_str()),
|
||||
response.metadata.get("owner_id").and_then(|v| v.as_str()),
|
||||
);
|
||||
let Some(user) = resolve_routine_notification_target(
|
||||
extension_manager.as_ref(),
|
||||
&response.metadata,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
tracing::warn!(
|
||||
notify_channel = ?notify_channel,
|
||||
"Skipping routine notification with no explicit target or owner scope"
|
||||
);
|
||||
continue;
|
||||
};
|
||||
|
||||
// Try the configured channel first, fall back to
|
||||
// broadcasting on all channels.
|
||||
let targeted_ok = if let Some(ref channel) = notify_channel {
|
||||
channels
|
||||
.broadcast(channel, &user, response.clone())
|
||||
.await
|
||||
.is_ok()
|
||||
match channels.broadcast(channel, &user, response.clone()).await {
|
||||
Ok(()) => true,
|
||||
Err(e) => {
|
||||
let should_fallback =
|
||||
should_fallback_routine_notification(&e);
|
||||
tracing::warn!(
|
||||
channel = %channel,
|
||||
user = %user,
|
||||
error = %e,
|
||||
should_fallback,
|
||||
"Failed to send routine notification to configured channel"
|
||||
);
|
||||
if !should_fallback {
|
||||
continue;
|
||||
}
|
||||
false
|
||||
}
|
||||
}
|
||||
} else {
|
||||
false
|
||||
};
|
||||
|
||||
if !targeted_ok {
|
||||
if !targeted_ok && let Some(user) = fallback_user {
|
||||
let results = channels.broadcast_all(&user, response).await;
|
||||
for (ch, result) in results {
|
||||
if let Err(e) = result {
|
||||
@@ -572,6 +693,29 @@ impl Agent {
|
||||
// Store successfully extracted document text in workspace for indexing
|
||||
self.store_extracted_documents(&message).await;
|
||||
|
||||
// Event-triggered routines consume plain user input before it enters
|
||||
// the normal chat/tool pipeline. This avoids a duplicate turn where
|
||||
// the main agent responds and the routine also fires on the same
|
||||
// inbound message.
|
||||
if !message.is_internal
|
||||
&& matches!(
|
||||
SubmissionParser::parse(&message.content),
|
||||
Submission::UserInput { .. }
|
||||
)
|
||||
&& let Some(ref engine) = routine_engine_for_loop
|
||||
{
|
||||
let fired = engine.check_event_triggers(&message).await;
|
||||
if fired > 0 {
|
||||
tracing::debug!(
|
||||
channel = %message.channel,
|
||||
user = %message.user_id,
|
||||
fired,
|
||||
"Consumed inbound user message with matching event-triggered routine(s)"
|
||||
);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
match self.handle_message(&message).await {
|
||||
Ok(Some(response)) if !response.is_empty() => {
|
||||
// Hook: BeforeOutbound — allow hooks to modify or suppress outbound
|
||||
@@ -644,14 +788,6 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Check event triggers (cheap in-memory regex, fires async if matched)
|
||||
if let Some(ref engine) = routine_engine_for_loop {
|
||||
let fired = engine.check_event_triggers(&message).await;
|
||||
if fired > 0 {
|
||||
tracing::debug!("Fired {} event-triggered routines", fired);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Cleanup
|
||||
@@ -768,10 +904,7 @@ impl Agent {
|
||||
// For Signal, use signal_target from metadata (group:ID or phone number),
|
||||
// otherwise fall back to user_id
|
||||
let target = message
|
||||
.metadata
|
||||
.get("signal_target")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string())
|
||||
.routing_target()
|
||||
.unwrap_or_else(|| message.user_id.clone());
|
||||
self.tools()
|
||||
.set_message_tool_context(Some(message.channel.clone()), Some(target))
|
||||
@@ -811,7 +944,7 @@ impl Agent {
|
||||
}
|
||||
|
||||
// Hydrate thread from DB if it's a historical thread not in memory
|
||||
if let Some(ref external_thread_id) = message.thread_id {
|
||||
if let Some(external_thread_id) = message.conversation_scope() {
|
||||
tracing::trace!(
|
||||
message_id = %message.id,
|
||||
thread_id = %external_thread_id,
|
||||
@@ -832,7 +965,7 @@ impl Agent {
|
||||
.resolve_thread(
|
||||
&message.user_id,
|
||||
&message.channel,
|
||||
message.thread_id.as_deref(),
|
||||
message.conversation_scope(),
|
||||
)
|
||||
.await;
|
||||
tracing::debug!(
|
||||
@@ -985,7 +1118,11 @@ impl Agent {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::truncate_for_preview;
|
||||
use super::{
|
||||
resolve_routine_notification_user, should_fallback_routine_notification,
|
||||
truncate_for_preview,
|
||||
};
|
||||
use crate::error::ChannelError;
|
||||
|
||||
#[test]
|
||||
fn test_truncate_short_input() {
|
||||
@@ -1048,4 +1185,55 @@ mod tests {
|
||||
// 'h','e','l','l','o',' ','世','界' = 8 chars
|
||||
assert_eq!(result, "hello 世界...");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_routine_notification_user_prefers_explicit_target() {
|
||||
let metadata = serde_json::json!({
|
||||
"notify_user": "12345",
|
||||
"owner_id": "owner-scope",
|
||||
});
|
||||
|
||||
let resolved = resolve_routine_notification_user(&metadata);
|
||||
assert_eq!(resolved.as_deref(), Some("12345")); // safety: test-only assertion
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_routine_notification_user_falls_back_to_owner_scope() {
|
||||
let metadata = serde_json::json!({
|
||||
"notify_user": null,
|
||||
"owner_id": "owner-scope",
|
||||
});
|
||||
|
||||
let resolved = resolve_routine_notification_user(&metadata);
|
||||
assert_eq!(resolved.as_deref(), Some("owner-scope")); // safety: test-only assertion
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolve_routine_notification_user_rejects_missing_values() {
|
||||
let metadata = serde_json::json!({
|
||||
"notify_user": " ",
|
||||
});
|
||||
|
||||
assert_eq!(resolve_routine_notification_user(&metadata), None); // safety: test-only assertion
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn targeted_routine_notifications_do_not_fallback_without_owner_route() {
|
||||
let error = ChannelError::MissingRoutingTarget {
|
||||
name: "telegram".to_string(),
|
||||
reason: "No stored owner routing target for channel 'telegram'.".to_string(),
|
||||
};
|
||||
|
||||
assert!(!should_fallback_routine_notification(&error)); // safety: test-only assertion
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn targeted_routine_notifications_may_fallback_for_other_errors() {
|
||||
let error = ChannelError::SendFailed {
|
||||
name: "telegram".to_string(),
|
||||
reason: "timeout talking to channel".to_string(),
|
||||
};
|
||||
|
||||
assert!(should_fallback_routine_notification(&error)); // safety: test-only assertion
|
||||
}
|
||||
}
|
||||
|
||||
@@ -836,7 +836,10 @@ impl Agent {
|
||||
// 1. Persist to DB if available.
|
||||
if let Some(store) = self.store() {
|
||||
let value = serde_json::Value::String(model.to_string());
|
||||
if let Err(e) = store.set_setting("default", "selected_model", &value).await {
|
||||
if let Err(e) = store
|
||||
.set_setting(self.owner_id(), "selected_model", &value)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to persist model to DB: {}", e);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -140,13 +140,15 @@ impl Agent {
|
||||
|
||||
// Create a JobContext for tool execution (chat doesn't have a real job)
|
||||
let mut job_ctx =
|
||||
JobContext::with_user(&message.user_id, "chat", "Interactive chat session");
|
||||
JobContext::with_user(&message.user_id, "chat", "Interactive chat session")
|
||||
.with_requester_id(&message.sender_id);
|
||||
job_ctx.http_interceptor = self.deps.http_interceptor.clone();
|
||||
job_ctx.user_timezone = user_tz.name().to_string();
|
||||
job_ctx.metadata = serde_json::json!({
|
||||
"notify_channel": message.channel,
|
||||
"notify_user": message.user_id,
|
||||
"notify_thread_id": message.thread_id,
|
||||
"notify_metadata": message.metadata,
|
||||
});
|
||||
|
||||
// Build system prompts once for this turn. Two variants: with tools
|
||||
@@ -1175,6 +1177,7 @@ mod tests {
|
||||
/// Build a minimal `Agent` for unit testing (no DB, no workspace, no extensions).
|
||||
fn make_test_agent() -> Agent {
|
||||
let deps = AgentDeps {
|
||||
owner_id: "default".to_string(),
|
||||
store: None,
|
||||
llm: Arc::new(StaticLlmProvider),
|
||||
cheap_llm: None,
|
||||
@@ -2014,6 +2017,7 @@ mod tests {
|
||||
/// `max_tool_iterations` override.
|
||||
fn make_test_agent_with_llm(llm: Arc<dyn LlmProvider>, max_tool_iterations: usize) -> Agent {
|
||||
let deps = AgentDeps {
|
||||
owner_id: "default".to_string(),
|
||||
store: None,
|
||||
llm,
|
||||
cheap_llm: None,
|
||||
@@ -2127,6 +2131,7 @@ mod tests {
|
||||
let max_iter = 3;
|
||||
let agent = {
|
||||
let deps = AgentDeps {
|
||||
owner_id: "default".to_string(),
|
||||
store: None,
|
||||
llm,
|
||||
cheap_llm: None,
|
||||
|
||||
+144
-11
@@ -26,6 +26,8 @@
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use chrono::TimeZone as _;
|
||||
use chrono_tz::Tz;
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use crate::channels::OutgoingResponse;
|
||||
@@ -37,7 +39,7 @@ use crate::workspace::hygiene::HygieneConfig;
|
||||
/// Configuration for the heartbeat runner.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct HeartbeatConfig {
|
||||
/// Interval between heartbeat checks.
|
||||
/// Interval between heartbeat checks (used when fire_at is not set).
|
||||
pub interval: Duration,
|
||||
/// Whether heartbeat is enabled.
|
||||
pub enabled: bool,
|
||||
@@ -47,11 +49,13 @@ pub struct HeartbeatConfig {
|
||||
pub notify_user_id: Option<String>,
|
||||
/// Channel to notify on heartbeat findings.
|
||||
pub notify_channel: Option<String>,
|
||||
/// Fixed time-of-day to fire (24h). When set, interval is ignored.
|
||||
pub fire_at: Option<chrono::NaiveTime>,
|
||||
/// Hour (0-23) when quiet hours start.
|
||||
pub quiet_hours_start: Option<u32>,
|
||||
/// Hour (0-23) when quiet hours end.
|
||||
pub quiet_hours_end: Option<u32>,
|
||||
/// Timezone for quiet hours evaluation (IANA name).
|
||||
/// Timezone for fire_at and quiet hours evaluation (IANA name).
|
||||
pub timezone: Option<String>,
|
||||
}
|
||||
|
||||
@@ -63,6 +67,7 @@ impl Default for HeartbeatConfig {
|
||||
max_failures: 3,
|
||||
notify_user_id: None,
|
||||
notify_channel: None,
|
||||
fire_at: None,
|
||||
quiet_hours_start: None,
|
||||
quiet_hours_end: None,
|
||||
timezone: None,
|
||||
@@ -109,6 +114,21 @@ impl HeartbeatConfig {
|
||||
self.notify_channel = Some(channel.into());
|
||||
self
|
||||
}
|
||||
|
||||
/// Set a fixed time-of-day to fire (overrides interval).
|
||||
pub fn with_fire_at(mut self, time: chrono::NaiveTime, tz: Option<String>) -> Self {
|
||||
self.fire_at = Some(time);
|
||||
self.timezone = tz;
|
||||
self
|
||||
}
|
||||
|
||||
/// Resolve timezone string to chrono_tz::Tz (defaults to UTC).
|
||||
fn resolved_tz(&self) -> Tz {
|
||||
self.timezone
|
||||
.as_deref()
|
||||
.and_then(crate::timezone::parse_timezone)
|
||||
.unwrap_or(chrono_tz::UTC)
|
||||
}
|
||||
}
|
||||
|
||||
/// Result of a heartbeat check.
|
||||
@@ -124,6 +144,33 @@ pub enum HeartbeatResult {
|
||||
Failed(String),
|
||||
}
|
||||
|
||||
/// Compute how long to sleep until the next occurrence of `fire_at` in `tz`.
|
||||
///
|
||||
/// If the target time today is still in the future, sleep until then.
|
||||
/// Otherwise sleep until the same time tomorrow.
|
||||
fn duration_until_next_fire(fire_at: chrono::NaiveTime, tz: Tz) -> Duration {
|
||||
let now = chrono::Utc::now().with_timezone(&tz);
|
||||
let today = now.date_naive();
|
||||
|
||||
// Try to build today's target datetime in the given timezone.
|
||||
// `.earliest()` picks the first occurrence if DST creates ambiguity.
|
||||
let candidate = tz.from_local_datetime(&today.and_time(fire_at)).earliest();
|
||||
|
||||
let target = match candidate {
|
||||
Some(t) if t > now => t,
|
||||
_ => {
|
||||
// Already past (or ambiguous) — schedule for tomorrow
|
||||
let tomorrow = today + chrono::Duration::days(1);
|
||||
tz.from_local_datetime(&tomorrow.and_time(fire_at))
|
||||
.earliest()
|
||||
.unwrap_or_else(|| now + chrono::Duration::days(1))
|
||||
}
|
||||
};
|
||||
|
||||
let secs = (target - now).num_seconds().max(1) as u64;
|
||||
Duration::from_secs(secs)
|
||||
}
|
||||
|
||||
/// Heartbeat runner for proactive periodic execution.
|
||||
pub struct HeartbeatRunner {
|
||||
config: HeartbeatConfig,
|
||||
@@ -175,17 +222,39 @@ impl HeartbeatRunner {
|
||||
return;
|
||||
}
|
||||
|
||||
tracing::info!(
|
||||
"Starting heartbeat loop with interval {:?}",
|
||||
self.config.interval
|
||||
);
|
||||
// Two scheduling modes:
|
||||
// fire_at → sleep until the next occurrence (recalculated each iteration)
|
||||
// interval → tokio::time::interval (drift-free, accounts for loop body time)
|
||||
let mut tick_interval = if self.config.fire_at.is_none() {
|
||||
let mut iv = tokio::time::interval(self.config.interval);
|
||||
// Don't fire immediately on startup.
|
||||
iv.tick().await;
|
||||
Some(iv)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let mut interval = tokio::time::interval(self.config.interval);
|
||||
// Don't run immediately on startup
|
||||
interval.tick().await;
|
||||
if let Some(fire_at) = self.config.fire_at {
|
||||
tracing::info!(
|
||||
"Starting heartbeat loop: fire daily at {:?} {:?}",
|
||||
fire_at,
|
||||
self.config.timezone
|
||||
);
|
||||
} else {
|
||||
tracing::info!(
|
||||
"Starting heartbeat loop with interval {:?}",
|
||||
self.config.interval
|
||||
);
|
||||
}
|
||||
|
||||
loop {
|
||||
interval.tick().await;
|
||||
if let Some(fire_at) = self.config.fire_at {
|
||||
let sleep_dur = duration_until_next_fire(fire_at, self.config.resolved_tz());
|
||||
tracing::info!("Next heartbeat in {:.1}h", sleep_dur.as_secs_f64() / 3600.0);
|
||||
tokio::time::sleep(sleep_dur).await;
|
||||
} else if let Some(ref mut iv) = tick_interval {
|
||||
iv.tick().await;
|
||||
}
|
||||
|
||||
// Skip during quiet hours
|
||||
if self.config.is_quiet_hours() {
|
||||
@@ -333,7 +402,11 @@ impl HeartbeatRunner {
|
||||
return;
|
||||
};
|
||||
|
||||
let user_id = self.config.notify_user_id.as_deref().unwrap_or("default");
|
||||
let user_id = self
|
||||
.config
|
||||
.notify_user_id
|
||||
.as_deref()
|
||||
.unwrap_or_else(|| self.workspace.user_id());
|
||||
|
||||
// Persist to heartbeat conversation and get thread_id
|
||||
let thread_id = if let Some(ref store) = self.store {
|
||||
@@ -362,6 +435,7 @@ impl HeartbeatRunner {
|
||||
attachments: Vec::new(),
|
||||
metadata: serde_json::json!({
|
||||
"source": "heartbeat",
|
||||
"owner_id": self.workspace.user_id(),
|
||||
}),
|
||||
};
|
||||
|
||||
@@ -656,4 +730,63 @@ mod tests {
|
||||
) -> tokio::task::JoinHandle<()> = spawn_heartbeat;
|
||||
let _ = _fn_ptr;
|
||||
}
|
||||
|
||||
// ==================== fire_at scheduling ====================
|
||||
|
||||
#[test]
|
||||
fn test_default_config_has_no_fire_at() {
|
||||
let config = HeartbeatConfig::default();
|
||||
assert!(config.fire_at.is_none());
|
||||
// Interval-based scheduling should be the default
|
||||
assert_eq!(config.interval, Duration::from_secs(30 * 60));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_with_fire_at_builder() {
|
||||
let time = chrono::NaiveTime::from_hms_opt(9, 0, 0).unwrap();
|
||||
let config =
|
||||
HeartbeatConfig::default().with_fire_at(time, Some("Pacific/Auckland".to_string()));
|
||||
assert_eq!(config.fire_at, Some(time));
|
||||
assert_eq!(config.timezone, Some("Pacific/Auckland".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_duration_until_next_fire_is_bounded() {
|
||||
// Result must always be between 1 second and ~24 hours
|
||||
let time = chrono::NaiveTime::from_hms_opt(14, 0, 0).unwrap();
|
||||
let dur = duration_until_next_fire(time, chrono_tz::UTC);
|
||||
assert!(dur.as_secs() >= 1, "duration must be at least 1 second");
|
||||
assert!(
|
||||
dur.as_secs() <= 86_401,
|
||||
"duration must be at most ~24 hours, got {}s",
|
||||
dur.as_secs()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_duration_until_next_fire_dst_timezone_no_panic() {
|
||||
// Use a timezone with DST (US Eastern) — should never panic
|
||||
let tz: Tz = "America/New_York".parse().unwrap();
|
||||
// Test a range of times including midnight boundaries
|
||||
for hour in [0, 2, 3, 12, 23] {
|
||||
let time = chrono::NaiveTime::from_hms_opt(hour, 30, 0).unwrap();
|
||||
let dur = duration_until_next_fire(time, tz);
|
||||
assert!(dur.as_secs() >= 1);
|
||||
assert!(dur.as_secs() <= 86_401);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_resolved_tz_defaults_to_utc() {
|
||||
let config = HeartbeatConfig::default();
|
||||
assert_eq!(config.resolved_tz(), chrono_tz::UTC);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_resolved_tz_parses_iana() {
|
||||
let time = chrono::NaiveTime::from_hms_opt(9, 0, 0).unwrap();
|
||||
let config =
|
||||
HeartbeatConfig::default().with_fire_at(time, Some("Europe/London".to_string()));
|
||||
assert_eq!(config.resolved_tz(), chrono_tz::Europe::London);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -422,8 +422,8 @@ impl Default for RoutineGuardrails {
|
||||
pub struct NotifyConfig {
|
||||
/// Channel to notify on (None = default/broadcast all).
|
||||
pub channel: Option<String>,
|
||||
/// User to notify.
|
||||
pub user: String,
|
||||
/// Explicit target to notify. None means "resolve the owner's last-seen target".
|
||||
pub user: Option<String>,
|
||||
/// Notify when routine produces actionable output.
|
||||
pub on_attention: bool,
|
||||
/// Notify when routine errors.
|
||||
@@ -436,7 +436,7 @@ impl Default for NotifyConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
channel: None,
|
||||
user: "default".to_string(),
|
||||
user: None,
|
||||
on_attention: true,
|
||||
on_failure: true,
|
||||
on_success: false,
|
||||
|
||||
@@ -172,6 +172,11 @@ impl RoutineEngine {
|
||||
EventMatcher::Message { routine, regex } => (routine, regex),
|
||||
EventMatcher::System { .. } => continue,
|
||||
};
|
||||
|
||||
if routine.user_id != message.user_id {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Channel filter
|
||||
if let Trigger::Event {
|
||||
channel: Some(ch), ..
|
||||
@@ -650,6 +655,7 @@ async fn execute_routine(ctx: EngineContext, routine: Routine, run: RoutineRun)
|
||||
send_notification(
|
||||
&ctx.notify_tx,
|
||||
&routine.notify,
|
||||
&routine.user_id,
|
||||
&routine.name,
|
||||
status,
|
||||
summary.as_deref(),
|
||||
@@ -694,7 +700,8 @@ async fn execute_full_job(
|
||||
reason: "scheduler not available".to_string(),
|
||||
})?;
|
||||
|
||||
let mut metadata = serde_json::json!({ "max_iterations": max_iterations });
|
||||
let mut metadata =
|
||||
serde_json::json!({ "max_iterations": max_iterations, "owner_id": routine.user_id });
|
||||
// Carry the routine's notify config in job metadata so the message tool
|
||||
// can resolve channel/target per-job without global state mutation.
|
||||
if let Some(channel) = &routine.notify.channel {
|
||||
@@ -1207,6 +1214,7 @@ async fn execute_routine_tool(
|
||||
async fn send_notification(
|
||||
tx: &mpsc::Sender<OutgoingResponse>,
|
||||
notify: &NotifyConfig,
|
||||
owner_id: &str,
|
||||
routine_name: &str,
|
||||
status: RunStatus,
|
||||
summary: Option<&str>,
|
||||
@@ -1243,6 +1251,7 @@ async fn send_notification(
|
||||
"source": "routine",
|
||||
"routine_name": routine_name,
|
||||
"status": status.to_string(),
|
||||
"owner_id": owner_id,
|
||||
"notify_user": notify.user,
|
||||
"notify_channel": notify.channel,
|
||||
}),
|
||||
|
||||
@@ -427,6 +427,14 @@ impl SubmissionResult {
|
||||
message: message.into(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Create a non-error status message (e.g., for blocking states like approval waiting).
|
||||
/// Uses Ok variant to avoid "Error:" prefix in rendering.
|
||||
pub fn pending(message: impl Into<String>) -> Self {
|
||||
Self::Ok {
|
||||
message: Some(message.into()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
||||
+115
-6
@@ -187,13 +187,18 @@ impl Agent {
|
||||
);
|
||||
|
||||
// First check thread state without holding lock during I/O
|
||||
let thread_state = {
|
||||
let (thread_state, approval_context) = {
|
||||
let sess = session.lock().await;
|
||||
let thread = sess
|
||||
.threads
|
||||
.get(&thread_id)
|
||||
.ok_or_else(|| Error::from(crate::error::JobError::NotFound { id: thread_id }))?;
|
||||
thread.state
|
||||
let approval_context = thread.pending_approval.as_ref().map(|a| {
|
||||
let desc_preview =
|
||||
crate::agent::agent_loop::truncate_for_preview(&a.description, 80);
|
||||
(a.tool_name.clone(), desc_preview)
|
||||
});
|
||||
(thread.state, approval_context)
|
||||
};
|
||||
|
||||
tracing::debug!(
|
||||
@@ -221,9 +226,13 @@ impl Agent {
|
||||
thread_id = %thread_id,
|
||||
"Thread awaiting approval, rejecting new input"
|
||||
);
|
||||
return Ok(SubmissionResult::error(
|
||||
"Waiting for approval. Use /interrupt to cancel.",
|
||||
));
|
||||
let msg = match approval_context {
|
||||
Some((tool_name, desc_preview)) => format!(
|
||||
"Waiting for approval: {tool_name} — {desc_preview}. Use /interrupt to cancel."
|
||||
),
|
||||
None => "Waiting for approval. Use /interrupt to cancel.".to_string(),
|
||||
};
|
||||
return Ok(SubmissionResult::pending(msg));
|
||||
}
|
||||
ThreadState::Completed => {
|
||||
tracing::warn!(
|
||||
@@ -924,7 +933,8 @@ impl Agent {
|
||||
|
||||
// Execute the approved tool and continue the loop
|
||||
let mut job_ctx =
|
||||
JobContext::with_user(&message.user_id, "chat", "Interactive chat session");
|
||||
JobContext::with_user(&message.user_id, "chat", "Interactive chat session")
|
||||
.with_requester_id(&message.sender_id);
|
||||
job_ctx.http_interceptor = self.deps.http_interceptor.clone();
|
||||
// Prefer a valid timezone from the approval message, fall back to the
|
||||
// resolved timezone stored when the approval was originally requested.
|
||||
@@ -1916,4 +1926,103 @@ mod tests {
|
||||
created_at: chrono::Utc::now(),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_awaiting_approval_rejection_includes_tool_context() {
|
||||
// Test that when a thread is in AwaitingApproval state and receives a new message,
|
||||
// process_user_input rejects it with a non-error status that includes tool context.
|
||||
use crate::agent::session::{PendingApproval, Session, Thread, ThreadState};
|
||||
use uuid::Uuid;
|
||||
|
||||
let session_id = Uuid::new_v4();
|
||||
let thread_id = Uuid::new_v4();
|
||||
let mut thread = Thread::with_id(thread_id, session_id);
|
||||
|
||||
// Set thread to AwaitingApproval with a pending tool approval
|
||||
let pending = PendingApproval {
|
||||
request_id: Uuid::new_v4(),
|
||||
tool_name: "shell".to_string(),
|
||||
parameters: serde_json::json!({"command": "echo hello"}),
|
||||
display_parameters: serde_json::json!({"command": "[REDACTED]"}),
|
||||
description: "Execute: echo hello".to_string(),
|
||||
tool_call_id: "call_0".to_string(),
|
||||
context_messages: vec![],
|
||||
deferred_tool_calls: vec![],
|
||||
user_timezone: None,
|
||||
};
|
||||
thread.await_approval(pending);
|
||||
|
||||
let mut session = Session::new("test-user");
|
||||
session.threads.insert(thread_id, thread);
|
||||
|
||||
// Verify thread is in AwaitingApproval state
|
||||
assert_eq!(
|
||||
session.threads[&thread_id].state,
|
||||
ThreadState::AwaitingApproval
|
||||
);
|
||||
|
||||
let result = extract_approval_message(&session, thread_id);
|
||||
|
||||
// Verify result is an Ok with a message (not an Error)
|
||||
match result {
|
||||
Ok(Some(msg)) => {
|
||||
// Should NOT start with "Error:"
|
||||
assert!(
|
||||
!msg.to_lowercase().starts_with("error:"),
|
||||
"Approval rejection should not have 'Error:' prefix. Got: {}",
|
||||
msg
|
||||
);
|
||||
|
||||
// Should contain "waiting for approval"
|
||||
assert!(
|
||||
msg.to_lowercase().contains("waiting for approval"),
|
||||
"Should contain 'waiting for approval'. Got: {}",
|
||||
msg
|
||||
);
|
||||
|
||||
// Should contain the tool name
|
||||
assert!(
|
||||
msg.contains("shell"),
|
||||
"Should contain tool name 'shell'. Got: {}",
|
||||
msg
|
||||
);
|
||||
|
||||
// Should contain the description (or truncated version)
|
||||
assert!(
|
||||
msg.contains("echo hello"),
|
||||
"Should contain description 'echo hello'. Got: {}",
|
||||
msg
|
||||
);
|
||||
}
|
||||
_ => panic!("Expected approval rejection message"),
|
||||
}
|
||||
}
|
||||
|
||||
// Helper function to extract the approval message without needing a full Agent instance
|
||||
fn extract_approval_message(
|
||||
session: &crate::agent::session::Session,
|
||||
thread_id: Uuid,
|
||||
) -> Result<Option<String>, crate::error::Error> {
|
||||
let thread = session.threads.get(&thread_id).ok_or_else(|| {
|
||||
crate::error::Error::from(crate::error::JobError::NotFound { id: thread_id })
|
||||
})?;
|
||||
|
||||
if thread.state == ThreadState::AwaitingApproval {
|
||||
let approval_context = thread.pending_approval.as_ref().map(|a| {
|
||||
let desc_preview =
|
||||
crate::agent::agent_loop::truncate_for_preview(&a.description, 80);
|
||||
(a.tool_name.clone(), desc_preview)
|
||||
});
|
||||
|
||||
let msg = match approval_context {
|
||||
Some((tool_name, desc_preview)) => format!(
|
||||
"Waiting for approval: {tool_name} — {desc_preview}. Use /interrupt to cancel."
|
||||
),
|
||||
None => "Waiting for approval. Use /interrupt to cancel.".to_string(),
|
||||
};
|
||||
Ok(Some(msg))
|
||||
} else {
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+19
-10
@@ -140,12 +140,14 @@ impl AppBuilder {
|
||||
self.handles = Some(handles);
|
||||
|
||||
// Post-init: migrate disk config, reload config from DB, attach session, cleanup
|
||||
if let Err(e) = crate::bootstrap::migrate_disk_to_db(db.as_ref(), "default").await {
|
||||
if let Err(e) =
|
||||
crate::bootstrap::migrate_disk_to_db(db.as_ref(), &self.config.owner_id).await
|
||||
{
|
||||
tracing::warn!("Disk-to-DB settings migration failed: {}", e);
|
||||
}
|
||||
|
||||
let toml_path = self.toml_path.as_deref();
|
||||
match Config::from_db_with_toml(db.as_ref(), "default", toml_path).await {
|
||||
match Config::from_db_with_toml(db.as_ref(), &self.config.owner_id, toml_path).await {
|
||||
Ok(db_config) => {
|
||||
self.config = db_config;
|
||||
tracing::debug!("Configuration reloaded from database");
|
||||
@@ -158,7 +160,9 @@ impl AppBuilder {
|
||||
}
|
||||
}
|
||||
|
||||
self.session.attach_store(db.clone(), "default").await;
|
||||
self.session
|
||||
.attach_store(db.clone(), &self.config.owner_id)
|
||||
.await;
|
||||
|
||||
// Fire-and-forget housekeeping — no need to block startup.
|
||||
let db_cleanup = db.clone();
|
||||
@@ -193,9 +197,10 @@ impl AppBuilder {
|
||||
let store: Option<&(dyn crate::db::SettingsStore + Sync)> =
|
||||
self.db.as_ref().map(|db| db.as_ref() as _);
|
||||
let toml_path = self.toml_path.as_deref();
|
||||
let owner_id = self.config.owner_id.clone();
|
||||
if let Err(e) = self
|
||||
.config
|
||||
.re_resolve_llm(store, "default", toml_path)
|
||||
.re_resolve_llm(store, &owner_id, toml_path)
|
||||
.await
|
||||
{
|
||||
tracing::warn!(
|
||||
@@ -224,15 +229,17 @@ impl AppBuilder {
|
||||
|
||||
if let Some(ref secrets) = store {
|
||||
// Inject LLM API keys from encrypted storage
|
||||
crate::config::inject_llm_keys_from_secrets(secrets.as_ref(), "default").await;
|
||||
crate::config::inject_llm_keys_from_secrets(secrets.as_ref(), &self.config.owner_id)
|
||||
.await;
|
||||
|
||||
// Re-resolve only the LLM config with newly available keys.
|
||||
let store: Option<&(dyn crate::db::SettingsStore + Sync)> =
|
||||
self.db.as_ref().map(|db| db.as_ref() as _);
|
||||
let toml_path = self.toml_path.as_deref();
|
||||
let owner_id = self.config.owner_id.clone();
|
||||
if let Err(e) = self
|
||||
.config
|
||||
.re_resolve_llm(store, "default", toml_path)
|
||||
.re_resolve_llm(store, &owner_id, toml_path)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to re-resolve LLM config after secret injection: {e}");
|
||||
@@ -304,7 +311,7 @@ impl AppBuilder {
|
||||
|
||||
// Register memory tools if database is available
|
||||
let workspace = if let Some(ref db) = self.db {
|
||||
let mut ws = Workspace::new_with_db("default", db.clone())
|
||||
let mut ws = Workspace::new_with_db(&self.config.owner_id, db.clone())
|
||||
.with_search_config(&self.config.search);
|
||||
if let Some(ref emb) = embeddings {
|
||||
ws = ws.with_embeddings(emb.clone());
|
||||
@@ -469,9 +476,10 @@ impl AppBuilder {
|
||||
let tools = Arc::clone(tools);
|
||||
let mcp_sm = Arc::clone(&mcp_session_manager);
|
||||
let pm = Arc::clone(&mcp_process_manager);
|
||||
let owner_id = self.config.owner_id.clone();
|
||||
async move {
|
||||
let servers_result = if let Some(ref d) = db {
|
||||
load_mcp_servers_from_db(d.as_ref(), "default").await
|
||||
load_mcp_servers_from_db(d.as_ref(), &owner_id).await
|
||||
} else {
|
||||
crate::tools::mcp::config::load_mcp_servers().await
|
||||
};
|
||||
@@ -491,6 +499,7 @@ impl AppBuilder {
|
||||
let secrets = secrets_store.clone();
|
||||
let tools = Arc::clone(&tools);
|
||||
let pm = Arc::clone(&pm);
|
||||
let owner_id = owner_id.clone();
|
||||
|
||||
join_set.spawn(async move {
|
||||
let server_name = server.name.clone();
|
||||
@@ -500,7 +509,7 @@ impl AppBuilder {
|
||||
&mcp_sm,
|
||||
&pm,
|
||||
secrets,
|
||||
"default",
|
||||
&owner_id,
|
||||
)
|
||||
.await
|
||||
{
|
||||
@@ -642,7 +651,7 @@ impl AppBuilder {
|
||||
self.config.wasm.tools_dir.clone(),
|
||||
self.config.channels.wasm_channels_dir.clone(),
|
||||
self.config.tunnel.public_url.clone(),
|
||||
"default".to_string(),
|
||||
self.config.owner_id.clone(),
|
||||
self.db.clone(),
|
||||
catalog_entries.clone(),
|
||||
));
|
||||
|
||||
+82
-6
@@ -67,14 +67,24 @@ pub struct IncomingMessage {
|
||||
pub id: Uuid,
|
||||
/// Channel this message came from.
|
||||
pub channel: String,
|
||||
/// User identifier within the channel.
|
||||
/// Storage/persistence scope for this interaction.
|
||||
///
|
||||
/// For owner-capable channels this is the stable instance owner ID when the
|
||||
/// configured owner is speaking; otherwise it can be a guest/sender-scoped
|
||||
/// identifier to preserve isolation.
|
||||
pub user_id: String,
|
||||
/// Stable instance owner scope for this IronClaw deployment.
|
||||
pub owner_id: String,
|
||||
/// Channel-specific sender/actor identifier.
|
||||
pub sender_id: String,
|
||||
/// Optional display name.
|
||||
pub user_name: Option<String>,
|
||||
/// Message content.
|
||||
pub content: String,
|
||||
/// Thread/conversation ID for threaded conversations.
|
||||
pub thread_id: Option<String>,
|
||||
/// Stable channel/chat/thread scope for this conversation.
|
||||
pub conversation_scope_id: Option<String>,
|
||||
/// When the message was received.
|
||||
pub received_at: DateTime<Utc>,
|
||||
/// Channel-specific metadata.
|
||||
@@ -84,9 +94,8 @@ pub struct IncomingMessage {
|
||||
/// File or media attachments on this message.
|
||||
pub attachments: Vec<IncomingAttachment>,
|
||||
/// Internal-only flag: message was generated inside the process (e.g. job
|
||||
/// monitor) and must bypass the normal user-input pipeline. This field is
|
||||
/// **not** settable via `with_metadata()` — only trusted code paths inside
|
||||
/// the binary can set it, preventing external channels from spoofing it.
|
||||
/// monitor) and must bypass the normal user-input pipeline. This field is
|
||||
/// not settable via metadata, so external channels cannot spoof it.
|
||||
pub(crate) is_internal: bool,
|
||||
}
|
||||
|
||||
@@ -97,13 +106,17 @@ impl IncomingMessage {
|
||||
user_id: impl Into<String>,
|
||||
content: impl Into<String>,
|
||||
) -> Self {
|
||||
let user_id = user_id.into();
|
||||
Self {
|
||||
id: Uuid::new_v4(),
|
||||
channel: channel.into(),
|
||||
user_id: user_id.into(),
|
||||
owner_id: user_id.clone(),
|
||||
sender_id: user_id.clone(),
|
||||
user_id,
|
||||
user_name: None,
|
||||
content: content.into(),
|
||||
thread_id: None,
|
||||
conversation_scope_id: None,
|
||||
received_at: Utc::now(),
|
||||
metadata: serde_json::Value::Null,
|
||||
timezone: None,
|
||||
@@ -114,7 +127,27 @@ impl IncomingMessage {
|
||||
|
||||
/// Set the thread ID.
|
||||
pub fn with_thread(mut self, thread_id: impl Into<String>) -> Self {
|
||||
self.thread_id = Some(thread_id.into());
|
||||
let thread_id = thread_id.into();
|
||||
self.conversation_scope_id = Some(thread_id.clone());
|
||||
self.thread_id = Some(thread_id);
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the stable owner scope for this message.
|
||||
pub fn with_owner_id(mut self, owner_id: impl Into<String>) -> Self {
|
||||
self.owner_id = owner_id.into();
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the channel-specific sender/actor identifier.
|
||||
pub fn with_sender_id(mut self, sender_id: impl Into<String>) -> Self {
|
||||
self.sender_id = sender_id.into();
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the conversation scope for this message.
|
||||
pub fn with_conversation_scope(mut self, scope_id: impl Into<String>) -> Self {
|
||||
self.conversation_scope_id = Some(scope_id.into());
|
||||
self
|
||||
}
|
||||
|
||||
@@ -147,6 +180,49 @@ impl IncomingMessage {
|
||||
self.is_internal = true;
|
||||
self
|
||||
}
|
||||
|
||||
/// Effective conversation scope, falling back to thread_id for legacy callers.
|
||||
pub fn conversation_scope(&self) -> Option<&str> {
|
||||
self.conversation_scope_id
|
||||
.as_deref()
|
||||
.or(self.thread_id.as_deref())
|
||||
}
|
||||
|
||||
/// Best-effort routing target for proactive replies on the current channel.
|
||||
pub fn routing_target(&self) -> Option<String> {
|
||||
routing_target_from_metadata(&self.metadata).or_else(|| {
|
||||
if self.sender_id.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(self.sender_id.clone())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Extract a channel-specific proactive routing target from message metadata.
|
||||
pub fn routing_target_from_metadata(metadata: &serde_json::Value) -> Option<String> {
|
||||
metadata
|
||||
.get("signal_target")
|
||||
.and_then(|value| match value {
|
||||
serde_json::Value::String(s) => Some(s.clone()),
|
||||
serde_json::Value::Number(n) => Some(n.to_string()),
|
||||
_ => None,
|
||||
})
|
||||
.or_else(|| {
|
||||
metadata.get("chat_id").and_then(|value| match value {
|
||||
serde_json::Value::String(s) => Some(s.clone()),
|
||||
serde_json::Value::Number(n) => Some(n.to_string()),
|
||||
_ => None,
|
||||
})
|
||||
})
|
||||
.or_else(|| {
|
||||
metadata.get("target").and_then(|value| match value {
|
||||
serde_json::Value::String(s) => Some(s.clone()),
|
||||
serde_json::Value::Number(n) => Some(n.to_string()),
|
||||
_ => None,
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
/// Stream of incoming messages.
|
||||
|
||||
+105
-11
@@ -133,7 +133,8 @@ impl HttpChannel {
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct WebhookRequest {
|
||||
/// User or client identifier (ignored, user is fixed by server config).
|
||||
/// Optional caller or client identifier for sender-scoped routing.
|
||||
/// The channel owner/storage scope remains fixed by server config.
|
||||
#[serde(default)]
|
||||
user_id: Option<String>,
|
||||
/// Message content.
|
||||
@@ -403,12 +404,38 @@ async fn process_authenticated_request(
|
||||
state: Arc<HttpChannelState>,
|
||||
req: WebhookRequest,
|
||||
) -> axum::response::Response {
|
||||
let _ = req.user_id.as_ref().map(|user_id| {
|
||||
tracing::debug!(
|
||||
provided_user_id = %user_id,
|
||||
"HTTP webhook request provided user_id, ignoring in favor of configured user_id"
|
||||
);
|
||||
});
|
||||
let normalized_user_id = req
|
||||
.user_id
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|user_id| !user_id.is_empty());
|
||||
|
||||
match (req.user_id.as_deref(), normalized_user_id) {
|
||||
(Some(raw_user_id), Some(user_id)) if raw_user_id != user_id => {
|
||||
tracing::debug!(
|
||||
provided_user_id = %raw_user_id,
|
||||
normalized_sender_id = %user_id,
|
||||
configured_owner_id = %state.user_id,
|
||||
"HTTP webhook request provided user_id; trimming and using it as sender_id while keeping the configured owner scope"
|
||||
);
|
||||
}
|
||||
(Some(user_id), Some(_)) => {
|
||||
tracing::debug!(
|
||||
provided_user_id = %user_id,
|
||||
configured_owner_id = %state.user_id,
|
||||
"HTTP webhook request provided user_id; using it as sender_id while keeping the configured owner scope"
|
||||
);
|
||||
}
|
||||
(Some(raw_user_id), None) => {
|
||||
tracing::debug!(
|
||||
provided_user_id = %raw_user_id,
|
||||
configured_owner_id = %state.user_id,
|
||||
"HTTP webhook request provided a blank user_id; falling back to the configured owner scope for sender_id"
|
||||
);
|
||||
}
|
||||
(None, None) => {}
|
||||
(None, Some(_)) => unreachable!("normalized user_id requires a raw user_id"),
|
||||
}
|
||||
|
||||
if req.content.len() > MAX_CONTENT_BYTES {
|
||||
return (
|
||||
@@ -514,11 +541,13 @@ async fn process_authenticated_request(
|
||||
Vec::new()
|
||||
};
|
||||
|
||||
let mut msg = IncomingMessage::new("http", &state.user_id, &req.content).with_metadata(
|
||||
serde_json::json!({
|
||||
let sender_id = normalized_user_id.unwrap_or(&state.user_id).to_string();
|
||||
let mut msg = IncomingMessage::new("http", &state.user_id, &req.content)
|
||||
.with_owner_id(&state.user_id)
|
||||
.with_sender_id(sender_id)
|
||||
.with_metadata(serde_json::json!({
|
||||
"wait_for_response": wait_for_response,
|
||||
}),
|
||||
);
|
||||
}));
|
||||
|
||||
if !attachments.is_empty() {
|
||||
msg = msg.with_attachments(attachments);
|
||||
@@ -682,6 +711,7 @@ mod tests {
|
||||
use axum::body::Body;
|
||||
use axum::http::{HeaderValue, Request};
|
||||
use secrecy::SecretString;
|
||||
use tokio_stream::StreamExt;
|
||||
use tower::ServiceExt;
|
||||
|
||||
use super::*;
|
||||
@@ -820,6 +850,70 @@ mod tests {
|
||||
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn webhook_blank_user_id_falls_back_to_owner_scope() {
|
||||
let secret = "test-secret-123";
|
||||
let channel = test_channel(Some(secret));
|
||||
let mut stream = channel.start().await.unwrap();
|
||||
let app = channel.routes();
|
||||
|
||||
let body = serde_json::json!({
|
||||
"content": "hello",
|
||||
"user_id": " "
|
||||
});
|
||||
let body_bytes = serde_json::to_vec(&body).unwrap();
|
||||
let signature = compute_signature(secret, &body_bytes);
|
||||
let req = Request::builder()
|
||||
.method("POST")
|
||||
.uri("/webhook")
|
||||
.header("content-type", "application/json")
|
||||
.header("x-hub-signature-256", signature)
|
||||
.body(Body::from(body_bytes))
|
||||
.unwrap();
|
||||
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let msg = tokio::time::timeout(std::time::Duration::from_secs(1), stream.next())
|
||||
.await
|
||||
.expect("timed out waiting for webhook message")
|
||||
.expect("stream should yield a webhook message");
|
||||
assert_eq!(msg.sender_id, "http");
|
||||
assert_eq!(msg.owner_id, "http");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn webhook_user_id_is_trimmed_before_becoming_sender_id() {
|
||||
let secret = "test-secret-123";
|
||||
let channel = test_channel(Some(secret));
|
||||
let mut stream = channel.start().await.unwrap();
|
||||
let app = channel.routes();
|
||||
|
||||
let body = serde_json::json!({
|
||||
"content": "hello",
|
||||
"user_id": " alice "
|
||||
});
|
||||
let body_bytes = serde_json::to_vec(&body).unwrap();
|
||||
let signature = compute_signature(secret, &body_bytes);
|
||||
let req = Request::builder()
|
||||
.method("POST")
|
||||
.uri("/webhook")
|
||||
.header("content-type", "application/json")
|
||||
.header("x-hub-signature-256", signature)
|
||||
.body(Body::from(body_bytes))
|
||||
.unwrap();
|
||||
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let msg = tokio::time::timeout(std::time::Duration::from_secs(1), stream.next())
|
||||
.await
|
||||
.expect("timed out waiting for webhook message")
|
||||
.expect("stream should yield a webhook message");
|
||||
assert_eq!(msg.sender_id, "alice");
|
||||
assert_eq!(msg.owner_id, "http");
|
||||
}
|
||||
|
||||
/// Regression test for issue #869: RwLock read guard was held across
|
||||
/// tx.send(msg).await in `process_message()`, blocking shutdown() from
|
||||
/// acquiring the write lock when the channel buffer was full.
|
||||
|
||||
+1
-1
@@ -39,7 +39,7 @@ mod webhook_server;
|
||||
|
||||
pub use channel::{
|
||||
AttachmentKind, Channel, ChannelSecretUpdater, IncomingAttachment, IncomingMessage,
|
||||
MessageStream, OutgoingResponse, StatusUpdate,
|
||||
MessageStream, OutgoingResponse, StatusUpdate, routing_target_from_metadata,
|
||||
};
|
||||
pub use http::{HttpChannel, HttpChannelState};
|
||||
pub use manager::ChannelManager;
|
||||
|
||||
+22
-7
@@ -200,6 +200,8 @@ fn format_json_params(params: &serde_json::Value, indent: &str) -> String {
|
||||
|
||||
/// REPL channel with line editing and markdown rendering.
|
||||
pub struct ReplChannel {
|
||||
/// Stable owner scope for this REPL instance.
|
||||
user_id: String,
|
||||
/// Optional single message to send (for -m flag).
|
||||
single_message: Option<String>,
|
||||
/// Debug mode flag (shared with input thread).
|
||||
@@ -213,7 +215,13 @@ pub struct ReplChannel {
|
||||
impl ReplChannel {
|
||||
/// Create a new REPL channel.
|
||||
pub fn new() -> Self {
|
||||
Self::with_user_id("default")
|
||||
}
|
||||
|
||||
/// Create a new REPL channel for a specific owner scope.
|
||||
pub fn with_user_id(user_id: impl Into<String>) -> Self {
|
||||
Self {
|
||||
user_id: user_id.into(),
|
||||
single_message: None,
|
||||
debug_mode: Arc::new(AtomicBool::new(false)),
|
||||
is_streaming: Arc::new(AtomicBool::new(false)),
|
||||
@@ -223,7 +231,13 @@ impl ReplChannel {
|
||||
|
||||
/// Create a REPL channel that sends a single message and exits.
|
||||
pub fn with_message(message: String) -> Self {
|
||||
Self::with_message_for_user("default", message)
|
||||
}
|
||||
|
||||
/// Create a REPL channel that sends a single message for a specific owner scope and exits.
|
||||
pub fn with_message_for_user(user_id: impl Into<String>, message: String) -> Self {
|
||||
Self {
|
||||
user_id: user_id.into(),
|
||||
single_message: Some(message),
|
||||
debug_mode: Arc::new(AtomicBool::new(false)),
|
||||
is_streaming: Arc::new(AtomicBool::new(false)),
|
||||
@@ -292,6 +306,7 @@ impl Channel for ReplChannel {
|
||||
async fn start(&self) -> Result<MessageStream, ChannelError> {
|
||||
let (tx, rx) = mpsc::channel(32);
|
||||
let single_message = self.single_message.clone();
|
||||
let user_id = self.user_id.clone();
|
||||
let debug_mode = Arc::clone(&self.debug_mode);
|
||||
let suppress_banner = Arc::clone(&self.suppress_banner);
|
||||
let esc_interrupt_triggered_for_thread = Arc::new(AtomicBool::new(false));
|
||||
@@ -301,11 +316,11 @@ impl Channel for ReplChannel {
|
||||
|
||||
// Single message mode: send it and return
|
||||
if let Some(msg) = single_message {
|
||||
let incoming = IncomingMessage::new("repl", "default", &msg).with_timezone(&sys_tz);
|
||||
let incoming = IncomingMessage::new("repl", &user_id, &msg).with_timezone(&sys_tz);
|
||||
let _ = tx.blocking_send(incoming);
|
||||
// Ensure the agent exits after handling exactly one turn in -m mode,
|
||||
// even when other channels (gateway/http) are enabled.
|
||||
let _ = tx.blocking_send(IncomingMessage::new("repl", "default", "/quit"));
|
||||
let _ = tx.blocking_send(IncomingMessage::new("repl", &user_id, "/quit"));
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -366,7 +381,7 @@ impl Channel for ReplChannel {
|
||||
"/quit" | "/exit" => {
|
||||
// Forward shutdown command so the agent loop exits even
|
||||
// when other channels (e.g. web gateway) are still active.
|
||||
let msg = IncomingMessage::new("repl", "default", "/quit")
|
||||
let msg = IncomingMessage::new("repl", &user_id, "/quit")
|
||||
.with_timezone(&sys_tz);
|
||||
let _ = tx.blocking_send(msg);
|
||||
break;
|
||||
@@ -389,7 +404,7 @@ impl Channel for ReplChannel {
|
||||
}
|
||||
|
||||
let msg =
|
||||
IncomingMessage::new("repl", "default", line).with_timezone(&sys_tz);
|
||||
IncomingMessage::new("repl", &user_id, line).with_timezone(&sys_tz);
|
||||
if tx.blocking_send(msg).is_err() {
|
||||
break;
|
||||
}
|
||||
@@ -397,14 +412,14 @@ impl Channel for ReplChannel {
|
||||
Err(ReadlineError::Interrupted) => {
|
||||
if esc_interrupt_triggered_for_thread.swap(false, Ordering::Relaxed) {
|
||||
// Esc: interrupt current operation and keep REPL open.
|
||||
let msg = IncomingMessage::new("repl", "default", "/interrupt")
|
||||
let msg = IncomingMessage::new("repl", &user_id, "/interrupt")
|
||||
.with_timezone(&sys_tz);
|
||||
if tx.blocking_send(msg).is_err() {
|
||||
break;
|
||||
}
|
||||
} else {
|
||||
// Ctrl+C (VINTR): request graceful shutdown.
|
||||
let msg = IncomingMessage::new("repl", "default", "/quit")
|
||||
let msg = IncomingMessage::new("repl", &user_id, "/quit")
|
||||
.with_timezone(&sys_tz);
|
||||
let _ = tx.blocking_send(msg);
|
||||
break;
|
||||
@@ -416,7 +431,7 @@ impl Channel for ReplChannel {
|
||||
// immediately — just drop the REPL thread silently so other
|
||||
// channels (gateway, telegram, …) keep running.
|
||||
if std::io::stdin().is_terminal() {
|
||||
let msg = IncomingMessage::new("repl", "default", "/quit")
|
||||
let msg = IncomingMessage::new("repl", &user_id, "/quit")
|
||||
.with_timezone(&sys_tz);
|
||||
let _ = tx.blocking_send(msg);
|
||||
}
|
||||
|
||||
@@ -27,6 +27,7 @@ pub struct WasmChannelLoader {
|
||||
pairing_store: Arc<PairingStore>,
|
||||
settings_store: Option<Arc<dyn SettingsStore>>,
|
||||
secrets_store: Option<Arc<dyn SecretsStore + Send + Sync>>,
|
||||
owner_scope_id: String,
|
||||
}
|
||||
|
||||
impl WasmChannelLoader {
|
||||
@@ -35,12 +36,14 @@ impl WasmChannelLoader {
|
||||
runtime: Arc<WasmChannelRuntime>,
|
||||
pairing_store: Arc<PairingStore>,
|
||||
settings_store: Option<Arc<dyn SettingsStore>>,
|
||||
owner_scope_id: impl Into<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
runtime,
|
||||
pairing_store,
|
||||
settings_store,
|
||||
secrets_store: None,
|
||||
owner_scope_id: owner_scope_id.into(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -149,6 +152,7 @@ impl WasmChannelLoader {
|
||||
self.runtime.clone(),
|
||||
prepared,
|
||||
capabilities,
|
||||
self.owner_scope_id.clone(),
|
||||
config_json,
|
||||
self.pairing_store.clone(),
|
||||
self.settings_store.clone(),
|
||||
@@ -487,7 +491,8 @@ mod tests {
|
||||
async fn test_loader_invalid_name() {
|
||||
let config = WasmChannelRuntimeConfig::for_testing();
|
||||
let runtime = Arc::new(WasmChannelRuntime::new(config).unwrap());
|
||||
let loader = WasmChannelLoader::new(runtime, Arc::new(PairingStore::new()), None);
|
||||
let loader =
|
||||
WasmChannelLoader::new(runtime, Arc::new(PairingStore::new()), None, "default");
|
||||
|
||||
let dir = TempDir::new().unwrap();
|
||||
let wasm_path = dir.path().join("test.wasm");
|
||||
@@ -505,7 +510,8 @@ mod tests {
|
||||
async fn load_from_dir_returns_empty_when_dir_missing() {
|
||||
let config = WasmChannelRuntimeConfig::for_testing();
|
||||
let runtime = Arc::new(WasmChannelRuntime::new(config).unwrap());
|
||||
let loader = WasmChannelLoader::new(runtime, Arc::new(PairingStore::new()), None);
|
||||
let loader =
|
||||
WasmChannelLoader::new(runtime, Arc::new(PairingStore::new()), None, "default");
|
||||
|
||||
let dir = TempDir::new().unwrap();
|
||||
let missing = dir.path().join("nonexistent_channels_dir");
|
||||
|
||||
@@ -69,7 +69,7 @@
|
||||
//! let runtime = WasmChannelRuntime::new(config)?;
|
||||
//!
|
||||
//! // Load channels from directory
|
||||
//! let loader = WasmChannelLoader::new(runtime);
|
||||
//! let loader = WasmChannelLoader::new(runtime, pairing_store, settings_store, owner_scope_id);
|
||||
//! let channels = loader.load_from_dir(Path::new("~/.ironclaw/channels/")).await?;
|
||||
//!
|
||||
//! // Add to channel manager
|
||||
@@ -90,6 +90,7 @@ pub mod setup;
|
||||
pub(crate) mod signature;
|
||||
#[allow(dead_code)]
|
||||
pub(crate) mod storage;
|
||||
mod telegram_host_config;
|
||||
mod wrapper;
|
||||
|
||||
// Core types
|
||||
@@ -107,4 +108,5 @@ pub use schema::{
|
||||
ChannelCapabilitiesFile, ChannelConfig, SecretSetupSchema, SetupSchema, WebhookSchema,
|
||||
};
|
||||
pub use setup::{WasmChannelSetup, inject_channel_credentials, setup_wasm_channels};
|
||||
pub(crate) use telegram_host_config::{TELEGRAM_CHANNEL_NAME, bot_username_setting_key};
|
||||
pub use wrapper::{HttpResponse, SharedWasmChannel, WasmChannel};
|
||||
|
||||
@@ -672,6 +672,7 @@ mod tests {
|
||||
runtime,
|
||||
prepared,
|
||||
capabilities,
|
||||
"default",
|
||||
"{}".to_string(),
|
||||
Arc::new(PairingStore::new()),
|
||||
None,
|
||||
|
||||
+38
-10
@@ -7,8 +7,9 @@ use std::collections::HashSet;
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::channels::wasm::{
|
||||
LoadedChannel, RegisteredEndpoint, SharedWasmChannel, WasmChannel, WasmChannelLoader,
|
||||
WasmChannelRouter, WasmChannelRuntime, WasmChannelRuntimeConfig, create_wasm_channel_router,
|
||||
LoadedChannel, RegisteredEndpoint, SharedWasmChannel, TELEGRAM_CHANNEL_NAME, WasmChannel,
|
||||
WasmChannelLoader, WasmChannelRouter, WasmChannelRuntime, WasmChannelRuntimeConfig,
|
||||
bot_username_setting_key, create_wasm_channel_router,
|
||||
};
|
||||
use crate::config::Config;
|
||||
use crate::db::Database;
|
||||
@@ -48,7 +49,8 @@ pub async fn setup_wasm_channels(
|
||||
let mut loader = WasmChannelLoader::new(
|
||||
Arc::clone(&runtime),
|
||||
Arc::clone(&pairing_store),
|
||||
settings_store,
|
||||
settings_store.clone(),
|
||||
config.owner_id.clone(),
|
||||
);
|
||||
if let Some(secrets) = secrets_store {
|
||||
loader = loader.with_secrets_store(Arc::clone(secrets));
|
||||
@@ -70,7 +72,14 @@ pub async fn setup_wasm_channels(
|
||||
let mut channel_names: Vec<String> = Vec::new();
|
||||
|
||||
for loaded in results.loaded {
|
||||
let (name, channel) = register_channel(loaded, config, secrets_store, &wasm_router).await;
|
||||
let (name, channel) = register_channel(
|
||||
loaded,
|
||||
config,
|
||||
secrets_store,
|
||||
settings_store.as_ref(),
|
||||
&wasm_router,
|
||||
)
|
||||
.await;
|
||||
channel_names.push(name.clone());
|
||||
channels.push((name, channel));
|
||||
}
|
||||
@@ -104,10 +113,16 @@ async fn register_channel(
|
||||
loaded: LoadedChannel,
|
||||
config: &Config,
|
||||
secrets_store: &Option<Arc<dyn SecretsStore + Send + Sync>>,
|
||||
settings_store: Option<&Arc<dyn crate::db::SettingsStore>>,
|
||||
wasm_router: &Arc<WasmChannelRouter>,
|
||||
) -> (String, Box<dyn crate::channels::Channel>) {
|
||||
let channel_name = loaded.name().to_string();
|
||||
tracing::info!("Loaded WASM channel: {}", channel_name);
|
||||
let owner_actor_id = config
|
||||
.channels
|
||||
.wasm_channel_owner_ids
|
||||
.get(channel_name.as_str())
|
||||
.map(ToString::to_string);
|
||||
|
||||
let secret_name = loaded.webhook_secret_name();
|
||||
let sig_key_secret_name = loaded.signature_key_secret_name();
|
||||
@@ -115,7 +130,7 @@ async fn register_channel(
|
||||
|
||||
let webhook_secret = if let Some(secrets) = secrets_store {
|
||||
secrets
|
||||
.get_decrypted("default", &secret_name)
|
||||
.get_decrypted(&config.owner_id, &secret_name)
|
||||
.await
|
||||
.ok()
|
||||
.map(|s| s.expose().to_string())
|
||||
@@ -133,7 +148,7 @@ async fn register_channel(
|
||||
require_secret: webhook_secret.is_some(),
|
||||
}];
|
||||
|
||||
let channel_arc = Arc::new(loaded.channel);
|
||||
let channel_arc = Arc::new(loaded.channel.with_owner_actor_id(owner_actor_id.clone()));
|
||||
|
||||
// Inject runtime config (tunnel URL, webhook secret, owner_id).
|
||||
{
|
||||
@@ -161,6 +176,15 @@ async fn register_channel(
|
||||
config_updates.insert("owner_id".to_string(), serde_json::json!(owner_id));
|
||||
}
|
||||
|
||||
if channel_name == TELEGRAM_CHANNEL_NAME
|
||||
&& let Some(store) = settings_store
|
||||
&& let Ok(Some(serde_json::Value::String(username))) = store
|
||||
.get_setting("default", &bot_username_setting_key(&channel_name))
|
||||
.await
|
||||
&& !username.trim().is_empty()
|
||||
{
|
||||
config_updates.insert("bot_username".to_string(), serde_json::json!(username));
|
||||
}
|
||||
// Inject channel-specific secrets into config for channels that need
|
||||
// credentials in API request bodies (e.g., Feishu token exchange).
|
||||
// The credential injection system only replaces placeholders in URLs
|
||||
@@ -198,7 +222,7 @@ async fn register_channel(
|
||||
// Register Ed25519 signature key if declared in capabilities.
|
||||
if let Some(ref sig_key_name) = sig_key_secret_name
|
||||
&& let Some(secrets) = secrets_store
|
||||
&& let Ok(key_secret) = secrets.get_decrypted("default", sig_key_name).await
|
||||
&& let Ok(key_secret) = secrets.get_decrypted(&config.owner_id, sig_key_name).await
|
||||
{
|
||||
match wasm_router
|
||||
.register_signature_key(&channel_name, key_secret.expose())
|
||||
@@ -216,7 +240,9 @@ async fn register_channel(
|
||||
// Register HMAC signing secret if declared in capabilities.
|
||||
if let Some(ref hmac_secret_name) = hmac_secret_name
|
||||
&& let Some(secrets) = secrets_store
|
||||
&& let Ok(secret) = secrets.get_decrypted("default", hmac_secret_name).await
|
||||
&& let Ok(secret) = secrets
|
||||
.get_decrypted(&config.owner_id, hmac_secret_name)
|
||||
.await
|
||||
{
|
||||
wasm_router
|
||||
.register_hmac_secret(&channel_name, secret.expose())
|
||||
@@ -231,6 +257,7 @@ async fn register_channel(
|
||||
.as_ref()
|
||||
.map(|s| s.as_ref() as &dyn SecretsStore),
|
||||
&channel_name,
|
||||
&config.owner_id,
|
||||
)
|
||||
.await
|
||||
{
|
||||
@@ -268,6 +295,7 @@ pub async fn inject_channel_credentials(
|
||||
channel: &Arc<WasmChannel>,
|
||||
secrets: Option<&dyn SecretsStore>,
|
||||
channel_name: &str,
|
||||
owner_id: &str,
|
||||
) -> anyhow::Result<usize> {
|
||||
if channel_name.trim().is_empty() {
|
||||
return Ok(0);
|
||||
@@ -279,7 +307,7 @@ pub async fn inject_channel_credentials(
|
||||
// 1. Try injecting from persistent secrets store if available
|
||||
if let Some(secrets) = secrets {
|
||||
let all_secrets = secrets
|
||||
.list("default")
|
||||
.list(owner_id)
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("Failed to list secrets: {}", e))?;
|
||||
|
||||
@@ -290,7 +318,7 @@ pub async fn inject_channel_credentials(
|
||||
continue;
|
||||
}
|
||||
|
||||
let decrypted = match secrets.get_decrypted("default", &secret_meta.name).await {
|
||||
let decrypted = match secrets.get_decrypted(owner_id, &secret_meta.name).await {
|
||||
Ok(d) => d,
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
pub const TELEGRAM_CHANNEL_NAME: &str = "telegram";
|
||||
const TELEGRAM_BOT_USERNAME_SETTING_PREFIX: &str = "channels.wasm_channel_bot_usernames";
|
||||
|
||||
pub fn bot_username_setting_key(channel_name: &str) -> String {
|
||||
format!("{TELEGRAM_BOT_USERNAME_SETTING_PREFIX}.{channel_name}")
|
||||
}
|
||||
+546
-189
File diff suppressed because it is too large
Load Diff
@@ -162,15 +162,30 @@ pub async fn chat_auth_token_handler(
|
||||
.await
|
||||
{
|
||||
Ok(result) => {
|
||||
clear_auth_mode(&state).await;
|
||||
let mut resp = ActionResponse::ok(result.message.clone());
|
||||
resp.activated = Some(result.activated);
|
||||
resp.auth_url = result.auth_url.clone();
|
||||
resp.verification = result.verification.clone();
|
||||
resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone());
|
||||
|
||||
state.sse.broadcast(SseEvent::AuthCompleted {
|
||||
extension_name: req.extension_name.clone(),
|
||||
success: true,
|
||||
message: result.message.clone(),
|
||||
});
|
||||
if result.verification.is_some() {
|
||||
state.sse.broadcast(SseEvent::AuthRequired {
|
||||
extension_name: req.extension_name.clone(),
|
||||
instructions: Some(result.message),
|
||||
auth_url: None,
|
||||
setup_url: None,
|
||||
});
|
||||
} else {
|
||||
clear_auth_mode(&state).await;
|
||||
|
||||
Ok(Json(ActionResponse::ok(result.message)))
|
||||
state.sse.broadcast(SseEvent::AuthCompleted {
|
||||
extension_name: req.extension_name.clone(),
|
||||
success: true,
|
||||
message: result.message,
|
||||
});
|
||||
}
|
||||
|
||||
Ok(Json(resp))
|
||||
}
|
||||
Err(e) => {
|
||||
let msg = e.to_string();
|
||||
|
||||
@@ -25,34 +25,34 @@ pub async fn extensions_list_handler(
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
let pairing_store = crate::pairing::PairingStore::new();
|
||||
let mut owner_bound_channels = std::collections::HashSet::new();
|
||||
for ext in &installed {
|
||||
if ext.kind == crate::extensions::ExtensionKind::WasmChannel
|
||||
&& ext_mgr.has_wasm_channel_owner_binding(&ext.name).await
|
||||
{
|
||||
owner_bound_channels.insert(ext.name.clone());
|
||||
}
|
||||
}
|
||||
let extensions = installed
|
||||
.into_iter()
|
||||
.map(|ext| {
|
||||
let activation_status = if ext.kind == crate::extensions::ExtensionKind::WasmChannel {
|
||||
Some(if ext.activation_error.is_some() {
|
||||
"failed".to_string()
|
||||
} else if !ext.authenticated {
|
||||
"installed".to_string()
|
||||
} else if ext.active {
|
||||
let has_paired = pairing_store
|
||||
.read_allow_from(&ext.name)
|
||||
.map(|list| !list.is_empty())
|
||||
.unwrap_or(false);
|
||||
if has_paired {
|
||||
"active".to_string()
|
||||
} else {
|
||||
"pairing".to_string()
|
||||
}
|
||||
} else {
|
||||
"configured".to_string()
|
||||
})
|
||||
let has_paired = pairing_store
|
||||
.read_allow_from(&ext.name)
|
||||
.map(|list| !list.is_empty())
|
||||
.unwrap_or(false);
|
||||
crate::channels::web::types::classify_wasm_channel_activation(
|
||||
&ext,
|
||||
has_paired,
|
||||
owner_bound_channels.contains(&ext.name),
|
||||
)
|
||||
} else if ext.kind == crate::extensions::ExtensionKind::ChannelRelay {
|
||||
Some(if ext.active {
|
||||
"active".to_string()
|
||||
crate::channels::web::types::ExtensionActivationStatus::Active
|
||||
} else if ext.authenticated {
|
||||
"configured".to_string()
|
||||
crate::channels::web::types::ExtensionActivationStatus::Configured
|
||||
} else {
|
||||
"installed".to_string()
|
||||
crate::channels::web::types::ExtensionActivationStatus::Installed
|
||||
})
|
||||
} else {
|
||||
None
|
||||
|
||||
+240
-115
@@ -26,7 +26,6 @@ use tower_http::set_header::SetResponseHeaderLayer;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::agent::SessionManager;
|
||||
use crate::agent::routine::{Trigger, next_cron_fire};
|
||||
use crate::bootstrap::ironclaw_base_dir;
|
||||
use crate::channels::IncomingMessage;
|
||||
use crate::channels::relay::DEFAULT_RELAY_NAME;
|
||||
@@ -36,6 +35,7 @@ use crate::channels::web::handlers::jobs::{
|
||||
jobs_events_handler, jobs_list_handler, jobs_prompt_handler, jobs_restart_handler,
|
||||
jobs_summary_handler,
|
||||
};
|
||||
use crate::channels::web::handlers::routines::{routines_delete_handler, routines_toggle_handler};
|
||||
use crate::channels::web::handlers::skills::{
|
||||
skills_install_handler, skills_list_handler, skills_remove_handler, skills_search_handler,
|
||||
};
|
||||
@@ -1163,19 +1163,43 @@ async fn chat_auth_token_handler(
|
||||
.configure_token(&req.extension_name, &req.token)
|
||||
.await
|
||||
{
|
||||
Ok(result) if result.activated => {
|
||||
// Clear auth mode on the active thread
|
||||
clear_auth_mode(&state).await;
|
||||
Ok(result) => {
|
||||
let mut resp = if result.verification.is_some() || result.activated {
|
||||
ActionResponse::ok(result.message.clone())
|
||||
} else {
|
||||
ActionResponse::fail(result.message.clone())
|
||||
};
|
||||
resp.activated = Some(result.activated);
|
||||
resp.auth_url = result.auth_url.clone();
|
||||
resp.verification = result.verification.clone();
|
||||
resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone());
|
||||
|
||||
state.sse.broadcast(SseEvent::AuthCompleted {
|
||||
extension_name: req.extension_name.clone(),
|
||||
success: true,
|
||||
message: result.message.clone(),
|
||||
});
|
||||
if result.verification.is_some() {
|
||||
state.sse.broadcast(SseEvent::AuthRequired {
|
||||
extension_name: req.extension_name.clone(),
|
||||
instructions: Some(result.message),
|
||||
auth_url: None,
|
||||
setup_url: None,
|
||||
});
|
||||
} else if result.activated {
|
||||
// Clear auth mode on the active thread
|
||||
clear_auth_mode(&state).await;
|
||||
|
||||
Ok(Json(ActionResponse::ok(result.message)))
|
||||
state.sse.broadcast(SseEvent::AuthCompleted {
|
||||
extension_name: req.extension_name.clone(),
|
||||
success: true,
|
||||
message: result.message,
|
||||
});
|
||||
} else {
|
||||
state.sse.broadcast(SseEvent::AuthCompleted {
|
||||
extension_name: req.extension_name.clone(),
|
||||
success: false,
|
||||
message: result.message,
|
||||
});
|
||||
}
|
||||
|
||||
Ok(Json(resp))
|
||||
}
|
||||
Ok(result) => Ok(Json(ActionResponse::fail(result.message))),
|
||||
Err(e) => {
|
||||
let msg = e.to_string();
|
||||
// Re-emit auth_required for retry on validation errors
|
||||
@@ -1818,29 +1842,34 @@ async fn extensions_list_handler(
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
let pairing_store = crate::pairing::PairingStore::new();
|
||||
let mut owner_bound_channels = std::collections::HashSet::new();
|
||||
for ext in &installed {
|
||||
if ext.kind == crate::extensions::ExtensionKind::WasmChannel
|
||||
&& ext_mgr.has_wasm_channel_owner_binding(&ext.name).await
|
||||
{
|
||||
owner_bound_channels.insert(ext.name.clone());
|
||||
}
|
||||
}
|
||||
let extensions = installed
|
||||
.into_iter()
|
||||
.map(|ext| {
|
||||
let activation_status = if ext.kind == crate::extensions::ExtensionKind::WasmChannel {
|
||||
Some(if ext.activation_error.is_some() {
|
||||
"failed".to_string()
|
||||
} else if !ext.authenticated {
|
||||
// No credentials configured yet.
|
||||
"installed".to_string()
|
||||
} else if ext.active {
|
||||
// Check pairing status for active channels.
|
||||
let has_paired = pairing_store
|
||||
.read_allow_from(&ext.name)
|
||||
.map(|list| !list.is_empty())
|
||||
.unwrap_or(false);
|
||||
if has_paired {
|
||||
"active".to_string()
|
||||
} else {
|
||||
"pairing".to_string()
|
||||
}
|
||||
let has_paired = pairing_store
|
||||
.read_allow_from(&ext.name)
|
||||
.map(|list| !list.is_empty())
|
||||
.unwrap_or(false);
|
||||
crate::channels::web::types::classify_wasm_channel_activation(
|
||||
&ext,
|
||||
has_paired,
|
||||
owner_bound_channels.contains(&ext.name),
|
||||
)
|
||||
} else if ext.kind == crate::extensions::ExtensionKind::ChannelRelay {
|
||||
Some(if ext.active {
|
||||
ExtensionActivationStatus::Active
|
||||
} else if ext.authenticated {
|
||||
ExtensionActivationStatus::Configured
|
||||
} else {
|
||||
// Authenticated but not yet active.
|
||||
"configured".to_string()
|
||||
ExtensionActivationStatus::Installed
|
||||
})
|
||||
} else {
|
||||
None
|
||||
@@ -2205,20 +2234,24 @@ async fn extensions_setup_submit_handler(
|
||||
|
||||
match ext_mgr.configure(&name, &req.secrets).await {
|
||||
Ok(result) => {
|
||||
// Broadcast completion status so chat UI can dismiss success cases while
|
||||
// leaving failed auth/configuration flows visible for correction.
|
||||
state.sse.broadcast(SseEvent::AuthCompleted {
|
||||
extension_name: name.clone(),
|
||||
success: result.activated,
|
||||
message: result.message.clone(),
|
||||
});
|
||||
let mut resp = if result.activated {
|
||||
let mut resp = if result.verification.is_some() || result.activated {
|
||||
ActionResponse::ok(result.message)
|
||||
} else {
|
||||
ActionResponse::fail(result.message)
|
||||
};
|
||||
resp.activated = Some(result.activated);
|
||||
resp.auth_url = result.auth_url;
|
||||
resp.auth_url = result.auth_url.clone();
|
||||
resp.verification = result.verification.clone();
|
||||
resp.instructions = result.verification.as_ref().map(|v| v.instructions.clone());
|
||||
if result.verification.is_none() {
|
||||
// Broadcast auth_completed so the chat UI can dismiss any in-progress
|
||||
// auth card or setup modal that was triggered by tool_auth/tool_activate.
|
||||
state.sse.broadcast(SseEvent::AuthCompleted {
|
||||
extension_name: name.clone(),
|
||||
success: result.activated,
|
||||
message: resp.message.clone(),
|
||||
});
|
||||
}
|
||||
Ok(Json(resp))
|
||||
}
|
||||
Err(e) => Ok(Json(ActionResponse::fail(e.to_string()))),
|
||||
@@ -2430,83 +2463,6 @@ async fn routines_trigger_handler(
|
||||
})))
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct ToggleRequest {
|
||||
enabled: Option<bool>,
|
||||
}
|
||||
|
||||
async fn routines_toggle_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
Path(id): Path<String>,
|
||||
body: Option<Json<ToggleRequest>>,
|
||||
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
|
||||
let store = state.store.as_ref().ok_or((
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"Database not available".to_string(),
|
||||
))?;
|
||||
|
||||
let routine_id = Uuid::parse_str(&id)
|
||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
|
||||
|
||||
let mut routine = store
|
||||
.get_routine(routine_id)
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?
|
||||
.ok_or((StatusCode::NOT_FOUND, "Routine not found".to_string()))?;
|
||||
|
||||
let was_enabled = routine.enabled;
|
||||
// If a specific value was provided, use it; otherwise toggle.
|
||||
routine.enabled = match body {
|
||||
Some(Json(req)) => req.enabled.unwrap_or(!routine.enabled),
|
||||
None => !routine.enabled,
|
||||
};
|
||||
|
||||
if routine.enabled
|
||||
&& !was_enabled
|
||||
&& let Trigger::Cron { schedule, timezone } = &routine.trigger
|
||||
{
|
||||
routine.next_fire_at = next_cron_fire(schedule, timezone.as_deref())
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
}
|
||||
|
||||
store
|
||||
.update_routine(&routine)
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"status": if routine.enabled { "enabled" } else { "disabled" },
|
||||
"routine_id": routine_id,
|
||||
})))
|
||||
}
|
||||
|
||||
async fn routines_delete_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 routine_id = Uuid::parse_str(&id)
|
||||
.map_err(|_| (StatusCode::BAD_REQUEST, "Invalid routine ID".to_string()))?;
|
||||
|
||||
let deleted = store
|
||||
.delete_routine(routine_id)
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
if deleted {
|
||||
Ok(Json(serde_json::json!({
|
||||
"status": "deleted",
|
||||
"routine_id": routine_id,
|
||||
})))
|
||||
} else {
|
||||
Err((StatusCode::NOT_FOUND, "Routine not found".to_string()))
|
||||
}
|
||||
}
|
||||
|
||||
async fn routines_runs_handler(
|
||||
State(state): State<Arc<GatewayState>>,
|
||||
Path(id): Path<String>,
|
||||
@@ -2743,7 +2699,11 @@ struct GatewayStatusResponse {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::channels::web::types::{
|
||||
ExtensionActivationStatus, classify_wasm_channel_activation,
|
||||
};
|
||||
use crate::cli::oauth_defaults;
|
||||
use crate::extensions::{ExtensionKind, InstalledExtension};
|
||||
use crate::testing::credentials::TEST_GATEWAY_CRYPTO_KEY;
|
||||
|
||||
#[test]
|
||||
@@ -2822,6 +2782,85 @@ mod tests {
|
||||
assert!(turns.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_wasm_channel_activation_status_owner_bound_counts_as_active() -> Result<(), String> {
|
||||
let ext = InstalledExtension {
|
||||
name: "telegram".to_string(),
|
||||
kind: ExtensionKind::WasmChannel,
|
||||
display_name: Some("Telegram".to_string()),
|
||||
description: None,
|
||||
url: None,
|
||||
authenticated: true,
|
||||
active: true,
|
||||
tools: Vec::new(),
|
||||
needs_setup: true,
|
||||
has_auth: false,
|
||||
installed: true,
|
||||
activation_error: None,
|
||||
version: None,
|
||||
};
|
||||
|
||||
let owner_bound = classify_wasm_channel_activation(&ext, false, true);
|
||||
if owner_bound != Some(ExtensionActivationStatus::Active) {
|
||||
return Err(format!(
|
||||
"owner-bound channel should be active, got {:?}",
|
||||
owner_bound
|
||||
));
|
||||
}
|
||||
|
||||
let unbound = classify_wasm_channel_activation(&ext, false, false);
|
||||
if unbound != Some(ExtensionActivationStatus::Pairing) {
|
||||
return Err(format!(
|
||||
"unbound channel should be pairing, got {:?}",
|
||||
unbound
|
||||
));
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_channel_relay_activation_status_is_preserved() -> Result<(), String> {
|
||||
let relay = InstalledExtension {
|
||||
name: "signal".to_string(),
|
||||
kind: ExtensionKind::ChannelRelay,
|
||||
display_name: Some("Signal".to_string()),
|
||||
description: None,
|
||||
url: None,
|
||||
authenticated: true,
|
||||
active: false,
|
||||
tools: Vec::new(),
|
||||
needs_setup: true,
|
||||
has_auth: false,
|
||||
installed: true,
|
||||
activation_error: None,
|
||||
version: None,
|
||||
};
|
||||
|
||||
let status = if relay.kind == crate::extensions::ExtensionKind::WasmChannel {
|
||||
classify_wasm_channel_activation(&relay, false, false)
|
||||
} else if relay.kind == crate::extensions::ExtensionKind::ChannelRelay {
|
||||
Some(if relay.active {
|
||||
ExtensionActivationStatus::Active
|
||||
} else if relay.authenticated {
|
||||
ExtensionActivationStatus::Configured
|
||||
} else {
|
||||
ExtensionActivationStatus::Installed
|
||||
})
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
if status != Some(ExtensionActivationStatus::Configured) {
|
||||
return Err(format!(
|
||||
"channel relay should retain configured status, got {:?}",
|
||||
status
|
||||
));
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// --- OAuth callback handler tests ---
|
||||
|
||||
/// Build a minimal `GatewayState` for testing the OAuth callback handler.
|
||||
@@ -2935,6 +2974,92 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_extensions_setup_submit_telegram_verification_does_not_broadcast_auth_required() {
|
||||
use axum::body::Body;
|
||||
use tokio::time::{Duration, timeout};
|
||||
use tower::ServiceExt;
|
||||
|
||||
let secrets = test_secrets_store();
|
||||
let (ext_mgr, _wasm_tools_dir, wasm_channels_dir) = test_ext_mgr(secrets);
|
||||
|
||||
std::fs::write(
|
||||
wasm_channels_dir.path().join("telegram.wasm"),
|
||||
b"\0asm fake",
|
||||
)
|
||||
.expect("write fake telegram wasm");
|
||||
let caps = serde_json::json!({
|
||||
"type": "channel",
|
||||
"name": "telegram",
|
||||
"setup": {
|
||||
"required_secrets": [
|
||||
{
|
||||
"name": "telegram_bot_token",
|
||||
"prompt": "Enter your Telegram Bot API token (from @BotFather)"
|
||||
}
|
||||
]
|
||||
}
|
||||
});
|
||||
std::fs::write(
|
||||
wasm_channels_dir.path().join("telegram.capabilities.json"),
|
||||
serde_json::to_string(&caps).expect("serialize telegram caps"),
|
||||
)
|
||||
.expect("write telegram caps");
|
||||
|
||||
ext_mgr
|
||||
.set_test_telegram_pending_verification("iclaw-7qk2m9", Some("test_hot_bot"))
|
||||
.await;
|
||||
|
||||
let state = test_gateway_state(Some(ext_mgr));
|
||||
let mut receiver = state.sse.sender().subscribe();
|
||||
let app = Router::new()
|
||||
.route(
|
||||
"/api/extensions/{name}/setup",
|
||||
post(extensions_setup_submit_handler),
|
||||
)
|
||||
.with_state(state);
|
||||
|
||||
let req_body = serde_json::json!({
|
||||
"secrets": {
|
||||
"telegram_bot_token": "123456789:ABCdefGhI"
|
||||
}
|
||||
});
|
||||
let req = axum::http::Request::builder()
|
||||
.method("POST")
|
||||
.uri("/api/extensions/telegram/setup")
|
||||
.header("content-type", "application/json")
|
||||
.body(Body::from(req_body.to_string()))
|
||||
.expect("request");
|
||||
|
||||
let resp = ServiceExt::<axum::http::Request<Body>>::oneshot(app, req)
|
||||
.await
|
||||
.expect("response");
|
||||
assert_eq!(resp.status(), StatusCode::OK);
|
||||
|
||||
let body = axum::body::to_bytes(resp.into_body(), 1024 * 64)
|
||||
.await
|
||||
.expect("body");
|
||||
let parsed: serde_json::Value = serde_json::from_slice(&body).expect("json response");
|
||||
assert_eq!(parsed["success"], serde_json::Value::Bool(true));
|
||||
assert_eq!(parsed["activated"], serde_json::Value::Bool(false));
|
||||
assert_eq!(parsed["verification"]["code"], "iclaw-7qk2m9");
|
||||
|
||||
let deadline = tokio::time::Instant::now() + Duration::from_millis(100);
|
||||
loop {
|
||||
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
|
||||
if remaining.is_zero() {
|
||||
break;
|
||||
}
|
||||
match timeout(remaining, receiver.recv()).await {
|
||||
Ok(Ok(crate::channels::web::types::SseEvent::AuthRequired { .. })) => {
|
||||
panic!("verification responses should not emit auth_required SSE events")
|
||||
}
|
||||
Ok(Ok(_)) => continue,
|
||||
Ok(Err(_)) | Err(_) => break,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn expired_flow_created_at() -> Option<std::time::Instant> {
|
||||
std::time::Instant::now()
|
||||
.checked_sub(oauth_defaults::OAUTH_FLOW_EXPIRY + std::time::Duration::from_secs(1))
|
||||
|
||||
+186
-10
@@ -527,7 +527,6 @@ function enableChatInput() {
|
||||
const btn = document.getElementById('send-btn');
|
||||
if (input) {
|
||||
input.disabled = false;
|
||||
input.placeholder = I18n.t('chat.inputPlaceholder');
|
||||
}
|
||||
if (btn) btn.disabled = false;
|
||||
}
|
||||
@@ -1205,11 +1204,13 @@ function showJobCard(data) {
|
||||
// --- Auth card ---
|
||||
|
||||
function handleAuthRequired(data) {
|
||||
setAuthFlowPending(true, data.instructions);
|
||||
if (data.auth_url) {
|
||||
setAuthFlowPending(true, data.instructions);
|
||||
// OAuth flow: show the global auth prompt with an OAuth button + optional token paste field.
|
||||
showAuthCard(data);
|
||||
} else {
|
||||
if (getConfigureOverlay(data.extension_name)) return;
|
||||
setAuthFlowPending(true, data.instructions);
|
||||
// Setup flow: fetch the extension's credential schema and show the multi-field
|
||||
// configure modal (the same UI used by the Extensions tab "Setup" button).
|
||||
showConfigureModal(data.extension_name);
|
||||
@@ -1433,13 +1434,11 @@ function setAuthFlowPending(pending, instructions) {
|
||||
if (authFlowPending) {
|
||||
input.disabled = true;
|
||||
btn.disabled = true;
|
||||
input.placeholder = instructions || 'Complete extension auth to continue chatting';
|
||||
return;
|
||||
}
|
||||
if (!currentThreadIsReadOnly) {
|
||||
input.disabled = false;
|
||||
btn.disabled = false;
|
||||
input.placeholder = I18n.t('chat.inputPlaceholder');
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2641,7 +2640,7 @@ function renderExtensionCard(ext) {
|
||||
pairingSection.className = 'ext-pairing';
|
||||
pairingSection.setAttribute('data-channel', ext.name);
|
||||
card.appendChild(pairingSection);
|
||||
loadPairingRequests(ext.name, pairingSection);
|
||||
loadPairingRequests(ext.name, pairingSection, ext.activation_status);
|
||||
}
|
||||
|
||||
return card;
|
||||
@@ -2712,8 +2711,11 @@ function renderConfigureModal(name, secrets) {
|
||||
const overlay = document.createElement('div');
|
||||
overlay.className = 'configure-overlay';
|
||||
overlay.setAttribute('data-extension-name', name);
|
||||
overlay.dataset.telegramVerificationState = 'idle';
|
||||
overlay.addEventListener('click', (e) => {
|
||||
if (e.target === overlay) closeConfigureModal();
|
||||
if (e.target !== overlay) return;
|
||||
if (name === 'telegram' && overlay.dataset.telegramVerificationState === 'waiting') return;
|
||||
closeConfigureModal();
|
||||
});
|
||||
|
||||
const modal = document.createElement('div');
|
||||
@@ -2723,6 +2725,13 @@ function renderConfigureModal(name, secrets) {
|
||||
header.textContent = I18n.t('config.title', { name: name });
|
||||
modal.appendChild(header);
|
||||
|
||||
if (name === 'telegram') {
|
||||
const hint = document.createElement('div');
|
||||
hint.className = 'configure-hint';
|
||||
hint.textContent = I18n.t('config.telegramOwnerHint');
|
||||
modal.appendChild(hint);
|
||||
}
|
||||
|
||||
const form = document.createElement('div');
|
||||
form.className = 'configure-form';
|
||||
|
||||
@@ -2730,6 +2739,7 @@ function renderConfigureModal(name, secrets) {
|
||||
for (const secret of secrets) {
|
||||
const field = document.createElement('div');
|
||||
field.className = 'configure-field';
|
||||
field.dataset.secretName = secret.name;
|
||||
|
||||
const label = document.createElement('label');
|
||||
label.textContent = secret.prompt;
|
||||
@@ -2774,6 +2784,16 @@ function renderConfigureModal(name, secrets) {
|
||||
|
||||
modal.appendChild(form);
|
||||
|
||||
const error = document.createElement('div');
|
||||
error.className = 'configure-inline-error';
|
||||
error.style.display = 'none';
|
||||
modal.appendChild(error);
|
||||
|
||||
const status = document.createElement('div');
|
||||
status.className = 'configure-inline-status';
|
||||
status.style.display = 'none';
|
||||
modal.appendChild(status);
|
||||
|
||||
const actions = document.createElement('div');
|
||||
actions.className = 'configure-actions';
|
||||
|
||||
@@ -2796,7 +2816,110 @@ function renderConfigureModal(name, secrets) {
|
||||
if (fields.length > 0) fields[0].input.focus();
|
||||
}
|
||||
|
||||
function submitConfigureModal(name, fields) {
|
||||
function renderTelegramVerificationChallenge(overlay, verification) {
|
||||
if (!overlay || !verification) return;
|
||||
const modal = overlay.querySelector('.configure-modal');
|
||||
if (!modal) return;
|
||||
const telegramField = modal.querySelector('.configure-field[data-secret-name="telegram_bot_token"]');
|
||||
|
||||
let panel = modal.querySelector('.configure-verification');
|
||||
if (!panel) {
|
||||
panel = document.createElement('div');
|
||||
panel.className = 'configure-verification';
|
||||
}
|
||||
if (telegramField && telegramField.parentNode) {
|
||||
telegramField.insertAdjacentElement('afterend', panel);
|
||||
} else {
|
||||
modal.insertBefore(
|
||||
panel,
|
||||
modal.querySelector('.configure-inline-error') || modal.querySelector('.configure-actions')
|
||||
);
|
||||
}
|
||||
|
||||
panel.innerHTML = '';
|
||||
|
||||
const title = document.createElement('div');
|
||||
title.className = 'configure-verification-title';
|
||||
title.textContent = I18n.t('config.telegramChallengeTitle');
|
||||
panel.appendChild(title);
|
||||
|
||||
const instructions = document.createElement('div');
|
||||
instructions.className = 'configure-verification-instructions';
|
||||
instructions.textContent = verification.instructions;
|
||||
panel.appendChild(instructions);
|
||||
|
||||
const commandLabel = document.createElement('div');
|
||||
commandLabel.className = 'configure-verification-instructions';
|
||||
commandLabel.textContent = I18n.t('config.telegramCommandLabel');
|
||||
panel.appendChild(commandLabel);
|
||||
|
||||
const command = document.createElement('code');
|
||||
command.className = 'configure-verification-code';
|
||||
command.textContent = '/start ' + verification.code;
|
||||
panel.appendChild(command);
|
||||
|
||||
if (verification.deep_link) {
|
||||
const link = document.createElement('a');
|
||||
link.className = 'configure-verification-link';
|
||||
link.href = verification.deep_link;
|
||||
link.target = '_blank';
|
||||
link.rel = 'noreferrer noopener';
|
||||
link.textContent = I18n.t('config.telegramOpenBot');
|
||||
panel.appendChild(link);
|
||||
}
|
||||
}
|
||||
|
||||
function getConfigurePrimaryButton(overlay) {
|
||||
return overlay && overlay.querySelector('.configure-actions button.btn-ext.activate');
|
||||
}
|
||||
|
||||
function getConfigureCancelButton(overlay) {
|
||||
return overlay && overlay.querySelector('.configure-actions button.btn-ext.remove');
|
||||
}
|
||||
|
||||
function setConfigureInlineError(overlay, message) {
|
||||
const error = overlay && overlay.querySelector('.configure-inline-error');
|
||||
if (!error) return;
|
||||
error.textContent = message || '';
|
||||
error.style.display = message ? 'block' : 'none';
|
||||
}
|
||||
|
||||
function clearConfigureInlineError(overlay) {
|
||||
setConfigureInlineError(overlay, '');
|
||||
}
|
||||
|
||||
function setConfigureInlineStatus(overlay, message) {
|
||||
const status = overlay && overlay.querySelector('.configure-inline-status');
|
||||
if (!status) return;
|
||||
status.textContent = message || '';
|
||||
status.style.display = message ? 'block' : 'none';
|
||||
}
|
||||
|
||||
function setTelegramConfigureState(overlay, fields, state) {
|
||||
if (!overlay) return;
|
||||
overlay.dataset.telegramVerificationState = state;
|
||||
|
||||
const primaryBtn = getConfigurePrimaryButton(overlay);
|
||||
const cancelBtn = getConfigureCancelButton(overlay);
|
||||
const waiting = state === 'waiting';
|
||||
const retry = state === 'retry';
|
||||
|
||||
setConfigureInlineStatus(overlay, waiting ? I18n.t('config.telegramOwnerWaiting') : '');
|
||||
|
||||
if (primaryBtn) {
|
||||
primaryBtn.style.display = waiting ? 'none' : '';
|
||||
primaryBtn.disabled = false;
|
||||
primaryBtn.textContent = retry ? I18n.t('config.telegramStartOver') : I18n.t('config.save');
|
||||
}
|
||||
if (cancelBtn) cancelBtn.disabled = waiting;
|
||||
}
|
||||
|
||||
function startTelegramAutoVerify(name, fields) {
|
||||
window.setTimeout(() => submitConfigureModal(name, fields, { telegramAutoVerify: true }), 0);
|
||||
}
|
||||
|
||||
function submitConfigureModal(name, fields, options) {
|
||||
options = options || {};
|
||||
const secrets = {};
|
||||
for (const f of fields) {
|
||||
if (f.input.value.trim()) {
|
||||
@@ -2804,10 +2927,16 @@ function submitConfigureModal(name, fields) {
|
||||
}
|
||||
}
|
||||
|
||||
// Disable buttons to prevent double-submit
|
||||
const overlay = getConfigureOverlay(name) || document.querySelector('.configure-overlay');
|
||||
const isTelegram = name === 'telegram';
|
||||
clearConfigureInlineError(overlay);
|
||||
|
||||
// Disable buttons to prevent double-submit
|
||||
var btns = overlay ? overlay.querySelectorAll('.configure-actions button') : [];
|
||||
btns.forEach(function(b) { b.disabled = true; });
|
||||
if (overlay && isTelegram) {
|
||||
setTelegramConfigureState(overlay, fields, 'waiting');
|
||||
}
|
||||
|
||||
apiFetch('/api/extensions/' + encodeURIComponent(name) + '/setup', {
|
||||
method: 'POST',
|
||||
@@ -2815,6 +2944,23 @@ function submitConfigureModal(name, fields) {
|
||||
})
|
||||
.then((res) => {
|
||||
if (res.success) {
|
||||
if (res.verification && isTelegram) {
|
||||
renderTelegramVerificationChallenge(overlay, res.verification);
|
||||
fields.forEach(function(f) { f.input.value = ''; });
|
||||
setTelegramConfigureState(overlay, fields, 'waiting');
|
||||
// Once the verification challenge is rendered inline, the global auth lock
|
||||
// should not keep the chat composer disabled for this setup-driven flow.
|
||||
setAuthFlowPending(false);
|
||||
enableChatInput();
|
||||
if (!options.telegramAutoVerify) {
|
||||
startTelegramAutoVerify(name, fields);
|
||||
return;
|
||||
}
|
||||
setTelegramConfigureState(overlay, fields, 'retry');
|
||||
setConfigureInlineError(overlay, I18n.t('config.telegramStartOverHint'));
|
||||
return;
|
||||
}
|
||||
|
||||
closeConfigureModal();
|
||||
if (res.auth_url) {
|
||||
showAuthCard({
|
||||
@@ -2830,11 +2976,29 @@ function submitConfigureModal(name, fields) {
|
||||
} else {
|
||||
// Keep modal open so the user can correct their input and retry.
|
||||
btns.forEach(function(b) { b.disabled = false; });
|
||||
setConfigureInlineError(overlay, res.message || 'Configuration failed');
|
||||
if (isTelegram) {
|
||||
const hasVerification = overlay && overlay.querySelector('.configure-verification');
|
||||
if (options.telegramAutoVerify || hasVerification) {
|
||||
setTelegramConfigureState(overlay, fields, 'retry');
|
||||
} else {
|
||||
setTelegramConfigureState(overlay, fields, 'idle');
|
||||
}
|
||||
}
|
||||
showToast(res.message || 'Configuration failed', 'error');
|
||||
}
|
||||
})
|
||||
.catch((err) => {
|
||||
btns.forEach(function(b) { b.disabled = false; });
|
||||
setConfigureInlineError(overlay, 'Configuration failed: ' + err.message);
|
||||
if (isTelegram) {
|
||||
const hasVerification = overlay && overlay.querySelector('.configure-verification');
|
||||
if (options.telegramAutoVerify || hasVerification) {
|
||||
setTelegramConfigureState(overlay, fields, 'retry');
|
||||
} else {
|
||||
setTelegramConfigureState(overlay, fields, 'idle');
|
||||
}
|
||||
}
|
||||
showToast('Configuration failed: ' + err.message, 'error');
|
||||
});
|
||||
}
|
||||
@@ -2843,6 +3007,10 @@ function closeConfigureModal(extensionName) {
|
||||
if (typeof extensionName !== 'string') extensionName = null;
|
||||
const existing = getConfigureOverlay(extensionName);
|
||||
if (existing) existing.remove();
|
||||
if (!document.querySelector('.configure-overlay') && !document.querySelector('.auth-card')) {
|
||||
setAuthFlowPending(false);
|
||||
enableChatInput();
|
||||
}
|
||||
}
|
||||
|
||||
// Validate that a server-supplied OAuth URL is HTTPS before opening a popup.
|
||||
@@ -2866,11 +3034,19 @@ function openOAuthUrl(url) {
|
||||
|
||||
// --- Pairing ---
|
||||
|
||||
function loadPairingRequests(channel, container) {
|
||||
function loadPairingRequests(channel, container, status) {
|
||||
apiFetch('/api/pairing/' + encodeURIComponent(channel))
|
||||
.then(data => {
|
||||
container.innerHTML = '';
|
||||
if (!data.requests || data.requests.length === 0) return;
|
||||
if (!data.requests || data.requests.length === 0) {
|
||||
if (status === 'pairing') {
|
||||
const hint = document.createElement('p');
|
||||
hint.className = 'pairing-hint';
|
||||
hint.textContent = 'Send any message to your bot to receive a pairing request here.';
|
||||
container.appendChild(hint);
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
const heading = document.createElement('div');
|
||||
heading.className = 'pairing-heading';
|
||||
|
||||
@@ -342,6 +342,13 @@ I18n.register('en', {
|
||||
|
||||
// Configure
|
||||
'config.title': 'Configure {name}',
|
||||
'config.telegramOwnerHint': 'After saving, IronClaw will show a one-time code. Send `/start CODE` to your bot in Telegram and IronClaw will finish setup automatically.',
|
||||
'config.telegramChallengeTitle': 'Telegram owner verification',
|
||||
'config.telegramOwnerWaiting': 'Waiting for Telegram owner verification...',
|
||||
'config.telegramCommandLabel': 'Send this in Telegram:',
|
||||
'config.telegramStartOver': 'Start over',
|
||||
'config.telegramStartOverHint': 'Telegram verification did not complete. Click Start over to generate a new code and try again.',
|
||||
'config.telegramOpenBot': 'Open bot in Telegram',
|
||||
'config.optional': ' (optional)',
|
||||
'config.alreadySet': '(already set — leave empty to keep)',
|
||||
'config.alreadyConfigured': 'Already configured',
|
||||
|
||||
@@ -342,6 +342,12 @@ I18n.register('zh-CN', {
|
||||
|
||||
// 配置
|
||||
'config.title': '配置 {name}',
|
||||
'config.telegramOwnerHint': '保存后,IronClaw 会显示一次性验证码。将 `/start CODE` 发送给你的 Telegram 机器人,IronClaw 会自动完成设置。',
|
||||
'config.telegramChallengeTitle': 'Telegram 所有者验证',
|
||||
'config.telegramOwnerWaiting': '正在等待 Telegram 所有者验证...',
|
||||
'config.telegramCommandLabel': '请在 Telegram 中发送:',
|
||||
'config.telegramStartOver': '重新开始',
|
||||
'config.telegramStartOverHint': 'Telegram 验证未完成。点击“重新开始”以生成新的验证码并重试。',
|
||||
'config.optional': '(可选)',
|
||||
'config.alreadySet': '(已设置 — 留空以保持不变)',
|
||||
'config.alreadyConfigured': '已配置',
|
||||
|
||||
@@ -2865,6 +2865,13 @@ body {
|
||||
flex: 1;
|
||||
}
|
||||
|
||||
.pairing-hint {
|
||||
color: var(--text-secondary);
|
||||
font-size: 13px;
|
||||
margin: 4px 0 8px;
|
||||
font-style: italic;
|
||||
}
|
||||
|
||||
/* Configure modal */
|
||||
.configure-overlay {
|
||||
position: fixed;
|
||||
@@ -2896,6 +2903,84 @@ body {
|
||||
color: var(--text-primary);
|
||||
}
|
||||
|
||||
.configure-hint {
|
||||
margin: 0 0 16px 0;
|
||||
padding: 10px 12px;
|
||||
border-radius: 8px;
|
||||
background: var(--bg-secondary);
|
||||
border: 1px solid var(--border);
|
||||
color: var(--text-secondary);
|
||||
font-size: 13px;
|
||||
line-height: 1.5;
|
||||
}
|
||||
|
||||
.configure-verification {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 10px;
|
||||
margin: 16px 0 0 0;
|
||||
padding: 12px;
|
||||
border-radius: 8px;
|
||||
background: var(--bg-secondary);
|
||||
border: 1px solid var(--border);
|
||||
}
|
||||
|
||||
.configure-verification-title {
|
||||
font-size: 13px;
|
||||
font-weight: 600;
|
||||
color: var(--text-primary);
|
||||
}
|
||||
|
||||
.configure-verification-instructions {
|
||||
font-size: 13px;
|
||||
line-height: 1.5;
|
||||
color: var(--text-secondary);
|
||||
}
|
||||
|
||||
.configure-verification-code {
|
||||
display: inline-block;
|
||||
width: fit-content;
|
||||
padding: 6px 10px;
|
||||
border-radius: 6px;
|
||||
background: rgba(255, 255, 255, 0.06);
|
||||
border: 1px solid var(--border);
|
||||
color: var(--text-primary);
|
||||
font-size: 13px;
|
||||
}
|
||||
|
||||
.configure-verification-link {
|
||||
width: fit-content;
|
||||
color: var(--accent, var(--text-link, #4ea3ff));
|
||||
font-size: 13px;
|
||||
text-decoration: none;
|
||||
}
|
||||
|
||||
.configure-verification-link:hover {
|
||||
text-decoration: underline;
|
||||
}
|
||||
|
||||
.configure-inline-error {
|
||||
margin: 16px 0 0 0;
|
||||
padding: 10px 12px;
|
||||
border-radius: 8px;
|
||||
background: rgba(220, 38, 38, 0.12);
|
||||
border: 1px solid rgba(220, 38, 38, 0.35);
|
||||
color: #fca5a5;
|
||||
font-size: 13px;
|
||||
line-height: 1.5;
|
||||
}
|
||||
|
||||
.configure-inline-status {
|
||||
margin: 16px 0 0 0;
|
||||
padding: 10px 12px;
|
||||
border-radius: 8px;
|
||||
background: var(--bg-secondary);
|
||||
border: 1px solid var(--border);
|
||||
color: var(--text-secondary);
|
||||
font-size: 13px;
|
||||
line-height: 1.5;
|
||||
}
|
||||
|
||||
.configure-form {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
|
||||
@@ -410,6 +410,40 @@ pub struct TransitionInfo {
|
||||
|
||||
// --- Extensions ---
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ExtensionActivationStatus {
|
||||
Installed,
|
||||
Configured,
|
||||
Pairing,
|
||||
Active,
|
||||
Failed,
|
||||
}
|
||||
|
||||
pub fn classify_wasm_channel_activation(
|
||||
ext: &crate::extensions::InstalledExtension,
|
||||
has_paired: bool,
|
||||
has_owner_binding: bool,
|
||||
) -> Option<ExtensionActivationStatus> {
|
||||
if ext.kind != crate::extensions::ExtensionKind::WasmChannel {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(if ext.activation_error.is_some() {
|
||||
ExtensionActivationStatus::Failed
|
||||
} else if !ext.authenticated {
|
||||
ExtensionActivationStatus::Installed
|
||||
} else if ext.active {
|
||||
if has_paired || has_owner_binding {
|
||||
ExtensionActivationStatus::Active
|
||||
} else {
|
||||
ExtensionActivationStatus::Pairing
|
||||
}
|
||||
} else {
|
||||
ExtensionActivationStatus::Configured
|
||||
})
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct ExtensionInfo {
|
||||
pub name: String,
|
||||
@@ -428,9 +462,9 @@ pub struct ExtensionInfo {
|
||||
/// Whether this extension has an auth configuration (OAuth or manual token).
|
||||
#[serde(default)]
|
||||
pub has_auth: bool,
|
||||
/// WASM channel activation status: "installed", "configured", "active", "failed".
|
||||
/// WASM channel activation status.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub activation_status: Option<String>,
|
||||
pub activation_status: Option<ExtensionActivationStatus>,
|
||||
/// Human-readable error when activation_status is "failed".
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub activation_error: Option<String>,
|
||||
@@ -503,6 +537,9 @@ pub struct ActionResponse {
|
||||
/// Whether the channel was successfully activated after setup.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub activated: Option<bool>,
|
||||
/// Pending manual verification challenge (for Telegram owner binding, etc.).
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub verification: Option<crate::extensions::VerificationChallenge>,
|
||||
}
|
||||
|
||||
impl ActionResponse {
|
||||
@@ -514,6 +551,7 @@ impl ActionResponse {
|
||||
awaiting_token: None,
|
||||
instructions: None,
|
||||
activated: None,
|
||||
verification: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -525,6 +563,7 @@ impl ActionResponse {
|
||||
awaiting_token: None,
|
||||
instructions: None,
|
||||
activated: None,
|
||||
verification: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+19
-8
@@ -265,14 +265,25 @@ async fn handle_client_message(
|
||||
if let Some(ref ext_mgr) = state.extension_manager {
|
||||
match ext_mgr.configure_token(&extension_name, &token).await {
|
||||
Ok(result) => {
|
||||
crate::channels::web::server::clear_auth_mode(state).await;
|
||||
state
|
||||
.sse
|
||||
.broadcast(crate::channels::web::types::SseEvent::AuthCompleted {
|
||||
extension_name,
|
||||
success: true,
|
||||
message: result.message,
|
||||
});
|
||||
if result.verification.is_some() {
|
||||
state.sse.broadcast(
|
||||
crate::channels::web::types::SseEvent::AuthRequired {
|
||||
extension_name: extension_name.clone(),
|
||||
instructions: Some(result.message),
|
||||
auth_url: None,
|
||||
setup_url: None,
|
||||
},
|
||||
);
|
||||
} else {
|
||||
crate::channels::web::server::clear_auth_mode(state).await;
|
||||
state.sse.broadcast(
|
||||
crate::channels::web::types::SseEvent::AuthCompleted {
|
||||
extension_name,
|
||||
success: true,
|
||||
message: result.message,
|
||||
},
|
||||
);
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
let msg = format!("Auth failed: {}", e);
|
||||
|
||||
+5
-4
@@ -405,10 +405,11 @@ fn check_routines_config() -> CheckResult {
|
||||
fn check_gateway_config(settings: &Settings) -> CheckResult {
|
||||
// Use the same resolve() path as runtime so invalid env values
|
||||
// (e.g. GATEWAY_PORT=abc) are caught here too.
|
||||
let tunnel_enabled = crate::config::TunnelConfig::resolve(settings)
|
||||
.map(|t| t.is_enabled())
|
||||
.unwrap_or(false);
|
||||
match crate::config::ChannelsConfig::resolve(settings, tunnel_enabled) {
|
||||
let owner_id = match crate::config::resolve_owner_id(settings) {
|
||||
Ok(owner_id) => owner_id,
|
||||
Err(e) => return CheckResult::Fail(format!("config error: {e}")),
|
||||
};
|
||||
match crate::config::ChannelsConfig::resolve(settings, &owner_id) {
|
||||
Ok(channels) => match channels.gateway {
|
||||
Some(gw) => {
|
||||
if gw.auth_token.is_some() {
|
||||
|
||||
+21
-7
@@ -292,6 +292,16 @@ async fn list(
|
||||
|
||||
// ── Create ──────────────────────────────────────────────────
|
||||
|
||||
fn cli_notify_config(notify_channel: Option<String>) -> NotifyConfig {
|
||||
NotifyConfig {
|
||||
channel: notify_channel,
|
||||
user: None,
|
||||
on_attention: true,
|
||||
on_failure: true,
|
||||
on_success: false,
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn create(
|
||||
db: &Arc<dyn Database>,
|
||||
@@ -338,13 +348,7 @@ async fn create(
|
||||
max_concurrent: 1,
|
||||
dedup_window: None,
|
||||
},
|
||||
notify: NotifyConfig {
|
||||
channel: notify_channel,
|
||||
user: user_id.to_string(),
|
||||
on_attention: true,
|
||||
on_failure: true,
|
||||
on_success: false,
|
||||
},
|
||||
notify: cli_notify_config(notify_channel),
|
||||
last_run_at: None,
|
||||
next_fire_at: next_fire,
|
||||
run_count: 0,
|
||||
@@ -729,4 +733,14 @@ mod tests {
|
||||
// Must be valid UTF-8 (would have panicked otherwise).
|
||||
assert!(result.is_char_boundary(result.len()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cli_notify_config_defaults_to_runtime_target_resolution() {
|
||||
let notify = cli_notify_config(Some("telegram".to_string()));
|
||||
assert_eq!(notify.channel.as_deref(), Some("telegram")); // safety: test-only assertion
|
||||
assert_eq!(notify.user, None); // safety: test-only assertion
|
||||
assert!(notify.on_attention); // safety: test-only assertion
|
||||
assert!(notify.on_failure); // safety: test-only assertion
|
||||
assert!(!notify.on_success); // safety: test-only assertion
|
||||
}
|
||||
}
|
||||
|
||||
+42
-6
@@ -32,13 +32,16 @@ impl Default for BuilderModeConfig {
|
||||
}
|
||||
|
||||
impl BuilderModeConfig {
|
||||
pub(crate) fn resolve() -> Result<Self, ConfigError> {
|
||||
pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result<Self, ConfigError> {
|
||||
let bs = &settings.builder;
|
||||
Ok(Self {
|
||||
enabled: parse_bool_env("BUILDER_ENABLED", true)?,
|
||||
build_dir: optional_env("BUILDER_DIR")?.map(PathBuf::from),
|
||||
max_iterations: parse_optional_env("BUILDER_MAX_ITERATIONS", 20)?,
|
||||
timeout_secs: parse_optional_env("BUILDER_TIMEOUT_SECS", 600)?,
|
||||
auto_register: parse_bool_env("BUILDER_AUTO_REGISTER", true)?,
|
||||
enabled: parse_bool_env("BUILDER_ENABLED", bs.enabled)?,
|
||||
build_dir: optional_env("BUILDER_DIR")?
|
||||
.map(PathBuf::from)
|
||||
.or_else(|| bs.build_dir.clone()),
|
||||
max_iterations: parse_optional_env("BUILDER_MAX_ITERATIONS", bs.max_iterations)?,
|
||||
timeout_secs: parse_optional_env("BUILDER_TIMEOUT_SECS", bs.timeout_secs)?,
|
||||
auto_register: parse_bool_env("BUILDER_AUTO_REGISTER", bs.auto_register)?,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -56,3 +59,36 @@ impl BuilderModeConfig {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::config::helpers::ENV_MUTEX;
|
||||
use crate::settings::Settings;
|
||||
|
||||
#[test]
|
||||
fn resolve_falls_back_to_settings() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let mut settings = Settings::default();
|
||||
settings.builder.max_iterations = 99;
|
||||
settings.builder.auto_register = false;
|
||||
|
||||
let cfg = BuilderModeConfig::resolve(&settings).expect("resolve");
|
||||
assert_eq!(cfg.max_iterations, 99);
|
||||
assert!(!cfg.auto_register);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn env_overrides_settings() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let mut settings = Settings::default();
|
||||
settings.builder.timeout_secs = 123;
|
||||
|
||||
// SAFETY: Under ENV_MUTEX, no concurrent env access.
|
||||
unsafe { std::env::set_var("BUILDER_TIMEOUT_SECS", "3") };
|
||||
let cfg = BuilderModeConfig::resolve(&settings).expect("resolve");
|
||||
unsafe { std::env::remove_var("BUILDER_TIMEOUT_SECS") };
|
||||
|
||||
assert_eq!(cfg.timeout_secs, 3);
|
||||
}
|
||||
}
|
||||
|
||||
+55
-335
@@ -91,36 +91,24 @@ pub struct SignalConfig {
|
||||
}
|
||||
|
||||
impl ChannelsConfig {
|
||||
/// Resolve channels config following `env > settings > default` for every field.
|
||||
pub(crate) fn resolve(settings: &Settings, tunnel_enabled: bool) -> Result<Self, ConfigError> {
|
||||
pub(crate) fn resolve(settings: &Settings, owner_id: &str) -> Result<Self, ConfigError> {
|
||||
let cs = &settings.channels;
|
||||
|
||||
// --- HTTP webhook ---
|
||||
// HTTP is enabled when env vars are set OR settings has it enabled.
|
||||
let http_enabled_by_env =
|
||||
optional_env("HTTP_PORT")?.is_some() || optional_env("HTTP_HOST")?.is_some();
|
||||
// When a tunnel is configured, default to loopback since external
|
||||
// traffic arrives through the tunnel. Without a tunnel the webhook
|
||||
// server needs to accept connections from the network directly.
|
||||
let default_host = if tunnel_enabled {
|
||||
"127.0.0.1"
|
||||
} else {
|
||||
"0.0.0.0"
|
||||
};
|
||||
let http = if http_enabled_by_env || cs.http_enabled {
|
||||
Some(HttpConfig {
|
||||
host: optional_env("HTTP_HOST")?
|
||||
.or_else(|| cs.http_host.clone())
|
||||
.unwrap_or_else(|| default_host.to_string()),
|
||||
.unwrap_or_else(|| "0.0.0.0".to_string()),
|
||||
port: parse_optional_env("HTTP_PORT", cs.http_port.unwrap_or(8080))?,
|
||||
webhook_secret: optional_env("HTTP_WEBHOOK_SECRET")?.map(SecretString::from),
|
||||
user_id: optional_env("HTTP_USER_ID")?.unwrap_or_else(|| "http".to_string()),
|
||||
user_id: owner_id.to_string(),
|
||||
})
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
// --- Web gateway ---
|
||||
let gateway_enabled = parse_bool_env("GATEWAY_ENABLED", cs.gateway_enabled)?;
|
||||
let gateway = if gateway_enabled {
|
||||
Some(GatewayConfig {
|
||||
@@ -133,33 +121,29 @@ impl ChannelsConfig {
|
||||
)?,
|
||||
auth_token: optional_env("GATEWAY_AUTH_TOKEN")?
|
||||
.or_else(|| cs.gateway_auth_token.clone()),
|
||||
user_id: optional_env("GATEWAY_USER_ID")?
|
||||
.or_else(|| cs.gateway_user_id.clone())
|
||||
.unwrap_or_else(|| "default".to_string()),
|
||||
user_id: owner_id.to_string(),
|
||||
})
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
// --- Signal ---
|
||||
let signal_url = optional_env("SIGNAL_HTTP_URL")?.or_else(|| cs.signal_http_url.clone());
|
||||
let signal = if let Some(http_url) = signal_url {
|
||||
let account = optional_env("SIGNAL_ACCOUNT")?
|
||||
.or_else(|| cs.signal_account.clone())
|
||||
.ok_or(ConfigError::InvalidValue {
|
||||
key: "SIGNAL_ACCOUNT".to_string(),
|
||||
message: "SIGNAL_ACCOUNT is required when Signal is enabled".to_string(),
|
||||
message: "SIGNAL_ACCOUNT is required when SIGNAL_HTTP_URL is set".to_string(),
|
||||
})?;
|
||||
let allow_from_str =
|
||||
optional_env("SIGNAL_ALLOW_FROM")?.or_else(|| cs.signal_allow_from.clone());
|
||||
let allow_from = match allow_from_str {
|
||||
None => vec![account.clone()],
|
||||
Some(s) => s
|
||||
.split(',')
|
||||
.map(|e| e.trim().to_string())
|
||||
.filter(|s| !s.is_empty())
|
||||
.collect(),
|
||||
};
|
||||
let allow_from =
|
||||
match optional_env("SIGNAL_ALLOW_FROM")?.or_else(|| cs.signal_allow_from.clone()) {
|
||||
None => vec![account.clone()],
|
||||
Some(s) => s
|
||||
.split(',')
|
||||
.map(|e| e.trim().to_string())
|
||||
.filter(|s| !s.is_empty())
|
||||
.collect(),
|
||||
};
|
||||
let dm_policy = optional_env("SIGNAL_DM_POLICY")?
|
||||
.or_else(|| cs.signal_dm_policy.clone())
|
||||
.unwrap_or_else(|| "pairing".to_string());
|
||||
@@ -201,18 +185,8 @@ impl ChannelsConfig {
|
||||
None
|
||||
};
|
||||
|
||||
// --- CLI ---
|
||||
let cli_enabled = parse_bool_env("CLI_ENABLED", cs.cli_enabled)?;
|
||||
|
||||
// --- WASM channels ---
|
||||
let wasm_channels_dir = optional_env("WASM_CHANNELS_DIR")?
|
||||
.map(PathBuf::from)
|
||||
.or_else(|| cs.wasm_channels_dir.clone())
|
||||
.unwrap_or_else(default_channels_dir);
|
||||
|
||||
let wasm_channels_enabled =
|
||||
parse_bool_env("WASM_CHANNELS_ENABLED", cs.wasm_channels_enabled)?;
|
||||
|
||||
Ok(Self {
|
||||
cli: CliConfig {
|
||||
enabled: cli_enabled,
|
||||
@@ -220,8 +194,14 @@ impl ChannelsConfig {
|
||||
http,
|
||||
gateway,
|
||||
signal,
|
||||
wasm_channels_dir,
|
||||
wasm_channels_enabled,
|
||||
wasm_channels_dir: optional_env("WASM_CHANNELS_DIR")?
|
||||
.map(PathBuf::from)
|
||||
.or_else(|| cs.wasm_channels_dir.clone())
|
||||
.unwrap_or_else(default_channels_dir),
|
||||
wasm_channels_enabled: parse_bool_env(
|
||||
"WASM_CHANNELS_ENABLED",
|
||||
cs.wasm_channels_enabled,
|
||||
)?,
|
||||
wasm_channel_owner_ids: {
|
||||
let mut ids = cs.wasm_channel_owner_ids.clone();
|
||||
// Backwards compat: TELEGRAM_OWNER_ID env var
|
||||
@@ -252,6 +232,8 @@ fn default_channels_dir() -> PathBuf {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::config::channels::*;
|
||||
use crate::config::helpers::ENV_MUTEX;
|
||||
use crate::settings::Settings;
|
||||
|
||||
#[test]
|
||||
fn cli_config_fields() {
|
||||
@@ -398,69 +380,6 @@ mod tests {
|
||||
assert!(!cfg.wasm_channels_enabled);
|
||||
}
|
||||
|
||||
/// When a tunnel is active and HTTP_HOST is not explicitly set, the
|
||||
/// webhook server should default to loopback to avoid unnecessary exposure.
|
||||
#[test]
|
||||
fn http_host_defaults_to_loopback_with_tunnel() {
|
||||
// Set HTTP_PORT to trigger HttpConfig creation, but leave HTTP_HOST unset
|
||||
// so the default kicks in.
|
||||
unsafe {
|
||||
std::env::set_var("HTTP_PORT", "9999");
|
||||
std::env::remove_var("HTTP_HOST");
|
||||
}
|
||||
let settings = crate::settings::Settings::default();
|
||||
let cfg = ChannelsConfig::resolve(&settings, true).unwrap();
|
||||
unsafe {
|
||||
std::env::remove_var("HTTP_PORT");
|
||||
}
|
||||
let http = cfg.http.expect("HttpConfig should be present");
|
||||
assert_eq!(
|
||||
http.host, "127.0.0.1",
|
||||
"tunnel active should default to loopback"
|
||||
);
|
||||
assert_eq!(http.port, 9999);
|
||||
}
|
||||
|
||||
/// Without a tunnel, the webhook server defaults to 0.0.0.0 so external
|
||||
/// services can reach it directly.
|
||||
#[test]
|
||||
fn http_host_defaults_to_all_interfaces_without_tunnel() {
|
||||
unsafe {
|
||||
std::env::set_var("HTTP_PORT", "9998");
|
||||
std::env::remove_var("HTTP_HOST");
|
||||
}
|
||||
let settings = crate::settings::Settings::default();
|
||||
let cfg = ChannelsConfig::resolve(&settings, false).unwrap();
|
||||
unsafe {
|
||||
std::env::remove_var("HTTP_PORT");
|
||||
}
|
||||
let http = cfg.http.expect("HttpConfig should be present");
|
||||
assert_eq!(
|
||||
http.host, "0.0.0.0",
|
||||
"no tunnel should default to all interfaces"
|
||||
);
|
||||
}
|
||||
|
||||
/// An explicit HTTP_HOST always wins regardless of tunnel state.
|
||||
#[test]
|
||||
fn explicit_http_host_overrides_tunnel_default() {
|
||||
unsafe {
|
||||
std::env::set_var("HTTP_PORT", "9997");
|
||||
std::env::set_var("HTTP_HOST", "192.168.1.50");
|
||||
}
|
||||
let settings = crate::settings::Settings::default();
|
||||
let cfg = ChannelsConfig::resolve(&settings, true).unwrap();
|
||||
unsafe {
|
||||
std::env::remove_var("HTTP_PORT");
|
||||
std::env::remove_var("HTTP_HOST");
|
||||
}
|
||||
let http = cfg.http.expect("HttpConfig should be present");
|
||||
assert_eq!(
|
||||
http.host, "192.168.1.50",
|
||||
"explicit host should override tunnel default"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_channels_dir_ends_with_channels() {
|
||||
let dir = default_channels_dir();
|
||||
@@ -471,242 +390,43 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn default_gateway_port_constant() {
|
||||
assert_eq!(DEFAULT_GATEWAY_PORT, 3000);
|
||||
}
|
||||
|
||||
/// With default settings and no env vars, gateway should use defaults.
|
||||
#[test]
|
||||
fn resolve_gateway_defaults_from_settings() {
|
||||
let _lock = crate::config::helpers::ENV_MUTEX.lock();
|
||||
// Clear env vars that would interfere
|
||||
unsafe {
|
||||
std::env::remove_var("GATEWAY_ENABLED");
|
||||
std::env::remove_var("GATEWAY_HOST");
|
||||
std::env::remove_var("GATEWAY_PORT");
|
||||
std::env::remove_var("GATEWAY_AUTH_TOKEN");
|
||||
std::env::remove_var("GATEWAY_USER_ID");
|
||||
std::env::remove_var("HTTP_PORT");
|
||||
std::env::remove_var("HTTP_HOST");
|
||||
std::env::remove_var("SIGNAL_HTTP_URL");
|
||||
std::env::remove_var("CLI_ENABLED");
|
||||
std::env::remove_var("WASM_CHANNELS_DIR");
|
||||
std::env::remove_var("WASM_CHANNELS_ENABLED");
|
||||
std::env::remove_var("TELEGRAM_OWNER_ID");
|
||||
}
|
||||
|
||||
let settings = crate::settings::Settings::default();
|
||||
let cfg = ChannelsConfig::resolve(&settings, false).unwrap();
|
||||
|
||||
let gw = cfg.gateway.expect("gateway should be enabled by default");
|
||||
assert_eq!(gw.host, "127.0.0.1");
|
||||
assert_eq!(gw.port, DEFAULT_GATEWAY_PORT);
|
||||
assert!(gw.auth_token.is_none());
|
||||
assert_eq!(gw.user_id, "default");
|
||||
}
|
||||
|
||||
/// Settings values should be used when no env vars are set.
|
||||
#[test]
|
||||
fn resolve_gateway_from_settings() {
|
||||
let _lock = crate::config::helpers::ENV_MUTEX.lock();
|
||||
unsafe {
|
||||
std::env::remove_var("GATEWAY_ENABLED");
|
||||
std::env::remove_var("GATEWAY_HOST");
|
||||
std::env::remove_var("GATEWAY_PORT");
|
||||
std::env::remove_var("GATEWAY_AUTH_TOKEN");
|
||||
std::env::remove_var("GATEWAY_USER_ID");
|
||||
std::env::remove_var("HTTP_PORT");
|
||||
std::env::remove_var("HTTP_HOST");
|
||||
std::env::remove_var("SIGNAL_HTTP_URL");
|
||||
std::env::remove_var("CLI_ENABLED");
|
||||
std::env::remove_var("WASM_CHANNELS_DIR");
|
||||
std::env::remove_var("WASM_CHANNELS_ENABLED");
|
||||
std::env::remove_var("TELEGRAM_OWNER_ID");
|
||||
}
|
||||
|
||||
let mut settings = crate::settings::Settings::default();
|
||||
settings.channels.gateway_port = Some(4000);
|
||||
settings.channels.gateway_host = Some("0.0.0.0".to_string());
|
||||
settings.channels.gateway_auth_token = Some("db-token-123".to_string());
|
||||
settings.channels.gateway_user_id = Some("myuser".to_string());
|
||||
|
||||
let cfg = ChannelsConfig::resolve(&settings, false).unwrap();
|
||||
let gw = cfg.gateway.expect("gateway should be enabled");
|
||||
assert_eq!(gw.port, 4000);
|
||||
assert_eq!(gw.host, "0.0.0.0");
|
||||
assert_eq!(gw.auth_token.as_deref(), Some("db-token-123"));
|
||||
assert_eq!(gw.user_id, "myuser");
|
||||
}
|
||||
|
||||
/// Env vars should override settings values.
|
||||
#[test]
|
||||
fn resolve_env_overrides_settings() {
|
||||
let _lock = crate::config::helpers::ENV_MUTEX.lock();
|
||||
unsafe {
|
||||
std::env::set_var("GATEWAY_PORT", "5000");
|
||||
std::env::set_var("GATEWAY_HOST", "10.0.0.1");
|
||||
std::env::set_var("GATEWAY_AUTH_TOKEN", "env-token");
|
||||
std::env::remove_var("GATEWAY_ENABLED");
|
||||
std::env::remove_var("GATEWAY_USER_ID");
|
||||
std::env::remove_var("HTTP_PORT");
|
||||
std::env::remove_var("HTTP_HOST");
|
||||
std::env::remove_var("SIGNAL_HTTP_URL");
|
||||
std::env::remove_var("CLI_ENABLED");
|
||||
std::env::remove_var("WASM_CHANNELS_DIR");
|
||||
std::env::remove_var("WASM_CHANNELS_ENABLED");
|
||||
std::env::remove_var("TELEGRAM_OWNER_ID");
|
||||
}
|
||||
|
||||
let mut settings = crate::settings::Settings::default();
|
||||
settings.channels.gateway_port = Some(4000);
|
||||
settings.channels.gateway_host = Some("0.0.0.0".to_string());
|
||||
settings.channels.gateway_auth_token = Some("db-token".to_string());
|
||||
|
||||
let cfg = ChannelsConfig::resolve(&settings, false).unwrap();
|
||||
let gw = cfg.gateway.expect("gateway should be enabled");
|
||||
assert_eq!(gw.port, 5000, "env should override settings");
|
||||
assert_eq!(gw.host, "10.0.0.1", "env should override settings");
|
||||
assert_eq!(
|
||||
gw.auth_token.as_deref(),
|
||||
Some("env-token"),
|
||||
"env should override settings"
|
||||
);
|
||||
|
||||
// Cleanup
|
||||
unsafe {
|
||||
std::env::remove_var("GATEWAY_PORT");
|
||||
std::env::remove_var("GATEWAY_HOST");
|
||||
std::env::remove_var("GATEWAY_AUTH_TOKEN");
|
||||
}
|
||||
}
|
||||
|
||||
/// CLI enabled should fall back to settings.
|
||||
#[test]
|
||||
fn resolve_cli_enabled_from_settings() {
|
||||
let _lock = crate::config::helpers::ENV_MUTEX.lock();
|
||||
unsafe {
|
||||
std::env::remove_var("CLI_ENABLED");
|
||||
std::env::remove_var("GATEWAY_ENABLED");
|
||||
std::env::remove_var("GATEWAY_HOST");
|
||||
std::env::remove_var("GATEWAY_PORT");
|
||||
std::env::remove_var("GATEWAY_AUTH_TOKEN");
|
||||
std::env::remove_var("GATEWAY_USER_ID");
|
||||
std::env::remove_var("HTTP_PORT");
|
||||
std::env::remove_var("HTTP_HOST");
|
||||
std::env::remove_var("SIGNAL_HTTP_URL");
|
||||
std::env::remove_var("WASM_CHANNELS_DIR");
|
||||
std::env::remove_var("WASM_CHANNELS_ENABLED");
|
||||
std::env::remove_var("TELEGRAM_OWNER_ID");
|
||||
}
|
||||
|
||||
let mut settings = crate::settings::Settings::default();
|
||||
settings.channels.cli_enabled = false;
|
||||
|
||||
let cfg = ChannelsConfig::resolve(&settings, false).unwrap();
|
||||
assert!(!cfg.cli.enabled, "settings should disable CLI");
|
||||
}
|
||||
|
||||
/// HTTP channel should activate when settings has it enabled.
|
||||
#[test]
|
||||
fn resolve_http_from_settings() {
|
||||
let _lock = crate::config::helpers::ENV_MUTEX.lock();
|
||||
unsafe {
|
||||
std::env::remove_var("HTTP_PORT");
|
||||
std::env::remove_var("HTTP_HOST");
|
||||
std::env::remove_var("HTTP_WEBHOOK_SECRET");
|
||||
std::env::remove_var("HTTP_USER_ID");
|
||||
std::env::remove_var("GATEWAY_ENABLED");
|
||||
std::env::remove_var("GATEWAY_HOST");
|
||||
std::env::remove_var("GATEWAY_PORT");
|
||||
std::env::remove_var("GATEWAY_AUTH_TOKEN");
|
||||
std::env::remove_var("GATEWAY_USER_ID");
|
||||
std::env::remove_var("SIGNAL_HTTP_URL");
|
||||
std::env::remove_var("CLI_ENABLED");
|
||||
std::env::remove_var("WASM_CHANNELS_DIR");
|
||||
std::env::remove_var("WASM_CHANNELS_ENABLED");
|
||||
std::env::remove_var("TELEGRAM_OWNER_ID");
|
||||
}
|
||||
|
||||
let mut settings = crate::settings::Settings::default();
|
||||
fn resolve_uses_settings_channel_values_with_owner_scope_user_ids() {
|
||||
let _guard = ENV_MUTEX.lock().unwrap_or_else(|e| e.into_inner());
|
||||
let mut settings = Settings::default();
|
||||
settings.channels.http_enabled = true;
|
||||
settings.channels.http_port = Some(9090);
|
||||
settings.channels.http_host = Some("10.0.0.1".to_string());
|
||||
settings.channels.http_host = Some("127.0.0.2".to_string());
|
||||
settings.channels.http_port = Some(8181);
|
||||
settings.channels.gateway_enabled = true;
|
||||
settings.channels.gateway_host = Some("127.0.0.3".to_string());
|
||||
settings.channels.gateway_port = Some(9191);
|
||||
settings.channels.gateway_auth_token = Some("tok".to_string());
|
||||
settings.channels.signal_http_url = Some("http://127.0.0.1:8080".to_string());
|
||||
settings.channels.signal_account = Some("+15551234567".to_string());
|
||||
settings.channels.signal_allow_from = Some("+15551234567,+15557654321".to_string());
|
||||
settings.channels.wasm_channels_dir = Some(PathBuf::from("/tmp/settings-channels"));
|
||||
settings.channels.wasm_channels_enabled = false;
|
||||
|
||||
let cfg = ChannelsConfig::resolve(&settings, false).unwrap();
|
||||
let http = cfg.http.expect("HTTP should be enabled from settings");
|
||||
assert_eq!(http.port, 9090);
|
||||
assert_eq!(http.host, "10.0.0.1");
|
||||
}
|
||||
let cfg = ChannelsConfig::resolve(&settings, "owner-scope").expect("resolve");
|
||||
|
||||
/// Settings round-trip through DB map for new gateway fields.
|
||||
#[test]
|
||||
fn settings_gateway_fields_db_roundtrip() {
|
||||
let mut settings = crate::settings::Settings::default();
|
||||
settings.channels.gateway_port = Some(4000);
|
||||
settings.channels.gateway_host = Some("0.0.0.0".to_string());
|
||||
settings.channels.gateway_auth_token = Some("tok-abc".to_string());
|
||||
settings.channels.gateway_user_id = Some("myuser".to_string());
|
||||
settings.channels.cli_enabled = false;
|
||||
let http = cfg.http.expect("http config");
|
||||
assert_eq!(http.host, "127.0.0.2");
|
||||
assert_eq!(http.port, 8181);
|
||||
assert_eq!(http.user_id, "owner-scope");
|
||||
|
||||
let map = settings.to_db_map();
|
||||
let restored = crate::settings::Settings::from_db_map(&map);
|
||||
let gateway = cfg.gateway.expect("gateway config");
|
||||
assert_eq!(gateway.host, "127.0.0.3");
|
||||
assert_eq!(gateway.port, 9191);
|
||||
assert_eq!(gateway.auth_token.as_deref(), Some("tok"));
|
||||
assert_eq!(gateway.user_id, "owner-scope");
|
||||
|
||||
let signal = cfg.signal.expect("signal config");
|
||||
assert_eq!(signal.account, "+15551234567");
|
||||
assert_eq!(signal.allow_from, vec!["+15551234567", "+15557654321"]);
|
||||
|
||||
assert_eq!(restored.channels.gateway_port, Some(4000));
|
||||
assert_eq!(restored.channels.gateway_host.as_deref(), Some("0.0.0.0"));
|
||||
assert_eq!(
|
||||
restored.channels.gateway_auth_token.as_deref(),
|
||||
Some("tok-abc")
|
||||
cfg.wasm_channels_dir,
|
||||
PathBuf::from("/tmp/settings-channels")
|
||||
);
|
||||
assert_eq!(restored.channels.gateway_user_id.as_deref(), Some("myuser"));
|
||||
assert!(!restored.channels.cli_enabled);
|
||||
}
|
||||
|
||||
/// Invalid boolean env values must produce errors, not silently degrade.
|
||||
#[test]
|
||||
fn resolve_rejects_invalid_bool_env() {
|
||||
let _lock = crate::config::helpers::ENV_MUTEX.lock();
|
||||
let settings = crate::settings::Settings::default();
|
||||
|
||||
// GATEWAY_ENABLED=maybe should error
|
||||
unsafe {
|
||||
std::env::set_var("GATEWAY_ENABLED", "maybe");
|
||||
std::env::remove_var("HTTP_PORT");
|
||||
std::env::remove_var("HTTP_HOST");
|
||||
std::env::remove_var("SIGNAL_HTTP_URL");
|
||||
std::env::remove_var("CLI_ENABLED");
|
||||
std::env::remove_var("WASM_CHANNELS_ENABLED");
|
||||
std::env::remove_var("GATEWAY_PORT");
|
||||
std::env::remove_var("GATEWAY_HOST");
|
||||
std::env::remove_var("GATEWAY_AUTH_TOKEN");
|
||||
std::env::remove_var("GATEWAY_USER_ID");
|
||||
std::env::remove_var("WASM_CHANNELS_DIR");
|
||||
std::env::remove_var("TELEGRAM_OWNER_ID");
|
||||
}
|
||||
let result = ChannelsConfig::resolve(&settings, false);
|
||||
assert!(result.is_err(), "GATEWAY_ENABLED=maybe should be rejected");
|
||||
|
||||
// CLI_ENABLED=on should error
|
||||
unsafe {
|
||||
std::env::remove_var("GATEWAY_ENABLED");
|
||||
std::env::set_var("CLI_ENABLED", "on");
|
||||
}
|
||||
let result = ChannelsConfig::resolve(&settings, false);
|
||||
assert!(result.is_err(), "CLI_ENABLED=on should be rejected");
|
||||
|
||||
// WASM_CHANNELS_ENABLED=yes should error
|
||||
unsafe {
|
||||
std::env::remove_var("CLI_ENABLED");
|
||||
std::env::set_var("WASM_CHANNELS_ENABLED", "yes");
|
||||
}
|
||||
let result = ChannelsConfig::resolve(&settings, false);
|
||||
assert!(
|
||||
result.is_err(),
|
||||
"WASM_CHANNELS_ENABLED=yes should be rejected"
|
||||
);
|
||||
|
||||
// Cleanup
|
||||
unsafe {
|
||||
std::env::remove_var("WASM_CHANNELS_ENABLED");
|
||||
}
|
||||
assert!(!cfg.wasm_channels_enabled);
|
||||
}
|
||||
}
|
||||
|
||||
+19
-2
@@ -7,17 +7,19 @@ use crate::settings::Settings;
|
||||
pub struct HeartbeatConfig {
|
||||
/// Whether heartbeat is enabled.
|
||||
pub enabled: bool,
|
||||
/// Interval between heartbeat checks in seconds.
|
||||
/// Interval between heartbeat checks in seconds (used when fire_at is not set).
|
||||
pub interval_secs: u64,
|
||||
/// Channel to notify on heartbeat findings.
|
||||
pub notify_channel: Option<String>,
|
||||
/// User ID to notify on heartbeat findings.
|
||||
pub notify_user: Option<String>,
|
||||
/// Fixed time-of-day to fire (HH:MM, 24h). When set, interval_secs is ignored.
|
||||
pub fire_at: Option<chrono::NaiveTime>,
|
||||
/// Hour (0-23) when quiet hours start.
|
||||
pub quiet_hours_start: Option<u32>,
|
||||
/// Hour (0-23) when quiet hours end.
|
||||
pub quiet_hours_end: Option<u32>,
|
||||
/// Timezone for quiet hours evaluation (IANA name).
|
||||
/// Timezone for fire_at and quiet hours evaluation (IANA name).
|
||||
pub timezone: Option<String>,
|
||||
}
|
||||
|
||||
@@ -28,6 +30,7 @@ impl Default for HeartbeatConfig {
|
||||
interval_secs: 1800, // 30 minutes
|
||||
notify_channel: None,
|
||||
notify_user: None,
|
||||
fire_at: None,
|
||||
quiet_hours_start: None,
|
||||
quiet_hours_end: None,
|
||||
timezone: None,
|
||||
@@ -37,6 +40,19 @@ impl Default for HeartbeatConfig {
|
||||
|
||||
impl HeartbeatConfig {
|
||||
pub(crate) fn resolve(settings: &Settings) -> Result<Self, ConfigError> {
|
||||
let fire_at_str =
|
||||
optional_env("HEARTBEAT_FIRE_AT")?.or_else(|| settings.heartbeat.fire_at.clone());
|
||||
let fire_at = fire_at_str
|
||||
.map(|s| {
|
||||
chrono::NaiveTime::parse_from_str(&s, "%H:%M").map_err(|e| {
|
||||
ConfigError::InvalidValue {
|
||||
key: "HEARTBEAT_FIRE_AT".to_string(),
|
||||
message: format!("must be HH:MM (24h), e.g. '14:00': {e}"),
|
||||
}
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
|
||||
Ok(Self {
|
||||
enabled: parse_bool_env("HEARTBEAT_ENABLED", settings.heartbeat.enabled)?,
|
||||
interval_secs: parse_optional_env(
|
||||
@@ -47,6 +63,7 @@ impl HeartbeatConfig {
|
||||
.or_else(|| settings.heartbeat.notify_channel.clone()),
|
||||
notify_user: optional_env("HEARTBEAT_NOTIFY_USER")?
|
||||
.or_else(|| settings.heartbeat.notify_user.clone()),
|
||||
fire_at,
|
||||
quiet_hours_start: parse_option_env::<u32>("HEARTBEAT_QUIET_START")?
|
||||
.or(settings.heartbeat.quiet_hours_start)
|
||||
.map(|h| {
|
||||
|
||||
+61
-19
@@ -9,7 +9,6 @@ use crate::llm::config::*;
|
||||
use crate::llm::registry::{ProviderProtocol, ProviderRegistry};
|
||||
use crate::llm::session::SessionConfig;
|
||||
use crate::settings::Settings;
|
||||
|
||||
impl LlmConfig {
|
||||
/// Create a test-friendly config without reading env vars.
|
||||
#[cfg(feature = "libsql")]
|
||||
@@ -39,6 +38,8 @@ impl LlmConfig {
|
||||
provider: None,
|
||||
bedrock: None,
|
||||
request_timeout_secs: 120,
|
||||
cheap_model: None,
|
||||
smart_routing_cascade: false,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -169,6 +170,14 @@ impl LlmConfig {
|
||||
|
||||
let request_timeout_secs = parse_optional_env("LLM_REQUEST_TIMEOUT_SECS", 120)?;
|
||||
|
||||
// Generic cheap model (works with any backend).
|
||||
// Falls back to NearAI-specific cheap_model in provider chain logic.
|
||||
let cheap_model = optional_env("LLM_CHEAP_MODEL")?;
|
||||
|
||||
// Generic smart routing cascade flag.
|
||||
// Defaults to true. Overrides NearAI-specific smart_routing_cascade.
|
||||
let smart_routing_cascade = parse_optional_env("SMART_ROUTING_CASCADE", true)?;
|
||||
|
||||
Ok(Self {
|
||||
backend: if is_nearai {
|
||||
"nearai".to_string()
|
||||
@@ -184,6 +193,8 @@ impl LlmConfig {
|
||||
provider,
|
||||
bedrock,
|
||||
request_timeout_secs,
|
||||
cheap_model,
|
||||
smart_routing_cascade,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -241,8 +252,30 @@ impl LlmConfig {
|
||||
)
|
||||
};
|
||||
|
||||
// Resolve API key from env
|
||||
let api_key = if let Some(env_var) = api_key_env {
|
||||
// Codex auth.json override: when LLM_USE_CODEX_AUTH=true,
|
||||
// credentials from the Codex CLI's auth.json take highest priority
|
||||
// (over env vars AND secrets store). In ChatGPT mode, the base URL
|
||||
// is also overridden to the private ChatGPT backend endpoint.
|
||||
let mut codex_base_url_override: Option<String> = None;
|
||||
let codex_creds = if parse_optional_env("LLM_USE_CODEX_AUTH", false)? {
|
||||
let path = optional_env("CODEX_AUTH_PATH")?
|
||||
.map(std::path::PathBuf::from)
|
||||
.unwrap_or_else(crate::llm::codex_auth::default_codex_auth_path);
|
||||
crate::llm::codex_auth::load_codex_credentials(&path)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let codex_refresh_token = codex_creds.as_ref().and_then(|c| c.refresh_token.clone());
|
||||
let codex_auth_path = codex_creds.as_ref().and_then(|c| c.auth_path.clone());
|
||||
|
||||
let api_key = if let Some(creds) = codex_creds {
|
||||
if creds.is_chatgpt_mode {
|
||||
codex_base_url_override = Some(creds.base_url().to_string());
|
||||
}
|
||||
Some(creds.token)
|
||||
} else if let Some(env_var) = api_key_env {
|
||||
// Resolve API key from env (including secrets store overlay)
|
||||
optional_env(env_var)?.map(SecretString::from)
|
||||
} else {
|
||||
None
|
||||
@@ -259,22 +292,28 @@ impl LlmConfig {
|
||||
}
|
||||
}
|
||||
|
||||
// Resolve base URL: env var > settings (backward compat) > registry default
|
||||
let base_url = if let Some(env_var) = base_url_env {
|
||||
optional_env(env_var)?
|
||||
} else {
|
||||
None
|
||||
}
|
||||
.or_else(|| {
|
||||
// Backward compat: check legacy settings fields
|
||||
match backend {
|
||||
"ollama" => settings.ollama_base_url.clone(),
|
||||
"openai_compatible" | "openrouter" => settings.openai_compatible_base_url.clone(),
|
||||
_ => None,
|
||||
}
|
||||
})
|
||||
.or_else(|| default_base_url.map(String::from))
|
||||
.unwrap_or_default();
|
||||
// Resolve base URL: codex override > env var > settings (backward compat) > registry default
|
||||
let is_codex_chatgpt = codex_base_url_override.is_some();
|
||||
let base_url = codex_base_url_override
|
||||
.or_else(|| {
|
||||
if let Some(env_var) = base_url_env {
|
||||
optional_env(env_var).ok().flatten()
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.or_else(|| {
|
||||
// Backward compat: check legacy settings fields
|
||||
match backend {
|
||||
"ollama" => settings.ollama_base_url.clone(),
|
||||
"openai_compatible" | "openrouter" => {
|
||||
settings.openai_compatible_base_url.clone()
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
})
|
||||
.or_else(|| default_base_url.map(String::from))
|
||||
.unwrap_or_default();
|
||||
|
||||
if base_url_required
|
||||
&& base_url.is_empty()
|
||||
@@ -340,6 +379,9 @@ impl LlmConfig {
|
||||
model,
|
||||
extra_headers,
|
||||
oauth_token,
|
||||
is_codex_chatgpt,
|
||||
refresh_token: codex_refresh_token,
|
||||
auth_path: codex_auth_path,
|
||||
cache_retention,
|
||||
unsupported_params,
|
||||
})
|
||||
|
||||
+51
-18
@@ -26,7 +26,7 @@ mod tunnel;
|
||||
mod wasm;
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::{LazyLock, Mutex};
|
||||
use std::sync::{LazyLock, Mutex, Once};
|
||||
|
||||
use crate::error::ConfigError;
|
||||
use crate::settings::Settings;
|
||||
@@ -74,10 +74,12 @@ pub use self::helpers::{env_or_override, set_runtime_env};
|
||||
/// 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()));
|
||||
static WARNED_EXPLICIT_DEFAULT_OWNER_ID: Once = Once::new();
|
||||
|
||||
/// Main configuration for the agent.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct Config {
|
||||
pub owner_id: String,
|
||||
pub database: DatabaseConfig,
|
||||
pub llm: LlmConfig,
|
||||
pub embeddings: EmbeddingsConfig,
|
||||
@@ -118,6 +120,7 @@ impl Config {
|
||||
installed_skills_dir: std::path::PathBuf,
|
||||
) -> Self {
|
||||
Self {
|
||||
owner_id: "default".to_string(),
|
||||
database: DatabaseConfig {
|
||||
backend: DatabaseBackend::LibSql,
|
||||
url: secrecy::SecretString::from("unused://test".to_string()),
|
||||
@@ -228,13 +231,7 @@ impl Config {
|
||||
pub async fn from_env_with_toml(
|
||||
toml_path: Option<&std::path::Path>,
|
||||
) -> Result<Self, ConfigError> {
|
||||
let _ = dotenvy::dotenv();
|
||||
crate::bootstrap::load_ironclaw_env();
|
||||
let mut settings = Settings::load();
|
||||
|
||||
// Overlay TOML config file (values win over JSON settings)
|
||||
Self::apply_toml_overlay(&mut settings, toml_path)?;
|
||||
|
||||
let settings = load_bootstrap_settings(toml_path)?;
|
||||
Self::build(&settings).await
|
||||
}
|
||||
|
||||
@@ -306,26 +303,25 @@ impl Config {
|
||||
|
||||
/// Build config from settings (shared by from_env and from_db).
|
||||
async fn build(settings: &Settings) -> Result<Self, ConfigError> {
|
||||
// Resolve tunnel first so channels can default to loopback when a
|
||||
// tunnel handles external exposure (no need to bind 0.0.0.0).
|
||||
let tunnel = TunnelConfig::resolve(settings)?;
|
||||
let owner_id = resolve_owner_id(settings)?;
|
||||
|
||||
Ok(Self {
|
||||
owner_id: owner_id.clone(),
|
||||
database: DatabaseConfig::resolve()?,
|
||||
llm: LlmConfig::resolve(settings)?,
|
||||
embeddings: EmbeddingsConfig::resolve(settings)?,
|
||||
channels: ChannelsConfig::resolve(settings, tunnel.is_enabled())?,
|
||||
tunnel,
|
||||
tunnel: TunnelConfig::resolve(settings)?,
|
||||
channels: ChannelsConfig::resolve(settings, &owner_id)?,
|
||||
agent: AgentConfig::resolve(settings)?,
|
||||
safety: resolve_safety_config()?,
|
||||
wasm: WasmConfig::resolve()?,
|
||||
safety: resolve_safety_config(settings)?,
|
||||
wasm: WasmConfig::resolve(settings)?,
|
||||
secrets: SecretsConfig::resolve().await?,
|
||||
builder: BuilderModeConfig::resolve()?,
|
||||
builder: BuilderModeConfig::resolve(settings)?,
|
||||
heartbeat: HeartbeatConfig::resolve(settings)?,
|
||||
hygiene: HygieneConfig::resolve()?,
|
||||
routines: RoutineConfig::resolve()?,
|
||||
sandbox: SandboxModeConfig::resolve()?,
|
||||
claude_code: ClaudeCodeConfig::resolve()?,
|
||||
sandbox: SandboxModeConfig::resolve(settings)?,
|
||||
claude_code: ClaudeCodeConfig::resolve(settings)?,
|
||||
skills: SkillsConfig::resolve()?,
|
||||
transcription: TranscriptionConfig::resolve(settings)?,
|
||||
search: WorkspaceSearchConfig::resolve()?,
|
||||
@@ -337,6 +333,43 @@ impl Config {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn load_bootstrap_settings(
|
||||
toml_path: Option<&std::path::Path>,
|
||||
) -> Result<Settings, ConfigError> {
|
||||
let _ = dotenvy::dotenv();
|
||||
crate::bootstrap::load_ironclaw_env();
|
||||
|
||||
let mut settings = Settings::load();
|
||||
Config::apply_toml_overlay(&mut settings, toml_path)?;
|
||||
Ok(settings)
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_owner_id(settings: &Settings) -> Result<String, ConfigError> {
|
||||
let env_owner_id = self::helpers::optional_env("IRONCLAW_OWNER_ID")?;
|
||||
let settings_owner_id = settings.owner_id.clone();
|
||||
let configured_owner_id = env_owner_id.clone().or(settings_owner_id.clone());
|
||||
|
||||
let owner_id = configured_owner_id
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or_else(|| "default".to_string());
|
||||
|
||||
if owner_id == "default"
|
||||
&& (env_owner_id.is_some()
|
||||
|| settings_owner_id
|
||||
.as_deref()
|
||||
.is_some_and(|value| !value.trim().is_empty()))
|
||||
{
|
||||
WARNED_EXPLICIT_DEFAULT_OWNER_ID.call_once(|| {
|
||||
tracing::warn!(
|
||||
"IRONCLAW_OWNER_ID resolved to the legacy 'default' scope explicitly; durable state will keep legacy owner behavior"
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
Ok(owner_id)
|
||||
}
|
||||
|
||||
/// Load API keys from the encrypted secrets store into a thread-safe overlay.
|
||||
///
|
||||
/// This bridges the gap between secrets stored during onboarding and the
|
||||
|
||||
+42
-3
@@ -3,9 +3,48 @@ use crate::error::ConfigError;
|
||||
|
||||
pub use ironclaw_safety::SafetyConfig;
|
||||
|
||||
pub(crate) fn resolve_safety_config() -> Result<SafetyConfig, ConfigError> {
|
||||
pub(crate) fn resolve_safety_config(
|
||||
settings: &crate::settings::Settings,
|
||||
) -> Result<SafetyConfig, ConfigError> {
|
||||
let ss = &settings.safety;
|
||||
Ok(SafetyConfig {
|
||||
max_output_length: parse_optional_env("SAFETY_MAX_OUTPUT_LENGTH", 100_000)?,
|
||||
injection_check_enabled: parse_bool_env("SAFETY_INJECTION_CHECK_ENABLED", true)?,
|
||||
max_output_length: parse_optional_env("SAFETY_MAX_OUTPUT_LENGTH", ss.max_output_length)?,
|
||||
injection_check_enabled: parse_bool_env(
|
||||
"SAFETY_INJECTION_CHECK_ENABLED",
|
||||
ss.injection_check_enabled,
|
||||
)?,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::config::helpers::ENV_MUTEX;
|
||||
use crate::settings::Settings;
|
||||
|
||||
#[test]
|
||||
fn resolve_falls_back_to_settings() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let mut settings = Settings::default();
|
||||
settings.safety.max_output_length = 42;
|
||||
settings.safety.injection_check_enabled = false;
|
||||
|
||||
let cfg = resolve_safety_config(&settings).expect("resolve");
|
||||
assert_eq!(cfg.max_output_length, 42);
|
||||
assert!(!cfg.injection_check_enabled);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn env_overrides_settings() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let mut settings = Settings::default();
|
||||
settings.safety.max_output_length = 42;
|
||||
|
||||
// SAFETY: Under ENV_MUTEX, no concurrent env access.
|
||||
unsafe { std::env::set_var("SAFETY_MAX_OUTPUT_LENGTH", "7") };
|
||||
let cfg = resolve_safety_config(&settings).expect("resolve");
|
||||
unsafe { std::env::remove_var("SAFETY_MAX_OUTPUT_LENGTH") };
|
||||
|
||||
assert_eq!(cfg.max_output_length, 7);
|
||||
}
|
||||
}
|
||||
|
||||
+121
-11
@@ -52,11 +52,20 @@ impl Default for SandboxModeConfig {
|
||||
}
|
||||
|
||||
impl SandboxModeConfig {
|
||||
pub(crate) fn resolve() -> Result<Self, ConfigError> {
|
||||
pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result<Self, ConfigError> {
|
||||
let ss = &settings.sandbox;
|
||||
|
||||
let extra_domains = optional_env("SANDBOX_EXTRA_DOMAINS")?
|
||||
.map(|s| s.split(',').map(|d| d.trim().to_string()).collect())
|
||||
.unwrap_or_default();
|
||||
.unwrap_or_else(|| {
|
||||
if ss.extra_allowed_domains.is_empty() {
|
||||
Vec::new()
|
||||
} else {
|
||||
ss.extra_allowed_domains.clone()
|
||||
}
|
||||
});
|
||||
|
||||
// reaper/orphan fields have no Settings counterpart — env > default only.
|
||||
let reaper_interval_secs: u64 = parse_optional_env("SANDBOX_REAPER_INTERVAL_SECS", 300)?;
|
||||
let orphan_threshold_secs: u64 = parse_optional_env("SANDBOX_ORPHAN_THRESHOLD_SECS", 600)?;
|
||||
|
||||
@@ -76,14 +85,15 @@ impl SandboxModeConfig {
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
enabled: parse_bool_env("SANDBOX_ENABLED", true)?,
|
||||
policy: parse_string_env("SANDBOX_POLICY", "readonly")?,
|
||||
enabled: parse_bool_env("SANDBOX_ENABLED", ss.enabled)?,
|
||||
policy: parse_string_env("SANDBOX_POLICY", ss.policy.clone())?,
|
||||
// allow_full_access has no Settings counterpart — env > default only.
|
||||
allow_full_access: parse_bool_env("SANDBOX_ALLOW_FULL_ACCESS", false)?,
|
||||
timeout_secs: parse_optional_env("SANDBOX_TIMEOUT_SECS", 120)?,
|
||||
memory_limit_mb: parse_optional_env("SANDBOX_MEMORY_LIMIT_MB", 2048)?,
|
||||
cpu_shares: parse_optional_env("SANDBOX_CPU_SHARES", 1024)?,
|
||||
image: parse_string_env("SANDBOX_IMAGE", "ironclaw-worker:latest")?,
|
||||
auto_pull_image: parse_bool_env("SANDBOX_AUTO_PULL", true)?,
|
||||
timeout_secs: parse_optional_env("SANDBOX_TIMEOUT_SECS", ss.timeout_secs)?,
|
||||
memory_limit_mb: parse_optional_env("SANDBOX_MEMORY_LIMIT_MB", ss.memory_limit_mb)?,
|
||||
cpu_shares: parse_optional_env("SANDBOX_CPU_SHARES", ss.cpu_shares)?,
|
||||
image: parse_string_env("SANDBOX_IMAGE", ss.image.clone())?,
|
||||
auto_pull_image: parse_bool_env("SANDBOX_AUTO_PULL", ss.auto_pull_image)?,
|
||||
extra_allowed_domains: extra_domains,
|
||||
reaper_interval_secs,
|
||||
orphan_threshold_secs,
|
||||
@@ -200,7 +210,7 @@ impl ClaudeCodeConfig {
|
||||
/// Load from environment variables only (used inside containers where
|
||||
/// there is no database or full config).
|
||||
pub fn from_env() -> Self {
|
||||
match Self::resolve() {
|
||||
match Self::resolve_env_only() {
|
||||
Ok(c) => c,
|
||||
Err(e) => {
|
||||
tracing::warn!("Failed to resolve ClaudeCodeConfig: {e}, using defaults");
|
||||
@@ -253,7 +263,33 @@ impl ClaudeCodeConfig {
|
||||
None
|
||||
}
|
||||
|
||||
pub(crate) fn resolve() -> Result<Self, ConfigError> {
|
||||
pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result<Self, ConfigError> {
|
||||
let defaults = Self::default();
|
||||
Ok(Self {
|
||||
// Use settings.sandbox.claude_code_enabled as fallback (written by setup wizard).
|
||||
enabled: parse_bool_env("CLAUDE_CODE_ENABLED", settings.sandbox.claude_code_enabled)?,
|
||||
config_dir: optional_env("CLAUDE_CONFIG_DIR")?
|
||||
.map(std::path::PathBuf::from)
|
||||
.unwrap_or(defaults.config_dir),
|
||||
model: parse_string_env("CLAUDE_CODE_MODEL", defaults.model)?,
|
||||
max_turns: parse_optional_env("CLAUDE_CODE_MAX_TURNS", defaults.max_turns)?,
|
||||
memory_limit_mb: parse_optional_env(
|
||||
"CLAUDE_CODE_MEMORY_LIMIT_MB",
|
||||
defaults.memory_limit_mb,
|
||||
)?,
|
||||
allowed_tools: optional_env("CLAUDE_CODE_ALLOWED_TOOLS")?
|
||||
.map(|s| {
|
||||
s.split(',')
|
||||
.map(|t| t.trim().to_string())
|
||||
.filter(|t| !t.is_empty())
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or(defaults.allowed_tools),
|
||||
})
|
||||
}
|
||||
|
||||
/// Resolve from env vars only, no Settings. Used inside containers.
|
||||
fn resolve_env_only() -> Result<Self, ConfigError> {
|
||||
let defaults = Self::default();
|
||||
Ok(Self {
|
||||
enabled: parse_bool_env("CLAUDE_CODE_ENABLED", defaults.enabled)?,
|
||||
@@ -554,6 +590,80 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
// ── Settings fallback tests ──────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn sandbox_resolve_falls_back_to_settings() {
|
||||
let _guard = crate::config::helpers::ENV_MUTEX
|
||||
.lock()
|
||||
.expect("env mutex poisoned");
|
||||
let mut settings = crate::settings::Settings::default();
|
||||
settings.sandbox.cpu_shares = 99;
|
||||
settings.sandbox.auto_pull_image = false;
|
||||
settings.sandbox.enabled = false;
|
||||
|
||||
let cfg = SandboxModeConfig::resolve(&settings).expect("resolve");
|
||||
assert!(!cfg.enabled);
|
||||
assert_eq!(cfg.cpu_shares, 99);
|
||||
assert!(!cfg.auto_pull_image);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sandbox_env_overrides_settings() {
|
||||
let _guard = crate::config::helpers::ENV_MUTEX
|
||||
.lock()
|
||||
.expect("env mutex poisoned");
|
||||
let mut settings = crate::settings::Settings::default();
|
||||
settings.sandbox.timeout_secs = 999;
|
||||
|
||||
// SAFETY: Under ENV_MUTEX, no concurrent env access.
|
||||
unsafe { std::env::set_var("SANDBOX_TIMEOUT_SECS", "5") };
|
||||
let cfg = SandboxModeConfig::resolve(&settings).expect("resolve");
|
||||
unsafe { std::env::remove_var("SANDBOX_TIMEOUT_SECS") };
|
||||
|
||||
assert_eq!(cfg.timeout_secs, 5);
|
||||
}
|
||||
|
||||
// ── ClaudeCodeConfig settings fallback tests ────────────────────
|
||||
|
||||
#[test]
|
||||
fn claude_code_resolve_uses_settings_enabled() {
|
||||
let _guard = crate::config::helpers::ENV_MUTEX
|
||||
.lock()
|
||||
.expect("env mutex poisoned");
|
||||
let mut settings = crate::settings::Settings::default();
|
||||
settings.sandbox.claude_code_enabled = true;
|
||||
|
||||
let cfg = ClaudeCodeConfig::resolve(&settings).expect("resolve");
|
||||
assert!(cfg.enabled);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_code_resolve_defaults_disabled() {
|
||||
let _guard = crate::config::helpers::ENV_MUTEX
|
||||
.lock()
|
||||
.expect("env mutex poisoned");
|
||||
let settings = crate::settings::Settings::default();
|
||||
let cfg = ClaudeCodeConfig::resolve(&settings).expect("resolve");
|
||||
assert!(!cfg.enabled);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn claude_code_env_overrides_settings() {
|
||||
let _guard = crate::config::helpers::ENV_MUTEX
|
||||
.lock()
|
||||
.expect("env mutex poisoned");
|
||||
let mut settings = crate::settings::Settings::default();
|
||||
settings.sandbox.claude_code_enabled = true;
|
||||
|
||||
// SAFETY: Under ENV_MUTEX, no concurrent env access.
|
||||
unsafe { std::env::set_var("CLAUDE_CODE_ENABLED", "false") };
|
||||
let cfg = ClaudeCodeConfig::resolve(&settings).expect("resolve");
|
||||
unsafe { std::env::remove_var("CLAUDE_CODE_ENABLED") };
|
||||
|
||||
assert!(!cfg.enabled);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_readonly_policy_unaffected() {
|
||||
let config = SandboxModeConfig {
|
||||
|
||||
+64
-13
@@ -9,11 +9,15 @@ use crate::settings::Settings;
|
||||
pub struct TranscriptionConfig {
|
||||
/// Whether audio transcription is enabled.
|
||||
pub enabled: bool,
|
||||
/// Provider: "openai" (default).
|
||||
/// Provider: "openai" (default) or "chat_completions".
|
||||
pub provider: String,
|
||||
/// OpenAI API key (reuses OPENAI_API_KEY).
|
||||
pub openai_api_key: Option<SecretString>,
|
||||
/// Model to use (default: "whisper-1").
|
||||
/// Explicit transcription API key (overrides provider-specific keys).
|
||||
pub api_key: Option<SecretString>,
|
||||
/// LLM API key (reuses LLM_API_KEY, used as fallback for chat_completions).
|
||||
pub llm_api_key: Option<SecretString>,
|
||||
/// Model to use (default depends on provider).
|
||||
pub model: String,
|
||||
/// Base URL override for the transcription API.
|
||||
pub base_url: Option<String>,
|
||||
@@ -25,6 +29,8 @@ impl Default for TranscriptionConfig {
|
||||
enabled: false,
|
||||
provider: "openai".to_string(),
|
||||
openai_api_key: None,
|
||||
api_key: None,
|
||||
llm_api_key: None,
|
||||
model: "whisper-1".to_string(),
|
||||
base_url: None,
|
||||
}
|
||||
@@ -42,8 +48,15 @@ impl TranscriptionConfig {
|
||||
optional_env("TRANSCRIPTION_PROVIDER")?.unwrap_or_else(|| "openai".to_string());
|
||||
|
||||
let openai_api_key = optional_env("OPENAI_API_KEY")?.map(SecretString::from);
|
||||
let api_key = optional_env("TRANSCRIPTION_API_KEY")?.map(SecretString::from);
|
||||
let llm_api_key = optional_env("LLM_API_KEY")?.map(SecretString::from);
|
||||
|
||||
let model = optional_env("TRANSCRIPTION_MODEL")?.unwrap_or_else(|| "whisper-1".to_string());
|
||||
let default_model = match provider.as_str() {
|
||||
"chat_completions" => "google/gemini-2.0-flash-001",
|
||||
_ => "whisper-1",
|
||||
};
|
||||
let model =
|
||||
optional_env("TRANSCRIPTION_MODEL")?.unwrap_or_else(|| default_model.to_string());
|
||||
|
||||
let base_url = optional_env("TRANSCRIPTION_BASE_URL")?;
|
||||
|
||||
@@ -51,29 +64,67 @@ impl TranscriptionConfig {
|
||||
enabled,
|
||||
provider,
|
||||
openai_api_key,
|
||||
api_key,
|
||||
llm_api_key,
|
||||
model,
|
||||
base_url,
|
||||
})
|
||||
}
|
||||
|
||||
/// Resolve the API key for the configured provider.
|
||||
///
|
||||
/// Priority: `TRANSCRIPTION_API_KEY` > provider-specific key.
|
||||
fn resolve_api_key(&self) -> Option<&SecretString> {
|
||||
self.api_key
|
||||
.as_ref()
|
||||
.or_else(|| match self.provider.as_str() {
|
||||
"chat_completions" => self.llm_api_key.as_ref().or(self.openai_api_key.as_ref()),
|
||||
_ => self.openai_api_key.as_ref(),
|
||||
})
|
||||
}
|
||||
|
||||
/// Create the transcription provider if enabled and configured.
|
||||
pub fn create_provider(&self) -> Option<Box<dyn crate::transcription::TranscriptionProvider>> {
|
||||
if !self.enabled {
|
||||
return None;
|
||||
}
|
||||
|
||||
// Currently only OpenAI Whisper is supported; more providers can be
|
||||
// added here with a match on self.provider.
|
||||
let api_key = self.openai_api_key.as_ref()?;
|
||||
tracing::info!(model = %self.model, "Audio transcription enabled via OpenAI Whisper");
|
||||
let api_key = self.resolve_api_key()?;
|
||||
|
||||
let mut provider = crate::transcription::OpenAiWhisperProvider::new(api_key.clone())
|
||||
.with_model(&self.model);
|
||||
match self.provider.as_str() {
|
||||
"chat_completions" => {
|
||||
tracing::info!(
|
||||
model = %self.model,
|
||||
"Audio transcription enabled via Chat Completions API"
|
||||
);
|
||||
|
||||
if let Some(ref base_url) = self.base_url {
|
||||
provider = provider.with_base_url(base_url);
|
||||
let mut provider = crate::transcription::ChatCompletionsTranscriptionProvider::new(
|
||||
api_key.clone(),
|
||||
)
|
||||
.with_model(&self.model);
|
||||
|
||||
if let Some(ref base_url) = self.base_url {
|
||||
provider = provider.with_base_url(base_url);
|
||||
}
|
||||
|
||||
Some(Box::new(provider))
|
||||
}
|
||||
_ => {
|
||||
tracing::info!(
|
||||
model = %self.model,
|
||||
"Audio transcription enabled via OpenAI Whisper"
|
||||
);
|
||||
|
||||
let mut provider =
|
||||
crate::transcription::OpenAiWhisperProvider::new(api_key.clone())
|
||||
.with_model(&self.model);
|
||||
|
||||
if let Some(ref base_url) = self.base_url {
|
||||
provider = provider.with_base_url(base_url);
|
||||
}
|
||||
|
||||
Some(Box::new(provider))
|
||||
}
|
||||
}
|
||||
|
||||
Some(Box::new(provider))
|
||||
}
|
||||
}
|
||||
|
||||
+50
-7
@@ -44,20 +44,30 @@ fn default_tools_dir() -> PathBuf {
|
||||
}
|
||||
|
||||
impl WasmConfig {
|
||||
pub(crate) fn resolve() -> Result<Self, ConfigError> {
|
||||
pub(crate) fn resolve(settings: &crate::settings::Settings) -> Result<Self, ConfigError> {
|
||||
let ws = &settings.wasm;
|
||||
Ok(Self {
|
||||
enabled: parse_bool_env("WASM_ENABLED", true)?,
|
||||
enabled: parse_bool_env("WASM_ENABLED", ws.enabled)?,
|
||||
tools_dir: optional_env("WASM_TOOLS_DIR")?
|
||||
.map(PathBuf::from)
|
||||
.or_else(|| ws.tools_dir.clone())
|
||||
.unwrap_or_else(default_tools_dir),
|
||||
default_memory_limit: parse_optional_env(
|
||||
"WASM_DEFAULT_MEMORY_LIMIT",
|
||||
10 * 1024 * 1024,
|
||||
ws.default_memory_limit,
|
||||
)?,
|
||||
default_timeout_secs: parse_optional_env("WASM_DEFAULT_TIMEOUT_SECS", 60)?,
|
||||
default_fuel_limit: parse_optional_env("WASM_DEFAULT_FUEL_LIMIT", 10_000_000)?,
|
||||
cache_compiled: parse_bool_env("WASM_CACHE_COMPILED", true)?,
|
||||
cache_dir: optional_env("WASM_CACHE_DIR")?.map(PathBuf::from),
|
||||
default_timeout_secs: parse_optional_env(
|
||||
"WASM_DEFAULT_TIMEOUT_SECS",
|
||||
ws.default_timeout_secs,
|
||||
)?,
|
||||
default_fuel_limit: parse_optional_env(
|
||||
"WASM_DEFAULT_FUEL_LIMIT",
|
||||
ws.default_fuel_limit,
|
||||
)?,
|
||||
cache_compiled: parse_bool_env("WASM_CACHE_COMPILED", ws.cache_compiled)?,
|
||||
cache_dir: optional_env("WASM_CACHE_DIR")?
|
||||
.map(PathBuf::from)
|
||||
.or_else(|| ws.cache_dir.clone()),
|
||||
})
|
||||
}
|
||||
|
||||
@@ -81,3 +91,36 @@ impl WasmConfig {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::config::helpers::ENV_MUTEX;
|
||||
use crate::settings::Settings;
|
||||
|
||||
#[test]
|
||||
fn resolve_falls_back_to_settings() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let mut settings = Settings::default();
|
||||
settings.wasm.default_memory_limit = 42;
|
||||
settings.wasm.cache_compiled = false;
|
||||
|
||||
let cfg = WasmConfig::resolve(&settings).expect("resolve");
|
||||
assert_eq!(cfg.default_memory_limit, 42);
|
||||
assert!(!cfg.cache_compiled);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn env_overrides_settings() {
|
||||
let _guard = ENV_MUTEX.lock().expect("env mutex poisoned");
|
||||
let mut settings = Settings::default();
|
||||
settings.wasm.default_fuel_limit = 42;
|
||||
|
||||
// SAFETY: Under ENV_MUTEX, no concurrent env access.
|
||||
unsafe { std::env::set_var("WASM_DEFAULT_FUEL_LIMIT", "7") };
|
||||
let cfg = WasmConfig::resolve(&settings).expect("resolve");
|
||||
unsafe { std::env::remove_var("WASM_DEFAULT_FUEL_LIMIT") };
|
||||
|
||||
assert_eq!(cfg.default_fuel_limit, 7);
|
||||
}
|
||||
}
|
||||
|
||||
+223
-3
@@ -46,11 +46,17 @@ impl ContextManager {
|
||||
description: impl Into<String>,
|
||||
) -> Result<Uuid, JobError> {
|
||||
// Hold write lock for the entire check-insert to prevent TOCTOU races
|
||||
// where two concurrent calls both pass the active_count check.
|
||||
// where two concurrent calls both pass the parallel_count check.
|
||||
let mut contexts = self.contexts.write().await;
|
||||
let active_count = contexts.values().filter(|c| c.state.is_active()).count();
|
||||
// Only count jobs that consume execution slots (Pending, InProgress, Stuck).
|
||||
// Completed and Submitted jobs are no longer actively executing and shouldn't
|
||||
// block new job creation.
|
||||
let parallel_count = contexts
|
||||
.values()
|
||||
.filter(|c| c.state.is_parallel_blocking())
|
||||
.count();
|
||||
|
||||
if active_count >= self.max_jobs {
|
||||
if parallel_count >= self.max_jobs {
|
||||
return Err(JobError::MaxJobsExceeded { max: self.max_jobs });
|
||||
}
|
||||
|
||||
@@ -965,4 +971,218 @@ mod tests {
|
||||
// And it's in the initial state (Pending), not modified by concurrent workers
|
||||
assert_eq!(returned_ctx.state, crate::context::JobState::Pending); // safety: test code
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sequential_routines_unlimited_completed_not_counted() {
|
||||
// TEST: Sequential (non-parallel) routines should NOT be limited by max_jobs.
|
||||
//
|
||||
// Completed/Submitted jobs should NOT count toward the parallel job limit,
|
||||
// since they're no longer actively consuming execution resources.
|
||||
//
|
||||
// Scenario: Create 10 sequential routines, each completing before the next starts.
|
||||
// Currently FAILS because Completed jobs still count as "active".
|
||||
// After fix, should PASS because only Pending/InProgress/Stuck count.
|
||||
|
||||
let manager = ContextManager::new(5); // max 5 truly parallel jobs
|
||||
|
||||
// Try to create and complete 10 sequential routines
|
||||
for i in 0..10 {
|
||||
let result = manager
|
||||
.create_job(format!("Sequential Routine {}", i), "one at a time")
|
||||
.await;
|
||||
|
||||
match result {
|
||||
Ok(job_id) => {
|
||||
// Simulate execution: Pending -> InProgress -> Completed
|
||||
manager
|
||||
.update_context(job_id, |ctx| {
|
||||
ctx.transition_to(crate::context::JobState::InProgress, None)
|
||||
})
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
|
||||
manager
|
||||
.update_context(job_id, |ctx| {
|
||||
ctx.transition_to(crate::context::JobState::Completed, None)
|
||||
})
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
|
||||
println!("✓ Routine {} created and completed", i);
|
||||
}
|
||||
Err(JobError::MaxJobsExceeded { max }) => {
|
||||
panic!(
|
||||
"✗ Routine {} FAILED to create: MaxJobsExceeded (max={}).\n\
|
||||
This shows the bug: Completed jobs from routines 0-4 are still counting \
|
||||
toward the limit even though they're not running.\n\
|
||||
After the fix, this test should pass because Completed jobs won't count.",
|
||||
i, max
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
panic!("Unexpected error for routine {}: {:?}", i, e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// If we reach here, all 10 routines succeeded (bug is fixed)
|
||||
assert_eq!(manager.all_jobs().await.len(), 10);
|
||||
println!("✓ SUCCESS: All 10 sequential routines created despite max_jobs=5 limit");
|
||||
println!(" This is correct: Completed jobs don't count toward parallel limit");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn parallel_jobs_limit_enforced_for_active_jobs() {
|
||||
// TEST: Parallel (simultaneous) jobs ARE limited by max_jobs.
|
||||
//
|
||||
// Jobs in Pending/InProgress/Stuck states consume execution slots.
|
||||
// The 6th truly-active job should fail because the limit is 5.
|
||||
//
|
||||
// This test verifies the limit DOES work correctly for parallel execution.
|
||||
|
||||
let manager = ContextManager::new(5); // max 5 parallel jobs
|
||||
|
||||
// Create 5 jobs and make them InProgress (simulating parallel execution)
|
||||
let mut job_ids = Vec::new();
|
||||
for i in 0..5 {
|
||||
let job_id = manager
|
||||
.create_job(format!("Parallel Job {}", i), "running in parallel")
|
||||
.await
|
||||
.expect("First 5 jobs should create successfully");
|
||||
job_ids.push(job_id);
|
||||
|
||||
// Transition to InProgress (simulating active execution)
|
||||
manager
|
||||
.update_context(job_id, |ctx| {
|
||||
ctx.transition_to(crate::context::JobState::InProgress, None)
|
||||
})
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
// Verify all 5 jobs are InProgress
|
||||
for job_id in &job_ids {
|
||||
let ctx = manager.get_context(*job_id).await.unwrap();
|
||||
assert_eq!(
|
||||
ctx.state,
|
||||
crate::context::JobState::InProgress,
|
||||
"All jobs should be InProgress"
|
||||
);
|
||||
}
|
||||
|
||||
// Check active count - should be 5 (all InProgress)
|
||||
let active_count = manager.active_count().await;
|
||||
assert_eq!(
|
||||
active_count, 5,
|
||||
"Active count should be 5 (all InProgress jobs count)"
|
||||
);
|
||||
|
||||
// Try to create a 6th job - should FAIL because limit is reached
|
||||
let result = manager.create_job("Parallel Job 6", "sixth job").await;
|
||||
|
||||
match result {
|
||||
Err(JobError::MaxJobsExceeded { max: 5 }) => {
|
||||
println!("✓ SUCCESS: Parallel job limit correctly enforced at 5 active jobs");
|
||||
println!("✓ 6th InProgress job correctly blocked when 5 are already running");
|
||||
}
|
||||
Ok(_) => {
|
||||
panic!(
|
||||
"FAILED: 6th parallel job should have been blocked \
|
||||
but was created. Limit enforcement is broken."
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
panic!(
|
||||
"UNEXPECTED ERROR: Expected MaxJobsExceeded but got: {:?}",
|
||||
e
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn completed_jobs_should_free_slots_after_fix() {
|
||||
// TEST: After the fix, Completed jobs should NOT count toward the limit.
|
||||
//
|
||||
// This test demonstrates that when a job transitions from InProgress -> Completed,
|
||||
// it should free up a slot in the parallel execution limit.
|
||||
//
|
||||
// Currently FAILS (bug not fixed), proving Completed jobs incorrectly stay in the limit.
|
||||
// After fix, this will PASS (Completed jobs freed their slot).
|
||||
|
||||
let manager = ContextManager::new(5); // max 5 parallel jobs
|
||||
|
||||
// Create 5 InProgress jobs (fill the limit)
|
||||
let mut job_ids = Vec::new();
|
||||
for i in 0..5 {
|
||||
let job_id = manager
|
||||
.create_job(format!("Job {}", i), "parallel")
|
||||
.await
|
||||
.unwrap();
|
||||
job_ids.push(job_id);
|
||||
|
||||
manager
|
||||
.update_context(job_id, |ctx| {
|
||||
ctx.transition_to(crate::context::JobState::InProgress, None)
|
||||
})
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
// Verify limit is hit
|
||||
let result = manager.create_job("Job 5", "should fail").await;
|
||||
assert!(
|
||||
matches!(result, Err(JobError::MaxJobsExceeded { max: 5 })),
|
||||
"Limit should be hit with 5 InProgress jobs"
|
||||
);
|
||||
println!("✓ Limit enforced: 5 InProgress jobs block 6th creation");
|
||||
|
||||
// Now transition job 0 from InProgress -> Completed
|
||||
manager
|
||||
.update_context(job_ids[0], |ctx| {
|
||||
ctx.transition_to(crate::context::JobState::Completed, None)
|
||||
})
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
|
||||
println!("✓ Job 0 transitioned: InProgress -> Completed");
|
||||
|
||||
// Try to create a 6th job - this will FAIL until the bug is fixed
|
||||
let result = manager
|
||||
.create_job("Job 5 (retry)", "after 1 Completed")
|
||||
.await;
|
||||
|
||||
match result {
|
||||
Ok(job_6) => {
|
||||
println!("✓ SUCCESS: 6th job created after job 0 completed");
|
||||
println!("✓ This proves Completed jobs don't count toward the limit (BUG FIXED)");
|
||||
|
||||
// Verify we can transition it to InProgress
|
||||
manager
|
||||
.update_context(job_6, |ctx| {
|
||||
ctx.transition_to(crate::context::JobState::InProgress, None)
|
||||
})
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
println!("✓ 6th job now InProgress: 4 remaining + 1 new = 5 limit reached");
|
||||
}
|
||||
Err(JobError::MaxJobsExceeded { max: 5 }) => {
|
||||
panic!(
|
||||
"✗ BUG NOT FIXED: 6th job creation still blocked after freeing slot.\n\
|
||||
State: 1 Completed (job 0) + 4 InProgress (jobs 1-4) = 5 active\n\
|
||||
BUG: Completed job 0 still counts toward limit\n\
|
||||
EXPECTED: Only 4 InProgress count, 1 slot free"
|
||||
);
|
||||
}
|
||||
Err(e) => {
|
||||
panic!("Unexpected error: {:?}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -48,6 +48,14 @@ impl JobState {
|
||||
pub fn can_transition_to(&self, target: JobState) -> bool {
|
||||
use JobState::*;
|
||||
|
||||
// Allow idempotent Completed -> Completed transition.
|
||||
// Both the execution loop and the worker wrapper may race to mark a
|
||||
// job complete; the second call should be a harmless no-op rather
|
||||
// than an error that masks the successful completion.
|
||||
if matches!((self, target), (Completed, Completed)) {
|
||||
return true;
|
||||
}
|
||||
|
||||
matches!(
|
||||
(self, target),
|
||||
// From Pending
|
||||
@@ -73,6 +81,15 @@ impl JobState {
|
||||
pub fn is_active(&self) -> bool {
|
||||
!self.is_terminal()
|
||||
}
|
||||
|
||||
/// Check if this job consumes a parallel execution slot.
|
||||
///
|
||||
/// Only jobs in Pending, InProgress, or Stuck states consume execution resources
|
||||
/// and should count toward the parallel job limit. Completed and Submitted jobs
|
||||
/// are in the state machine but are no longer actively executing.
|
||||
pub fn is_parallel_blocking(&self) -> bool {
|
||||
matches!(self, Self::Pending | Self::InProgress | Self::Stuck)
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for JobState {
|
||||
@@ -113,6 +130,9 @@ pub struct JobContext {
|
||||
pub state: JobState,
|
||||
/// User ID that owns this job (for workspace scoping).
|
||||
pub user_id: String,
|
||||
/// Channel-specific requester/actor ID, when different from the owner scope.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub requester_id: Option<String>,
|
||||
/// Conversation ID if linked to a conversation.
|
||||
pub conversation_id: Option<Uuid>,
|
||||
/// Job title.
|
||||
@@ -194,6 +214,7 @@ impl JobContext {
|
||||
job_id: Uuid::new_v4(),
|
||||
state: JobState::Pending,
|
||||
user_id: user_id.into(),
|
||||
requester_id: None,
|
||||
conversation_id: None,
|
||||
title: title.into(),
|
||||
description: description.into(),
|
||||
@@ -225,6 +246,12 @@ impl JobContext {
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the channel-specific requester/actor ID.
|
||||
pub fn with_requester_id(mut self, requester_id: impl Into<String>) -> Self {
|
||||
self.requester_id = Some(requester_id.into());
|
||||
self
|
||||
}
|
||||
|
||||
/// Transition to a new state.
|
||||
pub fn transition_to(
|
||||
&mut self,
|
||||
@@ -238,6 +265,18 @@ impl JobContext {
|
||||
));
|
||||
}
|
||||
|
||||
// Idempotent: already in the target state, skip recording a duplicate
|
||||
// transition. This handles the Completed -> Completed race between
|
||||
// execution_loop and the worker wrapper.
|
||||
if self.state == new_state {
|
||||
tracing::debug!(
|
||||
job_id = %self.job_id,
|
||||
state = %self.state,
|
||||
"idempotent state transition (already in target state), skipping"
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let transition = StateTransition {
|
||||
from: self.state,
|
||||
to: new_state,
|
||||
@@ -340,6 +379,45 @@ mod tests {
|
||||
assert!(!JobState::Accepted.can_transition_to(JobState::InProgress));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_completed_to_completed_is_idempotent() {
|
||||
// Regression test for the race condition where both execution_loop
|
||||
// and the worker wrapper call mark_completed(). The second call
|
||||
// must succeed without error and must not record a duplicate
|
||||
// transition.
|
||||
let mut ctx = JobContext::new("Test", "Idempotent completion test");
|
||||
ctx.transition_to(JobState::InProgress, None).unwrap();
|
||||
ctx.transition_to(JobState::Completed, Some("first".into()))
|
||||
.unwrap();
|
||||
assert_eq!(ctx.state, JobState::Completed);
|
||||
let transitions_before = ctx.transitions.len();
|
||||
|
||||
// Second Completed -> Completed must be a no-op
|
||||
let result = ctx.transition_to(JobState::Completed, Some("duplicate".into()));
|
||||
assert!(
|
||||
result.is_ok(),
|
||||
"Completed -> Completed should be idempotent"
|
||||
);
|
||||
assert_eq!(ctx.state, JobState::Completed);
|
||||
assert_eq!(
|
||||
ctx.transitions.len(),
|
||||
transitions_before,
|
||||
"idempotent transition should not record a new history entry"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_other_self_transitions_still_rejected() {
|
||||
// Ensure we only allow Completed -> Completed, not arbitrary X -> X.
|
||||
assert!(!JobState::Pending.can_transition_to(JobState::Pending));
|
||||
assert!(!JobState::InProgress.can_transition_to(JobState::InProgress));
|
||||
assert!(!JobState::Failed.can_transition_to(JobState::Failed));
|
||||
assert!(!JobState::Stuck.can_transition_to(JobState::Stuck));
|
||||
assert!(!JobState::Submitted.can_transition_to(JobState::Submitted));
|
||||
assert!(!JobState::Accepted.can_transition_to(JobState::Accepted));
|
||||
assert!(!JobState::Cancelled.can_transition_to(JobState::Cancelled));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_terminal_states() {
|
||||
assert!(JobState::Accepted.is_terminal());
|
||||
|
||||
@@ -106,6 +106,7 @@ impl JobStore for LibSqlBackend {
|
||||
job_id: get_text(&row, 0).parse().unwrap_or_default(),
|
||||
state,
|
||||
user_id: get_text(&row, 6),
|
||||
requester_id: None,
|
||||
conversation_id: get_opt_text(&row, 1).and_then(|s| s.parse().ok()),
|
||||
title: get_text(&row, 2),
|
||||
description: get_text(&row, 3),
|
||||
|
||||
+23
-2
@@ -247,6 +247,17 @@ pub(crate) fn opt_text_owned(s: Option<String>) -> libsql::Value {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn normalize_notify_user(value: Option<String>) -> Option<String> {
|
||||
value.and_then(|value| {
|
||||
let trimmed = value.trim();
|
||||
if trimmed.is_empty() || trimmed == "default" {
|
||||
None
|
||||
} else {
|
||||
Some(trimmed.to_string())
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/// Extract an i64 column, defaulting to 0.
|
||||
pub(crate) fn get_i64(row: &libsql::Row, idx: i32) -> i64 {
|
||||
row.get::<i64>(idx).unwrap_or(0)
|
||||
@@ -378,7 +389,7 @@ pub(crate) fn row_to_routine_libsql(row: &libsql::Row) -> Result<Routine, Databa
|
||||
},
|
||||
notify: NotifyConfig {
|
||||
channel: get_opt_text(row, 12),
|
||||
user: get_text(row, 13),
|
||||
user: normalize_notify_user(get_opt_text(row, 13)),
|
||||
on_success: get_i64(row, 14) != 0,
|
||||
on_failure: get_i64(row, 15) != 0,
|
||||
on_attention: get_i64(row, 16) != 0,
|
||||
@@ -419,7 +430,17 @@ mod tests {
|
||||
use chrono::{TimeZone, Utc};
|
||||
|
||||
use crate::db::Database;
|
||||
use crate::db::libsql::{LibSqlBackend, parse_timestamp};
|
||||
use crate::db::libsql::{LibSqlBackend, normalize_notify_user, parse_timestamp};
|
||||
|
||||
#[test]
|
||||
fn test_normalize_notify_user_treats_legacy_default_as_missing() {
|
||||
assert_eq!(normalize_notify_user(None), None); // safety: test-only assertion
|
||||
assert_eq!(normalize_notify_user(Some(String::new())), None); // safety: test-only assertion
|
||||
assert_eq!(normalize_notify_user(Some(" ".to_string())), None); // safety: test-only assertion
|
||||
assert_eq!(normalize_notify_user(Some("default".to_string())), None); // safety: test-only assertion
|
||||
let normalized = normalize_notify_user(Some("123456789".to_string()));
|
||||
assert_eq!(normalized, Some("123456789".to_string())); // safety: test-only assertion
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_timestamp_accepts_rfc3339_and_legacy_naive_formats() {
|
||||
|
||||
@@ -57,7 +57,7 @@ impl RoutineStore for LibSqlBackend {
|
||||
max_concurrent,
|
||||
dedup_window_secs,
|
||||
opt_text(routine.notify.channel.as_deref()),
|
||||
routine.notify.user.as_str(),
|
||||
opt_text(routine.notify.user.as_deref()),
|
||||
routine.notify.on_success as i64,
|
||||
routine.notify.on_failure as i64,
|
||||
routine.notify.on_attention as i64,
|
||||
@@ -250,7 +250,7 @@ impl RoutineStore for LibSqlBackend {
|
||||
max_concurrent,
|
||||
dedup_window_secs,
|
||||
opt_text(routine.notify.channel.as_deref()),
|
||||
routine.notify.user.as_str(),
|
||||
opt_text(routine.notify.user.as_deref()),
|
||||
routine.notify.on_success as i64,
|
||||
routine.notify.on_failure as i64,
|
||||
routine.notify.on_attention as i64,
|
||||
|
||||
@@ -462,7 +462,7 @@ CREATE TABLE IF NOT EXISTS routines (
|
||||
max_concurrent INTEGER NOT NULL DEFAULT 1,
|
||||
dedup_window_secs INTEGER,
|
||||
notify_channel TEXT,
|
||||
notify_user TEXT NOT NULL DEFAULT 'default',
|
||||
notify_user TEXT,
|
||||
notify_on_success INTEGER NOT NULL DEFAULT 0,
|
||||
notify_on_failure INTEGER NOT NULL DEFAULT 1,
|
||||
notify_on_attention INTEGER NOT NULL DEFAULT 1,
|
||||
@@ -546,7 +546,9 @@ CREATE INDEX IF NOT EXISTS idx_tool_failures_unrepaired ON tool_failures(tool_na
|
||||
|
||||
-- routines
|
||||
CREATE INDEX IF NOT EXISTS idx_routines_next_fire ON routines(next_fire_at);
|
||||
CREATE INDEX IF NOT EXISTS idx_routines_event_triggers ON routines(user_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_routines_event_triggers
|
||||
ON routines(trigger_type, user_id)
|
||||
WHERE enabled = 1 AND trigger_type IN ('event', 'system_event');
|
||||
|
||||
-- routine_runs
|
||||
CREATE INDEX IF NOT EXISTS idx_routine_runs_status ON routine_runs(status);
|
||||
@@ -654,6 +656,74 @@ END;
|
||||
r#"
|
||||
ALTER TABLE agent_jobs ADD COLUMN max_tokens INTEGER NOT NULL DEFAULT 0;
|
||||
ALTER TABLE agent_jobs ADD COLUMN total_tokens_used INTEGER NOT NULL DEFAULT 0;
|
||||
"#,
|
||||
),
|
||||
(
|
||||
13,
|
||||
"routine_notify_user_nullable",
|
||||
// Remove the legacy 'default' sentinel from routine notify_user.
|
||||
// SQLite cannot drop NOT NULL / DEFAULT constraints in place, so we
|
||||
// rebuild the table and normalize existing 'default' values to NULL.
|
||||
r#"
|
||||
PRAGMA foreign_keys=OFF;
|
||||
|
||||
CREATE TABLE IF NOT EXISTS routines_new (
|
||||
id TEXT PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
description TEXT NOT NULL DEFAULT '',
|
||||
user_id TEXT NOT NULL,
|
||||
enabled INTEGER NOT NULL DEFAULT 1,
|
||||
trigger_type TEXT NOT NULL,
|
||||
trigger_config TEXT NOT NULL,
|
||||
action_type TEXT NOT NULL,
|
||||
action_config TEXT NOT NULL,
|
||||
cooldown_secs INTEGER NOT NULL DEFAULT 300,
|
||||
max_concurrent INTEGER NOT NULL DEFAULT 1,
|
||||
dedup_window_secs INTEGER,
|
||||
notify_channel TEXT,
|
||||
notify_user TEXT,
|
||||
notify_on_success INTEGER NOT NULL DEFAULT 0,
|
||||
notify_on_failure INTEGER NOT NULL DEFAULT 1,
|
||||
notify_on_attention INTEGER NOT NULL DEFAULT 1,
|
||||
state TEXT NOT NULL DEFAULT '{}',
|
||||
last_run_at TEXT,
|
||||
next_fire_at TEXT,
|
||||
run_count INTEGER NOT NULL DEFAULT 0,
|
||||
consecutive_failures INTEGER NOT NULL DEFAULT 0,
|
||||
created_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now')),
|
||||
updated_at TEXT NOT NULL DEFAULT (strftime('%Y-%m-%dT%H:%M:%fZ', 'now')),
|
||||
UNIQUE (user_id, name)
|
||||
);
|
||||
|
||||
INSERT INTO routines_new (
|
||||
id, name, description, user_id, enabled,
|
||||
trigger_type, trigger_config, action_type, action_config,
|
||||
cooldown_secs, max_concurrent, dedup_window_secs,
|
||||
notify_channel, notify_user, notify_on_success, notify_on_failure, notify_on_attention,
|
||||
state, last_run_at, next_fire_at, run_count, consecutive_failures,
|
||||
created_at, updated_at
|
||||
)
|
||||
SELECT
|
||||
id, name, description, user_id, enabled,
|
||||
trigger_type, trigger_config, action_type, action_config,
|
||||
cooldown_secs, max_concurrent, dedup_window_secs,
|
||||
notify_channel,
|
||||
CASE WHEN notify_user = 'default' THEN NULL ELSE notify_user END,
|
||||
notify_on_success, notify_on_failure, notify_on_attention,
|
||||
state, last_run_at, next_fire_at, run_count, consecutive_failures,
|
||||
created_at, updated_at
|
||||
FROM routines;
|
||||
|
||||
DROP TABLE routines;
|
||||
ALTER TABLE routines_new RENAME TO routines;
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_routines_user ON routines(user_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_routines_next_fire ON routines(next_fire_at);
|
||||
CREATE INDEX IF NOT EXISTS idx_routines_event_triggers
|
||||
ON routines(trigger_type, user_id)
|
||||
WHERE enabled = 1 AND trigger_type IN ('event', 'system_event');
|
||||
|
||||
PRAGMA foreign_keys=ON;
|
||||
"#,
|
||||
),
|
||||
];
|
||||
|
||||
+27
-9
@@ -5,13 +5,22 @@
|
||||
//! certificates — the same TLS stack that `reqwest` already uses for HTTP.
|
||||
|
||||
use deadpool_postgres::{Pool, Runtime};
|
||||
use thiserror::Error;
|
||||
use tokio_postgres::NoTls;
|
||||
use tokio_postgres_rustls::MakeRustlsConnect;
|
||||
|
||||
use crate::config::SslMode;
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum CreatePoolError {
|
||||
#[error("{0}")]
|
||||
Pool(#[from] deadpool_postgres::CreatePoolError),
|
||||
#[error("postgres TLS configuration failed: {0}")]
|
||||
TlsConfig(#[from] rustls::Error),
|
||||
}
|
||||
|
||||
/// Build a rustls-based TLS connector using the platform's root certificate store.
|
||||
fn make_rustls_connector() -> MakeRustlsConnect {
|
||||
fn make_rustls_connector() -> Result<MakeRustlsConnect, rustls::Error> {
|
||||
let mut root_store = rustls::RootCertStore::empty();
|
||||
let native = rustls_native_certs::load_native_certs();
|
||||
for e in &native.errors {
|
||||
@@ -25,10 +34,15 @@ fn make_rustls_connector() -> MakeRustlsConnect {
|
||||
if root_store.is_empty() {
|
||||
tracing::error!("no system root certificates found -- TLS connections will fail");
|
||||
}
|
||||
let config = rustls::ClientConfig::builder()
|
||||
.with_root_certificates(root_store)
|
||||
.with_no_client_auth();
|
||||
MakeRustlsConnect::new(config)
|
||||
// `--all-features` brings in both aws-lc-rs and ring-backed rustls providers.
|
||||
// Pick the same ring provider reqwest already uses so postgres TLS setup stays deterministic.
|
||||
let config = rustls::ClientConfig::builder_with_provider(
|
||||
rustls::crypto::ring::default_provider().into(),
|
||||
)
|
||||
.with_safe_default_protocol_versions()?
|
||||
.with_root_certificates(root_store)
|
||||
.with_no_client_auth();
|
||||
Ok(MakeRustlsConnect::new(config))
|
||||
}
|
||||
|
||||
/// Create a [`deadpool_postgres::Pool`] with the appropriate TLS connector.
|
||||
@@ -45,12 +59,16 @@ fn make_rustls_connector() -> MakeRustlsConnect {
|
||||
pub fn create_pool(
|
||||
config: &deadpool_postgres::Config,
|
||||
ssl_mode: SslMode,
|
||||
) -> Result<Pool, deadpool_postgres::CreatePoolError> {
|
||||
) -> Result<Pool, CreatePoolError> {
|
||||
match ssl_mode {
|
||||
SslMode::Disable => config.create_pool(Some(Runtime::Tokio1), NoTls),
|
||||
SslMode::Disable => config
|
||||
.create_pool(Some(Runtime::Tokio1), NoTls)
|
||||
.map_err(CreatePoolError::from),
|
||||
SslMode::Prefer | SslMode::Require => {
|
||||
let tls = make_rustls_connector();
|
||||
config.create_pool(Some(Runtime::Tokio1), tls)
|
||||
let tls = make_rustls_connector()?;
|
||||
config
|
||||
.create_pool(Some(Runtime::Tokio1), tls)
|
||||
.map_err(CreatePoolError::from)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -122,6 +122,9 @@ pub enum ChannelError {
|
||||
#[error("Failed to send response on channel {name}: {reason}")]
|
||||
SendFailed { name: String, reason: String },
|
||||
|
||||
#[error("Channel {name} is missing a routing target: {reason}")]
|
||||
MissingRoutingTarget { name: String, reason: String },
|
||||
|
||||
#[error("Invalid message format: {0}")]
|
||||
InvalidMessage(String),
|
||||
|
||||
|
||||
+1588
-69
File diff suppressed because it is too large
Load Diff
@@ -453,6 +453,17 @@ pub struct ActivateResult {
|
||||
///
|
||||
/// Returned by `ExtensionManager::configure()`, the single entrypoint
|
||||
/// for providing secrets to any extension (chat auth, gateway setup, etc.).
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct VerificationChallenge {
|
||||
/// One-time code the user must send back to the integration.
|
||||
pub code: String,
|
||||
/// Human-readable instructions for completing verification.
|
||||
pub instructions: String,
|
||||
/// Deep-link or shortcut URL that prefills the verification payload when supported.
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub deep_link: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ConfigureResult {
|
||||
/// Human-readable status message.
|
||||
@@ -461,6 +472,8 @@ pub struct ConfigureResult {
|
||||
pub activated: bool,
|
||||
/// OAuth authorization URL (if OAuth flow was started).
|
||||
pub auth_url: Option<String>,
|
||||
/// Pending manual verification challenge (for Telegram owner binding, etc.).
|
||||
pub verification: Option<VerificationChallenge>,
|
||||
}
|
||||
|
||||
fn default_true() -> bool {
|
||||
|
||||
@@ -227,6 +227,7 @@ impl Store {
|
||||
job_id: row.get("id"),
|
||||
state,
|
||||
user_id: row.get::<_, String>("user_id"),
|
||||
requester_id: None,
|
||||
conversation_id: row.get("conversation_id"),
|
||||
title: row.get("title"),
|
||||
description: row.get("description"),
|
||||
|
||||
@@ -7,8 +7,12 @@ Multi-provider LLM integration with circuit breaker, retry, failover, and respon
|
||||
| File | Role |
|
||||
|------|------|
|
||||
| `mod.rs` | Provider factory (`create_llm_provider`, `build_provider_chain`); `LlmBackend` enum |
|
||||
| `config.rs` | LLM config types (`LlmConfig`, `RegistryProviderConfig`, `NearAiConfig`, `BedrockConfig`) |
|
||||
| `error.rs` | `LlmError` enum used by all providers |
|
||||
| `provider.rs` | `LlmProvider` trait, `ChatMessage`, `ToolCall`, `CompletionRequest`, `sanitize_tool_messages` |
|
||||
| `nearai_chat.rs` | NEAR AI Chat Completions provider (dual auth: session token or API key) |
|
||||
| `codex_auth.rs` | Reads Codex CLI `auth.json`, extracts tokens, refreshes ChatGPT OAuth access tokens |
|
||||
| `codex_chatgpt.rs` | Custom Responses API provider for Codex ChatGPT backend (`/backend-api/codex`) |
|
||||
| `reasoning.rs` | `Reasoning` struct, `ReasoningContext`, `RespondResult`, `ActionPlan`, `ToolSelection`; thinking-tag stripping; `SILENT_REPLY_TOKEN` |
|
||||
| `session.rs` | NEAR AI session token management with disk + DB persistence, OAuth login flow |
|
||||
| `circuit_breaker.rs` | Circuit breaker: Closed → Open → HalfOpen state machine |
|
||||
@@ -35,6 +39,12 @@ Set via `LLM_BACKEND` env var:
|
||||
| `tinfoil` | Tinfoil TEE inference | `TINFOIL_API_KEY`, `TINFOIL_MODEL` |
|
||||
| `bedrock` | AWS Bedrock (requires `--features bedrock`) | `BEDROCK_REGION`, `BEDROCK_MODEL`, `AWS_PROFILE` |
|
||||
|
||||
Codex auth reuse:
|
||||
- Set `LLM_USE_CODEX_AUTH=true` to load credentials from `~/.codex/auth.json` (override with `CODEX_AUTH_PATH`).
|
||||
- If Codex is logged in with API-key mode, IronClaw uses the standard OpenAI endpoint.
|
||||
- If Codex is logged in with ChatGPT OAuth mode, IronClaw routes to the private `chatgpt.com/backend-api/codex` Responses API via `codex_chatgpt.rs`.
|
||||
- ChatGPT mode supports one automatic 401 refresh using the refresh token persisted in `auth.json`.
|
||||
|
||||
## AWS Bedrock Provider
|
||||
|
||||
Uses the native Converse API via `aws-sdk-bedrockruntime` (`bedrock.rs`). Requires `--features bedrock` at build time — not in default features due to heavy AWS SDK dependencies.
|
||||
|
||||
+126
-4
@@ -34,7 +34,9 @@ const DEFAULT_MAX_TOKENS: u32 = 8192;
|
||||
/// Anthropic provider using OAuth Bearer authentication.
|
||||
pub struct AnthropicOAuthProvider {
|
||||
client: Client,
|
||||
token: SecretString,
|
||||
/// OAuth token, wrapped in RwLock so it can be updated after a successful
|
||||
/// Keychain refresh (fixes #1136: stale token reuse after expiry).
|
||||
token: std::sync::RwLock<SecretString>,
|
||||
model: String,
|
||||
base_url: Option<String>,
|
||||
active_model: std::sync::RwLock<String>,
|
||||
@@ -71,7 +73,7 @@ impl AnthropicOAuthProvider {
|
||||
|
||||
Ok(Self {
|
||||
client,
|
||||
token,
|
||||
token: std::sync::RwLock::new(token),
|
||||
model: config.model.clone(),
|
||||
base_url,
|
||||
active_model,
|
||||
@@ -98,6 +100,22 @@ impl AnthropicOAuthProvider {
|
||||
}
|
||||
}
|
||||
|
||||
/// Read the current token from the RwLock.
|
||||
fn current_token(&self) -> String {
|
||||
match self.token.read() {
|
||||
Ok(guard) => guard.expose_secret().to_string(),
|
||||
Err(poisoned) => poisoned.into_inner().expose_secret().to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Update the stored token after a successful Keychain refresh.
|
||||
fn update_token(&self, new_token: SecretString) {
|
||||
match self.token.write() {
|
||||
Ok(mut guard) => *guard = new_token,
|
||||
Err(poisoned) => *poisoned.into_inner() = new_token,
|
||||
}
|
||||
}
|
||||
|
||||
async fn send_request<R: for<'de> Deserialize<'de>>(
|
||||
&self,
|
||||
body: &AnthropicRequest,
|
||||
@@ -109,7 +127,7 @@ impl AnthropicOAuthProvider {
|
||||
let response = self
|
||||
.client
|
||||
.post(&url)
|
||||
.bearer_auth(self.token.expose_secret())
|
||||
.bearer_auth(self.current_token())
|
||||
.header("anthropic-version", ANTHROPIC_API_VERSION)
|
||||
.header("anthropic-beta", ANTHROPIC_OAUTH_BETA)
|
||||
.header("Content-Type", "application/json")
|
||||
@@ -125,12 +143,14 @@ impl AnthropicOAuthProvider {
|
||||
|
||||
if !status.is_success() {
|
||||
// Parse Retry-After header before consuming the body.
|
||||
// Falls back to 60s if header is missing or unparseable (prevents "retry after None" errors).
|
||||
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);
|
||||
.map(std::time::Duration::from_secs)
|
||||
.or(Some(std::time::Duration::from_secs(60)));
|
||||
|
||||
let response_text = response
|
||||
.text()
|
||||
@@ -141,6 +161,11 @@ impl AnthropicOAuthProvider {
|
||||
// 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.
|
||||
//
|
||||
// Brief delay to give Claude Code time to complete its async
|
||||
// Keychain refresh write (fixes race in #1136).
|
||||
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
|
||||
|
||||
if let Some(fresh) = crate::config::ClaudeCodeConfig::extract_oauth_token() {
|
||||
let fresh_token = SecretString::from(fresh);
|
||||
// Retry once with the refreshed token
|
||||
@@ -159,6 +184,11 @@ impl AnthropicOAuthProvider {
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
if retry.status().is_success() {
|
||||
// Persist the refreshed token so subsequent requests
|
||||
// don't hit 401 again (fixes #1136).
|
||||
self.update_token(fresh_token);
|
||||
tracing::info!("Anthropic OAuth token refreshed from credential store");
|
||||
|
||||
let text = retry.text().await.map_err(|e| LlmError::RequestFailed {
|
||||
provider: "anthropic_oauth".to_string(),
|
||||
reason: format!("Failed to read response body: {}", e),
|
||||
@@ -659,4 +689,96 @@ mod tests {
|
||||
assert_eq!(tool_calls.len(), 1);
|
||||
assert_eq!(tool_calls[0].name, "search");
|
||||
}
|
||||
|
||||
/// Regression test for #1136: token field must be mutable via RwLock
|
||||
/// so that a refreshed token persists across subsequent requests.
|
||||
#[test]
|
||||
fn test_token_update_persists() {
|
||||
let original = SecretString::from("old_token".to_string());
|
||||
let token = std::sync::RwLock::new(original);
|
||||
|
||||
// Read the original
|
||||
assert_eq!(token.read().unwrap().expose_secret(), "old_token");
|
||||
|
||||
// Simulate a successful refresh
|
||||
let refreshed = SecretString::from("new_token".to_string());
|
||||
*token.write().unwrap() = refreshed;
|
||||
|
||||
// Subsequent reads see the updated token
|
||||
assert_eq!(token.read().unwrap().expose_secret(), "new_token");
|
||||
}
|
||||
|
||||
// -- Retry-After header parsing tests (regression for rate limit "None" bug) --
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_parsing_delay_seconds() {
|
||||
// Verify delay-seconds format is parsed correctly
|
||||
let header_value = "45";
|
||||
let duration = parse_retry_after_anthropic_for_test(header_value);
|
||||
assert_eq!(
|
||||
duration,
|
||||
Some(std::time::Duration::from_secs(45)),
|
||||
"Should parse delay-seconds format"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_fallback_missing_header() {
|
||||
// Regression test: When Retry-After header is missing,
|
||||
// should fall back to 60s instead of None
|
||||
let duration = parse_retry_after_anthropic_for_test("");
|
||||
assert_eq!(
|
||||
duration,
|
||||
Some(std::time::Duration::from_secs(60)),
|
||||
"Missing header should fallback to 60s"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_fallback_invalid_format() {
|
||||
// Regression test: When Retry-After header is in unexpected format,
|
||||
// should fall back to 60s instead of None
|
||||
let invalid_formats = vec![
|
||||
"invalid",
|
||||
"not-a-number",
|
||||
"30.5", // float instead of int
|
||||
"abc123",
|
||||
"Mon, 02 Mar 2026 18:00:00 GMT", // RFC2822 not supported in anthropic version
|
||||
];
|
||||
|
||||
for format in invalid_formats {
|
||||
let duration = parse_retry_after_anthropic_for_test(format);
|
||||
assert_eq!(
|
||||
duration,
|
||||
Some(std::time::Duration::from_secs(60)),
|
||||
"Invalid format '{}' should fallback to 60s",
|
||||
format
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_zero_seconds_accepted() {
|
||||
// Verify zero seconds is a valid retry delay
|
||||
let duration = parse_retry_after_anthropic_for_test("0");
|
||||
assert_eq!(duration, Some(std::time::Duration::ZERO));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_large_number() {
|
||||
// Verify large numbers are accepted
|
||||
let duration = parse_retry_after_anthropic_for_test("7200"); // 2 hours
|
||||
assert_eq!(duration, Some(std::time::Duration::from_secs(7200)));
|
||||
}
|
||||
|
||||
/// Helper function to test Retry-After header parsing logic for Anthropic
|
||||
/// (simulates the parsing done in send_request without actual HTTP, including fallback)
|
||||
fn parse_retry_after_anthropic_for_test(header_value: &str) -> Option<std::time::Duration> {
|
||||
header_value
|
||||
.trim()
|
||||
.parse::<u64>()
|
||||
.ok()
|
||||
.map(std::time::Duration::from_secs)
|
||||
.or(Some(std::time::Duration::from_secs(60)))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,377 @@
|
||||
//! Read Codex CLI credentials for LLM authentication.
|
||||
//!
|
||||
//! When `LLM_USE_CODEX_AUTH=true`, IronClaw reads the Codex CLI's
|
||||
//! `auth.json` file (default: `~/.codex/auth.json`) and extracts
|
||||
//! credentials. This lets IronClaw piggyback on a Codex login without
|
||||
//! implementing its own OAuth flow.
|
||||
//!
|
||||
//! Codex supports two auth modes:
|
||||
//! - **API key** (`auth_mode: "apiKey"`) → uses `OPENAI_API_KEY` field
|
||||
//! against `api.openai.com/v1`.
|
||||
//! - **ChatGPT** (`auth_mode: "chatgpt"`) → uses `tokens.access_token`
|
||||
//! (OAuth JWT) against `chatgpt.com/backend-api/codex`.
|
||||
//!
|
||||
//! When in ChatGPT mode, the provider supports automatic token refresh
|
||||
//! on 401 responses using the `refresh_token` from `auth.json`.
|
||||
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use secrecy::{ExposeSecret, SecretString};
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// ChatGPT backend API endpoint used by Codex in ChatGPT auth mode.
|
||||
const CHATGPT_BACKEND_URL: &str = "https://chatgpt.com/backend-api/codex";
|
||||
|
||||
/// Standard OpenAI API endpoint used by Codex in API key mode.
|
||||
const OPENAI_API_URL: &str = "https://api.openai.com/v1";
|
||||
|
||||
/// OAuth token refresh endpoint (same as Codex CLI).
|
||||
const REFRESH_TOKEN_URL: &str = "https://auth.openai.com/oauth/token";
|
||||
|
||||
/// OAuth client ID used for token refresh (same as Codex CLI).
|
||||
const CLIENT_ID: &str = "app_EMoamEEZ73f0CkXaXp7hrann";
|
||||
|
||||
/// Credentials extracted from Codex's `auth.json`.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct CodexCredentials {
|
||||
/// The bearer token (API key or ChatGPT access_token).
|
||||
pub token: SecretString,
|
||||
/// Whether this is a ChatGPT OAuth token (vs. an OpenAI API key).
|
||||
pub is_chatgpt_mode: bool,
|
||||
/// OAuth refresh token (only present in ChatGPT mode).
|
||||
pub refresh_token: Option<SecretString>,
|
||||
/// Path to the auth.json file (for persisting refreshed tokens).
|
||||
pub auth_path: Option<PathBuf>,
|
||||
}
|
||||
|
||||
impl CodexCredentials {
|
||||
/// Returns the correct base URL for the auth mode.
|
||||
///
|
||||
/// - ChatGPT mode → `https://chatgpt.com/backend-api/codex`
|
||||
/// - API key mode → `https://api.openai.com/v1`
|
||||
pub fn base_url(&self) -> &'static str {
|
||||
if self.is_chatgpt_mode {
|
||||
CHATGPT_BACKEND_URL
|
||||
} else {
|
||||
OPENAI_API_URL
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Partial representation of Codex's `$CODEX_HOME/auth.json`.
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct CodexAuthJson {
|
||||
auth_mode: Option<String>,
|
||||
#[serde(rename = "OPENAI_API_KEY")]
|
||||
openai_api_key: Option<String>,
|
||||
tokens: Option<CodexTokens>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct CodexTokens {
|
||||
access_token: SecretString,
|
||||
refresh_token: Option<SecretString>,
|
||||
}
|
||||
|
||||
/// Request body for OAuth token refresh.
|
||||
#[derive(Serialize)]
|
||||
struct RefreshRequest<'a> {
|
||||
client_id: &'a str,
|
||||
grant_type: &'a str,
|
||||
refresh_token: &'a str,
|
||||
}
|
||||
|
||||
/// Response from the OAuth token refresh endpoint.
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct RefreshResponse {
|
||||
access_token: SecretString,
|
||||
refresh_token: Option<SecretString>,
|
||||
}
|
||||
|
||||
/// Default path used by Codex CLI: `~/.codex/auth.json`.
|
||||
pub fn default_codex_auth_path() -> PathBuf {
|
||||
let home_dir = dirs::home_dir().unwrap_or_else(|| {
|
||||
tracing::warn!(
|
||||
"Could not determine home directory; falling back to current working directory for Codex auth.json path"
|
||||
);
|
||||
PathBuf::from(".")
|
||||
});
|
||||
|
||||
home_dir.join(".codex").join("auth.json")
|
||||
}
|
||||
|
||||
/// Load credentials from a Codex `auth.json` file.
|
||||
///
|
||||
/// Returns `None` if the file is missing, unreadable, or contains
|
||||
/// no usable credentials.
|
||||
pub fn load_codex_credentials(path: &Path) -> Option<CodexCredentials> {
|
||||
let content = match std::fs::read_to_string(path) {
|
||||
Ok(c) => c,
|
||||
Err(e) => {
|
||||
tracing::debug!("Could not read Codex auth file {}: {}", path.display(), e);
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
let auth: CodexAuthJson = match serde_json::from_str(&content) {
|
||||
Ok(a) => a,
|
||||
Err(e) => {
|
||||
tracing::warn!("Failed to parse Codex auth file {}: {}", path.display(), e);
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
let is_chatgpt = auth
|
||||
.auth_mode
|
||||
.as_deref()
|
||||
.map(|m| m == "chatgpt" || m == "chatgptAuthTokens")
|
||||
.unwrap_or(false);
|
||||
|
||||
// API key mode: use OPENAI_API_KEY field.
|
||||
if !is_chatgpt {
|
||||
if let Some(key) = auth.openai_api_key.filter(|k| !k.is_empty()) {
|
||||
tracing::info!("Loaded API key from Codex auth.json (API key mode)");
|
||||
return Some(CodexCredentials {
|
||||
token: SecretString::from(key),
|
||||
is_chatgpt_mode: false,
|
||||
refresh_token: None,
|
||||
auth_path: None,
|
||||
});
|
||||
}
|
||||
// If auth_mode was explicitly `apiKey`, do not fall back to checking for a token.
|
||||
if auth.auth_mode.is_some() {
|
||||
return None;
|
||||
}
|
||||
}
|
||||
|
||||
// ChatGPT mode: use access_token as bearer token.
|
||||
if let Some(tokens) = auth.tokens
|
||||
&& !tokens.access_token.expose_secret().is_empty()
|
||||
{
|
||||
tracing::info!(
|
||||
"Loaded access token from Codex auth.json (ChatGPT mode, base_url={})",
|
||||
CHATGPT_BACKEND_URL
|
||||
);
|
||||
return Some(CodexCredentials {
|
||||
token: tokens.access_token,
|
||||
is_chatgpt_mode: true,
|
||||
refresh_token: tokens.refresh_token,
|
||||
auth_path: Some(path.to_path_buf()),
|
||||
});
|
||||
}
|
||||
|
||||
tracing::debug!(
|
||||
"Codex auth.json at {} contains no usable credentials",
|
||||
path.display()
|
||||
);
|
||||
None
|
||||
}
|
||||
|
||||
/// Attempt to refresh an expired access token using the refresh token.
|
||||
///
|
||||
/// On success, returns the new `access_token` and persists the refreshed
|
||||
/// tokens back to `auth.json`. This follows the same OAuth protocol as
|
||||
/// Codex CLI (`POST https://auth.openai.com/oauth/token`).
|
||||
///
|
||||
/// Returns `None` if the refresh token is missing, the request fails,
|
||||
/// or the response is malformed.
|
||||
pub async fn refresh_access_token(
|
||||
client: &reqwest::Client,
|
||||
refresh_token: &SecretString,
|
||||
auth_path: Option<&Path>,
|
||||
) -> Option<SecretString> {
|
||||
let req = RefreshRequest {
|
||||
client_id: CLIENT_ID,
|
||||
grant_type: "refresh_token",
|
||||
refresh_token: refresh_token.expose_secret(),
|
||||
};
|
||||
|
||||
tracing::info!("Attempting to refresh Codex OAuth access token");
|
||||
|
||||
let resp = match client
|
||||
.post(REFRESH_TOKEN_URL)
|
||||
.header("Content-Type", "application/json")
|
||||
.json(&req)
|
||||
.timeout(std::time::Duration::from_secs(10))
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
tracing::warn!("Token refresh request failed: {e}");
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
if !resp.status().is_success() {
|
||||
let status = resp.status();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
tracing::warn!("Token refresh failed: HTTP {status}: {body}");
|
||||
if status.as_u16() == 401 {
|
||||
tracing::warn!(
|
||||
"Refresh token may be expired or revoked. \
|
||||
Please re-authenticate with: codex --login"
|
||||
);
|
||||
}
|
||||
return None;
|
||||
}
|
||||
|
||||
let refresh_resp: RefreshResponse = match resp.json().await {
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
tracing::warn!("Failed to parse token refresh response: {e}");
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
let new_access_token = refresh_resp.access_token.clone();
|
||||
|
||||
// Persist refreshed tokens back to auth.json
|
||||
if let Some(path) = auth_path {
|
||||
if let Err(e) = persist_refreshed_tokens(
|
||||
path,
|
||||
refresh_resp.access_token.expose_secret(),
|
||||
refresh_resp
|
||||
.refresh_token
|
||||
.as_ref()
|
||||
.map(ExposeSecret::expose_secret),
|
||||
) {
|
||||
tracing::warn!(
|
||||
"Failed to persist refreshed tokens to {}: {e}",
|
||||
path.display()
|
||||
);
|
||||
} else {
|
||||
tracing::info!("Refreshed tokens persisted to {}", path.display());
|
||||
}
|
||||
}
|
||||
|
||||
Some(new_access_token)
|
||||
}
|
||||
|
||||
/// Update `auth.json` with refreshed tokens, preserving other fields.
|
||||
fn persist_refreshed_tokens(
|
||||
path: &Path,
|
||||
new_access_token: &str,
|
||||
new_refresh_token: Option<&str>,
|
||||
) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let content = std::fs::read_to_string(path)?;
|
||||
let mut json: serde_json::Value = serde_json::from_str(&content)?;
|
||||
|
||||
if let Some(tokens) = json.get_mut("tokens") {
|
||||
tokens["access_token"] = serde_json::Value::String(new_access_token.to_string());
|
||||
if let Some(rt) = new_refresh_token {
|
||||
tokens["refresh_token"] = serde_json::Value::String(rt.to_string());
|
||||
}
|
||||
}
|
||||
|
||||
let updated = serde_json::to_string_pretty(&json)?;
|
||||
let tmp_path = path.with_extension("json.tmp");
|
||||
std::fs::write(&tmp_path, updated)?;
|
||||
if let Err(e) = std::fs::rename(&tmp_path, path) {
|
||||
let _ = std::fs::remove_file(&tmp_path);
|
||||
return Err(Box::new(e));
|
||||
}
|
||||
set_auth_file_permissions(path)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(unix)]
|
||||
fn set_auth_file_permissions(path: &Path) -> Result<(), Box<dyn std::error::Error>> {
|
||||
use std::os::unix::fs::PermissionsExt;
|
||||
|
||||
std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(not(unix))]
|
||||
fn set_auth_file_permissions(_path: &Path) -> Result<(), Box<dyn std::error::Error>> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::io::Write;
|
||||
use tempfile::NamedTempFile;
|
||||
|
||||
#[test]
|
||||
fn loads_api_key_mode() {
|
||||
let mut f = NamedTempFile::new().unwrap();
|
||||
writeln!(
|
||||
f,
|
||||
r#"{{"auth_mode":"apiKey","OPENAI_API_KEY":"sk-test-123"}}"#
|
||||
)
|
||||
.unwrap();
|
||||
let creds = load_codex_credentials(f.path()).expect("should load");
|
||||
assert_eq!(creds.token.expose_secret(), "sk-test-123");
|
||||
assert!(!creds.is_chatgpt_mode);
|
||||
assert_eq!(creds.base_url(), OPENAI_API_URL);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn loads_chatgpt_mode() {
|
||||
let mut f = NamedTempFile::new().unwrap();
|
||||
writeln!(
|
||||
f,
|
||||
r#"{{"auth_mode":"chatgpt","tokens":{{"id_token":{{}},"access_token":"eyJ-test","refresh_token":"rt-x"}}}}"#
|
||||
)
|
||||
.unwrap();
|
||||
let creds = load_codex_credentials(f.path()).expect("should load");
|
||||
assert_eq!(creds.token.expose_secret(), "eyJ-test");
|
||||
assert!(creds.is_chatgpt_mode);
|
||||
assert_eq!(
|
||||
creds
|
||||
.refresh_token
|
||||
.as_ref()
|
||||
.expect("refresh token should be present")
|
||||
.expose_secret(),
|
||||
"rt-x"
|
||||
);
|
||||
assert_eq!(creds.base_url(), CHATGPT_BACKEND_URL);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn api_key_mode_ignores_tokens() {
|
||||
let mut f = NamedTempFile::new().unwrap();
|
||||
writeln!(
|
||||
f,
|
||||
r#"{{"auth_mode":"apiKey","OPENAI_API_KEY":"sk-priority","tokens":{{"id_token":{{}},"access_token":"eyJ-fallback","refresh_token":"rt-x"}}}}"#
|
||||
)
|
||||
.unwrap();
|
||||
let creds = load_codex_credentials(f.path()).expect("should load");
|
||||
assert_eq!(creds.token.expose_secret(), "sk-priority");
|
||||
assert!(!creds.is_chatgpt_mode);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn returns_none_for_missing_file() {
|
||||
assert!(load_codex_credentials(Path::new("/tmp/nonexistent_codex_auth.json")).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn returns_none_for_empty_json() {
|
||||
let mut f = NamedTempFile::new().unwrap();
|
||||
writeln!(f, "{{}}").unwrap();
|
||||
assert!(load_codex_credentials(f.path()).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn returns_none_for_empty_key() {
|
||||
let mut f = NamedTempFile::new().unwrap();
|
||||
writeln!(f, r#"{{"auth_mode":"apiKey","OPENAI_API_KEY":""}}"#).unwrap();
|
||||
assert!(load_codex_credentials(f.path()).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn api_key_mode_missing_key_does_not_fallback_to_chatgpt() {
|
||||
// Bug: if auth_mode is "apiKey" but key is missing, the old code would
|
||||
// fall through to check for a ChatGPT token, returning is_chatgpt_mode: true.
|
||||
let mut f = NamedTempFile::new().unwrap();
|
||||
writeln!(
|
||||
f,
|
||||
r#"{{"auth_mode":"apiKey","OPENAI_API_KEY":"","tokens":{{"id_token":{{}},"access_token":"eyJ-bad","refresh_token":"rt-x"}}}}"#
|
||||
)
|
||||
.unwrap();
|
||||
assert!(load_codex_credentials(f.path()).is_none());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,932 @@
|
||||
//! Codex ChatGPT Responses API provider.
|
||||
//!
|
||||
//! Implements `LlmProvider` by speaking the OpenAI Responses API protocol
|
||||
//! (`POST /responses`) used by the ChatGPT backend at
|
||||
//! `chatgpt.com/backend-api/codex`. This bypasses `rig-core`'s Chat
|
||||
//! Completions path, which is incompatible with this endpoint.
|
||||
//!
|
||||
//! # Warning
|
||||
//!
|
||||
//! The ChatGPT backend endpoint (`chatgpt.com/backend-api/codex`) is a
|
||||
//! **private, undocumented API**. Using subscriber OAuth tokens from a
|
||||
//! third-party application may violate the token's intended scope or
|
||||
//! OpenAI's Terms of Service. This feature is provided as-is for
|
||||
//! convenience and may break without notice.
|
||||
|
||||
use async_trait::async_trait;
|
||||
use eventsource_stream::Eventsource;
|
||||
use futures::{Stream, StreamExt};
|
||||
use reqwest::Client;
|
||||
use rust_decimal::Decimal;
|
||||
use secrecy::{ExposeSecret, SecretString};
|
||||
use serde_json::{Value, json};
|
||||
use std::path::PathBuf;
|
||||
use std::time::Duration;
|
||||
use tokio::sync::{Mutex, RwLock};
|
||||
|
||||
use super::codex_auth;
|
||||
use crate::error::LlmError;
|
||||
|
||||
use super::provider::{
|
||||
ChatMessage, CompletionRequest, CompletionResponse, ContentPart, FinishReason, LlmProvider,
|
||||
Role, ToolCall, ToolCompletionRequest, ToolCompletionResponse, ToolDefinition,
|
||||
};
|
||||
|
||||
/// Provider that speaks the Responses API protocol against the ChatGPT backend.
|
||||
pub struct CodexChatGptProvider {
|
||||
client: Client,
|
||||
base_url: String,
|
||||
api_key: RwLock<SecretString>,
|
||||
/// User-configured model name (or empty/"default" for auto-detect).
|
||||
configured_model: String,
|
||||
/// Lazily resolved model name (populated on first LLM call).
|
||||
resolved_model: tokio::sync::OnceCell<String>,
|
||||
/// OAuth refresh token for automatic 401 retry.
|
||||
refresh_token: Option<SecretString>,
|
||||
/// Path to auth.json for persisting refreshed tokens.
|
||||
auth_path: Option<PathBuf>,
|
||||
/// Timeout for actual `/responses` requests.
|
||||
request_timeout: Duration,
|
||||
/// Prevent concurrent 401 handlers from racing the same refresh token.
|
||||
refresh_lock: Mutex<()>,
|
||||
}
|
||||
|
||||
impl CodexChatGptProvider {
|
||||
#[cfg(test)]
|
||||
fn new(base_url: &str, api_key: &str, model: &str) -> Self {
|
||||
Self {
|
||||
client: Client::new(),
|
||||
base_url: base_url.trim_end_matches('/').to_string(),
|
||||
api_key: RwLock::new(SecretString::from(api_key.to_string())),
|
||||
configured_model: model.to_string(),
|
||||
resolved_model: tokio::sync::OnceCell::const_new(),
|
||||
refresh_token: None,
|
||||
auth_path: None,
|
||||
request_timeout: Duration::from_secs(120),
|
||||
refresh_lock: Mutex::new(()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Create a provider with lazy model detection.
|
||||
///
|
||||
/// The model is **not** resolved during construction. Instead, it is
|
||||
/// resolved on the first LLM call via [`resolve_model`], avoiding the
|
||||
/// need for `block_in_place` / `block_on` during provider setup.
|
||||
///
|
||||
/// **Model selection priority** (applied at resolution time):
|
||||
/// 1. If `configured_model` is non-empty, validate it against the
|
||||
/// `/models` endpoint. If it isn't in the supported list, log a
|
||||
/// warning with available models and fall back to the top model.
|
||||
/// 2. If `configured_model` is empty (or a generic placeholder like
|
||||
/// "default"), auto-detect the highest-priority model from the API.
|
||||
pub fn with_lazy_model(
|
||||
base_url: &str,
|
||||
api_key: SecretString,
|
||||
configured_model: &str,
|
||||
refresh_token: Option<SecretString>,
|
||||
auth_path: Option<PathBuf>,
|
||||
request_timeout_secs: u64,
|
||||
) -> Self {
|
||||
tracing::warn!(
|
||||
"Codex ChatGPT provider uses a private, undocumented API \
|
||||
(chatgpt.com/backend-api/codex). This may violate OpenAI's \
|
||||
Terms of Service and could break without notice."
|
||||
);
|
||||
|
||||
Self {
|
||||
client: Client::new(),
|
||||
base_url: base_url.trim_end_matches('/').to_string(),
|
||||
api_key: RwLock::new(api_key),
|
||||
configured_model: configured_model.to_string(),
|
||||
resolved_model: tokio::sync::OnceCell::const_new(),
|
||||
refresh_token,
|
||||
auth_path,
|
||||
request_timeout: Duration::from_secs(request_timeout_secs),
|
||||
refresh_lock: Mutex::new(()),
|
||||
}
|
||||
}
|
||||
|
||||
/// Resolve the model to use, lazily on first call.
|
||||
///
|
||||
/// Uses `OnceCell` so the `/models` fetch happens at most once.
|
||||
async fn resolve_model(&self) -> &str {
|
||||
self.resolved_model
|
||||
.get_or_init(|| async {
|
||||
let api_key = self.api_key.read().await.clone();
|
||||
let available = Self::fetch_available_models(&self.client, &self.base_url, &api_key)
|
||||
.await;
|
||||
|
||||
let configured = &self.configured_model;
|
||||
if !configured.is_empty() && configured != "default" {
|
||||
// User explicitly configured a model — validate it
|
||||
if available.is_empty() {
|
||||
tracing::warn!(
|
||||
"Could not fetch model list; using configured model '{configured}'"
|
||||
);
|
||||
return configured.clone();
|
||||
}
|
||||
if available.iter().any(|m| m == configured) {
|
||||
tracing::info!(model = %configured, "Codex ChatGPT: using configured model");
|
||||
return configured.clone();
|
||||
}
|
||||
tracing::warn!(
|
||||
configured = %configured,
|
||||
available = ?available,
|
||||
"Configured model not found in supported list, falling back to top model"
|
||||
);
|
||||
available
|
||||
.into_iter()
|
||||
.next()
|
||||
.unwrap_or_else(|| configured.clone())
|
||||
} else {
|
||||
// No user preference — auto-detect
|
||||
if let Some(top) = available.into_iter().next() {
|
||||
tracing::info!(model = %top, "Codex ChatGPT: auto-detected model");
|
||||
top
|
||||
} else {
|
||||
tracing::warn!(
|
||||
"Could not auto-detect model, using fallback '{configured}'"
|
||||
);
|
||||
configured.clone()
|
||||
}
|
||||
}
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
/// Query `/models?client_version=0.111.0` and return the list of available
|
||||
/// model slugs, ordered by priority (highest first).
|
||||
async fn fetch_available_models(
|
||||
client: &Client,
|
||||
base_url: &str,
|
||||
api_key: &SecretString,
|
||||
) -> Vec<String> {
|
||||
let url = format!("{base_url}/models?client_version=0.111.0");
|
||||
let resp = match client
|
||||
.get(&url)
|
||||
.bearer_auth(api_key.expose_secret())
|
||||
.timeout(Duration::from_secs(10))
|
||||
.send()
|
||||
.await
|
||||
{
|
||||
Ok(r) => r,
|
||||
Err(e) => {
|
||||
tracing::warn!("Failed to fetch Codex models: {e}");
|
||||
return Vec::new();
|
||||
}
|
||||
};
|
||||
if !resp.status().is_success() {
|
||||
tracing::warn!(status = %resp.status(), "Failed to fetch Codex models");
|
||||
return Vec::new();
|
||||
}
|
||||
let body: Value = match resp.json().await {
|
||||
Ok(v) => v,
|
||||
Err(_) => return Vec::new(),
|
||||
};
|
||||
// The response has { "models": [ { "slug": "...", ... }, ... ] }
|
||||
body.get("models")
|
||||
.and_then(|m| m.as_array())
|
||||
.map(|models| {
|
||||
models
|
||||
.iter()
|
||||
.filter_map(|m| {
|
||||
m.get("slug")
|
||||
.and_then(|s| s.as_str())
|
||||
.map(|s| s.to_string())
|
||||
})
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
/// Convert IronClaw messages to Responses API request JSON.
|
||||
fn build_request_body(
|
||||
&self,
|
||||
model: &str,
|
||||
messages: &[ChatMessage],
|
||||
tools: &[ToolDefinition],
|
||||
tool_choice: Option<&str>,
|
||||
) -> Value {
|
||||
// Extract system instructions
|
||||
let instructions: String = messages
|
||||
.iter()
|
||||
.filter(|m| m.role == Role::System)
|
||||
.map(|m| m.content.as_str())
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n\n");
|
||||
|
||||
// Convert non-system messages to Responses API input items
|
||||
let input: Vec<Value> = messages
|
||||
.iter()
|
||||
.filter(|m| m.role != Role::System)
|
||||
.flat_map(Self::message_to_input_items)
|
||||
.collect();
|
||||
|
||||
// Convert tool definitions
|
||||
let api_tools: Vec<Value> = tools
|
||||
.iter()
|
||||
.map(|t| {
|
||||
json!({
|
||||
"type": "function",
|
||||
"name": t.name,
|
||||
"description": t.description,
|
||||
"parameters": t.parameters,
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
|
||||
let mut body = json!({
|
||||
"model": model,
|
||||
"instructions": instructions,
|
||||
"input": input,
|
||||
"stream": true,
|
||||
"store": false,
|
||||
});
|
||||
|
||||
if !api_tools.is_empty() {
|
||||
body["tools"] = json!(api_tools);
|
||||
body["tool_choice"] = json!(tool_choice.unwrap_or("auto"));
|
||||
}
|
||||
|
||||
body
|
||||
}
|
||||
|
||||
/// Convert a single ChatMessage to one or more Responses API input items.
|
||||
fn message_to_input_items(msg: &ChatMessage) -> Vec<Value> {
|
||||
let mut items = Vec::new();
|
||||
|
||||
match msg.role {
|
||||
Role::User => {
|
||||
// Build content array: if content_parts is populated, use it
|
||||
// to include multimodal content (images). Otherwise fall back
|
||||
// to the plain text content field.
|
||||
let content = if !msg.content_parts.is_empty() {
|
||||
msg.content_parts
|
||||
.iter()
|
||||
.map(|part| match part {
|
||||
ContentPart::Text { text } => json!({
|
||||
"type": "input_text",
|
||||
"text": text,
|
||||
}),
|
||||
ContentPart::ImageUrl { image_url } => json!({
|
||||
"type": "input_image",
|
||||
"image_url": image_url.url,
|
||||
}),
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
} else {
|
||||
vec![json!({
|
||||
"type": "input_text",
|
||||
"text": msg.content,
|
||||
})]
|
||||
};
|
||||
|
||||
items.push(json!({
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": content,
|
||||
}));
|
||||
}
|
||||
Role::Assistant => {
|
||||
// If the assistant message has tool calls, emit function_call items
|
||||
if let Some(ref tool_calls) = msg.tool_calls {
|
||||
// Emit the assistant text as a message if non-empty
|
||||
if !msg.content.is_empty() {
|
||||
items.push(json!({
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{
|
||||
"type": "output_text",
|
||||
"text": msg.content,
|
||||
}],
|
||||
}));
|
||||
}
|
||||
for tc in tool_calls {
|
||||
let args = if tc.arguments.is_string() {
|
||||
tc.arguments.as_str().unwrap_or("{}").to_string()
|
||||
} else {
|
||||
serde_json::to_string(&tc.arguments).unwrap_or_default()
|
||||
};
|
||||
items.push(json!({
|
||||
"type": "function_call",
|
||||
"name": tc.name,
|
||||
"arguments": args,
|
||||
"call_id": tc.id,
|
||||
}));
|
||||
}
|
||||
} else {
|
||||
items.push(json!({
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [{
|
||||
"type": "output_text",
|
||||
"text": msg.content,
|
||||
}],
|
||||
}));
|
||||
}
|
||||
}
|
||||
Role::Tool => {
|
||||
items.push(json!({
|
||||
"type": "function_call_output",
|
||||
"call_id": msg.tool_call_id.as_deref().unwrap_or(""),
|
||||
"output": msg.content,
|
||||
}));
|
||||
}
|
||||
Role::System => {
|
||||
// System messages are handled via `instructions` field
|
||||
}
|
||||
}
|
||||
|
||||
items
|
||||
}
|
||||
|
||||
/// Send a request and parse the SSE response.
|
||||
///
|
||||
/// On HTTP 401, if a refresh token is available, attempts to refresh
|
||||
/// the access token and retry the request once.
|
||||
async fn send_request(&self, body: Value) -> Result<ResponsesResult, LlmError> {
|
||||
let url = format!("{}/responses", self.base_url);
|
||||
|
||||
tracing::debug!(
|
||||
url = %url,
|
||||
model = %body.get("model").and_then(|m| m.as_str()).unwrap_or("?"),
|
||||
"Codex ChatGPT: sending request"
|
||||
);
|
||||
|
||||
let api_key = self.api_key.read().await.clone();
|
||||
let resp =
|
||||
Self::send_http_request(&self.client, &url, &api_key, &body, self.request_timeout)
|
||||
.await?;
|
||||
|
||||
let status = resp.status();
|
||||
if status.as_u16() == 401 {
|
||||
// Attempt token refresh if we have a refresh token
|
||||
if let Some(ref rt) = self.refresh_token {
|
||||
let _refresh_guard = self.refresh_lock.lock().await;
|
||||
let current_token = self.api_key.read().await.clone();
|
||||
|
||||
if current_token.expose_secret() != api_key.expose_secret() {
|
||||
tracing::info!("Received 401, but another request already refreshed the token");
|
||||
let retry_resp = Self::send_http_request(
|
||||
&self.client,
|
||||
&url,
|
||||
¤t_token,
|
||||
&body,
|
||||
self.request_timeout,
|
||||
)
|
||||
.await?;
|
||||
let retry_status = retry_resp.status();
|
||||
if !retry_status.is_success() {
|
||||
let body_text =
|
||||
tokio::time::timeout(Duration::from_secs(5), retry_resp.text())
|
||||
.await
|
||||
.unwrap_or(Ok(String::new()))
|
||||
.unwrap_or_default();
|
||||
return Err(LlmError::RequestFailed {
|
||||
provider: "codex_chatgpt".to_string(),
|
||||
reason: format!(
|
||||
"HTTP {retry_status} from {url} (after concurrent token refresh): {body_text}"
|
||||
),
|
||||
});
|
||||
}
|
||||
return Self::parse_sse_response_stream(retry_resp, self.request_timeout).await;
|
||||
}
|
||||
|
||||
tracing::info!("Received 401, attempting token refresh");
|
||||
if let Some(new_token) =
|
||||
codex_auth::refresh_access_token(&self.client, rt, self.auth_path.as_deref())
|
||||
.await
|
||||
{
|
||||
// Update stored api_key
|
||||
*self.api_key.write().await = new_token.clone();
|
||||
tracing::info!("Token refreshed, retrying request");
|
||||
|
||||
// Retry the request with the new token
|
||||
let retry_resp = Self::send_http_request(
|
||||
&self.client,
|
||||
&url,
|
||||
&new_token,
|
||||
&body,
|
||||
self.request_timeout,
|
||||
)
|
||||
.await?;
|
||||
|
||||
let retry_status = retry_resp.status();
|
||||
if !retry_status.is_success() {
|
||||
let body_text =
|
||||
tokio::time::timeout(Duration::from_secs(5), retry_resp.text())
|
||||
.await
|
||||
.unwrap_or(Ok(String::new()))
|
||||
.unwrap_or_default();
|
||||
return Err(LlmError::RequestFailed {
|
||||
provider: "codex_chatgpt".to_string(),
|
||||
reason: format!(
|
||||
"HTTP {retry_status} from {url} (after token refresh): {body_text}"
|
||||
),
|
||||
});
|
||||
}
|
||||
|
||||
return Self::parse_sse_response_stream(retry_resp, self.request_timeout).await;
|
||||
} else {
|
||||
tracing::warn!(
|
||||
"Token refresh failed. Please re-authenticate with: codex --login"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// No refresh token or refresh failed — return the 401 error
|
||||
// Drain the response body to release the connection
|
||||
let _ = resp.text().await;
|
||||
return Err(LlmError::AuthFailed {
|
||||
provider: "codex_chatgpt".to_string(),
|
||||
});
|
||||
}
|
||||
|
||||
if !status.is_success() {
|
||||
// Read the error body with a timeout to avoid hanging
|
||||
let body_text = tokio::time::timeout(Duration::from_secs(5), resp.text())
|
||||
.await
|
||||
.unwrap_or(Ok(String::new()))
|
||||
.unwrap_or_default();
|
||||
return Err(LlmError::RequestFailed {
|
||||
provider: "codex_chatgpt".to_string(),
|
||||
reason: format!("HTTP {status} from {url}: {body_text}",),
|
||||
});
|
||||
}
|
||||
|
||||
Self::parse_sse_response_stream(resp, self.request_timeout).await
|
||||
}
|
||||
|
||||
/// Low-level HTTP POST to the /responses endpoint.
|
||||
async fn send_http_request(
|
||||
client: &Client,
|
||||
url: &str,
|
||||
api_key: &SecretString,
|
||||
body: &Value,
|
||||
timeout: Duration,
|
||||
) -> Result<reqwest::Response, LlmError> {
|
||||
client
|
||||
.post(url)
|
||||
.bearer_auth(api_key.expose_secret())
|
||||
.header("Content-Type", "application/json")
|
||||
.header("Accept", "text/event-stream")
|
||||
.json(body)
|
||||
.timeout(timeout)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| LlmError::RequestFailed {
|
||||
provider: "codex_chatgpt".to_string(),
|
||||
reason: format!("HTTP request failed: {e}"),
|
||||
})
|
||||
}
|
||||
|
||||
async fn parse_sse_response_stream(
|
||||
resp: reqwest::Response,
|
||||
idle_timeout: Duration,
|
||||
) -> Result<ResponsesResult, LlmError> {
|
||||
let stream = resp
|
||||
.bytes_stream()
|
||||
.map(|chunk| chunk.map_err(|e| e.to_string()));
|
||||
Self::parse_sse_stream(stream, idle_timeout).await
|
||||
}
|
||||
|
||||
async fn parse_sse_stream<S>(
|
||||
stream: S,
|
||||
idle_timeout: Duration,
|
||||
) -> Result<ResponsesResult, LlmError>
|
||||
where
|
||||
S: Stream<Item = Result<bytes::Bytes, String>> + Unpin,
|
||||
{
|
||||
let mut result = ResponsesResult::default();
|
||||
let mut stream = stream.eventsource();
|
||||
|
||||
loop {
|
||||
match tokio::time::timeout(idle_timeout, stream.next()).await {
|
||||
Ok(Some(Ok(event))) => {
|
||||
let data = event.data.trim();
|
||||
if data.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let parsed: Value = match serde_json::from_str(data) {
|
||||
Ok(v) => v,
|
||||
Err(_) => continue,
|
||||
};
|
||||
|
||||
if Self::handle_sse_event(&mut result, event.event.as_str(), &parsed) {
|
||||
return Ok(result);
|
||||
}
|
||||
}
|
||||
Ok(Some(Err(e))) => {
|
||||
return Err(LlmError::RequestFailed {
|
||||
provider: "codex_chatgpt".to_string(),
|
||||
reason: format!("Failed to read SSE stream: {e}"),
|
||||
});
|
||||
}
|
||||
Ok(None) => return Ok(result),
|
||||
Err(_) => {
|
||||
return Err(LlmError::RequestFailed {
|
||||
provider: "codex_chatgpt".to_string(),
|
||||
reason: format!(
|
||||
"Timed out waiting for SSE event after {}s",
|
||||
idle_timeout.as_secs()
|
||||
),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Parse SSE events from the response text.
|
||||
#[cfg(test)]
|
||||
fn parse_sse_response(sse_text: &str) -> Result<ResponsesResult, LlmError> {
|
||||
let mut result = ResponsesResult::default();
|
||||
let mut current_event_type = String::new();
|
||||
|
||||
for line in sse_text.lines() {
|
||||
if let Some(event) = line.strip_prefix("event: ") {
|
||||
current_event_type = event.trim().to_string();
|
||||
continue;
|
||||
}
|
||||
|
||||
if let Some(data) = line.strip_prefix("data: ") {
|
||||
let data = data.trim();
|
||||
if data.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let parsed: Value = match serde_json::from_str(data) {
|
||||
Ok(v) => v,
|
||||
Err(_) => continue,
|
||||
};
|
||||
|
||||
if Self::handle_sse_event(&mut result, current_event_type.as_str(), &parsed) {
|
||||
return Ok(result);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
fn handle_sse_event(result: &mut ResponsesResult, event_type: &str, parsed: &Value) -> bool {
|
||||
match event_type {
|
||||
"response.output_text.delta" => {
|
||||
if let Some(delta) = parsed.get("delta").and_then(|d| d.as_str()) {
|
||||
result.text.push_str(delta);
|
||||
}
|
||||
}
|
||||
"response.output_item.added" => {
|
||||
// Capture function call metadata when the item is first added.
|
||||
// The item has: id (item_id), call_id, name, type.
|
||||
let item = parsed.get("item").unwrap_or(parsed);
|
||||
if item.get("type").and_then(|t| t.as_str()) == Some("function_call") {
|
||||
let item_id = item
|
||||
.get("id")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
let call_id = item
|
||||
.get("call_id")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
let name = item
|
||||
.get("name")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("")
|
||||
.to_string();
|
||||
|
||||
result
|
||||
.pending_tool_calls
|
||||
.entry(item_id)
|
||||
.or_insert_with(|| PendingToolCall {
|
||||
call_id,
|
||||
name,
|
||||
arguments: String::new(),
|
||||
});
|
||||
}
|
||||
}
|
||||
"response.function_call_arguments.delta" => {
|
||||
// Delta events use `item_id` (not `call_id`)
|
||||
if let Some(item_id) = parsed.get("item_id").and_then(|v| v.as_str())
|
||||
&& let Some(entry) = result.pending_tool_calls.get_mut(item_id)
|
||||
&& let Some(delta) = parsed.get("delta").and_then(|d| d.as_str())
|
||||
{
|
||||
entry.arguments.push_str(delta);
|
||||
}
|
||||
}
|
||||
"response.completed" => {
|
||||
if let Some(response) = parsed.get("response")
|
||||
&& let Some(usage) = response.get("usage")
|
||||
{
|
||||
result.input_tokens = usage
|
||||
.get("input_tokens")
|
||||
.and_then(|v| v.as_u64())
|
||||
.unwrap_or(0) as u32;
|
||||
result.output_tokens = usage
|
||||
.get("output_tokens")
|
||||
.and_then(|v| v.as_u64())
|
||||
.unwrap_or(0) as u32;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
|
||||
false
|
||||
}
|
||||
|
||||
/// Remove keys with empty-string values from a JSON object.
|
||||
///
|
||||
/// gpt-5.2-codex fills optional tool parameters with `""` (e.g.
|
||||
/// `"timestamp": ""`). IronClaw's tool validation treats these as
|
||||
/// invalid "non-empty input expected". Stripping them makes the
|
||||
/// tool see only the actually-provided values.
|
||||
fn strip_empty_string_values(value: Value) -> Value {
|
||||
match value {
|
||||
Value::Object(map) => {
|
||||
let cleaned: serde_json::Map<String, Value> = map
|
||||
.into_iter()
|
||||
.filter(|(_, v)| !matches!(v, Value::String(s) if s.is_empty()))
|
||||
.map(|(k, v)| (k, Self::strip_empty_string_values(v)))
|
||||
.collect();
|
||||
Value::Object(cleaned)
|
||||
}
|
||||
other => other,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
struct ResponsesResult {
|
||||
text: String,
|
||||
/// Keyed by item_id (the SSE item identifier, e.g. "fc_...").
|
||||
pending_tool_calls: std::collections::HashMap<String, PendingToolCall>,
|
||||
input_tokens: u32,
|
||||
output_tokens: u32,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct PendingToolCall {
|
||||
/// The call_id from the API (e.g. "call_..."), used to match results.
|
||||
call_id: String,
|
||||
name: String,
|
||||
arguments: String,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LlmProvider for CodexChatGptProvider {
|
||||
fn model_name(&self) -> &str {
|
||||
// Return resolved model if available, otherwise the configured name.
|
||||
self.resolved_model
|
||||
.get()
|
||||
.map(|s| s.as_str())
|
||||
.unwrap_or(&self.configured_model)
|
||||
}
|
||||
|
||||
fn cost_per_token(&self) -> (Decimal, Decimal) {
|
||||
// ChatGPT backend doesn't expose per-token pricing
|
||||
(Decimal::ZERO, Decimal::ZERO)
|
||||
}
|
||||
|
||||
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, LlmError> {
|
||||
let model = self.resolve_model().await;
|
||||
let body = self.build_request_body(model, &request.messages, &[], None);
|
||||
let result = self.send_request(body).await?;
|
||||
|
||||
Ok(CompletionResponse {
|
||||
content: result.text,
|
||||
input_tokens: result.input_tokens,
|
||||
output_tokens: result.output_tokens,
|
||||
finish_reason: FinishReason::Stop,
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
})
|
||||
}
|
||||
|
||||
async fn complete_with_tools(
|
||||
&self,
|
||||
request: ToolCompletionRequest,
|
||||
) -> Result<ToolCompletionResponse, LlmError> {
|
||||
let model = self.resolve_model().await;
|
||||
let body = self.build_request_body(
|
||||
model,
|
||||
&request.messages,
|
||||
&request.tools,
|
||||
request.tool_choice.as_deref(),
|
||||
);
|
||||
let result = self.send_request(body).await?;
|
||||
|
||||
let tool_calls: Vec<ToolCall> = result
|
||||
.pending_tool_calls
|
||||
.into_values()
|
||||
.map(|tc| {
|
||||
let args: Value =
|
||||
serde_json::from_str(&tc.arguments).unwrap_or_else(|_| json!(tc.arguments));
|
||||
// gpt-5.2-codex fills optional parameters with empty strings (e.g.
|
||||
// `"timestamp": ""`), which IronClaw's tool validation rejects.
|
||||
// Strip them so only actually-provided values reach the tool.
|
||||
let args = Self::strip_empty_string_values(args);
|
||||
ToolCall {
|
||||
id: tc.call_id,
|
||||
name: tc.name,
|
||||
arguments: args,
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
|
||||
let finish_reason = if tool_calls.is_empty() {
|
||||
FinishReason::Stop
|
||||
} else {
|
||||
FinishReason::ToolUse
|
||||
};
|
||||
|
||||
Ok(ToolCompletionResponse {
|
||||
content: if result.text.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(result.text)
|
||||
},
|
||||
tool_calls,
|
||||
input_tokens: result.input_tokens,
|
||||
output_tokens: result.output_tokens,
|
||||
finish_reason,
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_input_tokens: 0,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use bytes::Bytes;
|
||||
use futures::stream;
|
||||
|
||||
#[test]
|
||||
fn test_message_conversion_user() {
|
||||
let items = CodexChatGptProvider::message_to_input_items(&ChatMessage::user("hello"));
|
||||
assert_eq!(items.len(), 1);
|
||||
assert_eq!(items[0]["type"], "message");
|
||||
assert_eq!(items[0]["role"], "user");
|
||||
assert_eq!(items[0]["content"][0]["type"], "input_text");
|
||||
assert_eq!(items[0]["content"][0]["text"], "hello");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_message_conversion_user_with_image() {
|
||||
use super::super::provider::ImageUrl;
|
||||
let parts = vec![
|
||||
ContentPart::Text {
|
||||
text: "What's in this image?".to_string(),
|
||||
},
|
||||
ContentPart::ImageUrl {
|
||||
image_url: ImageUrl {
|
||||
url: "data:image/png;base64,iVBOR...".to_string(),
|
||||
detail: None,
|
||||
},
|
||||
},
|
||||
];
|
||||
let msg = ChatMessage::user_with_parts("", parts);
|
||||
let items = CodexChatGptProvider::message_to_input_items(&msg);
|
||||
assert_eq!(items.len(), 1);
|
||||
assert_eq!(items[0]["type"], "message");
|
||||
assert_eq!(items[0]["role"], "user");
|
||||
let content = items[0]["content"].as_array().unwrap();
|
||||
assert_eq!(content.len(), 2);
|
||||
assert_eq!(content[0]["type"], "input_text");
|
||||
assert_eq!(content[0]["text"], "What's in this image?");
|
||||
assert_eq!(content[1]["type"], "input_image");
|
||||
assert_eq!(content[1]["image_url"], "data:image/png;base64,iVBOR...");
|
||||
}
|
||||
#[test]
|
||||
fn test_message_conversion_assistant() {
|
||||
let items = CodexChatGptProvider::message_to_input_items(&ChatMessage::assistant("hi"));
|
||||
assert_eq!(items.len(), 1);
|
||||
assert_eq!(items[0]["type"], "message");
|
||||
assert_eq!(items[0]["role"], "assistant");
|
||||
assert_eq!(items[0]["content"][0]["type"], "output_text");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_message_conversion_tool_result() {
|
||||
let msg = ChatMessage::tool_result("call_1", "search", "result text");
|
||||
let items = CodexChatGptProvider::message_to_input_items(&msg);
|
||||
assert_eq!(items.len(), 1);
|
||||
assert_eq!(items[0]["type"], "function_call_output");
|
||||
assert_eq!(items[0]["call_id"], "call_1");
|
||||
assert_eq!(items[0]["output"], "result text");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_message_conversion_assistant_with_tool_calls() {
|
||||
let tc = ToolCall {
|
||||
id: "call_1".to_string(),
|
||||
name: "search".to_string(),
|
||||
arguments: json!({"query": "rust"}),
|
||||
};
|
||||
let msg = ChatMessage::assistant_with_tool_calls(Some("thinking...".into()), vec![tc]);
|
||||
let items = CodexChatGptProvider::message_to_input_items(&msg);
|
||||
// Should produce: 1 text message + 1 function_call
|
||||
assert_eq!(items.len(), 2);
|
||||
assert_eq!(items[0]["type"], "message");
|
||||
assert_eq!(items[1]["type"], "function_call");
|
||||
assert_eq!(items[1]["name"], "search");
|
||||
assert_eq!(items[1]["call_id"], "call_1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_request_extracts_system_as_instructions() {
|
||||
let provider = CodexChatGptProvider::new("https://example.com", "key", "gpt-4o");
|
||||
let messages = vec![
|
||||
ChatMessage::system("You are helpful."),
|
||||
ChatMessage::user("hello"),
|
||||
];
|
||||
let body = provider.build_request_body("gpt-4o", &messages, &[], None);
|
||||
assert_eq!(body["instructions"], "You are helpful.");
|
||||
// input should only contain the user message, not the system message
|
||||
assert_eq!(body["input"].as_array().unwrap().len(), 1);
|
||||
// store must be false for ChatGPT backend
|
||||
assert_eq!(body["store"], false);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_sse_text_response() {
|
||||
let sse = r#"event: response.output_text.delta
|
||||
data: {"delta":"Hello"}
|
||||
|
||||
event: response.output_text.delta
|
||||
data: {"delta":" world!"}
|
||||
|
||||
event: response.completed
|
||||
data: {"response":{"usage":{"input_tokens":10,"output_tokens":5}}}
|
||||
|
||||
"#;
|
||||
let result = CodexChatGptProvider::parse_sse_response(sse).unwrap();
|
||||
assert_eq!(result.text, "Hello world!");
|
||||
assert_eq!(result.input_tokens, 10);
|
||||
assert_eq!(result.output_tokens, 5);
|
||||
assert!(result.pending_tool_calls.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_parse_sse_tool_call() {
|
||||
// Real API format: output_item.added has item.id (item_id) + item.call_id,
|
||||
// delta events use item_id (not call_id)
|
||||
let sse = r#"event: response.output_item.added
|
||||
data: {"item":{"id":"fc_1","type":"function_call","call_id":"call_1","name":"search"}}
|
||||
|
||||
event: response.function_call_arguments.delta
|
||||
data: {"item_id":"fc_1","delta":"{\"query\":"}
|
||||
|
||||
event: response.function_call_arguments.delta
|
||||
data: {"item_id":"fc_1","delta":"\"rust\"}"}
|
||||
|
||||
event: response.completed
|
||||
data: {"response":{"usage":{"input_tokens":20,"output_tokens":15}}}
|
||||
|
||||
"#;
|
||||
let result = CodexChatGptProvider::parse_sse_response(sse).unwrap();
|
||||
assert!(result.text.is_empty());
|
||||
assert_eq!(result.pending_tool_calls.len(), 1);
|
||||
let tc = result.pending_tool_calls.get("fc_1").unwrap();
|
||||
assert_eq!(tc.call_id, "call_1");
|
||||
assert_eq!(tc.name, "search");
|
||||
assert_eq!(tc.arguments, "{\"query\":\"rust\"}");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_parse_sse_stream_response() {
|
||||
let stream = stream::iter(vec![
|
||||
Ok(Bytes::from_static(
|
||||
b"event: response.output_text.delta\ndata: {\"delta\":\"Hello\"}\n\n",
|
||||
)),
|
||||
Ok(Bytes::from_static(
|
||||
b"event: response.output_text.delta\ndata: {\"delta\":\" world\"}\n\n",
|
||||
)),
|
||||
Ok(Bytes::from_static(
|
||||
b"event: response.completed\ndata: {\"response\":{\"usage\":{\"input_tokens\":3,\"output_tokens\":2}}}\n\n",
|
||||
)),
|
||||
]);
|
||||
|
||||
let result = CodexChatGptProvider::parse_sse_stream(stream, Duration::from_secs(1))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(result.text, "Hello world");
|
||||
assert_eq!(result.input_tokens, 3);
|
||||
assert_eq!(result.output_tokens, 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_strip_empty_string_values() {
|
||||
let input = json!({
|
||||
"format": "%Y-%m-%d",
|
||||
"operation": "now",
|
||||
"timestamp": "",
|
||||
"timestamp2": "",
|
||||
});
|
||||
let cleaned = CodexChatGptProvider::strip_empty_string_values(input);
|
||||
assert_eq!(cleaned, json!({"format": "%Y-%m-%d", "operation": "now"}));
|
||||
}
|
||||
}
|
||||
@@ -5,6 +5,8 @@
|
||||
//! extracted into a standalone crate. Resolution logic (reading env vars,
|
||||
//! settings) lives in `crate::config::llm`.
|
||||
|
||||
use std::path::PathBuf;
|
||||
|
||||
use secrecy::SecretString;
|
||||
|
||||
use crate::llm::registry::ProviderProtocol;
|
||||
@@ -85,6 +87,13 @@ pub struct RegistryProviderConfig {
|
||||
/// 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>,
|
||||
/// When true, route OpenAI-compatible traffic to the Codex ChatGPT
|
||||
/// Responses API provider instead of rig-core's Chat Completions path.
|
||||
pub is_codex_chatgpt: bool,
|
||||
/// OAuth refresh token for Codex ChatGPT token refresh.
|
||||
pub refresh_token: Option<SecretString>,
|
||||
/// Path to Codex auth.json for persisting refreshed tokens.
|
||||
pub auth_path: Option<PathBuf>,
|
||||
/// Prompt cache retention (Anthropic-specific).
|
||||
pub cache_retention: CacheRetention,
|
||||
/// Parameter names that this provider does not support (e.g., `["temperature"]`).
|
||||
@@ -129,6 +138,30 @@ pub struct LlmConfig {
|
||||
/// Default: 120. Increase for local LLMs (Ollama, vLLM, LM Studio) that
|
||||
/// need more time for prompt evaluation on consumer hardware.
|
||||
pub request_timeout_secs: u64,
|
||||
/// Generic cheap/fast model for lightweight tasks (heartbeat, routing, evaluation).
|
||||
/// Works with any backend. Set via `LLM_CHEAP_MODEL` env var.
|
||||
/// When set, takes priority over the NearAI-specific `NEARAI_CHEAP_MODEL`.
|
||||
pub cheap_model: Option<String>,
|
||||
/// Enable cascade mode for smart routing (retry with primary if cheap model
|
||||
/// response seems uncertain). Default: true. Set via `SMART_ROUTING_CASCADE`.
|
||||
pub smart_routing_cascade: bool,
|
||||
}
|
||||
|
||||
impl LlmConfig {
|
||||
/// Resolve the effective cheap model name.
|
||||
///
|
||||
/// Resolution order:
|
||||
/// 1. `LLM_CHEAP_MODEL` (generic, works with any backend)
|
||||
/// 2. `NEARAI_CHEAP_MODEL` (NearAI-only, backward compatibility)
|
||||
pub fn cheap_model_name(&self) -> Option<&str> {
|
||||
self.cheap_model.as_deref().or_else(|| {
|
||||
if self.backend == "nearai" {
|
||||
self.nearai.cheap_model.as_deref()
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// NEAR AI configuration.
|
||||
|
||||
+161
-29
@@ -12,6 +12,8 @@ mod anthropic_oauth;
|
||||
#[cfg(feature = "bedrock")]
|
||||
mod bedrock;
|
||||
pub mod circuit_breaker;
|
||||
pub(crate) mod codex_auth;
|
||||
mod codex_chatgpt;
|
||||
pub mod config;
|
||||
pub mod costs;
|
||||
pub mod error;
|
||||
@@ -102,7 +104,7 @@ pub async fn create_llm_provider(
|
||||
provider: config.backend.clone(),
|
||||
})?;
|
||||
|
||||
create_registry_provider(reg_config)
|
||||
create_registry_provider(reg_config, timeout)
|
||||
}
|
||||
|
||||
/// Create an LLM provider from a `NearAiConfig` directly.
|
||||
@@ -140,7 +142,13 @@ pub fn create_llm_provider_with_config(
|
||||
/// `create_*_provider` functions.
|
||||
fn create_registry_provider(
|
||||
config: &RegistryProviderConfig,
|
||||
request_timeout_secs: u64,
|
||||
) -> Result<Arc<dyn LlmProvider>, LlmError> {
|
||||
// Codex ChatGPT mode: use the Responses API provider
|
||||
if config.is_codex_chatgpt {
|
||||
return create_codex_chatgpt_from_registry(config, request_timeout_secs);
|
||||
}
|
||||
|
||||
match config.protocol {
|
||||
ProviderProtocol::OpenAiCompletions => create_openai_compat_from_registry(config),
|
||||
ProviderProtocol::Anthropic => create_anthropic_from_registry(config),
|
||||
@@ -148,6 +156,36 @@ fn create_registry_provider(
|
||||
}
|
||||
}
|
||||
|
||||
fn create_codex_chatgpt_from_registry(
|
||||
config: &RegistryProviderConfig,
|
||||
request_timeout_secs: u64,
|
||||
) -> Result<Arc<dyn LlmProvider>, LlmError> {
|
||||
let api_key = config
|
||||
.api_key
|
||||
.as_ref()
|
||||
.cloned()
|
||||
.ok_or_else(|| LlmError::AuthFailed {
|
||||
provider: "codex_chatgpt".to_string(),
|
||||
})?;
|
||||
|
||||
tracing::info!(
|
||||
configured_model = %config.model,
|
||||
base_url = %config.base_url,
|
||||
"Using Codex ChatGPT provider (Responses API) — model detection deferred to first call"
|
||||
);
|
||||
|
||||
let provider = codex_chatgpt::CodexChatGptProvider::with_lazy_model(
|
||||
&config.base_url,
|
||||
api_key,
|
||||
&config.model,
|
||||
config.refresh_token.clone(),
|
||||
config.auth_path.clone(),
|
||||
request_timeout_secs,
|
||||
);
|
||||
|
||||
Ok(Arc::new(provider))
|
||||
}
|
||||
|
||||
#[cfg(feature = "bedrock")]
|
||||
async fn create_bedrock_provider(config: &LlmConfig) -> Result<Arc<dyn LlmProvider>, LlmError> {
|
||||
let br = config
|
||||
@@ -163,6 +201,7 @@ async fn create_bedrock_provider(config: &LlmConfig) -> Result<Arc<dyn LlmProvid
|
||||
br.region,
|
||||
provider.active_model_name(),
|
||||
);
|
||||
|
||||
Ok(Arc::new(provider))
|
||||
}
|
||||
|
||||
@@ -337,32 +376,61 @@ fn create_ollama_from_registry(
|
||||
|
||||
/// Create a cheap/fast LLM provider for lightweight tasks (heartbeat, routing, evaluation).
|
||||
///
|
||||
/// Uses `NEARAI_CHEAP_MODEL` if set, otherwise falls back to the main provider.
|
||||
/// Currently only supports NEAR AI backend.
|
||||
/// Resolution order:
|
||||
/// 1. `LLM_CHEAP_MODEL` (generic, works with any backend)
|
||||
/// 2. `NEARAI_CHEAP_MODEL` (NearAI-only, backward compatibility)
|
||||
///
|
||||
/// Returns `None` if no cheap model is configured.
|
||||
pub fn create_cheap_llm_provider(
|
||||
config: &LlmConfig,
|
||||
session: Arc<SessionManager>,
|
||||
) -> Result<Option<Arc<dyn LlmProvider>>, LlmError> {
|
||||
let Some(ref cheap_model) = config.nearai.cheap_model else {
|
||||
let Some(cheap_model) = config.cheap_model_name() else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
if config.backend != "nearai" {
|
||||
tracing::warn!(
|
||||
"NEARAI_CHEAP_MODEL is set but LLM_BACKEND is '{}', not nearai. \
|
||||
Cheap model setting will be ignored.",
|
||||
config.backend
|
||||
);
|
||||
return Ok(None);
|
||||
create_cheap_provider_for_backend(config, session, cheap_model)
|
||||
}
|
||||
|
||||
/// Create a cheap provider for a specific backend.
|
||||
///
|
||||
/// Handles backend-specific provider construction:
|
||||
/// - `nearai` — clones NearAiConfig, swaps model, uses `create_llm_provider_with_config`
|
||||
/// - `bedrock` — returns error (smart routing not yet supported)
|
||||
/// - All others — clones `RegistryProviderConfig`, swaps model, uses `create_registry_provider`
|
||||
fn create_cheap_provider_for_backend(
|
||||
config: &LlmConfig,
|
||||
session: Arc<SessionManager>,
|
||||
cheap_model: &str,
|
||||
) -> Result<Option<Arc<dyn LlmProvider>>, LlmError> {
|
||||
if config.backend == "nearai" {
|
||||
let mut cheap_config = config.nearai.clone();
|
||||
cheap_config.model = cheap_model.to_string();
|
||||
let provider =
|
||||
create_llm_provider_with_config(&cheap_config, session, config.request_timeout_secs)?;
|
||||
return Ok(Some(provider));
|
||||
}
|
||||
|
||||
let mut cheap_config = config.nearai.clone();
|
||||
cheap_config.model = cheap_model.clone();
|
||||
if config.backend == "bedrock" {
|
||||
return Err(LlmError::RequestFailed {
|
||||
provider: "bedrock".to_string(),
|
||||
reason: "Smart routing with cheap model is not supported for Bedrock yet".to_string(),
|
||||
});
|
||||
}
|
||||
|
||||
Ok(Some(Arc::new(NearAiChatProvider::new(
|
||||
cheap_config,
|
||||
session,
|
||||
)?)))
|
||||
// Registry-based provider: clone config and swap model
|
||||
let reg_config = config.provider.as_ref().ok_or_else(|| LlmError::RequestFailed {
|
||||
provider: config.backend.clone(),
|
||||
reason: format!(
|
||||
"Cannot create cheap provider for backend '{}': no registry provider config available",
|
||||
config.backend
|
||||
),
|
||||
})?;
|
||||
|
||||
let mut cheap_reg_config = reg_config.clone();
|
||||
cheap_reg_config.model = cheap_model.to_string();
|
||||
let provider = create_registry_provider(&cheap_reg_config, config.request_timeout_secs)?;
|
||||
Ok(Some(provider))
|
||||
}
|
||||
|
||||
/// Build the full LLM provider chain with all configured wrappers.
|
||||
@@ -410,14 +478,15 @@ pub async fn build_provider_chain(
|
||||
};
|
||||
|
||||
// 2. Smart routing (cheap/primary split)
|
||||
let llm: Arc<dyn LlmProvider> = if let Some(ref cheap_model) = config.nearai.cheap_model {
|
||||
let mut cheap_config = config.nearai.clone();
|
||||
cheap_config.model = cheap_model.clone();
|
||||
let cheap = create_llm_provider_with_config(
|
||||
&cheap_config,
|
||||
session.clone(),
|
||||
config.request_timeout_secs,
|
||||
)?;
|
||||
let llm: Arc<dyn LlmProvider> = if let Some(cheap_model) = config.cheap_model_name() {
|
||||
let cheap = create_cheap_provider_for_backend(config, session.clone(), cheap_model)?
|
||||
.ok_or_else(|| LlmError::RequestFailed {
|
||||
provider: config.backend.clone(),
|
||||
reason: format!(
|
||||
"Failed to create cheap provider for model '{cheap_model}' on backend '{}'",
|
||||
config.backend
|
||||
),
|
||||
})?;
|
||||
let cheap: Arc<dyn LlmProvider> = if retry_config.max_retries > 0 {
|
||||
Arc::new(RetryProvider::new(cheap, retry_config.clone()))
|
||||
} else {
|
||||
@@ -432,7 +501,7 @@ pub async fn build_provider_chain(
|
||||
llm,
|
||||
cheap,
|
||||
SmartRoutingConfig {
|
||||
cascade_enabled: config.nearai.smart_routing_cascade,
|
||||
cascade_enabled: config.smart_routing_cascade,
|
||||
..SmartRoutingConfig::default()
|
||||
},
|
||||
))
|
||||
@@ -561,6 +630,8 @@ mod tests {
|
||||
provider: None,
|
||||
bedrock: None,
|
||||
request_timeout_secs: 120,
|
||||
cheap_model: None,
|
||||
smart_routing_cascade: true,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -575,7 +646,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_cheap_llm_provider_creates_provider_when_configured() {
|
||||
fn test_create_cheap_llm_provider_creates_provider_with_nearai_cheap_model() {
|
||||
let mut config = test_llm_config();
|
||||
config.nearai.cheap_model = Some("cheap-test-model".to_string());
|
||||
|
||||
@@ -589,7 +660,26 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_cheap_llm_provider_ignored_for_non_nearai_backend() {
|
||||
fn test_create_cheap_llm_provider_generic_overrides_nearai() {
|
||||
let mut config = test_llm_config();
|
||||
config.nearai.cheap_model = Some("nearai-cheap".to_string());
|
||||
config.cheap_model = Some("generic-cheap".to_string());
|
||||
|
||||
let session = Arc::new(SessionManager::new(SessionConfig::default()));
|
||||
let result = create_cheap_llm_provider(&config, session);
|
||||
|
||||
assert!(result.is_ok());
|
||||
let provider = result.unwrap();
|
||||
assert!(provider.is_some());
|
||||
assert_eq!(
|
||||
provider.unwrap().model_name(),
|
||||
"generic-cheap",
|
||||
"LLM_CHEAP_MODEL should take priority over NEARAI_CHEAP_MODEL"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_cheap_llm_provider_nearai_cheap_ignored_for_non_nearai_backend() {
|
||||
let mut config = test_llm_config();
|
||||
config.backend = "openai".to_string();
|
||||
config.nearai.cheap_model = Some("cheap-test-model".to_string());
|
||||
@@ -598,6 +688,48 @@ mod tests {
|
||||
let result = create_cheap_llm_provider(&config, session);
|
||||
|
||||
assert!(result.is_ok());
|
||||
assert!(result.unwrap().is_none());
|
||||
assert!(
|
||||
result.unwrap().is_none(),
|
||||
"NEARAI_CHEAP_MODEL should be ignored when backend is not nearai"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_cheap_llm_provider_bedrock_returns_error() {
|
||||
let mut config = test_llm_config();
|
||||
config.backend = "bedrock".to_string();
|
||||
config.cheap_model = Some("cheap-model".to_string());
|
||||
|
||||
let session = Arc::new(SessionManager::new(SessionConfig::default()));
|
||||
let result = create_cheap_llm_provider(&config, session);
|
||||
|
||||
assert!(
|
||||
result.is_err(),
|
||||
"Bedrock should return an error for cheap model"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_cheap_model_name_resolution() {
|
||||
// Generic takes priority
|
||||
let mut config = test_llm_config();
|
||||
config.cheap_model = Some("generic".to_string());
|
||||
config.nearai.cheap_model = Some("nearai".to_string());
|
||||
assert_eq!(config.cheap_model_name(), Some("generic"));
|
||||
|
||||
// NearAI fallback when backend is nearai
|
||||
let mut config = test_llm_config();
|
||||
config.nearai.cheap_model = Some("nearai".to_string());
|
||||
assert_eq!(config.cheap_model_name(), Some("nearai"));
|
||||
|
||||
// NearAI ignored for non-nearai backend
|
||||
let mut config = test_llm_config();
|
||||
config.backend = "openai".to_string();
|
||||
config.nearai.cheap_model = Some("nearai".to_string());
|
||||
assert_eq!(config.cheap_model_name(), None);
|
||||
|
||||
// None when nothing configured
|
||||
let config = test_llm_config();
|
||||
assert_eq!(config.cheap_model_name(), None);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -345,5 +345,7 @@ pub(crate) fn build_nearai_model_fetch_config() -> crate::config::LlmConfig {
|
||||
provider: None,
|
||||
bedrock: None,
|
||||
request_timeout_secs: 120,
|
||||
cheap_model: None,
|
||||
smart_routing_cascade: false,
|
||||
}
|
||||
}
|
||||
|
||||
+114
-1
@@ -244,6 +244,7 @@ impl NearAiChatProvider {
|
||||
let status = response.status();
|
||||
// Extract Retry-After header before consuming the response body.
|
||||
// Supports both delay-seconds (RFC 7231 §7.1.3) and HTTP-date formats.
|
||||
// Falls back to 60s if header is missing or unparseable (prevents "retry after None" errors).
|
||||
let retry_after_header = response
|
||||
.headers()
|
||||
.get("retry-after")
|
||||
@@ -264,7 +265,8 @@ impl NearAiChatProvider {
|
||||
));
|
||||
}
|
||||
None
|
||||
});
|
||||
})
|
||||
.or(Some(std::time::Duration::from_secs(60)));
|
||||
let response_text = response.text().await.map_err(|e| LlmError::RequestFailed {
|
||||
provider: "nearai_chat".to_string(),
|
||||
reason: format!("Failed to read response body: {}", e),
|
||||
@@ -2216,4 +2218,115 @@ mod tests {
|
||||
"http://example.com/api/proxy/v1/chat/completions"
|
||||
);
|
||||
}
|
||||
|
||||
// -- Retry-After header parsing tests (regression for rate limit "None" bug) --
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_parsing_delay_seconds() {
|
||||
// Verify delay-seconds format (most common) is parsed correctly
|
||||
let header_value = "30";
|
||||
let duration = parse_retry_after_for_test(header_value);
|
||||
assert_eq!(duration, Some(std::time::Duration::from_secs(30)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_parsing_rfc2822_date() {
|
||||
// Verify HTTP-date (RFC 2822) format is parsed correctly
|
||||
// Use a date 60 seconds in the future
|
||||
let now = chrono::Utc::now();
|
||||
let future = now + chrono::Duration::seconds(60);
|
||||
let date_str = future.to_rfc2822();
|
||||
|
||||
let duration = parse_retry_after_for_test(&date_str);
|
||||
assert!(duration.is_some());
|
||||
let d = duration.unwrap();
|
||||
// Allow ±5 seconds of drift due to processing time
|
||||
assert!(
|
||||
d.as_secs() >= 55 && d.as_secs() <= 65,
|
||||
"Expected ~60s, got {}s",
|
||||
d.as_secs()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_fallback_missing_header() {
|
||||
// Regression test: When Retry-After header is missing,
|
||||
// should fall back to 60s instead of None
|
||||
let duration = parse_retry_after_for_test("");
|
||||
assert_eq!(
|
||||
duration,
|
||||
Some(std::time::Duration::from_secs(60)),
|
||||
"Missing header should fallback to 60s"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_fallback_invalid_format() {
|
||||
// Regression test: When Retry-After header is in unexpected format,
|
||||
// should fall back to 60s instead of None
|
||||
let invalid_formats = vec![
|
||||
"invalid",
|
||||
"not-a-number",
|
||||
"30.5", // float instead of int
|
||||
"abc123",
|
||||
];
|
||||
|
||||
for format in invalid_formats {
|
||||
let duration = parse_retry_after_for_test(format);
|
||||
assert_eq!(
|
||||
duration,
|
||||
Some(std::time::Duration::from_secs(60)),
|
||||
"Invalid format '{}' should fallback to 60s",
|
||||
format
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_past_date_returns_zero() {
|
||||
// When HTTP-date is in the past, should return Duration::ZERO
|
||||
// (not None, which would trigger immediate retry)
|
||||
let past = chrono::Utc::now() - chrono::Duration::seconds(60);
|
||||
let past_date_str = past.to_rfc2822();
|
||||
|
||||
let duration = parse_retry_after_for_test(&past_date_str);
|
||||
assert_eq!(
|
||||
duration,
|
||||
Some(std::time::Duration::ZERO),
|
||||
"Past date should return Duration::ZERO, not None"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_zero_seconds_accepted() {
|
||||
// Verify zero seconds is a valid retry delay
|
||||
let duration = parse_retry_after_for_test("0");
|
||||
assert_eq!(duration, Some(std::time::Duration::ZERO));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_large_number() {
|
||||
// Verify large numbers are accepted
|
||||
let duration = parse_retry_after_for_test("3600"); // 1 hour
|
||||
assert_eq!(duration, Some(std::time::Duration::from_secs(3600)));
|
||||
}
|
||||
|
||||
/// Helper function to test Retry-After header parsing logic
|
||||
/// (simulates the parsing done in send_request without actual HTTP, including fallback)
|
||||
fn parse_retry_after_for_test(header_value: &str) -> Option<std::time::Duration> {
|
||||
let trimmed = header_value.trim();
|
||||
let parsed = if let Ok(secs) = trimmed.parse::<u64>() {
|
||||
Some(std::time::Duration::from_secs(secs))
|
||||
} else if let Ok(dt) = chrono::DateTime::parse_from_rfc2822(trimmed) {
|
||||
let now = chrono::Utc::now();
|
||||
let delta = dt.signed_duration_since(now);
|
||||
Some(std::time::Duration::from_secs(
|
||||
delta.num_seconds().max(0) as u64
|
||||
))
|
||||
} else {
|
||||
None
|
||||
};
|
||||
// Apply fallback to 60s if parsing failed (matches actual code behavior)
|
||||
parsed.or(Some(std::time::Duration::from_secs(60)))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -394,4 +394,31 @@ mod tests {
|
||||
assert_eq!(retry.cost_per_token(), (Decimal::ZERO, Decimal::ZERO));
|
||||
assert_eq!(retry.calculate_cost(100, 50), Decimal::ZERO);
|
||||
}
|
||||
|
||||
// Regression test: Rate limiter fallback when Retry-After header is missing
|
||||
//
|
||||
// Verifies that RateLimited errors always have a duration (never None)
|
||||
// due to the 60-second fallback applied in all rate limit error creation sites
|
||||
// (nearai_chat.rs, anthropic_oauth.rs, embeddings.rs).
|
||||
#[test]
|
||||
fn rate_limited_error_always_has_duration() {
|
||||
let err = LlmError::RateLimited {
|
||||
provider: "test".to_string(),
|
||||
retry_after: Some(std::time::Duration::from_secs(60)),
|
||||
};
|
||||
|
||||
if let LlmError::RateLimited { retry_after, .. } = err {
|
||||
assert!(
|
||||
retry_after.is_some(),
|
||||
"Rate limited error should always have retry_after duration"
|
||||
);
|
||||
assert_eq!(
|
||||
retry_after,
|
||||
Some(std::time::Duration::from_secs(60)),
|
||||
"Fallback should be 60 seconds"
|
||||
);
|
||||
} else {
|
||||
panic!("Expected RateLimited error");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+87
-7
@@ -357,15 +357,31 @@ fn convert_messages(messages: &[ChatMessage]) -> (Option<String>, Vec<RigMessage
|
||||
}
|
||||
}
|
||||
crate::llm::Role::Tool => {
|
||||
// Tool result message: wrap as User { ToolResult }
|
||||
// Tool result message: wrap as User { ToolResult }.
|
||||
// Merge consecutive tool results into a single User message
|
||||
// so the API sees one multi-result message instead of
|
||||
// multiple consecutive User messages (which Anthropic rejects).
|
||||
let tool_id = normalized_tool_call_id(msg.tool_call_id.as_deref(), history.len());
|
||||
history.push(RigMessage::User {
|
||||
content: OneOrMany::one(UserContent::ToolResult(RigToolResult {
|
||||
id: tool_id.clone(),
|
||||
call_id: Some(tool_id),
|
||||
content: OneOrMany::one(ToolResultContent::text(&msg.content)),
|
||||
})),
|
||||
let tool_result = UserContent::ToolResult(RigToolResult {
|
||||
id: tool_id.clone(),
|
||||
call_id: Some(tool_id),
|
||||
content: OneOrMany::one(ToolResultContent::text(&msg.content)),
|
||||
});
|
||||
|
||||
let should_merge = matches!(
|
||||
history.last(),
|
||||
Some(RigMessage::User { content }) if content.iter().all(|c| matches!(c, UserContent::ToolResult(_)))
|
||||
);
|
||||
|
||||
if should_merge {
|
||||
if let Some(RigMessage::User { content }) = history.last_mut() {
|
||||
content.push(tool_result);
|
||||
}
|
||||
} else {
|
||||
history.push(RigMessage::User {
|
||||
content: OneOrMany::one(tool_result),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1280,4 +1296,68 @@ mod tests {
|
||||
|
||||
assert!(adapter.unsupported_params.is_empty());
|
||||
}
|
||||
|
||||
/// Regression test: consecutive tool_result messages from parallel tool
|
||||
/// execution must be merged into a single User message with multiple
|
||||
/// ToolResult content items. Without merging, APIs like Anthropic reject
|
||||
/// the request due to consecutive User messages.
|
||||
#[test]
|
||||
fn test_consecutive_tool_results_merged_into_single_user_message() {
|
||||
let tc1 = IronToolCall {
|
||||
id: "call_a".to_string(),
|
||||
name: "search".to_string(),
|
||||
arguments: serde_json::json!({"q": "rust"}),
|
||||
};
|
||||
let tc2 = IronToolCall {
|
||||
id: "call_b".to_string(),
|
||||
name: "fetch".to_string(),
|
||||
arguments: serde_json::json!({"url": "https://example.com"}),
|
||||
};
|
||||
let assistant = ChatMessage::assistant_with_tool_calls(None, vec![tc1, tc2]);
|
||||
let result_a = ChatMessage::tool_result("call_a", "search", "search results");
|
||||
let result_b = ChatMessage::tool_result("call_b", "fetch", "fetch results");
|
||||
|
||||
let messages = vec![assistant, result_a, result_b];
|
||||
let (_preamble, history) = convert_messages(&messages);
|
||||
|
||||
// Should be: 1 assistant + 1 merged user (not 1 assistant + 2 users)
|
||||
assert_eq!(
|
||||
history.len(),
|
||||
2,
|
||||
"Expected 2 messages (assistant + merged user), got {}",
|
||||
history.len()
|
||||
);
|
||||
|
||||
// The second message should contain both tool results
|
||||
match &history[1] {
|
||||
RigMessage::User { content } => {
|
||||
assert_eq!(
|
||||
content.len(),
|
||||
2,
|
||||
"Expected 2 tool results in merged user message, got {}",
|
||||
content.len()
|
||||
);
|
||||
for item in content.iter() {
|
||||
assert!(
|
||||
matches!(item, UserContent::ToolResult(_)),
|
||||
"Expected ToolResult content"
|
||||
);
|
||||
}
|
||||
}
|
||||
other => panic!("Expected User message, got: {:?}", other),
|
||||
}
|
||||
}
|
||||
|
||||
/// Verify that a tool_result after a non-tool User message is NOT merged.
|
||||
#[test]
|
||||
fn test_tool_result_after_user_text_not_merged() {
|
||||
let user_msg = ChatMessage::user("hello");
|
||||
let tool_msg = ChatMessage::tool_result("call_1", "search", "results");
|
||||
|
||||
let messages = vec![user_msg, tool_msg];
|
||||
let (_preamble, history) = convert_messages(&messages);
|
||||
|
||||
// Should be 2 separate User messages (text user + tool result user)
|
||||
assert_eq!(history.len(), 2);
|
||||
}
|
||||
}
|
||||
|
||||
+46
-17
@@ -27,6 +27,8 @@ use ironclaw::{
|
||||
webhooks::{self, ToolWebhookState},
|
||||
};
|
||||
|
||||
#[cfg(unix)]
|
||||
use ironclaw::channels::ChannelSecretUpdater;
|
||||
#[cfg(any(feature = "postgres", feature = "libsql"))]
|
||||
use ironclaw::setup::{SetupConfig, SetupWizard};
|
||||
|
||||
@@ -151,7 +153,8 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
provider_only: *provider_only,
|
||||
quick: *quick,
|
||||
};
|
||||
let mut wizard = SetupWizard::with_config(config);
|
||||
let mut wizard =
|
||||
SetupWizard::try_with_config_and_toml(config, cli.config.as_deref())?;
|
||||
wizard.run().await?;
|
||||
}
|
||||
#[cfg(not(any(feature = "postgres", feature = "libsql")))]
|
||||
@@ -193,10 +196,13 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
{
|
||||
println!("Onboarding needed: {}", reason);
|
||||
println!();
|
||||
let mut wizard = SetupWizard::with_config(SetupConfig {
|
||||
quick: true,
|
||||
..Default::default()
|
||||
});
|
||||
let mut wizard = SetupWizard::try_with_config_and_toml(
|
||||
SetupConfig {
|
||||
quick: true,
|
||||
..Default::default()
|
||||
},
|
||||
cli.config.as_deref(),
|
||||
)?;
|
||||
wizard.run().await?;
|
||||
}
|
||||
|
||||
@@ -280,9 +286,12 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
|
||||
// Create CLI channel
|
||||
let repl_channel = if let Some(ref msg) = cli.message {
|
||||
Some(ReplChannel::with_message(msg.clone()))
|
||||
Some(ReplChannel::with_message_for_user(
|
||||
config.owner_id.clone(),
|
||||
msg.clone(),
|
||||
))
|
||||
} else if config.channels.cli.enabled {
|
||||
let repl = ReplChannel::new();
|
||||
let repl = ReplChannel::with_user_id(config.owner_id.clone());
|
||||
repl.suppress_banner();
|
||||
Some(repl)
|
||||
} else {
|
||||
@@ -309,12 +318,7 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
webhook_routes.push(webhooks::routes(ToolWebhookState {
|
||||
tools: Arc::clone(&components.tools),
|
||||
routine_engine: Arc::clone(&shared_routine_engine_slot),
|
||||
user_id: config
|
||||
.channels
|
||||
.gateway
|
||||
.as_ref()
|
||||
.map(|g| g.user_id.clone())
|
||||
.unwrap_or_else(|| "default".to_string()),
|
||||
user_id: config.owner_id.clone(),
|
||||
secrets_store: components.secrets_store.clone(),
|
||||
}));
|
||||
|
||||
@@ -521,6 +525,30 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
}
|
||||
}
|
||||
|
||||
// Persist auto-generated auth token so it survives restarts.
|
||||
// Write to the "default" settings namespace, which is the namespace
|
||||
// Config::from_db() reads from — NOT the gateway channel's user_id.
|
||||
if gw_config.auth_token.is_none() {
|
||||
let token_to_persist = gw.auth_token().to_string();
|
||||
if let Some(ref db) = components.db {
|
||||
let db = db.clone();
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = db
|
||||
.set_setting(
|
||||
"default",
|
||||
"channels.gateway_auth_token",
|
||||
&serde_json::Value::String(token_to_persist),
|
||||
)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to persist auto-generated gateway auth token: {e}");
|
||||
} else {
|
||||
tracing::debug!("Persisted auto-generated gateway auth token to settings");
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
gateway_url = Some(format!(
|
||||
"http://{}:{}/?token={}",
|
||||
gw_config.host,
|
||||
@@ -592,7 +620,7 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
// Register message tool for sending messages to connected channels
|
||||
components
|
||||
.tools
|
||||
.register_message_tools(Arc::clone(&channels))
|
||||
.register_message_tools(Arc::clone(&channels), components.extension_manager.clone())
|
||||
.await;
|
||||
|
||||
// Wire up channel runtime for hot-activation of WASM channels.
|
||||
@@ -677,6 +705,7 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
.map(|db| Arc::clone(db) as Arc<dyn ironclaw::db::SettingsStore>);
|
||||
|
||||
let deps = AgentDeps {
|
||||
owner_id: config.owner_id.clone(),
|
||||
store: components.db,
|
||||
llm: components.llm,
|
||||
cheap_llm: components.cheap_llm,
|
||||
@@ -740,7 +769,6 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
|
||||
#[cfg(unix)]
|
||||
{
|
||||
use ironclaw::channels::ChannelSecretUpdater;
|
||||
// Collect all channels that support secret updates
|
||||
let mut secret_updaters: Vec<Arc<dyn ChannelSecretUpdater>> = Vec::new();
|
||||
if let Some(ref state) = http_channel_state {
|
||||
@@ -750,6 +778,7 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
let sighup_webhook_server = webhook_server.clone();
|
||||
let sighup_settings_store_clone = sighup_settings_store.clone();
|
||||
let sighup_secrets_store = components.secrets_store.clone();
|
||||
let sighup_owner_id = config.owner_id.clone();
|
||||
let mut shutdown_rx = shutdown_tx.subscribe();
|
||||
|
||||
tokio::spawn(async move {
|
||||
@@ -780,7 +809,7 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
if let Some(ref secrets_store) = sighup_secrets_store {
|
||||
// Inject HTTP webhook secret from encrypted store
|
||||
if let Ok(webhook_secret) = secrets_store
|
||||
.get_decrypted("default", "http_webhook_secret")
|
||||
.get_decrypted(&sighup_owner_id, "http_webhook_secret")
|
||||
.await
|
||||
{
|
||||
// Thread-safe: Uses INJECTED_VARS mutex instead of unsafe std::env::set_var
|
||||
@@ -796,7 +825,7 @@ async fn async_main() -> anyhow::Result<()> {
|
||||
// Reload config (now with secrets injected into environment)
|
||||
let new_config = match &sighup_settings_store_clone {
|
||||
Some(store) => {
|
||||
ironclaw::config::Config::from_db(store.as_ref(), "default").await
|
||||
ironclaw::config::Config::from_db(store.as_ref(), &sighup_owner_id).await
|
||||
}
|
||||
None => ironclaw::config::Config::from_env().await,
|
||||
};
|
||||
|
||||
+49
-2
@@ -51,6 +51,15 @@ use crate::db::Database;
|
||||
use crate::llm::LlmProvider;
|
||||
use crate::secrets::SecretsStore;
|
||||
|
||||
/// Resolve the orchestrator port from the `ORCHESTRATOR_PORT` environment
|
||||
/// variable, falling back to 50051.
|
||||
fn resolve_orchestrator_port() -> u16 {
|
||||
std::env::var("ORCHESTRATOR_PORT")
|
||||
.ok()
|
||||
.and_then(|v| v.parse().ok())
|
||||
.unwrap_or(50051)
|
||||
}
|
||||
|
||||
/// Result of orchestrator setup, containing all handles needed by the agent.
|
||||
pub struct OrchestratorSetup {
|
||||
pub container_job_manager: Option<Arc<ContainerJobManager>>,
|
||||
@@ -101,11 +110,12 @@ pub async fn setup_orchestrator(
|
||||
let job_event_tx = Some(tx);
|
||||
|
||||
let token_store = TokenStore::new();
|
||||
let orchestrator_port = resolve_orchestrator_port();
|
||||
let job_config = ContainerJobConfig {
|
||||
image: config.sandbox.image.clone(),
|
||||
memory_limit_mb: config.sandbox.memory_limit_mb,
|
||||
cpu_shares: config.sandbox.cpu_shares,
|
||||
orchestrator_port: 50051,
|
||||
orchestrator_port,
|
||||
claude_code_api_key: std::env::var("ANTHROPIC_API_KEY").ok(),
|
||||
claude_code_oauth_token: crate::config::ClaudeCodeConfig::extract_oauth_token(),
|
||||
claude_code_model: config.claude_code.model.clone(),
|
||||
@@ -127,7 +137,7 @@ pub async fn setup_orchestrator(
|
||||
};
|
||||
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = OrchestratorApi::start(orchestrator_state, 50051).await {
|
||||
if let Err(e) = OrchestratorApi::start(orchestrator_state, orchestrator_port).await {
|
||||
tracing::error!("Orchestrator API failed: {}", e);
|
||||
}
|
||||
});
|
||||
@@ -151,3 +161,40 @@ pub async fn setup_orchestrator(
|
||||
docker_status,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Mutex;
|
||||
|
||||
use super::*;
|
||||
|
||||
/// Serialize access to `ORCHESTRATOR_PORT` env var across test threads.
|
||||
static ENV_LOCK: Mutex<()> = Mutex::new(());
|
||||
|
||||
#[test]
|
||||
fn resolve_orchestrator_port_from_env() {
|
||||
let _guard = ENV_LOCK.lock().unwrap();
|
||||
|
||||
// Safety: env-var mutation requires unsafe in edition 2024;
|
||||
// ENV_LOCK serializes concurrent access from other test threads.
|
||||
|
||||
// Absent env var → default 50051
|
||||
unsafe { std::env::remove_var("ORCHESTRATOR_PORT") };
|
||||
assert_eq!(resolve_orchestrator_port(), 50051);
|
||||
|
||||
// Valid custom port
|
||||
unsafe { std::env::set_var("ORCHESTRATOR_PORT", "50052") };
|
||||
assert_eq!(resolve_orchestrator_port(), 50052);
|
||||
|
||||
// Non-numeric value → fallback to default
|
||||
unsafe { std::env::set_var("ORCHESTRATOR_PORT", "not_a_port") };
|
||||
assert_eq!(resolve_orchestrator_port(), 50051);
|
||||
|
||||
// Out of u16 range → fallback to default
|
||||
unsafe { std::env::set_var("ORCHESTRATOR_PORT", "99999") };
|
||||
assert_eq!(resolve_orchestrator_port(), 50051);
|
||||
|
||||
// Cleanup
|
||||
unsafe { std::env::remove_var("ORCHESTRATOR_PORT") };
|
||||
}
|
||||
}
|
||||
|
||||
@@ -440,6 +440,39 @@ impl PairingStore {
|
||||
Ok(file.allow_from)
|
||||
}
|
||||
|
||||
/// Clear the allow-from list for a channel.
|
||||
///
|
||||
/// Called on credential refresh so that existing users must re-approve
|
||||
/// after a bot token change.
|
||||
pub fn clear_allow_from(&self, channel: &str) -> Result<(), PairingStoreError> {
|
||||
let path = allow_from_path(&self.base_dir, channel)?;
|
||||
if let Some(parent) = path.parent() {
|
||||
fs::create_dir_all(parent)?;
|
||||
}
|
||||
let file = fs::OpenOptions::new()
|
||||
.read(true)
|
||||
.write(true)
|
||||
.create(true)
|
||||
.truncate(true)
|
||||
.open(&path)?;
|
||||
file.lock_exclusive()?;
|
||||
let store = AllowFromStoreFile {
|
||||
version: 1,
|
||||
allow_from: Vec::new(),
|
||||
};
|
||||
let json = serde_json::to_string_pretty(&store)?;
|
||||
fs::write(&path, json)?;
|
||||
fs4::FileExt::unlock(&file)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Clear all pending pairing requests for a channel.
|
||||
///
|
||||
/// Called on credential refresh so stale requests don't confuse users.
|
||||
pub fn clear_pending(&self, channel: &str) -> Result<(), PairingStoreError> {
|
||||
self.write_pairing_file(channel, &[])
|
||||
}
|
||||
|
||||
/// Check if a sender is allowed (by id or username).
|
||||
pub fn is_sender_allowed(
|
||||
&self,
|
||||
@@ -517,6 +550,11 @@ impl PairingStore {
|
||||
requests: &[PairingRequest],
|
||||
) -> Result<(), PairingStoreError> {
|
||||
let path = pairing_path(&self.base_dir, channel)?;
|
||||
let parent = path.parent().ok_or_else(|| {
|
||||
PairingStoreError::InvalidPath(format!("path has no parent: {}", path.display()))
|
||||
})?;
|
||||
fs::create_dir_all(parent)?;
|
||||
|
||||
let mut file = fs::OpenOptions::new()
|
||||
.write(true)
|
||||
.create(true)
|
||||
@@ -717,4 +755,117 @@ mod tests {
|
||||
store.list_pending("").unwrap_err();
|
||||
store.upsert_request("", "u1", None).unwrap_err();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_clear_allow_from_removes_all_entries() {
|
||||
let (store, _) = test_store();
|
||||
let r1 = store.upsert_request("telegram", "user1", None).unwrap();
|
||||
store.approve("telegram", &r1.code).unwrap();
|
||||
|
||||
let list = store.read_allow_from("telegram").unwrap();
|
||||
assert_eq!(list.len(), 1);
|
||||
|
||||
store.clear_allow_from("telegram").unwrap();
|
||||
let list = store.read_allow_from("telegram").unwrap();
|
||||
assert!(list.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_clear_pending_removes_all_requests() {
|
||||
let (store, _) = test_store();
|
||||
store
|
||||
.upsert_request("telegram", "user1", Some(serde_json::json!({"chat_id": 1})))
|
||||
.unwrap();
|
||||
store
|
||||
.upsert_request("telegram", "user2", Some(serde_json::json!({"chat_id": 2})))
|
||||
.unwrap();
|
||||
|
||||
let requests = store.list_pending("telegram").unwrap();
|
||||
assert_eq!(requests.len(), 2);
|
||||
|
||||
store.clear_pending("telegram").unwrap();
|
||||
let requests = store.list_pending("telegram").unwrap();
|
||||
assert!(requests.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_clear_allow_from_allows_new_approval() {
|
||||
let (store, _) = test_store();
|
||||
let r1 = store.upsert_request("telegram", "user1", None).unwrap();
|
||||
store.approve("telegram", &r1.code).unwrap();
|
||||
|
||||
assert!(store.is_sender_allowed("telegram", "user1", None).unwrap());
|
||||
|
||||
store.clear_allow_from("telegram").unwrap();
|
||||
|
||||
assert!(!store.is_sender_allowed("telegram", "user1", None).unwrap());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_clear_allow_from_on_nonexistent_file() {
|
||||
let (store, _) = test_store();
|
||||
// No requests created, so allow_from file doesn't exist
|
||||
let result = store.clear_allow_from("telegram");
|
||||
assert!(result.is_ok());
|
||||
|
||||
// After clearing, should return empty list
|
||||
let list = store.read_allow_from("telegram").unwrap();
|
||||
assert!(list.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_clear_pending_on_nonexistent_file() {
|
||||
let (store, _) = test_store();
|
||||
// No requests created, so pairing file doesn't exist
|
||||
let result = store.clear_pending("telegram");
|
||||
assert!(result.is_ok());
|
||||
|
||||
// After clearing, should return empty list
|
||||
let requests = store.list_pending("telegram").unwrap();
|
||||
assert!(requests.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_clear_and_reapprove_workflow() {
|
||||
let (store, _) = test_store();
|
||||
|
||||
// Step 1: Create and approve user1
|
||||
let r1 = store.upsert_request("telegram", "user1", None).unwrap();
|
||||
store.approve("telegram", &r1.code).unwrap();
|
||||
assert!(store.is_sender_allowed("telegram", "user1", None).unwrap());
|
||||
|
||||
// Step 2: Simulate credential refresh by clearing pairing state
|
||||
store.clear_allow_from("telegram").unwrap();
|
||||
store.clear_pending("telegram").unwrap();
|
||||
|
||||
// Step 3: Verify user1 is no longer approved and no pending requests exist
|
||||
assert!(!store.is_sender_allowed("telegram", "user1", None).unwrap());
|
||||
let requests = store.list_pending("telegram").unwrap();
|
||||
assert!(requests.is_empty());
|
||||
|
||||
// Step 4: Create new pairing request and approve user1 again
|
||||
let r2 = store.upsert_request("telegram", "user1", None).unwrap();
|
||||
assert!(r2.created); // Should be a new request
|
||||
store.approve("telegram", &r2.code).unwrap();
|
||||
assert!(store.is_sender_allowed("telegram", "user1", None).unwrap());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_clear_one_channel_doesnt_affect_other() {
|
||||
let (store, _) = test_store();
|
||||
|
||||
// Approve users on two channels
|
||||
let r1 = store.upsert_request("telegram", "user1", None).unwrap();
|
||||
store.approve("telegram", &r1.code).unwrap();
|
||||
|
||||
let r2 = store.upsert_request("discord", "user2", None).unwrap();
|
||||
store.approve("discord", &r2.code).unwrap();
|
||||
|
||||
// Clear only telegram
|
||||
store.clear_allow_from("telegram").unwrap();
|
||||
|
||||
// Verify telegram is cleared but discord is not
|
||||
assert!(!store.is_sender_allowed("telegram", "user1", None).unwrap());
|
||||
assert!(store.is_sender_allowed("discord", "user2", None).unwrap());
|
||||
}
|
||||
}
|
||||
|
||||
+100
-3
@@ -236,14 +236,59 @@ impl SandboxManager {
|
||||
self.initialize().await?;
|
||||
}
|
||||
|
||||
// Get proxy port if running
|
||||
// Retry transient container failures (Docker daemon glitches, container
|
||||
// creation races) up to MAX_SANDBOX_RETRIES times with exponential backoff.
|
||||
const MAX_SANDBOX_RETRIES: u32 = 2;
|
||||
let mut last_err: Option<SandboxError> = None;
|
||||
|
||||
for attempt in 0..=MAX_SANDBOX_RETRIES {
|
||||
if attempt > 0 {
|
||||
let delay = std::time::Duration::from_secs(1 << attempt); // 2s, 4s
|
||||
tracing::warn!(
|
||||
attempt = attempt + 1,
|
||||
max_attempts = MAX_SANDBOX_RETRIES + 1,
|
||||
delay_secs = delay.as_secs(),
|
||||
"Retrying sandbox execution after transient failure"
|
||||
);
|
||||
tokio::time::sleep(delay).await;
|
||||
}
|
||||
|
||||
match self
|
||||
.try_execute_in_container(command, cwd, policy, env.clone())
|
||||
.await
|
||||
{
|
||||
Ok(output) => return Ok(output),
|
||||
Err(e) if is_transient_sandbox_error(&e) => {
|
||||
tracing::warn!(
|
||||
attempt = attempt + 1,
|
||||
error = %e,
|
||||
"Transient sandbox error, will retry"
|
||||
);
|
||||
last_err = Some(e);
|
||||
}
|
||||
Err(e) => return Err(e),
|
||||
}
|
||||
}
|
||||
|
||||
Err(last_err.unwrap_or_else(|| SandboxError::ExecutionFailed {
|
||||
reason: "all retry attempts exhausted".to_string(),
|
||||
}))
|
||||
}
|
||||
|
||||
/// Single attempt at container execution (no retry logic).
|
||||
async fn try_execute_in_container(
|
||||
&self,
|
||||
command: &str,
|
||||
cwd: &Path,
|
||||
policy: SandboxPolicy,
|
||||
env: HashMap<String, String>,
|
||||
) -> Result<ExecOutput> {
|
||||
let proxy_port = if let Some(proxy) = self.proxy.read().await.as_ref() {
|
||||
proxy.addr().await.map(|a| a.port()).unwrap_or(0)
|
||||
} else {
|
||||
0
|
||||
};
|
||||
|
||||
// Reuse the stored Docker connection, create a runner with the current proxy port
|
||||
let docker =
|
||||
self.docker
|
||||
.read()
|
||||
@@ -262,7 +307,6 @@ impl SandboxManager {
|
||||
};
|
||||
|
||||
let container_output = runner.execute(command, cwd, policy, &limits, env).await?;
|
||||
|
||||
Ok(container_output.into())
|
||||
}
|
||||
|
||||
@@ -373,6 +417,20 @@ impl Drop for SandboxManager {
|
||||
}
|
||||
}
|
||||
|
||||
/// Check whether a sandbox error is transient and worth retrying.
|
||||
///
|
||||
/// Transient errors are those caused by Docker daemon glitches, container
|
||||
/// creation race conditions, or container start failures — not by command
|
||||
/// execution failures, timeouts, or policy violations.
|
||||
fn is_transient_sandbox_error(err: &SandboxError) -> bool {
|
||||
matches!(
|
||||
err,
|
||||
SandboxError::DockerNotAvailable { .. }
|
||||
| SandboxError::ContainerCreationFailed { .. }
|
||||
| SandboxError::ContainerStartFailed { .. }
|
||||
)
|
||||
}
|
||||
|
||||
/// Builder for creating a sandbox manager.
|
||||
pub struct SandboxManagerBuilder {
|
||||
config: SandboxConfig,
|
||||
@@ -597,4 +655,43 @@ mod tests {
|
||||
assert!(output.truncated);
|
||||
assert!(output.stdout.len() <= 32 * 1024);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transient_errors_are_retryable() {
|
||||
assert!(super::is_transient_sandbox_error(
|
||||
&SandboxError::DockerNotAvailable {
|
||||
reason: "daemon restarting".to_string()
|
||||
}
|
||||
));
|
||||
assert!(super::is_transient_sandbox_error(
|
||||
&SandboxError::ContainerCreationFailed {
|
||||
reason: "image pull glitch".to_string()
|
||||
}
|
||||
));
|
||||
assert!(super::is_transient_sandbox_error(
|
||||
&SandboxError::ContainerStartFailed {
|
||||
reason: "cgroup race".to_string()
|
||||
}
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn non_transient_errors_are_not_retryable() {
|
||||
assert!(!super::is_transient_sandbox_error(&SandboxError::Timeout(
|
||||
std::time::Duration::from_secs(30)
|
||||
)));
|
||||
assert!(!super::is_transient_sandbox_error(
|
||||
&SandboxError::ExecutionFailed {
|
||||
reason: "exit code 1".to_string()
|
||||
}
|
||||
));
|
||||
assert!(!super::is_transient_sandbox_error(
|
||||
&SandboxError::NetworkBlocked {
|
||||
reason: "policy violation".to_string()
|
||||
}
|
||||
));
|
||||
assert!(!super::is_transient_sandbox_error(&SandboxError::Config {
|
||||
reason: "bad config".to_string()
|
||||
}));
|
||||
}
|
||||
}
|
||||
|
||||
+19
-1
@@ -16,6 +16,14 @@ pub struct Settings {
|
||||
#[serde(default, alias = "setup_completed")]
|
||||
pub onboard_completed: bool,
|
||||
|
||||
/// Stable owner scope for this IronClaw instance.
|
||||
///
|
||||
/// This is bootstrap configuration loaded from env / disk / TOML. We do
|
||||
/// not persist it in the per-user DB settings table because the DB lookup
|
||||
/// itself already requires the owner scope to be known.
|
||||
#[serde(default)]
|
||||
pub owner_id: Option<String>,
|
||||
|
||||
// === Step 1: Database ===
|
||||
/// Database backend: "postgres" or "libsql".
|
||||
#[serde(default)]
|
||||
@@ -360,6 +368,10 @@ pub struct HeartbeatSettings {
|
||||
#[serde(default)]
|
||||
pub notify_user: Option<String>,
|
||||
|
||||
/// Fixed time-of-day to fire (HH:MM, 24h). When set, interval_secs is ignored.
|
||||
#[serde(default)]
|
||||
pub fire_at: Option<String>,
|
||||
|
||||
/// Hour (0-23) when quiet hours start (heartbeat skipped).
|
||||
#[serde(default)]
|
||||
pub quiet_hours_start: Option<u32>,
|
||||
@@ -368,7 +380,7 @@ pub struct HeartbeatSettings {
|
||||
#[serde(default)]
|
||||
pub quiet_hours_end: Option<u32>,
|
||||
|
||||
/// Timezone for quiet hours evaluation (IANA name, e.g. "America/New_York").
|
||||
/// Timezone for fire_at and quiet hours (IANA name, e.g. "Pacific/Auckland").
|
||||
#[serde(default)]
|
||||
pub timezone: Option<String>,
|
||||
}
|
||||
@@ -384,6 +396,7 @@ impl Default for HeartbeatSettings {
|
||||
interval_secs: default_heartbeat_interval(),
|
||||
notify_channel: None,
|
||||
notify_user: None,
|
||||
fire_at: None,
|
||||
quiet_hours_start: None,
|
||||
quiet_hours_end: None,
|
||||
timezone: None,
|
||||
@@ -728,6 +741,10 @@ impl Settings {
|
||||
let mut settings = Self::default();
|
||||
|
||||
for (key, value) in map {
|
||||
if key == "owner_id" {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Convert the JSONB value to a string for the existing set() method
|
||||
let value_str = match value {
|
||||
serde_json::Value::String(s) => s.clone(),
|
||||
@@ -767,6 +784,7 @@ impl Settings {
|
||||
|
||||
let mut map = std::collections::HashMap::new();
|
||||
collect_settings_json(&json, String::new(), &mut map);
|
||||
map.remove("owner_id");
|
||||
map
|
||||
}
|
||||
|
||||
|
||||
+708
-268
File diff suppressed because it is too large
Load Diff
+3
-2
@@ -439,6 +439,7 @@ impl TestHarnessBuilder {
|
||||
};
|
||||
|
||||
let deps = AgentDeps {
|
||||
owner_id: "default".to_string(),
|
||||
store: Some(Arc::clone(&db)),
|
||||
llm,
|
||||
cheap_llm: None,
|
||||
@@ -1077,7 +1078,7 @@ mod tests {
|
||||
},
|
||||
notify: NotifyConfig {
|
||||
channel: None,
|
||||
user: "user1".to_string(),
|
||||
user: Some("user1".to_string()),
|
||||
on_attention: true,
|
||||
on_failure: true,
|
||||
on_success: false,
|
||||
@@ -1210,7 +1211,7 @@ mod tests {
|
||||
},
|
||||
notify: NotifyConfig {
|
||||
channel: None,
|
||||
user: "user1".to_string(),
|
||||
user: Some("user1".to_string()),
|
||||
on_attention: false,
|
||||
on_failure: false,
|
||||
on_success: false,
|
||||
|
||||
+125
-26
@@ -10,6 +10,7 @@ use async_trait::async_trait;
|
||||
use crate::bootstrap::ironclaw_base_dir;
|
||||
use crate::channels::{ChannelManager, OutgoingResponse};
|
||||
use crate::context::JobContext;
|
||||
use crate::extensions::ExtensionManager;
|
||||
use crate::tools::tool::{
|
||||
ApprovalRequirement, Tool, ToolError, ToolOutput, ToolRateLimitConfig, require_str,
|
||||
};
|
||||
@@ -17,6 +18,7 @@ use crate::tools::tool::{
|
||||
/// Tool for sending messages to channels.
|
||||
pub struct MessageTool {
|
||||
channel_manager: Arc<ChannelManager>,
|
||||
extension_manager: Option<Arc<ExtensionManager>>,
|
||||
/// Default channel for current conversation (set per-turn).
|
||||
/// Uses std::sync::RwLock because requires_approval() is sync and called from async context.
|
||||
default_channel: Arc<RwLock<Option<String>>>,
|
||||
@@ -32,12 +34,18 @@ impl MessageTool {
|
||||
|
||||
Self {
|
||||
channel_manager,
|
||||
extension_manager: None,
|
||||
default_channel: Arc::new(RwLock::new(None)),
|
||||
default_target: Arc::new(RwLock::new(None)),
|
||||
base_dir,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_extension_manager(mut self, extension_manager: Arc<ExtensionManager>) -> Self {
|
||||
self.extension_manager = Some(extension_manager);
|
||||
self
|
||||
}
|
||||
|
||||
/// Set the base directory for attachment validation.
|
||||
/// This is primarily used for testing or future configuration.
|
||||
pub fn with_base_dir(mut self, dir: PathBuf) -> Self {
|
||||
@@ -111,39 +119,76 @@ impl Tool for MessageTool {
|
||||
|
||||
let content = require_str(¶ms, "content")?;
|
||||
|
||||
let explicit_channel = params
|
||||
.get("channel")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|value| value.to_string());
|
||||
let default_channel = self
|
||||
.default_channel
|
||||
.read()
|
||||
.unwrap_or_else(|e| e.into_inner())
|
||||
.clone();
|
||||
let metadata_channel = ctx
|
||||
.metadata
|
||||
.get("notify_channel")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|value| value.to_string());
|
||||
|
||||
// Get channel: use param → conversation default → job metadata → None (broadcast all)
|
||||
let channel: Option<String> =
|
||||
if let Some(c) = params.get("channel").and_then(|v| v.as_str()) {
|
||||
Some(c.to_string())
|
||||
} else if let Some(c) = self
|
||||
.default_channel
|
||||
let channel: Option<String> = explicit_channel
|
||||
.clone()
|
||||
.or_else(|| default_channel.clone())
|
||||
.or_else(|| metadata_channel.clone());
|
||||
|
||||
let can_use_default_target = match (explicit_channel.as_deref(), default_channel.as_deref())
|
||||
{
|
||||
(None, _) => true,
|
||||
(Some(explicit), Some(current)) if explicit == current => true,
|
||||
_ => false,
|
||||
};
|
||||
let can_use_metadata_target = match (channel.as_deref(), metadata_channel.as_deref()) {
|
||||
(None, _) => true,
|
||||
(Some(resolved), Some(current)) if resolved == current => true,
|
||||
_ => false,
|
||||
};
|
||||
|
||||
// Get target: use param → conversation default → job metadata → owner scope
|
||||
// fallback when a specific channel is known.
|
||||
let target = if let Some(t) = params.get("target").and_then(|v| v.as_str()) {
|
||||
Some(t.to_string())
|
||||
} else if can_use_default_target
|
||||
&& let Some(t) = self
|
||||
.default_target
|
||||
.read()
|
||||
.unwrap_or_else(|e| e.into_inner())
|
||||
.clone()
|
||||
{
|
||||
Some(c)
|
||||
} else {
|
||||
ctx.metadata
|
||||
.get("notify_channel")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|c| c.to_string())
|
||||
};
|
||||
|
||||
// Get target: use param → conversation default → job metadata
|
||||
let target = if let Some(t) = params.get("target").and_then(|v| v.as_str()) {
|
||||
t.to_string()
|
||||
} else if let Some(t) = self
|
||||
.default_target
|
||||
.read()
|
||||
.unwrap_or_else(|e| e.into_inner())
|
||||
.clone()
|
||||
{
|
||||
t
|
||||
} else if let Some(t) = ctx.metadata.get("notify_user").and_then(|v| v.as_str()) {
|
||||
t.to_string()
|
||||
Some(t)
|
||||
} else if can_use_metadata_target
|
||||
&& let Some(t) = ctx.metadata.get("notify_user").and_then(|v| v.as_str())
|
||||
{
|
||||
Some(t.to_string())
|
||||
} else if channel.is_some() {
|
||||
if let Some(channel_name) = channel.as_deref() {
|
||||
if let Some(extension_manager) = self.extension_manager.as_ref()
|
||||
&& let Some(target) = extension_manager
|
||||
.notification_target_for_channel(channel_name)
|
||||
.await
|
||||
{
|
||||
Some(target)
|
||||
} else {
|
||||
Some(ctx.user_id.clone())
|
||||
}
|
||||
} else {
|
||||
Some(ctx.user_id.clone())
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let Some(target) = target else {
|
||||
return Err(ToolError::ExecutionFailed(
|
||||
"No target specified and no active conversation. Provide target parameter."
|
||||
"No target specified and no channel-scoped routing target could be resolved. Provide target parameter."
|
||||
.to_string(),
|
||||
));
|
||||
};
|
||||
@@ -659,6 +704,31 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn message_tool_falls_back_to_ctx_user_when_channel_known() {
|
||||
// Regression for owner-scoped notifications: a channel can be known
|
||||
// even when the concrete delivery target is omitted, so the message
|
||||
// tool should pass ctx.user_id through to the channel layer.
|
||||
let tool = MessageTool::new(Arc::new(ChannelManager::new()));
|
||||
|
||||
let mut ctx =
|
||||
crate::context::JobContext::with_user("owner-scope", "routine-job", "price alert");
|
||||
ctx.metadata = serde_json::json!({
|
||||
"notify_channel": "telegram",
|
||||
});
|
||||
|
||||
let result = tool
|
||||
.execute(serde_json::json!({"content": "NEAR price is $5"}), &ctx)
|
||||
.await;
|
||||
|
||||
assert!(result.is_err()); // safety: test-only assertion
|
||||
let err = result.unwrap_err().to_string();
|
||||
let mentions_missing_target = err.contains("No target specified");
|
||||
assert!(!mentions_missing_target); // safety: test-only assertion
|
||||
let mentions_missing_channel = err.contains("No channel specified");
|
||||
assert!(!mentions_missing_channel); // safety: test-only assertion
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn message_tool_no_metadata_still_errors() {
|
||||
// When neither conversation context nor metadata is set, should still
|
||||
@@ -710,4 +780,33 @@ mod tests {
|
||||
err
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn message_tool_does_not_apply_metadata_target_to_different_default_channel() {
|
||||
let tool = MessageTool::new(Arc::new(ChannelManager::new()));
|
||||
tool.set_context(Some("telegram".to_string()), None).await;
|
||||
|
||||
let mut ctx = crate::context::JobContext::with_user("owner-scope", "test", "test");
|
||||
ctx.metadata = serde_json::json!({
|
||||
"notify_channel": "signal",
|
||||
"notify_user": "metadata-user",
|
||||
});
|
||||
|
||||
let result = tool
|
||||
.execute(serde_json::json!({"content": "hello"}), &ctx)
|
||||
.await;
|
||||
|
||||
assert!(result.is_err());
|
||||
let err = result.unwrap_err().to_string();
|
||||
assert!(
|
||||
!err.contains("metadata-user"),
|
||||
"metadata target should not be applied to a different default channel: {}",
|
||||
err
|
||||
);
|
||||
assert!(
|
||||
err.contains("owner-scope"),
|
||||
"expected owner-scope fallback target when metadata channel differs: {}",
|
||||
err
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -106,7 +106,7 @@ pub(crate) fn routine_create_parameters_schema() -> serde_json::Value {
|
||||
},
|
||||
"notify_user": {
|
||||
"type": "string",
|
||||
"description": "User or destination to notify, for example a username or chat ID."
|
||||
"description": "Optional explicit user or destination to notify, for example a username or chat ID. Omit it to use the configured owner's last-seen target for that channel."
|
||||
},
|
||||
"timezone": {
|
||||
"type": "string",
|
||||
@@ -387,8 +387,7 @@ impl Tool for RoutineCreateTool {
|
||||
user: params
|
||||
.get("notify_user")
|
||||
.and_then(|v| v.as_str())
|
||||
.unwrap_or("default")
|
||||
.to_string(),
|
||||
.map(String::from),
|
||||
..NotifyConfig::default()
|
||||
},
|
||||
last_run_at: None,
|
||||
|
||||
@@ -501,9 +501,14 @@ impl ToolRegistry {
|
||||
pub async fn register_message_tools(
|
||||
&self,
|
||||
channel_manager: Arc<crate::channels::ChannelManager>,
|
||||
extension_manager: Option<Arc<crate::extensions::ExtensionManager>>,
|
||||
) {
|
||||
use crate::tools::builtin::MessageTool;
|
||||
let tool = Arc::new(MessageTool::new(channel_manager));
|
||||
let mut tool = MessageTool::new(channel_manager);
|
||||
if let Some(extension_manager) = extension_manager {
|
||||
tool = tool.with_extension_manager(extension_manager);
|
||||
}
|
||||
let tool = Arc::new(tool);
|
||||
*self.message_tool.write().await = Some(Arc::clone(&tool));
|
||||
self.tools
|
||||
.write()
|
||||
|
||||
+188
-8
@@ -841,13 +841,7 @@ impl Tool for WasmToolWrapper {
|
||||
// Pre-resolve host credentials from secrets store (async, before blocking task).
|
||||
// This decrypts the secrets once so the sync http_request() host function
|
||||
// can inject them without needing async access.
|
||||
//
|
||||
// BUG FIX: ExtensionManager stores OAuth tokens under user_id "default"
|
||||
// (hardcoded at construction in app.rs), but this was previously looking
|
||||
// them up under ctx.user_id — which could be a Telegram user ID, web
|
||||
// gateway user, etc. — causing credential resolution to silently fail.
|
||||
// Must match the storage key until per-user credential isolation is added.
|
||||
let credential_user_id = "default";
|
||||
let credential_user_id = &ctx.user_id;
|
||||
let host_credentials = resolve_host_credentials(
|
||||
&self.capabilities,
|
||||
self.secrets_store.as_deref(),
|
||||
@@ -1165,6 +1159,13 @@ async fn resolve_host_credentials(
|
||||
let secret = match store.get_decrypted(user_id, &mapping.secret_name).await {
|
||||
Ok(s) => Some(s),
|
||||
Err(e) => {
|
||||
tracing::trace!(
|
||||
user_id = %user_id,
|
||||
secret_name = %mapping.secret_name,
|
||||
error = %e,
|
||||
"No matching host credential resolved for WASM tool in the requested scope"
|
||||
);
|
||||
|
||||
// If lookup fails and we're not already looking up "default", try "default" as fallback
|
||||
if user_id != "default" {
|
||||
tracing::debug!(
|
||||
@@ -1385,7 +1386,16 @@ fn build_tool_usage_hint(tool_name: &str, schema: &serde_json::Value) -> String
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::context::JobContext;
|
||||
use crate::secrets::{
|
||||
CreateSecretParams, DecryptedSecret, InMemorySecretsStore, Secret, SecretError, SecretRef,
|
||||
SecretsStore,
|
||||
};
|
||||
|
||||
use crate::testing::credentials::{
|
||||
TEST_BEARER_TOKEN_123, TEST_GOOGLE_OAUTH_FRESH, TEST_GOOGLE_OAUTH_LEGACY,
|
||||
@@ -1396,6 +1406,78 @@ mod tests {
|
||||
use crate::tools::wasm::capabilities::Capabilities;
|
||||
use crate::tools::wasm::runtime::{WasmRuntimeConfig, WasmToolRuntime};
|
||||
|
||||
struct RecordingSecretsStore {
|
||||
inner: InMemorySecretsStore,
|
||||
get_decrypted_lookups: Mutex<Vec<(String, String)>>,
|
||||
}
|
||||
|
||||
impl RecordingSecretsStore {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
inner: test_secrets_store(),
|
||||
get_decrypted_lookups: Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
fn decrypted_lookups(&self) -> Vec<(String, String)> {
|
||||
self.get_decrypted_lookups.lock().unwrap().clone()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl SecretsStore for RecordingSecretsStore {
|
||||
async fn create(
|
||||
&self,
|
||||
user_id: &str,
|
||||
params: CreateSecretParams,
|
||||
) -> Result<Secret, SecretError> {
|
||||
self.inner.create(user_id, params).await
|
||||
}
|
||||
|
||||
async fn get(&self, user_id: &str, name: &str) -> Result<Secret, SecretError> {
|
||||
self.inner.get(user_id, name).await
|
||||
}
|
||||
|
||||
async fn get_decrypted(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<DecryptedSecret, SecretError> {
|
||||
self.get_decrypted_lookups
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push((user_id.to_string(), name.to_string()));
|
||||
self.inner.get_decrypted(user_id, name).await
|
||||
}
|
||||
|
||||
async fn exists(&self, user_id: &str, name: &str) -> Result<bool, SecretError> {
|
||||
self.inner.exists(user_id, name).await
|
||||
}
|
||||
|
||||
async fn list(&self, user_id: &str) -> Result<Vec<SecretRef>, SecretError> {
|
||||
self.inner.list(user_id).await
|
||||
}
|
||||
|
||||
async fn delete(&self, user_id: &str, name: &str) -> Result<bool, SecretError> {
|
||||
self.inner.delete(user_id, name).await
|
||||
}
|
||||
|
||||
async fn record_usage(&self, secret_id: Uuid) -> Result<(), SecretError> {
|
||||
self.inner.record_usage(secret_id).await
|
||||
}
|
||||
|
||||
async fn is_accessible(
|
||||
&self,
|
||||
user_id: &str,
|
||||
secret_name: &str,
|
||||
allowed_secrets: &[String],
|
||||
) -> Result<bool, SecretError> {
|
||||
self.inner
|
||||
.is_accessible(user_id, secret_name, allowed_secrets)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_wrapper_creation() {
|
||||
// This test verifies the runtime can be created
|
||||
@@ -1691,6 +1773,104 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_resolve_host_credentials_owner_scope_bearer() {
|
||||
use std::collections::HashMap;
|
||||
|
||||
use crate::secrets::{
|
||||
CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore,
|
||||
};
|
||||
use crate::tools::wasm::capabilities::HttpCapability;
|
||||
use crate::tools::wasm::wrapper::resolve_host_credentials;
|
||||
|
||||
let store = test_secrets_store();
|
||||
let ctx = JobContext::with_user("owner-scope", "owner-scope test", "owner-scope test");
|
||||
|
||||
store
|
||||
.create(
|
||||
&ctx.user_id,
|
||||
CreateSecretParams::new("google_oauth_token", TEST_GOOGLE_OAUTH_TOKEN),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut credentials = HashMap::new();
|
||||
credentials.insert(
|
||||
"google_oauth_token".to_string(),
|
||||
CredentialMapping {
|
||||
secret_name: "google_oauth_token".to_string(),
|
||||
location: CredentialLocation::AuthorizationBearer,
|
||||
host_patterns: vec!["www.googleapis.com".to_string()],
|
||||
},
|
||||
);
|
||||
|
||||
let caps = Capabilities {
|
||||
http: Some(HttpCapability {
|
||||
credentials,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let result = resolve_host_credentials(&caps, Some(&store), &ctx.user_id, None).await;
|
||||
assert_eq!(result.len(), 1);
|
||||
assert_eq!(
|
||||
result[0].headers.get("Authorization"),
|
||||
Some(&format!("Bearer {TEST_GOOGLE_OAUTH_TOKEN}"))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_execute_resolves_host_credentials_from_owner_scope_context() {
|
||||
use std::collections::HashMap;
|
||||
|
||||
use crate::secrets::{CredentialLocation, CredentialMapping};
|
||||
use crate::tools::wasm::capabilities::HttpCapability;
|
||||
|
||||
let runtime = Arc::new(WasmToolRuntime::new(WasmRuntimeConfig::for_testing()).unwrap());
|
||||
let prepared = runtime
|
||||
.prepare("search", b"\0asm\x0d\0\x01\0", None)
|
||||
.await
|
||||
.unwrap();
|
||||
let store = Arc::new(RecordingSecretsStore::new());
|
||||
let ctx = JobContext::with_user("owner-scope", "owner-scope test", "owner-scope test");
|
||||
|
||||
store
|
||||
.create(
|
||||
&ctx.user_id,
|
||||
CreateSecretParams::new("google_oauth_token", TEST_GOOGLE_OAUTH_TOKEN),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut credentials = HashMap::new();
|
||||
credentials.insert(
|
||||
"google_oauth_token".to_string(),
|
||||
CredentialMapping {
|
||||
secret_name: "google_oauth_token".to_string(),
|
||||
location: CredentialLocation::AuthorizationBearer,
|
||||
host_patterns: vec!["www.googleapis.com".to_string()],
|
||||
},
|
||||
);
|
||||
|
||||
let caps = Capabilities {
|
||||
http: Some(HttpCapability {
|
||||
credentials,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let wrapper = super::WasmToolWrapper::new(Arc::clone(&runtime), prepared, caps)
|
||||
.with_secrets_store(store.clone());
|
||||
let result = wrapper.execute(serde_json::json!({}), &ctx).await;
|
||||
assert!(result.is_err());
|
||||
|
||||
let lookups = store.decrypted_lookups();
|
||||
assert!(lookups.contains(&("owner-scope".to_string(), "google_oauth_token".to_string())));
|
||||
assert!(!lookups.contains(&("default".to_string(), "google_oauth_token".to_string())));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_resolve_host_credentials_missing_secret() {
|
||||
use std::collections::HashMap;
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
//! Chat Completions-based transcription provider.
|
||||
//!
|
||||
//! Uses the `/v1/chat/completions` endpoint with `input_audio` content type
|
||||
//! to transcribe audio. Compatible with OpenRouter, OpenAI GPT-4o-audio, and
|
||||
//! any provider that supports audio input via the Chat Completions API.
|
||||
|
||||
use async_trait::async_trait;
|
||||
use base64::Engine;
|
||||
use secrecy::{ExposeSecret, SecretString};
|
||||
|
||||
use super::{AudioFormat, TranscriptionError, TranscriptionProvider};
|
||||
|
||||
/// Transcription provider that sends audio via the Chat Completions API.
|
||||
///
|
||||
/// Unlike the Whisper provider (which uses `/v1/audio/transcriptions` with
|
||||
/// multipart upload), this provider sends base64-encoded audio as an
|
||||
/// `input_audio` content part in a chat message, enabling use with
|
||||
/// OpenRouter and other providers that only expose audio through the
|
||||
/// Chat Completions API.
|
||||
pub struct ChatCompletionsTranscriptionProvider {
|
||||
client: reqwest::Client,
|
||||
api_key: SecretString,
|
||||
model: String,
|
||||
base_url: String,
|
||||
}
|
||||
|
||||
impl ChatCompletionsTranscriptionProvider {
|
||||
/// Create a new provider with the given API key.
|
||||
pub fn new(api_key: SecretString) -> Self {
|
||||
Self {
|
||||
client: match reqwest::Client::builder()
|
||||
.timeout(std::time::Duration::from_secs(120))
|
||||
.build()
|
||||
{
|
||||
Ok(c) => c,
|
||||
Err(e) => {
|
||||
tracing::error!(
|
||||
"Failed to build HTTP client with timeout, falling back to default: {e}"
|
||||
);
|
||||
reqwest::Client::default()
|
||||
}
|
||||
},
|
||||
api_key,
|
||||
model: "google/gemini-2.0-flash-001".to_string(),
|
||||
base_url: "https://openrouter.ai/api".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Override the base URL.
|
||||
pub fn with_base_url(mut self, base_url: impl Into<String>) -> Self {
|
||||
self.base_url = base_url.into().trim_end_matches('/').to_string();
|
||||
self
|
||||
}
|
||||
|
||||
/// Override the model name.
|
||||
pub fn with_model(mut self, model: impl Into<String>) -> Self {
|
||||
self.model = model.into();
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
/// Map [`AudioFormat`] to the format string expected by the Chat Completions API.
|
||||
fn audio_format_str(format: AudioFormat) -> &'static str {
|
||||
match format {
|
||||
AudioFormat::Ogg => "ogg",
|
||||
AudioFormat::Mp3 => "mp3",
|
||||
AudioFormat::Mp4 => "mp4",
|
||||
AudioFormat::Wav => "wav",
|
||||
AudioFormat::Webm => "webm",
|
||||
AudioFormat::Flac => "flac",
|
||||
AudioFormat::M4a => "m4a",
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl TranscriptionProvider for ChatCompletionsTranscriptionProvider {
|
||||
async fn transcribe(
|
||||
&self,
|
||||
audio_data: &[u8],
|
||||
format: AudioFormat,
|
||||
) -> Result<String, TranscriptionError> {
|
||||
if audio_data.is_empty() {
|
||||
return Err(TranscriptionError::EmptyAudio);
|
||||
}
|
||||
|
||||
let b64 = base64::engine::general_purpose::STANDARD.encode(audio_data);
|
||||
|
||||
let body = serde_json::json!({
|
||||
"model": self.model,
|
||||
"messages": [{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Transcribe this audio. Return only the transcript text, nothing else."
|
||||
},
|
||||
{
|
||||
"type": "input_audio",
|
||||
"input_audio": {
|
||||
"data": b64,
|
||||
"format": audio_format_str(format)
|
||||
}
|
||||
}
|
||||
]
|
||||
}]
|
||||
});
|
||||
|
||||
let url = format!("{}/v1/chat/completions", self.base_url);
|
||||
|
||||
let response = self
|
||||
.client
|
||||
.post(&url)
|
||||
.header(
|
||||
"Authorization",
|
||||
format!("Bearer {}", self.api_key.expose_secret()),
|
||||
)
|
||||
.json(&body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| TranscriptionError::RequestFailed(e.to_string()))?;
|
||||
|
||||
let status = response.status();
|
||||
if !status.is_success() {
|
||||
let body = response
|
||||
.text()
|
||||
.await
|
||||
.unwrap_or_else(|_| "unknown error".to_string());
|
||||
return Err(TranscriptionError::RequestFailed(format!(
|
||||
"HTTP {}: {}",
|
||||
status, body
|
||||
)));
|
||||
}
|
||||
|
||||
let json: serde_json::Value = response
|
||||
.json()
|
||||
.await
|
||||
.map_err(|e| TranscriptionError::RequestFailed(e.to_string()))?;
|
||||
|
||||
// Extract text from the standard Chat Completions response format:
|
||||
// { "choices": [{ "message": { "content": "..." } }] }
|
||||
let text = json
|
||||
.get("choices")
|
||||
.and_then(|c| c.get(0))
|
||||
.and_then(|c| c.get("message"))
|
||||
.and_then(|m| m.get("content"))
|
||||
.and_then(|c| c.as_str())
|
||||
.ok_or_else(|| {
|
||||
TranscriptionError::RequestFailed(
|
||||
"unexpected response format: missing choices[0].message.content".to_string(),
|
||||
)
|
||||
})?;
|
||||
|
||||
Ok(text.trim().to_string())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn audio_format_str_maps_all_variants() {
|
||||
assert_eq!(audio_format_str(AudioFormat::Ogg), "ogg");
|
||||
assert_eq!(audio_format_str(AudioFormat::Mp3), "mp3");
|
||||
assert_eq!(audio_format_str(AudioFormat::Mp4), "mp4");
|
||||
assert_eq!(audio_format_str(AudioFormat::Wav), "wav");
|
||||
assert_eq!(audio_format_str(AudioFormat::Webm), "webm");
|
||||
assert_eq!(audio_format_str(AudioFormat::Flac), "flac");
|
||||
assert_eq!(audio_format_str(AudioFormat::M4a), "m4a");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rejects_empty_audio() {
|
||||
let provider =
|
||||
ChatCompletionsTranscriptionProvider::new(SecretString::from("test-key".to_string()));
|
||||
let result = provider.transcribe(&[], AudioFormat::Ogg).await;
|
||||
assert!(matches!(result, Err(TranscriptionError::EmptyAudio)));
|
||||
}
|
||||
}
|
||||
@@ -4,8 +4,10 @@
|
||||
//! backends and a [`TranscriptionMiddleware`] that detects audio attachments
|
||||
//! on incoming messages and replaces them with transcribed text.
|
||||
|
||||
mod chat_completions;
|
||||
mod openai;
|
||||
|
||||
pub use self::chat_completions::ChatCompletionsTranscriptionProvider;
|
||||
pub use self::openai::OpenAiWhisperProvider;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
+144
-4
@@ -1170,11 +1170,16 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
|
||||
// Reset counter after a successful LLM call
|
||||
self.consecutive_rate_limits
|
||||
.store(0, std::sync::atomic::Ordering::Relaxed);
|
||||
// Preserve the LLM's reasoning text so it appears in the
|
||||
// assistant_with_tool_calls message pushed by execute_tool_calls.
|
||||
let reasoning_text = s
|
||||
.iter()
|
||||
.find_map(|sel| (!sel.reasoning.is_empty()).then_some(sel.reasoning.clone()));
|
||||
let tool_calls: Vec<ToolCall> = selections_to_tool_calls(&s);
|
||||
return Ok(crate::llm::RespondOutput {
|
||||
result: RespondResult::ToolCalls {
|
||||
tool_calls,
|
||||
content: None,
|
||||
content: reasoning_text,
|
||||
},
|
||||
usage: crate::llm::TokenUsage::default(),
|
||||
});
|
||||
@@ -1586,7 +1591,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_mark_completed_twice_returns_error() {
|
||||
async fn test_mark_completed_twice_is_idempotent() {
|
||||
let worker = make_worker(vec![]).await;
|
||||
|
||||
worker
|
||||
@@ -1607,11 +1612,22 @@ mod tests {
|
||||
.unwrap();
|
||||
assert_eq!(ctx.state, JobState::Completed);
|
||||
|
||||
// Second mark_completed should succeed (idempotent) rather than
|
||||
// erroring, matching the fix for the execution_loop / worker wrapper
|
||||
// race condition.
|
||||
let result = worker.mark_completed().await;
|
||||
assert!(
|
||||
result.is_err(),
|
||||
"Completed → Completed transition should be rejected by state machine"
|
||||
result.is_ok(),
|
||||
"Completed -> Completed transition should be idempotent"
|
||||
);
|
||||
|
||||
// State should still be Completed
|
||||
let ctx = worker
|
||||
.context_manager()
|
||||
.get_context(worker.job_id)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(ctx.state, JobState::Completed);
|
||||
}
|
||||
|
||||
/// Build a Worker with the given approval context.
|
||||
@@ -1849,4 +1865,128 @@ mod tests {
|
||||
"Iteration cap should transition to Failed, not Stuck"
|
||||
);
|
||||
}
|
||||
|
||||
/// Regression test: selections_to_tool_calls must preserve tool_call_id
|
||||
/// so that tool_result messages match the assistant_with_tool_calls message
|
||||
/// and are not treated as orphaned by sanitize_tool_messages.
|
||||
#[test]
|
||||
fn test_selections_to_tool_calls_preserves_ids() {
|
||||
let selections = vec![
|
||||
ToolSelection {
|
||||
tool_name: "search".into(),
|
||||
parameters: serde_json::json!({"q": "test"}),
|
||||
reasoning: "Need to search".into(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: "call_abc".into(),
|
||||
},
|
||||
ToolSelection {
|
||||
tool_name: "fetch".into(),
|
||||
parameters: serde_json::json!({"url": "https://example.com"}),
|
||||
reasoning: "Need to fetch".into(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: "call_def".into(),
|
||||
},
|
||||
];
|
||||
|
||||
let tool_calls = selections_to_tool_calls(&selections);
|
||||
|
||||
assert_eq!(tool_calls.len(), 2);
|
||||
assert_eq!(tool_calls[0].id, "call_abc");
|
||||
assert_eq!(tool_calls[0].name, "search");
|
||||
assert_eq!(tool_calls[1].id, "call_def");
|
||||
assert_eq!(tool_calls[1].name, "fetch");
|
||||
}
|
||||
|
||||
/// Regression test: when select_tools returns selections with reasoning,
|
||||
/// the reasoning text should be preserved as content in the RespondResult
|
||||
/// so it appears in the assistant_with_tool_calls message. Without this,
|
||||
/// the LLM's reasoning context is lost and subsequent turns lack context.
|
||||
#[test]
|
||||
fn test_reasoning_text_extraction_from_selections() {
|
||||
// Simulate what call_llm does: extract first non-empty reasoning
|
||||
let selections = [
|
||||
ToolSelection {
|
||||
tool_name: "search".into(),
|
||||
parameters: serde_json::json!({}),
|
||||
reasoning: "I need to search for relevant information".into(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: "call_1".into(),
|
||||
},
|
||||
ToolSelection {
|
||||
tool_name: "fetch".into(),
|
||||
parameters: serde_json::json!({}),
|
||||
reasoning: "I need to search for relevant information".into(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: "call_2".into(),
|
||||
},
|
||||
];
|
||||
|
||||
let reasoning_text = selections
|
||||
.iter()
|
||||
.find_map(|sel| (!sel.reasoning.is_empty()).then_some(sel.reasoning.clone()));
|
||||
|
||||
assert_eq!(
|
||||
reasoning_text.as_deref(),
|
||||
Some("I need to search for relevant information"),
|
||||
"Reasoning text should be extracted from first non-empty selection"
|
||||
);
|
||||
|
||||
// Empty reasoning should result in None
|
||||
let empty_selections = [ToolSelection {
|
||||
tool_name: "echo".into(),
|
||||
parameters: serde_json::json!({}),
|
||||
reasoning: String::new(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: "call_3".into(),
|
||||
}];
|
||||
|
||||
let empty_reasoning = empty_selections
|
||||
.iter()
|
||||
.find_map(|sel| (!sel.reasoning.is_empty()).then_some(sel.reasoning.clone()));
|
||||
|
||||
assert!(
|
||||
empty_reasoning.is_none(),
|
||||
"Empty reasoning should not be included as content"
|
||||
);
|
||||
}
|
||||
|
||||
/// When the first selection has empty reasoning but a subsequent one has
|
||||
/// non-empty reasoning, find_map should skip the empty one and return the
|
||||
/// first non-empty reasoning.
|
||||
#[test]
|
||||
fn test_reasoning_text_skips_empty_first_selection() {
|
||||
let selections = [
|
||||
ToolSelection {
|
||||
tool_name: "echo".into(),
|
||||
parameters: serde_json::json!({}),
|
||||
reasoning: String::new(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: "call_1".into(),
|
||||
},
|
||||
ToolSelection {
|
||||
tool_name: "search".into(),
|
||||
parameters: serde_json::json!({}),
|
||||
reasoning: "Found the answer in the second selection".into(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: "call_2".into(),
|
||||
},
|
||||
ToolSelection {
|
||||
tool_name: "fetch".into(),
|
||||
parameters: serde_json::json!({}),
|
||||
reasoning: "Third selection reasoning".into(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: "call_3".into(),
|
||||
},
|
||||
];
|
||||
|
||||
let reasoning_text = selections
|
||||
.iter()
|
||||
.find_map(|sel| (!sel.reasoning.is_empty()).then_some(sel.reasoning.clone()));
|
||||
|
||||
assert_eq!(
|
||||
reasoning_text.as_deref(),
|
||||
Some("Found the answer in the second selection"),
|
||||
"Should skip empty first reasoning and return the first non-empty one"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -231,7 +231,8 @@ impl EmbeddingProvider for OpenAiEmbeddings {
|
||||
.get("retry-after")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|s| s.parse::<u64>().ok())
|
||||
.map(std::time::Duration::from_secs);
|
||||
.map(std::time::Duration::from_secs)
|
||||
.or(Some(std::time::Duration::from_secs(60)));
|
||||
return Err(EmbeddingError::RateLimited { retry_after });
|
||||
}
|
||||
|
||||
@@ -372,7 +373,8 @@ impl EmbeddingProvider for NearAiEmbeddings {
|
||||
.get("retry-after")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|s| s.parse::<u64>().ok())
|
||||
.map(std::time::Duration::from_secs);
|
||||
.map(std::time::Duration::from_secs)
|
||||
.or(Some(std::time::Duration::from_secs(60)));
|
||||
return Err(EmbeddingError::RateLimited { retry_after });
|
||||
}
|
||||
|
||||
@@ -646,4 +648,48 @@ mod tests {
|
||||
let provider = OpenAiEmbeddings::new("test-key").with_base_url("custom.example.com/v1");
|
||||
assert_eq!(provider.base_url, "https://custom.example.com/v1");
|
||||
}
|
||||
|
||||
// -- Retry-After header parsing tests (regression for rate limit "None" bug) --
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_parsing_delay_seconds() {
|
||||
// Verify delay-seconds format is parsed correctly
|
||||
let header_value = "120";
|
||||
let duration = parse_retry_after_embeddings_for_test(header_value);
|
||||
assert_eq!(
|
||||
duration,
|
||||
Some(std::time::Duration::from_secs(120)),
|
||||
"Should parse delay-seconds format"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_fallback_missing_header() {
|
||||
// Regression test: When Retry-After header is missing,
|
||||
// should fall back to 60s instead of None
|
||||
let duration = parse_retry_after_embeddings_for_test("");
|
||||
assert_eq!(
|
||||
duration,
|
||||
Some(std::time::Duration::from_secs(60)),
|
||||
"Missing header should fallback to 60s"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_retry_after_zero_seconds_accepted() {
|
||||
// Verify zero seconds is a valid retry delay
|
||||
let duration = parse_retry_after_embeddings_for_test("0");
|
||||
assert_eq!(duration, Some(std::time::Duration::ZERO));
|
||||
}
|
||||
|
||||
/// Helper function to test Retry-After header parsing logic for embeddings
|
||||
/// (simulates the parsing done in embed without actual HTTP, including fallback)
|
||||
fn parse_retry_after_embeddings_for_test(header_value: &str) -> Option<std::time::Duration> {
|
||||
header_value
|
||||
.trim()
|
||||
.parse::<u64>()
|
||||
.ok()
|
||||
.map(std::time::Duration::from_secs)
|
||||
.or(Some(std::time::Duration::from_secs(60)))
|
||||
}
|
||||
}
|
||||
|
||||
+2
-2
@@ -52,7 +52,7 @@ HEADED=1 pytest scenarios/
|
||||
| `test_html_injection.py` | XSS vectors injected directly via `page.evaluate("addMessage('assistant', ...)")` are sanitized by `renderMarkdown`; user messages are shown as escaped plain text |
|
||||
| `test_skills.py` | Skills tab UI visibility, ClawHub search (skipped if registry unreachable), install + remove lifecycle |
|
||||
| `test_sse_reconnect.py` | SSE reconnects after programmatic `eventSource.close()` + `connectSSE()`; history is reloaded after reconnect |
|
||||
| `test_tool_approval.py` | Approval card appears, buttons disable on approve/deny, parameters toggle; all triggered via `page.evaluate("showApproval(...)")` — no real tool call needed |
|
||||
| `test_tool_approval.py` | Approval card appears, buttons disable on approve/deny, parameters toggle via `page.evaluate("showApproval(...)")`; the waiting-approval regression uses a real HTTP tool call |
|
||||
|
||||
## `helpers.py`
|
||||
|
||||
@@ -164,7 +164,7 @@ async def test_my_ui_feature(page):
|
||||
- **`asyncio_default_fixture_loop_scope = "session"`** — all async fixtures share one event loop. Do not use `asyncio.run()` inside fixtures; use `await` directly.
|
||||
- **The `page` fixture navigates with `/?token=e2e-test-token` and waits for `#auth-screen` to be hidden.** Tests receive a page that is already past the auth screen and has SSE connected.
|
||||
- **`test_skills.py` makes real network calls to ClawHub.** Tests skip (not fail) if the registry is unreachable via `pytest.skip()`.
|
||||
- **`test_html_injection.py` and `test_tool_approval.py` inject state via `page.evaluate(...)`.** They test the browser-side rendering pipeline and do not depend on the LLM or backend tool execution.
|
||||
- **`test_html_injection.py` injects state via `page.evaluate(...)`, and most of `test_tool_approval.py` does too.** The waiting-approval regression in `test_tool_approval.py` intentionally uses a real tool approval flow so it can verify backend thread-state handling.
|
||||
- **Browser is Chromium only.** `conftest.py` uses `p.chromium.launch()`; there is no Firefox or WebKit variant.
|
||||
- **Default timeout is 120 seconds** (pyproject.toml). Individual `wait_for` calls inside tests use shorter timeouts (5–20s) for faster failure messages.
|
||||
- **The libsql database is a temp directory** created fresh per `pytest` invocation; tests do not share state across runs.
|
||||
|
||||
+4
-2
@@ -164,5 +164,7 @@ await page.evaluate("""
|
||||
""")
|
||||
```
|
||||
|
||||
This is the pattern used in `test_tool_approval.py` and parts of
|
||||
`test_extensions.py` (auth card, configure modal).
|
||||
This is the pattern used in most of `test_tool_approval.py` and parts of
|
||||
`test_extensions.py` (auth card, configure modal). The waiting-approval
|
||||
regression in `test_tool_approval.py` uses a real tool call instead so it can
|
||||
exercise backend approval state.
|
||||
|
||||
+120
-22
@@ -15,7 +15,13 @@ from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from helpers import AUTH_TOKEN, wait_for_port_line, wait_for_ready
|
||||
from helpers import (
|
||||
AUTH_TOKEN,
|
||||
HTTP_WEBHOOK_SECRET,
|
||||
OWNER_SCOPE_ID,
|
||||
wait_for_port_line,
|
||||
wait_for_ready,
|
||||
)
|
||||
|
||||
# Project root (two levels up from tests/e2e/)
|
||||
ROOT = Path(__file__).resolve().parent.parent.parent
|
||||
@@ -39,6 +45,9 @@ except Exception:
|
||||
# Temp directory for the libSQL database file (cleaned up automatically)
|
||||
_DB_TMPDIR = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-")
|
||||
|
||||
# Temp HOME so pairing/allowFrom state never touches the developer's real ~/.ironclaw
|
||||
_HOME_TMPDIR = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-home-")
|
||||
|
||||
# Temp directories for WASM extensions. These start empty and are populated by
|
||||
# the install pipeline during tests; fixtures do not pre-populate dev build
|
||||
# artifacts into them.
|
||||
@@ -46,6 +55,42 @@ _WASM_TOOLS_TMPDIR = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-wasm-tools
|
||||
_WASM_CHANNELS_TMPDIR = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-wasm-channels-")
|
||||
|
||||
|
||||
def _latest_mtime(path: Path) -> float:
|
||||
"""Return the newest mtime under a file or directory."""
|
||||
if not path.exists():
|
||||
return 0.0
|
||||
if path.is_file():
|
||||
return path.stat().st_mtime
|
||||
|
||||
latest = path.stat().st_mtime
|
||||
for root, dirnames, filenames in os.walk(path):
|
||||
dirnames[:] = [dirname for dirname in dirnames if dirname != "target"]
|
||||
for name in filenames:
|
||||
child = Path(root) / name
|
||||
try:
|
||||
latest = max(latest, child.stat().st_mtime)
|
||||
except FileNotFoundError:
|
||||
continue
|
||||
return latest
|
||||
|
||||
|
||||
def _binary_needs_rebuild(binary: Path) -> bool:
|
||||
"""Rebuild when the binary is missing or older than embedded sources."""
|
||||
if not binary.exists():
|
||||
return True
|
||||
|
||||
binary_mtime = binary.stat().st_mtime
|
||||
inputs = [
|
||||
ROOT / "Cargo.toml",
|
||||
ROOT / "Cargo.lock",
|
||||
ROOT / "build.rs",
|
||||
ROOT / "providers.json",
|
||||
ROOT / "src",
|
||||
ROOT / "channels-src",
|
||||
]
|
||||
return any(_latest_mtime(path) > binary_mtime for path in inputs)
|
||||
|
||||
|
||||
def _find_free_port() -> int:
|
||||
"""Bind to port 0 and return the OS-assigned port."""
|
||||
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
|
||||
@@ -53,11 +98,26 @@ def _find_free_port() -> int:
|
||||
return s.getsockname()[1]
|
||||
|
||||
|
||||
def _reserve_loopback_sockets(count: int) -> list[socket.socket]:
|
||||
"""Bind loopback sockets and keep them open until the server starts."""
|
||||
sockets: list[socket.socket] = []
|
||||
try:
|
||||
while len(sockets) < count:
|
||||
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
sock.bind(("127.0.0.1", 0))
|
||||
sockets.append(sock)
|
||||
return sockets
|
||||
except Exception:
|
||||
for sock in sockets:
|
||||
sock.close()
|
||||
raise
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def ironclaw_binary():
|
||||
"""Ensure ironclaw binary is built. Returns the binary path."""
|
||||
binary = ROOT / "target" / "debug" / "ironclaw"
|
||||
if not binary.exists():
|
||||
if _binary_needs_rebuild(binary):
|
||||
print("Building ironclaw (this may take a while)...")
|
||||
subprocess.run(
|
||||
["cargo", "build", "--no-default-features", "--features", "libsql"],
|
||||
@@ -69,6 +129,21 @@ def ironclaw_binary():
|
||||
return str(binary)
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def server_ports():
|
||||
"""Reserve dynamic ports for the gateway and HTTP webhook channel."""
|
||||
reserved = _reserve_loopback_sockets(2)
|
||||
try:
|
||||
yield {
|
||||
"gateway": reserved[0].getsockname()[1],
|
||||
"http": reserved[1].getsockname()[1],
|
||||
"sockets": reserved,
|
||||
}
|
||||
finally:
|
||||
for sock in reserved:
|
||||
sock.close()
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
async def mock_llm_server():
|
||||
"""Start the mock LLM server. Yields the base URL."""
|
||||
@@ -138,20 +213,35 @@ def _wasm_build_symlinks():
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
async def ironclaw_server(ironclaw_binary, mock_llm_server, wasm_tools_dir):
|
||||
async def ironclaw_server(
|
||||
ironclaw_binary,
|
||||
mock_llm_server,
|
||||
wasm_tools_dir,
|
||||
server_ports,
|
||||
):
|
||||
"""Start the ironclaw gateway. Yields the base URL."""
|
||||
gateway_port = _find_free_port()
|
||||
home_dir = _HOME_TMPDIR.name
|
||||
gateway_port = server_ports["gateway"]
|
||||
http_port = server_ports["http"]
|
||||
for sock in server_ports["sockets"]:
|
||||
if sock.fileno() != -1:
|
||||
sock.close()
|
||||
env = {
|
||||
# Minimal env: PATH for process spawning, HOME for Rust/cargo defaults
|
||||
"PATH": os.environ.get("PATH", "/usr/bin:/bin"),
|
||||
"HOME": os.environ.get("HOME", "/tmp"),
|
||||
"HOME": home_dir,
|
||||
"IRONCLAW_BASE_DIR": os.path.join(home_dir, ".ironclaw"),
|
||||
"RUST_LOG": "ironclaw=info",
|
||||
"RUST_BACKTRACE": "1",
|
||||
"IRONCLAW_OWNER_ID": OWNER_SCOPE_ID,
|
||||
"GATEWAY_ENABLED": "true",
|
||||
"GATEWAY_HOST": "127.0.0.1",
|
||||
"GATEWAY_PORT": str(gateway_port),
|
||||
"GATEWAY_AUTH_TOKEN": AUTH_TOKEN,
|
||||
"GATEWAY_USER_ID": "e2e-tester",
|
||||
"GATEWAY_USER_ID": "e2e-web-sender",
|
||||
"HTTP_HOST": "127.0.0.1",
|
||||
"HTTP_PORT": str(http_port),
|
||||
"HTTP_WEBHOOK_SECRET": HTTP_WEBHOOK_SECRET,
|
||||
"CLI_ENABLED": "false",
|
||||
"LLM_BACKEND": "openai_compatible",
|
||||
"LLM_BASE_URL": mock_llm_server,
|
||||
@@ -221,15 +311,22 @@ async def ironclaw_server(ironclaw_binary, mock_llm_server, wasm_tools_dir):
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
async def ironclaw_server_with_webhook_secret(ironclaw_binary, mock_llm_server, wasm_tools_dir):
|
||||
"""Start ironclaw with HTTP_WEBHOOK_SECRET configured for webhook tests.
|
||||
async def http_channel_server(ironclaw_server, server_ports):
|
||||
"""HTTP webhook channel base URL."""
|
||||
base_url = f"http://127.0.0.1:{server_ports['http']}"
|
||||
await wait_for_ready(f"{base_url}/health", timeout=30)
|
||||
return base_url
|
||||
|
||||
Yields a dict with:
|
||||
- 'url': base URL of the gateway
|
||||
- 'secret': the webhook secret value
|
||||
"""
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
async def http_channel_server_without_secret(
|
||||
ironclaw_binary,
|
||||
mock_llm_server,
|
||||
wasm_tools_dir,
|
||||
):
|
||||
"""Start the HTTP webhook channel without a configured secret."""
|
||||
gateway_port = _find_free_port()
|
||||
webhook_secret = "test-webhook-secret-e2e-12345"
|
||||
http_port = _find_free_port()
|
||||
env = {
|
||||
# Minimal env: PATH for process spawning, HOME for Rust/cargo defaults
|
||||
"PATH": os.environ.get("PATH", "/usr/bin:/bin"),
|
||||
@@ -241,13 +338,14 @@ async def ironclaw_server_with_webhook_secret(ironclaw_binary, mock_llm_server,
|
||||
"GATEWAY_PORT": str(gateway_port),
|
||||
"GATEWAY_AUTH_TOKEN": AUTH_TOKEN,
|
||||
"GATEWAY_USER_ID": "e2e-tester",
|
||||
"HTTP_WEBHOOK_SECRET": webhook_secret,
|
||||
"HTTP_HOST": "127.0.0.1",
|
||||
"HTTP_PORT": str(http_port),
|
||||
"CLI_ENABLED": "false",
|
||||
"LLM_BACKEND": "openai_compatible",
|
||||
"LLM_BASE_URL": mock_llm_server,
|
||||
"LLM_MODEL": "mock-model",
|
||||
"DATABASE_BACKEND": "libsql",
|
||||
"LIBSQL_PATH": os.path.join(_DB_TMPDIR.name, "e2e-webhook.db"),
|
||||
"LIBSQL_PATH": os.path.join(_DB_TMPDIR.name, "e2e-webhook-no-secret.db"),
|
||||
"SANDBOX_ENABLED": "false",
|
||||
"SKILLS_ENABLED": "true",
|
||||
"ROUTINES_ENABLED": "false",
|
||||
@@ -277,13 +375,12 @@ async def ironclaw_server_with_webhook_secret(ironclaw_binary, mock_llm_server,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
env=env,
|
||||
)
|
||||
base_url = f"http://127.0.0.1:{gateway_port}"
|
||||
gateway_url = f"http://127.0.0.1:{gateway_port}"
|
||||
http_base_url = f"http://127.0.0.1:{http_port}"
|
||||
try:
|
||||
await wait_for_ready(f"{base_url}/api/health", timeout=60)
|
||||
yield {
|
||||
"url": base_url,
|
||||
"secret": webhook_secret,
|
||||
}
|
||||
await wait_for_ready(f"{gateway_url}/api/health", timeout=60)
|
||||
await wait_for_ready(f"{http_base_url}/health", timeout=30)
|
||||
yield http_base_url
|
||||
except TimeoutError:
|
||||
# Dump stderr so CI logs show why the server failed to start
|
||||
returncode = proc.returncode
|
||||
@@ -296,7 +393,8 @@ async def ironclaw_server_with_webhook_secret(ironclaw_binary, mock_llm_server,
|
||||
stderr_text = stderr_bytes.decode("utf-8", errors="replace")
|
||||
proc.kill()
|
||||
pytest.fail(
|
||||
f"ironclaw server with webhook secret failed to start on port {gateway_port} "
|
||||
f"ironclaw server without webhook secret failed to start on ports "
|
||||
f"gateway={gateway_port}, http={http_port} "
|
||||
f"(returncode={returncode}).\nstderr:\n{stderr_text}"
|
||||
)
|
||||
finally:
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
"""Shared helpers for E2E tests."""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import hmac
|
||||
import re
|
||||
import time
|
||||
|
||||
@@ -95,12 +97,21 @@ SEL = {
|
||||
"toast_success": ".toast.toast-success",
|
||||
"toast_error": ".toast.toast-error",
|
||||
"toast_info": ".toast.toast-info",
|
||||
# Jobs / routines
|
||||
"jobs_tbody": "#jobs-tbody",
|
||||
"job_row": "#jobs-tbody .job-row",
|
||||
"jobs_empty": "#jobs-empty",
|
||||
"routines_tbody": "#routines-tbody",
|
||||
"routine_row": "#routines-tbody .routine-row",
|
||||
"routines_empty": "#routines-empty",
|
||||
}
|
||||
|
||||
TABS = ["chat", "memory", "jobs", "routines", "extensions", "skills"]
|
||||
|
||||
# Auth token used across all tests
|
||||
AUTH_TOKEN = "e2e-test-token"
|
||||
OWNER_SCOPE_ID = "e2e-owner-scope"
|
||||
HTTP_WEBHOOK_SECRET = "e2e-http-webhook-secret"
|
||||
|
||||
|
||||
async def wait_for_ready(url: str, *, timeout: float = 60, interval: float = 0.5):
|
||||
@@ -162,3 +173,16 @@ async def api_post(base_url: str, path: str, **kwargs) -> httpx.Response:
|
||||
timeout=kwargs.pop("timeout", 10),
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
|
||||
def signed_http_webhook_headers(body: bytes) -> dict[str, str]:
|
||||
"""Return headers for the owner-scoped HTTP webhook channel."""
|
||||
digest = hmac.new(
|
||||
HTTP_WEBHOOK_SECRET.encode("utf-8"),
|
||||
body,
|
||||
hashlib.sha256,
|
||||
).hexdigest()
|
||||
return {
|
||||
"Content-Type": "application/json",
|
||||
"X-Hub-Signature-256": f"sha256={digest}",
|
||||
}
|
||||
|
||||
@@ -12,11 +12,17 @@ scenarios/test_csp.py
|
||||
scenarios/test_extension_oauth.py
|
||||
scenarios/test_extensions.py
|
||||
scenarios/test_html_injection.py
|
||||
scenarios/test_mcp_auth_flow.py
|
||||
scenarios/test_oauth_credential_fallback.py
|
||||
scenarios/test_owner_scope.py
|
||||
scenarios/test_pairing.py
|
||||
scenarios/test_routine_event_batch.py
|
||||
scenarios/test_routine_oauth_credential_injection.py
|
||||
scenarios/test_skills.py
|
||||
scenarios/test_sse_reconnect.py
|
||||
scenarios/test_telegram_hot_activation.py
|
||||
scenarios/test_telegram_token_validation.py
|
||||
scenarios/test_tool_approval.py
|
||||
scenarios/test_tool_execution.py
|
||||
scenarios/test_wasm_lifecycle.py
|
||||
scenarios/test_wasm_lifecycle.py
|
||||
scenarios/test_webhook.py
|
||||
@@ -25,7 +25,69 @@ DEFAULT_RESPONSE = "I understand your request."
|
||||
|
||||
TOOL_CALL_PATTERNS = [
|
||||
(re.compile(r"echo (.+)", re.IGNORECASE), "echo", lambda m: {"message": m.group(1)}),
|
||||
(
|
||||
re.compile(r"make approval post (?P<label>[a-z0-9_-]+)", re.IGNORECASE),
|
||||
"http",
|
||||
lambda m: {
|
||||
"method": "POST",
|
||||
"url": f"https://example.com/{m.group('label')}",
|
||||
"body": {"label": m.group("label")},
|
||||
},
|
||||
),
|
||||
(re.compile(r"what time|current time", re.IGNORECASE), "time", lambda _: {"operation": "now"}),
|
||||
(
|
||||
re.compile(
|
||||
r"create lightweight owner routine (?P<name>[a-z0-9][a-z0-9_-]*)",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
"routine_create",
|
||||
lambda m: {
|
||||
"name": m.group("name"),
|
||||
"description": f"Owner-scope routine {m.group('name')}",
|
||||
"trigger_type": "manual",
|
||||
"prompt": f"Confirm that {m.group('name')} executed.",
|
||||
"action_type": "lightweight",
|
||||
"use_tools": False,
|
||||
},
|
||||
),
|
||||
(
|
||||
re.compile(
|
||||
r"create full[- ]job owner routine (?P<name>[a-z0-9][a-z0-9_-]*)",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
"routine_create",
|
||||
lambda m: {
|
||||
"name": m.group("name"),
|
||||
"description": f"Owner-scope full-job routine {m.group('name')}",
|
||||
"trigger_type": "manual",
|
||||
"prompt": f"Complete the routine job for {m.group('name')}.",
|
||||
"action_type": "full_job",
|
||||
},
|
||||
),
|
||||
(
|
||||
re.compile(
|
||||
r"create event routine (?P<name>[a-z0-9][a-z0-9_-]*) "
|
||||
r"channel (?P<channel>[a-z0-9_-]+) pattern (?P<pattern>[a-z0-9_|-]+)",
|
||||
re.IGNORECASE,
|
||||
),
|
||||
"routine_create",
|
||||
lambda m: {
|
||||
"name": m.group("name"),
|
||||
"description": f"Event routine {m.group('name')}",
|
||||
"trigger_type": "event",
|
||||
"event_channel": None if m.group("channel").lower() == "any" else m.group("channel"),
|
||||
"event_pattern": m.group("pattern"),
|
||||
"prompt": f"Acknowledge that {m.group('name')} fired.",
|
||||
"action_type": "lightweight",
|
||||
"use_tools": False,
|
||||
"cooldown_secs": 0,
|
||||
},
|
||||
),
|
||||
(
|
||||
re.compile(r"list owner routines", re.IGNORECASE),
|
||||
"routine_list",
|
||||
lambda _: {},
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -885,6 +885,36 @@ async def test_auth_and_configure_helpers_escape_selector_sensitive_extension_na
|
||||
assert result["configureStillPresent"] is False
|
||||
|
||||
|
||||
async def test_auth_required_does_not_reopen_existing_configure_modal(page):
|
||||
"""Regression: auth_required SSE should not clobber an already-open configure modal."""
|
||||
result = await page.evaluate(
|
||||
"""() => {
|
||||
const overlay = document.createElement('div');
|
||||
overlay.className = 'configure-overlay';
|
||||
overlay.setAttribute('data-extension-name', 'telegram');
|
||||
document.body.appendChild(overlay);
|
||||
|
||||
const originalShowConfigureModal = window.showConfigureModal;
|
||||
const originalSetAuthFlowPending = window.setAuthFlowPending;
|
||||
let showCalls = 0;
|
||||
let pendingCalls = 0;
|
||||
|
||||
window.showConfigureModal = () => { showCalls += 1; };
|
||||
window.setAuthFlowPending = () => { pendingCalls += 1; };
|
||||
|
||||
handleAuthRequired({ extension_name: 'telegram', instructions: 'pending', auth_url: null });
|
||||
|
||||
window.showConfigureModal = originalShowConfigureModal;
|
||||
window.setAuthFlowPending = originalSetAuthFlowPending;
|
||||
overlay.remove();
|
||||
return { showCalls, pendingCalls };
|
||||
}"""
|
||||
)
|
||||
|
||||
assert result["showCalls"] == 0
|
||||
assert result["pendingCalls"] == 0
|
||||
|
||||
|
||||
async def test_auth_completed_sse_dismisses_card(page):
|
||||
"""Simulating the auth_completed SSE event removes the auth card."""
|
||||
await _show_auth_card(page, extension_name="myext")
|
||||
|
||||
@@ -0,0 +1,226 @@
|
||||
"""Owner-scope end-to-end scenarios.
|
||||
|
||||
These tests exercise the explicit owner model across:
|
||||
- the web gateway chat UI
|
||||
- the owner-scoped HTTP webhook channel
|
||||
- routine tools / routines tab
|
||||
- job creation via routine execution / jobs tab
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import uuid
|
||||
|
||||
import httpx
|
||||
|
||||
from helpers import SEL, AUTH_TOKEN, signed_http_webhook_headers
|
||||
|
||||
|
||||
async def _send_and_get_response(
|
||||
page,
|
||||
message: str,
|
||||
*,
|
||||
expected_fragment: str,
|
||||
timeout: int = 30000,
|
||||
) -> str:
|
||||
"""Send a chat message and return the newest assistant response text."""
|
||||
chat_input = page.locator(SEL["chat_input"])
|
||||
await chat_input.wait_for(state="visible", timeout=5000)
|
||||
|
||||
assistant_sel = SEL["message_assistant"]
|
||||
before_count = await page.locator(assistant_sel).count()
|
||||
|
||||
await chat_input.fill(message)
|
||||
await chat_input.press("Enter")
|
||||
|
||||
expected = before_count + 1
|
||||
await page.wait_for_function(
|
||||
"""({ assistantSelector, expectedCount, expectedFragment }) => {
|
||||
const messages = document.querySelectorAll(assistantSelector);
|
||||
if (messages.length < expectedCount) return false;
|
||||
const text = (messages[messages.length - 1].innerText || '').trim().toLowerCase();
|
||||
return text.includes(expectedFragment.toLowerCase());
|
||||
}""",
|
||||
arg={
|
||||
"assistantSelector": assistant_sel,
|
||||
"expectedCount": expected,
|
||||
"expectedFragment": expected_fragment,
|
||||
},
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
return await page.locator(assistant_sel).last.inner_text()
|
||||
|
||||
|
||||
async def _post_http_webhook(
|
||||
http_channel_server: str,
|
||||
*,
|
||||
content: str,
|
||||
sender_id: str,
|
||||
thread_id: str,
|
||||
) -> str:
|
||||
"""Send a signed request to the owner-scoped HTTP webhook channel."""
|
||||
payload = {
|
||||
"user_id": sender_id,
|
||||
"thread_id": thread_id,
|
||||
"content": content,
|
||||
"wait_for_response": True,
|
||||
}
|
||||
body = json.dumps(payload).encode("utf-8")
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.post(
|
||||
f"{http_channel_server}/webhook",
|
||||
content=body,
|
||||
headers=signed_http_webhook_headers(body),
|
||||
timeout=90,
|
||||
)
|
||||
|
||||
assert response.status_code == 200, (
|
||||
f"HTTP webhook failed: {response.status_code} {response.text[:400]}"
|
||||
)
|
||||
data = response.json()
|
||||
assert data["status"] == "accepted", f"Unexpected webhook response: {data}"
|
||||
assert data["response"], f"Expected synchronous response body, got: {data}"
|
||||
return data["response"]
|
||||
|
||||
|
||||
async def _open_tab(page, tab: str) -> None:
|
||||
btn = page.locator(SEL["tab_button"].format(tab=tab))
|
||||
await btn.click()
|
||||
await page.locator(SEL["tab_panel"].format(tab=tab)).wait_for(
|
||||
state="visible",
|
||||
timeout=5000,
|
||||
)
|
||||
|
||||
|
||||
async def _wait_for_routine(base_url: str, name: str, timeout: float = 20.0) -> dict:
|
||||
"""Poll the routines API until the named routine exists."""
|
||||
async with httpx.AsyncClient() as client:
|
||||
for _ in range(int(timeout * 2)):
|
||||
response = await client.get(
|
||||
f"{base_url}/api/routines",
|
||||
headers={"Authorization": f"Bearer {AUTH_TOKEN}"},
|
||||
timeout=10,
|
||||
)
|
||||
response.raise_for_status()
|
||||
routines = response.json()["routines"]
|
||||
for routine in routines:
|
||||
if routine["name"] == name:
|
||||
return routine
|
||||
await _poll_sleep()
|
||||
raise AssertionError(f"Routine '{name}' was not created within {timeout}s")
|
||||
|
||||
|
||||
async def _wait_for_job(base_url: str, title: str, timeout: float = 30.0) -> dict:
|
||||
"""Poll the jobs API until the named job exists."""
|
||||
async with httpx.AsyncClient() as client:
|
||||
for _ in range(int(timeout * 2)):
|
||||
response = await client.get(
|
||||
f"{base_url}/api/jobs",
|
||||
headers={"Authorization": f"Bearer {AUTH_TOKEN}"},
|
||||
timeout=10,
|
||||
)
|
||||
response.raise_for_status()
|
||||
jobs = response.json()["jobs"]
|
||||
for job in jobs:
|
||||
if job["title"] == title:
|
||||
return job
|
||||
await _poll_sleep()
|
||||
raise AssertionError(f"Job '{title}' was not created within {timeout}s")
|
||||
|
||||
|
||||
async def _poll_sleep() -> None:
|
||||
"""Small shared backoff for API polling loops."""
|
||||
await asyncio.sleep(0.5)
|
||||
|
||||
|
||||
async def test_http_channel_created_routine_is_visible_in_web_routines_tab(
|
||||
page,
|
||||
ironclaw_server,
|
||||
http_channel_server,
|
||||
):
|
||||
"""A routine created from the HTTP channel is visible in the web owner UI."""
|
||||
routine_name = f"owner-http-{uuid.uuid4().hex[:8]}"
|
||||
|
||||
response_text = await _post_http_webhook(
|
||||
http_channel_server,
|
||||
content=f"create lightweight owner routine {routine_name}",
|
||||
sender_id="external-sender-alpha",
|
||||
thread_id="http-owner-routine-thread",
|
||||
)
|
||||
assert routine_name in response_text
|
||||
|
||||
await _wait_for_routine(ironclaw_server, routine_name)
|
||||
|
||||
await _open_tab(page, "routines")
|
||||
await page.locator(SEL["routine_row"]).filter(has_text=routine_name).first.wait_for(
|
||||
state="visible",
|
||||
timeout=15000,
|
||||
)
|
||||
|
||||
|
||||
async def test_web_created_routine_is_listed_from_http_channel_across_senders(
|
||||
page,
|
||||
ironclaw_server,
|
||||
http_channel_server,
|
||||
):
|
||||
"""Routines created in web chat remain owner-global across HTTP senders/threads."""
|
||||
routine_name = f"owner-web-{uuid.uuid4().hex[:8]}"
|
||||
|
||||
assistant_text = await _send_and_get_response(
|
||||
page,
|
||||
f"create lightweight owner routine {routine_name}",
|
||||
expected_fragment=routine_name,
|
||||
)
|
||||
assert routine_name in assistant_text
|
||||
|
||||
await _wait_for_routine(ironclaw_server, routine_name)
|
||||
|
||||
first_sender_text = await _post_http_webhook(
|
||||
http_channel_server,
|
||||
content="list owner routines",
|
||||
sender_id="http-sender-one",
|
||||
thread_id="owner-list-thread-a",
|
||||
)
|
||||
second_sender_text = await _post_http_webhook(
|
||||
http_channel_server,
|
||||
content="list owner routines",
|
||||
sender_id="http-sender-two",
|
||||
thread_id="owner-list-thread-b",
|
||||
)
|
||||
|
||||
assert routine_name in first_sender_text, first_sender_text
|
||||
assert routine_name in second_sender_text, second_sender_text
|
||||
|
||||
|
||||
async def test_http_created_full_job_routine_can_be_run_from_web_and_shows_in_jobs(
|
||||
page,
|
||||
ironclaw_server,
|
||||
http_channel_server,
|
||||
):
|
||||
"""A full-job routine created via HTTP can be run from the web UI and create a job."""
|
||||
routine_name = f"owner-job-{uuid.uuid4().hex[:8]}"
|
||||
|
||||
response_text = await _post_http_webhook(
|
||||
http_channel_server,
|
||||
content=f"create full-job owner routine {routine_name}",
|
||||
sender_id="http-job-sender",
|
||||
thread_id="owner-job-thread",
|
||||
)
|
||||
assert routine_name in response_text
|
||||
|
||||
await _wait_for_routine(ironclaw_server, routine_name)
|
||||
|
||||
await _open_tab(page, "routines")
|
||||
routine_row = page.locator(SEL["routine_row"]).filter(has_text=routine_name).first
|
||||
await routine_row.wait_for(state="visible", timeout=15000)
|
||||
await routine_row.locator('button[data-action="trigger-routine"]').click()
|
||||
|
||||
await _wait_for_job(ironclaw_server, routine_name, timeout=45.0)
|
||||
|
||||
await _open_tab(page, "jobs")
|
||||
await page.locator(SEL["job_row"]).filter(has_text=routine_name).first.wait_for(
|
||||
state="visible",
|
||||
timeout=20000,
|
||||
)
|
||||
@@ -1,534 +1,317 @@
|
||||
"""
|
||||
E2E tests for event-triggered routines with batch loading.
|
||||
|
||||
These tests verify that the N+1 query fix correctly:
|
||||
1. Fires event-triggered routines on matching messages
|
||||
2. Enforces concurrent limits via batch-loaded counts
|
||||
3. Maintains performance with multiple simultaneous triggers
|
||||
4. Works correctly through the full UI and agent loop
|
||||
|
||||
Playwright-based UI tests + SSE verification.
|
||||
"""
|
||||
"""E2E tests for event-triggered routines over the HTTP channel."""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import uuid
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from datetime import datetime, timedelta
|
||||
from typing import List, Dict, Any
|
||||
|
||||
from playwright.async_api import async_playwright, Page, Browser, BrowserContext
|
||||
from helpers import AUTH_TOKEN, SEL, signed_http_webhook_headers
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def browser_and_context():
|
||||
"""Create a Playwright browser and context for testing."""
|
||||
async with async_playwright() as p:
|
||||
browser = await p.chromium.launch(headless=True)
|
||||
context = await browser.new_context()
|
||||
yield browser, context
|
||||
await context.close()
|
||||
await browser.close()
|
||||
async def _send_chat_message(page, message: str) -> None:
|
||||
"""Send a chat message and wait for the assistant turn to appear."""
|
||||
chat_input = page.locator(SEL["chat_input"])
|
||||
await chat_input.wait_for(state="visible", timeout=5000)
|
||||
assistant_messages = page.locator(SEL["message_assistant"])
|
||||
before_count = await assistant_messages.count()
|
||||
|
||||
await chat_input.fill(message)
|
||||
await chat_input.press("Enter")
|
||||
|
||||
await page.wait_for_function(
|
||||
"""({ selector, expectedCount }) => {
|
||||
return document.querySelectorAll(selector).length >= expectedCount;
|
||||
}""",
|
||||
arg={
|
||||
"selector": SEL["message_assistant"],
|
||||
"expectedCount": before_count + 1,
|
||||
},
|
||||
timeout=30000,
|
||||
)
|
||||
|
||||
|
||||
class EventTriggerHelper:
|
||||
"""Helper methods for event trigger testing."""
|
||||
async def _create_event_routine(
|
||||
page,
|
||||
base_url: str,
|
||||
*,
|
||||
name: str,
|
||||
pattern: str,
|
||||
channel: str = "http",
|
||||
) -> dict:
|
||||
"""Create an event routine through chat and return its API record."""
|
||||
await _send_chat_message(
|
||||
page,
|
||||
f"create event routine {name} channel {channel} pattern {pattern}",
|
||||
)
|
||||
return await _wait_for_routine(base_url, name)
|
||||
|
||||
def __init__(self, page: Page):
|
||||
self.page = page
|
||||
|
||||
async def navigate_to_routines(self):
|
||||
"""Navigate to the routines page."""
|
||||
await self.page.goto("http://localhost:8000/routines")
|
||||
await self.page.wait_for_load_state("networkidle")
|
||||
async def _post_http_message(
|
||||
http_channel_server: str,
|
||||
*,
|
||||
content: str,
|
||||
sender_id: str | None = None,
|
||||
thread_id: str | None = None,
|
||||
) -> dict:
|
||||
"""Send a signed HTTP-channel message and return the JSON body."""
|
||||
payload = {
|
||||
"user_id": sender_id or f"sender-{uuid.uuid4().hex[:8]}",
|
||||
"thread_id": thread_id or f"thread-{uuid.uuid4().hex[:8]}",
|
||||
"content": content,
|
||||
"wait_for_response": True,
|
||||
}
|
||||
body = json.dumps(payload).encode("utf-8")
|
||||
|
||||
async def create_event_routine(
|
||||
self,
|
||||
name: str,
|
||||
trigger_regex: str,
|
||||
channel: str = "slack",
|
||||
max_concurrent: int = 1,
|
||||
) -> str:
|
||||
"""
|
||||
Create an event-triggered routine via UI.
|
||||
Returns the routine ID.
|
||||
"""
|
||||
await self.navigate_to_routines()
|
||||
|
||||
# Click "New Routine" button
|
||||
await self.page.click('button:has-text("New Routine")')
|
||||
await self.page.wait_for_selector('input[name="routine_name"]')
|
||||
|
||||
# Fill routine details
|
||||
await self.page.fill('input[name="routine_name"]', name)
|
||||
await self.page.fill(
|
||||
'textarea[name="routine_description"]',
|
||||
f"Test routine: {name}",
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.post(
|
||||
f"{http_channel_server}/webhook",
|
||||
content=body,
|
||||
headers=signed_http_webhook_headers(body),
|
||||
timeout=90,
|
||||
)
|
||||
|
||||
# Select "Event Trigger" type
|
||||
await self.page.click('label:has-text("Event Trigger")')
|
||||
await self.page.wait_for_selector('input[name="trigger_regex"]')
|
||||
|
||||
# Fill trigger details
|
||||
await self.page.fill('input[name="trigger_regex"]', trigger_regex)
|
||||
await self.page.select_option('select[name="trigger_channel"]', channel)
|
||||
|
||||
# Set guardrails
|
||||
await self.page.fill('input[name="max_concurrent"]', str(max_concurrent))
|
||||
|
||||
# Select lightweight action
|
||||
await self.page.click('label:has-text("Lightweight")')
|
||||
await self.page.fill(
|
||||
'textarea[name="lightweight_prompt"]',
|
||||
"Acknowledge the message and confirm trigger worked.",
|
||||
)
|
||||
|
||||
# Save routine
|
||||
await self.page.click('button:has-text("Save Routine")')
|
||||
await self.page.wait_for_selector('text=Routine created successfully')
|
||||
|
||||
# Extract routine ID from success message or URL
|
||||
routine_id = await self.page.locator('data-testid=routine-id').text_content()
|
||||
return routine_id.strip() if routine_id else None
|
||||
|
||||
async def create_multiple_routines(
|
||||
self, base_name: str, count: int, trigger_regex: str = None
|
||||
) -> List[str]:
|
||||
"""Create multiple event-triggered routines."""
|
||||
routine_ids = []
|
||||
for i in range(count):
|
||||
name = f"{base_name}_{i}"
|
||||
regex = trigger_regex or f"({i}|{base_name})"
|
||||
routine_id = await self.create_event_routine(name, regex)
|
||||
routine_ids.append(routine_id)
|
||||
await asyncio.sleep(0.1) # Small delay between creations
|
||||
return routine_ids
|
||||
|
||||
async def send_chat_message(self, message: str) -> List[str]:
|
||||
"""
|
||||
Send a chat message and return SSE events received.
|
||||
Captures all routine firing events.
|
||||
"""
|
||||
await self.page.goto("http://localhost:8000/chat")
|
||||
await self.page.wait_for_selector('input[placeholder*="message"]', timeout=5000)
|
||||
|
||||
# Collect SSE events
|
||||
sse_events = []
|
||||
|
||||
async def capture_sse(response):
|
||||
"""Intercept SSE events."""
|
||||
if "event-stream" in response.headers.get("content-type", ""):
|
||||
text = await response.text()
|
||||
for line in text.split("\n"):
|
||||
if line.startswith("data:"):
|
||||
try:
|
||||
event = json.loads(line[5:])
|
||||
sse_events.append(event)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
self.page.on("response", capture_sse)
|
||||
|
||||
# Send message
|
||||
await self.page.fill('input[placeholder*="message"]', message)
|
||||
await self.page.press('input[placeholder*="message"]', "Enter")
|
||||
|
||||
# Wait for response
|
||||
await self.page.wait_for_selector('text=Message processed', timeout=10000)
|
||||
await asyncio.sleep(0.5) # Allow time for SSE events
|
||||
|
||||
self.page.remove_listener("response", capture_sse)
|
||||
return sse_events
|
||||
|
||||
async def get_routine_execution_log(self, routine_id: str) -> List[Dict]:
|
||||
"""Get execution log entries for a routine."""
|
||||
await self.page.goto(f"http://localhost:8000/routines/{routine_id}/executions")
|
||||
await self.page.wait_for_load_state("networkidle")
|
||||
|
||||
# Extract log entries from table
|
||||
rows = await self.page.locator("tbody tr").all()
|
||||
executions = []
|
||||
|
||||
for row in rows:
|
||||
cells = await row.locator("td").all()
|
||||
if len(cells) >= 3:
|
||||
execution = {
|
||||
"timestamp": await cells[0].text_content(),
|
||||
"status": await cells[1].text_content(),
|
||||
"details": await cells[2].text_content(),
|
||||
}
|
||||
executions.append(execution)
|
||||
|
||||
return executions
|
||||
|
||||
async def check_database_queries_in_logs(
|
||||
self, max_queries_expected: int = 1
|
||||
) -> int:
|
||||
"""Check debug logs for database query count."""
|
||||
await self.page.goto("http://localhost:8000/debug/logs?filter=database")
|
||||
await self.page.wait_for_load_state("networkidle")
|
||||
|
||||
# Count batch queries
|
||||
log_lines = await self.page.locator("tr:has-text('batch')").all()
|
||||
batch_count = len(log_lines)
|
||||
|
||||
# Count individual COUNT queries (should be 0 after fix)
|
||||
count_queries = await self.page.locator("tr:has-text('COUNT')").all()
|
||||
count_query_count = len(count_queries)
|
||||
|
||||
return batch_count, count_query_count
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Tests
|
||||
# =============================================================================
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_event_trigger_routine(browser_and_context):
|
||||
"""Test creating an event-triggered routine via UI."""
|
||||
browser, context = browser_and_context
|
||||
page = await context.new_page()
|
||||
helper = EventTriggerHelper(page)
|
||||
|
||||
try:
|
||||
routine_id = await helper.create_event_routine(
|
||||
name="Test Trigger",
|
||||
trigger_regex="test|demo",
|
||||
channel="slack",
|
||||
max_concurrent=1,
|
||||
)
|
||||
|
||||
assert routine_id is not None, "Routine ID should be returned"
|
||||
assert len(routine_id) > 0, "Routine ID should not be empty"
|
||||
|
||||
finally:
|
||||
await page.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_event_trigger_fires_on_matching_message(browser_and_context):
|
||||
"""Test that event-triggered routine fires when message matches."""
|
||||
browser, context = browser_and_context
|
||||
page = await context.new_page()
|
||||
helper = EventTriggerHelper(page)
|
||||
|
||||
try:
|
||||
# Create routine
|
||||
routine_id = await helper.create_event_routine(
|
||||
name="Alert Handler",
|
||||
trigger_regex="urgent|critical|alert",
|
||||
channel="slack",
|
||||
)
|
||||
|
||||
# Send matching message
|
||||
sse_events = await helper.send_chat_message("URGENT: Server down!")
|
||||
|
||||
# Verify routine fired (look for event in SSE stream)
|
||||
routine_fired = any(
|
||||
event.get("type") == "routine_fired" and event.get("routine_id") == routine_id
|
||||
for event in sse_events
|
||||
)
|
||||
assert routine_fired, "Routine should fire on matching message"
|
||||
|
||||
# Check execution log
|
||||
executions = await helper.get_routine_execution_log(routine_id)
|
||||
assert len(executions) > 0, "Execution should be logged"
|
||||
assert "success" in executions[0]["status"].lower()
|
||||
|
||||
finally:
|
||||
await page.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_event_trigger_skips_non_matching_message(browser_and_context):
|
||||
"""Test that event-triggered routine skips when message doesn't match."""
|
||||
browser, context = browser_and_context
|
||||
page = await context.new_page()
|
||||
helper = EventTriggerHelper(page)
|
||||
|
||||
try:
|
||||
# Create routine
|
||||
routine_id = await helper.create_event_routine(
|
||||
name="Alert Handler",
|
||||
trigger_regex="urgent|critical|alert",
|
||||
channel="slack",
|
||||
)
|
||||
|
||||
# Send non-matching message
|
||||
sse_events = await helper.send_chat_message("Hello, how are you?")
|
||||
|
||||
# Verify routine did NOT fire
|
||||
routine_fired = any(
|
||||
event.get("type") == "routine_fired" and event.get("routine_id") == routine_id
|
||||
for event in sse_events
|
||||
)
|
||||
assert not routine_fired, "Routine should not fire on non-matching message"
|
||||
|
||||
finally:
|
||||
await page.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multiple_routines_fire_on_matching_message(browser_and_context):
|
||||
"""Test that multiple event-triggered routines fire on same message."""
|
||||
browser, context = browser_and_context
|
||||
page = await context.new_page()
|
||||
helper = EventTriggerHelper(page)
|
||||
|
||||
try:
|
||||
# Create 3 overlapping routines
|
||||
routine_ids = await helper.create_multiple_routines(
|
||||
base_name="Handler", count=3, trigger_regex="alert|warning|error"
|
||||
)
|
||||
|
||||
# Send matching message
|
||||
sse_events = await helper.send_chat_message("ERROR: Database connection failed")
|
||||
|
||||
# Verify all 3 routines fired
|
||||
fired_count = sum(
|
||||
1
|
||||
for event in sse_events
|
||||
if event.get("type") == "routine_fired" and event.get("routine_id") in routine_ids
|
||||
)
|
||||
|
||||
assert (
|
||||
fired_count >= 3
|
||||
), f"Expected all 3 routines to fire, got {fired_count}"
|
||||
|
||||
finally:
|
||||
await page.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_limit_prevents_additional_fires(browser_and_context):
|
||||
"""Test that concurrent limit is enforced via batch counts."""
|
||||
browser, context = browser_and_context
|
||||
page = await context.new_page()
|
||||
helper = EventTriggerHelper(page)
|
||||
|
||||
try:
|
||||
# Create routine with max_concurrent=1
|
||||
routine_id = await helper.create_event_routine(
|
||||
name="Limited Handler",
|
||||
trigger_regex="process|task",
|
||||
max_concurrent=1,
|
||||
)
|
||||
|
||||
# Trigger first message
|
||||
await helper.send_chat_message("Process message 1")
|
||||
await asyncio.sleep(1)
|
||||
|
||||
# Check first execution logged
|
||||
executions_1 = await helper.get_routine_execution_log(routine_id)
|
||||
assert len(executions_1) >= 1
|
||||
|
||||
# Trigger second message while first is still running
|
||||
sse_events = await helper.send_chat_message("Process message 2")
|
||||
|
||||
# Second routine should be skipped (concurrent limit)
|
||||
routine_skipped = any(
|
||||
event.get("type") == "routine_skipped"
|
||||
and event.get("reason") == "max_concurrent_reached"
|
||||
and event.get("routine_id") == routine_id
|
||||
for event in sse_events
|
||||
)
|
||||
assert routine_skipped, "Routine should be skipped when concurrent limit reached"
|
||||
|
||||
finally:
|
||||
await page.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rapid_messages_with_multiple_triggers_efficiency(browser_and_context):
|
||||
"""Test efficiency of batch loading with multiple rapid messages."""
|
||||
browser, context = browser_and_context
|
||||
page = await context.new_page()
|
||||
helper = EventTriggerHelper(page)
|
||||
|
||||
try:
|
||||
# Create 5 overlapping routines
|
||||
routine_ids = await helper.create_multiple_routines(
|
||||
base_name="Rapid", count=5, trigger_regex="test|demo|check"
|
||||
)
|
||||
|
||||
# Send 10 matching messages rapidly
|
||||
for i in range(10):
|
||||
message = f"test message {i}"
|
||||
await helper.send_chat_message(message)
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
# Check database logs for query efficiency
|
||||
batch_count, count_query_count = await helper.check_database_queries_in_logs()
|
||||
|
||||
# After fix: should have ~10 batch queries (1 per message)
|
||||
# Before fix: would have ~50 individual COUNT queries (5 routines × 10 messages)
|
||||
assert (
|
||||
count_query_count == 0
|
||||
), f"Should have 0 individual COUNT queries after fix, got {count_query_count}"
|
||||
assert (
|
||||
batch_count <= 15
|
||||
), f"Should have <=15 batch queries for 10 messages, got {batch_count}"
|
||||
|
||||
finally:
|
||||
await page.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_channel_filter_applied_correctly(browser_and_context):
|
||||
"""Test that channel filter prevents non-matching messages."""
|
||||
browser, context = browser_and_context
|
||||
page = await context.new_page()
|
||||
helper = EventTriggerHelper(page)
|
||||
|
||||
try:
|
||||
# Create routine for Slack channel
|
||||
slack_routine_id = await helper.create_event_routine(
|
||||
name="Slack Handler",
|
||||
trigger_regex="alert",
|
||||
channel="slack",
|
||||
)
|
||||
|
||||
# Simulate message from Telegram channel
|
||||
# (Note: In real UI, would need to change channel context)
|
||||
page.goto(
|
||||
"http://localhost:8000/chat?channel=telegram"
|
||||
) # Switch channel
|
||||
await helper.send_chat_message("alert: something urgent")
|
||||
|
||||
# Routine should not fire (different channel)
|
||||
executions = await helper.get_routine_execution_log(slack_routine_id)
|
||||
|
||||
# Check if any recent execution (last 5 min) exists
|
||||
recent = [
|
||||
e
|
||||
for e in executions
|
||||
if (datetime.now() - datetime.fromisoformat(e["timestamp"])).total_seconds()
|
||||
< 300
|
||||
]
|
||||
assert (
|
||||
len(recent) == 0
|
||||
), "Routine should not fire for different channel"
|
||||
|
||||
finally:
|
||||
await page.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_batch_query_failure_handling(browser_and_context):
|
||||
"""Test graceful handling of batch query failures."""
|
||||
browser, context = browser_and_context
|
||||
page = await context.new_page()
|
||||
helper = EventTriggerHelper(page)
|
||||
|
||||
try:
|
||||
# Create routine
|
||||
routine_id = await helper.create_event_routine(
|
||||
name="Error Handler",
|
||||
trigger_regex="test",
|
||||
)
|
||||
|
||||
# Simulate database error in logs (if possible with test hooks)
|
||||
# For now, just verify error handling doesn't crash UI
|
||||
await helper.send_chat_message("test message")
|
||||
|
||||
# Check that UI remains responsive
|
||||
assert await page.locator("text=Message processed").is_visible()
|
||||
|
||||
finally:
|
||||
await page.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_routine_execution_history_display(browser_and_context):
|
||||
"""Test that execution history correctly displays routine firings."""
|
||||
browser, context = browser_and_context
|
||||
page = await context.new_page()
|
||||
helper = EventTriggerHelper(page)
|
||||
|
||||
try:
|
||||
# Create routine
|
||||
routine_id = await helper.create_event_routine(
|
||||
name="History Test",
|
||||
trigger_regex="test",
|
||||
)
|
||||
|
||||
# Trigger routine 3 times
|
||||
for i in range(3):
|
||||
await helper.send_chat_message(f"test message {i}")
|
||||
await asyncio.sleep(0.2)
|
||||
|
||||
# Check execution log
|
||||
executions = await helper.get_routine_execution_log(routine_id)
|
||||
assert len(executions) >= 3, "Should have at least 3 executions logged"
|
||||
|
||||
# Verify all are recent (within last 5 minutes)
|
||||
for execution in executions[:3]:
|
||||
timestamp = datetime.fromisoformat(execution["timestamp"])
|
||||
age = datetime.now() - timestamp
|
||||
assert age < timedelta(minutes=5), "Execution should be recent"
|
||||
|
||||
finally:
|
||||
await page.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_batch_loads_independent(browser_and_context):
|
||||
"""Test that concurrent messages each get independent batch queries."""
|
||||
browser, context = browser_and_context
|
||||
page = await context.new_page()
|
||||
helper = EventTriggerHelper(page)
|
||||
|
||||
try:
|
||||
# Create 5 routines matching different patterns
|
||||
r1_id = await helper.create_event_routine(
|
||||
name="Pattern A", trigger_regex="alpha|alpha_only"
|
||||
)
|
||||
r2_id = await helper.create_event_routine(
|
||||
name="Pattern B", trigger_regex="beta|beta_only"
|
||||
)
|
||||
r3_id = await helper.create_event_routine(
|
||||
name="Pattern AB", trigger_regex="alpha|beta|common"
|
||||
)
|
||||
|
||||
# Send overlapping messages
|
||||
# Message 1: matches r1, r3
|
||||
sse1 = await helper.send_chat_message("alpha common")
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
# Message 2: matches r2, r3
|
||||
sse2 = await helper.send_chat_message("beta common")
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
# Verify correct routines fired
|
||||
r1_fired_msg1 = any(
|
||||
e.get("routine_id") == r1_id for e in sse1 if e.get("type") == "routine_fired"
|
||||
)
|
||||
r2_fired_msg2 = any(
|
||||
e.get("routine_id") == r2_id for e in sse2 if e.get("type") == "routine_fired"
|
||||
)
|
||||
r3_fired_both = (
|
||||
any(
|
||||
e.get("routine_id") == r3_id for e in sse1 if e.get("type") == "routine_fired"
|
||||
assert response.status_code == 200, (
|
||||
f"HTTP webhook failed: {response.status_code} {response.text[:400]}"
|
||||
)
|
||||
return response.json()
|
||||
|
||||
|
||||
async def _wait_for_routine(base_url: str, name: str, timeout: float = 20.0) -> dict:
|
||||
"""Poll the routines API until the named routine exists."""
|
||||
async with httpx.AsyncClient() as client:
|
||||
for _ in range(int(timeout * 2)):
|
||||
response = await client.get(
|
||||
f"{base_url}/api/routines",
|
||||
headers={"Authorization": f"Bearer {AUTH_TOKEN}"},
|
||||
timeout=10,
|
||||
)
|
||||
and any(
|
||||
e.get("routine_id") == r3_id for e in sse2 if e.get("type") == "routine_fired"
|
||||
response.raise_for_status()
|
||||
for routine in response.json()["routines"]:
|
||||
if routine["name"] == name:
|
||||
return routine
|
||||
await asyncio.sleep(0.5)
|
||||
raise AssertionError(f"Routine '{name}' was not created within {timeout}s")
|
||||
|
||||
|
||||
async def _get_routine_runs(base_url: str, routine_id: str) -> list[dict]:
|
||||
"""Fetch recent routine runs from the web API."""
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.get(
|
||||
f"{base_url}/api/routines/{routine_id}/runs",
|
||||
headers={"Authorization": f"Bearer {AUTH_TOKEN}"},
|
||||
timeout=10,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()["runs"]
|
||||
|
||||
|
||||
async def _wait_for_run_count(
|
||||
base_url: str,
|
||||
routine_id: str,
|
||||
*,
|
||||
expected_at_least: int,
|
||||
timeout: float = 20.0,
|
||||
) -> list[dict]:
|
||||
"""Poll until the routine has at least the expected run count."""
|
||||
for _ in range(int(timeout * 2)):
|
||||
runs = await _get_routine_runs(base_url, routine_id)
|
||||
if len(runs) >= expected_at_least:
|
||||
return runs
|
||||
await asyncio.sleep(0.5)
|
||||
raise AssertionError(
|
||||
f"Routine '{routine_id}' did not reach {expected_at_least} runs within {timeout}s"
|
||||
)
|
||||
|
||||
|
||||
async def _wait_for_completed_run(
|
||||
base_url: str,
|
||||
routine_id: str,
|
||||
*,
|
||||
timeout: float = 30.0,
|
||||
) -> dict:
|
||||
"""Poll until the newest run is no longer marked running."""
|
||||
for _ in range(int(timeout * 2)):
|
||||
runs = await _get_routine_runs(base_url, routine_id)
|
||||
if runs and runs[0]["status"].lower() != "running":
|
||||
return runs[0]
|
||||
await asyncio.sleep(0.5)
|
||||
raise AssertionError(f"Routine '{routine_id}' did not complete within {timeout}s")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_event_trigger_routine(page, ironclaw_server):
|
||||
"""Event routines can be created through the supported chat flow."""
|
||||
name = f"evt-{uuid.uuid4().hex[:8]}"
|
||||
routine = await _create_event_routine(
|
||||
page,
|
||||
ironclaw_server,
|
||||
name=name,
|
||||
pattern="test|demo",
|
||||
)
|
||||
|
||||
assert routine["id"]
|
||||
assert routine["trigger_type"] == "event"
|
||||
assert "test|demo" in routine["trigger_summary"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_event_trigger_fires_on_matching_message(
|
||||
page,
|
||||
ironclaw_server,
|
||||
http_channel_server,
|
||||
):
|
||||
"""Matching HTTP-channel messages create routine runs."""
|
||||
name = f"evt-{uuid.uuid4().hex[:8]}"
|
||||
routine = await _create_event_routine(
|
||||
page,
|
||||
ironclaw_server,
|
||||
name=name,
|
||||
pattern="urgent|critical|alert",
|
||||
)
|
||||
|
||||
response = await _post_http_message(
|
||||
http_channel_server,
|
||||
content="urgent: server down",
|
||||
)
|
||||
assert response["status"] == "accepted"
|
||||
|
||||
await _wait_for_run_count(
|
||||
ironclaw_server,
|
||||
routine["id"],
|
||||
expected_at_least=1,
|
||||
)
|
||||
completed_run = await _wait_for_completed_run(ironclaw_server, routine["id"])
|
||||
|
||||
assert completed_run["status"].lower() == "attention"
|
||||
assert completed_run["trigger_type"] == "event"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_event_trigger_skips_non_matching_message(
|
||||
page,
|
||||
ironclaw_server,
|
||||
http_channel_server,
|
||||
):
|
||||
"""Non-matching messages do not create routine runs."""
|
||||
name = f"evt-{uuid.uuid4().hex[:8]}"
|
||||
routine = await _create_event_routine(
|
||||
page,
|
||||
ironclaw_server,
|
||||
name=name,
|
||||
pattern="urgent|critical|alert",
|
||||
)
|
||||
|
||||
await _post_http_message(
|
||||
http_channel_server,
|
||||
content="hello there",
|
||||
)
|
||||
await asyncio.sleep(2)
|
||||
|
||||
assert await _get_routine_runs(ironclaw_server, routine["id"]) == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multiple_routines_fire_on_matching_message(
|
||||
page,
|
||||
ironclaw_server,
|
||||
http_channel_server,
|
||||
):
|
||||
"""A single matching message can fire multiple event routines."""
|
||||
routines = []
|
||||
for _ in range(3):
|
||||
name = f"evt-{uuid.uuid4().hex[:8]}"
|
||||
routines.append(
|
||||
await _create_event_routine(
|
||||
page,
|
||||
ironclaw_server,
|
||||
name=name,
|
||||
pattern="error|warning|alert",
|
||||
)
|
||||
)
|
||||
|
||||
assert r1_fired_msg1, "Routine 1 should fire on message 1"
|
||||
assert r2_fired_msg2, "Routine 2 should fire on message 2"
|
||||
assert r3_fired_both, "Routine 3 should fire on both messages"
|
||||
await _post_http_message(
|
||||
http_channel_server,
|
||||
content="error: database connection failed",
|
||||
)
|
||||
|
||||
finally:
|
||||
await page.close()
|
||||
for routine in routines:
|
||||
await _wait_for_run_count(
|
||||
ironclaw_server,
|
||||
routine["id"],
|
||||
expected_at_least=1,
|
||||
)
|
||||
completed_run = await _wait_for_completed_run(ironclaw_server, routine["id"])
|
||||
assert completed_run["status"].lower() == "attention"
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Integration with existing test patterns
|
||||
# =============================================================================
|
||||
@pytest.mark.asyncio
|
||||
async def test_channel_filter_applied_correctly(
|
||||
page,
|
||||
ironclaw_server,
|
||||
http_channel_server,
|
||||
):
|
||||
"""Channel filters prevent HTTP messages from firing non-HTTP routines."""
|
||||
http_routine = await _create_event_routine(
|
||||
page,
|
||||
ironclaw_server,
|
||||
name=f"evt-{uuid.uuid4().hex[:8]}",
|
||||
pattern="alert",
|
||||
channel="http",
|
||||
)
|
||||
telegram_routine = await _create_event_routine(
|
||||
page,
|
||||
ironclaw_server,
|
||||
name=f"evt-{uuid.uuid4().hex[:8]}",
|
||||
pattern="alert",
|
||||
channel="telegram",
|
||||
)
|
||||
|
||||
await _post_http_message(
|
||||
http_channel_server,
|
||||
content="alert from webhook",
|
||||
)
|
||||
|
||||
await _wait_for_run_count(
|
||||
ironclaw_server,
|
||||
http_routine["id"],
|
||||
expected_at_least=1,
|
||||
)
|
||||
http_run = await _wait_for_completed_run(ironclaw_server, http_routine["id"])
|
||||
await asyncio.sleep(2)
|
||||
telegram_runs = await _get_routine_runs(ironclaw_server, telegram_routine["id"])
|
||||
|
||||
assert http_run["status"].lower() == "attention"
|
||||
assert telegram_runs == []
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Run tests with: pytest tests/e2e/scenarios/test_routine_event_batch.py -v
|
||||
pytest.main([__file__, "-v", "-s"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_routine_execution_history_is_available(
|
||||
page,
|
||||
ironclaw_server,
|
||||
http_channel_server,
|
||||
):
|
||||
"""Routine run history is exposed by the routines runs API."""
|
||||
routine = await _create_event_routine(
|
||||
page,
|
||||
ironclaw_server,
|
||||
name=f"evt-{uuid.uuid4().hex[:8]}",
|
||||
pattern="history",
|
||||
)
|
||||
|
||||
await _post_http_message(
|
||||
http_channel_server,
|
||||
content="history event",
|
||||
)
|
||||
|
||||
await _wait_for_run_count(
|
||||
ironclaw_server,
|
||||
routine["id"],
|
||||
expected_at_least=1,
|
||||
)
|
||||
completed_run = await _wait_for_completed_run(ironclaw_server, routine["id"])
|
||||
|
||||
assert completed_run["id"]
|
||||
assert completed_run["started_at"]
|
||||
assert completed_run["status"].lower() == "attention"
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user