use std::time::Duration;
use serde_json::{Value, json};
use crate::tools::ToolSpec;
use crate::{
Capabilities, FinishReason, Provider, ProviderError, Request, Response, SchemaDialect,
Strategy, Usage,
};
const MAX_RETRY_AFTER: Duration = Duration::from_secs(60);
const ALL_STRATEGIES: [Strategy; 4] = [
Strategy::NativeSchema,
Strategy::ToolCall,
Strategy::JsonMode,
Strategy::PromptOnly,
];
#[derive(Debug, Clone)]
pub struct OpenAiCompatible {
client: reqwest::Client,
base_url: String,
model: String,
api_key: Option<String>,
capabilities: Capabilities,
max_transport_retries: u32,
initial_backoff: Duration,
}
impl OpenAiCompatible {
pub fn new(base_url: impl Into<String>, model: impl Into<String>) -> Self {
Self {
client: reqwest::Client::new(),
base_url: base_url.into().trim_end_matches('/').to_string(),
model: model.into(),
api_key: None,
capabilities: Capabilities::new(ALL_STRATEGIES.to_vec(), SchemaDialect::OpenAiStrict),
max_transport_retries: 2,
initial_backoff: Duration::from_millis(500),
}
}
pub fn ollama(model: impl Into<String>) -> Self {
Self::new("http://localhost:11434/v1", model).capabilities(
Capabilities::new(ALL_STRATEGIES.to_vec(), SchemaDialect::Generic)
.native_schema_visible(false),
)
}
pub fn api_key(mut self, key: impl Into<String>) -> Self {
self.api_key = Some(key.into());
self
}
pub fn client(mut self, client: reqwest::Client) -> Self {
self.client = client;
self
}
pub fn capabilities(mut self, capabilities: Capabilities) -> Self {
self.capabilities = capabilities;
self
}
pub fn max_transport_retries(mut self, retries: u32) -> Self {
self.max_transport_retries = retries;
self
}
pub fn initial_backoff(mut self, backoff: Duration) -> Self {
self.initial_backoff = backoff;
self
}
fn body(&self, request: &Request) -> Value {
let mut messages = vec![json!({ "role": "system", "content": request.instructions })];
for demo in &request.demonstrations {
messages.push(json!({ "role": "user", "content": demo.input.to_string() }));
messages.push(json!({ "role": "assistant", "content": demo.output.to_string() }));
}
messages.push(json!({ "role": "user", "content": request.input.to_string() }));
for turn in &request.repair {
messages.push(json!({ "role": "assistant", "content": turn.response }));
messages.push(json!({ "role": "user", "content": turn.feedback }));
}
let mut body = json!({ "model": self.model, "messages": messages });
let strict = self.capabilities.dialect == SchemaDialect::OpenAiStrict;
match request.strategy {
Strategy::NativeSchema => {
body["response_format"] = json!({
"type": "json_schema",
"json_schema": { "name": "output", "schema": request.output_schema, "strict": strict }
});
}
Strategy::ToolCall if !request.tools.is_empty() => {
let tools: Vec<Value> = request
.tools
.iter()
.map(|tool| {
let mut function = json!({
"name": tool.name,
"parameters": tool.parameters,
"strict": strict
});
if let Some(description) = &tool.description {
function["description"] = json!(description);
}
json!({ "type": "function", "function": function })
})
.collect();
body["tools"] = json!(tools);
body["tool_choice"] = json!("required");
}
Strategy::ToolCall => {
body["tools"] = json!([{
"type": "function",
"function": {
"name": "respond",
"description": "Return the result.",
"parameters": request.output_schema,
"strict": strict
}
}]);
body["tool_choice"] =
json!({ "type": "function", "function": { "name": "respond" } });
}
Strategy::JsonMode => body["response_format"] = json!({ "type": "json_object" }),
Strategy::PromptOnly => {}
}
if let Some(t) = request.options.temperature {
body["temperature"] = json!(t);
}
if let Some(m) = request.options.max_tokens {
body["max_tokens"] = json!(m);
}
body
}
async fn post(&self, body: &Value) -> Result<Value, ProviderError> {
let url = format!("{}/chat/completions", self.base_url);
let payload = serde_json::to_vec(body).expect("JSON values always serialize");
let mut attempt = 0;
loop {
let mut builder = self
.client
.post(&url)
.header(reqwest::header::CONTENT_TYPE, "application/json")
.body(payload.clone());
if let Some(key) = &self.api_key {
builder = builder.bearer_auth(key);
}
let (error, retry_after) = match builder.send().await {
Ok(response) if response.status().is_success() => {
let bytes = response
.bytes()
.await
.map_err(|e| ProviderError::Transport(e.to_string()))?;
return serde_json::from_slice(&bytes)
.map_err(|e| ProviderError::InvalidResponse(e.to_string()));
}
Ok(response) => {
let code = response.status().as_u16();
let retry_after = response
.headers()
.get(reqwest::header::RETRY_AFTER)
.and_then(|v| v.to_str().ok())
.and_then(|v| v.trim().parse::<u64>().ok())
.map(|s| Duration::from_secs(s).min(MAX_RETRY_AFTER));
let body = response.text().await.unwrap_or_default();
let error = ProviderError::Status { code, body };
if code != 429 && code < 500 {
return Err(error);
}
(error, retry_after)
}
Err(e) => (ProviderError::Transport(e.to_string()), None),
};
if attempt >= self.max_transport_retries {
return Err(error);
}
let delay = retry_after.unwrap_or(self.initial_backoff * 2u32.pow(attempt));
tokio::time::sleep(delay).await;
attempt += 1;
}
}
fn parse(&self, body: &Value, request: &Request) -> Result<Response, ProviderError> {
let strategy = request.strategy;
let choice = body
.pointer("/choices/0")
.ok_or_else(|| ProviderError::InvalidResponse("no choices in response".into()))?;
let message = &choice["message"];
let refusal = message["refusal"].as_str().filter(|r| !r.is_empty());
let call = message.pointer("/tool_calls/0/function");
let content = match (refusal, call) {
(Some(refusal), _) => refusal.to_string(),
(None, Some(call)) => {
let name = call["name"].as_str().unwrap_or_default();
let arguments = call["arguments"].as_str().unwrap_or("{}");
tool_output(&request.tools, name, arguments)
}
_ => {
let text = message["content"].as_str().unwrap_or_default();
match strategy {
Strategy::ToolCall => {
tool_call_in_text(&request.tools, text).unwrap_or_else(|| text.to_string())
}
_ => text.to_string(),
}
}
};
let finish = match (refusal, choice["finish_reason"].as_str()) {
(Some(_), _) | (_, Some("content_filter")) => FinishReason::Refusal,
(_, Some("length")) => FinishReason::Length,
_ => FinishReason::Stop,
};
Ok(Response {
content,
model: body["model"].as_str().unwrap_or(&self.model).to_string(),
finish,
usage: Usage {
input_tokens: body
.pointer("/usage/prompt_tokens")
.and_then(Value::as_u64)
.unwrap_or(0),
output_tokens: body
.pointer("/usage/completion_tokens")
.and_then(Value::as_u64)
.unwrap_or(0),
},
})
}
}
impl Provider for OpenAiCompatible {
fn capabilities(&self) -> Capabilities {
self.capabilities.clone()
}
async fn complete(&self, request: Request) -> Result<Response, ProviderError> {
let body = self.post(&self.body(&request)).await?;
self.parse(&body, &request)
}
}
fn tool_output(tools: &[ToolSpec], name: &str, arguments: &str) -> String {
match tools.iter().find(|tool| tool.name == name) {
Some(tool) => {
let arguments = serde_json::from_str(arguments).unwrap_or(Value::Null);
tool.output(arguments).to_string()
}
None => arguments.to_string(),
}
}
fn tool_call_in_text(tools: &[ToolSpec], text: &str) -> Option<String> {
let value = crate::json::extract_json(text).ok()?;
let name = value.get("name")?.as_str()?;
if name != "respond" && !tools.iter().any(|tool| tool.name == name) {
return None;
}
let arguments = match value.get("arguments")? {
Value::String(arguments) => arguments.clone(),
arguments @ Value::Object(_) => arguments.to_string(),
_ => return None,
};
Some(tool_output(tools, name, &arguments))
}