use std::{
collections::VecDeque,
fmt,
future::{Future, poll_fn},
pin::Pin,
task::{Context, Poll},
time::Duration,
};
use bytes::Bytes;
use futures_core::Stream;
use tokio::time::{Instant, Sleep};
use crate::{
Completion, Error, ErrorKind, Event,
http::{Backend, Hyper},
openai::chat::{Decoder, truncated},
};
type Timer = Pin<Box<Sleep>>;
type Inspector = Box<dyn FnMut(&Result<Event, Error>) + Send>;
pub struct RawStream<B: Backend = Hyper> {
body: Option<B::Body>,
idle: Option<Idle>,
}
struct Idle {
limit: Duration,
timer: Timer,
}
impl<B: Backend> RawStream<B> {
pub(crate) fn new(body: B::Body, idle: Option<Duration>) -> Self {
let mut stream = Self {
body: Some(body),
idle: None,
};
if let Some(limit) = idle {
stream.tighten_idle(limit);
}
stream
}
pub async fn next(&mut self) -> Option<Result<Bytes, Error>> {
poll_fn(|cx| Pin::new(&mut *self).poll_next(cx)).await
}
fn tighten_idle(&mut self, limit: Duration) {
if self.idle.as_ref().is_some_and(|idle| idle.limit <= limit) {
return;
}
self.idle = Some(Idle {
limit,
timer: Box::pin(tokio::time::sleep(limit)),
});
}
}
impl<B: Backend> Stream for RawStream<B> {
type Item = Result<Bytes, Error>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = &mut *self;
let Some(body) = this.body.as_mut() else {
return Poll::Ready(None);
};
match Pin::new(body).poll_next(cx) {
Poll::Ready(Some(Ok(bytes))) => {
if let Some(idle) = &mut this.idle {
idle.timer.as_mut().reset(Instant::now() + idle.limit);
}
Poll::Ready(Some(Ok(bytes)))
}
Poll::Ready(Some(Err(error))) => {
this.body = None;
Poll::Ready(Some(Err(error)))
}
Poll::Ready(None) => {
this.body = None;
Poll::Ready(None)
}
Poll::Pending => {
let stalled = this
.idle
.as_mut()
.is_some_and(|idle| idle.timer.as_mut().poll(cx).is_ready());
if !stalled {
return Poll::Pending;
}
this.body = None;
Poll::Ready(Some(Err(timeout("the response stalled"))))
}
}
}
}
impl<B: Backend> fmt::Debug for RawStream<B> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RawStream").finish_non_exhaustive()
}
}
pub struct EventStream<B: Backend = Hyper> {
bytes: Option<RawStream<B>>,
decoder: Option<Decoder>,
queue: VecDeque<Result<Event, Error>>,
first_token: Option<Timer>,
complete: Option<Timer>,
inspectors: Vec<Inspector>,
}
impl<B: Backend> EventStream<B> {
pub(crate) fn new(bytes: RawStream<B>, decoder: Decoder) -> Self {
Self {
bytes: Some(bytes),
decoder: Some(decoder),
queue: VecDeque::new(),
first_token: None,
complete: None,
inspectors: Vec::new(),
}
}
pub async fn next(&mut self) -> Option<Result<Event, Error>> {
poll_fn(|cx| Pin::new(&mut *self).poll_next(cx)).await
}
pub async fn completion(mut self) -> Result<Completion, Error> {
while let Some(event) = self.next().await {
if let Event::Completed(done) = event? {
return Ok(done);
}
}
Err(truncated())
}
pub fn first_token_by(mut self, deadline: std::time::Instant) -> Self {
earlier(&mut self.first_token, deadline);
self
}
pub fn complete_by(mut self, deadline: std::time::Instant) -> Self {
earlier(&mut self.complete, deadline);
self
}
pub fn idle_timeout(mut self, limit: Duration) -> Self {
if let Some(bytes) = self.bytes.as_mut() {
bytes.tighten_idle(limit);
}
self
}
pub fn inspect(mut self, watch: impl FnMut(&Result<Event, Error>) + Send + 'static) -> Self {
self.inspectors.push(Box::new(watch));
self
}
fn deliver(&mut self, item: Result<Event, Error>) -> Poll<Option<Result<Event, Error>>> {
if item.is_ok() {
self.first_token = None;
}
if matches!(item, Ok(Event::Completed(_)) | Err(_)) {
self.bytes = None;
self.decoder = None;
self.queue.clear();
self.first_token = None;
self.complete = None;
}
for watch in &mut self.inspectors {
watch(&item);
}
Poll::Ready(Some(item))
}
}
impl<B: Backend> Stream for EventStream<B> {
type Item = Result<Event, Error>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = &mut *self;
loop {
if let Some(item) = this.queue.pop_front() {
return this.deliver(item);
}
let (Some(bytes), Some(decoder)) = (this.bytes.as_mut(), this.decoder.as_mut()) else {
return Poll::Ready(None);
};
match Pin::new(bytes).poll_next(cx) {
Poll::Ready(Some(Ok(chunk))) => this.queue.extend(decoder.push(&chunk)),
Poll::Ready(Some(Err(error))) => this.queue.push_back(Err(error)),
Poll::Ready(None) => {
this.bytes = None;
if let Some(Err(error)) = this.decoder.take().map(Decoder::finish) {
this.queue.push_back(Err(error));
}
}
Poll::Pending => {
let late = if fired(&mut this.first_token, cx) {
"nothing of the answer arrived in time"
} else if fired(&mut this.complete, cx) {
"the answer was not complete in time"
} else {
return Poll::Pending;
};
return this.deliver(Err(timeout(late)));
}
}
}
}
}
impl<B: Backend> fmt::Debug for EventStream<B> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("EventStream").finish_non_exhaustive()
}
}
fn earlier(timer: &mut Option<Timer>, deadline: std::time::Instant) {
let deadline = Instant::from_std(deadline);
if timer.as_ref().is_some_and(|set| set.deadline() <= deadline) {
return;
}
*timer = Some(Box::pin(tokio::time::sleep_until(deadline)));
}
fn fired(timer: &mut Option<Timer>, cx: &mut Context<'_>) -> bool {
timer
.as_mut()
.is_some_and(|timer| timer.as_mut().poll(cx).is_ready())
}
fn timeout(detail: &'static str) -> Error {
Error::new(ErrorKind::Timeout).with_detail(detail)
}