use crate::error::{Result, ShimError};
use crate::provider::{Provider, ProviderRequest};
use crate::vision;
use serde_json::{json, Value};
pub struct OpenAiCompatible {
pub name: String,
pub base_url: String,
pub api_key: Option<String>,
}
impl OpenAiCompatible {
pub fn new(
name: impl Into<String>,
base_url: impl Into<String>,
api_key: Option<String>,
) -> Self {
Self {
name: name.into(),
base_url: base_url.into(),
api_key,
}
}
}
fn sanitize_messages(messages: &[Value]) -> Vec<Value> {
messages
.iter()
.map(|msg| {
let mut out = msg.clone();
if let Some(obj) = out.as_object_mut() {
obj.remove("reasoning_content"); obj.remove("reasoning_signature"); obj.remove("redacted_reasoning_content"); obj.remove("annotations");
obj.remove("refusal");
}
if let Some(content) = out.get("content").cloned() {
if content.is_array() {
let translated =
vision::translate_content_blocks(&content, vision::to_openai_chat);
out["content"] = vision::text_blocks_to_chat(&translated);
}
}
out
})
.collect()
}
fn normalize_reasoning(obj: &mut serde_json::Map<String, Value>) {
if obj.contains_key("reasoning_content") {
return;
}
if let Some(r) = obj
.get("reasoning")
.and_then(|r| r.as_str())
.filter(|s| !s.is_empty())
.map(str::to_string)
{
obj.insert("reasoning_content".to_string(), json!(r));
}
}
impl Provider for OpenAiCompatible {
fn name(&self) -> &str {
&self.name
}
fn transform_request(&self, model: &str, request: &Value) -> Result<ProviderRequest> {
let obj = request.as_object().ok_or(ShimError::MissingModel)?;
let messages = obj
.get("messages")
.and_then(|m| m.as_array())
.ok_or(ShimError::MissingModel)?;
let mut body = json!({
"model": model,
"messages": sanitize_messages(messages),
});
let body_obj = body.as_object_mut().unwrap();
for key in [
"max_tokens",
"max_completion_tokens",
"temperature",
"top_p",
"frequency_penalty",
"presence_penalty",
"stop",
"seed",
"stream",
"stream_options",
"tools",
"tool_choice",
"parallel_tool_calls",
"response_format",
"logprobs",
"top_logprobs",
"n",
"reasoning_effort",
] {
if let Some(v) = obj.get(key) {
body_obj.insert(key.to_string(), v.clone());
}
}
let ns = format!("x-{}", self.name);
if let Some(ext) = obj.get(&ns).and_then(|e| e.as_object()) {
for (k, v) in ext {
body_obj.insert(k.clone(), v.clone());
}
}
let mut headers = vec![("Content-Type".to_string(), "application/json".to_string())];
if let Some(key) = &self.api_key {
if !key.is_empty() {
headers.push(("Authorization".to_string(), format!("Bearer {key}")));
}
}
let url = format!("{}/chat/completions", self.base_url.trim_end_matches('/'));
Ok(ProviderRequest { url, headers, body })
}
fn transform_response(&self, _model: &str, mut response: Value) -> Result<Value> {
if let Some(err) = response.get("error") {
if !err.is_null() {
let message = err
.get("message")
.and_then(|m| m.as_str())
.unwrap_or("unknown error")
.to_string();
let status = err.get("code").and_then(|c| c.as_u64()).unwrap_or(400) as u16;
return Err(ShimError::ProviderError {
status,
body: message,
});
}
}
if let Some(choices) = response.get_mut("choices").and_then(|c| c.as_array_mut()) {
for choice in choices {
if let Some(msg) = choice.get_mut("message").and_then(|m| m.as_object_mut()) {
normalize_reasoning(msg);
}
}
}
Ok(response)
}
fn transform_stream_chunk(&self, _model: &str, chunk: &str) -> Result<Option<String>> {
let mut parsed: Value = match serde_json::from_str(chunk) {
Ok(v) => v,
Err(_) => return Ok(None),
};
if let Some(choices) = parsed.get_mut("choices").and_then(|c| c.as_array_mut()) {
for choice in choices {
if let Some(delta) = choice.get_mut("delta").and_then(|d| d.as_object_mut()) {
normalize_reasoning(delta);
}
}
}
Ok(Some(serde_json::to_string(&parsed)?))
}
}