use async_trait::async_trait;
use crate::types::{AiLibError, ChatCompletionRequest, ChatCompletionResponse};
pub mod breaker;
pub mod default;
pub mod rate_limit;
pub mod retry;
pub mod timeout;
pub use breaker::CircuitBreakerInterceptor;
pub use default::{create_default_interceptors, DefaultInterceptorsBuilder};
pub use rate_limit::RateLimitInterceptor;
pub use retry::RetryInterceptor;
pub use timeout::TimeoutInterceptor;
#[derive(Debug, Clone)]
pub struct RequestContext {
pub provider: String,
pub model: String,
}
#[derive(Debug, Clone)]
pub struct ResponseContext {
pub success: bool,
}
#[async_trait]
pub trait Interceptor: Send + Sync {
async fn on_request(&self, _ctx: &RequestContext, _req: &ChatCompletionRequest) {}
async fn on_response(
&self,
_ctx: &RequestContext,
_req: &ChatCompletionRequest,
_resp: &ChatCompletionResponse,
) {
}
async fn on_error(
&self,
_ctx: &RequestContext,
_req: &ChatCompletionRequest,
_err: &AiLibError,
) {
}
}
pub struct InterceptorPipeline {
pub(crate) interceptors: Vec<Box<dyn Interceptor>>,
}
impl InterceptorPipeline {
pub fn new() -> Self {
Self {
interceptors: Vec::new(),
}
}
pub fn with<I: Interceptor + 'static>(mut self, interceptor: I) -> Self {
self.interceptors.push(Box::new(interceptor));
self
}
pub async fn execute<F, Fut>(
&self,
ctx: &RequestContext,
req: &ChatCompletionRequest,
f: F,
) -> Result<ChatCompletionResponse, AiLibError>
where
F: FnOnce() -> Fut,
Fut: std::future::Future<Output = Result<ChatCompletionResponse, AiLibError>>,
{
for ic in &self.interceptors {
ic.on_request(ctx, req).await;
}
match f().await {
Ok(resp) => {
for ic in &self.interceptors {
ic.on_response(ctx, req, &resp).await;
}
Ok(resp)
}
Err(err) => {
for ic in &self.interceptors {
ic.on_error(ctx, req, &err).await;
}
Err(err)
}
}
}
}
impl Default for InterceptorPipeline {
fn default() -> Self {
Self::new()
}
}