use std::{
collections::VecDeque,
sync::{Arc, Mutex, MutexGuard},
};
use crate::driver::{Exchange, Model, Opened, Opening, Transport};
use crate::error::{EncodeError, ProviderError};
use crate::operation::Completion;
use crate::wire::{Capabilities, Descriptor, Mode, Wire};
use crate::{
completion::{AssistantContent, CompletionRequest, CompletionResponse, Usage},
message::{ToolCall, ToolFunction},
};
use super::streaming::{MOCK_PROVIDER, MockDecoder, MockDocument, MockFrame, MockStreamEvent};
#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
pub enum MockError {
Provider(String),
Request(String),
ProviderResponse(crate::provider_response::ProviderResponseError),
}
impl MockError {
pub fn provider(message: impl Into<String>) -> Self {
Self::Provider(message.into())
}
pub fn request(message: impl Into<String>) -> Self {
Self::Request(message.into())
}
pub(crate) fn into_completion_error(self) -> ProviderError {
match self {
Self::Provider(message) => ProviderError::Provider(message),
Self::Request(message) => ProviderError::request(message),
Self::ProviderResponse(response) => ProviderError::ProviderResponse(response),
}
}
}
#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct MockTurn {
response: Result<MockTurnResponse, MockError>,
}
#[derive(Clone, Debug, PartialEq, serde::Serialize, serde::Deserialize)]
struct MockTurnResponse {
choice: Vec<AssistantContent>,
usage: Usage,
response_id: Option<String>,
provider_request_id: Option<String>,
finish_reason: Option<crate::completion::FinishReason>,
#[serde(
default,
skip_serializing_if = "Option::is_none",
deserialize_with = "deserialize_scripted_raw"
)]
raw: Option<serde_json::Value>,
}
fn deserialize_scripted_raw<'de, D: serde::Deserializer<'de>>(
deserializer: D,
) -> Result<Option<serde_json::Value>, D::Error> {
serde::Deserialize::deserialize(deserializer).map(Some)
}
impl MockTurn {
pub fn text(text: impl Into<String>) -> Self {
Self::from_content(AssistantContent::text(text.into()))
}
pub fn tool_call(
id: impl Into<String>,
name: impl Into<String>,
arguments: serde_json::Value,
) -> Self {
match crate::message::ToolName::new(name) {
Ok(name) => Self::from_content(AssistantContent::ToolCall(ToolCall::from_wire(
id,
ToolFunction::new(name, arguments),
))),
Err(error) => Self::error(error.to_string()),
}
}
pub fn error(message: impl Into<String>) -> Self {
Self {
response: Err(MockError::provider(message)),
}
}
pub fn provider_response_error(
status: http::StatusCode,
body: impl Into<String>,
request_id: impl Into<String>,
) -> Self {
Self {
response: Err(MockError::ProviderResponse(
crate::provider_response::ProviderResponseError::new(status, body)
.with_provider_request_id(Some(request_id.into())),
)),
}
}
pub fn request_error(message: impl Into<String>) -> Self {
Self {
response: Err(MockError::request(message)),
}
}
pub fn from_content(content: AssistantContent) -> Self {
Self {
response: Ok(MockTurnResponse {
choice: vec![content],
usage: Usage::default(),
response_id: None,
provider_request_id: None,
finish_reason: None,
raw: None,
}),
}
}
pub fn from_contents(content: impl IntoIterator<Item = AssistantContent>) -> Self {
Self {
response: Ok(MockTurnResponse {
choice: content.into_iter().collect(),
usage: Usage::default(),
response_id: None,
provider_request_id: None,
finish_reason: None,
raw: None,
}),
}
}
pub fn with_call_id(mut self, call_id: impl Into<String>) -> Self {
let call_id = call_id.into();
if let Ok(response) = &mut self.response {
for content in response.choice.iter_mut() {
if let AssistantContent::ToolCall(tool_call) = content {
tool_call.id = crate::message::CallId::from_wire(call_id);
break;
}
}
}
self
}
pub fn with_usage(mut self, usage: Usage) -> Self {
if let Ok(response) = &mut self.response {
response.usage = usage;
}
self
}
pub fn with_response_id(mut self, response_id: impl Into<String>) -> Self {
if let Ok(response) = &mut self.response {
response.response_id = Some(response_id.into());
}
self
}
pub fn with_provider_request_id(mut self, request_id: impl Into<String>) -> Self {
if let Ok(response) = &mut self.response {
response.provider_request_id = Some(request_id.into());
}
self
}
pub fn with_finish_reason(mut self, finish_reason: crate::completion::FinishReason) -> Self {
if let Ok(response) = &mut self.response {
response.finish_reason = Some(finish_reason);
}
self
}
pub fn with_raw(mut self, raw: serde_json::Value) -> Self {
if let Ok(response) = &mut self.response {
response.raw = Some(raw);
}
self
}
pub fn raw(&self) -> Result<serde_json::Value, ProviderError> {
let response = self
.response
.as_ref()
.map_err(|error| error.clone().into_completion_error())?;
match &response.raw {
Some(raw) => Ok(raw.clone()),
None => Ok(serde_json::to_value(response)?),
}
}
fn into_completion_response(self) -> Result<CompletionResponse, ProviderError> {
let raw = self.raw()?;
let response = self.response.map_err(MockError::into_completion_error)?;
let mut origin = crate::message::Origin::new(MOCK_API, MOCK_PROVIDER, "");
origin.response_id = response.response_id;
let mut completion = CompletionResponse::new(response.choice, response.usage, origin, raw)
.with_optional_finish_reason(response.finish_reason);
completion.provider_request_id = response.provider_request_id;
Ok(completion)
}
}
type MockInvocation = (CompletionRequest, Option<crate::observe::AdapterContext>);
#[derive(Default)]
struct MockScriptState {
turns: Mutex<VecDeque<MockTurn>>,
stream_turns: Mutex<VecDeque<Vec<MockStreamEvent>>>,
requests: Mutex<Vec<MockInvocation>>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct MockScript {
name: String,
id: Option<String>,
capabilities: Capabilities,
}
impl MockScript {
pub fn new(name: impl Into<String>) -> Self {
Self {
name: name.into(),
id: None,
capabilities: Capabilities::default(),
}
}
pub fn with_id(mut self, id: impl Into<String>) -> Self {
self.id = Some(id.into());
self
}
pub fn with_capabilities(mut self, capabilities: Capabilities) -> Self {
self.capabilities = capabilities;
self
}
}
pub const MOCK_API: crate::message::Api = crate::message::Api::from_static("mock.script");
pub const MOCK_MODEL: &str = "mock-model";
impl crate::completion::ReplayTarget for MockScript {
fn api(&self) -> crate::message::Api {
MOCK_API
}
fn map_options(
&self,
_request: &crate::completion::CompletionRequest,
fields: crate::completion::options::OptionFields<'_>,
) -> crate::completion::options::OptionMap {
use crate::completion::options::{Mapping, OptionFields, OptionMap};
let OptionFields {
reasoning,
cache,
service_tier,
verbosity,
parallel_tool_calls,
top_p,
seed,
stop,
} = fields;
let taken = |set: bool| match set {
true => Mapping::Omit("a scripted reply ignores options"),
false => Mapping::Nothing,
};
OptionMap {
reasoning: taken(reasoning.is_some()),
cache: taken(cache.is_some()),
service_tier: taken(service_tier.is_some()),
verbosity: taken(verbosity.is_some()),
parallel_tool_calls: taken(parallel_tool_calls.is_some()),
top_p: taken(top_p.is_some()),
seed: taken(seed.is_some()),
stop: taken(!stop.is_empty()),
}
}
fn states_finish_reason(&self) -> bool {
false
}
fn provider(&self) -> &str {
&self.name
}
fn model(&self) -> &str {
self.id.as_deref().unwrap_or(MOCK_MODEL)
}
fn accepts(&self, _model: &str) -> crate::completion::Accepts {
crate::completion::Accepts::ALL
}
}
impl Default for MockScript {
fn default() -> Self {
Self::new(MOCK_PROVIDER)
}
}
impl Wire for MockScript {
type Op = Completion;
type Payload = CompletionRequest;
type Frame = MockFrame;
type Decoder<'id> = MockDecoder<'id>;
type Reassembler = MockDocument;
fn describe(&self) -> Descriptor<'_> {
Descriptor::new(&self.name)
.model(self.id.as_deref())
.capabilities(self.capabilities)
.replay(self)
}
fn encode(
&self,
request: CompletionRequest,
_mode: Mode,
) -> Result<CompletionRequest, EncodeError> {
Ok(request)
}
fn decoder<'id>(&self) -> MockDecoder<'id> {
MockDecoder::default()
}
}
#[derive(Clone, Default)]
pub struct MockRuntime {
state: Arc<MockScriptState>,
}
pub type MockCompletionModel = Model<MockScript, MockRuntime>;
impl MockCompletionModel {
pub fn text(text: impl Into<String>) -> Self {
Self::from_turns([MockTurn::text(text)])
}
pub fn from_turns(turns: impl IntoIterator<Item = MockTurn>) -> Self {
Self::scripted(turns.into_iter().collect(), VecDeque::new())
}
pub fn from_stream_turns(
stream_turns: impl IntoIterator<Item = impl IntoIterator<Item = MockStreamEvent>>,
) -> Self {
Self::scripted(
VecDeque::new(),
stream_turns
.into_iter()
.map(|turn| turn.into_iter().collect())
.collect(),
)
}
fn scripted(turns: VecDeque<MockTurn>, stream_turns: VecDeque<Vec<MockStreamEvent>>) -> Self {
Model::new(
MockScript::default(),
MockRuntime {
state: Arc::new(MockScriptState {
turns: Mutex::new(turns),
stream_turns: Mutex::new(stream_turns),
requests: Mutex::new(Vec::new()),
}),
},
)
}
pub fn requests(&self) -> Vec<CompletionRequest> {
self.transport
.requests_guard()
.iter()
.map(|(request, _)| request.clone())
.collect()
}
pub fn contexts(&self) -> Vec<Option<crate::observe::AdapterContext>> {
self.transport
.requests_guard()
.iter()
.map(|(_, context)| context.clone())
.collect()
}
pub fn request_count(&self) -> usize {
self.transport.requests_guard().len()
}
pub fn script(&self) -> Vec<MockTurn> {
lock(&self.transport.state.turns).iter().cloned().collect()
}
pub fn stream_script(&self) -> Vec<Vec<MockStreamEvent>> {
lock(&self.transport.state.stream_turns)
.iter()
.cloned()
.collect()
}
}
impl MockRuntime {
fn requests_guard(&self) -> MutexGuard<'_, Vec<MockInvocation>> {
lock(&self.state.requests)
}
}
fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
match mutex.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
}
}
impl Transport<MockScript> for MockRuntime {
fn send(&self, request: CompletionRequest, exchange: Exchange) -> Opening<MockFrame> {
let mode = exchange.mode;
self.requests_guard().push((request, exchange.observation));
match mode {
Mode::Unary => {
let Some(turn) = lock(&self.state.turns).pop_front() else {
return Opening::failed(ProviderError::Provider(
"mock completion model has no scripted completion turn".to_string(),
));
};
match turn.into_completion_response() {
Ok(response) => {
let document = response.raw.clone();
let request_id = response.provider_request_id.clone();
Opening::ready(
Opened::new(futures::stream::iter([Ok(MockFrame::Response(
Box::new(response),
))]))
.with_document(document)
.with_request_id(request_id),
)
}
Err(error) => Opening::ready(Opened::failed(error)),
}
}
Mode::Streaming => {
let Some(turn) = lock(&self.state.stream_turns).pop_front() else {
return Opening::failed(ProviderError::Provider(
"mock completion model has no scripted streaming turn".to_string(),
));
};
let request_id = turn.iter().find_map(|event| match event {
MockStreamEvent::RequestId(id) => Some(id.clone()),
_ => None,
});
Opening::ready(
Opened::new(futures::stream::iter(
turn.into_iter()
.filter(|event| !matches!(event, MockStreamEvent::RequestId(_)))
.map(|event| Ok(MockFrame::Event(event))),
))
.with_request_id(request_id),
)
}
}
}
}
#[cfg(test)]
mod tests;
pub fn refuse_options(
fields: crate::completion::options::OptionFields<'_>,
) -> crate::completion::options::OptionMap {
use crate::completion::options::{Mapping, OptionFields, OptionMap};
let OptionFields {
reasoning,
cache,
service_tier,
verbosity,
parallel_tool_calls,
top_p,
seed,
stop,
} = fields;
let refused = |set: bool| match set {
true => Mapping::unsupported("the test target takes no options"),
false => Mapping::Nothing,
};
OptionMap {
reasoning: refused(reasoning.is_some()),
cache: refused(cache.is_some()),
service_tier: refused(service_tier.is_some()),
verbosity: refused(verbosity.is_some()),
parallel_tool_calls: refused(parallel_tool_calls.is_some()),
top_p: refused(top_p.is_some()),
seed: refused(seed.is_some()),
stop: refused(!stop.is_empty()),
}
}