mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 00:17:45 +08:00
feat(routing): make client disconnect behavior strategy-scoped
This commit is contained in:
@@ -62,6 +62,7 @@ flate2.workspace = true
|
||||
futures-util.workspace = true
|
||||
hmac.workspace = true
|
||||
http.workspace = true
|
||||
http-body = "1"
|
||||
http-body-util = "0.1"
|
||||
hyper = { version = "1", features = ["client", "server", "http1", "http2"] }
|
||||
hyper-util = { version = "0.1", features = ["client-legacy", "client-pool", "server-auto", "service", "tokio"] }
|
||||
|
||||
@@ -13607,10 +13607,12 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn execute_stream_from_frame_stream_cancels_upstream_when_client_drops_body() {
|
||||
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
|
||||
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||
let state = AppState::new()
|
||||
async fn execute_stream_from_frame_stream_honors_client_disconnect_policy() {
|
||||
for cancel_on_client_disconnect in [false, true] {
|
||||
let usage_repository = Arc::new(InMemoryUsageReadRepository::default());
|
||||
let request_candidate_repository =
|
||||
Arc::new(InMemoryRequestCandidateRepository::default());
|
||||
let state = AppState::new()
|
||||
.expect("app state should build")
|
||||
.with_data_state_for_tests(
|
||||
crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests(
|
||||
@@ -13622,39 +13624,39 @@ mod tests {
|
||||
enabled: true,
|
||||
..UsageRuntimeConfig::default()
|
||||
});
|
||||
let plan = ExecutionPlan {
|
||||
request_id: "req-client-drop-cancels-upstream".into(),
|
||||
candidate_id: Some("cand-client-drop-cancels-upstream".into()),
|
||||
provider_name: Some("openai".into()),
|
||||
provider_id: "prov-1".into(),
|
||||
endpoint_id: "ep-1".into(),
|
||||
key_id: "key-1".into(),
|
||||
method: "POST".into(),
|
||||
url: "https://example.com/v1/chat/completions".into(),
|
||||
headers: BTreeMap::from([
|
||||
("content-type".into(), "application/json".into()),
|
||||
("accept".into(), "text/event-stream".into()),
|
||||
]),
|
||||
content_type: Some("application/json".into()),
|
||||
content_encoding: None,
|
||||
body: RequestBody::from_json(json!({
|
||||
"model": "gpt-5.4",
|
||||
"messages": [],
|
||||
"stream": true
|
||||
})),
|
||||
stream: true,
|
||||
client_api_format: "openai:chat".into(),
|
||||
provider_api_format: "openai:chat".into(),
|
||||
model_name: Some("gpt-5.4".into()),
|
||||
proxy: None,
|
||||
transport_profile: None,
|
||||
timeouts: None,
|
||||
};
|
||||
let release_terminal = Arc::new(Notify::new());
|
||||
let terminal_frame_drained = Arc::new(Notify::new());
|
||||
let release_terminal_for_stream = Arc::clone(&release_terminal);
|
||||
let terminal_frame_drained_for_stream = Arc::clone(&terminal_frame_drained);
|
||||
let frame_stream = stream! {
|
||||
let plan = ExecutionPlan {
|
||||
request_id: "req-client-drop-cancels-upstream".into(),
|
||||
candidate_id: Some("cand-client-drop-cancels-upstream".into()),
|
||||
provider_name: Some("openai".into()),
|
||||
provider_id: "prov-1".into(),
|
||||
endpoint_id: "ep-1".into(),
|
||||
key_id: "key-1".into(),
|
||||
method: "POST".into(),
|
||||
url: "https://example.com/v1/chat/completions".into(),
|
||||
headers: BTreeMap::from([
|
||||
("content-type".into(), "application/json".into()),
|
||||
("accept".into(), "text/event-stream".into()),
|
||||
]),
|
||||
content_type: Some("application/json".into()),
|
||||
content_encoding: None,
|
||||
body: RequestBody::from_json(json!({
|
||||
"model": "gpt-5.4",
|
||||
"messages": [],
|
||||
"stream": true
|
||||
})),
|
||||
stream: true,
|
||||
client_api_format: "openai:chat".into(),
|
||||
provider_api_format: "openai:chat".into(),
|
||||
model_name: Some("gpt-5.4".into()),
|
||||
proxy: None,
|
||||
transport_profile: None,
|
||||
timeouts: None,
|
||||
};
|
||||
let release_terminal = Arc::new(Notify::new());
|
||||
let terminal_frame_drained = Arc::new(Notify::new());
|
||||
let release_terminal_for_stream = Arc::clone(&release_terminal);
|
||||
let terminal_frame_drained_for_stream = Arc::clone(&terminal_frame_drained);
|
||||
let frame_stream = stream! {
|
||||
yield Ok::<Bytes, std::io::Error>(Bytes::from_static(
|
||||
b"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n",
|
||||
));
|
||||
@@ -13669,119 +13671,157 @@ mod tests {
|
||||
}
|
||||
.boxed();
|
||||
|
||||
let response = execute_stream_from_frame_stream(
|
||||
&state,
|
||||
plan,
|
||||
"trace-client-drop-cancels-upstream",
|
||||
&test_decision(),
|
||||
"openai_chat_stream",
|
||||
None,
|
||||
Some(json!({
|
||||
"request_id": "req-client-drop-cancels-upstream",
|
||||
"candidate_id": "cand-client-drop-cancels-upstream",
|
||||
"candidate_index": 0,
|
||||
"retry_index": 0,
|
||||
"provider_api_format": "openai:chat",
|
||||
"client_api_format": "openai:chat"
|
||||
})),
|
||||
crate::clock::current_unix_ms(),
|
||||
Instant::now(),
|
||||
RequestStageTrace::from_env(),
|
||||
true,
|
||||
frame_stream,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.expect("execution should succeed")
|
||||
.expect("execution should return a client response");
|
||||
let response = crate::request_lifecycle::run_request(async move {
|
||||
crate::request_lifecycle::configure_client_disconnect(
|
||||
aether_routing_core::RoutingExecutionPolicy {
|
||||
cancel_on_client_disconnect,
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
execute_stream_from_frame_stream(
|
||||
&state,
|
||||
plan,
|
||||
"trace-client-drop-cancels-upstream",
|
||||
&test_decision(),
|
||||
"openai_chat_stream",
|
||||
None,
|
||||
Some(json!({
|
||||
"request_id": "req-client-drop-cancels-upstream",
|
||||
"candidate_id": "cand-client-drop-cancels-upstream",
|
||||
"candidate_index": 0,
|
||||
"retry_index": 0,
|
||||
"provider_api_format": "openai:chat",
|
||||
"client_api_format": "openai:chat"
|
||||
})),
|
||||
crate::clock::current_unix_ms(),
|
||||
Instant::now(),
|
||||
RequestStageTrace::from_env(),
|
||||
true,
|
||||
frame_stream,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
.map(|response| response.expect("execution should return a client response"))
|
||||
})
|
||||
.await
|
||||
.expect("execution should succeed");
|
||||
|
||||
let mut body_stream = response.into_body().into_data_stream();
|
||||
let first = tokio::time::timeout(Duration::from_secs(1), async {
|
||||
loop {
|
||||
let chunk = body_stream
|
||||
.next()
|
||||
.await
|
||||
.expect("body should yield first chunk")
|
||||
.expect("first chunk should be ok");
|
||||
if chunk.as_ref() != b": aether-keepalive\n\n" {
|
||||
break chunk;
|
||||
let mut body_stream = response.into_body().into_data_stream();
|
||||
let first = tokio::time::timeout(Duration::from_secs(1), async {
|
||||
loop {
|
||||
let chunk = body_stream
|
||||
.next()
|
||||
.await
|
||||
.expect("body should yield first chunk")
|
||||
.expect("first chunk should be ok");
|
||||
if chunk.as_ref() != b": aether-keepalive\n\n" {
|
||||
break chunk;
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("first business chunk should arrive");
|
||||
assert_eq!(
|
||||
})
|
||||
.await
|
||||
.expect("first business chunk should arrive");
|
||||
assert_eq!(
|
||||
first.as_ref(),
|
||||
b"data: {\"id\":\"first\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"hello\"}}]}\n\n"
|
||||
);
|
||||
tokio::time::sleep(Duration::from_millis(30)).await;
|
||||
drop(body_stream);
|
||||
let candidates = tokio::time::timeout(Duration::from_secs(1), async {
|
||||
loop {
|
||||
let candidates = request_candidate_repository
|
||||
.list_by_request_id("req-client-drop-cancels-upstream")
|
||||
.await
|
||||
.expect("request candidates should read");
|
||||
if candidates
|
||||
.first()
|
||||
.is_some_and(|candidate| candidate.status == RequestCandidateStatus::Cancelled)
|
||||
{
|
||||
break candidates;
|
||||
}
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
tokio::time::sleep(Duration::from_millis(30)).await;
|
||||
drop(body_stream);
|
||||
if !cancel_on_client_disconnect {
|
||||
release_terminal.notify_one();
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("candidate should be marked cancelled");
|
||||
assert_eq!(candidates[0].status_code, Some(499));
|
||||
assert_eq!(
|
||||
candidates[0].error_type.as_deref(),
|
||||
Some("downstream_disconnect")
|
||||
);
|
||||
|
||||
let stored_usage = tokio::time::timeout(Duration::from_secs(1), async {
|
||||
loop {
|
||||
let usage = usage_repository
|
||||
.find_by_request_id("req-client-drop-cancels-upstream")
|
||||
.await
|
||||
.expect("usage should read");
|
||||
if usage
|
||||
.as_ref()
|
||||
.is_some_and(|usage| usage.status == "cancelled")
|
||||
{
|
||||
break usage.expect("cancelled usage should exist");
|
||||
let expected_candidate_status = if cancel_on_client_disconnect {
|
||||
RequestCandidateStatus::Cancelled
|
||||
} else {
|
||||
RequestCandidateStatus::Success
|
||||
};
|
||||
let expected_usage_status = if cancel_on_client_disconnect {
|
||||
"cancelled"
|
||||
} else {
|
||||
"completed"
|
||||
};
|
||||
let candidates = tokio::time::timeout(Duration::from_secs(1), async {
|
||||
loop {
|
||||
let candidates = request_candidate_repository
|
||||
.list_by_request_id("req-client-drop-cancels-upstream")
|
||||
.await
|
||||
.expect("request candidates should read");
|
||||
if candidates
|
||||
.first()
|
||||
.is_some_and(|candidate| candidate.status == expected_candidate_status)
|
||||
{
|
||||
break candidates;
|
||||
}
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
}
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("usage should be marked cancelled");
|
||||
assert_eq!(stored_usage.billing_status, "void");
|
||||
assert_eq!(stored_usage.status_code, Some(499));
|
||||
assert_eq!(stored_usage.input_tokens, 0);
|
||||
assert_eq!(stored_usage.output_tokens, 0);
|
||||
assert_eq!(stored_usage.total_tokens, 0);
|
||||
let first_byte_time_ms = stored_usage
|
||||
.first_byte_time_ms
|
||||
.expect("cancelled stream should retain first byte time");
|
||||
let response_time_ms = stored_usage
|
||||
.response_time_ms
|
||||
.expect("cancelled stream should record terminal duration");
|
||||
assert!(
|
||||
response_time_ms > first_byte_time_ms,
|
||||
"terminal duration should include time after the first byte"
|
||||
);
|
||||
|
||||
release_terminal.notify_one();
|
||||
assert!(
|
||||
tokio::time::timeout(
|
||||
Duration::from_millis(100),
|
||||
terminal_frame_drained.notified()
|
||||
)
|
||||
})
|
||||
.await
|
||||
.is_err(),
|
||||
"upstream frame stream should stop when the client disconnects"
|
||||
);
|
||||
.expect("candidate should be marked cancelled");
|
||||
assert_eq!(
|
||||
candidates[0].status_code,
|
||||
Some(if cancel_on_client_disconnect {
|
||||
499
|
||||
} else {
|
||||
200
|
||||
})
|
||||
);
|
||||
assert_eq!(
|
||||
candidates[0].error_type.as_deref(),
|
||||
cancel_on_client_disconnect.then_some("downstream_disconnect")
|
||||
);
|
||||
|
||||
let stored_usage = tokio::time::timeout(Duration::from_secs(1), async {
|
||||
loop {
|
||||
let usage = usage_repository
|
||||
.find_by_request_id("req-client-drop-cancels-upstream")
|
||||
.await
|
||||
.expect("usage should read");
|
||||
if usage
|
||||
.as_ref()
|
||||
.is_some_and(|usage| usage.status == expected_usage_status)
|
||||
{
|
||||
break usage.expect("cancelled usage should exist");
|
||||
}
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("usage should be marked cancelled");
|
||||
if !cancel_on_client_disconnect {
|
||||
assert_ne!(stored_usage.billing_status, "void");
|
||||
assert_eq!(stored_usage.status_code, Some(200));
|
||||
assert_eq!(stored_usage.input_tokens, 7);
|
||||
assert_eq!(stored_usage.output_tokens, 11);
|
||||
assert_eq!(stored_usage.total_tokens, 18);
|
||||
continue;
|
||||
}
|
||||
assert_eq!(stored_usage.billing_status, "void");
|
||||
assert_eq!(stored_usage.status_code, Some(499));
|
||||
assert_eq!(stored_usage.input_tokens, 0);
|
||||
assert_eq!(stored_usage.output_tokens, 0);
|
||||
assert_eq!(stored_usage.total_tokens, 0);
|
||||
let first_byte_time_ms = stored_usage
|
||||
.first_byte_time_ms
|
||||
.expect("cancelled stream should retain first byte time");
|
||||
let response_time_ms = stored_usage
|
||||
.response_time_ms
|
||||
.expect("cancelled stream should record terminal duration");
|
||||
assert!(
|
||||
response_time_ms > first_byte_time_ms,
|
||||
"terminal duration should include time after the first byte"
|
||||
);
|
||||
|
||||
release_terminal.notify_one();
|
||||
assert!(
|
||||
tokio::time::timeout(
|
||||
Duration::from_millis(100),
|
||||
terminal_frame_drained.notified()
|
||||
)
|
||||
.await
|
||||
.is_err(),
|
||||
"upstream frame stream should stop when the client disconnects"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -822,15 +822,20 @@ where
|
||||
let started_at = Instant::now();
|
||||
let (tx, rx) = mpsc::channel::<Result<Bytes, IoError>>(1);
|
||||
let request_diagnostics = current_request_diagnostics();
|
||||
let cancel_on_disconnect = crate::request_lifecycle::cancel_on_client_disconnect();
|
||||
|
||||
tokio::spawn(async move {
|
||||
scope_request_diagnostics_with(request_diagnostics, async move {
|
||||
let bytes = standard_text_sync_heartbeat_final_bytes(
|
||||
let completion = standard_text_sync_heartbeat_final_bytes(
|
||||
client_api_format.as_str(),
|
||||
redaction_slot.as_ref(),
|
||||
execute(state, parts, trace_id, decision, plan_kind, started_at).await,
|
||||
)
|
||||
.await;
|
||||
tokio::select! {
|
||||
biased;
|
||||
_ = tx.closed(), if cancel_on_disconnect => return,
|
||||
result = execute(state, parts, trace_id, decision, plan_kind, started_at) => result,
|
||||
},
|
||||
);
|
||||
let bytes = completion.await;
|
||||
let _ = tx.send(Ok(Bytes::from(bytes))).await;
|
||||
})
|
||||
.await;
|
||||
@@ -1097,23 +1102,26 @@ fn build_openai_image_sync_heartbeat_shell_response(
|
||||
let started_at = Instant::now();
|
||||
let (tx, rx) = mpsc::channel::<Result<Bytes, IoError>>(1);
|
||||
let request_diagnostics = current_request_diagnostics();
|
||||
let cancel_on_disconnect = crate::request_lifecycle::cancel_on_client_disconnect();
|
||||
|
||||
tokio::spawn(async move {
|
||||
scope_request_diagnostics_with(request_diagnostics, async move {
|
||||
let bytes = openai_image_sync_heartbeat_final_bytes(
|
||||
execute_openai_image_sync_heartbeat_attempts(
|
||||
state,
|
||||
request_path,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
attempts,
|
||||
transfer_tracker,
|
||||
started_at,
|
||||
)
|
||||
.await,
|
||||
)
|
||||
.await;
|
||||
let execution = execute_openai_image_sync_heartbeat_attempts(
|
||||
state,
|
||||
request_path,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
attempts,
|
||||
transfer_tracker,
|
||||
started_at,
|
||||
);
|
||||
let outcome = tokio::select! {
|
||||
biased;
|
||||
_ = tx.closed(), if cancel_on_disconnect => return,
|
||||
result = execution => result,
|
||||
};
|
||||
let bytes = openai_image_sync_heartbeat_final_bytes(outcome).await;
|
||||
let _ = tx.send(Ok(Bytes::from(bytes))).await;
|
||||
})
|
||||
.await;
|
||||
@@ -2331,6 +2339,45 @@ mod tests {
|
||||
.expect("background completion should release admission");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn standard_text_sync_heartbeat_cancels_when_routing_policy_enables_it() {
|
||||
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
|
||||
let (mut release_tx, release_rx) = tokio::sync::oneshot::channel::<()>();
|
||||
let response = crate::request_lifecycle::run_request(async move {
|
||||
crate::request_lifecycle::configure_client_disconnect(
|
||||
aether_routing_core::RoutingExecutionPolicy {
|
||||
cancel_on_client_disconnect: true,
|
||||
..Default::default()
|
||||
},
|
||||
);
|
||||
let (parts, _) = http::Request::builder()
|
||||
.method("POST")
|
||||
.uri("/v1/responses")
|
||||
.body(())
|
||||
.unwrap()
|
||||
.into_parts();
|
||||
build_standard_text_sync_heartbeat_shell_response(
|
||||
AppState::new().unwrap(),
|
||||
parts,
|
||||
"trace-heartbeat-disconnect".to_string(),
|
||||
test_standard_text_heartbeat_decision(),
|
||||
TEST_STANDARD_TEXT_SYNC_PLAN_KIND.to_string(),
|
||||
move |_, _, _, _, _, _| async move {
|
||||
started_tx.send(()).unwrap();
|
||||
release_rx.await.unwrap();
|
||||
Ok(LocalExecutionRequestOutcome::NoPath)
|
||||
},
|
||||
)
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
started_rx.await.unwrap();
|
||||
drop(response);
|
||||
tokio::time::timeout(Duration::from_secs(1), release_tx.closed())
|
||||
.await
|
||||
.expect("heartbeat must drop upstream execution immediately");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn standard_text_sync_heartbeat_propagates_request_diagnostics_to_terminal_usage() {
|
||||
let (state, usage_repository) = heartbeat_usage_test_state(json!({
|
||||
|
||||
@@ -1035,7 +1035,7 @@ pub(crate) async fn proxy_request(
|
||||
ConnectInfo(remote_addr): ConnectInfo<std::net::SocketAddr>,
|
||||
request: Request,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
crate::request_diagnostics::scope_request_diagnostics(Box::pin(proxy_request_inner(
|
||||
crate::request_lifecycle::run_request(Box::pin(proxy_request_inner(
|
||||
state,
|
||||
remote_addr,
|
||||
request,
|
||||
|
||||
@@ -68,7 +68,12 @@ pub(super) async fn relay_bound_connection(
|
||||
state: &AppState,
|
||||
context: &WebSocketRequestContext,
|
||||
) {
|
||||
let mut client_connected = true;
|
||||
loop {
|
||||
if !client_connected && !bound.turn_state.response_in_flight() {
|
||||
close_bound_upstream(bound).await;
|
||||
break;
|
||||
}
|
||||
let active_turn_deadline = bound.turn_state.attempt().map(|turn| turn.deadline());
|
||||
tokio::select! {
|
||||
_ = wait_for_optional_deadline(active_turn_deadline.map(|deadline| deadline.deadline)) => {
|
||||
@@ -100,8 +105,12 @@ pub(super) async fn relay_bound_connection(
|
||||
).await;
|
||||
break;
|
||||
}
|
||||
client_message = client_socket.next() => {
|
||||
client_message = client_socket.next(), if client_connected => {
|
||||
let Some(client_message) = client_message else {
|
||||
if retain_disconnected_turn(bound) {
|
||||
client_connected = false;
|
||||
continue;
|
||||
}
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
@@ -111,6 +120,10 @@ pub(super) async fn relay_bound_connection(
|
||||
break;
|
||||
};
|
||||
let Ok(client_message) = client_message else {
|
||||
if retain_disconnected_turn(bound) {
|
||||
client_connected = false;
|
||||
continue;
|
||||
}
|
||||
warn!(
|
||||
event_name = "responses_websocket_client_receive_failed",
|
||||
log_type = "ops",
|
||||
@@ -127,6 +140,12 @@ pub(super) async fn relay_bound_connection(
|
||||
close_bound_upstream(bound).await;
|
||||
break;
|
||||
};
|
||||
if matches!(client_message, AxumWsMessage::Close(_))
|
||||
&& retain_disconnected_turn(bound)
|
||||
{
|
||||
client_connected = false;
|
||||
continue;
|
||||
}
|
||||
match Box::pin(forward_client_message(
|
||||
client_message,
|
||||
bound,
|
||||
@@ -559,6 +578,7 @@ pub(super) async fn relay_bound_connection(
|
||||
let mut relay_send_error = None;
|
||||
let mut relay_serialization_failed = false;
|
||||
match relay_directive {
|
||||
_ if !client_connected => {}
|
||||
Some(ResponsesWebSocketRelayDirective::ForwardOriginal) => {
|
||||
let client_frame = match parsed_upstream_frame.as_ref().map(|frame| {
|
||||
bound
|
||||
@@ -673,6 +693,10 @@ pub(super) async fn relay_bound_connection(
|
||||
break;
|
||||
}
|
||||
if let Some(error) = relay_send_error {
|
||||
if terminal_outcome.is_none() && retain_disconnected_turn(bound) {
|
||||
client_connected = false;
|
||||
continue;
|
||||
}
|
||||
warn!(
|
||||
event_name = "responses_websocket_client_send_failed",
|
||||
log_type = "ops",
|
||||
@@ -737,6 +761,20 @@ pub(super) async fn relay_bound_connection(
|
||||
}
|
||||
}
|
||||
|
||||
fn retain_disconnected_turn(bound: &mut BoundResponsesConnection) -> bool {
|
||||
if bound
|
||||
.turn_state
|
||||
.attempt()
|
||||
.is_none_or(|attempt| attempt.cancel_on_client_disconnect())
|
||||
{
|
||||
return false;
|
||||
}
|
||||
bound
|
||||
.turn_state
|
||||
.record_client_delivery_aborted(CLIENT_DELIVERY_FAILED_REASON);
|
||||
true
|
||||
}
|
||||
|
||||
struct PendingContinuationRegistration {
|
||||
user_id: String,
|
||||
api_key_id: String,
|
||||
|
||||
@@ -845,6 +845,13 @@ fn websocket_auth_rejection_error(rejection: GatewayLocalAuthRejection) -> Gatew
|
||||
}
|
||||
|
||||
impl ResponsesProviderAttempt {
|
||||
pub(super) fn cancel_on_client_disconnect(&self) -> bool {
|
||||
crate::orchestration::routing_execution_policy_from_report_context(
|
||||
self.lifecycle.report_context(),
|
||||
)
|
||||
.is_some_and(|policy| policy.cancel_on_client_disconnect)
|
||||
}
|
||||
|
||||
/// Releases all per-turn capacity before terminal persistence starts.
|
||||
/// Provider-pool runtime tokens normally use an awaited removal. The
|
||||
/// bounded wait prevents a broken runtime backend from stalling the relay;
|
||||
|
||||
@@ -71,6 +71,7 @@ mod rate_limit;
|
||||
mod request_candidate_queue;
|
||||
mod request_candidate_runtime;
|
||||
mod request_diagnostics;
|
||||
mod request_lifecycle;
|
||||
mod roles;
|
||||
mod router;
|
||||
mod routing;
|
||||
|
||||
@@ -0,0 +1,336 @@
|
||||
use std::future::Future;
|
||||
use std::pin::Pin;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::task::{Context, Poll};
|
||||
|
||||
use aether_routing_core::RoutingExecutionPolicy;
|
||||
use axum::body::{Body, Bytes, HttpBody};
|
||||
use http::Response;
|
||||
use http_body::{Frame, SizeHint};
|
||||
use http_body_util::BodyExt;
|
||||
|
||||
use crate::request_diagnostics::{scope_request_diagnostics_with, RequestDiagnostics};
|
||||
use crate::GatewayError;
|
||||
|
||||
tokio::task_local! {
|
||||
static CANCEL_ON_CLIENT_DISCONNECT: Arc<AtomicBool>;
|
||||
}
|
||||
|
||||
pub(crate) fn configure_client_disconnect(policy: RoutingExecutionPolicy) {
|
||||
let _ = CANCEL_ON_CLIENT_DISCONNECT.try_with(|cancel| {
|
||||
cancel.store(policy.cancel_on_client_disconnect, Ordering::Release);
|
||||
});
|
||||
}
|
||||
|
||||
pub(crate) fn cancel_on_client_disconnect() -> bool {
|
||||
CANCEL_ON_CLIENT_DISCONNECT
|
||||
.try_with(|cancel| cancel.load(Ordering::Acquire))
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
pub(crate) async fn run_request<F>(future: F) -> Result<Response<Body>, GatewayError>
|
||||
where
|
||||
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
|
||||
{
|
||||
let cancel = Arc::new(AtomicBool::new(true));
|
||||
let diagnostics = Arc::new(RequestDiagnostics::default());
|
||||
let cancel_for_response = Arc::clone(&cancel);
|
||||
let future = CANCEL_ON_CLIENT_DISCONNECT.scope(
|
||||
Arc::clone(&cancel),
|
||||
scope_request_diagnostics_with(Some(Arc::clone(&diagnostics)), async move {
|
||||
let response = future.await?;
|
||||
if cancel_for_response.load(Ordering::Acquire) {
|
||||
return Ok(response);
|
||||
}
|
||||
Ok(response.map(|body| {
|
||||
Body::new(CompleteOnDisconnectBody {
|
||||
body: Some(body),
|
||||
diagnostics,
|
||||
})
|
||||
}))
|
||||
}),
|
||||
);
|
||||
CompleteOnDisconnectRequest {
|
||||
future: Some(Box::pin(future)),
|
||||
cancel,
|
||||
}
|
||||
.await
|
||||
}
|
||||
|
||||
struct CompleteOnDisconnectRequest<F>
|
||||
where
|
||||
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
|
||||
{
|
||||
future: Option<Pin<Box<F>>>,
|
||||
cancel: Arc<AtomicBool>,
|
||||
}
|
||||
|
||||
impl<F> Future for CompleteOnDisconnectRequest<F>
|
||||
where
|
||||
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
|
||||
{
|
||||
type Output = Result<Response<Body>, GatewayError>;
|
||||
|
||||
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
|
||||
let result = self
|
||||
.future
|
||||
.as_mut()
|
||||
.expect("request future")
|
||||
.as_mut()
|
||||
.poll(context);
|
||||
if result.is_ready() {
|
||||
self.future.take();
|
||||
}
|
||||
result
|
||||
}
|
||||
}
|
||||
|
||||
impl<F> Drop for CompleteOnDisconnectRequest<F>
|
||||
where
|
||||
F: Future<Output = Result<Response<Body>, GatewayError>> + Send + 'static,
|
||||
{
|
||||
fn drop(&mut self) {
|
||||
if self.cancel.load(Ordering::Acquire) {
|
||||
return;
|
||||
}
|
||||
if let (Some(future), Ok(runtime)) =
|
||||
(self.future.take(), tokio::runtime::Handle::try_current())
|
||||
{
|
||||
runtime.spawn(async move {
|
||||
if let Ok(response) = future.await {
|
||||
drain_body(response.into_body()).await;
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct CompleteOnDisconnectBody {
|
||||
body: Option<Body>,
|
||||
diagnostics: Arc<RequestDiagnostics>,
|
||||
}
|
||||
|
||||
impl HttpBody for CompleteOnDisconnectBody {
|
||||
type Data = Bytes;
|
||||
type Error = axum::Error;
|
||||
|
||||
fn poll_frame(
|
||||
mut self: Pin<&mut Self>,
|
||||
context: &mut Context<'_>,
|
||||
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
|
||||
let Some(body) = self.body.as_mut() else {
|
||||
return Poll::Ready(None);
|
||||
};
|
||||
let result = Pin::new(body).poll_frame(context);
|
||||
if matches!(result, Poll::Ready(None | Some(Err(_)))) {
|
||||
self.body.take();
|
||||
}
|
||||
result
|
||||
}
|
||||
|
||||
fn is_end_stream(&self) -> bool {
|
||||
self.body.as_ref().is_none_or(HttpBody::is_end_stream)
|
||||
}
|
||||
|
||||
fn size_hint(&self) -> SizeHint {
|
||||
self.body
|
||||
.as_ref()
|
||||
.map(HttpBody::size_hint)
|
||||
.unwrap_or_else(|| SizeHint::with_exact(0))
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for CompleteOnDisconnectBody {
|
||||
fn drop(&mut self) {
|
||||
let Some(body) = self.body.take().filter(|body| !body.is_end_stream()) else {
|
||||
return;
|
||||
};
|
||||
if let Ok(runtime) = tokio::runtime::Handle::try_current() {
|
||||
runtime.spawn(scope_request_diagnostics_with(
|
||||
Some(Arc::clone(&self.diagnostics)),
|
||||
drain_body(body),
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn drain_body(mut body: Body) {
|
||||
while let Some(frame) = body.frame().await {
|
||||
if frame.is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::io;
|
||||
use std::time::Duration;
|
||||
|
||||
use futures_util::stream;
|
||||
use http::HeaderMap;
|
||||
use http_body_util::StreamBody;
|
||||
use tokio::sync::{mpsc, oneshot};
|
||||
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn disconnected_request_finishes_and_keeps_admission_and_diagnostics() {
|
||||
let gate = aether_runtime::ConcurrencyGate::new("disconnect_request", 1);
|
||||
let permit = gate.try_acquire().unwrap();
|
||||
let (started_tx, started_rx) = oneshot::channel();
|
||||
let (release_tx, release_rx) = oneshot::channel();
|
||||
let (finished_tx, finished_rx) = oneshot::channel();
|
||||
let request = tokio::spawn(run_request(async move {
|
||||
let _permit = permit;
|
||||
configure_client_disconnect(RoutingExecutionPolicy::default());
|
||||
started_tx.send(()).unwrap();
|
||||
release_rx.await.unwrap();
|
||||
assert!(crate::request_diagnostics::current_request_diagnostics().is_some());
|
||||
finished_tx.send(()).unwrap();
|
||||
Ok(Response::new(Body::empty()))
|
||||
}));
|
||||
started_rx.await.unwrap();
|
||||
request.abort();
|
||||
assert!(request.await.unwrap_err().is_cancelled());
|
||||
assert_eq!(gate.snapshot().in_flight, 1);
|
||||
release_tx.send(()).unwrap();
|
||||
tokio::time::timeout(Duration::from_secs(1), finished_rx)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(gate.snapshot().in_flight, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn enabled_cancellation_and_unresolved_requests_drop_immediately() {
|
||||
for resolve_policy in [false, true] {
|
||||
let (started_tx, started_rx) = oneshot::channel();
|
||||
let (release_tx, release_rx) = oneshot::channel::<()>();
|
||||
let request = tokio::spawn(run_request(async move {
|
||||
if resolve_policy {
|
||||
configure_client_disconnect(RoutingExecutionPolicy {
|
||||
cancel_on_client_disconnect: true,
|
||||
..Default::default()
|
||||
});
|
||||
}
|
||||
started_tx.send(()).unwrap();
|
||||
release_rx.await.unwrap();
|
||||
Ok(Response::new(Body::empty()))
|
||||
}));
|
||||
started_rx.await.unwrap();
|
||||
request.abort();
|
||||
assert!(request.await.unwrap_err().is_cancelled());
|
||||
assert!(release_tx.send(()).is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn disconnected_body_drains_without_buffering_and_holds_admission() {
|
||||
for consume_first_chunk in [false, true] {
|
||||
let gate = aether_runtime::ConcurrencyGate::new("disconnect_body", 1);
|
||||
let permit = gate.try_acquire().unwrap();
|
||||
let (sender, receiver) = mpsc::channel(1);
|
||||
let (finished_tx, finished_rx) = oneshot::channel();
|
||||
let response = run_request(async move {
|
||||
configure_client_disconnect(RoutingExecutionPolicy::default());
|
||||
let body = Body::from_stream(stream::unfold(
|
||||
(receiver, finished_tx, permit),
|
||||
|(mut receiver, finished_tx, permit)| async move {
|
||||
match receiver.recv().await {
|
||||
Some(bytes) => {
|
||||
Some((Ok::<_, io::Error>(bytes), (receiver, finished_tx, permit)))
|
||||
}
|
||||
None => {
|
||||
assert!(crate::request_diagnostics::current_request_diagnostics()
|
||||
.is_some());
|
||||
finished_tx.send(()).unwrap();
|
||||
None
|
||||
}
|
||||
}
|
||||
},
|
||||
));
|
||||
Ok(Response::new(body))
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let mut body = response.into_body();
|
||||
if consume_first_chunk {
|
||||
sender.send(Bytes::from_static(b"first")).await.unwrap();
|
||||
assert_eq!(
|
||||
body.frame().await.unwrap().unwrap().into_data().unwrap(),
|
||||
"first"
|
||||
);
|
||||
}
|
||||
drop(body);
|
||||
assert_eq!(gate.snapshot().in_flight, 1);
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
for _ in 0..100 {
|
||||
sender.send(Bytes::from_static(b"remaining")).await.unwrap();
|
||||
}
|
||||
drop(sender);
|
||||
finished_rx.await.unwrap();
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(gate.snapshot().in_flight, 0);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn enabled_cancellation_drops_stream_receiver() {
|
||||
let (sender, receiver) = mpsc::channel::<Result<Bytes, io::Error>>(1);
|
||||
let response = run_request(async move {
|
||||
configure_client_disconnect(RoutingExecutionPolicy {
|
||||
cancel_on_client_disconnect: true,
|
||||
..Default::default()
|
||||
});
|
||||
Ok(Response::new(Body::from_stream(stream::unfold(
|
||||
receiver,
|
||||
|mut receiver| async { receiver.recv().await.map(|item| (item, receiver)) },
|
||||
))))
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
drop(response);
|
||||
assert!(sender.is_closed());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn connected_response_preserves_headers_size_hint_and_trailers() {
|
||||
let response = run_request(async {
|
||||
configure_client_disconnect(RoutingExecutionPolicy::default());
|
||||
Ok(Response::builder()
|
||||
.status(201)
|
||||
.header("x-test", "unchanged")
|
||||
.body(Body::from("hello"))
|
||||
.unwrap())
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(response.status(), 201);
|
||||
assert_eq!(response.headers()["x-test"], "unchanged");
|
||||
assert_eq!(response.body().size_hint().exact(), Some(5));
|
||||
assert_eq!(
|
||||
response.into_body().collect().await.unwrap().to_bytes(),
|
||||
"hello"
|
||||
);
|
||||
|
||||
let mut trailers = HeaderMap::new();
|
||||
trailers.insert("x-finished", "yes".parse().unwrap());
|
||||
let response = run_request(async move {
|
||||
configure_client_disconnect(RoutingExecutionPolicy::default());
|
||||
let frames = stream::iter([
|
||||
Ok::<_, io::Error>(Frame::data(Bytes::from_static(b"hello"))),
|
||||
Ok(Frame::trailers(trailers)),
|
||||
]);
|
||||
Ok(Response::new(Body::new(StreamBody::new(frames))))
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let collected = response.into_body().collect().await.unwrap();
|
||||
assert_eq!(collected.trailers().unwrap()["x-finished"], "yes");
|
||||
assert_eq!(collected.to_bytes(), "hello");
|
||||
}
|
||||
}
|
||||
@@ -57,7 +57,7 @@ pub(crate) fn resolve_gateway_routing_policy(
|
||||
|
||||
let config = serde_json::from_value::<RoutingGroupConfig>(input.group_config_json.clone())
|
||||
.map_err(|_| invalid_routing_group_config())?;
|
||||
resolve_routing_policy(
|
||||
let policy = resolve_routing_policy(
|
||||
&config,
|
||||
RoutingPolicyInput {
|
||||
group_id: input.group_id,
|
||||
@@ -73,7 +73,9 @@ pub(crate) fn resolve_gateway_routing_policy(
|
||||
phase: input.phase,
|
||||
},
|
||||
)
|
||||
.map_err(routing_policy_error)
|
||||
.map_err(routing_policy_error)?;
|
||||
crate::request_lifecycle::configure_client_disconnect(policy.execution_policy);
|
||||
Ok(policy)
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_gateway_static_default_routing_policy(
|
||||
@@ -82,6 +84,7 @@ pub(crate) fn resolve_gateway_static_default_routing_policy(
|
||||
let Some(default_policy) = static_default_policy_fields(input.group_config_json)? else {
|
||||
return Ok(None);
|
||||
};
|
||||
crate::request_lifecycle::configure_client_disconnect(default_policy.execution_policy);
|
||||
|
||||
Ok(Some(ResolvedRoutingPolicy {
|
||||
group_id: input.group_id.map(str::to_string),
|
||||
@@ -163,6 +166,10 @@ fn static_default_policy_fields(
|
||||
default_policy.get("cyber_continue_failover"),
|
||||
"cyber_continue_failover",
|
||||
)?,
|
||||
cancel_on_client_disconnect: routing_bool_field(
|
||||
default_policy.get("cancel_on_client_disconnect"),
|
||||
"cancel_on_client_disconnect",
|
||||
)?,
|
||||
};
|
||||
|
||||
Ok(Some(RoutingDefaultPolicy {
|
||||
@@ -238,7 +245,8 @@ mod tests {
|
||||
"default_policy": {
|
||||
"priority_mode": "global_key",
|
||||
"scheduling_mode": "load_balance",
|
||||
"keep_priority_on_conversion": true
|
||||
"keep_priority_on_conversion": true,
|
||||
"cancel_on_client_disconnect": true
|
||||
},
|
||||
"allowed_models": ["legacy-model"],
|
||||
"model_policies": [],
|
||||
|
||||
@@ -55,6 +55,24 @@ fn hash_api_key(value: &str) -> String {
|
||||
format!("{:x}", hasher.finalize())
|
||||
}
|
||||
|
||||
async fn build_cancelling_gateway(state: crate::AppState) -> Router {
|
||||
state
|
||||
.data
|
||||
.update_routing_group(
|
||||
"system-default",
|
||||
aether_data_contracts::repository::routing_profiles::UpdateRoutingGroupRecord {
|
||||
config_json: Some(json!({"default_policy": {"cancel_on_client_disconnect": true}})),
|
||||
version: Some(2),
|
||||
updated_at: 2,
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("routing policy should update")
|
||||
.expect("default strategy should exist");
|
||||
build_router_with_state(state)
|
||||
}
|
||||
|
||||
fn sample_local_openai_auth_snapshot(api_key_id: &str, user_id: &str) -> StoredAuthApiKeySnapshot {
|
||||
StoredAuthApiKeySnapshot::new(
|
||||
user_id.to_string(),
|
||||
@@ -384,7 +402,7 @@ async fn gateway_stops_execution_runtime_stream_when_client_disconnects_impl() {
|
||||
vec![sample_local_openai_key()],
|
||||
));
|
||||
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||
let gateway = build_router_with_state(
|
||||
let gateway = build_cancelling_gateway(
|
||||
build_state_with_execution_runtime_override(execution_runtime_url)
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
|
||||
@@ -395,7 +413,7 @@ async fn gateway_stops_execution_runtime_stream_when_client_disconnects_impl() {
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
),
|
||||
),
|
||||
);
|
||||
).await;
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
@@ -467,7 +485,7 @@ async fn gateway_settles_stream_attempt_when_client_disconnects_before_first_byt
|
||||
vec![sample_local_openai_key()],
|
||||
));
|
||||
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||
let gateway = build_router_with_state(
|
||||
let gateway = build_cancelling_gateway(
|
||||
build_state_with_execution_runtime_override(execution_runtime_url)
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
|
||||
@@ -478,7 +496,7 @@ async fn gateway_settles_stream_attempt_when_client_disconnects_before_first_byt
|
||||
DEVELOPMENT_ENCRYPTION_KEY,
|
||||
),
|
||||
),
|
||||
);
|
||||
).await;
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let request = reqwest::Client::new()
|
||||
|
||||
Reference in New Issue
Block a user