chore: resolve conflicts

This commit is contained in:
italic-jinxin
2026-03-25 19:23:08 +08:00
47 changed files with 2973 additions and 757 deletions
Generated
+9
View File
@@ -3428,6 +3428,7 @@ dependencies = [
"hyper-util",
"iana-time-zone",
"insta",
"ironclaw_common",
"ironclaw_safety",
"json5",
"libsql",
@@ -3485,6 +3486,14 @@ dependencies = [
"zip",
]
[[package]]
name = "ironclaw_common"
version = "0.1.0"
dependencies = [
"serde",
"serde_json",
]
[[package]]
name = "ironclaw_safety"
version = "0.1.0"
+4 -1
View File
@@ -1,5 +1,5 @@
[workspace]
members = [".", "crates/ironclaw_safety"]
members = [".", "crates/ironclaw_common", "crates/ironclaw_safety"]
exclude = [
"channels-src/discord",
"channels-src/telegram",
@@ -100,6 +100,9 @@ tower-http = { version = "0.6", features = ["trace", "cors", "set-header"] }
# Cron scheduling for routines
cron = "0.13"
# Shared types
ironclaw_common = { path = "crates/ironclaw_common", version = "0.1.0" }
# Safety/sanitization
ironclaw_safety = { path = "crates/ironclaw_safety", version = "0.1.0" }
regex = "1"
+18
View File
@@ -0,0 +1,18 @@
[package]
name = "ironclaw_common"
version = "0.1.0"
edition = "2024"
rust-version = "1.92"
description = "Shared types and utilities for the IronClaw workspace"
authors = ["NEAR AI <[email protected]>"]
license = "MIT OR Apache-2.0"
homepage = "https://github.com/nearai/ironclaw"
repository = "https://github.com/nearai/ironclaw"
publish = false
[package.metadata.dist]
dist = false
[dependencies]
serde = { version = "1", features = ["derive"] }
serde_json = "1"
+338
View File
@@ -0,0 +1,338 @@
//! Application-wide event types.
//!
//! `AppEvent` is the real-time event protocol used across the entire
//! application. The web gateway serialises these to SSE / WebSocket
//! frames, but other subsystems (agent loop, orchestrator, extensions)
//! produce and consume them too.
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum AppEvent {
#[serde(rename = "response")]
Response { content: String, thread_id: String },
#[serde(rename = "thinking")]
Thinking {
message: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "tool_started")]
ToolStarted {
name: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "tool_completed")]
ToolCompleted {
name: String,
success: bool,
#[serde(skip_serializing_if = "Option::is_none")]
error: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
parameters: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "tool_result")]
ToolResult {
name: String,
preview: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "stream_chunk")]
StreamChunk {
content: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "status")]
Status {
message: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "job_started")]
JobStarted {
job_id: String,
title: String,
browse_url: String,
},
#[serde(rename = "approval_needed")]
ApprovalNeeded {
request_id: String,
tool_name: String,
description: String,
parameters: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
/// Whether the "always" auto-approve option should be shown.
allow_always: bool,
},
#[serde(rename = "auth_required")]
AuthRequired {
extension_name: String,
#[serde(skip_serializing_if = "Option::is_none")]
instructions: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
auth_url: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
setup_url: Option<String>,
},
#[serde(rename = "auth_completed")]
AuthCompleted {
extension_name: String,
success: bool,
message: String,
},
#[serde(rename = "error")]
Error {
message: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "heartbeat")]
Heartbeat,
// Sandbox job streaming events (worker + Claude Code bridge)
#[serde(rename = "job_message")]
JobMessage {
job_id: String,
role: String,
content: String,
},
#[serde(rename = "job_tool_use")]
JobToolUse {
job_id: String,
tool_name: String,
input: serde_json::Value,
},
#[serde(rename = "job_tool_result")]
JobToolResult {
job_id: String,
tool_name: String,
output: String,
},
#[serde(rename = "job_status")]
JobStatus { job_id: String, message: String },
#[serde(rename = "job_result")]
JobResult {
job_id: String,
status: String,
#[serde(skip_serializing_if = "Option::is_none")]
session_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
fallback_deliverable: Option<serde_json::Value>,
},
/// An image was generated by a tool.
#[serde(rename = "image_generated")]
ImageGenerated {
data_url: String,
#[serde(skip_serializing_if = "Option::is_none")]
path: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
/// Suggested follow-up messages for the user.
#[serde(rename = "suggestions")]
Suggestions {
suggestions: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
/// Per-turn token usage and cost summary.
#[serde(rename = "turn_cost")]
TurnCost {
input_tokens: u64,
output_tokens: u64,
cost_usd: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
/// Extension activation status change (WASM channels).
#[serde(rename = "extension_status")]
ExtensionStatus {
extension_name: String,
status: String,
#[serde(skip_serializing_if = "Option::is_none")]
message: Option<String>,
},
}
impl AppEvent {
/// The wire-format event type string (matches the `#[serde(rename)]` value).
pub fn event_type(&self) -> &'static str {
match self {
Self::Response { .. } => "response",
Self::Thinking { .. } => "thinking",
Self::ToolStarted { .. } => "tool_started",
Self::ToolCompleted { .. } => "tool_completed",
Self::ToolResult { .. } => "tool_result",
Self::StreamChunk { .. } => "stream_chunk",
Self::Status { .. } => "status",
Self::JobStarted { .. } => "job_started",
Self::ApprovalNeeded { .. } => "approval_needed",
Self::AuthRequired { .. } => "auth_required",
Self::AuthCompleted { .. } => "auth_completed",
Self::Error { .. } => "error",
Self::Heartbeat => "heartbeat",
Self::JobMessage { .. } => "job_message",
Self::JobToolUse { .. } => "job_tool_use",
Self::JobToolResult { .. } => "job_tool_result",
Self::JobStatus { .. } => "job_status",
Self::JobResult { .. } => "job_result",
Self::ImageGenerated { .. } => "image_generated",
Self::Suggestions { .. } => "suggestions",
Self::TurnCost { .. } => "turn_cost",
Self::ExtensionStatus { .. } => "extension_status",
}
}
}
#[cfg(test)]
mod tests {
use super::*;
/// Verify that `event_type()` returns the same string as the serde
/// `"type"` field for every variant. This catches drift between the
/// `#[serde(rename)]` attributes and the manual match arms.
#[test]
fn event_type_matches_serde_type_field() {
let variants: Vec<AppEvent> = vec![
AppEvent::Response {
content: String::new(),
thread_id: String::new(),
},
AppEvent::Thinking {
message: String::new(),
thread_id: None,
},
AppEvent::ToolStarted {
name: String::new(),
thread_id: None,
},
AppEvent::ToolCompleted {
name: String::new(),
success: true,
error: None,
parameters: None,
thread_id: None,
},
AppEvent::ToolResult {
name: String::new(),
preview: String::new(),
thread_id: None,
},
AppEvent::StreamChunk {
content: String::new(),
thread_id: None,
},
AppEvent::Status {
message: String::new(),
thread_id: None,
},
AppEvent::JobStarted {
job_id: String::new(),
title: String::new(),
browse_url: String::new(),
},
AppEvent::ApprovalNeeded {
request_id: String::new(),
tool_name: String::new(),
description: String::new(),
parameters: String::new(),
thread_id: None,
allow_always: false,
},
AppEvent::AuthRequired {
extension_name: String::new(),
instructions: None,
auth_url: None,
setup_url: None,
},
AppEvent::AuthCompleted {
extension_name: String::new(),
success: true,
message: String::new(),
},
AppEvent::Error {
message: String::new(),
thread_id: None,
},
AppEvent::Heartbeat,
AppEvent::JobMessage {
job_id: String::new(),
role: String::new(),
content: String::new(),
},
AppEvent::JobToolUse {
job_id: String::new(),
tool_name: String::new(),
input: serde_json::Value::Null,
},
AppEvent::JobToolResult {
job_id: String::new(),
tool_name: String::new(),
output: String::new(),
},
AppEvent::JobStatus {
job_id: String::new(),
message: String::new(),
},
AppEvent::JobResult {
job_id: String::new(),
status: String::new(),
session_id: None,
fallback_deliverable: None,
},
AppEvent::ImageGenerated {
data_url: String::new(),
path: None,
thread_id: None,
},
AppEvent::Suggestions {
suggestions: vec![],
thread_id: None,
},
AppEvent::TurnCost {
input_tokens: 0,
output_tokens: 0,
cost_usd: String::new(),
thread_id: None,
},
AppEvent::ExtensionStatus {
extension_name: String::new(),
status: String::new(),
message: None,
},
];
for variant in &variants {
let json: serde_json::Value = serde_json::to_value(variant).unwrap();
let serde_type = json["type"].as_str().unwrap();
assert_eq!(
variant.event_type(),
serde_type,
"event_type() mismatch for variant: {:?}",
variant
);
}
}
#[test]
fn round_trip_deserialize() {
let original = AppEvent::Response {
content: "hello".to_string(),
thread_id: "t1".to_string(),
};
let json = serde_json::to_string(&original).unwrap();
let deserialized: AppEvent = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized.event_type(), "response");
}
}
+7
View File
@@ -0,0 +1,7 @@
//! Shared types and utilities for the IronClaw workspace.
mod event;
mod util;
pub use event::AppEvent;
pub use util::truncate_preview;
+100
View File
@@ -0,0 +1,100 @@
//! Shared utility functions.
/// Truncate a string to at most `max_bytes` bytes at a char boundary, appending "...".
///
/// If the input is wrapped in `<tool_output ...>...</tool_output>` and truncation
/// removes the closing tag, the tag is re-appended so downstream XML parsers
/// never see an unclosed element.
pub fn truncate_preview(s: &str, max_bytes: usize) -> String {
if s.len() <= max_bytes {
return s.to_string();
}
// Walk backwards from max_bytes to find a valid char boundary
let mut end = max_bytes;
while end > 0 && !s.is_char_boundary(end) {
end -= 1;
}
let mut result = format!("{}...", &s[..end]);
// Re-close <tool_output> if truncation cut through the closing tag.
if s.starts_with("<tool_output") && !result.ends_with("</tool_output>") {
result.push_str("\n</tool_output>");
}
result
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_truncate_preview_short_string() {
assert_eq!(truncate_preview("hello", 10), "hello");
}
#[test]
fn test_truncate_preview_exact_boundary() {
assert_eq!(truncate_preview("hello", 5), "hello");
}
#[test]
fn test_truncate_preview_truncates_ascii() {
assert_eq!(truncate_preview("hello world", 5), "hello...");
}
#[test]
fn test_truncate_preview_empty_string() {
assert_eq!(truncate_preview("", 10), "");
}
#[test]
fn test_truncate_preview_multibyte_char_boundary() {
let s = "a\u{20AC}b";
let result = truncate_preview(s, 3);
assert_eq!(result, "a...");
}
#[test]
fn test_truncate_preview_emoji() {
let s = "hi\u{1F980}";
let result = truncate_preview(s, 4);
assert_eq!(result, "hi...");
}
#[test]
fn test_truncate_preview_cjk() {
let s = "\u{4F60}\u{597D}\u{4E16}\u{754C}";
let result = truncate_preview(s, 7);
assert_eq!(result, "\u{4F60}\u{597D}...");
}
#[test]
fn test_truncate_preview_zero_max_bytes() {
assert_eq!(truncate_preview("hello", 0), "...");
}
#[test]
fn test_truncate_preview_closes_tool_output_tag() {
let s = "<tool_output name=\"search\">\nSome very long content here\n</tool_output>";
let result = truncate_preview(s, 60);
assert!(result.ends_with("</tool_output>"));
assert!(result.contains("..."));
}
#[test]
fn test_truncate_preview_no_extra_close_when_intact() {
let s = "<tool_output name=\"echo\">\nshort\n</tool_output>";
let result = truncate_preview(s, 500);
assert_eq!(result, s);
assert_eq!(result.matches("</tool_output>").count(), 1);
}
#[test]
fn test_truncate_preview_non_xml_unaffected() {
let s = "Just a plain long string that gets truncated";
let result = truncate_preview(s, 10);
assert_eq!(result, "Just a pla...");
assert!(!result.contains("</tool_output>"));
}
}
+2 -1
View File
@@ -1055,10 +1055,11 @@ impl Agent {
} else {
drop(sess);
self.session_manager
.resolve_thread(
.resolve_thread_with_parsed_uuid(
&message.user_id,
&message.channel,
message.conversation_scope(),
approval_thread_uuid,
)
.await
}
+24 -24
View File
@@ -21,8 +21,8 @@ use tokio::task::JoinHandle;
use uuid::Uuid;
use crate::channels::IncomingMessage;
use crate::channels::web::types::SseEvent;
use crate::context::{ContextManager, JobState};
use ironclaw_common::AppEvent;
/// Route context for forwarding job monitor events back to the user's channel.
#[derive(Debug, Clone)]
@@ -36,15 +36,15 @@ pub struct JobMonitorRoute {
/// injects assistant messages into the agent loop.
///
/// The monitor forwards:
/// - `SseEvent::JobMessage` (assistant role): injected as incoming messages so
/// - `AppEvent::JobMessage` (assistant role): injected as incoming messages so
/// the main agent can read and relay to the user.
/// - `SseEvent::JobResult`: injected as a completion notice, then the task exits.
/// - `AppEvent::JobResult`: injected as a completion notice, then the task exits.
///
/// Tool use/result and status events are intentionally skipped (too noisy for
/// the main agent's context window).
pub fn spawn_job_monitor(
job_id: Uuid,
event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>,
event_rx: broadcast::Receiver<(Uuid, String, AppEvent)>,
inject_tx: mpsc::Sender<IncomingMessage>,
route: JobMonitorRoute,
) -> JoinHandle<()> {
@@ -56,7 +56,7 @@ pub fn spawn_job_monitor(
/// jobs don't stay `InProgress` forever in the `ContextManager`.
pub fn spawn_job_monitor_with_context(
job_id: Uuid,
mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>,
mut event_rx: broadcast::Receiver<(Uuid, String, AppEvent)>,
inject_tx: mpsc::Sender<IncomingMessage>,
route: JobMonitorRoute,
context_manager: Option<Arc<ContextManager>>,
@@ -74,7 +74,7 @@ pub fn spawn_job_monitor_with_context(
}
match event {
SseEvent::JobMessage { role, content, .. } if role == "assistant" => {
AppEvent::JobMessage { role, content, .. } if role == "assistant" => {
let mut msg = IncomingMessage::new(
route.channel.clone(),
route.user_id.clone(),
@@ -92,7 +92,7 @@ pub fn spawn_job_monitor_with_context(
break;
}
}
SseEvent::JobResult { status, .. } => {
AppEvent::JobResult { status, .. } => {
// Transition in-memory state so the job frees its
// max_jobs slot and query tools show the final state.
if let Some(ref cm) = context_manager {
@@ -162,7 +162,7 @@ pub fn spawn_job_monitor_with_context(
/// inject messages into) but we still need to free the `max_jobs` slot.
pub fn spawn_completion_watcher(
job_id: Uuid,
mut event_rx: broadcast::Receiver<(Uuid, String, SseEvent)>,
mut event_rx: broadcast::Receiver<(Uuid, String, AppEvent)>,
context_manager: Arc<ContextManager>,
) -> JoinHandle<()> {
let short_id = job_id.to_string()[..8].to_string();
@@ -170,7 +170,7 @@ pub fn spawn_completion_watcher(
tokio::spawn(async move {
loop {
match event_rx.recv().await {
Ok((ev_job_id, _user_id, SseEvent::JobResult { status, .. }))
Ok((ev_job_id, _user_id, AppEvent::JobResult { status, .. }))
if ev_job_id == job_id =>
{
let target = if status == "completed" {
@@ -229,7 +229,7 @@ mod tests {
#[tokio::test]
async fn test_monitor_forwards_assistant_messages() {
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4();
@@ -240,7 +240,7 @@ mod tests {
.send((
job_id,
"test-user".to_string(),
SseEvent::JobMessage {
AppEvent::JobMessage {
job_id: job_id.to_string(),
role: "assistant".to_string(),
content: "I found a bug".to_string(),
@@ -262,7 +262,7 @@ mod tests {
#[tokio::test]
async fn test_monitor_ignores_other_jobs() {
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4();
@@ -274,7 +274,7 @@ mod tests {
.send((
other_job_id,
"test-user".to_string(),
SseEvent::JobMessage {
AppEvent::JobMessage {
job_id: other_job_id.to_string(),
role: "assistant".to_string(),
content: "wrong job".to_string(),
@@ -293,7 +293,7 @@ mod tests {
#[tokio::test]
async fn test_monitor_exits_on_job_result() {
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4();
@@ -304,7 +304,7 @@ mod tests {
.send((
job_id,
"test-user".to_string(),
SseEvent::JobResult {
AppEvent::JobResult {
job_id: job_id.to_string(),
status: "completed".to_string(),
session_id: None,
@@ -329,7 +329,7 @@ mod tests {
#[tokio::test]
async fn test_monitor_skips_tool_events() {
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let job_id = Uuid::new_v4();
@@ -340,7 +340,7 @@ mod tests {
.send((
job_id,
"test-user".to_string(),
SseEvent::JobToolUse {
AppEvent::JobToolUse {
job_id: job_id.to_string(),
tool_name: "shell".to_string(),
input: serde_json::json!({"command": "ls"}),
@@ -353,7 +353,7 @@ mod tests {
.send((
job_id,
"test-user".to_string(),
SseEvent::JobMessage {
AppEvent::JobMessage {
job_id: job_id.to_string(),
role: "user".to_string(),
content: "user prompt".to_string(),
@@ -409,7 +409,7 @@ mod tests {
.await
.unwrap();
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let handle = spawn_job_monitor_with_context(
@@ -425,7 +425,7 @@ mod tests {
.send((
job_id,
"test-user".to_string(),
SseEvent::JobResult {
AppEvent::JobResult {
job_id: job_id.to_string(),
status: "completed".to_string(),
session_id: None,
@@ -458,7 +458,7 @@ mod tests {
.await
.unwrap();
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16);
let (inject_tx, mut inject_rx) = mpsc::channel::<IncomingMessage>(16);
let handle = spawn_job_monitor_with_context(
@@ -474,7 +474,7 @@ mod tests {
.send((
job_id,
"test-user".to_string(),
SseEvent::JobResult {
AppEvent::JobResult {
job_id: job_id.to_string(),
status: "failed".to_string(),
session_id: None,
@@ -507,14 +507,14 @@ mod tests {
.await
.unwrap();
let (event_tx, _) = broadcast::channel::<(Uuid, String, SseEvent)>(16);
let (event_tx, _) = broadcast::channel::<(Uuid, String, AppEvent)>(16);
let handle = spawn_completion_watcher(job_id, event_tx.subscribe(), Arc::clone(&cm));
event_tx
.send((
job_id,
"test-user".to_string(),
SseEvent::JobResult {
AppEvent::JobResult {
job_id: job_id.to_string(),
status: "completed".to_string(),
session_id: None,
+4 -1
View File
@@ -1541,7 +1541,10 @@ async fn execute_lightweight_with_tools(
let force_text = iteration >= max_iterations;
if force_text {
// Final iteration: no tools, just get text response
// Final iteration: no tools, just get text response.
// Claude 4.6 rejects assistant prefill; NEAR AI rejects any non-user-ending
// conversation. Ensure the last message is user-role.
crate::util::ensure_ends_with_user_message(&mut messages);
let request = CompletionRequest::new(messages)
.with_max_tokens(effective_max_tokens)
.with_temperature(0.3);
+1 -1
View File
@@ -16,8 +16,8 @@ use chrono::{DateTime, TimeDelta, Utc};
use serde::{Deserialize, Serialize};
use uuid::Uuid;
use crate::channels::web::util::truncate_preview;
use crate::llm::{ChatMessage, ToolCall, generate_tool_call_id};
use ironclaw_common::truncate_preview;
/// A session containing one or more threads.
#[derive(Debug, Clone, Serialize, Deserialize)]
+189 -34
View File
@@ -102,11 +102,30 @@ impl SessionManager {
/// Resolve an external thread ID to an internal thread.
///
/// Returns the session and thread ID. Creates both if they don't exist.
/// Delegates to [`resolve_thread_with_parsed_uuid`](Self::resolve_thread_with_parsed_uuid)
/// with `parsed_uuid: None`.
pub async fn resolve_thread(
&self,
user_id: &str,
channel: &str,
external_thread_id: Option<&str>,
) -> (Arc<Mutex<Session>>, Uuid) {
self.resolve_thread_with_parsed_uuid(user_id, channel, external_thread_id, None)
.await
}
/// Like [`resolve_thread`](Self::resolve_thread), but accepts a pre-parsed
/// UUID to skip redundant parsing when the caller has already validated
/// the external thread ID as a UUID (e.g. the approval routing path).
///
/// Uses a single read-lock acquisition for both the key lookup and the UUID
/// adoption check to reduce contention under concurrent approval load.
pub async fn resolve_thread_with_parsed_uuid(
&self,
user_id: &str,
channel: &str,
external_thread_id: Option<&str>,
parsed_uuid: Option<Uuid>,
) -> (Arc<Mutex<Session>>, Uuid) {
let session = self.get_or_create_session(user_id).await;
@@ -116,51 +135,65 @@ impl SessionManager {
external_thread_id: external_thread_id.map(String::from),
};
// Check if we have a mapping
{
// Use pre-parsed UUID if available, otherwise parse from string.
let ext_uuid = parsed_uuid
.or_else(|| external_thread_id.and_then(|ext_tid| Uuid::parse_str(ext_tid).ok()));
// Validate that parsed_uuid (if provided) is consistent with external_thread_id.
#[cfg(debug_assertions)]
if let (Some(parsed), Some(ext_tid)) = (&parsed_uuid, external_thread_id) {
debug_assert_eq!(
Uuid::parse_str(ext_tid).ok().as_ref(),
Some(parsed),
"parsed_uuid must be the parsed form of external_thread_id"
);
}
// Single read lock for both the key lookup and UUID adoption check
let adoptable_uuid = {
let thread_map = self.thread_map.read().await;
// Fast path: exact key match
if let Some(&thread_id) = thread_map.get(&key) {
// Verify thread still exists in session
let sess = session.lock().await;
if sess.threads.contains_key(&thread_id) {
return (Arc::clone(&session), thread_id);
}
}
}
// Check if external_thread_id is itself a known thread UUID that
// exists in the session but was never registered in the thread_map
// (e.g. created by chat_new_thread_handler or hydrated from DB).
// We only adopt it if no thread_map entry maps to this UUID —
// otherwise it belongs to a different channel scope.
if let Some(ext_tid) = external_thread_id
&& let Ok(ext_uuid) = Uuid::parse_str(ext_tid)
{
let thread_map = self.thread_map.read().await;
let mapped_elsewhere = thread_map.values().any(|&v| v == ext_uuid);
drop(thread_map);
// UUID adoption check (still under the same read lock).
// If external_thread_id is a valid UUID not mapped elsewhere,
// it may be a thread created by chat_new_thread_handler or
// hydrated from DB that we can adopt.
// Only attempt adoption when external_thread_id is Some, preserving
// the invariant that None external_thread_id never triggers adoption.
if external_thread_id.is_some() {
ext_uuid.filter(|&uuid| !thread_map.values().any(|&v| v == uuid))
} else {
None
}
}; // Single read lock dropped here
if !mapped_elsewhere {
let sess = session.lock().await;
if sess.threads.contains_key(&ext_uuid) {
drop(sess);
// If we found an adoptable UUID, verify it exists in session and acquire write lock
if let Some(ext_uuid) = adoptable_uuid {
let sess = session.lock().await;
if sess.threads.contains_key(&ext_uuid) {
drop(sess);
let mut thread_map = self.thread_map.write().await;
// Re-check after acquiring write lock to prevent race condition
// where another task mapped this UUID between our read and write.
if !thread_map.values().any(|&v| v == ext_uuid) {
thread_map.insert(key, ext_uuid);
drop(thread_map);
// Ensure undo manager exists
let mut undo_managers = self.undo_managers.write().await;
undo_managers
.entry(ext_uuid)
.or_insert_with(|| Arc::new(Mutex::new(UndoManager::new())));
return (session, ext_uuid);
}
// If it was mapped elsewhere while we were unlocked, fall through
// to create a new thread, preserving channel isolation.
let mut thread_map = self.thread_map.write().await;
// Re-check after acquiring write lock to prevent race condition
// where another task mapped this UUID between our read and write.
if !thread_map.values().any(|&v| v == ext_uuid) {
thread_map.insert(key, ext_uuid);
drop(thread_map);
// Ensure undo manager exists
let mut undo_managers = self.undo_managers.write().await;
undo_managers
.entry(ext_uuid)
.or_insert_with(|| Arc::new(Mutex::new(UndoManager::new())));
return (session, ext_uuid);
}
// If mapped elsewhere while unlocked, fall through to create new thread
}
}
@@ -909,6 +942,44 @@ mod tests {
}
}
#[tokio::test]
async fn test_resolve_thread_consolidates_read_path() {
// Verify that resolve_thread still correctly handles:
// 1. Fast path: key exists in thread_map
// 2. UUID adoption: external_thread_id is a UUID in session but not in map
// 3. New thread: neither path matches
use crate::agent::session::Thread;
let manager = SessionManager::new();
// Case 1: Normal resolution creates thread and maps it
let (session1, tid1) = manager
.resolve_thread("user1", "chan1", Some("ext-1"))
.await;
// Resolving again with same key should return same thread (fast path)
let (_, tid1_again) = manager
.resolve_thread("user1", "chan1", Some("ext-1"))
.await;
assert_eq!(tid1, tid1_again);
// Case 2: UUID adoption - insert a thread directly into session
let adopted_id = Uuid::new_v4();
{
let mut sess = session1.lock().await;
let thread = Thread::with_id(adopted_id, sess.id);
sess.threads.insert(adopted_id, thread);
}
// Resolve with the UUID as external_thread_id -- should adopt it
let (_, resolved) = manager
.resolve_thread("user1", "chan1", Some(&adopted_id.to_string()))
.await;
assert_eq!(resolved, adopted_id);
// Case 3: Different channel gets different thread
let (_, tid2) = manager.resolve_thread("user1", "chan2", None).await;
assert_ne!(tid1, tid2);
}
#[tokio::test]
async fn test_resolve_thread_finds_existing_session_thread_by_uuid() {
use crate::agent::session::{Session, Thread};
@@ -947,4 +1018,88 @@ mod tests {
"should have exactly 1 thread, not a duplicate"
);
}
#[tokio::test]
async fn test_resolve_thread_with_pre_parsed_uuid_adopts_thread() {
use crate::agent::session::Thread;
let manager = SessionManager::new();
let (session, _) = manager.resolve_thread("user1", "chan1", None).await;
// Manually insert a thread with a known UUID
let known_id = Uuid::new_v4();
{
let mut sess = session.lock().await;
let thread = Thread::with_id(known_id, sess.id);
sess.threads.insert(known_id, thread);
}
// Resolve with pre-parsed UUID -- should adopt it without re-parsing
let (_, resolved) = manager
.resolve_thread_with_parsed_uuid(
"user1",
"chan1",
Some(&known_id.to_string()),
Some(known_id),
)
.await;
assert_eq!(resolved, known_id);
}
#[tokio::test]
async fn test_resolve_thread_with_parsed_uuid_none_delegates_to_parse() {
use crate::agent::session::Thread;
let manager = SessionManager::new();
let (session, _) = manager.resolve_thread("user2", "chan2", None).await;
// Insert a thread with a known UUID
let known_id = Uuid::new_v4();
{
let mut sess = session.lock().await;
let thread = Thread::with_id(known_id, sess.id);
sess.threads.insert(known_id, thread);
}
// Resolve with parsed_uuid=None but a valid UUID string -- should
// fall back to parsing the string and still adopt the thread
let (_, resolved) = manager
.resolve_thread_with_parsed_uuid("user2", "chan2", Some(&known_id.to_string()), None)
.await;
assert_eq!(resolved, known_id);
}
#[tokio::test]
async fn test_resolve_thread_with_none_external_thread_id_does_not_adopt() {
use crate::agent::session::Thread;
let manager = SessionManager::new();
let (session, default_tid) = manager.resolve_thread("user3", "chan3", None).await;
// Manually insert a thread with a known UUID (simulating a thread
// created by chat_new_thread_handler)
let known_id = Uuid::new_v4();
{
let mut sess = session.lock().await;
let thread = Thread::with_id(known_id, sess.id);
sess.threads.insert(known_id, thread);
}
// Resolve with external_thread_id=None but parsed_uuid=Some.
// This should NOT adopt the UUID — the old code prevented adoption
// when external_thread_id was None, and we preserve that invariant.
let (_, resolved) = manager
.resolve_thread_with_parsed_uuid("user3", "chan3", None, Some(known_id))
.await;
// Should return the existing default thread, not the injected UUID
assert_eq!(
resolved, default_tid,
"should return existing default thread when external_thread_id is None"
);
assert_ne!(
resolved, known_id,
"should NOT adopt UUID when external_thread_id is None"
);
}
}
+1 -1
View File
@@ -16,12 +16,12 @@ use crate::agent::dispatcher::{
};
use crate::agent::session::{MAX_PENDING_MESSAGES, PendingApproval, Session, ThreadState};
use crate::agent::submission::SubmissionResult;
use crate::channels::web::util::truncate_preview;
use crate::channels::{IncomingMessage, StatusUpdate};
use crate::context::JobContext;
use crate::error::Error;
use crate::llm::{ChatMessage, ToolCall};
use crate::tools::redact_params;
use ironclaw_common::truncate_preview;
const FORGED_THREAD_ID_ERROR: &str = "Invalid or unauthorized thread ID.";
+1 -7
View File
@@ -329,13 +329,7 @@ impl AppBuilder {
.create_provider(&self.config.llm.nearai.base_url, self.session.clone());
// Register memory tools if database is available
let workspace_user_id = self
.config
.channels
.gateway
.as_ref()
.map(|gw| gw.user_id.as_str())
.unwrap_or("default");
let workspace_user_id = self.config.owner_id.as_str();
let workspace = if let Some(ref db) = self.db {
let emb_cache_config = EmbeddingCacheConfig {
max_entries: self.config.embeddings.cache_size,
+3 -3
View File
@@ -175,7 +175,7 @@ pub async fn chat_auth_token_handler(
if result.verification.is_some() {
state.sse.broadcast_for_user(
&user.user_id,
SseEvent::AuthRequired {
AppEvent::AuthRequired {
extension_name: req.extension_name.clone(),
instructions: Some(result.message),
auth_url: None,
@@ -187,7 +187,7 @@ pub async fn chat_auth_token_handler(
state.sse.broadcast_for_user(
&user.user_id,
SseEvent::AuthCompleted {
AppEvent::AuthCompleted {
extension_name: req.extension_name.clone(),
success: true,
message: result.message,
@@ -202,7 +202,7 @@ pub async fn chat_auth_token_handler(
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
state.sse.broadcast_for_user(
&user.user_id,
SseEvent::AuthRequired {
AppEvent::AuthRequired {
extension_name: req.extension_name.clone(),
instructions: Some(msg.clone()),
auth_url: None,
+2 -1
View File
@@ -677,7 +677,8 @@ mod tests {
job_manager: None,
prompt_queue: None,
scheduler: None,
default_user_id: "test".to_string(),
owner_id: "test".to_string(),
default_sender_id: "test".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: None,
llm_provider: None,
+38 -19
View File
@@ -58,7 +58,7 @@ use self::log_layer::{LogBroadcaster, LogLevelHandle};
use self::auth::MultiAuthState;
use self::server::GatewayState;
use self::sse::SseManager;
use self::types::SseEvent;
use self::types::AppEvent;
/// Web gateway channel implementing the Channel trait.
pub struct GatewayChannel {
@@ -98,7 +98,8 @@ impl GatewayChannel {
job_manager: None,
prompt_queue: None,
scheduler: None,
default_user_id: config.user_id.clone(),
owner_id: config.user_id.clone(),
default_sender_id: config.user_id.clone(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())),
llm_provider: None,
@@ -122,6 +123,22 @@ impl GatewayChannel {
}
}
/// Rebind the single-user auth identity to the durable owner scope while
/// preserving the configured gateway sender/routing identity.
pub fn with_owner_scope(mut self, owner_id: impl Into<String>) -> Self {
let owner_id = owner_id.into();
let single_user_token = if self.config.user_tokens.is_none() {
self.auth.first_token().map(ToOwned::to_owned)
} else {
None
};
if let Some(token) = single_user_token {
self.auth = MultiAuthState::single(token, owner_id.clone());
}
self.rebuild_state(|s| s.owner_id = owner_id);
self
}
/// Create a gateway channel with a pre-built multi-user auth state.
pub fn new_multi_auth(config: GatewayConfig, auth: MultiAuthState) -> Self {
let state = Arc::new(GatewayState {
@@ -138,7 +155,8 @@ impl GatewayChannel {
job_manager: None,
prompt_queue: None,
scheduler: None,
default_user_id: config.user_id.clone(),
owner_id: config.user_id.clone(),
default_sender_id: config.user_id.clone(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(ws::WsConnectionTracker::new())),
llm_provider: None,
@@ -179,7 +197,8 @@ impl GatewayChannel {
job_manager: self.state.job_manager.clone(),
prompt_queue: self.state.prompt_queue.clone(),
scheduler: self.state.scheduler.clone(),
default_user_id: self.state.default_user_id.clone(),
owner_id: self.state.owner_id.clone(),
default_sender_id: self.state.default_sender_id.clone(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: self.state.ws_tracker.clone(),
llm_provider: self.state.llm_provider.clone(),
@@ -379,7 +398,7 @@ impl Channel for GatewayChannel {
self.state.sse.broadcast_for_user(
&msg.user_id,
SseEvent::Response {
AppEvent::Response {
content: response.content,
thread_id,
},
@@ -398,11 +417,11 @@ impl Channel for GatewayChannel {
.and_then(|v| v.as_str())
.map(String::from);
let event = match status {
StatusUpdate::Thinking(msg) => SseEvent::Thinking {
StatusUpdate::Thinking(msg) => AppEvent::Thinking {
message: msg,
thread_id: thread_id.clone(),
},
StatusUpdate::ToolStarted { name } => SseEvent::ToolStarted {
StatusUpdate::ToolStarted { name } => AppEvent::ToolStarted {
name,
thread_id: thread_id.clone(),
},
@@ -411,23 +430,23 @@ impl Channel for GatewayChannel {
success,
error,
parameters,
} => SseEvent::ToolCompleted {
} => AppEvent::ToolCompleted {
name,
success,
error,
parameters,
thread_id: thread_id.clone(),
},
StatusUpdate::ToolResult { name, preview } => SseEvent::ToolResult {
StatusUpdate::ToolResult { name, preview } => AppEvent::ToolResult {
name,
preview,
thread_id: thread_id.clone(),
},
StatusUpdate::StreamChunk(content) => SseEvent::StreamChunk {
StatusUpdate::StreamChunk(content) => AppEvent::StreamChunk {
content,
thread_id: thread_id.clone(),
},
StatusUpdate::Status(msg) => SseEvent::Status {
StatusUpdate::Status(msg) => AppEvent::Status {
message: msg,
thread_id: thread_id.clone(),
},
@@ -435,7 +454,7 @@ impl Channel for GatewayChannel {
job_id,
title,
browse_url,
} => SseEvent::JobStarted {
} => AppEvent::JobStarted {
job_id,
title,
browse_url,
@@ -446,7 +465,7 @@ impl Channel for GatewayChannel {
description,
parameters,
allow_always,
} => SseEvent::ApprovalNeeded {
} => AppEvent::ApprovalNeeded {
request_id,
tool_name,
description,
@@ -460,7 +479,7 @@ impl Channel for GatewayChannel {
instructions,
auth_url,
setup_url,
} => SseEvent::AuthRequired {
} => AppEvent::AuthRequired {
extension_name,
instructions,
auth_url,
@@ -470,17 +489,17 @@ impl Channel for GatewayChannel {
extension_name,
success,
message,
} => SseEvent::AuthCompleted {
} => AppEvent::AuthCompleted {
extension_name,
success,
message,
},
StatusUpdate::ImageGenerated { data_url, path } => SseEvent::ImageGenerated {
StatusUpdate::ImageGenerated { data_url, path } => AppEvent::ImageGenerated {
data_url,
path,
thread_id: thread_id.clone(),
},
StatusUpdate::Suggestions { suggestions } => SseEvent::Suggestions {
StatusUpdate::Suggestions { suggestions } => AppEvent::Suggestions {
suggestions,
thread_id,
},
@@ -488,7 +507,7 @@ impl Channel for GatewayChannel {
input_tokens,
output_tokens,
cost_usd,
} => SseEvent::TurnCost {
} => AppEvent::TurnCost {
input_tokens,
output_tokens,
cost_usd,
@@ -524,7 +543,7 @@ impl Channel for GatewayChannel {
};
self.state.sse.broadcast_for_user(
user_id,
SseEvent::Response {
AppEvent::Response {
content: response.content,
thread_id,
},
+29 -27
View File
@@ -350,8 +350,10 @@ pub struct GatewayState {
pub job_manager: Option<Arc<ContainerJobManager>>,
/// Prompt queue for Claude Code follow-up prompts.
pub prompt_queue: Option<PromptQueue>,
/// Default user ID (fallback for non-request contexts like heartbeat/routines).
pub default_user_id: String,
/// Durable owner scope for persistence and unauthenticated callback flows.
pub owner_id: String,
/// Default sender/routing identity for gateway-originated messages.
pub default_sender_id: String,
/// Shutdown signal sender.
pub shutdown_tx: tokio::sync::RwLock<Option<oneshot::Sender<()>>>,
/// WebSocket connection tracker.
@@ -800,7 +802,7 @@ async fn oauth_callback_handler(
error = %error,
"OAuth callback received with malformed state"
);
clear_auth_mode(&state, &state.default_user_id).await;
clear_auth_mode(&state, &state.owner_id).await;
return oauth_error_page("IronClaw");
}
};
@@ -836,7 +838,7 @@ async fn oauth_callback_handler(
if let Some(ref sse) = flow.sse_manager {
sse.broadcast_for_user(
&flow.user_id,
SseEvent::AuthCompleted {
AppEvent::AuthCompleted {
extension_name: flow.extension_name.clone(),
success: false,
message: "OAuth flow expired. Please try again.".to_string(),
@@ -974,11 +976,11 @@ async fn oauth_callback_handler(
message
};
// Broadcast SSE event to notify the web UI
// Broadcast event to notify the web UI
if let Some(ref sse) = flow.sse_manager {
sse.broadcast_for_user(
&flow.user_id,
SseEvent::AuthCompleted {
AppEvent::AuthCompleted {
extension_name: flow.extension_name,
success,
message: final_message.clone(),
@@ -1161,7 +1163,7 @@ async fn slack_relay_oauth_callback_handler(
let state_key = format!("relay:{}:oauth_state", DEFAULT_RELAY_NAME);
let stored_state = match ext_mgr
.secrets()
.get_decrypted(&state.default_user_id, &state_key)
.get_decrypted(&state.owner_id, &state_key)
.await
{
Ok(secret) => secret.expose().to_string(),
@@ -1185,10 +1187,7 @@ async fn slack_relay_oauth_callback_handler(
}
// Delete the nonce (one-time use)
let _ = ext_mgr
.secrets()
.delete(&state.default_user_id, &state_key)
.await;
let _ = ext_mgr.secrets().delete(&state.owner_id, &state_key).await;
let result: Result<(), String> = async {
let store = state.store.as_ref().ok_or_else(|| {
@@ -1199,16 +1198,12 @@ async fn slack_relay_oauth_callback_handler(
// Store team_id in settings
let team_id_key = format!("relay:{}:team_id", DEFAULT_RELAY_NAME);
let _ = store
.set_setting(
&state.default_user_id,
&team_id_key,
&serde_json::json!(team_id),
)
.set_setting(&state.owner_id, &team_id_key, &serde_json::json!(team_id))
.await;
// Activate the relay channel
ext_mgr
.activate_stored_relay(DEFAULT_RELAY_NAME, &state.default_user_id)
.activate_stored_relay(DEFAULT_RELAY_NAME, &state.owner_id)
.await
.map_err(|e| format!("Failed to activate relay channel: {}", e))?;
@@ -1227,8 +1222,8 @@ async fn slack_relay_oauth_callback_handler(
}
};
// Broadcast SSE event to notify the web UI
state.sse.broadcast(SseEvent::AuthCompleted {
// Broadcast event to notify the web UI
state.sse.broadcast(AppEvent::AuthCompleted {
extension_name: DEFAULT_RELAY_NAME.to_string(),
success,
message: message.clone(),
@@ -1328,6 +1323,9 @@ async fn chat_send_handler(
}
let mut msg = IncomingMessage::new("gateway", &user.user_id, &req.content);
if state.owner_id != state.default_sender_id && user.user_id == state.owner_id {
msg = msg.with_sender_id(&state.default_sender_id);
}
// Prefer timezone from JSON body, fall back to X-Timezone header
let tz = req
.timezone
@@ -1429,6 +1427,9 @@ async fn chat_approval_handler(
})?;
let mut msg = IncomingMessage::new("gateway", &user.user_id, content);
if state.owner_id != state.default_sender_id && user.user_id == state.owner_id {
msg = msg.with_sender_id(&state.default_sender_id);
}
if let Some(ref thread_id) = req.thread_id {
msg = msg.with_thread(thread_id);
@@ -1495,7 +1496,7 @@ async fn chat_auth_token_handler(
if result.verification.is_some() {
state.sse.broadcast_for_user(
&user.user_id,
SseEvent::AuthRequired {
AppEvent::AuthRequired {
extension_name: req.extension_name.clone(),
instructions: Some(result.message),
auth_url: None,
@@ -1508,7 +1509,7 @@ async fn chat_auth_token_handler(
state.sse.broadcast_for_user(
&user.user_id,
SseEvent::AuthCompleted {
AppEvent::AuthCompleted {
extension_name: req.extension_name.clone(),
success: true,
message: result.message,
@@ -1517,7 +1518,7 @@ async fn chat_auth_token_handler(
} else {
state.sse.broadcast_for_user(
&user.user_id,
SseEvent::AuthCompleted {
AppEvent::AuthCompleted {
extension_name: req.extension_name.clone(),
success: false,
message: result.message,
@@ -1533,7 +1534,7 @@ async fn chat_auth_token_handler(
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
state.sse.broadcast_for_user(
&user.user_id,
SseEvent::AuthRequired {
AppEvent::AuthRequired {
extension_name: req.extension_name.clone(),
instructions: Some(msg.clone()),
auth_url: None,
@@ -2501,7 +2502,7 @@ async fn extensions_setup_submit_handler(
// auth card or setup modal that was triggered by tool_auth/tool_activate.
state.sse.broadcast_for_user(
&user.user_id,
SseEvent::AuthCompleted {
AppEvent::AuthCompleted {
extension_name: name.clone(),
success: result.activated,
message: resp.message.clone(),
@@ -3403,7 +3404,8 @@ mod tests {
store: None,
job_manager: None,
prompt_queue: None,
default_user_id: "test".to_string(),
owner_id: "test".to_string(),
default_sender_id: "test".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: None,
llm_provider: None,
@@ -3595,7 +3597,7 @@ mod tests {
Ok(Ok(scoped))
if matches!(
scoped.event,
crate::channels::web::types::SseEvent::AuthRequired { .. }
crate::channels::web::types::AppEvent::AuthRequired { .. }
) =>
{
panic!("verification responses should not emit auth_required SSE events")
@@ -3877,7 +3879,7 @@ mod tests {
assert_eq!(resp.status(), StatusCode::OK);
match receiver.recv().await.expect("auth_completed event").event {
crate::channels::web::types::SseEvent::AuthCompleted {
crate::channels::web::types::AppEvent::AuthCompleted {
extension_name,
success,
message,
+19 -42
View File
@@ -11,7 +11,7 @@ use tokio::sync::broadcast;
use tokio_stream::StreamExt;
use tokio_stream::wrappers::BroadcastStream;
use crate::channels::web::types::SseEvent;
use crate::channels::web::types::AppEvent;
/// Maximum number of concurrent SSE/WebSocket connections.
/// Prevents resource exhaustion from connection flooding.
@@ -25,7 +25,7 @@ const MAX_CONNECTIONS: u64 = 100;
#[derive(Debug, Clone)]
pub(crate) struct ScopedEvent {
pub(crate) user_id: Option<String>,
pub(crate) event: SseEvent,
pub(crate) event: AppEvent,
}
/// Manages SSE broadcast to all connected browser tabs.
@@ -75,7 +75,7 @@ impl SseManager {
}
/// Broadcast an event to all connected clients (global/unscoped).
pub fn broadcast(&self, event: SseEvent) {
pub fn broadcast(&self, event: AppEvent) {
let _ = self.tx.send(ScopedEvent {
user_id: None,
event,
@@ -86,7 +86,7 @@ impl SseManager {
///
/// Only subscribers for this user_id (or unscoped subscribers) will
/// receive the event.
pub fn broadcast_for_user(&self, user_id: &str, event: SseEvent) {
pub fn broadcast_for_user(&self, user_id: &str, event: AppEvent) {
let _ = self.tx.send(ScopedEvent {
user_id: Some(user_id.to_string()),
event,
@@ -108,7 +108,7 @@ impl SseManager {
pub fn subscribe_raw(
&self,
user_id: Option<String>,
) -> Option<impl Stream<Item = SseEvent> + Send + 'static + use<>> {
) -> Option<impl Stream<Item = AppEvent> + Send + 'static + use<>> {
// Atomically increment only if below the limit. This prevents
// concurrent callers from overshooting max_connections.
let counter = Arc::clone(&self.connection_count);
@@ -186,30 +186,7 @@ impl SseManager {
return None;
}
};
let event_type = match &event {
SseEvent::Response { .. } => "response",
SseEvent::Thinking { .. } => "thinking",
SseEvent::ToolStarted { .. } => "tool_started",
SseEvent::ToolCompleted { .. } => "tool_completed",
SseEvent::ToolResult { .. } => "tool_result",
SseEvent::StreamChunk { .. } => "stream_chunk",
SseEvent::Status { .. } => "status",
SseEvent::ApprovalNeeded { .. } => "approval_needed",
SseEvent::AuthRequired { .. } => "auth_required",
SseEvent::AuthCompleted { .. } => "auth_completed",
SseEvent::Error { .. } => "error",
SseEvent::JobStarted { .. } => "job_started",
SseEvent::JobMessage { .. } => "job_message",
SseEvent::JobToolUse { .. } => "job_tool_use",
SseEvent::JobToolResult { .. } => "job_tool_result",
SseEvent::JobStatus { .. } => "job_status",
SseEvent::JobResult { .. } => "job_result",
SseEvent::Heartbeat => "heartbeat",
SseEvent::ImageGenerated { .. } => "image_generated",
SseEvent::Suggestions { .. } => "suggestions",
SseEvent::TurnCost { .. } => "turn_cost",
SseEvent::ExtensionStatus { .. } => "extension_status",
};
let event_type = event.event_type();
Some(Ok(Event::default().event(event_type).data(data)))
});
@@ -272,7 +249,7 @@ mod tests {
fn test_broadcast_without_receivers() {
let manager = SseManager::new();
// Should not panic even with no receivers
manager.broadcast(SseEvent::Heartbeat);
manager.broadcast(AppEvent::Heartbeat);
}
#[tokio::test]
@@ -280,14 +257,14 @@ mod tests {
let manager = SseManager::new();
let mut stream = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
manager.broadcast(SseEvent::Status {
manager.broadcast(AppEvent::Status {
message: "test".to_string(),
thread_id: None,
});
let event = stream.next().await.unwrap();
match event {
SseEvent::Status { message, .. } => assert_eq!(message, "test"),
AppEvent::Status { message, .. } => assert_eq!(message, "test"),
_ => panic!("unexpected event type"),
}
}
@@ -299,14 +276,14 @@ mod tests {
assert_eq!(manager.connection_count(), 1);
manager.broadcast(SseEvent::Thinking {
manager.broadcast(AppEvent::Thinking {
message: "working".to_string(),
thread_id: None,
});
let event = stream.next().await.unwrap();
match event {
SseEvent::Thinking { message, .. } => assert_eq!(message, "working"),
AppEvent::Thinking { message, .. } => assert_eq!(message, "working"),
_ => panic!("Expected Thinking event"),
}
}
@@ -329,12 +306,12 @@ mod tests {
let mut s2 = Box::pin(manager.subscribe_raw(None).expect("should subscribe"));
assert_eq!(manager.connection_count(), 2);
manager.broadcast(SseEvent::Heartbeat);
manager.broadcast(AppEvent::Heartbeat);
let e1 = s1.next().await.unwrap();
let e2 = s2.next().await.unwrap();
assert!(matches!(e1, SseEvent::Heartbeat));
assert!(matches!(e2, SseEvent::Heartbeat));
assert!(matches!(e1, AppEvent::Heartbeat));
assert!(matches!(e2, AppEvent::Heartbeat));
drop(s1);
assert_eq!(manager.connection_count(), 1);
@@ -373,25 +350,25 @@ mod tests {
// Send event scoped to alice
manager.broadcast_for_user(
"alice",
SseEvent::Status {
AppEvent::Status {
message: "alice only".to_string(),
thread_id: None,
},
);
// Send global event
manager.broadcast(SseEvent::Heartbeat);
manager.broadcast(AppEvent::Heartbeat);
// Alice gets her scoped event
let e = alice.next().await.unwrap();
assert!(matches!(e, SseEvent::Status { .. }));
assert!(matches!(e, AppEvent::Status { .. }));
// Alice also gets the global heartbeat
let e = alice.next().await.unwrap();
assert!(matches!(e, SseEvent::Heartbeat));
assert!(matches!(e, AppEvent::Heartbeat));
// Bob only gets the global heartbeat (alice's event was filtered)
let e = bob.next().await.unwrap(); // safety: test-only
assert!(matches!(e, SseEvent::Heartbeat)); // safety: test assertion
assert!(matches!(e, AppEvent::Heartbeat)); // safety: test assertion
}
}
+2 -1
View File
@@ -76,7 +76,8 @@ impl TestGatewayBuilder {
store: None,
job_manager: None,
prompt_queue: None,
default_user_id: self.user_id,
owner_id: self.user_id.clone(),
default_sender_id: self.user_id,
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: self.llm_provider,
+38 -1
View File
@@ -16,6 +16,7 @@ use axum::routing::{delete, get, post};
use tower::ServiceExt;
use uuid::Uuid;
use crate::channels::web::GatewayChannel;
use crate::channels::web::auth::{
AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware,
};
@@ -23,6 +24,7 @@ use crate::channels::web::server::{
ActiveConfigSnapshot, GatewayState, PerUserRateLimiter, PromptQueue, RateLimiter, WorkspacePool,
};
use crate::channels::web::sse::SseManager;
use crate::config::GatewayConfig;
// ── Helpers ────────────────────────────────────────────────────────────
@@ -64,7 +66,8 @@ fn build_state(
store,
job_manager: None,
prompt_queue,
default_user_id: "test".to_string(),
owner_id: "test".to_string(),
default_sender_id: "test".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: None,
llm_provider: None,
@@ -83,6 +86,40 @@ fn build_state(
})
}
fn gateway_config() -> GatewayConfig {
GatewayConfig {
host: "127.0.0.1".to_string(),
port: 3000,
auth_token: Some("gateway-auth".to_string()),
user_id: "gateway-sender".to_string(),
workspace_read_scopes: Vec::new(),
memory_layers: Vec::new(),
user_tokens: None,
}
}
#[test]
fn with_owner_scope_updates_gateway_owner_scope_in_multi_user_mode() {
let mut gateway = GatewayChannel::new(gateway_config());
gateway.auth = two_user_auth();
gateway.config.user_tokens = Some(HashMap::new());
let gateway = gateway.with_owner_scope("owner-scope");
assert_eq!(gateway.state.owner_id, "owner-scope");
assert_eq!(gateway.state.default_sender_id, "gateway-sender");
let alice = gateway
.auth
.authenticate("tok-alice")
.expect("alice token should remain valid");
let bob = gateway
.auth
.authenticate("tok-bob")
.expect("bob token should remain valid");
assert_eq!(alice.user_id, "alice");
assert_eq!(bob.user_id, "bob");
}
/// Create a libSQL-backed test database in a temporary directory.
///
/// Returns the database and a `TempDir` guard — the database file is
+27 -206
View File
@@ -114,165 +114,9 @@ pub struct ApprovalRequest {
pub thread_id: Option<String>,
}
// --- SSE Event Types ---
// --- App Event (re-exported from ironclaw_common) ---
#[derive(Debug, Clone, Serialize)]
#[serde(tag = "type")]
pub enum SseEvent {
#[serde(rename = "response")]
Response { content: String, thread_id: String },
#[serde(rename = "thinking")]
Thinking {
message: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "tool_started")]
ToolStarted {
name: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "tool_completed")]
ToolCompleted {
name: String,
success: bool,
#[serde(skip_serializing_if = "Option::is_none")]
error: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
parameters: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "tool_result")]
ToolResult {
name: String,
preview: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "stream_chunk")]
StreamChunk {
content: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "status")]
Status {
message: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "job_started")]
JobStarted {
job_id: String,
title: String,
browse_url: String,
},
#[serde(rename = "approval_needed")]
ApprovalNeeded {
request_id: String,
tool_name: String,
description: String,
parameters: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
/// Whether the "always" auto-approve option should be shown.
allow_always: bool,
},
#[serde(rename = "auth_required")]
AuthRequired {
extension_name: String,
#[serde(skip_serializing_if = "Option::is_none")]
instructions: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
auth_url: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
setup_url: Option<String>,
},
#[serde(rename = "auth_completed")]
AuthCompleted {
extension_name: String,
success: bool,
message: String,
},
#[serde(rename = "error")]
Error {
message: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
#[serde(rename = "heartbeat")]
Heartbeat,
// Sandbox job streaming events (worker + Claude Code bridge)
#[serde(rename = "job_message")]
JobMessage {
job_id: String,
role: String,
content: String,
},
#[serde(rename = "job_tool_use")]
JobToolUse {
job_id: String,
tool_name: String,
input: serde_json::Value,
},
#[serde(rename = "job_tool_result")]
JobToolResult {
job_id: String,
tool_name: String,
output: String,
},
#[serde(rename = "job_status")]
JobStatus { job_id: String, message: String },
#[serde(rename = "job_result")]
JobResult {
job_id: String,
status: String,
#[serde(skip_serializing_if = "Option::is_none")]
session_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
fallback_deliverable: Option<serde_json::Value>,
},
/// An image was generated by a tool.
#[serde(rename = "image_generated")]
ImageGenerated {
data_url: String,
#[serde(skip_serializing_if = "Option::is_none")]
path: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
/// Suggested follow-up messages for the user.
#[serde(rename = "suggestions")]
Suggestions {
suggestions: Vec<String>,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
/// Per-turn token usage and cost summary.
#[serde(rename = "turn_cost")]
TurnCost {
input_tokens: u64,
output_tokens: u64,
cost_usd: String,
#[serde(skip_serializing_if = "Option::is_none")]
thread_id: Option<String>,
},
/// Extension activation status change (WASM channels).
#[serde(rename = "extension_status")]
ExtensionStatus {
extension_name: String,
status: String,
#[serde(skip_serializing_if = "Option::is_none")]
message: Option<String>,
},
}
pub use ironclaw_common::AppEvent;
// --- Memory ---
@@ -784,32 +628,9 @@ pub enum WsServerMessage {
}
impl WsServerMessage {
/// Create a WsServerMessage from an SseEvent.
pub fn from_sse_event(event: &SseEvent) -> Self {
let event_type = match event {
SseEvent::Response { .. } => "response",
SseEvent::Thinking { .. } => "thinking",
SseEvent::ToolStarted { .. } => "tool_started",
SseEvent::ToolCompleted { .. } => "tool_completed",
SseEvent::ToolResult { .. } => "tool_result",
SseEvent::StreamChunk { .. } => "stream_chunk",
SseEvent::Status { .. } => "status",
SseEvent::JobStarted { .. } => "job_started",
SseEvent::ApprovalNeeded { .. } => "approval_needed",
SseEvent::AuthRequired { .. } => "auth_required",
SseEvent::AuthCompleted { .. } => "auth_completed",
SseEvent::Error { .. } => "error",
SseEvent::Heartbeat => "heartbeat",
SseEvent::JobMessage { .. } => "job_message",
SseEvent::JobToolUse { .. } => "job_tool_use",
SseEvent::JobToolResult { .. } => "job_tool_result",
SseEvent::JobStatus { .. } => "job_status",
SseEvent::JobResult { .. } => "job_result",
SseEvent::ImageGenerated { .. } => "image_generated",
SseEvent::Suggestions { .. } => "suggestions",
SseEvent::TurnCost { .. } => "turn_cost",
SseEvent::ExtensionStatus { .. } => "extension_status",
};
/// Create a WsServerMessage from an AppEvent.
pub fn from_app_event(event: &AppEvent) -> Self {
let event_type = event.event_type();
let data = serde_json::to_value(event).unwrap_or(serde_json::Value::Null);
WsServerMessage::Event {
event_type: event_type.to_string(),
@@ -1101,12 +922,12 @@ mod tests {
}
#[test]
fn test_ws_server_from_sse_response() {
let sse = SseEvent::Response {
fn test_ws_server_from_app_event_response() {
let event = AppEvent::Response {
content: "hello".to_string(),
thread_id: "t1".to_string(),
};
let ws = WsServerMessage::from_sse_event(&sse);
let ws = WsServerMessage::from_app_event(&event);
match ws {
WsServerMessage::Event { event_type, data } => {
assert_eq!(event_type, "response");
@@ -1118,12 +939,12 @@ mod tests {
}
#[test]
fn test_ws_server_from_sse_thinking() {
let sse = SseEvent::Thinking {
fn test_ws_server_from_app_event_thinking() {
let event = AppEvent::Thinking {
message: "reasoning...".to_string(),
thread_id: None,
};
let ws = WsServerMessage::from_sse_event(&sse);
let ws = WsServerMessage::from_app_event(&event);
match ws {
WsServerMessage::Event { event_type, data } => {
assert_eq!(event_type, "thinking");
@@ -1134,8 +955,8 @@ mod tests {
}
#[test]
fn test_ws_server_from_sse_approval_needed() {
let sse = SseEvent::ApprovalNeeded {
fn test_ws_server_from_app_event_approval_needed() {
let event = AppEvent::ApprovalNeeded {
request_id: "r1".to_string(),
tool_name: "shell".to_string(),
description: "Run ls".to_string(),
@@ -1143,7 +964,7 @@ mod tests {
thread_id: Some("t1".to_string()),
allow_always: true,
};
let ws = WsServerMessage::from_sse_event(&sse);
let ws = WsServerMessage::from_app_event(&event);
match ws {
WsServerMessage::Event { event_type, data } => {
assert_eq!(event_type, "approval_needed");
@@ -1155,9 +976,9 @@ mod tests {
}
#[test]
fn test_ws_server_from_sse_heartbeat() {
let sse = SseEvent::Heartbeat;
let ws = WsServerMessage::from_sse_event(&sse);
fn test_ws_server_from_app_event_heartbeat() {
let event = AppEvent::Heartbeat;
let ws = WsServerMessage::from_app_event(&event);
match ws {
WsServerMessage::Event { event_type, .. } => {
assert_eq!(event_type, "heartbeat");
@@ -1197,8 +1018,8 @@ mod tests {
}
#[test]
fn test_sse_auth_required_serialize() {
let event = SseEvent::AuthRequired {
fn test_app_event_auth_required_serialize() {
let event = AppEvent::AuthRequired {
extension_name: "notion".to_string(),
instructions: Some("Get your token from...".to_string()),
auth_url: None,
@@ -1214,8 +1035,8 @@ mod tests {
}
#[test]
fn test_sse_auth_completed_serialize() {
let event = SseEvent::AuthCompleted {
fn test_app_event_auth_completed_serialize() {
let event = AppEvent::AuthCompleted {
extension_name: "notion".to_string(),
success: true,
message: "notion authenticated (3 tools loaded)".to_string(),
@@ -1228,14 +1049,14 @@ mod tests {
}
#[test]
fn test_ws_server_from_sse_auth_required() {
let sse = SseEvent::AuthRequired {
fn test_ws_server_from_app_event_auth_required() {
let event = AppEvent::AuthRequired {
extension_name: "openai".to_string(),
instructions: Some("Enter API key".to_string()),
auth_url: None,
setup_url: None,
};
let ws = WsServerMessage::from_sse_event(&sse);
let ws = WsServerMessage::from_app_event(&event);
match ws {
WsServerMessage::Event { event_type, data } => {
assert_eq!(event_type, "auth_required");
@@ -1246,13 +1067,13 @@ mod tests {
}
#[test]
fn test_ws_server_from_sse_auth_completed() {
let sse = SseEvent::AuthCompleted {
fn test_ws_server_from_app_event_auth_completed() {
let event = AppEvent::AuthCompleted {
extension_name: "slack".to_string(),
success: false,
message: "Invalid token".to_string(),
};
let ws = WsServerMessage::from_sse_event(&sse);
let ws = WsServerMessage::from_app_event(&event);
match ws {
WsServerMessage::Event { event_type, data } => {
assert_eq!(event_type, "auth_completed");
+1 -105
View File
@@ -2,29 +2,7 @@
use crate::channels::web::types::{ToolCallInfo, TurnInfo};
/// Truncate a string to at most `max_bytes` bytes at a char boundary, appending "...".
///
/// If the input is wrapped in `<tool_output …>…</tool_output>` and truncation
/// removes the closing tag, the tag is re-appended so downstream XML parsers
/// never see an unclosed element.
pub fn truncate_preview(s: &str, max_bytes: usize) -> String {
if s.len() <= max_bytes {
return s.to_string();
}
// Walk backwards from max_bytes to find a valid char boundary
let mut end = max_bytes;
while end > 0 && !s.is_char_boundary(end) {
end -= 1;
}
let mut result = format!("{}...", &s[..end]);
// Re-close <tool_output> if truncation cut through the closing tag.
if s.starts_with("<tool_output") && !result.ends_with("</tool_output>") {
result.push_str("\n</tool_output>");
}
result
}
pub use ironclaw_common::truncate_preview;
/// Build TurnInfo pairs from flat DB messages (user/tool_calls/assistant triples).
///
@@ -118,88 +96,6 @@ mod tests {
use super::*;
use uuid::Uuid;
// ---- truncate_preview tests ----
#[test]
fn test_truncate_preview_short_string() {
assert_eq!(truncate_preview("hello", 10), "hello");
}
#[test]
fn test_truncate_preview_exact_boundary() {
assert_eq!(truncate_preview("hello", 5), "hello");
}
#[test]
fn test_truncate_preview_truncates_ascii() {
assert_eq!(truncate_preview("hello world", 5), "hello...");
}
#[test]
fn test_truncate_preview_empty_string() {
assert_eq!(truncate_preview("", 10), "");
}
#[test]
fn test_truncate_preview_multibyte_char_boundary() {
// '€' is 3 bytes (E2 82 AC). "a€b" = [61, E2, 82, AC, 62] = 5 bytes
// Truncating at max_bytes=3 should not split the euro sign.
let s = "a€b";
let result = truncate_preview(s, 3);
// max_bytes=3 lands mid-€, so it walks back to byte 1 ("a")
assert_eq!(result, "a...");
}
#[test]
fn test_truncate_preview_emoji() {
// '🦀' is 4 bytes. "hi🦀" = 6 bytes
let s = "hi🦀";
let result = truncate_preview(s, 4);
// max_bytes=4 lands mid-🦀, walks back to byte 2 ("hi")
assert_eq!(result, "hi...");
}
#[test]
fn test_truncate_preview_cjk() {
// CJK characters are 3 bytes each. "你好世界" = 12 bytes
let s = "你好世界";
let result = truncate_preview(s, 7);
// max_bytes=7 lands mid-character (byte 7 is inside 世), walks back to 6 ("你好")
assert_eq!(result, "你好...");
}
#[test]
fn test_truncate_preview_zero_max_bytes() {
assert_eq!(truncate_preview("hello", 0), "...");
}
#[test]
fn test_truncate_preview_closes_tool_output_tag() {
let s = "<tool_output name=\"search\">\nSome very long content here\n</tool_output>";
// Truncate so it cuts before the closing tag
let result = truncate_preview(s, 60);
assert!(result.ends_with("</tool_output>"));
assert!(result.contains("..."));
}
#[test]
fn test_truncate_preview_no_extra_close_when_intact() {
let s = "<tool_output name=\"echo\">\nshort\n</tool_output>";
// The string is short enough not to be truncated
let result = truncate_preview(s, 500);
assert_eq!(result, s);
// Should not have a duplicate closing tag
assert_eq!(result.matches("</tool_output>").count(), 1);
}
#[test]
fn test_truncate_preview_non_xml_unaffected() {
let s = "Just a plain long string that gets truncated";
let result = truncate_preview(s, 10);
assert_eq!(result, "Just a pla...");
assert!(!result.contains("</tool_output>"));
}
// ---- build_turns_from_db_messages tests ----
fn make_msg(role: &str, content: &str, offset_ms: i64) -> crate::history::ConversationMessage {
+6 -5
View File
@@ -97,7 +97,7 @@ pub async fn handle_ws_connection(
let msg = tokio::select! {
event = event_stream.next() => {
match event {
Some(sse_event) => WsServerMessage::from_sse_event(&sse_event),
Some(app_event) => WsServerMessage::from_app_event(&app_event),
None => break, // Broadcast channel closed
}
}
@@ -275,7 +275,7 @@ async fn handle_client_message(
if result.verification.is_some() {
state.sse.broadcast_for_user(
user_id,
crate::channels::web::types::SseEvent::AuthRequired {
crate::channels::web::types::AppEvent::AuthRequired {
extension_name: extension_name.clone(),
instructions: Some(result.message),
auth_url: None,
@@ -286,7 +286,7 @@ async fn handle_client_message(
crate::channels::web::server::clear_auth_mode(state, user_id).await;
state.sse.broadcast_for_user(
user_id,
crate::channels::web::types::SseEvent::AuthCompleted {
crate::channels::web::types::AppEvent::AuthCompleted {
extension_name,
success: true,
message: result.message,
@@ -299,7 +299,7 @@ async fn handle_client_message(
if matches!(e, crate::extensions::ExtensionError::ValidationFailed(_)) {
state.sse.broadcast_for_user(
user_id,
crate::channels::web::types::SseEvent::AuthRequired {
crate::channels::web::types::AppEvent::AuthRequired {
extension_name: extension_name.clone(),
instructions: Some(msg.clone()),
auth_url: None,
@@ -520,7 +520,8 @@ mod tests {
job_manager: None,
prompt_queue: None,
scheduler: None,
default_user_id: "test".to_string(),
owner_id: "test".to_string(),
default_sender_id: "test".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: None,
+400 -24
View File
@@ -62,6 +62,30 @@ pub fn builtin_client_id_override_env(secret_name: &str) -> Option<&'static str>
}
}
/// Suppress the baked-in desktop OAuth client secret when a hosted proxy is configured.
///
/// In hosted deployments, IronClaw may resolve the platform Google client ID from
/// environment variables while still falling back to the baked-in desktop secret.
/// That client_id/client_secret mismatch breaks Google token exchange and refresh.
///
/// When the proxy is configured, the platform will inject the correct server-side
/// secret for matching platform credentials, so the baked-in secret must be omitted.
pub fn hosted_proxy_client_secret(
client_secret: &Option<String>,
builtin: Option<&OAuthCredentials>,
exchange_proxy_configured: bool,
) -> Option<String> {
if !exchange_proxy_configured {
return client_secret.clone();
}
let builtin_secret = builtin.map(|credentials| credentials.client_secret);
match (client_secret, builtin_secret) {
(Some(resolved), Some(baked_in)) if resolved == baked_in => None,
_ => client_secret.clone(),
}
}
// ── Shared callback server ──────────────────────────────────────────────
// Core OAuth callback infrastructure is defined in `crate::llm::oauth_helpers`
@@ -661,6 +685,48 @@ pub struct ProxyTokenExchangeRequest<'a> {
pub extra_token_params: &'a HashMap<String, String>,
}
pub struct ProxyRefreshTokenRequest<'a> {
pub proxy_url: &'a str,
pub gateway_token: &'a str,
pub token_url: &'a str,
pub client_id: &'a str,
pub client_secret: Option<&'a str>,
pub refresh_token: &'a str,
pub provider: Option<&'a str>,
}
fn oauth_token_response_from_json(
token_data: serde_json::Value,
access_token_field: &str,
) -> Result<OAuthTokenResponse, OAuthCallbackError> {
let access_token = token_data
.get(access_token_field)
.and_then(|v| v.as_str())
.ok_or_else(|| {
let fields: Vec<&str> = token_data
.as_object()
.map(|o| o.keys().map(|k| k.as_str()).collect())
.unwrap_or_default();
OAuthCallbackError::Io(format!(
"No '{}' field in proxy response (fields present: {:?})",
access_token_field, fields
))
})?
.to_string();
let refresh_token = token_data
.get("refresh_token")
.and_then(|v| v.as_str())
.map(String::from);
let expires_in = token_data.get("expires_in").and_then(|v| v.as_u64());
Ok(OAuthTokenResponse {
access_token,
refresh_token,
expires_in,
})
}
/// Exchange an OAuth authorization code via the platform's token exchange proxy.
///
/// Authenticated via the gateway auth token (Bearer header). The caller may
@@ -682,6 +748,7 @@ pub async fn exchange_via_proxy(
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(60))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| OAuthCallbackError::Io(format!("Failed to build HTTP client: {}", e)))?;
let mut params = vec![
@@ -724,41 +791,350 @@ pub async fn exchange_via_proxy(
.json()
.await
.map_err(|e| OAuthCallbackError::Io(format!("Failed to parse proxy response: {}", e)))?;
oauth_token_response_from_json(token_data, request.access_token_field)
}
let access_token = token_data
.get(request.access_token_field)
.and_then(|v| v.as_str())
.ok_or_else(|| {
let fields: Vec<&str> = token_data
.as_object()
.map(|o| o.keys().map(|k| k.as_str()).collect())
.unwrap_or_default();
OAuthCallbackError::Io(format!(
"No '{}' field in proxy response (fields present: {:?})",
request.access_token_field, fields
))
})?
.to_string();
/// Refresh an OAuth access token via the platform's token refresh proxy.
///
/// Authenticated via the gateway auth token (Bearer header). The caller may
/// either rely on proxy-side secret lookup or forward a `client_secret` when
/// the provider requires it.
pub async fn refresh_token_via_proxy(
request: ProxyRefreshTokenRequest<'_>,
) -> Result<OAuthTokenResponse, OAuthCallbackError> {
if request.gateway_token.is_empty() {
return Err(OAuthCallbackError::Io(
"Gateway auth token is required for proxy token refresh".to_string(),
));
}
let refresh_token = token_data
.get("refresh_token")
.and_then(|v| v.as_str())
.map(String::from);
let expires_in = token_data.get("expires_in").and_then(|v| v.as_u64());
let refresh_url = format!("{}/oauth/refresh", request.proxy_url.trim_end_matches('/'));
let client = reqwest::Client::builder()
.timeout(Duration::from_secs(15))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(|e| OAuthCallbackError::Io(format!("Failed to build HTTP client: {}", e)))?;
Ok(OAuthTokenResponse {
access_token,
refresh_token,
expires_in,
})
let mut params = vec![
("refresh_token", request.refresh_token.to_string()),
("token_url", request.token_url.to_string()),
("client_id", request.client_id.to_string()),
];
if let Some(secret) = request.client_secret {
params.push(("client_secret", secret.to_string()));
}
if let Some(provider) = request.provider {
params.push(("provider", provider.to_string()));
}
let response = client
.post(&refresh_url)
.bearer_auth(request.gateway_token)
.form(&params)
.send()
.await
.map_err(|e| {
OAuthCallbackError::Io(format!("Token refresh proxy request failed: {}", e))
})?;
if !response.status().is_success() {
let status = response.status();
let body = response.text().await.unwrap_or_default();
return Err(OAuthCallbackError::Io(format!(
"Token refresh proxy failed: {} - {}",
status, body
)));
}
let token_data: serde_json::Value = response
.json()
.await
.map_err(|e| OAuthCallbackError::Io(format!("Failed to parse proxy response: {}", e)))?;
oauth_token_response_from_json(token_data, "access_token")
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::Arc;
use axum::extract::{Form, State};
use axum::http::HeaderMap;
use axum::response::Redirect;
use axum::routing::post;
use axum::{Json, Router};
use serde_json::json;
use tokio::net::TcpListener;
use tokio::sync::{Mutex, oneshot};
use crate::cli::oauth_defaults::{
builtin_credentials, callback_host, callback_url, is_loopback_host, landing_html,
};
use crate::config::helpers::lock_env;
use crate::testing::credentials::{TEST_OAUTH_CLIENT_ID, TEST_OAUTH_CLIENT_SECRET};
#[derive(Clone, Debug, PartialEq, Eq)]
struct RecordedProxyRequest {
authorization: Option<String>,
form: HashMap<String, String>,
}
#[derive(Clone)]
struct MockProxyState {
requests: Arc<Mutex<Vec<RecordedProxyRequest>>>,
exchange_redirect_target: String,
refresh_redirect_target: String,
}
struct MockProxyServer {
addr: SocketAddr,
requests: Arc<Mutex<Vec<RecordedProxyRequest>>>,
shutdown_tx: Option<oneshot::Sender<()>>,
server_task: Option<tokio::task::JoinHandle<()>>,
}
impl MockProxyServer {
async fn start() -> Self {
async fn exchange_handler(
State(state): State<MockProxyState>,
headers: HeaderMap,
Form(form): Form<HashMap<String, String>>,
) -> Json<serde_json::Value> {
state.requests.lock().await.push(RecordedProxyRequest {
authorization: headers
.get(axum::http::header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.map(str::to_string),
form,
});
Json(json!({
"access_token": "proxy-access-token",
"refresh_token": "proxy-refresh-token",
"expires_in": 7200
}))
}
async fn refresh_handler(
State(state): State<MockProxyState>,
headers: HeaderMap,
Form(form): Form<HashMap<String, String>>,
) -> Json<serde_json::Value> {
state.requests.lock().await.push(RecordedProxyRequest {
authorization: headers
.get(axum::http::header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.map(str::to_string),
form,
});
Json(json!({
"access_token": "proxy-access-token",
"refresh_token": "proxy-refresh-token",
"expires_in": 7200
}))
}
async fn exchange_redirect_handler(State(state): State<MockProxyState>) -> Redirect {
Redirect::temporary(&state.exchange_redirect_target)
}
async fn refresh_redirect_handler(State(state): State<MockProxyState>) -> Redirect {
Redirect::temporary(&state.refresh_redirect_target)
}
let requests = Arc::new(Mutex::new(Vec::new()));
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("bind mock proxy");
let addr = listener.local_addr().expect("read mock proxy addr");
let exchange_redirect_target = format!("http://{addr}/oauth/exchange");
let refresh_redirect_target = format!("http://{addr}/oauth/refresh");
let app = Router::new()
.route("/oauth/exchange", post(exchange_handler))
.route("/oauth/refresh", post(refresh_handler))
.route("/redirect/oauth/exchange", post(exchange_redirect_handler))
.route("/redirect/oauth/refresh", post(refresh_redirect_handler))
.with_state(MockProxyState {
requests: Arc::clone(&requests),
exchange_redirect_target,
refresh_redirect_target,
});
let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>();
let server_task = tokio::spawn(async move {
let _ = axum::serve(listener, app)
.with_graceful_shutdown(async {
let _ = shutdown_rx.await;
})
.await;
});
Self {
addr,
requests,
shutdown_tx: Some(shutdown_tx),
server_task: Some(server_task),
}
}
fn base_url(&self) -> String {
format!("http://{}", self.addr)
}
fn redirecting_base_url(&self) -> String {
format!("{}/redirect", self.base_url())
}
async fn requests(&self) -> Vec<RecordedProxyRequest> {
self.requests.lock().await.clone()
}
async fn shutdown(mut self) {
if let Some(tx) = self.shutdown_tx.take() {
let _ = tx.send(());
}
if let Some(task) = self.server_task.take() {
let _ = task.await;
}
}
}
impl Drop for MockProxyServer {
fn drop(&mut self) {
if let Some(tx) = self.shutdown_tx.take() {
let _ = tx.send(());
}
if let Some(task) = self.server_task.take() {
task.abort();
}
}
}
#[test]
fn test_hosted_proxy_client_secret_suppresses_builtin_secret() {
let builtin = builtin_credentials("google_oauth_token").expect("google builtin creds");
let client_secret = Some(builtin.client_secret.to_string());
let result = super::hosted_proxy_client_secret(&client_secret, Some(&builtin), true);
assert_eq!(result, None);
}
#[test]
fn test_hosted_proxy_client_secret_preserves_explicit_secret() {
let builtin = builtin_credentials("google_oauth_token").expect("google builtin creds");
let client_secret = Some("hosted-server-secret".to_string());
let result = super::hosted_proxy_client_secret(&client_secret, Some(&builtin), true);
assert_eq!(result, client_secret);
}
#[tokio::test]
async fn test_refresh_token_via_proxy_sends_auth_and_form() {
let server = MockProxyServer::start().await;
let response = super::refresh_token_via_proxy(super::ProxyRefreshTokenRequest {
proxy_url: &server.base_url(),
gateway_token: "gateway-test-token",
token_url: "https://oauth2.googleapis.com/token",
client_id: TEST_OAUTH_CLIENT_ID,
client_secret: Some(TEST_OAUTH_CLIENT_SECRET),
refresh_token: "refresh-token-123",
provider: Some("google"),
})
.await
.expect("proxy refresh succeeds");
assert_eq!(response.access_token, "proxy-access-token");
assert_eq!(
response.refresh_token.as_deref(),
Some("proxy-refresh-token")
);
assert_eq!(response.expires_in, Some(7200));
let requests = server.requests().await;
assert_eq!(requests.len(), 1);
assert_eq!(
requests[0].authorization.as_deref(),
Some("Bearer gateway-test-token")
);
assert_eq!(
requests[0].form.get("token_url").map(String::as_str),
Some("https://oauth2.googleapis.com/token")
);
assert_eq!(
requests[0].form.get("client_id").map(String::as_str),
Some(TEST_OAUTH_CLIENT_ID)
);
assert_eq!(
requests[0].form.get("client_secret").map(String::as_str),
Some(TEST_OAUTH_CLIENT_SECRET)
);
assert_eq!(
requests[0].form.get("refresh_token").map(String::as_str),
Some("refresh-token-123")
);
assert_eq!(
requests[0].form.get("provider").map(String::as_str),
Some("google")
);
server.shutdown().await;
}
#[tokio::test]
async fn test_exchange_via_proxy_does_not_follow_redirects() {
let server = MockProxyServer::start().await;
let error = match super::exchange_via_proxy(super::ProxyTokenExchangeRequest {
proxy_url: &server.redirecting_base_url(),
gateway_token: "gateway-test-token",
code: "auth-code-123",
redirect_uri: "http://localhost:3000/oauth/callback",
token_url: "https://oauth2.googleapis.com/token",
client_id: TEST_OAUTH_CLIENT_ID,
client_secret: Some(TEST_OAUTH_CLIENT_SECRET),
access_token_field: "access_token",
code_verifier: Some("code-verifier-123"),
extra_token_params: &HashMap::new(),
})
.await
{
Ok(_) => panic!("redirected proxy exchange should fail"),
Err(error) => error,
};
assert!(error.to_string().contains("307"));
assert!(server.requests().await.is_empty());
server.shutdown().await;
}
#[tokio::test]
async fn test_refresh_token_via_proxy_does_not_follow_redirects() {
let server = MockProxyServer::start().await;
let error = match super::refresh_token_via_proxy(super::ProxyRefreshTokenRequest {
proxy_url: &server.redirecting_base_url(),
gateway_token: "gateway-test-token",
token_url: "https://oauth2.googleapis.com/token",
client_id: TEST_OAUTH_CLIENT_ID,
client_secret: Some(TEST_OAUTH_CLIENT_SECRET),
refresh_token: "refresh-token-123",
provider: Some("google"),
})
.await
{
Ok(_) => panic!("redirected proxy refresh should fail"),
Err(error) => error,
};
assert!(error.to_string().contains("307"));
assert!(server.requests().await.is_empty());
server.shutdown().await;
}
#[test]
fn test_is_loopback_host() {
+270 -16
View File
@@ -2,6 +2,7 @@
//!
//! Commands for installing, listing, removing, and authenticating WASM tools.
use std::collections::{HashMap, HashSet};
use std::io::Write;
use std::path::{Path, PathBuf};
use std::sync::Arc;
@@ -79,6 +80,10 @@ pub enum ToolCommand {
/// Directory to look for tool (default: ~/.ironclaw/tools/)
#[arg(short, long)]
dir: Option<PathBuf>,
/// User ID for checking credential status (default: "default")
#[arg(short, long, default_value = "default")]
user: String,
},
/// Configure authentication for a tool
@@ -124,7 +129,11 @@ pub async fn run_tool_command(cmd: ToolCommand) -> anyhow::Result<()> {
} => install_tool(path, name, capabilities, target, release, skip_build, force).await,
ToolCommand::List { dir, verbose } => list_tools(dir, verbose).await,
ToolCommand::Remove { name, dir } => remove_tool(name, dir).await,
ToolCommand::Info { name_or_path, dir } => show_tool_info(name_or_path, dir).await,
ToolCommand::Info {
name_or_path,
dir,
user,
} => show_tool_info(name_or_path, dir, user).await,
ToolCommand::Auth { name, dir, user } => auth_tool(name, dir, user).await,
ToolCommand::Setup { name, dir, user } => setup_tool(name, dir, user).await,
}
@@ -388,7 +397,11 @@ async fn remove_tool(name: String, dir: Option<PathBuf>) -> anyhow::Result<()> {
}
/// Show information about a tool.
async fn show_tool_info(name_or_path: String, dir: Option<PathBuf>) -> anyhow::Result<()> {
async fn show_tool_info(
name_or_path: String,
dir: Option<PathBuf>,
user_id: String,
) -> anyhow::Result<()> {
let wasm_path = if name_or_path.ends_with(".wasm") {
PathBuf::from(&name_or_path)
} else {
@@ -423,7 +436,37 @@ async fn show_tool_info(name_or_path: String, dir: Option<PathBuf>) -> anyhow::R
println!("\nCapabilities ({}):", caps_path.display());
let content = fs::read_to_string(&caps_path).await?;
match CapabilitiesFile::from_json(&content) {
Ok(caps) => print_capabilities_detail(&caps),
Ok(caps) => {
// Lazily init secrets store only when auth secrets need checking.
let has_auth = caps.auth.is_some()
|| caps
.setup
.as_ref()
.is_some_and(|s| !s.required_secrets.is_empty())
|| caps
.http
.as_ref()
.is_some_and(|h| !h.credentials.is_empty());
let secrets_store = if has_auth {
match init_secrets_store().await {
Ok(store) => Some(store),
Err(e) => {
eprintln!(" Warning: could not init secrets store: {}", e);
None
}
}
} else {
None
};
print_capabilities_detail(
&caps,
secrets_store
.as_ref()
.map(|s| s.as_ref() as &(dyn SecretsStore + Send + Sync)),
&user_id,
)
.await;
}
Err(e) => println!(" Error parsing: {}", e),
}
} else {
@@ -476,8 +519,89 @@ fn print_capabilities_summary(caps: &CapabilitiesFile) {
}
}
/// Per-secret info collected from all auth-related capability sections.
struct AuthSecretInfo {
secret_name: String,
/// Human-readable label (from auth.display_name or setup prompt).
description: Option<String>,
/// Injection location (from http.credentials).
location: Option<String>,
}
/// Collected auth secrets and the set of secret names they cover.
struct CollectedAuthSecrets {
secrets: Vec<AuthSecretInfo>,
/// Secret names present in `secrets`, for filtering the Secrets capability section.
seen_names: HashSet<String>,
}
/// Collect and deduplicate auth secrets from all auth-related capability sections.
///
/// Priority for the description label: auth.display_name > setup.required_secrets.prompt.
/// Injection location is merged from http.credentials.
fn collect_auth_secrets(caps: &CapabilitiesFile) -> CollectedAuthSecrets {
let mut secrets: Vec<AuthSecretInfo> = Vec::new();
let mut seen: HashMap<String, usize> = HashMap::new();
// auth.display_name is the best label — seed first.
if let Some(ref auth) = caps.auth {
let index = secrets.len();
seen.insert(auth.secret_name.clone(), index);
secrets.push(AuthSecretInfo {
secret_name: auth.secret_name.clone(),
description: auth.display_name.clone(),
location: None,
});
}
// setup.required_secrets.prompt is second-best label.
if let Some(ref setup) = caps.setup {
for secret in &setup.required_secrets {
if !seen.contains_key(&secret.name) {
let index = secrets.len();
seen.insert(secret.name.clone(), index);
secrets.push(AuthSecretInfo {
secret_name: secret.name.clone(),
description: Some(secret.prompt.clone()),
location: None,
});
}
}
}
// Merge injection location from http.credentials.
if let Some(ref http) = caps.http {
for cred in http.credentials.values() {
let loc = format!("{:?}", cred.location);
if let Some(&index) = seen.get(&cred.secret_name) {
secrets[index].location = Some(loc);
} else {
let index = secrets.len();
seen.insert(cred.secret_name.clone(), index);
secrets.push(AuthSecretInfo {
secret_name: cred.secret_name.clone(),
description: None,
location: Some(loc),
});
}
}
}
let seen_names = seen.into_keys().collect();
CollectedAuthSecrets {
secrets,
seen_names,
}
}
/// Print detailed capabilities.
fn print_capabilities_detail(caps: &CapabilitiesFile) {
async fn print_capabilities_detail(
caps: &CapabilitiesFile,
secrets_store: Option<&(dyn SecretsStore + Send + Sync)>,
user_id: &str,
) {
let mut collected = collect_auth_secrets(caps);
if let Some(ref http) = caps.http {
println!(" HTTP:");
for endpoint in &http.allowlist {
@@ -490,13 +614,6 @@ fn print_capabilities_detail(caps: &CapabilitiesFile) {
println!(" {} {} {}", methods, endpoint.host, path);
}
if !http.credentials.is_empty() {
println!(" Credentials:");
for (key, cred) in &http.credentials {
println!(" {}: {} -> {:?}", key, cred.secret_name, cred.location);
}
}
if let Some(ref rate) = http.rate_limit {
println!(
" Rate limit: {}/min, {}/hour",
@@ -505,12 +622,24 @@ fn print_capabilities_detail(caps: &CapabilitiesFile) {
}
}
// Filter secrets already covered by the auth section (always rendered when non-empty).
if let Some(ref secrets) = caps.secrets
&& !secrets.allowed_names.is_empty()
{
println!(" Secrets (existence check only):");
for name in &secrets.allowed_names {
println!(" {}", name);
let extra: Vec<_> = if collected.secrets.is_empty() {
secrets.allowed_names.iter().collect()
} else {
secrets
.allowed_names
.iter()
.filter(|name| !collected.seen_names.contains(name.as_str()))
.collect()
};
if !extra.is_empty() {
println!(" Secrets (existence check only):");
for name in extra {
println!(" {}", name);
}
}
}
@@ -531,6 +660,38 @@ fn print_capabilities_detail(caps: &CapabilitiesFile) {
println!(" {}", prefix);
}
}
// Consolidated auth status — sorted by secret name for deterministic output.
if !collected.secrets.is_empty() {
collected
.secrets
.sort_by(|a, b| a.secret_name.cmp(&b.secret_name));
println!(" Auth:");
for info in &collected.secrets {
let (icon, label) = match secrets_store {
Some(store) => match store.exists(user_id, &info.secret_name).await {
Ok(true) => ("\u{2713}", "configured"),
Ok(false) => ("\u{2717}", "missing"),
Err(e) => {
eprintln!(
" Warning: failed to check secret `{}`: {}",
info.secret_name, e
);
("?", "unknown")
}
},
None => ("?", "unknown"),
};
let mut parts = info.secret_name.clone();
if let Some(ref desc) = info.description {
parts = format!("{} ({})", parts, desc);
}
if let Some(ref loc) = info.location {
parts = format!("{} -> {}", parts, loc);
}
println!(" {} {} {}", parts, icon, label);
}
}
}
/// Validate a tool name to prevent path traversal.
@@ -677,8 +838,7 @@ async fn combine_provider_scopes(
secret_name: &str,
base_oauth: &crate::tools::wasm::OAuthConfigSchema,
) -> crate::tools::wasm::OAuthConfigSchema {
let mut all_scopes: std::collections::HashSet<String> =
base_oauth.scopes.iter().cloned().collect();
let mut all_scopes: HashSet<String> = base_oauth.scopes.iter().cloned().collect();
if let Ok(mut entries) = tokio::fs::read_dir(tools_dir).await {
while let Ok(Some(entry)) = entries.next_entry().await {
@@ -1127,6 +1287,8 @@ async fn setup_tool(name: String, dir: Option<PathBuf>, user_id: String) -> anyh
#[cfg(test)]
mod tests {
use super::*;
use crate::secrets::{CreateSecretParams, SecretsStore};
use crate::testing::credentials::test_secrets_store;
#[test]
fn test_format_size() {
@@ -1143,4 +1305,96 @@ mod tests {
assert!(dir.to_string_lossy().contains(".ironclaw"));
assert!(dir.to_string_lossy().contains("tools"));
}
/// Verify that auth secrets are deduplicated across auth, setup, and http.credentials,
/// and that credential status is checked against the secrets store.
#[tokio::test]
async fn test_auth_secret_dedup_and_status() {
let caps = CapabilitiesFile::from_json(
r#"{
"auth": {
"secret_name": "gh_token",
"display_name": "GitHub"
},
"setup": {
"required_secrets": [
{ "name": "gh_token", "prompt": "GitHub PAT" },
{ "name": "extra_key", "prompt": "Extra API Key" }
]
},
"http": {
"allowlist": [{ "host": "api.github.com" }],
"credentials": {
"github": {
"secret_name": "gh_token",
"location": { "type": "bearer" },
"host_patterns": ["api.github.com"]
}
}
},
"secrets": {
"allowed_names": ["gh_token", "gh_*"]
}
}"#,
)
.unwrap();
let collected = collect_auth_secrets(&caps);
// gh_token should appear once (from auth), with location merged from credentials.
// extra_key should appear once (from setup).
assert_eq!(collected.secrets.len(), 2);
let gh = collected
.secrets
.iter()
.find(|s| s.secret_name == "gh_token")
.unwrap();
assert_eq!(gh.description.as_deref(), Some("GitHub"));
assert!(
gh.location.is_some(),
"location should be merged from http.credentials"
);
let extra = collected
.secrets
.iter()
.find(|s| s.secret_name == "extra_key")
.unwrap();
assert_eq!(extra.description.as_deref(), Some("Extra API Key"));
assert!(extra.location.is_none());
// Secrets section should filter gh_token (in seen_names) but keep gh_* (wildcard).
let secrets = caps.secrets.as_ref().unwrap();
let extra_secrets: Vec<_> = secrets
.allowed_names
.iter()
.filter(|name| !collected.seen_names.contains(name.as_str()))
.collect();
assert_eq!(extra_secrets, vec!["gh_*"]);
// Verify store check: missing secret -> exists returns false.
let store = test_secrets_store();
assert!(!store.exists("default", "gh_token").await.unwrap());
// Store gh_token and verify it's found.
store
.create(
"default",
CreateSecretParams::new("gh_token", "ghp_test123"),
)
.await
.unwrap();
assert!(store.exists("default", "gh_token").await.unwrap());
// extra_key still missing.
assert!(!store.exists("default", "extra_key").await.unwrap());
}
/// No auth sections → collect_auth_secrets returns empty.
#[test]
fn test_collect_auth_secrets_empty_caps() {
let caps = CapabilitiesFile::default();
let collected = collect_auth_secrets(&caps);
assert!(collected.secrets.is_empty());
assert!(collected.seen_names.is_empty());
}
}
+5 -7
View File
@@ -341,13 +341,11 @@ impl Config {
let tunnel = TunnelConfig::resolve(settings)?;
let channels = ChannelsConfig::resolve(settings, &owner_id)?;
// Resolve workspace config using the gateway user_id for default layers.
let workspace_user_id = channels
.gateway
.as_ref()
.map(|gw| gw.user_id.as_str())
.unwrap_or("default");
let workspace = WorkspaceConfig::resolve(workspace_user_id)?;
// Resolve the startup workspace against the durable owner scope. The
// gateway may expose a distinct sender identity, but the base runtime
// workspace stays owner-scoped and per-user gateway workspaces are
// handled separately by WorkspacePool.
let workspace = WorkspaceConfig::resolve(&owner_id)?;
Ok(Self {
owner_id: owner_id.clone(),
+64 -79
View File
@@ -53,22 +53,6 @@ struct HostedOAuthFlowStart {
flow: crate::cli::oauth_defaults::PendingOAuthFlow,
}
fn hosted_proxy_client_secret(
client_secret: &Option<String>,
builtin: Option<&crate::cli::oauth_defaults::OAuthCredentials>,
exchange_proxy_configured: bool,
) -> Option<String> {
if !exchange_proxy_configured {
return client_secret.clone();
}
let builtin_secret = builtin.map(|credentials| credentials.client_secret);
match (client_secret, builtin_secret) {
(Some(resolved), Some(baked_in)) if resolved == baked_in => None,
_ => client_secret.clone(),
}
}
fn normalize_oauth_callback_path(path: &str) -> String {
let trimmed_path = path.trim_end_matches('/');
if trimmed_path.is_empty() {
@@ -891,24 +875,27 @@ impl ExtensionManager {
*self.relay_channel_manager.write().await = Some(channel_manager);
}
/// Check if a channel name corresponds to a relay extension (has stored stream token
/// Check if a channel name corresponds to a relay extension (has stored team_id
/// or is tracked in the installed relay extensions set).
pub async fn is_relay_channel(&self, name: &str, user_id: &str) -> bool {
// Check in-memory installed set first (supports no-store mode)
if self.installed_relay_extensions.read().await.contains(name) {
return true;
}
// Then check for stored stream token
self.secrets
.exists(user_id, &format!("relay:{}:stream_token", name))
.await
.unwrap_or(false)
// Check for stored team_id (persisted across restarts by the OAuth callback)
if let Some(ref store) = self.store {
let key = format!("relay:{}:team_id", name);
if let Ok(Some(v)) = store.get_setting(user_id, &key).await {
return v.as_str().is_some_and(|s| !s.is_empty());
}
}
false
}
/// Restore persisted relay channels after startup.
///
/// Loads the persisted active channel list, filters to relay types (those with
/// a stored stream token), and activates each via `activate_stored_relay()`.
/// a stored team_id setting), and activates each via `activate_stored_relay()`.
/// Skips channels that are already active.
///
/// Call this only after `set_relay_channel_manager()` or `set_channel_runtime()`.
@@ -1131,7 +1118,7 @@ impl ExtensionManager {
/// Broadcast an extension status change to the web UI via SSE.
async fn broadcast_extension_status(&self, name: &str, status: &str, message: Option<&str>) {
if let Some(ref sse) = *self.sse_manager.read().await {
sse.broadcast(crate::channels::web::types::SseEvent::ExtensionStatus {
sse.broadcast(ironclaw_common::AppEvent::ExtensionStatus {
extension_name: name.to_string(),
status: status.to_string(),
message: message.map(|m| m.to_string()),
@@ -1428,9 +1415,11 @@ impl ExtensionManager {
if kind_filter.is_none() || kind_filter == Some(ExtensionKind::ChannelRelay) {
let installed = self.installed_relay_extensions.read().await;
let active_names = self.active_channel_names.read().await;
let errors = self.activation_errors.read().await;
for name in installed.iter() {
let active = active_names.contains(name);
let has_token = self.is_relay_channel(name, user_id).await;
let authenticated = self.is_relay_channel(name, user_id).await;
let activation_error = errors.get(name).cloned();
let registry_entry = self
.registry
.get_with_kind(name, Some(ExtensionKind::ChannelRelay))
@@ -1443,13 +1432,13 @@ impl ExtensionManager {
display_name,
description,
url: None,
authenticated: has_token,
authenticated,
active,
tools: Vec::new(),
needs_setup: false,
has_auth: true,
installed: true,
activation_error: None,
activation_error,
version: None,
});
}
@@ -1626,7 +1615,22 @@ impl ExtensionManager {
self.persist_active_channels(user_id).await;
self.activation_errors.write().await.remove(name);
// Remove stored stream token
// Remove stored team_id setting and clean up secrets
if let Some(ref store) = self.store
&& let Err(e) = store
.delete_setting(user_id, &format!("relay:{}:team_id", name))
.await
{
tracing::warn!(error = %e, name, "Failed to delete relay team_id setting on removal");
}
if let Err(e) = self
.secrets
.delete(user_id, &format!("relay:{}:oauth_state", name))
.await
{
tracing::warn!(error = %e, name, "Failed to delete relay oauth_state secret on removal");
}
// Clean up legacy stream_token secret from pre-webhook installs
let _ = self
.secrets
.delete(user_id, &format!("relay:{}:stream_token", name))
@@ -3179,7 +3183,7 @@ impl ExtensionManager {
// apps. Sending the desktop secret would cause a client_id/secret
// mismatch because the container's GOOGLE_OAUTH_CLIENT_ID is the web
// app, not the desktop app.
let proxy_client_secret = hosted_proxy_client_secret(
let proxy_client_secret = oauth_defaults::hosted_proxy_client_secret(
&client_secret,
builtin.as_ref(),
oauth_defaults::exchange_proxy_url().is_some(),
@@ -3284,7 +3288,7 @@ impl ExtensionManager {
}
.await;
// Broadcast SSE event
// Broadcast auth result event
let (success, message) = match result {
Ok(()) => (true, format!("{} authenticated successfully", display_name)),
Err(ref e) => (
@@ -3310,7 +3314,7 @@ impl ExtensionManager {
}
if let Some(ref sse) = sse_manager {
sse.broadcast(crate::channels::web::types::SseEvent::AuthCompleted {
sse.broadcast(ironclaw_common::AppEvent::AuthCompleted {
extension_name: ext_name,
success,
message,
@@ -4181,13 +4185,13 @@ impl ExtensionManager {
///
/// For Slack: initiates OAuth flow (redirect-based).
/// For Telegram: accepts a bot token, registers it with channel-relay,
/// and stores the returned stream token.
/// and stores the team_id setting.
async fn auth_channel_relay(
&self,
name: &str,
user_id: &str,
) -> Result<AuthResult, ExtensionError> {
// Check if already authenticated (stream token exists)
// Check if already authenticated (team_id setting exists)
if self.is_relay_channel(name, user_id).await {
return Ok(AuthResult::authenticated(name, ExtensionKind::ChannelRelay));
}
@@ -4233,19 +4237,9 @@ impl ExtensionManager {
name: &str,
user_id: &str,
) -> Result<ActivateResult, ExtensionError> {
let token_key = format!("relay:{}:stream_token", name);
let team_id_key = format!("relay:{}:team_id", name);
// Check if we have a stream token
// Verify auth: stream token must exist (even though we don't use it in this constructor path)
let _stream_token = match self.secrets.get_decrypted(user_id, &token_key).await {
Ok(secret) => secret.expose().to_string(),
Err(_) => {
return Err(ExtensionError::AuthRequired);
}
};
// Get team_id from settings
// Get team_id from settings (stored by the OAuth callback)
let team_id = if let Some(ref store) = self.store {
store
.get_setting(user_id, &team_id_key)
@@ -4258,6 +4252,10 @@ impl ExtensionManager {
String::new()
};
if team_id.is_empty() {
return Err(ExtensionError::AuthRequired);
}
// Use relay config captured at startup
let relay_config = self.relay_config()?;
@@ -4367,11 +4365,11 @@ impl ExtensionManager {
return Ok(ExtensionKind::WasmChannel);
}
// Check channel-relay extensions (installed in memory or has stored token)
// Check channel-relay extensions (installed in memory or has stored team_id)
if self.installed_relay_extensions.read().await.contains(name) {
return Ok(ExtensionKind::ChannelRelay);
}
// Also check if there's a stored stream token (persisted across restarts)
// Also check if there's a stored team_id setting (persisted across restarts)
if self.is_relay_channel(name, user_id).await {
return Ok(ExtensionKind::ChannelRelay);
}
@@ -4999,11 +4997,7 @@ impl ExtensionManager {
names.insert(server.token_secret_name());
(names, Vec::new())
}
ExtensionKind::ChannelRelay => {
let mut names = std::collections::HashSet::new();
names.insert(format!("relay:{}:stream_token", name));
(names, Vec::new())
}
ExtensionKind::ChannelRelay => (std::collections::HashSet::new(), Vec::new()),
};
let allowed_fields: std::collections::HashSet<String> =
@@ -5434,7 +5428,9 @@ impl ExtensionManager {
.map_err(|e| ExtensionError::NotInstalled(e.to_string()))?;
server.token_secret_name()
}
ExtensionKind::ChannelRelay => format!("relay:{}:stream_token", name),
ExtensionKind::ChannelRelay => {
return Err(ExtensionError::AuthRequired);
}
};
let mut secrets = std::collections::HashMap::new();
@@ -5702,7 +5698,7 @@ mod tests {
use crate::extensions::manager::{
ChannelRuntimeState, FallbackDecision, TelegramBindingData, TelegramBindingResult,
TelegramOwnerBindingState, build_wasm_channel_runtime_config_updates,
combine_install_errors, fallback_decision, hosted_proxy_client_secret, infer_kind_from_url,
combine_install_errors, fallback_decision, infer_kind_from_url,
normalize_hosted_callback_url, send_telegram_text_message,
telegram_message_matches_verification_code,
};
@@ -7043,7 +7039,7 @@ mod tests {
let dir = tempfile::tempdir().expect("temp dir");
let mgr = make_test_manager(None, dir.path().to_path_buf());
// No token stored → not a relay channel
// No store configured, no team_id → not a relay channel
assert!(!mgr.is_relay_channel("slack-relay", "test").await);
}
@@ -7862,19 +7858,13 @@ mod tests {
.await
.insert("test-relay".to_string());
// configure() should dispatch to activate_channel_relay(), not
// activate_wasm_channel(). Both will fail (no runtime configured),
// but the error should be about relay config, not WASM channels.
let mut secrets = std::collections::HashMap::new();
secrets.insert(
"relay:test-relay:stream_token".to_string(),
"tok".to_string(),
);
// configure() with empty secrets should dispatch to
// activate_channel_relay(), not activate_wasm_channel(). Relay auth
// is OAuth-only so there are no manual secrets to pass.
let result = mgr
.configure(
"test-relay",
&secrets,
&std::collections::HashMap::new(),
&std::collections::HashMap::new(),
"test",
)
@@ -7886,7 +7876,6 @@ mod tests {
);
let result = result.unwrap();
// Activation will fail (no relay config), but secrets should still be stored
assert!(
!result.activated,
"activation should fail without relay config"
@@ -7896,15 +7885,6 @@ mod tests {
"error should not mention WASM — got: {}",
result.message
);
// Verify the secret was stored
assert!(
mgr.secrets
.exists("test", "relay:test-relay:stream_token")
.await
.unwrap_or(false),
"configure should have stored the relay stream token"
);
}
#[test]
fn test_validation_failed_is_distinct_error_variant() {
@@ -7970,7 +7950,8 @@ mod tests {
let builtin_ref = builtin.as_ref();
let secret = Some(builtin_ref.unwrap().client_secret.to_string());
let result = hosted_proxy_client_secret(&secret, builtin_ref, true);
let result =
crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin_ref, true);
assert_eq!(
result, None,
"built-in desktop secret must be suppressed when the exchange proxy is configured"
@@ -7982,7 +7963,8 @@ mod tests {
let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token");
let secret = Some("user-entered-custom-secret".to_string());
let result = hosted_proxy_client_secret(&secret, builtin.as_ref(), true);
let result =
crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin.as_ref(), true);
assert_eq!(
result,
Some("user-entered-custom-secret".to_string()),
@@ -7996,7 +7978,8 @@ mod tests {
let builtin_ref = builtin.as_ref();
let secret = Some(builtin_ref.unwrap().client_secret.to_string());
let result = hosted_proxy_client_secret(&secret, builtin_ref, false);
let result =
crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin_ref, false);
assert_eq!(
result, secret,
"built-in secret must be kept when the callback will exchange directly"
@@ -8007,7 +7990,8 @@ mod tests {
fn test_proxy_client_secret_none_stays_none() {
let builtin = crate::cli::oauth_defaults::builtin_credentials("google_oauth_token");
let result = hosted_proxy_client_secret(&None, builtin.as_ref(), true);
let result =
crate::cli::oauth_defaults::hosted_proxy_client_secret(&None, builtin.as_ref(), true);
assert_eq!(
result, None,
"None secret stays None even when the exchange proxy is configured"
@@ -8021,7 +8005,8 @@ mod tests {
assert!(builtin.is_none());
let secret = Some("dcr-secret".to_string());
let result = hosted_proxy_client_secret(&secret, builtin.as_ref(), true);
let result =
crate::cli::oauth_defaults::hosted_proxy_client_secret(&secret, builtin.as_ref(), true);
assert_eq!(
result,
Some("dcr-secret".to_string()),
+68 -2
View File
@@ -463,8 +463,15 @@ impl LlmProvider for NearAiChatProvider {
let model = req.model.unwrap_or_else(|| self.active_model_name());
let mut raw_messages = req.messages;
crate::llm::provider::sanitize_tool_messages(&mut raw_messages);
let messages: Vec<ChatCompletionMessage> =
raw_messages.into_iter().map(|m| m.into()).collect();
let raw: Vec<ChatCompletionMessage> = raw_messages.into_iter().map(|m| m.into()).collect();
// NEAR AI rejects `role:"tool"` messages even on text-only completion paths.
// Apply the same flattening used by complete_with_tools().
let messages = if self.flatten_tool_messages {
flatten_tool_messages(raw)
} else {
raw
};
let request = ChatCompletionRequest {
model,
@@ -2193,6 +2200,65 @@ mod tests {
assert_eq!(deserialized.function.arguments, r#"{"city":"London"}"#);
}
// -- flatten_tool_messages in complete() path ----------------------------
#[test]
fn test_flatten_applied_on_text_only_path() {
// Verify that flatten_tool_messages converts tool-role messages to user
// messages (mirrors the complete_with_tools path).
let messages = vec![
ChatCompletionMessage {
role: "user".to_string(),
content: Some(MessageContent::Text("run it".to_string())),
tool_call_id: None,
name: None,
tool_calls: None,
},
ChatCompletionMessage {
role: "tool".to_string(),
content: Some(MessageContent::Text("ok".to_string())),
tool_call_id: Some("call_1".to_string()),
name: Some("run_cmd".to_string()),
tool_calls: None,
},
];
let flattened = flatten_tool_messages(messages);
assert_eq!(flattened.len(), 2);
assert_eq!(flattened[1].role, "user");
let text = flattened[1]
.content
.as_ref()
.and_then(|c| c.as_text())
.unwrap();
assert!(text.contains("run_cmd"), "should reference tool name");
assert!(text.contains("ok"), "should include tool result");
}
#[test]
fn test_no_flatten_when_no_tool_messages() {
// When there are no tool-role messages, flatten_tool_messages is a no-op.
let messages = vec![
ChatCompletionMessage {
role: "user".to_string(),
content: Some(MessageContent::Text("hi".to_string())),
tool_call_id: None,
name: None,
tool_calls: None,
},
ChatCompletionMessage {
role: "assistant".to_string(),
content: Some(MessageContent::Text("hello".to_string())),
tool_call_id: None,
name: None,
tool_calls: None,
},
];
let result = flatten_tool_messages(messages);
// No tool messages → unchanged roles
assert_eq!(result[0].role, "user");
assert_eq!(result[1].role, "assistant");
}
// -- api_url edge cases ---------------------------------------------------
#[test]
+1
View File
@@ -611,6 +611,7 @@ async fn async_main() -> anyhow::Result<()> {
} else {
GatewayChannel::new(gw_config.clone())
};
gw = gw.with_owner_scope(config.owner_id.clone());
gw = gw.with_llm_provider(Arc::clone(&components.llm));
if let Some(ref ws) = components.workspace {
gw = gw.with_workspace(Arc::clone(ws));
+14 -14
View File
@@ -14,7 +14,6 @@ use serde::{Deserialize, Serialize};
use tokio::sync::{Mutex, broadcast};
use uuid::Uuid;
use crate::channels::web::types::SseEvent;
use crate::db::Database;
use crate::llm::{CompletionRequest, LlmProvider, ToolCompletionRequest};
use crate::orchestrator::auth::{TokenStore, worker_auth_middleware};
@@ -25,6 +24,7 @@ use crate::worker::api::{
CompletionReport, CredentialResponse, JobDescription, ProxyCompletionRequest,
ProxyCompletionResponse, ProxyToolCompletionRequest, ProxyToolCompletionResponse, StatusUpdate,
};
use ironclaw_common::AppEvent;
/// A follow-up prompt queued for a Claude Code bridge.
#[derive(Debug, Clone, Serialize, Deserialize)]
@@ -41,7 +41,7 @@ pub struct OrchestratorState {
pub token_store: TokenStore,
/// Broadcast channel for job events (consumed by the web gateway SSE).
/// Tuple: (job_id, user_id, event).
pub job_event_tx: Option<broadcast::Sender<(Uuid, String, SseEvent)>>,
pub job_event_tx: Option<broadcast::Sender<(Uuid, String, AppEvent)>>,
/// Buffered follow-up prompts for sandbox jobs, keyed by job_id.
pub prompt_queue: Arc<Mutex<HashMap<Uuid, VecDeque<PendingPrompt>>>>,
/// Database handle for persisting job events.
@@ -277,10 +277,10 @@ async fn job_event_handler(
});
}
// Convert to SSE event and broadcast
// Convert to app event and broadcast
let job_id_str = job_id.to_string();
let sse_event = match payload.event_type.as_str() {
"message" => SseEvent::JobMessage {
let app_event = match payload.event_type.as_str() {
"message" => AppEvent::JobMessage {
job_id: job_id_str,
role: payload
.data
@@ -295,7 +295,7 @@ async fn job_event_handler(
.unwrap_or("")
.to_string(),
},
"tool_use" => SseEvent::JobToolUse {
"tool_use" => AppEvent::JobToolUse {
job_id: job_id_str,
tool_name: payload
.data
@@ -309,7 +309,7 @@ async fn job_event_handler(
.cloned()
.unwrap_or(serde_json::Value::Null),
},
"tool_result" => SseEvent::JobToolResult {
"tool_result" => AppEvent::JobToolResult {
job_id: job_id_str,
tool_name: payload
.data
@@ -324,7 +324,7 @@ async fn job_event_handler(
.unwrap_or("")
.to_string(),
},
"result" => SseEvent::JobResult {
"result" => AppEvent::JobResult {
job_id: job_id_str,
status: payload
.data
@@ -344,7 +344,7 @@ async fn job_event_handler(
// gain context/memory tracking capabilities.
fallback_deliverable: payload.data.get("fallback_deliverable").cloned(),
},
_ => SseEvent::JobStatus {
_ => AppEvent::JobStatus {
job_id: job_id_str,
message: payload
.data
@@ -390,9 +390,9 @@ async fn job_event_handler(
};
if user_id.is_empty() {
let _ = tx.send((job_id, String::new(), sse_event));
let _ = tx.send((job_id, String::new(), app_event));
} else {
let _ = tx.send((job_id, user_id, sse_event));
let _ = tx.send((job_id, user_id, app_event));
}
}
@@ -817,7 +817,7 @@ mod tests {
// No store configured, so user_id falls back to empty string.
assert_eq!(recv_uid, "");
match event {
SseEvent::JobMessage {
AppEvent::JobMessage {
job_id: jid,
role,
content,
@@ -872,7 +872,7 @@ mod tests {
let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap();
match event {
SseEvent::JobToolUse { tool_name, .. } => {
AppEvent::JobToolUse { tool_name, .. } => {
assert_eq!(tool_name, "shell");
}
other => panic!("Expected JobToolUse, got {:?}", other),
@@ -918,7 +918,7 @@ mod tests {
let (_recv_id, _recv_uid, event) = rx.recv().await.unwrap();
// Unknown event types fall through to JobStatus
assert!(matches!(event, SseEvent::JobStatus { .. }));
assert!(matches!(event, AppEvent::JobStatus { .. }));
}
// -- Status update test --
+2 -2
View File
@@ -46,10 +46,10 @@ use std::sync::Arc;
use tokio::sync::{Mutex, broadcast};
use uuid::Uuid;
use crate::channels::web::types::SseEvent;
use crate::db::Database;
use crate::llm::LlmProvider;
use crate::secrets::SecretsStore;
use ironclaw_common::AppEvent;
/// Resolve the orchestrator port from the `ORCHESTRATOR_PORT` environment
/// variable, falling back to 50051.
@@ -63,7 +63,7 @@ fn resolve_orchestrator_port() -> u16 {
/// Result of orchestrator setup, containing all handles needed by the agent.
pub struct OrchestratorSetup {
pub container_job_manager: Option<Arc<ContainerJobManager>>,
pub job_event_tx: Option<broadcast::Sender<(Uuid, String, SseEvent)>>,
pub job_event_tx: Option<broadcast::Sender<(Uuid, String, AppEvent)>>,
pub prompt_queue: Arc<Mutex<HashMap<Uuid, VecDeque<api::PendingPrompt>>>>,
pub docker_status: crate::sandbox::DockerStatus,
}
+3 -3
View File
@@ -17,7 +17,6 @@ use uuid::Uuid;
use crate::bootstrap::ironclaw_base_dir;
use crate::channels::IncomingMessage;
use crate::channels::web::types::SseEvent;
use crate::context::{ContextManager, JobContext, JobState};
use crate::db::Database;
use crate::history::SandboxJobRecord;
@@ -25,6 +24,7 @@ use crate::orchestrator::auth::CredentialGrant;
use crate::orchestrator::job_manager::{ContainerJobManager, JobMode};
use crate::secrets::SecretsStore;
use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput, require_str};
use ironclaw_common::AppEvent;
/// Lazy scheduler reference, filled after Agent::new creates the Scheduler.
///
@@ -85,7 +85,7 @@ pub struct CreateJobTool {
job_manager: Option<Arc<ContainerJobManager>>,
store: Option<Arc<dyn Database>>,
/// Broadcast sender for job events (used to subscribe a monitor).
event_tx: Option<tokio::sync::broadcast::Sender<(Uuid, String, SseEvent)>>,
event_tx: Option<tokio::sync::broadcast::Sender<(Uuid, String, AppEvent)>>,
/// Injection channel for pushing messages into the agent loop.
inject_tx: Option<tokio::sync::mpsc::Sender<IncomingMessage>>,
/// Encrypted secrets store for validating credential grants.
@@ -120,7 +120,7 @@ impl CreateJobTool {
/// monitor that forwards Claude Code output to the main agent loop.
pub fn with_monitor_deps(
mut self,
event_tx: tokio::sync::broadcast::Sender<(Uuid, String, SseEvent)>,
event_tx: tokio::sync::broadcast::Sender<(Uuid, String, AppEvent)>,
inject_tx: tokio::sync::mpsc::Sender<IncomingMessage>,
) -> Self {
self.event_tx = Some(event_tx);
+1 -5
View File
@@ -383,11 +383,7 @@ impl ToolRegistry {
job_manager: Option<Arc<ContainerJobManager>>,
store: Option<Arc<dyn Database>>,
job_event_tx: Option<
tokio::sync::broadcast::Sender<(
uuid::Uuid,
String,
crate::channels::web::types::SseEvent,
)>,
tokio::sync::broadcast::Sender<(uuid::Uuid, String, ironclaw_common::AppEvent)>,
>,
inject_tx: Option<tokio::sync::mpsc::Sender<crate::channels::IncomingMessage>>,
prompt_queue: Option<PromptQueue>,
+141
View File
@@ -418,6 +418,7 @@ fn resolve_oauth_refresh_config(cap_file: &CapabilitiesFile) -> Option<OAuthRefr
let oauth = auth.oauth.as_ref()?;
let builtin = crate::cli::oauth_defaults::builtin_credentials(&auth.secret_name);
let exchange_proxy_url = crate::cli::oauth_defaults::exchange_proxy_url();
let client_id = oauth
.client_id
@@ -440,11 +441,21 @@ fn resolve_oauth_refresh_config(cap_file: &CapabilitiesFile) -> Option<OAuthRefr
.and_then(|env| std::env::var(env).ok())
})
.or_else(|| builtin.as_ref().map(|c| c.client_secret.to_string()));
let client_secret = crate::cli::oauth_defaults::hosted_proxy_client_secret(
&client_secret,
builtin.as_ref(),
exchange_proxy_url.is_some(),
);
let gateway_token = crate::config::helpers::env_or_override("GATEWAY_AUTH_TOKEN")
.map(|token| token.trim().to_string())
.filter(|token| !token.is_empty());
Some(OAuthRefreshConfig {
token_url: oauth.token_url.clone(),
client_id,
client_secret,
exchange_proxy_url,
gateway_token,
secret_name: auth.secret_name.clone(),
provider: auth.provider.clone(),
})
@@ -711,9 +722,44 @@ mod tests {
use tempfile::TempDir;
use crate::config::helpers::lock_env;
use crate::testing::credentials::{TEST_OAUTH_CLIENT_ID, TEST_OAUTH_CLIENT_SECRET};
use crate::tools::wasm::loader::{WasmLoadError, check_wit_version_compat, discover_tools};
/// Restores a test-scoped env var override on drop.
struct EnvVarGuard {
key: String,
previous: Option<String>,
}
impl Drop for EnvVarGuard {
fn drop(&mut self) {
// SAFETY: Tests use lock_env() to serialize environment access.
unsafe {
if let Some(ref value) = self.previous {
std::env::set_var(&self.key, value);
} else {
std::env::remove_var(&self.key);
}
}
}
}
fn set_env_var(key: &str, value: Option<&str>) -> EnvVarGuard {
let previous = std::env::var(key).ok();
// SAFETY: Tests use lock_env() to serialize environment access.
unsafe {
match value {
Some(value) => std::env::set_var(key, value),
None => std::env::remove_var(key),
}
}
EnvVarGuard {
key: key.to_string(),
previous,
}
}
#[test]
fn wit_version_compat_none_is_ok() {
// Pre-versioning extensions (no wit_version declared) should always pass
@@ -871,6 +917,8 @@ mod tests {
config.client_secret,
Some(TEST_OAUTH_CLIENT_SECRET.to_string())
);
assert_eq!(config.exchange_proxy_url, None);
assert_eq!(config.gateway_token, None);
assert_eq!(config.secret_name, "google_oauth_token");
assert_eq!(config.provider, Some("google".to_string()));
}
@@ -931,6 +979,10 @@ mod tests {
AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema,
};
let _guard = lock_env();
let _proxy_guard = set_env_var("IRONCLAW_OAUTH_EXCHANGE_URL", None);
let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", None);
// google_oauth_token should fall back to built-in credentials
let caps = CapabilitiesFile {
auth: Some(AuthCapabilitySchema {
@@ -952,6 +1004,95 @@ mod tests {
let config = config.unwrap();
assert!(!config.client_id.is_empty());
assert!(config.client_secret.is_some());
assert_eq!(config.exchange_proxy_url, None);
assert_eq!(config.gateway_token, None);
}
#[test]
fn test_resolve_oauth_refresh_config_hosted_proxy_populates_env_and_suppresses_builtin_secret()
{
use crate::tools::wasm::capabilities_schema::{
AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema,
};
let _guard = lock_env();
let _proxy_guard = set_env_var(
"IRONCLAW_OAUTH_EXCHANGE_URL",
Some("https://compose-api.example.com"),
);
let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-test-token"));
let _client_id_guard =
set_env_var("GOOGLE_OAUTH_CLIENT_ID", Some("hosted-google-client-id"));
let caps = CapabilitiesFile {
auth: Some(AuthCapabilitySchema {
secret_name: "google_oauth_token".to_string(),
provider: Some("google".to_string()),
oauth: Some(OAuthConfigSchema {
authorization_url: "https://accounts.google.com/o/oauth2/v2/auth".to_string(),
token_url: "https://oauth2.googleapis.com/token".to_string(),
client_id_env: Some("GOOGLE_OAUTH_CLIENT_ID".to_string()),
..Default::default()
}),
..Default::default()
}),
..Default::default()
};
let config = super::resolve_oauth_refresh_config(&caps).expect("hosted oauth config");
assert_eq!(config.client_id, "hosted-google-client-id");
assert_eq!(config.client_secret, None);
assert_eq!(
config.exchange_proxy_url.as_deref(),
Some("https://compose-api.example.com")
);
assert_eq!(config.gateway_token.as_deref(), Some("gateway-test-token"));
}
#[test]
fn test_resolve_oauth_refresh_config_hosted_proxy_preserves_explicit_secret() {
use crate::tools::wasm::capabilities_schema::{
AuthCapabilitySchema, CapabilitiesFile, OAuthConfigSchema,
};
let _guard = lock_env();
let _proxy_guard = set_env_var(
"IRONCLAW_OAUTH_EXCHANGE_URL",
Some("https://compose-api.example.com"),
);
let _gateway_token_guard = set_env_var("GATEWAY_AUTH_TOKEN", Some("gateway-test-token"));
let _client_id_guard =
set_env_var("GOOGLE_OAUTH_CLIENT_ID", Some("hosted-google-client-id"));
let _client_secret_guard =
set_env_var("GOOGLE_OAUTH_CLIENT_SECRET", Some("hosted-server-secret"));
let caps = CapabilitiesFile {
auth: Some(AuthCapabilitySchema {
secret_name: "google_oauth_token".to_string(),
provider: Some("google".to_string()),
oauth: Some(OAuthConfigSchema {
authorization_url: "https://accounts.google.com/o/oauth2/v2/auth".to_string(),
token_url: "https://oauth2.googleapis.com/token".to_string(),
client_id_env: Some("GOOGLE_OAUTH_CLIENT_ID".to_string()),
client_secret_env: Some("GOOGLE_OAUTH_CLIENT_SECRET".to_string()),
..Default::default()
}),
..Default::default()
}),
..Default::default()
};
let config = super::resolve_oauth_refresh_config(&caps).expect("hosted oauth config");
assert_eq!(config.client_id, "hosted-google-client-id");
assert_eq!(
config.client_secret.as_deref(),
Some("hosted-server-secret")
);
assert_eq!(
config.exchange_proxy_url.as_deref(),
Some("https://compose-api.example.com")
);
assert_eq!(config.gateway_token.as_deref(), Some("gateway-test-token"));
}
// ---------------------------------------------------------------
+443 -38
View File
@@ -19,7 +19,7 @@ use wasmtime_wasi::{ResourceTable, WasiCtx, WasiCtxBuilder, WasiView};
use crate::context::JobContext;
use crate::llm::recording::{HttpExchangeRequest, HttpExchangeResponse, HttpInterceptor};
use crate::safety::LeakDetector;
use crate::secrets::SecretsStore;
use crate::secrets::{DecryptedSecret, SecretsStore};
use crate::tools::tool::{Tool, ToolError, ToolOutput};
use crate::tools::wasm::capabilities::Capabilities;
use crate::tools::wasm::credential_injector::{
@@ -44,6 +44,7 @@ wasmtime::component::bindgen!({
});
// Alias the export interface types for convenience.
use crate::cli::oauth_defaults;
use exports::near::agent::tool as wit_tool;
/// Configuration needed to refresh an expired OAuth access token.
@@ -59,6 +60,10 @@ pub struct OAuthRefreshConfig {
pub client_id: String,
/// OAuth client_secret (optional, some providers use PKCE without a secret).
pub client_secret: Option<String>,
/// Hosted OAuth proxy base URL (e.g., "http://host.docker.internal:8080").
pub exchange_proxy_url: Option<String>,
/// Gateway auth token for authenticating with the hosted OAuth proxy.
pub gateway_token: Option<String>,
/// Secret name of the access token (e.g., "google_oauth_token").
/// The refresh token lives at `{secret_name}_refresh_token`.
pub secret_name: String,
@@ -1210,6 +1215,53 @@ async fn refresh_oauth_token(
user_id: &str,
config: &OAuthRefreshConfig,
) -> bool {
let refresh_name = format!("{}_refresh_token", config.secret_name);
if let Some(proxy_url) = config.exchange_proxy_url.as_deref() {
let Some(gateway_token) = config.gateway_token.as_deref() else {
tracing::warn!(
"OAuth refresh proxy is configured, but no gateway auth token is available"
);
return false;
};
// In hosted mode, the configured exchange proxy owns the outbound token
// refresh and validation policy for the provider token_url. Direct-mode
// HTTPS/private-IP checks remain in place for self-hosted refreshes below.
let refresh_secret = match load_oauth_refresh_secret(store, user_id, &refresh_name).await {
Some(secret) => secret,
None => return false,
};
let token_response = match oauth_defaults::refresh_token_via_proxy(
oauth_defaults::ProxyRefreshTokenRequest {
proxy_url,
gateway_token,
token_url: &config.token_url,
client_id: &config.client_id,
client_secret: config.client_secret.as_deref(),
refresh_token: refresh_secret.expose(),
provider: config.provider.as_deref(),
},
)
.await
{
Ok(response) => response,
Err(error) => {
tracing::warn!(error = %error, "OAuth token refresh via proxy failed");
return false;
}
};
return persist_refreshed_oauth_tokens(
store,
user_id,
config,
&refresh_name,
token_response,
)
.await;
}
// SSRF defense: token_url comes from the tool's capabilities file.
if !config.token_url.starts_with("https://") {
tracing::warn!(
@@ -1227,19 +1279,6 @@ async fn refresh_oauth_token(
return false;
}
let refresh_name = format!("{}_refresh_token", config.secret_name);
let refresh_secret = match store.get_decrypted(user_id, &refresh_name).await {
Ok(s) => s,
Err(e) => {
tracing::debug!(
secret_name = %refresh_name,
error = %e,
"No refresh token available, skipping token refresh"
);
return false;
}
};
let client = match reqwest::Client::builder()
.timeout(Duration::from_secs(15))
.redirect(reqwest::redirect::Policy::none())
@@ -1252,6 +1291,10 @@ async fn refresh_oauth_token(
}
};
let refresh_secret = match load_oauth_refresh_secret(store, user_id, &refresh_name).await {
Some(secret) => secret,
None => return false,
};
let mut params = vec![
("grant_type", "refresh_token".to_string()),
("refresh_token", refresh_secret.expose().to_string()),
@@ -1287,22 +1330,55 @@ async fn refresh_oauth_token(
return false;
}
};
let new_access_token = match token_data.get("access_token").and_then(|v| v.as_str()) {
Some(t) => t,
let token_response = match token_data.get("access_token").and_then(|v| v.as_str()) {
Some(access_token) => oauth_defaults::OAuthTokenResponse {
access_token: access_token.to_string(),
refresh_token: token_data
.get("refresh_token")
.and_then(|v| v.as_str())
.map(str::to_string),
expires_in: token_data.get("expires_in").and_then(|v| v.as_u64()),
},
None => {
tracing::warn!("Token refresh response missing access_token field");
return false;
}
};
// Store the new access token with expiry
persist_refreshed_oauth_tokens(store, user_id, config, &refresh_name, token_response).await
}
async fn load_oauth_refresh_secret(
store: &(dyn SecretsStore + Send + Sync),
user_id: &str,
refresh_name: &str,
) -> Option<DecryptedSecret> {
match store.get_decrypted(user_id, refresh_name).await {
Ok(secret) => Some(secret),
Err(error) => {
tracing::debug!(
secret_name = %refresh_name,
error = %error,
"No refresh token available, skipping token refresh"
);
None
}
}
}
async fn persist_refreshed_oauth_tokens(
store: &(dyn SecretsStore + Send + Sync),
user_id: &str,
config: &OAuthRefreshConfig,
refresh_name: &str,
token_response: oauth_defaults::OAuthTokenResponse,
) -> bool {
let mut access_params =
crate::secrets::CreateSecretParams::new(&config.secret_name, new_access_token);
crate::secrets::CreateSecretParams::new(&config.secret_name, &token_response.access_token);
if let Some(ref provider) = config.provider {
access_params = access_params.with_provider(provider);
}
if let Some(expires_in) = token_data.get("expires_in").and_then(|v| v.as_u64()) {
if let Some(expires_in) = token_response.expires_in {
let expires_at = chrono::Utc::now() + chrono::Duration::seconds(expires_in as i64);
access_params = access_params.with_expiry(expires_at);
}
@@ -1312,10 +1388,8 @@ async fn refresh_oauth_token(
return false;
}
// Store rotated refresh token if the provider sent a new one
if let Some(new_refresh) = token_data.get("refresh_token").and_then(|v| v.as_str()) {
let mut refresh_params =
crate::secrets::CreateSecretParams::new(&refresh_name, new_refresh);
if let Some(new_refresh) = token_response.refresh_token.as_deref() {
let mut refresh_params = crate::secrets::CreateSecretParams::new(refresh_name, new_refresh);
if let Some(ref provider) = config.provider {
refresh_params = refresh_params.with_provider(provider);
}
@@ -1664,9 +1738,18 @@ fn build_tool_usage_hint(tool_name: &str, schema: &serde_json::Value) -> String
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use std::net::SocketAddr;
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use axum::extract::{Form, State};
use axum::http::HeaderMap;
use axum::routing::post;
use axum::{Json, Router};
use serde_json::json;
use tokio::net::TcpListener;
use tokio::sync::{Mutex as AsyncMutex, oneshot};
use uuid::Uuid;
use crate::context::JobContext;
@@ -1756,6 +1839,95 @@ mod tests {
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
struct RecordedProxyRequest {
authorization: Option<String>,
form: HashMap<String, String>,
}
struct MockProxyServer {
addr: SocketAddr,
requests: Arc<AsyncMutex<Vec<RecordedProxyRequest>>>,
shutdown_tx: Option<oneshot::Sender<()>>,
server_task: Option<tokio::task::JoinHandle<()>>,
}
impl MockProxyServer {
async fn start() -> Self {
async fn refresh_handler(
State(requests): State<Arc<AsyncMutex<Vec<RecordedProxyRequest>>>>,
headers: HeaderMap,
Form(form): Form<HashMap<String, String>>,
) -> Json<serde_json::Value> {
requests.lock().await.push(RecordedProxyRequest {
authorization: headers
.get(axum::http::header::AUTHORIZATION)
.and_then(|value| value.to_str().ok())
.map(str::to_string),
form,
});
Json(json!({
"access_token": "mock-refreshed-access-token",
"refresh_token": "mock-rotated-refresh-token",
"expires_in": 3600
}))
}
let requests = Arc::new(AsyncMutex::new(Vec::new()));
let app = Router::new()
.route("/oauth/refresh", post(refresh_handler))
.with_state(Arc::clone(&requests));
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("bind mock proxy");
let addr = listener.local_addr().expect("read mock proxy addr");
let (shutdown_tx, shutdown_rx) = oneshot::channel::<()>();
let server_task = tokio::spawn(async move {
let _ = axum::serve(listener, app)
.with_graceful_shutdown(async {
let _ = shutdown_rx.await;
})
.await;
});
Self {
addr,
requests,
shutdown_tx: Some(shutdown_tx),
server_task: Some(server_task),
}
}
fn base_url(&self) -> String {
format!("http://{}", self.addr)
}
async fn requests(&self) -> Vec<RecordedProxyRequest> {
self.requests.lock().await.clone()
}
async fn shutdown(mut self) {
if let Some(tx) = self.shutdown_tx.take() {
let _ = tx.send(());
}
if let Some(task) = self.server_task.take() {
let _ = task.await;
}
}
}
impl Drop for MockProxyServer {
fn drop(&mut self) {
if let Some(tx) = self.shutdown_tx.take() {
let _ = tx.send(());
}
if let Some(task) = self.server_task.take() {
task.abort();
}
}
}
#[test]
fn test_wrapper_creation() {
// This test verifies the runtime can be created
@@ -2094,8 +2266,6 @@ mod tests {
#[tokio::test]
async fn test_resolve_host_credentials_bearer() {
use std::collections::HashMap;
use crate::secrets::{
CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore,
};
@@ -2141,8 +2311,6 @@ mod tests {
#[tokio::test]
async fn test_resolve_host_credentials_owner_scope_bearer() {
use std::collections::HashMap;
use crate::secrets::{
CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore,
};
@@ -2188,8 +2356,6 @@ mod tests {
#[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;
@@ -2239,8 +2405,6 @@ mod tests {
#[tokio::test]
async fn test_resolve_host_credentials_missing_secret() {
use std::collections::HashMap;
use crate::secrets::{CredentialLocation, CredentialMapping};
use crate::tools::wasm::capabilities::HttpCapability;
use crate::tools::wasm::wrapper::resolve_host_credentials;
@@ -2272,8 +2436,6 @@ mod tests {
#[tokio::test]
async fn test_resolve_host_credentials_skips_refresh_when_not_expired() {
use std::collections::HashMap;
use crate::secrets::{
CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore,
};
@@ -2315,6 +2477,8 @@ mod tests {
token_url: "https://oauth2.googleapis.com/token".to_string(),
client_id: TEST_OAUTH_CLIENT_ID.to_string(),
client_secret: Some(TEST_OAUTH_CLIENT_SECRET.to_string()),
exchange_proxy_url: None,
gateway_token: None,
secret_name: "google_oauth_token".to_string(),
provider: Some("google".to_string()),
};
@@ -2331,8 +2495,6 @@ mod tests {
#[tokio::test]
async fn test_resolve_host_credentials_skips_refresh_no_config() {
use std::collections::HashMap;
use crate::secrets::{
CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore,
};
@@ -2376,8 +2538,6 @@ mod tests {
#[tokio::test]
async fn test_resolve_host_credentials_skips_refresh_no_expires_at() {
use std::collections::HashMap;
use crate::secrets::{
CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore,
};
@@ -2417,6 +2577,8 @@ mod tests {
token_url: "https://oauth2.googleapis.com/token".to_string(),
client_id: TEST_OAUTH_CLIENT_ID.to_string(),
client_secret: Some(TEST_OAUTH_CLIENT_SECRET.to_string()),
exchange_proxy_url: None,
gateway_token: None,
secret_name: "google_oauth_token".to_string(),
provider: Some("google".to_string()),
};
@@ -2431,6 +2593,249 @@ mod tests {
);
}
#[tokio::test]
async fn test_resolve_host_credentials_refreshes_via_proxy_without_direct_token_url_validation()
{
use crate::secrets::{
CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore,
};
use crate::tools::wasm::capabilities::HttpCapability;
use crate::tools::wasm::wrapper::{OAuthRefreshConfig, resolve_host_credentials};
let proxy = MockProxyServer::start().await;
let store = test_secrets_store();
store
.create(
"user1",
CreateSecretParams::new("google_oauth_token", "expired-access-token")
.with_expiry(chrono::Utc::now() - chrono::Duration::hours(1)),
)
.await
.unwrap();
store
.create(
"user1",
CreateSecretParams::new("google_oauth_token_refresh_token", "stored-refresh-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 oauth_config = OAuthRefreshConfig {
token_url: "http://127.0.0.1:9/provider-token-endpoint".to_string(),
client_id: "hosted-google-client-id".to_string(),
client_secret: None,
exchange_proxy_url: Some(proxy.base_url()),
gateway_token: Some("gateway-test-token".to_string()),
secret_name: "google_oauth_token".to_string(),
provider: Some("google".to_string()),
};
let resolved =
resolve_host_credentials(&caps, Some(&store), "user1", Some(&oauth_config)).await;
assert_eq!(resolved.len(), 1);
assert_eq!(
resolved[0].headers.get("Authorization"),
Some(&"Bearer mock-refreshed-access-token".to_string())
);
let access_secret = store.get("user1", "google_oauth_token").await.unwrap();
assert!(
access_secret
.expires_at
.expect("refreshed access token expiry")
> chrono::Utc::now()
);
let access_value = store
.get_decrypted("user1", "google_oauth_token")
.await
.unwrap();
assert_eq!(access_value.expose(), "mock-refreshed-access-token");
let refresh_value = store
.get_decrypted("user1", "google_oauth_token_refresh_token")
.await
.unwrap();
assert_eq!(refresh_value.expose(), "mock-rotated-refresh-token");
let requests = proxy.requests().await;
assert_eq!(requests.len(), 1);
assert_eq!(
requests[0].authorization.as_deref(),
Some("Bearer gateway-test-token")
);
assert_eq!(
requests[0].form.get("client_id").map(String::as_str),
Some("hosted-google-client-id")
);
assert_eq!(
requests[0].form.get("token_url").map(String::as_str),
Some("http://127.0.0.1:9/provider-token-endpoint")
);
assert_eq!(
requests[0].form.get("refresh_token").map(String::as_str),
Some("stored-refresh-token")
);
assert_eq!(
requests[0].form.get("provider").map(String::as_str),
Some("google")
);
assert!(!requests[0].form.contains_key("client_secret"));
proxy.shutdown().await;
}
#[tokio::test]
async fn test_resolve_host_credentials_skips_refresh_token_lookup_without_gateway_token() {
use crate::secrets::{
CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore,
};
use crate::tools::wasm::capabilities::HttpCapability;
use crate::tools::wasm::wrapper::{OAuthRefreshConfig, resolve_host_credentials};
let store = RecordingSecretsStore::new();
store
.create(
"user1",
CreateSecretParams::new("google_oauth_token", "expired-access-token")
.with_expiry(chrono::Utc::now() - chrono::Duration::hours(1)),
)
.await
.unwrap();
store
.create(
"user1",
CreateSecretParams::new("google_oauth_token_refresh_token", "stored-refresh-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 oauth_config = OAuthRefreshConfig {
token_url: "https://oauth2.googleapis.com/token".to_string(),
client_id: "hosted-google-client-id".to_string(),
client_secret: None,
exchange_proxy_url: Some("https://compose-api.example.com".to_string()),
gateway_token: None,
secret_name: "google_oauth_token".to_string(),
provider: Some("google".to_string()),
};
let resolved =
resolve_host_credentials(&caps, Some(&store), "user1", Some(&oauth_config)).await;
assert!(resolved.is_empty());
let lookups = store.decrypted_lookups();
assert!(lookups.contains(&("user1".to_string(), "google_oauth_token".to_string())));
assert!(!lookups.contains(&(
"user1".to_string(),
"google_oauth_token_refresh_token".to_string(),
)));
}
#[tokio::test]
async fn test_resolve_host_credentials_skips_refresh_token_lookup_for_invalid_direct_token_url()
{
use crate::secrets::{
CreateSecretParams, CredentialLocation, CredentialMapping, SecretsStore,
};
use crate::tools::wasm::capabilities::HttpCapability;
use crate::tools::wasm::wrapper::{OAuthRefreshConfig, resolve_host_credentials};
let store = RecordingSecretsStore::new();
store
.create(
"user1",
CreateSecretParams::new("google_oauth_token", "expired-access-token")
.with_expiry(chrono::Utc::now() - chrono::Duration::hours(1)),
)
.await
.unwrap();
store
.create(
"user1",
CreateSecretParams::new("google_oauth_token_refresh_token", "stored-refresh-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 oauth_config = OAuthRefreshConfig {
token_url: "http://127.0.0.1:9/provider-token-endpoint".to_string(),
client_id: TEST_OAUTH_CLIENT_ID.to_string(),
client_secret: Some(TEST_OAUTH_CLIENT_SECRET.to_string()),
exchange_proxy_url: None,
gateway_token: None,
secret_name: "google_oauth_token".to_string(),
provider: Some("google".to_string()),
};
let resolved =
resolve_host_credentials(&caps, Some(&store), "user1", Some(&oauth_config)).await;
assert!(resolved.is_empty());
let lookups = store.decrypted_lookups();
assert!(lookups.contains(&("user1".to_string(), "google_oauth_token".to_string())));
assert!(!lookups.contains(&(
"user1".to_string(),
"google_oauth_token_refresh_token".to_string(),
)));
}
#[test]
fn test_is_private_ip_v4() {
use std::net::IpAddr;
+4 -4
View File
@@ -429,8 +429,8 @@ mod tests {
port: 3000,
auth_token: None,
user_id: "test".to_string(),
workspace_read_scopes: vec![],
memory_layers: vec![],
workspace_read_scopes: Vec::new(),
memory_layers: Vec::new(),
user_tokens: None,
});
c
@@ -443,8 +443,8 @@ mod tests {
port,
auth_token: None,
user_id: "test".to_string(),
workspace_read_scopes: vec![],
memory_layers: vec![],
workspace_read_scopes: Vec::new(),
memory_layers: Vec::new(),
user_tokens: None,
});
c
+51 -1
View File
@@ -1,5 +1,7 @@
//! Shared utility functions used across the codebase.
use crate::llm::{ChatMessage, Role};
/// Find the largest valid UTF-8 char boundary at or before `pos`.
///
/// Polyfill for `str::floor_char_boundary` (nightly-only). Use when
@@ -16,6 +18,17 @@ pub fn floor_char_boundary(s: &str, pos: usize) -> usize {
i
}
/// Ensure the last message in `messages` is a user-role message.
///
/// NEAR AI rejects conversations that don't end with a user message;
/// Claude 4.6 rejects assistant prefill. Call this before any LLM
/// completion request to satisfy both requirements.
pub fn ensure_ends_with_user_message(messages: &mut Vec<ChatMessage>) {
if !matches!(messages.last(), Some(m) if m.role == Role::User) {
messages.push(ChatMessage::user("Continue."));
}
}
/// Check if an LLM response explicitly signals that a job/task is complete.
///
/// Uses phrase-level matching to avoid false positives from bare words like
@@ -72,7 +85,8 @@ pub fn llm_signals_completion(response: &str) -> bool {
#[cfg(test)]
mod tests {
use crate::util::{floor_char_boundary, llm_signals_completion};
use crate::llm::ChatMessage;
use crate::util::{ensure_ends_with_user_message, floor_char_boundary, llm_signals_completion};
// ── floor_char_boundary ──
@@ -103,6 +117,42 @@ mod tests {
assert_eq!(floor_char_boundary("", 5), 0);
}
// ── ensure_ends_with_user_message ──
#[test]
fn ensure_user_message_injects_when_empty() {
let mut msgs: Vec<ChatMessage> = vec![];
ensure_ends_with_user_message(&mut msgs);
assert_eq!(msgs.len(), 1);
assert_eq!(msgs[0].role, crate::llm::Role::User);
}
#[test]
fn ensure_user_message_injects_after_assistant() {
let mut msgs = vec![ChatMessage::user("hi"), ChatMessage::assistant("hello")];
ensure_ends_with_user_message(&mut msgs);
assert_eq!(msgs.len(), 3);
assert_eq!(msgs[2].role, crate::llm::Role::User);
}
#[test]
fn ensure_user_message_injects_after_tool_result() {
let mut msgs = vec![
ChatMessage::user("run tool"),
ChatMessage::tool_result("call_1", "my_tool", "result"),
];
ensure_ends_with_user_message(&mut msgs);
assert_eq!(msgs.len(), 3);
assert_eq!(msgs[2].role, crate::llm::Role::User);
}
#[test]
fn ensure_user_message_no_op_when_already_user() {
let mut msgs = vec![ChatMessage::user("hello")];
ensure_ends_with_user_message(&mut msgs);
assert_eq!(msgs.len(), 1);
}
// ── llm_signals_completion ──
#[test]
+5 -1
View File
@@ -151,7 +151,7 @@ Job: {}
Description: {}
You have tools for shell commands, file operations, and code editing.
Work independently to complete this job. Report when done."#,
Work independently to complete this job. When finished, your final message MUST include the phrase "The job is complete" to signal termination."#,
job.title, job.description
)));
@@ -373,6 +373,10 @@ impl LoopDelegate for ContainerDelegate {
// Poll for follow-up prompts from the user
self.poll_and_inject_prompt(reason_ctx).await;
// Claude 4.6 rejects assistant prefill; NEAR AI rejects any non-user-ending
// conversation. Ensure the last message is user-role before calling the LLM.
crate::util::ensure_ends_with_user_message(&mut reason_ctx.messages);
// Refresh tools (in case WASM tools were built)
reason_ctx.available_tools = self.tools.tool_definitions().await;
+12 -7
View File
@@ -18,7 +18,6 @@ use crate::agent::agentic_loop::{
};
use crate::agent::scheduler::WorkerMessage;
use crate::agent::task::TaskOutput;
use crate::channels::web::types::SseEvent;
use crate::context::{ContextManager, JobState};
use crate::db::Database;
use crate::error::Error;
@@ -33,6 +32,7 @@ use crate::tools::rate_limiter::RateLimitResult;
use crate::tools::{
ApprovalContext, ToolRegistry, autonomous_unavailable_error, prepare_tool_params, redact_params,
};
use ironclaw_common::AppEvent;
/// Shared dependencies for worker execution.
///
@@ -48,7 +48,7 @@ pub struct WorkerDeps {
pub hooks: Arc<HookRegistry>,
pub timeout: Duration,
pub use_planning: bool,
/// SSE manager for live job event streaming to the web gateway.
/// Broadcast sender for live job event streaming to the web gateway.
pub sse_tx: Option<Arc<crate::channels::web::sse::SseManager>>,
/// Approval context for tool execution. When `None`, all non-`Never` tools are
/// blocked (legacy behavior). When `Some`, the context determines which tools
@@ -141,7 +141,7 @@ impl Worker {
if let Some(ref sse) = self.deps.sse_tx {
let job_id_str = job_id.to_string();
let event = match event_type {
"message" => Some(SseEvent::JobMessage {
"message" => Some(AppEvent::JobMessage {
job_id: job_id_str,
role: data
.get("role")
@@ -154,7 +154,7 @@ impl Worker {
.unwrap_or("")
.to_string(),
}),
"tool_use" => Some(SseEvent::JobToolUse {
"tool_use" => Some(AppEvent::JobToolUse {
job_id: job_id_str,
tool_name: data
.get("tool_name")
@@ -166,7 +166,7 @@ impl Worker {
.cloned()
.unwrap_or(serde_json::Value::Null),
}),
"tool_result" => Some(SseEvent::JobToolResult {
"tool_result" => Some(AppEvent::JobToolResult {
job_id: job_id_str,
tool_name: data
.get("tool_name")
@@ -179,7 +179,7 @@ impl Worker {
.unwrap_or("")
.to_string(),
}),
"status" => Some(SseEvent::JobStatus {
"status" => Some(AppEvent::JobStatus {
job_id: job_id_str,
message: data
.get("message")
@@ -187,7 +187,7 @@ impl Worker {
.unwrap_or("")
.to_string(),
}),
"result" => Some(SseEvent::JobResult {
"result" => Some(AppEvent::JobResult {
job_id: job_id_str,
status: data
.get("status")
@@ -1232,6 +1232,11 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
) -> Option<LoopOutcome> {
// Refresh tool definitions so newly built tools become visible
reason_ctx.available_tools = self.worker.tools().tool_definitions().await;
// Claude 4.6 rejects assistant prefill; NEAR AI rejects any non-user-ending
// conversation. Ensure the last message is user-role before calling the LLM.
crate::util::ensure_ends_with_user_message(&mut reason_ctx.messages);
None
}
+9
View File
@@ -53,6 +53,7 @@ HEADED=1 pytest scenarios/
| `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 via `page.evaluate("showApproval(...)")`; the waiting-approval regression uses a real HTTP tool call |
| `test_oauth_refresh.py` | Hosted Gmail OAuth regression: complete setup via `/oauth/callback`, expire the stored access token in libSQL, trigger a real `gmail` tool call through `/api/chat/send`, and verify refresh goes through the mock `/oauth/refresh` proxy without forwarding `client_secret` |
## `helpers.py`
@@ -75,6 +76,7 @@ All fixtures are defined in `tests/e2e/conftest.py`. Running `pytest scenarios/`
| `ironclaw_binary` | Checks `target/debug/ironclaw`; if absent, runs `cargo build --no-default-features --features libsql` (timeout 600s). |
| `mock_llm_server` | Starts `mock_llm.py --port 0`, reads the assigned port from stdout, waits for `/v1/models` to return 200. Yields the base URL. |
| `ironclaw_server` | Starts the ironclaw binary with a minimal env (see below), waits for `/api/health` (timeout 60s). Yields the base URL. On teardown sends **SIGINT** (not SIGTERM) so the tokio ctrl_c handler triggers a graceful shutdown and LLVM coverage data is flushed. |
| `hosted_oauth_refresh_server` | Starts a second ironclaw instance with a dedicated libSQL DB and `GOOGLE_OAUTH_CLIENT_ID=hosted-google-client-id`, while still pointing `IRONCLAW_OAUTH_EXCHANGE_URL` at `mock_llm.py`. Yields a dict with `base_url`, `db_path`, `gateway_user_id`, and `mock_llm_url` for the hosted refresh regression scenario. |
| `browser` | Launches a single Chromium instance (headless by default; set `HEADED=1` for headed). Shared across all tests. |
### Function-scoped fixtures
@@ -100,6 +102,8 @@ EMBEDDING_ENABLED=false, SKILLS_ENABLED=true
ONBOARD_COMPLETED=true # prevents setup wizard
```
The `hosted_oauth_refresh_server` fixture uses the same baseline, but with its own DB/home tempdirs and `GOOGLE_OAUTH_CLIENT_ID=hosted-google-client-id` so hosted OAuth flows exercise proxy credential injection instead of the baked-in desktop Google app.
The binary is also started with `--no-onboard`. Coverage env vars (`CARGO_LLVM_COV*`, `LLVM_*`, `CARGO_ENCODED_RUSTFLAGS`, `CARGO_INCREMENTAL`) are forwarded from the outer environment when present.
## Mock LLM (`mock_llm.py`)
@@ -113,6 +117,11 @@ python mock_llm.py --port 0
It serves `POST /v1/chat/completions` (streaming + non-streaming) and `GET /v1/models`. Responses are pattern-matched from `CANNED_RESPONSES` against the last user message. Unmatched messages return `"I understand your request."`. The model name reported is always `"mock-model"`.
It also hosts OAuth test endpoints:
- `POST /oauth/exchange` for hosted auth-code exchange
- `POST /oauth/refresh` for hosted refresh-token exchange
- `GET /__mock/oauth/state` and `POST /__mock/oauth/reset` so HTTP E2E scenarios can assert exact proxy payloads and reset counters between setup and refresh assertions
To add a new canned response:
```python
# In mock_llm.py
+169 -34
View File
@@ -112,6 +112,39 @@ def _reserve_loopback_sockets(count: int) -> list[socket.socket]:
sock.close()
raise
async def _stop_process(
proc: asyncio.subprocess.Process, *, sig: int | None = None, timeout: float
) -> None:
"""Signal a subprocess and wait briefly without masking exit races."""
if proc.returncode is not None:
return
try:
if sig is None:
proc.kill()
else:
proc.send_signal(sig)
except ProcessLookupError:
try:
await asyncio.wait_for(proc.wait(), timeout=timeout)
except asyncio.TimeoutError:
pass
return
try:
await asyncio.wait_for(proc.wait(), timeout=timeout)
except asyncio.TimeoutError:
pass
def _forward_coverage_env(env: dict[str, str]) -> None:
"""Forward cargo-llvm-cov env vars into child processes when present."""
cov_env_prefixes = ("CARGO_LLVM_COV", "LLVM_")
cov_env_extras = ("CARGO_ENCODED_RUSTFLAGS", "CARGO_INCREMENTAL")
for key, val in os.environ.items():
if key.startswith(cov_env_prefixes) or key in cov_env_extras:
env[key] = val
@pytest.fixture(scope="session")
def ironclaw_binary():
@@ -264,14 +297,7 @@ async def ironclaw_server(
"IRONCLAW_OAUTH_CALLBACK_URL": "https://oauth.test.example/oauth/callback",
"IRONCLAW_OAUTH_EXCHANGE_URL": mock_llm_server,
}
# Forward LLVM coverage instrumentation env vars when present
# (allows cargo-llvm-cov to collect profraw data from E2E runs).
# Use prefix matching to stay resilient to cargo-llvm-cov changes.
COV_ENV_PREFIXES = ("CARGO_LLVM_COV", "LLVM_")
COV_ENV_EXTRAS = ("CARGO_ENCODED_RUSTFLAGS", "CARGO_INCREMENTAL")
for key, val in os.environ.items():
if key.startswith(COV_ENV_PREFIXES) or key in COV_ENV_EXTRAS:
env[key] = val
_forward_coverage_env(env)
proc = await asyncio.create_subprocess_exec(
ironclaw_binary, "--no-onboard",
stdin=asyncio.subprocess.DEVNULL,
@@ -279,35 +305,145 @@ async def ironclaw_server(
stderr=asyncio.subprocess.PIPE,
env=env,
)
startup_kill_attempted = False
base_url = f"http://127.0.0.1:{gateway_port}"
try:
await wait_for_ready(f"{base_url}/api/health", timeout=60)
yield base_url
except TimeoutError:
# Dump stderr so CI logs show why the server failed to start
if proc.returncode is None:
startup_kill_attempted = True
await _stop_process(proc, timeout=2)
returncode = proc.returncode
stderr_bytes = b""
if proc.stderr:
try:
stderr_bytes = await asyncio.wait_for(proc.stderr.read(8192), timeout=2)
except (asyncio.TimeoutError, Exception):
except asyncio.TimeoutError:
pass
stderr_text = stderr_bytes.decode("utf-8", errors="replace")
proc.kill()
pytest.fail(
f"ironclaw server failed to start on port {gateway_port} "
f"(returncode={returncode}).\nstderr:\n{stderr_text}"
)
finally:
if proc.returncode is None:
# Use SIGINT (not SIGTERM) so tokio's ctrl_c handler triggers a
# graceful shutdown. This lets the LLVM coverage runtime run its
# atexit handler and flush .profraw files for cargo-llvm-cov.
proc.send_signal(signal.SIGINT)
try:
await asyncio.wait_for(proc.wait(), timeout=10)
except asyncio.TimeoutError:
proc.kill()
if startup_kill_attempted:
await _stop_process(proc, timeout=2)
else:
# Use SIGINT (not SIGTERM) so tokio's ctrl_c handler triggers a
# graceful shutdown. This lets the LLVM coverage runtime run its
# atexit handler and flush .profraw files for cargo-llvm-cov.
await _stop_process(proc, sig=signal.SIGINT, timeout=10)
if proc.returncode is None:
await _stop_process(proc, timeout=2)
@pytest.fixture(scope="session")
async def hosted_oauth_refresh_server(
ironclaw_binary,
mock_llm_server,
wasm_tools_dir,
):
"""Start a hosted-mode ironclaw instance for OAuth refresh regression tests."""
reserved = _reserve_loopback_sockets(2)
db_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-hosted-oauth-db-")
home_tmpdir = tempfile.TemporaryDirectory(prefix="ironclaw-e2e-hosted-oauth-home-")
try:
gateway_port = reserved[0].getsockname()[1]
http_port = reserved[1].getsockname()[1]
for sock in reserved:
if sock.fileno() != -1:
sock.close()
db_path = os.path.join(db_tmpdir.name, "hosted-oauth-refresh.db")
home_dir = home_tmpdir.name
env = {
"PATH": os.environ.get("PATH", "/usr/bin:/bin"),
"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": OWNER_SCOPE_ID,
"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,
"LLM_MODEL": "mock-model",
"DATABASE_BACKEND": "libsql",
"LIBSQL_PATH": db_path,
"SECRETS_MASTER_KEY": "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef",
"SANDBOX_ENABLED": "false",
"SKILLS_ENABLED": "true",
"ROUTINES_ENABLED": "true",
"HEARTBEAT_ENABLED": "false",
"EMBEDDING_ENABLED": "false",
"WASM_ENABLED": "true",
"WASM_TOOLS_DIR": wasm_tools_dir,
"WASM_CHANNELS_DIR": _WASM_CHANNELS_TMPDIR.name,
"ONBOARD_COMPLETED": "true",
"IRONCLAW_OAUTH_CALLBACK_URL": "https://oauth.test.example/oauth/callback",
"IRONCLAW_OAUTH_EXCHANGE_URL": mock_llm_server,
"GOOGLE_OAUTH_CLIENT_ID": "hosted-google-client-id",
}
_forward_coverage_env(env)
proc = await asyncio.create_subprocess_exec(
ironclaw_binary, "--no-onboard",
stdin=asyncio.subprocess.DEVNULL,
stdout=asyncio.subprocess.PIPE,
stderr=asyncio.subprocess.PIPE,
env=env,
)
startup_kill_attempted = False
base_url = f"http://127.0.0.1:{gateway_port}"
try:
await wait_for_ready(f"{base_url}/api/health", timeout=60)
yield {
"base_url": base_url,
"db_path": db_path,
"gateway_user_id": OWNER_SCOPE_ID,
"mock_llm_url": mock_llm_server,
}
except TimeoutError:
if proc.returncode is None:
startup_kill_attempted = True
await _stop_process(proc, timeout=2)
returncode = proc.returncode
stderr_bytes = b""
if proc.stderr:
try:
stderr_bytes = await asyncio.wait_for(proc.stderr.read(8192), timeout=2)
except asyncio.TimeoutError:
pass
stderr_text = stderr_bytes.decode("utf-8", errors="replace")
pytest.fail(
f"hosted oauth refresh server failed to start on port {gateway_port} "
f"(returncode={returncode}).\nstderr:\n{stderr_text}"
)
finally:
if proc.returncode is None:
if startup_kill_attempted:
await _stop_process(proc, timeout=2)
else:
await _stop_process(proc, sig=signal.SIGINT, timeout=10)
if proc.returncode is None:
await _stop_process(proc, timeout=2)
finally:
for sock in reserved:
if sock.fileno() != -1:
sock.close()
db_tmpdir.cleanup()
home_tmpdir.cleanup()
@pytest.fixture(scope="session")
@@ -362,12 +498,7 @@ async def http_channel_server_without_secret(
"IRONCLAW_OAUTH_CALLBACK_URL": "https://oauth.test.example/oauth/callback",
"IRONCLAW_OAUTH_EXCHANGE_URL": mock_llm_server,
}
# Forward LLVM coverage instrumentation env vars when present
COV_ENV_PREFIXES = ("CARGO_LLVM_COV", "LLVM_")
COV_ENV_EXTRAS = ("CARGO_ENCODED_RUSTFLAGS", "CARGO_INCREMENTAL")
for key, val in os.environ.items():
if key.startswith(COV_ENV_PREFIXES) or key in COV_ENV_EXTRAS:
env[key] = val
_forward_coverage_env(env)
proc = await asyncio.create_subprocess_exec(
ironclaw_binary, "--no-onboard",
stdin=asyncio.subprocess.DEVNULL,
@@ -375,6 +506,7 @@ async def http_channel_server_without_secret(
stderr=asyncio.subprocess.PIPE,
env=env,
)
startup_kill_attempted = False
gateway_url = f"http://127.0.0.1:{gateway_port}"
http_base_url = f"http://127.0.0.1:{http_port}"
try:
@@ -383,15 +515,17 @@ async def http_channel_server_without_secret(
yield http_base_url
except TimeoutError:
# Dump stderr so CI logs show why the server failed to start
if proc.returncode is None:
startup_kill_attempted = True
await _stop_process(proc, timeout=2)
returncode = proc.returncode
stderr_bytes = b""
if proc.stderr:
try:
stderr_bytes = await asyncio.wait_for(proc.stderr.read(8192), timeout=2)
except (asyncio.TimeoutError, Exception):
except asyncio.TimeoutError:
pass
stderr_text = stderr_bytes.decode("utf-8", errors="replace")
proc.kill()
pytest.fail(
f"ironclaw server without webhook secret failed to start on ports "
f"gateway={gateway_port}, http={http_port} "
@@ -399,14 +533,15 @@ async def http_channel_server_without_secret(
)
finally:
if proc.returncode is None:
# Use SIGINT (not SIGTERM) so tokio's ctrl_c handler triggers a
# graceful shutdown. This lets the LLVM coverage runtime run its
# atexit handler and flush .profraw files for cargo-llvm-cov.
proc.send_signal(signal.SIGINT)
try:
await asyncio.wait_for(proc.wait(), timeout=10)
except asyncio.TimeoutError:
proc.kill()
if startup_kill_attempted:
await _stop_process(proc, timeout=2)
else:
# Use SIGINT (not SIGTERM) so tokio's ctrl_c handler triggers a
# graceful shutdown. This lets the LLVM coverage runtime run its
# atexit handler and flush .profraw files for cargo-llvm-cov.
await _stop_process(proc, sig=signal.SIGINT, timeout=10)
if proc.returncode is None:
await _stop_process(proc, timeout=2)
@pytest.fixture(scope="session")
+61
View File
@@ -34,6 +34,15 @@ TOOL_CALL_PATTERNS = [
"body": {"label": m.group("label")},
},
),
(
re.compile(r"check gmail unread|gmail unread", re.IGNORECASE),
"gmail",
lambda _: {
"action": "list_messages",
"query": "is:unread",
"max_results": 1,
},
),
(re.compile(r"what time|current time", re.IGNORECASE), "time", lambda _: {"operation": "now"}),
(
re.compile(
@@ -91,6 +100,15 @@ TOOL_CALL_PATTERNS = [
]
def _new_oauth_state() -> dict:
return {
"exchange_count": 0,
"refresh_count": 0,
"last_exchange": None,
"last_refresh": None,
}
def _last_user_content(messages: list[dict]) -> str:
for msg in reversed(messages):
if msg.get("role") == "user":
@@ -272,6 +290,12 @@ async def oauth_exchange(request: web.Request) -> web.Response:
specific token params such as RFC 8707 `resource` are forwarded here.
"""
data = await request.post()
oauth_state = request.app["oauth_state"]
oauth_state["exchange_count"] += 1
oauth_state["last_exchange"] = {
"authorization": request.headers.get("Authorization"),
"form": dict(data),
}
code = data.get("code", "")
access_token_field = data.get("access_token_field", "access_token")
@@ -290,6 +314,39 @@ async def oauth_exchange(request: web.Request) -> web.Response:
})
async def oauth_refresh(request: web.Request) -> web.Response:
"""Mock OAuth token refresh proxy for hosted refresh E2E tests."""
data = await request.post()
oauth_state = request.app["oauth_state"]
oauth_state["refresh_count"] += 1
oauth_state["last_refresh"] = {
"authorization": request.headers.get("Authorization"),
"form": dict(data),
}
if request.headers.get("Authorization") != "Bearer e2e-test-token":
return web.json_response({"error": "invalid_gateway_auth"}, status=401)
if data.get("client_id") != "hosted-google-client-id":
return web.json_response({"error": "invalid_client_id"}, status=400)
if "client_secret" in data:
return web.json_response({"error": "unexpected_client_secret"}, status=400)
return web.json_response({
"access_token": "mock-refreshed-access-token",
"refresh_token": "mock-rotated-refresh-token",
"expires_in": 3600,
})
async def oauth_state_handler(request: web.Request) -> web.Response:
return web.json_response(request.app["oauth_state"])
async def oauth_reset(request: web.Request) -> web.Response:
request.app["oauth_state"] = _new_oauth_state()
return web.json_response({"ok": True})
async def models(_request: web.Request) -> web.Response:
return web.json_response({
"object": "list",
@@ -424,12 +481,16 @@ def main():
parser.add_argument("--port", type=int, default=0)
args = parser.parse_args()
app = web.Application()
app["oauth_state"] = _new_oauth_state()
# Register both /v1/ and non-/v1/ paths (rig-core omits the /v1/ prefix)
app.router.add_post("/v1/chat/completions", chat_completions)
app.router.add_post("/chat/completions", chat_completions)
app.router.add_get("/v1/models", models)
app.router.add_get("/models", models)
app.router.add_post("/oauth/exchange", oauth_exchange)
app.router.add_post("/oauth/refresh", oauth_refresh)
app.router.add_get("/__mock/oauth/state", oauth_state_handler)
app.router.add_post("/__mock/oauth/reset", oauth_reset)
# Mock MCP server endpoints
app.router.add_post("/mcp", mcp_endpoint)
app.router.add_post("/mcp-400", mcp_endpoint_400)
+227
View File
@@ -0,0 +1,227 @@
"""Hosted OAuth refresh HTTP regression test.
Runs a real ironclaw binary in hosted mode, expires a stored Gmail access
token in the libSQL database, triggers a real gmail tool call through the
chat API, and verifies that refresh uses the hosted proxy endpoint.
"""
import asyncio
import sqlite3
from datetime import datetime, timezone
from urllib.parse import parse_qs, urlparse
import httpx
from helpers import api_get, api_post
def _extract_state(auth_url: str) -> str:
parsed = urlparse(auth_url)
state = parse_qs(parsed.query).get("state", [None])[0]
assert state, f"auth_url should include state: {auth_url}"
return state
def _parse_timestamp(value: str | None) -> datetime | None:
if value is None:
return None
return datetime.fromisoformat(value.replace("Z", "+00:00"))
def _expire_access_token(db_path: str, user_id: str, secret_name: str) -> None:
with sqlite3.connect(db_path) as conn:
cursor = conn.execute(
"""
UPDATE secrets
SET expires_at = strftime('%Y-%m-%dT%H:%M:%fZ', 'now', '-1 hour')
WHERE user_id = ?1 AND name = ?2
""",
(user_id, secret_name),
)
conn.commit()
assert cursor.rowcount == 1, f"Expected one secret row for {user_id}/{secret_name}"
def _find_secret_row(
db_path: str,
secret_name: str,
) -> tuple[str, str | None, str | None]:
with sqlite3.connect(db_path) as conn:
row = conn.execute(
"""
SELECT user_id, expires_at, updated_at
FROM secrets
WHERE name = ?1
ORDER BY updated_at DESC
LIMIT 1
""",
(secret_name,),
).fetchone()
assert row is not None, f"Missing secret row for {secret_name}"
return row[0], row[1], row[2]
async def _get_extension(base_url: str, name: str) -> dict | None:
response = await api_get(base_url, "/api/extensions", timeout=15)
response.raise_for_status()
for extension in response.json().get("extensions", []):
if extension["name"] == name:
return extension
return None
async def _reset_mock_oauth_state(mock_base_url: str) -> None:
async with httpx.AsyncClient() as client:
response = await client.post(f"{mock_base_url}/__mock/oauth/reset", timeout=10)
response.raise_for_status()
async def _get_mock_oauth_state(mock_base_url: str) -> dict:
async with httpx.AsyncClient() as client:
response = await client.get(f"{mock_base_url}/__mock/oauth/state", timeout=10)
response.raise_for_status()
return response.json()
async def _approve_pending_request(base_url: str, thread_id: str, request_id: str) -> None:
response = await api_post(
base_url,
"/api/chat/approval",
json={"request_id": request_id, "action": "approve", "thread_id": thread_id},
timeout=15,
)
assert response.status_code == 202, (
f"Approval submission failed: {response.status_code} {response.text[:400]}"
)
async def _wait_for_gmail_tool_call(base_url: str, thread_id: str, timeout: float = 30.0) -> dict:
approved_request_ids = set()
for _ in range(int(timeout * 2)):
response = await api_get(
base_url,
f"/api/chat/history?thread_id={thread_id}",
timeout=15,
)
response.raise_for_status()
history = response.json()
pending = history.get("pending_approval")
if pending and pending["request_id"] not in approved_request_ids:
await _approve_pending_request(base_url, thread_id, pending["request_id"])
approved_request_ids.add(pending["request_id"])
for turn in history.get("turns", []):
for tool_call in turn.get("tool_calls", []):
if tool_call.get("name") == "gmail":
return history
await asyncio.sleep(0.5)
raise AssertionError(f"Timed out waiting for gmail tool call in thread {thread_id}")
async def _wait_for_refresh_request(mock_base_url: str, timeout: float = 20.0) -> dict:
for _ in range(int(timeout * 2)):
state = await _get_mock_oauth_state(mock_base_url)
if state.get("refresh_count") == 1:
return state
await asyncio.sleep(0.5)
raise AssertionError("Timed out waiting for exactly one OAuth refresh request")
async def test_hosted_gmail_oauth_refresh_uses_proxy(hosted_oauth_refresh_server):
server = hosted_oauth_refresh_server["base_url"]
db_path = hosted_oauth_refresh_server["db_path"]
mock_base_url = hosted_oauth_refresh_server["mock_llm_url"]
install_response = await api_post(
server,
"/api/extensions/install",
json={"name": "gmail"},
timeout=180,
)
assert install_response.status_code == 200, install_response.text
assert install_response.json().get("success") is True
setup_response = await api_post(
server,
"/api/extensions/gmail/setup",
json={"secrets": {}},
timeout=30,
)
assert setup_response.status_code == 200, setup_response.text
setup_data = setup_response.json()
assert setup_data.get("success") is True, setup_data
auth_url = setup_data.get("auth_url")
assert auth_url, setup_data
auth_params = parse_qs(urlparse(auth_url).query)
assert auth_params.get("client_id") == ["hosted-google-client-id"]
async with httpx.AsyncClient() as client:
callback_response = await client.get(
f"{server}/oauth/callback",
params={"code": "mock_auth_code", "state": _extract_state(auth_url)},
timeout=30,
follow_redirects=True,
)
assert callback_response.status_code == 200, callback_response.text[:400]
callback_body = callback_response.text.lower()
assert "connected" in callback_body or "success" in callback_body
gmail = await _get_extension(server, "gmail")
assert gmail is not None, "gmail should be installed"
assert gmail["authenticated"] is True, gmail
assert "gmail" in gmail.get("tools", []), gmail
await _reset_mock_oauth_state(mock_base_url)
stored_user_id, expires_before, updated_before = _find_secret_row(
db_path, "google_oauth_token"
)
assert _parse_timestamp(expires_before) is not None
assert _parse_timestamp(updated_before) is not None
await asyncio.sleep(0.1)
_expire_access_token(db_path, stored_user_id, "google_oauth_token")
thread_response = await api_post(server, "/api/chat/thread/new", timeout=15)
assert thread_response.status_code == 200, thread_response.text
thread_id = thread_response.json()["id"]
send_response = await api_post(
server,
"/api/chat/send",
json={"content": "check gmail unread", "thread_id": thread_id},
timeout=30,
)
assert send_response.status_code == 202, send_response.text
history = await _wait_for_gmail_tool_call(server, thread_id)
assert any(
tool_call.get("name") == "gmail"
for turn in history.get("turns", [])
for tool_call in turn.get("tool_calls", [])
), history
oauth_state = await _wait_for_refresh_request(mock_base_url)
assert oauth_state["refresh_count"] == 1, oauth_state
last_refresh = oauth_state["last_refresh"]
assert last_refresh is not None, oauth_state
assert last_refresh["authorization"] == "Bearer e2e-test-token"
assert last_refresh["form"]["client_id"] == "hosted-google-client-id"
assert "client_secret" not in last_refresh["form"], last_refresh
refreshed_user_id, expires_after, updated_after = _find_secret_row(
db_path, "google_oauth_token"
)
assert refreshed_user_id == stored_user_id
expires_after_dt = _parse_timestamp(expires_after)
updated_after_dt = _parse_timestamp(updated_after)
updated_before_dt = _parse_timestamp(updated_before)
assert expires_after_dt is not None
assert updated_after_dt is not None
assert updated_before_dt is not None
assert expires_after_dt > datetime.now(timezone.utc)
assert updated_after_dt > updated_before_dt
+143 -27
View File
@@ -19,10 +19,13 @@ use axum::middleware;
use axum::routing::{get, post};
use tower::ServiceExt;
use ironclaw::channels::IncomingMessage;
use ironclaw::channels::web::auth::{
AuthenticatedUser, MultiAuthState, UserIdentity, auth_middleware,
};
use ironclaw::channels::web::server::{GatewayState, PerUserRateLimiter, RateLimiter};
use ironclaw::channels::web::server::{
GatewayState, PerUserRateLimiter, RateLimiter, start_server,
};
use ironclaw::channels::web::sse::SseManager;
use ironclaw::channels::web::test_helpers::TestGatewayBuilder;
use ironclaw::channels::web::ws::WsConnectionTracker;
@@ -37,6 +40,9 @@ const ALICE_TOKEN: &str = "tok-alice-secret";
const BOB_TOKEN: &str = "tok-bob-secret";
const ALICE_USER_ID: &str = "alice";
const BOB_USER_ID: &str = "bob";
const OWNER_TOKEN: &str = "tok-owner-secret";
const OWNER_SCOPE_ID: &str = "owner-scope";
const GATEWAY_SENDER_ID: &str = "gateway-sender";
/// Build a MultiAuthState with two users.
fn two_user_auth() -> MultiAuthState {
@@ -301,7 +307,7 @@ fn per_user_rate_limiter_single_user_mode() {
#[tokio::test]
async fn sse_scoped_event_only_delivered_to_target_user() {
use ironclaw::channels::web::types::SseEvent;
use ironclaw_common::AppEvent;
use tokio_stream::StreamExt;
let manager = SseManager::new();
@@ -319,34 +325,34 @@ async fn sse_scoped_event_only_delivered_to_target_user() {
// Send event scoped to alice
manager.broadcast_for_user(
ALICE_USER_ID,
SseEvent::Status {
AppEvent::Status {
message: "alice's event".to_string(),
thread_id: None,
},
);
// Send global heartbeat (both should get it)
manager.broadcast(SseEvent::Heartbeat);
manager.broadcast(AppEvent::Heartbeat);
// Alice gets her scoped event first
let e = alice_stream.next().await.unwrap();
match &e {
SseEvent::Status { message, .. } => assert_eq!(message, "alice's event"),
AppEvent::Status { message, .. } => assert_eq!(message, "alice's event"),
_ => panic!("Expected Status, got {:?}", e),
}
// Alice also gets heartbeat
let e = alice_stream.next().await.unwrap();
assert!(matches!(e, SseEvent::Heartbeat));
assert!(matches!(e, AppEvent::Heartbeat));
// Bob only gets the heartbeat (alice's event was filtered)
let e = bob_stream.next().await.unwrap();
assert!(matches!(e, SseEvent::Heartbeat));
assert!(matches!(e, AppEvent::Heartbeat));
}
#[tokio::test]
async fn sse_global_event_delivered_to_all_users() {
use ironclaw::channels::web::types::SseEvent;
use ironclaw_common::AppEvent;
use tokio_stream::StreamExt;
let manager = SseManager::new();
@@ -361,7 +367,7 @@ async fn sse_global_event_delivered_to_all_users() {
.expect("subscribe"),
);
manager.broadcast(SseEvent::Status {
manager.broadcast(AppEvent::Status {
message: "global announcement".to_string(),
thread_id: None,
});
@@ -369,7 +375,7 @@ async fn sse_global_event_delivered_to_all_users() {
let ea = alice.next().await.unwrap();
let eb = bob.next().await.unwrap();
match (&ea, &eb) {
(SseEvent::Status { message: a, .. }, SseEvent::Status { message: b, .. }) => {
(AppEvent::Status { message: a, .. }, AppEvent::Status { message: b, .. }) => {
assert_eq!(a, "global announcement");
assert_eq!(b, "global announcement");
}
@@ -379,7 +385,7 @@ async fn sse_global_event_delivered_to_all_users() {
#[tokio::test]
async fn sse_user_b_event_not_visible_to_user_a() {
use ironclaw::channels::web::types::SseEvent;
use ironclaw_common::AppEvent;
use tokio_stream::StreamExt;
let manager = SseManager::new();
@@ -392,19 +398,19 @@ async fn sse_user_b_event_not_visible_to_user_a() {
// Send event for bob only
manager.broadcast_for_user(
BOB_USER_ID,
SseEvent::Response {
AppEvent::Response {
content: "bob's secret".to_string(),
thread_id: "t1".to_string(),
},
);
// Send heartbeat so alice has something to receive
manager.broadcast(SseEvent::Heartbeat);
manager.broadcast(AppEvent::Heartbeat);
// Alice should only get heartbeat, not bob's response
let e = alice.next().await.unwrap();
assert!(
matches!(e, SseEvent::Heartbeat),
matches!(e, AppEvent::Heartbeat),
"Expected Heartbeat, got {:?}",
e
);
@@ -412,7 +418,7 @@ async fn sse_user_b_event_not_visible_to_user_a() {
#[tokio::test]
async fn sse_unscoped_subscriber_receives_all_events() {
use ironclaw::channels::web::types::SseEvent;
use ironclaw_common::AppEvent;
use tokio_stream::StreamExt;
let manager = SseManager::new();
@@ -421,19 +427,19 @@ async fn sse_unscoped_subscriber_receives_all_events() {
manager.broadcast_for_user(
ALICE_USER_ID,
SseEvent::Status {
AppEvent::Status {
message: "alice only".to_string(),
thread_id: None,
},
);
manager.broadcast_for_user(
BOB_USER_ID,
SseEvent::Status {
AppEvent::Status {
message: "bob only".to_string(),
thread_id: None,
},
);
manager.broadcast(SseEvent::Heartbeat);
manager.broadcast(AppEvent::Heartbeat);
// Unscoped subscriber gets ALL three events
let e1 = stream.next().await.unwrap();
@@ -441,14 +447,14 @@ async fn sse_unscoped_subscriber_receives_all_events() {
let e3 = stream.next().await.unwrap();
match &e1 {
SseEvent::Status { message, .. } => assert_eq!(message, "alice only"),
AppEvent::Status { message, .. } => assert_eq!(message, "alice only"),
_ => panic!("Expected alice's Status"),
}
match &e2 {
SseEvent::Status { message, .. } => assert_eq!(message, "bob only"),
AppEvent::Status { message, .. } => assert_eq!(message, "bob only"),
_ => panic!("Expected bob's Status"),
}
assert!(matches!(e3, SseEvent::Heartbeat));
assert!(matches!(e3, AppEvent::Heartbeat));
}
// ===========================================================================
@@ -537,7 +543,8 @@ fn gateway_state_has_multi_tenant_fields() {
job_manager: None,
prompt_queue: None,
scheduler: None,
default_user_id: "fallback".to_string(), // Multi-tenant: renamed from user_id
owner_id: "fallback".to_string(),
default_sender_id: "fallback".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: None,
@@ -554,7 +561,8 @@ fn gateway_state_has_multi_tenant_fields() {
secrets_store: None,
};
assert_eq!(state.default_user_id, "fallback");
assert_eq!(state.owner_id, "fallback");
assert_eq!(state.default_sender_id, "fallback");
assert!(state.workspace_pool.is_none());
}
@@ -573,6 +581,70 @@ async fn start_multi_user_server() -> (SocketAddr, Arc<GatewayState>) {
.expect("Failed to start multi-user test server")
}
async fn start_owner_scoped_sender_server() -> (
SocketAddr,
Arc<GatewayState>,
tokio::sync::mpsc::Receiver<IncomingMessage>,
) {
let (agent_tx, agent_rx) = tokio::sync::mpsc::channel(64);
let mut tokens = HashMap::new();
tokens.insert(
OWNER_TOKEN.to_string(),
UserIdentity {
user_id: OWNER_SCOPE_ID.to_string(),
workspace_read_scopes: Vec::new(),
},
);
tokens.insert(
BOB_TOKEN.to_string(),
UserIdentity {
user_id: BOB_USER_ID.to_string(),
workspace_read_scopes: Vec::new(),
},
);
let state = Arc::new(GatewayState {
msg_tx: tokio::sync::RwLock::new(Some(agent_tx)),
sse: Arc::new(SseManager::new()),
workspace: None,
workspace_pool: None,
session_manager: None,
log_broadcaster: None,
log_level_handle: None,
extension_manager: None,
tool_registry: None,
store: None,
job_manager: None,
prompt_queue: None,
scheduler: None,
owner_id: OWNER_SCOPE_ID.to_string(),
default_sender_id: GATEWAY_SENDER_ID.to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: None,
skill_registry: None,
skill_catalog: None,
chat_rate_limiter: PerUserRateLimiter::new(30, 60),
oauth_rate_limiter: RateLimiter::new(10, 60),
webhook_rate_limiter: RateLimiter::new(10, 60),
registry_entries: Vec::new(),
cost_guard: None,
routine_engine: Arc::new(tokio::sync::RwLock::new(None)),
startup_time: std::time::Instant::now(),
active_config: Default::default(),
secrets_store: None,
});
let auth = MultiAuthState::multi(tokens);
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
let bound = start_server(addr, state.clone(), auth)
.await
.expect("Failed to start owner-scoped sender test server");
(bound, state, agent_rx)
}
#[tokio::test]
async fn full_server_alice_can_access_protected_endpoint() {
let (addr, _state) = start_multi_user_server().await;
@@ -678,6 +750,49 @@ async fn full_server_chat_send_accepted_for_alice() {
assert_eq!(msg.channel, "gateway");
}
#[tokio::test]
async fn full_server_chat_send_rewrites_sender_only_for_owner_scope_rebind() {
let (addr, _state, mut agent_rx) = start_owner_scoped_sender_server().await;
let client = reqwest::Client::new();
let owner_resp = client
.post(format!("http://{}/api/chat/send", addr))
.header("Authorization", format!("Bearer {}", OWNER_TOKEN))
.header("Content-Type", "application/json")
.body(r#"{"content":"hello from owner"}"#)
.send()
.await
.unwrap();
assert_eq!(owner_resp.status(), 202);
let owner_msg = tokio::time::timeout(Duration::from_secs(2), agent_rx.recv())
.await
.expect("Timed out waiting for owner message")
.expect("Agent channel closed");
assert_eq!(owner_msg.user_id, OWNER_SCOPE_ID);
assert_eq!(owner_msg.sender_id, GATEWAY_SENDER_ID);
assert_eq!(owner_msg.content, "hello from owner");
let other_resp = client
.post(format!("http://{}/api/chat/send", addr))
.header("Authorization", format!("Bearer {}", BOB_TOKEN))
.header("Content-Type", "application/json")
.body(r#"{"content":"hello from bob"}"#)
.send()
.await
.unwrap();
assert_eq!(other_resp.status(), 202);
let other_msg = tokio::time::timeout(Duration::from_secs(2), agent_rx.recv())
.await
.expect("Timed out waiting for non-owner message")
.expect("Agent channel closed");
assert_eq!(other_msg.user_id, BOB_USER_ID);
assert_eq!(other_msg.sender_id, BOB_USER_ID);
assert_eq!(other_msg.content, "hello from bob");
}
#[tokio::test]
async fn full_server_chat_send_rejected_without_auth() {
let (addr, _state) = start_multi_user_server().await;
@@ -768,7 +883,7 @@ async fn full_server_jobs_endpoint_rejected_without_auth() {
#[tokio::test]
async fn full_server_ws_multi_user_event_isolation() {
use futures::StreamExt;
use ironclaw::channels::web::types::SseEvent;
use ironclaw_common::AppEvent;
use tokio_tungstenite::tungstenite::Message;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
@@ -801,14 +916,14 @@ async fn full_server_ws_multi_user_event_isolation() {
// Broadcast an event scoped to Alice only
state.sse.broadcast_for_user(
ALICE_USER_ID,
SseEvent::Status {
AppEvent::Status {
message: "alice-only-event".to_string(),
thread_id: None,
},
);
// Broadcast a global heartbeat so Bob has something to receive
state.sse.broadcast(SseEvent::Heartbeat);
state.sse.broadcast(AppEvent::Heartbeat);
// Alice should get her scoped event
let alice_msg = tokio::time::timeout(Duration::from_secs(2), alice_ws.next())
@@ -889,7 +1004,8 @@ async fn start_multi_user_server_with_db() -> (
job_manager: None,
prompt_queue: None,
scheduler: None,
default_user_id: ALICE_USER_ID.to_string(),
owner_id: ALICE_USER_ID.to_string(),
default_sender_id: ALICE_USER_ID.to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: None,
+4 -2
View File
@@ -203,7 +203,8 @@ async fn start_test_server_with_provider(
job_manager: None,
prompt_queue: None,
scheduler: None,
default_user_id: "test-user".to_string(),
owner_id: "test-user".to_string(),
default_sender_id: "test-user".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: Some(llm_provider),
@@ -702,7 +703,8 @@ async fn test_no_llm_provider_returns_503() {
job_manager: None,
prompt_queue: None,
scheduler: None,
default_user_id: "test-user".to_string(),
owner_id: "test-user".to_string(),
default_sender_id: "test-user".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: None, // No LLM!
+2 -1
View File
@@ -226,7 +226,8 @@ impl GatewayWorkflowHarness {
job_manager: None,
prompt_queue: None,
scheduler: Some(scheduler_slot.clone()),
default_user_id: user_id.clone(),
owner_id: user_id.clone(),
default_sender_id: user_id.clone(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: Some(Arc::clone(&components.llm)),
+11 -10
View File
@@ -5,7 +5,7 @@
//! - WebSocket upgrade with auth
//! - Ping/pong
//! - Client message → agent msg_tx
//! - Broadcast SSE event → WebSocket client
//! - Broadcast AppEvent → WebSocket client
//! - Connection tracking (counter increment/decrement)
//! - Gateway status endpoint
@@ -22,8 +22,8 @@ use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use ironclaw::channels::IncomingMessage;
use ironclaw::channels::web::server::{GatewayState, start_server};
use ironclaw::channels::web::sse::SseManager;
use ironclaw::channels::web::types::SseEvent;
use ironclaw::channels::web::ws::WsConnectionTracker;
use ironclaw_common::AppEvent;
const AUTH_TOKEN: &str = "test-token-12345";
const TIMEOUT: Duration = Duration::from_secs(5);
@@ -51,7 +51,8 @@ async fn start_test_server() -> (
job_manager: None,
prompt_queue: None,
scheduler: None,
default_user_id: "test-user".to_string(),
owner_id: "test-user".to_string(),
default_sender_id: "test-user".to_string(),
shutdown_tx: tokio::sync::RwLock::new(None),
ws_tracker: Some(Arc::new(WsConnectionTracker::new())),
llm_provider: None,
@@ -164,8 +165,8 @@ async fn test_ws_broadcast_event_received() {
// Give the connection a moment to fully establish
tokio::time::sleep(Duration::from_millis(50)).await;
// Broadcast an SSE event (simulates agent sending a response)
state.sse.broadcast(SseEvent::Response {
// Broadcast an event (simulates agent sending a response)
state.sse.broadcast(AppEvent::Response {
content: "agent says hi".to_string(),
thread_id: "t1".to_string(),
});
@@ -186,7 +187,7 @@ async fn test_ws_thinking_event() {
let mut ws = connect_ws(addr).await;
tokio::time::sleep(Duration::from_millis(50)).await;
state.sse.broadcast(SseEvent::Thinking {
state.sse.broadcast(AppEvent::Thinking {
message: "analyzing...".to_string(),
thread_id: None,
});
@@ -311,22 +312,22 @@ async fn test_ws_multiple_events_in_sequence() {
tokio::time::sleep(Duration::from_millis(50)).await;
// Broadcast multiple events rapidly
state.sse.broadcast(SseEvent::Thinking {
state.sse.broadcast(AppEvent::Thinking {
message: "step 1".to_string(),
thread_id: None,
});
state.sse.broadcast(SseEvent::ToolStarted {
state.sse.broadcast(AppEvent::ToolStarted {
name: "shell".to_string(),
thread_id: None,
});
state.sse.broadcast(SseEvent::ToolCompleted {
state.sse.broadcast(AppEvent::ToolCompleted {
name: "shell".to_string(),
success: true,
error: None,
parameters: None,
thread_id: None,
});
state.sse.broadcast(SseEvent::Response {
state.sse.broadcast(AppEvent::Response {
content: "done".to_string(),
thread_id: "t1".to_string(),
});