pub mod error;
use crate::message::Message;
use crate::stream::{StreamEvent, StreamStopReason, Usage};
use crate::tool::ToolSchema;
use error::ApiError;
use futures::Stream;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
#[derive(Debug, Clone)]
pub struct StreamRequest {
pub messages: Vec<Message>,
pub system: Option<String>,
pub tools: Option<Vec<ToolSchema>>,
}
impl StreamRequest {
#[must_use]
pub fn new(messages: Vec<Message>) -> Self {
Self {
messages,
system: None,
tools: None,
}
}
#[must_use]
pub fn with_system(mut self, system: Option<String>) -> Self {
self.system = system;
self
}
#[must_use]
pub fn with_tools(mut self, tools: Option<Vec<ToolSchema>>) -> Self {
self.tools = tools;
self
}
}
#[derive(Debug, Clone)]
pub struct NonStreamingResponse {
pub message: Message,
pub stop_reason: StreamStopReason,
pub usage: Option<Usage>,
}
pub trait ApiClient: Send + Sync {
fn model(&self) -> String;
fn set_model(&self, _model: &str) -> bool {
false
}
fn base_url(&self) -> String {
String::new()
}
fn stream_messages(
&self,
request: &StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>;
fn create_message(
&self,
request: &StreamRequest,
) -> Pin<Box<dyn Future<Output = Result<NonStreamingResponse, ApiError>> + Send + '_>>;
fn stream_messages_with_options(
&self,
request: &StreamRequest,
options: crate::structured::RequestOptions,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>> {
if let Some(err) = unsupported_options_error(&options) {
return Box::pin(futures::stream::once(async move { Err(err) }));
}
self.stream_messages(request)
}
fn create_message_with_options(
&self,
request: &StreamRequest,
options: crate::structured::RequestOptions,
) -> Pin<Box<dyn Future<Output = Result<NonStreamingResponse, ApiError>> + Send + '_>> {
if let Some(err) = unsupported_options_error(&options) {
return Box::pin(async move { Err(err) });
}
self.create_message(request)
}
fn extract_structured(&self, message: &Message) -> serde_json::Value {
if let Some((_, _, input)) = message.tool_call_parts().into_iter().next() {
return input.clone();
}
let text = message.text_content();
crate::structured::parse_json_lenient(&text).unwrap_or(serde_json::Value::String(text))
}
}
pub(crate) fn unsupported_options_error(
options: &crate::structured::RequestOptions,
) -> Option<ApiError> {
if let Some(err) = unsupported_structured_output_error(options) {
return Some(err);
}
if options.model.is_some() {
return Some(ApiError::config(
"this client does not support per-request model overrides (model)",
));
}
None
}
pub(crate) fn unsupported_structured_output_error(
options: &crate::structured::RequestOptions,
) -> Option<ApiError> {
if options.response_format.is_some() {
return Some(ApiError::config(
"this client does not support structured output (response_format)",
));
}
if !matches!(
options.tool_constraint,
crate::structured::ToolConstraint::None
) {
return Some(ApiError::config(
"this client does not support tool-call constraints (tool_constraint)",
));
}
None
}
pub type BoxedApiClient = Box<dyn ApiClient>;
pub type SharedApiClient = Arc<dyn ApiClient>;
#[cfg(test)]
mod tests {
use super::*;
use crate::stream::Usage;
use futures::StreamExt;
struct MockClient {
model_name: String,
}
impl MockClient {
fn new(model: &str) -> Self {
Self {
model_name: model.to_string(),
}
}
}
impl ApiClient for MockClient {
fn model(&self) -> String {
self.model_name.clone()
}
fn stream_messages(
&self,
_request: &StreamRequest,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>> {
let events: Vec<Result<StreamEvent, ApiError>> = vec![
Ok(StreamEvent::MessageStart(crate::stream::MessageStart {
message: crate::stream::MessageMetadata {
id: "msg_test".to_string(),
role: "assistant".to_string(),
model: self.model_name.clone(),
},
})),
Ok(StreamEvent::PartStart(crate::stream::PartStart {
index: 0,
part: Some(crate::message::MessagePart::text("Hello!")),
})),
Ok(StreamEvent::MessageDelta(crate::stream::MessageDelta {
delta: crate::stream::MessageDeltaPayload {
stop_reason: Some("end_turn".to_string()),
},
usage: Some(Usage::new(10, 5)),
})),
Ok(StreamEvent::MessageStop),
];
Box::pin(futures::stream::iter(events))
}
fn create_message(
&self,
_request: &StreamRequest,
) -> Pin<Box<dyn Future<Output = Result<NonStreamingResponse, ApiError>> + Send + '_>>
{
Box::pin(async {
Ok(NonStreamingResponse {
message: Message::assistant("Hello!"),
stop_reason: StreamStopReason::EndTurn,
usage: Some(Usage::default()),
})
})
}
}
#[test]
fn test_mock_client_model() {
let client = MockClient::new("test-model");
assert_eq!(client.model(), "test-model");
}
#[tokio::test]
async fn test_mock_client_stream() {
let client = MockClient::new("test-model");
let stream = client.stream_messages(&StreamRequest {
messages: vec![Message::user("Hi")],
system: None,
tools: None,
});
let events: Vec<_> = stream.collect().await;
assert_eq!(events.len(), 4);
assert!(matches!(
events[0].as_ref().unwrap(),
StreamEvent::MessageStart(_)
));
assert!(matches!(
events[1].as_ref().unwrap(),
StreamEvent::PartStart(_)
));
assert!(matches!(
events[2].as_ref().unwrap(),
StreamEvent::MessageDelta(_)
));
assert!(matches!(
events[3].as_ref().unwrap(),
StreamEvent::MessageStop
));
}
#[tokio::test]
async fn test_mock_client_create_message() {
let client = MockClient::new("test-model");
let result = client
.create_message(&StreamRequest {
messages: vec![Message::user("Hi")],
system: None,
tools: None,
})
.await;
assert!(result.is_ok());
let response = result.unwrap();
assert!(!response.message.parts.is_empty());
assert_eq!(response.stop_reason, StreamStopReason::EndTurn);
}
#[tokio::test]
async fn default_with_options_rejects_every_unsupported_field() {
let client = MockClient::new("primary-model");
let req = StreamRequest::new(vec![]);
let model_override =
crate::structured::RequestOptions::default().with_model("fallback-model");
let mut stream = client.stream_messages_with_options(&req, model_override.clone());
let first = stream.next().await;
drop(stream);
assert!(
matches!(&first, Some(Err(err)) if err.to_string().contains("model")),
"a model override the default impl cannot forward must fail loudly, got: {first:?}"
);
let result = client
.create_message_with_options(&req, model_override)
.await;
assert!(
result.is_err(),
"the non-streaming default must reject the same override"
);
let strict = crate::structured::RequestOptions::default()
.with_tool_constraint(crate::structured::ToolConstraint::Strict);
let mut stream = client.stream_messages_with_options(&req, strict.clone());
let first = stream.next().await;
drop(stream);
assert!(
matches!(&first, Some(Err(err)) if err.to_string().contains("tool_constraint")),
"a tool constraint the default impl cannot forward must fail loudly, got: {first:?}"
);
let result = client.create_message_with_options(&req, strict).await;
assert!(
result.is_err(),
"the non-streaming default must reject the same constraint"
);
let format = crate::structured::RequestOptions::default().with_response_format(
crate::structured::ResponseFormat::new("action", serde_json::json!({})),
);
let result = client.create_message_with_options(&req, format).await;
assert!(
result.is_err(),
"the non-streaming default must reject an unsupported response format"
);
let plain = crate::structured::RequestOptions::default();
let mut stream = client.stream_messages_with_options(&req, plain);
let first = stream.next().await;
drop(stream);
let delegated = first.is_some_and(|item| item.is_ok());
assert!(
delegated,
"default options must still delegate to the plain method"
);
}
#[cfg(feature = "testing")]
#[tokio::test]
async fn mock_client_serves_the_turn_under_the_routed_model() {
let mock = crate::testing::MockApiClient::new("primary").with_text_response("hi");
let opts = crate::structured::RequestOptions::default().with_model("fallback");
let mut stream =
mock.stream_messages_with_options(&StreamRequest::new(vec![Message::user("q")]), opts);
let first = futures::StreamExt::next(&mut stream).await;
match first {
Some(Ok(StreamEvent::MessageStart(start))) => assert_eq!(
start.message.model, "fallback",
"the mock — like the real providers — must serve the turn under the routed model"
),
other => panic!("expected a clean MessageStart under the override, got {other:?}"),
}
}
#[test]
fn stream_request_new_defaults() {
let req = StreamRequest::new(vec![Message::user("hi")]);
assert_eq!(req.messages.len(), 1);
assert!(req.system.is_none());
assert!(req.tools.is_none());
}
#[test]
fn stream_request_system_builder() {
let req = StreamRequest::new(vec![]).with_system(Some("be brief".to_string()));
assert_eq!(req.system.as_deref(), Some("be brief"));
let req = StreamRequest::new(vec![]).with_system(None);
assert!(req.system.is_none());
}
#[test]
fn stream_request_tools_builder() {
let tools = vec![crate::tool::ToolSchema {
tool: "search".into(),
description: "Search".into(),
input_schema: serde_json::json!({"type": "object"}),
}];
let req = StreamRequest::new(vec![]).with_tools(Some(tools));
assert_eq!(req.tools.as_ref().unwrap().len(), 1);
let req = StreamRequest::new(vec![]).with_tools(None);
assert!(req.tools.is_none());
}
#[test]
fn test_boxed_client() {
let client: BoxedApiClient = Box::new(MockClient::new("boxed"));
assert_eq!(client.model(), "boxed");
}
#[test]
fn test_shared_client() {
let client: SharedApiClient = Arc::new(MockClient::new("shared"));
assert_eq!(client.model(), "shared");
}
#[test]
fn default_set_model_returns_false() {
let client = MockClient::new("test-model");
assert!(!client.set_model("other-model"));
assert_eq!(client.model(), "test-model");
}
#[test]
fn extract_structured_default_returns_tool_call_input() {
let client = MockClient::new("m");
let message = Message::new(
crate::message::Role::Assistant,
vec![crate::message::MessagePart::tool_call(
"tc_1",
"search",
serde_json::json!({"q": "rust"}),
)],
);
let value = client.extract_structured(&message);
assert_eq!(value, serde_json::json!({"q": "rust"}));
}
#[test]
fn extract_structured_default_returns_first_tool_call_input() {
let client = MockClient::new("m");
let message = Message::new(
crate::message::Role::Assistant,
vec![
crate::message::MessagePart::tool_call(
"tc_1",
"first",
serde_json::json!({"order": 1}),
),
crate::message::MessagePart::tool_call(
"tc_2",
"second",
serde_json::json!({"order": 2}),
),
],
);
let value = client.extract_structured(&message);
assert_eq!(value, serde_json::json!({"order": 1}));
}
#[test]
fn extract_structured_default_parses_text_as_json() {
let client = MockClient::new("m");
let message = Message::assistant(r#"{"tool": "write", "args": {}}"#);
let value = client.extract_structured(&message);
assert_eq!(value["tool"], "write");
}
#[test]
fn extract_structured_default_lenient_parses_embedded_json() {
let client = MockClient::new("m");
let message = Message::assistant(r#"Here is the result: {"answer": 42}"#);
let value = client.extract_structured(&message);
assert_eq!(value["answer"], 42);
}
#[test]
fn extract_structured_default_prose_falls_back_to_string() {
let client = MockClient::new("m");
let message = Message::assistant("just prose, no json here");
let value = client.extract_structured(&message);
assert_eq!(value, serde_json::json!("just prose, no json here"));
}
#[test]
fn extract_structured_default_empty_message_falls_back_to_empty_string() {
let client = MockClient::new("m");
let message = Message::assistant("");
let value = client.extract_structured(&message);
assert_eq!(value, serde_json::json!(""));
}
#[test]
fn non_streaming_response_fields_are_accessible() {
let response = NonStreamingResponse {
message: Message::assistant("hello"),
stop_reason: StreamStopReason::EndTurn,
usage: Some(Usage::new(100, 50)),
};
assert_eq!(response.message.text_content(), "hello");
assert_eq!(response.stop_reason, StreamStopReason::EndTurn);
let usage = response.usage.expect("usage present");
assert_eq!(usage.input_tokens, 100);
assert_eq!(usage.output_tokens, 50);
assert_eq!(usage.total_tokens(), 150);
assert!(!response.message.parts.is_empty());
}
#[test]
fn non_streaming_response_usage_can_be_none() {
let response = NonStreamingResponse {
message: Message::assistant("hello"),
stop_reason: StreamStopReason::EndTurn,
usage: None,
};
assert!(response.usage.is_none());
}
#[tokio::test]
async fn default_with_options_rejects_unsupported_tool_constraint() {
struct PlainClient;
impl ApiClient for PlainClient {
fn model(&self) -> String {
"plain".to_string()
}
fn stream_messages(
&self,
_request: &StreamRequest,
) -> Pin<
Box<
dyn futures::Stream<Item = Result<StreamEvent, crate::api::error::ApiError>>
+ Send,
>,
> {
Box::pin(futures::stream::empty())
}
fn create_message(
&self,
_request: &StreamRequest,
) -> Pin<
Box<
dyn std::future::Future<
Output = Result<NonStreamingResponse, crate::api::error::ApiError>,
> + Send
+ '_,
>,
> {
Box::pin(async {
Ok(NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: StreamStopReason::EndTurn,
usage: None,
})
})
}
}
let options = crate::structured::RequestOptions::new()
.with_tool_constraint(crate::structured::ToolConstraint::Strict);
let mut stream =
PlainClient.stream_messages_with_options(&StreamRequest::new(vec![]), options);
let first = stream.next().await;
assert!(
matches!(&first, Some(Err(crate::api::error::ApiError::Config(_)))),
"doc: unsupported option fields yield an ApiError::config error as the first item; got {first:?}"
);
}
}