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(crate) fn append_extra_body_map(
57 mut form: reqwest::multipart::Form,
58 extra_body_map: &Option<serde_json::Map<String, serde_json::Value>>,
59) -> reqwest::multipart::Form {
60 if let Some(map) = extra_body_map {
61 for (key, value) in map {
62 let text = match value {
63 serde_json::Value::String(s) => s.clone(),
64 other => other.to_string(),
65 };
66 form = form.text(key.clone(), text);
67 }
68 }
69 form
70}
71
72pub trait PostNoStream: Post + Serialize + Sync + Send {
73 type Response: DeserializeOwned + FromStr<Err = OapiError> + Send + Sync;
74
75 fn get_response_string(
80 &self,
81 client: &reqwest::Client,
82 base_url: &str,
83 options: &RequestOptions,
84 ) -> impl Future<Output = Result<String, OapiError>> + Send + Sync {
85 async move {
86 if self.is_streaming() {
87 return Err(OapiError::NonStreamingViolation);
88 }
89
90 let response = apply_options(
91 client
92 .post(self.build_url(base_url)?)
93 .header("Accept", "application/json"),
94 options,
95 )
96 .json(self)
97 .send()
98 .await?;
99
100 crate::rest::response_text_checked(response).await
101 }
102 }
103
104 fn get_response(
106 &self,
107 client: &reqwest::Client,
108 url: &str,
109 options: &RequestOptions,
110 ) -> impl Future<Output = Result<Self::Response, OapiError>> + Send + Sync {
111 async move {
112 let text = self.get_response_string(client, url, options).await?;
113 let result = Self::Response::from_str(&text)?;
114 Ok(result)
115 }
116 }
117}
118
119pub trait PostStream: Post + Serialize + Sync + Send {
120 type Response: DeserializeOwned + FromStr<Err = OapiError> + Send + Sync + 'static;
121
122 fn get_stream_response_string(
174 &self,
175 client: &reqwest::Client,
176 base_url: &str,
177 options: &RequestOptions,
178 ) -> impl Future<
179 Output = Result<
180 impl Stream<Item = Result<String, OapiError>> + Send + Unpin + 'static,
181 OapiError,
182 >,
183 > + Send
184 + Sync {
185 async move {
186 if !self.is_streaming() {
187 return Err(OapiError::StreamingViolation);
188 }
189
190 let response = apply_options(
191 client
192 .post(self.build_url(base_url)?)
193 .header("Accept", "text/event-stream"),
194 options,
195 )
196 .json(self)
197 .send()
198 .await?;
199
200 let stream = crate::rest::check_status(response)
203 .await?
204 .bytes_stream()
205 .eventsource()
206 .map(|event| match event {
207 Ok(event) => Ok(event.data),
208 Err(e) => Err(OapiError::SseParseError(format!("SSE parse error: {e}"))),
209 });
210
211 Ok(stream)
212 }
213 }
214
215 fn get_stream_response(
223 &self,
224 client: &reqwest::Client,
225 base_url: &str,
226 options: &RequestOptions,
227 ) -> impl Future<
228 Output = Result<
229 impl Stream<Item = Result<Self::Response, OapiError>> + Send + Unpin + 'static,
230 OapiError,
231 >,
232 > + Send
233 + Sync {
234 async move {
235 let stream = self
236 .get_stream_response_string(client, base_url, options)
237 .await?;
238
239 let parsed_stream = stream
240 .take_while(|result| {
246 std::future::ready(!matches!(result, Ok(data) if data == "[DONE]"))
247 })
248 .and_then(|data| std::future::ready(Self::Response::from_str(&data)));
249
250 Ok(parsed_stream)
251 }
252 }
253}
254
255pub trait PostBinary: Post + Serialize + Sync + Send {
260 fn get_response_bytes(
266 &self,
267 client: &reqwest::Client,
268 base_url: &str,
269 options: &RequestOptions,
270 ) -> impl Future<Output = Result<Vec<u8>, OapiError>> + Send + Sync {
271 async move {
272 let response = apply_options(
273 client
274 .post(self.build_url(base_url)?)
275 .header("Accept", "application/octet-stream"),
276 options,
277 )
278 .json(self)
279 .send()
280 .await?;
281
282 crate::rest::response_bytes_checked(response).await
283 }
284 }
285
286 fn get_stream_response_bytes(
299 &self,
300 client: &reqwest::Client,
301 base_url: &str,
302 options: &RequestOptions,
303 ) -> impl Future<
304 Output = Result<
305 impl Stream<Item = Result<Vec<u8>, OapiError>> + Send + Unpin + 'static,
306 OapiError,
307 >,
308 > + Send
309 + Sync {
310 async move {
311 let response = apply_options(
312 client
313 .post(self.build_url(base_url)?)
314 .header("Accept", "application/octet-stream"),
315 options,
316 )
317 .json(self)
318 .send()
319 .await?;
320
321 let response = crate::rest::check_status(response).await?;
322
323 Ok(response.bytes_stream().map(|chunk| {
324 chunk
325 .map(|bytes| bytes.to_vec())
326 .map_err(OapiError::SendError)
327 }))
328 }
329 }
330}
331
332#[cfg(test)]
333mod test {
334 use futures_util::StreamExt;
335 use serde::Deserialize;
336 use serde_json::json;
337 use wiremock::matchers::{method, path};
338 use wiremock::{Mock, MockServer, ResponseTemplate};
339
340 use super::*;
341 use crate::chat::create::response::streaming::ChatCompletionChunk;
342 use crate::rest::{RequestOptions, skip_deserialization_errors};
343
344 fn sse_body() -> String {
348 const CHUNK_A: &str = r#"{"id":"1","choices":[{"index":0,"delta":{"content":"a"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#;
349 const CHUNK_B: &str = r#"{"id":"1","choices":[{"index":0,"delta":{"content":"b"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#;
350 format!("data: {CHUNK_A}\n\ndata: oops\n\ndata: {CHUNK_B}\n\ndata: [DONE]\n\n")
351 }
352
353 fn chunk_content(chunk: &ChatCompletionChunk) -> String {
354 chunk.choices[0].delta.content.clone().unwrap_or_default()
355 }
356
357 #[derive(Serialize)]
358 struct TestJsonRequest;
359
360 impl Post for TestJsonRequest {
361 fn is_streaming(&self) -> bool {
362 false
363 }
364 fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
365 Ok(format!("{}/test", base_url.trim_end_matches('/')))
366 }
367 }
368
369 impl PostNoStream for TestJsonRequest {
370 type Response = TestResponse;
371 }
372
373 #[derive(Debug, Deserialize)]
374 struct TestResponse {
375 #[allow(dead_code)]
376 id: String,
377 }
378
379 crate::impl_from_str!(TestResponse);
380
381 #[derive(Serialize)]
382 struct TestStreamRequest;
383
384 impl Post for TestStreamRequest {
385 fn is_streaming(&self) -> bool {
386 true
387 }
388 fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
389 Ok(format!("{}/test", base_url.trim_end_matches('/')))
390 }
391 }
392
393 impl PostStream for TestStreamRequest {
394 type Response = ChatCompletionChunk;
395 }
396
397 #[tokio::test]
398 async fn sends_bearer_auth_and_extra_headers() {
399 let server = MockServer::start().await;
400 Mock::given(method("POST"))
401 .and(path("/test"))
402 .respond_with(
403 ResponseTemplate::new(200).set_body_string(json!({"id": "x"}).to_string()),
404 )
405 .mount(&server)
406 .await;
407
408 let options = RequestOptions::bearer("sk-test")
409 .with_header("OpenAI-Organization", "org-1")
410 .unwrap();
411 TestJsonRequest
412 .get_response_string(&crate::rest::default_client(), &server.uri(), &options)
413 .await
414 .expect("request must succeed");
415
416 let requests = server.received_requests().await.expect("recorded requests");
417 assert_eq!(requests.len(), 1);
418 let headers = &requests[0].headers;
419 assert_eq!(headers.get("authorization").unwrap(), "Bearer sk-test");
420 assert_eq!(headers.get("openai-organization").unwrap(), "org-1");
421 }
422
423 #[tokio::test]
424 async fn header_only_auth_sends_no_authorization() {
425 let server = MockServer::start().await;
426 Mock::given(method("POST"))
427 .and(path("/test"))
428 .respond_with(
429 ResponseTemplate::new(200).set_body_string(json!({"id": "x"}).to_string()),
430 )
431 .mount(&server)
432 .await;
433
434 let options = RequestOptions::new()
436 .with_header("api-key", "azure-key")
437 .unwrap();
438 TestJsonRequest
439 .get_response_string(&crate::rest::default_client(), &server.uri(), &options)
440 .await
441 .expect("request must succeed");
442
443 let requests = server.received_requests().await.expect("recorded requests");
444 assert_eq!(requests.len(), 1);
445 let headers = &requests[0].headers;
446 assert!(headers.get("authorization").is_none());
447 assert_eq!(headers.get("api-key").unwrap(), "azure-key");
448 }
449
450 #[tokio::test]
451 async fn multipart_upload_uses_build_url_and_options() {
452 let server = MockServer::start().await;
453 Mock::given(method("POST"))
454 .and(path("/files"))
455 .respond_with(
456 ResponseTemplate::new(200).set_body_string(json!({"id": "x"}).to_string()),
457 )
458 .mount(&server)
459 .await;
460
461 let path = std::env::temp_dir().join("openai_interface_multipart_test.txt");
462 std::fs::write(&path, b"hello").expect("test file must be writable");
463
464 let mut extra = serde_json::Map::new();
465 extra.insert(
466 "vendor_extension".to_string(),
467 serde_json::json!({"depth": 3}),
468 );
469
470 let request = crate::files::create::request::CreateFileRequest {
471 file: path,
472 purpose: crate::files::create::request::FilePurpose::Batch,
473 expires_after: None,
474 extra_body_map: Some(extra),
475 };
476 let options = RequestOptions::new()
477 .with_header("api-key", "azure-key")
478 .unwrap();
479 request
480 .get_response_string(&crate::rest::default_client(), &server.uri(), &options)
481 .await
482 .expect("upload must succeed");
483
484 let requests = server.received_requests().await.expect("recorded requests");
485 assert_eq!(requests.len(), 1);
486 let body = String::from_utf8_lossy(&requests[0].body);
488 assert!(
489 body.contains("vendor_extension") && body.contains("\"depth\":3"),
490 "extra body map must reach the multipart body: {body}"
491 );
492 let headers = &requests[0].headers;
493 assert_eq!(headers.get("api-key").unwrap(), "azure-key");
496 assert_eq!(requests[0].url.path(), "/files");
497 }
498
499 #[tokio::test]
500 async fn raw_stream_yields_every_data_item_including_done() {
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_string(
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<String, OapiError>> = stream.collect().await;
517 let data: Vec<&str> = items
518 .iter()
519 .map(|item| item.as_ref().expect("raw items must not fail").as_str())
520 .collect();
521 assert_eq!(data.len(), 4);
522 assert_eq!(data[3], "[DONE]");
523 }
524
525 #[tokio::test]
526 async fn parsed_stream_surfaces_bad_chunks_without_ending() {
527 let server = MockServer::start().await;
528 Mock::given(method("POST"))
529 .respond_with(ResponseTemplate::new(200).set_body_string(sse_body()))
530 .mount(&server)
531 .await;
532
533 let stream = TestStreamRequest
534 .get_stream_response(
535 &crate::rest::default_client(),
536 &server.uri(),
537 &RequestOptions::bearer("sk-test"),
538 )
539 .await
540 .expect("stream must start");
541
542 let items: Vec<Result<ChatCompletionChunk, OapiError>> = stream.collect().await;
543 assert_eq!(items.len(), 3, "good, bad, good; [DONE] ends the stream");
544 assert_eq!(chunk_content(items[0].as_ref().unwrap()), "a");
545 assert!(matches!(items[1], Err(OapiError::DeserializationError(_))));
546 assert_eq!(chunk_content(items[2].as_ref().unwrap()), "b");
547 }
548
549 #[tokio::test]
550 async fn skip_deserialization_errors_drops_bad_chunks() {
551 let server = MockServer::start().await;
552 Mock::given(method("POST"))
553 .respond_with(ResponseTemplate::new(200).set_body_string(sse_body()))
554 .mount(&server)
555 .await;
556
557 let stream = TestStreamRequest
558 .get_stream_response(
559 &crate::rest::default_client(),
560 &server.uri(),
561 &RequestOptions::bearer("sk-test"),
562 )
563 .await
564 .expect("stream must start");
565 let mut stream = skip_deserialization_errors(stream);
566
567 let mut contents = Vec::new();
568 while let Some(item) = stream.next().await {
569 contents.push(chunk_content(&item.expect("errors must be skipped")));
570 }
571 assert_eq!(contents, ["a", "b"]);
572 }
573}