mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-11 11:49:50 +08:00
feat(routing): make client disconnect behavior strategy-scoped
This commit is contained in:
@@ -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;
|
||||
|
||||
Reference in New Issue
Block a user