use crate::error::ProviderError;
use crate::providers::embed::EmbeddingResponse;
use crate::providers::types::build_http_client;
use crate::providers::types::ServiceType;
use potato_type::openai::v1::OpenAIChatResponse;
use potato_type::openai::v1::{OpenAIEmbeddingRequest, OpenAIEmbeddingResponse};
use potato_type::prompt::Prompt;
use potato_type::{Common, Provider};
use reqwest::Client;
use reqwest::Response;
use serde::Serialize;
use serde_json::Value;
use tracing::{debug, error, instrument};
#[derive(Debug, PartialEq)]
pub enum OpenAIAuth {
ApiKey(String),
NotSet,
}
impl OpenAIAuth {
#[instrument(skip_all)]
pub fn from_env() -> Self {
let api_key =
std::env::var("OPENAI_API_KEY").unwrap_or_else(|_| Common::Undefined.to_string());
if api_key != Common::Undefined.to_string() {
debug!("Using OpenAI API key from environment variable");
return Self::ApiKey(api_key);
}
Self::NotSet
}
}
struct OpenAIPaths {}
impl OpenAIPaths {
fn base_url() -> String {
"https://api.openai.com/v1".to_string()
}
}
#[derive(Debug, PartialEq)]
pub struct OpenAIApiConfig {
base_url: String,
service_type: ServiceType,
auth: OpenAIAuth,
}
impl OpenAIApiConfig {
fn new(service_type: ServiceType) -> Result<Self, ProviderError> {
let base_url = std::env::var("OPENAI_API_URL").unwrap_or(
std::env::var("OPENAI_BASE_URL").unwrap_or_else(|_| OpenAIPaths::base_url()),
);
let auth = OpenAIAuth::from_env();
Ok(Self {
base_url,
service_type,
auth,
})
}
fn build_url(&self) -> String {
let endpoint = self.get_endpoint();
format!("{}/{}", self.base_url, endpoint)
}
async fn set_auth_header(
&self,
req: reqwest::RequestBuilder,
) -> Result<reqwest::RequestBuilder, ProviderError> {
match &self.auth {
OpenAIAuth::ApiKey(api_key) => Ok(req.bearer_auth(api_key)),
OpenAIAuth::NotSet => Ok(req),
}
}
fn get_endpoint(&self) -> &'static str {
self.service_type.openai_endpoint()
}
}
#[derive(Debug)]
pub struct OpenAIClient {
client: Client,
config: OpenAIApiConfig,
pub provider: Provider,
}
impl PartialEq for OpenAIClient {
fn eq(&self, other: &Self) -> bool {
self.config == other.config && self.provider == other.provider
}
}
impl OpenAIClient {
pub fn new(service_type: ServiceType) -> Result<Self, ProviderError> {
let client = build_http_client(None)?;
let config = OpenAIApiConfig::new(service_type)?;
Ok(Self {
client,
config,
provider: Provider::OpenAI,
})
}
async fn make_request(&self, object: &Value) -> Result<Response, ProviderError> {
let request = self.client.post(self.config.build_url()).json(&object);
let request = self.config.set_auth_header(request).await?;
let response = request.send().await.map_err(ProviderError::RequestError)?;
let status = response.status();
if !status.is_success() {
error!("OpenAI API request failed with status: {}", status);
let body = response
.text()
.await
.unwrap_or_else(|_| "No response body".to_string());
return Err(ProviderError::CompletionError(body, status));
}
Ok(response)
}
#[instrument(skip_all)]
pub async fn chat_completion(
&self,
prompt: &Prompt,
) -> Result<OpenAIChatResponse, ProviderError> {
if let OpenAIAuth::NotSet = self.config.auth {
return Err(ProviderError::MissingAuthenticationError);
}
let request_body = prompt.request.to_request(&self.provider)?;
debug!(
"Sending chat completion request to OpenAI API: {:?}",
request_body
);
let response = self.make_request(&request_body).await?;
let chat_response: OpenAIChatResponse = response.json().await?;
debug!("Chat completion successful");
Ok(chat_response)
}
#[instrument(skip_all)]
pub async fn create_embedding<T>(
&self,
inputs: Vec<String>,
config: &T,
) -> Result<EmbeddingResponse, ProviderError>
where
T: Serialize,
{
if let OpenAIAuth::NotSet = self.config.auth {
return Err(ProviderError::MissingAuthenticationError);
}
let request = serde_json::to_value(OpenAIEmbeddingRequest::new(inputs, config))
.map_err(ProviderError::SerializationError)?;
debug!("Sending embedding request to OpenAI API: {:?}", request);
let response = self.make_request(&request).await?;
let embedding_response: OpenAIEmbeddingResponse = response.json().await?;
Ok(EmbeddingResponse::OpenAI(embedding_response))
}
}