use std::io::{BufRead, BufReader};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, OnceLock};
use std::time::{Duration, Instant};
#[cfg(feature = "chat-backend")]
use crate::backend::ChatBackend;
use crate::backend::{GenerationResult, InferenceParams, LlmBackend, TokenCallback};
use crate::error::{CoreError, CoreResult};
#[cfg(feature = "chat-backend")]
use crate::messages::{Message, Role};
const DEFAULT_TIMEOUT: Duration = Duration::from_secs(120);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum OpenAiEndpoint {
#[default]
Chat,
Completions,
}
#[derive(Debug, Clone)]
pub struct OpenAiConfig {
pub base_url: String,
pub model: String,
pub api_key: Option<String>,
pub endpoint: OpenAiEndpoint,
pub timeout: Duration,
}
impl OpenAiConfig {
pub fn new(base_url: impl Into<String>, model: impl Into<String>) -> Self {
Self {
base_url: base_url.into(),
model: model.into(),
api_key: None,
endpoint: OpenAiEndpoint::default(),
timeout: DEFAULT_TIMEOUT,
}
}
pub fn with_api_key(mut self, api_key: impl Into<String>) -> Self {
self.api_key = Some(api_key.into());
self
}
pub fn with_endpoint(mut self, endpoint: OpenAiEndpoint) -> Self {
self.endpoint = endpoint;
self
}
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.timeout = timeout;
self
}
}
pub struct OpenAiHttpBackend {
blocking: OnceLock<reqwest::blocking::Client>,
#[cfg(feature = "chat-backend")]
streaming: reqwest::Client,
config: OpenAiConfig,
}
impl OpenAiHttpBackend {
pub fn new(config: OpenAiConfig) -> CoreResult<Self> {
Ok(Self {
blocking: OnceLock::new(),
#[cfg(feature = "chat-backend")]
streaming: reqwest::Client::builder()
.timeout(config.timeout)
.build()
.map_err(|e| CoreError::Backend(format!("failed to build HTTP client: {e}")))?,
config,
})
}
fn blocking(&self) -> CoreResult<&reqwest::blocking::Client> {
if let Some(client) = self.blocking.get() {
return Ok(client);
}
let built = reqwest::blocking::Client::builder()
.timeout(self.config.timeout)
.build()
.map_err(|e| CoreError::Backend(format!("failed to build HTTP client: {e}")))?;
let _ = self.blocking.set(built);
self.blocking
.get()
.ok_or_else(|| CoreError::Backend("HTTP client went missing".into()))
}
}
impl LlmBackend for OpenAiHttpBackend {
fn generate(
&self,
prompt: &str,
params: &InferenceParams,
abort: Arc<AtomicBool>,
mut on_token: TokenCallback,
) -> CoreResult<GenerationResult> {
let mut body = serde_json::json!({
"model": self.config.model,
"max_tokens": params.max_tokens,
"temperature": params.temperature,
"stream": true,
"stream_options": { "include_usage": true },
});
let path = match self.config.endpoint {
OpenAiEndpoint::Chat => {
body["messages"] = serde_json::json!([{ "role": "user", "content": prompt }]);
"chat/completions"
}
OpenAiEndpoint::Completions => {
body["prompt"] = serde_json::json!(prompt);
"completions"
}
};
let url = format!("{}/{path}", self.config.base_url);
let mut req = self.blocking()?.post(&url).json(&body);
if let Some(key) = &self.config.api_key {
req = req.bearer_auth(key);
}
let resp = req
.send()
.map_err(|e| CoreError::BackendUnreachable(format!("request to {url} failed: {e}")))?;
let status = resp.status();
if !status.is_success() {
let detail = resp.text().unwrap_or_default();
return Err(CoreError::Backend(format!(
"endpoint returned HTTP {}: {}",
status.as_u16(),
detail.trim()
)));
}
let start = Instant::now();
let mut ttft_ms = 0.0;
let mut text = String::new();
let mut streamed: u32 = 0;
let mut usage_prompt: Option<u32> = None;
let mut usage_completion: Option<u32> = None;
let reader = BufReader::new(resp);
for line in reader.lines() {
if abort.load(Ordering::Relaxed) {
return Err(CoreError::Aborted);
}
let line = line
.map_err(|e| CoreError::BackendUnreachable(format!("stream read failed: {e}")))?;
let Some(data) = line.strip_prefix("data: ") else {
continue;
};
let piece = read_event(data, self.config.endpoint);
usage_prompt = piece.prompt_tokens.or(usage_prompt);
usage_completion = piece.completion_tokens.or(usage_completion);
if piece.done {
break;
}
if let Some(said) = piece.text {
if streamed == 0 {
ttft_ms = start.elapsed().as_secs_f64() * 1000.0;
}
streamed += 1;
text.push_str(&said);
let elapsed = start.elapsed().as_secs_f64().max(1e-6);
on_token(&said, streamed, streamed as f64 / elapsed);
}
}
let gen_ms = start.elapsed().as_secs_f64() * 1000.0;
let tokens_generated = usage_completion.unwrap_or(streamed);
let prompt_tokens = usage_prompt.unwrap_or(0);
Ok(GenerationResult {
text,
tokens_generated,
prompt_tokens,
tokens_per_sec: tokens_generated as f64 / (gen_ms / 1000.0).max(1e-6),
time_to_first_token_ms: ttft_ms,
generation_time_ms: gen_ms,
})
}
fn tokenize_count(&self, text: &str) -> CoreResult<u32> {
Ok((text.chars().count() as u32 / 4).max(1))
}
fn is_ready(&self) -> bool {
true
}
}
#[derive(Debug, Default)]
struct Piece {
text: Option<String>,
prompt_tokens: Option<u32>,
completion_tokens: Option<u32>,
done: bool,
}
fn read_event(data: &str, endpoint: OpenAiEndpoint) -> Piece {
let data = data.trim();
if data == "[DONE]" {
return Piece {
done: true,
..Piece::default()
};
}
let Ok(chunk) = serde_json::from_str::<serde_json::Value>(data) else {
return Piece::default();
};
let usage = chunk.get("usage").filter(|usage| !usage.is_null());
let count = |name: &str| {
usage
.and_then(|usage| usage.get(name))
.and_then(serde_json::Value::as_u64)
.map(|n| n as u32)
};
let text = match endpoint {
OpenAiEndpoint::Chat => chunk["choices"][0]["delta"]["content"].as_str(),
OpenAiEndpoint::Completions => chunk["choices"][0]["text"].as_str(),
};
Piece {
text: text.filter(|piece| !piece.is_empty()).map(str::to_string),
prompt_tokens: count("prompt_tokens"),
completion_tokens: count("completion_tokens"),
done: false,
}
}
#[cfg(feature = "chat-backend")]
fn wire(system: &str, messages: &[Message]) -> Vec<serde_json::Value> {
let mut turns = Vec::with_capacity(messages.len() + 1);
if !system.trim().is_empty() {
turns.push(serde_json::json!({ "role": "system", "content": system }));
}
for message in messages {
let (role, content) = match message.role {
Role::System => ("system", message.content.clone()),
Role::User => ("user", message.content.clone()),
Role::Assistant | Role::ToolCall => ("assistant", message.content.clone()),
Role::ToolResult => ("user", format!("[Tool result]\n{}", message.content)),
};
turns.push(serde_json::json!({ "role": role, "content": content }));
}
turns
}
#[cfg(feature = "chat-backend")]
#[async_trait::async_trait]
impl ChatBackend for OpenAiHttpBackend {
async fn chat(
&self,
system: &str,
messages: &[Message],
params: &InferenceParams,
abort: Arc<AtomicBool>,
mut on_token: TokenCallback,
) -> CoreResult<GenerationResult> {
let body = serde_json::json!({
"model": self.config.model,
"messages": wire(system, messages),
"max_tokens": params.max_tokens,
"temperature": params.temperature,
"stream": true,
"stream_options": { "include_usage": true },
});
let url = format!("{}/chat/completions", self.config.base_url);
let mut request = self.streaming.post(&url).json(&body);
if let Some(key) = &self.config.api_key {
request = request.bearer_auth(key);
}
let mut response = request
.send()
.await
.map_err(|e| CoreError::BackendUnreachable(format!("request to {url} failed: {e}")))?;
let status = response.status();
if !status.is_success() {
let detail = response.text().await.unwrap_or_default();
return Err(CoreError::Backend(format!(
"endpoint returned HTTP {}: {}",
status.as_u16(),
detail.trim()
)));
}
let start = Instant::now();
let mut ttft_ms = 0.0;
let mut text = String::new();
let mut streamed: u32 = 0;
let mut usage_prompt: Option<u32> = None;
let mut usage_completion: Option<u32> = None;
let mut pending = String::new();
let mut finished = false;
while !finished {
if abort.load(Ordering::Relaxed) {
return Err(CoreError::Aborted);
}
let Some(bytes) = response
.chunk()
.await
.map_err(|e| CoreError::BackendUnreachable(format!("stream read failed: {e}")))?
else {
break;
};
pending.push_str(&String::from_utf8_lossy(&bytes));
while let Some(at) = pending.find('\n') {
let line: String = pending.drain(..=at).collect();
let Some(data) = line.trim_end().strip_prefix("data: ") else {
continue;
};
let piece = read_event(data, OpenAiEndpoint::Chat);
usage_prompt = piece.prompt_tokens.or(usage_prompt);
usage_completion = piece.completion_tokens.or(usage_completion);
if piece.done {
finished = true;
break;
}
if let Some(said) = piece.text {
if streamed == 0 {
ttft_ms = start.elapsed().as_secs_f64() * 1000.0;
}
streamed += 1;
text.push_str(&said);
let elapsed = start.elapsed().as_secs_f64().max(1e-6);
on_token(&said, streamed, streamed as f64 / elapsed);
}
}
}
let gen_ms = start.elapsed().as_secs_f64() * 1000.0;
let tokens_generated = usage_completion.unwrap_or(streamed);
Ok(GenerationResult {
text,
tokens_generated,
prompt_tokens: usage_prompt.unwrap_or(0),
tokens_per_sec: tokens_generated as f64 / (gen_ms / 1000.0).max(1e-6),
time_to_first_token_ms: ttft_ms,
generation_time_ms: gen_ms,
})
}
}