use std::collections::HashMap;
use crate::completion::{CompletionResponse, Usage};
use crate::error::ProviderError;
use crate::operation::{Block, CallFragment, Completion, Finish};
use crate::wire::{Decoder, Flow, Out, WireEvent};
pub const MOCK_PROVIDER: &str = "mock";
pub fn mock_final(usage: Usage) -> Finish {
Finish {
usage,
..Finish::default()
}
}
fn fixture_item(value: serde_json::Value) -> Result<Option<serde_json::Value>, ProviderError> {
match value {
serde_json::Value::Null => Ok(None),
serde_json::Value::Object(map) if map.is_empty() => Ok(None),
serde_json::Value::Object(map) => Ok(Some(serde_json::Value::Object(map))),
other => Err(ProviderError::Provider(format!(
"mock stream fixture provider item must be a JSON object, got: {other}"
))),
}
}
pub fn mock_final_with_total_tokens(total_tokens: u64) -> Finish {
mock_final(Usage {
total_tokens: Some(total_tokens),
..Default::default()
})
}
#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
pub enum MockStreamEvent {
Text(String),
TextStart {
id: String,
additional_params: Option<serde_json::Value>,
},
TextAdditionalParams(serde_json::Value),
ToolCall {
id: String,
name: String,
arguments: serde_json::Value,
call_id: Option<String>,
},
ToolCallNameDelta { id: String, name: String },
ToolCallArgumentsDelta { id: String, arguments: String },
ToolCallEnd { id: String },
Reasoning { id: String, text: String },
ReasoningDelta { id: String, reasoning: String },
Unknown(serde_json::Value),
RequestId(String),
FinalResponse(Finish),
Error(MockError),
}
use super::completion::MockError;
fn fixture_provider_id(id: &str) -> Option<&str> {
let unnamed = ["reasoning-", "block-", "output-", "tool-", "text-"]
.iter()
.any(|namespace| {
id.strip_prefix(namespace)
.is_some_and(|rest| rest.parse::<u64>().is_ok())
});
(!id.is_empty() && !unnamed).then_some(id)
}
impl MockStreamEvent {
pub fn text(text: impl Into<String>) -> Self {
Self::Text(text.into())
}
pub fn text_start(id: impl Into<String>, additional_params: Option<serde_json::Value>) -> Self {
Self::TextStart {
id: id.into(),
additional_params,
}
}
pub fn text_additional_params(additional_params: serde_json::Value) -> Self {
Self::TextAdditionalParams(additional_params)
}
pub fn tool_call(
id: impl Into<String>,
name: impl Into<String>,
arguments: serde_json::Value,
) -> Self {
Self::ToolCall {
id: id.into(),
name: name.into(),
arguments,
call_id: None,
}
}
pub fn with_call_id(mut self, call_id: impl Into<String>) -> Self {
if let Self::ToolCall { call_id: id, .. } = &mut self {
*id = Some(call_id.into());
}
self
}
pub fn tool_call_name_delta(id: impl Into<String>, name: impl Into<String>) -> Self {
Self::ToolCallNameDelta {
id: id.into(),
name: name.into(),
}
}
pub fn tool_call_arguments_delta(id: impl Into<String>, arguments: impl Into<String>) -> Self {
Self::ToolCallArgumentsDelta {
id: id.into(),
arguments: arguments.into(),
}
}
pub fn tool_call_end(id: impl Into<String>) -> Self {
Self::ToolCallEnd { id: id.into() }
}
pub fn reasoning(reasoning: impl Into<String>) -> Self {
Self::Reasoning {
id: "reasoning-0".to_string(),
text: reasoning.into(),
}
}
pub fn with_reasoning_id(mut self, reasoning_id: impl Into<String>) -> Self {
if let Self::Reasoning { id, .. } = &mut self {
*id = reasoning_id.into();
}
self
}
pub fn reasoning_delta(reasoning: impl Into<String>) -> Self {
Self::reasoning_delta_with_id("reasoning-0", reasoning)
}
pub fn reasoning_delta_with_id(id: impl Into<String>, reasoning: impl Into<String>) -> Self {
Self::ReasoningDelta {
id: id.into(),
reasoning: reasoning.into(),
}
}
pub fn unknown(value: serde_json::Value) -> Self {
Self::Unknown(value)
}
pub fn final_response(usage: Usage) -> Self {
Self::FinalResponse(mock_final(usage))
}
pub fn final_response_with_default_usage() -> Self {
Self::FinalResponse(mock_final(Usage::default()))
}
pub fn final_response_with_total_tokens(total_tokens: u64) -> Self {
Self::FinalResponse(mock_final_with_total_tokens(total_tokens))
}
pub fn error(message: impl Into<String>) -> Self {
Self::Error(MockError::provider(message))
}
}
#[derive(Clone, Debug)]
pub enum MockFrame {
Event(MockStreamEvent),
Response(Box<CompletionResponse>),
}
#[derive(Debug, Default)]
pub struct MockDocument {
document: Option<serde_json::Value>,
failed: bool,
}
impl crate::wire::document::Serves<crate::operation::Completion> for MockDocument {}
impl crate::wire::document::Reassemble<MockFrame> for MockDocument {
fn absorb(&mut self, frame: &MockFrame) {
if self.document.is_some() || self.failed {
return;
}
match frame {
MockFrame::Response(response) => self.document = Some(response.raw.clone()),
MockFrame::Event(MockStreamEvent::FinalResponse(finish)) => {
self.document = serde_json::to_value(finish).ok();
}
MockFrame::Event(MockStreamEvent::Error(_)) => self.failed = true,
MockFrame::Event(_) => {}
}
}
fn finish(self) -> serde_json::Value {
self.document.unwrap_or(serde_json::Value::Null)
}
}
#[derive(Default)]
pub struct MockDecoder<'id> {
reasoning: Vec<(String, usize)>,
calls: HashMap<String, usize>,
brand: std::marker::PhantomData<fn(&'id ()) -> &'id ()>,
}
fn id_item(id: &str) -> serde_json::Value {
fixture_provider_id(id).map_or(
serde_json::Value::Null,
|id| serde_json::json!({ "id": id }),
)
}
impl<'id> MockDecoder<'id> {
fn call_index(&mut self, out: &mut Out<'id, Completion>, id: &str) -> usize {
if let Some(index) = self.calls.get(id) {
return *index;
}
let index = out.fresh_index();
if !id.is_empty() {
self.calls.insert(id.to_owned(), index);
}
index
}
fn event(
&mut self,
event: MockStreamEvent,
mut out: Out<'id, Completion>,
) -> Result<Flow, ProviderError> {
match event {
MockStreamEvent::Text(text) => {
out.run(Block::Text, &text)?;
}
MockStreamEvent::TextStart {
id: _,
additional_params,
} => {
out.end_run()?;
let index = out.run(Block::Text, "")?;
if let Some(item) = additional_params.map(fixture_item).transpose()?.flatten() {
out.edit(index, |slot| *slot = item)?;
}
}
MockStreamEvent::TextAdditionalParams(additional_params) => {
let Some(serde_json::Value::Object(fields)) = fixture_item(additional_params)?
else {
return Err(ProviderError::Provider(
"mock stream fixture `TextAdditionalParams` carries no data — \
drop the event instead"
.to_string(),
));
};
let index = out.run(Block::Text, "")?;
out.edit(index, |item| {
if !item.is_object() {
*item = serde_json::Value::Object(serde_json::Map::new());
}
if let Some(item) = item.as_object_mut() {
item.extend(fields);
}
})?;
}
MockStreamEvent::ToolCall {
id,
name,
arguments,
call_id,
} => {
out.end_run()?;
let index = match self.calls.remove(&id) {
Some(index) => index,
None => out.fresh_index(),
};
out.fragment(
Some(index),
CallFragment {
id: call_id.as_deref().or(fixture_provider_id(&id)),
name: Some(name.as_str()),
..CallFragment::default()
},
)?;
out.announce(index, arguments)?;
out.finish(index)?;
}
MockStreamEvent::ToolCallNameDelta { id, name } => {
out.end_run()?;
let index = self.call_index(&mut out, &id);
out.fragment(
Some(index),
CallFragment {
id: fixture_provider_id(&id),
name: Some(name.as_str()),
..CallFragment::default()
},
)?;
}
MockStreamEvent::ToolCallArgumentsDelta { id, arguments } => {
out.end_run()?;
let index = self.call_index(&mut out, &id);
out.fragment(
Some(index),
CallFragment {
id: fixture_provider_id(&id),
arguments: Some(arguments.as_str()),
..CallFragment::default()
},
)?;
}
MockStreamEvent::ToolCallEnd { id } => {
let index = self.call_index(&mut out, &id);
self.calls.remove(&id);
out.finish(index)?;
}
MockStreamEvent::Reasoning { id, text } => {
out.end_run()?;
match self.reasoning.iter().position(|(open, _)| *open == id) {
Some(at) => {
let (_, index) = self.reasoning.remove(at);
out.finish(index)?;
}
None => {
let index = out.fresh_index();
out.whole(
index,
Block::Reasoning { redacted: false },
id_item(&id),
&text,
)?;
}
}
}
MockStreamEvent::ReasoningDelta { id, reasoning } => {
out.end_run()?;
let index = match self.reasoning.iter().find(|(open, _)| *open == id) {
Some((_, index)) => *index,
None => {
let index = out.fresh_index();
out.open(index, Block::Reasoning { redacted: false }, id_item(&id))?;
self.reasoning.push((id, index));
index
}
};
out.push(index, &reasoning)?;
}
MockStreamEvent::Unknown(value) => out.unknown(value.into()),
MockStreamEvent::RequestId(_) => {}
MockStreamEvent::FinalResponse(finish) => {
out.end_run()?;
for (_, index) in std::mem::take(&mut self.reasoning) {
out.finish(index)?;
}
return Ok(out.end(finish));
}
MockStreamEvent::Error(error) => return Err(error.into_completion_error()),
}
Ok(Flow::More)
}
fn response(
&mut self,
response: CompletionResponse,
mut out: Out<'id, Completion>,
) -> Result<Flow, ProviderError> {
for content in response.choice.iter().cloned() {
out.content(content)?;
}
Ok(out.end(Finish {
usage: response.usage,
reason: response.finish_reason(),
response_id: response.response_id().map(str::to_owned),
model: response.model().map(str::to_owned),
error: response.error.clone(),
}))
}
}
impl<'id> Decoder<'id, Completion, MockFrame> for MockDecoder<'id> {
type Event = MockFrame;
fn classify(&self, frame: MockFrame) -> WireEvent<MockFrame> {
WireEvent::Known(frame)
}
fn decode(
&mut self,
frame: MockFrame,
out: Out<'id, Completion>,
) -> Result<Flow, ProviderError> {
match frame {
MockFrame::Event(event) => self.event(event, out),
MockFrame::Response(response) => self.response(*response, out),
}
}
}