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
52pub trait PostNoStream: Post + Serialize + Sync + Send {
53    type Response: DeserializeOwned + FromStr<Err = OapiError> + Send + Sync;
54
55    /// Sends a POST request and returns the raw response body.
56    ///
57    /// The `client` is supplied by the caller; see [`crate::rest::default_client`]
58    /// for a sensible default.
59    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    /// Sends a POST request and deserializes the response.
85    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    /// Sends a streaming POST request and returns the raw SSE data items as
103    /// strings.
104    ///
105    /// The returned stream yields every SSE data payload verbatim, including
106    /// the final `"[DONE]"` sentinel; use
107    /// [`Self::get_stream_response`] when you want the sentinel consumed and
108    /// the items deserialized. SSE framing errors surface as
109    /// [`OapiError::SseParseError`] items.
110    ///
111    /// The `client` is supplied by the caller; see [`crate::rest::default_client`]
112    /// for a sensible default.
113    ///
114    /// # Example
115    ///
116    /// ```rust,no_run
117    /// use std::sync::LazyLock;
118    /// use futures_util::StreamExt;
119    /// use openai_interface::chat::create::request::{Message, RequestBody};
120    /// use openai_interface::rest::{default_client, post::PostStream, RequestOptions};
121    ///
122    /// const DEEPSEEK_API_KEY: &str = "YOUR_API_KEY";
123    /// const DEEPSEEK_CHAT_URL: &'static str = "https://api.deepseek.com";
124    /// const DEEPSEEK_MODEL: &'static str = "deepseek-chat";
125    ///
126    /// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
127    /// // Needs the `ferritls` cargo feature; drop this line if you install
128    /// // your own rustls crypto provider (see `openai_interface::rest`).
129    /// # #[cfg(feature = "ferritls")]
130    /// openai_interface::rest::install_crypto_provider().ok();
131    ///
132    /// let request = RequestBody {
133    ///     messages: vec![
134    ///         Message::System {
135    ///             content: "This is a request of test purpose. Reply briefly".into(),
136    ///             name: None,
137    ///         },
138    ///         Message::User {
139    ///             content: "What's your name?".into(),
140    ///             name: None,
141    ///         },
142    ///     ],
143    ///     model: DEEPSEEK_MODEL.to_string(),
144    ///     stream: Some(true),
145    ///     ..Default::default()
146    /// };
147    ///
148    /// let options = RequestOptions::bearer(DEEPSEEK_API_KEY);
149    /// let mut response = request
150    ///     .get_stream_response_string(&default_client(), DEEPSEEK_CHAT_URL, &options)
151    ///     .await?;
152    ///
153    /// while let Some(chunk) = response.next().await {
154    ///     println!("{}", chunk?);
155    /// }
156    /// # Ok(())
157    /// # }
158    /// ```
159    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            // Parse the body as an SSE stream. Framing errors become stream
187            // items rather than terminating the stream.
188            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    /// Sends a streaming POST request and deserializes every SSE data item
202    /// into [`Self::Response`].
203    ///
204    /// The stream ends after the `data: [DONE]` sentinel. Deserialization
205    /// failures surface as [`OapiError::DeserializationError`] items without
206    /// ending the stream; wrap the stream with
207    /// [`crate::rest::skip_deserialization_errors`] to drop them instead.
208    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                // The `data: [DONE]` sentinel is not JSON; end the stream
227                // there. Everything else — including SSE framing errors —
228                // flows through to the consumer.
229                // `future::ready` keeps the combinators `Unpin` (async blocks
230                // are not).
231                .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
241/// Trait for POST requests whose response body is binary rather than JSON.
242///
243/// Endpoints such as `POST /audio/speech` return the raw audio bytes instead
244/// of a JSON object, so the body is returned as raw bytes.
245pub trait PostBinary: Post + Serialize + Sync + Send {
246    /// Sends a POST request with a JSON body and returns the raw response
247    /// body bytes.
248    ///
249    /// The `client` is supplied by the caller; see [`crate::rest::default_client`]
250    /// for a sensible default.
251    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    /// Sends a POST request with a JSON body and streams the raw response
273    /// body as byte chunks.
274    ///
275    /// Unlike [`Self::get_response_bytes`], the body is not buffered: each
276    /// item is one transport-level chunk. This matches the official SDK's
277    /// streaming binary responses (e.g. `POST /audio/speech` with
278    /// `stream_format`, where the chunks carry either raw audio or SSE
279    /// frames). Parsing any framing inside the byte stream is up to the
280    /// caller.
281    ///
282    /// The `client` is supplied by the caller; see [`crate::rest::default_client`]
283    /// for a sensible default.
284    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    /// Builds the default client and pins the SSE body used by the stream
331    /// tests: one good chunk, one chunk that is not valid JSON, one good
332    /// chunk, then the `[DONE]` sentinel.
333    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        // Azure-style: credentials travel in `api-key`, not `Authorization`.
421        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        // The multipart overrides must route through `build_url`, so the
468        // caller passes a bare base URL.
469        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}