Skip to main content

ironflow_core/providers/http/
adapter.rs

1//! Core adapter trait and generic provider wrapper for HTTP-based LLM APIs.
2
3use std::sync::Arc;
4use std::time::{Duration, Instant};
5
6use futures_util::future::join_all;
7use reqwest::Client;
8use serde_json::{Value, json};
9use tokio::sync::Semaphore;
10use tracing::{debug, info, warn};
11
12use crate::error::AgentError;
13use crate::provider::{
14    AgentConfig, AgentOutput, AgentProvider, DebugMessage, DebugToolCall, DebugToolResult,
15    InvokeFuture, LogSink, ToolProfile,
16};
17use crate::providers::http::sse::{SseDelta, collect_sse_stream};
18use crate::providers::http::tools::ToolRegistry;
19use crate::providers::http::tools::profiles::ToolProfiles;
20use crate::providers::http::tools::routing::route_tool_call;
21
22/// Normalized result of one API turn (one HTTP request/response cycle).
23#[derive(Debug)]
24pub struct TurnResult {
25    /// Free-form text content from the model.
26    pub text: Option<String>,
27    /// Tool calls requested by the model in this turn (unused in V1 - no tool execution).
28    #[allow(dead_code)]
29    pub tool_calls: Vec<HttpToolCall>,
30    /// Whether this is the final turn.
31    pub is_final: bool,
32    /// Extracted structured JSON value when a schema was requested.
33    pub structured_value: Option<Value>,
34    /// Token usage reported by the provider.
35    pub usage: HttpUsage,
36    /// Concrete model identifier returned by the provider.
37    pub model: Option<String>,
38}
39
40/// A single tool call requested by the model.
41#[derive(Debug, Clone)]
42#[allow(dead_code)]
43pub struct HttpToolCall {
44    /// Provider-assigned call identifier.
45    pub id: String,
46    /// Tool name.
47    pub name: String,
48    /// Input arguments as JSON.
49    pub input: Value,
50}
51
52/// Token usage from a single turn.
53#[derive(Debug, Default)]
54pub struct HttpUsage {
55    /// Uncached input/prompt tokens consumed (excludes cache reads and writes).
56    pub input_tokens: Option<u64>,
57    /// Input tokens served from the prompt cache.
58    pub cache_read_input_tokens: Option<u64>,
59    /// Input tokens written to the prompt cache.
60    pub cache_creation_input_tokens: Option<u64>,
61    /// Output/completion tokens generated.
62    pub output_tokens: Option<u64>,
63}
64
65/// Internal trait implemented by each HTTP provider backend.
66///
67/// The generic [`HttpAgentProvider`] calls these methods to build requests,
68/// parse responses, and configure authentication. The agentic loop, retry,
69/// and timeout are handled by the wrapper.
70pub trait HttpAgentAdapter: Send + Sync + 'static {
71    /// Provider name for logging and errors.
72    fn provider_name(&self) -> &'static str;
73
74    /// Full endpoint URL for the given model.
75    fn endpoint_url(&self, model: &str) -> String;
76
77    /// Authentication and provider-specific headers.
78    fn auth_headers(&self) -> Vec<(String, String)>;
79
80    /// Build the initial JSON request body from an [`AgentConfig`].
81    fn build_request(&self, config: &AgentConfig) -> Result<Value, AgentError>;
82
83    /// Parse a non-streaming response body into a [`TurnResult`].
84    fn parse_response(&self, body: &Value, config: &AgentConfig) -> Result<TurnResult, AgentError>;
85
86    /// Parse a single SSE `data:` line into a streaming delta.
87    fn parse_sse_line(&self, line: &str) -> Option<SseDelta>;
88
89    /// Fold accumulated SSE deltas into a complete [`TurnResult`].
90    fn fold_sse_deltas(
91        &self,
92        deltas: Vec<SseDelta>,
93        config: &AgentConfig,
94    ) -> Result<TurnResult, AgentError>;
95
96    /// Compute cost in USD from the token usage of one turn, pricing
97    /// uncached input, cache reads, cache writes and output separately.
98    /// Returns `None` if unknown.
99    fn compute_cost(&self, model: &str, usage: &HttpUsage) -> Option<f64>;
100
101    /// Resolve model alias (e.g. "sonnet") to a provider-specific model ID.
102    fn resolve_model(&self, model: &str) -> String;
103}
104
105/// Default timeout for HTTP provider requests.
106const DEFAULT_TIMEOUT: Duration = Duration::from_secs(120);
107
108/// Generic HTTP provider that wraps any [`HttpAgentAdapter`].
109///
110/// Implements [`AgentProvider`] by delegating request construction and response
111/// parsing to the adapter while handling the HTTP transport, timeout, and
112/// agentic execution loop.
113///
114/// When a [`ToolRegistry`] is attached via [`with_tools`](Self::with_tools),
115/// the provider runs a multi-turn agentic loop: executing tool calls locally
116/// and feeding results back to the model until it produces a final response
117/// (or hits `max_turns` / `max_budget_usd` limits).
118///
119/// Without a registry, the provider behaves as single-turn (backward-compatible).
120///
121/// Steps that need different tools pick a named profile registered with
122/// [`with_tool_profile`](Self::with_tool_profile).
123pub struct HttpAgentProvider<A: HttpAgentAdapter> {
124    adapter: A,
125    client: Client,
126    timeout: Duration,
127    tools: ToolProfiles,
128}
129
130impl<A: HttpAgentAdapter> HttpAgentProvider<A> {
131    /// Create a new HTTP provider with the given adapter.
132    pub fn new(adapter: A) -> Self {
133        let client = Client::builder()
134            .timeout(DEFAULT_TIMEOUT)
135            .build()
136            .expect("failed to build reqwest client");
137        Self {
138            adapter,
139            client,
140            timeout: DEFAULT_TIMEOUT,
141            tools: ToolProfiles::default(),
142        }
143    }
144
145    /// Attach the default tool registry to enable multi-turn agentic execution.
146    ///
147    /// The default tools go to every step that selects no tool profile. When
148    /// tools are selected, the provider will:
149    /// 1. Include the tools in every request (OpenAI `tools` format).
150    /// 2. Execute tool calls returned by the model.
151    /// 3. Loop until the model produces a final response or limits are hit.
152    pub fn with_tools(mut self, registry: ToolRegistry) -> Self {
153        self.tools.set_default(registry);
154        self
155    }
156
157    /// Register a named tool profile that steps select with
158    /// [`AgentConfig::tool_profile`]. Call it once per profile.
159    ///
160    /// A step only sees the tools of the profile it selects, and only those
161    /// can run. To share an MCP server between profiles without opening it
162    /// twice, see `register_shared_mcp_tools` (feature `tool-mcp`).
163    ///
164    /// # Panics
165    ///
166    /// Panics if `profile` is already registered.
167    ///
168    /// # Examples
169    ///
170    /// ```
171    /// use ironflow_core::provider::ToolProfile;
172    /// use ironflow_core::providers::http::HttpAgentProvider;
173    /// use ironflow_core::providers::http::adapter::HttpAgentAdapter;
174    /// use ironflow_core::providers::http::tools::ToolRegistry;
175    ///
176    /// const SUGGESTION: ToolProfile = ToolProfile::new("suggestion");
177    /// const BUG: ToolProfile = ToolProfile::new("bug");
178    ///
179    /// fn with_profiles<A: HttpAgentAdapter>(provider: HttpAgentProvider<A>) -> HttpAgentProvider<A> {
180    ///     provider
181    ///         .with_tool_profile(SUGGESTION, ToolRegistry::new())
182    ///         .with_tool_profile(BUG, ToolRegistry::new())
183    /// }
184    /// ```
185    pub fn with_tool_profile(mut self, profile: ToolProfile, registry: ToolRegistry) -> Self {
186        self.tools.insert(profile, registry);
187        self
188    }
189
190    /// Override the request timeout.
191    pub fn with_timeout(mut self, timeout: Duration) -> Self {
192        self.timeout = timeout;
193        self.client = Client::builder()
194            .timeout(timeout)
195            .build()
196            .expect("failed to build reqwest client");
197        self
198    }
199
200    async fn execute_turn(
201        &self,
202        request_body: &Value,
203        config: &AgentConfig,
204    ) -> Result<TurnResult, AgentError> {
205        let model = self.adapter.resolve_model(&config.model);
206        let url = self.adapter.endpoint_url(&model);
207        let headers = self.adapter.auth_headers();
208
209        let mut req = self.client.post(&url).json(request_body);
210        for (key, value) in &headers {
211            req = req.header(key, value);
212        }
213        if let Some(ref ctx) = config.trace_context {
214            req = req.header("traceparent", ctx.to_traceparent());
215        }
216
217        let response = tokio::time::timeout(self.timeout, req.send())
218            .await
219            .map_err(|_| AgentError::Timeout {
220                limit: self.timeout,
221            })?
222            .map_err(|e| {
223                if e.is_timeout() {
224                    AgentError::Timeout {
225                        limit: self.timeout,
226                    }
227                } else {
228                    AgentError::HttpProvider {
229                        provider: self.adapter.provider_name().to_string(),
230                        status_code: 0,
231                        message: format!("connection failed: {e}"),
232                    }
233                }
234            })?;
235
236        let status = response.status().as_u16();
237
238        if status == 429 {
239            let retry_after = response
240                .headers()
241                .get("retry-after")
242                .and_then(|v| v.to_str().ok())
243                .and_then(|v| v.parse::<u64>().ok());
244            return Err(AgentError::RateLimited {
245                provider: self.adapter.provider_name().to_string(),
246                retry_after_secs: retry_after,
247            });
248        }
249
250        if status >= 400 {
251            let body_text = response.text().await.unwrap_or_default();
252            let message = serde_json::from_str::<Value>(&body_text)
253                .ok()
254                .and_then(|v| {
255                    v.get("error")
256                        .and_then(|e| e.get("message"))
257                        .and_then(|m| m.as_str())
258                        .map(String::from)
259                })
260                .unwrap_or(body_text);
261            return Err(AgentError::HttpProvider {
262                provider: self.adapter.provider_name().to_string(),
263                status_code: status,
264                message,
265            });
266        }
267
268        if config.verbose {
269            let deltas = collect_sse_stream(&self.adapter, response, self.timeout).await?;
270            self.adapter.fold_sse_deltas(deltas, config)
271        } else {
272            let body: Value = response
273                .json()
274                .await
275                .map_err(|e| AgentError::HttpProvider {
276                    provider: self.adapter.provider_name().to_string(),
277                    status_code: 0,
278                    message: format!("failed to parse response JSON: {e}"),
279                })?;
280            self.adapter.parse_response(&body, config)
281        }
282    }
283}
284
285/// Accumulates usage across turns and builds the final [`AgentOutput`].
286struct LoopState {
287    start: Instant,
288    total_input_tokens: u64,
289    total_cache_read_tokens: Option<u64>,
290    total_cache_creation_tokens: Option<u64>,
291    total_output_tokens: u64,
292    total_cost: f64,
293    model_name: Option<String>,
294    debug_messages: Vec<DebugMessage>,
295    verbose: bool,
296}
297
298impl LoopState {
299    fn new(start: Instant, verbose: bool) -> Self {
300        Self {
301            start,
302            total_input_tokens: 0,
303            total_cache_read_tokens: None,
304            total_cache_creation_tokens: None,
305            total_output_tokens: 0,
306            total_cost: 0.0,
307            model_name: None,
308            debug_messages: Vec::new(),
309            verbose,
310        }
311    }
312
313    fn into_output(self, value: Value) -> AgentOutput {
314        AgentOutput {
315            value,
316            session_id: None,
317            cost_usd: if self.total_cost > 0.0 {
318                Some(self.total_cost)
319            } else {
320                None
321            },
322            input_tokens: Some(self.total_input_tokens),
323            cache_read_input_tokens: self.total_cache_read_tokens,
324            cache_creation_input_tokens: self.total_cache_creation_tokens,
325            output_tokens: Some(self.total_output_tokens),
326            model: self.model_name,
327            duration_ms: self.start.elapsed().as_millis() as u64,
328            debug_messages: if self.verbose {
329                Some(self.debug_messages)
330            } else {
331                None
332            },
333            account_id: None,
334            environment_id: None,
335        }
336    }
337}
338
339/// Extract the final value from a turn result (structured or text).
340fn extract_value(turn_result: &TurnResult) -> Value {
341    if let Some(ref structured) = turn_result.structured_value {
342        structured.clone()
343    } else {
344        turn_result
345            .text
346            .as_ref()
347            .map(|t| Value::String(t.clone()))
348            .unwrap_or(Value::String(String::new()))
349    }
350}
351
352/// Extract the text value from a turn result (ignoring structured).
353fn extract_text_value(turn_result: &TurnResult) -> Value {
354    turn_result
355        .text
356        .as_ref()
357        .map(|t| Value::String(t.clone()))
358        .unwrap_or(Value::String(String::new()))
359}
360
361/// Execute a single tool call and return its `(content, is_error)` result.
362///
363/// Mirrors the model's routing rules: MCP-prefixed names are resolved
364/// through `route_tool_call` when connectors are registered, otherwise the
365/// name is looked up directly. An unknown tool or a routing failure is
366/// reported back to the model as an error, never aborts the turn.
367async fn execute_tool_call(
368    tc: &HttpToolCall,
369    registry: &ToolRegistry,
370    provider_name: &'static str,
371) -> (String, bool) {
372    debug!(
373        provider = provider_name,
374        tool = %tc.name,
375        call_id = %tc.id,
376        "executing tool call"
377    );
378
379    let connectors = registry.connectors();
380    let registry_key = if connectors.is_empty() {
381        tc.name.clone()
382    } else {
383        match route_tool_call(&tc.name, connectors) {
384            Ok(routed) => routed.registry_key,
385            Err(routing_err) => return (routing_err.to_string(), true),
386        }
387    };
388
389    match registry.execute(&registry_key, tc.input.clone()).await {
390        Some(Ok(output)) => (output.content, output.is_error),
391        Some(Err(err)) => (format!("Tool execution error: {err}"), true),
392        None => (format!("Unknown tool: {}", tc.name), true),
393    }
394}
395
396/// Execute every tool call of one turn.
397///
398/// Calls are processed in the model's order. A maximal run of consecutive
399/// read-only calls (see [`Tool::read_only`](crate::providers::http::tools::Tool::read_only))
400/// is executed concurrently via `join_all`, bounded by `max_parallel` permits
401/// on a [`Semaphore`]. A non-read-only call is a barrier: it waits for the
402/// previous group to finish, runs alone, and the next group only starts once
403/// it completes. An unknown tool name (not present in the registry) is
404/// treated as non-read-only.
405///
406/// Returns `(content, is_error)` pairs in the same order as `tool_calls`,
407/// regardless of completion order within a parallel group.
408///
409/// # Panics
410///
411/// Panics if the internal semaphore is closed, which cannot happen since it
412/// is never closed.
413async fn execute_turn_tool_calls(
414    tool_calls: &[HttpToolCall],
415    registry: &ToolRegistry,
416    provider_name: &'static str,
417    max_parallel: usize,
418) -> Vec<(String, bool)> {
419    let max_parallel = max_parallel.max(1);
420    let mut results = Vec::with_capacity(tool_calls.len());
421    let mut idx = 0;
422
423    while idx < tool_calls.len() {
424        if registry.is_read_only(&tool_calls[idx].name) {
425            let end = tool_calls[idx..]
426                .iter()
427                .position(|tc| !registry.is_read_only(&tc.name))
428                .map(|offset| idx + offset)
429                .unwrap_or(tool_calls.len());
430
431            let semaphore = Arc::new(Semaphore::new(max_parallel));
432            let group_results = join_all(tool_calls[idx..end].iter().map(|tc| {
433                let semaphore = Arc::clone(&semaphore);
434                async move {
435                    let _permit = semaphore
436                        .acquire()
437                        .await
438                        .expect("semaphore closed unexpectedly");
439                    execute_tool_call(tc, registry, provider_name).await
440                }
441            }))
442            .await;
443
444            results.extend(group_results);
445            idx = end;
446        } else {
447            results.push(execute_tool_call(&tool_calls[idx], registry, provider_name).await);
448            idx += 1;
449        }
450    }
451
452    results
453}
454
455impl<A: HttpAgentAdapter> AgentProvider for HttpAgentProvider<A> {
456    fn invoke<'a>(&'a self, config: &'a AgentConfig) -> InvokeFuture<'a> {
457        Box::pin(async move {
458            // Resolved before any request: an unknown profile never reaches the model.
459            let selection = self.tools.select(config.tool_profile.as_ref())?;
460            if !self.tools.is_empty() {
461                info!(
462                    provider = self.adapter.provider_name(),
463                    profile = %selection.label(),
464                    "{}",
465                    selection.describe()
466                );
467            }
468            let tool_registry = selection.registry();
469            let mut request_body = self.adapter.build_request(config)?;
470
471            // Inject tools into the request if a registry is available
472            if let Some(registry) = tool_registry
473                && !registry.is_empty()
474            {
475                let tools_array = registry.to_openai_tools();
476                request_body["tools"] = Value::Array(tools_array);
477            }
478
479            let max_turns = config.max_turns.unwrap_or(25) as usize;
480            let max_budget = config.max_budget_usd.unwrap_or(f64::MAX);
481            let mut state = LoopState::new(Instant::now(), config.verbose);
482
483            // Messages array for multi-turn
484            let mut messages: Vec<Value> = request_body
485                .get("messages")
486                .and_then(|m| m.as_array())
487                .cloned()
488                .unwrap_or_default();
489
490            for turn in 0..max_turns {
491                request_body["messages"] = Value::Array(messages.clone());
492                let turn_result = self.execute_turn(&request_body, config).await?;
493
494                // Accumulate usage
495                let turn_input = turn_result.usage.input_tokens.unwrap_or(0);
496                let turn_output = turn_result.usage.output_tokens.unwrap_or(0);
497                state.total_input_tokens += turn_input;
498                state.total_output_tokens += turn_output;
499                if let Some(v) = turn_result.usage.cache_read_input_tokens {
500                    state.total_cache_read_tokens =
501                        Some(state.total_cache_read_tokens.unwrap_or(0) + v);
502                }
503                if let Some(v) = turn_result.usage.cache_creation_input_tokens {
504                    state.total_cache_creation_tokens =
505                        Some(state.total_cache_creation_tokens.unwrap_or(0) + v);
506                }
507
508                if state.model_name.is_none() {
509                    state.model_name = turn_result.model.clone();
510                }
511
512                if let Some(ref model) = state.model_name
513                    && let Some(turn_cost) = self.adapter.compute_cost(model, &turn_result.usage)
514                {
515                    state.total_cost += turn_cost;
516                }
517
518                // Record debug trace for this turn
519                if config.verbose {
520                    let tool_calls_debug: Vec<DebugToolCall> = turn_result
521                        .tool_calls
522                        .iter()
523                        .map(|tc| DebugToolCall {
524                            id: Some(tc.id.clone()),
525                            name: tc.name.clone(),
526                            input: tc.input.clone(),
527                        })
528                        .collect();
529
530                    state.debug_messages.push(DebugMessage {
531                        text: turn_result.text.clone(),
532                        thinking: None,
533                        thinking_redacted: false,
534                        tool_calls: tool_calls_debug,
535                        tool_results: Vec::new(),
536                        stop_reason: if turn_result.is_final {
537                            Some("end_turn".to_string())
538                        } else {
539                            Some("tool_use".to_string())
540                        },
541                        input_tokens: Some(turn_input),
542                        output_tokens: Some(turn_output),
543                    });
544                }
545
546                // Final response (no tool calls) -> return
547                if turn_result.is_final || turn_result.tool_calls.is_empty() {
548                    info!(
549                        provider = self.adapter.provider_name(),
550                        turns = turn + 1,
551                        duration_ms = state.start.elapsed().as_millis() as u64,
552                        input_tokens = state.total_input_tokens,
553                        cache_read_input_tokens = state.total_cache_read_tokens,
554                        cache_creation_input_tokens = state.total_cache_creation_tokens,
555                        output_tokens = state.total_output_tokens,
556                        "invocation complete"
557                    );
558                    return Ok(state.into_output(extract_value(&turn_result)));
559                }
560
561                // Tool calls but no registry -> return text (backward compat)
562                let registry = match tool_registry {
563                    Some(r) => r,
564                    None => {
565                        warn!(
566                            provider = self.adapter.provider_name(),
567                            tool_calls = turn_result.tool_calls.len(),
568                            "model requested tool calls but no registry attached, returning text"
569                        );
570                        return Ok(state.into_output(extract_text_value(&turn_result)));
571                    }
572                };
573
574                // Budget exceeded -> stop
575                if state.total_cost >= max_budget {
576                    warn!(
577                        provider = self.adapter.provider_name(),
578                        cost = state.total_cost,
579                        budget = max_budget,
580                        "budget exceeded, stopping agentic loop"
581                    );
582                    return Ok(state.into_output(extract_text_value(&turn_result)));
583                }
584
585                // Build assistant message with tool_calls for conversation history
586                let assistant_tool_calls: Vec<Value> = turn_result
587                    .tool_calls
588                    .iter()
589                    .map(|tc| {
590                        json!({
591                            "id": tc.id,
592                            "type": "function",
593                            "function": {
594                                "name": tc.name,
595                                "arguments": tc.input.to_string()
596                            }
597                        })
598                    })
599                    .collect();
600
601                let mut assistant_msg = json!({"role": "assistant"});
602                if let Some(ref text) = turn_result.text {
603                    assistant_msg["content"] = Value::String(text.clone());
604                } else {
605                    assistant_msg["content"] = Value::Null;
606                }
607                assistant_msg["tool_calls"] = Value::Array(assistant_tool_calls);
608                messages.push(assistant_msg);
609
610                // Execute tool calls (consecutive read-only calls run concurrently)
611                let max_parallel = config.max_parallel_tools.max(1);
612                let results = execute_turn_tool_calls(
613                    &turn_result.tool_calls,
614                    registry,
615                    self.adapter.provider_name(),
616                    max_parallel,
617                )
618                .await;
619
620                let mut tool_results_debug: Vec<DebugToolResult> = Vec::new();
621                for (tc, (content, is_error)) in turn_result.tool_calls.iter().zip(results) {
622                    messages.push(json!({
623                        "role": "tool",
624                        "tool_call_id": tc.id,
625                        "content": content
626                    }));
627
628                    if config.verbose {
629                        tool_results_debug.push(DebugToolResult {
630                            tool_use_id: Some(tc.id.clone()),
631                            content: Value::String(content.clone()),
632                            is_error,
633                        });
634                    }
635                }
636
637                if config.verbose
638                    && let Some(last_msg) = state.debug_messages.last_mut()
639                {
640                    last_msg.tool_results = tool_results_debug;
641                }
642
643                info!(
644                    provider = self.adapter.provider_name(),
645                    turn = turn + 1,
646                    tools_executed = turn_result.tool_calls.len(),
647                    "turn complete, continuing loop"
648                );
649            }
650
651            warn!(
652                provider = self.adapter.provider_name(),
653                max_turns, "max turns reached, returning last state"
654            );
655            Ok(state.into_output(Value::String(String::new())))
656        })
657    }
658
659    /// Same as [`invoke`](AgentProvider::invoke), and records on `log_sink`
660    /// the tool profile of the step with the tools it exposes.
661    fn invoke_with_logs<'a>(
662        &'a self,
663        config: &'a AgentConfig,
664        log_sink: Arc<dyn LogSink>,
665    ) -> InvokeFuture<'a> {
666        // An unknown profile is not logged here: `invoke` fails with it.
667        if !self.tools.is_empty()
668            && let Ok(selection) = self.tools.select(config.tool_profile.as_ref())
669        {
670            log_sink.log("system", &selection.describe());
671        }
672        self.invoke(config)
673    }
674}
675
676#[cfg(test)]
677mod tests {
678    use std::future::Future;
679    use std::pin::Pin;
680    use std::sync::Mutex;
681    use std::sync::atomic::{AtomicUsize, Ordering};
682
683    use tokio::time::sleep;
684
685    use crate::providers::http::tools::{Tool, ToolError, ToolOutput};
686
687    use super::*;
688
689    type CallLog = Arc<Mutex<Vec<(String, Instant, Instant)>>>;
690
691    struct DelayTool {
692        name: String,
693        read_only: bool,
694        delay: Duration,
695        concurrent: Arc<AtomicUsize>,
696        max_concurrent: Arc<AtomicUsize>,
697        log: Option<CallLog>,
698    }
699
700    impl DelayTool {
701        fn new(name: &str, read_only: bool, delay_ms: u64) -> Self {
702            Self {
703                name: name.to_string(),
704                read_only,
705                delay: Duration::from_millis(delay_ms),
706                concurrent: Arc::new(AtomicUsize::new(0)),
707                max_concurrent: Arc::new(AtomicUsize::new(0)),
708                log: None,
709            }
710        }
711
712        fn with_counters(mut self, concurrent: &Arc<AtomicUsize>, max: &Arc<AtomicUsize>) -> Self {
713            self.concurrent = Arc::clone(concurrent);
714            self.max_concurrent = Arc::clone(max);
715            self
716        }
717
718        fn with_log(mut self, log: &CallLog) -> Self {
719            self.log = Some(Arc::clone(log));
720            self
721        }
722    }
723
724    impl Tool for DelayTool {
725        fn name(&self) -> &str {
726            &self.name
727        }
728
729        fn description(&self) -> &str {
730            "Sleeps then returns its name"
731        }
732
733        fn parameters_schema(&self) -> Value {
734            json!({"type": "object", "properties": {}})
735        }
736
737        fn read_only(&self) -> bool {
738            self.read_only
739        }
740
741        fn execute(
742            &self,
743            _input: Value,
744        ) -> Pin<Box<dyn Future<Output = Result<ToolOutput, ToolError>> + Send + '_>> {
745            Box::pin(async move {
746                let start = Instant::now();
747                let now = self.concurrent.fetch_add(1, Ordering::SeqCst) + 1;
748                self.max_concurrent.fetch_max(now, Ordering::SeqCst);
749                sleep(self.delay).await;
750                self.concurrent.fetch_sub(1, Ordering::SeqCst);
751                let end = Instant::now();
752                if let Some(ref log) = self.log {
753                    log.lock()
754                        .expect("log mutex poisoned")
755                        .push((self.name.clone(), start, end));
756                }
757                Ok(ToolOutput::success(self.name.clone()))
758            })
759        }
760    }
761
762    fn call(id: &str, name: &str) -> HttpToolCall {
763        HttpToolCall {
764            id: id.to_string(),
765            name: name.to_string(),
766            input: json!({}),
767        }
768    }
769
770    fn entry(log: &CallLog, name: &str) -> (Instant, Instant) {
771        let entries = log.lock().expect("log mutex poisoned");
772        let (_, start, end) = entries
773            .iter()
774            .find(|(n, _, _)| n == name)
775            .unwrap_or_else(|| panic!("no log entry for {name}"));
776        (*start, *end)
777    }
778
779    #[tokio::test]
780    async fn parallel_read_only_tools() {
781        let registry = ToolRegistry::new()
782            .register(DelayTool::new("read_a", true, 200))
783            .register(DelayTool::new("read_b", true, 200));
784        let calls = vec![call("1", "read_a"), call("2", "read_b")];
785
786        let started = Instant::now();
787        let results = execute_turn_tool_calls(&calls, &registry, "test", 4).await;
788        let elapsed = started.elapsed();
789
790        assert!(
791            elapsed < Duration::from_millis(350),
792            "read-only calls should run concurrently, took {elapsed:?}"
793        );
794        assert_eq!(results.len(), 2);
795        assert!(results.iter().all(|(_, is_error)| !is_error));
796    }
797
798    #[tokio::test]
799    async fn write_tool_is_barrier() {
800        let log: CallLog = Arc::new(Mutex::new(Vec::new()));
801        let registry = ToolRegistry::new()
802            .register(DelayTool::new("read1", true, 100).with_log(&log))
803            .register(DelayTool::new("write", false, 100).with_log(&log))
804            .register(DelayTool::new("read2", true, 100).with_log(&log));
805        let calls = vec![call("1", "read1"), call("2", "write"), call("3", "read2")];
806
807        let results = execute_turn_tool_calls(&calls, &registry, "test", 4).await;
808        assert_eq!(results.len(), 3);
809
810        let (_, read1_end) = entry(&log, "read1");
811        let (write_start, write_end) = entry(&log, "write");
812        let (read2_start, _) = entry(&log, "read2");
813
814        assert!(write_start >= read1_end, "write must wait for read1");
815        assert!(read2_start >= write_end, "read2 must wait for write");
816    }
817
818    #[tokio::test]
819    async fn tool_results_keep_call_order() {
820        let registry = ToolRegistry::new()
821            .register(DelayTool::new("a", true, 150))
822            .register(DelayTool::new("b", true, 50))
823            .register(DelayTool::new("c", true, 100));
824        let calls = vec![call("1", "a"), call("2", "b"), call("3", "c")];
825
826        let results = execute_turn_tool_calls(&calls, &registry, "test", 4).await;
827
828        assert_eq!(
829            results,
830            vec![
831                ("a".to_string(), false),
832                ("b".to_string(), false),
833                ("c".to_string(), false),
834            ]
835        );
836    }
837
838    #[tokio::test]
839    async fn max_parallel_tools_one_is_sequential() {
840        let cur = Arc::new(AtomicUsize::new(0));
841        let peak = Arc::new(AtomicUsize::new(0));
842        let registry = ToolRegistry::new()
843            .register(DelayTool::new("a", true, 50).with_counters(&cur, &peak))
844            .register(DelayTool::new("b", true, 50).with_counters(&cur, &peak))
845            .register(DelayTool::new("c", true, 50).with_counters(&cur, &peak));
846        let calls = vec![call("1", "a"), call("2", "b"), call("3", "c")];
847
848        let results = execute_turn_tool_calls(&calls, &registry, "test", 1).await;
849
850        assert_eq!(results.len(), 3);
851        assert_eq!(peak.load(Ordering::SeqCst), 1);
852    }
853
854    #[tokio::test]
855    async fn read_only_group_respects_max_parallel() {
856        let cur = Arc::new(AtomicUsize::new(0));
857        let peak = Arc::new(AtomicUsize::new(0));
858        let registry = ToolRegistry::new()
859            .register(DelayTool::new("a", true, 50).with_counters(&cur, &peak))
860            .register(DelayTool::new("b", true, 50).with_counters(&cur, &peak))
861            .register(DelayTool::new("c", true, 50).with_counters(&cur, &peak));
862        let calls = vec![call("1", "a"), call("2", "b"), call("3", "c")];
863
864        let results = execute_turn_tool_calls(&calls, &registry, "test", 2).await;
865
866        assert_eq!(results.len(), 3);
867        assert_eq!(peak.load(Ordering::SeqCst), 2);
868    }
869
870    #[tokio::test]
871    async fn max_parallel_zero_is_floored_to_one() {
872        let registry = ToolRegistry::new().register(DelayTool::new("a", true, 10));
873        let calls = vec![call("1", "a")];
874
875        let results = execute_turn_tool_calls(&calls, &registry, "test", 0).await;
876
877        assert_eq!(results, vec![("a".to_string(), false)]);
878    }
879
880    #[tokio::test]
881    async fn unknown_tool_call_is_treated_as_barrier() {
882        let registry = ToolRegistry::new();
883        let calls = vec![call("1", "missing")];
884
885        let results = execute_turn_tool_calls(&calls, &registry, "test", 4).await;
886
887        assert_eq!(results, vec![("Unknown tool: missing".to_string(), true)]);
888    }
889
890    #[tokio::test]
891    async fn empty_tool_calls_return_empty_results() {
892        let registry = ToolRegistry::new();
893
894        let results = execute_turn_tool_calls(&[], &registry, "test", 4).await;
895
896        assert!(results.is_empty());
897    }
898}