Skip to main content

rig_core/test_utils/
streaming.rs

1//! Streaming helpers for [`MockCompletionModel`](super::MockCompletionModel):
2//! the script a mock reply is written from, and the decoder that writes it
3//! through the completion writer, as a wire's decoder does.
4
5use std::collections::HashMap;
6
7use crate::completion::{CompletionResponse, Usage};
8use crate::error::ProviderError;
9use crate::operation::{Block, CallFragment, Completion, Finish};
10use crate::wire::{Decoder, Flow, Out, WireEvent};
11
12/// Provider descriptor name reported by the test doubles.
13pub const MOCK_PROVIDER: &str = "mock";
14
15/// The end the mock model's reply finishes with, carrying `usage`.
16pub fn mock_final(usage: Usage) -> Finish {
17    Finish {
18        usage,
19        ..Finish::default()
20    }
21}
22
23/// A fixture's provider item: a non-empty JSON object, or `None` for
24/// `null` and `{}`. Any other value is a scripting mistake surfaced as a
25/// stream error.
26fn fixture_item(value: serde_json::Value) -> Result<Option<serde_json::Value>, ProviderError> {
27    match value {
28        serde_json::Value::Null => Ok(None),
29        serde_json::Value::Object(map) if map.is_empty() => Ok(None),
30        serde_json::Value::Object(map) => Ok(Some(serde_json::Value::Object(map))),
31        other => Err(ProviderError::Provider(format!(
32            "mock stream fixture provider item must be a JSON object, got: {other}"
33        ))),
34    }
35}
36
37/// The end of a reply whose usage has only `total_tokens` set.
38pub fn mock_final_with_total_tokens(total_tokens: u64) -> Finish {
39    mock_final(Usage {
40        total_tokens: Some(total_tokens),
41        ..Default::default()
42    })
43}
44
45/// Scripted streaming event yielded by [`MockCompletionModel`](super::MockCompletionModel).
46#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
47pub enum MockStreamEvent {
48    /// Text chunk.
49    Text(String),
50    /// Start a new text part, with the provider item it is decoded from.
51    TextStart {
52        id: String,
53        additional_params: Option<serde_json::Value>,
54    },
55    /// Fields merged into the current text part's provider item.
56    TextAdditionalParams(serde_json::Value),
57    /// Complete tool call event.
58    ToolCall {
59        id: String,
60        name: String,
61        arguments: serde_json::Value,
62        call_id: Option<String>,
63    },
64    /// A tool call's name, as a wire streams it.
65    ToolCallNameDelta { id: String, name: String },
66    /// A fragment of a tool call's arguments.
67    ToolCallArgumentsDelta { id: String, arguments: String },
68    /// The end of a tool call streamed as fragments: the call closes with
69    /// what they carried, as a wire's step-stop does.
70    ToolCallEnd { id: String },
71    /// Complete reasoning event.
72    Reasoning { id: String, text: String },
73    /// Reasoning delta event.
74    ReasoningDelta { id: String, reasoning: String },
75    /// Provider-native output item that Rig does not model.
76    Unknown(serde_json::Value),
77    /// The transport request id of this turn. The mock transport reports
78    /// it, as a real transport reports a response header; the decoder never
79    /// sees it. Its position in the turn does not matter; when a turn
80    /// scripts more than one, the first wins.
81    RequestId(String),
82    /// The provider's end of the reply.
83    FinalResponse(Finish),
84    /// A failure, which ends the reply.
85    Error(MockError),
86}
87
88use super::completion::MockError;
89
90/// The provider id a fixture spells, when it spells one. Corpus fixtures
91/// are plain data: the renderings `reasoning-0`, `block-3`, `output-1`,
92/// `tool-2` and `text-0` name a part the wire gave no id, and anything else
93/// is the wire's own id.
94fn fixture_provider_id(id: &str) -> Option<&str> {
95    let unnamed = ["reasoning-", "block-", "output-", "tool-", "text-"]
96        .iter()
97        .any(|namespace| {
98            id.strip_prefix(namespace)
99                .is_some_and(|rest| rest.parse::<u64>().is_ok())
100        });
101    (!id.is_empty() && !unnamed).then_some(id)
102}
103
104impl MockStreamEvent {
105    /// Create a text chunk.
106    pub fn text(text: impl Into<String>) -> Self {
107        Self::Text(text.into())
108    }
109
110    /// Start a new text content block identified by `id`.
111    pub fn text_start(id: impl Into<String>, additional_params: Option<serde_json::Value>) -> Self {
112        Self::TextStart {
113            id: id.into(),
114            additional_params,
115        }
116    }
117
118    /// Add provider-specific metadata to the current text content block.
119    pub fn text_additional_params(additional_params: serde_json::Value) -> Self {
120        Self::TextAdditionalParams(additional_params)
121    }
122
123    /// Create a complete tool call event.
124    pub fn tool_call(
125        id: impl Into<String>,
126        name: impl Into<String>,
127        arguments: serde_json::Value,
128    ) -> Self {
129        Self::ToolCall {
130            id: id.into(),
131            name: name.into(),
132            arguments,
133            call_id: None,
134        }
135    }
136
137    /// Attach a provider-specific call ID to a complete tool call event.
138    pub fn with_call_id(mut self, call_id: impl Into<String>) -> Self {
139        if let Self::ToolCall { call_id: id, .. } = &mut self {
140            *id = Some(call_id.into());
141        }
142        self
143    }
144
145    /// Create a tool call name delta.
146    pub fn tool_call_name_delta(id: impl Into<String>, name: impl Into<String>) -> Self {
147        Self::ToolCallNameDelta {
148            id: id.into(),
149            name: name.into(),
150        }
151    }
152
153    /// Create a tool call arguments delta.
154    pub fn tool_call_arguments_delta(id: impl Into<String>, arguments: impl Into<String>) -> Self {
155        Self::ToolCallArgumentsDelta {
156            id: id.into(),
157            arguments: arguments.into(),
158        }
159    }
160
161    /// Create the end of a tool call streamed as deltas.
162    pub fn tool_call_end(id: impl Into<String>) -> Self {
163        Self::ToolCallEnd { id: id.into() }
164    }
165
166    /// Create a complete reasoning event with the default mock id
167    /// (`"reasoning-0"`). Use [`Self::with_reasoning_id`] for tests that
168    /// need distinct reasoning items.
169    pub fn reasoning(reasoning: impl Into<String>) -> Self {
170        Self::Reasoning {
171            id: "reasoning-0".to_string(),
172            text: reasoning.into(),
173        }
174    }
175
176    /// Attach a provider-specific reasoning ID to a complete reasoning event.
177    pub fn with_reasoning_id(mut self, reasoning_id: impl Into<String>) -> Self {
178        if let Self::Reasoning { id, .. } = &mut self {
179            *id = reasoning_id.into();
180        }
181        self
182    }
183
184    /// Create a reasoning delta event with the default mock id
185    /// (`"reasoning-0"`). Use [`Self::reasoning_delta_with_id`] for tests
186    /// that need distinct reasoning items.
187    pub fn reasoning_delta(reasoning: impl Into<String>) -> Self {
188        Self::reasoning_delta_with_id("reasoning-0", reasoning)
189    }
190
191    /// Create a reasoning delta event with an explicit reasoning item id.
192    pub fn reasoning_delta_with_id(id: impl Into<String>, reasoning: impl Into<String>) -> Self {
193        Self::ReasoningDelta {
194            id: id.into(),
195            reasoning: reasoning.into(),
196        }
197    }
198
199    /// Create an unmodeled provider output item.
200    pub fn unknown(value: serde_json::Value) -> Self {
201        Self::Unknown(value)
202    }
203
204    /// Create the provider's end of the reply, with usage.
205    pub fn final_response(usage: Usage) -> Self {
206        Self::FinalResponse(mock_final(usage))
207    }
208
209    /// Create a final response event whose usage reports no counter.
210    pub fn final_response_with_default_usage() -> Self {
211        Self::FinalResponse(mock_final(Usage::default()))
212    }
213
214    /// Create a final response event whose usage has only `total_tokens` set.
215    pub fn final_response_with_total_tokens(total_tokens: u64) -> Self {
216        Self::FinalResponse(mock_final_with_total_tokens(total_tokens))
217    }
218
219    /// Create a stream error event.
220    pub fn error(message: impl Into<String>) -> Self {
221        Self::Error(MockError::provider(message))
222    }
223}
224
225/// One step of a mock reply: a scripted event, or a whole response.
226#[derive(Clone, Debug)]
227pub enum MockFrame {
228    /// A scripted event.
229    Event(MockStreamEvent),
230    /// A whole response, as a unary turn answers.
231    Response(Box<CompletionResponse>),
232}
233
234/// The reply document of [`MockScript`](super::MockScript): a whole
235/// response's `raw`, or a stream's scripted end, serialized. A stream that
236/// failed before its end has none.
237#[derive(Debug, Default)]
238pub struct MockDocument {
239    document: Option<serde_json::Value>,
240    failed: bool,
241}
242
243impl crate::wire::document::Serves<crate::operation::Completion> for MockDocument {}
244
245impl crate::wire::document::Reassemble<MockFrame> for MockDocument {
246    fn absorb(&mut self, frame: &MockFrame) {
247        if self.document.is_some() || self.failed {
248            return;
249        }
250        match frame {
251            MockFrame::Response(response) => self.document = Some(response.raw.clone()),
252            MockFrame::Event(MockStreamEvent::FinalResponse(finish)) => {
253                self.document = serde_json::to_value(finish).ok();
254            }
255            MockFrame::Event(MockStreamEvent::Error(_)) => self.failed = true,
256            MockFrame::Event(_) => {}
257        }
258    }
259
260    fn finish(self) -> serde_json::Value {
261        self.document.unwrap_or(serde_json::Value::Null)
262    }
263}
264
265/// The decoder of [`MockScript`](super::MockScript): each scripted step is
266/// written through the completion writer.
267#[derive(Default)]
268pub struct MockDecoder<'id> {
269    /// The reasoning item each scripted id streams into, in start order.
270    reasoning: Vec<(String, usize)>,
271    /// The item index each scripted call id streams under.
272    calls: HashMap<String, usize>,
273    brand: std::marker::PhantomData<fn(&'id ()) -> &'id ()>,
274}
275
276/// The provider item a scripted id names, when it names one.
277fn id_item(id: &str) -> serde_json::Value {
278    fixture_provider_id(id).map_or(
279        serde_json::Value::Null,
280        |id| serde_json::json!({ "id": id }),
281    )
282}
283
284impl<'id> MockDecoder<'id> {
285    fn call_index(&mut self, out: &mut Out<'id, Completion>, id: &str) -> usize {
286        if let Some(index) = self.calls.get(id) {
287            return *index;
288        }
289        let index = out.fresh_index();
290        if !id.is_empty() {
291            self.calls.insert(id.to_owned(), index);
292        }
293        index
294    }
295
296    fn event(
297        &mut self,
298        event: MockStreamEvent,
299        mut out: Out<'id, Completion>,
300    ) -> Result<Flow, ProviderError> {
301        match event {
302            MockStreamEvent::Text(text) => {
303                out.run(Block::Text, &text)?;
304            }
305            MockStreamEvent::TextStart {
306                id: _,
307                additional_params,
308            } => {
309                out.end_run()?;
310                let index = out.run(Block::Text, "")?;
311                if let Some(item) = additional_params.map(fixture_item).transpose()?.flatten() {
312                    out.edit(index, |slot| *slot = item)?;
313                }
314            }
315            MockStreamEvent::TextAdditionalParams(additional_params) => {
316                // An empty fixture object is a scripting mistake, not a no-op.
317                let Some(serde_json::Value::Object(fields)) = fixture_item(additional_params)?
318                else {
319                    return Err(ProviderError::Provider(
320                        "mock stream fixture `TextAdditionalParams` carries no data — \
321                         drop the event instead"
322                            .to_string(),
323                    ));
324                };
325                let index = out.run(Block::Text, "")?;
326                out.edit(index, |item| {
327                    if !item.is_object() {
328                        *item = serde_json::Value::Object(serde_json::Map::new());
329                    }
330                    if let Some(item) = item.as_object_mut() {
331                        item.extend(fields);
332                    }
333                })?;
334            }
335            MockStreamEvent::ToolCall {
336                id,
337                name,
338                arguments,
339                call_id,
340            } => {
341                out.end_run()?;
342                // An id-less call is a wire that sends none: rig issues the
343                // id. A call scripted as fragments under this id is the one
344                // this restatement closes.
345                let index = match self.calls.remove(&id) {
346                    Some(index) => index,
347                    None => out.fresh_index(),
348                };
349                out.fragment(
350                    Some(index),
351                    CallFragment {
352                        id: call_id.as_deref().or(fixture_provider_id(&id)),
353                        name: Some(name.as_str()),
354                        ..CallFragment::default()
355                    },
356                )?;
357                out.announce(index, arguments)?;
358                out.finish(index)?;
359            }
360            MockStreamEvent::ToolCallNameDelta { id, name } => {
361                out.end_run()?;
362                let index = self.call_index(&mut out, &id);
363                out.fragment(
364                    Some(index),
365                    CallFragment {
366                        id: fixture_provider_id(&id),
367                        name: Some(name.as_str()),
368                        ..CallFragment::default()
369                    },
370                )?;
371            }
372            MockStreamEvent::ToolCallArgumentsDelta { id, arguments } => {
373                out.end_run()?;
374                let index = self.call_index(&mut out, &id);
375                out.fragment(
376                    Some(index),
377                    CallFragment {
378                        id: fixture_provider_id(&id),
379                        arguments: Some(arguments.as_str()),
380                        ..CallFragment::default()
381                    },
382                )?;
383            }
384            MockStreamEvent::ToolCallEnd { id } => {
385                let index = self.call_index(&mut out, &id);
386                self.calls.remove(&id);
387                out.finish(index)?;
388            }
389            MockStreamEvent::Reasoning { id, text } => {
390                out.end_run()?;
391                // A whole reasoning closes the part streamed under its id.
392                match self.reasoning.iter().position(|(open, _)| *open == id) {
393                    Some(at) => {
394                        let (_, index) = self.reasoning.remove(at);
395                        out.finish(index)?;
396                    }
397                    None => {
398                        let index = out.fresh_index();
399                        out.whole(
400                            index,
401                            Block::Reasoning { redacted: false },
402                            id_item(&id),
403                            &text,
404                        )?;
405                    }
406                }
407            }
408            MockStreamEvent::ReasoningDelta { id, reasoning } => {
409                out.end_run()?;
410                let index = match self.reasoning.iter().find(|(open, _)| *open == id) {
411                    Some((_, index)) => *index,
412                    None => {
413                        let index = out.fresh_index();
414                        out.open(index, Block::Reasoning { redacted: false }, id_item(&id))?;
415                        self.reasoning.push((id, index));
416                        index
417                    }
418                };
419                out.push(index, &reasoning)?;
420            }
421            MockStreamEvent::Unknown(value) => out.unknown(value.into()),
422            MockStreamEvent::RequestId(_) => {}
423            MockStreamEvent::FinalResponse(finish) => {
424                out.end_run()?;
425                for (_, index) in std::mem::take(&mut self.reasoning) {
426                    out.finish(index)?;
427                }
428                return Ok(out.end(finish));
429            }
430            MockStreamEvent::Error(error) => return Err(error.into_completion_error()),
431        }
432        Ok(Flow::More)
433    }
434
435    fn response(
436        &mut self,
437        response: CompletionResponse,
438        mut out: Out<'id, Completion>,
439    ) -> Result<Flow, ProviderError> {
440        for content in response.choice.iter().cloned() {
441            out.content(content)?;
442        }
443        Ok(out.end(Finish {
444            usage: response.usage,
445            reason: response.finish_reason(),
446            response_id: response.response_id().map(str::to_owned),
447            model: response.model().map(str::to_owned),
448            error: response.error.clone(),
449        }))
450    }
451}
452
453impl<'id> Decoder<'id, Completion, MockFrame> for MockDecoder<'id> {
454    type Event = MockFrame;
455
456    fn classify(&self, frame: MockFrame) -> WireEvent<MockFrame> {
457        WireEvent::Known(frame)
458    }
459
460    fn decode(
461        &mut self,
462        frame: MockFrame,
463        out: Out<'id, Completion>,
464    ) -> Result<Flow, ProviderError> {
465        match frame {
466            MockFrame::Event(event) => self.event(event, out),
467            MockFrame::Response(response) => self.response(*response, out),
468        }
469    }
470}