use std::pin::Pin;
use futures::Stream;
use reqwest::{Client, RequestBuilder};
use serde::Deserialize;
use tracing::{debug, error, info, instrument};
use super::openai::{ChatRequestOptions, build_chat_request, parse_chat_completions_sse};
use crate::{
config::{Config, EndpointCapabilities},
error::{Capability, Error, ProviderKind, Result},
message::{Message, Prompt, Response, ToolCall, ToolDefinition, Usage},
model::OpenAICompatibleModel,
};
pub struct OpenAICompatibleProvider {
client: Client,
base_url: String,
api_key: Option<String>,
capabilities: EndpointCapabilities,
}
impl std::fmt::Debug for OpenAICompatibleProvider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OpenAICompatibleProvider")
.field("base_url", &self.base_url)
.field(
"api_key",
match &self.api_key {
Some(_) => &"[REDACTED]",
None => &"[UNSET]",
},
)
.field("capabilities", &self.capabilities)
.finish()
}
}
impl OpenAICompatibleProvider {
pub fn new(config: &Config) -> Result<Self> {
let base_url = config
.openai_compatible_base_url()
.ok_or(Error::ProviderNotConfigured(ProviderKind::OpenAICompatible))?;
let client = super::http_client_builder()
.timeout(std::time::Duration::from_secs(config.timeout()))
.build()
.map_err(|e| Error::Config(format!("Failed to create HTTP client: {e}")))?;
Ok(Self {
client,
base_url,
api_key: config.openai_compatible_key(),
capabilities: config.openai_compatible_capabilities(),
})
}
pub fn base_url(&self) -> &str {
&self.base_url
}
pub fn capabilities(&self) -> EndpointCapabilities {
self.capabilities
}
#[instrument(skip(self, prompt, config))]
pub async fn generate(
&self,
model: &OpenAICompatibleModel,
prompt: &Prompt,
config: &crate::generation::GenerationConfig,
tool_definitions: Option<&[ToolDefinition]>,
) -> Result<Response> {
info!(
model = model.as_str(),
base_url = %self.base_url,
"Generating completion with an OpenAI-compatible endpoint"
);
let requested = RequestedCapabilities::of(config, tool_definitions);
self.check_declared_capabilities(requested)?;
let request = self.build_request(model, prompt, config, false, tool_definitions);
let url = self.chat_completions_url();
debug!(url = %url, "Sending request to the OpenAI-compatible endpoint");
debug!(request_payload = ?request, "OpenAI-compatible request payload");
let response = self.authorized(&url).json(&request).send().await?;
let status = response.status();
if !status.is_success() {
let error_body = response.text().await.unwrap_or_default();
error!(status = %status, body = %error_body, "OpenAI-compatible request failed");
return Err(self.parse_error(status.as_u16(), &error_body, requested));
}
let response_body: ChatCompletionsResponse = response.json().await?;
debug!(response_payload = ?response_body, "OpenAI-compatible response payload");
self.parse_response(model, response_body)
}
#[instrument(skip(self, prompt, config))]
pub async fn generate_stream(
&self,
model: &OpenAICompatibleModel,
prompt: &Prompt,
config: &crate::generation::GenerationConfig,
) -> Result<Pin<Box<dyn Stream<Item = Result<crate::provider::ProviderStreamEvent>> + Send>>>
{
let requested = RequestedCapabilities::of(config, None);
self.check_declared_capabilities(requested)?;
let request = self.build_request(model, prompt, config, true, None);
let url = self.chat_completions_url();
debug!(url = %url, "Sending streaming request to the OpenAI-compatible endpoint");
let response = self.authorized(&url).json(&request).send().await?;
let status = response.status();
if !status.is_success() {
let error_body = response.text().await.unwrap_or_default();
error!(status = %status, body = %error_body, "OpenAI-compatible streaming request failed");
return Err(self.parse_error(status.as_u16(), &error_body, requested));
}
let byte_stream = response.bytes_stream();
Ok(Box::pin(parse_chat_completions_sse(byte_stream)))
}
fn chat_completions_url(&self) -> String {
format!("{}/chat/completions", self.base_url)
}
fn authorized(&self, url: &str) -> RequestBuilder {
let request = self
.client
.post(url)
.header("Content-Type", "application/json");
match &self.api_key {
Some(key) => request.header("Authorization", format!("Bearer {key}")),
None => request,
}
}
fn build_request(
&self,
model: &OpenAICompatibleModel,
prompt: &Prompt,
config: &crate::generation::GenerationConfig,
stream: bool,
tool_definitions: Option<&[ToolDefinition]>,
) -> super::openai::OpenAIRequest {
build_chat_request(
prompt,
config,
ChatRequestOptions {
model: model.as_str(),
stream,
tool_definitions,
drop_sampling_params: false,
strict_json_schema: false,
},
)
}
fn check_declared_capabilities(&self, requested: RequestedCapabilities) -> Result<()> {
if requested.tool_calling && !self.capabilities.tool_calling {
return Err(self.capability_unsupported(
Capability::ToolCalling,
"the endpoint was configured as not supporting tool calling",
));
}
if requested.structured_output && !self.capabilities.structured_output {
return Err(self.capability_unsupported(
Capability::StructuredOutput,
"the endpoint was configured as not supporting structured output",
));
}
Ok(())
}
fn capability_unsupported(&self, capability: Capability, message: &str) -> Error {
Error::CapabilityUnsupported {
provider: ProviderKind::OpenAICompatible,
capability,
base_url: self.base_url.clone(),
message: message.to_string(),
}
}
fn parse_response(
&self,
model: &OpenAICompatibleModel,
response: ChatCompletionsResponse,
) -> Result<Response> {
let choice = response.choices.first().ok_or_else(|| Error::Request {
provider: ProviderKind::OpenAICompatible,
message: "No choices in response".to_string(),
})?;
let content = choice.message.content.clone().unwrap_or_default();
let tool_calls = choice
.message
.tool_calls
.clone()
.unwrap_or_default()
.into_iter()
.map(|tc| {
let arguments = serde_json::from_str(&tc.function.arguments).map_err(|error| {
Error::Request {
provider: ProviderKind::OpenAICompatible,
message: format!(
"Invalid tool arguments returned for '{}': {error}",
tc.function.name
),
}
})?;
Ok(ToolCall {
id: tc.id,
name: tc.function.name,
arguments,
})
})
.collect::<Result<Vec<_>>>()?;
let usage = response.usage.map(|u| Usage {
prompt_tokens: Some(u.prompt_tokens),
completion_tokens: Some(u.completion_tokens),
total_tokens: Some(u.total_tokens),
});
info!(
model = model.as_str(),
finish_reason = ?choice.finish_reason,
"Received response from the OpenAI-compatible endpoint"
);
Ok(Response {
messages: vec![Message::assistant_with_tool_calls(content, tool_calls)],
usage,
model: model.as_str().to_string(),
provider: ProviderKind::OpenAICompatible,
finish_reason: choice.finish_reason.clone(),
})
}
fn parse_error(&self, status: u16, body: &str, requested: RequestedCapabilities) -> Error {
let message = super::openai::parse_error_message(body);
if let Some(message) = &message {
if CAPABILITY_REJECTION_STATUSES.contains(&status) {
if requested.tool_calling && mentions_unsupported(message, TOOL_SUBJECTS) {
return self.capability_unsupported(Capability::ToolCalling, message);
}
if requested.structured_output
&& mentions_unsupported(message, STRUCTURED_OUTPUT_SUBJECTS)
{
return self.capability_unsupported(Capability::StructuredOutput, message);
}
}
}
let Some(message) = message else {
return Error::Request {
provider: ProviderKind::OpenAICompatible,
message: format!("HTTP {status}: {body}"),
};
};
match status {
401 | 403 => Error::Auth {
provider: ProviderKind::OpenAICompatible,
message,
},
429 => Error::RateLimit {
provider: ProviderKind::OpenAICompatible,
message,
},
400 => Error::InvalidRequest(message),
_ => Error::Request {
provider: ProviderKind::OpenAICompatible,
message,
},
}
}
}
#[derive(Debug, Clone, Copy)]
struct RequestedCapabilities {
tool_calling: bool,
structured_output: bool,
}
impl RequestedCapabilities {
fn of(
config: &crate::generation::GenerationConfig,
tool_definitions: Option<&[ToolDefinition]>,
) -> Self {
Self {
tool_calling: tool_definitions.is_some_and(|tools| !tools.is_empty()),
structured_output: config.json_schema.is_some() || config.json_mode == Some(true),
}
}
}
const CAPABILITY_REJECTION_STATUSES: &[u16] = &[400, 404, 422, 501];
const TOOL_SUBJECTS: &[&str] = &["tool", "function call", "function_call"];
const STRUCTURED_OUTPUT_SUBJECTS: &[&str] = &[
"response_format",
"response format",
"json_schema",
"json schema",
"structured output",
"guided decoding",
"grammar",
];
fn mentions_unsupported(message: &str, subjects: &[&str]) -> bool {
const UNAVAILABLE: &[&str] = &[
"not support",
"unsupported",
"not implemented",
"unrecognized",
"unknown",
"no support",
"does not accept",
];
let message = message.to_ascii_lowercase();
subjects.iter().any(|subject| message.contains(subject))
&& UNAVAILABLE.iter().any(|marker| message.contains(marker))
}
#[derive(Debug, Deserialize)]
struct ChatCompletionsResponse {
choices: Vec<ChatCompletionsChoice>,
usage: Option<ChatCompletionsUsage>,
}
#[derive(Debug, Deserialize)]
struct ChatCompletionsChoice {
message: ChatCompletionsMessage,
finish_reason: Option<String>,
}
#[derive(Debug, Deserialize)]
struct ChatCompletionsMessage {
content: Option<String>,
tool_calls: Option<Vec<ChatCompletionsToolCall>>,
}
#[derive(Debug, Clone, Deserialize)]
struct ChatCompletionsToolCall {
id: String,
function: ChatCompletionsFunctionCall,
}
#[derive(Debug, Clone, Deserialize)]
struct ChatCompletionsFunctionCall {
name: String,
arguments: String,
}
#[derive(Debug, Deserialize)]
struct ChatCompletionsUsage {
prompt_tokens: i32,
completion_tokens: i32,
total_tokens: i32,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::generation::GenerationConfig;
fn test_provider(capabilities: EndpointCapabilities) -> OpenAICompatibleProvider {
OpenAICompatibleProvider {
client: Client::new(),
base_url: "http://localhost:11434/v1".to_string(),
api_key: None,
capabilities,
}
}
fn tool() -> ToolDefinition {
ToolDefinition {
name: "get_weather".to_string(),
description: None,
input_schema: serde_json::json!({ "type": "object" }),
}
}
#[test]
fn requested_capabilities_track_what_the_request_uses() {
let plain = RequestedCapabilities::of(&GenerationConfig::new(), None);
assert!(!plain.tool_calling);
assert!(!plain.structured_output);
let empty_tools = RequestedCapabilities::of(&GenerationConfig::new(), Some(&[]));
assert!(!empty_tools.tool_calling);
let with_tools = RequestedCapabilities::of(&GenerationConfig::new(), Some(&[tool()]));
assert!(with_tools.tool_calling);
let json_mode =
RequestedCapabilities::of(&GenerationConfig::new().with_json_mode(true), None);
assert!(json_mode.structured_output);
let schema = RequestedCapabilities::of(
&GenerationConfig::new().with_json_schema(serde_json::json!({ "type": "object" })),
None,
);
assert!(schema.structured_output);
}
#[test]
fn declared_gaps_are_refused_before_any_request() {
let provider = test_provider(EndpointCapabilities::text_only());
let error = provider
.check_declared_capabilities(RequestedCapabilities::of(
&GenerationConfig::new(),
Some(&[tool()]),
))
.expect_err("tool calling should be refused");
assert_eq!(
error.unsupported_capability(),
Some(Capability::ToolCalling)
);
let error = provider
.check_declared_capabilities(RequestedCapabilities::of(
&GenerationConfig::new().with_json_mode(true),
None,
))
.expect_err("structured output should be refused");
assert_eq!(
error.unsupported_capability(),
Some(Capability::StructuredOutput)
);
}
#[test]
fn a_fully_capable_endpoint_refuses_nothing() {
let provider = test_provider(EndpointCapabilities::default());
provider
.check_declared_capabilities(RequestedCapabilities::of(
&GenerationConfig::new().with_json_mode(true),
Some(&[tool()]),
))
.expect("a fully capable endpoint should accept everything");
}
#[test]
fn unsupported_markers_need_both_a_subject_and_a_denial() {
assert!(mentions_unsupported(
"registry/llama3.2:1b does not support tools",
TOOL_SUBJECTS
));
assert!(mentions_unsupported(
"unsupported parameter: response_format",
STRUCTURED_OUTPUT_SUBJECTS
));
assert!(!mentions_unsupported(
"invalid arguments for tool 'get_weather'",
TOOL_SUBJECTS
));
assert!(!mentions_unsupported(
"unknown model 'llama3.1:8b'",
TOOL_SUBJECTS
));
}
#[test]
fn capability_classification_only_applies_to_what_was_sent() {
let provider = test_provider(EndpointCapabilities::default());
let body = serde_json::json!({
"error": { "message": "this model does not support tools" }
})
.to_string();
let with_tools = provider.parse_error(
400,
&body,
RequestedCapabilities {
tool_calling: true,
structured_output: false,
},
);
assert_eq!(
with_tools.unsupported_capability(),
Some(Capability::ToolCalling)
);
let without_tools = provider.parse_error(
400,
&body,
RequestedCapabilities {
tool_calling: false,
structured_output: false,
},
);
assert!(without_tools.unsupported_capability().is_none());
assert!(matches!(without_tools, Error::InvalidRequest(_)));
}
#[test]
fn transport_failures_are_not_mistaken_for_capability_gaps() {
let provider = test_provider(EndpointCapabilities::default());
let body = serde_json::json!({
"error": { "message": "tools are not supported right now" }
})
.to_string();
let error = provider.parse_error(
503,
&body,
RequestedCapabilities {
tool_calling: true,
structured_output: false,
},
);
assert!(error.unsupported_capability().is_none());
assert!(matches!(error, Error::Request { .. }));
}
#[test]
fn non_json_error_bodies_keep_their_status_and_text() {
let provider = test_provider(EndpointCapabilities::default());
let error = provider.parse_error(
502,
"<html>bad gateway</html>",
RequestedCapabilities {
tool_calling: false,
structured_output: false,
},
);
assert!(error.to_string().contains("HTTP 502"));
}
#[test]
fn structured_output_does_not_request_strict_schemas() {
let provider = test_provider(EndpointCapabilities::default());
let prompt = Prompt::single(Message::user("hello"));
let config =
GenerationConfig::new().with_json_schema(serde_json::json!({ "type": "object" }));
let request = provider.build_request(
&OpenAICompatibleModel::new("llama3.1:8b"),
&prompt,
&config,
false,
None,
);
let body = serde_json::to_value(&request).expect("request should serialize");
assert_eq!(body["response_format"]["json_schema"]["strict"], false);
assert_eq!(body["model"], "llama3.1:8b");
}
}