openai_protocol/
multipart.rs1#[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#[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 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}