use crate::client::ClientRequest;
use crate::errors::SDKError;
use crate::response::ClientResponse;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
pub type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
pub type NextFn<'a> = Box<
dyn Fn(ClientRequest) -> BoxFuture<'a, Result<ClientResponse, SDKError>> + Send + Sync + 'a,
>;
pub trait Middleware: Send + Sync {
fn execute<'a>(
&'a self,
request: ClientRequest,
next: NextFn<'a>,
) -> BoxFuture<'a, Result<ClientResponse, SDKError>>;
}
pub struct MiddlewarePipeline {
middlewares: Vec<Arc<dyn Middleware>>,
}
impl Default for MiddlewarePipeline {
fn default() -> Self {
Self::new()
}
}
impl MiddlewarePipeline {
#[must_use]
pub fn new() -> Self {
Self {
middlewares: Vec::new(),
}
}
pub fn add<M: Middleware + 'static>(&mut self, middleware: M) {
self.middlewares.push(Arc::new(middleware));
}
#[must_use]
pub fn len(&self) -> usize {
self.middlewares.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.middlewares.is_empty()
}
#[allow(clippy::missing_errors_doc)]
pub async fn execute<F, Fut>(
&self,
request: ClientRequest,
final_handler: F,
) -> Result<ClientResponse, SDKError>
where
F: Fn(ClientRequest) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<ClientResponse, SDKError>> + Send + 'static,
{
let mut next: NextFn<'_> = Box::new(move |req| Box::pin(final_handler(req)));
for middleware in self.middlewares.iter().rev() {
let next_shared: Arc<NextFn<'_>> = Arc::new(next);
let m: &dyn Middleware = &**middleware;
next = Box::new(move |req| {
let next_shared = Arc::clone(&next_shared);
m.execute(req, Box::new(move |req2| next_shared(req2)))
});
}
next(request).await
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::conversation::Conversation;
use crate::message::{Message, Role};
use std::sync::atomic::{AtomicUsize, Ordering};
struct TestMiddleware {
calls: Arc<AtomicUsize>,
}
impl Middleware for TestMiddleware {
fn execute<'a>(
&'a self,
request: ClientRequest,
next: NextFn<'a>,
) -> BoxFuture<'a, Result<ClientResponse, SDKError>> {
self.calls.fetch_add(1, Ordering::SeqCst);
Box::pin(async move {
let mut res = next(request).await?;
let mut new_msg = res.message.clone();
let mut content = new_msg.content().to_string();
content.push_str(" (modified)");
new_msg = crate::message::Message::new(new_msg.role().clone(), content);
res.message = new_msg;
Ok(res)
})
}
}
#[tokio::test]
async fn test_middleware_pipeline() {
let mut pipeline = MiddlewarePipeline::new();
let calls = Arc::new(AtomicUsize::new(0));
pipeline.add(TestMiddleware {
calls: Arc::clone(&calls),
});
let request = ClientRequest::new(Conversation::new());
let handler = |_req: ClientRequest| async move {
Ok(ClientResponse::new(Message::new(
Role::Assistant,
"Original",
)))
};
let res = pipeline.execute(request, handler).await.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 1);
assert_eq!(res.message.content(), "Original (modified)");
}
struct RecordingMiddleware {
name: &'static str,
order: Arc<std::sync::Mutex<Vec<String>>>,
}
impl Middleware for RecordingMiddleware {
fn execute<'a>(
&'a self,
request: ClientRequest,
next: NextFn<'a>,
) -> BoxFuture<'a, Result<ClientResponse, SDKError>> {
self.order
.lock()
.unwrap()
.push(format!("{}-enter", self.name));
Box::pin(async move {
let res = next(request).await;
self.order
.lock()
.unwrap()
.push(format!("{}-exit", self.name));
res
})
}
}
#[tokio::test]
async fn test_middleware_pipeline_multiple_middlewares_order() {
let order = Arc::new(std::sync::Mutex::new(Vec::new()));
let mut pipeline = MiddlewarePipeline::new();
pipeline.add(RecordingMiddleware {
name: "A",
order: Arc::clone(&order),
});
pipeline.add(RecordingMiddleware {
name: "B",
order: Arc::clone(&order),
});
let request = ClientRequest::new(Conversation::new());
let handler = |_req: ClientRequest| async move {
Ok(ClientResponse::new(Message::new(
Role::Assistant,
"Success",
)))
};
pipeline.execute(request, handler).await.unwrap();
assert_eq!(
*order.lock().unwrap(),
vec!["A-enter", "B-enter", "B-exit", "A-exit"]
);
}
#[tokio::test]
async fn test_middleware_pipeline_retry_with_logging_middleware() {
use crate::middleware::logging::LoggingMiddleware;
use crate::middleware::retry::{RetryConfig, RetryMiddleware};
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
let mut pipeline = MiddlewarePipeline::new();
pipeline.add(LoggingMiddleware::new());
pipeline.add(RetryMiddleware::with_config(RetryConfig {
max_attempts: 3,
delay: Duration::from_millis(1),
backoff_multiplier: 1.0,
}));
let request = ClientRequest::new(Conversation::new());
let calls = Arc::new(AtomicUsize::new(0));
let calls_clone = Arc::clone(&calls);
let handler = move |_req: ClientRequest| {
let calls = Arc::clone(&calls_clone);
async move {
let attempt = calls.fetch_add(1, Ordering::SeqCst) + 1;
if attempt < 3 {
Err(SDKError::Api(crate::errors::ApiError::with_status(
"server error",
500,
)))
} else {
Ok(ClientResponse::new(Message::new(
Role::Assistant,
"Success",
)))
}
}
};
let res = pipeline.execute(request, handler).await.unwrap();
assert_eq!(res.message.content(), "Success");
assert_eq!(calls.load(Ordering::SeqCst), 3);
}
}
pub mod logging;
pub mod retry;