use async_trait::async_trait;
use serde::Deserialize;
use serde_json::{Value, json};
use crate::core::Secret;
use super::wire::{RESPOND_TOOL, classify_status, classify_transport, structured};
use super::{
Completion, ModelError, ModelId, ModelProvider, Request, SchemaMode, Usage,
chat_completions_stream, sse,
};
const PROVIDER: &str = "chat-completions";
pub struct ChatCompletions {
http: reqwest::Client,
key: Option<Secret>,
base: String,
default_schema_mode: SchemaMode,
schema_modes: std::collections::BTreeMap<String, SchemaMode>,
stream: bool,
egress: Option<crate::core::Egress>,
timeout: std::time::Duration,
}
impl std::fmt::Debug for ChatCompletions {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ChatCompletions")
.field("base", &self.base)
.field("key", &self.key.as_ref().map(|_| "<redacted>"))
.finish_non_exhaustive()
}
}
impl ChatCompletions {
pub const DEFAULT_TIMEOUT: std::time::Duration = std::time::Duration::from_mins(5);
pub fn new(base: impl Into<String>) -> Result<Self, ModelError> {
let http = reqwest::Client::builder()
.build()
.map_err(|e| ModelError::Unreachable {
model: ModelId::new(PROVIDER, "*"),
detail: format!("could not build an HTTP client: {e}"),
})?;
let mut base = base.into();
while base.ends_with('/') {
base.pop();
}
if !base.ends_with("/v1") {
base.push_str("/v1");
}
Ok(Self {
http,
key: None,
base,
default_schema_mode: SchemaMode::ForcedTool,
schema_modes: std::collections::BTreeMap::new(),
stream: true,
egress: None,
timeout: Self::DEFAULT_TIMEOUT,
})
}
#[must_use]
pub fn bearer(mut self, key: impl Into<String>) -> Self {
self.key = Some(Secret::new(key));
self
}
#[must_use]
pub const fn timeout(mut self, timeout: std::time::Duration) -> Self {
self.timeout = timeout;
self
}
#[must_use]
pub fn structured_via(mut self, mode: SchemaMode) -> Self {
self.default_schema_mode = mode;
self
}
#[must_use]
pub fn structured_via_for(mut self, model: impl Into<String>, mode: SchemaMode) -> Self {
self.schema_modes.insert(model.into(), mode);
self
}
#[must_use]
pub fn egress(mut self, egress: crate::core::Egress) -> Self {
self.egress = Some(egress);
self
}
#[must_use]
pub const fn buffered(mut self) -> Self {
self.stream = false;
self
}
fn check_egress(&self, model: &ModelId) -> Result<(), ModelError> {
let Some(egress) = &self.egress else {
return Ok(());
};
let host = reqwest::Url::parse(&self.base)
.ok()
.and_then(|u| u.host_str().map(ToOwned::to_owned));
egress
.permits(host.as_deref())
.map_err(|e| ModelError::Egress {
model: model.clone(),
detail: e.to_string(),
})
}
fn mode_for(&self, model: &ModelId) -> SchemaMode {
self.schema_modes
.get(&model.model)
.copied()
.unwrap_or(self.default_schema_mode)
}
}
#[derive(Debug, Default, Deserialize)]
struct PromptDetails {
#[serde(default)]
cached_tokens: u64,
}
#[derive(Debug, Default, Deserialize)]
struct ApiUsage {
#[serde(default)]
prompt_tokens: u64,
#[serde(default)]
completion_tokens: u64,
#[serde(default)]
prompt_tokens_details: Option<PromptDetails>,
}
impl ApiUsage {
fn normalised(&self) -> Usage {
Usage {
input_tokens: self.prompt_tokens,
output_tokens: self.completion_tokens,
cache_write_tokens: 0,
cache_read_tokens: self
.prompt_tokens_details
.as_ref()
.map_or(0, |d| d.cached_tokens),
minor_units: 0,
}
}
}
#[derive(Debug, Default, Deserialize)]
struct ApiFunction {
#[serde(default)]
name: String,
#[serde(default)]
arguments: String,
}
#[derive(Debug, Default, Deserialize)]
struct ApiToolCall {
#[serde(default)]
id: String,
#[serde(default)]
function: ApiFunction,
}
#[derive(Debug, Default)]
struct ApiMessage {
content: Option<String>,
tool_calls: Vec<ApiToolCall>,
refusal: Option<String>,
raw: Value,
}
impl<'de> Deserialize<'de> for ApiMessage {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
#[derive(Deserialize)]
struct Fields {
#[serde(default)]
content: Option<String>,
#[serde(default)]
tool_calls: Vec<ApiToolCall>,
#[serde(default)]
refusal: Option<String>,
}
let raw = Value::deserialize(deserializer)?;
let fields: Fields =
serde_json::from_value(raw.clone()).map_err(serde::de::Error::custom)?;
Ok(Self {
content: fields.content,
tool_calls: fields.tool_calls,
refusal: fields.refusal,
raw,
})
}
}
#[derive(Debug, Deserialize)]
struct ApiChoice {
#[serde(default)]
message: ApiMessage,
#[serde(default)]
finish_reason: Option<String>,
}
#[derive(Debug, Deserialize)]
struct ApiResponse {
#[serde(default)]
choices: Vec<ApiChoice>,
#[serde(default)]
usage: Option<ApiUsage>,
}
impl ApiResponse {
fn usage(&self) -> Usage {
self.usage
.as_ref()
.map(ApiUsage::normalised)
.unwrap_or_default()
}
}
fn messages(prompt: &Value) -> Vec<Value> {
if let Value::Array(turns) = prompt {
return turns.clone();
}
let mut out = Vec::new();
if let Some(system) = prompt.get("system").filter(|s| !s.is_null()) {
let content = match system {
Value::String(s) => s.clone(),
other => other.to_string(),
};
out.push(json!({ "role": "system", "content": content }));
}
if let Some(Value::Array(turns)) = prompt.get("messages") {
out.extend(turns.iter().cloned());
return out;
}
let user = match prompt {
Value::String(s) => s.clone(),
other => match other.get("input") {
Some(Value::String(s)) => s.clone(),
Some(other) => other.to_string(),
None => {
let mut rest = other.clone();
if let Some(map) = rest.as_object_mut() {
map.remove("system");
}
rest.to_string()
}
},
};
out.push(json!({ "role": "user", "content": user }));
out
}
fn continue_with(
out: &mut Vec<Value>,
exchanges: &[super::ToolExchange],
continuation: Option<&super::ProviderContinuation>,
) {
if exchanges.is_empty() {
return;
}
if let Some(state) = continuation.and_then(|c| c.state.as_array()) {
out.extend(state.iter().cloned());
} else {
out.push(json!({
"role": "assistant",
"tool_calls": exchanges.iter().map(|e| json!({
"id": e.call.id,
"type": "function",
"function": {
"name": e.call.name,
"arguments": e.call.arguments.to_string(),
},
})).collect::<Vec<_>>(),
}));
}
for e in exchanges {
out.push(json!({
"role": "tool",
"tool_call_id": e.call.id,
"content": match &e.output {
Value::String(s) => s.clone(),
other => other.to_string(),
},
}));
}
}
fn tool_messages(exchanges: &[super::ToolExchange]) -> impl Iterator<Item = Value> + '_ {
exchanges.iter().map(|e| {
json!({
"role": "tool",
"tool_call_id": e.call.id,
"content": match &e.output {
Value::String(s) => s.clone(),
other => other.to_string(),
},
})
})
}
impl ChatCompletions {
#[allow(clippy::too_many_arguments)]
fn body(
&self,
model: &ModelId,
prompt: &Value,
max_output_tokens: u32,
schema: Option<&Value>,
tools: &[super::ToolDeclaration],
exchanges: &[super::ToolExchange],
continuation: Option<&super::ProviderContinuation>,
) -> Result<Value, ModelError> {
if let Some(state) = continuation
&& (state.provider != PROVIDER || !state.state.is_array())
{
return Err(ModelError::Refused {
model: model.clone(),
detail: "the continuation was not a chat-completions message array".to_owned(),
});
}
if continuation.is_some() && exchanges.is_empty() {
return Err(ModelError::Refused {
model: model.clone(),
detail: "a continuation without tool exchanges has no request to follow".to_owned(),
});
}
let mut msgs = messages(prompt);
continue_with(&mut msgs, exchanges, continuation);
let mut body = json!({
"model": model.model,
"messages": msgs,
"max_tokens": max_output_tokens,
});
if let Some(schema) = schema {
if self.mode_for(model) == SchemaMode::ForcedTool && !tools.is_empty() {
return Err(ModelError::Refused {
model: model.clone(),
detail: format!(
"model '{}' obtains a declared response schema by forcing a \
synthetic tool, which cannot be combined with the {} tool(s) \
this request declares. Use a server with native `json_schema` \
support, or drop the schema and validate the answer yourself",
model.model,
tools.len()
),
});
}
match self.mode_for(model) {
SchemaMode::Native => {
body["response_format"] = json!({
"type": "json_schema",
"json_schema": {
"name": RESPOND_TOOL,
"strict": true,
"schema": schema,
},
});
}
SchemaMode::ForcedTool => {
body["tools"] = json!([{
"type": "function",
"function": {
"name": RESPOND_TOOL,
"description": "Return the answer in the required shape.",
"parameters": schema,
},
}]);
body["tool_choice"] =
json!({ "type": "function", "function": { "name": RESPOND_TOOL } });
}
}
}
if !tools.is_empty() {
body["tools"] = Value::Array(
tools
.iter()
.map(|t| {
json!({
"type": "function",
"function": {
"name": t.name,
"description": t.description,
"parameters": t.parameters,
},
})
})
.collect(),
);
}
if self.stream {
body["stream"] = json!(true);
body["stream_options"] = json!({ "include_usage": true });
}
Ok(body)
}
fn interpret(
&self,
parsed: &ApiResponse,
model: &ModelId,
schema: Option<&Value>,
) -> Result<Completion, ModelError> {
let usage = parsed.usage();
let Some(choice) = parsed.choices.first() else {
return Err(ModelError::Unusable {
model: model.clone(),
usage,
detail: "the response carried no choices".to_owned(),
});
};
if let Some(why) = choice.message.refusal.as_deref().filter(|r| !r.is_empty()) {
return Err(ModelError::Unusable {
model: model.clone(),
usage,
detail: format!("the model declined to answer: {why}"),
});
}
let text = choice.message.content.clone().unwrap_or_default();
let truncated = choice.finish_reason.as_deref() == Some("length");
let emulating = schema.is_some() && self.mode_for(model) == SchemaMode::ForcedTool;
let mut calls = Vec::new();
let mut forced: Option<String> = None;
for c in &choice.message.tool_calls {
if c.function.name == RESPOND_TOOL {
forced = Some(c.function.arguments.clone());
continue;
}
let arguments =
serde_json::from_str(&c.function.arguments).map_err(|e| ModelError::Unusable {
model: model.clone(),
usage,
detail: format!(
"tool call '{}' for '{}' carried malformed JSON arguments: {e}",
c.id, c.function.name
),
})?;
calls.push(super::ToolCall {
id: c.id.clone(),
name: c.function.name.clone(),
arguments,
});
}
if text.is_empty() && calls.is_empty() && !truncated && !emulating {
return Err(ModelError::Unusable {
model: model.clone(),
usage,
detail: format!(
"the answer carried no content (finish_reason {:?})",
choice.finish_reason
),
});
}
let (text, structured_value) = if emulating {
let Some(raw) = forced else {
return Err(ModelError::Unusable {
model: model.clone(),
usage,
detail: "a tool call was forced and the answer carried none — \
the model did not honour `tool_choice`"
.to_owned(),
});
};
let parsed_schema = structured(schema, &raw, &[], model, usage)?;
(raw, parsed_schema)
} else {
let parsed_schema = structured(schema, &text, &calls, model, usage)?;
(text, parsed_schema)
};
let continuation = (!calls.is_empty()).then(|| {
let mut message = choice.message.raw.clone();
match message.as_object_mut() {
Some(object) => {
object.insert("role".to_owned(), json!("assistant"));
}
None => {
message = json!({
"role": "assistant",
"content": choice.message.content,
"tool_calls": choice.message.tool_calls.iter().map(|c| json!({
"id": c.id,
"type": "function",
"function": {
"name": c.function.name,
"arguments": c.function.arguments,
},
})).collect::<Vec<_>>(),
});
}
}
super::ProviderContinuation::new(PROVIDER, json!([message]))
});
Ok(Completion {
structured: structured_value,
tool_calls: calls,
text,
usage,
stop_reason: choice.finish_reason.clone(),
truncated,
continuation,
})
}
async fn read_buffered(
&self,
response: reqwest::Response,
model: &ModelId,
schema: Option<&Value>,
) -> Result<Completion, ModelError> {
let body = crate::netguard::intake::read(response, crate::netguard::intake::ANSWER)
.await
.map_err(|e| super::wire::classify_intake(model, Usage::default(), &e))?;
let parsed: ApiResponse =
serde_json::from_slice(&body).map_err(|e| ModelError::Unusable {
model: model.clone(),
usage: Usage::default(),
detail: format!("the response body did not parse: {e}"),
})?;
self.interpret(&parsed, model, schema)
}
async fn read_streamed(
&self,
response: reqwest::Response,
model: &ModelId,
schema: Option<&Value>,
observer: Option<(&dyn super::ModelStreamObserver, &crate::core::Label)>,
) -> Result<Completion, ModelError> {
use futures_util::StreamExt;
let mut decoder = sse::Decoder::new();
let mut acc = chat_completions_stream::Accumulator::new();
let mut body = response.bytes_stream();
let mut meter = crate::netguard::intake::Meter::new(crate::netguard::intake::ANSWER);
while let Some(chunk) = body.next().await {
let chunk = match chunk {
Ok(chunk) => chunk,
Err(e) => return Err(severed(model, &acc, &e.to_string())),
};
if let Err(e) = meter.charge(chunk.len()) {
return Err(super::wire::classify_intake(model, Usage::default(), &e));
}
let events = decoder
.push(&chunk)
.map_err(|error| severed(model, &acc, &error.to_string()))?;
for event in events {
if let Some(delta) = acc.push(&event.data)
&& let Some((observer, label)) = observer
{
observer.event(crate::core::Tainted::with_label(
super::ModelStreamEvent::TextDelta(delta),
label.clone(),
));
}
}
if acc.done() {
break;
}
}
if !acc.done() {
return Err(severed(
model,
&acc,
"the stream ended before its `[DONE]` terminal",
));
}
let parsed: ApiResponse =
serde_json::from_value(acc.into_response()).map_err(|e| ModelError::Unusable {
model: model.clone(),
usage: Usage::default(),
detail: format!("the reassembled stream did not parse: {e}"),
})?;
let completion = self.interpret(&parsed, model, schema)?;
if let Some((observer, label)) = observer {
observer.event(crate::core::Tainted::with_label(
super::ModelStreamEvent::Usage(completion.usage),
label.clone(),
));
}
Ok(completion)
}
}
fn severed(
model: &ModelId,
acc: &chat_completions_stream::Accumulator,
detail: &str,
) -> ModelError {
if acc.generated() {
return ModelError::Unaccounted {
model: model.clone(),
detail: detail.to_owned(),
};
}
ModelError::Unavailable {
model: model.clone(),
detail: format!("the stream ended before it generated: {detail}"),
}
}
fn accumulate_continuation(
completion: &mut Completion,
prior: Option<&super::ProviderContinuation>,
exchanges: &[super::ToolExchange],
) {
let Some(current) = completion.continuation.as_mut() else {
return;
};
let mut state = prior
.and_then(|value| value.state.as_array())
.cloned()
.unwrap_or_default();
state.extend(tool_messages(exchanges));
if let Some(items) = current.state.as_array() {
state.extend(items.iter().cloned());
}
current.state = Value::Array(state);
}
#[async_trait]
impl ModelProvider for ChatCompletions {
fn request_profile(&self, model: &ModelId) -> Value {
let schema_mode = match self.mode_for(model) {
SchemaMode::Native => "native",
SchemaMode::ForcedTool => "forced-tool",
};
json!({
"driver": "chat-completions/v1",
"base": self.base,
"schema_mode": schema_mode,
"stream": self.stream,
"timeout_ms": self.timeout.as_millis(),
})
}
async fn complete(&self, request: Request<'_>) -> Result<Completion, ModelError> {
let Request {
model,
prompt,
max_output_tokens,
reasoning_effort,
schema,
tools,
exchanges,
continuation,
stream,
} = request;
super::refuse_provider_side_media(prompt, model)?;
super::refuse_in_thread_instructions(prompt, model)?;
if reasoning_effort.is_some() {
return Err(ModelError::Refused {
model: model.clone(),
detail: "the chat-completions wire has no neutral reasoning-effort \
mapping; configure the model's own default instead"
.to_owned(),
});
}
self.check_egress(model)?;
let body = self.body(
model,
prompt,
max_output_tokens,
schema,
tools,
exchanges,
continuation,
)?;
let mut http = self
.http
.post(format!("{}/chat/completions", self.base))
.timeout(self.timeout)
.json(&body);
if let Some(key) = &self.key {
http = http.bearer_auth(key.expose());
}
let response = http
.send()
.await
.map_err(|e| classify_transport(model, &e))?;
let status = response.status();
if !status.is_success() {
let headers = response.headers().clone();
let text =
crate::netguard::intake::read_text(response, crate::netguard::intake::METADATA)
.await
.unwrap_or_default();
return Err(classify_status(model, status.as_u16(), &headers, &text));
}
let mut completion = if self.stream {
self.read_streamed(response, model, schema, stream).await?
} else {
self.read_buffered(response, model, schema).await?
};
if !self.stream
&& let Some((observer, label)) = stream
{
observer.event(crate::core::Tainted::with_label(
super::ModelStreamEvent::Usage(completion.usage),
label.clone(),
));
}
accumulate_continuation(&mut completion, continuation, exchanges);
Ok(completion)
}
}