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        }
334    }
335}
336
337/// Extract the final value from a turn result (structured or text).
338fn extract_value(turn_result: &TurnResult) -> Value {
339    if let Some(ref structured) = turn_result.structured_value {
340        structured.clone()
341    } else {
342        turn_result
343            .text
344            .as_ref()
345            .map(|t| Value::String(t.clone()))
346            .unwrap_or(Value::String(String::new()))
347    }
348}
349
350/// Extract the text value from a turn result (ignoring structured).
351fn extract_text_value(turn_result: &TurnResult) -> Value {
352    turn_result
353        .text
354        .as_ref()
355        .map(|t| Value::String(t.clone()))
356        .unwrap_or(Value::String(String::new()))
357}
358
359/// Execute a single tool call and return its `(content, is_error)` result.
360///
361/// Mirrors the model's routing rules: MCP-prefixed names are resolved
362/// through `route_tool_call` when connectors are registered, otherwise the
363/// name is looked up directly. An unknown tool or a routing failure is
364/// reported back to the model as an error, never aborts the turn.
365async fn execute_tool_call(
366    tc: &HttpToolCall,
367    registry: &ToolRegistry,
368    provider_name: &'static str,
369) -> (String, bool) {
370    debug!(
371        provider = provider_name,
372        tool = %tc.name,
373        call_id = %tc.id,
374        "executing tool call"
375    );
376
377    let connectors = registry.connectors();
378    let registry_key = if connectors.is_empty() {
379        tc.name.clone()
380    } else {
381        match route_tool_call(&tc.name, connectors) {
382            Ok(routed) => routed.registry_key,
383            Err(routing_err) => return (routing_err.to_string(), true),
384        }
385    };
386
387    match registry.execute(&registry_key, tc.input.clone()).await {
388        Some(Ok(output)) => (output.content, output.is_error),
389        Some(Err(err)) => (format!("Tool execution error: {err}"), true),
390        None => (format!("Unknown tool: {}", tc.name), true),
391    }
392}
393
394/// Execute every tool call of one turn.
395///
396/// Calls are processed in the model's order. A maximal run of consecutive
397/// read-only calls (see [`Tool::read_only`](crate::providers::http::tools::Tool::read_only))
398/// is executed concurrently via `join_all`, bounded by `max_parallel` permits
399/// on a [`Semaphore`]. A non-read-only call is a barrier: it waits for the
400/// previous group to finish, runs alone, and the next group only starts once
401/// it completes. An unknown tool name (not present in the registry) is
402/// treated as non-read-only.
403///
404/// Returns `(content, is_error)` pairs in the same order as `tool_calls`,
405/// regardless of completion order within a parallel group.
406///
407/// # Panics
408///
409/// Panics if the internal semaphore is closed, which cannot happen since it
410/// is never closed.
411async fn execute_turn_tool_calls(
412    tool_calls: &[HttpToolCall],
413    registry: &ToolRegistry,
414    provider_name: &'static str,
415    max_parallel: usize,
416) -> Vec<(String, bool)> {
417    let max_parallel = max_parallel.max(1);
418    let mut results = Vec::with_capacity(tool_calls.len());
419    let mut idx = 0;
420
421    while idx < tool_calls.len() {
422        if registry.is_read_only(&tool_calls[idx].name) {
423            let end = tool_calls[idx..]
424                .iter()
425                .position(|tc| !registry.is_read_only(&tc.name))
426                .map(|offset| idx + offset)
427                .unwrap_or(tool_calls.len());
428
429            let semaphore = Arc::new(Semaphore::new(max_parallel));
430            let group_results = join_all(tool_calls[idx..end].iter().map(|tc| {
431                let semaphore = Arc::clone(&semaphore);
432                async move {
433                    let _permit = semaphore
434                        .acquire()
435                        .await
436                        .expect("semaphore closed unexpectedly");
437                    execute_tool_call(tc, registry, provider_name).await
438                }
439            }))
440            .await;
441
442            results.extend(group_results);
443            idx = end;
444        } else {
445            results.push(execute_tool_call(&tool_calls[idx], registry, provider_name).await);
446            idx += 1;
447        }
448    }
449
450    results
451}
452
453impl<A: HttpAgentAdapter> AgentProvider for HttpAgentProvider<A> {
454    fn invoke<'a>(&'a self, config: &'a AgentConfig) -> InvokeFuture<'a> {
455        Box::pin(async move {
456            // Resolved before any request: an unknown profile never reaches the model.
457            let selection = self.tools.select(config.tool_profile.as_ref())?;
458            if !self.tools.is_empty() {
459                info!(
460                    provider = self.adapter.provider_name(),
461                    profile = %selection.label(),
462                    "{}",
463                    selection.describe()
464                );
465            }
466            let tool_registry = selection.registry();
467            let mut request_body = self.adapter.build_request(config)?;
468
469            // Inject tools into the request if a registry is available
470            if let Some(registry) = tool_registry
471                && !registry.is_empty()
472            {
473                let tools_array = registry.to_openai_tools();
474                request_body["tools"] = Value::Array(tools_array);
475            }
476
477            let max_turns = config.max_turns.unwrap_or(25) as usize;
478            let max_budget = config.max_budget_usd.unwrap_or(f64::MAX);
479            let mut state = LoopState::new(Instant::now(), config.verbose);
480
481            // Messages array for multi-turn
482            let mut messages: Vec<Value> = request_body
483                .get("messages")
484                .and_then(|m| m.as_array())
485                .cloned()
486                .unwrap_or_default();
487
488            for turn in 0..max_turns {
489                request_body["messages"] = Value::Array(messages.clone());
490                let turn_result = self.execute_turn(&request_body, config).await?;
491
492                // Accumulate usage
493                let turn_input = turn_result.usage.input_tokens.unwrap_or(0);
494                let turn_output = turn_result.usage.output_tokens.unwrap_or(0);
495                state.total_input_tokens += turn_input;
496                state.total_output_tokens += turn_output;
497                if let Some(v) = turn_result.usage.cache_read_input_tokens {
498                    state.total_cache_read_tokens =
499                        Some(state.total_cache_read_tokens.unwrap_or(0) + v);
500                }
501                if let Some(v) = turn_result.usage.cache_creation_input_tokens {
502                    state.total_cache_creation_tokens =
503                        Some(state.total_cache_creation_tokens.unwrap_or(0) + v);
504                }
505
506                if state.model_name.is_none() {
507                    state.model_name = turn_result.model.clone();
508                }
509
510                if let Some(ref model) = state.model_name
511                    && let Some(turn_cost) = self.adapter.compute_cost(model, &turn_result.usage)
512                {
513                    state.total_cost += turn_cost;
514                }
515
516                // Record debug trace for this turn
517                if config.verbose {
518                    let tool_calls_debug: Vec<DebugToolCall> = turn_result
519                        .tool_calls
520                        .iter()
521                        .map(|tc| DebugToolCall {
522                            id: Some(tc.id.clone()),
523                            name: tc.name.clone(),
524                            input: tc.input.clone(),
525                        })
526                        .collect();
527
528                    state.debug_messages.push(DebugMessage {
529                        text: turn_result.text.clone(),
530                        thinking: None,
531                        thinking_redacted: false,
532                        tool_calls: tool_calls_debug,
533                        tool_results: Vec::new(),
534                        stop_reason: if turn_result.is_final {
535                            Some("end_turn".to_string())
536                        } else {
537                            Some("tool_use".to_string())
538                        },
539                        input_tokens: Some(turn_input),
540                        output_tokens: Some(turn_output),
541                    });
542                }
543
544                // Final response (no tool calls) -> return
545                if turn_result.is_final || turn_result.tool_calls.is_empty() {
546                    info!(
547                        provider = self.adapter.provider_name(),
548                        turns = turn + 1,
549                        duration_ms = state.start.elapsed().as_millis() as u64,
550                        input_tokens = state.total_input_tokens,
551                        cache_read_input_tokens = state.total_cache_read_tokens,
552                        cache_creation_input_tokens = state.total_cache_creation_tokens,
553                        output_tokens = state.total_output_tokens,
554                        "invocation complete"
555                    );
556                    return Ok(state.into_output(extract_value(&turn_result)));
557                }
558
559                // Tool calls but no registry -> return text (backward compat)
560                let registry = match tool_registry {
561                    Some(r) => r,
562                    None => {
563                        warn!(
564                            provider = self.adapter.provider_name(),
565                            tool_calls = turn_result.tool_calls.len(),
566                            "model requested tool calls but no registry attached, returning text"
567                        );
568                        return Ok(state.into_output(extract_text_value(&turn_result)));
569                    }
570                };
571
572                // Budget exceeded -> stop
573                if state.total_cost >= max_budget {
574                    warn!(
575                        provider = self.adapter.provider_name(),
576                        cost = state.total_cost,
577                        budget = max_budget,
578                        "budget exceeded, stopping agentic loop"
579                    );
580                    return Ok(state.into_output(extract_text_value(&turn_result)));
581                }
582
583                // Build assistant message with tool_calls for conversation history
584                let assistant_tool_calls: Vec<Value> = turn_result
585                    .tool_calls
586                    .iter()
587                    .map(|tc| {
588                        json!({
589                            "id": tc.id,
590                            "type": "function",
591                            "function": {
592                                "name": tc.name,
593                                "arguments": tc.input.to_string()
594                            }
595                        })
596                    })
597                    .collect();
598
599                let mut assistant_msg = json!({"role": "assistant"});
600                if let Some(ref text) = turn_result.text {
601                    assistant_msg["content"] = Value::String(text.clone());
602                } else {
603                    assistant_msg["content"] = Value::Null;
604                }
605                assistant_msg["tool_calls"] = Value::Array(assistant_tool_calls);
606                messages.push(assistant_msg);
607
608                // Execute tool calls (consecutive read-only calls run concurrently)
609                let max_parallel = config.max_parallel_tools.max(1);
610                let results = execute_turn_tool_calls(
611                    &turn_result.tool_calls,
612                    registry,
613                    self.adapter.provider_name(),
614                    max_parallel,
615                )
616                .await;
617
618                let mut tool_results_debug: Vec<DebugToolResult> = Vec::new();
619                for (tc, (content, is_error)) in turn_result.tool_calls.iter().zip(results) {
620                    messages.push(json!({
621                        "role": "tool",
622                        "tool_call_id": tc.id,
623                        "content": content
624                    }));
625
626                    if config.verbose {
627                        tool_results_debug.push(DebugToolResult {
628                            tool_use_id: Some(tc.id.clone()),
629                            content: Value::String(content.clone()),
630                            is_error,
631                        });
632                    }
633                }
634
635                if config.verbose
636                    && let Some(last_msg) = state.debug_messages.last_mut()
637                {
638                    last_msg.tool_results = tool_results_debug;
639                }
640
641                info!(
642                    provider = self.adapter.provider_name(),
643                    turn = turn + 1,
644                    tools_executed = turn_result.tool_calls.len(),
645                    "turn complete, continuing loop"
646                );
647            }
648
649            warn!(
650                provider = self.adapter.provider_name(),
651                max_turns, "max turns reached, returning last state"
652            );
653            Ok(state.into_output(Value::String(String::new())))
654        })
655    }
656
657    /// Same as [`invoke`](AgentProvider::invoke), and records on `log_sink`
658    /// the tool profile of the step with the tools it exposes.
659    fn invoke_with_logs<'a>(
660        &'a self,
661        config: &'a AgentConfig,
662        log_sink: Arc<dyn LogSink>,
663    ) -> InvokeFuture<'a> {
664        // An unknown profile is not logged here: `invoke` fails with it.
665        if !self.tools.is_empty()
666            && let Ok(selection) = self.tools.select(config.tool_profile.as_ref())
667        {
668            log_sink.log("system", &selection.describe());
669        }
670        self.invoke(config)
671    }
672}
673
674#[cfg(test)]
675mod tests {
676    use std::future::Future;
677    use std::pin::Pin;
678    use std::sync::Mutex;
679    use std::sync::atomic::{AtomicUsize, Ordering};
680
681    use tokio::time::sleep;
682
683    use crate::providers::http::tools::{Tool, ToolError, ToolOutput};
684
685    use super::*;
686
687    type CallLog = Arc<Mutex<Vec<(String, Instant, Instant)>>>;
688
689    struct DelayTool {
690        name: String,
691        read_only: bool,
692        delay: Duration,
693        concurrent: Arc<AtomicUsize>,
694        max_concurrent: Arc<AtomicUsize>,
695        log: Option<CallLog>,
696    }
697
698    impl DelayTool {
699        fn new(name: &str, read_only: bool, delay_ms: u64) -> Self {
700            Self {
701                name: name.to_string(),
702                read_only,
703                delay: Duration::from_millis(delay_ms),
704                concurrent: Arc::new(AtomicUsize::new(0)),
705                max_concurrent: Arc::new(AtomicUsize::new(0)),
706                log: None,
707            }
708        }
709
710        fn with_counters(mut self, concurrent: &Arc<AtomicUsize>, max: &Arc<AtomicUsize>) -> Self {
711            self.concurrent = Arc::clone(concurrent);
712            self.max_concurrent = Arc::clone(max);
713            self
714        }
715
716        fn with_log(mut self, log: &CallLog) -> Self {
717            self.log = Some(Arc::clone(log));
718            self
719        }
720    }
721
722    impl Tool for DelayTool {
723        fn name(&self) -> &str {
724            &self.name
725        }
726
727        fn description(&self) -> &str {
728            "Sleeps then returns its name"
729        }
730
731        fn parameters_schema(&self) -> Value {
732            json!({"type": "object", "properties": {}})
733        }
734
735        fn read_only(&self) -> bool {
736            self.read_only
737        }
738
739        fn execute(
740            &self,
741            _input: Value,
742        ) -> Pin<Box<dyn Future<Output = Result<ToolOutput, ToolError>> + Send + '_>> {
743            Box::pin(async move {
744                let start = Instant::now();
745                let now = self.concurrent.fetch_add(1, Ordering::SeqCst) + 1;
746                self.max_concurrent.fetch_max(now, Ordering::SeqCst);
747                sleep(self.delay).await;
748                self.concurrent.fetch_sub(1, Ordering::SeqCst);
749                let end = Instant::now();
750                if let Some(ref log) = self.log {
751                    log.lock()
752                        .expect("log mutex poisoned")
753                        .push((self.name.clone(), start, end));
754                }
755                Ok(ToolOutput::success(self.name.clone()))
756            })
757        }
758    }
759
760    fn call(id: &str, name: &str) -> HttpToolCall {
761        HttpToolCall {
762            id: id.to_string(),
763            name: name.to_string(),
764            input: json!({}),
765        }
766    }
767
768    fn entry(log: &CallLog, name: &str) -> (Instant, Instant) {
769        let entries = log.lock().expect("log mutex poisoned");
770        let (_, start, end) = entries
771            .iter()
772            .find(|(n, _, _)| n == name)
773            .unwrap_or_else(|| panic!("no log entry for {name}"));
774        (*start, *end)
775    }
776
777    #[tokio::test]
778    async fn parallel_read_only_tools() {
779        let registry = ToolRegistry::new()
780            .register(DelayTool::new("read_a", true, 200))
781            .register(DelayTool::new("read_b", true, 200));
782        let calls = vec![call("1", "read_a"), call("2", "read_b")];
783
784        let started = Instant::now();
785        let results = execute_turn_tool_calls(&calls, &registry, "test", 4).await;
786        let elapsed = started.elapsed();
787
788        assert!(
789            elapsed < Duration::from_millis(350),
790            "read-only calls should run concurrently, took {elapsed:?}"
791        );
792        assert_eq!(results.len(), 2);
793        assert!(results.iter().all(|(_, is_error)| !is_error));
794    }
795
796    #[tokio::test]
797    async fn write_tool_is_barrier() {
798        let log: CallLog = Arc::new(Mutex::new(Vec::new()));
799        let registry = ToolRegistry::new()
800            .register(DelayTool::new("read1", true, 100).with_log(&log))
801            .register(DelayTool::new("write", false, 100).with_log(&log))
802            .register(DelayTool::new("read2", true, 100).with_log(&log));
803        let calls = vec![call("1", "read1"), call("2", "write"), call("3", "read2")];
804
805        let results = execute_turn_tool_calls(&calls, &registry, "test", 4).await;
806        assert_eq!(results.len(), 3);
807
808        let (_, read1_end) = entry(&log, "read1");
809        let (write_start, write_end) = entry(&log, "write");
810        let (read2_start, _) = entry(&log, "read2");
811
812        assert!(write_start >= read1_end, "write must wait for read1");
813        assert!(read2_start >= write_end, "read2 must wait for write");
814    }
815
816    #[tokio::test]
817    async fn tool_results_keep_call_order() {
818        let registry = ToolRegistry::new()
819            .register(DelayTool::new("a", true, 150))
820            .register(DelayTool::new("b", true, 50))
821            .register(DelayTool::new("c", true, 100));
822        let calls = vec![call("1", "a"), call("2", "b"), call("3", "c")];
823
824        let results = execute_turn_tool_calls(&calls, &registry, "test", 4).await;
825
826        assert_eq!(
827            results,
828            vec![
829                ("a".to_string(), false),
830                ("b".to_string(), false),
831                ("c".to_string(), false),
832            ]
833        );
834    }
835
836    #[tokio::test]
837    async fn max_parallel_tools_one_is_sequential() {
838        let cur = Arc::new(AtomicUsize::new(0));
839        let peak = Arc::new(AtomicUsize::new(0));
840        let registry = ToolRegistry::new()
841            .register(DelayTool::new("a", true, 50).with_counters(&cur, &peak))
842            .register(DelayTool::new("b", true, 50).with_counters(&cur, &peak))
843            .register(DelayTool::new("c", true, 50).with_counters(&cur, &peak));
844        let calls = vec![call("1", "a"), call("2", "b"), call("3", "c")];
845
846        let results = execute_turn_tool_calls(&calls, &registry, "test", 1).await;
847
848        assert_eq!(results.len(), 3);
849        assert_eq!(peak.load(Ordering::SeqCst), 1);
850    }
851
852    #[tokio::test]
853    async fn read_only_group_respects_max_parallel() {
854        let cur = Arc::new(AtomicUsize::new(0));
855        let peak = Arc::new(AtomicUsize::new(0));
856        let registry = ToolRegistry::new()
857            .register(DelayTool::new("a", true, 50).with_counters(&cur, &peak))
858            .register(DelayTool::new("b", true, 50).with_counters(&cur, &peak))
859            .register(DelayTool::new("c", true, 50).with_counters(&cur, &peak));
860        let calls = vec![call("1", "a"), call("2", "b"), call("3", "c")];
861
862        let results = execute_turn_tool_calls(&calls, &registry, "test", 2).await;
863
864        assert_eq!(results.len(), 3);
865        assert_eq!(peak.load(Ordering::SeqCst), 2);
866    }
867
868    #[tokio::test]
869    async fn max_parallel_zero_is_floored_to_one() {
870        let registry = ToolRegistry::new().register(DelayTool::new("a", true, 10));
871        let calls = vec![call("1", "a")];
872
873        let results = execute_turn_tool_calls(&calls, &registry, "test", 0).await;
874
875        assert_eq!(results, vec![("a".to_string(), false)]);
876    }
877
878    #[tokio::test]
879    async fn unknown_tool_call_is_treated_as_barrier() {
880        let registry = ToolRegistry::new();
881        let calls = vec![call("1", "missing")];
882
883        let results = execute_turn_tool_calls(&calls, &registry, "test", 4).await;
884
885        assert_eq!(results, vec![("Unknown tool: missing".to_string(), true)]);
886    }
887
888    #[tokio::test]
889    async fn empty_tool_calls_return_empty_results() {
890        let registry = ToolRegistry::new();
891
892        let results = execute_turn_tool_calls(&[], &registry, "test", 4).await;
893
894        assert!(results.is_empty());
895    }
896}