mirror of
https://github.com/fabro-sh/fabro.git
synced 2026-09-06 08:18:58 +00:00
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:
parent
dba48f866e
commit
e4fc345912
3 changed files with 159 additions and 14 deletions
|
|
@ -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 { .. }))));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)))
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -78,6 +78,9 @@ pub enum EventData {
|
|||
Error {
|
||||
error: String,
|
||||
},
|
||||
TextDelta {
|
||||
delta: String,
|
||||
},
|
||||
ContextWarning {
|
||||
estimated_tokens: usize,
|
||||
context_window_size: usize,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue