use crate::chat_client::{
client::{ChatClient, TokenUsage},
error::Error,
openai_api::{chat_completions::StreamingChunk, message::Content},
};
use eventsource_stream::{Event, EventStreamError};
use futures::{
ready,
stream::{FusedStream, Stream, StreamExt},
task::Poll,
};
use std::pin::Pin;
pub enum Delta {
Reasoning(String),
Content(String),
Usage(TokenUsage),
}
#[derive(Debug)]
enum State {
WaitingForData,
ReceivingReasoning,
ReceivingContent { accumulated_response: String },
WaitingForDone,
WaitingForEndOfStream,
Terminated,
}
impl State {
fn finalize(&mut self, new_state: Self) -> Option<String> {
let old_state = std::mem::replace(self, new_state);
match old_state {
Self::ReceivingContent {
accumulated_response,
} => (!accumulated_response.is_empty()).then_some(accumulated_response),
_ => None,
}
}
}
pub struct CompletionStream<'a, S> {
client: &'a mut ChatClient,
stream: S,
state: State,
request: Content,
}
impl<'a, S> CompletionStream<'a, S> {
pub(crate) fn new(client: &'a mut ChatClient, stream: S, request: Content) -> Self {
Self {
client,
stream,
state: State::WaitingForData,
request,
}
}
}
impl<'a, S> Stream for CompletionStream<'a, S>
where
S: Stream<Item = Result<Event, EventStreamError<reqwest::Error>>> + Unpin,
{
type Item = Result<Delta, Error>;
fn poll_next(
self: Pin<&mut Self>,
cx: &mut futures::task::Context,
) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
if matches!(this.state, State::Terminated) {
return Poll::Ready(None);
}
loop {
let event = match ready!(this.stream.poll_next_unpin(cx)) {
Some(Ok(event)) => {
if event.data == "[DONE]" {
if let Some(response) = this.state.finalize(State::WaitingForEndOfStream) {
this.client.extend_context(
this.request.clone(),
Content::Text(response),
&Default::default(),
);
}
continue;
}
event
}
Some(Err(e)) => {
if let Some(response) = this.state.finalize(State::Terminated) {
this.client.extend_context(
this.request.clone(),
Content::Text(response),
&Default::default(),
);
}
return Poll::Ready(Some(Err(Error::from(e))));
}
None => {
if let Some(response) = this.state.finalize(State::Terminated) {
this.client.extend_context(
this.request.clone(),
Content::Text(response),
&Default::default(),
);
}
return Poll::Ready(None);
}
};
let delta = match parse_stream_chunk(&event.data) {
Ok(Some(delta)) => delta,
Ok(None) => continue,
Err(e) => {
if let Some(response) = this.state.finalize(State::Terminated) {
this.client.extend_context(
this.request.clone(),
Content::Text(response),
&Default::default(),
);
}
return Poll::Ready(Some(Err(e)));
}
};
match this.state {
State::WaitingForData | State::ReceivingReasoning => match delta {
Delta::Reasoning(_) => {
this.state = State::ReceivingReasoning;
}
Delta::Content(ref content) => {
this.state = State::ReceivingContent {
accumulated_response: content.clone(),
};
}
Delta::Usage(_) => {
this.state = State::WaitingForDone;
}
},
State::ReceivingContent {
ref mut accumulated_response,
} => match delta {
Delta::Reasoning(_) => {
if let Some(response) = this.state.finalize(State::Terminated) {
this.client.extend_context(
this.request.clone(),
Content::Text(response),
&Default::default(),
);
}
return Poll::Ready(Some(Err(Error::UnexpectedStreamEvent(
"reasoning after content",
))));
}
Delta::Content(ref content) => {
accumulated_response.push_str(content);
}
Delta::Usage(ref usage) => {
if let Some(response) = this.state.finalize(State::WaitingForDone) {
this.client.extend_context(
this.request.clone(),
Content::Text(response),
usage,
);
}
}
},
State::WaitingForDone => {
this.state = State::Terminated;
match delta {
Delta::Reasoning(_) => {
return Poll::Ready(Some(Err(Error::UnexpectedStreamEvent(
"reasoning after usage",
))))
}
Delta::Content(_) => {
return Poll::Ready(Some(Err(Error::UnexpectedStreamEvent(
"content after usage",
))))
}
Delta::Usage(_) => {
return Poll::Ready(Some(Err(Error::UnexpectedStreamEvent(
"duplicate usage",
))))
}
}
}
State::WaitingForEndOfStream => {
this.state = State::Terminated;
return Poll::Ready(None);
}
State::Terminated => unreachable!("terminated state is handled by early return"),
}
return Poll::Ready(Some(Ok(delta)));
}
}
}
fn parse_stream_chunk(event: &str) -> Result<Option<Delta>, Error> {
let mut chunk: StreamingChunk = serde_json::from_str(event)?;
let choice = match chunk.choices.pop() {
Some(choice) => choice,
None => {
if let Some(usage) = chunk.usage {
return Ok(Some(Delta::Usage(usage.into())));
} else {
return Err(Error::NoChoices);
}
}
};
if let Some(reasoning) = choice.delta.reasoning {
Ok(Some(Delta::Reasoning(reasoning)))
} else if let Some(content) = choice.delta.content {
if content.is_empty() {
if let Some(usage) = chunk.usage {
return Ok(Some(Delta::Usage(usage.into())));
}
}
Ok(Some(Delta::Content(content)))
} else if let Some(refusal) = choice.delta.refusal {
Err(Error::Refusal(refusal))
} else if choice.finish_reason.is_some() {
Ok(None)
} else {
Err(Error::NoContent)
}
}
impl<'a, S> FusedStream for CompletionStream<'a, S>
where
S: Stream<Item = Result<Event, EventStreamError<reqwest::Error>>> + Unpin,
{
fn is_terminated(&self) -> bool {
matches!(self.state, State::Terminated)
}
}