use super::{ApiClient, BareLoop, LoopError, Message, Run};
use crate::capabilities::StreamCapable;
use crate::observer::{TextDeltaContext, ThinkingDeltaContext};
use crate::stream::handler::{HandlerEvent, StreamHandlerError};
use crate::stream::{StreamAccumulator, StreamEvent, StreamStopReason, Usage};
use futures::StreamExt;
impl<C: ApiClient> BareLoop<C> {
pub(super) async fn stream_turn(
&self,
contributor_messages: Vec<Message>,
) -> Result<(Message, Option<Usage>, StreamStopReason), LoopError> {
let handler = self.managers.stream_handler();
let mut messages = contributor_messages;
messages.extend(self.machine.full_history());
let request = crate::api::StreamRequest::new(messages)
.with_system(self.session.config.system_prompt.clone())
.with_tools(self.build_tool_schemas());
let mut stream = handler.stream_turn(
&*self.client,
&request,
self.request_options.clone(),
&self.cancelled,
);
let mut accumulator = StreamAccumulator::new();
let mut stop_reason = StreamStopReason::EndTurn;
while let Some(result) = stream.next().await {
match result.map_err(Self::map_handler_error)? {
HandlerEvent::Stream(ev) => {
self.dispatch_stream_event(&ev, &mut accumulator, &mut stop_reason)?;
}
HandlerEvent::AttemptReset => {
accumulator = StreamAccumulator::new();
stop_reason = StreamStopReason::EndTurn;
}
HandlerEvent::Fallback {
message,
stop_reason: fallback_stop_reason,
usage: fallback_usage,
} => {
return Ok((message, fallback_usage, fallback_stop_reason));
}
}
}
let usage = accumulator.usage().copied();
Ok((accumulator.build(), usage, stop_reason))
}
fn dispatch_stream_event(
&self,
event: &StreamEvent,
accumulator: &mut StreamAccumulator,
stop_reason: &mut StreamStopReason,
) -> Result<(), LoopError> {
if let StreamEvent::IndexedDelta(d) = event
&& let crate::stream::DeltaPart::Text { text } = &d.delta
{
if let Some(streamer) = &self.text_streamer {
streamer(text.as_str());
}
self.managers.observers().on_text_delta(&TextDeltaContext {
turn: self.current_run().map_or(0, Run::turn_count),
delta: text.clone(),
});
}
if let StreamEvent::IndexedDelta(d) = event
&& let crate::stream::DeltaPart::Thinking { text } = &d.delta
{
self.managers
.observers()
.on_thinking_delta(&ThinkingDeltaContext {
turn: self.current_run().map_or(0, Run::turn_count),
delta: text.clone(),
});
}
if let StreamEvent::MessageDelta(d) = event
&& let Some(reason_str) = &d.delta.stop_reason
{
*stop_reason = StreamStopReason::from_api_str(reason_str).unwrap_or(*stop_reason);
}
accumulator
.process(event)
.map_err(|e| LoopError::Api(format!("stream accumulation error: {e}")))
}
fn map_handler_error(error: StreamHandlerError) -> LoopError {
match error {
StreamHandlerError::Cancelled => LoopError::Cancelled,
StreamHandlerError::InitFailed(outcome) => {
LoopError::Api(format!("stream init failed: {outcome}"))
}
StreamHandlerError::StreamFailed(outcome) => {
LoopError::Api(format!("stream failed: {outcome}"))
}
StreamHandlerError::FallbackFailed {
stream_outcome,
fallback_error,
} => LoopError::Api(format!(
"stream ({stream_outcome}) and fallback failed: {fallback_error}"
)),
StreamHandlerError::RateLimitEscalation {
attempts,
retry_after,
} => LoopError::RateLimitEscalation {
attempts,
retry_after,
},
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::api::error::ApiError;
struct StubClient;
impl ApiClient for StubClient {
fn model(&self) -> String {
"stub".to_string()
}
fn stream_messages(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<dyn futures::Stream<Item = Result<StreamEvent, ApiError>> + Send + 'static>,
> {
Box::pin(futures::stream::empty())
}
fn create_message(
&self,
_request: &crate::api::StreamRequest,
) -> std::pin::Pin<
Box<
dyn std::future::Future<Output = Result<crate::api::NonStreamingResponse, ApiError>>
+ Send
+ '_,
>,
> {
Box::pin(async {
Ok(crate::api::NonStreamingResponse {
message: crate::message::Message::assistant(""),
stop_reason: crate::stream::StreamStopReason::EndTurn,
usage: Some(crate::stream::Usage::default()),
})
})
}
}
#[test]
fn map_handler_error_escalation() {
let mapped =
BareLoop::<StubClient>::map_handler_error(StreamHandlerError::RateLimitEscalation {
attempts: 3,
retry_after: Some(std::time::Duration::from_secs(12)),
});
match mapped {
LoopError::RateLimitEscalation {
attempts,
retry_after,
} => {
assert_eq!(attempts, 3);
assert_eq!(retry_after, Some(std::time::Duration::from_secs(12)));
}
other => panic!("expected RateLimitEscalation, got {other:?}"),
}
}
}