use std::{
future::Future,
pin::Pin,
time::{Duration, Instant},
};
use crate::{
Client, Error, ErrorKind, EventStream, Request,
http::{Backend, Hyper},
};
const BACKOFF_LIMIT: Duration = Duration::from_secs(30);
pub trait Layer<B: Backend = Hyper>: Send + Sync + 'static {
fn call(
&self,
request: Request,
next: Next<B>,
) -> impl Future<Output = Result<EventStream<B>, Error>> + Send;
}
pub struct Next<B: Backend = Hyper> {
client: Client<B>,
index: usize,
}
impl<B: Backend> Next<B> {
pub(crate) fn new(client: Client<B>, index: usize) -> Self {
Self { client, index }
}
pub async fn run(self, request: Request) -> Result<EventStream<B>, Error> {
self.client.run_from(self.index, request).await
}
}
impl<B: Backend> Clone for Next<B> {
fn clone(&self) -> Self {
Self {
client: self.client.clone(),
index: self.index,
}
}
}
impl<B: Backend> std::fmt::Debug for Next<B> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Next").field("index", &self.index).finish()
}
}
pub(crate) trait Erased<B: Backend>: Send + Sync {
fn call<'a>(
&'a self,
request: Request,
next: Next<B>,
) -> Pin<Box<dyn Future<Output = Result<EventStream<B>, Error>> + Send + 'a>>;
}
impl<B: Backend, L: Layer<B>> Erased<B> for L {
fn call<'a>(
&'a self,
request: Request,
next: Next<B>,
) -> Pin<Box<dyn Future<Output = Result<EventStream<B>, Error>> + Send + 'a>> {
Box::pin(Layer::call(self, request, next))
}
}
pub(crate) struct Wrap<F>(pub(crate) F);
impl<B, F, Fut> Layer<B> for Wrap<F>
where
B: Backend,
F: Fn(Request, Next<B>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<EventStream<B>, Error>> + Send,
{
fn call(
&self,
request: Request,
next: Next<B>,
) -> impl Future<Output = Result<EventStream<B>, Error>> + Send {
(self.0)(request, next)
}
}
#[derive(Debug, Clone)]
pub struct Retry {
retries: u32,
transient: bool,
backoff: Duration,
}
impl Retry {
pub fn connect(retries: u32) -> Self {
Self {
retries,
transient: false,
backoff: Duration::from_millis(500),
}
}
pub fn transient(retries: u32) -> Self {
Self {
transient: true,
..Self::connect(retries)
}
}
pub fn backoff(mut self, initial: Duration) -> Self {
self.backoff = initial;
self
}
fn covers(&self, error: &Error) -> bool {
error.is_unsent() || (self.transient && error.is_retryable())
}
}
impl<B: Backend> Layer<B> for Retry {
async fn call(&self, request: Request, next: Next<B>) -> Result<EventStream<B>, Error> {
let mut backoff = self.backoff;
for _ in 0..self.retries {
let error = match next.clone().run(request.clone()).await {
Ok(stream) => return Ok(stream),
Err(error) if self.covers(&error) => error,
Err(error) => return Err(error),
};
let wait = error.retry_after().unwrap_or(backoff).min(BACKOFF_LIMIT);
tokio::time::sleep(wait).await;
backoff = backoff.saturating_mul(2).min(BACKOFF_LIMIT);
}
next.run(request).await
}
}
#[derive(Debug, Clone, Copy)]
pub struct Timeout {
limit: Duration,
what: Limited,
}
#[derive(Debug, Clone, Copy)]
enum Limited {
FirstToken,
Idle,
Total,
}
impl Timeout {
pub fn first_token(limit: Duration) -> Self {
Self {
limit,
what: Limited::FirstToken,
}
}
pub fn idle(limit: Duration) -> Self {
Self {
limit,
what: Limited::Idle,
}
}
pub fn total(limit: Duration) -> Self {
Self {
limit,
what: Limited::Total,
}
}
}
impl<B: Backend> Layer<B> for Timeout {
async fn call(&self, request: Request, next: Next<B>) -> Result<EventStream<B>, Error> {
let deadline = Instant::now() + self.limit;
let stream = tokio::time::timeout(self.limit, next.run(request))
.await
.map_err(|_| {
Error::new(ErrorKind::Timeout).with_detail("the server did not answer in time")
})??;
Ok(match self.what {
Limited::FirstToken => stream.first_token_by(deadline),
Limited::Idle => stream.idle_timeout(self.limit),
Limited::Total => stream.complete_by(deadline),
})
}
}
#[cfg(feature = "tracing")]
#[derive(Debug, Clone, Copy, Default)]
pub struct Trace;
#[cfg(feature = "tracing")]
impl<B: Backend> Layer<B> for Trace {
async fn call(&self, request: Request, next: Next<B>) -> Result<EventStream<B>, Error> {
use crate::Event;
let started = Instant::now();
let model = request.model.clone();
tracing::debug!(model = %model, messages = request.messages.len(), "request");
let stream = match next.run(request).await {
Ok(stream) => stream,
Err(error) => {
tracing::warn!(
model = %model,
kind = ?error.kind(),
elapsed_ms = started.elapsed().as_millis() as u64,
"request failed"
);
return Err(error);
}
};
tracing::debug!(
model = %model,
elapsed_ms = started.elapsed().as_millis() as u64,
"response started"
);
Ok(stream.inspect(move |item| match item {
Ok(Event::Completed(done)) => tracing::info!(
model = %model,
finish = ?done.finish,
input_tokens = done.usage.map(|usage| usage.input),
output_tokens = done.usage.map(|usage| usage.output),
elapsed_ms = started.elapsed().as_millis() as u64,
"answer complete"
),
Err(error) => tracing::warn!(
model = %model,
kind = ?error.kind(),
elapsed_ms = started.elapsed().as_millis() as u64,
"answer failed"
),
Ok(_) => {}
}))
}
}