1use std::{future::Future, str::FromStr};
2
3use eventsource_stream::Eventsource;
4use futures_util::{Stream, StreamExt, TryStreamExt};
5use serde::{Serialize, de::DeserializeOwned};
6
7use crate::errors::OapiError;
8use crate::rest::{Auth, RequestOptions};
9
10pub trait Post {
11 fn is_streaming(&self) -> bool;
12 fn build_url(&self, base_url: &str) -> Result<String, OapiError>;
16}
17
18pub(crate) fn apply_options(
22 builder: reqwest::RequestBuilder,
23 options: &RequestOptions,
24) -> reqwest::RequestBuilder {
25 let builder = match &options.auth {
26 Auth::Bearer(token) => builder.bearer_auth(token),
27 Auth::None => builder,
28 };
29 builder.headers(options.extra_headers.clone())
30}
31
32pub(crate) async fn post_multipart_json(
36 client: &reqwest::Client,
37 url: String,
38 form: reqwest::multipart::Form,
39 options: &RequestOptions,
40) -> Result<String, OapiError> {
41 let response = apply_options(
42 client.post(url).header("Accept", "application/json"),
43 options,
44 )
45 .multipart(form)
46 .send()
47 .await?;
48
49 crate::rest::response_text_checked(response).await
50}
51
52pub trait PostNoStream: Post + Serialize + Sync + Send {
53 type Response: DeserializeOwned + FromStr<Err = OapiError> + Send + Sync;
54
55 fn get_response_string(
60 &self,
61 client: &reqwest::Client,
62 base_url: &str,
63 options: &RequestOptions,
64 ) -> impl Future<Output = Result<String, OapiError>> + Send + Sync {
65 async move {
66 if self.is_streaming() {
67 return Err(OapiError::NonStreamingViolation);
68 }
69
70 let response = apply_options(
71 client
72 .post(self.build_url(base_url)?)
73 .header("Accept", "application/json"),
74 options,
75 )
76 .json(self)
77 .send()
78 .await?;
79
80 crate::rest::response_text_checked(response).await
81 }
82 }
83
84 fn get_response(
86 &self,
87 client: &reqwest::Client,
88 url: &str,
89 options: &RequestOptions,
90 ) -> impl Future<Output = Result<Self::Response, OapiError>> + Send + Sync {
91 async move {
92 let text = self.get_response_string(client, url, options).await?;
93 let result = Self::Response::from_str(&text)?;
94 Ok(result)
95 }
96 }
97}
98
99pub trait PostStream: Post + Serialize + Sync + Send {
100 type Response: DeserializeOwned + FromStr<Err = OapiError> + Send + Sync + 'static;
101
102 fn get_stream_response_string(
160 &self,
161 client: &reqwest::Client,
162 base_url: &str,
163 options: &RequestOptions,
164 ) -> impl Future<
165 Output = Result<
166 impl Stream<Item = Result<String, OapiError>> + Send + Unpin + 'static,
167 OapiError,
168 >,
169 > + Send
170 + Sync {
171 async move {
172 if !self.is_streaming() {
173 return Err(OapiError::StreamingViolation);
174 }
175
176 let response = apply_options(
177 client
178 .post(self.build_url(base_url)?)
179 .header("Accept", "text/event-stream"),
180 options,
181 )
182 .json(self)
183 .send()
184 .await?;
185
186 let stream = crate::rest::check_status(response)
189 .await?
190 .bytes_stream()
191 .eventsource()
192 .map(|event| match event {
193 Ok(event) => Ok(event.data),
194 Err(e) => Err(OapiError::SseParseError(format!("SSE parse error: {e}"))),
195 });
196
197 Ok(stream)
198 }
199 }
200
201 fn get_stream_response(
209 &self,
210 client: &reqwest::Client,
211 base_url: &str,
212 options: &RequestOptions,
213 ) -> impl Future<
214 Output = Result<
215 impl Stream<Item = Result<Self::Response, OapiError>> + Send + Unpin + 'static,
216 OapiError,
217 >,
218 > + Send
219 + Sync {
220 async move {
221 let stream = self
222 .get_stream_response_string(client, base_url, options)
223 .await?;
224
225 let parsed_stream = stream
226 .take_while(|result| {
232 std::future::ready(!matches!(result, Ok(data) if data == "[DONE]"))
233 })
234 .and_then(|data| std::future::ready(Self::Response::from_str(&data)));
235
236 Ok(parsed_stream)
237 }
238 }
239}
240
241pub trait PostBinary: Post + Serialize + Sync + Send {
246 fn get_response_bytes(
252 &self,
253 client: &reqwest::Client,
254 base_url: &str,
255 options: &RequestOptions,
256 ) -> impl Future<Output = Result<Vec<u8>, OapiError>> + Send + Sync {
257 async move {
258 let response = apply_options(
259 client
260 .post(self.build_url(base_url)?)
261 .header("Accept", "application/octet-stream"),
262 options,
263 )
264 .json(self)
265 .send()
266 .await?;
267
268 crate::rest::response_bytes_checked(response).await
269 }
270 }
271
272 fn get_stream_response_bytes(
285 &self,
286 client: &reqwest::Client,
287 base_url: &str,
288 options: &RequestOptions,
289 ) -> impl Future<
290 Output = Result<
291 impl Stream<Item = Result<Vec<u8>, OapiError>> + Send + Unpin + 'static,
292 OapiError,
293 >,
294 > + Send
295 + Sync {
296 async move {
297 let response = apply_options(
298 client
299 .post(self.build_url(base_url)?)
300 .header("Accept", "application/octet-stream"),
301 options,
302 )
303 .json(self)
304 .send()
305 .await?;
306
307 let response = crate::rest::check_status(response).await?;
308
309 Ok(response.bytes_stream().map(|chunk| {
310 chunk
311 .map(|bytes| bytes.to_vec())
312 .map_err(OapiError::SendError)
313 }))
314 }
315 }
316}
317
318#[cfg(test)]
319mod test {
320 use futures_util::StreamExt;
321 use serde::Deserialize;
322 use serde_json::json;
323 use wiremock::matchers::{method, path};
324 use wiremock::{Mock, MockServer, ResponseTemplate};
325
326 use super::*;
327 use crate::chat::create::response::streaming::ChatCompletionChunk;
328 use crate::rest::{RequestOptions, skip_deserialization_errors};
329
330 fn sse_body() -> String {
334 const CHUNK_A: &str = r#"{"id":"1","choices":[{"index":0,"delta":{"content":"a"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#;
335 const CHUNK_B: &str = r#"{"id":"1","choices":[{"index":0,"delta":{"content":"b"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#;
336 format!("data: {CHUNK_A}\n\ndata: oops\n\ndata: {CHUNK_B}\n\ndata: [DONE]\n\n")
337 }
338
339 fn chunk_content(chunk: &ChatCompletionChunk) -> String {
340 chunk.choices[0].delta.content.clone().unwrap_or_default()
341 }
342
343 #[derive(Serialize)]
344 struct TestJsonRequest;
345
346 impl Post for TestJsonRequest {
347 fn is_streaming(&self) -> bool {
348 false
349 }
350 fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
351 Ok(format!("{}/test", base_url.trim_end_matches('/')))
352 }
353 }
354
355 impl PostNoStream for TestJsonRequest {
356 type Response = TestResponse;
357 }
358
359 #[derive(Debug, Deserialize)]
360 struct TestResponse {
361 #[allow(dead_code)]
362 id: String,
363 }
364
365 crate::impl_from_str!(TestResponse);
366
367 #[derive(Serialize)]
368 struct TestStreamRequest;
369
370 impl Post for TestStreamRequest {
371 fn is_streaming(&self) -> bool {
372 true
373 }
374 fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
375 Ok(format!("{}/test", base_url.trim_end_matches('/')))
376 }
377 }
378
379 impl PostStream for TestStreamRequest {
380 type Response = ChatCompletionChunk;
381 }
382
383 #[tokio::test]
384 async fn sends_bearer_auth_and_extra_headers() {
385 let server = MockServer::start().await;
386 Mock::given(method("POST"))
387 .and(path("/test"))
388 .respond_with(
389 ResponseTemplate::new(200).set_body_string(json!({"id": "x"}).to_string()),
390 )
391 .mount(&server)
392 .await;
393
394 let options = RequestOptions::bearer("sk-test")
395 .with_header("OpenAI-Organization", "org-1")
396 .unwrap();
397 TestJsonRequest
398 .get_response_string(&crate::rest::default_client(), &server.uri(), &options)
399 .await
400 .expect("request must succeed");
401
402 let requests = server.received_requests().await.expect("recorded requests");
403 assert_eq!(requests.len(), 1);
404 let headers = &requests[0].headers;
405 assert_eq!(headers.get("authorization").unwrap(), "Bearer sk-test");
406 assert_eq!(headers.get("openai-organization").unwrap(), "org-1");
407 }
408
409 #[tokio::test]
410 async fn header_only_auth_sends_no_authorization() {
411 let server = MockServer::start().await;
412 Mock::given(method("POST"))
413 .and(path("/test"))
414 .respond_with(
415 ResponseTemplate::new(200).set_body_string(json!({"id": "x"}).to_string()),
416 )
417 .mount(&server)
418 .await;
419
420 let options = RequestOptions::new()
422 .with_header("api-key", "azure-key")
423 .unwrap();
424 TestJsonRequest
425 .get_response_string(&crate::rest::default_client(), &server.uri(), &options)
426 .await
427 .expect("request must succeed");
428
429 let requests = server.received_requests().await.expect("recorded requests");
430 assert_eq!(requests.len(), 1);
431 let headers = &requests[0].headers;
432 assert!(headers.get("authorization").is_none());
433 assert_eq!(headers.get("api-key").unwrap(), "azure-key");
434 }
435
436 #[tokio::test]
437 async fn multipart_upload_uses_build_url_and_options() {
438 let server = MockServer::start().await;
439 Mock::given(method("POST"))
440 .and(path("/files"))
441 .respond_with(
442 ResponseTemplate::new(200).set_body_string(json!({"id": "x"}).to_string()),
443 )
444 .mount(&server)
445 .await;
446
447 let path = std::env::temp_dir().join("openai_interface_multipart_test.txt");
448 std::fs::write(&path, b"hello").expect("test file must be writable");
449
450 let request = crate::files::create::request::CreateFileRequest {
451 file: path,
452 purpose: crate::files::create::request::FilePurpose::Batch,
453 expires_after: None,
454 extra_body: None,
455 };
456 let options = RequestOptions::new()
457 .with_header("api-key", "azure-key")
458 .unwrap();
459 request
460 .get_response_string(&crate::rest::default_client(), &server.uri(), &options)
461 .await
462 .expect("upload must succeed");
463
464 let requests = server.received_requests().await.expect("recorded requests");
465 assert_eq!(requests.len(), 1);
466 let headers = &requests[0].headers;
467 assert_eq!(headers.get("api-key").unwrap(), "azure-key");
470 assert_eq!(requests[0].url.path(), "/files");
471 }
472
473 #[tokio::test]
474 async fn raw_stream_yields_every_data_item_including_done() {
475 let server = MockServer::start().await;
476 Mock::given(method("POST"))
477 .respond_with(ResponseTemplate::new(200).set_body_string(sse_body()))
478 .mount(&server)
479 .await;
480
481 let stream = TestStreamRequest
482 .get_stream_response_string(
483 &crate::rest::default_client(),
484 &server.uri(),
485 &RequestOptions::bearer("sk-test"),
486 )
487 .await
488 .expect("stream must start");
489
490 let items: Vec<Result<String, OapiError>> = stream.collect().await;
491 let data: Vec<&str> = items
492 .iter()
493 .map(|item| item.as_ref().expect("raw items must not fail").as_str())
494 .collect();
495 assert_eq!(data.len(), 4);
496 assert_eq!(data[3], "[DONE]");
497 }
498
499 #[tokio::test]
500 async fn parsed_stream_surfaces_bad_chunks_without_ending() {
501 let server = MockServer::start().await;
502 Mock::given(method("POST"))
503 .respond_with(ResponseTemplate::new(200).set_body_string(sse_body()))
504 .mount(&server)
505 .await;
506
507 let stream = TestStreamRequest
508 .get_stream_response(
509 &crate::rest::default_client(),
510 &server.uri(),
511 &RequestOptions::bearer("sk-test"),
512 )
513 .await
514 .expect("stream must start");
515
516 let items: Vec<Result<ChatCompletionChunk, OapiError>> = stream.collect().await;
517 assert_eq!(items.len(), 3, "good, bad, good; [DONE] ends the stream");
518 assert_eq!(chunk_content(items[0].as_ref().unwrap()), "a");
519 assert!(matches!(items[1], Err(OapiError::DeserializationError(_))));
520 assert_eq!(chunk_content(items[2].as_ref().unwrap()), "b");
521 }
522
523 #[tokio::test]
524 async fn skip_deserialization_errors_drops_bad_chunks() {
525 let server = MockServer::start().await;
526 Mock::given(method("POST"))
527 .respond_with(ResponseTemplate::new(200).set_body_string(sse_body()))
528 .mount(&server)
529 .await;
530
531 let stream = TestStreamRequest
532 .get_stream_response(
533 &crate::rest::default_client(),
534 &server.uri(),
535 &RequestOptions::bearer("sk-test"),
536 )
537 .await
538 .expect("stream must start");
539 let mut stream = skip_deserialization_errors(stream);
540
541 let mut contents = Vec::new();
542 while let Some(item) = stream.next().await {
543 contents.push(chunk_content(&item.expect("errors must be skipped")));
544 }
545 assert_eq!(contents, ["a", "b"]);
546 }
547}