mod retry;
pub use retry::{RetryLayer, RetryLayerService, TelemetryLayer, TelemetryLayerService};
use std::pin::Pin;
use async_trait::async_trait;
use futures::Stream;
use crate::error::ProviderError;
use crate::traits::LlmProvider;
use crate::types::{CompletionRequest, CompletionResponse, ModelInfo, RequestOptions, StreamEvent};
pub trait ProviderLayer<S: LlmProvider> {
type Stack: LlmProvider;
fn wrap(self, inner: S) -> Self::Stack;
}
pub struct Layered<L, S> {
layer: L,
inner: S,
}
impl<L, S: LlmProvider> Layered<L, S> {
pub fn new(layer: L, inner: S) -> Self {
Self { layer, inner }
}
}
impl<L, S: LlmProvider> std::fmt::Debug for Layered<L, S>
where
S: std::fmt::Debug,
{
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Layered")
.field("inner", &self.inner)
.finish()
}
}
#[async_trait]
impl<L, S: LlmProvider> LlmProvider for Layered<L, S>
where
L: LayerService<S> + Send + Sync,
{
async fn complete(
&self,
request: CompletionRequest,
options: RequestOptions,
) -> Result<CompletionResponse, ProviderError> {
self.layer.complete(&self.inner, request, options).await
}
async fn complete_stream(
&self,
request: CompletionRequest,
options: RequestOptions,
) -> Result<
Pin<Box<dyn Stream<Item = Result<StreamEvent, ProviderError>> + Send>>,
ProviderError,
> {
self.layer.complete_stream(&self.inner, request, options).await
}
fn models(&self) -> &[ModelInfo] {
self.inner.models()
}
fn name(&self) -> &str {
self.inner.name()
}
}
#[async_trait]
pub trait LayerService<S: LlmProvider>: Send + Sync {
async fn complete(
&self,
inner: &S,
request: CompletionRequest,
options: RequestOptions,
) -> Result<CompletionResponse, ProviderError>;
async fn complete_stream(
&self,
inner: &S,
request: CompletionRequest,
options: RequestOptions,
) -> Result<
Pin<Box<dyn Stream<Item = Result<StreamEvent, ProviderError>> + Send>>,
ProviderError,
>;
}