use std::{fmt, sync::Arc, time::Duration};
use base64::{Engine as _, engine::general_purpose::STANDARD};
use reqwest::{
Client, RequestBuilder, Response, StatusCode, Url,
header::{AUTHORIZATION, HeaderValue},
multipart::{Form, Part},
redirect::Policy,
};
use serde_json::Value;
use zeroize::Zeroizing;
use crate::{
Error, GPT_4O_TRANSCRIBE, GeneratedImage, ImageFormat, ImageGeneration, ImageGenerationRequest,
ImageQuality, ImageTokenDetails, ImageUsage, Result, Transcription, TranscriptionRequest,
TranscriptionTokenDetails, TranscriptionTokenUsage, TranscriptionUsage,
error::{clean_message, transport},
};
const API_BASE: &str = "https://api.openai.com/v1/";
const DEFAULT_TIMEOUT: Duration = Duration::from_secs(5 * 60);
const MAX_RESPONSE_BYTES: usize = 128 * 1024 * 1024;
struct ApiKey(Zeroizing<String>);
impl ApiKey {
fn new(value: impl Into<String>) -> Result<Self> {
let supplied = Zeroizing::new(value.into());
let value = supplied.trim();
if value.is_empty() {
return Err(Error::InvalidApiKey);
}
let authorization = Zeroizing::new(format!("Bearer {value}"));
if HeaderValue::from_str(&authorization).is_err() {
return Err(Error::InvalidApiKey);
}
Ok(Self(Zeroizing::new(value.to_owned())))
}
fn sensitive_authorization(&self) -> Result<HeaderValue> {
let authorization = Zeroizing::new(format!("Bearer {}", self.0.as_str()));
let mut value = HeaderValue::from_str(&authorization).map_err(|_| Error::InvalidApiKey)?;
value.set_sensitive(true);
Ok(value)
}
}
impl fmt::Debug for ApiKey {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("ApiKey([REDACTED])")
}
}
#[derive(Clone)]
pub struct OpenAi {
api_key: Arc<ApiKey>,
client: Client,
transcription_endpoint: Url,
image_generation_endpoint: Url,
}
impl fmt::Debug for OpenAi {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("OpenAi")
.field("api_key", &self.api_key)
.field("api_base", &API_BASE)
.finish_non_exhaustive()
}
}
impl OpenAi {
pub fn open(api_key: impl Into<String>) -> Result<Self> {
let api_key = ApiKey::new(api_key)?;
let base = Url::parse(API_BASE)
.map_err(|_| Error::Protocol("compiled API base URL is invalid".into()))?;
let client = Client::builder()
.timeout(DEFAULT_TIMEOUT)
.redirect(Policy::none())
.retry(reqwest::retry::never())
.referer(false)
.no_proxy()
.https_only(true)
.user_agent(concat!("kcode-openai-api/", env!("CARGO_PKG_VERSION")))
.build()
.map_err(transport)?;
Ok(Self {
api_key: Arc::new(api_key),
client,
transcription_endpoint: base
.join("audio/transcriptions")
.map_err(|_| Error::Protocol("compiled transcription URL is invalid".into()))?,
image_generation_endpoint: base
.join("images/generations")
.map_err(|_| Error::Protocol("compiled image generation URL is invalid".into()))?,
})
}
pub async fn transcribe(&self, request: TranscriptionRequest) -> Result<Transcription> {
request.validate()?;
let TranscriptionRequest {
audio,
prompt,
language,
} = request;
let (file_name, mime_type, data) = audio.into_parts();
let part = Part::bytes(data)
.file_name(file_name)
.mime_str(&mime_type)
.map_err(|_| Error::InvalidInput("audio MIME type is invalid".into()))?;
let mut form = Form::new()
.part("file", part)
.text("model", GPT_4O_TRANSCRIBE)
.text("response_format", "json");
if let Some(prompt) = prompt {
form = form.text("prompt", prompt);
}
if let Some(language) = language {
form = form.text("language", language);
}
let (payload, request_id) = self
.execute(
self.client
.post(self.transcription_endpoint.clone())
.multipart(form),
)
.await?;
parse_transcription(&payload, request_id)
}
pub async fn generate_image(&self, request: ImageGenerationRequest) -> Result<ImageGeneration> {
request.validate()?;
let requested_format = request.output_format;
let (payload, request_id) = self
.execute(
self.client
.post(self.image_generation_endpoint.clone())
.json(&request.payload()),
)
.await?;
parse_image_generation(&payload, requested_format, request_id)
}
async fn execute(&self, request: RequestBuilder) -> Result<(Value, Option<String>)> {
let response = request
.header(AUTHORIZATION, self.api_key.sensitive_authorization()?)
.send()
.await
.map_err(transport)?;
let status = response.status();
let request_id = response
.headers()
.get("x-request-id")
.and_then(|value| value.to_str().ok())
.map(|value| clean_message(value, 200));
let body = bounded_body(response).await?;
if !status.is_success() {
return Err(provider_error(status, &body, request_id));
}
let payload = serde_json::from_slice(&body)
.map_err(|_| Error::Protocol("response was not valid JSON".into()))?;
Ok((payload, request_id))
}
}
async fn bounded_body(mut response: Response) -> Result<Vec<u8>> {
if response
.content_length()
.is_some_and(|value| value > MAX_RESPONSE_BYTES as u64)
{
return Err(Error::Protocol("response exceeded 128 MiB".into()));
}
let initial_capacity = response
.content_length()
.and_then(|value| usize::try_from(value).ok())
.unwrap_or(0)
.min(MAX_RESPONSE_BYTES);
let mut body = Vec::with_capacity(initial_capacity);
while let Some(chunk) = response.chunk().await.map_err(transport)? {
let length = body
.len()
.checked_add(chunk.len())
.ok_or_else(|| Error::Protocol("response exceeded 128 MiB".into()))?;
if length > MAX_RESPONSE_BYTES {
return Err(Error::Protocol("response exceeded 128 MiB".into()));
}
body.extend_from_slice(&chunk);
}
Ok(body)
}
fn parse_transcription(payload: &Value, request_id: Option<String>) -> Result<Transcription> {
let text = payload
.get("text")
.and_then(Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.ok_or_else(|| Error::Protocol("transcription response omitted non-empty text".into()))?
.to_owned();
let usage = payload
.get("usage")
.filter(|value| !value.is_null())
.map(parse_transcription_usage)
.transpose()?;
Ok(Transcription {
text,
usage,
request_id,
})
}
fn parse_transcription_usage(value: &Value) -> Result<TranscriptionUsage> {
let usage_type = value.get("type").and_then(Value::as_str);
if usage_type == Some("duration") {
let seconds = value
.get("seconds")
.and_then(Value::as_f64)
.ok_or_else(|| {
Error::Protocol("duration transcription usage omitted seconds".into())
})?;
if !seconds.is_finite() || seconds < 0.0 {
return Err(Error::Protocol(
"duration transcription usage contained invalid seconds".into(),
));
}
return Ok(TranscriptionUsage::DurationSeconds(seconds));
}
if !matches!(usage_type, None | Some("tokens")) {
return Err(Error::Protocol(
"transcription usage returned an unsupported type".into(),
));
}
let input_tokens = required_u64(value, "input_tokens", "transcription usage")?;
let output_tokens = required_u64(value, "output_tokens", "transcription usage")?;
let total_tokens = required_u64(value, "total_tokens", "transcription usage")?;
let input_details = value
.get("input_token_details")
.filter(|details| !details.is_null())
.map(|details| {
Ok(TranscriptionTokenDetails {
audio_tokens: optional_u64(details, "audio_tokens", "transcription usage")?,
text_tokens: optional_u64(details, "text_tokens", "transcription usage")?,
})
})
.transpose()?;
Ok(TranscriptionUsage::Tokens(TranscriptionTokenUsage {
input_tokens,
output_tokens,
total_tokens,
input_details,
}))
}
fn parse_image_generation(
payload: &Value,
requested_format: ImageFormat,
request_id: Option<String>,
) -> Result<ImageGeneration> {
let created = required_u64(payload, "created", "image generation response")?;
let data = payload
.get("data")
.and_then(Value::as_array)
.ok_or_else(|| Error::Protocol("image generation response omitted image data".into()))?;
if data.len() != 1 {
return Err(Error::Protocol(
"single-image request did not return exactly one image".into(),
));
}
let encoded = data[0]
.get("b64_json")
.and_then(Value::as_str)
.ok_or_else(|| Error::Protocol("generated image omitted base64 data".into()))?;
let decoded = STANDARD
.decode(encoded)
.map_err(|_| Error::Protocol("generated image contained invalid base64".into()))?;
if decoded.is_empty() {
return Err(Error::Protocol("generated image was empty".into()));
}
let format = match payload.get("output_format").and_then(Value::as_str) {
Some(value) => ImageFormat::parse(value)
.ok_or_else(|| Error::Protocol("generated image used an unknown format".into()))?,
None => requested_format,
};
let quality = payload
.get("quality")
.and_then(Value::as_str)
.and_then(ImageQuality::parse);
let size = payload
.get("size")
.and_then(Value::as_str)
.map(|value| clean_message(value, 40));
let usage = payload
.get("usage")
.filter(|value| !value.is_null())
.map(parse_image_usage)
.transpose()?;
Ok(ImageGeneration {
created,
image: GeneratedImage {
data: decoded,
format,
},
size,
quality,
usage,
request_id,
})
}
fn parse_image_usage(value: &Value) -> Result<ImageUsage> {
Ok(ImageUsage {
input_tokens: required_u64(value, "input_tokens", "image usage")?,
output_tokens: required_u64(value, "output_tokens", "image usage")?,
total_tokens: required_u64(value, "total_tokens", "image usage")?,
input_details: parse_image_token_details(
value
.get("input_tokens_details")
.ok_or_else(|| Error::Protocol("image usage omitted input token details".into()))?,
)?,
output_details: value
.get("output_tokens_details")
.filter(|details| !details.is_null())
.map(parse_image_token_details)
.transpose()?,
})
}
fn parse_image_token_details(value: &Value) -> Result<ImageTokenDetails> {
Ok(ImageTokenDetails {
text_tokens: required_u64(value, "text_tokens", "image token details")?,
image_tokens: required_u64(value, "image_tokens", "image token details")?,
})
}
fn required_u64(value: &Value, field: &str, context: &str) -> Result<u64> {
value
.get(field)
.and_then(Value::as_u64)
.ok_or_else(|| Error::Protocol(format!("{context} omitted {field}")))
}
fn optional_u64(value: &Value, field: &str, context: &str) -> Result<Option<u64>> {
match value.get(field) {
None | Some(Value::Null) => Ok(None),
Some(value) => value
.as_u64()
.map(Some)
.ok_or_else(|| Error::Protocol(format!("{context} returned invalid {field}"))),
}
}
fn provider_error(status: StatusCode, body: &[u8], request_id: Option<String>) -> Error {
let payload = serde_json::from_slice::<Value>(body).ok();
let code = payload
.as_ref()
.and_then(|value| {
value
.pointer("/error/code")
.and_then(Value::as_str)
.or_else(|| value.pointer("/error/type").and_then(Value::as_str))
})
.map(|value| clean_message(value, 100));
let message = payload
.as_ref()
.and_then(|value| value.pointer("/error/message"))
.and_then(Value::as_str)
.map(|value| clean_message(value, 400))
.unwrap_or_else(|| format!("provider request failed with HTTP {status}"));
Error::Provider {
status: status.as_u16(),
code,
message,
request_id,
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
use crate::{ImageGenerationRequest, ImageSize};
#[test]
fn debug_and_authorization_header_redact_api_key() {
let client = OpenAi::open("secret-api-key").unwrap();
let debug = format!("{client:?}");
assert!(debug.contains("[REDACTED]"));
assert!(!debug.contains("secret-api-key"));
let header = client.api_key.sensitive_authorization().unwrap();
assert!(header.is_sensitive());
let request = client
.client
.post(client.transcription_endpoint.clone())
.header(AUTHORIZATION, header);
assert!(!format!("{request:?}").contains("secret-api-key"));
}
#[test]
fn transcription_response_normalizes_token_usage() {
let payload = json!({
"text": " hello world ",
"usage": {
"type": "tokens",
"input_tokens": 12,
"output_tokens": 3,
"total_tokens": 15,
"input_token_details": {"audio_tokens": 10, "text_tokens": 2}
}
});
let parsed = parse_transcription(&payload, Some("req_123".into())).unwrap();
assert_eq!(parsed.text, "hello world");
assert_eq!(parsed.request_id.as_deref(), Some("req_123"));
assert_eq!(
parsed.usage,
Some(TranscriptionUsage::Tokens(TranscriptionTokenUsage {
input_tokens: 12,
output_tokens: 3,
total_tokens: 15,
input_details: Some(TranscriptionTokenDetails {
audio_tokens: Some(10),
text_tokens: Some(2),
}),
}))
);
}
#[test]
fn image_response_decodes_bytes_and_usage() {
let payload = json!({
"created": 1_721_000_000_u64,
"background": "opaque",
"output_format": "png",
"quality": "high",
"size": "2048x2048",
"data": [{"b64_json": "AQID"}],
"usage": {
"input_tokens": 10,
"output_tokens": 20,
"total_tokens": 30,
"input_tokens_details": {"text_tokens": 10, "image_tokens": 0},
"output_tokens_details": {"text_tokens": 0, "image_tokens": 20}
}
});
let parsed = parse_image_generation(&payload, ImageFormat::Png, None).unwrap();
assert_eq!(parsed.image.data, vec![1, 2, 3]);
assert_eq!(parsed.image.format, ImageFormat::Png);
assert_eq!(parsed.size.as_deref(), Some("2048x2048"));
assert_eq!(parsed.quality, Some(ImageQuality::High));
assert_eq!(
parsed.usage.unwrap().output_details.unwrap().image_tokens,
20
);
}
#[test]
fn gpt_image_payload_supports_flexible_dimensions() {
let mut request = ImageGenerationRequest::new("draw a quiet library");
request.size = ImageSize::dimensions(1536, 864).unwrap();
let payload = request.payload();
assert_eq!(payload["model"], "gpt-image-2");
assert_eq!(payload["size"], "1536x864");
}
#[test]
fn provider_errors_are_sanitized_and_keep_request_ids() {
let error = provider_error(
StatusCode::BAD_REQUEST,
br#"{"error":{"code":"moderation_blocked","message":"bad\nrequest"}}"#,
Some("req_456".into()),
);
assert!(matches!(
error,
Error::Provider {
status: 400,
code: Some(code),
message,
request_id: Some(request_id),
} if code == "moderation_blocked" && message == "bad request" && request_id == "req_456"
));
}
}