Files
Aether/apps/aether-gateway/src/handlers/proxy/websocket/responses/connection.rs
T
AAEE86 247e7105a2 fix(ws): bill a provider-reached terminal even when client delivery fails
评审第 5 条后半:provider 终态已经到达、只是 gateway 写客户端 socket 失败时,
relay loop 用 client_disconnected() 覆盖了结算信号,于是一条供应商已经完成推理
并消耗了 token 的响应被记成 void billing、candidate 记 Cancelled、不投射供应商
效果、也不提交 execution report。上游成本凭空消失。

结算表只改一行:作废账单的条件从
    provider.cancelled_by_provider() || delivery.is_aborted()
收紧为
    provider.cancelled_by_provider() || (delivery.is_aborted() && !provider.is_terminal())

于是 Terminal{cancelled=false} + delivery Aborted 与 delivery Complete 落在同一侧:
Billed、candidate Success 或 Failed、投射供应商效果、提交 execution report。
状态码随之变成纯 provider 事实(不再把 200 改写成 499);作废分支的 provider
状态码本身就是 499,取值不变。

依据:供应商已经完成推理并消耗 token,客户端还能用 previous_response_id 续取
这条响应。供应商没给出终态时(客户端先走了)仍然作废,这一侧未改。

配套改动:
- connection.rs 写客户端失败处改为 record_client_delivery_aborted(reason) +
  settle_signal_for_client_delivery_failure(terminal_outcome):provider 终态已到达
  就用那条终态作结算信号,不再无条件覆盖。投递失败原因也不再谎称
  「客户端在终态前断开」。
- 投递结果记在 attempt 上而非 logical turn 上:结算按 attempt 进行,且配额透明
  重试时各 attempt 的投递结果彼此独立。
- report_context 新增 websocket_client_delivery="aborted" 与
  websocket_client_delivery_reason,只增字段不改既有字段,便于事后区分
  「客户端拿到了」和「客户端没拿到但已计费」。
- candidate error_type 新增 client_delivery_failed(原先这个场景写的是
  websocket_cancelled)。它排在供应商侧分类之前:这条记录之所以特别正是因为
  内容没送到客户端,供应商侧判定仍由 candidate_status 与 error_message 保留。
- finish_summary 改用作废判定而非「投递失败」判定:provider 终态已到达时摘要
  必须保留真实的 finish_reason 与 usage,否则计费记录会被写坏。

e2e 期望值变化:client_disconnect_mid_turn_still_settles_the_usage_row 改名为
client_disconnect_before_any_provider_output_settles_a_void_row,并补上
「不计费 + status=cancelled + status_code=499」的断言。原用例的 mock 行为是
StallAfterCreated(只发 response.created 就静默),provider 从未给出终态,所以
它走的是未改动的作废一侧;原来的文档注释说「must still be billed」与实际语义
不符,一并纠正。真正被修正的那一行无法在 e2e 里确定性触发——它取决于 relay
loop 的 select! 先观察到上游终态帧还是先观察到已关闭的客户端 socket,是构造性
竞态——因此由 relay 级单测确定性覆盖,e2e 里以注释指向这两个单测。

新增 7 个测试:结算表修正行(并与「投递成功」逐字段对照,只有 candidate 错误
分类不同)、无终态时仍作废、供应商声明取消即使送达也不计费、结算信号选择、
已记录的投递失败不被结算信号覆盖、relay 级「终态到达 + 客户端已关闭 ⇒ Billed /
Success / ProviderSuccess / 已提交 report 且 usage 完整保留」及其镜像、
report_context 只增不改。
2026-08-17 14:53:12 +08:00

555 lines
25 KiB
Rust

//! Connection-level Responses WebSocket FSM.
use std::time::Duration;
use axum::extract::ws::WebSocket;
use futures_util::{SinkExt, StreamExt};
use serde_json::Value;
use tokio::time::sleep;
use wreq::ws::message::Message as WreqWsMessage;
use super::client::{adapter_drain_ready, forward_client_message, RelayDisposition};
use super::frame::ParsedResponsesWebSocketFrame;
use super::lifecycle::{
await_pending_adapter_observation, finalize_active_turn, queue_turn_finalization,
ActiveProviderAttempt,
};
use super::quota::{
active_continuation_can_retry_from_full_input, detach_exhausted_upstream,
is_usage_limit_error_event, mark_active_response_retry_unsafe,
observe_active_response_rebind_safety, retry_active_turn_after_quota_exhaustion,
send_previous_response_not_found, should_request_full_continuation_retry,
};
use super::relay_policy::{
classify_quota_relay, fatal_relay_policy, FatalRelaySignal, QuotaRelayAction, QuotaRelayFacts,
};
use super::settlement::settle_signal_for_client_delivery_failure;
use super::state::BoundResponsesConnection;
use super::turn::{
ResponsesProviderAttempt, ResponsesWebSocketTurnObservation, ResponsesWebSocketTurnOutcome,
};
use super::upstream::{close_bound_upstream, receive_optional_upstream};
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
use crate::handlers::proxy::websocket::session::{
wait_for_optional_deadline, CLOSE_INTERNAL_ERROR, CLOSE_TRY_AGAIN,
RESPONSES_WEBSOCKET_SESSION_LIMITS, WEBSOCKET_LOG_TRANSPORT,
};
use crate::handlers::proxy::websocket::transport::{
close_client_socket, send_client_message, send_gateway_error_with_status,
send_responses_websocket_error, upstream_message_to_client,
};
use crate::AppState;
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
/// 写客户端 socket 失败时记录的投递失败原因。刻意不说「客户端在终态前断开」:
/// 供应商的终态可能已经到达,只是最后一跳没送出去。
const CLIENT_DELIVERY_FAILED_REASON: &str =
"gateway could not relay the provider event to the client";
macro_rules! debug {
($($arg:tt)*) => {
tracing::debug!(target: LOG_TARGET, $($arg)*)
};
}
macro_rules! warn {
($($arg:tt)*) => {
tracing::warn!(target: LOG_TARGET, $($arg)*)
};
}
pub(super) async fn relay_bound_connection(
client_socket: &mut WebSocket,
bound: &mut BoundResponsesConnection,
state: &AppState,
context: &WebSocketRequestContext,
connection_permit: Option<aether_runtime::AdmissionPermit>,
) {
let connection_deadline = sleep(RESPONSES_WEBSOCKET_SESSION_LIMITS.max_connection_duration);
tokio::pin!(connection_deadline);
loop {
let active_turn_deadline = bound.turn_state.attempt().map(|turn| turn.deadline());
tokio::select! {
_ = &mut connection_deadline => {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::connection_limit_reached(),
).await;
send_gateway_error_with_status(
client_socket,
503,
"websocket_connection_limit_reached",
"WebSocket connection duration limit reached; reconnect to continue",
).await;
close_bound_upstream(bound).await;
close_client_socket(client_socket, CLOSE_TRY_AGAIN, "connection_limit_reached").await;
break;
}
_ = wait_for_optional_deadline(active_turn_deadline.map(|deadline| deadline.deadline)) => {
let Some(turn_deadline) = active_turn_deadline else {
continue;
};
warn!(
event_name = "responses_websocket_turn_timeout",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
timeout_phase = ?turn_deadline.phase,
timeout_ms = turn_deadline.timeout.as_millis() as u64,
"Responses WebSocket response did not reach its configured deadline"
);
finalize_active_turn(bound, state, turn_deadline.phase.outcome()).await;
send_gateway_error_with_status(
client_socket,
504,
turn_deadline.phase.error_code(),
turn_deadline.phase.client_message(),
).await;
close_bound_upstream(bound).await;
close_client_socket(
client_socket,
CLOSE_TRY_AGAIN,
turn_deadline.phase.error_code(),
).await;
break;
}
_ = wait_for_connection_permit_loss(connection_permit.as_ref()) => {
let policy = fatal_relay_policy(FatalRelaySignal::ConnectionAdmissionLost);
warn!(
event_name = "responses_websocket_connection_admission_lost",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"gateway closed Responses WebSocket after its connection admission became unhealthy"
);
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::connection_admission_lost(),
).await;
close_bound_upstream(bound).await;
send_gateway_error_with_status(
client_socket,
policy.status_code,
policy.error_code,
policy.client_message,
).await;
close_client_socket(client_socket, policy.close_code, policy.close_reason).await;
break;
}
client_message = client_socket.next() => {
let Some(client_message) = client_message else {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::client_disconnected(),
).await;
close_bound_upstream(bound).await;
break;
};
let Ok(client_message) = client_message else {
warn!(
event_name = "responses_websocket_client_receive_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"client WebSocket receive failed"
);
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::client_disconnected(),
).await;
close_bound_upstream(bound).await;
break;
};
match forward_client_message(client_message, bound, client_socket, state, context).await {
RelayDisposition::Continue => {}
RelayDisposition::Close => {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::client_disconnected(),
).await;
break;
}
RelayDisposition::UpstreamError(code) => {
warn!(
event_name = "responses_websocket_upstream_send_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = code,
"Upstream WebSocket send failed"
);
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::upstream_send_failed(),
).await;
send_gateway_error_with_status(
client_socket,
502,
code,
"Gateway could not forward the WebSocket event upstream",
).await;
close_bound_upstream(bound).await;
close_client_socket(client_socket, CLOSE_INTERNAL_ERROR, code).await;
break;
}
}
}
upstream_message = receive_optional_upstream(&mut bound.upstream) => {
let Some(upstream_message) = upstream_message else {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::upstream_closed(),
).await;
bound.upstream = None;
close_client_socket(client_socket, 1000, "upstream_closed").await;
break;
};
let Ok(upstream_message) = upstream_message else {
warn!(
event_name = "responses_websocket_upstream_receive_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
"Upstream WebSocket receive failed"
);
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::upstream_receive_failed(),
).await;
send_gateway_error_with_status(
client_socket,
502,
"responses_websocket_receive_failed",
"Provider connection closed unexpectedly",
).await;
bound.upstream = None;
close_client_socket(client_socket, CLOSE_INTERNAL_ERROR, "upstream_receive_failed").await;
break;
};
let parsed_upstream_frame = match &upstream_message {
WreqWsMessage::Text(text) => {
ParsedResponsesWebSocketFrame::parse(text.as_str()).ok()
}
_ => None,
};
let parsed_upstream_event = parsed_upstream_frame
.as_ref()
.map(ParsedResponsesWebSocketFrame::event);
if let WreqWsMessage::Text(text) = &upstream_message {
debug!(
event_name = "responses_websocket_upstream_event",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
event_type = %parsed_upstream_frame
.as_ref()
.map(ParsedResponsesWebSocketFrame::event_type_for_log)
.unwrap_or_else(|| "invalid_json".to_string()),
frame_bytes = text.len(),
chunked = parsed_upstream_frame
.as_ref()
.is_some_and(ParsedResponsesWebSocketFrame::is_chunked),
active_turn = bound.turn_state.response_in_flight(),
"gateway received Responses WebSocket event"
);
}
if matches!(&upstream_message, WreqWsMessage::Binary(_)) {
mark_active_response_retry_unsafe(bound, "upstream_binary_frame");
} else if matches!(&upstream_message, WreqWsMessage::Text(_))
&& parsed_upstream_event.is_none()
{
mark_active_response_retry_unsafe(bound, "invalid_upstream_event");
}
if let Some(event) = parsed_upstream_event {
observe_active_response_rebind_safety(bound, event);
if bound.pending_adapter_drain.is_none()
&& bound.adapter.observes_upstream_events()
{
let adapter = bound.adapter;
if let Some(observation) = adapter.observe_upstream_event(event) {
let directive = observation.drain;
await_pending_adapter_observation(bound).await;
let state_for_observation = state.clone();
let trace_id = context.trace_id.clone();
let report_context = bound.decision_template.report_context.clone();
bound.pending_adapter_observation = Some(tokio::spawn(async move {
adapter
.persist_upstream_observation(
&state_for_observation,
&trace_id,
report_context.as_ref(),
observation,
)
.await;
}));
if let Some(directive) = directive {
bound.pending_adapter_drain = Some(directive);
// A definitive quota signal must be visible to
// the next planner before a transparent retry.
await_pending_adapter_observation(bound).await;
}
}
}
}
let observation = match &upstream_message {
WreqWsMessage::Text(text) => {
let adapter = bound.adapter;
match parsed_upstream_frame.as_ref() {
Some(frame) => bound
.turn_state
.attempt_mut()
.and_then(|turn| turn.observe_upstream_frame(frame, adapter)),
None => {
if let Some(turn) = bound.turn_state.attempt_mut() {
turn.observe_invalid_upstream_text(text.as_str())
}
else {
None
}
}
}
}
_ => None,
};
if matches!(
observation,
Some(ResponsesWebSocketTurnObservation::Started)
| Some(ResponsesWebSocketTurnObservation::Terminal(_))
) {
if let Some(turn) = bound.turn_state.attempt_mut() {
turn.mark_stream_started(state).await;
}
}
let terminal_outcome = match observation {
Some(ResponsesWebSocketTurnObservation::Terminal(outcome)) => Some(outcome),
_ => None,
};
if matches!(&upstream_message, WreqWsMessage::Text(_))
&& parsed_upstream_frame.is_none()
{
let policy = fatal_relay_policy(FatalRelaySignal::InvalidUpstreamText);
finalize_active_turn(
bound,
state,
terminal_outcome.unwrap_or_else(
ResponsesWebSocketTurnOutcome::upstream_receive_failed,
),
)
.await;
send_responses_websocket_error(
client_socket,
policy.status_code,
"server_error",
policy.error_code,
policy.client_message,
)
.await;
close_bound_upstream(bound).await;
close_client_socket(
client_socket,
policy.close_code,
policy.close_reason,
)
.await;
break;
}
let is_close = matches!(upstream_message, WreqWsMessage::Close(_));
let drain_for_adapter = adapter_drain_ready(
bound.pending_adapter_drain,
bound.turn_state.response_in_flight(),
observation,
is_close,
);
let quota_facts = QuotaRelayFacts {
drain_ready: drain_for_adapter,
retry_current_turn: bound
.pending_adapter_drain
.is_some_and(|directive| directive.retry_current_turn),
transparent_retry_failed: false,
usage_limit_error: parsed_upstream_event.is_some_and(is_usage_limit_error_event),
continuation_retry_eligible: active_continuation_can_retry_from_full_input(bound),
upstream_closed: is_close,
};
let mut quota_relay_action = classify_quota_relay(quota_facts);
if matches!(quota_relay_action, QuotaRelayAction::AttemptTransparentRetry) {
// detach_attempt 保留 logical turn:重试是同一轮请求的下一个 attempt。
let mut retry_turn = bound.turn_state.detach_attempt().map(ActiveProviderAttempt::disarm);
if let Some(turn) = retry_turn.as_mut() {
turn.release_admission().await;
}
if retry_active_turn_after_quota_exhaustion(bound, state, context).await {
if let Some(turn) = retry_turn {
queue_turn_finalization(
bound,
state,
turn,
terminal_outcome.unwrap_or_else(
ResponsesWebSocketTurnOutcome::upstream_closed,
),
)
.await;
}
continue;
}
if let Some(turn) = retry_turn {
let restored = bound
.turn_state
.resume(ActiveProviderAttempt::new(state, turn));
debug_assert!(
restored.is_ok(),
"a failed transparent retry must be able to restore its attempt"
);
drop(restored);
}
quota_relay_action = classify_quota_relay(QuotaRelayFacts {
retry_current_turn: false,
transparent_retry_failed: true,
..quota_facts
});
}
if matches!(
quota_relay_action,
QuotaRelayAction::RequestFullContinuationRetry
) {
let directive = bound
.pending_adapter_drain
.expect("adapter drain state should be present");
debug!(
event_name = "responses_websocket_continuation_retry_required",
log_type = "event",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = "previous_response_not_found",
"gateway will ask the client to retry the continuation with complete input"
);
let mut turn = bound.turn_state.end().map(ActiveProviderAttempt::disarm);
if let Some(active_turn) = turn.as_mut() {
active_turn.release_admission().await;
}
send_previous_response_not_found(client_socket).await;
if let Some(turn) = turn {
queue_turn_finalization(
bound,
state,
turn,
terminal_outcome.unwrap_or_else(
ResponsesWebSocketTurnOutcome::upstream_closed,
),
)
.await;
}
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
if matches!(quota_relay_action, QuotaRelayAction::ForwardQuotaAndDetach) {
let directive = bound
.pending_adapter_drain
.expect("adapter drain state should be present");
finalize_active_turn(
bound,
state,
terminal_outcome
.unwrap_or_else(ResponsesWebSocketTurnOutcome::provider_quota_exhausted),
)
.await;
send_gateway_error_with_status(
client_socket,
429,
directive.error_code,
"Provider connection closed after reporting exhausted quota; send a new response.create to select another Provider connection",
)
.await;
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
if let Err(error) = send_client_message(
client_socket,
upstream_message_to_client(upstream_message.clone()),
).await {
warn!(
event_name = "responses_websocket_client_send_failed",
log_type = "ops",
transport = WEBSOCKET_LOG_TRANSPORT,
websocket = true,
trace_id = %context.trace_id,
error_code = error.as_str(),
provider_terminal_reached = terminal_outcome.is_some(),
"gateway could not relay a provider event to the client"
);
// 投递失败是独立事实,不能覆盖已经到达的 provider 终态:
// 供应商已经完成推理并消耗 token,账单按它的终态计。
bound
.turn_state
.record_client_delivery_aborted(CLIENT_DELIVERY_FAILED_REASON);
finalize_active_turn(
bound,
state,
settle_signal_for_client_delivery_failure(terminal_outcome),
).await;
close_bound_upstream(bound).await;
break;
}
if let (Some(turn), Some(frame)) =
(bound.turn_state.attempt_mut(), parsed_upstream_frame.as_ref())
{
turn.capture_client_frame(frame.event());
}
if let Some(outcome) = terminal_outcome {
finalize_active_turn(bound, state, outcome).await;
} else if is_close {
finalize_active_turn(
bound,
state,
ResponsesWebSocketTurnOutcome::upstream_closed(),
)
.await;
}
if drain_for_adapter {
let directive = bound
.pending_adapter_drain
.expect("adapter drain state should be present");
detach_exhausted_upstream(bound, directive, &context.trace_id).await;
continue;
}
if is_close {
bound.upstream = None;
break;
}
}
}
}
}
async fn wait_for_connection_permit_loss(permit: Option<&aether_runtime::AdmissionPermit>) {
let Some(permit) = permit else {
std::future::pending::<()>().await;
return;
};
let mut health = tokio::time::interval(Duration::from_secs(1));
health.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
health.tick().await;
if !permit.is_healthy() {
return;
}
}
}