use crate::chat_client::{
context::Context,
error::Error,
openai_api::{
chat_completions::{ChatCompletionsRequest, OpenRouterReasoning, StreamOptions, Usage},
client::{Auth, OpenAiClient, OpenAiClientConfig},
message::{Content, ContentPart, ResponseAssistantMessage},
},
stream::CompletionStream,
};
use eventsource_stream::{Event, EventStreamError};
use futures::stream::Stream;
use regex::Regex;
use serde_json::{json, Value};
use std::time::Duration;
#[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>,
pdf_engine: Option<String>,
image_generation: bool,
},
}
impl ApiOptions {
fn as_openai_reasoning_effort(&self) -> Option<String> {
match self {
ApiOptions::OpenAi { reasoning_effort } => reasoning_effort.clone(),
_ => None,
}
}
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,
}
}
fn as_openrouter_plugins(&self) -> Option<Vec<Value>> {
match self {
ApiOptions::OpenRouter {
pdf_engine: Some(engine),
..
} => Some(vec![json!({
"id": "file-parser",
"pdf": {
"engine": engine,
}
})]),
_ => None,
}
}
fn as_openrouter_modalities(&self) -> Option<Value> {
match self {
ApiOptions::OpenRouter {
image_generation, ..
} => image_generation.then_some(json!(["image", "text"])),
_ => None,
}
}
}
#[derive(Debug)]
pub struct ChatClientConfig {
pub auth: Auth,
pub api_url: String,
pub api_options: ApiOptions,
pub api_version: Option<String>,
pub http_timeout: Duration,
pub model: String,
pub system_message: Option<String>,
pub min_history_tokens: Option<usize>,
pub max_history_tokens: Option<usize>,
pub verbosity: Option<String>,
pub sanitize_links: bool,
pub extra_params: Option<serde_json::map::Map<String, Value>>,
}
impl ChatClientConfig {
pub fn default_with_auth(auth: Auth) -> Self {
Self {
auth,
api_url: String::from("https://api.openai.com/v1/"),
api_options: ApiOptions::OpenAi {
reasoning_effort: None,
},
api_version: None,
http_timeout: Duration::from_secs(300),
model: String::from("gpt-4o-mini"),
system_message: None,
min_history_tokens: None,
max_history_tokens: None,
verbosity: None,
sanitize_links: false,
extra_params: None,
}
}
}
#[derive(Debug, Default)]
pub struct TokenUsage {
pub tokens_in: usize,
pub tokens_in_cached: Option<usize>,
pub tokens_out: usize,
pub tokens_reasoning: Option<usize>,
}
impl From<Usage> for TokenUsage {
fn from(usage: Usage) -> Self {
Self {
tokens_in: usage.prompt_tokens,
tokens_in_cached: usage.prompt_tokens_details.and_then(|d| d.cached_tokens),
tokens_out: usage.completion_tokens,
tokens_reasoning: usage
.completion_tokens_details
.and_then(|d| d.reasoning_tokens),
}
}
}
#[derive(Debug)]
pub struct Completion {
pub response: Content,
pub reasoning: Option<String>,
pub token_usage: TokenUsage,
}
#[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,
sanitize_links: bool,
sanitize_re: regex::Regex,
extra_params: Option<serde_json::map::Map<String, Value>>,
}
impl ChatClient {
pub fn new(config: ChatClientConfig) -> Result<Self, Error> {
Self::new_with_client(config, reqwest::Client::new())
}
pub fn new_with_client(
config: ChatClientConfig,
client: reqwest::Client,
) -> Result<Self, Error> {
let system_tokens = match config.system_message {
None => 0,
Some(ref system_message) => {
let tokenizer =
tiktoken_rs::o200k_base().map_err(|e| Error::TokenizerInit(format!("{e}")))?;
tokenizer.encode_with_special_tokens(system_message).len()
}
};
Self::new_with_client_and_system_tokens(config, client, system_tokens)
}
pub fn new_with_client_and_system_tokens(
config: ChatClientConfig,
client: reqwest::Client,
system_tokens: usize,
) -> Result<Self, Error> {
let ChatClientConfig {
auth,
api_url,
api_options,
api_version,
http_timeout,
model,
system_message,
min_history_tokens,
max_history_tokens,
verbosity,
sanitize_links,
extra_params,
} = config;
let client = OpenAiClient::new(OpenAiClientConfig {
client,
auth,
base_url: ensure_trailing_slash(api_url),
api_version,
timeout: http_timeout,
})?;
let context = Context::new(
system_message,
system_tokens,
min_history_tokens,
max_history_tokens,
);
let sanitize_re = Regex::new(r"(?:\?|\&)utm_source=openai\)").expect("to be valid regex");
Ok(Self {
client,
model_config: ModelConfig {
model,
api_options,
verbosity,
},
context,
sanitize_links,
sanitize_re,
extra_params,
})
}
pub async fn ask(&mut self, request: String) -> Result<String, Error> {
self.request_completion(Content::Text(request))
.await
.and_then(|c| match c.response {
Content::Text(text) => Ok(text),
Content::ContentParts(_) => Err(Error::NonTextContent),
})
}
pub async fn request_completion(&mut self, request: Content) -> Result<Completion, Error> {
let mut completion = self
.client
.chat_completions(Self::body(
self.model_config.clone(),
&self.context,
request.clone(),
false,
self.extra_params.clone(),
)?)
.await?;
let choice = completion.choices.pop().ok_or(Error::NoChoices)?;
let assistant_message = ResponseAssistantMessage::try_from(choice.message)?;
let response = assistant_message
.content
.map(|response| self.sanitize_links(response))
.ok_or(
assistant_message
.refusal
.map_or(Error::NoContent, Error::Refusal),
)?;
let response = if let Some(images) = assistant_message.images {
Content::ContentParts(
(!response.is_empty())
.then_some(ContentPart::Text(response))
.into_iter()
.chain(images.into_iter().map(ContentPart::Image))
.collect(),
)
} else {
Content::Text(response)
};
let token_usage = completion.usage.into();
self.extend_context(request, response.clone(), &token_usage);
Ok(Completion {
response,
reasoning: assistant_message.reasoning,
token_usage,
})
}
pub async fn stream_completion<'a>(
&'a mut self,
request: Content,
) -> Result<
CompletionStream<'a, impl Stream<Item = Result<Event, EventStreamError<reqwest::Error>>>>,
Error,
> {
let stream = self
.client
.chat_completions_stream(Self::body(
self.model_config.clone(),
&self.context,
request.clone(),
true,
self.extra_params.clone(),
)?)
.await?;
Ok(CompletionStream::new(self, stream, request))
}
pub(crate) fn extend_context(
&mut self,
request: Content,
response: Content,
usage: &TokenUsage,
) {
let request_tokens = usage.tokens_in.saturating_sub(self.context.tokens());
let response_tokens = usage
.tokens_out
.saturating_sub(usage.tokens_reasoning.unwrap_or_default());
self.context
.push(request, response, request_tokens + response_tokens);
}
fn body(
ModelConfig {
model,
api_options,
verbosity,
}: ModelConfig,
context: &Context,
request: Content,
stream: bool,
extra_params: Option<serde_json::map::Map<String, Value>>,
) -> Result<Value, Error> {
let mut json = serde_json::to_value(ChatCompletionsRequest {
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,
stream: Some(stream),
stream_options: stream.then_some(StreamOptions {
include_obfuscation: None,
include_usage: Some(true),
}),
plugins: api_options.as_openrouter_plugins(),
modalities: api_options.as_openrouter_modalities(),
..Default::default()
})
.map_err(Error::BodySerializationError)?;
if let Some(extra_params) = extra_params {
let json_object = json
.as_object_mut()
.ok_or(Error::InternalError("chat completions body is not a JSON"))?;
json_object.extend(extra_params);
}
Ok(json)
}
fn sanitize_links(&self, response: String) -> String {
if self.sanitize_links {
self.sanitize_re.replace_all(&response, ")").to_string()
} else {
response
}
}
}
fn ensure_trailing_slash(url: String) -> String {
if url.ends_with('/') {
url
} else {
url + "/"
}
}