Skip to main content

openai_protocol/
multipart.rs

1//! Axum extractors for multipart/form-data inference endpoints.
2//!
3//! JSON endpoints get [`crate::validated::ValidatedJson`]; the
4//! `/v1/audio/transcriptions` endpoint uses multipart/form-data and gets
5//! [`AudioTranscriptionMultipart`], which parses the form into a typed
6//! `(TranscriptionRequest, AudioFile)` pair before the handler runs.
7
8#[cfg(feature = "axum")]
9use axum::{
10    extract::{multipart::MultipartError, FromRequest, Multipart, Request},
11    http::StatusCode,
12    response::{IntoResponse, Response},
13};
14
15#[cfg(feature = "axum")]
16use crate::transcription::{AudioFile, TranscriptionRequest};
17
18/// Extractor for `/v1/audio/transcriptions` requests.
19///
20/// Parses `multipart/form-data` into a [`TranscriptionRequest`] (text fields)
21/// plus an [`AudioFile`] (the `file` part). Returns `400 Bad Request` on
22/// malformed parts, missing/empty `file`, missing/blank `model`, or
23/// out-of-range `temperature`.
24#[cfg(feature = "axum")]
25pub struct AudioTranscriptionMultipart {
26    pub request: TranscriptionRequest,
27    pub audio: AudioFile,
28}
29
30#[cfg(feature = "axum")]
31impl<S: Send + Sync> FromRequest<S> for AudioTranscriptionMultipart {
32    type Rejection = Response;
33
34    async fn from_request(req: Request, state: &S) -> Result<Self, Self::Rejection> {
35        let mut multipart = Multipart::from_request(req, state)
36            .await
37            .map_err(IntoResponse::into_response)?;
38
39        let mut file_bytes: Option<bytes::Bytes> = None;
40        let mut file_name: Option<String> = None;
41        let mut file_content_type: Option<String> = None;
42        let mut request = TranscriptionRequest::default();
43        let mut timestamp_granularities: Vec<String> = Vec::new();
44
45        loop {
46            let field = match multipart.next_field().await {
47                Ok(Some(f)) => f,
48                Ok(None) => break,
49                Err(e) => {
50                    return Err(bad_request(format!("Failed to read multipart field: {e}")));
51                }
52            };
53
54            let name = field.name().unwrap_or("").to_string();
55            match name.as_str() {
56                "file" => {
57                    file_name = field.file_name().map(str::to_string);
58                    file_content_type = field.content_type().map(str::to_string);
59                    match field.bytes().await {
60                        Ok(b) => file_bytes = Some(b),
61                        Err(e) => {
62                            return Err(bad_request(format!(
63                                "Failed to read audio file bytes: {e}"
64                            )));
65                        }
66                    }
67                }
68                "model" => match field.text().await {
69                    Ok(t) => request.model = t,
70                    Err(e) => return Err(bad_text_field("model", e)),
71                },
72                "language" => match field.text().await {
73                    Ok(t) => request.language = Some(t),
74                    Err(e) => return Err(bad_text_field("language", e)),
75                },
76                "prompt" => match field.text().await {
77                    Ok(t) => request.prompt = Some(t),
78                    Err(e) => return Err(bad_text_field("prompt", e)),
79                },
80                "response_format" => match field.text().await {
81                    Ok(t) => request.response_format = Some(t),
82                    Err(e) => return Err(bad_text_field("response_format", e)),
83                },
84                "temperature" => match field.text().await {
85                    Ok(t) => match t.trim().parse::<f32>() {
86                        Ok(v) if v.is_finite() && (0.0..=1.0).contains(&v) => {
87                            request.temperature = Some(v);
88                        }
89                        Ok(v) => {
90                            return Err(bad_request(format!(
91                                "Invalid 'temperature' value: {v} (must be a finite number in [0.0, 1.0])"
92                            )));
93                        }
94                        Err(e) => {
95                            return Err(bad_request(format!("Invalid 'temperature' value: {e}")));
96                        }
97                    },
98                    Err(e) => return Err(bad_text_field("temperature", e)),
99                },
100                "timestamp_granularities" | "timestamp_granularities[]" => {
101                    match field.text().await {
102                        Ok(t) => timestamp_granularities.push(t),
103                        Err(e) => return Err(bad_text_field("timestamp_granularities", e)),
104                    }
105                }
106                "stream" => match field.text().await {
107                    Ok(t) => match t.as_str() {
108                        "true" | "True" | "TRUE" | "1" => request.stream = Some(true),
109                        "false" | "False" | "FALSE" | "0" => request.stream = Some(false),
110                        other => {
111                            return Err(bad_request(format!(
112                                "Invalid 'stream' value: '{other}' (expected true/false/1/0)"
113                            )));
114                        }
115                    },
116                    Err(e) => return Err(bad_text_field("stream", e)),
117                },
118                _ => {
119                    // Unknown field; drain to free resources but otherwise ignore.
120                    let _ = field.bytes().await;
121                }
122            }
123        }
124
125        if request.model.trim().is_empty() {
126            return Err(bad_request("Missing required 'model' field".to_string()));
127        }
128        request.model = request.model.trim().to_string();
129
130        let bytes = match file_bytes {
131            Some(b) if !b.is_empty() => b,
132            Some(_) => {
133                return Err(bad_request("Uploaded 'file' part is empty".to_string()));
134            }
135            None => {
136                return Err(bad_request("Missing required 'file' part".to_string()));
137            }
138        };
139
140        if !timestamp_granularities.is_empty() {
141            request.timestamp_granularities = Some(timestamp_granularities);
142        }
143
144        let audio = AudioFile {
145            bytes,
146            file_name: file_name.unwrap_or_else(|| "audio".to_string()),
147            content_type: file_content_type,
148        };
149
150        Ok(AudioTranscriptionMultipart { request, audio })
151    }
152}
153
154#[cfg(feature = "axum")]
155fn bad_request(message: String) -> Response {
156    (StatusCode::BAD_REQUEST, message).into_response()
157}
158
159#[cfg(feature = "axum")]
160fn bad_text_field(field: &str, e: MultipartError) -> Response {
161    bad_request(format!("Failed to read '{field}' field: {e}"))
162}