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 <noreply@anthropic.com>
This commit is contained in:
Bryan Helmkamp 2026-02-23 13:10:54 -05:00
parent dba48f866e
commit e4fc345912
3 changed files with 159 additions and 14 deletions

View file

@ -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<dyn ProviderAdapter>).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 { .. }))));
}
}

View file

@ -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<StreamEventStream, SdkError> {
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<Result<StreamEvent, SdkError>> = 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<StreamEventStream, SdkError> {
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<StreamEventStream, SdkError> {
*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<Response, SdkError> {
Err(self.error.clone())
}
async fn stream(&self, _request: &Request) -> Result<StreamEventStream, SdkError> {
Err(SdkError::Configuration {
message: "streaming not supported in mock".into(),
})
let events: Vec<Result<StreamEvent, SdkError>> = vec![
Ok(StreamEvent::text_delta(self.partial_text.clone(), None)),
Err(self.error.clone()),
];
Ok(Box::pin(futures::stream::iter(events)))
}
}

View file

@ -78,6 +78,9 @@ pub enum EventData {
Error {
error: String,
},
TextDelta {
delta: String,
},
ContextWarning {
estimated_tokens: usize,
context_window_size: usize,