use crate::{
completion::{CompletionError, CompletionModel, CompletionRequest, CompletionResponse, ModelChoice, Usage},
embeddings::{Embedding, EmbeddingError, EmbeddingModel as EmbeddingModelTrait},
http::{HttpClient, HttpRequest},
message::{AssistantContent, Message, ToolCall, UserContent},
tool::ToolDefinition,
};
use serde::{Deserialize, Serialize};
pub const GPT_5_6: &str = "gpt-5.6-sol";
pub const GPT_5_6_TERRA: &str = "gpt-5.6-terra";
pub const GPT_5_6_LUNA: &str = "gpt-5.6-luna";
pub const GPT_5_6_CYBER: &str = "gpt-5.6-cyber";
pub const GPT_5_3_CODEX: &str = "gpt-5.3-codex";
pub const GPT_5: &str = "gpt-5";
pub const GPT_5_MINI: &str = "gpt-5-mini";
pub const GPT_5_NANO: &str = "gpt-5-nano";
pub const GPT_4_1: &str = "gpt-4.1";
pub const GPT_4_1_MINI: &str = "gpt-4.1-mini";
pub const GPT_4_1_NANO: &str = "gpt-4.1-nano";
pub const GPT_4O: &str = "gpt-4o";
pub const GPT_4O_MINI: &str = "gpt-4o-mini";
pub const O3: &str = "o3";
pub const O3_MINI: &str = "o3-mini";
pub const O4_MINI: &str = "o4-mini";
pub const GPT_4_TURBO: &str = "gpt-4-turbo";
pub const GPT_35_TURBO: &str = "gpt-3.5-turbo";
#[deprecated(note = "retired by OpenAI on 2025-07-28; use O3 instead")]
pub const O1: &str = "o1";
#[deprecated(note = "retired by OpenAI on 2025-10-27; use O4_MINI instead")]
pub const O1_MINI: &str = "o1-mini";
pub const TEXT_EMBEDDING_3_LARGE: &str = "text-embedding-3-large";
pub const TEXT_EMBEDDING_3_SMALL: &str = "text-embedding-3-small";
pub const TEXT_EMBEDDING_ADA_002: &str = "text-embedding-ada-002";
const BASE_URL: &str = "https://api.openai.com/v1";
pub struct Client<H> {
http: H,
api_key: String,
base_url: String,
}
impl<H: HttpClient + Clone> Client<H> {
pub fn new(http: H, api_key: impl Into<String>) -> Self {
Self { http, api_key: api_key.into(), base_url: BASE_URL.to_owned() }
}
pub fn with_base_url(mut self, url: impl Into<String>) -> Self {
self.base_url = url.into();
self
}
pub fn embedding_model(
&self,
model: impl Into<String>,
) -> EmbeddingModel<H> {
let model = model.into();
let ndims = default_ndims(&model);
EmbeddingModel {
http: self.http.clone(),
api_key: self.api_key.clone(),
base_url: self.base_url.clone(),
ndims,
model,
dimensions: None,
}
}
pub fn model(&self, model: impl Into<String>) -> Model<H> {
Model {
http: self.http.clone(),
api_key: self.api_key.clone(),
base_url: self.base_url.clone(),
model: model.into(),
}
}
}
pub struct Model<H> {
http: H,
api_key: String,
base_url: String,
model: String,
}
impl<H: HttpClient> CompletionModel for Model<H> {
type Error = CompletionError;
async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, CompletionError> {
let body = build_request(&self.model, request)?;
let bytes = serde_json::to_vec(&body)?;
let http_req = HttpRequest::new(format!("{}/chat/completions", self.base_url))
.header("Authorization", format!("Bearer {}", self.api_key))
.json_body(bytes);
let resp = self.http.post(http_req).await
.map_err(|e| CompletionError::Http(e.to_string()))?;
if !resp.is_success() {
let message = String::from_utf8_lossy(&resp.body).into_owned();
return Err(CompletionError::Provider { status: resp.status, message });
}
let api_resp: ApiResponse = resp.json()?;
parse_response(api_resp)
}
}
#[derive(Serialize)]
struct ApiRequest {
model: String,
messages: Vec<serde_json::Value>,
#[serde(skip_serializing_if = "Vec::is_empty")]
tools: Vec<ApiTool>,
#[serde(skip_serializing_if = "Option::is_none")]
temperature: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
max_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
reasoning_effort: Option<&'static str>,
}
#[derive(Serialize)]
struct ApiTool {
#[serde(rename = "type")]
kind: &'static str,
function: ApiFunction,
}
#[derive(Serialize)]
struct ApiFunction {
name: String,
description: String,
parameters: serde_json::Value,
}
#[derive(Deserialize)]
struct ApiResponse {
choices: Vec<ApiChoice>,
usage: Option<ApiUsage>,
}
#[derive(Deserialize)]
struct ApiChoice {
message: ApiMessage,
}
#[derive(Deserialize)]
struct ApiMessage {
content: Option<String>,
#[serde(default)]
tool_calls: Vec<ApiToolCall>,
}
#[derive(Deserialize)]
struct ApiToolCall {
id: String,
function: ApiToolCallFunction,
}
#[derive(Deserialize)]
struct ApiToolCallFunction {
name: String,
arguments: String,
}
#[derive(Deserialize)]
struct ApiUsage {
prompt_tokens: u32,
completion_tokens: u32,
}
fn build_request(model: &str, req: CompletionRequest) -> Result<ApiRequest, CompletionError> {
let messages = convert_messages(req.messages)?;
let tools = req.tools.into_iter().map(convert_tool).collect();
Ok(ApiRequest {
model: model.to_owned(),
messages,
tools,
temperature: req.temperature,
max_tokens: req.max_tokens,
reasoning_effort: req.thinking.map(|enabled| if enabled { "high" } else { "minimal" }),
})
}
fn convert_messages(messages: Vec<Message>) -> Result<Vec<serde_json::Value>, CompletionError> {
let mut out = Vec::new();
for msg in messages {
match msg {
Message::System { content } => {
out.push(serde_json::json!({ "role": "system", "content": content }));
}
Message::User { content } => {
let mut text_parts: Vec<serde_json::Value> = Vec::new();
for part in content {
match part {
UserContent::Text(t) => {
text_parts
.push(serde_json::json!({ "type": "text", "text": t.text }));
}
UserContent::ToolResult(r) => {
if !text_parts.is_empty() {
let parts = std::mem::take(&mut text_parts);
out.push(serde_json::json!({ "role": "user", "content": parts }));
}
out.push(serde_json::json!({
"role": "tool",
"tool_call_id": r.call_id,
"content": r.content,
}));
}
}
}
if !text_parts.is_empty() {
out.push(serde_json::json!({ "role": "user", "content": text_parts }));
}
}
Message::Assistant { content } => {
let mut text: Option<String> = None;
let mut tool_calls: Vec<serde_json::Value> = Vec::new();
for part in content {
match part {
AssistantContent::Text(t) => {
text = Some(t.text);
}
AssistantContent::ToolCall(c) => {
let arguments = serde_json::to_string(&c.arguments)?;
tool_calls.push(serde_json::json!({
"id": c.id,
"type": "function",
"function": { "name": c.name, "arguments": arguments },
}));
}
}
}
let mut msg = serde_json::json!({ "role": "assistant", "content": text });
if !tool_calls.is_empty() {
msg["tool_calls"] = serde_json::json!(tool_calls);
}
out.push(msg);
}
}
}
Ok(out)
}
fn convert_tool(def: ToolDefinition) -> ApiTool {
ApiTool {
kind: "function",
function: ApiFunction {
name: def.name,
description: def.description,
parameters: def.parameters,
},
}
}
fn parse_response(resp: ApiResponse) -> Result<CompletionResponse, CompletionError> {
let choice = resp
.choices
.into_iter()
.next()
.ok_or_else(|| CompletionError::Response("no choices in response".into()))?;
let model_choice = if !choice.message.tool_calls.is_empty() {
let calls = choice
.message
.tool_calls
.into_iter()
.map(|c| {
let arguments: serde_json::Value =
serde_json::from_str(&c.function.arguments)?;
Ok(ToolCall { id: c.id, name: c.function.name, arguments })
})
.collect::<Result<Vec<_>, serde_json::Error>>()?;
ModelChoice::ToolCall(calls)
} else {
let text = choice
.message
.content
.ok_or_else(|| CompletionError::Response("no content and no tool_calls".into()))?;
ModelChoice::Message(text)
};
let usage = resp.usage.map(|u| Usage {
prompt_tokens: u.prompt_tokens,
completion_tokens: u.completion_tokens,
});
Ok(CompletionResponse { choice: model_choice, reasoning: None, usage })
}
fn default_ndims(model: &str) -> usize {
match model {
TEXT_EMBEDDING_3_LARGE => 3072,
TEXT_EMBEDDING_3_SMALL | TEXT_EMBEDDING_ADA_002 => 1536,
_ => 0,
}
}
pub struct EmbeddingModel<H> {
http: H,
api_key: String,
base_url: String,
model: String,
ndims: usize,
dimensions: Option<usize>,
}
impl<H: HttpClient + Clone> EmbeddingModel<H> {
pub fn with_dimensions(mut self, dims: usize) -> Self {
self.dimensions = Some(dims);
self.ndims = dims;
self
}
}
#[derive(Serialize)]
struct EmbedRequest<'a> {
model: &'a str,
input: &'a [String],
#[serde(skip_serializing_if = "Option::is_none")]
dimensions: Option<usize>,
}
#[derive(Deserialize)]
struct EmbedResponse {
data: Vec<EmbedData>,
}
#[derive(Deserialize)]
struct EmbedData {
embedding: Vec<f64>,
index: usize,
}
impl<H: HttpClient> EmbeddingModelTrait for EmbeddingModel<H> {
const MAX_DOCUMENTS: usize = 2048;
type Error = EmbeddingError;
fn ndims(&self) -> usize {
self.ndims
}
async fn embed_texts(&self, texts: Vec<String>) -> Result<Vec<Embedding>, EmbeddingError> {
let dimensions = if self.model == TEXT_EMBEDDING_ADA_002 {
None
} else {
self.dimensions
};
let body = serde_json::to_vec(&EmbedRequest {
model: &self.model,
input: &texts,
dimensions,
})?;
let req = HttpRequest::new(format!("{}/embeddings", self.base_url))
.header("Authorization", format!("Bearer {}", self.api_key))
.json_body(body);
let resp = self.http.post(req).await
.map_err(|e| EmbeddingError::Http(e.to_string()))?;
if !resp.is_success() {
let message = String::from_utf8_lossy(&resp.body).into_owned();
return Err(EmbeddingError::Provider { status: resp.status, message });
}
let mut api_resp: EmbedResponse = resp.json()?;
if api_resp.data.len() != texts.len() {
return Err(EmbeddingError::Response(format!(
"expected {} embeddings, got {}",
texts.len(),
api_resp.data.len(),
)));
}
api_resp.data.sort_by_key(|d| d.index);
Ok(api_resp
.data
.into_iter()
.zip(texts)
.map(|(data, document)| Embedding { document, vec: data.embedding })
.collect())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn thinking_toggle_serializes_expected_shape() {
let mut on = CompletionRequest::new(vec![Message::user("hi")]);
on.thinking = Some(true);
let json = serde_json::to_value(build_request(GPT_5_6, on).unwrap()).unwrap();
assert_eq!(json["reasoning_effort"], "high");
let mut off = CompletionRequest::new(vec![Message::user("hi")]);
off.thinking = Some(false);
let json = serde_json::to_value(build_request(GPT_5_6, off).unwrap()).unwrap();
assert_eq!(json["reasoning_effort"], "minimal");
let unset = CompletionRequest::new(vec![Message::user("hi")]);
let json = serde_json::to_value(build_request(GPT_5_6, unset).unwrap()).unwrap();
assert!(json.get("reasoning_effort").is_none());
}
}