Merge branch 'staging' into feat/nearai-mcp

# Conflicts:
#	src/app.rs
This commit is contained in:
Coffee
2026-03-17 13:37:41 +08:00
98 changed files with 7765 additions and 2308 deletions
+6 -2
View File
@@ -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_telegram_token_validation.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
+30 -3
View File
@@ -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
+6
View File
@@ -222,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
View File
@@ -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 |
+133 -189
View File
@@ -102,7 +102,6 @@ struct TelegramMessage {
sticker: Option<TelegramSticker>,
/// Forum topic ID. Present when the message is sent inside a forum topic.
/// https://core.telegram.org/bots/api#message
#[serde(default)]
message_thread_id: Option<i64>,
@@ -207,10 +206,6 @@ struct TelegramChat {
/// Title for groups/channels.
title: Option<String>,
/// True when the supergroup has topics (forum mode) enabled.
#[serde(default)]
is_forum: Option<bool>,
/// Username for private chats.
username: Option<String>,
}
@@ -508,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
@@ -697,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), metadata.message_thread_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> {
@@ -734,8 +733,6 @@ impl Guest for TelegramChannel {
"action": "typing"
});
// sendChatAction requires message_thread_id even for the General
// topic (id=1), unlike sendMessage which rejects it.
if let Some(thread_id) = metadata.message_thread_id {
payload["message_thread_id"] = serde_json::Value::Number(thread_id.into());
}
@@ -766,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, metadata.message_thread_id)
{
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!(
@@ -777,7 +778,13 @@ impl Guest for TelegramChannel {
),
);
if let Err(retry_err) = send_message(metadata.chat_id, &prompt, None, None, metadata.message_thread_id) {
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!(
@@ -822,9 +829,8 @@ impl std::fmt::Display for SendError {
/// Normalize `message_thread_id` for outbound API calls.
///
/// Telegram rejects `sendMessage` (and other send methods) when
/// `message_thread_id = 1` (the "General" topic). Return `None` in that
/// case so the field is omitted from the payload.
/// 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)
}
@@ -950,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!(
@@ -992,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,
@@ -1088,7 +1091,14 @@ fn send_photo(
data.len()
),
);
return send_document(chat_id, filename, mime_type, data, reply_to_message_id, message_thread_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());
@@ -1096,10 +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_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());
@@ -1151,10 +1171,20 @@ fn send_document(
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_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());
@@ -1191,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.
///
@@ -1218,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"), message_thread_id) {
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, message_thread_id)
.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()),
}
}
@@ -1392,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(())
@@ -1412,7 +1450,7 @@ fn send_pairing_reply(chat_id: i64, code: &str) -> Result<(), String> {
),
None,
Some("Markdown"),
None, // Pairing happens in DMs, not forum topics
None,
)
.map(|_| ())
.map_err(|e| e.to_string())
@@ -1494,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)),
@@ -1507,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)),
@@ -1520,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)),
@@ -1745,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());
@@ -1830,8 +1865,6 @@ fn handle_message(message: TelegramMessage) {
}
}
let bot_username = channel_host::workspace_read(BOT_USERNAME_PATH).unwrap_or_default();
// For group chats, only respond if bot was mentioned or respond_to_all is enabled
if !is_private {
let respond_to_all = channel_host::workspace_read(RESPOND_TO_ALL_GROUP_PATH)
@@ -1841,6 +1874,7 @@ fn handle_message(message: TelegramMessage) {
if !respond_to_all {
let has_command = content.starts_with('/');
let bot_username = channel_host::workspace_read(BOT_USERNAME_PATH).unwrap_or_default();
let has_bot_mention = if bot_username.is_empty() {
content.contains('@')
} else {
@@ -1876,18 +1910,7 @@ fn handle_message(message: TelegramMessage) {
let metadata_json = serde_json::to_string(&metadata).unwrap_or_else(|_| "{}".to_string());
// Compute thread_id for forum topics: "chat_id:topic_id" to prevent
// collisions across different groups (topic IDs are only unique per chat).
// Only use message_thread_id when the chat is a forum — non-forum groups
// also carry message_thread_id for reply threads, which are not topics.
let thread_id = if message.chat.is_forum == Some(true) {
message.message_thread_id.map(|topic_id| {
format!("{}:{}", message.chat.id, topic_id)
})
} else {
None
};
let bot_username = channel_host::workspace_read(BOT_USERNAME_PATH).unwrap_or_default();
let content_to_emit = match content_to_emit_for_agent(
&content,
if bot_username.is_empty() {
@@ -1907,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,
thread_id: Some(message.chat.id.to_string()),
metadata_json,
attachments,
});
@@ -2507,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]
@@ -2559,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]
@@ -2707,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]
@@ -2726,100 +2766,4 @@ mod tests {
// Verify the constant is 20 MB, matching the Slack channel limit
assert_eq!(MAX_DOWNLOAD_SIZE_BYTES, 20 * 1024 * 1024);
}
// === Forum Topics (thread_id) tests ===
#[test]
fn test_parse_forum_message_with_thread_id() {
let json = r#"{
"message_id": 100,
"message_thread_id": 42,
"is_topic_message": true,
"from": {"id": 1, "is_bot": false, "first_name": "A"},
"chat": {"id": -1001234567890, "type": "supergroup", "is_forum": true},
"text": "Hello from a topic"
}"#;
let msg: TelegramMessage = serde_json::from_str(json).unwrap();
assert_eq!(msg.message_thread_id, Some(42));
assert_eq!(msg.is_topic_message, Some(true));
assert_eq!(msg.chat.is_forum, Some(true));
}
#[test]
fn test_parse_non_forum_message_backward_compat() {
let json = r#"{
"message_id": 1,
"from": {"id": 1, "is_bot": false, "first_name": "A"},
"chat": {"id": 1, "type": "private"},
"text": "Hello"
}"#;
let msg: TelegramMessage = serde_json::from_str(json).unwrap();
assert_eq!(msg.message_thread_id, None);
assert_eq!(msg.is_topic_message, None);
assert_eq!(msg.chat.is_forum, None);
}
#[test]
fn test_metadata_with_message_thread_id() {
let metadata = TelegramMessageMetadata {
chat_id: -1001234567890,
message_id: 100,
user_id: 42,
is_private: false,
message_thread_id: Some(7),
};
let json = serde_json::to_string(&metadata).unwrap();
let parsed: TelegramMessageMetadata = serde_json::from_str(&json).unwrap();
assert_eq!(parsed.message_thread_id, Some(7));
}
#[test]
fn test_metadata_backward_compat_no_thread_id() {
// Old metadata JSON without message_thread_id should deserialize with None
let json = r#"{"chat_id":123,"message_id":1,"user_id":42,"is_private":true}"#;
let metadata: TelegramMessageMetadata = serde_json::from_str(json).unwrap();
assert_eq!(metadata.message_thread_id, None);
}
#[test]
fn test_metadata_thread_id_not_serialized_when_none() {
let metadata = TelegramMessageMetadata {
chat_id: 123,
message_id: 1,
user_id: 42,
is_private: true,
message_thread_id: None,
};
let json = serde_json::to_string(&metadata).unwrap();
assert!(!json.contains("message_thread_id"));
}
#[test]
fn test_thread_id_composition() {
// Verify "chat_id:topic_id" format for forum topics
let chat_id: i64 = -1001234567890;
let topic_id: i64 = 42;
let thread_id = format!("{}:{}", chat_id, topic_id);
assert_eq!(thread_id, "-1001234567890:42");
}
#[test]
fn test_normalize_thread_id_general_topic() {
// General topic (id=1) must be omitted — Telegram rejects sendMessage
// with message_thread_id=1.
assert_eq!(normalize_thread_id(Some(1)), None);
}
#[test]
fn test_normalize_thread_id_regular_topic() {
// Non-General topics pass through unchanged
assert_eq!(normalize_thread_id(Some(42)), Some(42));
assert_eq!(normalize_thread_id(Some(123)), Some(123));
}
#[test]
fn test_normalize_thread_id_none() {
// None stays None
assert_eq!(normalize_thread_id(None), None);
}
}
+7 -7
View File
@@ -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';
+1 -1
View File
@@ -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,
+1 -1
View File
@@ -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": [
+1 -1
View File
@@ -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": [
+1 -1
View File
@@ -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": [
+1 -1
View File
@@ -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": [
+1 -1
View File
@@ -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
View File
@@ -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
}
}
+4 -1
View File
@@ -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);
}
}
+6 -1
View File
@@ -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,
+6 -1
View File
@@ -402,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 {
@@ -431,6 +435,7 @@ impl HeartbeatRunner {
attachments: Vec::new(),
metadata: serde_json::json!({
"source": "heartbeat",
"owner_id": self.workspace.user_id(),
}),
};
+3 -3
View File
@@ -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,
+10 -1
View File
@@ -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,
}),
+8
View File
@@ -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
View File
@@ -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
View File
@@ -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());
@@ -471,10 +478,11 @@ 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();
let companion_mcp_server = companion_mcp_server.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
};
@@ -505,6 +513,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();
@@ -516,7 +525,7 @@ impl AppBuilder {
nearai_api_key,
&pm,
secrets,
"default",
&owner_id,
)
.await
{
@@ -660,7 +669,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(),
companion_mcp_server,
catalog_entries.clone(),
+82 -6
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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);
}
+8 -2
View File
@@ -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");
+3 -1
View File
@@ -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};
+1
View File
@@ -672,6 +672,7 @@ mod tests {
runtime,
prepared,
capabilities,
"default",
"{}".to_string(),
Arc::new(PairingStore::new()),
None,
+38 -10
View File
@@ -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}")
}
File diff suppressed because it is too large Load Diff
+22 -7
View File
@@ -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();
+20 -20
View File
@@ -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
View File
@@ -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))
+175 -7
View File
@@ -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');
}
}
@@ -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.
+7
View File
@@ -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',
+6
View File
@@ -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': '已配置',
+78
View File
@@ -2896,6 +2896,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;
+41 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
}
}
+55 -335
View File
@@ -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);
}
}
+12
View File
@@ -38,6 +38,8 @@ impl LlmConfig {
provider: None,
bedrock: None,
request_timeout_secs: 120,
cheap_model: None,
smart_routing_cascade: false,
}
}
@@ -168,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()
@@ -183,6 +193,8 @@ impl LlmConfig {
provider,
bedrock,
request_timeout_secs,
cheap_model,
smart_routing_cascade,
})
}
+46 -13
View File
@@ -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,16 +303,15 @@ 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(settings)?,
wasm: WasmConfig::resolve(settings)?,
@@ -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
+64 -13
View File
@@ -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))
}
}
+223 -3
View File
@@ -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);
}
}
}
}
+19
View File
@@ -81,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 {
@@ -121,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.
@@ -202,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(),
@@ -233,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,
+1
View File
@@ -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
View File
@@ -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() {
+2 -2
View File
@@ -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,
+72 -2
View File
@@ -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
View File
@@ -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)
}
}
}
+3
View File
@@ -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),
+1513 -68
View File
File diff suppressed because it is too large Load Diff
+13
View File
@@ -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 {
+1
View File
@@ -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"),
+77 -1
View File
@@ -143,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()
@@ -705,4 +707,78 @@ mod tests {
// 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)))
}
}
+24
View File
@@ -138,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.
+121 -28
View File
@@ -378,32 +378,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.
@@ -451,14 +480,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 {
@@ -473,7 +503,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()
},
))
@@ -602,6 +632,8 @@ mod tests {
provider: None,
bedrock: None,
request_timeout_secs: 120,
cheap_model: None,
smart_routing_cascade: true,
}
}
@@ -616,7 +648,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());
@@ -630,7 +662,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());
@@ -639,6 +690,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);
}
}
+2
View File
@@ -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
View File
@@ -215,6 +215,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")
@@ -235,7 +236,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),
@@ -2187,4 +2189,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)))
}
}
+27
View File
@@ -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");
}
}
}
+20 -16
View File
@@ -153,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")))]
@@ -195,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?;
}
@@ -282,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 {
@@ -311,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(),
}));
@@ -618,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.
@@ -703,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,
@@ -775,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 {
@@ -805,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
@@ -821,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
View File
@@ -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") };
}
}
+13
View File
@@ -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)]
@@ -733,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(),
@@ -772,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
View File
File diff suppressed because it is too large Load Diff
+3 -2
View File
@@ -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
View File
@@ -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(&params, "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
);
}
}
+2 -3
View File
@@ -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,
+6 -1
View File
@@ -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
View File
@@ -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;
+179
View File
@@ -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)));
}
}
+2
View File
@@ -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;
+48 -2
View File
@@ -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
View File
@@ -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 (520s) 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
View File
@@ -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
View File
@@ -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:
+24
View File
@@ -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}",
}
+7 -1
View File
@@ -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
+62
View File
@@ -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 _: {},
),
]
+30
View File
@@ -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")
+226
View File
@@ -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,
)
+294 -511
View File
@@ -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"
@@ -0,0 +1,243 @@
"""Telegram hot-activation UI coverage."""
import asyncio
import json
from helpers import SEL
_CONFIGURE_SECRET_INPUT = "input[type='password']"
_CONFIGURE_SAVE_BUTTON = ".configure-actions button.btn-ext.activate"
_TELEGRAM_INSTALLED = {
"name": "telegram",
"display_name": "Telegram",
"kind": "wasm_channel",
"description": "Telegram Bot API channel",
"url": None,
"active": False,
"authenticated": False,
"has_auth": False,
"needs_setup": True,
"tools": [],
"activation_status": "installed",
"activation_error": None,
}
_TELEGRAM_ACTIVE = {
**_TELEGRAM_INSTALLED,
"active": True,
"authenticated": True,
"needs_setup": False,
"activation_status": "active",
}
async def go_to_extensions(page):
await page.locator(SEL["tab_button"].format(tab="extensions")).click()
await page.locator(SEL["tab_panel"].format(tab="extensions")).wait_for(
state="visible", timeout=5000
)
await page.locator(
f"{SEL['extensions_list']} .empty-state, {SEL['ext_card_installed']}"
).first.wait_for(state="visible", timeout=8000)
async def mock_extension_lists(page, ext_handler):
async def handle_ext_list(route):
path = route.request.url.split("?")[0]
if path.endswith("/api/extensions"):
await ext_handler(route)
else:
await route.continue_()
async def handle_tools(route):
await route.fulfill(
status=200,
content_type="application/json",
body=json.dumps({"tools": []}),
)
async def handle_registry(route):
await route.fulfill(
status=200,
content_type="application/json",
body=json.dumps({"entries": []}),
)
# Register the broad route first so the specific endpoints below win.
await page.route("**/api/extensions*", handle_ext_list)
await page.route("**/api/extensions/tools", handle_tools)
await page.route("**/api/extensions/registry", handle_registry)
async def wait_for_toast(page, text: str, *, timeout: int = 5000):
await page.locator(SEL["toast"], has_text=text).wait_for(
state="visible", timeout=timeout
)
async def test_telegram_setup_modal_shows_bot_token_field(page):
async def handle_ext_list(route):
await route.fulfill(
status=200,
content_type="application/json",
body=json.dumps({"extensions": [_TELEGRAM_INSTALLED]}),
)
async def handle_setup(route):
await route.fulfill(
status=200,
content_type="application/json",
body=json.dumps(
{
"secrets": [
{
"name": "telegram_bot_token",
"prompt": "Enter your Telegram Bot API token (from @BotFather)",
"provided": False,
"optional": False,
"auto_generate": False,
}
]
}
),
)
await mock_extension_lists(page, handle_ext_list)
await page.route("**/api/extensions/telegram/setup", handle_setup)
await go_to_extensions(page)
card = page.locator(SEL["ext_card_installed"]).first
await card.locator(SEL["ext_configure_btn"], has_text="Setup").click()
modal = page.locator(SEL["configure_modal"])
await modal.wait_for(state="visible", timeout=5000)
assert "Telegram Bot API token" in await modal.text_content()
assert "IronClaw will show a one-time code" in (
await modal.text_content()
)
input_el = modal.locator(_CONFIGURE_SECRET_INPUT)
assert await input_el.count() == 1
async def test_telegram_hot_activation_transitions_installed_to_active(page):
phase = {"value": "installed"}
captured_setup_payloads = []
post_count = {"value": 0}
second_request_started = asyncio.Event()
allow_second_response = asyncio.Event()
async def handle_ext_list(route):
extensions = {
"installed": [_TELEGRAM_INSTALLED],
"active": [_TELEGRAM_ACTIVE],
}[phase["value"]]
await route.fulfill(
status=200,
content_type="application/json",
body=json.dumps({"extensions": extensions}),
)
async def handle_setup(route):
if route.request.method == "GET":
await route.fulfill(
status=200,
content_type="application/json",
body=json.dumps(
{
"secrets": [
{
"name": "telegram_bot_token",
"prompt": "Enter your Telegram Bot API token (from @BotFather)",
"provided": False,
"optional": False,
"auto_generate": False,
}
]
}
),
)
return
payload = json.loads(route.request.post_data or "{}")
captured_setup_payloads.append(payload)
post_count["value"] += 1
await asyncio.sleep(0.05)
if post_count["value"] == 1:
await route.fulfill(
status=200,
content_type="application/json",
body=json.dumps(
{
"success": True,
"activated": False,
"message": "Configuration saved for 'telegram'. Send `/start iclaw-7qk2m9` to @test_hot_bot in Telegram. IronClaw will finish setup automatically.",
"verification": {
"code": "iclaw-7qk2m9",
"instructions": "Send `/start iclaw-7qk2m9` to @test_hot_bot in Telegram. IronClaw will finish setup automatically.",
"deep_link": "https://t.me/test_hot_bot?start=iclaw-7qk2m9",
},
}
),
)
else:
second_request_started.set()
await allow_second_response.wait()
await route.fulfill(
status=200,
content_type="application/json",
body=json.dumps(
{
"success": True,
"activated": True,
"message": "Configuration saved, Telegram owner verified, and 'telegram' activated. Hot-activated WASM channel",
}
),
)
await mock_extension_lists(page, handle_ext_list)
await page.route("**/api/extensions/telegram/setup", handle_setup)
await go_to_extensions(page)
card = page.locator(SEL["ext_card_installed"]).first
await card.locator(SEL["ext_configure_btn"], has_text="Setup").click()
modal = page.locator(SEL["configure_modal"])
await modal.wait_for(state="visible", timeout=5000)
await modal.locator(_CONFIGURE_SECRET_INPUT).fill("123456789:ABCdefGhI")
await modal.locator(_CONFIGURE_SAVE_BUTTON).click()
await second_request_started.wait()
await modal.locator(".configure-inline-status", has_text="Waiting for Telegram owner verification...").wait_for(
state="visible", timeout=5000
)
assert "iclaw-7qk2m9" in (await modal.text_content())
assert "/start iclaw-7qk2m9" in (await modal.text_content())
assert await modal.locator(".configure-verification-link").count() == 1
await modal.locator(_CONFIGURE_SAVE_BUTTON).wait_for(state="hidden", timeout=5000)
await page.locator(SEL["configure_overlay"]).click(position={"x": 1, "y": 1})
assert await page.locator(SEL["configure_overlay"]).is_visible()
allow_second_response.set()
await page.locator(SEL["configure_overlay"]).wait_for(state="hidden", timeout=5000)
phase["value"] = "active"
await page.evaluate(
"""
handleAuthCompleted({
extension_name: 'telegram',
success: true,
message: "Configuration saved, Telegram owner verified, and 'telegram' activated. Hot-activated WASM channel",
});
"""
)
await wait_for_toast(page, "Telegram owner verified")
await card.locator(SEL["ext_active_label"]).wait_for(state="visible", timeout=5000)
assert await card.locator(SEL["ext_pairing_label"]).count() == 0
assert captured_setup_payloads == [
{"secrets": {"telegram_bot_token": "123456789:ABCdefGhI"}},
{"secrets": {}},
]
+56
View File
@@ -130,3 +130,59 @@ async def test_approval_params_toggle(page):
await toggle.click()
await page.wait_for_timeout(300)
assert await params.is_hidden(), "Parameters should be hidden after second toggle"
async def test_waiting_for_approval_message_no_error_prefix(page):
"""Verify that input submitted while awaiting approval shows non-error status with tool context.
Trigger a real approval-needed tool call, then attempt to send another message while
approval is pending. The backend should reject the second input with a non-error
status that includes the pending tool context.
"""
assistant_messages = page.locator(SEL["message_assistant"])
chat_input = page.locator(SEL["chat_input"])
await chat_input.wait_for(state="visible", timeout=5000)
# Trigger a real HTTP tool call that pauses for approval in the default E2E harness.
await chat_input.fill("make approval post approval-required")
await chat_input.press("Enter")
card = page.locator(SEL["approval_card"]).last
await card.wait_for(state="visible", timeout=10000)
tool_name = await card.locator(".approval-tool-name").text_content()
desc_text = await card.locator(".approval-description").text_content()
assert tool_name == "http"
assert desc_text is not None and "HTTP requests to external APIs" in desc_text
# With the thread now genuinely awaiting approval, the next message should be rejected
# as a non-error pending status.
initial_count = await assistant_messages.count()
await chat_input.fill("send another message now")
await chat_input.press("Enter")
await page.wait_for_function(
f"() => document.querySelectorAll('{SEL['message_assistant']}').length > {initial_count}",
timeout=10000,
)
last_msg = assistant_messages.last.locator(".message-content")
msg_text = await last_msg.inner_text()
# Verify no "Error:" prefix
assert not msg_text.lower().startswith("error:"), (
f"Approval rejection must NOT have 'Error:' prefix. Got: {msg_text!r}"
)
# Verify it contains "waiting for approval"
assert "waiting for approval" in msg_text.lower(), (
f"Expected 'Waiting for approval' text. Got: {msg_text!r}"
)
# Verify it contains the tool name and description
assert "http" in msg_text.lower(), (
f"Expected tool name 'http' in message. Got: {msg_text!r}"
)
assert "HTTP requests to external APIs" in msg_text, (
f"Expected tool description in message. Got: {msg_text!r}"
)
+118 -255
View File
@@ -7,7 +7,7 @@ import json
import httpx
import pytest
from helpers import AUTH_TOKEN
from helpers import HTTP_WEBHOOK_SECRET
def compute_signature(secret: str, body: bytes) -> str:
@@ -16,325 +16,188 @@ def compute_signature(secret: str, body: bytes) -> str:
return f"sha256={mac.hexdigest()}"
@pytest.mark.asyncio
async def test_webhook_requires_http_webhook_secret_configured(ironclaw_server):
"""
Webhook endpoint rejects requests when HTTP_WEBHOOK_SECRET is not configured.
This tests the fail-closed security posture.
"""
headers = {"Authorization": f"Bearer {AUTH_TOKEN}"}
async def _post_webhook(
base_url: str,
body_data: dict,
*,
signature: str | None = None,
content_type: str = "application/json",
) -> httpx.Response:
"""Send a raw webhook request with optional signature."""
body_bytes = json.dumps(body_data).encode()
headers = {"Content-Type": content_type}
if signature is not None:
headers["X-Hub-Signature-256"] = signature
async with httpx.AsyncClient() as client:
# When no webhook secret is configured on the server, all requests fail
r = await client.post(
f"{ironclaw_server}/webhook",
json={"content": "test message"},
return await client.post(
f"{base_url}/webhook",
content=body_bytes,
headers=headers,
)
# Server should reject with 503 Service Unavailable (fail closed)
assert r.status_code in (401, 503)
@pytest.mark.asyncio
async def test_webhook_hmac_signature_valid(ironclaw_server_with_webhook_secret):
async def test_webhook_requires_http_webhook_secret_configured(
http_channel_server_without_secret,
):
"""Webhook fails closed when no secret is configured."""
response = await _post_webhook(
http_channel_server_without_secret,
{"content": "test message"},
)
assert response.status_code == 503
data = response.json()
assert data["status"] == "error"
assert "Webhook authentication not configured" in data.get("response", "")
@pytest.mark.asyncio
async def test_webhook_hmac_signature_valid(http_channel_server):
"""Valid X-Hub-Signature-256 HMAC signature is accepted."""
secret = ironclaw_server_with_webhook_secret["secret"]
base_url = ironclaw_server_with_webhook_secret["url"]
body = {"content": "hello from webhook"}
signature = compute_signature(HTTP_WEBHOOK_SECRET, json.dumps(body).encode())
headers = {"Authorization": f"Bearer {AUTH_TOKEN}"}
body_data = {"content": "hello from webhook"}
body_bytes = json.dumps(body_data).encode()
signature = compute_signature(secret, body_bytes)
response = await _post_webhook(http_channel_server, body, signature=signature)
async with httpx.AsyncClient() as client:
r = await client.post(
f"{base_url}/webhook",
content=body_bytes,
headers={
**headers,
"Content-Type": "application/json",
"X-Hub-Signature-256": signature,
},
)
assert r.status_code == 200, f"Expected 200, got {r.status_code}: {r.text}"
resp = r.json()
assert resp["status"] == "ok"
assert response.status_code == 200, (
f"Expected 200, got {response.status_code}: {response.text}"
)
data = response.json()
assert data["status"] == "accepted"
@pytest.mark.asyncio
async def test_webhook_invalid_hmac_signature_rejected(
ironclaw_server_with_webhook_secret,
):
async def test_webhook_invalid_hmac_signature_rejected(http_channel_server):
"""Invalid X-Hub-Signature-256 signature is rejected with 401."""
base_url = ironclaw_server_with_webhook_secret["url"]
response = await _post_webhook(
http_channel_server,
{"content": "hello"},
signature="sha256=0000000000000000000000000000000000000000000000000000000000000000",
)
headers = {"Authorization": f"Bearer {AUTH_TOKEN}"}
body_data = {"content": "hello"}
body_bytes = json.dumps(body_data).encode()
invalid_signature = "sha256=0000000000000000000000000000000000000000000000000000000000000000"
async with httpx.AsyncClient() as client:
r = await client.post(
f"{base_url}/webhook",
content=body_bytes,
headers={
**headers,
"Content-Type": "application/json",
"X-Hub-Signature-256": invalid_signature,
},
)
assert r.status_code == 401, f"Expected 401, got {r.status_code}"
resp = r.json()
assert resp["status"] == "error"
assert "Invalid webhook signature" in resp.get("response", "")
assert response.status_code == 401
data = response.json()
assert data["status"] == "error"
assert "Invalid webhook signature" in data.get("response", "")
@pytest.mark.asyncio
async def test_webhook_wrong_secret_rejected(ironclaw_server_with_webhook_secret):
async def test_webhook_wrong_secret_rejected(http_channel_server):
"""Signature computed with wrong secret is rejected."""
base_url = ironclaw_server_with_webhook_secret["url"]
body = {"content": "hello"}
signature = compute_signature("wrong-secret", json.dumps(body).encode())
headers = {"Authorization": f"Bearer {AUTH_TOKEN}"}
body_data = {"content": "hello"}
body_bytes = json.dumps(body_data).encode()
# Compute signature with wrong secret
wrong_signature = compute_signature("wrong-secret", body_bytes)
response = await _post_webhook(http_channel_server, body, signature=signature)
async with httpx.AsyncClient() as client:
r = await client.post(
f"{base_url}/webhook",
content=body_bytes,
headers={
**headers,
"Content-Type": "application/json",
"X-Hub-Signature-256": wrong_signature,
},
)
assert r.status_code == 401
resp = r.json()
assert resp["status"] == "error"
assert response.status_code == 401
assert response.json()["status"] == "error"
@pytest.mark.asyncio
async def test_webhook_malformed_signature_rejected(
ironclaw_server_with_webhook_secret,
):
"""Malformed X-Hub-Signature-256 header is rejected."""
base_url = ironclaw_server_with_webhook_secret["url"]
async def test_webhook_missing_signature_header_rejected(http_channel_server):
"""Missing X-Hub-Signature-256 header is rejected when no body secret is provided."""
response = await _post_webhook(http_channel_server, {"content": "hello"})
headers = {"Authorization": f"Bearer {AUTH_TOKEN}"}
body_data = {"content": "hello"}
body_bytes = json.dumps(body_data).encode()
async with httpx.AsyncClient() as client:
# Missing sha256= prefix
r = await client.post(
f"{base_url}/webhook",
content=body_bytes,
headers={
**headers,
"Content-Type": "application/json",
"X-Hub-Signature-256": "deadbeef",
},
)
assert r.status_code == 401
assert response.status_code == 401
data = response.json()
assert "Webhook authentication required" in data.get("response", "")
assert "X-Hub-Signature-256" in data.get("response", "")
@pytest.mark.asyncio
async def test_webhook_missing_signature_header_rejected(
ironclaw_server_with_webhook_secret,
):
"""Missing X-Hub-Signature-256 header is rejected when no body secret provided."""
base_url = ironclaw_server_with_webhook_secret["url"]
async def test_webhook_deprecated_body_secret_still_works(http_channel_server):
"""Deprecated body secret support still accepts old clients."""
response = await _post_webhook(
http_channel_server,
{"content": "hello", "secret": HTTP_WEBHOOK_SECRET},
)
headers = {"Authorization": f"Bearer {AUTH_TOKEN}"}
body_data = {"content": "hello"}
body_bytes = json.dumps(body_data).encode()
async with httpx.AsyncClient() as client:
# No X-Hub-Signature-256 header and no body secret
r = await client.post(
f"{base_url}/webhook",
content=body_bytes,
headers={
**headers,
"Content-Type": "application/json",
},
)
assert r.status_code == 401
resp = r.json()
assert "Webhook authentication required" in resp.get("response", "")
assert "X-Hub-Signature-256" in resp.get("response", "")
assert response.status_code == 200, (
f"Expected 200, got {response.status_code}: {response.text}"
)
assert response.json()["status"] == "accepted"
@pytest.mark.asyncio
async def test_webhook_deprecated_body_secret_still_works(
ironclaw_server_with_webhook_secret,
):
"""
Deprecated: body 'secret' field still works for backward compatibility.
This test ensures we don't break existing clients during the migration period.
"""
secret = ironclaw_server_with_webhook_secret["secret"]
base_url = ironclaw_server_with_webhook_secret["url"]
async def test_webhook_header_takes_precedence_over_body_secret(http_channel_server):
"""Header signature wins when both header and body secret are provided."""
body = {"content": "hello", "secret": "wrong-secret-in-body"}
signature = compute_signature(HTTP_WEBHOOK_SECRET, json.dumps(body).encode())
headers = {"Authorization": f"Bearer {AUTH_TOKEN}"}
# Old-style request with secret in body
body_data = {"content": "hello", "secret": secret}
body_bytes = json.dumps(body_data).encode()
response = await _post_webhook(http_channel_server, body, signature=signature)
async with httpx.AsyncClient() as client:
r = await client.post(
f"{base_url}/webhook",
content=body_bytes,
headers={
**headers,
"Content-Type": "application/json",
},
)
# Should succeed (backward compatibility)
assert r.status_code == 200, f"Expected 200, got {r.status_code}: {r.text}"
resp = r.json()
assert resp["status"] == "ok"
assert response.status_code == 200
assert response.json()["status"] == "accepted"
@pytest.mark.asyncio
async def test_webhook_header_takes_precedence_over_body_secret(
ironclaw_server_with_webhook_secret,
):
"""
When both X-Hub-Signature-256 header and body secret are provided,
header takes precedence.
"""
secret = ironclaw_server_with_webhook_secret["secret"]
base_url = ironclaw_server_with_webhook_secret["url"]
headers = {"Authorization": f"Bearer {AUTH_TOKEN}"}
body_data = {"content": "hello", "secret": "wrong-secret-in-body"}
body_bytes = json.dumps(body_data).encode()
# Compute signature with correct secret
signature = compute_signature(secret, body_bytes)
async def test_webhook_case_insensitive_header_lookup(http_channel_server):
"""HTTP headers are treated case-insensitively."""
body = {"content": "hello"}
body_bytes = json.dumps(body).encode()
signature = compute_signature(HTTP_WEBHOOK_SECRET, body_bytes)
async with httpx.AsyncClient() as client:
r = await client.post(
f"{base_url}/webhook",
response = await client.post(
f"{http_channel_server}/webhook",
content=body_bytes,
headers={
**headers,
"Content-Type": "application/json",
"X-Hub-Signature-256": signature,
},
)
# Should succeed because header signature is valid (takes precedence)
assert r.status_code == 200
resp = r.json()
assert resp["status"] == "ok"
@pytest.mark.asyncio
async def test_webhook_case_insensitive_header_lookup(
ironclaw_server_with_webhook_secret,
):
"""HTTP headers are case-insensitive. Test with different cases."""
secret = ironclaw_server_with_webhook_secret["secret"]
base_url = ironclaw_server_with_webhook_secret["url"]
headers = {"Authorization": f"Bearer {AUTH_TOKEN}"}
body_data = {"content": "hello"}
body_bytes = json.dumps(body_data).encode()
signature = compute_signature(secret, body_bytes)
async with httpx.AsyncClient() as client:
# Try with lowercase
r = await client.post(
f"{base_url}/webhook",
content=body_bytes,
headers={
**headers,
"Content-Type": "application/json",
"x-hub-signature-256": signature,
},
)
assert r.status_code == 200
assert response.status_code == 200
@pytest.mark.asyncio
async def test_webhook_wrong_content_type_rejected(
ironclaw_server_with_webhook_secret,
):
async def test_webhook_wrong_content_type_rejected(http_channel_server):
"""Webhook only accepts application/json Content-Type."""
secret = ironclaw_server_with_webhook_secret["secret"]
base_url = ironclaw_server_with_webhook_secret["url"]
body = {"content": "hello"}
signature = compute_signature(HTTP_WEBHOOK_SECRET, json.dumps(body).encode())
headers = {"Authorization": f"Bearer {AUTH_TOKEN}"}
body_data = {"content": "hello"}
body_bytes = json.dumps(body_data).encode()
signature = compute_signature(secret, body_bytes)
response = await _post_webhook(
http_channel_server,
body,
signature=signature,
content_type="text/plain",
)
async with httpx.AsyncClient() as client:
r = await client.post(
f"{base_url}/webhook",
content=body_bytes,
headers={
**headers,
"Content-Type": "text/plain",
"X-Hub-Signature-256": signature,
},
)
assert r.status_code == 415 # Unsupported Media Type
resp = r.json()
assert "application/json" in resp.get("response", "")
assert response.status_code == 415
assert "application/json" in response.json().get("response", "")
@pytest.mark.asyncio
async def test_webhook_invalid_json_rejected(ironclaw_server_with_webhook_secret):
async def test_webhook_invalid_json_rejected(http_channel_server):
"""Invalid JSON in body is rejected."""
secret = ironclaw_server_with_webhook_secret["secret"]
base_url = ironclaw_server_with_webhook_secret["url"]
headers = {"Authorization": f"Bearer {AUTH_TOKEN}"}
body_bytes = b"not valid json"
signature = compute_signature(secret, body_bytes)
signature = compute_signature(HTTP_WEBHOOK_SECRET, body_bytes)
async with httpx.AsyncClient() as client:
r = await client.post(
f"{base_url}/webhook",
response = await client.post(
f"{http_channel_server}/webhook",
content=body_bytes,
headers={
**headers,
"Content-Type": "application/json",
"X-Hub-Signature-256": signature,
},
)
assert r.status_code == 401 or r.status_code == 400
assert response.status_code in (400, 401)
@pytest.mark.asyncio
async def test_webhook_message_queued_for_processing(
ironclaw_server_with_webhook_secret,
):
"""Message via webhook is queued and can be retrieved."""
secret = ironclaw_server_with_webhook_secret["secret"]
base_url = ironclaw_server_with_webhook_secret["url"]
async def test_webhook_message_queued_for_processing(http_channel_server):
"""Accepted webhook requests return a real message id."""
body = {"content": "webhook test message 12345"}
signature = compute_signature(HTTP_WEBHOOK_SECRET, json.dumps(body).encode())
headers = {"Authorization": f"Bearer {AUTH_TOKEN}"}
test_message = "webhook test message 12345"
body_data = {"content": test_message}
body_bytes = json.dumps(body_data).encode()
signature = compute_signature(secret, body_bytes)
response = await _post_webhook(http_channel_server, body, signature=signature)
async with httpx.AsyncClient() as client:
r = await client.post(
f"{base_url}/webhook",
content=body_bytes,
headers={
**headers,
"Content-Type": "application/json",
"X-Hub-Signature-256": signature,
},
)
assert r.status_code == 200
resp = r.json()
assert resp["status"] == "ok"
# Message ID should be present
assert "message_id" in resp
assert resp["message_id"] != "00000000-0000-0000-0000-000000000000"
assert response.status_code == 200
data = response.json()
assert data["status"] == "accepted"
assert "message_id" in data
assert data["message_id"] != "00000000-0000-0000-0000-000000000000"
+33 -1
View File
@@ -442,6 +442,9 @@ mod advanced {
other => panic!("expected event trigger, got {other:?}"),
}
rig.clear().await;
let llm_calls_before = rig.llm_call_count();
rig.send_incoming(IncomingMessage::new(
"telegram",
"test-user",
@@ -451,8 +454,18 @@ mod advanced {
let runs = wait_for_routine_run(rig.database(), routine.id, TIMEOUT).await;
assert_eq!(runs[0].trigger_type, "event");
assert_eq!(
rig.llm_call_count(),
llm_calls_before + 1,
"matching event message should only trigger the routine LLM call"
);
let responses = rig.wait_for_responses(3, TIMEOUT).await;
let responses = rig.wait_for_responses(1, TIMEOUT).await;
assert_eq!(
responses.len(),
1,
"expected only the routine notification after the matching event"
);
assert!(
responses.iter().any(|response| {
response
@@ -505,6 +518,9 @@ mod advanced {
other => panic!("expected event trigger, got {other:?}"),
}
rig.clear().await;
let llm_calls_before = rig.llm_call_count();
rig.send_incoming(IncomingMessage::new(
"telegram",
"test-user",
@@ -514,6 +530,22 @@ mod advanced {
let runs = wait_for_routine_run(rig.database(), routine.id, TIMEOUT).await;
assert_eq!(runs[0].trigger_type, "event");
assert_eq!(
rig.llm_call_count(),
llm_calls_before + 1,
"matching event message should only trigger the routine LLM call"
);
let responses = rig.wait_for_responses(1, TIMEOUT).await;
assert_eq!(
responses.len(),
1,
"expected only the routine notification after the matching event"
);
assert!(
responses[0].content.contains("Bug report detected"),
"expected routine notification, got: {responses:?}"
);
rig.shutdown();
}
+1 -1
View File
@@ -155,7 +155,7 @@ mod tests {
}
assert_eq!(routine.notify.channel.as_deref(), Some("telegram"));
assert_eq!(routine.notify.user, "ops-team");
assert_eq!(routine.notify.user.as_deref(), Some("ops-team"));
assert_eq!(routine.guardrails.cooldown.as_secs(), 600);
rig.shutdown();
+245 -4
View File
@@ -48,6 +48,19 @@ mod tests {
Arc::new(Workspace::new_with_db("default", db.clone()))
}
fn make_message(
channel: &str,
user_id: &str,
owner_id: &str,
sender_id: &str,
content: &str,
) -> IncomingMessage {
IncomingMessage::new(channel, user_id, content)
.with_owner_id(owner_id)
.with_sender_id(sender_id)
.with_metadata(serde_json::json!({}))
}
/// Helper to insert a routine directly into the database.
fn make_routine(name: &str, trigger: Trigger, prompt: &str) -> Routine {
Routine {
@@ -218,7 +231,13 @@ mod tests {
engine.refresh_event_cache().await;
// Positive match: message containing "deploy to production".
let matching_msg = IncomingMessage::new("test", "default", "deploy to production now");
let matching_msg = make_message(
"test",
"default",
"default",
"default",
"deploy to production now",
);
let fired = engine.check_event_triggers(&matching_msg).await;
assert!(
fired >= 1,
@@ -229,12 +248,114 @@ mod tests {
tokio::time::sleep(Duration::from_millis(500)).await;
// Negative match: message that doesn't match.
let non_matching_msg =
IncomingMessage::new("test", "default", "check the staging environment");
let non_matching_msg = make_message(
"test",
"default",
"default",
"default",
"check the staging environment",
);
let fired_neg = engine.check_event_triggers(&non_matching_msg).await;
assert_eq!(fired_neg, 0, "Expected 0 routines fired on non-match");
}
#[tokio::test]
async fn event_trigger_respects_message_user_scope() {
let (db, _tmp) = create_test_db().await;
let ws = create_workspace(&db);
let trace = LlmTrace::single_turn(
"test-event-user-scope",
"deploy",
vec![TraceStep {
request_hint: None,
response: TraceResponse::Text {
content: "Owner event handled".to_string(),
input_tokens: 50,
output_tokens: 8,
},
expected_tool_results: vec![],
}],
);
let llm = Arc::new(TraceLlm::from_trace(trace));
let (notify_tx, _notify_rx) = tokio::sync::mpsc::channel(16);
let tools = Arc::new(ToolRegistry::new());
let safety = Arc::new(SafetyLayer::new(&SafetyConfig {
max_output_length: 100_000,
injection_check_enabled: true,
}));
let engine = Arc::new(RoutineEngine::new(
RoutineConfig::default(),
db.clone(),
llm,
ws,
notify_tx,
None,
tools,
safety,
));
let routine = make_routine(
"owner-deploy-watcher",
Trigger::Event {
channel: None,
pattern: "deploy.*production".to_string(),
},
"Report on deployment.",
);
db.create_routine(&routine).await.expect("create_routine");
engine.refresh_event_cache().await;
let guest_msg = make_message(
"telegram",
"guest",
"default",
"guest-sender",
"deploy to production now",
);
let guest_fired = engine.check_event_triggers(&guest_msg).await;
assert_eq!(
guest_fired, 0,
"Guest scope must not fire owner event routines"
);
tokio::time::sleep(Duration::from_millis(200)).await;
let guest_runs = db
.list_routine_runs(routine.id, 10)
.await
.expect("list_routine_runs after guest message");
assert!(
guest_runs.is_empty(),
"Guest message should not create routine runs"
);
let owner_msg = make_message(
"telegram",
"default",
"default",
"owner-sender",
"deploy to production now",
);
let owner_fired = engine.check_event_triggers(&owner_msg).await;
assert!(
owner_fired >= 1,
"Owner scope should fire matching owner event routine"
);
tokio::time::sleep(Duration::from_millis(500)).await;
let owner_runs = db
.list_routine_runs(routine.id, 10)
.await
.expect("list_routine_runs after owner message");
assert_eq!(
owner_runs.len(),
1,
"Owner message should create exactly one run"
);
}
// -----------------------------------------------------------------------
// Test 3: system_event_trigger_matches_and_filters
// -----------------------------------------------------------------------
@@ -434,7 +555,13 @@ mod tests {
engine.refresh_event_cache().await;
// First fire should work.
let msg = IncomingMessage::new("test", "default", "test-cooldown trigger");
let msg = make_message(
"test",
"default",
"default",
"default",
"test-cooldown trigger",
);
let fired1 = engine.check_event_triggers(&msg).await;
assert!(fired1 >= 1, "First fire should work");
@@ -553,4 +680,118 @@ mod tests {
"Expected Skipped for empty checklist, got: {result:?}"
);
}
/// Helper to set up a test environment for routine engine mutation tests.
/// Returns the engine, database, and temp directory.
async fn setup_routine_mutation_test()
-> (Arc<RoutineEngine>, Arc<dyn Database>, tempfile::TempDir) {
let (db, dir) = create_test_db().await;
let ws = create_workspace(&db);
let (notify_tx, _rx) = tokio::sync::mpsc::channel(16);
let tools = Arc::new(ToolRegistry::new());
let safety_config = SafetyConfig {
max_output_length: 100_000,
injection_check_enabled: true,
};
let safety = Arc::new(SafetyLayer::new(&safety_config));
let trace = LlmTrace::single_turn(
"test-routine-mutation",
"test",
vec![TraceStep {
request_hint: None,
response: TraceResponse::Text {
content: "ROUTINE_OK".to_string(),
input_tokens: 50,
output_tokens: 5,
},
expected_tool_results: vec![],
}],
);
let llm = Arc::new(TraceLlm::from_trace(trace));
let engine = Arc::new(RoutineEngine::new(
RoutineConfig::default(),
Arc::clone(&db),
llm,
ws,
notify_tx,
None,
tools,
safety,
));
(engine, db, dir)
}
/// Regression test for issue #1076: disabling an event routine via a DB mutation
/// followed by refresh_event_cache() (the path now taken by the web toggle handler)
/// must immediately stop the routine from firing.
#[tokio::test]
async fn toggle_disabling_event_routine_removes_from_cache() {
let (engine, db, _dir) = setup_routine_mutation_test().await;
// Create and cache an event routine.
let mut routine = make_routine(
"disable-me",
Trigger::Event {
pattern: "DISABLE_ME".to_string(),
channel: None,
},
"Handle DISABLE_ME event",
);
db.create_routine(&routine).await.expect("create_routine");
engine.refresh_event_cache().await;
let msg = IncomingMessage::new("test", "default", "DISABLE_ME");
let fired_before = engine.check_event_triggers(&msg).await;
assert!(fired_before >= 1, "Expected routine to fire before disable");
// Simulate what routines_toggle_handler now does: update DB, then refresh.
routine.enabled = false;
routine.updated_at = Utc::now();
db.update_routine(&routine).await.expect("update_routine");
engine.refresh_event_cache().await;
let fired_after = engine.check_event_triggers(&msg).await;
assert_eq!(
fired_after, 0,
"Disabled routine must not fire after cache refresh"
);
}
/// Regression test for issue #1076: deleting an event routine via a DB mutation
/// followed by refresh_event_cache() must immediately stop the routine from firing.
#[tokio::test]
async fn delete_event_routine_removes_from_cache() {
let (engine, db, _dir) = setup_routine_mutation_test().await;
let routine = make_routine(
"delete-me",
Trigger::Event {
pattern: "DELETE_ME".to_string(),
channel: None,
},
"Handle DELETE_ME event",
);
db.create_routine(&routine).await.expect("create_routine");
engine.refresh_event_cache().await;
let msg = IncomingMessage::new("test", "default", "DELETE_ME");
assert!(
engine.check_event_triggers(&msg).await >= 1,
"Expected routine to fire before delete"
);
// Simulate what routines_delete_handler now does: delete from DB, then refresh.
db.delete_routine(routine.id).await.expect("delete_routine");
engine.refresh_event_cache().await;
assert_eq!(
engine.check_event_triggers(&msg).await,
0,
"Deleted routine must not fire after cache refresh"
);
}
}
+353
View File
@@ -0,0 +1,353 @@
//! E2E tests for Telegram message routing through the real agent + message tool.
#[cfg(feature = "libsql")]
mod support;
#[cfg(feature = "libsql")]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use futures::StreamExt;
use ironclaw::agent::{Agent, AgentDeps};
use ironclaw::app::{AppBuilder, AppBuilderFlags};
use ironclaw::channels::web::log_layer::LogBroadcaster;
use ironclaw::channels::{
Channel, ChannelManager, IncomingMessage, MessageStream, OutgoingResponse, StatusUpdate,
};
use ironclaw::config::Config;
use ironclaw::db::{Database, libsql::LibSqlBackend};
use ironclaw::error::ChannelError;
use ironclaw::llm::{LlmProvider, SessionConfig, SessionManager};
use tokio::sync::{Mutex, mpsc};
use tokio_stream::wrappers::ReceiverStream;
use crate::support::test_channel::{TestChannel, TestChannelHandle};
use crate::support::trace_llm::{LlmTrace, TraceLlm, TraceResponse, TraceStep, TraceToolCall};
type TelegramCaptures = Arc<Mutex<Vec<(String, OutgoingResponse)>>>;
struct RecordingTelegramChannel {
captures: TelegramCaptures,
}
impl RecordingTelegramChannel {
fn new() -> (Self, TelegramCaptures) {
let captures = Arc::new(Mutex::new(Vec::new()));
(
Self {
captures: Arc::clone(&captures),
},
captures,
)
}
}
#[async_trait]
impl Channel for RecordingTelegramChannel {
fn name(&self) -> &str {
"telegram"
}
async fn start(&self) -> Result<MessageStream, ChannelError> {
let (_tx, rx) = mpsc::channel::<IncomingMessage>(1);
Ok(ReceiverStream::new(rx).boxed())
}
async fn respond(
&self,
_msg: &IncomingMessage,
response: OutgoingResponse,
) -> Result<(), ChannelError> {
self.captures
.lock()
.await
.push(("respond".to_string(), response));
Ok(())
}
async fn send_status(
&self,
_status: StatusUpdate,
_metadata: &serde_json::Value,
) -> Result<(), ChannelError> {
Ok(())
}
async fn broadcast(
&self,
user_id: &str,
response: OutgoingResponse,
) -> Result<(), ChannelError> {
self.captures
.lock()
.await
.push((user_id.to_string(), response));
Ok(())
}
async fn health_check(&self) -> Result<(), ChannelError> {
Ok(())
}
}
struct Harness {
gateway: Arc<TestChannel>,
telegram_captures: Arc<Mutex<Vec<(String, OutgoingResponse)>>>,
db: Arc<dyn Database>,
owner_id: String,
_temp_dir: tempfile::TempDir,
agent_handle: Option<tokio::task::JoinHandle<()>>,
}
impl Harness {
async fn store_telegram_owner_binding(&self, owner_id: i64) {
for scope in [&self.owner_id, "test-user"] {
self.db
.set_setting(
scope,
"channels.wasm_channel_owner_ids.telegram",
&serde_json::json!(owner_id),
)
.await
.expect("failed to store telegram owner binding");
}
}
async fn wait_for_telegram_broadcasts(
&self,
expected: usize,
timeout: Duration,
) -> Vec<(String, OutgoingResponse)> {
let deadline = tokio::time::Instant::now() + timeout;
loop {
let snapshot = self.telegram_captures.lock().await.clone();
if snapshot.len() >= expected || tokio::time::Instant::now() >= deadline {
return snapshot;
}
tokio::time::sleep(Duration::from_millis(50)).await;
}
}
}
impl Drop for Harness {
fn drop(&mut self) {
self.gateway.signal_shutdown();
if let Some(handle) = self.agent_handle.take() {
handle.abort();
}
}
}
async fn build_harness(trace: LlmTrace) -> Harness {
let temp_dir = tempfile::tempdir().expect("failed to create temp dir");
let db_path = temp_dir.path().join("telegram_message_routing.db");
let backend = LibSqlBackend::new_local(&db_path)
.await
.expect("failed to create test LibSqlBackend");
backend
.run_migrations()
.await
.expect("failed to run migrations");
let db: Arc<dyn Database> = Arc::new(backend);
let skills_dir = temp_dir.path().join("skills");
let installed_skills_dir = temp_dir.path().join("installed_skills");
let _ = std::fs::create_dir_all(&skills_dir);
let _ = std::fs::create_dir_all(&installed_skills_dir);
let mut config = Config::for_testing(db_path, skills_dir, installed_skills_dir);
config.agent.auto_approve_tools = true;
let session = Arc::new(SessionManager::new(SessionConfig::default()));
let log_broadcaster = Arc::new(LogBroadcaster::new());
let llm: Arc<dyn LlmProvider> = Arc::new(TraceLlm::from_trace(trace));
let mut builder = AppBuilder::new(
config,
AppBuilderFlags::default(),
None,
session,
log_broadcaster,
);
builder.with_database(Arc::clone(&db));
builder.with_llm(llm);
let mut components = builder
.build_all()
.await
.expect("AppBuilder::build_all() failed");
components.config.agent.auto_approve_tools = true;
components.config.agent.allow_local_tools = true;
let deps = AgentDeps {
owner_id: components.config.owner_id.clone(),
store: components.db.clone(),
llm: components.llm.clone(),
cheap_llm: components.cheap_llm.clone(),
safety: components.safety.clone(),
tools: components.tools.clone(),
workspace: components.workspace.clone(),
extension_manager: components.extension_manager.clone(),
skill_registry: components.skill_registry.clone(),
skill_catalog: components.skill_catalog.clone(),
skills_config: components.config.skills.clone(),
hooks: components.hooks.clone(),
cost_guard: components.cost_guard.clone(),
sse_tx: None,
http_interceptor: None,
transcription: None,
document_extraction: None,
};
let gateway = Arc::new(TestChannel::new());
let gateway_handle = TestChannelHandle::new(Arc::clone(&gateway));
let (telegram_channel, telegram_captures) = RecordingTelegramChannel::new();
let channel_manager = ChannelManager::new();
channel_manager.add(Box::new(gateway_handle)).await;
channel_manager.add(Box::new(telegram_channel)).await;
let channels = Arc::new(channel_manager);
deps.tools
.register_message_tools(Arc::clone(&channels), deps.extension_manager.clone())
.await;
let agent = Agent::new(
components.config.agent.clone(),
deps,
channels,
None,
None,
None,
Some(Arc::clone(&components.context_manager)),
None,
);
let agent_handle = tokio::spawn(async move {
if let Err(err) = agent.run().await {
eprintln!("[telegram routing e2e] Agent exited with error: {err}");
}
});
if let Some(rx) = gateway.take_ready_rx().await {
let _ = tokio::time::timeout(Duration::from_secs(5), rx).await;
}
Harness {
gateway,
telegram_captures,
db,
owner_id: components.config.owner_id.clone(),
_temp_dir: temp_dir,
agent_handle: Some(agent_handle),
}
}
fn single_message_trace(arguments: serde_json::Value, final_text: &str) -> LlmTrace {
LlmTrace::single_turn(
"telegram-message-routing",
"send a reminder",
vec![
TraceStep {
request_hint: None,
response: TraceResponse::ToolCalls {
tool_calls: vec![TraceToolCall {
id: "call_message_1".to_string(),
name: "message".to_string(),
arguments,
}],
input_tokens: 32,
output_tokens: 12,
},
expected_tool_results: Vec::new(),
},
TraceStep {
request_hint: None,
response: TraceResponse::Text {
content: final_text.to_string(),
input_tokens: 24,
output_tokens: 8,
},
expected_tool_results: Vec::new(),
},
],
)
}
#[tokio::test]
async fn telegram_message_tool_uses_bound_owner_target_when_target_omitted() {
let harness = build_harness(single_message_trace(
serde_json::json!({
"content": "Walk Conan",
"channel": "telegram",
}),
"Sent on Telegram.",
))
.await;
harness.store_telegram_owner_binding(424242).await;
harness
.gateway
.send_message("remind me to walk conan")
.await;
let responses = harness
.gateway
.wait_for_responses(1, Duration::from_secs(10))
.await;
assert!(
responses
.iter()
.any(|response| response.content.contains("Sent on Telegram")),
"expected assistant confirmation, got: {:?}",
responses
.iter()
.map(|response| &response.content)
.collect::<Vec<_>>()
);
let broadcasts = harness
.wait_for_telegram_broadcasts(1, Duration::from_secs(10))
.await;
assert_eq!(
broadcasts.len(),
1,
"expected exactly one telegram broadcast"
);
assert_eq!(broadcasts[0].0, "424242");
assert_eq!(broadcasts[0].1.content, "Walk Conan");
}
#[tokio::test]
async fn telegram_message_tool_prefers_explicit_target_over_bound_owner_target() {
let harness = build_harness(single_message_trace(
serde_json::json!({
"content": "Walk Conan",
"channel": "telegram",
"target": "999999",
}),
"Sent on Telegram.",
))
.await;
harness.store_telegram_owner_binding(424242).await;
harness.gateway.send_message("send the reminder").await;
let _ = harness
.gateway
.wait_for_responses(1, Duration::from_secs(10))
.await;
let broadcasts = harness
.wait_for_telegram_broadcasts(1, Duration::from_secs(10))
.await;
assert_eq!(
broadcasts.len(),
1,
"expected exactly one telegram broadcast"
);
assert_eq!(broadcasts[0].0, "999999");
assert_eq!(broadcasts[0].1.content, "Walk Conan");
}
}
@@ -34,14 +34,6 @@
"output_tokens": 18
}
},
{
"response": {
"type": "text",
"content": "I saw the Telegram message.",
"input_tokens": 90,
"output_tokens": 12
}
},
{
"response": {
"type": "text",
@@ -35,14 +35,6 @@
"output_tokens": 20
}
},
{
"response": {
"type": "text",
"content": "I saw the Telegram message.",
"input_tokens": 90,
"output_tokens": 12
}
},
{
"response": {
"type": "text",
@@ -239,6 +239,7 @@ impl GatewayWorkflowHarness {
let mut agent = Agent::new(
components.config.agent.clone(),
AgentDeps {
owner_id: components.config.owner_id.clone(),
store: components.db,
llm: components.llm,
cheap_llm: components.cheap_llm,
+2 -1
View File
@@ -612,6 +612,7 @@ impl TestRigBuilder {
// 7. Construct AgentDeps from AppComponents (mirrors main.rs).
let deps = AgentDeps {
owner_id: components.config.owner_id.clone(),
store: components.db,
llm: components.llm,
cheap_llm: components.cheap_llm,
@@ -652,7 +653,7 @@ impl TestRigBuilder {
// 7b. Register message tool so routines can send messages to channels.
deps.tools
.register_message_tools(Arc::clone(&channels))
.register_message_tools(Arc::clone(&channels), deps.extension_manager.clone())
.await;
// 8. Create Agent.
+117 -20
View File
@@ -6,17 +6,24 @@
//! 1. When owner_id is null and dm_policy is "allowlist", unauthorized users in
//! group chats are dropped even if they @mention the bot
//! 2. When owner_id is null and dm_policy is "open", all users can interact
//! 3. When owner_id is set, only that user can interact
//! 3. When owner_id is set, the owner gets instance-global access while
//! non-owner senders remain channel-scoped guests subject to authorization
//! 4. Authorization works correctly for both private and group chats
use std::collections::HashMap;
use std::sync::Arc;
#[cfg(feature = "integration")]
use futures::StreamExt;
#[cfg(feature = "integration")]
use ironclaw::channels::Channel;
use ironclaw::channels::wasm::{
ChannelCapabilities, PreparedChannelModule, WasmChannel, WasmChannelRuntime,
WasmChannelRuntimeConfig,
};
use ironclaw::pairing::PairingStore;
#[cfg(feature = "integration")]
use tokio::time::{Duration, timeout};
/// Skip the test if the Telegram WASM module hasn't been built.
/// In CI (detected via the `CI` env var), panic instead of skipping so a
@@ -40,8 +47,31 @@ macro_rules! require_telegram_wasm {
/// Path to the built Telegram WASM module
fn telegram_wasm_path() -> std::path::PathBuf {
std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("channels-src/telegram/target/wasm32-wasip2/release/telegram_channel.wasm")
let local = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("channels-src/telegram/target/wasm32-wasip2/release/telegram_channel.wasm");
if local.exists() {
return local;
}
if let Ok(output) = std::process::Command::new("git")
.args(["worktree", "list", "--porcelain"])
.output()
&& output.status.success()
{
let stdout = String::from_utf8_lossy(&output.stdout);
for line in stdout.lines() {
if let Some(path) = line.strip_prefix("worktree ") {
let candidate = std::path::PathBuf::from(path).join(
"channels-src/telegram/target/wasm32-wasip2/release/telegram_channel.wasm",
);
if candidate.exists() {
return candidate;
}
}
}
}
local
}
/// Create a test runtime for WASM channel operations.
@@ -74,6 +104,14 @@ async fn load_telegram_module(
async fn create_telegram_channel(
runtime: Arc<WasmChannelRuntime>,
config_json: &str,
) -> WasmChannel {
create_telegram_channel_with_store(runtime, config_json, Arc::new(PairingStore::new())).await
}
async fn create_telegram_channel_with_store(
runtime: Arc<WasmChannelRuntime>,
config_json: &str,
pairing_store: Arc<PairingStore>,
) -> WasmChannel {
let module = load_telegram_module(&runtime)
.await
@@ -83,8 +121,9 @@ async fn create_telegram_channel(
runtime,
module,
ChannelCapabilities::for_channel("telegram").with_path("/webhook/telegram"),
"default",
config_json.to_string(),
Arc::new(PairingStore::new()),
pairing_store,
None,
)
}
@@ -222,31 +261,29 @@ async fn test_group_message_authorized_user_allowed() {
}
#[tokio::test]
async fn test_group_message_with_owner_id_set() {
async fn test_private_message_with_owner_id_set_uses_guest_pairing_flow() {
require_telegram_wasm!();
let runtime = create_test_runtime();
let dir = tempfile::tempdir().expect("tempdir");
let pairing_store = Arc::new(PairingStore::with_base_dir(dir.path().to_path_buf()));
// Config: owner_id=123 (only this user can interact)
// Config: owner_id=123, non-owner private DMs should enter the guest
// pairing flow instead of being rejected solely for not being the owner.
let config = serde_json::json!({
"bot_username": "test_bot",
"bot_username": null,
"owner_id": 123,
"dm_policy": "allowlist",
"allow_from": ["anyone"], // ignored when owner_id is set
"dm_policy": "pairing",
"allow_from": [],
"respond_to_all_group_messages": false
})
.to_string();
let channel = create_telegram_channel(runtime, &config).await;
let channel = create_telegram_channel_with_store(runtime, &config, pairing_store.clone()).await;
// Message from different user (should be dropped)
// Non-owner private message should produce a pairing request.
let update = build_telegram_update(
3,
102,
-123456789,
"group",
999, // Not the owner
"Other",
"Hey @test_bot hello",
3, 102, 999, "private", 999, // Not the owner
"Other", "hello",
);
let response = channel
@@ -263,8 +300,68 @@ async fn test_group_message_with_owner_id_set() {
assert_eq!(response.status, 200);
// REGRESSION TEST: Non-owner messages are dropped when owner_id is set
// This behavior is consistent and not affected by the fix
let pending = pairing_store
.list_pending("telegram")
.expect("pairing store should be readable");
assert_eq!(pending.len(), 1);
assert_eq!(pending[0].id, "999");
}
#[tokio::test]
#[cfg(feature = "integration")]
async fn test_private_messages_use_chat_id_as_thread_scope() {
require_telegram_wasm!();
let runtime = create_test_runtime();
let config = serde_json::json!({
"bot_username": null,
"owner_id": null,
"dm_policy": "open",
"allow_from": [],
"respond_to_all_group_messages": false
})
.to_string();
let channel = create_telegram_channel(runtime, &config).await;
let mut stream = channel
.start_message_stream_for_test()
.await
.expect("Failed to bootstrap test message stream");
for (update_id, message_id, text) in [(6, 105, "first"), (7, 106, "second")] {
let update = build_telegram_update(
update_id,
message_id,
999,
"private",
999,
"ThreadUser",
text,
);
let response = channel
.call_on_http_request(
"POST",
"/webhook/telegram",
&HashMap::new(),
&HashMap::new(),
&update,
true,
)
.await
.expect("HTTP callback failed");
assert_eq!(response.status, 200);
let msg = timeout(Duration::from_secs(1), stream.next())
.await
.expect("message should arrive")
.expect("stream should yield a message");
assert_eq!(msg.thread_id.as_deref(), Some("999"));
assert_eq!(msg.conversation_scope(), Some("999"));
}
channel.shutdown().await.expect("Shutdown failed");
}
#[tokio::test]
+1
View File
@@ -43,6 +43,7 @@ fn create_test_channel(
runtime,
prepared,
capabilities,
"default",
"{}".to_string(),
Arc::new(PairingStore::new()),
None,