#[cfg(feature = "audio")]
use super::audio_generation::AudioGenerationModel;
use super::completion::CompletionModel;
use super::embedding::{
EmbeddingModel, TEXT_EMBEDDING_3_LARGE, TEXT_EMBEDDING_3_SMALL, TEXT_EMBEDDING_ADA_002,
};
#[cfg(feature = "image")]
use super::image_generation::ImageGenerationModel;
use super::transcription::TranscriptionModel;
use crate::client::{CompletionClient, EmbeddingsClient, ProviderClient, TranscriptionClient};
#[cfg(feature = "audio")]
use crate::client::AudioGenerationClient;
#[cfg(feature = "image")]
use crate::client::ImageGenerationClient;
use serde::Deserialize;
const OPENAI_API_BASE_URL: &str = "https://api.openai.com/v1";
#[derive(Clone, Debug)]
pub struct Client {
base_url: String,
http_client: reqwest::Client,
}
impl Client {
pub fn new(api_key: &str) -> Self {
Self::from_url(api_key, OPENAI_API_BASE_URL)
}
pub fn from_url(api_key: &str, base_url: &str) -> Self {
Self {
base_url: base_url.to_string(),
http_client: reqwest::Client::builder()
.default_headers({
let mut headers = reqwest::header::HeaderMap::new();
headers.insert(
"Authorization",
format!("Bearer {api_key}")
.parse()
.expect("Bearer token should parse"),
);
headers
})
.build()
.expect("OpenAI reqwest client should build"),
}
}
pub(crate) fn post(&self, path: &str) -> reqwest::RequestBuilder {
let url = format!("{}/{}", self.base_url, path).replace("//", "/");
self.http_client.post(url)
}
}
impl ProviderClient for Client {
fn from_env() -> Self {
let api_key = std::env::var("OPENAI_API_KEY").expect("OPENAI_API_KEY not set");
Self::new(&api_key)
}
}
impl CompletionClient for Client {
type CompletionModel = CompletionModel;
fn completion_model(&self, model: &str) -> CompletionModel {
CompletionModel::new(self.clone(), model)
}
}
impl EmbeddingsClient for Client {
type EmbeddingModel = EmbeddingModel;
fn embedding_model(&self, model: &str) -> Self::EmbeddingModel {
let ndims = match model {
TEXT_EMBEDDING_3_LARGE => 3072,
TEXT_EMBEDDING_3_SMALL | TEXT_EMBEDDING_ADA_002 => 1536,
_ => 0,
};
EmbeddingModel::new(self.clone(), model, ndims)
}
fn embedding_model_with_ndims(&self, model: &str, ndims: usize) -> Self::EmbeddingModel {
EmbeddingModel::new(self.clone(), model, ndims)
}
}
impl TranscriptionClient for Client {
type TranscriptionModel = TranscriptionModel;
fn transcription_model(&self, model: &str) -> TranscriptionModel {
TranscriptionModel::new(self.clone(), model)
}
}
#[cfg(feature = "image")]
impl ImageGenerationClient for Client {
type ImageGenerationModel = ImageGenerationModel;
fn image_generation_model(&self, model: &str) -> Self::ImageGenerationModel {
ImageGenerationModel::new(self.clone(), model)
}
}
#[cfg(feature = "audio")]
impl AudioGenerationClient for Client {
type AudioGenerationModel = AudioGenerationModel;
fn audio_generation_model(&self, model: &str) -> Self::AudioGenerationModel {
AudioGenerationModel::new(self.clone(), model)
}
}
#[derive(Clone, Debug, Deserialize)]
pub struct Usage {
pub prompt_tokens: usize,
pub total_tokens: usize,
}
impl std::fmt::Display for Usage {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"Prompt tokens: {} Total tokens: {}",
self.prompt_tokens, self.total_tokens
)
}
}
#[derive(Debug, Deserialize)]
pub struct ApiErrorResponse {
pub(crate) message: String,
}
#[derive(Debug, Deserialize)]
#[serde(untagged)]
pub(crate) enum ApiResponse<T> {
Ok(T),
Err(ApiErrorResponse),
}
#[cfg(test)]
mod tests {
use crate::message::ImageDetail;
use crate::providers::openai::{
AssistantContent, Function, ImageUrl, Message, ToolCall, ToolType, UserContent,
};
use crate::{message, OneOrMany};
use serde_path_to_error::deserialize;
#[test]
fn test_deserialize_message() {
let assistant_message_json = r#"
{
"role": "assistant",
"content": "\n\nHello there, how may I assist you today?"
}
"#;
let assistant_message_json2 = r#"
{
"role": "assistant",
"content": [
{
"type": "text",
"text": "\n\nHello there, how may I assist you today?"
}
],
"tool_calls": null
}
"#;
let assistant_message_json3 = r#"
{
"role": "assistant",
"tool_calls": [
{
"id": "call_h89ipqYUjEpCPI6SxspMnoUU",
"type": "function",
"function": {
"name": "subtract",
"arguments": "{\"x\": 2, \"y\": 5}"
}
}
],
"content": null,
"refusal": null
}
"#;
let user_message_json = r#"
{
"role": "user",
"content": [
{
"type": "text",
"text": "What's in this image?"
},
{
"type": "image_url",
"image_url": {
"url": "https://upload.wikimedia.org/wikipedia/commons/thumb/d/dd/Gfp-wisconsin-madison-the-nature-boardwalk.jpg/2560px-Gfp-wisconsin-madison-the-nature-boardwalk.jpg"
}
},
{
"type": "audio",
"input_audio": {
"data": "...",
"format": "mp3"
}
}
]
}
"#;
let assistant_message: Message = {
let jd = &mut serde_json::Deserializer::from_str(assistant_message_json);
deserialize(jd).unwrap_or_else(|err| {
panic!(
"Deserialization error at {} ({}:{}): {}",
err.path(),
err.inner().line(),
err.inner().column(),
err
);
})
};
let assistant_message2: Message = {
let jd = &mut serde_json::Deserializer::from_str(assistant_message_json2);
deserialize(jd).unwrap_or_else(|err| {
panic!(
"Deserialization error at {} ({}:{}): {}",
err.path(),
err.inner().line(),
err.inner().column(),
err
);
})
};
let assistant_message3: Message = {
let jd: &mut serde_json::Deserializer<serde_json::de::StrRead<'_>> =
&mut serde_json::Deserializer::from_str(assistant_message_json3);
deserialize(jd).unwrap_or_else(|err| {
panic!(
"Deserialization error at {} ({}:{}): {}",
err.path(),
err.inner().line(),
err.inner().column(),
err
);
})
};
let user_message: Message = {
let jd = &mut serde_json::Deserializer::from_str(user_message_json);
deserialize(jd).unwrap_or_else(|err| {
panic!(
"Deserialization error at {} ({}:{}): {}",
err.path(),
err.inner().line(),
err.inner().column(),
err
);
})
};
match assistant_message {
Message::Assistant { content, .. } => {
assert_eq!(
content[0],
AssistantContent::Text {
text: "\n\nHello there, how may I assist you today?".to_string()
}
);
}
_ => panic!("Expected assistant message"),
}
match assistant_message2 {
Message::Assistant {
content,
tool_calls,
..
} => {
assert_eq!(
content[0],
AssistantContent::Text {
text: "\n\nHello there, how may I assist you today?".to_string()
}
);
assert_eq!(tool_calls, vec![]);
}
_ => panic!("Expected assistant message"),
}
match assistant_message3 {
Message::Assistant {
content,
tool_calls,
refusal,
..
} => {
assert!(content.is_empty());
assert!(refusal.is_none());
assert_eq!(
tool_calls[0],
ToolCall {
id: "call_h89ipqYUjEpCPI6SxspMnoUU".to_string(),
r#type: ToolType::Function,
function: Function {
name: "subtract".to_string(),
arguments: serde_json::json!({"x": 2, "y": 5}),
},
}
);
}
_ => panic!("Expected assistant message"),
}
match user_message {
Message::User { content, .. } => {
let (first, second) = {
let mut iter = content.into_iter();
(iter.next().unwrap(), iter.next().unwrap())
};
assert_eq!(
first,
UserContent::Text {
text: "What's in this image?".to_string()
}
);
assert_eq!(second, UserContent::Image { image_url: ImageUrl { url: "https://upload.wikimedia.org/wikipedia/commons/thumb/d/dd/Gfp-wisconsin-madison-the-nature-boardwalk.jpg/2560px-Gfp-wisconsin-madison-the-nature-boardwalk.jpg".to_string(), detail: ImageDetail::default() } });
}
_ => panic!("Expected user message"),
}
}
#[test]
fn test_message_to_message_conversion() {
let user_message = message::Message::User {
content: OneOrMany::one(message::UserContent::text("Hello")),
};
let assistant_message = message::Message::Assistant {
content: OneOrMany::one(message::AssistantContent::text("Hi there!")),
};
let converted_user_message: Vec<Message> = user_message.clone().try_into().unwrap();
let converted_assistant_message: Vec<Message> =
assistant_message.clone().try_into().unwrap();
match converted_user_message[0].clone() {
Message::User { content, .. } => {
assert_eq!(
content.first(),
UserContent::Text {
text: "Hello".to_string()
}
);
}
_ => panic!("Expected user message"),
}
match converted_assistant_message[0].clone() {
Message::Assistant { content, .. } => {
assert_eq!(
content[0].clone(),
AssistantContent::Text {
text: "Hi there!".to_string()
}
);
}
_ => panic!("Expected assistant message"),
}
let original_user_message: message::Message =
converted_user_message[0].clone().try_into().unwrap();
let original_assistant_message: message::Message =
converted_assistant_message[0].clone().try_into().unwrap();
assert_eq!(original_user_message, user_message);
assert_eq!(original_assistant_message, assistant_message);
}
#[test]
fn test_message_from_message_conversion() {
let user_message = Message::User {
content: OneOrMany::one(UserContent::Text {
text: "Hello".to_string(),
}),
name: None,
};
let assistant_message = Message::Assistant {
content: vec![AssistantContent::Text {
text: "Hi there!".to_string(),
}],
refusal: None,
audio: None,
name: None,
tool_calls: vec![],
};
let converted_user_message: message::Message = user_message.clone().try_into().unwrap();
let converted_assistant_message: message::Message =
assistant_message.clone().try_into().unwrap();
match converted_user_message.clone() {
message::Message::User { content } => {
assert_eq!(content.first(), message::UserContent::text("Hello"));
}
_ => panic!("Expected user message"),
}
match converted_assistant_message.clone() {
message::Message::Assistant { content } => {
assert_eq!(
content.first(),
message::AssistantContent::text("Hi there!")
);
}
_ => panic!("Expected assistant message"),
}
let original_user_message: Vec<Message> = converted_user_message.try_into().unwrap();
let original_assistant_message: Vec<Message> =
converted_assistant_message.try_into().unwrap();
assert_eq!(original_user_message[0], user_message);
assert_eq!(original_assistant_message[0], assistant_message);
}
}