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#[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 StreamAssistantItem(StreamedAssistantContent<R>),
71 ToolExecutionCommitted {
81 tool_call: rig_core::message::ToolCall,
87 internal_call_id: String,
91 },
92 StreamUserItem(StreamedUserContent),
98 CompletionCall(CompletionCall),
116 ModelTurnRetried {
123 turn: usize,
125 },
126 FinalResponse(PromptResponse),
129}
130
131fn 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
191async 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
230fn 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
284pub struct StreamingPromptRequest<M>
292where
293 M: CompletionModel,
294{
295 runner: AgentRunner<M>,
297}
298
299impl<M> StreamingPromptRequest<M>
300where
301 M: CompletionModel + 'static,
302 <M as CompletionModel>::StreamingResponse: WasmCompatSend + GetTokenUsage,
303{
304 pub fn new(agent: Arc<Agent<M>>, prompt: impl Into<Message>) -> StreamingPromptRequest<M> {
307 Self::from_agent(agent.as_ref(), prompt)
308 }
309
310 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 pub fn max_turns(mut self, turns: usize) -> Self {
327 self.runner = self.runner.max_turns(turns);
328 self
329 }
330
331 pub fn tool_concurrency(mut self, concurrency: usize) -> Self {
339 self.runner = self.runner.tool_concurrency(concurrency);
340 self
341 }
342
343 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#[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#[allow(clippy::large_enum_variant)]
387pub(crate) enum DriveItem<R> {
388 Item(MultiTurnStreamItem<R>),
392 Done(Box<PromptResponse>),
396}
397
398pub(crate) trait TurnSource<M>: WasmCompatSend
406where
407 M: CompletionModel,
408{
409 type Raw: WasmCompatSend;
411
412 fn open_chat_span(
415 &self,
416 runner: &AgentRunner<M>,
417 effective_preamble: Option<&str>,
418 ) -> tracing::Span;
419
420 #[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 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 fn record_run_level_telemetry(
449 &self,
450 agent_span: &tracing::Span,
451 response: &PromptResponse,
452 created_agent_span: bool,
453 );
454
455 fn final_item(&self, response: &PromptResponse) -> Option<MultiTurnStreamItem<Self::Raw>>;
458}
459
460pub(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
479pub(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 let hook_ctx = HookContext::new(is_streaming, runner.agent_name.clone());
504 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 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 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 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 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
675pub(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 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 enum ToolSurface {
729 Executed(Box<rig_core::message::ToolCall>),
731 Skipped,
732 Preresolved,
733 }
734 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 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 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 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 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 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 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 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 if let Some((_, err)) = first_error {
915 yield Err(StreamingError::Prompt(Box::new(err)));
916 return;
917 }
918
919 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 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
982pub(crate) struct StreamingTurnSource {
985 last_final_choice: OneOrMany<AssistantContent>,
988 last_message_id: Option<String>,
989 agent_name: String,
991 created_agent_span: bool,
995 record_telemetry_content: bool,
997 observes_text_delta: bool,
1000 observes_tool_call_delta: bool,
1001 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 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 let mut turn_recovered = false;
1080
1081 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 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 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 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 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: ¤t_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 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 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 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 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 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 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 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 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 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 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 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>; type IntoFuture = WasmBoxedFuture<'static, Self::Output>;
1580
1581 fn into_future(self) -> Self::IntoFuture {
1582 Box::pin(async move { self.send().await })
1584 }
1585}
1586
1587pub 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 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 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 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 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 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 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 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 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 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 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 #[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 #[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 #[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 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 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 #[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 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 tokio::time::sleep(Duration::from_millis(100)).await;
6353
6354 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.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 #[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 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 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 anyhow::ensure!(
6455 history.iter().any(|m| matches!(m, Message::User { .. })),
6456 "History should contain the user message"
6457 );
6458
6459 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}