Skip to main content

rig_core/test_utils/
completion.rs

1//! Completion helpers for deterministic agent-loop tests.
2
3use std::{
4    collections::VecDeque,
5    sync::{Arc, Mutex, MutexGuard},
6};
7
8use crate::driver::{Exchange, Model, Opened, Opening, Transport};
9use crate::error::{EncodeError, ProviderError};
10use crate::operation::Completion;
11use crate::wire::{Capabilities, Descriptor, Mode, Wire};
12use crate::{
13    completion::{AssistantContent, CompletionRequest, CompletionResponse, Usage},
14    message::{ToolCall, ToolFunction},
15};
16
17use super::streaming::{MOCK_PROVIDER, MockDecoder, MockFrame, MockStreamEvent};
18
19/// Scripted error returned by [`MockCompletionModel`].
20#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
21pub enum MockError {
22    /// Provider error.
23    Provider(String),
24    /// Request construction error.
25    Request(String),
26    /// A preserved provider error response (rig#2314), id included.
27    ProviderResponse(crate::provider_response::ProviderResponseError),
28}
29
30impl MockError {
31    /// Create a provider error.
32    pub fn provider(message: impl Into<String>) -> Self {
33        Self::Provider(message.into())
34    }
35
36    /// Create a request error.
37    pub fn request(message: impl Into<String>) -> Self {
38        Self::Request(message.into())
39    }
40
41    pub(crate) fn into_completion_error(self) -> ProviderError {
42        match self {
43            Self::Provider(message) => ProviderError::Provider(message),
44            Self::Request(message) => ProviderError::request(message),
45            Self::ProviderResponse(response) => ProviderError::ProviderResponse(response),
46        }
47    }
48}
49
50/// A scripted non-streaming mock completion turn.
51///
52/// A turn is data: a script serializes, so a scripted model can be written
53/// to a fixture and read back (see `MockCompletionModel::script`).
54#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
55pub struct MockTurn {
56    response: Result<MockTurnResponse, MockError>,
57}
58
59#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
60struct MockTurnResponse {
61    choice: Vec<AssistantContent>,
62    usage: Usage,
63    message_id: Option<String>,
64    response_id: Option<String>,
65    provider_request_id: Option<String>,
66    finish_reason: Option<crate::completion::FinishReason>,
67    /// A scripted provider document, when the test supplies one; otherwise
68    /// the turn itself, serialized, is the mock's document. Absent from the
69    /// serialized turn when unscripted, so that document never nests
70    /// itself; a scripted one survives a serde round trip of the script.
71    #[serde(
72        default,
73        skip_serializing_if = "Option::is_none",
74        deserialize_with = "deserialize_scripted_raw"
75    )]
76    raw: Option<serde_json::Value>,
77}
78
79fn deserialize_scripted_raw<'de, D: serde::Deserializer<'de>>(
80    deserializer: D,
81) -> Result<Option<serde_json::Value>, D::Error> {
82    // Only a missing field means unscripted; explicit JSON null is a document.
83    serde::Deserialize::deserialize(deserializer).map(Some)
84}
85
86impl MockTurn {
87    /// Create a text response turn.
88    pub fn text(text: impl Into<String>) -> Self {
89        Self::from_content(AssistantContent::text(text.into()))
90    }
91
92    /// Create a tool-call response turn. An empty `name` scripts a provider
93    /// error instead: no call can be built without a name.
94    pub fn tool_call(
95        id: impl Into<String>,
96        name: impl Into<String>,
97        arguments: serde_json::Value,
98    ) -> Self {
99        match crate::message::ToolName::new(name) {
100            Ok(name) => Self::from_content(AssistantContent::ToolCall(ToolCall::from_wire(
101                id,
102                ToolFunction::new(name, arguments),
103            ))),
104            Err(error) => Self::error(error.to_string()),
105        }
106    }
107
108    /// Create a provider-error response turn.
109    pub fn error(message: impl Into<String>) -> Self {
110        Self {
111            response: Err(MockError::provider(message)),
112        }
113    }
114
115    /// Create a provider-response error turn carrying a transport request id
116    /// (rig#2314): the scripted failure a test uses to assert error-identity
117    /// attribution.
118    pub fn provider_response_error(
119        status: http::StatusCode,
120        body: impl Into<String>,
121        request_id: impl Into<String>,
122    ) -> Self {
123        Self {
124            response: Err(MockError::ProviderResponse(
125                crate::provider_response::ProviderResponseError::new(status, body)
126                    .with_provider_request_id(Some(request_id.into())),
127            )),
128        }
129    }
130
131    /// Create a request-error response turn.
132    pub fn request_error(message: impl Into<String>) -> Self {
133        Self {
134            response: Err(MockError::request(message)),
135        }
136    }
137
138    /// Create a response turn from one assistant content item.
139    pub fn from_content(content: AssistantContent) -> Self {
140        Self {
141            response: Ok(MockTurnResponse {
142                choice: vec![content],
143                usage: Usage::default(),
144                message_id: None,
145                response_id: None,
146                provider_request_id: None,
147                finish_reason: None,
148                raw: None,
149            }),
150        }
151    }
152
153    /// Create a response turn from assistant content items.
154    ///
155    /// Infallible now that content is a `Vec`: an empty turn is a shape a
156    /// provider can genuinely return, so it is a value to build, not an error.
157    pub fn from_contents(content: impl IntoIterator<Item = AssistantContent>) -> Self {
158        Self {
159            response: Ok(MockTurnResponse {
160                choice: content.into_iter().collect(),
161                usage: Usage::default(),
162                message_id: None,
163                response_id: None,
164                provider_request_id: None,
165                finish_reason: None,
166                raw: None,
167            }),
168        }
169    }
170
171    /// Attach a provider-specific call ID to a tool-call response turn.
172    pub fn with_call_id(mut self, call_id: impl Into<String>) -> Self {
173        let call_id = call_id.into();
174        if let Ok(response) = &mut self.response {
175            for content in response.choice.iter_mut() {
176                if let AssistantContent::ToolCall(tool_call) = content {
177                    tool_call.id = crate::message::CallId::from_wire(call_id);
178                    break;
179                }
180            }
181        }
182        self
183    }
184
185    /// Override usage for this turn.
186    pub fn with_usage(mut self, usage: Usage) -> Self {
187        if let Ok(response) = &mut self.response {
188            response.usage = usage;
189        }
190        self
191    }
192
193    /// Set a provider-assigned assistant message ID for this turn.
194    pub fn with_message_id(mut self, message_id: impl Into<String>) -> Self {
195        if let Ok(response) = &mut self.response {
196            response.message_id = Some(message_id.into());
197        }
198        self
199    }
200
201    /// Set a provider-assigned response-scoped ID for this turn.
202    pub fn with_response_id(mut self, response_id: impl Into<String>) -> Self {
203        if let Ok(response) = &mut self.response {
204            response.response_id = Some(response_id.into());
205        }
206        self
207    }
208
209    /// Set a provider transport request id for this turn.
210    pub fn with_provider_request_id(mut self, request_id: impl Into<String>) -> Self {
211        if let Ok(response) = &mut self.response {
212            response.provider_request_id = Some(request_id.into());
213        }
214        self
215    }
216
217    /// Set the terminal finish reason for this turn.
218    ///
219    /// Without this, a mocked blocking turn always reports `None`, which
220    /// leaves the whole blocking half of the truncation contract (rig#2322)
221    /// unexercisable — the streamed mock could script a reason and the
222    /// blocking one could not.
223    pub fn with_finish_reason(mut self, finish_reason: crate::completion::FinishReason) -> Self {
224        if let Ok(response) = &mut self.response {
225            response.finish_reason = Some(finish_reason);
226        }
227        self
228    }
229
230    /// Script the provider's own response for this turn — what a real seam
231    /// would serialize from its raw type. Attached to the response as-is, so
232    /// agent tests can prove the payload reaches every observer of the turn
233    /// without a live provider. A turn without a scripted payload carries
234    /// the scripted turn itself, serialized — the mock's own document, the
235    /// same capture every real adapter performs — so a scripted value in a
236    /// test is distinguishable from the mock's default by content.
237    pub fn with_raw(mut self, raw: serde_json::Value) -> Self {
238        if let Ok(response) = &mut self.response {
239            response.raw = Some(raw);
240        }
241        self
242    }
243
244    /// The provider document the mock attaches to this turn's response: the
245    /// scripted payload when one was supplied, otherwise the turn itself,
246    /// serialized. Public so a test can state the expected `raw` of a
247    /// recorded call without repeating the mock's serialization. An error
248    /// turn has no document.
249    pub fn raw(&self) -> Result<serde_json::Value, ProviderError> {
250        let response = self
251            .response
252            .as_ref()
253            .map_err(|error| error.clone().into_completion_error())?;
254        match &response.raw {
255            Some(raw) => Ok(raw.clone()),
256            None => Ok(serde_json::to_value(response)?),
257        }
258    }
259
260    fn into_completion_response(self) -> Result<CompletionResponse, ProviderError> {
261        let raw = self.raw()?;
262        let response = self.response.map_err(MockError::into_completion_error)?;
263        let mut completion =
264            CompletionResponse::new(response.choice, response.usage, MOCK_PROVIDER, raw)
265                .with_optional_finish_reason(response.finish_reason);
266        completion.message_id = response.message_id;
267        completion.response_id = response.response_id;
268        completion.provider_request_id = response.provider_request_id;
269        Ok(completion)
270    }
271}
272
273type MockInvocation = (CompletionRequest, Option<crate::observe::AdapterContext>);
274
275#[derive(Default)]
276struct MockScriptState {
277    turns: Mutex<VecDeque<MockTurn>>,
278    stream_turns: Mutex<VecDeque<Vec<MockStreamEvent>>>,
279    requests: Mutex<Vec<MockInvocation>>,
280}
281
282/// The scripted completion wire: its payload is the request, and each frame
283/// is a scripted step its decoder writes through the completion writer's
284/// part handles, as any wire's decoder does. A runtime that answers from a
285/// script ([`MockRuntime`], or a test's own) is its transport.
286#[derive(Clone, Debug, PartialEq, Eq)]
287pub struct MockScript {
288    name: String,
289    id: Option<String>,
290    capabilities: Capabilities,
291}
292
293impl MockScript {
294    /// A scripted wire named `name`, as records and telemetry name it, with
295    /// default capabilities.
296    pub fn new(name: impl Into<String>) -> Self {
297        Self {
298            name: name.into(),
299            id: None,
300            capabilities: Capabilities::default(),
301        }
302    }
303
304    /// The same wire, addressing the model `id`.
305    pub fn with_id(mut self, id: impl Into<String>) -> Self {
306        self.id = Some(id.into());
307        self
308    }
309
310    /// The same wire, reporting `capabilities`.
311    pub fn with_capabilities(mut self, capabilities: Capabilities) -> Self {
312        self.capabilities = capabilities;
313        self
314    }
315}
316
317impl Default for MockScript {
318    fn default() -> Self {
319        Self::new(MOCK_PROVIDER)
320    }
321}
322
323impl Wire for MockScript {
324    type Op = Completion;
325    type Payload = CompletionRequest;
326    type Frame = MockFrame;
327    type Decoder<'id> = MockDecoder<'id>;
328
329    fn describe(&self) -> Descriptor<'_> {
330        Descriptor::new(&self.name)
331            .model(self.id.as_deref())
332            .capabilities(self.capabilities)
333    }
334
335    /// A completion replays only the reasoning this runtime issued, as every
336    /// completion wire's encode scopes it.
337    fn encode(
338        &self,
339        request: CompletionRequest,
340        _mode: Mode,
341    ) -> Result<CompletionRequest, EncodeError> {
342        let issuers = [crate::message::Issuer::from(self.name.clone())];
343        let mut request = request.replayable_to(&issuers)?;
344        for message in request.chat_history.iter_mut() {
345            // `replayable_to` left only messages with a part to keep.
346            if let crate::message::Message::Assistant { content, .. } = message {
347                content.retain(|part| match part {
348                    AssistantContent::Reasoning(reasoning) => {
349                        reasoning.open_for(&issuers).is_some()
350                    }
351                    _ => true,
352                });
353            }
354        }
355        Ok(request)
356    }
357
358    fn decoder<'id>(&self) -> MockDecoder<'id> {
359        MockDecoder::default()
360    }
361}
362
363/// The scripted runtime behind [`MockCompletionModel`]: the transport of a
364/// [`MockScript`] wire named [`MOCK_PROVIDER`].
365///
366/// Each call consumes exactly one scripted turn. If no turn is available,
367/// the call fails with [`ProviderError::Provider`] and a clear message
368/// instead of repeating previous responses.
369#[derive(Clone, Default)]
370pub struct MockRuntime {
371    state: Arc<MockScriptState>,
372}
373
374/// A cloneable scripted completion model for tests: the scripted wire over
375/// its runtime. Clones share the script and the recorded requests.
376pub type MockCompletionModel = Model<MockScript, MockRuntime>;
377
378impl MockCompletionModel {
379    /// Create a mock model that returns one text completion.
380    pub fn text(text: impl Into<String>) -> Self {
381        Self::from_turns([MockTurn::text(text)])
382    }
383
384    /// Create a mock model from scripted non-streaming turns.
385    pub fn from_turns(turns: impl IntoIterator<Item = MockTurn>) -> Self {
386        Self::scripted(turns.into_iter().collect(), VecDeque::new())
387    }
388
389    /// Create a mock model from scripted streaming turns.
390    pub fn from_stream_turns(
391        stream_turns: impl IntoIterator<Item = impl IntoIterator<Item = MockStreamEvent>>,
392    ) -> Self {
393        Self::scripted(
394            VecDeque::new(),
395            stream_turns
396                .into_iter()
397                .map(|turn| turn.into_iter().collect())
398                .collect(),
399        )
400    }
401
402    fn scripted(turns: VecDeque<MockTurn>, stream_turns: VecDeque<Vec<MockStreamEvent>>) -> Self {
403        Model::new(
404            MockScript::default(),
405            MockRuntime {
406                state: Arc::new(MockScriptState {
407                    turns: Mutex::new(turns),
408                    stream_turns: Mutex::new(stream_turns),
409                    requests: Mutex::new(Vec::new()),
410                }),
411            },
412        )
413    }
414
415    /// Return cloned requests received by this model.
416    pub fn requests(&self) -> Vec<CompletionRequest> {
417        self.transport
418            .requests_guard()
419            .iter()
420            .map(|(request, _)| request.clone())
421            .collect()
422    }
423
424    /// Return invocation contexts in the same order as the captured requests.
425    pub fn contexts(&self) -> Vec<Option<crate::observe::AdapterContext>> {
426        self.transport
427            .requests_guard()
428            .iter()
429            .map(|(_, context)| context.clone())
430            .collect()
431    }
432
433    /// Return the number of requests received by this model.
434    pub fn request_count(&self) -> usize {
435        self.transport.requests_guard().len()
436    }
437
438    /// The non-streaming turns not yet consumed, in order — the read-back
439    /// half of the script, so a script is serde in and serde out.
440    pub fn script(&self) -> Vec<MockTurn> {
441        lock(&self.transport.state.turns).iter().cloned().collect()
442    }
443
444    /// The streaming turns not yet consumed, in order.
445    pub fn stream_script(&self) -> Vec<Vec<MockStreamEvent>> {
446        lock(&self.transport.state.stream_turns)
447            .iter()
448            .cloned()
449            .collect()
450    }
451}
452
453impl MockRuntime {
454    fn requests_guard(&self) -> MutexGuard<'_, Vec<MockInvocation>> {
455        lock(&self.state.requests)
456    }
457}
458
459fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
460    match mutex.lock() {
461        Ok(guard) => guard,
462        Err(poisoned) => poisoned.into_inner(),
463    }
464}
465
466impl Transport<MockScript> for MockRuntime {
467    fn send(&self, request: CompletionRequest, exchange: Exchange) -> Opening<MockFrame> {
468        let mode = exchange.mode;
469        self.requests_guard().push((request, exchange.observation));
470        match mode {
471            // A whole turn is one response, its document the response's
472            // `raw`; a scripted failure fails the reply.
473            Mode::Unary => {
474                let Some(turn) = lock(&self.state.turns).pop_front() else {
475                    return Opening::failed(ProviderError::Provider(
476                        "mock completion model has no scripted completion turn".to_string(),
477                    ));
478                };
479                match turn.into_completion_response() {
480                    Ok(response) => {
481                        let document = response.raw.clone();
482                        let request_id = response.provider_request_id.clone();
483                        Opening::ready(
484                            Opened::new(futures::stream::iter([Ok(MockFrame::Response(
485                                Box::new(response),
486                            ))]))
487                            .with_document(document)
488                            .with_request_id(request_id),
489                        )
490                    }
491                    Err(error) => Opening::ready(Opened::failed(error)),
492                }
493            }
494            Mode::Streaming => {
495                let Some(turn) = lock(&self.state.stream_turns).pop_front() else {
496                    return Opening::failed(ProviderError::Provider(
497                        "mock completion model has no scripted streaming turn".to_string(),
498                    ));
499                };
500                let request_id = turn.iter().find_map(|event| match event {
501                    MockStreamEvent::RequestId(id) => Some(id.clone()),
502                    _ => None,
503                });
504                Opening::ready(
505                    Opened::new(futures::stream::iter(
506                        turn.into_iter()
507                            .filter(|event| !matches!(event, MockStreamEvent::RequestId(_)))
508                            .map(|event| Ok(MockFrame::Event(event))),
509                    ))
510                    .with_request_id(request_id),
511                )
512            }
513        }
514    }
515}
516
517#[cfg(test)]
518mod tests;