1use 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#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
21pub enum MockError {
22 Provider(String),
24 Request(String),
26 ProviderResponse(crate::provider_response::ProviderResponseError),
28}
29
30impl MockError {
31 pub fn provider(message: impl Into<String>) -> Self {
33 Self::Provider(message.into())
34 }
35
36 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#[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 #[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 serde::Deserialize::deserialize(deserializer).map(Some)
83}
84
85impl MockTurn {
86 pub fn text(text: impl Into<String>) -> Self {
88 Self::from_content(AssistantContent::text(text.into()))
89 }
90
91 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 pub fn error(message: impl Into<String>) -> Self {
109 Self {
110 response: Err(MockError::provider(message)),
111 }
112 }
113
114 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 pub fn request_error(message: impl Into<String>) -> Self {
132 Self {
133 response: Err(MockError::request(message)),
134 }
135 }
136
137 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 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 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 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 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 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 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 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 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#[derive(Clone, Debug, PartialEq, Eq)]
275pub struct MockScript {
276 name: String,
277 id: Option<String>,
278 capabilities: Capabilities,
279}
280
281impl MockScript {
282 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 pub fn with_id(mut self, id: impl Into<String>) -> Self {
294 self.id = Some(id.into());
295 self
296 }
297
298 pub fn with_capabilities(mut self, capabilities: Capabilities) -> Self {
300 self.capabilities = capabilities;
301 self
302 }
303}
304
305pub const MOCK_API: crate::message::Api = crate::message::Api::from_static("mock.script");
307
308pub 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 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 fn states_finish_reason(&self) -> bool {
352 false
353 }
354
355 fn provider(&self) -> &str {
356 &self.name
357 }
358
359 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 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#[derive(Clone, Default)]
411pub struct MockRuntime {
412 state: Arc<MockScriptState>,
413}
414
415pub type MockCompletionModel = Model<MockScript, MockRuntime>;
418
419impl MockCompletionModel {
420 pub fn text(text: impl Into<String>) -> Self {
422 Self::from_turns([MockTurn::text(text)])
423 }
424
425 pub fn from_turns(turns: impl IntoIterator<Item = MockTurn>) -> Self {
427 Self::scripted(turns.into_iter().collect(), VecDeque::new())
428 }
429
430 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 pub fn requests(&self) -> Vec<CompletionRequest> {
458 self.transport
459 .requests_guard()
460 .iter()
461 .map(|(request, _)| request.clone())
462 .collect()
463 }
464
465 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 pub fn request_count(&self) -> usize {
476 self.transport.requests_guard().len()
477 }
478
479 pub fn script(&self) -> Vec<MockTurn> {
482 lock(&self.transport.state.turns).iter().cloned().collect()
483 }
484
485 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 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
561pub 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}