use std::{
collections::{BTreeMap, HashMap},
pin::Pin,
sync::{Arc, LazyLock},
time::Duration,
};
use async_trait::async_trait;
use enumset::enum_set;
use futures::stream::{Stream, StreamExt};
use moka::future::Cache;
use serde::Deserialize;
use crate::{
capabilities::{
MediaKind, MediaSupport, ModelCapabilities, ReasoningCapability, ReasoningMode,
ReasoningParamConflicts,
},
config::ReasoningEffort,
error::LanguageModelError,
identifiers::{ApiKey, ModelId},
media::SourceKind,
openai_wire::{
OpenAiResponse, OpenAiStreamOptions, build_request, convert_openai_sse_event,
map_openai_error, parse_openai_response,
},
provider::{ChatModelInfo, GenerateRequest, LanguageModelProvider, ResponseFormatKind},
response::{LanguageModelResponse, StreamDelta},
retry::{RetryConfig, with_retry},
sse::parse_sse_stream,
};
pub struct OpenAiLanguageModel {
client: Arc<reqwest::Client>,
api_key: String,
base_url: String,
retry_config: RetryConfig,
list_models_cache: Cache<(), Vec<ChatModelInfo>>,
}
pub struct OpenAiConfig {
pub api_key: ApiKey,
pub base_url: String,
pub retry_config: RetryConfig,
}
impl OpenAiConfig {
pub const DEFAULT_BASE_URL: &'static str = "https://api.openai.com/v1";
}
pub struct OpenAiDeps {
pub client: Arc<reqwest::Client>,
}
impl OpenAiLanguageModel {
#[must_use]
pub fn new(deps: OpenAiDeps, config: OpenAiConfig) -> Self {
Self {
client: deps.client,
api_key: config.api_key.into_string(),
base_url: config.base_url,
retry_config: config.retry_config,
list_models_cache: Cache::builder()
.time_to_live(Duration::from_secs(3600))
.max_capacity(1)
.build(),
}
}
async fn execute(
&self,
model: &str,
messages: &[crate::message::Message],
config: &crate::config::LanguageModelConfig,
) -> Result<LanguageModelResponse, LanguageModelError> {
let request_body = build_request(
model,
messages,
config,
openai_uses_completion_tokens(model),
);
let url = format!("{}/chat/completions", self.base_url);
let response = self
.client
.post(&url)
.header("authorization", format!("Bearer {}", self.api_key))
.header("content-type", "application/json")
.json(&request_body)
.send()
.await
.map_err(|e| LanguageModelError::provider(e.to_string()))?;
let status = response.status();
if !status.is_success() {
let headers = response.headers().clone();
let body =
crate::http_error::read_error_body_or_warn(response, "openai", status.as_u16())
.await;
return Err(map_openai_error(status.as_u16(), &body, &headers));
}
let api_response: OpenAiResponse = response
.json()
.await
.map_err(|e| LanguageModelError::provider(format!("failed to parse response: {e}")))?;
parse_openai_response(api_response)
}
async fn stream(
&self,
model: &str,
messages: &[crate::message::Message],
config: &crate::config::LanguageModelConfig,
) -> Result<
Pin<Box<dyn Stream<Item = Result<StreamDelta, LanguageModelError>> + Send>>,
LanguageModelError,
> {
let mut request_body = build_request(
model,
messages,
config,
openai_uses_completion_tokens(model),
);
request_body.stream = Some(true);
request_body.stream_options = Some(OpenAiStreamOptions {
include_usage: true,
});
let url = format!("{}/chat/completions", self.base_url);
let response = self
.client
.post(&url)
.header("authorization", format!("Bearer {}", self.api_key))
.header("content-type", "application/json")
.json(&request_body)
.send()
.await
.map_err(|e| LanguageModelError::provider(e.to_string()))?;
let status = response.status();
if !status.is_success() {
let headers = response.headers().clone();
let body =
crate::http_error::read_error_body_or_warn(response, "openai", status.as_u16())
.await;
return Err(map_openai_error(status.as_u16(), &body, &headers));
}
let byte_stream = response.bytes_stream();
let sse_stream = parse_sse_stream(byte_stream);
let delta_stream = sse_stream
.filter_map(|event_result| async move { convert_openai_sse_event(event_result) });
Ok(Box::pin(delta_stream))
}
async fn fetch_and_merge_models(&self) -> Result<Vec<ChatModelInfo>, LanguageModelError> {
let url = format!("{}/models", self.base_url);
let response = self
.client
.get(&url)
.header("authorization", format!("Bearer {}", self.api_key))
.send()
.await
.map_err(|e| LanguageModelError::provider(e.to_string()))?;
let status = response.status();
if !status.is_success() {
let headers = response.headers().clone();
let body =
crate::http_error::read_error_body_or_warn(response, "openai", status.as_u16())
.await;
return Err(map_openai_error(status.as_u16(), &body, &headers));
}
let api_response: OpenAiListModelsResponse = response
.json()
.await
.map_err(|e| LanguageModelError::provider(format!("failed to parse models: {e}")))?;
let mut out = Vec::with_capacity(api_response.data.len());
for entry in api_response.data {
let caps = MODEL_CAPABILITIES.get(entry.id.as_str()).cloned();
if caps.is_none() {
tracing::warn!(
provider = "openai",
model_id = %entry.id,
"model returned by upstream but no local capability metadata; \
ChatModelInfo will have minimal fields. \
Update the local capability table when ready."
);
}
let mut formats = vec![ResponseFormatKind::Text];
if caps.as_ref().is_some_and(|c| c.supports_json_schema) {
formats.push(ResponseFormatKind::JsonObject);
formats.push(ResponseFormatKind::JsonSchema);
}
out.push(ChatModelInfo {
id: ModelId::new(entry.id),
display_name: None,
context_window: caps.as_ref().map(|c| c.context_window),
supports_streaming: caps.as_ref().is_some_and(|c| c.supports_streaming),
supported_response_formats: formats,
media_support: caps.as_ref().map(openai_media_support).unwrap_or_default(),
reasoning: caps.as_ref().and_then(|c| c.reasoning.clone()),
});
}
Ok(out)
}
}
#[async_trait]
impl LanguageModelProvider for OpenAiLanguageModel {
fn name(&self) -> &'static str {
"openai"
}
fn capabilities(&self, model: &ModelId) -> Option<ModelCapabilities> {
let caps = MODEL_CAPABILITIES.get(model.as_str())?;
Some(ModelCapabilities {
model_id: model.as_str().to_owned(),
media_support: openai_media_support(caps),
reasoning: caps.reasoning.clone(),
latency_optimized_supported: false,
extended_cache_ttl_supported: false,
})
}
#[tracing::instrument(skip(self), fields(provider = "openai", model_count = tracing::field::Empty), err(Display))]
async fn list_models(&self) -> Result<Vec<ChatModelInfo>, LanguageModelError> {
let result = self
.list_models_cache
.try_get_with((), self.fetch_and_merge_models())
.await
.map_err(|arc_err| (*arc_err).clone());
if let Ok(ref v) = result {
tracing::Span::current().record("model_count", v.len());
}
result
}
#[tracing::instrument(
skip(self, request),
fields(
provider = "openai",
model = %request.model,
messages = request.messages.len(),
prompt_tokens = tracing::field::Empty,
completion_tokens = tracing::field::Empty,
total_tokens = tracing::field::Empty,
cache_creation_input_tokens = tracing::field::Empty,
cache_read_input_tokens = tracing::field::Empty,
),
err(Display),
)]
async fn generate(
&self,
request: GenerateRequest<'_>,
) -> Result<LanguageModelResponse, LanguageModelError> {
self.validate_request(&request)?;
let model = request.model.as_str();
let response = with_retry(&self.retry_config, || {
self.execute(model, request.messages, request.config)
})
.await?;
let span = tracing::Span::current();
if let Some(usage) = &response.usage {
span.record("prompt_tokens", usage.input_tokens);
span.record("completion_tokens", usage.output_tokens);
span.record("total_tokens", usage.input_tokens + usage.output_tokens);
span.record(
"cache_creation_input_tokens",
usage.cache_creation_input_tokens,
);
span.record("cache_read_input_tokens", usage.cache_read_input_tokens);
}
Ok(response)
}
#[tracing::instrument(
skip(self, request),
fields(
provider = "openai",
model = %request.model,
messages = request.messages.len(),
first_token_ms = tracing::field::Empty,
prompt_tokens = tracing::field::Empty,
completion_tokens = tracing::field::Empty,
total_tokens = tracing::field::Empty,
cache_creation_input_tokens = tracing::field::Empty,
cache_read_input_tokens = tracing::field::Empty,
),
err(Display),
)]
async fn generate_stream(
&self,
request: GenerateRequest<'_>,
) -> Result<
Pin<Box<dyn Stream<Item = Result<StreamDelta, LanguageModelError>> + Send>>,
LanguageModelError,
> {
let started_at = std::time::Instant::now();
self.validate_request(&request)?;
let inner = self
.stream(request.model.as_str(), request.messages, request.config)
.await?;
let wrapped =
crate::streaming_timing::instrument_stream(tracing::Span::current(), started_at, inner);
Ok(Box::pin(wrapped))
}
}
fn openai_media_support(_caps: &ChatModelCapabilities) -> BTreeMap<MediaKind, MediaSupport> {
let mut m = BTreeMap::new();
m.insert(
MediaKind::Image,
MediaSupport {
sources: enum_set!(SourceKind::Url | SourceKind::InlineBytes),
formats: &["png", "jpeg", "webp", "gif"],
max_bytes: Some(20 * 1024 * 1024),
max_count_per_message: None,
},
);
m.insert(
MediaKind::Document,
MediaSupport {
sources: enum_set!(SourceKind::InlineBytes | SourceKind::ProviderFile),
formats: &["pdf"],
max_bytes: Some(50 * 1024 * 1024),
max_count_per_message: None,
},
);
m
}
#[derive(Debug, Clone)]
struct ChatModelCapabilities {
context_window: u32,
supports_streaming: bool,
supports_json_schema: bool,
reasoning: Option<ReasoningCapability>,
}
const fn openai_reasoning_conflicts() -> ReasoningParamConflicts {
ReasoningParamConflicts {
temperature_forbidden: true,
top_k_forbidden: false,
top_p_allowed_range: None,
}
}
const fn openai_gpt5_reasoning() -> ReasoningCapability {
ReasoningCapability {
supported_modes: enum_set!(ReasoningMode::Adaptive),
supported_efforts: enum_set!(
ReasoningEffort::None
| ReasoningEffort::Low
| ReasoningEffort::Medium
| ReasoningEffort::High
| ReasoningEffort::XHigh
),
manual_budget_range: None,
conflicts: openai_reasoning_conflicts(),
sampling_params_removed: false,
}
}
const fn openai_o_series_reasoning() -> ReasoningCapability {
ReasoningCapability {
supported_modes: enum_set!(ReasoningMode::Adaptive),
supported_efforts: enum_set!(
ReasoningEffort::Low | ReasoningEffort::Medium | ReasoningEffort::High
),
manual_budget_range: None,
conflicts: openai_reasoning_conflicts(),
sampling_params_removed: false,
}
}
static MODEL_CAPABILITIES: LazyLock<HashMap<&'static str, ChatModelCapabilities>> =
LazyLock::new(|| {
let mut m = HashMap::new();
for id in [
"gpt-4o",
"gpt-4o-mini",
"gpt-4o-2024-11-20",
"gpt-4o-2024-08-06",
] {
m.insert(
id,
ChatModelCapabilities {
context_window: 128_000,
supports_streaming: true,
supports_json_schema: true,
reasoning: None,
},
);
}
for id in ["o1", "o1-mini", "o1-preview", "o3", "o3-mini", "o4-mini"] {
m.insert(
id,
ChatModelCapabilities {
context_window: 128_000,
supports_streaming: true,
supports_json_schema: true,
reasoning: Some(openai_o_series_reasoning()),
},
);
}
for id in ["gpt-5", "gpt-5-2025-08-07"] {
m.insert(
id,
ChatModelCapabilities {
context_window: 256_000,
supports_streaming: true,
supports_json_schema: true,
reasoning: None,
},
);
}
for id in ["gpt-5.5", "gpt-5.4"] {
m.insert(
id,
ChatModelCapabilities {
context_window: 1_000_000,
supports_streaming: true,
supports_json_schema: true,
reasoning: Some(openai_gpt5_reasoning()),
},
);
}
m.insert(
"gpt-5.4-mini",
ChatModelCapabilities {
context_window: 400_000,
supports_streaming: true,
supports_json_schema: true,
reasoning: Some(openai_gpt5_reasoning()),
},
);
m
});
fn openai_uses_completion_tokens(model: &str) -> bool {
MODEL_CAPABILITIES
.get(model)
.is_some_and(|c| c.reasoning.is_some())
}
#[derive(Deserialize)]
struct OpenAiListModelsResponse {
data: Vec<OpenAiListModelEntry>,
}
#[derive(Deserialize)]
struct OpenAiListModelEntry {
id: String,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn capabilities_known_model_returns_image_support() {
let lm = OpenAiLanguageModel::new(
OpenAiDeps {
client: Arc::new(reqwest::Client::new()),
},
OpenAiConfig {
api_key: ApiKey::parse("test").unwrap(),
base_url: OpenAiConfig::DEFAULT_BASE_URL.to_owned(),
retry_config: RetryConfig::default(),
},
);
let caps = lm.capabilities(&ModelId::new("gpt-4o")).unwrap();
let image = caps.media_support.get(&MediaKind::Image).unwrap();
assert!(image.sources.contains(SourceKind::Url));
assert!(image.sources.contains(SourceKind::InlineBytes));
assert!(!image.sources.contains(SourceKind::ProviderFile));
}
#[test]
fn custom_base_url() {
let lm = OpenAiLanguageModel::new(
OpenAiDeps {
client: Arc::new(reqwest::Client::new()),
},
OpenAiConfig {
base_url: "https://my-vllm.example.com/v1".into(),
api_key: ApiKey::parse("key").unwrap(),
retry_config: RetryConfig::default(),
},
);
assert_eq!(lm.base_url, "https://my-vllm.example.com/v1");
}
#[test]
fn gpt5_family_uses_completion_tokens() {
assert!(openai_uses_completion_tokens("gpt-5.5"));
assert!(openai_uses_completion_tokens("gpt-5.4"));
assert!(openai_uses_completion_tokens("gpt-5.4-mini"));
assert!(openai_uses_completion_tokens("o3-mini"));
assert!(!openai_uses_completion_tokens("gpt-4o"));
assert!(!openai_uses_completion_tokens("some-future-model"));
}
}