mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-25 14:53:34 +00:00
Compare commits
9
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bd77b8b14b | ||
|
|
2b625ef3df | ||
|
|
7d745d5479 | ||
|
|
1bc10fe4ca | ||
|
|
f53c1bb10b | ||
|
|
bc6725205a | ||
|
|
275bcfb658 | ||
|
|
7776d267f8 | ||
|
|
e805ec61aa |
@@ -78,15 +78,55 @@ jobs:
|
||||
- name: Check lints
|
||||
run: cargo clippy --all --benches --tests --examples ${{ matrix.flags }} -- -D warnings
|
||||
|
||||
no-panics:
|
||||
name: No panics in production code
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@v6
|
||||
with:
|
||||
fetch-depth: 0
|
||||
- name: Check for .unwrap(), .expect(), assert!() in production code
|
||||
run: |
|
||||
BASE="${{ github.event.pull_request.base.sha }}"
|
||||
# Get added lines in .rs files (production only, exclude tests/)
|
||||
ADDED=$(git diff "$BASE"...HEAD -- 'src/**/*.rs' 'crates/**/*.rs' \
|
||||
| grep -E '^\+[^+]' || true)
|
||||
|
||||
if [ -z "$ADDED" ]; then
|
||||
echo "No production Rust changes detected."
|
||||
exit 0
|
||||
fi
|
||||
|
||||
# Match panic-inducing patterns, excluding test code and safety suppressions
|
||||
VIOLATIONS=$(echo "$ADDED" \
|
||||
| grep -E '\.(unwrap|expect)\(|[^_]assert(_eq|_ne)?!' \
|
||||
| grep -Ev 'debug_assert|// safety:|#\[cfg\(test\)\]|#\[test\]|mod tests' \
|
||||
|| true)
|
||||
|
||||
if [ -n "$VIOLATIONS" ]; then
|
||||
echo "::error::Found .unwrap(), .expect(), or assert!() in production code."
|
||||
echo "Production code must use proper error handling instead of panicking."
|
||||
echo "Suppress false positives with an inline '// safety: <reason>' comment."
|
||||
echo ""
|
||||
echo "$VIOLATIONS" | head -20
|
||||
echo ""
|
||||
COUNT=$(echo "$VIOLATIONS" | wc -l | tr -d ' ')
|
||||
echo "Total: $COUNT violation(s)"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
echo "OK: No panic-inducing calls in changed production code."
|
||||
|
||||
# Roll-up job for branch protection
|
||||
code-style:
|
||||
name: Code Style (fmt + clippy + deny)
|
||||
runs-on: ubuntu-latest
|
||||
if: always()
|
||||
needs: [format, clippy, clippy-windows, deny-check]
|
||||
needs: [format, clippy, clippy-windows, deny-check, no-panics]
|
||||
steps:
|
||||
- run: |
|
||||
if [[ "${{ needs.format.result }}" != "success" || "${{ needs.clippy.result }}" != "success" || "${{ needs.deny-check.result }}" != "success" ]]; then
|
||||
if [[ "${{ needs.format.result }}" != "success" || "${{ needs.clippy.result }}" != "success" || "${{ needs.deny-check.result }}" != "success" || "${{ needs.no-panics.result }}" != "success" ]]; then
|
||||
echo "One or more jobs failed"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
@@ -14,6 +14,10 @@
|
||||
|
||||
target/
|
||||
|
||||
# Python
|
||||
__pycache__/
|
||||
*.pyc
|
||||
|
||||
# Benchmark results (local runs, not committed)
|
||||
bench-results/
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "discord",
|
||||
"display_name": "Discord Channel",
|
||||
"kind": "channel",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Talk to your agent in Discord",
|
||||
"keywords": [
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "github",
|
||||
"display_name": "GitHub",
|
||||
"kind": "tool",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "GitHub integration for issues, PRs, repos, and code search",
|
||||
"keywords": [
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
"name": "web-search",
|
||||
"display_name": "Web Search",
|
||||
"kind": "tool",
|
||||
"version": "0.2.0",
|
||||
"version": "0.2.1",
|
||||
"wit_version": "0.3.0",
|
||||
"description": "Search the web using Brave Search API",
|
||||
"keywords": [
|
||||
|
||||
@@ -70,19 +70,21 @@ echo
|
||||
# This is a WARNING, not a hard violation.
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
echo "--- Check 2: .unwrap() / .expect() in production code ---"
|
||||
echo "--- Check 2: .unwrap() / .expect() / assert!() in production code ---"
|
||||
|
||||
# Collect raw matches excluding obvious test-only files and lines
|
||||
raw_results=$(grep -rn '\.unwrap()\|\.expect(' src/ \
|
||||
# Collect raw matches excluding obvious test-only files and lines.
|
||||
# Also catches assert!(), assert_eq!(), assert_ne!() but NOT debug_assert variants.
|
||||
raw_results=$(grep -rnE '\.(unwrap|expect)\(|[^_]assert(_eq|_ne)?!' src/ \
|
||||
--include='*.rs' \
|
||||
| grep -v 'src/main.rs' \
|
||||
| grep -v 'src/testing.rs' \
|
||||
| grep -v 'src/setup/' \
|
||||
| grep -Ev 'debug_assert|// safety:' \
|
||||
|| true)
|
||||
|
||||
if [ -n "$raw_results" ]; then
|
||||
total=$(echo "$raw_results" | wc -l | tr -d ' ')
|
||||
echo "WARNING: ~$total .unwrap()/.expect() calls found in src/ (excluding main/testing/setup)."
|
||||
echo "WARNING: ~$total .unwrap()/.expect()/assert!() calls found in src/ (excluding main/testing/setup)."
|
||||
echo "Many are in test modules; a per-file breakdown helps triage:"
|
||||
echo
|
||||
# Show per-file counts, sorted by count descending, top 15
|
||||
|
||||
@@ -10,6 +10,7 @@
|
||||
# 3. Hardcoded /tmp paths in tests (flaky in parallel runs)
|
||||
# 4. Tool parameters logged without redaction (secret leaks)
|
||||
# 5. Multi-step DB operations without transaction wrapping
|
||||
# 6. .unwrap(), .expect(), assert!() in production code (panics)
|
||||
#
|
||||
# Suppress individual lines with an inline "// safety: <reason>" comment.
|
||||
|
||||
@@ -128,6 +129,24 @@ if [ -n "$DIFF_W_OUTPUT" ]; then
|
||||
fi
|
||||
fi
|
||||
|
||||
# 6. .unwrap(), .expect(), assert!() in production code
|
||||
# Matches added lines containing panic-inducing calls.
|
||||
# Excludes test files, test modules, and debug_assert (compiled out in release).
|
||||
# Suppress with "// safety: <reason>".
|
||||
PROD_DIFF="$DIFF_OUTPUT"
|
||||
# Strip hunks from test-only files (tests/ directory, *_test.rs, test_*.rs)
|
||||
PROD_DIFF=$(echo "$PROD_DIFF" | grep -v '^+++ b/tests/' || true)
|
||||
if echo "$PROD_DIFF" | grep -nE '^\+' \
|
||||
| grep -E '\.(unwrap|expect)\(|[^_]assert(_eq|_ne)?!' \
|
||||
| grep -vE 'debug_assert|// safety:|#\[cfg\(test\)\]|#\[test\]|mod tests' \
|
||||
| head -5 | grep -q .; then
|
||||
warn "PANIC" "Production code must not use .unwrap(), .expect(), or assert!(). Use proper error handling."
|
||||
echo "$PROD_DIFF" | grep -nE '^\+' \
|
||||
| grep -E '\.(unwrap|expect)\(|[^_]assert(_eq|_ne)?!' \
|
||||
| grep -vE 'debug_assert|// safety:|#\[cfg\(test\)\]|#\[test\]|mod tests' \
|
||||
| head -5 | sed 's/^/ /'
|
||||
fi
|
||||
|
||||
if [ "$WARNINGS" -gt 0 ]; then
|
||||
echo ""
|
||||
echo "Found $WARNINGS potential issue(s). Fix them or add '// safety: <reason>' to suppress."
|
||||
|
||||
+31
-12
@@ -93,19 +93,26 @@ impl RoutineEngine {
|
||||
let mut cache = Vec::new();
|
||||
for routine in routines {
|
||||
match &routine.trigger {
|
||||
Trigger::Event { pattern, .. } => match Regex::new(pattern) {
|
||||
Ok(re) => cache.push(EventMatcher::Message {
|
||||
routine: routine.clone(),
|
||||
regex: re,
|
||||
}),
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
routine = %routine.name,
|
||||
"Invalid event regex '{}': {}",
|
||||
pattern, e
|
||||
);
|
||||
Trigger::Event { pattern, .. } => {
|
||||
// Use RegexBuilder with size limit to prevent ReDoS
|
||||
// from user-supplied patterns (issue #825).
|
||||
match regex::RegexBuilder::new(pattern)
|
||||
.size_limit(64 * 1024) // 64KB compiled size limit
|
||||
.build()
|
||||
{
|
||||
Ok(re) => cache.push(EventMatcher::Message {
|
||||
routine: routine.clone(),
|
||||
regex: re,
|
||||
}),
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
routine = %routine.name,
|
||||
"Invalid or too complex event regex '{}': {}",
|
||||
pattern, e
|
||||
);
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
Trigger::SystemEvent { .. } => {
|
||||
cache.push(EventMatcher::System {
|
||||
routine: routine.clone(),
|
||||
@@ -973,6 +980,18 @@ async fn execute_lightweight_with_tools(
|
||||
}
|
||||
};
|
||||
|
||||
// Truncate oversized tool output to prevent unbounded context growth.
|
||||
// Routine tool loops are lightweight and should not accumulate
|
||||
// large payloads across iterations.
|
||||
const MAX_TOOL_OUTPUT_CHARS: usize = 8192;
|
||||
let result_content = if result_content.len() > MAX_TOOL_OUTPUT_CHARS {
|
||||
let truncated = &result_content
|
||||
[..result_content.floor_char_boundary(MAX_TOOL_OUTPUT_CHARS)];
|
||||
format!("{truncated}\n... [output truncated to {MAX_TOOL_OUTPUT_CHARS} chars]")
|
||||
} else {
|
||||
result_content
|
||||
};
|
||||
|
||||
// Add tool result to context
|
||||
messages.push(ChatMessage::tool_result(&tc.id, &tc.name, &result_content));
|
||||
}
|
||||
|
||||
+48
-15
@@ -269,21 +269,24 @@ async fn webhook_handler(
|
||||
let mut fallback_req = None;
|
||||
{
|
||||
let webhook_secret = state.webhook_secret.read().await;
|
||||
let Some(expected_secret) = webhook_secret.as_ref() else {
|
||||
return (
|
||||
StatusCode::UNAUTHORIZED,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some(
|
||||
"Webhook authentication required: HTTP webhook secret is not configured."
|
||||
.to_string(),
|
||||
),
|
||||
}),
|
||||
)
|
||||
.into_response();
|
||||
let expected_secret = match webhook_secret.as_ref() {
|
||||
Some(secret) => secret.expose_secret(),
|
||||
None => {
|
||||
// No secret configured — reject all requests. This guards against
|
||||
// the secret being cleared at runtime via update_secret(None).
|
||||
// The start() method also prevents startup without a secret, but
|
||||
// this is defense-in-depth for the SIGHUP hot-swap path.
|
||||
return (
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
Json(WebhookResponse {
|
||||
message_id: Uuid::nil(),
|
||||
status: "error".to_string(),
|
||||
response: Some("Webhook authentication not configured".to_string()),
|
||||
}),
|
||||
)
|
||||
.into_response();
|
||||
}
|
||||
};
|
||||
let expected_secret = expected_secret.expose_secret();
|
||||
|
||||
match headers.get("x-ironclaw-signature") {
|
||||
Some(raw_signature) => match raw_signature.to_str() {
|
||||
@@ -1085,7 +1088,7 @@ mod tests {
|
||||
.unwrap();
|
||||
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
|
||||
assert_eq!(resp.status(), StatusCode::SERVICE_UNAVAILABLE); // safety: test assertion
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -1206,4 +1209,34 @@ mod tests {
|
||||
let body = b"test body content";
|
||||
assert!(!verify_hmac_signature(secret, body, "sha256=not-hex!"));
|
||||
}
|
||||
|
||||
/// Regression test for issue #1033: when the webhook secret is cleared at
|
||||
/// runtime via update_secret(None), subsequent requests must be rejected
|
||||
/// instead of being processed without authentication.
|
||||
#[tokio::test]
|
||||
async fn webhook_rejects_when_secret_cleared_at_runtime() {
|
||||
let channel = test_channel(Some("initial-secret"));
|
||||
let _stream = channel.start().await.unwrap();
|
||||
|
||||
// Clear the secret at runtime (simulates a bad SIGHUP config reload)
|
||||
channel.update_secret(None).await;
|
||||
|
||||
let app = channel.routes();
|
||||
let body = serde_json::json!({
|
||||
"content": "hello"
|
||||
});
|
||||
let req = Request::builder()
|
||||
.method("POST")
|
||||
.uri("/webhook")
|
||||
.header("content-type", "application/json")
|
||||
.body(Body::from(serde_json::to_vec(&body).unwrap()))
|
||||
.unwrap();
|
||||
|
||||
let resp = app.oneshot(req).await.unwrap();
|
||||
assert_eq!(
|
||||
resp.status(),
|
||||
StatusCode::SERVICE_UNAVAILABLE,
|
||||
"requests must be rejected when webhook secret is cleared at runtime"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -294,6 +294,8 @@ impl Channel for RelayChannel {
|
||||
match client.connect_stream(&token, stream_timeout_secs).await {
|
||||
Ok((new_stream, new_parser)) => {
|
||||
tracing::info!("Relay SSE stream reconnected");
|
||||
consecutive_failures = 0;
|
||||
backoff_ms = backoff_initial_ms;
|
||||
current_stream = new_stream;
|
||||
// Abort old parser before replacing
|
||||
if let Some(old) = parser_handle.write().await.take() {
|
||||
@@ -312,6 +314,8 @@ impl Channel for RelayChannel {
|
||||
tracing::info!(
|
||||
"Relay SSE stream reconnected with new token"
|
||||
);
|
||||
consecutive_failures = 0;
|
||||
backoff_ms = backoff_initial_ms;
|
||||
current_stream = new_stream;
|
||||
if let Some(old) = parser_handle.write().await.take() {
|
||||
old.abort();
|
||||
|
||||
@@ -190,12 +190,21 @@ pub async fn routines_toggle_handler(
|
||||
None => !routine.enabled,
|
||||
};
|
||||
|
||||
// When re-enabling a cron routine, recompute next_fire_at so the cron
|
||||
// ticker can pick it up. Mirrors the CLI behavior (issue #1077).
|
||||
if routine.enabled
|
||||
&& !was_enabled
|
||||
&& let Trigger::Cron { schedule, timezone } = &routine.trigger
|
||||
&& let Trigger::Cron {
|
||||
ref schedule,
|
||||
ref timezone,
|
||||
} = routine.trigger
|
||||
{
|
||||
routine.next_fire_at = next_cron_fire(schedule, timezone.as_deref())
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
routine.next_fire_at = next_cron_fire(schedule, timezone.as_deref()).map_err(|e| {
|
||||
(
|
||||
StatusCode::INTERNAL_SERVER_ERROR,
|
||||
format!("Failed to compute next fire: {e}"),
|
||||
)
|
||||
})?;
|
||||
}
|
||||
|
||||
store
|
||||
@@ -203,6 +212,12 @@ pub async fn routines_toggle_handler(
|
||||
.await
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
// Refresh the in-memory event trigger cache so event/system_event
|
||||
// routines reflect the new enabled state immediately (issue #1076).
|
||||
if let Some(engine) = state.routine_engine.read().await.as_ref() {
|
||||
engine.refresh_event_cache().await;
|
||||
}
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"status": if routine.enabled { "enabled" } else { "disabled" },
|
||||
"routine_id": routine_id,
|
||||
@@ -227,6 +242,12 @@ pub async fn routines_delete_handler(
|
||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||
|
||||
if deleted {
|
||||
// Refresh the in-memory event trigger cache so deleted event/system_event
|
||||
// routines stop firing immediately (issue #1076).
|
||||
if let Some(engine) = state.routine_engine.read().await.as_ref() {
|
||||
engine.refresh_event_cache().await;
|
||||
}
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"status": "deleted",
|
||||
"routine_id": routine_id,
|
||||
|
||||
@@ -214,6 +214,7 @@ fn is_disallowed_ipv4(v4: &Ipv4Addr) -> bool {
|
||||
|| v4.is_multicast()
|
||||
|| v4.is_unspecified()
|
||||
|| *v4 == Ipv4Addr::new(169, 254, 169, 254)
|
||||
|| (v4.octets()[0] == 100 && (v4.octets()[1] & 0xC0) == 64)
|
||||
}
|
||||
|
||||
fn is_disallowed_ip(ip: &IpAddr) -> bool {
|
||||
@@ -913,6 +914,8 @@ mod tests {
|
||||
assert!(is_disallowed_ip(&IpAddr::V4(Ipv4Addr::new(
|
||||
169, 254, 169, 254
|
||||
))));
|
||||
// Carrier-grade NAT
|
||||
assert!(is_disallowed_ip(&IpAddr::V4(Ipv4Addr::new(100, 64, 0, 1))));
|
||||
// Public
|
||||
assert!(!is_disallowed_ip(&IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8))));
|
||||
}
|
||||
|
||||
+252
-121
@@ -24,6 +24,132 @@ use crate::context::JobContext;
|
||||
use crate::db::Database;
|
||||
use crate::tools::tool::{ApprovalRequirement, Tool, ToolError, ToolOutput, require_str};
|
||||
|
||||
pub(crate) fn routine_create_parameters_schema() -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {
|
||||
"type": "string",
|
||||
"description": "Unique routine name, for example 'daily-pr-review'."
|
||||
},
|
||||
"description": {
|
||||
"type": "string",
|
||||
"description": "Short summary of what the routine is for."
|
||||
},
|
||||
"trigger_type": {
|
||||
"type": "string",
|
||||
"enum": ["cron", "event", "system_event", "manual"],
|
||||
"description": "When the routine fires: 'cron' for schedules, 'event' for incoming messages, 'system_event' for structured emitted events, or 'manual' for explicit runs."
|
||||
},
|
||||
"schedule": {
|
||||
"type": "string",
|
||||
"description": "Cron schedule for 'cron' triggers. Uses 6 fields: second minute hour day month weekday."
|
||||
},
|
||||
"event_pattern": {
|
||||
"type": "string",
|
||||
"description": "Regex matched against incoming message text for 'event' triggers, for example '^bug\\\\b'."
|
||||
},
|
||||
"event_channel": {
|
||||
"type": "string",
|
||||
"description": "Optional platform filter for 'event' triggers, for example 'telegram'. Omit to match any channel. Not a chat or thread ID."
|
||||
},
|
||||
"event_source": {
|
||||
"type": "string",
|
||||
"description": "Structured event source for 'system_event' triggers, for example 'github'."
|
||||
},
|
||||
"event_type": {
|
||||
"type": "string",
|
||||
"description": "Structured event type for 'system_event' triggers, for example 'issue.opened'."
|
||||
},
|
||||
"event_filters": {
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"additionalProperties": {
|
||||
"type": ["string", "number", "boolean"]
|
||||
},
|
||||
"description": "Optional exact-match payload filters for 'system_event' triggers. Values can be strings, numbers, or booleans."
|
||||
},
|
||||
"prompt": {
|
||||
"type": "string",
|
||||
"description": "Instructions for what the routine should do after it fires."
|
||||
},
|
||||
"context_paths": {
|
||||
"type": "array",
|
||||
"items": { "type": "string" },
|
||||
"description": "Workspace paths to load as extra context before running the routine."
|
||||
},
|
||||
"action_type": {
|
||||
"type": "string",
|
||||
"enum": ["lightweight", "full_job"],
|
||||
"description": "Execution mode: 'lightweight' for one LLM turn or 'full_job' for a multi-step job with tools."
|
||||
},
|
||||
"use_tools": {
|
||||
"type": "boolean",
|
||||
"description": "Enable safe tool use in 'lightweight' mode. Ignored for 'full_job'."
|
||||
},
|
||||
"max_tool_rounds": {
|
||||
"type": "integer",
|
||||
"description": "Maximum tool-call rounds in 'lightweight' mode when 'use_tools' is true."
|
||||
},
|
||||
"cooldown_secs": {
|
||||
"type": "integer",
|
||||
"description": "Minimum seconds between fires."
|
||||
},
|
||||
"tool_permissions": {
|
||||
"type": "array",
|
||||
"items": { "type": "string" },
|
||||
"description": "Pre-authorized tool names for 'full_job' routines."
|
||||
},
|
||||
"notify_channel": {
|
||||
"type": "string",
|
||||
"description": "Where routine output should be sent, for example 'telegram' or 'slack'. This does not control what triggers the routine."
|
||||
},
|
||||
"notify_user": {
|
||||
"type": "string",
|
||||
"description": "User or destination to notify, for example a username or chat ID."
|
||||
},
|
||||
"timezone": {
|
||||
"type": "string",
|
||||
"description": "IANA timezone used to evaluate 'cron' schedules, for example 'America/New_York'."
|
||||
}
|
||||
},
|
||||
"required": ["name", "trigger_type", "prompt"]
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn routine_update_parameters_schema() -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {
|
||||
"type": "string",
|
||||
"description": "Name of the routine to update."
|
||||
},
|
||||
"enabled": {
|
||||
"type": "boolean",
|
||||
"description": "Set to true to enable the routine or false to disable it."
|
||||
},
|
||||
"prompt": {
|
||||
"type": "string",
|
||||
"description": "Replace the routine instructions for what it should do after it fires."
|
||||
},
|
||||
"schedule": {
|
||||
"type": "string",
|
||||
"description": "New cron schedule for existing 'cron' routines only. This does not convert other trigger types."
|
||||
},
|
||||
"timezone": {
|
||||
"type": "string",
|
||||
"description": "New IANA timezone for existing 'cron' routines only, for example 'America/New_York'."
|
||||
},
|
||||
"description": {
|
||||
"type": "string",
|
||||
"description": "Replace the routine summary."
|
||||
}
|
||||
},
|
||||
"required": ["name"]
|
||||
})
|
||||
}
|
||||
|
||||
// ==================== routine_create ====================
|
||||
|
||||
pub struct RoutineCreateTool {
|
||||
@@ -50,92 +176,7 @@ impl Tool for RoutineCreateTool {
|
||||
}
|
||||
|
||||
fn parameters_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {
|
||||
"type": "string",
|
||||
"description": "Unique name for the routine (e.g. 'daily-pr-review')"
|
||||
},
|
||||
"description": {
|
||||
"type": "string",
|
||||
"description": "What this routine does"
|
||||
},
|
||||
"trigger_type": {
|
||||
"type": "string",
|
||||
"enum": ["cron", "event", "system_event", "manual"],
|
||||
"description": "When the routine fires"
|
||||
},
|
||||
"schedule": {
|
||||
"type": "string",
|
||||
"description": "Cron expression (for cron trigger). E.g. '0 9 * * MON-FRI' for weekdays at 9am. Uses 6-field cron (sec min hour day month weekday)."
|
||||
},
|
||||
"event_pattern": {
|
||||
"type": "string",
|
||||
"description": "Regex pattern to match messages (for event trigger)"
|
||||
},
|
||||
"event_channel": {
|
||||
"type": "string",
|
||||
"description": "Optional channel filter for event trigger (e.g. 'telegram')"
|
||||
},
|
||||
"event_source": {
|
||||
"type": "string",
|
||||
"description": "Event source for system_event triggers (e.g. 'github')"
|
||||
},
|
||||
"event_type": {
|
||||
"type": "string",
|
||||
"description": "Event type for system_event triggers (e.g. 'issue.opened')"
|
||||
},
|
||||
"event_filters": {
|
||||
"type": "object",
|
||||
"description": "Optional exact-match filters against payload fields for system_event triggers. Values can be strings, numbers, or booleans."
|
||||
},
|
||||
"prompt": {
|
||||
"type": "string",
|
||||
"description": "The prompt/instructions for the routine"
|
||||
},
|
||||
"context_paths": {
|
||||
"type": "array",
|
||||
"items": { "type": "string" },
|
||||
"description": "Workspace paths to load as context (e.g. ['context/priorities.md'])"
|
||||
},
|
||||
"action_type": {
|
||||
"type": "string",
|
||||
"enum": ["lightweight", "full_job"],
|
||||
"description": "Execution mode: 'lightweight' (single LLM call, default) or 'full_job' (multi-turn with tools)"
|
||||
},
|
||||
"use_tools": {
|
||||
"type": "boolean",
|
||||
"description": "Enable tool access in lightweight mode (default: false). Only safe tools (no approval required) are available. Ignored for full_job mode."
|
||||
},
|
||||
"max_tool_rounds": {
|
||||
"type": "integer",
|
||||
"description": "Max tool call rounds in lightweight mode (default: 3). Only used when use_tools is true."
|
||||
},
|
||||
"cooldown_secs": {
|
||||
"type": "integer",
|
||||
"description": "Minimum seconds between fires (default: 300)"
|
||||
},
|
||||
"tool_permissions": {
|
||||
"type": "array",
|
||||
"items": { "type": "string" },
|
||||
"description": "Tool names pre-authorized for Always-approval tools in full_job mode (e.g. ['shell']). UnlessAutoApproved tools are automatically permitted in routines."
|
||||
},
|
||||
"notify_channel": {
|
||||
"type": "string",
|
||||
"description": "Channel to send results to (e.g. 'telegram', 'slack', 'tui'). Sets the default channel for message tool calls in routine jobs."
|
||||
},
|
||||
"notify_user": {
|
||||
"type": "string",
|
||||
"description": "User/target to notify (e.g. username, chat ID). Defaults to 'default'."
|
||||
},
|
||||
"timezone": {
|
||||
"type": "string",
|
||||
"description": "IANA timezone for cron schedule evaluation (e.g. 'America/New_York'). Defaults to UTC."
|
||||
}
|
||||
},
|
||||
"required": ["name", "trigger_type", "prompt"]
|
||||
})
|
||||
routine_create_parameters_schema()
|
||||
}
|
||||
|
||||
async fn execute(
|
||||
@@ -199,9 +240,13 @@ impl Tool for RoutineCreateTool {
|
||||
"event trigger requires 'event_pattern'".to_string(),
|
||||
)
|
||||
})?;
|
||||
// Validate regex
|
||||
regex::Regex::new(pattern)
|
||||
.map_err(|e| ToolError::InvalidParameters(format!("invalid regex: {e}")))?;
|
||||
// Validate regex with size limit to prevent ReDoS (issue #825)
|
||||
regex::RegexBuilder::new(pattern)
|
||||
.size_limit(64 * 1024)
|
||||
.build()
|
||||
.map_err(|e| {
|
||||
ToolError::InvalidParameters(format!("invalid or too complex regex: {e}"))
|
||||
})?;
|
||||
let channel = params
|
||||
.get("event_channel")
|
||||
.and_then(|v| v.as_str())
|
||||
@@ -478,41 +523,13 @@ impl Tool for RoutineUpdateTool {
|
||||
}
|
||||
|
||||
fn description(&self) -> &str {
|
||||
"Update an existing routine. Can modify trigger, prompt, schedule, or toggle enabled state. \
|
||||
Pass the routine name and only the fields you want to change."
|
||||
"Update an existing routine. Can change prompt, description, enabled state, or cron timing. \
|
||||
Pass the routine name and only the fields you want to change. \
|
||||
This does not convert one trigger type into another."
|
||||
}
|
||||
|
||||
fn parameters_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {
|
||||
"type": "string",
|
||||
"description": "Name of the routine to update"
|
||||
},
|
||||
"enabled": {
|
||||
"type": "boolean",
|
||||
"description": "Enable or disable the routine"
|
||||
},
|
||||
"prompt": {
|
||||
"type": "string",
|
||||
"description": "New prompt/instructions"
|
||||
},
|
||||
"schedule": {
|
||||
"type": "string",
|
||||
"description": "New cron schedule (for cron triggers)"
|
||||
},
|
||||
"timezone": {
|
||||
"type": "string",
|
||||
"description": "IANA timezone for cron schedule (e.g. 'America/New_York'). Only valid for cron triggers."
|
||||
},
|
||||
"description": {
|
||||
"type": "string",
|
||||
"description": "New description"
|
||||
}
|
||||
},
|
||||
"required": ["name"]
|
||||
})
|
||||
routine_update_parameters_schema()
|
||||
}
|
||||
|
||||
async fn execute(
|
||||
@@ -953,3 +970,117 @@ impl Tool for EventEmitTool {
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{routine_create_parameters_schema, routine_update_parameters_schema};
|
||||
use crate::tools::validate_tool_schema;
|
||||
|
||||
fn property<'a>(schema: &'a serde_json::Value, name: &str) -> &'a serde_json::Value {
|
||||
schema
|
||||
.get("properties")
|
||||
.and_then(|props| props.get(name))
|
||||
.unwrap_or_else(|| panic!("missing schema property {name}"))
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn routine_create_schema_exposes_all_trigger_and_delivery_fields() {
|
||||
let schema = routine_create_parameters_schema();
|
||||
let errors = validate_tool_schema(&schema, "routine_create");
|
||||
assert!(
|
||||
errors.is_empty(),
|
||||
"routine_create schema should validate cleanly: {errors:?}"
|
||||
);
|
||||
|
||||
for field in [
|
||||
"trigger_type",
|
||||
"schedule",
|
||||
"event_pattern",
|
||||
"event_channel",
|
||||
"event_source",
|
||||
"event_type",
|
||||
"event_filters",
|
||||
"action_type",
|
||||
"use_tools",
|
||||
"max_tool_rounds",
|
||||
"tool_permissions",
|
||||
"notify_channel",
|
||||
"notify_user",
|
||||
"timezone",
|
||||
] {
|
||||
let _ = property(&schema, field);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn routine_create_schema_descriptions_cover_event_trigger_gotchas() {
|
||||
let schema = routine_create_parameters_schema();
|
||||
|
||||
let trigger_type = property(&schema, "trigger_type")
|
||||
.get("description")
|
||||
.and_then(|value| value.as_str())
|
||||
.expect("trigger_type description");
|
||||
assert!(trigger_type.contains("incoming messages"));
|
||||
assert!(trigger_type.contains("structured emitted events"));
|
||||
|
||||
let event_pattern = property(&schema, "event_pattern")
|
||||
.get("description")
|
||||
.and_then(|value| value.as_str())
|
||||
.expect("event_pattern description");
|
||||
assert!(event_pattern.contains("incoming message text"));
|
||||
assert!(event_pattern.contains("^bug\\\\b"));
|
||||
|
||||
let event_channel = property(&schema, "event_channel")
|
||||
.get("description")
|
||||
.and_then(|value| value.as_str())
|
||||
.expect("event_channel description");
|
||||
assert!(event_channel.contains("Omit to match any channel"));
|
||||
assert!(event_channel.contains("Not a chat or thread ID"));
|
||||
|
||||
let notify_channel = property(&schema, "notify_channel")
|
||||
.get("description")
|
||||
.and_then(|value| value.as_str())
|
||||
.expect("notify_channel description");
|
||||
assert!(notify_channel.contains("does not control what triggers"));
|
||||
|
||||
let prompt = property(&schema, "prompt")
|
||||
.get("description")
|
||||
.and_then(|value| value.as_str())
|
||||
.expect("prompt description");
|
||||
assert!(prompt.contains("after it fires"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn routine_update_schema_exposes_supported_fields_and_limits() {
|
||||
let schema = routine_update_parameters_schema();
|
||||
let errors = validate_tool_schema(&schema, "routine_update");
|
||||
assert!(
|
||||
errors.is_empty(),
|
||||
"routine_update schema should validate cleanly: {errors:?}"
|
||||
);
|
||||
|
||||
for field in [
|
||||
"name",
|
||||
"enabled",
|
||||
"prompt",
|
||||
"schedule",
|
||||
"timezone",
|
||||
"description",
|
||||
] {
|
||||
let _ = property(&schema, field);
|
||||
}
|
||||
|
||||
let schedule = property(&schema, "schedule")
|
||||
.get("description")
|
||||
.and_then(|value| value.as_str())
|
||||
.expect("schedule description");
|
||||
assert!(schedule.contains("existing 'cron' routines only"));
|
||||
assert!(schedule.contains("does not convert other trigger types"));
|
||||
|
||||
let timezone = property(&schema, "timezone")
|
||||
.get("description")
|
||||
.and_then(|value| value.as_str())
|
||||
.expect("timezone description");
|
||||
assert!(timezone.contains("existing 'cron' routines only"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -247,7 +247,11 @@ fn resolve_timezone_for_output(
|
||||
params: &serde_json::Value,
|
||||
ctx: &JobContext,
|
||||
) -> Result<Option<(Tz, String)>, ToolError> {
|
||||
if let Some(name) = params.get("timezone").and_then(|v| v.as_str()) {
|
||||
if let Some(name) = params
|
||||
.get("timezone")
|
||||
.and_then(|v| v.as_str())
|
||||
.filter(|s| !s.is_empty())
|
||||
{
|
||||
let tz = parse_timezone(name)?;
|
||||
return Ok(Some((tz, tz.to_string())));
|
||||
}
|
||||
@@ -286,7 +290,11 @@ fn context_timezone(ctx: &JobContext) -> Result<Option<(Tz, String)>, ToolError>
|
||||
|
||||
fn optional_timezone(params: &serde_json::Value, keys: &[&str]) -> Result<Option<Tz>, ToolError> {
|
||||
for key in keys {
|
||||
if let Some(value) = params.get(*key).and_then(|v| v.as_str()) {
|
||||
if let Some(value) = params
|
||||
.get(*key)
|
||||
.and_then(|v| v.as_str())
|
||||
.filter(|s| !s.is_empty())
|
||||
{
|
||||
return parse_timezone(value).map(Some);
|
||||
}
|
||||
}
|
||||
@@ -534,4 +542,48 @@ mod tests {
|
||||
|
||||
assert_eq!(dt.to_rfc3339(), "2026-03-08T07:30:00+00:00");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_now_with_empty_timezone_string_does_not_error() {
|
||||
// LLMs sometimes pass "" for optional fields instead of omitting them.
|
||||
// Empty timezone should be treated as absent and fall back to UTC.
|
||||
let tool = TimeTool;
|
||||
let ctx = JobContext::with_user("test", "chat", "test");
|
||||
|
||||
let output = tool
|
||||
.execute(
|
||||
serde_json::json!({
|
||||
"operation": "now",
|
||||
"timezone": ""
|
||||
}),
|
||||
&ctx,
|
||||
)
|
||||
.await
|
||||
.expect("empty timezone string should not error");
|
||||
|
||||
assert!(output.result.get("iso").is_some(), "should have iso");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_convert_with_empty_from_timezone_string_does_not_error() {
|
||||
// LLMs sometimes pass "" for optional fields instead of omitting them.
|
||||
// Empty from_timezone should be treated as absent.
|
||||
let tool = TimeTool;
|
||||
let ctx = JobContext::with_user("test", "chat", "test");
|
||||
|
||||
let output = tool
|
||||
.execute(
|
||||
serde_json::json!({
|
||||
"operation": "convert",
|
||||
"timestamp": "2026-03-08T12:00:00Z",
|
||||
"to_timezone": "America/New_York",
|
||||
"from_timezone": ""
|
||||
}),
|
||||
&ctx,
|
||||
)
|
||||
.await
|
||||
.expect("empty from_timezone string should not error");
|
||||
|
||||
assert!(output.result.get("output").is_some(), "should have output");
|
||||
}
|
||||
}
|
||||
|
||||
+60
-40
@@ -18,6 +18,44 @@ use crate::cli::oauth_defaults::{self, OAUTH_CALLBACK_PORT};
|
||||
use crate::secrets::{CreateSecretParams, SecretsStore};
|
||||
use crate::tools::mcp::config::McpServerConfig;
|
||||
|
||||
/// Shared HTTP client for all OAuth/discovery requests.
|
||||
///
|
||||
/// Redirects are disabled for security (prevents redirect-based SSRF).
|
||||
/// Per-request timeouts can override the default via `.timeout()` on
|
||||
/// the request builder.
|
||||
fn oauth_http_client() -> Result<&'static reqwest::Client, AuthError> {
|
||||
static CLIENT: std::sync::OnceLock<Result<reqwest::Client, String>> =
|
||||
std::sync::OnceLock::new();
|
||||
CLIENT
|
||||
.get_or_init(|| {
|
||||
reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(30))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.map_err(|e| e.to_string())
|
||||
})
|
||||
.as_ref()
|
||||
.map_err(|e| AuthError::Http(e.clone()))
|
||||
}
|
||||
|
||||
/// Log a debug message when a discovery/auth response is a redirect.
|
||||
/// Helps users diagnose configuration issues when legitimate servers
|
||||
/// redirect and our no-redirect policy causes a failure.
|
||||
fn log_redirect_if_applicable(url: &str, response: &reqwest::Response) {
|
||||
if response.status().is_redirection() {
|
||||
let location = response
|
||||
.headers()
|
||||
.get("location")
|
||||
.and_then(|v| v.to_str().ok());
|
||||
tracing::debug!(
|
||||
"OAuth request to '{}' returned redirect {} -> {:?} (redirects disabled for security)",
|
||||
url,
|
||||
response.status(),
|
||||
location
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
/// OAuth authorization error.
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum AuthError {
|
||||
@@ -287,10 +325,8 @@ async fn validate_url_safe(url: &str) -> Result<(), AuthError> {
|
||||
)));
|
||||
}
|
||||
if scheme == "http" {
|
||||
let host = parsed.host_str().unwrap_or("");
|
||||
let is_localhost =
|
||||
host == "localhost" || host == "127.0.0.1" || host == "::1" || host == "[::1]";
|
||||
if !is_localhost {
|
||||
if !crate::tools::mcp::config::is_localhost_url(url) {
|
||||
let host = parsed.host_str().unwrap_or("");
|
||||
return Err(AuthError::DiscoveryFailed(format!(
|
||||
"HTTP is only allowed for localhost; use HTTPS for '{}'",
|
||||
host
|
||||
@@ -382,18 +418,17 @@ fn parse_resource_metadata_url(www_authenticate: &str) -> Option<String> {
|
||||
async fn fetch_resource_metadata(url: &str) -> Result<ProtectedResourceMetadata, AuthError> {
|
||||
validate_url_safe(url).await?;
|
||||
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(10))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.map_err(|e| AuthError::Http(e.to_string()))?;
|
||||
let client = oauth_http_client()?;
|
||||
|
||||
let response = client
|
||||
.get(url)
|
||||
.timeout(Duration::from_secs(10))
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?;
|
||||
|
||||
log_redirect_if_applicable(url, &response);
|
||||
|
||||
if !response.status().is_success() {
|
||||
return Err(AuthError::DiscoveryFailed(format!(
|
||||
"HTTP {}",
|
||||
@@ -411,20 +446,19 @@ async fn fetch_resource_metadata(url: &str) -> Result<ProtectedResourceMetadata,
|
||||
async fn discover_via_401(server_url: &str) -> Result<AuthorizationServerMetadata, AuthError> {
|
||||
validate_url_safe(server_url).await?;
|
||||
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(10))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.map_err(|e| AuthError::Http(e.to_string()))?;
|
||||
let client = oauth_http_client()?;
|
||||
|
||||
let response = client
|
||||
.post(server_url)
|
||||
.timeout(Duration::from_secs(10))
|
||||
.header("Content-Type", "application/json")
|
||||
.body("{}")
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?;
|
||||
|
||||
log_redirect_if_applicable(server_url, &response);
|
||||
|
||||
if response.status().as_u16() != 401 {
|
||||
return Err(AuthError::DiscoveryFailed(format!(
|
||||
"Expected 401, got {}",
|
||||
@@ -472,20 +506,19 @@ pub async fn discover_protected_resource(
|
||||
) -> Result<ProtectedResourceMetadata, AuthError> {
|
||||
validate_url_safe(server_url).await?;
|
||||
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(10))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.map_err(|e| AuthError::Http(e.to_string()))?;
|
||||
let client = oauth_http_client()?;
|
||||
|
||||
let well_known_url = build_well_known_uri(server_url, "oauth-protected-resource")?;
|
||||
|
||||
let response = client
|
||||
.get(&well_known_url)
|
||||
.timeout(Duration::from_secs(10))
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?;
|
||||
|
||||
log_redirect_if_applicable(&well_known_url, &response);
|
||||
|
||||
if !response.status().is_success() {
|
||||
return Err(AuthError::NotSupported);
|
||||
}
|
||||
@@ -502,20 +535,19 @@ pub async fn discover_authorization_server(
|
||||
) -> Result<AuthorizationServerMetadata, AuthError> {
|
||||
validate_url_safe(auth_server_url).await?;
|
||||
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(10))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.map_err(|e| AuthError::Http(e.to_string()))?;
|
||||
let client = oauth_http_client()?;
|
||||
|
||||
let well_known_url = build_well_known_uri(auth_server_url, "oauth-authorization-server")?;
|
||||
|
||||
let response = client
|
||||
.get(&well_known_url)
|
||||
.timeout(Duration::from_secs(10))
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| AuthError::DiscoveryFailed(e.to_string()))?;
|
||||
|
||||
log_redirect_if_applicable(&well_known_url, &response);
|
||||
|
||||
if !response.status().is_success() {
|
||||
return Err(AuthError::DiscoveryFailed(format!(
|
||||
"HTTP {}",
|
||||
@@ -595,11 +627,7 @@ pub async fn register_client(
|
||||
) -> Result<ClientRegistrationResponse, AuthError> {
|
||||
validate_url_safe(registration_endpoint).await?;
|
||||
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(30))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.map_err(|e| AuthError::Http(e.to_string()))?;
|
||||
let client = oauth_http_client()?;
|
||||
|
||||
let request = ClientRegistrationRequest {
|
||||
client_name: "IronClaw".to_string(),
|
||||
@@ -813,7 +841,7 @@ pub fn build_authorization_url(
|
||||
if let Some(pkce) = pkce {
|
||||
url.push_str(&format!(
|
||||
"&code_challenge={}&code_challenge_method=S256",
|
||||
pkce.challenge
|
||||
urlencoding::encode(&pkce.challenge)
|
||||
));
|
||||
}
|
||||
|
||||
@@ -863,11 +891,7 @@ pub async fn exchange_code_for_token(
|
||||
) -> Result<AccessToken, AuthError> {
|
||||
validate_url_safe(token_url).await?;
|
||||
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(30))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.map_err(|e| AuthError::Http(e.to_string()))?;
|
||||
let client = oauth_http_client()?;
|
||||
|
||||
let mut params = vec![
|
||||
("grant_type", "authorization_code".to_string()),
|
||||
@@ -1054,11 +1078,7 @@ pub async fn refresh_access_token(
|
||||
|
||||
validate_url_safe(&token_url).await?;
|
||||
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(30))
|
||||
.redirect(reqwest::redirect::Policy::none())
|
||||
.build()
|
||||
.map_err(|e| AuthError::Http(e.to_string()))?;
|
||||
let client = oauth_http_client()?;
|
||||
|
||||
// Compute canonical resource URI for RFC 8707
|
||||
let resource = canonical_resource_uri(&server_config.url);
|
||||
|
||||
+205
-63
@@ -5,7 +5,7 @@
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use tokio::sync::RwLock;
|
||||
@@ -58,9 +58,10 @@ pub struct McpClient {
|
||||
/// Custom headers to include in every request.
|
||||
custom_headers: HashMap<String, String>,
|
||||
|
||||
/// Whether the MCP initialize handshake has completed.
|
||||
/// Used as a local idempotency guard when no session_manager is present.
|
||||
initialized: AtomicBool,
|
||||
/// Ensures the MCP initialize handshake runs exactly once.
|
||||
/// Uses `OnceCell` to serialize concurrent callers so only one
|
||||
/// actually sends the request; subsequent calls return immediately.
|
||||
initialized: tokio::sync::OnceCell<InitializeResult>,
|
||||
}
|
||||
|
||||
impl McpClient {
|
||||
@@ -83,7 +84,7 @@ impl McpClient {
|
||||
user_id: "default".to_string(),
|
||||
server_config: None,
|
||||
custom_headers: HashMap::new(),
|
||||
initialized: AtomicBool::new(false),
|
||||
initialized: tokio::sync::OnceCell::new(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -106,7 +107,7 @@ impl McpClient {
|
||||
user_id: "default".to_string(),
|
||||
server_config: None,
|
||||
custom_headers: HashMap::new(),
|
||||
initialized: AtomicBool::new(false),
|
||||
initialized: tokio::sync::OnceCell::new(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -114,20 +115,24 @@ impl McpClient {
|
||||
///
|
||||
/// Use this when you have an `McpServerConfig` with custom headers but no OAuth.
|
||||
/// The config must use HTTP transport (the default); for stdio/UDS use `new_with_transport`.
|
||||
pub fn new_with_config(config: McpServerConfig) -> Self {
|
||||
assert!(
|
||||
matches!(
|
||||
config.effective_transport(),
|
||||
crate::tools::mcp::config::EffectiveTransport::Http
|
||||
),
|
||||
"new_with_config only supports HTTP transport; use new_with_transport for stdio/UDS"
|
||||
);
|
||||
///
|
||||
/// Returns an error if the config uses a non-HTTP transport.
|
||||
pub fn new_with_config(config: McpServerConfig) -> Result<Self, ToolError> {
|
||||
if !matches!(
|
||||
config.effective_transport(),
|
||||
crate::tools::mcp::config::EffectiveTransport::Http
|
||||
) {
|
||||
return Err(ToolError::InvalidParameters(
|
||||
"new_with_config only supports HTTP transport; use new_with_transport for stdio/UDS"
|
||||
.to_string(),
|
||||
));
|
||||
}
|
||||
let transport = Arc::new(HttpMcpTransport::new(
|
||||
config.url.clone(),
|
||||
config.name.clone(),
|
||||
));
|
||||
|
||||
Self {
|
||||
Ok(Self {
|
||||
transport,
|
||||
server_url: config.url.clone(),
|
||||
server_name: config.name.clone(),
|
||||
@@ -137,9 +142,9 @@ impl McpClient {
|
||||
secrets: None,
|
||||
user_id: "default".to_string(),
|
||||
custom_headers: config.headers.clone(),
|
||||
initialized: AtomicBool::new(false),
|
||||
initialized: tokio::sync::OnceCell::new(),
|
||||
server_config: Some(config),
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/// Create a new authenticated MCP client.
|
||||
@@ -169,7 +174,7 @@ impl McpClient {
|
||||
user_id: user_id.into(),
|
||||
server_config: Some(config),
|
||||
custom_headers,
|
||||
initialized: AtomicBool::new(false),
|
||||
initialized: tokio::sync::OnceCell::new(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -205,7 +210,7 @@ impl McpClient {
|
||||
user_id: user_id.into(),
|
||||
server_config,
|
||||
custom_headers,
|
||||
initialized: AtomicBool::new(false),
|
||||
initialized: tokio::sync::OnceCell::new(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -336,53 +341,64 @@ impl McpClient {
|
||||
}
|
||||
|
||||
/// Initialize the connection to the MCP server.
|
||||
///
|
||||
/// Uses `OnceCell` to guarantee that exactly one caller performs the
|
||||
/// handshake, even under concurrent access. Subsequent calls return
|
||||
/// immediately.
|
||||
pub async fn initialize(&self) -> Result<InitializeResult, ToolError> {
|
||||
// Fast path: already initialized (local flag or session manager)
|
||||
if self.initialized.load(Ordering::Relaxed) {
|
||||
return Ok(InitializeResult::default());
|
||||
}
|
||||
if let Some(ref session_manager) = self.session_manager
|
||||
&& session_manager.is_initialized(&self.server_name).await
|
||||
{
|
||||
self.initialized.store(true, Ordering::Relaxed);
|
||||
return Ok(InitializeResult::default());
|
||||
}
|
||||
if let Some(ref session_manager) = self.session_manager {
|
||||
session_manager
|
||||
.get_or_create(&self.server_name, &self.server_url)
|
||||
.await;
|
||||
}
|
||||
let result = self
|
||||
.initialized
|
||||
.get_or_try_init(|| async {
|
||||
if let Some(ref session_manager) = self.session_manager
|
||||
&& session_manager.is_initialized(&self.server_name).await
|
||||
{
|
||||
return Ok(InitializeResult::default());
|
||||
}
|
||||
if let Some(ref session_manager) = self.session_manager {
|
||||
session_manager
|
||||
.get_or_create(&self.server_name, &self.server_url)
|
||||
.await;
|
||||
}
|
||||
|
||||
let request = McpRequest::initialize(self.next_request_id());
|
||||
let response = self.send_request(request).await?;
|
||||
let request = McpRequest::initialize(self.next_request_id());
|
||||
let response = self.send_request(request).await?;
|
||||
|
||||
if let Some(error) = response.error {
|
||||
return Err(ToolError::ExternalService(format!(
|
||||
"MCP initialization error: {} (code {})",
|
||||
error.message, error.code
|
||||
)));
|
||||
}
|
||||
if let Some(error) = response.error {
|
||||
return Err(ToolError::ExternalService(format!(
|
||||
"MCP initialization error: {} (code {})",
|
||||
error.message, error.code
|
||||
)));
|
||||
}
|
||||
|
||||
let result: InitializeResult = response
|
||||
.result
|
||||
.ok_or_else(|| {
|
||||
ToolError::ExternalService("No result in initialize response".to_string())
|
||||
let init_result: InitializeResult = response
|
||||
.result
|
||||
.ok_or_else(|| {
|
||||
ToolError::ExternalService("No result in initialize response".to_string())
|
||||
})
|
||||
.and_then(|r| {
|
||||
serde_json::from_value(r).map_err(|e| {
|
||||
ToolError::ExternalService(format!("Invalid initialize result: {}", e))
|
||||
})
|
||||
})?;
|
||||
|
||||
if let Some(ref session_manager) = self.session_manager {
|
||||
session_manager.mark_initialized(&self.server_name).await;
|
||||
}
|
||||
|
||||
let notification = McpRequest::initialized_notification();
|
||||
if let Err(e) = self.send_request(notification).await {
|
||||
tracing::debug!(
|
||||
"Failed to send initialized notification to '{}': {}",
|
||||
self.server_name,
|
||||
e
|
||||
);
|
||||
}
|
||||
|
||||
Ok(init_result)
|
||||
})
|
||||
.and_then(|r| {
|
||||
serde_json::from_value(r).map_err(|e| {
|
||||
ToolError::ExternalService(format!("Invalid initialize result: {}", e))
|
||||
})
|
||||
})?;
|
||||
.await?;
|
||||
|
||||
if let Some(ref session_manager) = self.session_manager {
|
||||
session_manager.mark_initialized(&self.server_name).await;
|
||||
}
|
||||
self.initialized.store(true, Ordering::Relaxed);
|
||||
|
||||
let notification = McpRequest::initialized_notification();
|
||||
let _ = self.send_request(notification).await;
|
||||
|
||||
Ok(result)
|
||||
Ok(result.clone())
|
||||
}
|
||||
|
||||
/// List available tools from the MCP server.
|
||||
@@ -471,6 +487,11 @@ impl McpClient {
|
||||
}
|
||||
}
|
||||
|
||||
/// Clone the client, resetting the tools cache and initialization state.
|
||||
/// The cloned client shares the same transport and session manager, so
|
||||
/// re-initialization will short-circuit via the session manager check if
|
||||
/// the source was already initialized. The `next_id` counter is copied
|
||||
/// so that cloned clients continue with monotonically increasing IDs.
|
||||
impl Clone for McpClient {
|
||||
fn clone(&self) -> Self {
|
||||
Self {
|
||||
@@ -484,7 +505,7 @@ impl Clone for McpClient {
|
||||
user_id: self.user_id.clone(),
|
||||
server_config: self.server_config.clone(),
|
||||
custom_headers: self.custom_headers.clone(),
|
||||
initialized: AtomicBool::new(self.initialized.load(Ordering::Relaxed)),
|
||||
initialized: tokio::sync::OnceCell::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -707,7 +728,7 @@ mod tests {
|
||||
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.clone());
|
||||
let client = McpClient::new_with_config(config.clone()).expect("HTTP config should work");
|
||||
|
||||
assert_eq!(client.server_name(), "test");
|
||||
assert_eq!(client.server_url(), "http://localhost:8080");
|
||||
@@ -719,7 +740,7 @@ mod tests {
|
||||
#[test]
|
||||
fn test_new_with_config_no_headers() {
|
||||
let config = McpServerConfig::new("bare", "http://localhost:9090");
|
||||
let client = McpClient::new_with_config(config);
|
||||
let client = McpClient::new_with_config(config).expect("HTTP config should work");
|
||||
|
||||
assert_eq!(client.server_name(), "bare");
|
||||
assert!(client.custom_headers.is_empty());
|
||||
@@ -971,4 +992,125 @@ mod tests {
|
||||
assert_eq!(obj.len(), 1);
|
||||
assert!(obj["outer"]["inner"].is_null());
|
||||
}
|
||||
|
||||
// --- Issue 1 regression: new_with_config rejects non-HTTP transport ---
|
||||
|
||||
#[test]
|
||||
fn test_new_with_config_rejects_stdio_transport() {
|
||||
let config = McpServerConfig::new_stdio(
|
||||
"stdio-server",
|
||||
"echo",
|
||||
vec!["hello".to_string()],
|
||||
HashMap::new(),
|
||||
);
|
||||
let result = McpClient::new_with_config(config);
|
||||
let err = result
|
||||
.err()
|
||||
.expect("stdio config must be rejected")
|
||||
.to_string();
|
||||
assert!(
|
||||
err.contains("new_with_config only supports HTTP"),
|
||||
"error should explain the restriction: {}",
|
||||
err
|
||||
);
|
||||
}
|
||||
|
||||
// --- Issue 13: McpToolWrapper unit tests ---
|
||||
|
||||
fn make_test_mcp_tool(destructive: bool) -> McpTool {
|
||||
use crate::tools::mcp::protocol::McpToolAnnotations;
|
||||
McpTool {
|
||||
name: "do_thing".to_string(),
|
||||
description: "Does a thing".to_string(),
|
||||
input_schema: serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"input": {"type": "string"}
|
||||
}
|
||||
}),
|
||||
annotations: if destructive {
|
||||
Some(McpToolAnnotations {
|
||||
destructive_hint: true,
|
||||
side_effects_hint: false,
|
||||
read_only_hint: false,
|
||||
execution_time_hint: None,
|
||||
})
|
||||
} else {
|
||||
None
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_mcp_tool_wrapper_name_is_prefixed() {
|
||||
let client = Arc::new(McpClient::new("http://localhost:8080"));
|
||||
let wrapper = McpToolWrapper {
|
||||
tool: make_test_mcp_tool(false),
|
||||
prefixed_name: "mcp__myserver__do_thing".to_string(),
|
||||
client,
|
||||
};
|
||||
assert_eq!(wrapper.name(), "mcp__myserver__do_thing");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_mcp_tool_wrapper_description() {
|
||||
let client = Arc::new(McpClient::new("http://localhost:8080"));
|
||||
let wrapper = McpToolWrapper {
|
||||
tool: make_test_mcp_tool(false),
|
||||
prefixed_name: "mcp__s__do_thing".to_string(),
|
||||
client,
|
||||
};
|
||||
assert_eq!(wrapper.description(), "Does a thing");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_mcp_tool_wrapper_parameters_schema() {
|
||||
let client = Arc::new(McpClient::new("http://localhost:8080"));
|
||||
let wrapper = McpToolWrapper {
|
||||
tool: make_test_mcp_tool(false),
|
||||
prefixed_name: "mcp__s__do_thing".to_string(),
|
||||
client,
|
||||
};
|
||||
let schema = wrapper.parameters_schema();
|
||||
assert_eq!(schema["type"], "object");
|
||||
assert!(schema["properties"]["input"].is_object());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_mcp_tool_wrapper_requires_sanitization() {
|
||||
let client = Arc::new(McpClient::new("http://localhost:8080"));
|
||||
let wrapper = McpToolWrapper {
|
||||
tool: make_test_mcp_tool(false),
|
||||
prefixed_name: "mcp__s__do_thing".to_string(),
|
||||
client,
|
||||
};
|
||||
assert!(
|
||||
wrapper.requires_sanitization(),
|
||||
"MCP tools should always require sanitization"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_mcp_tool_wrapper_approval_destructive() {
|
||||
let client = Arc::new(McpClient::new("http://localhost:8080"));
|
||||
let wrapper = McpToolWrapper {
|
||||
tool: make_test_mcp_tool(true),
|
||||
prefixed_name: "mcp__s__do_thing".to_string(),
|
||||
client,
|
||||
};
|
||||
let approval = wrapper.requires_approval(&serde_json::json!({}));
|
||||
assert_eq!(approval, ApprovalRequirement::UnlessAutoApproved);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_mcp_tool_wrapper_approval_non_destructive() {
|
||||
let client = Arc::new(McpClient::new("http://localhost:8080"));
|
||||
let wrapper = McpToolWrapper {
|
||||
tool: make_test_mcp_tool(false),
|
||||
prefixed_name: "mcp__s__do_thing".to_string(),
|
||||
client,
|
||||
};
|
||||
let approval = wrapper.requires_approval(&serde_json::json!({}));
|
||||
assert_eq!(approval, ApprovalRequirement::Never);
|
||||
}
|
||||
}
|
||||
|
||||
+38
-6
@@ -163,10 +163,8 @@ impl McpServerConfig {
|
||||
}
|
||||
|
||||
// Remote servers must use HTTPS (localhost is allowed for development)
|
||||
let url_lower = self.url.to_lowercase();
|
||||
let is_localhost =
|
||||
url_lower.contains("localhost") || url_lower.contains("127.0.0.1");
|
||||
if !is_localhost && !url_lower.starts_with("https://") {
|
||||
let is_localhost = is_localhost_url(&self.url);
|
||||
if !is_localhost && !self.url.to_lowercase().starts_with("https://") {
|
||||
return Err(ConfigError::InvalidConfig {
|
||||
reason: "Remote MCP servers must use HTTPS".to_string(),
|
||||
});
|
||||
@@ -442,7 +440,12 @@ pub async fn save_mcp_servers_to(
|
||||
}
|
||||
|
||||
let content = serde_json::to_string_pretty(config)?;
|
||||
fs::write(path, content).await?;
|
||||
|
||||
// Write to a temporary file first, then atomically rename to avoid
|
||||
// corrupting the config if the process crashes during the write.
|
||||
let tmp_path = path.with_extension("json.tmp");
|
||||
fs::write(&tmp_path, content).await?;
|
||||
fs::rename(&tmp_path, path).await?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -570,7 +573,7 @@ pub async fn remove_mcp_server_db(
|
||||
///
|
||||
/// Uses `url::Url` for proper parsing so edge cases (IPv6, userinfo, ports)
|
||||
/// are handled correctly without manual string splitting.
|
||||
fn is_localhost_url(url: &str) -> bool {
|
||||
pub(crate) fn is_localhost_url(url: &str) -> bool {
|
||||
let Ok(parsed) = url::Url::parse(url) else {
|
||||
return false;
|
||||
};
|
||||
@@ -1125,4 +1128,33 @@ mod tests {
|
||||
assert!(parsed.transport.is_none());
|
||||
assert_eq!(parsed.headers.get("X-Custom").unwrap(), "value");
|
||||
}
|
||||
|
||||
// --- Issue 3 regression: is_localhost_url rejects attacker subdomains ---
|
||||
|
||||
#[test]
|
||||
fn test_is_localhost_url_rejects_attacker_subdomain() {
|
||||
// Before the fix, url.contains("localhost") matched this.
|
||||
assert!(
|
||||
!is_localhost_url("http://evil.localhost.attacker.com:8080/mcp"),
|
||||
"attacker subdomain containing 'localhost' must not be treated as local"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_localhost_url_accepts_real_localhost() {
|
||||
assert!(is_localhost_url("http://localhost:8080/mcp"));
|
||||
assert!(is_localhost_url("https://localhost/path"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_localhost_url_accepts_loopback_ip() {
|
||||
assert!(is_localhost_url("http://127.0.0.1:3000"));
|
||||
assert!(is_localhost_url("http://[::1]:3000"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_localhost_url_rejects_remote() {
|
||||
assert!(!is_localhost_url("https://mcp.example.com"));
|
||||
assert!(!is_localhost_url("http://192.168.1.1:8080"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -18,6 +18,8 @@ pub enum McpFactoryError {
|
||||
UnixConnect { name: String, reason: String },
|
||||
#[error("Unix socket transport is not supported on this platform (server '{name}')")]
|
||||
UnixNotSupported { name: String },
|
||||
#[error("Invalid configuration for MCP server '{name}': {reason}")]
|
||||
InvalidConfig { name: String, reason: String },
|
||||
}
|
||||
|
||||
/// Create an `McpClient` from a server configuration, dispatching on the
|
||||
@@ -89,10 +91,18 @@ pub async fn create_client_from_config(
|
||||
))
|
||||
} else {
|
||||
Ok(McpClient::new_with_config(server)
|
||||
.map_err(|e| McpFactoryError::InvalidConfig {
|
||||
name: server_name.clone(),
|
||||
reason: e.to_string(),
|
||||
})?
|
||||
.with_session_manager(Arc::clone(session_manager)))
|
||||
}
|
||||
} else {
|
||||
Ok(McpClient::new_with_config(server)
|
||||
.map_err(|e| McpFactoryError::InvalidConfig {
|
||||
name: server_name,
|
||||
reason: e.to_string(),
|
||||
})?
|
||||
.with_session_manager(Arc::clone(session_manager)))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -139,7 +139,7 @@ impl McpTransport for HttpMcpTransport {
|
||||
.to_string();
|
||||
|
||||
if content_type.contains("text/event-stream") {
|
||||
self.parse_sse_response(response).await
|
||||
self.parse_sse_response(response, request.id).await
|
||||
} else {
|
||||
response.json().await.map_err(|e| {
|
||||
ToolError::ExternalService(format!(
|
||||
@@ -161,11 +161,14 @@ impl McpTransport for HttpMcpTransport {
|
||||
}
|
||||
|
||||
impl HttpMcpTransport {
|
||||
/// Parse a Server-Sent Events response, returning the first valid JSON-RPC
|
||||
/// `data:` line as an [`McpResponse`].
|
||||
/// Parse a Server-Sent Events response, returning the JSON-RPC response
|
||||
/// whose `id` matches `request_id`. Non-matching events (e.g. server
|
||||
/// notifications or progress updates) are skipped so that the caller
|
||||
/// receives the actual result for its request.
|
||||
async fn parse_sse_response(
|
||||
&self,
|
||||
response: reqwest::Response,
|
||||
request_id: Option<u64>,
|
||||
) -> Result<McpResponse, ToolError> {
|
||||
use futures::StreamExt;
|
||||
|
||||
@@ -202,9 +205,10 @@ impl HttpMcpTransport {
|
||||
remaining_start = i + 1;
|
||||
|
||||
if let Some(json_str) = line.strip_prefix("data: ")
|
||||
&& let Ok(response) = serde_json::from_str::<McpResponse>(json_str)
|
||||
&& let Ok(resp) = serde_json::from_str::<McpResponse>(json_str)
|
||||
&& resp.id == request_id
|
||||
{
|
||||
return Ok(response);
|
||||
return Ok(resp);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -216,14 +220,15 @@ impl HttpMcpTransport {
|
||||
|
||||
// Process any remaining data without a trailing newline.
|
||||
if let Some(json_str) = buffer.strip_prefix("data: ")
|
||||
&& let Ok(response) = serde_json::from_str::<McpResponse>(json_str.trim())
|
||||
&& let Ok(resp) = serde_json::from_str::<McpResponse>(json_str.trim())
|
||||
&& resp.id == request_id
|
||||
{
|
||||
return Ok(response);
|
||||
return Ok(resp);
|
||||
}
|
||||
|
||||
Err(ToolError::ExternalService(format!(
|
||||
"[{}] No valid data in SSE response: {}",
|
||||
self.server_name, buffer
|
||||
"[{}] No matching response (id={:?}) in SSE stream",
|
||||
self.server_name, request_id
|
||||
)))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -14,7 +14,7 @@ use tokio::sync::{Mutex, oneshot};
|
||||
use tokio::task::JoinHandle;
|
||||
|
||||
use crate::tools::mcp::protocol::{McpRequest, McpResponse};
|
||||
use crate::tools::mcp::transport::{McpTransport, spawn_jsonrpc_reader, write_jsonrpc_line};
|
||||
use crate::tools::mcp::transport::{McpTransport, spawn_jsonrpc_reader, stream_transport_send};
|
||||
use crate::tools::tool::ToolError;
|
||||
|
||||
/// MCP transport that communicates with a child process over stdin/stdout.
|
||||
@@ -118,63 +118,14 @@ impl McpTransport for StdioMcpTransport {
|
||||
request: &McpRequest,
|
||||
_headers: &HashMap<String, String>,
|
||||
) -> Result<McpResponse, ToolError> {
|
||||
// JSON-RPC notifications (no id) are fire-and-forget: the server
|
||||
// will not send a response, so we must not wait for one.
|
||||
if request.id.is_none() {
|
||||
let mut stdin = self.stdin.lock().await;
|
||||
write_jsonrpc_line(&mut *stdin, request).await?;
|
||||
return Ok(McpResponse {
|
||||
jsonrpc: "2.0".to_string(),
|
||||
id: None,
|
||||
result: None,
|
||||
error: None,
|
||||
});
|
||||
}
|
||||
|
||||
let id = request.id.unwrap_or(0);
|
||||
let (tx, rx) = oneshot::channel();
|
||||
|
||||
// Register the pending response handler before writing the request,
|
||||
// so we don't miss a fast response from the child.
|
||||
{
|
||||
let mut pending = self.pending.lock().await;
|
||||
pending.insert(id, tx);
|
||||
}
|
||||
|
||||
// Write the request to stdin.
|
||||
{
|
||||
let mut stdin = self.stdin.lock().await;
|
||||
if let Err(e) = write_jsonrpc_line(&mut *stdin, request).await {
|
||||
// Remove the pending entry on write failure.
|
||||
let mut pending = self.pending.lock().await;
|
||||
pending.remove(&id);
|
||||
return Err(e);
|
||||
}
|
||||
}
|
||||
|
||||
// Wait for the response with a timeout.
|
||||
let timeout = Duration::from_secs(30);
|
||||
match tokio::time::timeout(timeout, rx).await {
|
||||
Ok(Ok(response)) => Ok(response),
|
||||
Ok(Err(_)) => {
|
||||
// Sender was dropped (reader task ended). Clean up pending entry.
|
||||
let mut pending = self.pending.lock().await;
|
||||
pending.remove(&id);
|
||||
Err(ToolError::ExternalService(format!(
|
||||
"[{}] MCP server closed connection before responding to request {:?}",
|
||||
self.server_name, request.id
|
||||
)))
|
||||
}
|
||||
Err(_) => {
|
||||
// Timeout: remove the pending entry.
|
||||
let mut pending = self.pending.lock().await;
|
||||
pending.remove(&id);
|
||||
Err(ToolError::ExternalService(format!(
|
||||
"[{}] Timeout waiting for response to request {:?} after {:?}",
|
||||
self.server_name, request.id, timeout
|
||||
)))
|
||||
}
|
||||
}
|
||||
stream_transport_send(
|
||||
&self.stdin,
|
||||
&self.pending,
|
||||
request,
|
||||
&self.server_name,
|
||||
Duration::from_secs(30),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn shutdown(&self) -> Result<(), ToolError> {
|
||||
|
||||
+105
-1
@@ -97,7 +97,13 @@ pub fn spawn_jsonrpc_reader<R: AsyncBufRead + Unpin + Send + 'static>(
|
||||
}
|
||||
};
|
||||
|
||||
let id = response.id.unwrap_or(0);
|
||||
let Some(id) = response.id else {
|
||||
tracing::debug!(
|
||||
"[{}] Received JSON-RPC notification (no id), skipping dispatch",
|
||||
server_name
|
||||
);
|
||||
continue;
|
||||
};
|
||||
let mut map = pending.lock().await;
|
||||
if let Some(tx) = map.remove(&id) {
|
||||
// Ignore send error — the receiver may have been dropped (timeout).
|
||||
@@ -115,6 +121,76 @@ pub fn spawn_jsonrpc_reader<R: AsyncBufRead + Unpin + Send + 'static>(
|
||||
})
|
||||
}
|
||||
|
||||
/// Send a JSON-RPC request over a stream-based transport (stdio / unix socket).
|
||||
///
|
||||
/// Handles notification fire-and-forget, pending response registration,
|
||||
/// write, timeout, and cleanup. Used by both [`StdioMcpTransport`] and
|
||||
/// [`UnixMcpTransport`] to avoid duplicating the send logic.
|
||||
pub(crate) async fn stream_transport_send<W: AsyncWrite + Unpin>(
|
||||
writer: &Mutex<W>,
|
||||
pending: &Mutex<HashMap<u64, oneshot::Sender<McpResponse>>>,
|
||||
request: &McpRequest,
|
||||
server_name: &str,
|
||||
timeout_duration: std::time::Duration,
|
||||
) -> Result<McpResponse, ToolError> {
|
||||
// JSON-RPC notifications (no id) are fire-and-forget: the server
|
||||
// will not send a response, so we must not wait for one.
|
||||
if request.id.is_none() {
|
||||
let mut w = writer.lock().await;
|
||||
write_jsonrpc_line(&mut *w, request).await?;
|
||||
return Ok(McpResponse {
|
||||
jsonrpc: "2.0".to_string(),
|
||||
id: None,
|
||||
result: None,
|
||||
error: None,
|
||||
});
|
||||
}
|
||||
|
||||
let id = request.id.unwrap_or(0);
|
||||
let (tx, rx) = oneshot::channel();
|
||||
|
||||
// Register the pending response handler before writing the request,
|
||||
// so we don't miss a fast response from the server.
|
||||
{
|
||||
let mut map = pending.lock().await;
|
||||
map.insert(id, tx);
|
||||
}
|
||||
|
||||
// Write the request.
|
||||
{
|
||||
let mut w = writer.lock().await;
|
||||
if let Err(e) = write_jsonrpc_line(&mut *w, request).await {
|
||||
// Remove the pending entry on write failure.
|
||||
let mut map = pending.lock().await;
|
||||
map.remove(&id);
|
||||
return Err(e);
|
||||
}
|
||||
}
|
||||
|
||||
// Wait for the response with a timeout.
|
||||
match tokio::time::timeout(timeout_duration, rx).await {
|
||||
Ok(Ok(response)) => Ok(response),
|
||||
Ok(Err(_)) => {
|
||||
// Sender was dropped (reader task ended). Clean up pending entry.
|
||||
let mut map = pending.lock().await;
|
||||
map.remove(&id);
|
||||
Err(ToolError::ExternalService(format!(
|
||||
"[{}] MCP server closed connection before responding to request {:?}",
|
||||
server_name, request.id
|
||||
)))
|
||||
}
|
||||
Err(_) => {
|
||||
// Timeout: remove the pending entry.
|
||||
let mut map = pending.lock().await;
|
||||
map.remove(&id);
|
||||
Err(ToolError::ExternalService(format!(
|
||||
"[{}] Timeout waiting for response to request {:?} after {:?}",
|
||||
server_name, request.id, timeout_duration
|
||||
)))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -193,4 +269,32 @@ mod tests {
|
||||
|
||||
handle.await.expect("reader task should finish");
|
||||
}
|
||||
|
||||
/// Issue 9 regression: a JSON-RPC notification (no id) must not resolve
|
||||
/// a pending request keyed by id 0 (the old `unwrap_or(0)` default).
|
||||
#[tokio::test]
|
||||
async fn test_notification_does_not_resolve_pending_id_zero() {
|
||||
// A notification response (no id), followed by a proper response for id 0.
|
||||
let notification = r#"{"jsonrpc":"2.0","method":"notifications/progress","params":{}}"#;
|
||||
let real_response = r#"{"jsonrpc":"2.0","id":0,"result":{"ok":true}}"#;
|
||||
let input = format!("{notification}\n{real_response}\n");
|
||||
|
||||
let reader = std::io::Cursor::new(input.into_bytes());
|
||||
let pending: Arc<Mutex<HashMap<u64, oneshot::Sender<McpResponse>>>> =
|
||||
Arc::new(Mutex::new(HashMap::new()));
|
||||
|
||||
let (tx, rx) = oneshot::channel();
|
||||
{
|
||||
let mut map = pending.lock().await;
|
||||
map.insert(0, tx);
|
||||
}
|
||||
|
||||
let handle = spawn_jsonrpc_reader(reader, pending.clone(), "test".into());
|
||||
|
||||
let resp = rx.await.expect("should receive the real id=0 response");
|
||||
assert_eq!(resp.id, Some(0));
|
||||
assert!(resp.result.is_some());
|
||||
|
||||
handle.await.expect("reader task should finish");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -15,7 +15,7 @@ use tokio::sync::{Mutex, oneshot};
|
||||
use tokio::task::JoinHandle;
|
||||
|
||||
use crate::tools::mcp::protocol::{McpRequest, McpResponse};
|
||||
use crate::tools::mcp::transport::{McpTransport, spawn_jsonrpc_reader, write_jsonrpc_line};
|
||||
use crate::tools::mcp::transport::{McpTransport, spawn_jsonrpc_reader, stream_transport_send};
|
||||
use crate::tools::tool::ToolError;
|
||||
|
||||
/// MCP transport that communicates over a Unix domain socket.
|
||||
@@ -91,63 +91,14 @@ impl McpTransport for UnixMcpTransport {
|
||||
request: &McpRequest,
|
||||
_headers: &HashMap<String, String>,
|
||||
) -> Result<McpResponse, ToolError> {
|
||||
// JSON-RPC notifications (no id) are fire-and-forget: the server
|
||||
// will not send a response, so we must not wait for one.
|
||||
if request.id.is_none() {
|
||||
let mut writer = self.writer.lock().await;
|
||||
write_jsonrpc_line(&mut *writer, request).await?;
|
||||
return Ok(McpResponse {
|
||||
jsonrpc: "2.0".to_string(),
|
||||
id: None,
|
||||
result: None,
|
||||
error: None,
|
||||
});
|
||||
}
|
||||
|
||||
let id = request.id.unwrap_or(0);
|
||||
let (tx, rx) = oneshot::channel();
|
||||
|
||||
// Register the pending response handler before writing the request,
|
||||
// so we don't miss a fast response from the server.
|
||||
{
|
||||
let mut pending = self.pending.lock().await;
|
||||
pending.insert(id, tx);
|
||||
}
|
||||
|
||||
// Write the request to the socket.
|
||||
{
|
||||
let mut writer = self.writer.lock().await;
|
||||
if let Err(e) = write_jsonrpc_line(&mut *writer, request).await {
|
||||
// Remove the pending entry on write failure.
|
||||
let mut pending = self.pending.lock().await;
|
||||
pending.remove(&id);
|
||||
return Err(e);
|
||||
}
|
||||
}
|
||||
|
||||
// Wait for the response with a timeout.
|
||||
let timeout = Duration::from_secs(30);
|
||||
match tokio::time::timeout(timeout, rx).await {
|
||||
Ok(Ok(response)) => Ok(response),
|
||||
Ok(Err(_)) => {
|
||||
// Sender was dropped (reader task ended). Clean up pending entry.
|
||||
let mut pending = self.pending.lock().await;
|
||||
pending.remove(&id);
|
||||
Err(ToolError::ExternalService(format!(
|
||||
"[{}] MCP server closed connection before responding to request {:?}",
|
||||
self.server_name, request.id
|
||||
)))
|
||||
}
|
||||
Err(_) => {
|
||||
// Timeout: remove the pending entry.
|
||||
let mut pending = self.pending.lock().await;
|
||||
pending.remove(&id);
|
||||
Err(ToolError::ExternalService(format!(
|
||||
"[{}] Timeout waiting for response to request {:?} after {:?}",
|
||||
self.server_name, request.id, timeout
|
||||
)))
|
||||
}
|
||||
}
|
||||
stream_transport_send(
|
||||
&self.writer,
|
||||
&self.pending,
|
||||
request,
|
||||
&self.server_name,
|
||||
Duration::from_secs(30),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn shutdown(&self) -> Result<(), ToolError> {
|
||||
|
||||
@@ -558,48 +558,7 @@ mod tests {
|
||||
// Routine tools
|
||||
(
|
||||
"routine_create",
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": { "type": "string", "description": "Routine name" },
|
||||
"description": { "type": "string", "description": "What it does" },
|
||||
"trigger_type": {
|
||||
"type": "string",
|
||||
"enum": ["cron", "event", "system_event", "manual"],
|
||||
"description": "When the routine fires"
|
||||
},
|
||||
"schedule": { "type": "string", "description": "Cron expression" },
|
||||
"event_pattern": { "type": "string", "description": "Regex pattern" },
|
||||
"event_channel": { "type": "string", "description": "Channel filter" },
|
||||
"event_source": { "type": "string", "description": "System event source" },
|
||||
"event_type": { "type": "string", "description": "System event type" },
|
||||
"event_filters": {
|
||||
"type": "object",
|
||||
"additionalProperties": { "type": "string" },
|
||||
"description": "Exact-match payload filters"
|
||||
},
|
||||
"prompt": { "type": "string", "description": "Instructions" },
|
||||
"context_paths": {
|
||||
"type": "array",
|
||||
"items": { "type": "string" },
|
||||
"description": "Workspace paths to load"
|
||||
},
|
||||
"action_type": {
|
||||
"type": "string",
|
||||
"enum": ["lightweight", "full_job"],
|
||||
"description": "Execution mode"
|
||||
},
|
||||
"cooldown_secs": { "type": "integer", "description": "Min seconds between fires" },
|
||||
"tool_permissions": {
|
||||
"type": "array",
|
||||
"items": { "type": "string" },
|
||||
"description": "Pre-authorized tools for full_job mode"
|
||||
},
|
||||
"notify_channel": { "type": "string", "description": "Channel for message tool" },
|
||||
"notify_user": { "type": "string", "description": "User/target to notify" }
|
||||
},
|
||||
"required": ["name", "trigger_type", "prompt"]
|
||||
}),
|
||||
crate::tools::builtin::routine::routine_create_parameters_schema(),
|
||||
),
|
||||
(
|
||||
"routine_list",
|
||||
@@ -611,17 +570,7 @@ mod tests {
|
||||
),
|
||||
(
|
||||
"routine_update",
|
||||
serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": { "type": "string", "description": "Name" },
|
||||
"enabled": { "type": "boolean", "description": "Toggle" },
|
||||
"prompt": { "type": "string", "description": "New prompt" },
|
||||
"schedule": { "type": "string", "description": "New cron schedule" },
|
||||
"description": { "type": "string", "description": "New description" }
|
||||
},
|
||||
"required": ["name"]
|
||||
}),
|
||||
crate::tools::builtin::routine::routine_update_parameters_schema(),
|
||||
),
|
||||
(
|
||||
"routine_delete",
|
||||
|
||||
+48
-3
@@ -430,9 +430,24 @@ pub fn redact_params(params: &serde_json::Value, sensitive: &[&str]) -> serde_js
|
||||
/// Properties without a `"type"` field are allowed (freeform/any-type).
|
||||
/// This is an intentional pattern used by tools like `json` and `http` for
|
||||
/// OpenAI compatibility, since union types with arrays require `items`.
|
||||
/// Maximum nesting depth for tool schema validation to prevent stack overflow
|
||||
/// on maliciously crafted schemas.
|
||||
const MAX_SCHEMA_DEPTH: usize = 16;
|
||||
|
||||
pub fn validate_tool_schema(schema: &serde_json::Value, path: &str) -> Vec<String> {
|
||||
validate_tool_schema_inner(schema, path, 0)
|
||||
}
|
||||
|
||||
fn validate_tool_schema_inner(schema: &serde_json::Value, path: &str, depth: usize) -> Vec<String> {
|
||||
let mut errors = Vec::new();
|
||||
|
||||
if depth > MAX_SCHEMA_DEPTH {
|
||||
errors.push(format!(
|
||||
"{path}: schema nesting exceeds maximum depth of {MAX_SCHEMA_DEPTH}"
|
||||
));
|
||||
return errors;
|
||||
}
|
||||
|
||||
// Rule 1: must have "type": "object" at this level
|
||||
match schema.get("type").and_then(|t| t.as_str()) {
|
||||
Some("object") => {}
|
||||
@@ -474,14 +489,17 @@ pub fn validate_tool_schema(schema: &serde_json::Value, path: &str) -> Vec<Strin
|
||||
if let Some(prop_type) = prop.get("type").and_then(|t| t.as_str()) {
|
||||
match prop_type {
|
||||
"object" => {
|
||||
errors.extend(validate_tool_schema(prop, &prop_path));
|
||||
errors.extend(validate_tool_schema_inner(prop, &prop_path, depth + 1));
|
||||
}
|
||||
"array" => {
|
||||
if let Some(items) = prop.get("items") {
|
||||
// If items is an object type, recurse
|
||||
if items.get("type").and_then(|t| t.as_str()) == Some("object") {
|
||||
errors
|
||||
.extend(validate_tool_schema(items, &format!("{prop_path}.items")));
|
||||
errors.extend(validate_tool_schema_inner(
|
||||
items,
|
||||
&format!("{prop_path}.items"),
|
||||
depth + 1,
|
||||
));
|
||||
}
|
||||
} else {
|
||||
errors.push(format!("{prop_path}: array property missing \"items\""));
|
||||
@@ -810,6 +828,33 @@ mod tests {
|
||||
assert!(errors[0].contains("\"missing_field\""));
|
||||
}
|
||||
|
||||
/// Regression test for issue #975: deeply nested schemas must not cause
|
||||
/// stack overflow. The validator should stop at MAX_SCHEMA_DEPTH and
|
||||
/// report an error instead of recursing infinitely.
|
||||
#[test]
|
||||
fn test_validate_schema_depth_limit() {
|
||||
// Build a schema nested 20 levels deep (exceeds MAX_SCHEMA_DEPTH=16)
|
||||
let mut schema = serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"leaf": { "type": "string" }
|
||||
}
|
||||
});
|
||||
for _ in 0..20 {
|
||||
schema = serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"nested": schema
|
||||
}
|
||||
});
|
||||
}
|
||||
let errors = validate_tool_schema(&schema, "test");
|
||||
assert!(
|
||||
errors.iter().any(|e| e.contains("maximum depth")),
|
||||
"expected depth limit error, got: {errors:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_approval_context_autonomous_allows_unless_auto_approved() {
|
||||
let ctx = ApprovalContext::autonomous();
|
||||
|
||||
@@ -101,24 +101,75 @@ pub struct CapabilitiesFile {
|
||||
pub capabilities: Option<Box<CapabilitiesFile>>,
|
||||
}
|
||||
|
||||
/// Maximum length for the description field to prevent memory abuse.
|
||||
const MAX_DESCRIPTION_CHARS: usize = 4096;
|
||||
/// Maximum serialized size of the parameters schema JSON.
|
||||
const MAX_PARAMETERS_SCHEMA_BYTES: usize = 64 * 1024;
|
||||
|
||||
impl CapabilitiesFile {
|
||||
/// Parse from JSON string.
|
||||
pub fn from_json(json: &str) -> Result<Self, serde_json::Error> {
|
||||
serde_json::from_str::<Self>(json).map(Self::resolve_nested)
|
||||
let mut caps = serde_json::from_str::<Self>(json).map(Self::resolve_nested)?;
|
||||
caps.enforce_limits();
|
||||
Ok(caps)
|
||||
}
|
||||
|
||||
/// Parse from JSON bytes.
|
||||
pub fn from_bytes(bytes: &[u8]) -> Result<Self, serde_json::Error> {
|
||||
serde_json::from_slice::<Self>(bytes).map(Self::resolve_nested)
|
||||
let mut caps = serde_json::from_slice::<Self>(bytes).map(Self::resolve_nested)?;
|
||||
caps.enforce_limits();
|
||||
Ok(caps)
|
||||
}
|
||||
|
||||
/// Truncate oversized fields to prevent unbounded memory usage.
|
||||
fn enforce_limits(&mut self) {
|
||||
// Truncate oversized description (issue #976)
|
||||
if let Some(ref desc) = self.description
|
||||
&& desc.len() > MAX_DESCRIPTION_CHARS
|
||||
{
|
||||
let truncated = &desc[..desc.floor_char_boundary(MAX_DESCRIPTION_CHARS)];
|
||||
tracing::warn!(
|
||||
"Capabilities description truncated from {} to {} chars",
|
||||
desc.len(),
|
||||
MAX_DESCRIPTION_CHARS,
|
||||
);
|
||||
self.description = Some(truncated.to_string());
|
||||
}
|
||||
// Drop oversized parameters schema (issue #977)
|
||||
if let Some(ref params) = self.parameters {
|
||||
let size = params.to_string().len();
|
||||
if size > MAX_PARAMETERS_SCHEMA_BYTES {
|
||||
tracing::warn!(
|
||||
"Capabilities parameters schema dropped ({} bytes exceeds {} limit)",
|
||||
size,
|
||||
MAX_PARAMETERS_SCHEMA_BYTES,
|
||||
);
|
||||
self.parameters = None;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Merge nested `capabilities` wrapper into top-level fields.
|
||||
///
|
||||
/// Channel-level JSON nests tool capabilities under `"capabilities"`.
|
||||
/// This promotes the inner fields so callers can access them uniformly.
|
||||
fn resolve_nested(mut self) -> Self {
|
||||
/// Maximum nesting depth for capabilities resolution.
|
||||
const MAX_NESTED_DEPTH: usize = 8;
|
||||
|
||||
fn resolve_nested(self) -> Self {
|
||||
self.resolve_nested_inner(0)
|
||||
}
|
||||
|
||||
fn resolve_nested_inner(mut self, depth: usize) -> Self {
|
||||
if depth > Self::MAX_NESTED_DEPTH {
|
||||
tracing::warn!(
|
||||
"Capabilities nesting exceeds maximum depth of {}, stopping resolution",
|
||||
Self::MAX_NESTED_DEPTH
|
||||
);
|
||||
return self;
|
||||
}
|
||||
if let Some(inner) = self.capabilities.take() {
|
||||
let inner = inner.resolve_nested();
|
||||
let inner = inner.resolve_nested_inner(depth + 1);
|
||||
self.description = self.description.or(inner.description);
|
||||
self.parameters = self.parameters.or(inner.parameters);
|
||||
self.http = self.http.or(inner.http);
|
||||
@@ -1383,4 +1434,63 @@ mod tests {
|
||||
"Outer description should take precedence over inner"
|
||||
);
|
||||
}
|
||||
|
||||
/// Regression test for issue #974: deeply nested capabilities wrappers
|
||||
/// must not cause stack overflow. resolve_nested should stop at
|
||||
/// MAX_NESTED_DEPTH and return gracefully.
|
||||
#[test]
|
||||
fn test_resolve_nested_depth_limit() {
|
||||
// Build a capabilities file nested beyond MAX_NESTED_DEPTH (8).
|
||||
// The description is at the innermost level which is beyond the limit,
|
||||
// so it won't be resolved — the key assertion is no stack overflow.
|
||||
let mut json = r#"{ "description": "leaf" }"#.to_string();
|
||||
for _ in 0..20 {
|
||||
json = format!(r#"{{ "capabilities": {json} }}"#);
|
||||
}
|
||||
// Should not stack overflow — this is the primary assertion.
|
||||
let _caps = CapabilitiesFile::from_json(&json).unwrap();
|
||||
}
|
||||
|
||||
/// Regression test for issue #976: oversized description strings are truncated.
|
||||
#[test]
|
||||
fn test_description_truncated_at_limit() {
|
||||
let long_desc = "x".repeat(10_000);
|
||||
let json = format!(r#"{{ "description": "{long_desc}" }}"#);
|
||||
let caps = CapabilitiesFile::from_json(&json).unwrap();
|
||||
let desc = caps.description.unwrap();
|
||||
assert!(
|
||||
desc.len() <= super::MAX_DESCRIPTION_CHARS + 50, // allow for minor overhead
|
||||
"description should be truncated to ~{} chars, got {}",
|
||||
super::MAX_DESCRIPTION_CHARS,
|
||||
desc.len()
|
||||
);
|
||||
}
|
||||
|
||||
/// Regression test for issue #977: oversized parameters schema is dropped.
|
||||
#[test]
|
||||
fn test_oversized_parameters_schema_dropped() {
|
||||
// Build a parameters schema larger than MAX_PARAMETERS_SCHEMA_BYTES
|
||||
let mut properties = serde_json::Map::new();
|
||||
for i in 0..2000 {
|
||||
properties.insert(
|
||||
format!("field_{i}"),
|
||||
serde_json::json!({
|
||||
"type": "string",
|
||||
"description": "x".repeat(50)
|
||||
}),
|
||||
);
|
||||
}
|
||||
let schema = serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": properties,
|
||||
});
|
||||
let json = serde_json::json!({
|
||||
"parameters": schema,
|
||||
});
|
||||
let caps = CapabilitiesFile::from_json(&json.to_string()).unwrap();
|
||||
assert!(
|
||||
caps.parameters.is_none(),
|
||||
"oversized parameters schema should be dropped"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
+4
-4
@@ -1108,9 +1108,10 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
|
||||
return LoopSignal::InjectMessage(content);
|
||||
}
|
||||
|
||||
// Check for terminal or non-progressing state. The loop should stop when the
|
||||
// job has been cancelled, failed, stuck, or already completed — not just the
|
||||
// three states that `is_terminal()` covers (Accepted/Failed/Cancelled).
|
||||
// Check for terminal or post-completion state. The loop should stop when the
|
||||
// job has been cancelled, failed, or already completed — but NOT when Stuck,
|
||||
// because Stuck is recoverable (Stuck -> InProgress via self-repair).
|
||||
// Stopping on Stuck would prevent recovery from resuming the worker (issue #892).
|
||||
if let Ok(ctx) = self
|
||||
.worker
|
||||
.context_manager()
|
||||
@@ -1120,7 +1121,6 @@ impl<'a> LoopDelegate for JobDelegate<'a> {
|
||||
ctx.state,
|
||||
JobState::Cancelled
|
||||
| JobState::Failed
|
||||
| JobState::Stuck
|
||||
| JobState::Completed
|
||||
| JobState::Submitted
|
||||
| JobState::Accepted
|
||||
|
||||
@@ -9,6 +9,10 @@ mod support;
|
||||
mod advanced {
|
||||
use std::time::Duration;
|
||||
|
||||
use ironclaw::agent::routine::Trigger;
|
||||
use ironclaw::channels::IncomingMessage;
|
||||
use ironclaw::db::Database;
|
||||
|
||||
use crate::support::cleanup::CleanupGuard;
|
||||
use crate::support::test_rig::TestRigBuilder;
|
||||
use crate::support::trace_llm::LlmTrace;
|
||||
@@ -19,6 +23,28 @@ mod advanced {
|
||||
);
|
||||
const TIMEOUT: Duration = Duration::from_secs(30);
|
||||
|
||||
async fn wait_for_routine_run(
|
||||
db: &std::sync::Arc<dyn Database>,
|
||||
routine_id: uuid::Uuid,
|
||||
timeout: Duration,
|
||||
) -> Vec<ironclaw::agent::routine::RoutineRun> {
|
||||
let deadline = tokio::time::Instant::now() + timeout;
|
||||
loop {
|
||||
let runs = db
|
||||
.list_routine_runs(routine_id, 10)
|
||||
.await
|
||||
.expect("list_routine_runs");
|
||||
if !runs.is_empty() {
|
||||
return runs;
|
||||
}
|
||||
assert!(
|
||||
tokio::time::Instant::now() < deadline,
|
||||
"timed out waiting for routine run"
|
||||
);
|
||||
tokio::time::sleep(Duration::from_millis(100)).await;
|
||||
}
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// 1. Multi-turn memory coherence
|
||||
// -----------------------------------------------------------------------
|
||||
@@ -380,6 +406,118 @@ mod advanced {
|
||||
rig.shutdown();
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// 6b. Event routine: Telegram-scoped trigger fires on matching message
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn routine_event_trigger_telegram_channel_fires() {
|
||||
let trace = LlmTrace::from_file(format!("{FIXTURES}/routine_event_telegram.json")).unwrap();
|
||||
let rig = TestRigBuilder::new()
|
||||
.with_trace(trace.clone())
|
||||
.with_routines()
|
||||
.with_auto_approve_tools(true)
|
||||
.build()
|
||||
.await;
|
||||
|
||||
rig.send_message(
|
||||
"Create a routine that watches Telegram messages starting with 'bug:' and alerts me.",
|
||||
)
|
||||
.await;
|
||||
let create_responses = rig.wait_for_responses(1, TIMEOUT).await;
|
||||
rig.verify_trace_expects(&trace, &create_responses);
|
||||
|
||||
let routine = rig
|
||||
.database()
|
||||
.get_routine_by_name("test-user", "telegram-bug-watcher")
|
||||
.await
|
||||
.expect("get_routine_by_name")
|
||||
.expect("telegram-bug-watcher should exist");
|
||||
|
||||
match &routine.trigger {
|
||||
Trigger::Event { channel, pattern } => {
|
||||
assert_eq!(channel.as_deref(), Some("telegram"));
|
||||
assert_eq!(pattern, "^bug\\b");
|
||||
}
|
||||
other => panic!("expected event trigger, got {other:?}"),
|
||||
}
|
||||
|
||||
rig.send_incoming(IncomingMessage::new(
|
||||
"telegram",
|
||||
"test-user",
|
||||
"bug: home button broken",
|
||||
))
|
||||
.await;
|
||||
|
||||
let runs = wait_for_routine_run(rig.database(), routine.id, TIMEOUT).await;
|
||||
assert_eq!(runs[0].trigger_type, "event");
|
||||
|
||||
let responses = rig.wait_for_responses(3, TIMEOUT).await;
|
||||
assert!(
|
||||
responses.iter().any(|response| {
|
||||
response
|
||||
.metadata
|
||||
.get("source")
|
||||
.and_then(|value| value.as_str())
|
||||
== Some("routine")
|
||||
&& response.content.contains("telegram-bug-watcher")
|
||||
&& response.content.contains("Bug report detected")
|
||||
}),
|
||||
"expected routine notification in responses: {responses:?}"
|
||||
);
|
||||
|
||||
rig.shutdown();
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// 6c. Event routine without channel filter still fires on Telegram
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn routine_event_trigger_without_channel_filter_still_fires() {
|
||||
let trace =
|
||||
LlmTrace::from_file(format!("{FIXTURES}/routine_event_any_channel.json")).unwrap();
|
||||
let rig = TestRigBuilder::new()
|
||||
.with_trace(trace)
|
||||
.with_routines()
|
||||
.with_auto_approve_tools(true)
|
||||
.build()
|
||||
.await;
|
||||
|
||||
rig.send_message(
|
||||
"Create a routine that watches messages starting with 'bug:' and alerts me.",
|
||||
)
|
||||
.await;
|
||||
let _ = rig.wait_for_responses(1, TIMEOUT).await;
|
||||
|
||||
let routine = rig
|
||||
.database()
|
||||
.get_routine_by_name("test-user", "any-channel-bug-watcher")
|
||||
.await
|
||||
.expect("get_routine_by_name")
|
||||
.expect("any-channel-bug-watcher should exist");
|
||||
|
||||
match &routine.trigger {
|
||||
Trigger::Event { channel, pattern } => {
|
||||
assert_eq!(channel, &None);
|
||||
assert_eq!(pattern, "^bug\\b");
|
||||
}
|
||||
other => panic!("expected event trigger, got {other:?}"),
|
||||
}
|
||||
|
||||
rig.send_incoming(IncomingMessage::new(
|
||||
"telegram",
|
||||
"test-user",
|
||||
"bug: login button broken",
|
||||
))
|
||||
.await;
|
||||
|
||||
let runs = wait_for_routine_run(rig.database(), routine.id, TIMEOUT).await;
|
||||
assert_eq!(runs[0].trigger_type, "event");
|
||||
|
||||
rig.shutdown();
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// 7. Prompt injection resilience
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
@@ -10,6 +10,8 @@ mod support;
|
||||
mod tests {
|
||||
use std::time::Duration;
|
||||
|
||||
use ironclaw::agent::routine::{RoutineAction, Trigger};
|
||||
|
||||
use crate::support::test_rig::TestRigBuilder;
|
||||
use crate::support::trace_llm::LlmTrace;
|
||||
|
||||
@@ -123,6 +125,39 @@ mod tests {
|
||||
"routine_list should succeed: {completed:?}"
|
||||
);
|
||||
|
||||
let routine = rig
|
||||
.database()
|
||||
.get_routine_by_name("test-user", "daily-check")
|
||||
.await
|
||||
.expect("get_routine_by_name")
|
||||
.expect("daily-check should exist");
|
||||
|
||||
match &routine.trigger {
|
||||
Trigger::Cron { schedule, timezone } => {
|
||||
assert_eq!(schedule, "0 0 9 * * *");
|
||||
assert_eq!(timezone.as_deref(), Some("America/New_York"));
|
||||
}
|
||||
other => panic!("expected cron trigger, got {other:?}"),
|
||||
}
|
||||
|
||||
match &routine.action {
|
||||
RoutineAction::Lightweight {
|
||||
context_paths,
|
||||
use_tools,
|
||||
max_tool_rounds,
|
||||
..
|
||||
} => {
|
||||
assert_eq!(context_paths, &vec!["context/priorities.md".to_string()]);
|
||||
assert!(*use_tools, "lightweight routine should keep use_tools=true");
|
||||
assert_eq!(*max_tool_rounds, 2);
|
||||
}
|
||||
other => panic!("expected lightweight action, got {other:?}"),
|
||||
}
|
||||
|
||||
assert_eq!(routine.notify.channel.as_deref(), Some("telegram"));
|
||||
assert_eq!(routine.notify.user, "ops-team");
|
||||
assert_eq!(routine.guardrails.cooldown.as_secs(), 600);
|
||||
|
||||
rig.shutdown();
|
||||
}
|
||||
|
||||
@@ -168,7 +203,48 @@ mod tests {
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Test 5: routine_history
|
||||
// Test 5: routine_manual_create
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
async fn routine_manual_create() {
|
||||
let trace = LlmTrace::from_file(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/tests/fixtures/llm_traces/tools/routine_manual_create.json"
|
||||
))
|
||||
.expect("failed to load routine_manual_create.json");
|
||||
|
||||
let rig = TestRigBuilder::new()
|
||||
.with_trace(trace.clone())
|
||||
.with_auto_approve_tools(true)
|
||||
.build()
|
||||
.await;
|
||||
|
||||
rig.send_message("Create a manual routine for bug triage")
|
||||
.await;
|
||||
let responses = rig.wait_for_responses(1, Duration::from_secs(15)).await;
|
||||
|
||||
rig.verify_trace_expects(&trace, &responses);
|
||||
|
||||
let routine = rig
|
||||
.database()
|
||||
.get_routine_by_name("test-user", "manual-triage")
|
||||
.await
|
||||
.expect("get_routine_by_name")
|
||||
.expect("manual-triage should exist");
|
||||
|
||||
assert!(matches!(routine.trigger, Trigger::Manual));
|
||||
assert!(
|
||||
matches!(&routine.action, RoutineAction::Lightweight { use_tools, .. } if !*use_tools),
|
||||
"manual routine should default to lightweight without tools: {:?}",
|
||||
routine.action
|
||||
);
|
||||
|
||||
rig.shutdown();
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Test 6: routine_history
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
@@ -205,7 +281,7 @@ mod tests {
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Test 6: routine_system_event_emit
|
||||
// Test 7: routine_system_event_emit
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
@@ -253,11 +329,47 @@ mod tests {
|
||||
emit_result.1
|
||||
);
|
||||
|
||||
let routine = rig
|
||||
.database()
|
||||
.get_routine_by_name("test-user", "gh-issue-emit-test")
|
||||
.await
|
||||
.expect("get_routine_by_name")
|
||||
.expect("gh-issue-emit-test should exist");
|
||||
|
||||
match &routine.trigger {
|
||||
Trigger::SystemEvent {
|
||||
source,
|
||||
event_type,
|
||||
filters,
|
||||
} => {
|
||||
assert_eq!(source, "github");
|
||||
assert_eq!(event_type, "issue.opened");
|
||||
assert_eq!(
|
||||
filters.get("repository").map(String::as_str),
|
||||
Some("nearai/ironclaw")
|
||||
);
|
||||
assert_eq!(filters.get("priority").map(String::as_str), Some("p1"));
|
||||
}
|
||||
other => panic!("expected system_event trigger, got {other:?}"),
|
||||
}
|
||||
|
||||
match &routine.action {
|
||||
RoutineAction::FullJob {
|
||||
description,
|
||||
tool_permissions,
|
||||
..
|
||||
} => {
|
||||
assert!(description.contains("Summarize the new issue"));
|
||||
assert_eq!(tool_permissions, &vec!["shell".to_string()]);
|
||||
}
|
||||
other => panic!("expected full_job action, got {other:?}"),
|
||||
}
|
||||
|
||||
rig.shutdown();
|
||||
}
|
||||
|
||||
// -----------------------------------------------------------------------
|
||||
// Test 7: skill_install_routine_webhook_sim
|
||||
// Test 8: skill_install_routine_webhook_sim
|
||||
// -----------------------------------------------------------------------
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
{
|
||||
"model_name": "advanced-routine-event-any-channel",
|
||||
"expects": {
|
||||
"tools_used": ["routine_create"],
|
||||
"all_tools_succeeded": true,
|
||||
"min_responses": 1
|
||||
},
|
||||
"steps": [
|
||||
{
|
||||
"response": {
|
||||
"type": "tool_calls",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_routine_create_event_any_channel",
|
||||
"name": "routine_create",
|
||||
"arguments": {
|
||||
"name": "any-channel-bug-watcher",
|
||||
"description": "Watch bug reports from any incoming channel.",
|
||||
"trigger_type": "event",
|
||||
"event_pattern": "^bug\\b",
|
||||
"prompt": "Summarize the bug report in one line."
|
||||
}
|
||||
}
|
||||
],
|
||||
"input_tokens": 130,
|
||||
"output_tokens": 38
|
||||
}
|
||||
},
|
||||
{
|
||||
"response": {
|
||||
"type": "text",
|
||||
"content": "Created the any-channel-bug-watcher routine for bug messages.",
|
||||
"input_tokens": 170,
|
||||
"output_tokens": 18
|
||||
}
|
||||
},
|
||||
{
|
||||
"response": {
|
||||
"type": "text",
|
||||
"content": "I saw the Telegram message.",
|
||||
"input_tokens": 90,
|
||||
"output_tokens": 12
|
||||
}
|
||||
},
|
||||
{
|
||||
"response": {
|
||||
"type": "text",
|
||||
"content": "Bug report detected: login button broken.",
|
||||
"input_tokens": 120,
|
||||
"output_tokens": 14
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
{
|
||||
"model_name": "advanced-routine-event-telegram",
|
||||
"expects": {
|
||||
"tools_used": ["routine_create"],
|
||||
"all_tools_succeeded": true,
|
||||
"min_responses": 1
|
||||
},
|
||||
"steps": [
|
||||
{
|
||||
"response": {
|
||||
"type": "tool_calls",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_routine_create_event_telegram",
|
||||
"name": "routine_create",
|
||||
"arguments": {
|
||||
"name": "telegram-bug-watcher",
|
||||
"description": "Watch Telegram bug reports and alert on them.",
|
||||
"trigger_type": "event",
|
||||
"event_channel": "telegram",
|
||||
"event_pattern": "^bug\\b",
|
||||
"prompt": "Summarize the bug report in one line."
|
||||
}
|
||||
}
|
||||
],
|
||||
"input_tokens": 140,
|
||||
"output_tokens": 40
|
||||
}
|
||||
},
|
||||
{
|
||||
"response": {
|
||||
"type": "text",
|
||||
"content": "Created the telegram-bug-watcher routine for Telegram bug messages.",
|
||||
"input_tokens": 180,
|
||||
"output_tokens": 20
|
||||
}
|
||||
},
|
||||
{
|
||||
"response": {
|
||||
"type": "text",
|
||||
"content": "I saw the Telegram message.",
|
||||
"input_tokens": 90,
|
||||
"output_tokens": 12
|
||||
}
|
||||
},
|
||||
{
|
||||
"response": {
|
||||
"type": "text",
|
||||
"content": "Bug report detected: home button broken.",
|
||||
"input_tokens": 120,
|
||||
"output_tokens": 14
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -18,8 +18,16 @@
|
||||
"name": "daily-check",
|
||||
"trigger_type": "cron",
|
||||
"schedule": "0 0 9 * * *",
|
||||
"timezone": "America/New_York",
|
||||
"prompt": "Check system status and report any issues.",
|
||||
"description": "Daily system health check"
|
||||
"description": "Daily system health check",
|
||||
"context_paths": ["context/priorities.md"],
|
||||
"action_type": "lightweight",
|
||||
"use_tools": true,
|
||||
"max_tool_rounds": 2,
|
||||
"cooldown_secs": 600,
|
||||
"notify_channel": "telegram",
|
||||
"notify_user": "ops-team"
|
||||
}
|
||||
}
|
||||
],
|
||||
|
||||
@@ -0,0 +1,36 @@
|
||||
{
|
||||
"model_name": "test-routine-manual-create",
|
||||
"expects": {
|
||||
"tools_used": ["routine_create"],
|
||||
"all_tools_succeeded": true,
|
||||
"min_responses": 1
|
||||
},
|
||||
"steps": [
|
||||
{
|
||||
"response": {
|
||||
"type": "tool_calls",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call_rc_manual_1",
|
||||
"name": "routine_create",
|
||||
"arguments": {
|
||||
"name": "manual-triage",
|
||||
"trigger_type": "manual",
|
||||
"prompt": "Summarize the latest bug reports when this routine is fired."
|
||||
}
|
||||
}
|
||||
],
|
||||
"input_tokens": 90,
|
||||
"output_tokens": 22
|
||||
}
|
||||
},
|
||||
{
|
||||
"response": {
|
||||
"type": "text",
|
||||
"content": "Created the manual-triage routine. It will only run when explicitly fired.",
|
||||
"input_tokens": 140,
|
||||
"output_tokens": 18
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -21,7 +21,12 @@
|
||||
"trigger_type": "system_event",
|
||||
"event_source": "github",
|
||||
"event_type": "issue.opened",
|
||||
"event_filters": {
|
||||
"repository": "nearai/ironclaw",
|
||||
"priority": "p1"
|
||||
},
|
||||
"action_type": "full_job",
|
||||
"tool_permissions": ["shell"],
|
||||
"prompt": "Summarize the new issue and propose next steps."
|
||||
}
|
||||
}
|
||||
@@ -42,6 +47,7 @@
|
||||
"event_type": "issue.opened",
|
||||
"payload": {
|
||||
"repository": "nearai/ironclaw",
|
||||
"priority": "p1",
|
||||
"issue_number": 123,
|
||||
"title": "Support event-driven project workflow"
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user