1use std::{fmt, sync::Arc, time::Duration};
2
3use base64::{Engine as _, engine::general_purpose::STANDARD};
4use reqwest::{
5 Client, RequestBuilder, Response, StatusCode, Url,
6 header::{AUTHORIZATION, HeaderValue},
7 multipart::{Form, Part},
8 redirect::Policy,
9};
10use serde_json::Value;
11use zeroize::Zeroizing;
12
13use crate::{
14 Error, GPT_4O_TRANSCRIBE, GeneratedImage, ImageAnalysis, ImageAnalysisRequest,
15 ImageAnalysisStatus, ImageAnalysisUsage, ImageFormat, ImageGeneration, ImageGenerationRequest,
16 ImageQuality, ImageTokenDetails, ImageUsage, Result, Transcription, TranscriptionRequest,
17 TranscriptionTokenDetails, TranscriptionTokenUsage, TranscriptionUsage,
18 error::{clean_message, transport},
19};
20
21const API_BASE: &str = "https://api.openai.com/v1/";
22const DEFAULT_TIMEOUT: Duration = Duration::from_secs(5 * 60);
23const MAX_RESPONSE_BYTES: usize = 128 * 1024 * 1024;
24
25struct ApiKey(Zeroizing<String>);
26
27impl ApiKey {
28 fn new(value: impl Into<String>) -> Result<Self> {
29 let supplied = Zeroizing::new(value.into());
30 let value = supplied.trim();
31 if value.is_empty() {
32 return Err(Error::InvalidApiKey);
33 }
34 let authorization = Zeroizing::new(format!("Bearer {value}"));
35 if HeaderValue::from_str(&authorization).is_err() {
36 return Err(Error::InvalidApiKey);
37 }
38 Ok(Self(Zeroizing::new(value.to_owned())))
39 }
40
41 fn sensitive_authorization(&self) -> Result<HeaderValue> {
42 let authorization = Zeroizing::new(format!("Bearer {}", self.0.as_str()));
43 let mut value = HeaderValue::from_str(&authorization).map_err(|_| Error::InvalidApiKey)?;
44 value.set_sensitive(true);
45 Ok(value)
46 }
47}
48
49impl fmt::Debug for ApiKey {
50 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
51 f.write_str("ApiKey([REDACTED])")
52 }
53}
54
55#[derive(Clone)]
57pub struct OpenAi {
58 api_key: Arc<ApiKey>,
59 client: Client,
60 transcription_endpoint: Url,
61 responses_endpoint: Url,
62 image_generation_endpoint: Url,
63}
64
65impl fmt::Debug for OpenAi {
66 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
67 f.debug_struct("OpenAi")
68 .field("api_key", &self.api_key)
69 .field("api_base", &API_BASE)
70 .finish_non_exhaustive()
71 }
72}
73
74impl OpenAi {
75 pub fn open(api_key: impl Into<String>) -> Result<Self> {
80 let api_key = ApiKey::new(api_key)?;
81 let base = Url::parse(API_BASE)
82 .map_err(|_| Error::Protocol("compiled API base URL is invalid".into()))?;
83 let client = Client::builder()
84 .timeout(DEFAULT_TIMEOUT)
85 .redirect(Policy::none())
86 .retry(reqwest::retry::never())
87 .referer(false)
88 .no_proxy()
89 .https_only(true)
90 .user_agent(concat!("kcode-openai-api/", env!("CARGO_PKG_VERSION")))
91 .build()
92 .map_err(transport)?;
93 Ok(Self {
94 api_key: Arc::new(api_key),
95 client,
96 transcription_endpoint: base
97 .join("audio/transcriptions")
98 .map_err(|_| Error::Protocol("compiled transcription URL is invalid".into()))?,
99 responses_endpoint: base
100 .join("responses")
101 .map_err(|_| Error::Protocol("compiled Responses API URL is invalid".into()))?,
102 image_generation_endpoint: base
103 .join("images/generations")
104 .map_err(|_| Error::Protocol("compiled image generation URL is invalid".into()))?,
105 })
106 }
107
108 pub async fn transcribe(&self, request: TranscriptionRequest) -> Result<Transcription> {
110 request.validate()?;
111 let TranscriptionRequest {
112 audio,
113 prompt,
114 language,
115 } = request;
116 let (file_name, mime_type, data) = audio.into_parts();
117 let part = Part::bytes(data)
118 .file_name(file_name)
119 .mime_str(&mime_type)
120 .map_err(|_| Error::InvalidInput("audio MIME type is invalid".into()))?;
121 let mut form = Form::new()
122 .part("file", part)
123 .text("model", GPT_4O_TRANSCRIBE)
124 .text("response_format", "json");
125 if let Some(prompt) = prompt {
126 form = form.text("prompt", prompt);
127 }
128 if let Some(language) = language {
129 form = form.text("language", language);
130 }
131 let (payload, request_id) = self
132 .execute(
133 self.client
134 .post(self.transcription_endpoint.clone())
135 .multipart(form),
136 )
137 .await?;
138 parse_transcription(&payload, request_id)
139 }
140
141 pub async fn analyze_image(&self, request: ImageAnalysisRequest) -> Result<ImageAnalysis> {
143 request.validate()?;
144 let (payload, request_id) = self
145 .execute(
146 self.client
147 .post(self.responses_endpoint.clone())
148 .json(&request.payload()),
149 )
150 .await?;
151 parse_image_analysis(&payload, request_id)
152 }
153
154 pub async fn generate_image(&self, request: ImageGenerationRequest) -> Result<ImageGeneration> {
156 request.validate()?;
157 let requested_format = request.output_format;
158 let (payload, request_id) = self
159 .execute(
160 self.client
161 .post(self.image_generation_endpoint.clone())
162 .json(&request.payload()),
163 )
164 .await?;
165 parse_image_generation(&payload, requested_format, request_id)
166 }
167
168 async fn execute(&self, request: RequestBuilder) -> Result<(Value, Option<String>)> {
169 let response = request
170 .header(AUTHORIZATION, self.api_key.sensitive_authorization()?)
171 .send()
172 .await
173 .map_err(transport)?;
174 let status = response.status();
175 let request_id = response
176 .headers()
177 .get("x-request-id")
178 .and_then(|value| value.to_str().ok())
179 .map(|value| clean_message(value, 200));
180 let body = bounded_body(response).await?;
181 if !status.is_success() {
182 return Err(provider_error(status, &body, request_id));
183 }
184 let payload = serde_json::from_slice(&body)
185 .map_err(|_| Error::Protocol("response was not valid JSON".into()))?;
186 Ok((payload, request_id))
187 }
188}
189
190async fn bounded_body(mut response: Response) -> Result<Vec<u8>> {
191 if response
192 .content_length()
193 .is_some_and(|value| value > MAX_RESPONSE_BYTES as u64)
194 {
195 return Err(Error::Protocol("response exceeded 128 MiB".into()));
196 }
197 let initial_capacity = response
198 .content_length()
199 .and_then(|value| usize::try_from(value).ok())
200 .unwrap_or(0)
201 .min(MAX_RESPONSE_BYTES);
202 let mut body = Vec::with_capacity(initial_capacity);
203 while let Some(chunk) = response.chunk().await.map_err(transport)? {
204 let length = body
205 .len()
206 .checked_add(chunk.len())
207 .ok_or_else(|| Error::Protocol("response exceeded 128 MiB".into()))?;
208 if length > MAX_RESPONSE_BYTES {
209 return Err(Error::Protocol("response exceeded 128 MiB".into()));
210 }
211 body.extend_from_slice(&chunk);
212 }
213 Ok(body)
214}
215
216fn parse_transcription(payload: &Value, request_id: Option<String>) -> Result<Transcription> {
217 let text = payload
218 .get("text")
219 .and_then(Value::as_str)
220 .map(str::trim)
221 .filter(|value| !value.is_empty())
222 .ok_or_else(|| Error::Protocol("transcription response omitted non-empty text".into()))?
223 .to_owned();
224 let usage = payload
225 .get("usage")
226 .filter(|value| !value.is_null())
227 .map(parse_transcription_usage)
228 .transpose()?;
229 Ok(Transcription {
230 text,
231 usage,
232 request_id,
233 })
234}
235
236fn parse_transcription_usage(value: &Value) -> Result<TranscriptionUsage> {
237 let usage_type = value.get("type").and_then(Value::as_str);
238 if usage_type == Some("duration") {
239 let seconds = value
240 .get("seconds")
241 .and_then(Value::as_f64)
242 .ok_or_else(|| {
243 Error::Protocol("duration transcription usage omitted seconds".into())
244 })?;
245 if !seconds.is_finite() || seconds < 0.0 {
246 return Err(Error::Protocol(
247 "duration transcription usage contained invalid seconds".into(),
248 ));
249 }
250 return Ok(TranscriptionUsage::DurationSeconds(seconds));
251 }
252 if !matches!(usage_type, None | Some("tokens")) {
253 return Err(Error::Protocol(
254 "transcription usage returned an unsupported type".into(),
255 ));
256 }
257 let input_tokens = required_u64(value, "input_tokens", "transcription usage")?;
258 let output_tokens = required_u64(value, "output_tokens", "transcription usage")?;
259 let total_tokens = required_u64(value, "total_tokens", "transcription usage")?;
260 let input_details = value
261 .get("input_token_details")
262 .filter(|details| !details.is_null())
263 .map(|details| {
264 Ok(TranscriptionTokenDetails {
265 audio_tokens: optional_u64(details, "audio_tokens", "transcription usage")?,
266 text_tokens: optional_u64(details, "text_tokens", "transcription usage")?,
267 })
268 })
269 .transpose()?;
270 Ok(TranscriptionUsage::Tokens(TranscriptionTokenUsage {
271 input_tokens,
272 output_tokens,
273 total_tokens,
274 input_details,
275 }))
276}
277
278fn parse_image_analysis(payload: &Value, request_id: Option<String>) -> Result<ImageAnalysis> {
279 let response_id = required_nonempty_string(payload, "id", "image-analysis response")?;
280 let model = required_nonempty_string(payload, "model", "image-analysis response")?;
281 let status = match payload.get("status").and_then(Value::as_str) {
282 Some("completed") => ImageAnalysisStatus::Completed,
283 Some("incomplete") => {
284 let reason = payload
285 .pointer("/incomplete_details/reason")
286 .and_then(Value::as_str)
287 .map(str::trim)
288 .filter(|value| !value.is_empty())
289 .map(|value| clean_message(value, 100));
290 ImageAnalysisStatus::Incomplete { reason }
291 }
292 _ => {
293 return Err(Error::Protocol(
294 "image-analysis response returned an unsupported status".into(),
295 ));
296 }
297 };
298
299 let output = payload
300 .get("output")
301 .and_then(Value::as_array)
302 .ok_or_else(|| Error::Protocol("image-analysis response omitted output".into()))?;
303 let mut fragments = Vec::new();
304 for item in output {
305 if item.get("type").and_then(Value::as_str) != Some("message")
306 || item.get("role").and_then(Value::as_str) != Some("assistant")
307 {
308 continue;
309 }
310 let Some(content) = item.get("content").and_then(Value::as_array) else {
311 continue;
312 };
313 for part in content {
314 if part.get("type").and_then(Value::as_str) != Some("output_text") {
315 continue;
316 }
317 if let Some(text) = part.get("text").and_then(Value::as_str) {
318 let text = text.trim();
319 if !text.is_empty() {
320 fragments.push(text);
321 }
322 }
323 }
324 }
325 let text = fragments.join("\n");
326 if text.is_empty() {
327 return Err(Error::Protocol(
328 "image-analysis response omitted non-empty assistant text".into(),
329 ));
330 }
331
332 let usage = payload
333 .get("usage")
334 .filter(|value| !value.is_null())
335 .map(parse_image_analysis_usage)
336 .transpose()?;
337 Ok(ImageAnalysis {
338 text,
339 response_id,
340 model,
341 status,
342 usage,
343 request_id,
344 })
345}
346
347fn parse_image_analysis_usage(value: &Value) -> Result<ImageAnalysisUsage> {
348 let (cached_input_tokens, cache_write_input_tokens) = match value.get("input_tokens_details") {
349 None | Some(Value::Null) => (None, None),
350 Some(details) => (
351 optional_u64(details, "cached_tokens", "image-analysis usage")?,
352 optional_u64(details, "cache_write_tokens", "image-analysis usage")?,
353 ),
354 };
355 let reasoning_output_tokens = match value.get("output_tokens_details") {
356 None | Some(Value::Null) => None,
357 Some(details) => optional_u64(details, "reasoning_tokens", "image-analysis usage")?,
358 };
359 Ok(ImageAnalysisUsage {
360 input_tokens: required_u64(value, "input_tokens", "image-analysis usage")?,
361 output_tokens: required_u64(value, "output_tokens", "image-analysis usage")?,
362 total_tokens: required_u64(value, "total_tokens", "image-analysis usage")?,
363 cached_input_tokens,
364 cache_write_input_tokens,
365 reasoning_output_tokens,
366 })
367}
368
369fn parse_image_generation(
370 payload: &Value,
371 requested_format: ImageFormat,
372 request_id: Option<String>,
373) -> Result<ImageGeneration> {
374 let created = required_u64(payload, "created", "image generation response")?;
375 let data = payload
376 .get("data")
377 .and_then(Value::as_array)
378 .ok_or_else(|| Error::Protocol("image generation response omitted image data".into()))?;
379 if data.len() != 1 {
380 return Err(Error::Protocol(
381 "single-image request did not return exactly one image".into(),
382 ));
383 }
384 let encoded = data[0]
385 .get("b64_json")
386 .and_then(Value::as_str)
387 .ok_or_else(|| Error::Protocol("generated image omitted base64 data".into()))?;
388 let decoded = STANDARD
389 .decode(encoded)
390 .map_err(|_| Error::Protocol("generated image contained invalid base64".into()))?;
391 if decoded.is_empty() {
392 return Err(Error::Protocol("generated image was empty".into()));
393 }
394 let format = match payload.get("output_format").and_then(Value::as_str) {
395 Some(value) => ImageFormat::parse(value)
396 .ok_or_else(|| Error::Protocol("generated image used an unknown format".into()))?,
397 None => requested_format,
398 };
399 let quality = payload
400 .get("quality")
401 .and_then(Value::as_str)
402 .and_then(ImageQuality::parse);
403 let size = payload
404 .get("size")
405 .and_then(Value::as_str)
406 .map(|value| clean_message(value, 40));
407 let usage = payload
408 .get("usage")
409 .filter(|value| !value.is_null())
410 .map(parse_image_usage)
411 .transpose()?;
412 Ok(ImageGeneration {
413 created,
414 image: GeneratedImage {
415 data: decoded,
416 format,
417 },
418 size,
419 quality,
420 usage,
421 request_id,
422 })
423}
424
425fn parse_image_usage(value: &Value) -> Result<ImageUsage> {
426 Ok(ImageUsage {
427 input_tokens: required_u64(value, "input_tokens", "image usage")?,
428 output_tokens: required_u64(value, "output_tokens", "image usage")?,
429 total_tokens: required_u64(value, "total_tokens", "image usage")?,
430 input_details: parse_image_token_details(
431 value
432 .get("input_tokens_details")
433 .ok_or_else(|| Error::Protocol("image usage omitted input token details".into()))?,
434 )?,
435 output_details: value
436 .get("output_tokens_details")
437 .filter(|details| !details.is_null())
438 .map(parse_image_token_details)
439 .transpose()?,
440 })
441}
442
443fn parse_image_token_details(value: &Value) -> Result<ImageTokenDetails> {
444 Ok(ImageTokenDetails {
445 text_tokens: required_u64(value, "text_tokens", "image token details")?,
446 image_tokens: required_u64(value, "image_tokens", "image token details")?,
447 })
448}
449
450fn required_nonempty_string(value: &Value, field: &str, context: &str) -> Result<String> {
451 value
452 .get(field)
453 .and_then(Value::as_str)
454 .map(str::trim)
455 .filter(|value| !value.is_empty())
456 .map(str::to_owned)
457 .ok_or_else(|| Error::Protocol(format!("{context} omitted non-empty {field}")))
458}
459
460fn required_u64(value: &Value, field: &str, context: &str) -> Result<u64> {
461 value
462 .get(field)
463 .and_then(Value::as_u64)
464 .ok_or_else(|| Error::Protocol(format!("{context} omitted {field}")))
465}
466
467fn optional_u64(value: &Value, field: &str, context: &str) -> Result<Option<u64>> {
468 match value.get(field) {
469 None | Some(Value::Null) => Ok(None),
470 Some(value) => value
471 .as_u64()
472 .map(Some)
473 .ok_or_else(|| Error::Protocol(format!("{context} returned invalid {field}"))),
474 }
475}
476
477fn provider_error(status: StatusCode, body: &[u8], request_id: Option<String>) -> Error {
478 let payload = serde_json::from_slice::<Value>(body).ok();
479 let code = payload
480 .as_ref()
481 .and_then(|value| {
482 value
483 .pointer("/error/code")
484 .and_then(Value::as_str)
485 .or_else(|| value.pointer("/error/type").and_then(Value::as_str))
486 })
487 .map(|value| clean_message(value, 100));
488 let message = payload
489 .as_ref()
490 .and_then(|value| value.pointer("/error/message"))
491 .and_then(Value::as_str)
492 .map(|value| clean_message(value, 400))
493 .unwrap_or_else(|| format!("provider request failed with HTTP {status}"));
494 Error::Provider {
495 status: status.as_u16(),
496 code,
497 message,
498 request_id,
499 }
500}
501
502#[cfg(test)]
503mod tests {
504 use serde_json::json;
505
506 use super::*;
507 use crate::{ImageGenerationRequest, ImageSize};
508
509 #[test]
510 fn debug_and_authorization_header_redact_api_key() {
511 let client = OpenAi::open("secret-api-key").unwrap();
512 let debug = format!("{client:?}");
513 assert!(debug.contains("[REDACTED]"));
514 assert!(!debug.contains("secret-api-key"));
515
516 let header = client.api_key.sensitive_authorization().unwrap();
517 assert!(header.is_sensitive());
518 let request = client
519 .client
520 .post(client.transcription_endpoint.clone())
521 .header(AUTHORIZATION, header);
522 assert!(!format!("{request:?}").contains("secret-api-key"));
523 }
524
525 #[test]
526 fn transcription_response_normalizes_token_usage() {
527 let payload = json!({
528 "text": " hello world ",
529 "usage": {
530 "type": "tokens",
531 "input_tokens": 12,
532 "output_tokens": 3,
533 "total_tokens": 15,
534 "input_token_details": {"audio_tokens": 10, "text_tokens": 2}
535 }
536 });
537 let parsed = parse_transcription(&payload, Some("req_123".into())).unwrap();
538 assert_eq!(parsed.text, "hello world");
539 assert_eq!(parsed.request_id.as_deref(), Some("req_123"));
540 assert_eq!(
541 parsed.usage,
542 Some(TranscriptionUsage::Tokens(TranscriptionTokenUsage {
543 input_tokens: 12,
544 output_tokens: 3,
545 total_tokens: 15,
546 input_details: Some(TranscriptionTokenDetails {
547 audio_tokens: Some(10),
548 text_tokens: Some(2),
549 }),
550 }))
551 );
552 }
553
554 #[test]
555 fn image_analysis_response_normalizes_ordered_text_and_usage() {
556 let payload = json!({
557 "id": "resp_123",
558 "status": "completed",
559 "model": "gpt-5.6-2026-07-01",
560 "output": [
561 {"type": "reasoning", "summary": []},
562 {
563 "type": "message",
564 "role": "assistant",
565 "content": [
566 {"type": "output_text", "text": " first observation "},
567 {"type": "refusal", "refusal": "ignored"}
568 ]
569 },
570 {
571 "type": "message",
572 "role": "assistant",
573 "content": [
574 {"type": "output_text", "text": "second observation"}
575 ]
576 }
577 ],
578 "usage": {
579 "input_tokens": 40,
580 "output_tokens": 12,
581 "total_tokens": 52,
582 "input_tokens_details": {
583 "cached_tokens": 3,
584 "cache_write_tokens": 2
585 },
586 "output_tokens_details": {"reasoning_tokens": 4}
587 }
588 });
589 let parsed = parse_image_analysis(&payload, Some("req_vision".into())).unwrap();
590 assert_eq!(parsed.text, "first observation\nsecond observation");
591 assert_eq!(parsed.response_id, "resp_123");
592 assert_eq!(parsed.model, "gpt-5.6-2026-07-01");
593 assert_eq!(parsed.status, ImageAnalysisStatus::Completed);
594 assert_eq!(parsed.request_id.as_deref(), Some("req_vision"));
595 assert_eq!(
596 parsed.usage,
597 Some(ImageAnalysisUsage {
598 input_tokens: 40,
599 output_tokens: 12,
600 total_tokens: 52,
601 cached_input_tokens: Some(3),
602 cache_write_input_tokens: Some(2),
603 reasoning_output_tokens: Some(4),
604 })
605 );
606 }
607
608 #[test]
609 fn image_analysis_response_labels_valid_partial_text() {
610 let payload = json!({
611 "id": "resp_partial",
612 "status": "incomplete",
613 "incomplete_details": {"reason": "content_filter"},
614 "model": "gpt-5.6",
615 "output": [{
616 "type": "message",
617 "role": "assistant",
618 "content": [{"type": "output_text", "text": "visible partial result"}]
619 }]
620 });
621 let parsed = parse_image_analysis(&payload, None).unwrap();
622 assert_eq!(parsed.text, "visible partial result");
623 assert_eq!(
624 parsed.status,
625 ImageAnalysisStatus::Incomplete {
626 reason: Some("content_filter".into())
627 }
628 );
629 assert_eq!(parsed.usage, None);
630 }
631
632 #[test]
633 fn image_response_decodes_bytes_and_usage() {
634 let payload = json!({
635 "created": 1_721_000_000_u64,
636 "background": "opaque",
637 "output_format": "png",
638 "quality": "high",
639 "size": "2048x2048",
640 "data": [{"b64_json": "AQID"}],
641 "usage": {
642 "input_tokens": 10,
643 "output_tokens": 20,
644 "total_tokens": 30,
645 "input_tokens_details": {"text_tokens": 10, "image_tokens": 0},
646 "output_tokens_details": {"text_tokens": 0, "image_tokens": 20}
647 }
648 });
649 let parsed = parse_image_generation(&payload, ImageFormat::Png, None).unwrap();
650 assert_eq!(parsed.image.data, vec![1, 2, 3]);
651 assert_eq!(parsed.image.format, ImageFormat::Png);
652 assert_eq!(parsed.size.as_deref(), Some("2048x2048"));
653 assert_eq!(parsed.quality, Some(ImageQuality::High));
654 assert_eq!(
655 parsed.usage.unwrap().output_details.unwrap().image_tokens,
656 20
657 );
658 }
659
660 #[test]
661 fn gpt_image_payload_supports_flexible_dimensions() {
662 let mut request = ImageGenerationRequest::new("draw a quiet library");
663 request.size = ImageSize::dimensions(1536, 864).unwrap();
664 let payload = request.payload();
665 assert_eq!(payload["model"], "gpt-image-2");
666 assert_eq!(payload["size"], "1536x864");
667 }
668
669 #[test]
670 fn provider_errors_are_sanitized_and_keep_request_ids() {
671 let error = provider_error(
672 StatusCode::BAD_REQUEST,
673 br#"{"error":{"code":"moderation_blocked","message":"bad\nrequest"}}"#,
674 Some("req_456".into()),
675 );
676 assert!(matches!(
677 error,
678 Error::Provider {
679 status: 400,
680 code: Some(code),
681 message,
682 request_id: Some(request_id),
683 } if code == "moderation_blocked" && message == "bad request" && request_id == "req_456"
684 ));
685 }
686}