From 6c60df214e32c727fefeba6f498f991c75e466da Mon Sep 17 00:00:00 2001 From: wsp1911 Date: Wed, 24 Jun 2026 19:05:07 +0800 Subject: [PATCH] fix(ai-adapters): cancel orphaned stream handlers on turn cancel - manage provider stream handler lifecycle in execute_sse_request - cancel background SSE handlers when the response stream is dropped - stop provider adapters from spawning detached stream tasks - add coverage for managed stream cancellation on drop --- src/crates/adapters/ai-adapters/Cargo.toml | 2 +- .../adapters/ai-adapters/src/client/sse.rs | 90 +++++++++++++++++-- .../src/providers/anthropic/request.rs | 4 +- .../src/providers/gemini/code_assist.rs | 4 +- .../src/providers/gemini/request.rs | 4 +- .../ai-adapters/src/providers/openai/chat.rs | 4 +- .../src/providers/openai/codex_chatgpt.rs | 4 +- .../src/providers/openai/responses.rs | 4 +- 8 files changed, 97 insertions(+), 19 deletions(-) diff --git a/src/crates/adapters/ai-adapters/Cargo.toml b/src/crates/adapters/ai-adapters/Cargo.toml index 9e4c717dbf..69f6ce66f1 100644 --- a/src/crates/adapters/ai-adapters/Cargo.toml +++ b/src/crates/adapters/ai-adapters/Cargo.toml @@ -23,9 +23,9 @@ serde = { workspace = true } serde_json = { workspace = true } tokio = { workspace = true } tokio-stream = { workspace = true } +tokio-util = { workspace = true } urlencoding = { workspace = true } [dev-dependencies] axum = { workspace = true } bitfun-events = { path = "../../contracts/events" } -tokio-util = { workspace = true } diff --git a/src/crates/adapters/ai-adapters/src/client/sse.rs b/src/crates/adapters/ai-adapters/src/client/sse.rs index 11a7f05ac3..d5e7cb3eef 100644 --- a/src/crates/adapters/ai-adapters/src/client/sse.rs +++ b/src/crates/adapters/ai-adapters/src/client/sse.rs @@ -4,13 +4,20 @@ use crate::stream::UnifiedResponse; use crate::trace::{ModelExchangeRequestAttempt, ModelExchangeTraceConfig}; use anyhow::{anyhow, Result}; use chrono::{DateTime, Utc}; +use futures::Stream; use log::{debug, error, warn}; use reqwest::{ header::{HeaderMap, RETRY_AFTER}, StatusCode, }; +use std::future::Future; +use std::pin::Pin; +use std::task::{Context, Poll}; use std::time::Duration; use tokio::sync::mpsc; +use tokio::task::JoinHandle; +use tokio_stream::wrappers::UnboundedReceiverStream; +use tokio_util::sync::CancellationToken; const BASE_RETRY_DELAY_MS: u64 = 500; /// Maximum delay applied to a `Retry-After` header value. @@ -107,7 +114,42 @@ fn retry_delay_ms(attempt: usize, headers: &HeaderMap) -> u64 { retry_after_delay_ms(headers).unwrap_or_else(|| exponential_retry_delay_ms(attempt)) } -pub(crate) async fn execute_sse_request( +struct ManagedResponseStream { + inner: UnboundedReceiverStream>, + handler_cancel: CancellationToken, + handler_task: Option>, +} + +impl ManagedResponseStream { + fn new( + rx: mpsc::UnboundedReceiver>, + handler_cancel: CancellationToken, + handler_task: JoinHandle<()>, + ) -> Self { + Self { + inner: UnboundedReceiverStream::new(rx), + handler_cancel, + handler_task: Some(handler_task), + } + } +} + +impl Stream for ManagedResponseStream { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.inner).poll_next(cx) + } +} + +impl Drop for ManagedResponseStream { + fn drop(&mut self) { + self.handler_cancel.cancel(); + let _ = self.handler_task.take(); + } +} + +pub(crate) async fn execute_sse_request( label: &str, url: &str, request_body: &serde_json::Value, @@ -115,16 +157,17 @@ pub(crate) async fn execute_sse_request( ttft_timeout: Option, trace: Option, build_request: BuildRequest, - spawn_handler: SpawnHandler, + build_handler: BuildHandler, ) -> Result where BuildRequest: Fn() -> reqwest::RequestBuilder, - SpawnHandler: Fn( + BuildHandler: Fn( reqwest::Response, mpsc::UnboundedSender>, Option>, Option, - ), + ) -> HandlerFuture, + HandlerFuture: Future + Send + 'static, { let mut last_error = None; for attempt in 0..max_tries { @@ -288,10 +331,18 @@ where let (tx, rx) = mpsc::unbounded_channel(); let (tx_raw, rx_raw) = mpsc::unbounded_channel(); let remaining_ttft_timeout = remaining_ttft_timeout(request_start_time, ttft_timeout); - spawn_handler(response, tx, Some(tx_raw), remaining_ttft_timeout); + let handler_cancel = CancellationToken::new(); + let handler_cancel_for_task = handler_cancel.clone(); + let handler_future = build_handler(response, tx, Some(tx_raw), remaining_ttft_timeout); + let handler_task = tokio::spawn(async move { + tokio::select! { + _ = handler_cancel_for_task.cancelled() => {} + _ = handler_future => {} + } + }); return Ok(StreamResponse { - stream: Box::pin(tokio_stream::wrappers::UnboundedReceiverStream::new(rx)), + stream: Box::pin(ManagedResponseStream::new(rx, handler_cancel, handler_task)), raw_sse_rx: Some(rx_raw), trace_handle, }); @@ -311,6 +362,10 @@ where mod tests { use super::*; use reqwest::header::HeaderValue; + use std::sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }; #[test] fn format_ttft_timeout_error_includes_timeout_seconds() { @@ -333,6 +388,29 @@ mod tests { assert!(remaining > Duration::from_secs(2)); } + #[tokio::test] + async fn managed_response_stream_drop_cancels_handler_task() { + let (_tx, rx) = mpsc::unbounded_channel(); + let handler_cancel = CancellationToken::new(); + let handler_cancel_for_task = handler_cancel.clone(); + let observed_cancel = Arc::new(AtomicBool::new(false)); + let observed_cancel_for_task = Arc::clone(&observed_cancel); + let handler_task = tokio::spawn(async move { + tokio::select! { + _ = handler_cancel_for_task.cancelled() => { + observed_cancel_for_task.store(true, Ordering::SeqCst); + } + _ = tokio::time::sleep(Duration::from_secs(60)) => {} + } + }); + + let stream = ManagedResponseStream::new(rx, handler_cancel, handler_task); + drop(stream); + + tokio::time::sleep(Duration::from_millis(20)).await; + assert!(observed_cancel.load(Ordering::SeqCst)); + } + #[test] fn retryable_http_statuses_include_rate_limit_and_server_errors() { assert!(is_retryable_http_status(StatusCode::TOO_MANY_REQUESTS)); diff --git a/src/crates/adapters/ai-adapters/src/providers/anthropic/request.rs b/src/crates/adapters/ai-adapters/src/providers/anthropic/request.rs index 888a14f740..bdbbd509c9 100644 --- a/src/crates/adapters/ai-adapters/src/providers/anthropic/request.rs +++ b/src/crates/adapters/ai-adapters/src/providers/anthropic/request.rs @@ -351,14 +351,14 @@ pub(crate) async fn send_stream( trace, || apply_headers(client, client.client.post(&url), &url), move |response, tx, tx_raw, remaining_ttft_timeout| { - tokio::spawn(handle_anthropic_stream( + handle_anthropic_stream( response, tx, tx_raw, inline_think_in_text, remaining_ttft_timeout, idle_timeout, - )); + ) }, ) .await diff --git a/src/crates/adapters/ai-adapters/src/providers/gemini/code_assist.rs b/src/crates/adapters/ai-adapters/src/providers/gemini/code_assist.rs index 9cad79192f..e3ab4a6781 100644 --- a/src/crates/adapters/ai-adapters/src/providers/gemini/code_assist.rs +++ b/src/crates/adapters/ai-adapters/src/providers/gemini/code_assist.rs @@ -183,13 +183,13 @@ pub(crate) async fn send_stream( trace, || apply_headers(client, client.client.post(&url)), move |response, tx, tx_raw, remaining_ttft_timeout| { - tokio::spawn(handle_gemini_stream( + handle_gemini_stream( response, tx, tx_raw, remaining_ttft_timeout, idle_timeout, - )); + ) }, ) .await diff --git a/src/crates/adapters/ai-adapters/src/providers/gemini/request.rs b/src/crates/adapters/ai-adapters/src/providers/gemini/request.rs index 8df79808cf..77f2006061 100644 --- a/src/crates/adapters/ai-adapters/src/providers/gemini/request.rs +++ b/src/crates/adapters/ai-adapters/src/providers/gemini/request.rs @@ -347,13 +347,13 @@ pub(crate) async fn send_stream( trace, || apply_headers(client, client.client.post(&url)), move |response, tx, tx_raw, remaining_ttft_timeout| { - tokio::spawn(handle_gemini_stream( + handle_gemini_stream( response, tx, tx_raw, remaining_ttft_timeout, idle_timeout, - )); + ) }, ) .await diff --git a/src/crates/adapters/ai-adapters/src/providers/openai/chat.rs b/src/crates/adapters/ai-adapters/src/providers/openai/chat.rs index be8a08c0ba..5128956364 100644 --- a/src/crates/adapters/ai-adapters/src/providers/openai/chat.rs +++ b/src/crates/adapters/ai-adapters/src/providers/openai/chat.rs @@ -101,14 +101,14 @@ pub(crate) async fn send_stream( trace, || common::apply_headers(client, client.client.post(&url)), move |response, tx, tx_raw, remaining_ttft_timeout| { - tokio::spawn(handle_openai_stream( + handle_openai_stream( response, tx, tx_raw, inline_think_in_text, remaining_ttft_timeout, idle_timeout, - )); + ) }, ) .await diff --git a/src/crates/adapters/ai-adapters/src/providers/openai/codex_chatgpt.rs b/src/crates/adapters/ai-adapters/src/providers/openai/codex_chatgpt.rs index 509e68c6c6..eda37530f7 100644 --- a/src/crates/adapters/ai-adapters/src/providers/openai/codex_chatgpt.rs +++ b/src/crates/adapters/ai-adapters/src/providers/openai/codex_chatgpt.rs @@ -171,13 +171,13 @@ pub(crate) async fn send_stream( trace, || common::apply_headers(client, client.client.post(&url)), move |response, tx, tx_raw, remaining_ttft_timeout| { - tokio::spawn(handle_responses_stream( + handle_responses_stream( response, tx, tx_raw, remaining_ttft_timeout, idle_timeout, - )); + ) }, ) .await diff --git a/src/crates/adapters/ai-adapters/src/providers/openai/responses.rs b/src/crates/adapters/ai-adapters/src/providers/openai/responses.rs index b5fc6527c5..044d59aa39 100644 --- a/src/crates/adapters/ai-adapters/src/providers/openai/responses.rs +++ b/src/crates/adapters/ai-adapters/src/providers/openai/responses.rs @@ -135,13 +135,13 @@ pub(crate) async fn send_stream( trace, || common::apply_headers(client, client.client.post(&url)), move |response, tx, tx_raw, remaining_ttft_timeout| { - tokio::spawn(handle_responses_stream( + handle_responses_stream( response, tx, tx_raw, remaining_ttft_timeout, idle_timeout, - )); + ) }, ) .await