use std::{future::Future, str::FromStr};
use eventsource_stream::Eventsource;
use futures_util::{Stream, StreamExt, TryStreamExt};
use serde::{Serialize, de::DeserializeOwned};
use crate::errors::OapiError;
use crate::rest::{Auth, RequestOptions};
pub trait Post {
fn is_streaming(&self) -> bool;
fn build_url(&self, base_url: &str) -> Result<String, OapiError>;
}
pub(crate) fn apply_options(
builder: reqwest::RequestBuilder,
options: &RequestOptions,
) -> reqwest::RequestBuilder {
let builder = match &options.auth {
Auth::Bearer(token) => builder.bearer_auth(token),
Auth::None => builder,
};
builder.headers(options.extra_headers.clone())
}
pub(crate) async fn post_multipart_json(
client: &reqwest::Client,
url: String,
form: reqwest::multipart::Form,
options: &RequestOptions,
) -> Result<String, OapiError> {
let response = apply_options(
client.post(url).header("Accept", "application/json"),
options,
)
.multipart(form)
.send()
.await?;
crate::rest::response_text_checked(response).await
}
pub(crate) fn append_extra_body_map(
mut form: reqwest::multipart::Form,
extra_body_map: &Option<serde_json::Map<String, serde_json::Value>>,
) -> reqwest::multipart::Form {
if let Some(map) = extra_body_map {
for (key, value) in map {
let text = match value {
serde_json::Value::String(s) => s.clone(),
other => other.to_string(),
};
form = form.text(key.clone(), text);
}
}
form
}
pub trait PostNoStream: Post + Serialize + Sync + Send {
type Response: DeserializeOwned + FromStr<Err = OapiError> + Send + Sync;
fn get_response_string(
&self,
client: &reqwest::Client,
base_url: &str,
options: &RequestOptions,
) -> impl Future<Output = Result<String, OapiError>> + Send + Sync {
async move {
if self.is_streaming() {
return Err(OapiError::NonStreamingViolation);
}
let response = apply_options(
client
.post(self.build_url(base_url)?)
.header("Accept", "application/json"),
options,
)
.json(self)
.send()
.await?;
crate::rest::response_text_checked(response).await
}
}
fn get_response(
&self,
client: &reqwest::Client,
url: &str,
options: &RequestOptions,
) -> impl Future<Output = Result<Self::Response, OapiError>> + Send + Sync {
async move {
let text = self.get_response_string(client, url, options).await?;
let result = Self::Response::from_str(&text)?;
Ok(result)
}
}
}
pub trait PostStream: Post + Serialize + Sync + Send {
type Response: DeserializeOwned + FromStr<Err = OapiError> + Send + Sync + 'static;
fn get_stream_response_string(
&self,
client: &reqwest::Client,
base_url: &str,
options: &RequestOptions,
) -> impl Future<
Output = Result<
impl Stream<Item = Result<String, OapiError>> + Send + Unpin + 'static,
OapiError,
>,
> + Send
+ Sync {
async move {
if !self.is_streaming() {
return Err(OapiError::StreamingViolation);
}
let response = apply_options(
client
.post(self.build_url(base_url)?)
.header("Accept", "text/event-stream"),
options,
)
.json(self)
.send()
.await?;
let stream = crate::rest::check_status(response)
.await?
.bytes_stream()
.eventsource()
.map(|event| match event {
Ok(event) => Ok(event.data),
Err(e) => Err(OapiError::SseParseError(format!("SSE parse error: {e}"))),
});
Ok(stream)
}
}
fn get_stream_response(
&self,
client: &reqwest::Client,
base_url: &str,
options: &RequestOptions,
) -> impl Future<
Output = Result<
impl Stream<Item = Result<Self::Response, OapiError>> + Send + Unpin + 'static,
OapiError,
>,
> + Send
+ Sync {
async move {
let stream = self
.get_stream_response_string(client, base_url, options)
.await?;
let parsed_stream = stream
.take_while(|result| {
std::future::ready(!matches!(result, Ok(data) if data == "[DONE]"))
})
.and_then(|data| std::future::ready(Self::Response::from_str(&data)));
Ok(parsed_stream)
}
}
}
pub trait PostBinary: Post + Serialize + Sync + Send {
fn get_response_bytes(
&self,
client: &reqwest::Client,
base_url: &str,
options: &RequestOptions,
) -> impl Future<Output = Result<Vec<u8>, OapiError>> + Send + Sync {
async move {
let response = apply_options(
client
.post(self.build_url(base_url)?)
.header("Accept", "application/octet-stream"),
options,
)
.json(self)
.send()
.await?;
crate::rest::response_bytes_checked(response).await
}
}
fn get_stream_response_bytes(
&self,
client: &reqwest::Client,
base_url: &str,
options: &RequestOptions,
) -> impl Future<
Output = Result<
impl Stream<Item = Result<Vec<u8>, OapiError>> + Send + Unpin + 'static,
OapiError,
>,
> + Send
+ Sync {
async move {
let response = apply_options(
client
.post(self.build_url(base_url)?)
.header("Accept", "application/octet-stream"),
options,
)
.json(self)
.send()
.await?;
let response = crate::rest::check_status(response).await?;
Ok(response.bytes_stream().map(|chunk| {
chunk
.map(|bytes| bytes.to_vec())
.map_err(OapiError::SendError)
}))
}
}
}
#[cfg(test)]
mod test {
use futures_util::StreamExt;
use serde::Deserialize;
use serde_json::json;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use super::*;
use crate::chat::create::response::streaming::ChatCompletionChunk;
use crate::rest::{RequestOptions, skip_deserialization_errors};
fn sse_body() -> String {
const CHUNK_A: &str = r#"{"id":"1","choices":[{"index":0,"delta":{"content":"a"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#;
const CHUNK_B: &str = r#"{"id":"1","choices":[{"index":0,"delta":{"content":"b"},"finish_reason":null}],"created":1,"model":"m","object":"chat.completion.chunk"}"#;
format!("data: {CHUNK_A}\n\ndata: oops\n\ndata: {CHUNK_B}\n\ndata: [DONE]\n\n")
}
fn chunk_content(chunk: &ChatCompletionChunk) -> String {
chunk.choices[0].delta.content.clone().unwrap_or_default()
}
#[derive(Serialize)]
struct TestJsonRequest;
impl Post for TestJsonRequest {
fn is_streaming(&self) -> bool {
false
}
fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
Ok(format!("{}/test", base_url.trim_end_matches('/')))
}
}
impl PostNoStream for TestJsonRequest {
type Response = TestResponse;
}
#[derive(Debug, Deserialize)]
struct TestResponse {
#[allow(dead_code)]
id: String,
}
crate::impl_from_str!(TestResponse);
#[derive(Serialize)]
struct TestStreamRequest;
impl Post for TestStreamRequest {
fn is_streaming(&self) -> bool {
true
}
fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
Ok(format!("{}/test", base_url.trim_end_matches('/')))
}
}
impl PostStream for TestStreamRequest {
type Response = ChatCompletionChunk;
}
#[tokio::test]
async fn sends_bearer_auth_and_extra_headers() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/test"))
.respond_with(
ResponseTemplate::new(200).set_body_string(json!({"id": "x"}).to_string()),
)
.mount(&server)
.await;
let options = RequestOptions::bearer("sk-test")
.with_header("OpenAI-Organization", "org-1")
.unwrap();
TestJsonRequest
.get_response_string(&crate::rest::default_client(), &server.uri(), &options)
.await
.expect("request must succeed");
let requests = server.received_requests().await.expect("recorded requests");
assert_eq!(requests.len(), 1);
let headers = &requests[0].headers;
assert_eq!(headers.get("authorization").unwrap(), "Bearer sk-test");
assert_eq!(headers.get("openai-organization").unwrap(), "org-1");
}
#[tokio::test]
async fn header_only_auth_sends_no_authorization() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/test"))
.respond_with(
ResponseTemplate::new(200).set_body_string(json!({"id": "x"}).to_string()),
)
.mount(&server)
.await;
let options = RequestOptions::new()
.with_header("api-key", "azure-key")
.unwrap();
TestJsonRequest
.get_response_string(&crate::rest::default_client(), &server.uri(), &options)
.await
.expect("request must succeed");
let requests = server.received_requests().await.expect("recorded requests");
assert_eq!(requests.len(), 1);
let headers = &requests[0].headers;
assert!(headers.get("authorization").is_none());
assert_eq!(headers.get("api-key").unwrap(), "azure-key");
}
#[tokio::test]
async fn multipart_upload_uses_build_url_and_options() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/files"))
.respond_with(
ResponseTemplate::new(200).set_body_string(json!({"id": "x"}).to_string()),
)
.mount(&server)
.await;
let path = std::env::temp_dir().join("openai_interface_multipart_test.txt");
std::fs::write(&path, b"hello").expect("test file must be writable");
let mut extra = serde_json::Map::new();
extra.insert(
"vendor_extension".to_string(),
serde_json::json!({"depth": 3}),
);
let request = crate::files::create::request::CreateFileRequest {
file: path,
purpose: crate::files::create::request::FilePurpose::Batch,
expires_after: None,
extra_body_map: Some(extra),
};
let options = RequestOptions::new()
.with_header("api-key", "azure-key")
.unwrap();
request
.get_response_string(&crate::rest::default_client(), &server.uri(), &options)
.await
.expect("upload must succeed");
let requests = server.received_requests().await.expect("recorded requests");
assert_eq!(requests.len(), 1);
let body = String::from_utf8_lossy(&requests[0].body);
assert!(
body.contains("vendor_extension") && body.contains("\"depth\":3"),
"extra body map must reach the multipart body: {body}"
);
let headers = &requests[0].headers;
assert_eq!(headers.get("api-key").unwrap(), "azure-key");
assert_eq!(requests[0].url.path(), "/files");
}
#[tokio::test]
async fn raw_stream_yields_every_data_item_including_done() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_string(sse_body()))
.mount(&server)
.await;
let stream = TestStreamRequest
.get_stream_response_string(
&crate::rest::default_client(),
&server.uri(),
&RequestOptions::bearer("sk-test"),
)
.await
.expect("stream must start");
let items: Vec<Result<String, OapiError>> = stream.collect().await;
let data: Vec<&str> = items
.iter()
.map(|item| item.as_ref().expect("raw items must not fail").as_str())
.collect();
assert_eq!(data.len(), 4);
assert_eq!(data[3], "[DONE]");
}
#[tokio::test]
async fn parsed_stream_surfaces_bad_chunks_without_ending() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_string(sse_body()))
.mount(&server)
.await;
let stream = TestStreamRequest
.get_stream_response(
&crate::rest::default_client(),
&server.uri(),
&RequestOptions::bearer("sk-test"),
)
.await
.expect("stream must start");
let items: Vec<Result<ChatCompletionChunk, OapiError>> = stream.collect().await;
assert_eq!(items.len(), 3, "good, bad, good; [DONE] ends the stream");
assert_eq!(chunk_content(items[0].as_ref().unwrap()), "a");
assert!(matches!(items[1], Err(OapiError::DeserializationError(_))));
assert_eq!(chunk_content(items[2].as_ref().unwrap()), "b");
}
#[tokio::test]
async fn skip_deserialization_errors_drops_bad_chunks() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_string(sse_body()))
.mount(&server)
.await;
let stream = TestStreamRequest
.get_stream_response(
&crate::rest::default_client(),
&server.uri(),
&RequestOptions::bearer("sk-test"),
)
.await
.expect("stream must start");
let mut stream = skip_deserialization_errors(stream);
let mut contents = Vec::new();
while let Some(item) = stream.next().await {
contents.push(chunk_content(&item.expect("errors must be skipped")));
}
assert_eq!(contents, ["a", "b"]);
}
}