use std::collections::BTreeMap;
use std::env;
use std::sync::Arc;
use futures_util::StreamExt;
use reqwest::Client;
use tracing::{info_span, instrument, Instrument};
use super::types::*;
use crate::providers::error::{CompletionError, ProviderError};
use crate::providers::http::build_http_client;
use crate::providers::zai::client::{
CompletionRequest, CompletionResponse, Message, ToolCall, Usage,
};
const DEFAULT_BASE_URL: &str = "https://api.openai.com/v1";
#[derive(Clone)]
pub struct OpenAIClient {
inner: Arc<OpenAIClientInner>,
}
struct OpenAIClientInner {
http_client: Client,
api_key: String,
base_url: String,
organization: Option<String>,
temperature: Option<f64>,
top_p: Option<f64>,
top_k: Option<u64>,
repetition_penalty: Option<f64>,
max_tokens: Option<u64>,
frequency_penalty: Option<f64>,
presence_penalty: Option<f64>,
parallel_tool_calls: Option<bool>,
enable_thinking: Option<bool>,
}
impl OpenAIClient {
pub fn from_env() -> Result<Self, ProviderError> {
let api_key = env::var("OPENAI_API_KEY")
.map_err(|_| ProviderError::EnvVarNotSet("OPENAI_API_KEY".to_string()))?;
let mut builder = OpenAIClientBuilder::new(&api_key);
if let Ok(org) = env::var("OPENAI_ORGANIZATION") {
builder = builder.organization(&org);
}
Ok(builder.build())
}
pub fn completion_model(&self, model_id: &str) -> OpenAICompletionModel {
OpenAICompletionModel {
client: self.clone(),
model_id: model_id.to_string(),
}
}
}
pub struct OpenAIClientBuilder {
api_key: String,
base_url: String,
organization: Option<String>,
temperature: Option<f64>,
top_p: Option<f64>,
top_k: Option<u64>,
repetition_penalty: Option<f64>,
max_tokens: Option<u64>,
frequency_penalty: Option<f64>,
presence_penalty: Option<f64>,
parallel_tool_calls: Option<bool>,
enable_thinking: Option<bool>,
}
impl OpenAIClientBuilder {
pub fn new(api_key: &str) -> Self {
Self {
api_key: api_key.to_string(),
base_url: DEFAULT_BASE_URL.to_string(),
organization: None,
temperature: None,
top_p: None,
top_k: None,
repetition_penalty: None,
max_tokens: None,
frequency_penalty: None,
presence_penalty: None,
parallel_tool_calls: Some(true),
enable_thinking: None,
}
}
pub fn base_url(mut self, url: &str) -> Self {
self.base_url = url.to_string();
self
}
pub fn organization(mut self, org: &str) -> Self {
self.organization = Some(org.to_string());
self
}
pub fn temperature(mut self, temp: f64) -> Self {
self.temperature = Some(temp.clamp(0.0, 2.0));
self
}
pub fn top_p(mut self, p: f64) -> Self {
self.top_p = Some(p.clamp(0.0, 1.0));
self
}
pub fn top_k(mut self, k: u64) -> Self {
self.top_k = Some(k);
self
}
pub fn repetition_penalty(mut self, penalty: f64) -> Self {
self.repetition_penalty = Some(penalty.clamp(0.5, 2.0));
self
}
pub fn max_tokens(mut self, tokens: u64) -> Self {
self.max_tokens = Some(tokens);
self
}
pub fn frequency_penalty(mut self, penalty: f64) -> Self {
self.frequency_penalty = Some(penalty.clamp(-2.0, 2.0));
self
}
pub fn presence_penalty(mut self, penalty: f64) -> Self {
self.presence_penalty = Some(penalty.clamp(-2.0, 2.0));
self
}
pub fn parallel_tool_calls(mut self, enabled: bool) -> Self {
self.parallel_tool_calls = Some(enabled);
self
}
pub fn enable_thinking(mut self, enabled: bool) -> Self {
self.enable_thinking = Some(enabled);
self
}
pub fn build(self) -> OpenAIClient {
OpenAIClient {
inner: Arc::new(OpenAIClientInner {
http_client: build_http_client(),
api_key: self.api_key,
base_url: self.base_url,
organization: self.organization,
temperature: self.temperature,
top_p: self.top_p,
top_k: self.top_k,
repetition_penalty: self.repetition_penalty,
max_tokens: self.max_tokens,
frequency_penalty: self.frequency_penalty,
presence_penalty: self.presence_penalty,
parallel_tool_calls: self.parallel_tool_calls,
enable_thinking: self.enable_thinking,
}),
}
}
}
#[derive(Clone)]
pub struct OpenAICompletionModel {
client: OpenAIClient,
model_id: String,
}
impl OpenAICompletionModel {
pub fn model_id(&self) -> &str {
&self.model_id
}
pub fn provider(&self) -> &str {
"openai"
}
#[instrument(skip(self, request), fields(model = %self.model_id, provider = "openai"))]
pub async fn completion(
&self,
request: CompletionRequest,
) -> Result<CompletionResponse<OpenAIResponse>, CompletionError> {
let inner = &self.client.inner;
let mut messages = Vec::new();
if let Some(preamble) = &request.preamble {
messages.push(OpenAIMessage {
role: "system".to_string(),
content: Some(preamble.clone()),
tool_calls: None,
tool_call_id: None,
name: None,
reasoning: None,
});
}
for msg in &request.messages {
messages.push(OpenAIMessage {
role: msg.role.clone(),
content: if msg.content.is_empty() {
None
} else {
Some(msg.content.clone())
},
tool_calls: msg.tool_calls.as_ref().map(|calls| {
calls
.iter()
.map(|tc| OpenAIToolCall {
id: tc.id.clone(),
call_type: "function".to_string(),
function: OpenAIFunctionCall {
name: tc.name.clone(),
arguments: tc.arguments.clone(),
},
})
.collect()
}),
tool_call_id: msg.tool_call_id.clone(),
name: None,
reasoning: msg.reasoning.clone(),
});
}
let tools = if request.tools.is_empty() {
None
} else {
Some(
request
.tools
.iter()
.map(|t| OpenAITool {
tool_type: "function".to_string(),
function: OpenAIFunction {
name: t.name.clone(),
description: Some(t.description.clone()),
parameters: Some(normalize_tool_parameters(t.parameters.clone())),
},
})
.collect(),
)
};
let chat_template_kwargs = inner
.enable_thinking
.map(|enabled| serde_json::json!({ "enable_thinking": enabled }));
let openai_request = OpenAIRequest {
model: self.model_id.clone(),
messages,
temperature: request.temperature.or(inner.temperature),
max_tokens: request.max_tokens.or(inner.max_tokens),
top_p: inner.top_p,
top_k: inner.top_k,
repetition_penalty: inner.repetition_penalty,
frequency_penalty: inner.frequency_penalty,
presence_penalty: inner.presence_penalty,
tools,
tool_choice: None,
parallel_tool_calls: inner.parallel_tool_calls,
user: None,
chat_template_kwargs,
stream: Some(true),
stream_options: Some(OpenAIStreamOptions {
include_usage: true,
}),
};
let url = format!("{}/chat/completions", inner.base_url);
if std::env::var("SAC_DUMP_REQUESTS").is_ok() {
if let Ok(body_json) = serde_json::to_string_pretty(&openai_request) {
let dir = std::env::var("SAC_DUMP_REQUESTS_DIR")
.unwrap_or_else(|_| "/tmp/lms".to_string());
let _ = std::fs::create_dir_all(&dir);
let ts = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_millis())
.unwrap_or(0);
let path = format!("{}/sac_request_{}.json", dir, ts);
let _ = std::fs::write(&path, body_json.as_bytes());
tracing::info!("SAC REQUEST DUMP → {} ({} bytes)", path, body_json.len());
}
}
let mut req = inner
.http_client
.post(&url)
.header("Authorization", format!("Bearer {}", inner.api_key))
.header("Content-Type", "application/json")
.header("Accept", "text/event-stream");
if let Some(org) = &inner.organization {
req = req.header("OpenAI-Organization", org);
}
let response = req
.json(&openai_request)
.send()
.instrument(info_span!("openai_http_request"))
.await
.map_err(ProviderError::Request)?;
let status = response.status();
if !status.is_success() {
let error_text = response.text().await.unwrap_or_default();
if status.as_u16() == 429 {
return Err(CompletionError::Provider(ProviderError::RateLimited {
retry_after_ms: None,
}));
}
if status.as_u16() == 401 {
return Err(CompletionError::Provider(ProviderError::Authentication(
error_text,
)));
}
return Err(CompletionError::Provider(ProviderError::Http {
status: status.as_u16(),
message: error_text,
}));
}
let openai_response = collect_openai_sse_stream(response, &self.model_id).await?;
let choice = openai_response.choices.first().ok_or_else(|| {
CompletionError::Provider(ProviderError::InvalidResponse(
"No choices in response".to_string(),
))
})?;
let tool_calls = choice.message.tool_calls.as_ref().map(|calls| {
let mut out: Vec<ToolCall> = Vec::with_capacity(calls.len());
for tc in calls {
let id = tc.id.clone();
let name = tc.function.name.clone();
let args = tc.function.arguments.clone();
let split = split_packed_kimi_args(&id, &name, &args);
out.extend(split);
}
out
});
let reasoning_text = choice.message.reasoning.clone();
let content = choice.message.content.clone().unwrap_or_default();
let message = Message {
role: choice.message.role.clone(),
content,
tool_calls,
tool_call_id: None,
reasoning: reasoning_text.clone(),
};
let cache_read_tokens = openai_response
.usage
.prompt_tokens_details
.as_ref()
.and_then(|d| d.cached_tokens)
.unwrap_or(0);
let finish_reason = choice.finish_reason.clone();
Ok(CompletionResponse {
message,
usage: Usage {
prompt_tokens: openai_response.usage.prompt_tokens,
completion_tokens: openai_response.usage.completion_tokens,
total_tokens: openai_response.usage.total_tokens,
cache_read_tokens,
cache_creation_tokens: 0,
},
raw: openai_response,
reasoning_content: reasoning_text,
finish_reason,
})
}
}
fn normalize_tool_parameters(mut value: serde_json::Value) -> serde_json::Value {
if let Some(obj) = value.as_object_mut() {
let is_object_type = obj
.get("type")
.and_then(|v| v.as_str())
.map(|s| s == "object")
.unwrap_or(false);
if is_object_type && !obj.contains_key("properties") {
obj.insert("properties".to_string(), serde_json::json!({}));
}
}
value
}
fn split_packed_kimi_args(id: &str, name: &str, args: &str) -> Vec<ToolCall> {
const SEP_CORE: &str = "<|tool_call_argument_begin|>";
if !args.contains(SEP_CORE) {
return vec![ToolCall {
id: id.to_string(),
name: name.to_string(),
arguments: args.to_string(),
}];
}
let mut boundaries: Vec<(usize, usize, String)> = Vec::new(); let bytes = args.as_bytes();
let sep_bytes = SEP_CORE.as_bytes();
let mut i = 0;
while i + sep_bytes.len() <= bytes.len() {
if bytes[i..i + sep_bytes.len()] == *sep_bytes {
let mut start = i;
while start > 0 && bytes[start - 1].is_ascii_digit() {
start -= 1;
}
if start == 0 || bytes[start - 1] != b':' {
i += 1;
continue;
}
start -= 1; let name_end = start;
while start > 0
&& (bytes[start - 1].is_ascii_alphanumeric() || bytes[start - 1] == b'_')
{
start -= 1;
}
let prefix_start = if start >= "functions.".len()
&& &args[start - "functions.".len()..start] == "functions."
{
start - "functions.".len()
} else {
start
};
let inner_name = args[start..name_end].to_string();
if inner_name.is_empty() {
i += 1;
continue;
}
boundaries.push((prefix_start, i + sep_bytes.len(), inner_name));
i += sep_bytes.len();
} else {
i += 1;
}
}
if boundaries.is_empty() {
return vec![ToolCall {
id: id.to_string(),
name: name.to_string(),
arguments: args.to_string(),
}];
}
let mut out: Vec<ToolCall> = Vec::with_capacity(boundaries.len() + 1);
let first_args_end = boundaries[0].0;
out.push(ToolCall {
id: id.to_string(),
name: name.to_string(),
arguments: args[..first_args_end].trim().to_string(),
});
for idx in 0..boundaries.len() {
let (_, args_start, inner_name) = &boundaries[idx];
let args_end = if idx + 1 < boundaries.len() {
boundaries[idx + 1].0
} else {
args.len()
};
out.push(ToolCall {
id: format!("{id}_split_{}", idx + 1),
name: inner_name.clone(),
arguments: args[*args_start..args_end].trim().to_string(),
});
}
out
}
struct PartialToolCall {
id: String,
name: String,
arguments: String,
}
async fn collect_openai_sse_stream(
response: reqwest::Response,
model_id: &str,
) -> Result<OpenAIResponse, CompletionError> {
let mut resp_id = String::new();
let mut resp_object = "chat.completion".to_string();
let mut resp_created: u64 = 0;
let mut resp_model = model_id.to_string();
let mut system_fingerprint: Option<String> = None;
let mut role = String::from("assistant");
let mut content = String::new();
let mut reasoning = String::new();
let mut tool_calls: BTreeMap<u64, PartialToolCall> = BTreeMap::new();
let mut finish_reason: Option<String> = None;
let mut usage = OpenAIUsage::default();
let mut stream = response.bytes_stream();
let mut buffer = String::new();
'outer: while let Some(chunk) = stream.next().await {
let bytes = chunk.map_err(ProviderError::Request)?;
buffer.push_str(&String::from_utf8_lossy(&bytes));
while let Some(idx) = buffer.find("\n\n") {
let event_raw = buffer[..idx].to_string();
buffer.drain(..idx + 2);
if process_openai_sse_event(
&event_raw,
&mut resp_id,
&mut resp_object,
&mut resp_created,
&mut resp_model,
&mut system_fingerprint,
&mut role,
&mut content,
&mut reasoning,
&mut tool_calls,
&mut finish_reason,
&mut usage,
)? {
break 'outer;
}
}
}
if !buffer.trim().is_empty() {
let leftover = std::mem::take(&mut buffer);
let _ = process_openai_sse_event(
&leftover,
&mut resp_id,
&mut resp_object,
&mut resp_created,
&mut resp_model,
&mut system_fingerprint,
&mut role,
&mut content,
&mut reasoning,
&mut tool_calls,
&mut finish_reason,
&mut usage,
);
}
let assembled_tool_calls: Vec<OpenAIToolCall> = tool_calls
.into_values()
.filter(|tc| !tc.id.is_empty() || !tc.name.is_empty() || !tc.arguments.is_empty())
.map(|tc| OpenAIToolCall {
id: tc.id,
call_type: "function".to_string(),
function: OpenAIFunctionCall {
name: tc.name,
arguments: tc.arguments,
},
})
.collect();
let message = OpenAIMessage {
role,
content: if content.is_empty() {
None
} else {
Some(content)
},
tool_calls: if assembled_tool_calls.is_empty() {
None
} else {
Some(assembled_tool_calls)
},
tool_call_id: None,
name: None,
reasoning: if reasoning.is_empty() {
None
} else {
Some(reasoning)
},
};
Ok(OpenAIResponse {
id: resp_id,
object: resp_object,
created: resp_created,
model: resp_model,
choices: vec![OpenAIChoice {
index: 0,
message,
finish_reason,
}],
usage,
system_fingerprint,
})
}
#[allow(clippy::too_many_arguments)]
fn process_openai_sse_event(
raw: &str,
resp_id: &mut String,
resp_object: &mut String,
resp_created: &mut u64,
resp_model: &mut String,
system_fingerprint: &mut Option<String>,
role: &mut String,
content: &mut String,
reasoning: &mut String,
tool_calls: &mut BTreeMap<u64, PartialToolCall>,
finish_reason: &mut Option<String>,
usage: &mut OpenAIUsage,
) -> Result<bool, CompletionError> {
let mut data = String::new();
for line in raw.lines() {
if let Some(rest) = line.strip_prefix("data:") {
if !data.is_empty() {
data.push('\n');
}
data.push_str(rest.trim_start());
}
}
if data.is_empty() {
return Ok(false);
}
if data == "[DONE]" {
return Ok(true);
}
let Ok(json) = serde_json::from_str::<serde_json::Value>(&data) else {
return Ok(false);
};
if let Some(v) = json.get("id").and_then(|v| v.as_str()) {
*resp_id = v.to_string();
}
if let Some(v) = json.get("object").and_then(|v| v.as_str()) {
*resp_object = v.to_string();
}
if let Some(v) = json.get("created").and_then(|v| v.as_u64()) {
*resp_created = v;
}
if let Some(v) = json.get("model").and_then(|v| v.as_str()) {
*resp_model = v.to_string();
}
if let Some(v) = json.get("system_fingerprint").and_then(|v| v.as_str()) {
*system_fingerprint = Some(v.to_string());
}
if let Some(u) = json.get("usage") {
if let Some(v) = u.get("prompt_tokens").and_then(|v| v.as_u64()) {
usage.prompt_tokens = v;
}
if let Some(v) = u.get("completion_tokens").and_then(|v| v.as_u64()) {
usage.completion_tokens = v;
}
if let Some(v) = u.get("total_tokens").and_then(|v| v.as_u64()) {
usage.total_tokens = v;
}
}
let Some(choices) = json.get("choices").and_then(|v| v.as_array()) else {
return Ok(false);
};
let Some(choice) = choices.first() else {
return Ok(false);
};
if let Some(fr) = choice.get("finish_reason").and_then(|v| v.as_str()) {
*finish_reason = Some(fr.to_string());
}
let Some(delta) = choice.get("delta") else {
return Ok(false);
};
if let Some(r) = delta.get("role").and_then(|v| v.as_str()) {
*role = r.to_string();
}
if let Some(c) = delta.get("content").and_then(|v| v.as_str()) {
content.push_str(c);
}
if let Some(r) = delta
.get("reasoning")
.and_then(|v| v.as_str())
.or_else(|| delta.get("reasoning_content").and_then(|v| v.as_str()))
{
reasoning.push_str(r);
}
if let Some(tcs) = delta.get("tool_calls").and_then(|v| v.as_array()) {
for tc in tcs {
let idx = tc.get("index").and_then(|v| v.as_u64()).unwrap_or(0);
let entry = tool_calls.entry(idx).or_insert_with(|| PartialToolCall {
id: String::new(),
name: String::new(),
arguments: String::new(),
});
if let Some(id) = tc.get("id").and_then(|v| v.as_str()) {
if !id.is_empty() {
entry.id = id.to_string();
}
}
if let Some(func) = tc.get("function") {
if let Some(name) = func.get("name").and_then(|v| v.as_str()) {
if !name.is_empty() {
entry.name = name.to_string();
}
}
if let Some(args) = func.get("arguments").and_then(|v| v.as_str()) {
entry.arguments.push_str(args);
}
}
}
}
Ok(false)
}