From e4fc345912650aeaf3fa4f100c12834cfb85a466 Mon Sep 17 00:00:00 2001 From: Bryan Helmkamp Date: Mon, 23 Feb 2026 13:10:54 -0500 Subject: [PATCH] Switch agent session from complete() to stream() to avoid request timeouts Long Opus responses hit the 120s request timeout with complete(). With stream(), the request timeout only covers initial connection + first chunk, then the 30s stream_read timeout guards against stalls between chunks. Also emits AssistantTextDelta events during streaming. Co-Authored-By: Claude Opus 4.6 --- crates/agent/src/session.rs | 90 ++++++++++++++++++++++++++++++-- crates/agent/src/test_support.rs | 80 ++++++++++++++++++++++++---- crates/agent/src/types.rs | 3 ++ 3 files changed, 159 insertions(+), 14 deletions(-) diff --git a/crates/agent/src/session.rs b/crates/agent/src/session.rs index dc7794f4f..eb87c09f0 100644 --- a/crates/agent/src/session.rs +++ b/crates/agent/src/session.rs @@ -14,9 +14,11 @@ use std::collections::VecDeque; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{Arc, Mutex}; use std::time::SystemTime; +use futures::StreamExt; use llm::client::Client; use llm::error::{ProviderErrorKind, SdkError}; -use llm::types::{Message, Request, ToolChoice, ToolResult}; +use llm::generate::StreamAccumulator; +use llm::types::{Message, Request, StreamEvent, ToolChoice, ToolResult}; pub struct Session { id: String, @@ -267,9 +269,9 @@ impl Session { EventData::Empty, ); - // Call LLM - let response = match self.llm_client.complete(&request).await { - Ok(resp) => resp, + // Call LLM (streaming) + let mut event_stream = match self.llm_client.stream(&request).await { + Ok(stream) => stream, Err(err) => { self.event_emitter.emit( EventKind::Error, @@ -285,6 +287,49 @@ impl Session { } }; + let mut accumulator = StreamAccumulator::new(); + + while let Some(event_result) = event_stream.next().await { + match event_result { + Ok(event) => { + if let StreamEvent::TextDelta { ref delta, .. } = event { + self.event_emitter.emit( + EventKind::AssistantTextDelta, + self.id.clone(), + EventData::TextDelta { + delta: delta.clone(), + }, + ); + } + accumulator.process(&event); + } + Err(err) => { + self.event_emitter.emit( + EventKind::Error, + self.id.clone(), + EventData::Error { + error: err.to_string(), + }, + ); + return Err(AgentError::Llm(err)); + } + } + + // Check abort flag between chunks + if self.abort_flag.load(Ordering::SeqCst) { + self.state = SessionState::Closed; + self.event_emitter + .emit(EventKind::SessionEnd, self.id.clone(), EventData::Empty); + return Err(AgentError::Aborted); + } + } + + let response = accumulator.response().cloned().ok_or_else(|| { + AgentError::Llm(SdkError::Stream { + message: "Stream ended without a Finish event".into(), + }) + })?; + // Record assistant turn let text = response.text(); let tool_calls = response.tool_calls(); @@ -1582,4 +1627,41 @@ mod tests { _ => panic!("Expected ToolCallEnd event data"), } } + + #[tokio::test] + async fn stream_emits_text_delta_events() { + let mut session = make_session(vec![text_response("Hello there!")]).await; + let mut rx = session.subscribe(); + + session.process_input("Hi").await.unwrap(); + + let mut deltas = Vec::new(); + while let Ok(event) = rx.try_recv() { + if event.kind == EventKind::AssistantTextDelta { + if let EventData::TextDelta { delta } = &event.data { + deltas.push(delta.clone()); + } + } + } + + assert_eq!(deltas.len(), 1); + assert_eq!(deltas[0], "Hello there!"); + } + + #[tokio::test] + async fn stream_mid_stream_error() { + let provider = Arc::new(MockMidStreamErrorProvider { + partial_text: "partial".into(), + error: SdkError::Stream { + message: "connection reset".into(), + }, + }); + let client = make_client(provider as Arc).await; + let profile = Arc::new(TestProfile::new()); + let env = Arc::new(MockExecutionEnvironment::default()); + let mut session = Session::new(client, profile, env, SessionConfig::default()); + + let result = session.process_input("Hello").await; + assert!(matches!(result, Err(AgentError::Llm(SdkError::Stream { .. })))); + } } diff --git a/crates/agent/src/test_support.rs b/crates/agent/src/test_support.rs index 8d7d58d98..c9a403e22 100644 --- a/crates/agent/src/test_support.rs +++ b/crates/agent/src/test_support.rs @@ -11,7 +11,7 @@ use std::sync::{Arc, Mutex}; use llm::client::Client; use llm::error::SdkError; use llm::provider::{ProviderAdapter, StreamEventStream}; -use llm::types::{FinishReason, Message, Request, Response, Usage}; +use llm::types::{ContentPart, FinishReason, Message, Request, Response, StreamEvent, Usage}; // --- MockExecutionEnvironment --- @@ -394,12 +394,45 @@ impl ProviderAdapter for MockLlmProvider { } async fn stream(&self, _request: &Request) -> Result { - Err(SdkError::Configuration { - message: "streaming not supported in mock".into(), - }) + let idx = self.call_index.fetch_add(1, Ordering::SeqCst); + let response = if idx < self.responses.len() { + self.responses[idx].clone() + } else { + self.responses[self.responses.len() - 1].clone() + }; + Ok(response_to_stream(response)) } } +/// Convert a canned `Response` into a `StreamEventStream` for mock streaming. +fn response_to_stream(response: Response) -> StreamEventStream { + let mut events: Vec> = Vec::new(); + + // Emit text deltas for text content + let text = response.text(); + if !text.is_empty() { + events.push(Ok(StreamEvent::text_delta(text, None))); + } + + // Emit tool call events + for part in &response.message.content { + if let ContentPart::ToolCall(tc) = part { + events.push(Ok(StreamEvent::ToolCallEnd { + tool_call: tc.clone(), + })); + } + } + + // Emit finish + events.push(Ok(StreamEvent::finish( + response.finish_reason.clone(), + response.usage.clone(), + response, + ))); + + Box::pin(futures::stream::iter(events)) +} + // --- Helper functions --- pub(crate) fn text_response(text: &str) -> Response { @@ -552,9 +585,7 @@ impl ProviderAdapter for MockErrorProvider { } async fn stream(&self, _request: &Request) -> Result { - Err(SdkError::Configuration { - message: "streaming not supported in mock".into(), - }) + Err(self.error.clone()) } } @@ -587,10 +618,39 @@ impl ProviderAdapter for CapturingLlmProvider { Ok(text_response("captured")) } + async fn stream(&self, request: &Request) -> Result { + *self + .captured_request + .lock() + .expect("captured_request lock poisoned") = Some(request.clone()); + Ok(response_to_stream(text_response("captured"))) + } +} + +// --- MockMidStreamErrorProvider --- + +/// A mock provider that yields some text deltas then an error mid-stream. +pub(crate) struct MockMidStreamErrorProvider { + pub partial_text: String, + pub error: SdkError, +} + +#[async_trait] +impl ProviderAdapter for MockMidStreamErrorProvider { + fn name(&self) -> &str { + "mock" + } + + async fn complete(&self, _request: &Request) -> Result { + Err(self.error.clone()) + } + async fn stream(&self, _request: &Request) -> Result { - Err(SdkError::Configuration { - message: "streaming not supported in mock".into(), - }) + let events: Vec> = vec![ + Ok(StreamEvent::text_delta(self.partial_text.clone(), None)), + Err(self.error.clone()), + ]; + Ok(Box::pin(futures::stream::iter(events))) } } diff --git a/crates/agent/src/types.rs b/crates/agent/src/types.rs index ab2a4fb0d..0966eb2fc 100644 --- a/crates/agent/src/types.rs +++ b/crates/agent/src/types.rs @@ -78,6 +78,9 @@ pub enum EventData { Error { error: String, }, + TextDelta { + delta: String, + }, ContextWarning { estimated_tokens: usize, context_window_size: usize,