Skip to main content

openai_interface/rest/
post.rs

1use std::{future::Future, str::FromStr};
2
3use eventsource_stream::Eventsource;
4use futures_util::{StreamExt, TryStreamExt, stream::BoxStream};
5use serde::{Serialize, de::DeserializeOwned};
6
7use crate::errors::OapiError;
8
9pub trait Post {
10    fn is_streaming(&self) -> bool;
11    /// Builds the URL for the request.
12    ///
13    /// `base_url` should be like <https://api.openai.com/v1>
14    fn build_url(&self, base_url: &str) -> Result<String, OapiError>;
15}
16
17pub trait PostNoStream: Post + Serialize + Sync + Send {
18    type Response: DeserializeOwned + FromStr<Err = OapiError> + Send + Sync;
19
20    /// Sends a POST request and returns the raw response body.
21    ///
22    /// The `client` is supplied by the caller; see [`crate::rest::default_client`]
23    /// for a sensible default.
24    fn get_response_string(
25        &self,
26        client: &reqwest::Client,
27        base_url: &str,
28        key: &str,
29    ) -> impl Future<Output = Result<String, OapiError>> + Send + Sync {
30        async move {
31            if self.is_streaming() {
32                return Err(OapiError::NonStreamingViolation);
33            }
34
35            let response = client
36                .post(self.build_url(base_url)?)
37                .header("Content-Type", "application/json")
38                .header("Accept", "application/json")
39                .bearer_auth(key)
40                .json(self)
41                .send()
42                .await?;
43
44            crate::rest::response_text_checked(response).await
45        }
46    }
47
48    /// Sends a POST request and deserializes the response.
49    fn get_response(
50        &self,
51        client: &reqwest::Client,
52        url: &str,
53        key: &str,
54    ) -> impl Future<Output = Result<Self::Response, OapiError>> + Send + Sync {
55        async move {
56            let text = self.get_response_string(client, url, key).await?;
57            let result = Self::Response::from_str(&text)?;
58            Ok(result)
59        }
60    }
61}
62
63pub trait PostStream: Post + Serialize + Sync + Send {
64    type Response: DeserializeOwned + FromStr<Err = OapiError> + Send + Sync;
65
66    /// Sends a streaming POST request and returns the raw SSE data items as
67    /// strings.
68    ///
69    /// The returned stream stops after the `data: [DONE]` sentinel. Errors
70    /// terminate the stream.
71    ///
72    /// The `client` is supplied by the caller; see [`crate::rest::default_client`]
73    /// for a sensible default.
74    ///
75    /// # Example
76    ///
77    /// ```rust,no_run
78    /// use std::sync::LazyLock;
79    /// use futures_util::StreamExt;
80    /// use openai_interface::chat::create::request::{Message, RequestBody};
81    /// use openai_interface::rest::{default_client, post::PostStream};
82    ///
83    /// const DEEPSEEK_API_KEY: &str = "YOUR_API_KEY";
84    /// const DEEPSEEK_CHAT_URL: &'static str = "https://api.deepseek.com";
85    /// const DEEPSEEK_MODEL: &'static str = "deepseek-chat";
86    ///
87    /// # async fn example() -> Result<(), Box<dyn std::error::Error>> {
88    /// let request = RequestBody {
89    ///     messages: vec![
90    ///         Message::System {
91    ///             content: "This is a request of test purpose. Reply briefly".to_string(),
92    ///             name: None,
93    ///         },
94    ///         Message::User {
95    ///             content: "What's your name?".into(),
96    ///             name: None,
97    ///         },
98    ///     ],
99    ///     model: DEEPSEEK_MODEL.to_string(),
100    ///     stream: Some(true),
101    ///     ..Default::default()
102    /// };
103    ///
104    /// let mut response = request
105    ///     .get_stream_response_string(&default_client(), DEEPSEEK_CHAT_URL, DEEPSEEK_API_KEY)
106    ///     .await?;
107    ///
108    /// while let Some(chunk) = response.next().await {
109    ///     println!("{}", chunk?);
110    /// }
111    /// # Ok(())
112    /// # }
113    /// ```
114    fn get_stream_response_string(
115        &self,
116        client: &reqwest::Client,
117        base_url: &str,
118        api_key: &str,
119    ) -> impl Future<Output = Result<BoxStream<'static, Result<String, OapiError>>, OapiError>>
120    + Send
121    + Sync {
122        async move {
123            if !self.is_streaming() {
124                return Err(OapiError::StreamingViolation);
125            }
126
127            let response = client
128                .post(self.build_url(base_url)?)
129                .header("Content-Type", "application/json")
130                .header("Accept", "text/event-stream")
131                .bearer_auth(api_key)
132                .json(self)
133                .send()
134                .await?;
135
136            // Parse the body as an SSE stream. Errors end the stream.
137            let stream = crate::rest::check_status(response)
138                .await?
139                .bytes_stream()
140                .eventsource()
141                .map(|event| match event {
142                    Ok(event) => Ok(event.data),
143                    Err(e) => Err(OapiError::SseParseError(format!("SSE parse error: {e}"))),
144                })
145                .boxed();
146
147            Ok(stream)
148        }
149    }
150
151    /// Sends a streaming POST request and deserializes every SSE data item
152    /// into [`Self::Response`].
153    ///
154    /// The stream ends after the `data: [DONE]` sentinel or on the first
155    /// error, whichever comes first.
156    fn get_stream_response(
157        &self,
158        client: &reqwest::Client,
159        base_url: &str,
160        api_key: &str,
161    ) -> impl Future<
162        Output = Result<BoxStream<'static, Result<Self::Response, OapiError>>, OapiError>,
163    > + Send
164    + Sync {
165        async move {
166            let stream = self
167                .get_stream_response_string(client, base_url, api_key)
168                .await?;
169
170            let parsed_stream = stream
171                // The `data: [DONE]` sentinel is not JSON; end the stream there.
172                .take_while(|result| {
173                    let should_continue = matches!(result, Ok(data) if data != "[DONE]");
174                    async move { should_continue }
175                })
176                .and_then(|data| async move { Self::Response::from_str(&data) });
177
178            Ok(Box::pin(parsed_stream) as BoxStream<'static, _>)
179        }
180    }
181}
182
183/// Trait for POST requests whose response body is binary rather than JSON.
184///
185/// Endpoints such as `POST /audio/speech` return the raw audio bytes instead
186/// of a JSON object, so the body is returned as raw bytes.
187pub trait PostBinary: Post + Serialize + Sync + Send {
188    /// Sends a POST request with a JSON body and returns the raw response
189    /// body bytes.
190    ///
191    /// The `client` is supplied by the caller; see [`crate::rest::default_client`]
192    /// for a sensible default.
193    fn get_response_bytes(
194        &self,
195        client: &reqwest::Client,
196        base_url: &str,
197        api_key: &str,
198    ) -> impl Future<Output = Result<Vec<u8>, OapiError>> + Send + Sync {
199        async move {
200            let response = client
201                .post(self.build_url(base_url)?)
202                .header("Content-Type", "application/json")
203                .header("Accept", "application/octet-stream")
204                .bearer_auth(api_key)
205                .json(self)
206                .send()
207                .await?;
208
209            crate::rest::response_bytes_checked(response).await
210        }
211    }
212
213    /// Sends a POST request with a JSON body and streams the raw response
214    /// body as byte chunks.
215    ///
216    /// Unlike [`Self::get_response_bytes`], the body is not buffered: each
217    /// item is one transport-level chunk. This matches the official SDK's
218    /// streaming binary responses (e.g. `POST /audio/speech` with
219    /// `stream_format`, where the chunks carry either raw audio or SSE
220    /// frames). Parsing any framing inside the byte stream is up to the
221    /// caller.
222    ///
223    /// The `client` is supplied by the caller; see [`crate::rest::default_client`]
224    /// for a sensible default.
225    fn get_stream_response_bytes(
226        &self,
227        client: &reqwest::Client,
228        base_url: &str,
229        api_key: &str,
230    ) -> impl Future<Output = Result<BoxStream<'static, Result<Vec<u8>, OapiError>>, OapiError>>
231    + Send
232    + Sync {
233        async move {
234            let response = client
235                .post(self.build_url(base_url)?)
236                .header("Content-Type", "application/json")
237                .header("Accept", "application/octet-stream")
238                .bearer_auth(api_key)
239                .json(self)
240                .send()
241                .await?;
242
243            let response = crate::rest::check_status(response).await?;
244
245            Ok(response
246                .bytes_stream()
247                .map(|chunk| {
248                    chunk
249                        .map(|bytes| bytes.to_vec())
250                        .map_err(OapiError::SendError)
251                })
252                .boxed())
253        }
254    }
255}