Skip to main content

rig_agent/agent/prompt_request/
streaming.rs

1use rig_core::{
2    OneOrMany,
3    message::{AssistantContent, UserContent},
4    wasm_compat::{WasmBoxedFuture, WasmCompatSend},
5};
6
7use crate::{
8    agent::completion::{PreparedCompletionRequest, build_prepared_completion_request},
9    agent::hook::{
10        AgentHook, HookContext, HookStack, InvalidToolCallAction, ModelTurnFinished, StepEventKind,
11        StreamResponseFinish, TextDelta, ToolCallDelta,
12    },
13    agent::prompt_request::{assistant_text_from_choice, is_empty_assistant_turn},
14    agent::run::{
15        AgentRun, AgentRunStep, PendingToolCall,
16        streamed::{StreamedResolution, StreamedTurnAssembler, StreamedTurnEvent},
17    },
18    agent::runner::{
19        AgentRunner, CompletionCallOutcome, ModelTurnDecision, ToolExecution, acquire_agent_span,
20        append_run_messages, build_chat_span, new_execute_tool_span, observe_action,
21        resolve_completion_call, resolve_model_turn_action, run_single_tool,
22    },
23    completion::GetTokenUsage,
24    streaming::{StreamedAssistantContent, StreamedUserContent, ToolCallDeltaContent},
25    tool::{ToolContext, server::ToolRegistrySnapshot},
26};
27use futures::{Stream, StreamExt, stream};
28use serde::{Deserialize, Serialize};
29use std::{collections::VecDeque, pin::Pin, sync::Arc};
30use tracing_futures::Instrument;
31
32use super::{CompletionCall, PromptResponse, forward_prompt_setters};
33use crate::{
34    agent::Agent,
35    completion::{CompletionError, CompletionModel, PromptError},
36};
37use rig_core::message::{Message, Text};
38
39// The `Send` bound is dropped exactly where `rig-core`'s `WasmCompat*` markers
40// go no-op — browser wasm. `rig-core` keys those markers on this same
41// predicate, so keep the two in step: a bare `target_arch = "wasm32"` would
42// also drop `Send` on WASI, where `rig-core` still requires it.
43#[cfg(not(all(target_arch = "wasm32", target_os = "unknown")))]
44pub type StreamingResult<R> =
45    Pin<Box<dyn Stream<Item = Result<MultiTurnStreamItem<R>, StreamingError>> + Send>>;
46
47#[cfg(all(target_arch = "wasm32", target_os = "unknown"))]
48pub type StreamingResult<R> =
49    Pin<Box<dyn Stream<Item = Result<MultiTurnStreamItem<R>, StreamingError>>>>;
50
51#[derive(Deserialize, Serialize, Debug, Clone)]
52#[serde(tag = "type", rename_all = "camelCase")]
53#[non_exhaustive]
54pub enum MultiTurnStreamItem<R> {
55    /// A streamed assistant content item — the content the **model emitted**:
56    /// text/reasoning deltas, tool-call deltas, and, when the model turn is
57    /// committed, the complete [`StreamedAssistantContent::ToolCall`] for each
58    /// tool call Rig routes to execution. Such a call is reported here whether or
59    /// not the tool body ultimately runs (a hook skip still reports it);
60    /// it is **not** an execution-lifecycle event (see
61    /// [`ToolExecutionCommitted`](Self::ToolExecutionCommitted)).
62    ///
63    /// Two kinds of model tool call are **not** re-emitted as a complete
64    /// `ToolCall` item here (their arguments still stream as tool-call deltas):
65    /// a call rejected and handled by invalid-tool-call recovery (surfaced via
66    /// that recovery path), and a structured-output Tool-mode output-tool call,
67    /// which finalizes the run directly — its structured result is surfaced in
68    /// the [`FinalResponse`](Self::FinalResponse) rather than as a completed
69    /// `ToolCall` item.
70    StreamAssistantItem(StreamedAssistantContent<R>),
71    /// Confirmation that Rig **executed and committed** a tool call. This is not
72    /// a real-time start notification: it is surfaced together with its
73    /// `ToolResult` only after the whole batch settles successfully. Use tool
74    /// hooks for live host-side start/result observation.
75    ///
76    /// This item is emitted only for a tool whose body actually ran (it passed
77    /// its `ToolCall` hook checks), never for a call dropped by a sibling's
78    /// termination, skipped by a hook, or resolved by invalid-call recovery.
79    /// Correlate it with the model call and result through `internal_call_id`.
80    ToolExecutionCommitted {
81        /// The tool call as **executed**: the model's call with any
82        /// [`ToolCallAction::Rewrite`](crate::agent::ToolCallAction::Rewrite) hook rewrite
83        /// applied (so a redaction rewrite is reflected here, not leaked). The
84        /// model's *original* call is reported via
85        /// [`StreamAssistantItem`](Self::StreamAssistantItem).
86        tool_call: rig_core::message::ToolCall,
87        /// Rig-generated id correlating this execution with the model tool call
88        /// ([`StreamedAssistantContent::ToolCall::internal_call_id`]) and the
89        /// resulting [`StreamedUserContent::ToolResult`].
90        internal_call_id: String,
91    },
92    /// A streamed user content item: the **result** of an executed (or
93    /// hook-skipped) tool call. The tool batch commits and surfaces atomically at
94    /// every `tool_concurrency` (including the sequential default): results are
95    /// surfaced (in call order) only after the whole batch settles successfully —
96    /// a run that terminates mid-batch surfaces no successful tool results.
97    StreamUserItem(StreamedUserContent),
98    /// Details for one successfully completed completion request made by this agent stream.
99    ///
100    /// This is emitted when a provider call finishes. Usage is the provider's
101    /// final usage for that completion request when available; it is not
102    /// incremental per streamed token.
103    ///
104    /// ```rust,ignore
105    /// match item {
106    ///     MultiTurnStreamItem::CompletionCall(completion_call) => {
107    ///         // Zero-valued usage means the provider reported no metrics.
108    ///         if completion_call.usage.has_values() {
109    ///             let context_tokens = completion_call.usage.input_tokens;
110    ///         }
111    ///     }
112    ///     _ => {}
113    /// }
114    /// ```
115    CompletionCall(CompletionCall),
116    /// The completed model turn was rejected by a hook for retry.
117    ///
118    /// Text and reasoning deltas emitted for this turn were provisional. A
119    /// consumer should discard or visually reset output associated with `turn`.
120    /// A subsequent attempt is made only if the run's total model-call budget
121    /// permits it.
122    ModelTurnRetried {
123        /// One-based model-call index of the rejected turn.
124        turn: usize,
125    },
126    /// The final result from the stream: the unified [`PromptResponse`] shared
127    /// with the blocking surface.
128    FinalResponse(PromptResponse),
129}
130
131/// Build the unified [`PromptResponse`] for the streaming surface from the
132/// final turn's structured content.
133fn final_response_from_content(
134    content: OneOrMany<AssistantContent>,
135    aggregated_usage: crate::completion::Usage,
136    completion_calls: Vec<CompletionCall>,
137    history: Option<Vec<Message>>,
138) -> PromptResponse {
139    let mut response = PromptResponse::new(assistant_text_from_choice(&content), aggregated_usage)
140        .with_content(content)
141        .with_completion_calls(completion_calls);
142    response.messages = history;
143    response
144}
145
146impl<R> MultiTurnStreamItem<R> {
147    pub(crate) fn stream_item(item: StreamedAssistantContent<R>) -> Self {
148        Self::StreamAssistantItem(item)
149    }
150
151    pub fn final_response(
152        content: OneOrMany<AssistantContent>,
153        aggregated_usage: crate::completion::Usage,
154    ) -> Self {
155        Self::FinalResponse(final_response_from_content(
156            content,
157            aggregated_usage,
158            Vec::new(),
159            None,
160        ))
161    }
162
163    pub fn final_response_with_history(
164        content: OneOrMany<AssistantContent>,
165        aggregated_usage: crate::completion::Usage,
166        history: Option<Vec<Message>>,
167    ) -> Self {
168        Self::FinalResponse(final_response_from_content(
169            content,
170            aggregated_usage,
171            Vec::new(),
172            history,
173        ))
174    }
175
176    pub(crate) fn final_response_with_completion_calls(
177        content: OneOrMany<AssistantContent>,
178        aggregated_usage: crate::completion::Usage,
179        completion_calls: Vec<CompletionCall>,
180        history: Option<Vec<Message>>,
181    ) -> Self {
182        Self::FinalResponse(final_response_from_content(
183            content,
184            aggregated_usage,
185            completion_calls,
186            history,
187        ))
188    }
189}
190
191/// Drain a provider stream abandoned by invalid tool-call recovery so the
192/// reported usage for the recovered completion call is not lost.
193async fn drain_stream_usage<R>(
194    stream: &mut crate::streaming::StreamingCompletionResponse<R>,
195) -> Result<crate::completion::Usage, StreamingError>
196where
197    R: Clone + Unpin + GetTokenUsage,
198{
199    while let Some(content) = stream.next().await {
200        match content {
201            Ok(StreamedAssistantContent::Final(final_resp)) => {
202                return Ok(final_resp.token_usage());
203            }
204            Ok(_) => {}
205            Err(err) => return Err(err.into()),
206        }
207    }
208
209    Ok(crate::completion::Usage::new())
210}
211
212pub(crate) fn record_usage_on_span(span: &tracing::Span, usage: crate::completion::Usage) {
213    span.record("gen_ai.usage.input_tokens", usage.input_tokens);
214    span.record("gen_ai.usage.output_tokens", usage.output_tokens);
215    span.record(
216        "gen_ai.usage.cache_read.input_tokens",
217        usage.cached_input_tokens,
218    );
219    span.record(
220        "gen_ai.usage.cache_creation.input_tokens",
221        usage.cache_creation_input_tokens,
222    );
223    span.record(
224        "gen_ai.usage.tool_use_prompt_tokens",
225        usage.tool_use_prompt_tokens,
226    );
227    span.record("gen_ai.usage.reasoning_tokens", usage.reasoning_tokens);
228}
229
230/// Build the final streamed content for a finished run (#1928).
231///
232/// When the finishing turn carries a tool call it is a Tool-mode output-tool
233/// call (a real tool call would have routed to `CallTools`, not `Done`). In that
234/// case the tool call AND the model's prose are dropped, any reasoning/image
235/// content is kept, and `output` is appended as the final text — so the streamed
236/// [`PromptResponse::output`] string is the structured output rather than the
237/// prose, with no unanswered tool_use, matching the non-streaming `output`. Note
238/// this shapes only the surfaced [`PromptResponse::content`]; the persisted
239/// message history is built by the state machine (which keeps the prose, like the
240/// blocking driver), so `content` and `messages` intentionally differ on prose in
241/// this case.
242/// Otherwise returns `None` and the caller surfaces the turn's content unchanged.
243fn finalize_streamed_choice(
244    last_final_choice: &OneOrMany<AssistantContent>,
245    output: &str,
246) -> Option<OneOrMany<AssistantContent>> {
247    let finalized_via_output_tool = last_final_choice
248        .iter()
249        .any(|item| matches!(item, AssistantContent::ToolCall(_)));
250    if !finalized_via_output_tool {
251        return None;
252    }
253    let mut items: Vec<AssistantContent> = last_final_choice
254        .iter()
255        .filter(|item| {
256            !matches!(
257                item,
258                AssistantContent::ToolCall(_) | AssistantContent::Text(_)
259            )
260        })
261        .cloned()
262        .collect();
263    items.push(AssistantContent::text(output.to_string()));
264    Some(
265        OneOrMany::from_iter_optional(items)
266            .unwrap_or_else(|| OneOrMany::one(AssistantContent::text(output.to_string()))),
267    )
268}
269
270#[derive(Debug, thiserror::Error)]
271pub enum StreamingError {
272    #[error("CompletionError: {0}")]
273    Completion(#[from] CompletionError),
274    #[error("PromptError: {0}")]
275    Prompt(#[from] Box<PromptError>),
276}
277
278impl From<rig_core::memory::MemoryError> for StreamingError {
279    fn from(err: rig_core::memory::MemoryError) -> Self {
280        Self::Prompt(Box::new(PromptError::MemoryError(err)))
281    }
282}
283
284/// A builder for creating prompt requests with customizable options.
285/// Uses generics to track which options have been set during the build process.
286///
287/// When the agent has no configured `default_max_turns`, the implicit budget is
288/// one model call. Use [`.max_turns()`](Self::max_turns) to override the agent's
289/// configured or implicit budget; a tool call followed by a model-authored final
290/// answer generally requires at least two model calls.
291pub struct StreamingPromptRequest<M>
292where
293    M: CompletionModel,
294{
295    /// The hook-aware driver this streaming request configures and runs.
296    runner: AgentRunner<M>,
297}
298
299impl<M> StreamingPromptRequest<M>
300where
301    M: CompletionModel + 'static,
302    <M as CompletionModel>::StreamingResponse: WasmCompatSend + GetTokenUsage,
303{
304    /// Create a new `StreamingPromptRequest` from an agent, including its
305    /// default hooks.
306    pub fn new(agent: Arc<Agent<M>>, prompt: impl Into<Message>) -> StreamingPromptRequest<M> {
307        Self::from_agent(agent.as_ref(), prompt)
308    }
309
310    /// Create a new StreamingPromptRequest from an agent, cloning the agent's
311    /// data and default hook stack.
312    pub fn from_agent(agent: &Agent<M>, prompt: impl Into<Message>) -> StreamingPromptRequest<M> {
313        StreamingPromptRequest {
314            runner: AgentRunner::from_agent(agent, prompt),
315        }
316    }
317
318    /// Set the total model-call budget, including the initial call and every
319    /// retry or continuation. Zero emits no model calls; one permits only the
320    /// initial call.
321    ///
322    /// Named to match the blocking
323    /// [`PromptRequest::max_turns`](super::PromptRequest::max_turns) and
324    /// [`TypedPromptRequest::max_turns`](super::TypedPromptRequest::max_turns)
325    /// builders so the same call reads identically on either surface.
326    pub fn max_turns(mut self, turns: usize) -> Self {
327        self.runner = self.runner.max_turns(turns);
328        self
329    }
330
331    /// Execute up to `concurrency` of a turn's tool calls at once (1 by default,
332    /// i.e. sequential). See [`AgentRunner::tool_concurrency`]: at any
333    /// `concurrency` the stream emits the model's `ToolCall` items (call order),
334    /// then — atomically, after the whole tool batch settles successfully — the
335    /// per-tool `ToolExecutionCommitted` + `ToolResult` items in **call order** (not
336    /// completion order). The streamed message history is unchanged at any
337    /// `concurrency`.
338    pub fn tool_concurrency(mut self, concurrency: usize) -> Self {
339        self.runner = self.runner.tool_concurrency(concurrency);
340        self
341    }
342
343    /// Append a hook to this request's hook stack (on top of any the agent
344    /// already carries). Hooks run in registration order; how their results
345    /// compose is event-dependent (`CompletionCall` request patches accumulate
346    /// and merge, `ToolCall`/`ToolResult` rewrites chain, while model-turn
347    /// steering and observe-only/recovery events use first-non-`Continue`-wins). See the
348    /// [`hook`](crate::agent::hook) module docs.
349    pub fn add_hook<H>(mut self, hook: H) -> Self
350    where
351        H: AgentHook + 'static,
352    {
353        self.runner = self.runner.add_hook(hook);
354        self
355    }
356
357    forward_prompt_setters!(runner);
358
359    async fn send(self) -> StreamingResult<M::StreamingResponse> {
360        self.runner.stream().await
361    }
362}
363
364/// A boxed, medium-specific item stream for one engine step (model turn or tool
365/// batch). Boxed so a generic [`drive_agent`] can forward it without the
366/// per-step future leaking into the engine's own (`Send`) inference.
367// Same browser-wasm predicate as `StreamingResult` above, for the same reason.
368#[cfg(not(all(target_arch = "wasm32", target_os = "unknown")))]
369pub(crate) type DriveStream<'a, R> =
370    Pin<Box<dyn Stream<Item = Result<MultiTurnStreamItem<R>, StreamingError>> + Send + 'a>>;
371
372#[cfg(all(target_arch = "wasm32", target_os = "unknown"))]
373pub(crate) type DriveStream<'a, R> =
374    Pin<Box<dyn Stream<Item = Result<MultiTurnStreamItem<R>, StreamingError>> + 'a>>;
375
376/// One item emitted by the shared engine [`drive_agent`].
377///
378/// `Item`s are forwarded to a streaming consumer (and ignored by the blocking
379/// fold); `Done` carries both the canonical [`PromptResponse`] the blocking
380/// surface returns and the medium-specific final stream item the streaming
381/// surface yields.
382// The large `Item` variant is the per-delta hot path (one per streamed token);
383// boxing it to shrink the variant spread would add an allocation per delta,
384// which the streaming path is specifically tuned to avoid. `Done` is yielded
385// once per run, so the wasted space on that rare variant is irrelevant.
386#[allow(clippy::large_enum_variant)]
387pub(crate) enum DriveItem<R> {
388    /// An intermediate stream item (assistant delta, tool call/result, a
389    /// per-call `CompletionCall`, or — last, for the streaming surface — the
390    /// final response item).
391    Item(MultiTurnStreamItem<R>),
392    /// The run finished; carries the canonical response the blocking fold
393    /// returns. The streaming surface has already received the final item as the
394    /// preceding `Item` and ignores this.
395    Done(Box<PromptResponse>),
396}
397
398/// The per-medium half of the agent loop: how a turn is fetched from the model,
399/// how its tools are executed, and how the run's spans/usage/final item are
400/// shaped. The medium-independent outer loop (turn counting, the `CompletionCall`
401/// hook, request preparation, memory) lives once in [`drive_agent`]; only the
402/// genuinely divergent pieces are behind this trait. Invalid-tool-call recovery
403/// is one of them — it lives inside each source's `run_model_turn` (end-of-turn
404/// for blocking, mid-stream for streaming), not in `drive_agent`.
405pub(crate) trait TurnSource<M>: WasmCompatSend
406where
407    M: CompletionModel,
408{
409    /// The raw provider response carried on per-delta stream items.
410    type Raw: WasmCompatSend;
411
412    /// Build this medium's per-turn `chat` span (name + parenting + any
413    /// `follows_from` chaining differ between blocking and streaming).
414    fn open_chat_span(
415        &self,
416        runner: &AgentRunner<M>,
417        effective_preamble: Option<&str>,
418    ) -> tracing::Span;
419
420    /// Run one model turn: issue the provider call, feed the result into the
421    /// sans-IO machine, and yield any intermediate items. Returning normally
422    /// advances the loop; yielding an `Err` terminates the run.
423    #[allow(clippy::too_many_arguments)]
424    fn run_model_turn<'a>(
425        &'a mut self,
426        runner: &'a AgentRunner<M>,
427        hook_ctx: &'a HookContext,
428        run: &'a mut AgentRun,
429        prepared: PreparedCompletionRequest<M>,
430        chat_span: tracing::Span,
431        agent_span: &'a tracing::Span,
432        prompt: Message,
433    ) -> DriveStream<'a, Self::Raw>;
434
435    /// Execute a turn's tool calls, feeding the results into the machine and
436    /// yielding any intermediate items.
437    fn run_tool_calls<'a>(
438        &'a self,
439        runner: &'a AgentRunner<M>,
440        hook_ctx: &'a HookContext,
441        run: &'a mut AgentRun,
442        calls: Vec<PendingToolCall>,
443        tool_snapshot: Arc<ToolRegistrySnapshot>,
444    ) -> DriveStream<'a, Self::Raw>;
445
446    /// Record run-level telemetry onto the agent span at `Done`. Gated on
447    /// `created_agent_span` so a caller-supplied outer span is never polluted.
448    fn record_run_level_telemetry(
449        &self,
450        agent_span: &tracing::Span,
451        response: &PromptResponse,
452        created_agent_span: bool,
453    );
454
455    /// Build the final stream item surfaced at `Done`, or `None` when the
456    /// surface discards it (the blocking fold) so the engine skips the work.
457    fn final_item(&self, response: &PromptResponse) -> Option<MultiTurnStreamItem<Self::Raw>>;
458}
459
460/// Convert a [`StreamingError`] back into a [`PromptError`] for the blocking
461/// surface ([`AgentRunner::run`]), which folds the shared engine. Lossless:
462/// every streaming error originates as one of these.
463pub(crate) fn streaming_error_into_prompt(err: StreamingError) -> PromptError {
464    match err {
465        StreamingError::Completion(err) => PromptError::CompletionError(err),
466        StreamingError::Prompt(err) => *err,
467    }
468}
469
470pub(crate) fn store_error_usage<M>(runner: &AgentRunner<M>, run: &AgentRun)
471where
472    M: CompletionModel,
473{
474    if let Some(usage) = &runner.error_usage {
475        *usage.lock().unwrap_or_else(|error| error.into_inner()) = run.usage();
476    }
477}
478
479/// The single agent drive loop, shared by the blocking and streaming surfaces.
480///
481/// Owns the medium-independent loop — `next_step` dispatch, the `CompletionCall`
482/// hook + request preparation, the `Done` memory append — and delegates the
483/// medium-specific model call, tool execution, span shaping and finalization to
484/// a [`TurnSource`]. The streaming surface forwards the yielded [`DriveItem`]s;
485/// the blocking surface folds them to `Done`.
486pub(crate) fn drive_agent<M, S>(
487    runner: AgentRunner<M>,
488    mut source: S,
489    mut run: AgentRun,
490    agent_span: tracing::Span,
491    created_agent_span: bool,
492    memory_handle: Option<(Arc<dyn rig_core::memory::ConversationMemory>, String)>,
493    is_streaming: bool,
494) -> impl Stream<Item = Result<DriveItem<S::Raw>, StreamingError>>
495where
496    M: CompletionModel,
497    S: TurnSource<M>,
498{
499    async_stream::stream! {
500        // Run-scoped hook context: minted once, shared by every hook event on
501        // both surfaces. `is_streaming` records which surface is driving; the
502        // per-turn index is advanced on each `CallModel` step below.
503        let hook_ctx = HookContext::new(is_streaming, runner.agent_name.clone());
504        // Set only after a model turn commits successfully and consumed by its
505        // immediately following CallTools step. This keeps the sans-IO run state
506        // serializable while pinning execution to the definitions sent that turn.
507        let mut pending_tool_snapshot: Option<Arc<ToolRegistrySnapshot>> = None;
508
509        'outer: loop {
510            let step = match run.next_step() {
511                Ok(step) => step,
512                Err(err) => {
513                    store_error_usage(&runner, &run);
514                    yield Err(Box::new(err).into());
515                    break 'outer;
516                }
517            };
518
519            match step {
520                AgentRunStep::CallModel { prompt, history, turn } => {
521                    drop(pending_tool_snapshot.take());
522                    if runner.max_turns > 1 {
523                        tracing::info!("Current conversation Turns: {}/{}", turn, runner.max_turns);
524                    }
525                    hook_ctx.set_turn(turn);
526
527                    let request_patch =
528                        match resolve_completion_call(&runner.hooks, &hook_ctx, &prompt, &history, turn).await {
529                            CompletionCallOutcome::Terminate(reason) => {
530                                store_error_usage(&runner, &run);
531                                yield Err(StreamingError::Prompt(Box::new(run.cancel_error(reason))));
532                                break 'outer;
533                            }
534                            CompletionCallOutcome::Proceed(request_patch) => request_patch,
535                        };
536
537                    // Record this turn's base system prompt — the patched-or-baseline
538                    // preamble, before any output-mode augmentation the request builder
539                    // appends. Borrow rather than clone since it only needs to outlive
540                    // span creation.
541                    let effective_preamble = request_patch
542                        .as_ref()
543                        .and_then(|o| o.preamble.as_deref())
544                        .or(runner.preamble.as_deref());
545
546                    let chat_span = source.open_chat_span(&runner, effective_preamble);
547
548                    // Pin Tool output mode once committed so later turns stay
549                    // consistent even if the per-turn tool set changes (#1928).
550                    let committed_output_tool = run.output_tool_name().map(str::to_owned);
551                    let mut prepared = match build_prepared_completion_request(
552                        &runner.model,
553                        prompt.clone(),
554                        &history,
555                        runner.preamble.as_deref(),
556                        &runner.static_context,
557                        runner.temperature,
558                        runner.max_tokens,
559                        runner.additional_params.as_ref(),
560                        runner.record_telemetry_content,
561                        runner.tool_choice.as_ref(),
562                        &runner.tool_server_handle,
563                        runner.output_schema.as_ref(),
564                        &runner.output_mode,
565                        committed_output_tool.as_deref(),
566                        runner.output_tool_description.as_deref(),
567                        runner.augment_output_preamble,
568                        request_patch.as_ref(),
569                    )
570                    .await
571                    {
572                        Ok(prepared) => prepared,
573                        Err(err) => {
574                            store_error_usage(&runner, &run);
575                            yield Err(err.into());
576                            break 'outer;
577                        }
578                    };
579                    run.set_output_tool_name(prepared.output_tool_name.clone());
580                    let turn_tool_snapshot = prepared.tool_snapshot.clone();
581                    if runner.record_telemetry_content {
582                        let input_messages = prepared.builder.messages_for_telemetry();
583                        rig_core::telemetry::record_model_input(&chat_span, &input_messages, true);
584                        prepared.builder = prepared.builder.record_content_telemetry(false);
585                    }
586
587                    let mut turn_stream = source.run_model_turn(
588                        &runner,
589                        &hook_ctx,
590                        &mut run,
591                        prepared,
592                        chat_span,
593                        &agent_span,
594                        prompt,
595                    );
596                    let mut turn_error = None;
597                    while let Some(item) = turn_stream.next().await {
598                        match item {
599                            Ok(item) => yield Ok(DriveItem::Item(item)),
600                            Err(err) => {
601                                turn_error = Some(err);
602                                break;
603                            }
604                        }
605                    }
606                    drop(turn_stream);
607                    if let Some(err) = turn_error {
608                        store_error_usage(&runner, &run);
609                        yield Err(err);
610                        break 'outer;
611                    }
612                    pending_tool_snapshot = Some(turn_tool_snapshot);
613                }
614                AgentRunStep::CallTools { calls } => {
615                    let Some(tool_snapshot) = pending_tool_snapshot.take() else {
616                        store_error_usage(&runner, &run);
617                        yield Err(StreamingError::Completion(CompletionError::ResponseError(
618                            "agent requested tool execution without a prepared registry snapshot"
619                                .to_string(),
620                        )));
621                        break 'outer;
622                    };
623                    let mut tool_stream = source.run_tool_calls(
624                        &runner,
625                        &hook_ctx,
626                        &mut run,
627                        calls,
628                        tool_snapshot,
629                    );
630                    let mut tool_error = None;
631                    while let Some(item) = tool_stream.next().await {
632                        match item {
633                            Ok(item) => yield Ok(DriveItem::Item(item)),
634                            Err(err) => {
635                                tool_error = Some(err);
636                                break;
637                            }
638                        }
639                    }
640                    drop(tool_stream);
641                    if let Some(err) = tool_error {
642                        store_error_usage(&runner, &run);
643                        yield Err(err);
644                        break 'outer;
645                    }
646                }
647                AgentRunStep::Done(response) => {
648                    // Run-completion marker, unifying the blocking and streaming
649                    // drivers' run-finished logs into one shared event.
650                    tracing::info!(
651                        turn = run.turn(),
652                        max_turns = runner.max_turns,
653                        "Agent run finished"
654                    );
655                    source.record_run_level_telemetry(&agent_span, &response, created_agent_span);
656                    append_run_messages(
657                        memory_handle.as_ref(),
658                        response.messages.as_deref().unwrap_or_default(),
659                    )
660                    .await;
661                    // Build the final item only when the surface forwards it
662                    // (streaming). The blocking fold discards it, so its source
663                    // returns `None` and the extra full-response clone is skipped.
664                    if let Some(final_item) = source.final_item(&response) {
665                        yield Ok(DriveItem::Item(final_item));
666                    }
667                    yield Ok(DriveItem::Done(Box::new(response)));
668                    break 'outer;
669                }
670            }
671        }
672    }
673}
674
675/// Execute a turn's tool calls **atomically per batch**, shared by both surfaces.
676///
677/// The batch commits and surfaces all-or-nothing:
678///
679/// - The model tool-call events ([`StreamedAssistantContent::ToolCall`]) are
680///   emitted up front — they report what the model emitted at turn commit.
681/// - Every tool then runs (sequentially at `tool_concurrency <= 1`, else
682///   concurrently bounded by it), with outcomes **collected, not surfaced**.
683/// - On the first hook termination / fail-closed error the batch fails fast: no
684///   new tool starts, not-yet-started concurrent siblings are dropped,
685///   already-started ones are drained, and the deterministic lowest call-index
686///   error is surfaced with **no** successful [`ToolExecutionCommitted`] /
687///   [`StreamUserItem`](MultiTurnStreamItem::StreamUserItem) items and **no**
688///   history commit.
689/// - Only if the whole batch settles successfully are the per-tool
690///   [`ToolExecutionCommitted`](MultiTurnStreamItem::ToolExecutionCommitted) + result
691///   items surfaced (in call order, only for tools whose body actually ran) and
692///   the results committed to run history.
693///
694/// When `forward_items` is `false` (the blocking fold) no stream items are built,
695/// but the collect/commit and fail-fast behavior is identical, so `run()` and
696/// `stream()` return the same terminal reason. `chain_tool_span` lets the
697/// blocking surface chain spans into its linear `follows_from` sequence.
698pub(crate) fn drive_tool_calls<'a, M, R, F>(
699    runner: &'a AgentRunner<M>,
700    hook_ctx: &'a HookContext,
701    run: &'a mut AgentRun,
702    calls: Vec<PendingToolCall>,
703    tool_snapshot: Arc<ToolRegistrySnapshot>,
704    chain_tool_span: F,
705    forward_items: bool,
706) -> DriveStream<'a, R>
707where
708    M: CompletionModel,
709    R: WasmCompatSend + 'a,
710    F: Fn(tracing::Span) -> tracing::Span + WasmCompatSend + 'a,
711{
712    // Per-call working state: a stable internal_call_id and the execute span,
713    // paired with the model's tool call. `span` is `Span::none()` for a
714    // preresolved (invalid-recovery) call, which never executes.
715    struct PreparedToolCall {
716        tool_call: rig_core::message::ToolCall,
717        preresolved_result: Option<UserContent>,
718        internal_call_id: String,
719        span: tracing::Span,
720    }
721    // How a settled tool call is surfaced on the stream once the batch succeeds:
722    //   - `Executed`: `ToolExecutionCommitted` (with the effective, hook-rewritten
723    //     call) + the `ToolResult`.
724    //   - `Skipped`: the `ToolResult` only (a `ToolCall` hook returned `Skip`, so
725    //     nothing ran — no execution commit — but the model still sees the result).
726    //   - `Preresolved`: neither (an invalid-recovery result, already surfaced
727    //     during the model turn); committed to history only.
728    enum ToolSurface {
729        // Boxed to keep this enum small next to the empty `Skipped`/`Preresolved`.
730        Executed(Box<rig_core::message::ToolCall>),
731        Skipped,
732        Preresolved,
733    }
734    // A collected tool outcome, held (not surfaced or committed) until the whole
735    // batch settles.
736    struct CollectedToolResult {
737        content: UserContent,
738        internal_call_id: String,
739        surface: ToolSurface,
740    }
741
742    Box::pin(async_stream::stream! {
743        let full_history_for_errors = run.full_history();
744        let call_count = calls.len();
745
746        // Assign each call a stable internal_call_id and, for calls that will
747        // actually execute, an execute span. Emit the MODEL tool-call events now,
748        // right after the turn committed: these report what the model emitted and
749        // are *not* execution-lifecycle events. A preresolved call emits no model
750        // tool-call event (its synthetic result was already surfaced during the
751        // model turn) and gets no execute span.
752        let mut prepared: Vec<PreparedToolCall> = Vec::with_capacity(call_count);
753        for pending in calls {
754            let internal_call_id = pending.internal_call_id.unwrap_or_else(rig_core::id::generate);
755            let (span, preresolved_result) = match pending.preresolved_result {
756                Some(result) => (tracing::Span::none(), Some(result)),
757                None => {
758                    if forward_items {
759                        yield Ok(MultiTurnStreamItem::stream_item(
760                            StreamedAssistantContent::ToolCall {
761                                tool_call: pending.tool_call.clone(),
762                                internal_call_id: internal_call_id.clone(),
763                            },
764                        ));
765                    }
766                    (chain_tool_span(new_execute_tool_span()), None)
767                }
768            };
769            prepared.push(PreparedToolCall {
770                tool_call: pending.tool_call,
771                preresolved_result,
772                internal_call_id,
773                span,
774            });
775        }
776
777        // Run all tools, COLLECTING outcomes in call order — nothing is surfaced
778        // or committed until the whole batch settles (atomic per-batch). On the
779        // first hook termination / fail-closed error we stop starting new tools;
780        // already-started ones are drained; the lowest call-index error wins; and
781        // no successful result is surfaced or committed.
782        let mut collected: Vec<Option<CollectedToolResult>> =
783            (0..call_count).map(|_| None).collect();
784        let mut first_error: Option<(usize, PromptError)> = None;
785
786        if runner.concurrency <= 1 {
787            // Sequential: run in call order, fail-fast on the first terminating
788            // error so the remaining tools never start.
789            for (index, call) in prepared.into_iter().enumerate() {
790                let PreparedToolCall { tool_call, preresolved_result, internal_call_id, span } = call;
791                if let Some(result) = preresolved_result {
792                    if let Some(slot) = collected.get_mut(index) {
793                        *slot = Some(CollectedToolResult {
794                            content: result,
795                            internal_call_id,
796                            surface: ToolSurface::Preresolved,
797                        });
798                    }
799                    continue;
800                }
801                let outcome = run_single_tool(
802                    runner,
803                    hook_ctx,
804                    &tool_snapshot,
805                    &tool_call,
806                    &internal_call_id,
807                    &full_history_for_errors,
808                )
809                .instrument(span)
810                .await;
811                match outcome {
812                    Ok(outcome) => {
813                        let surface = match outcome.execution {
814                            ToolExecution::Executed(effective) => ToolSurface::Executed(effective),
815                            ToolExecution::Skipped => ToolSurface::Skipped,
816                        };
817                        if let Some(slot) = collected.get_mut(index) {
818                            *slot = Some(CollectedToolResult {
819                                content: outcome.content,
820                                internal_call_id,
821                                surface,
822                            });
823                        }
824                    }
825                    Err(err) => {
826                        first_error = Some((index, err));
827                        break;
828                    }
829                }
830            }
831        } else {
832            // Concurrent: bounded by `tool_concurrency`. A shared `terminating`
833            // flag makes a not-yet-started sibling skip (its side effect never
834            // runs) once any sibling terminates — avoiding the Semantic-Kernel
835            // fail-open — while already-in-flight siblings are drained so the
836            // lowest call-index terminator wins and no task is left detached.
837            let terminating = Arc::new(std::sync::atomic::AtomicBool::new(false));
838            let unordered = stream::iter(prepared.into_iter().enumerate())
839                .map(|(index, call)| {
840                    let PreparedToolCall { tool_call, preresolved_result, internal_call_id, span } = call;
841                    let tool_snapshot = &tool_snapshot;
842                    let full_history_for_errors = &full_history_for_errors;
843                    let terminating = terminating.clone();
844                    async move {
845                        if let Some(result) = preresolved_result {
846                            return (
847                                index,
848                                Some(Ok(CollectedToolResult {
849                                    content: result,
850                                    internal_call_id,
851                                    surface: ToolSurface::Preresolved,
852                                })),
853                            );
854                        }
855                        // `None` marks a dropped (never-started) sibling.
856                        if terminating.load(std::sync::atomic::Ordering::SeqCst) {
857                            return (index, None);
858                        }
859                        let outcome = run_single_tool(
860                            runner,
861                            hook_ctx,
862                            tool_snapshot,
863                            &tool_call,
864                            &internal_call_id,
865                            full_history_for_errors,
866                        )
867                        .await;
868                        let mapped = outcome.map(|o| {
869                            let surface = match o.execution {
870                                ToolExecution::Executed(effective) => {
871                                    ToolSurface::Executed(effective)
872                                }
873                                ToolExecution::Skipped => ToolSurface::Skipped,
874                            };
875                            CollectedToolResult {
876                                content: o.content,
877                                internal_call_id,
878                                surface,
879                            }
880                        });
881                        (index, Some(mapped))
882                    }
883                    .instrument(span)
884                })
885                .buffer_unordered(runner.concurrency);
886            futures::pin_mut!(unordered);
887
888            while let Some((index, outcome)) = unordered.next().await {
889                // A dropped sibling records nothing.
890                let result = match outcome {
891                    Some(result) => result,
892                    None => continue,
893                };
894                match result {
895                    Ok(collected_result) => {
896                        if let Some(slot) = collected.get_mut(index) {
897                            *slot = Some(collected_result);
898                        }
899                    }
900                    Err(err) => {
901                        // Fail-fast: stop starting new siblings; keep draining
902                        // in-flight ones so the lowest call-index terminator wins.
903                        terminating.store(true, std::sync::atomic::Ordering::SeqCst);
904                        if first_error.as_ref().is_none_or(|(i, _)| index < *i) {
905                            first_error = Some((index, err));
906                        }
907                    }
908                }
909            }
910        }
911
912        // Settle. On termination: surface only the deterministic error — no
913        // execution commit, no result, no history commit (all-or-nothing).
914        if let Some((_, err)) = first_error {
915            yield Err(StreamingError::Prompt(Box::new(err)));
916            return;
917        }
918
919        // Success: prepare each call's stream items and results in call order,
920        // commit the results, then surface the buffered items. An executed call
921        // surfaces `ToolExecutionCommitted`
922        // (with the effective, hook-rewritten call) then its `ToolResult`; a
923        // hook-skipped call surfaces its `ToolResult` only (nothing ran); a
924        // preresolved call surfaces nothing (already surfaced during the model
925        // turn) but is still committed. Every non-dropped slot is filled; a
926        // dropped slot only occurs after a termination, handled above.
927        let mut committed: Vec<UserContent> = Vec::with_capacity(call_count);
928        let mut surface_items: Vec<MultiTurnStreamItem<R>> =
929            Vec::with_capacity(call_count.saturating_mul(2));
930        for slot in collected {
931            let CollectedToolResult { content, internal_call_id, surface } = match slot {
932                Some(collected_result) => collected_result,
933                None => {
934                    yield Err(StreamingError::Prompt(Box::new(PromptError::CompletionError(
935                        CompletionError::ResponseError(
936                            "tool execution finished without producing every result".to_string(),
937                        ),
938                    ))));
939                    return;
940                }
941            };
942            if forward_items {
943                // An executed call also surfaces its execution commit; a skipped
944                // call surfaces only its result; a preresolved call surfaces
945                // nothing here.
946                let surface_result = match surface {
947                    ToolSurface::Executed(tool_call) => {
948                        surface_items.push(MultiTurnStreamItem::ToolExecutionCommitted {
949                            tool_call: *tool_call,
950                            internal_call_id: internal_call_id.clone(),
951                        });
952                        true
953                    }
954                    ToolSurface::Skipped => true,
955                    ToolSurface::Preresolved => false,
956                };
957                if surface_result
958                    && let UserContent::ToolResult(tool_result) = &content
959                {
960                    surface_items.push(MultiTurnStreamItem::StreamUserItem(
961                        StreamedUserContent::ToolResult {
962                            tool_result: tool_result.clone(),
963                            internal_call_id,
964                        },
965                    ));
966                }
967            }
968            committed.push(content);
969        }
970
971        if let Err(err) = run.tool_results(committed) {
972            yield Err(Box::new(err).into());
973            return;
974        }
975
976        for item in surface_items {
977            yield Ok(item);
978        }
979    })
980}
981
982/// [`TurnSource`] for the streaming surface: each turn opens a provider stream,
983/// drives a [`StreamedTurnAssembler`], and yields assistant/tool deltas.
984pub(crate) struct StreamingTurnSource {
985    /// The raw provider choice of the most recent turn; the final response
986    /// surfaces it as-is, even when canonical reordering was recorded in history.
987    last_final_choice: OneOrMany<AssistantContent>,
988    last_message_id: Option<String>,
989    /// Resolved agent name, kept only for the empty-turn diagnostic warning.
990    agent_name: String,
991    /// Whether we created the agent span (vs. adopting a caller's ambient span);
992    /// gates recording `gen_ai.completion` onto it, matching the blocking source
993    /// so neither surface pollutes a caller-supplied span.
994    created_agent_span: bool,
995    /// Whether sensitive run-level prompt and completion content may be recorded.
996    record_telemetry_content: bool,
997    /// Hot-path interest gates, computed once: skip building/dispatching the
998    /// high-frequency delta events when no hook observes them.
999    observes_text_delta: bool,
1000    observes_tool_call_delta: bool,
1001    /// Whether any hook is present — gates building the (history-cloning)
1002    /// invalid-tool diagnostic context.
1003    has_hooks: bool,
1004}
1005
1006impl StreamingTurnSource {
1007    pub(crate) fn new(
1008        hooks: &HookStack,
1009        agent_name: String,
1010        created_agent_span: bool,
1011        record_telemetry_content: bool,
1012    ) -> Self {
1013        Self {
1014            last_final_choice: OneOrMany::one(AssistantContent::text("")),
1015            last_message_id: None,
1016            agent_name,
1017            created_agent_span,
1018            record_telemetry_content,
1019            observes_text_delta: hooks.observes(StepEventKind::TextDelta),
1020            observes_tool_call_delta: hooks.observes(StepEventKind::ToolCallDelta),
1021            has_hooks: !hooks.is_empty(),
1022        }
1023    }
1024}
1025
1026impl<M> TurnSource<M> for StreamingTurnSource
1027where
1028    M: CompletionModel,
1029    <M as CompletionModel>::StreamingResponse: WasmCompatSend + GetTokenUsage,
1030{
1031    type Raw = M::StreamingResponse;
1032
1033    fn open_chat_span(
1034        &self,
1035        runner: &AgentRunner<M>,
1036        effective_preamble: Option<&str>,
1037    ) -> tracing::Span {
1038        build_chat_span!(runner, effective_preamble, "chat_streaming", "chat")
1039    }
1040
1041    fn run_model_turn<'a>(
1042        &'a mut self,
1043        runner: &'a AgentRunner<M>,
1044        hook_ctx: &'a HookContext,
1045        run: &'a mut AgentRun,
1046        prepared: PreparedCompletionRequest<M>,
1047        chat_span: tracing::Span,
1048        agent_span: &'a tracing::Span,
1049        current_prompt: Message,
1050    ) -> DriveStream<'a, M::StreamingResponse> {
1051        Box::pin(async_stream::stream! {
1052            let mut stream = match prepared
1053                .builder
1054                .stream()
1055                .instrument(chat_span.clone())
1056                .await
1057            {
1058                Ok(stream) => stream,
1059                Err(err) => {
1060                    yield Err(err.into());
1061                    return;
1062                }
1063            };
1064            // Captured from each completion-call emission so the normalized
1065            // `ModelTurnFinished` event carries the turn's usage.
1066            let mut last_usage = crate::completion::Usage::new();
1067
1068            let mut assembler = StreamedTurnAssembler::new(
1069                prepared.executable_tool_names.clone(),
1070                prepared.allowed_tool_names.clone(),
1071            );
1072            let mut completion_call_emitted = false;
1073            let mut turn_abandoned = false;
1074            let mut provider_final_seen = false;
1075            let mut pending_final = None;
1076            // Mirrors the blocking driver's `response_hook_suppressed`: a turn
1077            // whose invalid tool call was repaired is a recovered turn, so its
1078            // response-finish hook is suppressed.
1079            let mut turn_recovered = false;
1080
1081            // Emit the turn's single `CompletionCall` exactly once, recording its
1082            // usage onto the chat span and into the run. Defined here (not a free
1083            // fn) so it captures `completion_call_emitted`/`chat_span`/`run`; the
1084            // `yield` stays at each call site because `async_stream::stream!`
1085            // cannot see a `yield` produced inside a nested macro expansion.
1086            // Returns the item to yield (`Some` the first time, `None` after), or
1087            // the terminal error to surface.
1088            macro_rules! emit_completion_call {
1089                ($usage:expr) => {{
1090                    let usage = $usage;
1091                    last_usage = usage;
1092                    if !completion_call_emitted {
1093                        if usage.has_values() {
1094                            record_usage_on_span(&chat_span, usage);
1095                        }
1096                        match run.record_streamed_completion_call(usage) {
1097                            Ok(call) => {
1098                                completion_call_emitted = true;
1099                                Ok(Some(MultiTurnStreamItem::CompletionCall(call)))
1100                            }
1101                            Err(err) => Err(Box::new(err).into()),
1102                        }
1103                    } else {
1104                        Ok(None)
1105                    }
1106                }};
1107            }
1108
1109            'turn: while let Some(item) = stream.next().await {
1110                let item = match item {
1111                    Ok(item) => item,
1112                    Err(err) => {
1113                        yield Err(err.into());
1114                        return;
1115                    }
1116                };
1117                if provider_final_seen {
1118                    yield Err(CompletionError::ResponseError(
1119                        "provider stream emitted visible assistant content after its final response"
1120                            .to_string(),
1121                    )
1122                    .into());
1123                    return;
1124                }
1125                let mut events: VecDeque<StreamedTurnEvent> = match assembler.ingest(&item) {
1126                    Ok(events) => events.into(),
1127                    Err(err) => {
1128                        yield Err(err.into());
1129                        return;
1130                    }
1131                };
1132                // At most one event per ingested item forwards the item itself;
1133                // moving it out of the slot avoids a clone per streamed delta.
1134                let mut item_slot = Some(item);
1135                while let Some(event) = events.pop_front() {
1136                    match event {
1137                        StreamedTurnEvent::EmitIngested => {
1138                            if self.observes_text_delta
1139                                && let Some(StreamedAssistantContent::Text(text)) =
1140                                    item_slot.as_ref()
1141                                && let Some(reason) = observe_action(
1142                                    runner
1143                                        .hooks
1144                                        .on_text_delta(
1145                                            hook_ctx,
1146                                            TextDelta {
1147                                                delta: &text.text,
1148                                                aggregated: assembler.aggregated_text(),
1149                                            },
1150                                        )
1151                                        .await,
1152                                )
1153                            {
1154                                yield Err(StreamingError::Prompt(Box::new(
1155                                    run.cancel_error(reason),
1156                                )));
1157                                return;
1158                            }
1159                            if let Some(item) = item_slot.take() {
1160                                yield Ok(MultiTurnStreamItem::stream_item(item));
1161                            }
1162                        }
1163                        StreamedTurnEvent::EmitToolCallDelta {
1164                            id,
1165                            internal_call_id,
1166                            content,
1167                        } => {
1168                            if self.observes_tool_call_delta {
1169                                let (delta_name, delta_text) = match &content {
1170                                    ToolCallDeltaContent::Name(name) => (Some(name.as_str()), ""),
1171                                    ToolCallDeltaContent::Delta(delta) => (None, delta.as_str()),
1172                                };
1173                                if let Some(reason) = observe_action(
1174                                    runner
1175                                        .hooks
1176                                        .on_tool_call_delta(
1177                                            hook_ctx,
1178                                            ToolCallDelta {
1179                                                tool_call_id: &id,
1180                                                internal_call_id: &internal_call_id,
1181                                                tool_name: delta_name,
1182                                                delta: delta_text,
1183                                            },
1184                                        )
1185                                        .await,
1186                                ) {
1187                                    yield Err(StreamingError::Prompt(Box::new(
1188                                        run.cancel_error(reason),
1189                                    )));
1190                                    return;
1191                                }
1192                            }
1193
1194                            yield Ok(MultiTurnStreamItem::StreamAssistantItem(
1195                                StreamedAssistantContent::ToolCallDelta {
1196                                    id,
1197                                    internal_call_id,
1198                                    content,
1199                                },
1200                            ));
1201                        }
1202                        StreamedTurnEvent::Completed { usage, emit_final } => {
1203                            match emit_completion_call!(usage) {
1204                                Ok(Some(item)) => yield Ok(item),
1205                                Ok(None) => {}
1206                                Err(err) => {
1207                                    yield Err(err);
1208                                    return;
1209                                }
1210                            }
1211                            provider_final_seen = true;
1212
1213                            if emit_final
1214                                && matches!(
1215                                    item_slot.as_ref(),
1216                                    Some(StreamedAssistantContent::Final(_))
1217                                )
1218                            {
1219                                pending_final = item_slot.take();
1220                            }
1221                        }
1222                        StreamedTurnEvent::InvalidToolCall(invalid) => {
1223                            let partial = assembler.partial_turn(stream.message_id.clone());
1224                            // Gated on `has_hooks`: building the diagnostic context
1225                            // clones the chat history, so an empty stack skips it and
1226                            // fails fast — identical to the blocking path.
1227                            let action = if self.has_hooks {
1228                                let context =
1229                                    run.streamed_invalid_tool_call_context(&partial, &invalid);
1230                                runner
1231                                    .hooks
1232                                    .on_invalid_tool_call(hook_ctx, &context)
1233                                    .await
1234                                    .unwrap_or_else(InvalidToolCallAction::fail)
1235                            } else {
1236                                InvalidToolCallAction::fail()
1237                            };
1238
1239                            let resolution =
1240                                match run.resolve_streamed_invalid_tool_call(&partial, &invalid, action) {
1241                                    Ok(resolution) => resolution,
1242                                    Err(err) => {
1243                                        yield Err(Box::new(err).into());
1244                                        return;
1245                                    }
1246                                };
1247
1248                            match resolution {
1249                                StreamedResolution::Repaired { .. } => {
1250                                    // Replayed deltas flow through the same event
1251                                    // handling above; the turn is now recovered, so
1252                                    // its response-finish hook is suppressed.
1253                                    turn_recovered = true;
1254                                    events.extend(assembler.resolve_pending_invalid(&resolution));
1255                                }
1256                                StreamedResolution::TurnAbandoned {
1257                                    ref skipped_tool_result,
1258                                } => {
1259                                    let skipped_tool_result = skipped_tool_result.clone();
1260                                    assembler.resolve_pending_invalid(&resolution);
1261
1262                                    if let Some(err) = assembler.pending_delta_error() {
1263                                        yield Err(err.into());
1264                                        return;
1265                                    }
1266                                    let drained_usage = match drain_stream_usage(&mut stream).await {
1267                                        Ok(usage) => usage,
1268                                        Err(err) => {
1269                                            yield Err(err);
1270                                            return;
1271                                        }
1272                                    };
1273                                    match emit_completion_call!(drained_usage) {
1274                                        Ok(Some(item)) => yield Ok(item),
1275                                        Ok(None) => {}
1276                                        Err(err) => {
1277                                            yield Err(err);
1278                                            return;
1279                                        }
1280                                    }
1281                                    if let Some(tool_result) = skipped_tool_result {
1282                                        yield Ok(MultiTurnStreamItem::StreamUserItem(
1283                                            StreamedUserContent::ToolResult {
1284                                                tool_result,
1285                                                internal_call_id: invalid.internal_call_id.clone(),
1286                                            },
1287                                        ));
1288                                    }
1289                                    turn_abandoned = true;
1290                                    break 'turn;
1291                                }
1292                            }
1293                        }
1294                    }
1295                }
1296            }
1297
1298            if turn_abandoned {
1299                return;
1300            }
1301
1302            if let Some(err) = assembler.pending_delta_error() {
1303                yield Err(err.into());
1304                return;
1305            }
1306
1307            // Final fallback: no usage was ever learned, so there is nothing to
1308            // record onto the span and this is the last read of the flag — kept
1309            // inline (not `emit_completion_call!`) so it doesn't emit a dead
1310            // `completion_call_emitted = true` write.
1311            if !completion_call_emitted {
1312                match run.record_streamed_completion_call(crate::completion::Usage::new()) {
1313                    Ok(call) => yield Ok(MultiTurnStreamItem::CompletionCall(call)),
1314                    Err(err) => {
1315                        yield Err(Box::new(err).into());
1316                        return;
1317                    }
1318                }
1319            }
1320
1321            let final_turn_content = stream.choice.clone();
1322            let streamed_turn = assembler.finish(stream.message_id.clone(), &final_turn_content);
1323            if pending_final.is_some()
1324                && !turn_recovered
1325                && let Some(reason) = observe_action(
1326                    runner
1327                        .hooks
1328                        .on_stream_response_finish(
1329                            hook_ctx,
1330                            StreamResponseFinish {
1331                                prompt: &current_prompt,
1332                                content: &streamed_turn.choice,
1333                                usage: last_usage,
1334                                message_id: streamed_turn.message_id.as_deref(),
1335                            },
1336                        )
1337                        .await,
1338                )
1339            {
1340                yield Err(StreamingError::Prompt(Box::new(run.cancel_error(reason))));
1341                return;
1342            }
1343            self.last_message_id = streamed_turn.message_id.clone();
1344            // The canonical assistant content: `finish` normalizes
1345            // reasoning/text/tool ordering, so this can differ from the raw
1346            // `stream.choice` aggregate. `ModelTurnFinished` — the normalized
1347            // per-turn event — carries this, matching what is recorded into run
1348            // history; the raw `stream.choice` is kept in `last_final_choice` for
1349            // the raw/final streaming behavior.
1350            let canonical_choice = streamed_turn.choice.clone();
1351            if let Err(err) = run.streamed_turn(streamed_turn) {
1352                yield Err(Box::new(err).into());
1353                return;
1354            }
1355            // Normalized per-turn event, fired once the turn is parked for
1356            // acceptance on the streaming surface — including tool-only /
1357            // reasoning-only turns that fire no `StreamResponseFinish`.
1358            // Suppressed for recovered turns, mirroring the blocking surface's
1359            // `Continue` arm.
1360            if !turn_recovered {
1361                let action = runner
1362                    .hooks
1363                    .on_model_turn_finished(
1364                        hook_ctx,
1365                        ModelTurnFinished {
1366                            turn: hook_ctx.turn(),
1367                            content: &canonical_choice,
1368                            usage: last_usage,
1369                        },
1370                    )
1371                    .await;
1372                match resolve_model_turn_action(run, action) {
1373                    Ok(ModelTurnDecision::Advance) => {}
1374                    Ok(ModelTurnDecision::Retried) => {
1375                        yield Ok(MultiTurnStreamItem::ModelTurnRetried {
1376                            turn: hook_ctx.turn(),
1377                        });
1378                        return;
1379                    }
1380                    Ok(ModelTurnDecision::Terminate(reason)) => {
1381                        // Before model-turn steering was added, Stop observed
1382                        // this already completed provider turn: its buffered
1383                        // final and content telemetry were visible before the
1384                        // cancellation. Preserve that behavior while Retry
1385                        // alone suppresses the provisional final.
1386                        if self.created_agent_span && self.record_telemetry_content {
1387                            agent_span.record(
1388                                "gen_ai.completion",
1389                                assistant_text_from_choice(&canonical_choice),
1390                            );
1391                        }
1392                        rig_core::telemetry::record_model_output(
1393                            &chat_span,
1394                            &canonical_choice,
1395                            runner.record_telemetry_content,
1396                        );
1397                        if let Some(item) = pending_final.take() {
1398                            yield Ok(MultiTurnStreamItem::stream_item(item));
1399                        }
1400                        yield Err(StreamingError::Prompt(Box::new(run.cancel_error(reason))));
1401                        return;
1402                    }
1403                    Err(err) => {
1404                        yield Err(StreamingError::Prompt(Box::new(err)));
1405                        return;
1406                    }
1407                }
1408            }
1409
1410            // Only hook-accepted canonical output belongs in content telemetry.
1411            // Keep caller-owned spans untouched, matching the blocking source.
1412            if self.created_agent_span && self.record_telemetry_content {
1413                agent_span.record(
1414                    "gen_ai.completion",
1415                    assistant_text_from_choice(&canonical_choice),
1416                );
1417            }
1418            rig_core::telemetry::record_model_output(
1419                &chat_span,
1420                &canonical_choice,
1421                runner.record_telemetry_content,
1422            );
1423
1424            if let Some(item) = pending_final {
1425                yield Ok(MultiTurnStreamItem::stream_item(item));
1426            }
1427            self.last_final_choice = final_turn_content;
1428        })
1429    }
1430
1431    fn run_tool_calls<'a>(
1432        &'a self,
1433        runner: &'a AgentRunner<M>,
1434        hook_ctx: &'a HookContext,
1435        run: &'a mut AgentRun,
1436        calls: Vec<PendingToolCall>,
1437        tool_snapshot: Arc<ToolRegistrySnapshot>,
1438    ) -> DriveStream<'a, M::StreamingResponse> {
1439        // The streaming surface chains nothing onto its tool spans, and forwards
1440        // the ToolCall/ToolResult items to the consumer.
1441        drive_tool_calls(
1442            runner,
1443            hook_ctx,
1444            run,
1445            calls,
1446            tool_snapshot,
1447            |span| span,
1448            true,
1449        )
1450    }
1451
1452    fn record_run_level_telemetry(
1453        &self,
1454        agent_span: &tracing::Span,
1455        response: &PromptResponse,
1456        created_agent_span: bool,
1457    ) {
1458        if created_agent_span {
1459            record_usage_on_span(agent_span, response.usage);
1460        }
1461    }
1462
1463    fn final_item(
1464        &self,
1465        response: &PromptResponse,
1466    ) -> Option<MultiTurnStreamItem<M::StreamingResponse>> {
1467        // Tool output mode (#1928): when the finishing turn made the output-tool
1468        // call, surface the run's structured output as the final content.
1469        let final_choice = finalize_streamed_choice(&self.last_final_choice, &response.output)
1470            .unwrap_or_else(|| {
1471                if is_empty_assistant_turn(&self.last_final_choice) {
1472                    tracing::warn!(
1473                        agent_name = self.agent_name.as_str(),
1474                        message_id = ?self.last_message_id,
1475                        "Streaming turn completed without assistant text; final response will be empty"
1476                    );
1477                }
1478                self.last_final_choice.clone()
1479            });
1480        // Always surface the accumulated messages (parity with the blocking
1481        // `run()`), regardless of whether the caller supplied input history.
1482        let final_messages: Option<Vec<Message>> =
1483            Some(response.messages.clone().unwrap_or_default());
1484        Some(MultiTurnStreamItem::final_response_with_completion_calls(
1485            final_choice,
1486            response.usage,
1487            response.completion_calls.clone(),
1488            final_messages,
1489        ))
1490    }
1491}
1492
1493impl<M> AgentRunner<M>
1494where
1495    M: CompletionModel + 'static,
1496    <M as CompletionModel>::StreamingResponse: WasmCompatSend + GetTokenUsage,
1497{
1498    /// Drive the agent loop, streaming assistant content, tool activity, and a
1499    /// final response. Hooks fire at every observable point, including streamed
1500    /// text and tool-call deltas. Returns the stream after loading any
1501    /// configured conversation memory.
1502    ///
1503    /// Shares the drive loop, run construction, tool execution and fail-closed
1504    /// hook handling with the blocking [`run`](AgentRunner::run) via
1505    /// `drive_agent`, so the two behave identically apart from the streamed
1506    /// delta events.
1507    pub async fn stream(self) -> StreamingResult<M::StreamingResponse> {
1508        let (agent_span, created_agent_span) = acquire_agent_span(
1509            self.agent_name_or_default(),
1510            self.preamble.as_deref(),
1511            self.record_telemetry_content,
1512        );
1513
1514        if self.record_telemetry_content
1515            && let Some(text) = self.prompt.rag_text()
1516        {
1517            agent_span.record("gen_ai.prompt", text);
1518        }
1519
1520        // When the caller passes explicit history, memory is fully bypassed for
1521        // this request (no load AND no save). Otherwise, if a memory backend and
1522        // conversation id are both configured, load prior history.
1523        let (history_override, memory_handle) = match &self.chat_history {
1524            Some(_) => (None, None),
1525            None => match (&self.memory, &self.conversation_id) {
1526                (Some(memory), Some(id)) => match memory.load(id).await {
1527                    Ok(loaded) => (Some(loaded), Some((memory.clone(), id.clone()))),
1528                    Err(err) => {
1529                        let stream = async_stream::stream! {
1530                            yield Err(StreamingError::from(err));
1531                        };
1532                        // Instrument under the agent span like the success path so
1533                        // a load failure stays tied to invoke_agent.
1534                        return Box::pin(stream.instrument(agent_span));
1535                    }
1536                },
1537                _ => (None, None),
1538            },
1539        };
1540
1541        let run = self.build_run(history_override);
1542        let source = StreamingTurnSource::new(
1543            &self.hooks,
1544            self.agent_name_or_default().to_string(),
1545            created_agent_span,
1546            self.record_telemetry_content,
1547        );
1548
1549        // The blocking surface folds this same engine; the streaming surface
1550        // forwards intermediate items (the final response item is the last one)
1551        // and ends on `Done`.
1552        let driver = drive_agent(
1553            self,
1554            source,
1555            run,
1556            agent_span.clone(),
1557            created_agent_span,
1558            memory_handle,
1559            true,
1560        )
1561        .filter_map(|item| {
1562            std::future::ready(match item {
1563                Ok(DriveItem::Item(item)) => Some(Ok(item)),
1564                Ok(DriveItem::Done(_)) => None,
1565                Err(err) => Some(Err(err)),
1566            })
1567        });
1568
1569        Box::pin(driver.instrument(agent_span))
1570    }
1571}
1572
1573impl<M> IntoFuture for StreamingPromptRequest<M>
1574where
1575    M: CompletionModel + 'static,
1576    <M as CompletionModel>::StreamingResponse: WasmCompatSend,
1577{
1578    type Output = StreamingResult<M::StreamingResponse>; // what `.await` returns
1579    type IntoFuture = WasmBoxedFuture<'static, Self::Output>;
1580
1581    fn into_future(self) -> Self::IntoFuture {
1582        // Wrap send() in a future, because send() returns a stream immediately
1583        Box::pin(async move { self.send().await })
1584    }
1585}
1586
1587/// Helper function to stream assistant-visible completion output to stdout.
1588///
1589/// This helper prints streamed assistant text and reasoning. Streaming metadata
1590/// events, such as `MultiTurnStreamItem::CompletionCall`, are not printed;
1591/// metadata is returned on the [`PromptResponse`] via accessors such as
1592/// [`PromptResponse::completion_calls`]. A model-turn retry prints a visible
1593/// boundary because text already written to stdout cannot be retracted.
1594pub async fn stream_to_stdout<R>(
1595    stream: &mut StreamingResult<R>,
1596) -> Result<PromptResponse, std::io::Error> {
1597    let mut final_res = PromptResponse::empty();
1598    print!("Response: ");
1599    while let Some(content) = stream.next().await {
1600        match content {
1601            Ok(MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Text(
1602                Text { text, .. },
1603            ))) => {
1604                print!("{text}");
1605                std::io::Write::flush(&mut std::io::stdout())?;
1606            }
1607            Ok(MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Reasoning(
1608                reasoning,
1609            ))) => {
1610                let reasoning = reasoning.display_text();
1611                print!("{reasoning}");
1612                std::io::Write::flush(&mut std::io::stdout())?;
1613            }
1614            Ok(MultiTurnStreamItem::FinalResponse(res)) => {
1615                final_res = res;
1616            }
1617            Ok(MultiTurnStreamItem::ModelTurnRetried { turn }) => {
1618                print!("\n[model turn {turn} rejected; retry requested]\nResponse: ");
1619                std::io::Write::flush(&mut std::io::stdout())?;
1620            }
1621            Err(err) => {
1622                eprintln!("Error: {err}");
1623            }
1624            _ => {}
1625        }
1626    }
1627
1628    Ok(final_res)
1629}
1630
1631#[cfg(test)]
1632#[allow(irrefutable_let_patterns, unreachable_patterns)]
1633mod migrated_tests {
1634    use crate::agent::{
1635        InvalidToolCallAction, InvalidToolCallContext, ObservationAction, StreamResponseFinish,
1636        TextDelta, ToolCall, ToolCallAction, ToolCallDelta,
1637    };
1638
1639    use super::*;
1640    use crate::agent::AgentBuilder;
1641    use crate::agent::hook::{AgentHook, HookContext};
1642    use crate::agent::prompt_request::{TOOL_NOT_EXECUTED_DUE_TO_INVALID_PEER, tool_result_output};
1643    use crate::agent::run::streamed::merge_reasoning_blocks;
1644    use crate::client::AgentClientExt;
1645    use crate::completion::{CompletionRequest, Prompt, PromptError, ToolDefinition, Usage};
1646    use crate::streaming::{StreamingPrompt, ToolCallDeltaContent};
1647    use crate::test_utils::{
1648        AppendFailingMemory, FailingMemory, MockAddTool, MockBarrierTool, MockCompletionModel,
1649        MockContextProbeTool, MockResponse, MockStreamEvent, MockSubtractTool, MockToolError,
1650        MockTurn, SessionId,
1651    };
1652    use crate::tool::{Tool, ToolContext};
1653    use futures::{StreamExt, TryStreamExt};
1654    use rig_core::client::ProviderClient;
1655    use rig_core::message::{
1656        AssistantContent, DocumentSourceKind, ImageMediaType, Message, ReasoningContent,
1657        ToolChoice, ToolResultContent, UserContent,
1658    };
1659    use rig_core::providers::anthropic;
1660    use serde::Deserialize;
1661    use std::collections::{BTreeSet, HashMap};
1662    use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
1663    use std::sync::{Arc, Mutex};
1664    use std::time::Duration;
1665    use tracing::field::{Field, Visit};
1666    use tracing::{Id, Subscriber};
1667    use tracing_subscriber::layer::{Context, SubscriberExt};
1668    use tracing_subscriber::{Layer, Registry, registry::LookupSpan};
1669
1670    fn reasoning(
1671        id: Option<&str>,
1672        content: impl IntoIterator<Item = ReasoningContent>,
1673    ) -> rig_core::message::Reasoning {
1674        let mut reasoning = rig_core::message::Reasoning::new("");
1675        reasoning.id = id.map(str::to_string);
1676        reasoning.content = content.into_iter().collect();
1677        reasoning
1678    }
1679
1680    struct StopAgentStreamingBeforeCompletion;
1681
1682    impl AgentHook for StopAgentStreamingBeforeCompletion {
1683        async fn on_completion_call(
1684            &self,
1685            _ctx: &HookContext,
1686            _event: crate::agent::CompletionCallEvent<'_>,
1687        ) -> crate::agent::CompletionCallAction {
1688            crate::agent::CompletionCallAction::stop("agent streaming stopped")
1689        }
1690    }
1691
1692    #[tokio::test]
1693    async fn public_streaming_request_constructor_preserves_agent_hooks() {
1694        let model = MockCompletionModel::from_stream_turns([[
1695            MockStreamEvent::text("should not run"),
1696            MockStreamEvent::final_response(Usage::new()),
1697        ]]);
1698        let agent = Arc::new(
1699            AgentBuilder::new(model.clone())
1700                .add_hook(StopAgentStreamingBeforeCompletion)
1701                .build(),
1702        );
1703
1704        let mut stream = StreamingPromptRequest::new(agent, "go").await;
1705        let error = stream
1706            .try_next()
1707            .await
1708            .expect_err("the configured agent hook should terminate the stream");
1709
1710        assert!(matches!(
1711            error,
1712            StreamingError::Prompt(error)
1713                if matches!(*error, PromptError::PromptCancelled { ref reason, .. }
1714                    if reason == "agent streaming stopped")
1715        ));
1716        assert_eq!(model.request_count(), 0);
1717    }
1718
1719    #[test]
1720    fn finalize_streamed_choice_surfaces_output_over_tool_call_and_prose() {
1721        use rig_core::message::{ToolCall, ToolFunction};
1722
1723        let output_call = AssistantContent::ToolCall(ToolCall::new(
1724            "c1".to_string(),
1725            ToolFunction::new(
1726                "final_result".to_string(),
1727                serde_json::json!({"city": "Tokyo"}),
1728            ),
1729        ));
1730
1731        // Prose + output-tool call (#1928): the streamed response text must be
1732        // the structured output, not the prose, with no orphan tool_use.
1733        let with_prose = OneOrMany::many(vec![
1734            AssistantContent::text("Sure, here is the weather:"),
1735            output_call.clone(),
1736        ])
1737        .expect("two items");
1738        let final_choice = finalize_streamed_choice(&with_prose, r#"{"city":"Tokyo"}"#)
1739            .expect("a turn with the output-tool call is finalized via it");
1740        assert_eq!(
1741            assistant_text_from_choice(&final_choice),
1742            r#"{"city":"Tokyo"}"#
1743        );
1744        assert!(
1745            !final_choice
1746                .iter()
1747                .any(|item| matches!(item, AssistantContent::ToolCall(_))),
1748            "no unanswered tool_use should remain in the final content"
1749        );
1750
1751        // Output-tool call only.
1752        let only_call = OneOrMany::one(output_call);
1753        let final_choice = finalize_streamed_choice(&only_call, r#"{"city":"Tokyo"}"#)
1754            .expect("finalized via output tool");
1755        assert_eq!(
1756            assistant_text_from_choice(&final_choice),
1757            r#"{"city":"Tokyo"}"#
1758        );
1759
1760        // A plain-text finalize (no tool call) is left to the caller.
1761        let text_only = OneOrMany::one(AssistantContent::text(r#"{"city":"Tokyo"}"#));
1762        assert!(finalize_streamed_choice(&text_only, r#"{"city":"Tokyo"}"#).is_none());
1763    }
1764
1765    #[test]
1766    fn merge_reasoning_blocks_preserves_order_and_signatures() {
1767        let mut accumulated = Vec::new();
1768        let first = reasoning(
1769            Some("rs_1"),
1770            [ReasoningContent::Text {
1771                text: "step-1".to_string(),
1772                signature: Some("sig-1".to_string()),
1773            }],
1774        );
1775        let second = reasoning(
1776            Some("rs_1"),
1777            [
1778                ReasoningContent::Text {
1779                    text: "step-2".to_string(),
1780                    signature: Some("sig-2".to_string()),
1781                },
1782                ReasoningContent::Summary("summary".to_string()),
1783            ],
1784        );
1785
1786        merge_reasoning_blocks(&mut accumulated, &first);
1787        merge_reasoning_blocks(&mut accumulated, &second);
1788
1789        assert_eq!(accumulated.len(), 1);
1790        let merged = accumulated.first().expect("expected accumulated reasoning");
1791        assert_eq!(merged.id.as_deref(), Some("rs_1"));
1792        assert_eq!(merged.content.len(), 3);
1793        assert!(matches!(
1794            merged.content.first(),
1795            Some(ReasoningContent::Text { text, signature: Some(sig) })
1796                if text == "step-1" && sig == "sig-1"
1797        ));
1798        assert!(matches!(
1799            merged.content.get(1),
1800            Some(ReasoningContent::Text { text, signature: Some(sig) })
1801                if text == "step-2" && sig == "sig-2"
1802        ));
1803    }
1804
1805    #[test]
1806    fn merge_reasoning_blocks_keeps_distinct_ids_as_separate_items() {
1807        let mut accumulated = vec![reasoning(
1808            Some("rs_a"),
1809            [ReasoningContent::Text {
1810                text: "step-1".to_string(),
1811                signature: None,
1812            }],
1813        )];
1814        let incoming = reasoning(
1815            Some("rs_b"),
1816            [ReasoningContent::Text {
1817                text: "step-2".to_string(),
1818                signature: None,
1819            }],
1820        );
1821
1822        merge_reasoning_blocks(&mut accumulated, &incoming);
1823        assert_eq!(accumulated.len(), 2);
1824        assert_eq!(
1825            accumulated.first().and_then(|r| r.id.as_deref()),
1826            Some("rs_a")
1827        );
1828        assert_eq!(
1829            accumulated.get(1).and_then(|r| r.id.as_deref()),
1830            Some("rs_b")
1831        );
1832    }
1833
1834    #[test]
1835    fn merge_reasoning_blocks_keeps_none_ids_separate_items() {
1836        let mut accumulated = vec![reasoning(
1837            None,
1838            [ReasoningContent::Text {
1839                text: "first".to_string(),
1840                signature: None,
1841            }],
1842        )];
1843        let incoming = reasoning(
1844            None,
1845            [ReasoningContent::Text {
1846                text: "second".to_string(),
1847                signature: None,
1848            }],
1849        );
1850
1851        merge_reasoning_blocks(&mut accumulated, &incoming);
1852        assert_eq!(accumulated.len(), 2);
1853        assert!(accumulated.first().is_some_and(|reasoning| {
1854            reasoning.id.is_none()
1855                && matches!(
1856                    reasoning.content.first(),
1857                    Some(ReasoningContent::Text { text, .. }) if text == "first"
1858                )
1859        }));
1860        assert!(accumulated.get(1).is_some_and(|reasoning| {
1861            reasoning.id.is_none()
1862                && matches!(
1863                    reasoning.content.first(),
1864                    Some(ReasoningContent::Text { text, .. }) if text == "second"
1865                )
1866        }));
1867    }
1868
1869    #[test]
1870    fn tool_result_output_preserves_multimodal_tool_output() {
1871        let instruction = serde_json::json!({
1872            "instruction": "Use the image part to answer."
1873        });
1874        let mut content = rig_core::OneOrMany::one(ToolResultContent::json(instruction.clone()));
1875        content.push(ToolResultContent::image_base64(
1876            "base64data==",
1877            Some(ImageMediaType::PNG),
1878            None,
1879        ));
1880        let user_content = tool_result_output(
1881            "tool_call_1".to_string(),
1882            Some("call_1".to_string()),
1883            crate::tool::ToolOutput::content(content),
1884        );
1885
1886        let tool_result = match user_content {
1887            UserContent::ToolResult(tool_result) => tool_result,
1888            other => panic!("expected tool result content, got {other:?}"),
1889        };
1890
1891        assert_eq!(tool_result.id, "tool_call_1");
1892        assert_eq!(tool_result.call_id.as_deref(), Some("call_1"));
1893        assert_eq!(tool_result.content.len(), 2);
1894
1895        let mut items = tool_result.content.iter();
1896        match items.next() {
1897            Some(ToolResultContent::Json { value }) => {
1898                assert_eq!(value, &instruction);
1899            }
1900            other => panic!("expected structured JSON payload first, got {other:?}"),
1901        }
1902
1903        match items.next() {
1904            Some(ToolResultContent::Image(image)) => {
1905                assert_eq!(image.media_type, Some(ImageMediaType::PNG));
1906                assert!(matches!(
1907                    image.data,
1908                    DocumentSourceKind::Base64(ref data) if data == "base64data=="
1909                ));
1910            }
1911            other => panic!("expected image payload second, got {other:?}"),
1912        }
1913    }
1914
1915    fn validate_follow_up_tool_history(request: &CompletionRequest) -> Result<(), String> {
1916        let history = request.chat_history.iter().cloned().collect::<Vec<_>>();
1917        if history.len() != 3 {
1918            return Err(format!(
1919                "follow-up request should contain [original user prompt, assistant tool call, user tool result]: {history:?}"
1920            ));
1921        }
1922
1923        if !matches!(
1924            history.first(),
1925            Some(Message::User { content })
1926                if matches!(
1927                    content.first(),
1928                    UserContent::Text(text) if text.text == "do tool work"
1929                )
1930        ) {
1931            return Err(format!(
1932                "follow-up request should begin with the original user prompt: {history:?}"
1933            ));
1934        }
1935
1936        if !matches!(
1937            history.get(1),
1938            Some(Message::Assistant { content, .. })
1939                if matches!(
1940                    content.first(),
1941                    AssistantContent::ToolCall(tool_call)
1942                        if tool_call.id == "tool_call_1"
1943                            && tool_call.call_id.as_deref() == Some("call_1")
1944                )
1945        ) {
1946            return Err(format!(
1947                "follow-up request is missing the assistant tool call in position 2: {history:?}"
1948            ));
1949        }
1950
1951        if !matches!(
1952            history.get(2),
1953            Some(Message::User { content })
1954                if matches!(
1955                    content.first(),
1956                    UserContent::ToolResult(tool_result)
1957                        if tool_result.id == "tool_call_1"
1958                            && tool_result.call_id.as_deref() == Some("call_1")
1959                )
1960        ) {
1961            return Err(format!(
1962                "follow-up request should end with the user tool result: {history:?}"
1963            ));
1964        }
1965
1966        Ok(())
1967    }
1968
1969    fn history_contains_tool_call(history: &[Message], tool_name: &str) -> bool {
1970        history.iter().any(|message| {
1971            matches!(
1972                message,
1973                Message::Assistant { content, .. }
1974                    if content.iter().any(|item| matches!(
1975                        item,
1976                        AssistantContent::ToolCall(tool_call)
1977                            if tool_call.function.name == tool_name
1978                    ))
1979            )
1980        })
1981    }
1982
1983    fn history_contains_text(history: &[Message], expected: &str) -> bool {
1984        history.iter().any(|message| {
1985            matches!(
1986                message,
1987                Message::Assistant { content, .. }
1988                    if content.iter().any(|item| matches!(
1989                        item,
1990                        AssistantContent::Text(text) if text.text == expected
1991                    ))
1992            )
1993        })
1994    }
1995
1996    fn assistant_reasoning_precedes_tool_call(
1997        history: &[Message],
1998        expected_reasoning: &str,
1999        tool_name: &str,
2000    ) -> bool {
2001        history.iter().any(|message| {
2002            let Message::Assistant { content, .. } = message else {
2003                return false;
2004            };
2005
2006            let reasoning_index = content.iter().position(|item| {
2007                matches!(
2008                    item,
2009                    AssistantContent::Reasoning(reasoning)
2010                        if reasoning.content.iter().any(|content| matches!(
2011                            content,
2012                            ReasoningContent::Text { text, .. }
2013                                if text == expected_reasoning
2014                        ))
2015                )
2016            });
2017            let tool_index = content.iter().position(|item| {
2018                matches!(
2019                    item,
2020                    AssistantContent::ToolCall(tool_call)
2021                        if tool_call.function.name == tool_name
2022                )
2023            });
2024
2025            matches!((reasoning_index, tool_index), (Some(reasoning), Some(tool)) if reasoning < tool)
2026        })
2027    }
2028
2029    fn assistant_reasoning_precedes_text_and_tool_call(
2030        history: &[Message],
2031        expected_reasoning: &str,
2032        expected_text: &str,
2033        tool_name: &str,
2034    ) -> bool {
2035        history.iter().any(|message| {
2036            let Message::Assistant { content, .. } = message else {
2037                return false;
2038            };
2039
2040            let reasoning_index = content.iter().position(|item| {
2041                matches!(
2042                    item,
2043                    AssistantContent::Reasoning(reasoning)
2044                        if reasoning.content.iter().any(|content| matches!(
2045                            content,
2046                            ReasoningContent::Text { text, .. }
2047                                if text == expected_reasoning
2048                        ))
2049                )
2050            });
2051            let text_index = content.iter().position(|item| {
2052                matches!(
2053                    item,
2054                    AssistantContent::Text(text) if text.text == expected_text
2055                )
2056            });
2057            let tool_index = content.iter().position(|item| {
2058                matches!(
2059                    item,
2060                    AssistantContent::ToolCall(tool_call)
2061                        if tool_call.function.name == tool_name
2062                )
2063            });
2064
2065            matches!(
2066                (reasoning_index, text_index, tool_index),
2067                (Some(reasoning), Some(text), Some(tool))
2068                    if reasoning < text && text < tool
2069            )
2070        })
2071    }
2072
2073    #[derive(Clone)]
2074    struct PanicOnUnknownToolHook;
2075
2076    impl AgentHook for PanicOnUnknownToolHook {
2077        async fn on_tool_call_delta(
2078            &self,
2079            _: &HookContext,
2080            _: ToolCallDelta<'_>,
2081        ) -> ObservationAction {
2082            panic!("unknown tool call delta should fail before delta hooks run")
2083        }
2084        async fn on_tool_call(&self, _: &HookContext, _: ToolCall<'_>) -> ToolCallAction {
2085            panic!("unknown tool call should fail before tool hooks run")
2086        }
2087        async fn on_stream_response_finish(
2088            &self,
2089            _: &HookContext,
2090            _: StreamResponseFinish<'_>,
2091        ) -> ObservationAction {
2092            panic!("unknown tool call should fail before stream finish hooks run")
2093        }
2094    }
2095
2096    #[derive(Clone)]
2097    struct CountingAddTool {
2098        calls: Arc<AtomicU32>,
2099    }
2100
2101    #[derive(Clone)]
2102    struct CountingSubtractTool {
2103        calls: Arc<AtomicU32>,
2104    }
2105
2106    #[derive(Deserialize)]
2107    struct CountingOperationArgs {
2108        x: i32,
2109        y: i32,
2110    }
2111
2112    fn arithmetic_tool_definition(name: &str, description: &str) -> ToolDefinition {
2113        ToolDefinition {
2114            name: name.to_string(),
2115            description: description.to_string(),
2116            parameters: serde_json::json!({
2117                "type": "object",
2118                "properties": {
2119                    "x": {
2120                        "type": "number",
2121                        "description": "The first operand"
2122                    },
2123                    "y": {
2124                        "type": "number",
2125                        "description": "The second operand"
2126                    }
2127                },
2128                "required": ["x", "y"],
2129            }),
2130        }
2131    }
2132
2133    impl Tool for CountingAddTool {
2134        const NAME: &'static str = "add";
2135        type Error = MockToolError;
2136        type Args = CountingOperationArgs;
2137        type Output = i32;
2138
2139        fn description(&self) -> String {
2140            "Add x and y together".to_string()
2141        }
2142
2143        fn parameters(&self) -> serde_json::Value {
2144            arithmetic_tool_definition(Self::NAME, "Add x and y together").parameters
2145        }
2146
2147        async fn call(
2148            &self,
2149            _context: &mut ToolContext,
2150            args: Self::Args,
2151        ) -> Result<Self::Output, Self::Error> {
2152            self.calls.fetch_add(1, Ordering::SeqCst);
2153            Ok(args.x + args.y)
2154        }
2155    }
2156
2157    impl Tool for CountingSubtractTool {
2158        const NAME: &'static str = "subtract";
2159        type Error = MockToolError;
2160        type Args = CountingOperationArgs;
2161        type Output = i32;
2162
2163        fn description(&self) -> String {
2164            "Subtract y from x".to_string()
2165        }
2166
2167        fn parameters(&self) -> serde_json::Value {
2168            arithmetic_tool_definition(Self::NAME, "Subtract y from x").parameters
2169        }
2170
2171        async fn call(
2172            &self,
2173            _context: &mut ToolContext,
2174            args: Self::Args,
2175        ) -> Result<Self::Output, Self::Error> {
2176            self.calls.fetch_add(1, Ordering::SeqCst);
2177            Ok(args.x - args.y)
2178        }
2179    }
2180
2181    fn streaming_tool_then_text_model() -> MockCompletionModel {
2182        MockCompletionModel::from_stream_turns([
2183            vec![
2184                MockStreamEvent::tool_call(
2185                    "tool_call_1",
2186                    "add",
2187                    serde_json::json!({"x": 1, "y": 2}),
2188                )
2189                .with_call_id("call_1"),
2190                MockStreamEvent::final_response_with_total_tokens(4),
2191            ],
2192            vec![
2193                MockStreamEvent::text("done"),
2194                MockStreamEvent::final_response_with_total_tokens(6),
2195            ],
2196        ])
2197    }
2198
2199    fn usage(input_tokens: u64, output_tokens: u64) -> Usage {
2200        Usage {
2201            input_tokens,
2202            output_tokens,
2203            total_tokens: input_tokens + output_tokens,
2204            cached_input_tokens: 0,
2205            cache_creation_input_tokens: 0,
2206            tool_use_prompt_tokens: 0,
2207            reasoning_tokens: 0,
2208        }
2209    }
2210
2211    #[tokio::test]
2212    async fn execution_commit_items_are_not_emitted_when_run_commit_fails() {
2213        let runner = AgentBuilder::new(MockCompletionModel::default())
2214            .build()
2215            .runner("go");
2216        let tool_snapshot = Arc::new(
2217            runner
2218                .tool_server_handle
2219                .snapshot_tool_defs(None)
2220                .await
2221                .expect("empty tool snapshot should build"),
2222        );
2223
2224        let mut run = AgentRun::new("go").max_turns(2);
2225        assert!(matches!(
2226            run.next_step().expect("initial model step"),
2227            AgentRunStep::CallModel { .. }
2228        ));
2229
2230        let tool_name = "missing".to_string();
2231        let advertised = BTreeSet::from([tool_name.clone()]);
2232        let turn = crate::agent::run::ModelTurn::new(
2233            None,
2234            OneOrMany::one(AssistantContent::ToolCall(
2235                rig_core::message::ToolCall::new(
2236                    "expected_call".to_string(),
2237                    rig_core::message::ToolFunction::new(tool_name, serde_json::json!({})),
2238                ),
2239            )),
2240            Usage::new(),
2241            advertised.clone(),
2242            advertised,
2243        );
2244        assert!(matches!(
2245            run.model_response(turn)
2246                .expect("tool turn should be accepted"),
2247            crate::agent::run::ModelTurnOutcome::Continue { .. }
2248        ));
2249
2250        let mut calls = match run.next_step().expect("tool step") {
2251            AgentRunStep::CallTools { calls } => calls,
2252            other => panic!("expected tool step, got {other:?}"),
2253        };
2254        // Corrupt only the driver's copy so execution settles successfully but
2255        // `AgentRun` rejects the result before any commit-labelled item escapes.
2256        calls[0].tool_call.id = "mismatched_call".to_string();
2257
2258        let hook_context = HookContext::new(true, None);
2259        hook_context.set_turn(1);
2260        let mut stream = drive_tool_calls::<MockCompletionModel, MockResponse, _>(
2261            &runner,
2262            &hook_context,
2263            &mut run,
2264            calls,
2265            tool_snapshot,
2266            |span| span,
2267            true,
2268        );
2269
2270        let mut saw_commit = false;
2271        let mut saw_result = false;
2272        let mut saw_error = false;
2273        while let Some(item) = stream.next().await {
2274            match item {
2275                Ok(MultiTurnStreamItem::ToolExecutionCommitted { .. }) => saw_commit = true,
2276                Ok(MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult {
2277                    ..
2278                })) => saw_result = true,
2279                Err(_) => saw_error = true,
2280                _ => {}
2281            }
2282        }
2283
2284        assert!(
2285            saw_error,
2286            "the mismatched result must fail run-state commit"
2287        );
2288        assert!(!saw_commit, "a failed run-state commit cannot be announced");
2289        assert!(!saw_result, "an uncommitted result cannot be surfaced");
2290    }
2291
2292    #[derive(Clone, Debug, Default)]
2293    struct CapturedSpan {
2294        id: u64,
2295        name: String,
2296        parent_id: Option<u64>,
2297        fields: HashMap<String, u64>,
2298        string_fields: HashMap<String, String>,
2299        record_counts: HashMap<String, usize>,
2300    }
2301
2302    #[derive(Clone, Default)]
2303    struct CapturedSpans(Arc<Mutex<Vec<CapturedSpan>>>);
2304
2305    impl CapturedSpans {
2306        fn clear(&self) {
2307            if let Ok(mut spans) = self.0.lock() {
2308                spans.clear();
2309            }
2310        }
2311
2312        fn insert(&self, id: &Id, name: &str, parent_id: Option<u64>) {
2313            let id = id.into_u64();
2314            if let Ok(mut spans) = self.0.lock() {
2315                spans.push(CapturedSpan {
2316                    id,
2317                    name: name.to_string(),
2318                    parent_id,
2319                    fields: HashMap::new(),
2320                    string_fields: HashMap::new(),
2321                    record_counts: HashMap::new(),
2322                });
2323            }
2324        }
2325
2326        fn record(&self, id: &Id, fields: Vec<CapturedField>) {
2327            if let Ok(mut spans) = self.0.lock()
2328                && let Some(span) = spans.iter_mut().rev().find(|span| span.id == id.into_u64())
2329            {
2330                for field in fields {
2331                    match field {
2332                        CapturedField::Number(name, value) => {
2333                            *span.record_counts.entry(name.clone()).or_insert(0) += 1;
2334                            span.fields.insert(name, value);
2335                        }
2336                        CapturedField::Text(name, value) => {
2337                            *span.record_counts.entry(name.clone()).or_insert(0) += 1;
2338                            span.fields.insert(name.clone(), 0);
2339                            span.string_fields.insert(name, value);
2340                        }
2341                    }
2342                }
2343            }
2344        }
2345
2346        fn record_strings(&self, id: &Id, fields: Vec<(String, String)>) {
2347            if let Ok(mut spans) = self.0.lock()
2348                && let Some(span) = spans.iter_mut().rev().find(|span| span.id == id.into_u64())
2349            {
2350                span.string_fields.extend(fields);
2351            }
2352        }
2353
2354        fn snapshot(&self) -> Vec<CapturedSpan> {
2355            self.0.lock().map(|spans| spans.clone()).unwrap_or_default()
2356        }
2357    }
2358
2359    struct SpanCaptureLayer {
2360        spans: CapturedSpans,
2361    }
2362
2363    impl<S> Layer<S> for SpanCaptureLayer
2364    where
2365        S: Subscriber,
2366        S: for<'lookup> LookupSpan<'lookup>,
2367    {
2368        fn on_new_span(&self, attrs: &tracing::span::Attributes<'_>, id: &Id, ctx: Context<'_, S>) {
2369            let parent_id = attrs
2370                .parent()
2371                .map(Id::into_u64)
2372                .or_else(|| ctx.current_span().id().map(Id::into_u64));
2373            self.spans.insert(id, attrs.metadata().name(), parent_id);
2374            let mut string_fields = Vec::new();
2375            attrs.record(&mut SpanStringCaptureVisitor {
2376                fields: &mut string_fields,
2377            });
2378            self.spans.record_strings(id, string_fields);
2379        }
2380
2381        fn on_record(&self, span: &Id, values: &tracing::span::Record<'_>, _ctx: Context<'_, S>) {
2382            let mut fields = Vec::new();
2383            values.record(&mut SpanFieldCaptureVisitor {
2384                fields: &mut fields,
2385            });
2386            self.spans.record(span, fields);
2387            let mut string_fields = Vec::new();
2388            values.record(&mut SpanStringCaptureVisitor {
2389                fields: &mut string_fields,
2390            });
2391            self.spans.record_strings(span, string_fields);
2392        }
2393    }
2394
2395    enum CapturedField {
2396        Number(String, u64),
2397        Text(String, String),
2398    }
2399
2400    struct SpanFieldCaptureVisitor<'a> {
2401        fields: &'a mut Vec<CapturedField>,
2402    }
2403
2404    struct SpanStringCaptureVisitor<'a> {
2405        fields: &'a mut Vec<(String, String)>,
2406    }
2407
2408    impl Visit for SpanStringCaptureVisitor<'_> {
2409        fn record_str(&mut self, field: &Field, value: &str) {
2410            self.fields
2411                .push((field.name().to_string(), value.to_string()));
2412        }
2413
2414        fn record_debug(&mut self, field: &Field, value: &dyn std::fmt::Debug) {
2415            self.fields
2416                .push((field.name().to_string(), format!("{value:?}")));
2417        }
2418    }
2419
2420    impl Visit for SpanFieldCaptureVisitor<'_> {
2421        fn record_u64(&mut self, field: &Field, value: u64) {
2422            self.fields
2423                .push(CapturedField::Number(field.name().to_string(), value));
2424        }
2425
2426        // Capture the *presence* of non-numeric fields (e.g. `gen_ai.completion`)
2427        // with a placeholder value so tests can assert whether they were recorded.
2428        fn record_str(&mut self, field: &Field, value: &str) {
2429            self.fields.push(CapturedField::Text(
2430                field.name().to_string(),
2431                value.to_string(),
2432            ));
2433        }
2434
2435        fn record_debug(&mut self, field: &Field, value: &dyn std::fmt::Debug) {
2436            self.fields.push(CapturedField::Text(
2437                field.name().to_string(),
2438                format!("{value:?}"),
2439            ));
2440        }
2441    }
2442
2443    async fn assert_stream_usage_recorded_on_chat_spans(
2444        agent: crate::agent::Agent<MockCompletionModel>,
2445        prompt: &str,
2446        max_turns: usize,
2447        expected_usages: &[Usage],
2448    ) {
2449        // Scoped-subscriber tests must not run concurrently; the warm-up
2450        // below explains the callsite-interest hazard this guards against.
2451        let _isolation = crate::test_utils::scoped_tracing_subscriber_guard().await;
2452        let spans = CapturedSpans::default();
2453        let subscriber = Registry::default().with(SpanCaptureLayer {
2454            spans: spans.clone(),
2455        });
2456        let _default = tracing::subscriber::set_default(subscriber);
2457
2458        // Span callsites in the driver are shared with every other test in
2459        // this binary. The FIRST thread to hit a callsite caches its interest
2460        // from that thread's dispatcher (`Dispatchers::Rebuilder::JustOne`
2461        // consults `dispatcher::get_default`), so a parallel test without a
2462        // subscriber can permanently cache `Interest::never` for the very
2463        // spans this harness asserts on. Defend in two steps, both under the
2464        // isolation guard: (1) warm the whole driver path from THIS thread so
2465        // unregistered callsites first-register against this subscriber, then
2466        // (2) rebuild the interest cache to heal callsites a foreign thread
2467        // already poisoned.
2468        let warmup_model = MockCompletionModel::from_stream_turns([[
2469            MockStreamEvent::text("warmup"),
2470            MockStreamEvent::final_response(Usage::default()),
2471        ]]);
2472        let warmup_agent = crate::agent::AgentBuilder::new(warmup_model).build();
2473        let mut warmup_stream = warmup_agent.stream_prompt("warmup").max_turns(1).await;
2474        while let Some(item) = warmup_stream
2475            .try_next()
2476            .await
2477            .expect("warmup stream should not error")
2478        {
2479            if matches!(item, MultiTurnStreamItem::FinalResponse(_)) {
2480                break;
2481            }
2482        }
2483        tracing::callsite::rebuild_interest_cache();
2484        spans.clear();
2485
2486        let empty_history: &[Message] = &[];
2487        // Declare the fields the guard protects so a regression (recording onto
2488        // a caller span) is actually observable, not silently a no-op.
2489        let outer_span = tracing::info_span!("outer", gen_ai.completion = tracing::field::Empty);
2490
2491        async {
2492            let mut stream = agent
2493                .stream_prompt(prompt)
2494                .history(empty_history)
2495                .max_turns(max_turns)
2496                .await;
2497
2498            while let Some(item) = stream.try_next().await.expect("stream should not error") {
2499                if matches!(item, MultiTurnStreamItem::FinalResponse(_)) {
2500                    break;
2501                }
2502            }
2503        }
2504        .instrument(outer_span)
2505        .await;
2506
2507        let span_snapshot = spans.snapshot();
2508        let outer_span_id = span_snapshot
2509            .iter()
2510            .find(|span| span.name == "outer")
2511            .map(|span| span.id)
2512            .expect("outer span should be captured");
2513        let chat_spans = span_snapshot
2514            .iter()
2515            .filter(|span| span.name == "chat_streaming")
2516            .collect::<Vec<_>>();
2517
2518        assert_eq!(chat_spans.len(), expected_usages.len());
2519        assert!(
2520            span_snapshot.iter().all(|span| span.name != "invoke_agent"),
2521            "outer span path should not create invoke_agent"
2522        );
2523
2524        for (chat_span, expected_usage) in chat_spans.into_iter().zip(expected_usages) {
2525            assert_eq!(chat_span.parent_id, Some(outer_span_id));
2526            assert_eq!(
2527                chat_span
2528                    .string_fields
2529                    .get("gen_ai.operation.name")
2530                    .map(String::as_str),
2531                Some("chat")
2532            );
2533            assert_eq!(
2534                chat_span.fields.get("gen_ai.usage.input_tokens"),
2535                Some(&expected_usage.input_tokens)
2536            );
2537            assert_eq!(
2538                chat_span.fields.get("gen_ai.usage.output_tokens"),
2539                Some(&expected_usage.output_tokens)
2540            );
2541            assert_eq!(
2542                chat_span.fields.get("gen_ai.usage.cache_read.input_tokens"),
2543                Some(&expected_usage.cached_input_tokens)
2544            );
2545            assert_eq!(
2546                chat_span
2547                    .fields
2548                    .get("gen_ai.usage.cache_creation.input_tokens"),
2549                Some(&expected_usage.cache_creation_input_tokens)
2550            );
2551            assert_eq!(
2552                chat_span.fields.get("gen_ai.usage.tool_use_prompt_tokens"),
2553                Some(&expected_usage.tool_use_prompt_tokens)
2554            );
2555            assert_eq!(
2556                chat_span.fields.get("gen_ai.usage.reasoning_tokens"),
2557                Some(&expected_usage.reasoning_tokens)
2558            );
2559        }
2560
2561        let outer_span = span_snapshot
2562            .iter()
2563            .find(|span| span.id == outer_span_id)
2564            .expect("outer span should be present");
2565        assert!(
2566            outer_span
2567                .fields
2568                .keys()
2569                .all(|field| !field.starts_with("gen_ai.usage.")),
2570            "usage should not be recorded onto the caller's outer span"
2571        );
2572        assert!(
2573            !outer_span.fields.contains_key("gen_ai.completion"),
2574            "gen_ai.completion should not be recorded onto the caller's outer span \
2575             (parity with the blocking driver)"
2576        );
2577    }
2578
2579    async fn capture_stream_message_telemetry(
2580        record_telemetry_content: bool,
2581    ) -> (CapturedSpan, Vec<CompletionRequest>) {
2582        let _isolation = crate::test_utils::scoped_tracing_subscriber_guard().await;
2583        let spans = CapturedSpans::default();
2584        let subscriber = Registry::default().with(SpanCaptureLayer {
2585            spans: spans.clone(),
2586        });
2587        let _default = tracing::subscriber::set_default(subscriber);
2588
2589        let warmup_model = MockCompletionModel::from_stream_turns([[
2590            MockStreamEvent::text("warmup"),
2591            MockStreamEvent::final_response(Usage::default()),
2592        ]]);
2593        let warmup_agent = crate::agent::AgentBuilder::new(warmup_model).build();
2594        let mut warmup_stream = warmup_agent.stream_prompt("warmup").max_turns(1).await;
2595        while let Some(item) = warmup_stream
2596            .try_next()
2597            .await
2598            .expect("warmup stream should not error")
2599        {
2600            if matches!(item, MultiTurnStreamItem::FinalResponse(_)) {
2601                break;
2602            }
2603        }
2604        tracing::callsite::rebuild_interest_cache();
2605        spans.clear();
2606
2607        let model = MockCompletionModel::from_stream_turns([[
2608            MockStreamEvent::text("stream response secret"),
2609            MockStreamEvent::final_response(Usage::default()),
2610        ]]);
2611        let recorded_model = model.clone();
2612        let builder = AgentBuilder::new(model);
2613        let agent = if record_telemetry_content {
2614            builder
2615                .record_content_telemetry(true)
2616                .context("static stream context secret")
2617                .build()
2618        } else {
2619            builder.context("static stream context secret").build()
2620        };
2621
2622        let mut stream = agent
2623            .stream_prompt("stream prompt secret")
2624            .max_turns(1)
2625            .await;
2626        while let Some(item) = stream.try_next().await.expect("stream should not error") {
2627            if matches!(item, MultiTurnStreamItem::FinalResponse(_)) {
2628                break;
2629            }
2630        }
2631
2632        let span = spans
2633            .snapshot()
2634            .into_iter()
2635            .find(|span| span.name == "chat_streaming")
2636            .expect("chat_streaming span should be captured");
2637        (span, recorded_model.requests())
2638    }
2639
2640    async fn capture_unary_message_telemetry(
2641        record_telemetry_content: bool,
2642    ) -> (CapturedSpan, CapturedSpan, Vec<CompletionRequest>) {
2643        let _isolation = crate::test_utils::scoped_tracing_subscriber_guard().await;
2644        let spans = CapturedSpans::default();
2645        let subscriber = Registry::default().with(SpanCaptureLayer {
2646            spans: spans.clone(),
2647        });
2648        let _default = tracing::subscriber::set_default(subscriber);
2649
2650        let warmup_agent =
2651            crate::agent::AgentBuilder::new(MockCompletionModel::text("warmup")).build();
2652        warmup_agent
2653            .prompt("warmup")
2654            .await
2655            .expect("warmup prompt should not error");
2656        tracing::callsite::rebuild_interest_cache();
2657        spans.clear();
2658
2659        let model = MockCompletionModel::text("blocking response secret");
2660        let recorded_model = model.clone();
2661        let builder = AgentBuilder::new(model).preamble("blocking system secret");
2662        let agent = if record_telemetry_content {
2663            builder.record_content_telemetry(true).build()
2664        } else {
2665            builder.build()
2666        };
2667
2668        agent
2669            .prompt("blocking prompt secret")
2670            .await
2671            .expect("prompt should not error");
2672
2673        let snapshot = spans.snapshot();
2674        let chat_span = snapshot
2675            .iter()
2676            .find(|span| span.name == "chat")
2677            .cloned()
2678            .expect("chat span should be captured");
2679        let agent_span = snapshot
2680            .into_iter()
2681            .find(|span| span.name == "invoke_agent")
2682            .expect("invoke_agent span should be captured");
2683        (chat_span, agent_span, recorded_model.requests())
2684    }
2685
2686    #[tokio::test]
2687    async fn stream_prompt_message_telemetry_is_opt_in() {
2688        let (default_span, default_requests) = capture_stream_message_telemetry(false).await;
2689        assert!(
2690            !default_span.fields.contains_key("gen_ai.input.messages"),
2691            "default streaming prompt should not record input message contents"
2692        );
2693        assert!(
2694            !default_span.fields.contains_key("gen_ai.output.messages"),
2695            "default streaming prompt should not record output message contents"
2696        );
2697
2698        assert_eq!(default_requests.len(), 1);
2699        assert!(
2700            !default_requests[0].record_telemetry_content,
2701            "default agent stream should keep provider request message telemetry disabled"
2702        );
2703
2704        let (opt_in_span, opt_in_requests) = capture_stream_message_telemetry(true).await;
2705        let input = opt_in_span
2706            .string_fields
2707            .get("gen_ai.input.messages")
2708            .expect("opt-in should record input messages");
2709        assert!(input.contains("stream prompt secret"));
2710        assert!(input.contains("static stream context secret"));
2711        let output = opt_in_span
2712            .string_fields
2713            .get("gen_ai.output.messages")
2714            .expect("opt-in should record output messages");
2715        assert!(output.contains("stream response secret"));
2716        assert_eq!(
2717            opt_in_span
2718                .record_counts
2719                .get("gen_ai.input.messages")
2720                .copied(),
2721            Some(1),
2722            "agent-owned input message telemetry should be recorded once"
2723        );
2724        assert_eq!(
2725            opt_in_span
2726                .record_counts
2727                .get("gen_ai.output.messages")
2728                .copied(),
2729            Some(1),
2730            "agent-owned output message telemetry should be recorded once"
2731        );
2732        assert_eq!(opt_in_requests.len(), 1);
2733        assert!(
2734            !opt_in_requests[0].record_telemetry_content,
2735            "agent-owned stream telemetry should clear the provider request flag"
2736        );
2737    }
2738
2739    #[tokio::test]
2740    async fn unary_prompt_message_telemetry_records_accepted_output_when_opted_in() {
2741        let (default_span, default_agent_span, default_requests) =
2742            capture_unary_message_telemetry(false).await;
2743        assert!(
2744            !default_span.fields.contains_key("gen_ai.input.messages"),
2745            "default blocking prompt should not record input message contents"
2746        );
2747        assert!(
2748            !default_span.fields.contains_key("gen_ai.output.messages"),
2749            "default blocking prompt should not record output message contents"
2750        );
2751        assert!(
2752            !default_span
2753                .string_fields
2754                .contains_key("gen_ai.system_instructions"),
2755            "default blocking prompt should not record system instructions"
2756        );
2757        assert!(
2758            !default_agent_span
2759                .string_fields
2760                .contains_key("gen_ai.prompt")
2761        );
2762        assert!(
2763            !default_agent_span
2764                .string_fields
2765                .contains_key("gen_ai.completion")
2766        );
2767        assert_eq!(default_requests.len(), 1);
2768        assert!(
2769            !default_requests[0].record_telemetry_content,
2770            "default blocking prompt should keep provider request message telemetry disabled"
2771        );
2772
2773        let (opt_in_span, opt_in_agent_span, opt_in_requests) =
2774            capture_unary_message_telemetry(true).await;
2775        let input = opt_in_span
2776            .string_fields
2777            .get("gen_ai.input.messages")
2778            .expect("opt-in should record blocking input messages");
2779        assert!(input.contains("blocking prompt secret"));
2780        let output = opt_in_span
2781            .string_fields
2782            .get("gen_ai.output.messages")
2783            .expect("opt-in should record blocking output messages");
2784        assert!(output.contains("blocking response secret"));
2785        assert_eq!(
2786            opt_in_span
2787                .string_fields
2788                .get("gen_ai.system_instructions")
2789                .map(String::as_str),
2790            Some(r#"[{"type":"text","content":"blocking system secret"}]"#)
2791        );
2792        assert_eq!(
2793            opt_in_agent_span
2794                .string_fields
2795                .get("gen_ai.prompt")
2796                .map(String::as_str),
2797            Some("blocking prompt secret")
2798        );
2799        assert_eq!(
2800            opt_in_agent_span
2801                .string_fields
2802                .get("gen_ai.completion")
2803                .map(String::as_str),
2804            Some("blocking response secret")
2805        );
2806        assert_eq!(opt_in_requests.len(), 1);
2807        assert!(
2808            !opt_in_requests[0].record_telemetry_content,
2809            "agent-owned blocking telemetry should clear the provider request flag"
2810        );
2811    }
2812
2813    async fn capture_tool_content_telemetry(record_telemetry_content: bool) -> CapturedSpan {
2814        let _isolation = crate::test_utils::scoped_tracing_subscriber_guard().await;
2815        let spans = CapturedSpans::default();
2816        let subscriber = Registry::default().with(SpanCaptureLayer {
2817            spans: spans.clone(),
2818        });
2819        let _default = tracing::subscriber::set_default(subscriber);
2820
2821        let warmup = AgentBuilder::new(MockCompletionModel::from_turns([
2822            MockTurn::tool_call("warmup", "add", serde_json::json!({"x": 1, "y": 2})),
2823            MockTurn::text("done"),
2824        ]))
2825        .tool(MockAddTool)
2826        .build();
2827        warmup
2828            .runner("warmup")
2829            .max_turns(2)
2830            .run()
2831            .await
2832            .expect("warmup tool run should succeed");
2833        tracing::callsite::rebuild_interest_cache();
2834        spans.clear();
2835
2836        let builder = AgentBuilder::new(MockCompletionModel::from_turns([
2837            MockTurn::tool_call(
2838                "secret-tool-call",
2839                "add",
2840                serde_json::json!({"x": 12345, "y": 67890}),
2841            ),
2842            MockTurn::text("done"),
2843        ]))
2844        .tool(MockAddTool);
2845        let agent = if record_telemetry_content {
2846            builder.record_content_telemetry(true).build()
2847        } else {
2848            builder.build()
2849        };
2850        agent
2851            .runner("use the tool")
2852            .max_turns(2)
2853            .run()
2854            .await
2855            .expect("tool run should succeed");
2856
2857        spans
2858            .snapshot()
2859            .into_iter()
2860            .find(|span| span.name == "execute_tool")
2861            .expect("execute_tool span should be captured")
2862    }
2863
2864    #[tokio::test]
2865    async fn tool_arguments_and_results_follow_content_telemetry_toggle() {
2866        let default_span = capture_tool_content_telemetry(false).await;
2867        assert!(
2868            !default_span
2869                .string_fields
2870                .contains_key("gen_ai.tool.call.arguments")
2871        );
2872        assert!(
2873            !default_span
2874                .string_fields
2875                .contains_key("gen_ai.tool.call.result")
2876        );
2877        assert_eq!(
2878            default_span
2879                .string_fields
2880                .get("gen_ai.tool.name")
2881                .map(String::as_str),
2882            Some("add"),
2883            "structural tool metadata should remain available"
2884        );
2885
2886        let opt_in_span = capture_tool_content_telemetry(true).await;
2887        assert!(
2888            opt_in_span
2889                .string_fields
2890                .get("gen_ai.tool.call.arguments")
2891                .is_some_and(|args| args.contains("12345") && args.contains("67890"))
2892        );
2893        assert!(
2894            opt_in_span
2895                .string_fields
2896                .get("gen_ai.tool.call.result")
2897                .is_some_and(|result| result.contains("80235"))
2898        );
2899    }
2900
2901    #[tokio::test]
2902    async fn streaming_rejected_message_telemetry_does_not_record_output() {
2903        let _isolation = crate::test_utils::scoped_tracing_subscriber_guard().await;
2904        let spans = CapturedSpans::default();
2905        let subscriber = Registry::default().with(SpanCaptureLayer {
2906            spans: spans.clone(),
2907        });
2908        let _default = tracing::subscriber::set_default(subscriber);
2909
2910        let warmup_model = MockCompletionModel::from_stream_turns([[
2911            MockStreamEvent::text("warmup"),
2912            MockStreamEvent::final_response(Usage::default()),
2913        ]]);
2914        let warmup_agent = crate::agent::AgentBuilder::new(warmup_model).build();
2915        let mut warmup_stream = warmup_agent.stream_prompt("warmup").max_turns(1).await;
2916        while let Some(item) = warmup_stream
2917            .try_next()
2918            .await
2919            .expect("warmup stream should not error")
2920        {
2921            if matches!(item, MultiTurnStreamItem::FinalResponse(_)) {
2922                break;
2923            }
2924        }
2925        tracing::callsite::rebuild_interest_cache();
2926        spans.clear();
2927
2928        let model = MockCompletionModel::from_stream_turns([[
2929            MockStreamEvent::text("rejected stream output secret"),
2930            MockStreamEvent::tool_call(
2931                "tool_call_1",
2932                "default_api",
2933                serde_json::json!({"x": 2, "y": 3}),
2934            ),
2935            MockStreamEvent::final_response(Usage::default()),
2936        ]]);
2937        let agent = AgentBuilder::new(model)
2938            .record_content_telemetry(true)
2939            .build();
2940
2941        let mut stream = agent
2942            .stream_prompt("stream rejection prompt")
2943            .max_turns(1)
2944            .await;
2945        let err = loop {
2946            match stream.try_next().await {
2947                Ok(Some(_)) => continue,
2948                Ok(None) => panic!("rejected stream should error"),
2949                Err(err) => break err,
2950            }
2951        };
2952        assert!(
2953            err.to_string().contains("default_api"),
2954            "expected invalid tool error, got {err}"
2955        );
2956
2957        let chat_span = spans
2958            .snapshot()
2959            .into_iter()
2960            .find(|span| span.name == "chat_streaming")
2961            .expect("chat_streaming span should be captured");
2962        assert!(
2963            chat_span.fields.contains_key("gen_ai.input.messages"),
2964            "opt-in rejected stream should still record input messages"
2965        );
2966        assert!(
2967            !chat_span.fields.contains_key("gen_ai.output.messages"),
2968            "rejected streaming turn must not record output message contents"
2969        );
2970    }
2971
2972    #[tokio::test]
2973    async fn unary_repaired_message_telemetry_records_canonical_output() {
2974        let _isolation = crate::test_utils::scoped_tracing_subscriber_guard().await;
2975        let spans = CapturedSpans::default();
2976        let subscriber = Registry::default().with(SpanCaptureLayer {
2977            spans: spans.clone(),
2978        });
2979        let _default = tracing::subscriber::set_default(subscriber);
2980
2981        let warmup_agent =
2982            crate::agent::AgentBuilder::new(MockCompletionModel::text("warmup")).build();
2983        warmup_agent
2984            .prompt("warmup")
2985            .await
2986            .expect("warmup prompt should not error");
2987        tracing::callsite::rebuild_interest_cache();
2988        spans.clear();
2989
2990        let model = MockCompletionModel::new([
2991            MockTurn::tool_call(
2992                "tool_call_1",
2993                "default_api",
2994                serde_json::json!({"x": 2, "y": 3}),
2995            ),
2996            MockTurn::text("done"),
2997        ]);
2998        let recorded_model = model.clone();
2999        let agent = AgentBuilder::new(model)
3000            .record_content_telemetry(true)
3001            .tool(MockAddTool)
3002            .build();
3003
3004        let output = agent
3005            .prompt("repair tool call")
3006            .add_hook(RepairDefaultApiHook)
3007            .max_turns(3)
3008            .await
3009            .expect("repaired tool call should complete");
3010        assert_eq!(output, "done");
3011
3012        let output_messages: Vec<String> = spans
3013            .snapshot()
3014            .into_iter()
3015            .filter(|span| span.name == "chat")
3016            .filter_map(|span| span.string_fields.get("gen_ai.output.messages").cloned())
3017            .collect();
3018        assert!(
3019            output_messages.iter().any(|output| output.contains("add")),
3020            "repaired accepted output should include canonical tool name: {output_messages:?}"
3021        );
3022        assert!(
3023            !output_messages
3024                .iter()
3025                .any(|output| output.contains("default_api")),
3026            "repaired output telemetry must not serialize stale raw tool name: {output_messages:?}"
3027        );
3028
3029        let requests = recorded_model.requests();
3030        assert_eq!(requests.len(), 2);
3031        assert!(
3032            requests
3033                .iter()
3034                .all(|request| !request.record_telemetry_content),
3035            "agent-owned repaired telemetry should clear provider request flags"
3036        );
3037    }
3038
3039    #[test]
3040    fn completion_calls_stream_item_serializes_and_deserializes_expected_shape() {
3041        let item: MultiTurnStreamItem<MockResponse> =
3042            MultiTurnStreamItem::CompletionCall(CompletionCall::new(2, usage(3, 4)));
3043
3044        let value = serde_json::to_value(&item).expect("serialize completion call event");
3045
3046        assert_eq!(
3047            value,
3048            serde_json::json!({
3049                "type": "completionCall",
3050                "call_index": 2,
3051                "usage": {
3052                    "input_tokens": 3,
3053                    "output_tokens": 4,
3054                    "total_tokens": 7,
3055                    "cached_input_tokens": 0,
3056                    "cache_creation_input_tokens": 0,
3057                    "tool_use_prompt_tokens": 0,
3058                    "reasoning_tokens": 0,
3059                }
3060            })
3061        );
3062
3063        let item: MultiTurnStreamItem<MockResponse> =
3064            serde_json::from_value(value).expect("deserialize completion call event");
3065        match item {
3066            MultiTurnStreamItem::CompletionCall(call_usage) => {
3067                assert_eq!(call_usage, CompletionCall::new(2, usage(3, 4)));
3068            }
3069            other => panic!("expected completion call event, got {other:?}"),
3070        }
3071
3072        let item: MultiTurnStreamItem<MockResponse> =
3073            MultiTurnStreamItem::CompletionCall(CompletionCall::new(3, Usage::new()));
3074        let value = serde_json::to_value(&item).expect("serialize missing usage event");
3075
3076        // Unreported usage serializes as a plain zero-valued object (Usage's
3077        // documented sentinel for missing provider metrics).
3078        assert_eq!(
3079            value,
3080            serde_json::json!({
3081                "type": "completionCall",
3082                "call_index": 3,
3083                "usage": {
3084                    "input_tokens": 0,
3085                    "output_tokens": 0,
3086                    "total_tokens": 0,
3087                    "cached_input_tokens": 0,
3088                    "cache_creation_input_tokens": 0,
3089                    "tool_use_prompt_tokens": 0,
3090                    "reasoning_tokens": 0,
3091                }
3092            })
3093        );
3094
3095        // Stream items serialized before the Option encoding was dropped used
3096        // `"usage": null`; they must still deserialize.
3097        let legacy: MultiTurnStreamItem<MockResponse> = serde_json::from_value(serde_json::json!({
3098            "type": "completionCall",
3099            "call_index": 3,
3100            "usage": null
3101        }))
3102        .expect("legacy null-usage event should deserialize");
3103        match legacy {
3104            MultiTurnStreamItem::CompletionCall(call) => {
3105                assert_eq!(call, CompletionCall::new(3, Usage::new()));
3106            }
3107            other => panic!("expected completion call event, got {other:?}"),
3108        }
3109    }
3110
3111    #[test]
3112    fn final_response_serializes_completion_calls_with_missing_usage() {
3113        let item: MultiTurnStreamItem<MockResponse> =
3114            MultiTurnStreamItem::final_response_with_completion_calls(
3115                OneOrMany::one(AssistantContent::text("done")),
3116                usage(3, 4),
3117                vec![
3118                    CompletionCall::new(0, Usage::new()),
3119                    CompletionCall::new(1, usage(3, 4)),
3120                ],
3121                None,
3122            );
3123
3124        if let MultiTurnStreamItem::FinalResponse(response) = &item {
3125            assert_eq!(response.requests(), 2);
3126        }
3127
3128        let value = serde_json::to_value(&item).expect("serialize final response");
3129
3130        assert_eq!(
3131            value.get("completion_calls"),
3132            Some(&serde_json::json!([
3133                {
3134                    "call_index": 0,
3135                    "usage": {
3136                        "input_tokens": 0,
3137                        "output_tokens": 0,
3138                        "total_tokens": 0,
3139                        "cached_input_tokens": 0,
3140                        "cache_creation_input_tokens": 0,
3141                        "tool_use_prompt_tokens": 0,
3142                        "reasoning_tokens": 0,
3143                    }
3144                },
3145                {
3146                    "call_index": 1,
3147                    "usage": {
3148                        "input_tokens": 3,
3149                        "output_tokens": 4,
3150                        "total_tokens": 7,
3151                        "cached_input_tokens": 0,
3152                        "cache_creation_input_tokens": 0,
3153                        "tool_use_prompt_tokens": 0,
3154                        "reasoning_tokens": 0,
3155                    }
3156                }
3157            ]))
3158        );
3159    }
3160
3161    fn streaming_text_then_final_model() -> MockCompletionModel {
3162        MockCompletionModel::from_stream_turns([[
3163            MockStreamEvent::text("hello"),
3164            MockStreamEvent::text(" world"),
3165            MockStreamEvent::final_response_with_total_tokens(3),
3166        ]])
3167    }
3168
3169    fn citation_metadata() -> serde_json::Value {
3170        serde_json::json!({
3171            "citations": [{
3172                "type": "web_search_result_location",
3173                "cited_text": "Claude Shannon was born in 1916.",
3174                "url": "https://example.com/shannon",
3175                "title": "Claude Shannon",
3176                "encrypted_index": "encrypted-reference"
3177            }]
3178        })
3179    }
3180
3181    fn streaming_cited_text_then_final_model() -> MockCompletionModel {
3182        MockCompletionModel::from_stream_turns([[
3183            MockStreamEvent::text_start(Some(citation_metadata())),
3184            MockStreamEvent::text("cited "),
3185            MockStreamEvent::text_start(None),
3186            MockStreamEvent::text("answer"),
3187            MockStreamEvent::final_response_with_total_tokens(3),
3188        ]])
3189    }
3190
3191    fn streaming_cited_text_then_tool_model() -> MockCompletionModel {
3192        MockCompletionModel::from_stream_turns([
3193            vec![
3194                MockStreamEvent::text_start(Some(citation_metadata())),
3195                MockStreamEvent::text("I need a tool. "),
3196                MockStreamEvent::tool_call(
3197                    "tool_call_1",
3198                    "add",
3199                    serde_json::json!({"x": 1, "y": 2}),
3200                )
3201                .with_call_id("call_1"),
3202                MockStreamEvent::final_response_with_total_tokens(4),
3203            ],
3204            vec![
3205                MockStreamEvent::text("done"),
3206                MockStreamEvent::final_response_with_total_tokens(6),
3207            ],
3208        ])
3209    }
3210
3211    fn streaming_final_only_model() -> MockCompletionModel {
3212        MockCompletionModel::from_stream_turns([[
3213            MockStreamEvent::final_response_with_total_tokens(1),
3214        ]])
3215    }
3216
3217    #[derive(Clone)]
3218    struct TerminateOnStreamFinish;
3219
3220    impl AgentHook for TerminateOnStreamFinish {
3221        async fn on_stream_response_finish(
3222            &self,
3223            _ctx: &HookContext,
3224            event: StreamResponseFinish<'_>,
3225        ) -> ObservationAction {
3226            match event {
3227                StreamResponseFinish { .. } => {
3228                    ObservationAction::stop("stop after completion call")
3229                }
3230                _ => ObservationAction::continue_run(),
3231            }
3232        }
3233    }
3234
3235    type RecordedToolCallDelta = (String, String, Option<String>, String);
3236
3237    #[derive(Clone)]
3238    struct RepairDefaultApiHook;
3239
3240    impl AgentHook for RepairDefaultApiHook {
3241        async fn on_invalid_tool_call(
3242            &self,
3243            _ctx: &HookContext,
3244            event: &InvalidToolCallContext,
3245        ) -> Option<InvalidToolCallAction> {
3246            Some(match event {
3247                context => {
3248                    assert_eq!(context.tool_name, "default_api");
3249                    InvalidToolCallAction::repair("add")
3250                }
3251                _ => InvalidToolCallAction::fail(),
3252            })
3253        }
3254    }
3255
3256    #[derive(Clone)]
3257    struct RetryDefaultApiHook;
3258
3259    impl AgentHook for RetryDefaultApiHook {
3260        async fn on_invalid_tool_call(
3261            &self,
3262            _ctx: &HookContext,
3263            event: &InvalidToolCallContext,
3264        ) -> Option<InvalidToolCallAction> {
3265            Some(match event {
3266                context => {
3267                    assert_eq!(context.tool_name, "default_api");
3268                    if let Some(args) = context.args.as_deref() {
3269                        assert!(!args.is_empty());
3270                    }
3271                    InvalidToolCallAction::retry("Use the add tool instead")
3272                }
3273                _ => InvalidToolCallAction::fail(),
3274            })
3275        }
3276    }
3277
3278    #[derive(Clone)]
3279    struct SkipDefaultApiHook;
3280
3281    impl AgentHook for SkipDefaultApiHook {
3282        async fn on_invalid_tool_call(
3283            &self,
3284            _ctx: &HookContext,
3285            event: &InvalidToolCallContext,
3286        ) -> Option<InvalidToolCallAction> {
3287            Some(match event {
3288                context => {
3289                    assert_eq!(context.tool_name, "default_api");
3290                    InvalidToolCallAction::skip("default_api was skipped")
3291                }
3292                _ => InvalidToolCallAction::fail(),
3293            })
3294        }
3295    }
3296
3297    #[derive(Clone, Default)]
3298    struct RecordingInvalidToolCallHook {
3299        contexts: Arc<Mutex<Vec<InvalidToolCallContext>>>,
3300    }
3301
3302    impl RecordingInvalidToolCallHook {
3303        fn observed(&self) -> Vec<InvalidToolCallContext> {
3304            self.contexts
3305                .lock()
3306                .expect("invalid tool context records mutex was poisoned")
3307                .clone()
3308        }
3309    }
3310
3311    impl AgentHook for RecordingInvalidToolCallHook {
3312        async fn on_invalid_tool_call(
3313            &self,
3314            _ctx: &HookContext,
3315            event: &InvalidToolCallContext,
3316        ) -> Option<InvalidToolCallAction> {
3317            Some(match event {
3318                context => {
3319                    self.contexts
3320                        .lock()
3321                        .expect("invalid tool context records mutex was poisoned")
3322                        .push(context.clone());
3323                    InvalidToolCallAction::fail()
3324                }
3325                _ => InvalidToolCallAction::fail(),
3326            })
3327        }
3328    }
3329
3330    #[derive(Clone, Default)]
3331    struct RecordingToolCallDeltaHook {
3332        deltas: Arc<Mutex<Vec<RecordedToolCallDelta>>>,
3333    }
3334
3335    impl RecordingToolCallDeltaHook {
3336        fn observed(&self) -> Vec<RecordedToolCallDelta> {
3337            self.deltas
3338                .lock()
3339                .expect("tool call delta hook records mutex was poisoned")
3340                .clone()
3341        }
3342    }
3343
3344    impl AgentHook for RecordingToolCallDeltaHook {
3345        async fn on_tool_call_delta(
3346            &self,
3347            _ctx: &HookContext,
3348            event: ToolCallDelta<'_>,
3349        ) -> ObservationAction {
3350            match event {
3351                ToolCallDelta {
3352                    tool_call_id,
3353                    internal_call_id,
3354                    tool_name,
3355                    delta,
3356                } => {
3357                    let record = (
3358                        tool_call_id.to_string(),
3359                        internal_call_id.to_string(),
3360                        tool_name.map(str::to_string),
3361                        delta.to_string(),
3362                    );
3363                    self.deltas
3364                        .lock()
3365                        .expect("tool call delta hook records mutex was poisoned")
3366                        .push(record);
3367                    ObservationAction::continue_run()
3368                }
3369                _ => ObservationAction::continue_run(),
3370            }
3371        }
3372    }
3373
3374    #[derive(Clone, Default)]
3375    struct RecordingTextDeltaHook {
3376        deltas: Arc<Mutex<Vec<(String, String)>>>,
3377    }
3378
3379    impl RecordingTextDeltaHook {
3380        fn observed(&self) -> Vec<(String, String)> {
3381            self.deltas
3382                .lock()
3383                .expect("text delta hook records mutex was poisoned")
3384                .clone()
3385        }
3386    }
3387
3388    impl AgentHook for RecordingTextDeltaHook {
3389        async fn on_text_delta(
3390            &self,
3391            _ctx: &HookContext,
3392            event: TextDelta<'_>,
3393        ) -> ObservationAction {
3394            match event {
3395                TextDelta { delta, aggregated } => {
3396                    let record = (delta.to_string(), aggregated.to_string());
3397                    self.deltas
3398                        .lock()
3399                        .expect("text delta hook records mutex was poisoned")
3400                        .push(record);
3401                    ObservationAction::continue_run()
3402                }
3403                _ => ObservationAction::continue_run(),
3404            }
3405        }
3406    }
3407
3408    #[derive(Clone)]
3409    struct RecordingTextAndSkipInvalidToolHook {
3410        text: RecordingTextDeltaHook,
3411    }
3412
3413    impl AgentHook for RecordingTextAndSkipInvalidToolHook {
3414        async fn on_text_delta(
3415            &self,
3416            ctx: &HookContext,
3417            event: TextDelta<'_>,
3418        ) -> ObservationAction {
3419            self.text.on_text_delta(ctx, event).await
3420        }
3421        async fn on_invalid_tool_call(
3422            &self,
3423            ctx: &HookContext,
3424            event: &InvalidToolCallContext,
3425        ) -> Option<InvalidToolCallAction> {
3426            SkipDefaultApiHook.on_invalid_tool_call(ctx, event).await
3427        }
3428    }
3429
3430    #[derive(Clone)]
3431    struct RecordingTextAndRetryInvalidToolHook {
3432        text: RecordingTextDeltaHook,
3433    }
3434
3435    impl AgentHook for RecordingTextAndRetryInvalidToolHook {
3436        async fn on_text_delta(
3437            &self,
3438            ctx: &HookContext,
3439            event: TextDelta<'_>,
3440        ) -> ObservationAction {
3441            self.text.on_text_delta(ctx, event).await
3442        }
3443        async fn on_invalid_tool_call(
3444            &self,
3445            ctx: &HookContext,
3446            event: &InvalidToolCallContext,
3447        ) -> Option<InvalidToolCallAction> {
3448            RetryDefaultApiHook.on_invalid_tool_call(ctx, event).await
3449        }
3450    }
3451
3452    #[derive(Clone)]
3453    struct RecordingDeltaAndRetryInvalidToolHook {
3454        delta: RecordingToolCallDeltaHook,
3455    }
3456
3457    impl AgentHook for RecordingDeltaAndRetryInvalidToolHook {
3458        async fn on_tool_call_delta(
3459            &self,
3460            ctx: &HookContext,
3461            event: ToolCallDelta<'_>,
3462        ) -> ObservationAction {
3463            self.delta.on_tool_call_delta(ctx, event).await
3464        }
3465        async fn on_invalid_tool_call(
3466            &self,
3467            ctx: &HookContext,
3468            event: &InvalidToolCallContext,
3469        ) -> Option<InvalidToolCallAction> {
3470            RetryDefaultApiHook.on_invalid_tool_call(ctx, event).await
3471        }
3472    }
3473
3474    #[derive(Clone)]
3475    struct RecordingDeltaAndSkipInvalidToolHook {
3476        delta: RecordingToolCallDeltaHook,
3477    }
3478
3479    impl AgentHook for RecordingDeltaAndSkipInvalidToolHook {
3480        async fn on_tool_call_delta(
3481            &self,
3482            ctx: &HookContext,
3483            event: ToolCallDelta<'_>,
3484        ) -> ObservationAction {
3485            self.delta.on_tool_call_delta(ctx, event).await
3486        }
3487        async fn on_invalid_tool_call(
3488            &self,
3489            ctx: &HookContext,
3490            event: &InvalidToolCallContext,
3491        ) -> Option<InvalidToolCallAction> {
3492            SkipDefaultApiHook.on_invalid_tool_call(ctx, event).await
3493        }
3494    }
3495
3496    #[derive(Clone, Default)]
3497    struct TerminatingToolCallDeltaHook {
3498        deltas: Arc<Mutex<Vec<RecordedToolCallDelta>>>,
3499    }
3500
3501    impl TerminatingToolCallDeltaHook {
3502        fn observed(&self) -> Vec<RecordedToolCallDelta> {
3503            self.deltas
3504                .lock()
3505                .expect("tool call delta hook records mutex was poisoned")
3506                .clone()
3507        }
3508    }
3509
3510    impl AgentHook for TerminatingToolCallDeltaHook {
3511        async fn on_tool_call_delta(
3512            &self,
3513            _ctx: &HookContext,
3514            event: ToolCallDelta<'_>,
3515        ) -> ObservationAction {
3516            match event {
3517                ToolCallDelta {
3518                    tool_call_id,
3519                    internal_call_id,
3520                    tool_name,
3521                    delta,
3522                } => {
3523                    let record = (
3524                        tool_call_id.to_string(),
3525                        internal_call_id.to_string(),
3526                        tool_name.map(str::to_string),
3527                        delta.to_string(),
3528                    );
3529                    self.deltas
3530                        .lock()
3531                        .expect("tool call delta hook records mutex was poisoned")
3532                        .push(record);
3533                    ObservationAction::stop("stop on tool call delta")
3534                }
3535                _ => ObservationAction::continue_run(),
3536            }
3537        }
3538    }
3539
3540    fn text_metadata(content: &OneOrMany<AssistantContent>) -> Option<&serde_json::Value> {
3541        content.iter().find_map(|item| match item {
3542            AssistantContent::Text(text) => text.additional_params.as_ref(),
3543            _ => None,
3544        })
3545    }
3546
3547    #[tokio::test]
3548    async fn stream_prompt_continues_after_tool_call_turn() {
3549        let model = streaming_tool_then_text_model();
3550        let recorded = model.clone();
3551        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
3552        let empty_history: &[Message] = &[];
3553
3554        let mut stream = agent
3555            .stream_prompt("do tool work")
3556            .history(empty_history)
3557            .max_turns(3)
3558            .await;
3559        let mut saw_tool_call = false;
3560        let mut saw_tool_result = false;
3561        let mut saw_final_response = false;
3562        let mut final_text = String::new();
3563        let mut final_response_text = None;
3564        let mut final_history = None;
3565
3566        while let Some(item) = stream.next().await {
3567            match item {
3568                Ok(MultiTurnStreamItem::StreamAssistantItem(
3569                    StreamedAssistantContent::ToolCall { .. },
3570                )) => {
3571                    saw_tool_call = true;
3572                }
3573                Ok(MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult {
3574                    ..
3575                })) => {
3576                    saw_tool_result = true;
3577                }
3578                Ok(MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Text(
3579                    text,
3580                ))) => {
3581                    final_text.push_str(&text.text);
3582                }
3583                Ok(MultiTurnStreamItem::FinalResponse(res)) => {
3584                    saw_final_response = true;
3585                    final_response_text = Some(res.output().to_owned());
3586                    final_history = res.messages().map(|history| history.to_vec());
3587                    break;
3588                }
3589                Ok(_) => {}
3590                Err(err) => panic!("unexpected streaming error: {err:?}"),
3591            }
3592        }
3593
3594        assert!(saw_tool_call);
3595        assert!(saw_tool_result);
3596        assert!(saw_final_response);
3597        assert_eq!(final_text, "done");
3598        assert_eq!(final_response_text.as_deref(), Some("done"));
3599        let history = final_history.expect("expected final response history");
3600        assert!(history.iter().any(|message| matches!(
3601            message,
3602            Message::Assistant { content, .. }
3603                if content.iter().any(|item| matches!(
3604                    item,
3605                    AssistantContent::Text(text) if text.text == "done"
3606                ))
3607        )));
3608        let requests = recorded.requests();
3609        assert_eq!(requests.len(), 2);
3610        assert!(validate_follow_up_tool_history(&requests[1]).is_ok());
3611    }
3612
3613    /// `StreamingPromptRequest::tool_concurrency` reaches the runner: two
3614    /// barrier-synchronized tools in a streamed turn only finish if they run
3615    /// concurrently. At `tool_concurrency(2)` the stream completes; sequential
3616    /// execution would block on the first tool forever, so the timeout asserts
3617    /// the public builder actually enables concurrency on the streaming path.
3618    #[tokio::test]
3619    async fn streaming_prompt_request_tool_concurrency_runs_tools_concurrently() {
3620        let barrier = Arc::new(tokio::sync::Barrier::new(2));
3621        let model = MockCompletionModel::from_stream_turns([
3622            vec![
3623                MockStreamEvent::tool_call("b1", "barrier_tool", serde_json::json!({})),
3624                MockStreamEvent::tool_call("b2", "barrier_tool", serde_json::json!({})),
3625                MockStreamEvent::final_response_with_total_tokens(0),
3626            ],
3627            vec![
3628                MockStreamEvent::text("done"),
3629                MockStreamEvent::final_response_with_total_tokens(0),
3630            ],
3631        ]);
3632        let agent = AgentBuilder::new(model)
3633            .tool(MockBarrierTool::new(barrier))
3634            .build();
3635
3636        let drive = async {
3637            let mut stream = agent
3638                .stream_prompt("hit the barrier twice")
3639                .max_turns(3)
3640                .tool_concurrency(2)
3641                .await;
3642            while let Some(item) = stream.next().await {
3643                item.unwrap_or_else(|err| panic!("unexpected streaming error: {err:?}"));
3644            }
3645        };
3646
3647        tokio::time::timeout(Duration::from_secs(5), drive)
3648            .await
3649            .expect("streamed tools must run concurrently, not deadlock at the barrier");
3650    }
3651
3652    /// The streaming driver threads the per-call `ToolContext` to executed
3653    /// tools, exactly like the blocking path.
3654    #[tokio::test]
3655    async fn tool_context_reaches_tool_through_streaming_loop() {
3656        let model = MockCompletionModel::from_stream_turns([
3657            vec![
3658                MockStreamEvent::tool_call("tool_call_1", "context_probe", serde_json::json!({}))
3659                    .with_call_id("call_1"),
3660                MockStreamEvent::final_response_with_total_tokens(4),
3661            ],
3662            vec![
3663                MockStreamEvent::text("done"),
3664                MockStreamEvent::final_response_with_total_tokens(6),
3665            ],
3666        ]);
3667        let probe = MockContextProbeTool::default();
3668        let agent = AgentBuilder::new(model).tool(probe.clone()).build();
3669        let empty_history: &[Message] = &[];
3670
3671        let mut tool_context = ToolContext::new();
3672        tool_context.insert(SessionId("xyz-789".to_string()));
3673
3674        let mut stream = agent
3675            .stream_prompt("do tool work")
3676            .tool_context(tool_context)
3677            .history(empty_history)
3678            .max_turns(3)
3679            .await;
3680
3681        while let Some(item) = stream.next().await {
3682            match item {
3683                Ok(MultiTurnStreamItem::FinalResponse(_)) => break,
3684                Err(err) => panic!("unexpected streaming error: {err:?}"),
3685                Ok(_) => {}
3686            }
3687        }
3688
3689        assert_eq!(probe.observed().as_deref(), Some("session:xyz-789"));
3690    }
3691
3692    /// Streaming counterpart of the blocking empty-context default: when no
3693    /// [`ToolContext`] is supplied, the tool still receives a fresh empty
3694    /// context (observing `no-session`), not a stale value.
3695    #[tokio::test]
3696    async fn streaming_tool_runs_with_empty_context_when_none_supplied() {
3697        let model = MockCompletionModel::from_stream_turns([
3698            vec![
3699                MockStreamEvent::tool_call("tool_call_1", "context_probe", serde_json::json!({}))
3700                    .with_call_id("call_1"),
3701                MockStreamEvent::final_response_with_total_tokens(4),
3702            ],
3703            vec![
3704                MockStreamEvent::text("done"),
3705                MockStreamEvent::final_response_with_total_tokens(6),
3706            ],
3707        ]);
3708        let probe = MockContextProbeTool::default();
3709        let agent = AgentBuilder::new(model).tool(probe.clone()).build();
3710        let empty_history: &[Message] = &[];
3711
3712        let mut stream = agent
3713            .stream_prompt("do tool work")
3714            .history(empty_history)
3715            .max_turns(3)
3716            .await;
3717
3718        while let Some(item) = stream.next().await {
3719            match item {
3720                Ok(MultiTurnStreamItem::FinalResponse(_)) => break,
3721                Err(err) => panic!("unexpected streaming error: {err:?}"),
3722                Ok(_) => {}
3723            }
3724        }
3725
3726        assert_eq!(probe.observed().as_deref(), Some("no-session"));
3727    }
3728
3729    #[tokio::test]
3730    async fn unknown_tool_call_fails_before_streaming_second_request() {
3731        let model = MockCompletionModel::from_stream_turns([
3732            vec![
3733                MockStreamEvent::tool_call(
3734                    "tool_call_1",
3735                    "default_api",
3736                    serde_json::json!({"x": 1, "y": 2}),
3737                ),
3738                MockStreamEvent::final_response_with_total_tokens(4),
3739            ],
3740            vec![
3741                MockStreamEvent::text("should not be requested"),
3742                MockStreamEvent::final_response_with_total_tokens(6),
3743            ],
3744        ]);
3745        let recorded = model.clone();
3746        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
3747
3748        let mut stream = agent
3749            .stream_prompt("use the tool")
3750            .add_hook(PanicOnUnknownToolHook)
3751            .max_turns(3)
3752            .await;
3753        let mut saw_tool_call = false;
3754        let mut error = None;
3755
3756        while let Some(item) = stream.next().await {
3757            match item {
3758                Ok(MultiTurnStreamItem::StreamAssistantItem(
3759                    StreamedAssistantContent::ToolCall { .. },
3760                )) => {
3761                    saw_tool_call = true;
3762                }
3763                Ok(_) => {}
3764                Err(err) => {
3765                    error = Some(err);
3766                    break;
3767                }
3768            }
3769        }
3770
3771        assert!(!saw_tool_call);
3772        let error = error.expect("unknown model-emitted tool should fail");
3773        match error {
3774            StreamingError::Prompt(err) => match *err {
3775                PromptError::UnknownToolCall {
3776                    tool_name,
3777                    available_tools,
3778                    allowed_tools,
3779                    chat_history,
3780                } => {
3781                    assert_eq!(tool_name, "default_api");
3782                    assert_eq!(available_tools, vec!["add".to_string()]);
3783                    assert_eq!(allowed_tools, vec!["add".to_string()]);
3784                    assert!(history_contains_tool_call(&chat_history, "default_api"));
3785                }
3786                other => panic!("expected UnknownToolCall, got {other:?}"),
3787            },
3788            other => panic!("expected prompt streaming error, got {other:?}"),
3789        }
3790        assert_eq!(recorded.request_count(), 1);
3791    }
3792
3793    #[tokio::test]
3794    async fn invalid_tool_call_hook_can_repair_streaming_tool_name() {
3795        let model = MockCompletionModel::from_stream_turns([
3796            vec![
3797                MockStreamEvent::tool_call(
3798                    "tool_call_1",
3799                    "default_api",
3800                    serde_json::json!({"x": 2, "y": 3}),
3801                ),
3802                MockStreamEvent::final_response_with_total_tokens(4),
3803            ],
3804            vec![
3805                MockStreamEvent::text("done"),
3806                MockStreamEvent::final_response_with_total_tokens(6),
3807            ],
3808        ]);
3809        let recorded = model.clone();
3810        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
3811
3812        let mut stream = agent
3813            .stream_prompt("use the tool")
3814            .add_hook(RepairDefaultApiHook)
3815            .max_turns(3)
3816            .history(Vec::<Message>::new())
3817            .await;
3818        let mut saw_repaired_tool_call = false;
3819        let mut saw_tool_result = false;
3820        let mut final_response_text = None;
3821
3822        while let Some(item) = stream.next().await {
3823            match item {
3824                Ok(MultiTurnStreamItem::StreamAssistantItem(
3825                    StreamedAssistantContent::ToolCall { tool_call, .. },
3826                )) => {
3827                    assert_eq!(tool_call.function.name, "add");
3828                    saw_repaired_tool_call = true;
3829                }
3830                Ok(MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult {
3831                    tool_result,
3832                    ..
3833                })) => {
3834                    assert!(tool_result.content.iter().any(|content| {
3835                        matches!(
3836                            content,
3837                            ToolResultContent::Json { value }
3838                                if value == &serde_json::json!(5)
3839                        )
3840                    }));
3841                    saw_tool_result = true;
3842                }
3843                Ok(MultiTurnStreamItem::FinalResponse(response)) => {
3844                    final_response_text = Some(response.output().to_string());
3845                    break;
3846                }
3847                Ok(_) => {}
3848                Err(err) => panic!("unexpected streaming error: {err:?}"),
3849            }
3850        }
3851
3852        assert!(saw_repaired_tool_call);
3853        assert!(saw_tool_result);
3854        assert_eq!(final_response_text.as_deref(), Some("done"));
3855        assert_eq!(recorded.request_count(), 2);
3856    }
3857
3858    #[tokio::test]
3859    async fn invalid_tool_call_context_uses_completed_streaming_tool_call_provider_id() {
3860        let invalid_hook = RecordingInvalidToolCallHook::default();
3861        let model = MockCompletionModel::from_stream_turns([
3862            vec![
3863                MockStreamEvent::tool_call(
3864                    "tool_call_1",
3865                    "default_api",
3866                    serde_json::json!({"x": 2, "y": 3}),
3867                )
3868                .with_call_id("provider_call_1"),
3869                MockStreamEvent::final_response_with_total_tokens(4),
3870            ],
3871            vec![
3872                MockStreamEvent::text("should not be requested"),
3873                MockStreamEvent::final_response_with_total_tokens(6),
3874            ],
3875        ]);
3876        let recorded = model.clone();
3877        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
3878
3879        let mut stream = agent
3880            .stream_prompt("use the tool")
3881            .add_hook(invalid_hook.clone())
3882            .max_turns(3)
3883            .await;
3884        let mut error = None;
3885
3886        while let Some(item) = stream.next().await {
3887            if let Err(err) = item {
3888                error = Some(err);
3889                break;
3890            }
3891        }
3892
3893        assert!(error.is_some(), "invalid tool should fail");
3894        assert_eq!(recorded.request_count(), 1);
3895        let contexts = invalid_hook.observed();
3896        assert_eq!(contexts.len(), 1);
3897        let context = &contexts[0];
3898        assert_eq!(context.tool_name, "default_api");
3899        assert_eq!(context.tool_call_id.as_deref(), Some("tool_call_1"));
3900        assert!(context.internal_call_id.is_some());
3901        assert!(context.is_streaming);
3902    }
3903
3904    #[tokio::test]
3905    async fn invalid_tool_call_hook_skip_emits_streaming_tool_result() {
3906        let add_calls = Arc::new(AtomicU32::new(0));
3907        let model = MockCompletionModel::from_stream_turns([
3908            vec![
3909                MockStreamEvent::tool_call(
3910                    "tool_call_1",
3911                    "default_api",
3912                    serde_json::json!({"x": 2, "y": 3}),
3913                )
3914                .with_call_id("call_1"),
3915                MockStreamEvent::final_response_with_total_tokens(4),
3916            ],
3917            vec![
3918                MockStreamEvent::text("continued"),
3919                MockStreamEvent::final_response_with_total_tokens(6),
3920            ],
3921        ]);
3922        let recorded = model.clone();
3923        let agent = AgentBuilder::new(model)
3924            .tool(CountingAddTool {
3925                calls: add_calls.clone(),
3926            })
3927            .build();
3928
3929        let mut stream = agent
3930            .stream_prompt("use the tool")
3931            .add_hook(SkipDefaultApiHook)
3932            .max_turns(3)
3933            .history(Vec::<Message>::new())
3934            .await;
3935        let mut skipped_tool_result = None;
3936        let mut final_response_text = None;
3937
3938        while let Some(item) = stream.next().await {
3939            match item {
3940                Ok(MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult {
3941                    tool_result,
3942                    internal_call_id,
3943                })) => {
3944                    assert!(!internal_call_id.is_empty());
3945                    skipped_tool_result = Some(tool_result);
3946                }
3947                Ok(MultiTurnStreamItem::FinalResponse(response)) => {
3948                    final_response_text = Some(response.output().to_string());
3949                    break;
3950                }
3951                Ok(_) => {}
3952                Err(err) => panic!("unexpected streaming error: {err:?}"),
3953            }
3954        }
3955
3956        let skipped_tool_result =
3957            skipped_tool_result.expect("skip recovery should emit a synthetic tool result");
3958        assert_eq!(skipped_tool_result.id, "tool_call_1");
3959        assert_eq!(skipped_tool_result.call_id.as_deref(), Some("call_1"));
3960        assert!(skipped_tool_result.content.iter().any(|content| matches!(
3961            content,
3962            ToolResultContent::Text(text) if text.text == "default_api was skipped"
3963        )));
3964        assert_eq!(final_response_text.as_deref(), Some("continued"));
3965        assert_eq!(add_calls.load(Ordering::SeqCst), 0);
3966
3967        let requests = recorded.requests();
3968        assert_eq!(requests.len(), 2);
3969        let follow_up_history = requests[1].chat_history.iter().cloned().collect::<Vec<_>>();
3970        assert!(matches!(
3971            follow_up_history.get(2),
3972            Some(Message::User { content })
3973                if content.iter().any(|item| matches!(
3974                    item,
3975                    UserContent::ToolResult(result)
3976                        if result.id == "tool_call_1"
3977                            && result.content.iter().any(|content| matches!(
3978                                content,
3979                                ToolResultContent::Text(text)
3980                                    if text.text == "default_api was skipped"
3981                            ))
3982                ))
3983        ));
3984    }
3985
3986    #[tokio::test]
3987    async fn invalid_tool_call_hook_retries_mixed_streaming_turn_without_executing_valid_call() {
3988        let add_calls = Arc::new(AtomicU32::new(0));
3989        let model = MockCompletionModel::from_stream_turns([
3990            vec![
3991                MockStreamEvent::text("checking "),
3992                MockStreamEvent::tool_call(
3993                    "tool_call_1",
3994                    "add",
3995                    serde_json::json!({"x": 2, "y": 3}),
3996                )
3997                .with_call_id("call_1"),
3998                MockStreamEvent::tool_call(
3999                    "tool_call_2",
4000                    "default_api",
4001                    serde_json::json!({"x": 4, "y": 5}),
4002                )
4003                .with_call_id("call_2"),
4004                MockStreamEvent::final_response_with_total_tokens(4),
4005            ],
4006            vec![
4007                MockStreamEvent::text("retried"),
4008                MockStreamEvent::final_response_with_total_tokens(6),
4009            ],
4010        ]);
4011        let recorded = model.clone();
4012        let agent = AgentBuilder::new(model)
4013            .tool(CountingAddTool {
4014                calls: add_calls.clone(),
4015            })
4016            .build();
4017
4018        let mut stream = agent
4019            .stream_prompt("use the tool")
4020            .add_hook(RetryDefaultApiHook)
4021            .max_turns(3)
4022            .history(Vec::<Message>::new())
4023            .max_invalid_tool_call_retries(1)
4024            .await;
4025        let mut completion_call_events = Vec::new();
4026        let mut final_response_text = None;
4027        let mut final_response_usage = Usage::new();
4028        let mut final_completion_calls = Vec::new();
4029
4030        while let Some(item) = stream.next().await {
4031            match item {
4032                Ok(MultiTurnStreamItem::CompletionCall(completion_call)) => {
4033                    completion_call_events.push(completion_call);
4034                }
4035                Ok(MultiTurnStreamItem::FinalResponse(response)) => {
4036                    final_response_text = Some(response.output().to_string());
4037                    final_response_usage = response.usage();
4038                    final_completion_calls = response.completion_calls().to_vec();
4039                    break;
4040                }
4041                Ok(_) => {}
4042                Err(err) => panic!("unexpected streaming error: {err:?}"),
4043            }
4044        }
4045
4046        assert_eq!(final_response_text.as_deref(), Some("retried"));
4047        assert_eq!(add_calls.load(Ordering::SeqCst), 0);
4048        let mut first_usage = Usage::new();
4049        first_usage.total_tokens = 4;
4050        let mut second_usage = Usage::new();
4051        second_usage.total_tokens = 6;
4052        let expected_completion_calls = vec![
4053            CompletionCall::new(0, first_usage),
4054            CompletionCall::new(1, second_usage),
4055        ];
4056        assert_eq!(completion_call_events, expected_completion_calls);
4057        assert_eq!(final_completion_calls, expected_completion_calls);
4058        assert_eq!(final_response_usage.total_tokens, 10);
4059
4060        let requests = recorded.requests();
4061        assert_eq!(requests.len(), 2);
4062        let retry_history = requests[1].chat_history.iter().cloned().collect::<Vec<_>>();
4063        assert_eq!(retry_history.len(), 3);
4064        assert!(matches!(
4065            retry_history.get(1),
4066            Some(Message::Assistant { content, .. })
4067                if content.iter().any(|item| matches!(
4068                    item,
4069                    AssistantContent::Text(text) if text.text == "checking "
4070                ))
4071                    && content.iter().any(|item| matches!(
4072                        item,
4073                        AssistantContent::ToolCall(tool_call)
4074                            if tool_call.id == "tool_call_1"
4075                                && tool_call.function.name == "add"
4076                    ))
4077                    && content.iter().any(|item| matches!(
4078                        item,
4079                        AssistantContent::ToolCall(tool_call)
4080                            if tool_call.id == "tool_call_2"
4081                                && tool_call.function.name == "default_api"
4082                    ))
4083        ));
4084        assert!(matches!(
4085            retry_history.get(2),
4086            Some(Message::User { content })
4087                if content.iter().filter(|item| matches!(item, UserContent::ToolResult(_))).count() == 2
4088                    && content.iter().any(|item| matches!(
4089                        item,
4090                        UserContent::ToolResult(result)
4091                            if result.id == "tool_call_1"
4092                                && result.content.iter().any(|content| matches!(
4093                                    content,
4094                                    ToolResultContent::Text(text)
4095                                        if text.text == TOOL_NOT_EXECUTED_DUE_TO_INVALID_PEER
4096                                ))
4097                    ))
4098                    && content.iter().any(|item| matches!(
4099                        item,
4100                        UserContent::ToolResult(result)
4101                            if result.id == "tool_call_2"
4102                                && result.content.iter().any(|content| matches!(
4103                                    content,
4104                                    ToolResultContent::Text(text)
4105                                        if text.text == "Use the add tool instead"
4106                                ))
4107                    ))
4108        ));
4109    }
4110
4111    #[tokio::test]
4112    async fn invalid_tool_call_hook_skips_mixed_streaming_turn_without_executing_valid_call() {
4113        let add_calls = Arc::new(AtomicU32::new(0));
4114        let model = MockCompletionModel::from_stream_turns([
4115            vec![
4116                MockStreamEvent::text("checking "),
4117                MockStreamEvent::tool_call(
4118                    "tool_call_1",
4119                    "add",
4120                    serde_json::json!({"x": 2, "y": 3}),
4121                )
4122                .with_call_id("call_1"),
4123                MockStreamEvent::tool_call(
4124                    "tool_call_2",
4125                    "default_api",
4126                    serde_json::json!({"x": 4, "y": 5}),
4127                )
4128                .with_call_id("call_2"),
4129                MockStreamEvent::final_response_with_total_tokens(4),
4130            ],
4131            vec![
4132                MockStreamEvent::text("continued"),
4133                MockStreamEvent::final_response_with_total_tokens(6),
4134            ],
4135        ]);
4136        let recorded = model.clone();
4137        let agent = AgentBuilder::new(model)
4138            .tool(CountingAddTool {
4139                calls: add_calls.clone(),
4140            })
4141            .build();
4142
4143        let mut stream = agent
4144            .stream_prompt("use the tool")
4145            .add_hook(SkipDefaultApiHook)
4146            .max_turns(3)
4147            .history(Vec::<Message>::new())
4148            .await;
4149        let mut skipped_tool_result = None;
4150        let mut final_response_text = None;
4151
4152        while let Some(item) = stream.next().await {
4153            match item {
4154                Ok(MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult {
4155                    tool_result,
4156                    ..
4157                })) => {
4158                    skipped_tool_result = Some(tool_result);
4159                }
4160                Ok(MultiTurnStreamItem::FinalResponse(response)) => {
4161                    final_response_text = Some(response.output().to_string());
4162                    break;
4163                }
4164                Ok(_) => {}
4165                Err(err) => panic!("unexpected streaming error: {err:?}"),
4166            }
4167        }
4168
4169        let skipped_tool_result =
4170            skipped_tool_result.expect("skip recovery should emit a synthetic tool result");
4171        assert_eq!(skipped_tool_result.id, "tool_call_2");
4172        assert_eq!(skipped_tool_result.call_id.as_deref(), Some("call_2"));
4173        assert_eq!(final_response_text.as_deref(), Some("continued"));
4174        assert_eq!(add_calls.load(Ordering::SeqCst), 0);
4175
4176        let requests = recorded.requests();
4177        assert_eq!(requests.len(), 2);
4178        let follow_up_history = requests[1].chat_history.iter().cloned().collect::<Vec<_>>();
4179        assert_eq!(follow_up_history.len(), 3);
4180        assert!(matches!(
4181            follow_up_history.get(1),
4182            Some(Message::Assistant { content, .. })
4183                if content.iter().any(|item| matches!(
4184                    item,
4185                    AssistantContent::Text(text) if text.text == "checking "
4186                ))
4187                    && content.iter().any(|item| matches!(
4188                        item,
4189                        AssistantContent::ToolCall(tool_call)
4190                            if tool_call.id == "tool_call_1"
4191                                && tool_call.function.name == "add"
4192                    ))
4193                    && content.iter().any(|item| matches!(
4194                        item,
4195                        AssistantContent::ToolCall(tool_call)
4196                            if tool_call.id == "tool_call_2"
4197                                && tool_call.function.name == "default_api"
4198                    ))
4199        ));
4200        assert!(matches!(
4201            follow_up_history.get(2),
4202            Some(Message::User { content })
4203                if content.iter().filter(|item| matches!(item, UserContent::ToolResult(_))).count() == 2
4204                    && content.iter().any(|item| matches!(
4205                        item,
4206                        UserContent::ToolResult(result)
4207                            if result.id == "tool_call_1"
4208                                && result.call_id.as_deref() == Some("call_1")
4209                                && result.content.iter().any(|content| matches!(
4210                                    content,
4211                                    ToolResultContent::Text(text)
4212                                        if text.text == TOOL_NOT_EXECUTED_DUE_TO_INVALID_PEER
4213                                ))
4214                    ))
4215                    && content.iter().any(|item| matches!(
4216                        item,
4217                        UserContent::ToolResult(result)
4218                            if result.id == "tool_call_2"
4219                                && result.call_id.as_deref() == Some("call_2")
4220                                && result.content.iter().any(|content| matches!(
4221                                    content,
4222                                    ToolResultContent::Text(text)
4223                                        if text.text == "default_api was skipped"
4224                                ))
4225            ))
4226        ));
4227    }
4228
4229    #[tokio::test]
4230    async fn invalid_completed_tool_call_skip_preserves_streaming_reasoning_history() {
4231        let model = MockCompletionModel::from_stream_turns([
4232            vec![
4233                MockStreamEvent::text("checking "),
4234                MockStreamEvent::reasoning("reasoned step").with_reasoning_id("rs_1"),
4235                MockStreamEvent::tool_call(
4236                    "tool_call_1",
4237                    "default_api",
4238                    serde_json::json!({"x": 2, "y": 3}),
4239                ),
4240                MockStreamEvent::final_response_with_total_tokens(4),
4241            ],
4242            vec![
4243                MockStreamEvent::text("continued"),
4244                MockStreamEvent::final_response_with_total_tokens(6),
4245            ],
4246        ]);
4247        let recorded = model.clone();
4248        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
4249
4250        let mut stream = agent
4251            .stream_prompt("use the tool")
4252            .add_hook(SkipDefaultApiHook)
4253            .max_turns(3)
4254            .history(Vec::<Message>::new())
4255            .await;
4256
4257        while let Some(item) = stream.next().await {
4258            match item {
4259                Ok(MultiTurnStreamItem::FinalResponse(_)) => break,
4260                Ok(_) => {}
4261                Err(err) => panic!("unexpected streaming error: {err:?}"),
4262            }
4263        }
4264
4265        let requests = recorded.requests();
4266        assert_eq!(requests.len(), 2);
4267        let follow_up_history = requests[1].chat_history.iter().cloned().collect::<Vec<_>>();
4268        assert!(history_contains_text(&follow_up_history, "checking "));
4269        assert!(assistant_reasoning_precedes_tool_call(
4270            &follow_up_history,
4271            "reasoned step",
4272            "default_api"
4273        ));
4274        assert!(
4275            assistant_reasoning_precedes_text_and_tool_call(
4276                &follow_up_history,
4277                "reasoned step",
4278                "checking ",
4279                "default_api"
4280            ),
4281            "{follow_up_history:?}"
4282        );
4283    }
4284
4285    #[tokio::test]
4286    async fn invalid_name_delta_retry_preserves_streaming_reasoning_history() {
4287        let model = MockCompletionModel::from_stream_turns([
4288            vec![
4289                MockStreamEvent::reasoning_delta(Some("rs_1"), "delta reason"),
4290                MockStreamEvent::tool_call_arguments_delta(
4291                    "tool_call_1",
4292                    "internal_1",
4293                    r#"{"x":2,"y":3}"#,
4294                ),
4295                MockStreamEvent::tool_call_name_delta("tool_call_1", "internal_1", "default_api"),
4296                MockStreamEvent::final_response_with_total_tokens(4),
4297            ],
4298            vec![
4299                MockStreamEvent::text("retried"),
4300                MockStreamEvent::final_response_with_total_tokens(6),
4301            ],
4302        ]);
4303        let recorded = model.clone();
4304        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
4305
4306        let mut stream = agent
4307            .stream_prompt("use the tool")
4308            .add_hook(RetryDefaultApiHook)
4309            .max_turns(3)
4310            .history(Vec::<Message>::new())
4311            .max_invalid_tool_call_retries(1)
4312            .await;
4313
4314        while let Some(item) = stream.next().await {
4315            match item {
4316                Ok(MultiTurnStreamItem::FinalResponse(_)) => break,
4317                Ok(_) => {}
4318                Err(err) => panic!("unexpected streaming error: {err:?}"),
4319            }
4320        }
4321
4322        let requests = recorded.requests();
4323        assert_eq!(requests.len(), 2);
4324        let retry_history = requests[1].chat_history.iter().cloned().collect::<Vec<_>>();
4325        assert!(assistant_reasoning_precedes_tool_call(
4326            &retry_history,
4327            "delta reason",
4328            "default_api"
4329        ));
4330    }
4331
4332    #[tokio::test]
4333    async fn invalid_tool_call_hook_skip_resets_streaming_text_delta_state() {
4334        let text_hook = RecordingTextDeltaHook::default();
4335        let model = MockCompletionModel::from_stream_turns([
4336            vec![
4337                MockStreamEvent::text("stale "),
4338                MockStreamEvent::tool_call(
4339                    "tool_call_1",
4340                    "default_api",
4341                    serde_json::json!({"x": 2, "y": 3}),
4342                ),
4343                MockStreamEvent::final_response_with_total_tokens(4),
4344            ],
4345            vec![
4346                MockStreamEvent::text("fresh"),
4347                MockStreamEvent::final_response_with_total_tokens(6),
4348            ],
4349        ]);
4350        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
4351
4352        let mut stream = agent
4353            .stream_prompt("use the tool")
4354            .add_hook(RecordingTextAndSkipInvalidToolHook {
4355                text: text_hook.clone(),
4356            })
4357            .max_turns(3)
4358            .history(Vec::<Message>::new())
4359            .await;
4360
4361        while let Some(item) = stream.next().await {
4362            match item {
4363                Ok(MultiTurnStreamItem::FinalResponse(_)) => break,
4364                Ok(_) => {}
4365                Err(err) => panic!("unexpected streaming error: {err:?}"),
4366            }
4367        }
4368
4369        assert_eq!(
4370            text_hook.observed(),
4371            vec![
4372                ("stale ".to_string(), "stale ".to_string()),
4373                ("fresh".to_string(), "fresh".to_string()),
4374            ]
4375        );
4376    }
4377
4378    #[tokio::test]
4379    async fn invalid_tool_call_delta_retry_uses_structured_tool_feedback() {
4380        let delta_hook = RecordingToolCallDeltaHook::default();
4381        let add_calls = Arc::new(AtomicU32::new(0));
4382        let model = MockCompletionModel::from_stream_turns([
4383            vec![
4384                MockStreamEvent::text("checking "),
4385                MockStreamEvent::reasoning_delta(Some("rs_1"), "diagnostic reason"),
4386                MockStreamEvent::tool_call(
4387                    "tool_call_0",
4388                    "add",
4389                    serde_json::json!({"x": 1, "y": 2}),
4390                )
4391                .with_call_id("call_0"),
4392                MockStreamEvent::tool_call_arguments_delta(
4393                    "tool_call_1",
4394                    "internal_1",
4395                    r#"{"x":2,"y":3}"#,
4396                ),
4397                MockStreamEvent::tool_call_name_delta("tool_call_1", "internal_1", "default_api"),
4398                MockStreamEvent::final_response_with_total_tokens(4),
4399            ],
4400            vec![
4401                MockStreamEvent::text("retried"),
4402                MockStreamEvent::final_response_with_total_tokens(6),
4403            ],
4404        ]);
4405        let recorded = model.clone();
4406        let agent = AgentBuilder::new(model)
4407            .tool(CountingAddTool {
4408                calls: add_calls.clone(),
4409            })
4410            .build();
4411
4412        let mut stream = agent
4413            .stream_prompt("use the tool")
4414            .add_hook(RecordingDeltaAndRetryInvalidToolHook {
4415                delta: delta_hook.clone(),
4416            })
4417            .max_turns(3)
4418            .history(Vec::<Message>::new())
4419            .max_invalid_tool_call_retries(1)
4420            .await;
4421        let mut completion_call_events = Vec::new();
4422        let mut final_response_text = None;
4423        let mut final_response_usage = Usage::new();
4424        let mut final_completion_calls = Vec::new();
4425
4426        while let Some(item) = stream.next().await {
4427            match item {
4428                Ok(MultiTurnStreamItem::CompletionCall(completion_call)) => {
4429                    completion_call_events.push(completion_call);
4430                }
4431                Ok(MultiTurnStreamItem::StreamAssistantItem(
4432                    StreamedAssistantContent::ToolCallDelta { .. },
4433                )) => panic!("invalid tool-call delta should not be emitted"),
4434                Ok(MultiTurnStreamItem::FinalResponse(response)) => {
4435                    final_response_text = Some(response.output().to_string());
4436                    final_response_usage = response.usage();
4437                    final_completion_calls = response.completion_calls().to_vec();
4438                    break;
4439                }
4440                Ok(_) => {}
4441                Err(err) => panic!("unexpected streaming error: {err:?}"),
4442            }
4443        }
4444
4445        assert_eq!(final_response_text.as_deref(), Some("retried"));
4446        assert!(delta_hook.observed().is_empty());
4447        assert_eq!(add_calls.load(Ordering::SeqCst), 0);
4448        let mut first_usage = Usage::new();
4449        first_usage.total_tokens = 4;
4450        let mut second_usage = Usage::new();
4451        second_usage.total_tokens = 6;
4452        let expected_completion_calls = vec![
4453            CompletionCall::new(0, first_usage),
4454            CompletionCall::new(1, second_usage),
4455        ];
4456        assert_eq!(completion_call_events, expected_completion_calls);
4457        assert_eq!(final_completion_calls, expected_completion_calls);
4458        assert_eq!(final_response_usage.total_tokens, 10);
4459
4460        let requests = recorded.requests();
4461        assert_eq!(requests.len(), 2);
4462        let retry_history = requests[1].chat_history.iter().cloned().collect::<Vec<_>>();
4463        assert!(matches!(
4464            retry_history.get(1),
4465            Some(Message::Assistant { content, .. })
4466                if content.iter().any(|item| matches!(
4467                    item,
4468                    AssistantContent::Text(text) if text.text == "checking "
4469                ))
4470                    && content.iter().any(|item| matches!(
4471                        item,
4472                        AssistantContent::ToolCall(tool_call)
4473                            if tool_call.id == "tool_call_0"
4474                                && tool_call.function.name == "add"
4475                    ))
4476                    && content.iter().any(|item| matches!(
4477                    item,
4478                    AssistantContent::ToolCall(tool_call)
4479                        if tool_call.id == "tool_call_1"
4480                            && tool_call.function.name == "default_api"
4481                            && tool_call.function.arguments == serde_json::json!({"x": 2, "y": 3})
4482                ))
4483        ));
4484        assert!(matches!(
4485            retry_history.get(2),
4486            Some(Message::User { content })
4487                if content.iter().filter(|item| matches!(item, UserContent::ToolResult(_))).count() == 2
4488                    && content.iter().any(|item| matches!(
4489                        item,
4490                        UserContent::ToolResult(result)
4491                            if result.id == "tool_call_0"
4492                                && result.call_id.as_deref() == Some("call_0")
4493                                && result.content.iter().any(|content| matches!(
4494                                    content,
4495                                    ToolResultContent::Text(text)
4496                                        if text.text == TOOL_NOT_EXECUTED_DUE_TO_INVALID_PEER
4497                                ))
4498                    ))
4499                    && content.iter().any(|item| matches!(
4500                    item,
4501                    UserContent::ToolResult(result)
4502                        if result.id == "tool_call_1"
4503                            && result.content.iter().any(|content| matches!(
4504                                content,
4505                                ToolResultContent::Text(text)
4506                                    if text.text == "Use the add tool instead"
4507                            ))
4508                ))
4509        ));
4510    }
4511
4512    #[tokio::test]
4513    async fn invalid_tool_call_delta_context_includes_same_turn_history_and_tool_call_id() {
4514        let invalid_hook = RecordingInvalidToolCallHook::default();
4515        let model = MockCompletionModel::from_stream_turns([
4516            vec![
4517                MockStreamEvent::text("checking "),
4518                MockStreamEvent::reasoning_delta(Some("rs_1"), "diagnostic reason"),
4519                MockStreamEvent::tool_call(
4520                    "tool_call_0",
4521                    "add",
4522                    serde_json::json!({"x": 1, "y": 2}),
4523                )
4524                .with_call_id("call_0"),
4525                MockStreamEvent::tool_call_arguments_delta(
4526                    "tool_call_1",
4527                    "internal_1",
4528                    r#"{"x":2,"y":3}"#,
4529                ),
4530                MockStreamEvent::tool_call_name_delta("tool_call_1", "internal_1", "default_api"),
4531                MockStreamEvent::final_response_with_total_tokens(4),
4532            ],
4533            vec![
4534                MockStreamEvent::text("should not be requested"),
4535                MockStreamEvent::final_response_with_total_tokens(6),
4536            ],
4537        ]);
4538        let recorded = model.clone();
4539        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
4540
4541        let mut stream = agent
4542            .stream_prompt("use the tool")
4543            .add_hook(invalid_hook.clone())
4544            .max_turns(3)
4545            .await;
4546        let mut error = None;
4547
4548        while let Some(item) = stream.next().await {
4549            if let Err(err) = item {
4550                error = Some(err);
4551                break;
4552            }
4553        }
4554
4555        assert!(error.is_some(), "invalid name delta should fail");
4556        assert_eq!(recorded.request_count(), 1);
4557        let contexts = invalid_hook.observed();
4558        assert_eq!(contexts.len(), 1);
4559        let context = &contexts[0];
4560        assert_eq!(context.tool_name, "default_api");
4561        assert_eq!(context.tool_call_id.as_deref(), Some("tool_call_1"));
4562        assert_eq!(context.internal_call_id.as_deref(), Some("internal_1"));
4563        assert!(context.is_streaming);
4564        assert!(history_contains_text(&context.chat_history, "checking "));
4565        assert!(
4566            assistant_reasoning_precedes_tool_call(
4567                &context.chat_history,
4568                "diagnostic reason",
4569                "add"
4570            ),
4571            "{:?}",
4572            context.chat_history
4573        );
4574        assert!(history_contains_tool_call(&context.chat_history, "add"));
4575        assert!(history_contains_tool_call(
4576            &context.chat_history,
4577            "default_api"
4578        ));
4579    }
4580
4581    #[tokio::test]
4582    async fn invalid_tool_call_delta_retry_resets_streaming_text_delta_state() {
4583        let text_hook = RecordingTextDeltaHook::default();
4584        let model = MockCompletionModel::from_stream_turns([
4585            vec![
4586                MockStreamEvent::text("stale "),
4587                MockStreamEvent::tool_call_arguments_delta(
4588                    "tool_call_1",
4589                    "internal_1",
4590                    r#"{"x":2,"y":3}"#,
4591                ),
4592                MockStreamEvent::tool_call_name_delta("tool_call_1", "internal_1", "default_api"),
4593                MockStreamEvent::final_response_with_total_tokens(4),
4594            ],
4595            vec![
4596                MockStreamEvent::text("fresh"),
4597                MockStreamEvent::final_response_with_total_tokens(6),
4598            ],
4599        ]);
4600        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
4601
4602        let mut stream = agent
4603            .stream_prompt("use the tool")
4604            .add_hook(RecordingTextAndRetryInvalidToolHook {
4605                text: text_hook.clone(),
4606            })
4607            .max_turns(3)
4608            .history(Vec::<Message>::new())
4609            .max_invalid_tool_call_retries(1)
4610            .await;
4611
4612        while let Some(item) = stream.next().await {
4613            match item {
4614                Ok(MultiTurnStreamItem::FinalResponse(_)) => break,
4615                Ok(_) => {}
4616                Err(err) => panic!("unexpected streaming error: {err:?}"),
4617            }
4618        }
4619
4620        assert_eq!(
4621            text_hook.observed(),
4622            vec![
4623                ("stale ".to_string(), "stale ".to_string()),
4624                ("fresh".to_string(), "fresh".to_string()),
4625            ]
4626        );
4627    }
4628
4629    #[tokio::test]
4630    async fn invalid_tool_call_delta_skip_uses_structured_tool_feedback() {
4631        let delta_hook = RecordingToolCallDeltaHook::default();
4632        let add_calls = Arc::new(AtomicU32::new(0));
4633        let model = MockCompletionModel::from_stream_turns([
4634            vec![
4635                MockStreamEvent::text("checking "),
4636                MockStreamEvent::tool_call(
4637                    "tool_call_0",
4638                    "add",
4639                    serde_json::json!({"x": 1, "y": 2}),
4640                )
4641                .with_call_id("call_0"),
4642                MockStreamEvent::tool_call_arguments_delta(
4643                    "tool_call_1",
4644                    "internal_1",
4645                    r#"{"x":2,"y":3}"#,
4646                ),
4647                MockStreamEvent::tool_call_name_delta("tool_call_1", "internal_1", "default_api"),
4648                MockStreamEvent::final_response_with_total_tokens(4),
4649            ],
4650            vec![
4651                MockStreamEvent::text("continued"),
4652                MockStreamEvent::final_response_with_total_tokens(6),
4653            ],
4654        ]);
4655        let recorded = model.clone();
4656        let agent = AgentBuilder::new(model)
4657            .tool(CountingAddTool {
4658                calls: add_calls.clone(),
4659            })
4660            .build();
4661
4662        let mut stream = agent
4663            .stream_prompt("use the tool")
4664            .add_hook(RecordingDeltaAndSkipInvalidToolHook {
4665                delta: delta_hook.clone(),
4666            })
4667            .max_turns(3)
4668            .history(Vec::<Message>::new())
4669            .await;
4670        let mut skipped_tool_result = None;
4671        let mut final_response_text = None;
4672
4673        while let Some(item) = stream.next().await {
4674            match item {
4675                Ok(MultiTurnStreamItem::StreamAssistantItem(
4676                    StreamedAssistantContent::ToolCallDelta { .. },
4677                )) => panic!("invalid tool-call delta should not be emitted"),
4678                Ok(MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult {
4679                    tool_result,
4680                    internal_call_id,
4681                })) => {
4682                    assert_eq!(internal_call_id, "internal_1");
4683                    skipped_tool_result = Some(tool_result);
4684                }
4685                Ok(MultiTurnStreamItem::FinalResponse(response)) => {
4686                    final_response_text = Some(response.output().to_string());
4687                    break;
4688                }
4689                Ok(_) => {}
4690                Err(err) => panic!("unexpected streaming error: {err:?}"),
4691            }
4692        }
4693
4694        let skipped_tool_result =
4695            skipped_tool_result.expect("skip recovery should emit a synthetic tool result");
4696        assert_eq!(skipped_tool_result.id, "tool_call_1");
4697        assert!(skipped_tool_result.call_id.is_none());
4698        assert!(skipped_tool_result.content.iter().any(|content| matches!(
4699            content,
4700            ToolResultContent::Text(text) if text.text == "default_api was skipped"
4701        )));
4702        assert_eq!(final_response_text.as_deref(), Some("continued"));
4703        assert!(delta_hook.observed().is_empty());
4704        assert_eq!(add_calls.load(Ordering::SeqCst), 0);
4705
4706        let requests = recorded.requests();
4707        assert_eq!(requests.len(), 2);
4708        let follow_up_history = requests[1].chat_history.iter().cloned().collect::<Vec<_>>();
4709        assert!(matches!(
4710            follow_up_history.get(1),
4711            Some(Message::Assistant { content, .. })
4712                if content.iter().any(|item| matches!(
4713                    item,
4714                    AssistantContent::Text(text) if text.text == "checking "
4715                ))
4716                    && content.iter().any(|item| matches!(
4717                        item,
4718                        AssistantContent::ToolCall(tool_call)
4719                            if tool_call.id == "tool_call_0"
4720                                && tool_call.function.name == "add"
4721                    ))
4722                    && content.iter().any(|item| matches!(
4723                    item,
4724                    AssistantContent::ToolCall(tool_call)
4725                        if tool_call.id == "tool_call_1"
4726                            && tool_call.function.name == "default_api"
4727                            && tool_call.function.arguments == serde_json::json!({"x": 2, "y": 3})
4728                ))
4729        ));
4730        assert!(matches!(
4731            follow_up_history.get(2),
4732            Some(Message::User { content })
4733                if content.iter().filter(|item| matches!(item, UserContent::ToolResult(_))).count() == 2
4734                    && content.iter().any(|item| matches!(
4735                        item,
4736                        UserContent::ToolResult(result)
4737                            if result.id == "tool_call_0"
4738                                && result.call_id.as_deref() == Some("call_0")
4739                                && result.content.iter().any(|content| matches!(
4740                                    content,
4741                                    ToolResultContent::Text(text)
4742                                        if text.text == TOOL_NOT_EXECUTED_DUE_TO_INVALID_PEER
4743                                ))
4744                    ))
4745                    && content.iter().any(|item| matches!(
4746                    item,
4747                    UserContent::ToolResult(result)
4748                        if result.id == "tool_call_1"
4749                            && result.content.iter().any(|content| matches!(
4750                                content,
4751                                ToolResultContent::Text(text)
4752                                    if text.text == "default_api was skipped"
4753                            ))
4754                ))
4755        ));
4756    }
4757
4758    #[tokio::test]
4759    async fn streaming_retry_budget_exhaustion_history_contains_invalid_tool_call() {
4760        let model = MockCompletionModel::from_stream_turns([
4761            vec![
4762                MockStreamEvent::tool_call(
4763                    "tool_call_1",
4764                    "default_api",
4765                    serde_json::json!({"x": 1, "y": 2}),
4766                ),
4767                MockStreamEvent::final_response_with_total_tokens(4),
4768            ],
4769            vec![
4770                MockStreamEvent::text("should not be requested"),
4771                MockStreamEvent::final_response_with_total_tokens(6),
4772            ],
4773        ]);
4774        let recorded = model.clone();
4775        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
4776
4777        let mut stream = agent
4778            .stream_prompt("use the tool")
4779            .add_hook(RetryDefaultApiHook)
4780            .max_turns(3)
4781            .max_invalid_tool_call_retries(0)
4782            .await;
4783        let mut error = None;
4784
4785        while let Some(item) = stream.next().await {
4786            if let Err(err) = item {
4787                error = Some(err);
4788                break;
4789            }
4790        }
4791
4792        let error = error.expect("retry budget exhaustion should fail");
4793        match error {
4794            StreamingError::Prompt(err) => match *err {
4795                PromptError::UnknownToolCall {
4796                    tool_name,
4797                    chat_history,
4798                    ..
4799                } => {
4800                    assert_eq!(tool_name, "default_api");
4801                    assert!(history_contains_tool_call(&chat_history, "default_api"));
4802                }
4803                other => panic!("expected UnknownToolCall, got {other:?}"),
4804            },
4805            other => panic!("expected prompt streaming error, got {other:?}"),
4806        }
4807        assert_eq!(recorded.request_count(), 1);
4808    }
4809
4810    #[tokio::test]
4811    async fn streaming_name_delta_retry_budget_exhaustion_history_includes_same_turn_context() {
4812        let model = MockCompletionModel::from_stream_turns([
4813            vec![
4814                MockStreamEvent::text("checking "),
4815                MockStreamEvent::tool_call(
4816                    "tool_call_0",
4817                    "add",
4818                    serde_json::json!({"x": 1, "y": 2}),
4819                )
4820                .with_call_id("call_0"),
4821                MockStreamEvent::tool_call_arguments_delta(
4822                    "tool_call_1",
4823                    "internal_1",
4824                    r#"{"x":2,"y":3}"#,
4825                ),
4826                MockStreamEvent::tool_call_name_delta("tool_call_1", "internal_1", "default_api"),
4827                MockStreamEvent::final_response_with_total_tokens(4),
4828            ],
4829            vec![
4830                MockStreamEvent::text("should not be requested"),
4831                MockStreamEvent::final_response_with_total_tokens(6),
4832            ],
4833        ]);
4834        let recorded = model.clone();
4835        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
4836
4837        let mut stream = agent
4838            .stream_prompt("use the tool")
4839            .add_hook(RetryDefaultApiHook)
4840            .max_turns(3)
4841            .max_invalid_tool_call_retries(0)
4842            .await;
4843        let mut error = None;
4844
4845        while let Some(item) = stream.next().await {
4846            if let Err(err) = item {
4847                error = Some(err);
4848                break;
4849            }
4850        }
4851
4852        let error = error.expect("retry budget exhaustion should fail");
4853        match error {
4854            StreamingError::Prompt(err) => match *err {
4855                PromptError::UnknownToolCall {
4856                    tool_name,
4857                    chat_history,
4858                    ..
4859                } => {
4860                    assert_eq!(tool_name, "default_api");
4861                    assert!(history_contains_text(&chat_history, "checking "));
4862                    assert!(history_contains_tool_call(&chat_history, "add"));
4863                    assert!(history_contains_tool_call(&chat_history, "default_api"));
4864                }
4865                other => panic!("expected UnknownToolCall, got {other:?}"),
4866            },
4867            other => panic!("expected prompt streaming error, got {other:?}"),
4868        }
4869        assert_eq!(recorded.request_count(), 1);
4870    }
4871
4872    #[tokio::test]
4873    async fn completed_unknown_tool_call_after_text_fails_before_finish_hook_or_later_emit() {
4874        let add_calls = Arc::new(AtomicU32::new(0));
4875        let model = MockCompletionModel::from_stream_turns([
4876            vec![
4877                MockStreamEvent::text("thinking "),
4878                MockStreamEvent::tool_call(
4879                    "tool_call_1",
4880                    "default_api",
4881                    serde_json::json!({"x": 1, "y": 2}),
4882                ),
4883                MockStreamEvent::final_response_with_total_tokens(4),
4884            ],
4885            vec![
4886                MockStreamEvent::text("should not be requested"),
4887                MockStreamEvent::final_response_with_total_tokens(6),
4888            ],
4889        ]);
4890        let recorded = model.clone();
4891        let agent = AgentBuilder::new(model)
4892            .tool(CountingAddTool {
4893                calls: add_calls.clone(),
4894            })
4895            .build();
4896
4897        let mut stream = agent
4898            .stream_prompt("use the tool")
4899            .add_hook(PanicOnUnknownToolHook)
4900            .max_turns(3)
4901            .await;
4902        let mut saw_text = false;
4903        let mut saw_completion_call = false;
4904        let mut saw_final_response = false;
4905        let mut saw_tool_call = false;
4906        let mut saw_tool_result = false;
4907        let mut error = None;
4908
4909        while let Some(item) = stream.next().await {
4910            match item {
4911                Ok(MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Text(_))) => {
4912                    saw_text = true;
4913                }
4914                Ok(MultiTurnStreamItem::CompletionCall(_)) => {
4915                    saw_completion_call = true;
4916                }
4917                Ok(MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Final(
4918                    _,
4919                )))
4920                | Ok(MultiTurnStreamItem::FinalResponse(_)) => {
4921                    saw_final_response = true;
4922                }
4923                Ok(MultiTurnStreamItem::StreamAssistantItem(
4924                    StreamedAssistantContent::ToolCall { .. },
4925                )) => {
4926                    saw_tool_call = true;
4927                }
4928                Ok(MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult {
4929                    ..
4930                })) => {
4931                    saw_tool_result = true;
4932                }
4933                Ok(_) => {}
4934                Err(err) => {
4935                    error = Some(err);
4936                    break;
4937                }
4938            }
4939        }
4940
4941        assert!(saw_text);
4942        assert!(!saw_completion_call);
4943        assert!(!saw_final_response);
4944        assert!(!saw_tool_call);
4945        assert!(!saw_tool_result);
4946        assert_eq!(add_calls.load(Ordering::SeqCst), 0);
4947        let error = error.expect("completed unknown tool call should fail immediately");
4948        match error {
4949            StreamingError::Prompt(err) => match *err {
4950                PromptError::UnknownToolCall {
4951                    tool_name,
4952                    available_tools,
4953                    allowed_tools,
4954                    chat_history,
4955                } => {
4956                    assert_eq!(tool_name, "default_api");
4957                    assert_eq!(available_tools, vec!["add".to_string()]);
4958                    assert_eq!(allowed_tools, vec!["add".to_string()]);
4959                    assert!(history_contains_tool_call(&chat_history, "default_api"));
4960                }
4961                other => panic!("expected UnknownToolCall, got {other:?}"),
4962            },
4963            other => panic!("expected prompt streaming error, got {other:?}"),
4964        }
4965        assert_eq!(recorded.request_count(), 1);
4966    }
4967
4968    #[tokio::test]
4969    async fn mixed_streaming_tool_calls_fail_before_any_tool_execution() {
4970        let add_calls = Arc::new(AtomicU32::new(0));
4971        let model = MockCompletionModel::from_stream_turns([
4972            vec![
4973                MockStreamEvent::tool_call(
4974                    "tool_call_1",
4975                    "add",
4976                    serde_json::json!({"x": 1, "y": 2}),
4977                )
4978                .with_call_id("call_1"),
4979                MockStreamEvent::tool_call(
4980                    "tool_call_2",
4981                    "default_api",
4982                    serde_json::json!({"x": 3, "y": 4}),
4983                ),
4984                MockStreamEvent::final_response_with_total_tokens(4),
4985            ],
4986            vec![
4987                MockStreamEvent::text("should not be requested"),
4988                MockStreamEvent::final_response_with_total_tokens(6),
4989            ],
4990        ]);
4991        let recorded = model.clone();
4992        let agent = AgentBuilder::new(model)
4993            .tool(CountingAddTool {
4994                calls: add_calls.clone(),
4995            })
4996            .build();
4997
4998        let mut stream = agent
4999            .stream_prompt("use tools")
5000            .add_hook(PanicOnUnknownToolHook)
5001            .max_turns(3)
5002            .await;
5003        let mut saw_completion_call = false;
5004        let mut saw_tool_call = false;
5005        let mut saw_tool_result = false;
5006        let mut error = None;
5007
5008        while let Some(item) = stream.next().await {
5009            match item {
5010                Ok(MultiTurnStreamItem::CompletionCall(_)) => {
5011                    saw_completion_call = true;
5012                }
5013                Ok(MultiTurnStreamItem::StreamAssistantItem(
5014                    StreamedAssistantContent::ToolCall { .. },
5015                )) => {
5016                    saw_tool_call = true;
5017                }
5018                Ok(MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult {
5019                    ..
5020                })) => {
5021                    saw_tool_result = true;
5022                }
5023                Ok(_) => {}
5024                Err(err) => {
5025                    error = Some(err);
5026                    break;
5027                }
5028            }
5029        }
5030
5031        assert!(!saw_completion_call);
5032        assert!(!saw_tool_call);
5033        assert!(!saw_tool_result);
5034        assert_eq!(add_calls.load(Ordering::SeqCst), 0);
5035        let error = error.expect("mixed unknown streamed tool call should fail");
5036        match error {
5037            StreamingError::Prompt(err) => match *err {
5038                PromptError::UnknownToolCall {
5039                    tool_name,
5040                    available_tools,
5041                    allowed_tools,
5042                    chat_history,
5043                } => {
5044                    assert_eq!(tool_name, "default_api");
5045                    assert_eq!(available_tools, vec!["add".to_string()]);
5046                    assert_eq!(allowed_tools, vec!["add".to_string()]);
5047                    assert!(history_contains_tool_call(&chat_history, "default_api"));
5048                }
5049                other => panic!("expected UnknownToolCall, got {other:?}"),
5050            },
5051            other => panic!("expected prompt streaming error, got {other:?}"),
5052        }
5053        assert_eq!(recorded.request_count(), 1);
5054    }
5055
5056    #[tokio::test]
5057    async fn multiple_valid_streaming_tool_calls_execute_after_batch_validation() {
5058        let add_calls = Arc::new(AtomicU32::new(0));
5059        let subtract_calls = Arc::new(AtomicU32::new(0));
5060        let model = MockCompletionModel::from_stream_turns([
5061            vec![
5062                MockStreamEvent::tool_call(
5063                    "tool_call_1",
5064                    "add",
5065                    serde_json::json!({"x": 1, "y": 2}),
5066                )
5067                .with_call_id("call_1"),
5068                MockStreamEvent::tool_call(
5069                    "tool_call_2",
5070                    "subtract",
5071                    serde_json::json!({"x": 8, "y": 3}),
5072                )
5073                .with_call_id("call_2"),
5074                MockStreamEvent::final_response_with_total_tokens(4),
5075            ],
5076            vec![
5077                MockStreamEvent::text("done"),
5078                MockStreamEvent::final_response_with_total_tokens(6),
5079            ],
5080        ]);
5081        let recorded = model.clone();
5082        let agent = AgentBuilder::new(model)
5083            .tool(CountingAddTool {
5084                calls: add_calls.clone(),
5085            })
5086            .tool(CountingSubtractTool {
5087                calls: subtract_calls.clone(),
5088            })
5089            .build();
5090
5091        let mut stream = agent.stream_prompt("use tools").max_turns(3).await;
5092        let mut tool_call_names = Vec::new();
5093        let mut tool_result_ids = Vec::new();
5094        let mut final_response_text = None;
5095
5096        while let Some(item) = stream.next().await {
5097            match item {
5098                Ok(MultiTurnStreamItem::StreamAssistantItem(
5099                    StreamedAssistantContent::ToolCall { tool_call, .. },
5100                )) => {
5101                    tool_call_names.push(tool_call.function.name);
5102                }
5103                Ok(MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult {
5104                    tool_result,
5105                    ..
5106                })) => {
5107                    tool_result_ids.push(tool_result.id);
5108                }
5109                Ok(MultiTurnStreamItem::FinalResponse(response)) => {
5110                    final_response_text = Some(response.output().to_owned());
5111                    break;
5112                }
5113                Ok(_) => {}
5114                Err(err) => panic!("unexpected streaming error: {err:?}"),
5115            }
5116        }
5117
5118        assert_eq!(
5119            tool_call_names,
5120            vec!["add".to_string(), "subtract".to_string()]
5121        );
5122        assert_eq!(
5123            tool_result_ids,
5124            vec!["tool_call_1".to_string(), "tool_call_2".to_string()]
5125        );
5126        assert_eq!(add_calls.load(Ordering::SeqCst), 1);
5127        assert_eq!(subtract_calls.load(Ordering::SeqCst), 1);
5128        assert_eq!(final_response_text.as_deref(), Some("done"));
5129        assert_eq!(recorded.request_count(), 2);
5130    }
5131
5132    #[tokio::test]
5133    async fn disallowed_specific_tool_call_fails_before_streaming_second_request() {
5134        let model = MockCompletionModel::from_stream_turns([
5135            vec![
5136                MockStreamEvent::tool_call(
5137                    "tool_call_1",
5138                    "subtract",
5139                    serde_json::json!({"x": 3, "y": 1}),
5140                ),
5141                MockStreamEvent::final_response_with_total_tokens(4),
5142            ],
5143            vec![
5144                MockStreamEvent::text("should not be requested"),
5145                MockStreamEvent::final_response_with_total_tokens(6),
5146            ],
5147        ]);
5148        let recorded = model.clone();
5149        let agent = AgentBuilder::new(model)
5150            .tool(MockAddTool)
5151            .tool(MockSubtractTool)
5152            .tool_choice(ToolChoice::Specific {
5153                function_names: vec!["add".to_string()],
5154            })
5155            .build();
5156
5157        let mut stream = agent
5158            .stream_prompt("use the allowed tool")
5159            .add_hook(PanicOnUnknownToolHook)
5160            .max_turns(3)
5161            .await;
5162        let mut saw_tool_call = false;
5163        let mut error = None;
5164
5165        while let Some(item) = stream.next().await {
5166            match item {
5167                Ok(MultiTurnStreamItem::StreamAssistantItem(
5168                    StreamedAssistantContent::ToolCall { .. },
5169                )) => {
5170                    saw_tool_call = true;
5171                }
5172                Ok(_) => {}
5173                Err(err) => {
5174                    error = Some(err);
5175                    break;
5176                }
5177            }
5178        }
5179
5180        assert!(!saw_tool_call);
5181        let error = error.expect("disallowed model-emitted tool should fail");
5182        match error {
5183            StreamingError::Prompt(err) => match *err {
5184                PromptError::UnknownToolCall {
5185                    tool_name,
5186                    available_tools,
5187                    allowed_tools,
5188                    chat_history,
5189                } => {
5190                    assert_eq!(tool_name, "subtract");
5191                    assert_eq!(
5192                        available_tools,
5193                        vec!["add".to_string(), "subtract".to_string()]
5194                    );
5195                    assert_eq!(allowed_tools, vec!["add".to_string()]);
5196                    assert!(history_contains_tool_call(&chat_history, "subtract"));
5197                }
5198                other => panic!("expected UnknownToolCall, got {other:?}"),
5199            },
5200            other => panic!("expected prompt streaming error, got {other:?}"),
5201        }
5202        assert_eq!(recorded.request_count(), 1);
5203    }
5204
5205    #[tokio::test]
5206    async fn mixed_specific_tool_calls_fail_before_any_tool_execution() {
5207        let add_calls = Arc::new(AtomicU32::new(0));
5208        let model = MockCompletionModel::from_stream_turns([
5209            vec![
5210                MockStreamEvent::tool_call(
5211                    "tool_call_1",
5212                    "add",
5213                    serde_json::json!({"x": 1, "y": 2}),
5214                ),
5215                MockStreamEvent::tool_call(
5216                    "tool_call_2",
5217                    "subtract",
5218                    serde_json::json!({"x": 3, "y": 1}),
5219                ),
5220                MockStreamEvent::final_response_with_total_tokens(4),
5221            ],
5222            vec![
5223                MockStreamEvent::text("should not be requested"),
5224                MockStreamEvent::final_response_with_total_tokens(6),
5225            ],
5226        ]);
5227        let recorded = model.clone();
5228        let agent = AgentBuilder::new(model)
5229            .tool(CountingAddTool {
5230                calls: add_calls.clone(),
5231            })
5232            .tool(MockSubtractTool)
5233            .tool_choice(ToolChoice::Specific {
5234                function_names: vec!["add".to_string()],
5235            })
5236            .build();
5237
5238        let mut stream = agent
5239            .stream_prompt("use the allowed tool")
5240            .add_hook(PanicOnUnknownToolHook)
5241            .max_turns(3)
5242            .await;
5243        let mut saw_tool_call = false;
5244        let mut saw_tool_result = false;
5245        let mut error = None;
5246
5247        while let Some(item) = stream.next().await {
5248            match item {
5249                Ok(MultiTurnStreamItem::StreamAssistantItem(
5250                    StreamedAssistantContent::ToolCall { .. },
5251                )) => {
5252                    saw_tool_call = true;
5253                }
5254                Ok(MultiTurnStreamItem::StreamUserItem(StreamedUserContent::ToolResult {
5255                    ..
5256                })) => {
5257                    saw_tool_result = true;
5258                }
5259                Ok(_) => {}
5260                Err(err) => {
5261                    error = Some(err);
5262                    break;
5263                }
5264            }
5265        }
5266
5267        assert!(!saw_tool_call);
5268        assert!(!saw_tool_result);
5269        assert_eq!(add_calls.load(Ordering::SeqCst), 0);
5270        let error = error.expect("mixed disallowed streamed tool call should fail");
5271        match error {
5272            StreamingError::Prompt(err) => match *err {
5273                PromptError::UnknownToolCall {
5274                    tool_name,
5275                    available_tools,
5276                    allowed_tools,
5277                    chat_history,
5278                } => {
5279                    assert_eq!(tool_name, "subtract");
5280                    assert_eq!(
5281                        available_tools,
5282                        vec!["add".to_string(), "subtract".to_string()]
5283                    );
5284                    assert_eq!(allowed_tools, vec!["add".to_string()]);
5285                    assert!(history_contains_tool_call(&chat_history, "subtract"));
5286                }
5287                other => panic!("expected UnknownToolCall, got {other:?}"),
5288            },
5289            other => panic!("expected prompt streaming error, got {other:?}"),
5290        }
5291        assert_eq!(recorded.request_count(), 1);
5292    }
5293
5294    #[tokio::test]
5295    async fn tool_choice_none_rejects_streaming_tool_call() {
5296        let model = MockCompletionModel::from_stream_turns([
5297            vec![
5298                MockStreamEvent::tool_call(
5299                    "tool_call_1",
5300                    "add",
5301                    serde_json::json!({"x": 1, "y": 2}),
5302                ),
5303                MockStreamEvent::final_response_with_total_tokens(4),
5304            ],
5305            vec![
5306                MockStreamEvent::text("should not be requested"),
5307                MockStreamEvent::final_response_with_total_tokens(6),
5308            ],
5309        ]);
5310        let recorded = model.clone();
5311        let agent = AgentBuilder::new(model)
5312            .tool(MockAddTool)
5313            .tool_choice(ToolChoice::None)
5314            .build();
5315
5316        let mut stream = agent
5317            .stream_prompt("do not use tools")
5318            .add_hook(PanicOnUnknownToolHook)
5319            .max_turns(3)
5320            .await;
5321        let mut saw_tool_call = false;
5322        let mut error = None;
5323
5324        while let Some(item) = stream.next().await {
5325            match item {
5326                Ok(MultiTurnStreamItem::StreamAssistantItem(
5327                    StreamedAssistantContent::ToolCall { .. },
5328                )) => {
5329                    saw_tool_call = true;
5330                }
5331                Ok(_) => {}
5332                Err(err) => {
5333                    error = Some(err);
5334                    break;
5335                }
5336            }
5337        }
5338
5339        assert!(!saw_tool_call);
5340        let error = error.expect("ToolChoice::None should reject returned tool calls");
5341        match error {
5342            StreamingError::Prompt(err) => match *err {
5343                PromptError::UnknownToolCall {
5344                    tool_name,
5345                    available_tools,
5346                    allowed_tools,
5347                    chat_history,
5348                } => {
5349                    assert_eq!(tool_name, "add");
5350                    assert_eq!(available_tools, vec!["add".to_string()]);
5351                    assert!(allowed_tools.is_empty());
5352                    assert!(history_contains_tool_call(&chat_history, "add"));
5353                }
5354                other => panic!("expected UnknownToolCall, got {other:?}"),
5355            },
5356            other => panic!("expected prompt streaming error, got {other:?}"),
5357        }
5358        assert_eq!(recorded.request_count(), 1);
5359    }
5360
5361    #[tokio::test]
5362    async fn tool_choice_none_rejects_streaming_tool_call_name_delta_before_hook_or_emit() {
5363        let model = MockCompletionModel::from_stream_turns([
5364            vec![
5365                MockStreamEvent::tool_call_name_delta("tool_1", "internal_1", "add"),
5366                MockStreamEvent::tool_call_arguments_delta("tool_1", "internal_1", "{\"x\":1}"),
5367                MockStreamEvent::final_response_with_total_tokens(4),
5368            ],
5369            vec![
5370                MockStreamEvent::text("should not be requested"),
5371                MockStreamEvent::final_response_with_total_tokens(6),
5372            ],
5373        ]);
5374        let recorded = model.clone();
5375        let agent = AgentBuilder::new(model)
5376            .tool(MockAddTool)
5377            .tool_choice(ToolChoice::None)
5378            .build();
5379
5380        let mut stream = agent
5381            .stream_prompt("do not use tools")
5382            .add_hook(PanicOnUnknownToolHook)
5383            .max_turns(3)
5384            .await;
5385        let mut saw_delta = false;
5386        let mut error = None;
5387
5388        while let Some(item) = stream.next().await {
5389            match item {
5390                Ok(MultiTurnStreamItem::StreamAssistantItem(
5391                    StreamedAssistantContent::ToolCallDelta { .. },
5392                )) => {
5393                    saw_delta = true;
5394                }
5395                Ok(_) => {}
5396                Err(err) => {
5397                    error = Some(err);
5398                    break;
5399                }
5400            }
5401        }
5402
5403        assert!(!saw_delta);
5404        let error = error.expect("ToolChoice::None should reject returned tool-call deltas");
5405        match error {
5406            StreamingError::Prompt(err) => match *err {
5407                PromptError::UnknownToolCall {
5408                    tool_name,
5409                    available_tools,
5410                    allowed_tools,
5411                    chat_history,
5412                } => {
5413                    assert_eq!(tool_name, "add");
5414                    assert_eq!(available_tools, vec!["add".to_string()]);
5415                    assert!(allowed_tools.is_empty());
5416                    assert!(history_contains_tool_call(&chat_history, "add"));
5417                }
5418                other => panic!("expected UnknownToolCall, got {other:?}"),
5419            },
5420            other => panic!("expected prompt streaming error, got {other:?}"),
5421        }
5422        assert_eq!(recorded.request_count(), 1);
5423    }
5424
5425    #[tokio::test]
5426    async fn unknown_tool_call_name_delta_fails_before_streaming_delta_hook_or_emit() {
5427        let model = MockCompletionModel::from_stream_turns([
5428            vec![
5429                MockStreamEvent::tool_call_name_delta("tool_1", "internal_1", "default_api"),
5430                MockStreamEvent::tool_call_arguments_delta("tool_1", "internal_1", "{\"x\":1}"),
5431                MockStreamEvent::final_response_with_total_tokens(4),
5432            ],
5433            vec![
5434                MockStreamEvent::text("should not be requested"),
5435                MockStreamEvent::final_response_with_total_tokens(6),
5436            ],
5437        ]);
5438        let recorded = model.clone();
5439        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
5440
5441        let mut stream = agent
5442            .stream_prompt("stream a bad tool call")
5443            .add_hook(PanicOnUnknownToolHook)
5444            .max_turns(3)
5445            .await;
5446        let mut saw_delta = false;
5447        let mut error = None;
5448
5449        while let Some(item) = stream.next().await {
5450            match item {
5451                Ok(MultiTurnStreamItem::StreamAssistantItem(
5452                    StreamedAssistantContent::ToolCallDelta { .. },
5453                )) => {
5454                    saw_delta = true;
5455                }
5456                Ok(_) => {}
5457                Err(err) => {
5458                    error = Some(err);
5459                    break;
5460                }
5461            }
5462        }
5463
5464        assert!(!saw_delta);
5465        let error = error.expect("unknown tool-call name delta should fail");
5466        match error {
5467            StreamingError::Prompt(err) => match *err {
5468                PromptError::UnknownToolCall {
5469                    tool_name,
5470                    available_tools,
5471                    allowed_tools,
5472                    chat_history,
5473                } => {
5474                    assert_eq!(tool_name, "default_api");
5475                    assert_eq!(available_tools, vec!["add".to_string()]);
5476                    assert_eq!(allowed_tools, vec!["add".to_string()]);
5477                    assert!(history_contains_tool_call(&chat_history, "default_api"));
5478                }
5479                other => panic!("expected UnknownToolCall, got {other:?}"),
5480            },
5481            other => panic!("expected prompt streaming error, got {other:?}"),
5482        }
5483        assert_eq!(recorded.request_count(), 1);
5484    }
5485
5486    #[tokio::test]
5487    async fn tool_call_args_delta_before_unknown_name_fails_before_hook_or_emit() {
5488        let model = MockCompletionModel::from_stream_turns([
5489            vec![
5490                MockStreamEvent::tool_call_arguments_delta("tool_1", "internal_1", "{\"x\":1}"),
5491                MockStreamEvent::tool_call_name_delta("tool_1", "internal_1", "default_api"),
5492                MockStreamEvent::final_response_with_total_tokens(4),
5493            ],
5494            vec![
5495                MockStreamEvent::text("should not be requested"),
5496                MockStreamEvent::final_response_with_total_tokens(6),
5497            ],
5498        ]);
5499        let recorded = model.clone();
5500        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
5501
5502        let mut stream = agent
5503            .stream_prompt("stream a bad tool call")
5504            .add_hook(PanicOnUnknownToolHook)
5505            .max_turns(3)
5506            .await;
5507        let mut saw_delta = false;
5508        let mut error = None;
5509
5510        while let Some(item) = stream.next().await {
5511            match item {
5512                Ok(MultiTurnStreamItem::StreamAssistantItem(
5513                    StreamedAssistantContent::ToolCallDelta { .. },
5514                )) => {
5515                    saw_delta = true;
5516                }
5517                Ok(_) => {}
5518                Err(err) => {
5519                    error = Some(err);
5520                    break;
5521                }
5522            }
5523        }
5524
5525        assert!(!saw_delta);
5526        let error = error.expect("unknown tool-call name should reject buffered args");
5527        match error {
5528            StreamingError::Prompt(err) => match *err {
5529                PromptError::UnknownToolCall {
5530                    tool_name,
5531                    available_tools,
5532                    allowed_tools,
5533                    chat_history,
5534                } => {
5535                    assert_eq!(tool_name, "default_api");
5536                    assert_eq!(available_tools, vec!["add".to_string()]);
5537                    assert_eq!(allowed_tools, vec!["add".to_string()]);
5538                    assert!(history_contains_tool_call(&chat_history, "default_api"));
5539                }
5540                other => panic!("expected UnknownToolCall, got {other:?}"),
5541            },
5542            other => panic!("expected prompt streaming error, got {other:?}"),
5543        }
5544        assert_eq!(recorded.request_count(), 1);
5545    }
5546
5547    #[tokio::test]
5548    async fn tool_call_args_delta_before_valid_name_buffers_then_emits_in_safe_order() {
5549        let model = MockCompletionModel::from_stream_turns([[
5550            MockStreamEvent::tool_call_arguments_delta("tool_1", "internal_1", "{\"x\":"),
5551            MockStreamEvent::tool_call_name_delta("tool_1", "internal_1", "add"),
5552            MockStreamEvent::tool_call_arguments_delta("tool_1", "internal_1", "1}"),
5553            MockStreamEvent::final_response_with_total_tokens(3),
5554        ]]);
5555        let hook = RecordingToolCallDeltaHook::default();
5556        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
5557
5558        let mut stream = agent
5559            .stream_prompt("stream a tool call")
5560            .add_hook(hook.clone())
5561            .await;
5562        let mut stream_deltas = Vec::new();
5563
5564        while let Some(item) = stream.next().await {
5565            match item {
5566                Ok(MultiTurnStreamItem::StreamAssistantItem(
5567                    StreamedAssistantContent::ToolCallDelta {
5568                        id,
5569                        internal_call_id,
5570                        content,
5571                    },
5572                )) => {
5573                    stream_deltas.push((id, internal_call_id, content));
5574                }
5575                Ok(MultiTurnStreamItem::FinalResponse(_)) => break,
5576                Ok(_) => {}
5577                Err(err) => panic!("unexpected streaming error: {err:?}"),
5578            }
5579        }
5580
5581        assert_eq!(
5582            hook.observed(),
5583            vec![
5584                (
5585                    "tool_1".to_string(),
5586                    "internal_1".to_string(),
5587                    Some("add".to_string()),
5588                    String::new()
5589                ),
5590                (
5591                    "tool_1".to_string(),
5592                    "internal_1".to_string(),
5593                    None,
5594                    "{\"x\":".to_string()
5595                ),
5596                (
5597                    "tool_1".to_string(),
5598                    "internal_1".to_string(),
5599                    None,
5600                    "1}".to_string()
5601                ),
5602            ]
5603        );
5604        assert_eq!(
5605            stream_deltas,
5606            vec![
5607                (
5608                    "tool_1".to_string(),
5609                    "internal_1".to_string(),
5610                    ToolCallDeltaContent::Name("add".to_string())
5611                ),
5612                (
5613                    "tool_1".to_string(),
5614                    "internal_1".to_string(),
5615                    ToolCallDeltaContent::Delta("{\"x\":".to_string())
5616                ),
5617                (
5618                    "tool_1".to_string(),
5619                    "internal_1".to_string(),
5620                    ToolCallDeltaContent::Delta("1}".to_string())
5621                ),
5622            ]
5623        );
5624    }
5625
5626    #[tokio::test]
5627    async fn tool_call_args_delta_without_name_errors_at_stream_end() {
5628        let model = MockCompletionModel::from_stream_turns([
5629            vec![
5630                MockStreamEvent::tool_call_arguments_delta("tool_1", "internal_1", "{\"x\":1}"),
5631                MockStreamEvent::final_response_with_total_tokens(4),
5632            ],
5633            vec![
5634                MockStreamEvent::text("should not be requested"),
5635                MockStreamEvent::final_response_with_total_tokens(6),
5636            ],
5637        ]);
5638        let recorded = model.clone();
5639        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
5640
5641        let mut stream = agent
5642            .stream_prompt("stream an incomplete tool call")
5643            .add_hook(PanicOnUnknownToolHook)
5644            .max_turns(3)
5645            .await;
5646        let mut saw_delta = false;
5647        let mut saw_completion_call = false;
5648        let mut saw_final_response = false;
5649        let mut error = None;
5650
5651        while let Some(item) = stream.next().await {
5652            match item {
5653                Ok(MultiTurnStreamItem::StreamAssistantItem(
5654                    StreamedAssistantContent::ToolCallDelta { .. },
5655                )) => {
5656                    saw_delta = true;
5657                }
5658                Ok(MultiTurnStreamItem::CompletionCall(_)) => {
5659                    saw_completion_call = true;
5660                }
5661                Ok(MultiTurnStreamItem::FinalResponse(_)) => {
5662                    saw_final_response = true;
5663                }
5664                Ok(_) => {}
5665                Err(err) => {
5666                    error = Some(err);
5667                    break;
5668                }
5669            }
5670        }
5671
5672        assert!(!saw_delta);
5673        assert!(!saw_completion_call);
5674        assert!(!saw_final_response);
5675        let error = error.expect("unterminated tool-call args delta should fail");
5676        match error {
5677            StreamingError::Completion(CompletionError::ResponseError(message)) => {
5678                assert!(
5679                    message.contains("streamed tool call arguments"),
5680                    "{message}"
5681                );
5682                assert!(message.contains("tool_1"), "{message}");
5683                assert!(message.contains("internal_1"), "{message}");
5684            }
5685            other => panic!("expected completion response error, got {other:?}"),
5686        }
5687        assert_eq!(recorded.request_count(), 1);
5688    }
5689
5690    #[tokio::test]
5691    async fn tool_choice_none_buffers_args_then_rejects_name_without_emit() {
5692        let model = MockCompletionModel::from_stream_turns([
5693            vec![
5694                MockStreamEvent::tool_call_arguments_delta("tool_1", "internal_1", "{\"x\":1}"),
5695                MockStreamEvent::tool_call_name_delta("tool_1", "internal_1", "add"),
5696                MockStreamEvent::final_response_with_total_tokens(4),
5697            ],
5698            vec![
5699                MockStreamEvent::text("should not be requested"),
5700                MockStreamEvent::final_response_with_total_tokens(6),
5701            ],
5702        ]);
5703        let recorded = model.clone();
5704        let agent = AgentBuilder::new(model)
5705            .tool(MockAddTool)
5706            .tool_choice(ToolChoice::None)
5707            .build();
5708
5709        let mut stream = agent
5710            .stream_prompt("do not use tools")
5711            .add_hook(PanicOnUnknownToolHook)
5712            .max_turns(3)
5713            .await;
5714        let mut saw_delta = false;
5715        let mut error = None;
5716
5717        while let Some(item) = stream.next().await {
5718            match item {
5719                Ok(MultiTurnStreamItem::StreamAssistantItem(
5720                    StreamedAssistantContent::ToolCallDelta { .. },
5721                )) => {
5722                    saw_delta = true;
5723                }
5724                Ok(_) => {}
5725                Err(err) => {
5726                    error = Some(err);
5727                    break;
5728                }
5729            }
5730        }
5731
5732        assert!(!saw_delta);
5733        let error = error.expect("ToolChoice::None should reject buffered tool-call deltas");
5734        match error {
5735            StreamingError::Prompt(err) => match *err {
5736                PromptError::UnknownToolCall {
5737                    tool_name,
5738                    available_tools,
5739                    allowed_tools,
5740                    chat_history,
5741                } => {
5742                    assert_eq!(tool_name, "add");
5743                    assert_eq!(available_tools, vec!["add".to_string()]);
5744                    assert!(allowed_tools.is_empty());
5745                    assert!(history_contains_tool_call(&chat_history, "add"));
5746                }
5747                other => panic!("expected UnknownToolCall, got {other:?}"),
5748            },
5749            other => panic!("expected prompt streaming error, got {other:?}"),
5750        }
5751        assert_eq!(recorded.request_count(), 1);
5752    }
5753
5754    #[tokio::test]
5755    async fn stream_prompt_emits_tool_call_deltas_without_hook() {
5756        let model = MockCompletionModel::from_stream_turns([[
5757            MockStreamEvent::tool_call_name_delta("tool_1", "internal_1", "add"),
5758            MockStreamEvent::tool_call_arguments_delta("tool_1", "internal_1", "{\"x\":"),
5759            MockStreamEvent::tool_call_arguments_delta("tool_1", "internal_1", "1}"),
5760            MockStreamEvent::final_response_with_total_tokens(3),
5761        ]]);
5762        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
5763
5764        let mut stream = agent.stream_prompt("stream a tool call").await;
5765        let mut deltas = Vec::new();
5766
5767        while let Some(item) = stream.next().await {
5768            match item {
5769                Ok(MultiTurnStreamItem::StreamAssistantItem(
5770                    StreamedAssistantContent::ToolCallDelta {
5771                        id,
5772                        internal_call_id,
5773                        content,
5774                    },
5775                )) => {
5776                    deltas.push((id, internal_call_id, content));
5777                }
5778                Ok(MultiTurnStreamItem::FinalResponse(_)) => break,
5779                Ok(_) => {}
5780                Err(err) => panic!("unexpected streaming error: {err:?}"),
5781            }
5782        }
5783
5784        assert_eq!(
5785            deltas,
5786            vec![
5787                (
5788                    "tool_1".to_string(),
5789                    "internal_1".to_string(),
5790                    ToolCallDeltaContent::Name("add".to_string())
5791                ),
5792                (
5793                    "tool_1".to_string(),
5794                    "internal_1".to_string(),
5795                    ToolCallDeltaContent::Delta("{\"x\":".to_string())
5796                ),
5797                (
5798                    "tool_1".to_string(),
5799                    "internal_1".to_string(),
5800                    ToolCallDeltaContent::Delta("1}".to_string())
5801                ),
5802            ]
5803        );
5804    }
5805
5806    #[tokio::test]
5807    async fn stream_prompt_emits_tool_call_deltas_after_hook_continue() {
5808        let model = MockCompletionModel::from_stream_turns([[
5809            MockStreamEvent::tool_call_name_delta("tool_1", "internal_1", "add"),
5810            MockStreamEvent::tool_call_arguments_delta("tool_1", "internal_1", "{\"x\":"),
5811            MockStreamEvent::tool_call_arguments_delta("tool_1", "internal_1", "1}"),
5812            MockStreamEvent::final_response_with_total_tokens(3),
5813        ]]);
5814        let hook = RecordingToolCallDeltaHook::default();
5815        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
5816
5817        let mut stream = agent
5818            .stream_prompt("stream a tool call")
5819            .add_hook(hook.clone())
5820            .await;
5821        let mut stream_deltas = Vec::new();
5822
5823        while let Some(item) = stream.next().await {
5824            match item {
5825                Ok(MultiTurnStreamItem::StreamAssistantItem(
5826                    StreamedAssistantContent::ToolCallDelta {
5827                        id,
5828                        internal_call_id,
5829                        content,
5830                    },
5831                )) => {
5832                    stream_deltas.push((id, internal_call_id, content));
5833                }
5834                Ok(MultiTurnStreamItem::FinalResponse(_)) => break,
5835                Ok(_) => {}
5836                Err(err) => panic!("unexpected streaming error: {err:?}"),
5837            }
5838        }
5839
5840        assert_eq!(
5841            hook.observed(),
5842            vec![
5843                (
5844                    "tool_1".to_string(),
5845                    "internal_1".to_string(),
5846                    Some("add".to_string()),
5847                    String::new()
5848                ),
5849                (
5850                    "tool_1".to_string(),
5851                    "internal_1".to_string(),
5852                    None,
5853                    "{\"x\":".to_string()
5854                ),
5855                (
5856                    "tool_1".to_string(),
5857                    "internal_1".to_string(),
5858                    None,
5859                    "1}".to_string()
5860                ),
5861            ]
5862        );
5863        assert_eq!(
5864            stream_deltas,
5865            vec![
5866                (
5867                    "tool_1".to_string(),
5868                    "internal_1".to_string(),
5869                    ToolCallDeltaContent::Name("add".to_string())
5870                ),
5871                (
5872                    "tool_1".to_string(),
5873                    "internal_1".to_string(),
5874                    ToolCallDeltaContent::Delta("{\"x\":".to_string())
5875                ),
5876                (
5877                    "tool_1".to_string(),
5878                    "internal_1".to_string(),
5879                    ToolCallDeltaContent::Delta("1}".to_string())
5880                ),
5881            ]
5882        );
5883    }
5884
5885    #[tokio::test]
5886    async fn stream_prompt_tool_call_deltas_hook_termination_prevents_delta_emit() {
5887        let model = MockCompletionModel::from_stream_turns([[
5888            MockStreamEvent::tool_call_name_delta("tool_1", "internal_1", "add"),
5889            MockStreamEvent::tool_call_arguments_delta("tool_1", "internal_1", "{\"x\":"),
5890            MockStreamEvent::final_response_with_total_tokens(3),
5891        ]]);
5892        let hook = TerminatingToolCallDeltaHook::default();
5893        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
5894
5895        let mut stream = agent
5896            .stream_prompt("stream a tool call")
5897            .add_hook(hook.clone())
5898            .await;
5899        let mut saw_delta = false;
5900        let mut saw_final_response = false;
5901        let mut error_message = None;
5902
5903        while let Some(item) = stream.next().await {
5904            match item {
5905                Ok(MultiTurnStreamItem::StreamAssistantItem(
5906                    StreamedAssistantContent::ToolCallDelta { .. },
5907                )) => {
5908                    saw_delta = true;
5909                }
5910                Ok(MultiTurnStreamItem::FinalResponse(_)) => {
5911                    saw_final_response = true;
5912                }
5913                Ok(_) => {}
5914                Err(err) => {
5915                    error_message = Some(err.to_string());
5916                    break;
5917                }
5918            }
5919        }
5920
5921        assert_eq!(
5922            hook.observed(),
5923            vec![(
5924                "tool_1".to_string(),
5925                "internal_1".to_string(),
5926                Some("add".to_string()),
5927                String::new()
5928            )]
5929        );
5930        assert!(!saw_delta);
5931        assert!(!saw_final_response);
5932        assert!(
5933            error_message
5934                .as_deref()
5935                .is_some_and(|message| message.contains("PromptCancelled: stop on tool call delta")),
5936            "expected hook termination error, got {error_message:?}"
5937        );
5938    }
5939
5940    #[tokio::test]
5941    async fn stream_prompt_exposes_completion_calls() {
5942        let first_call_usage = usage(10, 2);
5943        let second_call_usage = usage(25, 5);
5944        let model = MockCompletionModel::from_stream_turns([
5945            vec![
5946                MockStreamEvent::tool_call(
5947                    "tool_call_1",
5948                    "add",
5949                    serde_json::json!({"x": 1, "y": 2}),
5950                )
5951                .with_call_id("call_1"),
5952                MockStreamEvent::final_response(first_call_usage),
5953            ],
5954            vec![
5955                MockStreamEvent::text("done"),
5956                MockStreamEvent::final_response(second_call_usage),
5957            ],
5958        ]);
5959        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
5960        let empty_history: &[Message] = &[];
5961
5962        let mut stream = agent
5963            .stream_prompt("do tool work")
5964            .history(empty_history)
5965            .max_turns(3)
5966            .await;
5967        let mut completion_calls_events = Vec::new();
5968        let mut final_response = None;
5969
5970        while let Some(item) = stream.next().await {
5971            match item {
5972                Ok(MultiTurnStreamItem::CompletionCall(call_usage)) => {
5973                    completion_calls_events.push(call_usage);
5974                }
5975                Ok(MultiTurnStreamItem::FinalResponse(response)) => {
5976                    final_response = Some(response);
5977                    break;
5978                }
5979                Ok(_) => {}
5980                Err(err) => panic!("unexpected streaming error: {err:?}"),
5981            }
5982        }
5983
5984        assert_eq!(
5985            completion_calls_events,
5986            vec![
5987                CompletionCall::new(0, first_call_usage),
5988                CompletionCall::new(1, second_call_usage)
5989            ]
5990        );
5991
5992        let final_response = final_response.expect("expected final response");
5993        assert_eq!(
5994            final_response.usage(),
5995            Usage {
5996                input_tokens: 35,
5997                output_tokens: 7,
5998                total_tokens: 42,
5999                cached_input_tokens: 0,
6000                cache_creation_input_tokens: 0,
6001                tool_use_prompt_tokens: 0,
6002                reasoning_tokens: 0,
6003            }
6004        );
6005        assert_eq!(
6006            final_response.completion_calls(),
6007            &[
6008                CompletionCall::new(0, first_call_usage),
6009                CompletionCall::new(1, second_call_usage)
6010            ]
6011        );
6012    }
6013
6014    #[tokio::test(flavor = "current_thread")]
6015    async fn stream_prompt_records_single_call_usage_on_chat_span_under_outer_span() {
6016        let call_usage = usage(10, 2);
6017        let model = MockCompletionModel::from_stream_turns([[
6018            MockStreamEvent::text("done"),
6019            MockStreamEvent::final_response(call_usage),
6020        ]]);
6021        let agent = AgentBuilder::new(model).build();
6022
6023        assert_stream_usage_recorded_on_chat_spans(agent, "say done", 1, &[call_usage]).await;
6024    }
6025
6026    #[tokio::test(flavor = "current_thread")]
6027    async fn stream_prompt_records_multi_turn_usage_on_chat_spans_under_outer_span() {
6028        let first_call_usage = usage(10, 2);
6029        let second_call_usage = usage(25, 5);
6030        let model = MockCompletionModel::from_stream_turns([
6031            vec![
6032                MockStreamEvent::tool_call(
6033                    "tool_call_1",
6034                    "add",
6035                    serde_json::json!({"x": 1, "y": 2}),
6036                )
6037                .with_call_id("call_1"),
6038                MockStreamEvent::final_response(first_call_usage),
6039            ],
6040            vec![
6041                MockStreamEvent::text("done"),
6042                MockStreamEvent::final_response(second_call_usage),
6043            ],
6044        ]);
6045        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
6046
6047        assert_stream_usage_recorded_on_chat_spans(
6048            agent,
6049            "do tool work",
6050            3,
6051            &[first_call_usage, second_call_usage],
6052        )
6053        .await;
6054    }
6055
6056    #[tokio::test]
6057    async fn stream_prompt_emits_completion_call_before_finish_hook_termination() {
6058        let call_usage = usage(10, 2);
6059        let model = MockCompletionModel::from_stream_turns([[
6060            MockStreamEvent::text("done"),
6061            MockStreamEvent::final_response(call_usage),
6062        ]]);
6063        let agent = AgentBuilder::new(model).build();
6064
6065        let mut stream = agent
6066            .stream_prompt("say done")
6067            .add_hook(TerminateOnStreamFinish)
6068            .await;
6069        let mut completion_calls = Vec::new();
6070        let mut saw_error = false;
6071
6072        while let Some(item) = stream.next().await {
6073            match item {
6074                Ok(MultiTurnStreamItem::CompletionCall(completion_call)) => {
6075                    completion_calls.push(completion_call);
6076                }
6077                Ok(MultiTurnStreamItem::FinalResponse(response)) => {
6078                    panic!("unexpected final response after hook termination: {response:?}");
6079                }
6080                Ok(_) => {}
6081                Err(_) => {
6082                    saw_error = true;
6083                    break;
6084                }
6085            }
6086        }
6087
6088        assert_eq!(completion_calls, vec![CompletionCall::new(0, call_usage)]);
6089        assert!(saw_error);
6090    }
6091
6092    #[tokio::test]
6093    async fn stream_prompt_completion_calls_records_unreported_usage() {
6094        let second_call_usage = usage(25, 5);
6095        let model = MockCompletionModel::from_stream_turns([
6096            vec![
6097                MockStreamEvent::tool_call(
6098                    "tool_call_1",
6099                    "add",
6100                    serde_json::json!({"x": 1, "y": 2}),
6101                )
6102                .with_call_id("call_1"),
6103            ],
6104            vec![
6105                MockStreamEvent::text("done"),
6106                MockStreamEvent::final_response(second_call_usage),
6107            ],
6108        ]);
6109        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
6110        let empty_history: &[Message] = &[];
6111
6112        let mut stream = agent
6113            .stream_prompt("do tool work")
6114            .history(empty_history)
6115            .max_turns(3)
6116            .await;
6117        let mut completion_calls_events = Vec::new();
6118        let mut final_response = None;
6119
6120        while let Some(item) = stream.next().await {
6121            match item {
6122                Ok(MultiTurnStreamItem::CompletionCall(call_usage)) => {
6123                    completion_calls_events.push(call_usage);
6124                }
6125                Ok(MultiTurnStreamItem::FinalResponse(response)) => {
6126                    final_response = Some(response);
6127                    break;
6128                }
6129                Ok(_) => {}
6130                Err(err) => panic!("unexpected streaming error: {err:?}"),
6131            }
6132        }
6133
6134        let expected_usage = vec![
6135            CompletionCall::new(0, Usage::new()),
6136            CompletionCall::new(1, second_call_usage),
6137        ];
6138        assert_eq!(completion_calls_events, expected_usage);
6139
6140        let final_response = final_response.expect("expected final response");
6141        assert_eq!(final_response.completion_calls(), expected_usage.as_slice());
6142    }
6143
6144    #[tokio::test]
6145    async fn final_response_matches_streamed_text_when_provider_final_is_textless() {
6146        let agent = AgentBuilder::new(streaming_text_then_final_model()).build();
6147
6148        let mut stream = agent.stream_prompt("say hello").await;
6149        let mut streamed_text = String::new();
6150        let mut final_response_text = None;
6151
6152        while let Some(item) = stream.next().await {
6153            match item {
6154                Ok(MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Text(
6155                    text,
6156                ))) => streamed_text.push_str(&text.text),
6157                Ok(MultiTurnStreamItem::FinalResponse(res)) => {
6158                    final_response_text = Some(res.output().to_owned());
6159                    break;
6160                }
6161                Ok(_) => {}
6162                Err(err) => panic!("unexpected streaming error: {err:?}"),
6163            }
6164        }
6165
6166        assert_eq!(streamed_text, "hello world");
6167        assert_eq!(final_response_text.as_deref(), Some("hello world"));
6168    }
6169
6170    #[tokio::test]
6171    async fn final_response_preserves_structured_text_metadata() {
6172        let agent = AgentBuilder::new(streaming_cited_text_then_final_model()).build();
6173
6174        let mut stream = agent.stream_prompt("answer with citations").await;
6175        let mut final_response = None;
6176
6177        while let Some(item) = stream.next().await {
6178            match item {
6179                Ok(MultiTurnStreamItem::FinalResponse(res)) => {
6180                    final_response = Some(res);
6181                    break;
6182                }
6183                Ok(_) => {}
6184                Err(err) => panic!("unexpected streaming error: {err:?}"),
6185            }
6186        }
6187
6188        let final_response = final_response.expect("expected final response");
6189        assert_eq!(final_response.output(), "cited answer");
6190        let metadata = text_metadata(final_response.content())
6191            .expect("expected text metadata in final content");
6192        assert_eq!(
6193            metadata["citations"][0]["encrypted_index"],
6194            "encrypted-reference"
6195        );
6196    }
6197
6198    #[tokio::test]
6199    async fn final_response_history_preserves_structured_text_metadata() {
6200        let agent = AgentBuilder::new(streaming_cited_text_then_final_model()).build();
6201
6202        let empty_history: &[Message] = &[];
6203        let mut stream = agent
6204            .stream_prompt("answer with citations")
6205            .history(empty_history)
6206            .await;
6207        let mut final_response = None;
6208
6209        while let Some(item) = stream.next().await {
6210            match item {
6211                Ok(MultiTurnStreamItem::FinalResponse(res)) => {
6212                    final_response = Some(res);
6213                    break;
6214                }
6215                Ok(_) => {}
6216                Err(err) => panic!("unexpected streaming error: {err:?}"),
6217            }
6218        }
6219
6220        let final_response = final_response.expect("expected final response");
6221        let history = final_response
6222            .messages()
6223            .expect("with_history should include final history");
6224        let assistant_content = history
6225            .iter()
6226            .find_map(|message| match message {
6227                Message::Assistant { content, .. } => Some(content),
6228                _ => None,
6229            })
6230            .expect("expected assistant message in history");
6231        let metadata =
6232            text_metadata(assistant_content).expect("expected text metadata in assistant history");
6233        assert_eq!(
6234            metadata["citations"][0]["encrypted_index"],
6235            "encrypted-reference"
6236        );
6237    }
6238
6239    #[tokio::test]
6240    async fn tool_follow_up_history_preserves_structured_text_metadata() {
6241        let model = streaming_cited_text_then_tool_model();
6242        let recorded = model.clone();
6243        let agent = AgentBuilder::new(model).tool(MockAddTool).build();
6244        let empty_history: &[Message] = &[];
6245
6246        let mut stream = agent
6247            .stream_prompt("use a tool with citations")
6248            .history(empty_history)
6249            .max_turns(3)
6250            .await;
6251
6252        while let Some(item) = stream.next().await {
6253            match item {
6254                Ok(MultiTurnStreamItem::FinalResponse(_)) => break,
6255                Ok(_) => {}
6256                Err(err) => panic!("unexpected streaming error: {err:?}"),
6257            }
6258        }
6259
6260        let requests = recorded.requests();
6261        assert_eq!(requests.len(), 2);
6262        let follow_up_history = requests[1].chat_history.iter().collect::<Vec<_>>();
6263        let assistant_content = follow_up_history
6264            .iter()
6265            .find_map(|message| match message {
6266                Message::Assistant { content, .. } => Some(content),
6267                _ => None,
6268            })
6269            .expect("expected assistant message in follow-up history");
6270        let metadata = text_metadata(assistant_content)
6271            .expect("expected citation metadata in follow-up assistant history");
6272        assert_eq!(
6273            metadata["citations"][0]["encrypted_index"],
6274            "encrypted-reference"
6275        );
6276    }
6277
6278    #[tokio::test]
6279    async fn final_response_can_remain_empty_for_truly_textless_turns() {
6280        let agent = AgentBuilder::new(streaming_final_only_model()).build();
6281
6282        let mut stream = agent.stream_prompt("say nothing").await;
6283        let mut streamed_text = String::new();
6284        let mut final_response_text = None;
6285
6286        while let Some(item) = stream.next().await {
6287            match item {
6288                Ok(MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Text(
6289                    text,
6290                ))) => streamed_text.push_str(&text.text),
6291                Ok(MultiTurnStreamItem::FinalResponse(res)) => {
6292                    final_response_text = Some(res.output().to_owned());
6293                    break;
6294                }
6295                Ok(_) => {}
6296                Err(err) => panic!("unexpected streaming error: {err:?}"),
6297            }
6298        }
6299
6300        assert!(streamed_text.is_empty());
6301        assert_eq!(final_response_text.as_deref(), Some(""));
6302    }
6303
6304    /// Background task that logs periodically to detect span leakage.
6305    /// If span leakage occurs, these logs will be prefixed with `invoke_agent{...}`.
6306    async fn background_logger(stop: Arc<AtomicBool>, leak_count: Arc<AtomicU32>) {
6307        let mut interval = tokio::time::interval(Duration::from_millis(50));
6308        let mut count = 0u32;
6309
6310        while !stop.load(Ordering::Relaxed) {
6311            interval.tick().await;
6312            count += 1;
6313
6314            tracing::event!(
6315                target: "background_logger",
6316                tracing::Level::INFO,
6317                count = count,
6318                "Background tick"
6319            );
6320
6321            // Check if we're inside an unexpected span
6322            let current = tracing::Span::current();
6323            if !current.is_disabled() && !current.is_none() {
6324                leak_count.fetch_add(1, Ordering::Relaxed);
6325            }
6326        }
6327
6328        tracing::info!(target: "background_logger", total_ticks = count, "Background logger stopped");
6329    }
6330
6331    /// Test that span context doesn't leak to concurrent tasks during streaming.
6332    ///
6333    /// This test verifies that using `.instrument()` instead of `span.enter()` in
6334    /// async_stream prevents thread-local span context from leaking to other tasks.
6335    ///
6336    /// Uses single-threaded runtime to force all tasks onto the same thread,
6337    /// making the span leak deterministic (it only occurs when tasks share a thread).
6338    #[tokio::test(flavor = "current_thread")]
6339    #[ignore = "This requires an API key"]
6340    async fn test_span_context_isolation() -> anyhow::Result<()> {
6341        let stop = Arc::new(AtomicBool::new(false));
6342        let leak_count = Arc::new(AtomicU32::new(0));
6343
6344        // Start background logger
6345        let bg_stop = stop.clone();
6346        let bg_leak = leak_count.clone();
6347        let bg_handle = tokio::spawn(async move {
6348            background_logger(bg_stop, bg_leak).await;
6349        });
6350
6351        // Small delay to let background logger start
6352        tokio::time::sleep(Duration::from_millis(100)).await;
6353
6354        // Make streaming request WITHOUT an outer span so rig creates its own invoke_agent span
6355        // (rig reuses current span if one exists, so we need to ensure there's no current span)
6356        let client = anthropic::Client::from_env()?;
6357        let agent = client
6358            .agent(anthropic::completion::CLAUDE_HAIKU_4_5)
6359            .preamble("You are a helpful assistant.")
6360            .temperature(0.1)
6361            .max_tokens(100)
6362            .build();
6363
6364        let mut stream = agent
6365            .stream_prompt("Say 'hello world' and nothing else.")
6366            .await;
6367
6368        let mut full_content = String::new();
6369        while let Some(item) = stream.next().await {
6370            match item {
6371                Ok(MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Text(
6372                    text,
6373                ))) => {
6374                    full_content.push_str(&text.text);
6375                }
6376                Ok(MultiTurnStreamItem::FinalResponse(_)) => {
6377                    break;
6378                }
6379                Err(e) => {
6380                    tracing::warn!("Error: {:?}", e);
6381                    break;
6382                }
6383                _ => {}
6384            }
6385        }
6386
6387        tracing::info!("Got response: {:?}", full_content);
6388
6389        // Stop background logger
6390        stop.store(true, Ordering::Relaxed);
6391        bg_handle.await?;
6392
6393        let leaks = leak_count.load(Ordering::Relaxed);
6394        anyhow::ensure!(
6395            leaks == 0,
6396            "SPAN LEAK DETECTED: Background logger was inside unexpected spans {leaks} times. \
6397             This indicates that span.enter() is being used inside async_stream instead of .instrument()"
6398        );
6399
6400        Ok(())
6401    }
6402
6403    /// Test that FinalResponse contains the updated chat history when a starting
6404    /// history is provided via `.history(..)`.
6405    ///
6406    /// This verifies that:
6407    /// 1. PromptResponse.messages() returns Some when a starting history was provided
6408    /// 2. The history contains both the user prompt and assistant response
6409    #[tokio::test]
6410    #[ignore = "This requires an API key"]
6411    async fn test_chat_history_in_final_response() -> anyhow::Result<()> {
6412        use rig_core::message::Message;
6413
6414        let client = anthropic::Client::from_env()?;
6415        let agent = client
6416            .agent(anthropic::completion::CLAUDE_HAIKU_4_5)
6417            .preamble("You are a helpful assistant. Keep responses brief.")
6418            .temperature(0.1)
6419            .max_tokens(50)
6420            .build();
6421
6422        // Send streaming request with history
6423        let empty_history: &[Message] = &[];
6424        let mut stream = agent
6425            .stream_prompt("Say 'hello' and nothing else.")
6426            .history(empty_history)
6427            .await;
6428
6429        // Consume the stream and collect FinalResponse
6430        let mut response_text = String::new();
6431        let mut final_history = None;
6432        while let Some(item) = stream.next().await {
6433            match item {
6434                Ok(MultiTurnStreamItem::StreamAssistantItem(StreamedAssistantContent::Text(
6435                    text,
6436                ))) => {
6437                    response_text.push_str(&text.text);
6438                }
6439                Ok(MultiTurnStreamItem::FinalResponse(res)) => {
6440                    final_history = res.messages().map(|h| h.to_vec());
6441                    break;
6442                }
6443                Err(e) => {
6444                    return Err(e.into());
6445                }
6446                _ => {}
6447            }
6448        }
6449
6450        let history = final_history
6451            .ok_or_else(|| anyhow::anyhow!("final response should include history"))?;
6452
6453        // Should contain at least the user message
6454        anyhow::ensure!(
6455            history.iter().any(|m| matches!(m, Message::User { .. })),
6456            "History should contain the user message"
6457        );
6458
6459        // Should contain the assistant response
6460        anyhow::ensure!(
6461            history
6462                .iter()
6463                .any(|m| matches!(m, Message::Assistant { .. })),
6464            "History should contain the assistant response"
6465        );
6466
6467        tracing::info!(
6468            "History after streaming: {} messages, response: {:?}",
6469            history.len(),
6470            response_text
6471        );
6472
6473        Ok(())
6474    }
6475
6476    #[tokio::test]
6477    async fn streaming_appends_to_memory_after_final_response() {
6478        use rig_core::memory::{ConversationMemory, InMemoryConversationMemory};
6479
6480        let memory = InMemoryConversationMemory::new();
6481        let agent = AgentBuilder::new(streaming_text_then_final_model())
6482            .memory(memory.clone())
6483            .build();
6484
6485        let mut stream = agent
6486            .stream_prompt("hi there")
6487            .conversation("stream-thread")
6488            .await;
6489
6490        let mut history_in_final = None;
6491        while let Some(item) = stream.next().await {
6492            match item {
6493                Ok(MultiTurnStreamItem::FinalResponse(res)) => {
6494                    history_in_final = res.messages().map(|h| h.to_vec());
6495                    break;
6496                }
6497                Ok(_) => {}
6498                Err(err) => panic!("unexpected streaming error: {err:?}"),
6499            }
6500        }
6501
6502        let final_history = history_in_final
6503            .expect("PromptResponse.messages should be populated when memory is configured");
6504        assert_eq!(
6505            final_history.len(),
6506            2,
6507            "user prompt + assistant response in final history: {final_history:?}"
6508        );
6509
6510        let stored = memory.load("stream-thread").await.unwrap();
6511        assert_eq!(stored.len(), 2, "memory should contain user + assistant");
6512    }
6513
6514    #[tokio::test]
6515    async fn streaming_reasoning_without_tools_does_not_duplicate_final_history() {
6516        let agent = AgentBuilder::new(MockCompletionModel::from_stream_turns([[
6517            MockStreamEvent::text("final answer"),
6518            MockStreamEvent::reasoning("reasoned step").with_reasoning_id("rs_1"),
6519            MockStreamEvent::final_response_with_total_tokens(3),
6520        ]]))
6521        .build();
6522
6523        let mut stream = agent
6524            .stream_prompt("think before answering")
6525            .history(Vec::<Message>::new())
6526            .await;
6527
6528        let mut history_in_final = None;
6529        while let Some(item) = stream.next().await {
6530            match item {
6531                Ok(MultiTurnStreamItem::FinalResponse(res)) => {
6532                    history_in_final = res.messages().map(|h| h.to_vec());
6533                    break;
6534                }
6535                Ok(_) => {}
6536                Err(err) => panic!("unexpected streaming error: {err:?}"),
6537            }
6538        }
6539
6540        let final_history = history_in_final
6541            .expect("PromptResponse.messages should be populated when with_history is used");
6542        assert_eq!(
6543            final_history.len(),
6544            2,
6545            "user prompt + one assistant response in final history: {final_history:?}"
6546        );
6547
6548        assert!(matches!(
6549            final_history.first(),
6550            Some(Message::User { content })
6551                if matches!(
6552                    content.first(),
6553                    UserContent::Text(text) if text.text == "think before answering"
6554                )
6555        ));
6556
6557        let assistant_messages = final_history
6558            .iter()
6559            .filter_map(|message| match message {
6560                Message::Assistant { content, .. } => Some(content),
6561                _ => None,
6562            })
6563            .collect::<Vec<_>>();
6564        assert_eq!(
6565            assistant_messages.len(),
6566            1,
6567            "reasoning turn should produce exactly one assistant history message: {final_history:?}"
6568        );
6569        let assistant_content = assistant_messages
6570            .first()
6571            .expect("expected assistant history message");
6572        assert!(assistant_content.iter().any(|item| matches!(
6573            item,
6574            AssistantContent::Text(text) if text.text == "final answer"
6575        )));
6576        assert!(assistant_content.iter().any(|item| matches!(
6577            item,
6578            AssistantContent::Reasoning(reasoning)
6579                if reasoning.id.as_deref() == Some("rs_1")
6580                    && reasoning.content.iter().any(|content| matches!(
6581                        content,
6582                        ReasoningContent::Text { text, .. } if text == "reasoned step"
6583                    ))
6584        )));
6585        let reasoning_index = assistant_content
6586            .iter()
6587            .position(|item| matches!(item, AssistantContent::Reasoning(_)))
6588            .expect("assistant history should contain reasoning");
6589        let text_index = assistant_content
6590            .iter()
6591            .position(|item| matches!(item, AssistantContent::Text(_)))
6592            .expect("assistant history should contain text");
6593        assert!(
6594            reasoning_index < text_index,
6595            "assistant reasoning must be stored before assistant text: {assistant_content:?}"
6596        );
6597    }
6598
6599    #[tokio::test]
6600    async fn streaming_with_history_overrides_memory() {
6601        use rig_core::memory::{ConversationMemory, InMemoryConversationMemory};
6602
6603        let memory = InMemoryConversationMemory::new();
6604        memory
6605            .append("t1", vec![Message::user("from-memory")])
6606            .await
6607            .unwrap();
6608
6609        let agent = AgentBuilder::new(streaming_text_then_final_model())
6610            .memory(memory.clone())
6611            .build();
6612
6613        let mut stream = agent
6614            .stream_prompt("hi")
6615            .conversation("t1")
6616            .history(vec![Message::user("from-caller")])
6617            .await;
6618
6619        while let Some(item) = stream.next().await {
6620            if let Ok(MultiTurnStreamItem::FinalResponse(_)) = item {
6621                break;
6622            }
6623        }
6624
6625        let stored = memory.load("t1").await.unwrap();
6626        assert_eq!(
6627            stored.len(),
6628            1,
6629            "with_history bypasses memory; only the pre-seeded entry remains: {stored:?}"
6630        );
6631    }
6632
6633    #[tokio::test]
6634    async fn streaming_without_memory_disables_for_request() {
6635        use rig_core::memory::{ConversationMemory, InMemoryConversationMemory};
6636
6637        let memory = InMemoryConversationMemory::new();
6638        let agent = AgentBuilder::new(streaming_text_then_final_model())
6639            .memory(memory.clone())
6640            .conversation("default")
6641            .build();
6642
6643        let mut stream = agent.stream_prompt("hi").without_memory().await;
6644
6645        while let Some(item) = stream.next().await {
6646            if let Ok(MultiTurnStreamItem::FinalResponse(_)) = item {
6647                break;
6648            }
6649        }
6650
6651        let stored = memory.load("default").await.unwrap();
6652        assert!(stored.is_empty(), "without_memory disables save");
6653    }
6654
6655    #[tokio::test]
6656    async fn streaming_load_error_yields_memory_error() {
6657        let agent = AgentBuilder::new(streaming_text_then_final_model())
6658            .memory(FailingMemory::default())
6659            .build();
6660
6661        let mut stream = agent.stream_prompt("hi").conversation("t1").await;
6662
6663        let first = stream.next().await.expect("at least one item");
6664        match first {
6665            Err(StreamingError::Prompt(err)) => match *err {
6666                PromptError::MemoryError(err) => {
6667                    assert!(err.to_string().contains("load boom"));
6668                }
6669                other => panic!("expected PromptError::MemoryError, got {other:?}"),
6670            },
6671            other => panic!("expected StreamingError::Prompt, got {other:?}"),
6672        }
6673    }
6674
6675    #[tokio::test]
6676    async fn streaming_with_filter_shapes_loaded_history() {
6677        use rig_core::memory::{ConversationMemory, InMemoryConversationMemory};
6678
6679        let memory = InMemoryConversationMemory::new()
6680            .with_filter(|msgs: Vec<Message>| msgs.into_iter().rev().take(2).rev().collect());
6681        memory
6682            .append(
6683                "t1",
6684                vec![
6685                    Message::user("1"),
6686                    Message::assistant("2"),
6687                    Message::user("3"),
6688                    Message::assistant("4"),
6689                ],
6690            )
6691            .await
6692            .unwrap();
6693
6694        let model = MockCompletionModel::from_stream_turns([[
6695            MockStreamEvent::text("ok"),
6696            MockStreamEvent::final_response_with_total_tokens(1),
6697        ]]);
6698        let recorded = model.clone();
6699        let agent = AgentBuilder::new(model).memory(memory).build();
6700
6701        let mut stream = agent.stream_prompt("ping").conversation("t1").await;
6702        while let Some(item) = stream.next().await {
6703            if let Ok(MultiTurnStreamItem::FinalResponse(_)) = item {
6704                break;
6705            }
6706        }
6707
6708        let received = recorded.requests()[0]
6709            .chat_history
6710            .iter()
6711            .cloned()
6712            .collect::<Vec<_>>();
6713        assert_eq!(
6714            received.len(),
6715            3,
6716            "window-truncated history (2) + current prompt: {received:?}"
6717        );
6718    }
6719
6720    #[tokio::test]
6721    async fn streaming_append_error_does_not_suppress_final_response() {
6722        let agent = AgentBuilder::new(streaming_text_then_final_model())
6723            .memory(AppendFailingMemory::default())
6724            .build();
6725
6726        let mut stream = agent.stream_prompt("hi").conversation("t1").await;
6727
6728        let mut saw_final = false;
6729        while let Some(item) = stream.next().await {
6730            if let Ok(MultiTurnStreamItem::FinalResponse(_)) = item {
6731                saw_final = true;
6732                break;
6733            }
6734        }
6735        assert!(
6736            saw_final,
6737            "FinalResponse must be yielded even when memory.append fails"
6738        );
6739    }
6740}