pub mod error;
use crate::api::error::ApiError;
use crate::message::Message;
use crate::stream::StreamEvent;
use crate::tool::ToolSchema;
use futures::Stream;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
pub trait ApiClient: Send + Sync {
fn model(&self) -> String;
fn set_model(&self, _model: &str) -> bool {
false
}
fn stream_messages(
&self,
messages: Vec<Message>,
system: Option<String>,
tools: Option<Vec<ToolSchema>>,
) -> Pin<Box<dyn Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>>;
fn create_message(
&self,
messages: Vec<Message>,
system: Option<String>,
tools: Option<Vec<ToolSchema>>,
) -> Pin<Box<dyn Future<Output = Result<serde_json::Value, ApiError>> + Send + '_>>;
}
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,
_messages: Vec<Message>,
_system: Option<String>,
_tools: Option<Vec<ToolSchema>>,
) -> 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,
_messages: Vec<Message>,
_system: Option<String>,
_tools: Option<Vec<ToolSchema>>,
) -> Pin<Box<dyn Future<Output = Result<serde_json::Value, ApiError>> + Send + '_>>
{
Box::pin(async {
Ok(serde_json::json!({
"content": [{"type": "text", "text": "Hello!"}]
}))
})
}
}
#[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(vec![Message::user("Hi")], None, 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(vec![Message::user("Hi")], None, None)
.await;
assert!(result.is_ok());
let json = result.unwrap();
assert!(json.get("content").is_some());
}
#[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");
}
}