1use 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
12pub const MOCK_PROVIDER: &str = "mock";
14
15pub fn mock_final(usage: Usage) -> Finish {
17 Finish {
18 usage,
19 ..Finish::default()
20 }
21}
22
23fn 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
37pub 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#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
47pub enum MockStreamEvent {
48 Text(String),
50 TextStart {
52 id: String,
53 additional_params: Option<serde_json::Value>,
54 },
55 TextAdditionalParams(serde_json::Value),
57 ToolCall {
59 id: String,
60 name: String,
61 arguments: serde_json::Value,
62 call_id: Option<String>,
63 },
64 ToolCallNameDelta { id: String, name: String },
66 ToolCallArgumentsDelta { id: String, arguments: String },
68 ToolCallEnd { id: String },
71 Reasoning { id: String, text: String },
73 ReasoningDelta { id: String, reasoning: String },
75 Unknown(serde_json::Value),
77 RequestId(String),
82 FinalResponse(Finish),
84 Error(MockError),
86}
87
88use super::completion::MockError;
89
90fn 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 pub fn text(text: impl Into<String>) -> Self {
107 Self::Text(text.into())
108 }
109
110 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 pub fn text_additional_params(additional_params: serde_json::Value) -> Self {
120 Self::TextAdditionalParams(additional_params)
121 }
122
123 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 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 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 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 pub fn tool_call_end(id: impl Into<String>) -> Self {
163 Self::ToolCallEnd { id: id.into() }
164 }
165
166 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 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 pub fn reasoning_delta(reasoning: impl Into<String>) -> Self {
188 Self::reasoning_delta_with_id("reasoning-0", reasoning)
189 }
190
191 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 pub fn unknown(value: serde_json::Value) -> Self {
201 Self::Unknown(value)
202 }
203
204 pub fn final_response(usage: Usage) -> Self {
206 Self::FinalResponse(mock_final(usage))
207 }
208
209 pub fn final_response_with_default_usage() -> Self {
211 Self::FinalResponse(mock_final(Usage::default()))
212 }
213
214 pub fn final_response_with_total_tokens(total_tokens: u64) -> Self {
216 Self::FinalResponse(mock_final_with_total_tokens(total_tokens))
217 }
218
219 pub fn error(message: impl Into<String>) -> Self {
221 Self::Error(MockError::provider(message))
222 }
223}
224
225#[derive(Clone, Debug)]
227pub enum MockFrame {
228 Event(MockStreamEvent),
230 Response(Box<CompletionResponse>),
232}
233
234#[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#[derive(Default)]
268pub struct MockDecoder<'id> {
269 reasoning: Vec<(String, usize)>,
271 calls: HashMap<String, usize>,
273 brand: std::marker::PhantomData<fn(&'id ()) -> &'id ()>,
274}
275
276fn 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 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 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 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}