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, MockDocument, 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    response_id: Option<String>,
64    provider_request_id: Option<String>,
65    finish_reason: Option<crate::completion::FinishReason>,
66    /// A scripted provider document, when the test supplies one; otherwise
67    /// the turn itself, serialized, is the mock's document. Absent from the
68    /// serialized turn when unscripted, so that document never nests
69    /// itself; a scripted one survives a serde round trip of the script.
70    #[serde(
71        default,
72        skip_serializing_if = "Option::is_none",
73        deserialize_with = "deserialize_scripted_raw"
74    )]
75    raw: Option<serde_json::Value>,
76}
77
78fn deserialize_scripted_raw<'de, D: serde::Deserializer<'de>>(
79    deserializer: D,
80) -> Result<Option<serde_json::Value>, D::Error> {
81    // Only a missing field means unscripted; explicit JSON null is a document.
82    serde::Deserialize::deserialize(deserializer).map(Some)
83}
84
85impl MockTurn {
86    /// Create a text response turn.
87    pub fn text(text: impl Into<String>) -> Self {
88        Self::from_content(AssistantContent::text(text.into()))
89    }
90
91    /// Create a tool-call response turn. An empty `name` scripts a provider
92    /// error instead: no call can be built without a name.
93    pub fn tool_call(
94        id: impl Into<String>,
95        name: impl Into<String>,
96        arguments: serde_json::Value,
97    ) -> Self {
98        match crate::message::ToolName::new(name) {
99            Ok(name) => Self::from_content(AssistantContent::ToolCall(ToolCall::from_wire(
100                id,
101                ToolFunction::new(name, arguments),
102            ))),
103            Err(error) => Self::error(error.to_string()),
104        }
105    }
106
107    /// Create a provider-error response turn.
108    pub fn error(message: impl Into<String>) -> Self {
109        Self {
110            response: Err(MockError::provider(message)),
111        }
112    }
113
114    /// Create a provider-response error turn carrying a transport request id
115    /// (rig#2314): the scripted failure a test uses to assert error-identity
116    /// attribution.
117    pub fn provider_response_error(
118        status: http::StatusCode,
119        body: impl Into<String>,
120        request_id: impl Into<String>,
121    ) -> Self {
122        Self {
123            response: Err(MockError::ProviderResponse(
124                crate::provider_response::ProviderResponseError::new(status, body)
125                    .with_provider_request_id(Some(request_id.into())),
126            )),
127        }
128    }
129
130    /// Create a request-error response turn.
131    pub fn request_error(message: impl Into<String>) -> Self {
132        Self {
133            response: Err(MockError::request(message)),
134        }
135    }
136
137    /// Create a response turn from one assistant content item.
138    pub fn from_content(content: AssistantContent) -> Self {
139        Self {
140            response: Ok(MockTurnResponse {
141                choice: vec![content],
142                usage: Usage::default(),
143                response_id: None,
144                provider_request_id: None,
145                finish_reason: None,
146                raw: None,
147            }),
148        }
149    }
150
151    /// Create a response turn from assistant content items.
152    ///
153    /// Infallible now that content is a `Vec`: an empty turn is a shape a
154    /// provider can genuinely return, so it is a value to build, not an error.
155    pub fn from_contents(content: impl IntoIterator<Item = AssistantContent>) -> Self {
156        Self {
157            response: Ok(MockTurnResponse {
158                choice: content.into_iter().collect(),
159                usage: Usage::default(),
160                response_id: None,
161                provider_request_id: None,
162                finish_reason: None,
163                raw: None,
164            }),
165        }
166    }
167
168    /// Attach a provider-specific call ID to a tool-call response turn.
169    pub fn with_call_id(mut self, call_id: impl Into<String>) -> Self {
170        let call_id = call_id.into();
171        if let Ok(response) = &mut self.response {
172            for content in response.choice.iter_mut() {
173                if let AssistantContent::ToolCall(tool_call) = content {
174                    tool_call.id = crate::message::CallId::from_wire(call_id);
175                    break;
176                }
177            }
178        }
179        self
180    }
181
182    /// Override usage for this turn.
183    pub fn with_usage(mut self, usage: Usage) -> Self {
184        if let Ok(response) = &mut self.response {
185            response.usage = usage;
186        }
187        self
188    }
189
190    /// Set a provider-assigned response-scoped ID for this turn.
191    pub fn with_response_id(mut self, response_id: impl Into<String>) -> Self {
192        if let Ok(response) = &mut self.response {
193            response.response_id = Some(response_id.into());
194        }
195        self
196    }
197
198    /// Set a provider transport request id for this turn.
199    pub fn with_provider_request_id(mut self, request_id: impl Into<String>) -> Self {
200        if let Ok(response) = &mut self.response {
201            response.provider_request_id = Some(request_id.into());
202        }
203        self
204    }
205
206    /// Set the terminal finish reason for this turn.
207    ///
208    /// Without this, a mocked blocking turn always reports `None`, which
209    /// leaves the whole blocking half of the truncation contract (rig#2322)
210    /// unexercisable — the streamed mock could script a reason and the
211    /// blocking one could not.
212    pub fn with_finish_reason(mut self, finish_reason: crate::completion::FinishReason) -> Self {
213        if let Ok(response) = &mut self.response {
214            response.finish_reason = Some(finish_reason);
215        }
216        self
217    }
218
219    /// Script the provider's own response for this turn — what a real seam
220    /// would serialize from its raw type. Attached to the response as-is, so
221    /// agent tests can prove the payload reaches every observer of the turn
222    /// without a live provider. A turn without a scripted payload carries
223    /// the scripted turn itself, serialized — the mock's own document, the
224    /// same capture every real adapter performs — so a scripted value in a
225    /// test is distinguishable from the mock's default by content.
226    pub fn with_raw(mut self, raw: serde_json::Value) -> Self {
227        if let Ok(response) = &mut self.response {
228            response.raw = Some(raw);
229        }
230        self
231    }
232
233    /// The provider document the mock attaches to this turn's response: the
234    /// scripted payload when one was supplied, otherwise the turn itself,
235    /// serialized. Public so a test can state the expected `raw` of a
236    /// recorded call without repeating the mock's serialization. An error
237    /// turn has no document.
238    pub fn raw(&self) -> Result<serde_json::Value, ProviderError> {
239        let response = self
240            .response
241            .as_ref()
242            .map_err(|error| error.clone().into_completion_error())?;
243        match &response.raw {
244            Some(raw) => Ok(raw.clone()),
245            None => Ok(serde_json::to_value(response)?),
246        }
247    }
248
249    fn into_completion_response(self) -> Result<CompletionResponse, ProviderError> {
250        let raw = self.raw()?;
251        let response = self.response.map_err(MockError::into_completion_error)?;
252        let mut origin = crate::message::Origin::new(MOCK_API, MOCK_PROVIDER, "");
253        origin.response_id = response.response_id;
254        let mut completion = CompletionResponse::new(response.choice, response.usage, origin, raw)
255            .with_optional_finish_reason(response.finish_reason);
256        completion.provider_request_id = response.provider_request_id;
257        Ok(completion)
258    }
259}
260
261type MockInvocation = (CompletionRequest, Option<crate::observe::AdapterContext>);
262
263#[derive(Default)]
264struct MockScriptState {
265    turns: Mutex<VecDeque<MockTurn>>,
266    stream_turns: Mutex<VecDeque<Vec<MockStreamEvent>>>,
267    requests: Mutex<Vec<MockInvocation>>,
268}
269
270/// The scripted completion wire: its payload is the request, and each frame
271/// is a scripted step its decoder writes through the completion writer's
272/// part handles, as any wire's decoder does. A runtime that answers from a
273/// script ([`MockRuntime`], or a test's own) is its transport.
274#[derive(Clone, Debug, PartialEq, Eq)]
275pub struct MockScript {
276    name: String,
277    id: Option<String>,
278    capabilities: Capabilities,
279}
280
281impl MockScript {
282    /// A scripted wire named `name`, as records and telemetry name it, with
283    /// default capabilities.
284    pub fn new(name: impl Into<String>) -> Self {
285        Self {
286            name: name.into(),
287            id: None,
288            capabilities: Capabilities::default(),
289        }
290    }
291
292    /// The same wire, addressing the model `id`.
293    pub fn with_id(mut self, id: impl Into<String>) -> Self {
294        self.id = Some(id.into());
295        self
296    }
297
298    /// The same wire, reporting `capabilities`.
299    pub fn with_capabilities(mut self, capabilities: Capabilities) -> Self {
300        self.capabilities = capabilities;
301        self
302    }
303}
304
305/// The wire format the scripted wire names.
306pub const MOCK_API: crate::message::Api = crate::message::Api::from_static("mock.script");
307
308/// The model a scripted wire addresses when it names none.
309pub const MOCK_MODEL: &str = "mock-model";
310
311impl crate::completion::ReplayTarget for MockScript {
312    fn api(&self) -> crate::message::Api {
313        MOCK_API
314    }
315
316    /// A script takes every option and sends nothing for it, so a test
317    /// reads the options it set on the request the script records.
318    fn map_options(
319        &self,
320        _request: &crate::completion::CompletionRequest,
321        fields: crate::completion::options::OptionFields<'_>,
322    ) -> crate::completion::options::OptionMap {
323        use crate::completion::options::{Mapping, OptionFields, OptionMap};
324        let OptionFields {
325            reasoning,
326            cache,
327            service_tier,
328            verbosity,
329            parallel_tool_calls,
330            top_p,
331            seed,
332            stop,
333        } = fields;
334        let taken = |set: bool| match set {
335            true => Mapping::Omit("a scripted reply ignores options"),
336            false => Mapping::Nothing,
337        };
338        OptionMap {
339            reasoning: taken(reasoning.is_some()),
340            cache: taken(cache.is_some()),
341            service_tier: taken(service_tier.is_some()),
342            verbosity: taken(verbosity.is_some()),
343            parallel_tool_calls: taken(parallel_tool_calls.is_some()),
344            top_p: taken(top_p.is_some()),
345            seed: taken(seed.is_some()),
346            stop: taken(!stop.is_empty()),
347        }
348    }
349
350    // Scripts state a finish reason only when a test is about one.
351    fn states_finish_reason(&self) -> bool {
352        false
353    }
354
355    fn provider(&self) -> &str {
356        &self.name
357    }
358
359    /// The model the script addresses: its id, or [`MOCK_MODEL`].
360    fn model(&self) -> &str {
361        self.id.as_deref().unwrap_or(MOCK_MODEL)
362    }
363
364    fn accepts(&self, _model: &str) -> crate::completion::Accepts {
365        crate::completion::Accepts::ALL
366    }
367}
368
369impl Default for MockScript {
370    fn default() -> Self {
371        Self::new(MOCK_PROVIDER)
372    }
373}
374
375impl Wire for MockScript {
376    type Op = Completion;
377    type Payload = CompletionRequest;
378    type Frame = MockFrame;
379    type Decoder<'id> = MockDecoder<'id>;
380    type Reassembler = MockDocument;
381
382    fn describe(&self) -> Descriptor<'_> {
383        Descriptor::new(&self.name)
384            .model(self.id.as_deref())
385            .capabilities(self.capabilities)
386            .replay(self)
387    }
388
389    /// The request as the runtime received it, its history already shaped
390    /// for this wire.
391    fn encode(
392        &self,
393        request: CompletionRequest,
394        _mode: Mode,
395    ) -> Result<CompletionRequest, EncodeError> {
396        Ok(request)
397    }
398
399    fn decoder<'id>(&self) -> MockDecoder<'id> {
400        MockDecoder::default()
401    }
402}
403
404/// The scripted runtime behind [`MockCompletionModel`]: the transport of a
405/// [`MockScript`] wire named [`MOCK_PROVIDER`].
406///
407/// Each call consumes exactly one scripted turn. If no turn is available,
408/// the call fails with [`ProviderError::Provider`] and a clear message
409/// instead of repeating previous responses.
410#[derive(Clone, Default)]
411pub struct MockRuntime {
412    state: Arc<MockScriptState>,
413}
414
415/// A cloneable scripted completion model for tests: the scripted wire over
416/// its runtime. Clones share the script and the recorded requests.
417pub type MockCompletionModel = Model<MockScript, MockRuntime>;
418
419impl MockCompletionModel {
420    /// Create a mock model that returns one text completion.
421    pub fn text(text: impl Into<String>) -> Self {
422        Self::from_turns([MockTurn::text(text)])
423    }
424
425    /// Create a mock model from scripted non-streaming turns.
426    pub fn from_turns(turns: impl IntoIterator<Item = MockTurn>) -> Self {
427        Self::scripted(turns.into_iter().collect(), VecDeque::new())
428    }
429
430    /// Create a mock model from scripted streaming turns.
431    pub fn from_stream_turns(
432        stream_turns: impl IntoIterator<Item = impl IntoIterator<Item = MockStreamEvent>>,
433    ) -> Self {
434        Self::scripted(
435            VecDeque::new(),
436            stream_turns
437                .into_iter()
438                .map(|turn| turn.into_iter().collect())
439                .collect(),
440        )
441    }
442
443    fn scripted(turns: VecDeque<MockTurn>, stream_turns: VecDeque<Vec<MockStreamEvent>>) -> Self {
444        Model::new(
445            MockScript::default(),
446            MockRuntime {
447                state: Arc::new(MockScriptState {
448                    turns: Mutex::new(turns),
449                    stream_turns: Mutex::new(stream_turns),
450                    requests: Mutex::new(Vec::new()),
451                }),
452            },
453        )
454    }
455
456    /// Return cloned requests received by this model.
457    pub fn requests(&self) -> Vec<CompletionRequest> {
458        self.transport
459            .requests_guard()
460            .iter()
461            .map(|(request, _)| request.clone())
462            .collect()
463    }
464
465    /// Return invocation contexts in the same order as the captured requests.
466    pub fn contexts(&self) -> Vec<Option<crate::observe::AdapterContext>> {
467        self.transport
468            .requests_guard()
469            .iter()
470            .map(|(_, context)| context.clone())
471            .collect()
472    }
473
474    /// Return the number of requests received by this model.
475    pub fn request_count(&self) -> usize {
476        self.transport.requests_guard().len()
477    }
478
479    /// The non-streaming turns not yet consumed, in order — the read-back
480    /// half of the script, so a script is serde in and serde out.
481    pub fn script(&self) -> Vec<MockTurn> {
482        lock(&self.transport.state.turns).iter().cloned().collect()
483    }
484
485    /// The streaming turns not yet consumed, in order.
486    pub fn stream_script(&self) -> Vec<Vec<MockStreamEvent>> {
487        lock(&self.transport.state.stream_turns)
488            .iter()
489            .cloned()
490            .collect()
491    }
492}
493
494impl MockRuntime {
495    fn requests_guard(&self) -> MutexGuard<'_, Vec<MockInvocation>> {
496        lock(&self.state.requests)
497    }
498}
499
500fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
501    match mutex.lock() {
502        Ok(guard) => guard,
503        Err(poisoned) => poisoned.into_inner(),
504    }
505}
506
507impl Transport<MockScript> for MockRuntime {
508    fn send(&self, request: CompletionRequest, exchange: Exchange) -> Opening<MockFrame> {
509        let mode = exchange.mode;
510        self.requests_guard().push((request, exchange.observation));
511        match mode {
512            // A whole turn is one response, its document the response's
513            // `raw`; a scripted failure fails the reply.
514            Mode::Unary => {
515                let Some(turn) = lock(&self.state.turns).pop_front() else {
516                    return Opening::failed(ProviderError::Provider(
517                        "mock completion model has no scripted completion turn".to_string(),
518                    ));
519                };
520                match turn.into_completion_response() {
521                    Ok(response) => {
522                        let document = response.raw.clone();
523                        let request_id = response.provider_request_id.clone();
524                        Opening::ready(
525                            Opened::new(futures::stream::iter([Ok(MockFrame::Response(
526                                Box::new(response),
527                            ))]))
528                            .with_document(document)
529                            .with_request_id(request_id),
530                        )
531                    }
532                    Err(error) => Opening::ready(Opened::failed(error)),
533                }
534            }
535            Mode::Streaming => {
536                let Some(turn) = lock(&self.state.stream_turns).pop_front() else {
537                    return Opening::failed(ProviderError::Provider(
538                        "mock completion model has no scripted streaming turn".to_string(),
539                    ));
540                };
541                let request_id = turn.iter().find_map(|event| match event {
542                    MockStreamEvent::RequestId(id) => Some(id.clone()),
543                    _ => None,
544                });
545                Opening::ready(
546                    Opened::new(futures::stream::iter(
547                        turn.into_iter()
548                            .filter(|event| !matches!(event, MockStreamEvent::RequestId(_)))
549                            .map(|event| Ok(MockFrame::Event(event))),
550                    ))
551                    .with_request_id(request_id),
552                )
553            }
554        }
555    }
556}
557
558#[cfg(test)]
559mod tests;
560
561/// Every set option refused, for a test target that takes none.
562pub fn refuse_options(
563    fields: crate::completion::options::OptionFields<'_>,
564) -> crate::completion::options::OptionMap {
565    use crate::completion::options::{Mapping, OptionFields, OptionMap};
566    let OptionFields {
567        reasoning,
568        cache,
569        service_tier,
570        verbosity,
571        parallel_tool_calls,
572        top_p,
573        seed,
574        stop,
575    } = fields;
576    let refused = |set: bool| match set {
577        true => Mapping::unsupported("the test target takes no options"),
578        false => Mapping::Nothing,
579    };
580    OptionMap {
581        reasoning: refused(reasoning.is_some()),
582        cache: refused(cache.is_some()),
583        service_tier: refused(service_tier.is_some()),
584        verbosity: refused(verbosity.is_some()),
585        parallel_tool_calls: refused(parallel_tool_calls.is_some()),
586        top_p: refused(top_p.is_some()),
587        seed: refused(seed.is_some()),
588        stop: refused(!stop.is_empty()),
589    }
590}