Everything is streaming (#7247)

Co-authored-by: Douwe Osinga <douwe@squareup.com>
Co-authored-by: Claude Sonnet 4.5 <noreply@anthropic.com>
This commit is contained in:
Douwe Osinga
2026-02-17 19:53:38 +01:00
committed by GitHub
parent 65d3a1550f
commit 9a91fccfb5
42 changed files with 801 additions and 1718 deletions
+67 -51
View File
@@ -5,7 +5,8 @@ use std::sync::Arc;
use tokio::sync::Mutex;
use super::base::{
LeadWorkerProviderTrait, Provider, ProviderDef, ProviderMetadata, ProviderUsage,
collect_stream, stream_from_single_message, LeadWorkerProviderTrait, MessageStream, Provider,
ProviderDef, ProviderMetadata, ProviderUsage,
};
use super::errors::ProviderError;
use crate::conversation::message::{Message, MessageContent};
@@ -356,14 +357,14 @@ impl Provider for LeadWorkerProvider {
self.lead_provider.get_model_config()
}
async fn complete_with_model(
async fn stream(
&self,
session_id: Option<&str>,
_model_config: &ModelConfig,
session_id: &str,
system: &str,
messages: &[Message],
tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> {
) -> Result<MessageStream, ProviderError> {
// Get the active provider
let provider = self.get_active_provider().await;
@@ -410,9 +411,13 @@ impl Provider for LeadWorkerProvider {
// Make the completion request
let model_config = provider.get_model_config();
let result = provider
.complete_with_model(session_id, &model_config, system, messages, tools)
let stream_result = provider
.stream(&model_config, session_id, system, messages, tools)
.await;
let result = match stream_result {
Ok(stream) => collect_stream(stream).await,
Err(e) => Err(e),
};
// For technical failures, try with default model (lead provider) instead
let final_result = match &result {
@@ -421,10 +426,14 @@ impl Provider for LeadWorkerProvider {
// Try with lead provider as the default/fallback for technical failures
let model_config = self.lead_provider.get_model_config();
let default_result = self
let default_stream_result = self
.lead_provider
.complete_with_model(session_id, &model_config, system, messages, tools)
.stream(&model_config, session_id, system, messages, tools)
.await;
let default_result = match default_stream_result {
Ok(stream) => collect_stream(stream).await,
Err(e) => Err(e),
};
match &default_result {
Ok(_) => {
@@ -445,7 +454,10 @@ impl Provider for LeadWorkerProvider {
// Handle the result and update tracking (only for successful completions)
self.handle_completion_result(&final_result).await;
final_result
match final_result {
Ok((message, usage)) => Ok(stream_from_single_message(message, usage)),
Err(e) => Err(e),
}
}
async fn fetch_supported_models(&self) -> Result<Vec<String>, ProviderError> {
@@ -514,28 +526,27 @@ mod tests {
self.model_config.clone()
}
async fn complete_with_model(
async fn stream(
&self,
_session_id: Option<&str>,
_model_config: &ModelConfig,
_session_id: &str,
_system: &str,
_messages: &[Message],
_tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> {
Ok((
Message::new(
Role::Assistant,
Utc::now().timestamp(),
vec![MessageContent::Text(
RawTextContent {
text: format!("Response from {}", self.name),
meta: None,
}
.no_annotation(),
)],
),
ProviderUsage::new(self.name.clone(), Usage::default()),
))
) -> Result<MessageStream, ProviderError> {
let message = Message::new(
Role::Assistant,
Utc::now().timestamp(),
vec![MessageContent::Text(
RawTextContent {
text: format!("Response from {}", self.name),
meta: None,
}
.no_annotation(),
)],
);
let usage = ProviderUsage::new(self.name.clone(), Usage::default());
Ok(stream_from_single_message(message, usage))
}
}
@@ -552,11 +563,12 @@ mod tests {
});
let provider = LeadWorkerProvider::new(lead_provider, worker_provider, Some(3));
let model_config = provider.get_model_config();
// First three turns should use lead provider
for i in 0..3 {
let (_message, usage) = provider
.complete("test-session-id", "system", &[], &[])
.complete(&model_config, "test-session-id", "system", &[], &[])
.await
.unwrap();
assert_eq!(usage.model, "lead");
@@ -567,7 +579,7 @@ mod tests {
// Subsequent turns should use worker provider
for i in 3..6 {
let (_message, usage) = provider
.complete("test-session-id", "system", &[], &[])
.complete(&model_config, "test-session-id", "system", &[], &[])
.await
.unwrap();
assert_eq!(usage.model, "worker");
@@ -582,7 +594,7 @@ mod tests {
assert!(!provider.is_in_fallback_mode().await);
let (_message, usage) = provider
.complete("test-session-id", "system", &[], &[])
.complete(&model_config, "test-session-id", "system", &[], &[])
.await
.unwrap();
assert_eq!(usage.model, "lead");
@@ -603,11 +615,12 @@ mod tests {
});
let provider = LeadWorkerProvider::new(lead_provider, worker_provider, Some(2));
let model_config = provider.get_model_config();
// First two turns use lead (should succeed)
for _i in 0..2 {
let result = provider
.complete("test-session-id", "system", &[], &[])
.complete(&model_config, "test-session-id", "system", &[], &[])
.await;
assert!(result.is_ok());
assert_eq!(result.unwrap().1.model, "lead");
@@ -615,8 +628,9 @@ mod tests {
}
// Next turn uses worker (will fail, but should retry with lead and succeed)
let model_config = provider.get_model_config();
let result = provider
.complete("test-session-id", "system", &[], &[])
.complete(&model_config, "test-session-id", "system", &[], &[])
.await;
assert!(result.is_ok()); // Should succeed because lead provider is used as fallback
assert_eq!(result.unwrap().1.model, "lead"); // Should be lead provider
@@ -624,8 +638,9 @@ mod tests {
assert!(!provider.is_in_fallback_mode().await); // Not in fallback mode
// Another turn - should still try worker first, then retry with lead
let model_config = provider.get_model_config();
let result = provider
.complete("test-session-id", "system", &[], &[])
.complete(&model_config, "test-session-id", "system", &[], &[])
.await;
assert!(result.is_ok()); // Should succeed because lead provider is used as fallback
assert_eq!(result.unwrap().1.model, "lead"); // Should be lead provider
@@ -663,16 +678,18 @@ mod tests {
}
// Should use lead provider in fallback mode
let model_config = provider.get_model_config();
let result = provider
.complete("test-session-id", "system", &[], &[])
.complete(&model_config, "test-session-id", "system", &[], &[])
.await;
assert!(result.is_ok());
assert_eq!(result.unwrap().1.model, "lead");
assert!(provider.is_in_fallback_mode().await);
// One more fallback turn
let model_config = provider.get_model_config();
let result = provider
.complete("test-session-id", "system", &[], &[])
.complete(&model_config, "test-session-id", "system", &[], &[])
.await;
assert!(result.is_ok());
assert_eq!(result.unwrap().1.model, "lead");
@@ -696,33 +713,32 @@ mod tests {
self.model_config.clone()
}
async fn complete_with_model(
async fn stream(
&self,
_session_id: Option<&str>,
_model_config: &ModelConfig,
_session_id: &str,
_system: &str,
_messages: &[Message],
_tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> {
) -> Result<MessageStream, ProviderError> {
if self.should_fail {
Err(ProviderError::ExecutionError(
"Simulated failure".to_string(),
))
} else {
Ok((
Message::new(
Role::Assistant,
Utc::now().timestamp(),
vec![MessageContent::Text(
RawTextContent {
text: format!("Response from {}", self.name),
meta: None,
}
.no_annotation(),
)],
),
ProviderUsage::new(self.name.clone(), Usage::default()),
))
let message = Message::new(
Role::Assistant,
Utc::now().timestamp(),
vec![MessageContent::Text(
RawTextContent {
text: format!("Response from {}", self.name),
meta: None,
}
.no_annotation(),
)],
);
let usage = ProviderUsage::new(self.name.clone(), Usage::default());
Ok(stream_from_single_message(message, usage))
}
}
}