Skip to main content

kcode_openai_api/
api.rs

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/// Cloneable asynchronous OpenAI client for transcription, image analysis, and image generation.
56#[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    /// Opens a client and retains the supplied OpenAI API key in memory.
76    ///
77    /// Clones share the key and HTTP connection pool. Dropping the last clone
78    /// discards the retained key with best-effort zeroization.
79    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    /// Transcribes one in-memory recording with `gpt-4o-transcribe`.
109    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    /// Analyzes one in-memory image with the fixed Responses API model.
142    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    /// Generates one image with the latest GPT Image model, `gpt-image-2`.
155    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}