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, 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 message_id: Option<String>,
64 response_id: Option<String>,
65 provider_request_id: Option<String>,
66 finish_reason: Option<crate::completion::FinishReason>,
67 #[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 serde::Deserialize::deserialize(deserializer).map(Some)
84}
85
86impl MockTurn {
87 pub fn text(text: impl Into<String>) -> Self {
89 Self::from_content(AssistantContent::text(text.into()))
90 }
91
92 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 pub fn error(message: impl Into<String>) -> Self {
110 Self {
111 response: Err(MockError::provider(message)),
112 }
113 }
114
115 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 pub fn request_error(message: impl Into<String>) -> Self {
133 Self {
134 response: Err(MockError::request(message)),
135 }
136 }
137
138 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 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 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 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 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 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 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 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 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 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#[derive(Clone, Debug, PartialEq, Eq)]
287pub struct MockScript {
288 name: String,
289 id: Option<String>,
290 capabilities: Capabilities,
291}
292
293impl MockScript {
294 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 pub fn with_id(mut self, id: impl Into<String>) -> Self {
306 self.id = Some(id.into());
307 self
308 }
309
310 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 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 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#[derive(Clone, Default)]
370pub struct MockRuntime {
371 state: Arc<MockScriptState>,
372}
373
374pub type MockCompletionModel = Model<MockScript, MockRuntime>;
377
378impl MockCompletionModel {
379 pub fn text(text: impl Into<String>) -> Self {
381 Self::from_turns([MockTurn::text(text)])
382 }
383
384 pub fn from_turns(turns: impl IntoIterator<Item = MockTurn>) -> Self {
386 Self::scripted(turns.into_iter().collect(), VecDeque::new())
387 }
388
389 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 pub fn requests(&self) -> Vec<CompletionRequest> {
417 self.transport
418 .requests_guard()
419 .iter()
420 .map(|(request, _)| request.clone())
421 .collect()
422 }
423
424 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 pub fn request_count(&self) -> usize {
435 self.transport.requests_guard().len()
436 }
437
438 pub fn script(&self) -> Vec<MockTurn> {
441 lock(&self.transport.state.turns).iter().cloned().collect()
442 }
443
444 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 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;