use std::{future::Future, str::FromStr};
use eventsource_stream::Eventsource;
use futures_util::{StreamExt, TryStreamExt, stream::BoxStream};
use serde::{Serialize, de::DeserializeOwned};
use crate::errors::OapiError;
pub trait Post {
fn is_streaming(&self) -> bool;
fn build_url(&self, base_url: &str) -> Result<String, OapiError>;
}
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,
key: &str,
) -> impl Future<Output = Result<String, OapiError>> + Send + Sync {
async move {
if self.is_streaming() {
return Err(OapiError::NonStreamingViolation);
}
let response = client
.post(self.build_url(base_url)?)
.header("Content-Type", "application/json")
.header("Accept", "application/json")
.bearer_auth(key)
.json(self)
.send()
.await?;
crate::rest::response_text_checked(response).await
}
}
fn get_response(
&self,
client: &reqwest::Client,
url: &str,
key: &str,
) -> impl Future<Output = Result<Self::Response, OapiError>> + Send + Sync {
async move {
let text = self.get_response_string(client, url, key).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;
fn get_stream_response_string(
&self,
client: &reqwest::Client,
base_url: &str,
api_key: &str,
) -> impl Future<Output = Result<BoxStream<'static, Result<String, OapiError>>, OapiError>>
+ Send
+ Sync {
async move {
if !self.is_streaming() {
return Err(OapiError::StreamingViolation);
}
let response = client
.post(self.build_url(base_url)?)
.header("Content-Type", "application/json")
.header("Accept", "text/event-stream")
.bearer_auth(api_key)
.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}"))),
})
.boxed();
Ok(stream)
}
}
fn get_stream_response(
&self,
client: &reqwest::Client,
base_url: &str,
api_key: &str,
) -> impl Future<
Output = Result<BoxStream<'static, Result<Self::Response, OapiError>>, OapiError>,
> + Send
+ Sync {
async move {
let stream = self
.get_stream_response_string(client, base_url, api_key)
.await?;
let parsed_stream = stream
.take_while(|result| {
let should_continue = matches!(result, Ok(data) if data != "[DONE]");
async move { should_continue }
})
.and_then(|data| async move { Self::Response::from_str(&data) });
Ok(Box::pin(parsed_stream) as BoxStream<'static, _>)
}
}
}
pub trait PostBinary: Post + Serialize + Sync + Send {
fn get_response_bytes(
&self,
client: &reqwest::Client,
base_url: &str,
api_key: &str,
) -> impl Future<Output = Result<Vec<u8>, OapiError>> + Send + Sync {
async move {
let response = client
.post(self.build_url(base_url)?)
.header("Content-Type", "application/json")
.header("Accept", "application/octet-stream")
.bearer_auth(api_key)
.json(self)
.send()
.await?;
crate::rest::response_bytes_checked(response).await
}
}
}