use crate::chat_client::{
context::Context,
openai_api::{
chat_completions::{ChatCompletionsBody, OpenRouterReasoning},
client::{Auth, Error as OpenAiClientError, OpenAiClient},
message::{self, AssistantMessage},
},
};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ReasoningSettings {
Effort(String),
Budget(i64),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ApiOptions {
OpenAi {
reasoning_effort: Option<String>,
},
OpenRouter {
reasoning: Option<ReasoningSettings>,
},
}
impl ApiOptions {
pub fn as_openai_reasoning_effort(&self) -> Option<String> {
match self {
ApiOptions::OpenAi { reasoning_effort } => reasoning_effort.clone(),
_ => None,
}
}
pub fn as_openrouter_reasoning_settings(&self) -> Option<OpenRouterReasoning> {
match self {
ApiOptions::OpenRouter { reasoning } => reasoning.as_ref().map(|r| match r {
ReasoningSettings::Effort(e) => OpenRouterReasoning::from_effort(e.clone()),
ReasoningSettings::Budget(b) => OpenRouterReasoning::from_budget(*b),
}),
_ => None,
}
}
}
#[derive(Debug)]
pub struct ChatClientConfig {
pub api_url: String,
pub api_options: ApiOptions,
pub api_version: Option<String>,
pub model: String,
pub system_message: Option<String>,
pub min_history_tokens: Option<usize>,
pub max_history_tokens: Option<usize>,
pub verbosity: Option<String>,
}
impl Default for ChatClientConfig {
fn default() -> Self {
Self {
api_url: String::from("https://api.openai.com/v1/"),
api_options: ApiOptions::OpenAi {
reasoning_effort: None,
},
api_version: None,
model: String::from("gpt-4o-mini"),
system_message: None,
min_history_tokens: None,
max_history_tokens: None,
verbosity: None,
}
}
}
#[derive(Debug)]
pub struct Completion {
pub response: String,
pub reasoning: Option<String>,
pub tokens_in: usize,
pub tokens_in_cached: Option<usize>,
pub tokens_out: usize,
pub tokens_reasoning: Option<usize>,
}
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error("API error: {0}")]
OpenAiClient(#[from] OpenAiClientError),
#[error("Response contains no choices")]
NoChoices,
#[error("Response contains no message")]
NoMessage,
#[error("Invalid message: {0}")]
InvalidMessage(#[from] message::Error),
#[error("Assistant message contains no `content`")]
NoContent,
#[error("Model refused the request: \"{0}\"")]
Refusal(String),
#[error("Failed to initialize tokenizer: {0}")]
TokenizerInit(String),
}
#[derive(Debug, Clone)]
pub struct ModelConfig {
pub model: String,
pub api_options: ApiOptions,
pub verbosity: Option<String>,
}
pub struct ChatClient {
client: OpenAiClient,
model_config: ModelConfig,
context: Context,
}
impl ChatClient {
pub fn new(auth: Auth, config: ChatClientConfig) -> Result<Self, Error> {
let ChatClientConfig {
api_url,
api_options,
api_version,
model,
system_message,
min_history_tokens,
max_history_tokens,
verbosity,
} = config;
let api_url = ensure_trailing_slash(api_url);
let context = create_context(system_message, min_history_tokens, max_history_tokens)?;
Ok(Self {
client: OpenAiClient::new(auth, api_url, api_version)?,
model_config: ModelConfig {
model,
api_options,
verbosity,
},
context,
})
}
pub fn new_with_client(
client: reqwest::Client,
config: ChatClientConfig,
) -> Result<Self, Error> {
let ChatClientConfig {
api_url,
api_options,
api_version,
model,
system_message,
min_history_tokens,
max_history_tokens,
verbosity,
} = config;
let api_url = ensure_trailing_slash(api_url);
let context = create_context(system_message, min_history_tokens, max_history_tokens)?;
Ok(Self {
client: OpenAiClient::new_with_client(client, api_url, api_version),
model_config: ModelConfig {
model,
api_options,
verbosity,
},
context,
})
}
pub async fn ask(&mut self, request: String) -> Result<String, Error> {
self.request_completion(request).await.map(|c| c.response)
}
pub async fn request_completion(&mut self, request: String) -> Result<Completion, Error> {
let mut completion = self
.client
.chat_completions(Self::body(
self.model_config.clone(),
&self.context,
request.clone(),
))
.await?;
let choice = completion.choices.pop().ok_or(Error::NoChoices)?;
let assistant_message = AssistantMessage::try_from(choice.message)?;
let response = assistant_message.content.ok_or(
assistant_message
.refusal
.map_or(Error::NoContent, Error::Refusal),
)?;
self.context.push(request, response.clone());
Ok(Completion {
response,
reasoning: assistant_message.reasoning,
tokens_in: completion.usage.prompt_tokens,
tokens_in_cached: completion
.usage
.prompt_tokens_details
.and_then(|d| d.cached_tokens),
tokens_out: completion.usage.completion_tokens,
tokens_reasoning: completion
.usage
.completion_tokens_details
.and_then(|d| d.reasoning_tokens),
})
}
fn body(
ModelConfig {
model,
api_options,
verbosity,
}: ModelConfig,
context: &Context,
request: String,
) -> ChatCompletionsBody {
ChatCompletionsBody {
model,
messages: context.with_request(request).map(Into::into).collect(),
reasoning_effort: api_options.as_openai_reasoning_effort(),
reasoning: api_options.as_openrouter_reasoning_settings(),
verbosity,
..Default::default()
}
}
}
fn ensure_trailing_slash(url: String) -> String {
if url.ends_with('/') {
url
} else {
url + "/"
}
}
fn create_context(
system_message: Option<String>,
min_history_tokens: Option<usize>,
max_history_tokens: Option<usize>,
) -> Result<Context, Error> {
let context = if min_history_tokens.is_some() || max_history_tokens.is_some() {
Context::new_with_rolling_window(
system_message,
tiktoken_rs::o200k_base().map_err(|e| Error::TokenizerInit(format!("{e}")))?,
min_history_tokens,
max_history_tokens,
)
} else {
Context::new(system_message)
};
Ok(context)
}