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, strict_schema_problem, structured,
};
use super::{
Completion, ModelError, ModelId, ModelProvider, Request, SchemaMode, Usage, openai_stream, sse,
};
pub struct OpenAi {
http: reqwest::Client,
key: Secret,
base: String,
max_output_tokens: u32,
default_schema_mode: SchemaMode,
schema_modes: std::collections::BTreeMap<String, SchemaMode>,
stream: bool,
egress: Option<crate::core::Egress>,
}
impl std::fmt::Debug for OpenAi {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OpenAi")
.field("base", &self.base)
.field("max_output_tokens", &self.max_output_tokens)
.field("key", &"<redacted>")
.finish_non_exhaustive()
}
}
impl OpenAi {
pub const DEFAULT_MAX_OUTPUT_TOKENS: u32 = 4096;
pub fn new(key: impl Into<String>) -> Result<Self, ModelError> {
let http = reqwest::Client::builder()
.build()
.map_err(|e| ModelError::Unreachable {
model: ModelId::new("openai", "*"),
detail: format!("could not build an HTTP client: {e}"),
})?;
Ok(Self {
http,
key: Secret::new(key),
base: "https://api.openai.com".to_owned(),
max_output_tokens: Self::DEFAULT_MAX_OUTPUT_TOKENS,
default_schema_mode: SchemaMode::Native,
schema_modes: std::collections::BTreeMap::new(),
stream: true,
egress: None,
})
}
#[must_use]
pub fn base(mut self, base: impl Into<String>) -> Self {
self.base = base.into();
self
}
#[must_use]
pub const fn max_output_tokens(mut self, n: u32) -> Self {
self.max_output_tokens = n;
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::Refused {
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)
}
fn apply_schema(
body: &mut Value,
schema: &Value,
model: &ModelId,
mode: SchemaMode,
) -> Result<(), ModelError> {
if let Some(problem) = strict_schema_problem(schema) {
return Err(ModelError::Refused {
model: model.clone(),
detail: format!(
"the schema cannot be used with strict constrained decoding: {problem}"
),
});
}
match mode {
SchemaMode::Native => {
body["text"] = json!({
"format": {
"type": "json_schema",
"name": RESPOND_TOOL,
"strict": true,
"schema": schema,
}
});
}
SchemaMode::ForcedTool => {
body["tools"] = json!([{
"type": "function",
"name": RESPOND_TOOL,
"description": "Return the answer in the required shape.",
"strict": true,
"parameters": schema,
}]);
body["tool_choice"] = json!({ "type": "function", "name": RESPOND_TOOL });
}
}
Ok(())
}
}
#[derive(Debug, Default, Deserialize)]
struct TokenDetails {
#[serde(default)]
reasoning_tokens: u64,
}
#[derive(Debug, Default, Deserialize)]
struct InputDetails {
#[serde(default)]
cached_tokens: u64,
}
#[derive(Debug, Deserialize)]
struct ApiUsage {
#[serde(default)]
input_tokens: u64,
#[serde(default)]
output_tokens: u64,
#[serde(default)]
input_tokens_details: Option<InputDetails>,
#[serde(default)]
output_tokens_details: Option<TokenDetails>,
}
#[derive(Debug, Deserialize)]
struct ContentPart {
#[serde(rename = "type")]
kind: String,
#[serde(default)]
text: String,
#[serde(default)]
refusal: String,
}
#[derive(Debug, Deserialize)]
struct OutputItem {
#[serde(rename = "type", default)]
kind: String,
#[serde(default)]
content: Vec<ContentPart>,
#[serde(default)]
arguments: Option<String>,
#[serde(default)]
name: Option<String>,
}
#[derive(Debug, Deserialize)]
struct Incomplete {
#[serde(default)]
reason: String,
}
#[derive(Debug, Deserialize)]
struct ApiResponse {
#[serde(default)]
status: String,
#[serde(default)]
output: Vec<OutputItem>,
#[serde(default)]
usage: Option<ApiUsage>,
#[serde(default)]
incomplete_details: Option<Incomplete>,
#[serde(default)]
error: Option<Value>,
}
impl ApiResponse {
fn usage(&self) -> Usage {
let u = self.usage.as_ref();
Usage {
input_tokens: u.map_or(0, |u| u.input_tokens),
output_tokens: u.map_or(0, |u| u.output_tokens),
cache_write_tokens: 0,
cache_read_tokens: u
.and_then(|u| u.input_tokens_details.as_ref())
.map_or(0, |d| d.cached_tokens),
minor_units: 0,
}
}
fn reasoning_tokens(&self) -> u64 {
self.usage
.as_ref()
.and_then(|u| u.output_tokens_details.as_ref())
.map_or(0, |d| d.reasoning_tokens)
}
fn text(&self) -> String {
self.output
.iter()
.flat_map(|i| i.content.iter())
.filter(|c| c.kind == "output_text")
.map(|c| c.text.as_str())
.collect::<Vec<_>>()
.join("")
}
fn forced_tool_arguments(&self) -> Option<&str> {
self.output
.iter()
.find(|i| i.kind == "function_call" && i.name.as_deref() == Some(RESPOND_TOOL))
.and_then(|i| i.arguments.as_deref())
}
fn refusal(&self) -> Option<&str> {
self.output
.iter()
.flat_map(|i| i.content.iter())
.find(|c| c.kind == "refusal")
.map(|c| c.refusal.as_str())
}
}
fn input(prompt: &Value) -> Value {
match prompt {
Value::String(s) => json!(s),
Value::Array(_) => prompt.clone(),
other => other.get("input").cloned().unwrap_or_else(|| {
let mut rest = other.clone();
if let Some(map) = rest.as_object_mut() {
map.remove("system");
}
json!(rest.to_string())
}),
}
}
fn instructions(prompt: &Value) -> Option<Value> {
prompt.get("system").cloned().filter(|s| !s.is_null())
}
impl OpenAi {
fn body(
&self,
model: &ModelId,
prompt: &Value,
schema: Option<&Value>,
) -> Result<Value, ModelError> {
let mut body = json!({
"model": model.model,
"max_output_tokens": self.max_output_tokens,
"input": input(prompt),
});
if let Some(system) = instructions(prompt) {
body["instructions"] = system;
}
if let Some(schema) = schema {
Self::apply_schema(&mut body, schema, model, self.mode_for(model))?;
}
if self.stream {
body["stream"] = json!(true);
}
Ok(body)
}
fn interpret(
&self,
parsed: &ApiResponse,
model: &ModelId,
schema: Option<&Value>,
) -> Result<Completion, ModelError> {
let usage = parsed.usage();
if let Some(why) = parsed.refusal() {
return Err(ModelError::Unusable {
model: model.clone(),
usage,
detail: format!("the model declined to answer: {why}"),
});
}
if parsed.status == "failed" {
return Err(ModelError::Unusable {
model: model.clone(),
usage,
detail: parsed
.error
.as_ref()
.map_or_else(|| "the response failed".to_owned(), ToString::to_string),
});
}
let text = parsed.text();
let truncated = parsed.status == "incomplete";
let emulating = schema.is_some() && self.mode_for(model) == SchemaMode::ForcedTool;
if text.is_empty() && !truncated && !emulating {
return Err(ModelError::Unusable {
model: model.clone(),
usage,
detail: format!(
"the answer carried no text content (status '{}', {} reasoning token(s))",
parsed.status,
parsed.reasoning_tokens()
),
});
}
let (text, structured_value) = if emulating {
let Some(raw) = parsed.forced_tool_arguments() 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(),
});
};
(raw.to_owned(), structured(schema, raw, model, usage)?)
} else {
let parsed_schema = structured(schema, &text, model, usage)?;
(text, parsed_schema)
};
Ok(Completion {
structured: structured_value,
text,
usage,
stop_reason: Some(parsed.incomplete_details.as_ref().map_or_else(
|| parsed.status.clone(),
|i| format!("incomplete:{}", i.reason),
)),
truncated,
})
}
async fn read_buffered(
&self,
response: reqwest::Response,
model: &ModelId,
schema: Option<&Value>,
) -> Result<Completion, ModelError> {
let parsed: ApiResponse = response.json().await.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>,
) -> Result<Completion, ModelError> {
use futures_util::StreamExt;
let mut decoder = sse::Decoder::new();
let mut acc = openai_stream::Accumulator::new();
let mut body = response.bytes_stream();
while let Some(chunk) = body.next().await {
let chunk = match chunk {
Ok(chunk) => chunk,
Err(e) => return Err(severed(model, &acc, &e.to_string())),
};
for event in decoder.push(&chunk) {
acc.event(&event.name, &event.data);
}
if let Some(message) = acc.error() {
return Err(severed(model, &acc, message));
}
if acc.outcome().is_some() {
break;
}
}
let Some(terminal) = acc.terminal() else {
return Err(severed(
model,
&acc,
"the stream ended before a terminal `response.*` event",
));
};
let parsed: ApiResponse =
serde_json::from_value(terminal.clone()).map_err(|e| ModelError::Unusable {
model: model.clone(),
usage: Usage::default(),
detail: format!("the terminal stream event did not parse: {e}"),
})?;
self.interpret(&parsed, model, schema)
}
}
fn severed(model: &ModelId, acc: &openai_stream::Accumulator, detail: &str) -> ModelError {
if acc.generated() {
return ModelError::Unaccounted {
model: model.clone(),
detail: match acc.id() {
Some(id) => {
format!("{detail} (response '{id}' can be read back to account for it)")
}
None => detail.to_owned(),
},
};
}
ModelError::Unavailable {
model: model.clone(),
detail: format!("the stream ended before it generated: {detail}"),
}
}
#[async_trait]
impl ModelProvider for OpenAi {
async fn complete(&self, request: Request<'_>) -> Result<Completion, ModelError> {
let Request {
model,
prompt,
schema,
} = request;
self.check_egress(model)?;
let response = self
.http
.post(format!("{}/v1/responses", self.base))
.bearer_auth(self.key.expose())
.json(&self.body(model, prompt, schema)?)
.send()
.await
.map_err(|e| classify_transport(model, &e))?;
let status = response.status();
if !status.is_success() {
let detail = response.text().await.unwrap_or_default();
return Err(classify_status(model, status.as_u16(), &detail));
}
if self.stream {
self.read_streamed(response, model, schema).await
} else {
self.read_buffered(response, model, schema).await
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn driver() -> OpenAi {
OpenAi::new("test-key").expect("build the driver")
}
#[test]
fn a_system_instruction_becomes_the_instructions_field() {
let body = driver()
.body(
&ModelId::new("openai", "gpt-x"),
&json!({ "system": "answer only in French", "input": "hi" }),
None,
)
.expect("body");
assert_eq!(
body["instructions"], "answer only in French",
"the system instruction must become `instructions`: {body}"
);
assert_eq!(body["input"], "hi", "the question must survive: {body}");
}
#[test]
fn a_system_instruction_is_not_shown_as_the_question() {
let body = driver()
.body(
&ModelId::new("openai", "gpt-x"),
&json!({ "system": "be terse", "ticket": "printer on fire" }),
None,
)
.expect("body");
let asked = body["input"].as_str().unwrap_or_default();
assert!(
!asked.contains("be terse"),
"the instruction leaked into the question: {asked}"
);
assert!(
asked.contains("printer on fire"),
"the actual content went missing: {asked}"
);
}
#[test]
fn a_prompt_without_a_system_sends_no_instructions() {
let body = driver()
.body(&ModelId::new("openai", "gpt-x"), &json!("hi"), None)
.expect("body");
assert!(
body.get("instructions").is_none(),
"an unset instruction must not become an empty one: {body}"
);
}
}