Skip to main content

openai_interface/rest/
post.rs

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    /// Builds the URL for the request.
13    ///
14    /// `base_url` should be like <https://api.openai.com/v1>
15    fn build_url(&self, base_url: &str) -> Result<String, OapiError>;
16}
17
18/// Applies the per-request authentication and extra headers to a request
19/// under construction. Headers in [`RequestOptions::extra_headers`] replace
20/// previously set headers of the same name.
21pub(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
32/// Sends a multipart form as a POST request and returns the raw response
33/// body. Shared by the endpoints that upload files (`files`, `images`,
34/// `audio`, `uploads`).
35pub(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
52/// Appends the request's `extra_body_map` entries to a multipart form as
53/// text fields: strings verbatim, any other JSON value in its serialized
54/// form. Multipart endpoints have no JSON body, so without this the
55/// flattened map would silently never reach the wire.
56pub(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    /// Sends a POST request and returns the raw response body.
76    ///
77    /// The `client` is supplied by the caller; see [`crate::rest::default_client`]
78    /// for a sensible default.
79    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    /// Sends a POST request and deserializes the response.
105    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    /// Sends a streaming POST request and returns the raw SSE data items as
123    /// strings.
124    ///
125    /// The returned stream yields every SSE data payload verbatim, including
126    /// the final `"[DONE]"` sentinel; use
127    /// [`Self::get_stream_response`] when you want the sentinel consumed and
128    /// the items deserialized. SSE framing errors surface as
129    /// [`OapiError::SseParseError`] items.
130    ///
131    /// The `client` is supplied by the caller; see [`crate::rest::default_client`]
132    /// for a sensible default.
133    ///
134    /// # Example
135    ///
136    /// ```rust,no_run
137    /// use std::sync::LazyLock;
138    /// use futures_util::StreamExt;
139    /// use openai_interface::chat::create::request::{Message, RequestBody};
140    /// use openai_interface::rest::{default_client, post::PostStream, RequestOptions};
141    ///
142    /// const DEEPSEEK_API_KEY: &str = "YOUR_API_KEY";
143    /// const DEEPSEEK_CHAT_URL: &'static str = "https://api.deepseek.com";
144    /// const DEEPSEEK_MODEL: &'static str = "deepseek-chat";
145    ///
146    /// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
147    /// // Needs the `ferritls` cargo feature; drop this line if you install
148    /// // your own rustls crypto provider (see `openai_interface::rest`).
149    /// # #[cfg(feature = "ferritls")]
150    /// openai_interface::rest::install_crypto_provider().ok();
151    ///
152    /// let request = RequestBody {
153    ///     messages: vec![
154    ///         Message::system("This is a request of test purpose. Reply briefly"),
155    ///         Message::user("What's your name?"),
156    ///     ],
157    ///     model: DEEPSEEK_MODEL.to_string(),
158    ///     stream: Some(true),
159    ///     ..Default::default()
160    /// };
161    ///
162    /// let options = RequestOptions::bearer(DEEPSEEK_API_KEY);
163    /// let mut response = request
164    ///     .get_stream_response_string(&default_client(), DEEPSEEK_CHAT_URL, &options)
165    ///     .await?;
166    ///
167    /// while let Some(chunk) = response.next().await {
168    ///     println!("{}", chunk?);
169    /// }
170    /// # Ok(())
171    /// # }
172    /// ```
173    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            // Parse the body as an SSE stream. Framing errors become stream
201            // items rather than terminating the stream.
202            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    /// Sends a streaming POST request and deserializes every SSE data item
216    /// into [`Self::Response`].
217    ///
218    /// The stream ends after the `data: [DONE]` sentinel. Deserialization
219    /// failures surface as [`OapiError::DeserializationError`] items without
220    /// ending the stream; wrap the stream with
221    /// [`crate::rest::skip_deserialization_errors`] to drop them instead.
222    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                // The `data: [DONE]` sentinel is not JSON; end the stream
241                // there. Everything else — including SSE framing errors —
242                // flows through to the consumer.
243                // `future::ready` keeps the combinators `Unpin` (async blocks
244                // are not).
245                .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
255/// Trait for POST requests whose response body is binary rather than JSON.
256///
257/// Endpoints such as `POST /audio/speech` return the raw audio bytes instead
258/// of a JSON object, so the body is returned as raw bytes.
259pub trait PostBinary: Post + Serialize + Sync + Send {
260    /// Sends a POST request with a JSON body and returns the raw response
261    /// body bytes.
262    ///
263    /// The `client` is supplied by the caller; see [`crate::rest::default_client`]
264    /// for a sensible default.
265    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    /// Sends a POST request with a JSON body and streams the raw response
287    /// body as byte chunks.
288    ///
289    /// Unlike [`Self::get_response_bytes`], the body is not buffered: each
290    /// item is one transport-level chunk. This matches the official SDK's
291    /// streaming binary responses (e.g. `POST /audio/speech` with
292    /// `stream_format`, where the chunks carry either raw audio or SSE
293    /// frames). Parsing any framing inside the byte stream is up to the
294    /// caller.
295    ///
296    /// The `client` is supplied by the caller; see [`crate::rest::default_client`]
297    /// for a sensible default.
298    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    /// Builds the default client and pins the SSE body used by the stream
345    /// tests: one good chunk, one chunk that is not valid JSON, one good
346    /// chunk, then the `[DONE]` sentinel.
347    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        // Azure-style: credentials travel in `api-key`, not `Authorization`.
435        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        // The extra body map is forwarded as multipart text fields.
487        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        // The multipart overrides must route through `build_url`, so the
494        // caller passes a bare base URL.
495        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}