Compare commits

..
Author SHA1 Message Date
ZakiandClaude Opus 4.6 5b0d261398 feat(mcp): support custom HTTP headers for MCP server auth (#639)
Some MCP servers (e.g., Browser-Use) require custom headers instead of
OAuth for authentication. Add a `headers` field to McpServerConfig that
injects custom HTTP headers into every request to that server.

Changes:
- Add `headers: HashMap<String, String>` to McpServerConfig with serde
  default/skip_serializing_if for backward compatibility
- Add `with_headers()` builder and `new_with_config()` constructor
- Inject custom headers in McpClient::send_request() before auth header
- Add `--header` / `-H` CLI flag to `ironclaw mcp add` (Key:Value format)
- Update app.rs and extensions/manager.rs to use new_with_config()
- Show custom header names in verbose `mcp list` output

Co-Authored-By: Claude Opus 4.6 <[email protected]>
2026-03-07 18:26:30 -08:00
11 changed files with 202 additions and 593 deletions
Generated
-29
View File
@@ -864,16 +864,6 @@ dependencies = [
"windows-link",
]
[[package]]
name = "chrono-tz"
version = "0.10.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a6139a8597ed92cf816dfb33f5dd6cf0bb93a6adc938f11039f371bc5bcd26c3"
dependencies = [
"chrono",
"phf 0.12.1",
]
[[package]]
name = "cipher"
version = "0.4.4"
@@ -2882,7 +2872,6 @@ dependencies = [
"bollard",
"bytes",
"chrono",
"chrono-tz",
"clap",
"clap_complete",
"cron",
@@ -3903,15 +3892,6 @@ dependencies = [
"phf_shared 0.11.3",
]
[[package]]
name = "phf"
version = "0.12.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "913273894cec178f401a31ec4b656318d95473527be05c0752cc41cdc32be8b7"
dependencies = [
"phf_shared 0.12.1",
]
[[package]]
name = "phf"
version = "0.13.1"
@@ -3986,15 +3966,6 @@ dependencies = [
"uncased",
]
[[package]]
name = "phf_shared"
version = "0.12.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "06005508882fb681fd97892ecff4b7fd0fee13ef1aa569f8695dae7ab9099981"
dependencies = [
"siphasher",
]
[[package]]
name = "phf_shared"
version = "0.13.1"
-1
View File
@@ -73,7 +73,6 @@ toml = "0.8"
# Core types
uuid = { version = "1", features = ["v4", "v5", "serde"] }
chrono = { version = "0.4", features = ["serde"] }
chrono-tz = "0.10"
rust_decimal = { version = "1", features = ["serde", "serde-with-str", "maths"] }
rust_decimal_macros = "1"
+1 -1
View File
@@ -542,7 +542,7 @@ impl AppBuilder {
server, mcp_sm, secrets, "default",
)
} else {
McpClient::new_with_name(&server_name, &server.url)
McpClient::new_with_config(server.clone())
};
match client.list_tools().await {
-1
View File
@@ -26,4 +26,3 @@ pub mod routines;
pub mod settings;
#[allow(dead_code)]
pub mod static_files;
pub mod webhooks;
-210
View File
@@ -1,210 +0,0 @@
//! Public webhook trigger endpoint for routine webhook triggers.
//!
//! `POST /api/webhooks/{path}` — matches the path against routines with
//! `Trigger::Webhook { path, secret }`, validates the secret via constant-time
//! comparison, and fires the matching routine through the message pipeline.
use std::sync::Arc;
use axum::{
Json,
extract::{Path, State},
http::{HeaderMap, StatusCode},
};
use subtle::ConstantTimeEq;
use crate::agent::routine::{RoutineAction, Trigger};
use crate::channels::IncomingMessage;
use crate::channels::web::server::GatewayState;
/// Handle incoming webhook POST to `/api/webhooks/{path}`.
///
/// This endpoint is **public** (no gateway auth token required) but protected
/// by the per-routine webhook secret sent via the `X-Webhook-Secret` header.
pub async fn webhook_trigger_handler(
State(state): State<Arc<GatewayState>>,
Path(path): Path<String>,
headers: HeaderMap,
) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
let store = state.store.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Database not available".to_string(),
))?;
// Load all routines and find one whose Trigger::Webhook path matches.
let routines = store
.list_all_routines()
.await
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
let matched = routines.into_iter().find(|r| {
if !r.enabled {
return false;
}
match &r.trigger {
Trigger::Webhook { path: Some(wp), .. } => *wp == path,
Trigger::Webhook { path: None, .. } => path == r.id.to_string(),
_ => false,
}
});
let routine = matched.ok_or((
StatusCode::NOT_FOUND,
"No routine matches this webhook path".to_string(),
))?;
// Validate the webhook secret if one is configured on the routine.
if let Trigger::Webhook {
secret: Some(expected_secret),
..
} = &routine.trigger
{
let provided_secret = headers
.get("x-webhook-secret")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
if !bool::from(provided_secret.as_bytes().ct_eq(expected_secret.as_bytes())) {
return Err((
StatusCode::UNAUTHORIZED,
"Invalid webhook secret".to_string(),
));
}
}
// Build the prompt from the routine action.
let prompt = match &routine.action {
RoutineAction::Lightweight { prompt, .. } => prompt.clone(),
RoutineAction::FullJob {
title, description, ..
} => format!("{}: {}", title, description),
};
let content = format!("[routine:{}] {}", routine.name, prompt);
let thread_id = format!(
"routine-{}-{}",
routine.id,
chrono::Utc::now().timestamp_millis()
);
let msg = IncomingMessage::new("gateway", &routine.user_id, content).with_thread(thread_id);
let tx_guard = state.msg_tx.read().await;
let tx = tx_guard.as_ref().ok_or((
StatusCode::SERVICE_UNAVAILABLE,
"Channel not started".to_string(),
))?;
tx.send(msg).await.map_err(|_| {
(
StatusCode::INTERNAL_SERVER_ERROR,
"Channel closed".to_string(),
)
})?;
Ok(Json(serde_json::json!({
"status": "triggered",
"routine_id": routine.id,
"routine_name": routine.name,
})))
}
#[cfg(test)]
mod tests {
use super::*;
/// Verify constant-time comparison logic for webhook secrets.
#[test]
fn test_webhook_secret_constant_time_comparison() {
let expected = "my-secret-token";
// Matching secret
let provided = "my-secret-token";
assert!(bool::from(provided.as_bytes().ct_eq(expected.as_bytes())));
// Wrong secret
let wrong = "wrong-secret";
assert!(!bool::from(wrong.as_bytes().ct_eq(expected.as_bytes())));
// Empty secret
let empty = "";
assert!(!bool::from(empty.as_bytes().ct_eq(expected.as_bytes())));
}
/// Verify that webhook path matching logic works for both explicit paths
/// and fallback to routine ID.
#[test]
fn test_webhook_path_matching() {
use chrono::Utc;
use uuid::Uuid;
let routine_id = Uuid::parse_str("550e8400-e29b-41d4-a716-446655440000").unwrap();
let routine = crate::agent::routine::Routine {
id: routine_id,
name: "test-routine".to_string(),
description: "A test routine".to_string(),
user_id: "test-user".to_string(),
enabled: true,
trigger: Trigger::Webhook {
path: Some("my-hook".to_string()),
secret: None,
},
action: RoutineAction::Lightweight {
prompt: "do stuff".to_string(),
context_paths: vec![],
max_tokens: 4096,
},
guardrails: crate::agent::routine::RoutineGuardrails::default(),
notify: crate::agent::routine::NotifyConfig::default(),
last_run_at: None,
next_fire_at: None,
run_count: 0,
consecutive_failures: 0,
state: serde_json::Value::Null,
created_at: Utc::now(),
updated_at: Utc::now(),
};
// Explicit path match
let matches_explicit = match &routine.trigger {
Trigger::Webhook { path: Some(wp), .. } => *wp == "my-hook",
_ => false,
};
assert!(matches_explicit);
// Should NOT match wrong path
let matches_wrong = match &routine.trigger {
Trigger::Webhook { path: Some(wp), .. } => *wp == "other-hook",
_ => false,
};
assert!(!matches_wrong);
// Routine with no explicit path falls back to ID
let routine_no_path = crate::agent::routine::Routine {
trigger: Trigger::Webhook {
path: None,
secret: None,
},
..routine
};
let matches_id = match &routine_no_path.trigger {
Trigger::Webhook { path: None, .. } => {
routine_no_path.id.to_string() == "550e8400-e29b-41d4-a716-446655440000"
}
_ => false,
};
assert!(matches_id);
// Disabled routine should not match
let disabled_routine = crate::agent::routine::Routine {
enabled: false,
trigger: Trigger::Webhook {
path: Some("my-hook".to_string()),
secret: None,
},
..routine_no_path
};
let should_skip = !disabled_routine.enabled;
assert!(should_skip);
}
}
+1 -3
View File
@@ -37,7 +37,6 @@ use crate::channels::web::handlers::jobs::{
use crate::channels::web::handlers::skills::{
skills_install_handler, skills_list_handler, skills_remove_handler, skills_search_handler,
};
use crate::channels::web::handlers::webhooks::webhook_trigger_handler;
use crate::channels::web::log_layer::LogBroadcaster;
use crate::channels::web::sse::SseManager;
use crate::channels::web::types::*;
@@ -201,8 +200,7 @@ pub async fn start_server(
// Public routes (no auth)
let public = Router::new()
.route("/api/health", get(health_handler))
.route("/oauth/callback", get(oauth_callback_handler))
.route("/api/webhooks/{path}", post(webhook_trigger_handler));
.route("/oauth/callback", get(oauth_callback_handler));
// Protected routes (require auth)
let auth_state = AuthState { token: auth_token };
+32 -1
View File
@@ -47,6 +47,10 @@ pub enum McpCommand {
/// Server description
#[arg(long)]
description: Option<String>,
/// Custom HTTP headers (format: "Key:Value", can be repeated)
#[arg(long = "header", short = 'H')]
headers: Vec<String>,
},
/// Remove an MCP server
@@ -108,6 +112,7 @@ pub async fn run_mcp_command(cmd: McpCommand) -> anyhow::Result<()> {
token_url,
scopes,
description,
headers,
} => {
add_server(
name,
@@ -117,6 +122,7 @@ pub async fn run_mcp_command(cmd: McpCommand) -> anyhow::Result<()> {
token_url,
scopes,
description,
headers,
)
.await
}
@@ -133,6 +139,7 @@ pub async fn run_mcp_command(cmd: McpCommand) -> anyhow::Result<()> {
}
/// Add a new MCP server.
#[allow(clippy::too_many_arguments)]
async fn add_server(
name: String,
url: String,
@@ -141,6 +148,7 @@ async fn add_server(
token_url: Option<String>,
scopes: Option<String>,
description: Option<String>,
headers: Vec<String>,
) -> anyhow::Result<()> {
let mut config = McpServerConfig::new(&name, &url);
@@ -148,6 +156,18 @@ async fn add_server(
config = config.with_description(desc);
}
// Parse custom headers (format: "Key:Value")
if !headers.is_empty() {
let mut header_map = std::collections::HashMap::new();
for h in &headers {
let (key, value) = h.split_once(':').ok_or_else(|| {
anyhow::anyhow!("Invalid header format '{}'. Expected 'Key:Value'.", h)
})?;
header_map.insert(key.trim().to_string(), value.trim().to_string());
}
config = config.with_headers(header_map);
}
// Track if auth is required
let requires_auth = client_id.is_some();
@@ -242,6 +262,17 @@ async fn list_servers(verbose: bool) -> anyhow::Result<()> {
if let Some(ref desc) = server.description {
println!(" Description: {}", desc);
}
if !server.headers.is_empty() {
println!(
" Custom headers: {}",
server
.headers
.keys()
.cloned()
.collect::<Vec<_>>()
.join(", ")
);
}
if let Some(ref oauth) = server.oauth {
println!(" OAuth Client ID: {}", oauth.client_id);
if !oauth.scopes.is_empty() {
@@ -374,7 +405,7 @@ async fn test_server(name: String, user_id: String) -> anyhow::Result<()> {
return Ok(());
} else {
// No OAuth and no tokens - try unauthenticated
McpClient::new_with_name(&server.name, &server.url)
McpClient::new_with_config(server.clone())
};
// Test connection
+1 -1
View File
@@ -2508,7 +2508,7 @@ impl ExtensionManager {
&self.user_id,
)
} else {
McpClient::new_with_name(&server.name, &server.url)
McpClient::new_with_config(server.clone())
};
// Try to list and create tools
+23 -345
View File
@@ -1,53 +1,11 @@
//! Time utility tool.
use async_trait::async_trait;
use chrono::{DateTime, FixedOffset, Utc};
use chrono_tz::Tz;
use chrono::{DateTime, Utc};
use crate::context::JobContext;
use crate::tools::tool::{Tool, ToolError, ToolOutput, require_str};
/// Parse a timezone string into a `chrono_tz::Tz`, returning a clear error.
fn parse_timezone(tz_str: &str) -> Result<Tz, ToolError> {
tz_str.parse::<Tz>().map_err(|_| {
ToolError::InvalidParameters(format!(
"Unknown timezone '{}'. Use IANA names like 'America/New_York' or 'Europe/London'.",
tz_str
))
})
}
/// Parse an input timestamp string. Accepts RFC 3339 with offset, or naive
/// datetime in `YYYY-MM-DDTHH:MM:SS` / `YYYY-MM-DD HH:MM:SS` format
/// (interpreted as UTC unless `default_tz` is provided).
fn parse_input_timestamp(
input: &str,
default_tz: Option<Tz>,
) -> Result<DateTime<FixedOffset>, ToolError> {
// Try RFC 3339 first (has offset info)
if let Ok(dt) = DateTime::parse_from_rfc3339(input) {
return Ok(dt);
}
// Try common formats without offset — interpret in default_tz or UTC
for fmt in &["%Y-%m-%dT%H:%M:%S", "%Y-%m-%d %H:%M:%S"] {
if let Ok(naive) = chrono::NaiveDateTime::parse_from_str(input, fmt) {
let tz = default_tz.unwrap_or(Tz::UTC);
let local = naive.and_local_timezone(tz).single().ok_or_else(|| {
ToolError::InvalidParameters(format!(
"Ambiguous or invalid datetime '{}' in timezone '{}'",
input, tz
))
})?;
return Ok(local.fixed_offset());
}
}
Err(ToolError::InvalidParameters(format!(
"Invalid timestamp '{}'. Use RFC 3339 (e.g. '2026-03-07T12:00:00Z') \
or 'YYYY-MM-DD HH:MM:SS' format.",
input
)))
}
/// Tool for getting current time and date operations.
pub struct TimeTool;
@@ -58,7 +16,7 @@ impl Tool for TimeTool {
}
fn description(&self) -> &str {
"Get current time, convert timezones, format timestamps, or calculate time differences."
"Get current time, convert timezones, or calculate time differences."
}
fn parameters_schema(&self) -> serde_json::Value {
@@ -67,28 +25,20 @@ impl Tool for TimeTool {
"properties": {
"operation": {
"type": "string",
"enum": ["now", "parse", "convert", "format", "diff"],
"enum": ["now", "parse", "format", "diff"],
"description": "The time operation to perform"
},
"timestamp": {
"type": "string",
"description": "ISO 8601 timestamp (for parse/convert/format/diff operations)"
"description": "ISO 8601 timestamp (for parse/format/diff operations)"
},
"format": {
"type": "string",
"description": "Output format string (for format operation)"
},
"timestamp2": {
"type": "string",
"description": "Second timestamp (for diff operation)"
},
"timezone": {
"type": "string",
"description": "IANA timezone name, e.g. 'America/New_York' (for now/convert/format/parse)"
},
"to_timezone": {
"type": "string",
"description": "Target IANA timezone for convert operation"
},
"format_string": {
"type": "string",
"description": "strftime format string (for format operation), default: '%Y-%m-%d %H:%M:%S %Z'"
}
},
"required": ["operation"]
@@ -107,91 +57,36 @@ impl Tool for TimeTool {
let result = match operation {
"now" => {
let now = Utc::now();
let mut result = serde_json::json!({
"utc_iso": now.to_rfc3339(),
serde_json::json!({
"iso": now.to_rfc3339(),
"unix": now.timestamp(),
"unix_millis": now.timestamp_millis()
});
if let Some(tz_str) = params.get("timezone").and_then(|v| v.as_str()) {
let tz = parse_timezone(tz_str)?;
let local = now.with_timezone(&tz);
result["local_iso"] = serde_json::json!(local.to_rfc3339());
result["timezone"] = serde_json::json!(tz_str);
}
result
})
}
"parse" => {
let timestamp = require_str(&params, "timestamp")?;
let tz = params
.get("timezone")
.and_then(|v| v.as_str())
.map(parse_timezone)
.transpose()?;
let dt = parse_input_timestamp(timestamp, tz)?;
let utc = dt.with_timezone(&Utc);
let mut result = serde_json::json!({
"iso": utc.to_rfc3339(),
"unix": utc.timestamp(),
"unix_millis": utc.timestamp_millis()
});
if let Some(tz) = tz {
let local = dt.with_timezone(&tz);
result["local_iso"] = serde_json::json!(local.to_rfc3339());
result["timezone"] = serde_json::json!(tz.to_string());
}
result
}
"convert" => {
let timestamp = require_str(&params, "timestamp")?;
let to_tz_str = require_str(&params, "to_timezone")?;
let to_tz = parse_timezone(to_tz_str)?;
let from_tz = params
.get("timezone")
.and_then(|v| v.as_str())
.map(parse_timezone)
.transpose()?;
let dt = parse_input_timestamp(timestamp, from_tz)?;
let converted = dt.with_timezone(&to_tz);
let dt: DateTime<Utc> = timestamp.parse().map_err(|e| {
ToolError::InvalidParameters(format!("invalid timestamp: {}", e))
})?;
serde_json::json!({
"input": timestamp,
"output": converted.to_rfc3339(),
"timezone": to_tz.to_string()
"iso": dt.to_rfc3339(),
"unix": dt.timestamp(),
"unix_millis": dt.timestamp_millis()
})
}
"format" => {
let timestamp = require_str(&params, "timestamp")?;
let fmt = params
.get("format_string")
.and_then(|v| v.as_str())
.unwrap_or("%Y-%m-%d %H:%M:%S %Z");
let tz = params
.get("timezone")
.and_then(|v| v.as_str())
.map(parse_timezone)
.transpose()?;
let dt = parse_input_timestamp(timestamp, None)?;
let formatted = if let Some(tz) = tz {
dt.with_timezone(&tz).format(fmt).to_string()
} else {
dt.format(fmt).to_string()
};
serde_json::json!({ "formatted": formatted })
}
"diff" => {
let ts1 = require_str(&params, "timestamp")?;
let ts2 = require_str(&params, "timestamp2")?;
let dt1 = parse_input_timestamp(ts1, None)?;
let dt2 = parse_input_timestamp(ts2, None)?;
let dt1: DateTime<Utc> = ts1.parse().map_err(|e| {
ToolError::InvalidParameters(format!("invalid timestamp: {}", e))
})?;
let dt2: DateTime<Utc> = ts2.parse().map_err(|e| {
ToolError::InvalidParameters(format!("invalid timestamp2: {}", e))
})?;
let diff = dt2.signed_duration_since(dt1);
@@ -217,220 +112,3 @@ impl Tool for TimeTool {
false // Internal tool, no external data
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::context::JobContext;
use serde_json::json;
fn test_ctx() -> JobContext {
JobContext::new("test-job", "test time tool")
}
#[tokio::test]
async fn test_now_utc() {
let tool = TimeTool;
let result = tool
.execute(json!({"operation": "now"}), &test_ctx())
.await
.unwrap();
let v: serde_json::Value = result.result.clone();
assert!(v["utc_iso"].as_str().is_some());
assert!(v["iso"].as_str().is_some());
assert!(v["unix"].as_i64().is_some());
// No timezone requested — no local_iso
assert!(v.get("local_iso").is_none());
}
#[tokio::test]
async fn test_now_with_timezone() {
let tool = TimeTool;
let result = tool
.execute(
json!({"operation": "now", "timezone": "America/New_York"}),
&test_ctx(),
)
.await
.unwrap();
let v: serde_json::Value = result.result.clone();
assert!(v["local_iso"].as_str().is_some());
assert_eq!(v["timezone"].as_str().unwrap(), "America/New_York");
// local_iso should contain a non-UTC offset
let local = v["local_iso"].as_str().unwrap();
assert!(!local.ends_with('Z') || local.contains("-04:00") || local.contains("-05:00"));
}
#[tokio::test]
async fn test_now_invalid_timezone() {
let tool = TimeTool;
let result = tool
.execute(
json!({"operation": "now", "timezone": "Not/A/Zone"}),
&test_ctx(),
)
.await;
assert!(result.is_err());
let err = result.unwrap_err();
assert!(err.to_string().contains("Unknown timezone"));
assert!(err.to_string().contains("Not/A/Zone"));
}
#[tokio::test]
async fn test_convert_timezone() {
let tool = TimeTool;
let result = tool
.execute(
json!({
"operation": "convert",
"timestamp": "2026-03-07T12:00:00Z",
"to_timezone": "Asia/Tokyo"
}),
&test_ctx(),
)
.await
.unwrap();
let v: serde_json::Value = result.result.clone();
// UTC 12:00 -> JST 21:00 (UTC+9)
let output = v["output"].as_str().unwrap();
assert!(output.contains("21:00:00"));
assert_eq!(v["timezone"].as_str().unwrap(), "Asia/Tokyo");
}
#[tokio::test]
async fn test_convert_dst_boundary() {
let tool = TimeTool;
// US spring forward: 2026-03-08 2:00 AM EST -> 3:00 AM EDT
// Before DST: EST = UTC-5, After: EDT = UTC-4
let result = tool
.execute(
json!({
"operation": "convert",
"timestamp": "2026-03-08T06:30:00Z",
"to_timezone": "America/New_York"
}),
&test_ctx(),
)
.await
.unwrap();
let v: serde_json::Value = result.result.clone();
// UTC 06:30 on Mar 8 -> after spring forward, EDT (UTC-4) = 02:30
// But DST springs forward at 2 AM -> 3 AM, so 06:30 UTC = 01:30 EST or 02:30 EDT
let output = v["output"].as_str().unwrap();
assert!(output.contains("2026-03-08"));
}
#[tokio::test]
async fn test_format_with_timezone() {
let tool = TimeTool;
let result = tool
.execute(
json!({
"operation": "format",
"timestamp": "2026-03-07T12:00:00Z",
"timezone": "Europe/London",
"format_string": "%Y-%m-%d %H:%M %Z"
}),
&test_ctx(),
)
.await
.unwrap();
let v: serde_json::Value = result.result.clone();
let formatted = v["formatted"].as_str().unwrap();
assert!(formatted.contains("2026-03-07"));
assert!(formatted.contains("12:00")); // London = UTC in March (before DST)
assert!(formatted.contains("GMT"));
}
#[tokio::test]
async fn test_format_default_format_string() {
let tool = TimeTool;
let result = tool
.execute(
json!({
"operation": "format",
"timestamp": "2026-06-15T18:30:00Z",
"timezone": "America/Los_Angeles"
}),
&test_ctx(),
)
.await
.unwrap();
let v: serde_json::Value = result.result.clone();
let formatted = v["formatted"].as_str().unwrap();
// UTC 18:30 -> PDT (UTC-7) = 11:30
assert!(formatted.contains("11:30:00"));
assert!(formatted.contains("PDT"));
}
#[tokio::test]
async fn test_parse_naive_with_timezone() {
let tool = TimeTool;
let result = tool
.execute(
json!({
"operation": "parse",
"timestamp": "2026-03-07 09:00:00",
"timezone": "America/New_York"
}),
&test_ctx(),
)
.await
.unwrap();
let v: serde_json::Value = result.result.clone();
// 09:00 EST = 14:00 UTC (EST = UTC-5 in March before DST)
let iso = v["iso"].as_str().unwrap();
assert!(iso.contains("14:00:00"));
assert_eq!(v["timezone"].as_str().unwrap(), "America/New_York");
}
#[tokio::test]
async fn test_diff() {
let tool = TimeTool;
let result = tool
.execute(
json!({
"operation": "diff",
"timestamp": "2026-03-07T00:00:00Z",
"timestamp2": "2026-03-07T02:30:00Z"
}),
&test_ctx(),
)
.await
.unwrap();
let v: serde_json::Value = result.result.clone();
assert_eq!(v["hours"].as_i64().unwrap(), 2);
assert_eq!(v["minutes"].as_i64().unwrap(), 150);
assert_eq!(v["seconds"].as_i64().unwrap(), 9000);
}
#[tokio::test]
async fn test_convert_missing_to_timezone() {
let tool = TimeTool;
let result = tool
.execute(
json!({
"operation": "convert",
"timestamp": "2026-03-07T12:00:00Z"
}),
&test_ctx(),
)
.await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_unknown_operation() {
let tool = TimeTool;
let result = tool
.execute(json!({"operation": "explode"}), &test_ctx())
.await;
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("unknown operation")
);
}
}
+55 -1
View File
@@ -3,6 +3,7 @@
//! Supports both local (unauthenticated) and hosted (OAuth-authenticated) servers.
//! Uses the Streamable HTTP transport with session management.
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
@@ -52,6 +53,9 @@ pub struct McpClient {
/// Server configuration (for token secret name lookup).
server_config: Option<McpServerConfig>,
/// Custom HTTP headers injected into every request.
custom_headers: HashMap<String, String>,
}
impl McpClient {
@@ -75,6 +79,7 @@ impl McpClient {
secrets: None,
user_id: "default".to_string(),
server_config: None,
custom_headers: HashMap::new(),
}
}
@@ -95,6 +100,28 @@ impl McpClient {
secrets: None,
user_id: "default".to_string(),
server_config: None,
custom_headers: HashMap::new(),
}
}
/// Create a new simple MCP client from a server configuration (no authentication).
///
/// Use this when you have an `McpServerConfig` with custom headers but no OAuth.
pub fn new_with_config(config: McpServerConfig) -> Self {
Self {
server_name: config.name.clone(),
server_url: config.url.clone(),
http_client: reqwest::Client::builder()
.timeout(Duration::from_secs(30))
.build()
.expect("Failed to create HTTP client"),
next_id: AtomicU64::new(1),
tools_cache: RwLock::new(None),
session_manager: None,
secrets: None,
user_id: "default".to_string(),
custom_headers: config.headers.clone(),
server_config: Some(config),
}
}
@@ -119,6 +146,7 @@ impl McpClient {
session_manager: Some(session_manager),
secrets: Some(secrets),
user_id: user_id.into(),
custom_headers: config.headers.clone(),
server_config: Some(config),
}
}
@@ -178,7 +206,12 @@ impl McpClient {
.header("Content-Type", "application/json")
.json(&request);
// Add Authorization header if we have a token
// Add custom headers from config
for (key, value) in &self.custom_headers {
req_builder = req_builder.header(key, value);
}
// Add Authorization header if we have a token (overrides custom Authorization)
if let Some(token) = self.get_access_token().await? {
req_builder = req_builder.header("Authorization", format!("Bearer {}", token));
}
@@ -474,6 +507,7 @@ impl Clone for McpClient {
secrets: self.secrets.clone(),
user_id: self.user_id.clone(),
server_config: self.server_config.clone(),
custom_headers: self.custom_headers.clone(),
}
}
}
@@ -692,6 +726,26 @@ mod tests {
assert_eq!(id3, 3);
}
#[test]
fn test_custom_headers_from_config() {
use std::collections::HashMap;
let mut headers = HashMap::new();
headers.insert("X-API-Key".to_string(), "secret".to_string());
headers.insert("X-Custom".to_string(), "value".to_string());
let config = McpServerConfig::new("test", "http://localhost:8080").with_headers(headers);
let client = McpClient::new_with_config(config);
assert_eq!(client.custom_headers.len(), 2);
assert_eq!(client.custom_headers.get("X-API-Key").unwrap(), "secret");
}
#[test]
fn test_new_has_no_custom_headers() {
let client = McpClient::new("http://localhost:8080");
assert!(client.custom_headers.is_empty());
}
#[test]
fn test_mcp_tool_requires_approval_destructive() {
use crate::tools::mcp::protocol::{McpTool, McpToolAnnotations};
+89
View File
@@ -32,6 +32,13 @@ pub struct McpServerConfig {
/// Optional description for the server.
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
/// Custom HTTP headers to send with every request to this server.
///
/// Useful for MCP servers that require non-OAuth authentication
/// (e.g., API keys via `X-API-Key` header).
#[serde(default, skip_serializing_if = "HashMap::is_empty")]
pub headers: HashMap<String, String>,
}
fn default_true() -> bool {
@@ -47,9 +54,16 @@ impl McpServerConfig {
oauth: None,
enabled: true,
description: None,
headers: HashMap::new(),
}
}
/// Set custom HTTP headers for this server.
pub fn with_headers(mut self, headers: HashMap<String, String>) -> Self {
self.headers = headers;
self
}
/// Set OAuth configuration.
pub fn with_oauth(mut self, oauth: OAuthConfig) -> Self {
self.oauth = Some(oauth);
@@ -593,4 +607,79 @@ mod tests {
let config = McpServerConfig::new("bad", "http://mcp.example.com");
assert!(!config.requires_auth());
}
#[test]
fn test_custom_headers_default_empty() {
let config = McpServerConfig::new("test", "http://localhost:8080");
assert!(config.headers.is_empty());
}
#[test]
fn test_custom_headers_with_builder() {
let mut headers = HashMap::new();
headers.insert("X-API-Key".to_string(), "secret123".to_string());
headers.insert("X-Custom".to_string(), "value".to_string());
let config = McpServerConfig::new("browser-use", "https://mcp.browser-use.com")
.with_headers(headers.clone());
assert_eq!(config.headers.len(), 2);
assert_eq!(config.headers.get("X-API-Key").unwrap(), "secret123");
}
#[test]
fn test_custom_headers_serde_roundtrip() {
let mut headers = HashMap::new();
headers.insert("Authorization".to_string(), "Bearer tok_123".to_string());
let config =
McpServerConfig::new("test-serde", "http://localhost:3000").with_headers(headers);
let json = serde_json::to_string(&config).unwrap();
let deserialized: McpServerConfig = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized.headers.len(), 1);
assert_eq!(
deserialized.headers.get("Authorization").unwrap(),
"Bearer tok_123"
);
}
#[test]
fn test_custom_headers_absent_in_json_defaults_empty() {
let json = serde_json::json!({
"name": "legacy",
"url": "http://localhost:8080"
});
let config: McpServerConfig = serde_json::from_value(json).unwrap();
assert!(config.headers.is_empty());
}
#[test]
fn test_custom_headers_skipped_when_empty_in_serialization() {
let config = McpServerConfig::new("minimal", "http://localhost:8080");
let json = serde_json::to_value(&config).unwrap();
// Empty headers map should not appear in serialized output
assert!(json.get("headers").is_none());
}
#[tokio::test]
async fn test_custom_headers_persist_to_disk() {
let dir = tempdir().unwrap();
let path = dir.path().join("mcp-headers-test.json");
let mut headers = HashMap::new();
headers.insert("X-API-Key".to_string(), "key123".to_string());
let mut config = McpServersFile::default();
config.upsert(
McpServerConfig::new("headered", "http://localhost:9090").with_headers(headers),
);
save_mcp_servers_to(&config, &path).await.unwrap();
let loaded = load_mcp_servers_from(&path).await.unwrap();
let server = loaded.get("headered").unwrap();
assert_eq!(server.headers.get("X-API-Key").unwrap(), "key123");
}
}