Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 15 additions & 2 deletions crates/goose/src/agents/agent.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2463,8 +2463,21 @@ impl Agent {
// reasoning without hiding final-only non-streaming thoughts.
let mut surfaced_thinking_in_turn = false;

while let Some(next) = stream.next().await {
if is_token_cancelled(&cancel_token) || exit_chat {
loop {
let next = if let Some(cancel_token) = &cancel_token {
tokio::select! {
biased;
_ = cancel_token.cancelled() => break,
next = stream.next() => next,
}
} else {
stream.next().await
};
let Some(next) = next else {
break;
};

if exit_chat {
break;
}

Expand Down
278 changes: 277 additions & 1 deletion crates/goose/src/agents/reply_parts.rs
Original file line number Diff line number Diff line change
Expand Up @@ -361,6 +361,7 @@ pub(crate) async fn stream_response_from_provider(

// Clone owned data to move into the async stream
let system_prompt = system_prompt.to_owned();
let session_id = session_id.to_owned();
let tools = tools.to_owned();
let toolshim_tools = toolshim_tools.to_owned();
let provider = provider.clone();
Expand All @@ -372,7 +373,7 @@ pub(crate) async fn stream_response_from_provider(
let request_started = std::time::Instant::now();
debug!("WAITING_LLM_STREAM_START");
let stream_result = crate::session_context::with_session_id(
Some(session_id.to_string()),
Some(session_id.clone()),
provider.stream(
&model_config,
system_prompt.as_str(),
Expand All @@ -397,6 +398,66 @@ pub(crate) async fn stream_response_from_provider(
};

Ok(Box::pin(try_stream! {
if !provider.manages_own_context() {
let retry_config = provider.retry_config().transient_only();
let mut attempts = 0;

loop {
match stream.next().await {
None => break,
Some(Ok(item)) => {
stream = Box::pin(
futures::stream::once(std::future::ready(Ok(item))).chain(stream),
);
break;
}
Some(Err(error))
if goose_providers::retry::should_retry(&error, &retry_config)
&& attempts < retry_config.max_retries =>
Comment on lines +415 to +416

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Avoid retrying permanent first-frame 4xx errors

When an OpenAI-compatible gateway sends a deterministic client/validation failure as the first SSE frame, this retry check treats it as transient because the existing parser maps choice-less status/statusCode/code >= 400 and detail error frames to ProviderError::ServerError (crates/goose-provider-types/src/formats/openai.rs:1113-1137). That means a bad model/payload/auth-shaped in-stream 4xx now replays the same invalid request until max_retries and delays surfacing the real error; discriminate 4xx/validation frames before applying the transient retry policy.

Useful? React with 👍 / 👎.

{
attempts += 1;
let delay = match &error {
ProviderError::RateLimitExceeded {
retry_delay: Some(provider_delay),
..
} => *provider_delay,
_ => retry_config.delay_for_attempt(attempts),
};
warn!(
"Provider stream failed before its first item, retrying ({}/{}): {:?}",
attempts, retry_config.max_retries, error
);

let skip_backoff = std::env::var("GOOSE_PROVIDER_SKIP_BACKOFF")
.unwrap_or_default()
.parse::<bool>()
.unwrap_or(false);
if !skip_backoff {
tokio::time::sleep(delay).await;
Comment on lines +435 to +436

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Make the retry backoff cancellable

When the first stream item is a transient error, this sleep prevents the legacy agent loop from responding to cancellation until the entire backoff finishes, because crates/goose/src/agents/agent.rs:2466-2468 awaits stream.next() before checking its cancellation token. A provider-supplied retry_delay can make Stop appear hung for an arbitrarily long time, while the state-machine path remains cancellable through the tokio::select! in ops_llm.rs:513-518; make this retry wait cancellation-aware in both agent-loop paths.

AGENTS.md reference: AGENTS.md:L19-L23

Useful? React with 👍 / 👎.

}

stream = match crate::session_context::with_session_id(
Some(session_id.clone()),
provider.stream(
&model_config,
system_prompt.as_str(),
messages_for_provider.messages(),
&tools,
),
)
.await
{
Ok(stream) => stream,
Err(error) => {
Err(enhance_model_error(error, &provider, config.toolshim).await)?
}
};
}
Some(Err(error)) => Err(error)?,
}
}
}

if config.toolshim {
// Toolshim mode: accumulate the full response before processing
// so that tool-use markers spanning multiple chunks are detected
Expand Down Expand Up @@ -1789,4 +1850,219 @@ mod tests {
);
assert!(stats.elapsed_ms.expect("elapsed_ms must be filled") >= 100);
}

type TestStreamItem = Result<(Option<Message>, Option<ProviderUsage>), ProviderError>;

struct SequencedProvider {
responses: Mutex<std::collections::VecDeque<Result<Vec<TestStreamItem>, ProviderError>>>,
calls: Arc<std::sync::atomic::AtomicUsize>,
manages_context: bool,
}

impl SequencedProvider {
fn new(responses: Vec<Result<Vec<TestStreamItem>, ProviderError>>) -> Self {
Self {
responses: Mutex::new(responses.into()),
calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
manages_context: false,
}
}

fn managing_context(mut self) -> Self {
self.manages_context = true;
self
}
}

#[async_trait]
impl Provider for SequencedProvider {
fn get_name(&self) -> &str {
"sequenced"
}

fn retry_config(&self) -> goose_providers::retry::RetryConfig {
goose_providers::retry::RetryConfig::new(2, 0, 1.0, 0)
}

fn manages_own_context(&self) -> bool {
self.manages_context
}

async fn stream(
&self,
_model_config: &ModelConfig,
_system: &str,
_messages: &[Message],
_tools: &[Tool],
) -> Result<MessageStream, ProviderError> {
self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let response = self.responses.lock().unwrap().pop_front().unwrap();
response.map(|items| Box::pin(futures::stream::iter(items)) as MessageStream)
}
}

fn successful_item() -> TestStreamItem {
Ok((Some(Message::assistant().with_text("ok")), None))
}

fn transient_error() -> ProviderError {
ProviderError::NetworkError("stream closed".into())
}

async fn stream_for_test(provider: Arc<dyn Provider>) -> MessageStream {
stream_response_from_provider(
provider,
ModelConfig::new("test-model"),
"session",
"system",
&[Message::user().with_text("hi")],
&[],
&[],
)
.await
.unwrap()
}

#[tokio::test]
async fn first_item_transient_error_retries() {
let provider = Arc::new(SequencedProvider::new(vec![
Ok(vec![Err(transient_error())]),
Ok(vec![successful_item()]),
]));
let calls = provider.calls.clone();
let mut stream = stream_for_test(provider).await;

assert_eq!(
stream
.next()
.await
.unwrap()
.unwrap()
.0
.unwrap()
.as_concat_text(),
"ok"
);
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 2);
}

#[tokio::test]
async fn first_item_non_transient_error_does_not_retry() {
let provider = Arc::new(SequencedProvider::new(vec![Ok(vec![Err(
ProviderError::ContextLengthExceeded("too long".into()),
)])]));
let calls = provider.calls.clone();
let mut stream = stream_for_test(provider).await;

assert!(matches!(
stream.next().await.unwrap(),
Err(ProviderError::ContextLengthExceeded(_))
));
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
}

#[tokio::test]
async fn first_item_retry_exhaustion_returns_last_error() {
let provider = Arc::new(SequencedProvider::new(vec![
Ok(vec![Err(transient_error())]),
Ok(vec![Err(transient_error())]),
Ok(vec![Err(transient_error())]),
]));
let calls = provider.calls.clone();
let mut stream = stream_for_test(provider).await;

assert!(matches!(
stream.next().await.unwrap(),
Err(ProviderError::NetworkError(_))
));
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 3);
}

#[tokio::test]
async fn provider_managing_context_does_not_retry() {
let provider = Arc::new(
SequencedProvider::new(vec![Ok(vec![Err(transient_error())])]).managing_context(),
);
let calls = provider.calls.clone();
let mut stream = stream_for_test(provider).await;

assert!(matches!(
stream.next().await.unwrap(),
Err(ProviderError::NetworkError(_))
));
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
}

#[tokio::test]
async fn replacement_stream_creation_error_does_not_retry() {
let provider = Arc::new(SequencedProvider::new(vec![
Ok(vec![Err(transient_error())]),
Err(ProviderError::ServerError("unavailable".into())),
]));
let calls = provider.calls.clone();
let mut stream = stream_for_test(provider).await;

assert!(matches!(
stream.next().await.unwrap(),
Err(ProviderError::ServerError(_))
));
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 2);
}

#[tokio::test]
async fn error_after_first_item_does_not_retry() {
let provider = Arc::new(SequencedProvider::new(vec![Ok(vec![
successful_item(),
Err(transient_error()),
])]));
let calls = provider.calls.clone();
let mut stream = stream_for_test(provider).await;

assert!(stream.next().await.unwrap().is_ok());
assert!(matches!(
stream.next().await.unwrap(),
Err(ProviderError::NetworkError(_))
));
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
}

#[tokio::test]
async fn empty_stream_is_not_retried() {
let provider = Arc::new(SequencedProvider::new(vec![Ok(vec![])]));
let calls = provider.calls.clone();
let mut stream = stream_for_test(provider).await;

assert!(stream.next().await.is_none());
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
}

struct PendingProvider;

#[async_trait]
impl Provider for PendingProvider {
fn get_name(&self) -> &str {
"pending"
}

async fn stream(
&self,
_model_config: &ModelConfig,
_system: &str,
_messages: &[Message],
_tools: &[Tool],
) -> Result<MessageStream, ProviderError> {
Ok(Box::pin(futures::stream::pending()))
}
}

#[tokio::test]
async fn pending_first_item_does_not_block_stream_creation() {
let result = tokio::time::timeout(
Duration::from_secs(1),
stream_for_test(Arc::new(PendingProvider)),
)
.await;

assert!(result.is_ok());
}
}