mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-05 00:47:48 +08:00
feat(gateway): Codex/OpenAI Responses WebSocket 代理模式
在 /v1/responses 上支持 WebSocket 升级,把客户端帧中继到上游 Codex / OpenAI Responses WebSocket 端点,同时保持既有的路由、鉴权、配额与用量 语义: - 路由与准入:control/route/ai.rs 识别 WebSocket 升级请求; websocket/ingress.rs 复用 API Key 鉴权、IP 规则与并发许可,并引入 独立的 WebSocket 连接许可 - 中继:websocket/responses/* 按 connection / session / turn 分层, 帧解析归一化、socket 写入有界、continuation 保持调度亲和性 - 配额:orchestration/codex_quota_breaker.rs 在账号配额耗尽时熔断并 自动恢复,不再直接断开客户端连接 - 用量:每个 turn 的终态用量落库,request_metadata 记录 websocket_mode / websocket_transport,管理端与 usage 视图暴露 is_websocket - 管理端:provider 可配置 Responses WebSocket 开关
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
mod body_buffer;
|
||||
mod local;
|
||||
mod websocket;
|
||||
|
||||
use self::body_buffer::{
|
||||
buffer_and_normalize_request_body, build_request_body_buffer_error_response,
|
||||
@@ -8,6 +9,7 @@ use self::body_buffer::{
|
||||
use self::local::{
|
||||
maybe_build_local_admin_proxy_response, maybe_build_local_internal_proxy_response,
|
||||
};
|
||||
pub(crate) use self::websocket::responses::responses_websocket;
|
||||
use super::internal::resolve_local_proxy_execution_path;
|
||||
pub(crate) use super::public::matches_model_mapping_for_models;
|
||||
use crate::ai_serving::api::{
|
||||
|
||||
@@ -0,0 +1,293 @@
|
||||
//! Authenticated public WebSocket upgrade admission shared by AI adapters.
|
||||
|
||||
use std::future::Future;
|
||||
use std::net::SocketAddr;
|
||||
|
||||
use axum::body::Body;
|
||||
use axum::extract::ws::{WebSocket, WebSocketUpgrade};
|
||||
use axum::http::{HeaderMap, Method, Response, StatusCode, Uri};
|
||||
use tracing::{info, warn};
|
||||
|
||||
use crate::api::response::{
|
||||
build_local_auth_rejection_response, build_local_http_error_response,
|
||||
build_local_overloaded_response,
|
||||
};
|
||||
use crate::control::{
|
||||
trusted_auth_local_rejection, GatewayControlDecision, GatewayLocalAuthRejection,
|
||||
};
|
||||
use crate::handlers::proxy::websocket::session::{WebSocketSessionLimits, WEBSOCKET_LOG_TRANSPORT};
|
||||
use crate::handlers::shared::ip_rules_allow;
|
||||
use crate::headers::{effective_client_ip, extract_or_generate_trace_id};
|
||||
use crate::router::RequestAdmissionError;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
/// Request facts that survive the HTTP Upgrade and are needed by a protocol
|
||||
/// adapter for planning, rate limiting, and connection-scoped audit logs.
|
||||
pub(crate) struct WebSocketRequestContext {
|
||||
pub(crate) trace_id: String,
|
||||
pub(crate) headers: HeaderMap,
|
||||
pub(crate) uri: Uri,
|
||||
pub(crate) remote_addr: SocketAddr,
|
||||
pub(crate) decision: GatewayControlDecision,
|
||||
pub(crate) rpm_bypassed: bool,
|
||||
/// Held for the lifetime of the upgraded socket. The Responses session
|
||||
/// polls its health and closes the client when a distributed lease is
|
||||
/// revoked or expires.
|
||||
pub(crate) websocket_connection_permit: Option<aether_runtime::AdmissionPermit>,
|
||||
}
|
||||
|
||||
/// Adapter-specific wording and event identifiers for generic upgrade checks.
|
||||
#[derive(Clone, Copy)]
|
||||
pub(crate) struct WebSocketIngressSpec {
|
||||
pub(crate) route_unavailable_message: &'static str,
|
||||
pub(crate) ip_whitelist_failure_event_name: &'static str,
|
||||
}
|
||||
|
||||
/// Performs the HTTP-only part of an AI WebSocket request.
|
||||
///
|
||||
/// The ordinary request permit covers only the HTTP Upgrade window. A
|
||||
/// dedicated WebSocket connection permit is held for the socket lifetime so
|
||||
/// idle clients cannot consume capacity reserved for normal HTTP requests.
|
||||
pub(crate) async fn upgrade_authenticated_ai_websocket<F, Fut>(
|
||||
state: AppState,
|
||||
remote_addr: SocketAddr,
|
||||
ws: WebSocketUpgrade,
|
||||
headers: HeaderMap,
|
||||
uri: Uri,
|
||||
limits: WebSocketSessionLimits,
|
||||
spec: WebSocketIngressSpec,
|
||||
run_session: F,
|
||||
) -> Result<Response<Body>, GatewayError>
|
||||
where
|
||||
F: FnOnce(WebSocket, AppState, WebSocketRequestContext) -> Fut + Send + 'static,
|
||||
Fut: Future<Output = ()> + Send + 'static,
|
||||
{
|
||||
let trace_id = extract_or_generate_trace_id(&headers);
|
||||
let client_ip = effective_client_ip(&headers, &remote_addr);
|
||||
if state.admin_security_ip_blacklisted(client_ip).await? {
|
||||
return build_local_http_error_response(
|
||||
&trace_id,
|
||||
None,
|
||||
StatusCode::FORBIDDEN,
|
||||
"当前 IP 已被禁止访问",
|
||||
);
|
||||
}
|
||||
|
||||
let request_context = crate::control::resolve_public_request_context(
|
||||
&state,
|
||||
&Method::GET,
|
||||
&uri,
|
||||
&headers,
|
||||
&trace_id,
|
||||
)
|
||||
.await?;
|
||||
let Some(decision) = request_context.control_decision else {
|
||||
return build_local_http_error_response(
|
||||
&trace_id,
|
||||
None,
|
||||
StatusCode::NOT_FOUND,
|
||||
spec.route_unavailable_message,
|
||||
);
|
||||
};
|
||||
if let Some(rejection) = trusted_auth_local_rejection(Some(&decision), &headers) {
|
||||
return build_local_auth_rejection_response(&trace_id, Some(&decision), &rejection);
|
||||
}
|
||||
let Some(auth_context) = decision.auth_context.as_ref() else {
|
||||
return build_local_auth_rejection_response(
|
||||
&trace_id,
|
||||
Some(&decision),
|
||||
&GatewayLocalAuthRejection::InvalidApiKey,
|
||||
);
|
||||
};
|
||||
if !auth_context.access_allowed {
|
||||
return build_local_auth_rejection_response(
|
||||
&trace_id,
|
||||
Some(&decision),
|
||||
&GatewayLocalAuthRejection::InvalidApiKey,
|
||||
);
|
||||
}
|
||||
if !ip_rules_allow(auth_context.ip_rules.as_deref(), client_ip) {
|
||||
return build_local_auth_rejection_response(
|
||||
&trace_id,
|
||||
Some(&decision),
|
||||
&GatewayLocalAuthRejection::IpNotAllowed {
|
||||
remote_ip: client_ip.to_string(),
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
let ip_whitelisted = match state.admin_security_ip_whitelisted(client_ip).await {
|
||||
Ok(value) => value,
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = spec.ip_whitelist_failure_event_name,
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %trace_id,
|
||||
client_ip = %client_ip,
|
||||
error = ?error,
|
||||
"gateway continued with WebSocket rate limiting after IP whitelist check error"
|
||||
);
|
||||
false
|
||||
}
|
||||
};
|
||||
let request_permit = match state.try_acquire_request_permit().await {
|
||||
Ok(permit) => permit,
|
||||
Err(error) => {
|
||||
return websocket_admission_error_response(
|
||||
&trace_id,
|
||||
&decision,
|
||||
Some(uri.path()),
|
||||
error,
|
||||
)
|
||||
}
|
||||
};
|
||||
let websocket_connection_permit = match state.try_acquire_websocket_connection_permit().await {
|
||||
Ok(permit) => permit,
|
||||
Err(error) => {
|
||||
return websocket_admission_error_response(
|
||||
&trace_id,
|
||||
&decision,
|
||||
Some(uri.path()),
|
||||
error,
|
||||
)
|
||||
}
|
||||
};
|
||||
|
||||
let context = WebSocketRequestContext {
|
||||
trace_id,
|
||||
headers,
|
||||
uri,
|
||||
remote_addr,
|
||||
decision,
|
||||
rpm_bypassed: ip_whitelisted,
|
||||
websocket_connection_permit,
|
||||
};
|
||||
Ok(ws
|
||||
.max_frame_size(limits.max_frame_size)
|
||||
.max_message_size(limits.max_message_size)
|
||||
.on_upgrade(move |socket| async move {
|
||||
drop(request_permit);
|
||||
run_session(socket, state, context).await;
|
||||
}))
|
||||
}
|
||||
|
||||
fn websocket_admission_error_response(
|
||||
trace_id: &str,
|
||||
decision: &GatewayControlDecision,
|
||||
request_path: Option<&str>,
|
||||
error: RequestAdmissionError,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
match error {
|
||||
RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Saturated {
|
||||
gate,
|
||||
limit,
|
||||
})
|
||||
| RequestAdmissionError::Distributed(
|
||||
aether_runtime_state::RuntimeSemaphoreError::Saturated { gate, limit },
|
||||
)
|
||||
| RequestAdmissionError::Distributed(
|
||||
aether_runtime_state::RuntimeSemaphoreError::Unavailable { gate, limit, .. },
|
||||
) => build_local_overloaded_response(trace_id, Some(decision), request_path, gate, limit),
|
||||
RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Closed { gate }) => Err(
|
||||
GatewayError::Internal(format!("gateway concurrency gate {gate} is closed")),
|
||||
),
|
||||
RequestAdmissionError::Distributed(
|
||||
aether_runtime_state::RuntimeSemaphoreError::InvalidConfiguration(message),
|
||||
) => Err(GatewayError::Internal(message)),
|
||||
}
|
||||
}
|
||||
|
||||
/// Connection-level access log fields which are independent of a protocol's
|
||||
/// per-turn usage lifecycle.
|
||||
#[derive(Clone, Copy)]
|
||||
pub(crate) struct WebSocketConnectionLogSpec {
|
||||
pub(crate) opened_event_name: &'static str,
|
||||
pub(crate) closed_event_name: &'static str,
|
||||
pub(crate) opened_message: &'static str,
|
||||
pub(crate) closed_message: &'static str,
|
||||
pub(crate) execution_path: &'static str,
|
||||
pub(crate) provider_type: &'static str,
|
||||
}
|
||||
|
||||
pub(crate) struct WebSocketConnectionLog {
|
||||
spec: WebSocketConnectionLogSpec,
|
||||
trace_id: String,
|
||||
remote_addr: SocketAddr,
|
||||
path: String,
|
||||
route_class: String,
|
||||
user_id: String,
|
||||
api_key_id: String,
|
||||
started_at: std::time::Instant,
|
||||
}
|
||||
|
||||
impl WebSocketConnectionLog {
|
||||
pub(crate) fn new(context: &WebSocketRequestContext, spec: WebSocketConnectionLogSpec) -> Self {
|
||||
let auth_context = context.decision.auth_context.as_ref();
|
||||
Self {
|
||||
spec,
|
||||
trace_id: context.trace_id.clone(),
|
||||
remote_addr: context.remote_addr,
|
||||
path: context.uri.path().to_string(),
|
||||
route_class: context
|
||||
.decision
|
||||
.route_class
|
||||
.as_deref()
|
||||
.unwrap_or("ai_public")
|
||||
.to_string(),
|
||||
user_id: auth_context
|
||||
.map(|auth_context| auth_context.user_id.clone())
|
||||
.unwrap_or_else(|| "-".to_string()),
|
||||
api_key_id: auth_context
|
||||
.map(|auth_context| auth_context.api_key_id.clone())
|
||||
.unwrap_or_else(|| "-".to_string()),
|
||||
started_at: std::time::Instant::now(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn log_opened(&self) {
|
||||
info!(
|
||||
event_name = self.spec.opened_event_name,
|
||||
log_type = "access",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
status = "upgraded",
|
||||
status_code = 101u16,
|
||||
trace_id = %self.trace_id,
|
||||
remote_addr = %self.remote_addr,
|
||||
method = "GET",
|
||||
path = %self.path,
|
||||
user_id = %self.user_id,
|
||||
api_key_id = %self.api_key_id,
|
||||
route_class = %self.route_class,
|
||||
execution_path = self.spec.execution_path,
|
||||
provider_type = self.spec.provider_type,
|
||||
message = self.spec.opened_message,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for WebSocketConnectionLog {
|
||||
fn drop(&mut self) {
|
||||
info!(
|
||||
event_name = self.spec.closed_event_name,
|
||||
log_type = "access",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
status = "closed",
|
||||
status_code = 101u16,
|
||||
trace_id = %self.trace_id,
|
||||
remote_addr = %self.remote_addr,
|
||||
method = "GET",
|
||||
path = %self.path,
|
||||
user_id = %self.user_id,
|
||||
api_key_id = %self.api_key_id,
|
||||
route_class = %self.route_class,
|
||||
execution_path = self.spec.execution_path,
|
||||
provider_type = self.spec.provider_type,
|
||||
elapsed_ms = self.started_at.elapsed().as_millis() as u64,
|
||||
message = self.spec.closed_message,
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
//! Shared infrastructure for public AI WebSocket bridges.
|
||||
//!
|
||||
//! Protocol adapters live below [`responses`]. This layer deliberately owns
|
||||
//! only transport concerns that are common to future adapters: authenticated
|
||||
//! upgrade admission, connection limits, upstream handshakes, and frame
|
||||
//! conversion. It does not interpret provider events or make routing
|
||||
//! decisions.
|
||||
|
||||
pub(crate) mod ingress;
|
||||
pub(crate) mod responses;
|
||||
pub(crate) mod session;
|
||||
pub(crate) mod transport;
|
||||
@@ -0,0 +1,186 @@
|
||||
//! Provider-specific hooks for the standard Responses WebSocket session.
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::adapters::CODEX_RESPONSES_WEBSOCKET_ADAPTER;
|
||||
use crate::ai_serving::AiExecutionDecision;
|
||||
use crate::handlers::proxy::websocket::transport::UpstreamWebSocketErrorCodes;
|
||||
use crate::orchestration::ResponsesWebSocketAdapter;
|
||||
use crate::AppState;
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub(super) struct ResponsesWebSocketDrainDirective {
|
||||
pub(super) error_code: &'static str,
|
||||
/// The terminal upstream event may be replayed only when the session has
|
||||
/// not exposed any standard Responses event to the client.
|
||||
pub(super) retry_current_turn: bool,
|
||||
/// When present, the exhausted provider key remains excluded from later
|
||||
/// turns on this client socket until the upstream's reported reset time.
|
||||
pub(super) retry_exclusion_until_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
/// Provider-specific observation produced while relaying an upstream frame.
|
||||
/// The session can make the retry/drain decision synchronously, while the
|
||||
/// optional persistence sink runs outside the frame-forwarding path.
|
||||
#[derive(Debug, Clone)]
|
||||
pub(super) struct ResponsesWebSocketAdapterObservation {
|
||||
pub(super) drain: Option<ResponsesWebSocketDrainDirective>,
|
||||
pub(super) quota_metadata: Option<Value>,
|
||||
}
|
||||
|
||||
/// Provider identity used by the shared session's temporary exclusion table.
|
||||
/// The session does not need to know how a provider derives its account id.
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub(super) struct ResponsesWebSocketExclusionIdentity {
|
||||
pub(super) account_id: Option<String>,
|
||||
}
|
||||
|
||||
/// Whether receiving an upstream event still leaves the active client turn
|
||||
/// safe to replay on a freshly bound upstream. The shared session keeps the
|
||||
/// conservative default; provider adapters may explicitly whitelist their
|
||||
/// documented, pre-response advisory events.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(super) enum ResponsesWebSocketRebindSafety {
|
||||
Safe,
|
||||
Unsafe { reason: &'static str },
|
||||
}
|
||||
|
||||
/// Boundary between the standard Responses protocol engine and provider
|
||||
/// behavior. Adapters receive already-planned provider requests; they never
|
||||
/// own public WebSocket parsing, turn accounting, or model scheduling.
|
||||
#[async_trait]
|
||||
pub(super) trait ResponsesWebSocketProtocolAdapter: Send + Sync {
|
||||
fn kind(&self) -> ResponsesWebSocketAdapter;
|
||||
|
||||
fn upstream_errors(&self) -> UpstreamWebSocketErrorCodes;
|
||||
|
||||
/// Adds provider-specific metadata to an otherwise standard Responses
|
||||
/// stream report. The event payload is never rewritten for the client.
|
||||
fn decorate_turn_report_context(&self, report_context: &mut Option<Value>, event: &Value);
|
||||
|
||||
/// Whether this adapter needs the shared session to parse each upstream
|
||||
/// text event before normal turn accounting runs.
|
||||
fn observes_upstream_events(&self) -> bool;
|
||||
|
||||
/// Classifies whether a received upstream event can be followed by a
|
||||
/// transparent quota-driven rebind. An adapter must return `Safe` only
|
||||
/// for events that neither create public Responses state nor make a replay
|
||||
/// observably ambiguous to the client.
|
||||
fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety;
|
||||
|
||||
/// Lets an adapter classify provider-only events. Returning a directive
|
||||
/// asks the shared session to drain after the active standard response.
|
||||
fn observe_upstream_event(&self, event: &Value)
|
||||
-> Option<ResponsesWebSocketAdapterObservation>;
|
||||
|
||||
fn exhaustion_exclusion_identity(
|
||||
&self,
|
||||
_decision: &AiExecutionDecision,
|
||||
) -> Option<ResponsesWebSocketExclusionIdentity> {
|
||||
None
|
||||
}
|
||||
|
||||
/// Persists an adapter observation outside the frame-forwarding path.
|
||||
async fn persist_upstream_observation(
|
||||
&self,
|
||||
state: &AppState,
|
||||
trace_id: &str,
|
||||
report_context: Option<&Value>,
|
||||
observation: ResponsesWebSocketAdapterObservation,
|
||||
);
|
||||
}
|
||||
|
||||
pub(super) fn resolve_responses_websocket_adapter(
|
||||
kind: ResponsesWebSocketAdapter,
|
||||
) -> &'static dyn ResponsesWebSocketProtocolAdapter {
|
||||
match kind {
|
||||
ResponsesWebSocketAdapter::Standard => &STANDARD_RESPONSES_WEBSOCKET_ADAPTER,
|
||||
ResponsesWebSocketAdapter::Codex => &CODEX_RESPONSES_WEBSOCKET_ADAPTER,
|
||||
}
|
||||
}
|
||||
|
||||
struct StandardResponsesWebSocketAdapter;
|
||||
|
||||
const STANDARD_UPSTREAM_WEBSOCKET_ERRORS: UpstreamWebSocketErrorCodes =
|
||||
UpstreamWebSocketErrorCodes {
|
||||
upstream_url_missing: "responses_upstream_url_missing",
|
||||
upstream_url_invalid: "responses_upstream_url_invalid",
|
||||
headers_invalid: "responses_websocket_headers_invalid",
|
||||
client_build_failed: "responses_websocket_client_build_failed",
|
||||
proxy_invalid: "responses_websocket_proxy_invalid",
|
||||
tunnel_proxy_unsupported: "responses_websocket_tunnel_proxy_unsupported",
|
||||
handshake_failed: "responses_websocket_handshake_failed",
|
||||
upgrade_rejected: "responses_websocket_upgrade_rejected",
|
||||
upgrade_failed: "responses_websocket_upgrade_failed",
|
||||
};
|
||||
|
||||
static STANDARD_RESPONSES_WEBSOCKET_ADAPTER: StandardResponsesWebSocketAdapter =
|
||||
StandardResponsesWebSocketAdapter;
|
||||
|
||||
#[async_trait]
|
||||
impl ResponsesWebSocketProtocolAdapter for StandardResponsesWebSocketAdapter {
|
||||
fn kind(&self) -> ResponsesWebSocketAdapter {
|
||||
ResponsesWebSocketAdapter::Standard
|
||||
}
|
||||
|
||||
fn upstream_errors(&self) -> UpstreamWebSocketErrorCodes {
|
||||
STANDARD_UPSTREAM_WEBSOCKET_ERRORS
|
||||
}
|
||||
|
||||
fn decorate_turn_report_context(&self, _report_context: &mut Option<Value>, _event: &Value) {}
|
||||
|
||||
fn observes_upstream_events(&self) -> bool {
|
||||
false
|
||||
}
|
||||
|
||||
fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety {
|
||||
let reason = if is_standard_responses_event(event) {
|
||||
"standard_response_event"
|
||||
} else {
|
||||
"unrecognized_upstream_event"
|
||||
};
|
||||
ResponsesWebSocketRebindSafety::Unsafe { reason }
|
||||
}
|
||||
|
||||
fn observe_upstream_event(
|
||||
&self,
|
||||
_event: &Value,
|
||||
) -> Option<ResponsesWebSocketAdapterObservation> {
|
||||
None
|
||||
}
|
||||
|
||||
async fn persist_upstream_observation(
|
||||
&self,
|
||||
_state: &AppState,
|
||||
_trace_id: &str,
|
||||
_report_context: Option<&Value>,
|
||||
_observation: ResponsesWebSocketAdapterObservation,
|
||||
) {
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn is_standard_responses_event(event: &Value) -> bool {
|
||||
event
|
||||
.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|event_type| event_type.starts_with("response."))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{resolve_responses_websocket_adapter, ResponsesWebSocketProtocolAdapter};
|
||||
use crate::orchestration::ResponsesWebSocketAdapter;
|
||||
|
||||
#[test]
|
||||
fn standard_adapter_has_no_codex_extensions() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
|
||||
assert_eq!(adapter.kind(), ResponsesWebSocketAdapter::Standard);
|
||||
assert!(!adapter.observes_upstream_events());
|
||||
assert_eq!(
|
||||
adapter.upstream_errors().handshake_failed,
|
||||
"responses_websocket_handshake_failed"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,290 @@
|
||||
//! Codex-specific extensions for the standard Responses WebSocket session.
|
||||
|
||||
use async_trait::async_trait;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::super::adapter::{
|
||||
is_standard_responses_event, ResponsesWebSocketAdapterObservation,
|
||||
ResponsesWebSocketDrainDirective, ResponsesWebSocketExclusionIdentity,
|
||||
ResponsesWebSocketProtocolAdapter, ResponsesWebSocketRebindSafety,
|
||||
};
|
||||
use crate::ai_serving::AiExecutionDecision;
|
||||
use crate::clock::current_unix_secs;
|
||||
use crate::handlers::proxy::websocket::transport::UpstreamWebSocketErrorCodes;
|
||||
use crate::orchestration::{
|
||||
codex_account_id_from_headers, codex_quota_exhaustion_reset_at,
|
||||
sync_codex_websocket_quota_metadata, ResponsesWebSocketAdapter,
|
||||
};
|
||||
use crate::AppState;
|
||||
|
||||
const CODEX_WEBSOCKET_LOG_TARGET: &str = "aether_gateway::handlers::proxy::codex_ws";
|
||||
const CODEX_WEBSOCKET_RATE_LIMITS_REPORT_CONTEXT_FIELD: &str = "codex_websocket_rate_limits";
|
||||
|
||||
const CODEX_UPSTREAM_WEBSOCKET_ERRORS: UpstreamWebSocketErrorCodes = UpstreamWebSocketErrorCodes {
|
||||
upstream_url_missing: "codex_upstream_url_missing",
|
||||
upstream_url_invalid: "codex_upstream_url_invalid",
|
||||
headers_invalid: "codex_websocket_headers_invalid",
|
||||
client_build_failed: "codex_websocket_client_build_failed",
|
||||
proxy_invalid: "codex_websocket_proxy_invalid",
|
||||
tunnel_proxy_unsupported: "codex_websocket_tunnel_proxy_unsupported",
|
||||
handshake_failed: "codex_websocket_handshake_failed",
|
||||
upgrade_rejected: "codex_websocket_upgrade_rejected",
|
||||
upgrade_failed: "codex_websocket_upgrade_failed",
|
||||
};
|
||||
|
||||
pub(crate) static CODEX_RESPONSES_WEBSOCKET_ADAPTER: CodexResponsesWebSocketAdapter =
|
||||
CodexResponsesWebSocketAdapter;
|
||||
|
||||
pub(crate) struct CodexResponsesWebSocketAdapter;
|
||||
|
||||
#[async_trait]
|
||||
impl ResponsesWebSocketProtocolAdapter for CodexResponsesWebSocketAdapter {
|
||||
fn kind(&self) -> ResponsesWebSocketAdapter {
|
||||
ResponsesWebSocketAdapter::Codex
|
||||
}
|
||||
|
||||
fn upstream_errors(&self) -> UpstreamWebSocketErrorCodes {
|
||||
CODEX_UPSTREAM_WEBSOCKET_ERRORS
|
||||
}
|
||||
|
||||
fn decorate_turn_report_context(&self, report_context: &mut Option<Value>, event: &Value) {
|
||||
let Some(rate_limits) = parse_codex_rate_limits(event) else {
|
||||
return;
|
||||
};
|
||||
let context = report_context.get_or_insert_with(|| Value::Object(Map::new()));
|
||||
let Some(context) = context.as_object_mut() else {
|
||||
return;
|
||||
};
|
||||
context.insert(
|
||||
CODEX_WEBSOCKET_RATE_LIMITS_REPORT_CONTEXT_FIELD.to_string(),
|
||||
rate_limits,
|
||||
);
|
||||
}
|
||||
|
||||
fn observes_upstream_events(&self) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn rebind_safety_for_upstream_event(&self, event: &Value) -> ResponsesWebSocketRebindSafety {
|
||||
if let Some(chunks) = event.get("chunks").and_then(Value::as_array) {
|
||||
if chunks.is_empty() {
|
||||
return ResponsesWebSocketRebindSafety::Unsafe {
|
||||
reason: "unrecognized_upstream_event",
|
||||
};
|
||||
}
|
||||
return chunks
|
||||
.iter()
|
||||
.map(codex_direct_rebind_safety)
|
||||
.find(|safety| matches!(safety, ResponsesWebSocketRebindSafety::Unsafe { .. }))
|
||||
.unwrap_or(ResponsesWebSocketRebindSafety::Safe);
|
||||
}
|
||||
codex_direct_rebind_safety(event)
|
||||
}
|
||||
|
||||
fn observe_upstream_event(
|
||||
&self,
|
||||
event: &Value,
|
||||
) -> Option<ResponsesWebSocketAdapterObservation> {
|
||||
let rate_limits = parse_codex_rate_limits(event)?;
|
||||
let exhausted =
|
||||
aether_admin::provider::quota::codex_rate_limit_metadata_exhausted(&rate_limits);
|
||||
let retry_exclusion_until_unix_secs =
|
||||
codex_quota_exhaustion_reset_at(&rate_limits, current_unix_secs());
|
||||
Some(ResponsesWebSocketAdapterObservation {
|
||||
drain: exhausted.then_some(ResponsesWebSocketDrainDirective {
|
||||
error_code: "codex_account_quota_exhausted",
|
||||
retry_current_turn: true,
|
||||
retry_exclusion_until_unix_secs,
|
||||
}),
|
||||
quota_metadata: Some(rate_limits),
|
||||
})
|
||||
}
|
||||
|
||||
fn exhaustion_exclusion_identity(
|
||||
&self,
|
||||
decision: &AiExecutionDecision,
|
||||
) -> Option<ResponsesWebSocketExclusionIdentity> {
|
||||
Some(ResponsesWebSocketExclusionIdentity {
|
||||
account_id: codex_account_id_from_headers(&decision.provider_request_headers)
|
||||
.map(str::to_string),
|
||||
})
|
||||
}
|
||||
|
||||
async fn persist_upstream_observation(
|
||||
&self,
|
||||
state: &AppState,
|
||||
trace_id: &str,
|
||||
report_context: Option<&Value>,
|
||||
observation: ResponsesWebSocketAdapterObservation,
|
||||
) {
|
||||
let Some(rate_limits) = observation.quota_metadata else {
|
||||
return;
|
||||
};
|
||||
if let Err(error) =
|
||||
sync_codex_websocket_quota_metadata(state, report_context, rate_limits).await
|
||||
{
|
||||
tracing::warn!(
|
||||
target: CODEX_WEBSOCKET_LOG_TARGET,
|
||||
event_name = "codex_websocket_quota_sync_failed",
|
||||
log_type = "ops",
|
||||
transport = "websocket",
|
||||
websocket = true,
|
||||
trace_id = %trace_id,
|
||||
error = ?error,
|
||||
"gateway failed to persist Codex WebSocket quota metadata"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn codex_direct_rebind_safety(event: &Value) -> ResponsesWebSocketRebindSafety {
|
||||
let event_type = event
|
||||
.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.unwrap_or_default();
|
||||
if matches!(event_type, "codex.rate_limits" | "codex.response.metadata") {
|
||||
// Codex emits these as pre-response advisory metadata. They do
|
||||
// not create a public `response.*` object, so a replacement
|
||||
// upstream can safely emit its own current snapshot.
|
||||
return ResponsesWebSocketRebindSafety::Safe;
|
||||
}
|
||||
if event_type == "error" && parse_codex_rate_limits(event).is_some() {
|
||||
// The quota error is withheld from the client when the shared
|
||||
// session successfully rebinds, therefore it remains replay-safe.
|
||||
return ResponsesWebSocketRebindSafety::Safe;
|
||||
}
|
||||
let reason = if is_standard_responses_event(event) {
|
||||
"standard_response_event"
|
||||
} else {
|
||||
"unrecognized_upstream_event"
|
||||
};
|
||||
ResponsesWebSocketRebindSafety::Unsafe { reason }
|
||||
}
|
||||
|
||||
fn parse_codex_rate_limits(event: &Value) -> Option<Value> {
|
||||
aether_admin::provider::quota::parse_codex_websocket_rate_limits_response(
|
||||
event,
|
||||
current_unix_secs(),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
CodexResponsesWebSocketAdapter, ResponsesWebSocketProtocolAdapter,
|
||||
ResponsesWebSocketRebindSafety,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn codex_rate_limit_chunk_is_kept_for_the_terminal_report() {
|
||||
let adapter = CodexResponsesWebSocketAdapter;
|
||||
assert!(adapter.observes_upstream_events());
|
||||
let mut context = Some(json!({"key_id": "codex-key"}));
|
||||
adapter.decorate_turn_report_context(
|
||||
&mut context,
|
||||
&json!({
|
||||
"chunks": [{
|
||||
"type": "codex.rate_limits",
|
||||
"plan_type": "free",
|
||||
"rate_limits": {
|
||||
"allowed": true,
|
||||
"limit_reached": false,
|
||||
"primary": {
|
||||
"used_percent": 91,
|
||||
"window_minutes": 43200,
|
||||
"reset_after_seconds": 2590791
|
||||
}
|
||||
}
|
||||
}]
|
||||
}),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
context.as_ref().and_then(
|
||||
|context| context.pointer("/codex_websocket_rate_limits/primary_used_percent")
|
||||
),
|
||||
Some(&json!(91.0))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn usage_limit_error_is_kept_for_the_terminal_report() {
|
||||
let adapter = CodexResponsesWebSocketAdapter;
|
||||
let mut context = Some(json!({"key_id": "codex-key"}));
|
||||
adapter.decorate_turn_report_context(
|
||||
&mut context,
|
||||
&json!({
|
||||
"type": "error",
|
||||
"error": {
|
||||
"type": "usage_limit_reached",
|
||||
"plan_type": "free",
|
||||
"resets_at": 1_787_274_385u64,
|
||||
},
|
||||
"status_code": 429,
|
||||
"headers": {
|
||||
"X-Codex-Primary-Used-Percent": "100",
|
||||
"X-Codex-Primary-Reset-At": "1787274385",
|
||||
},
|
||||
}),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
context
|
||||
.as_ref()
|
||||
.and_then(|context| context.pointer("/codex_websocket_rate_limits/allowed")),
|
||||
Some(&json!(false))
|
||||
);
|
||||
assert_eq!(
|
||||
context.as_ref().and_then(|context| {
|
||||
context.pointer("/codex_websocket_rate_limits/primary_used_percent")
|
||||
}),
|
||||
Some(&json!(100.0))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn only_known_codex_pre_response_metadata_is_safe_to_rebind() {
|
||||
let adapter = CodexResponsesWebSocketAdapter;
|
||||
|
||||
assert_eq!(
|
||||
adapter.rebind_safety_for_upstream_event(&json!({
|
||||
"type": "codex.rate_limits",
|
||||
"rate_limits": {"allowed": true}
|
||||
})),
|
||||
ResponsesWebSocketRebindSafety::Safe
|
||||
);
|
||||
assert_eq!(
|
||||
adapter.rebind_safety_for_upstream_event(&json!({
|
||||
"type": "codex.response.metadata"
|
||||
})),
|
||||
ResponsesWebSocketRebindSafety::Safe
|
||||
);
|
||||
assert_eq!(
|
||||
adapter.rebind_safety_for_upstream_event(&json!({
|
||||
"chunks": [
|
||||
{"type": "codex.rate_limits", "rate_limits": {"allowed": true}},
|
||||
{"type": "codex.response.metadata"}
|
||||
]
|
||||
})),
|
||||
ResponsesWebSocketRebindSafety::Safe
|
||||
);
|
||||
assert_eq!(
|
||||
adapter.rebind_safety_for_upstream_event(&json!({
|
||||
"type": "response.created"
|
||||
})),
|
||||
ResponsesWebSocketRebindSafety::Unsafe {
|
||||
reason: "standard_response_event"
|
||||
}
|
||||
);
|
||||
assert_eq!(
|
||||
adapter.rebind_safety_for_upstream_event(&json!({
|
||||
"type": "codex.unknown"
|
||||
})),
|
||||
ResponsesWebSocketRebindSafety::Unsafe {
|
||||
reason: "unrecognized_upstream_event"
|
||||
}
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
//! Provider-specific Responses WebSocket adapters.
|
||||
|
||||
mod codex;
|
||||
|
||||
pub(super) use codex::CODEX_RESPONSES_WEBSOCKET_ADAPTER;
|
||||
@@ -0,0 +1,80 @@
|
||||
//! Per-turn resource admission for the Responses WebSocket bridge.
|
||||
//!
|
||||
//! A WebSocket connection may live for a long time, but each `response.create`
|
||||
//! is still one active upstream execution. Keep the resource leases attached
|
||||
//! to the turn instead of the socket so idle connections do not consume
|
||||
//! upstream capacity.
|
||||
|
||||
use std::time::Instant;
|
||||
|
||||
use aether_contracts::ExecutionPlan;
|
||||
|
||||
use crate::execution_runtime::acquire_upstream_execution_gate;
|
||||
use crate::provider_pool_demand::{
|
||||
acquire_provider_pool_in_flight_guard, ProviderPoolInFlightGuard,
|
||||
};
|
||||
use crate::upstream_admission::UpstreamTargetAdmissionPermit;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
pub(super) struct ResponsesWebSocketTurnAdmission {
|
||||
upstream_execution: Option<aether_runtime::ConcurrencyPermit>,
|
||||
upstream_target: Option<UpstreamTargetAdmissionPermit>,
|
||||
provider_pool: Option<ProviderPoolInFlightGuard>,
|
||||
acquired_at: Instant,
|
||||
}
|
||||
|
||||
impl ResponsesWebSocketTurnAdmission {
|
||||
pub(super) async fn acquire(
|
||||
state: &AppState,
|
||||
plan: &ExecutionPlan,
|
||||
trace_id: &str,
|
||||
) -> Result<Self, GatewayError> {
|
||||
let upstream_execution = acquire_upstream_execution_gate(state, trace_id).await?;
|
||||
let upstream_target = match state
|
||||
.upstream_target_admission
|
||||
.acquire(plan, trace_id)
|
||||
.await
|
||||
{
|
||||
Ok(permit) => permit,
|
||||
Err(error) => {
|
||||
drop(upstream_execution);
|
||||
return Err(error);
|
||||
}
|
||||
};
|
||||
let provider_pool = acquire_provider_pool_in_flight_guard(
|
||||
state.runtime_state.clone(),
|
||||
&plan.provider_id,
|
||||
&plan.request_id,
|
||||
plan.candidate_id.as_deref(),
|
||||
&plan.key_id,
|
||||
)
|
||||
.await;
|
||||
|
||||
Ok(Self {
|
||||
upstream_execution,
|
||||
upstream_target,
|
||||
provider_pool,
|
||||
acquired_at: Instant::now(),
|
||||
})
|
||||
}
|
||||
|
||||
/// Release the distributed provider token before the turn's persistence
|
||||
/// work. The remaining permits are local RAII guards and are dropped with
|
||||
/// this value.
|
||||
pub(super) async fn release(mut self) {
|
||||
if let Some(provider_pool) = self.provider_pool.take() {
|
||||
provider_pool.release().await;
|
||||
}
|
||||
drop(self.upstream_target.take());
|
||||
drop(self.upstream_execution.take());
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for ResponsesWebSocketTurnAdmission {
|
||||
fn drop(&mut self) {
|
||||
crate::stage_metrics::observe_gateway_stage_ms(
|
||||
"websocket_turn_admission_held",
|
||||
self.acquired_at.elapsed().as_millis() as u64,
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,413 @@
|
||||
//! Identity of the physical upstream connection backing a Responses session.
|
||||
//!
|
||||
//! A Responses continuation carries state that lives on one provider socket.
|
||||
//! Comparing only the selected key is therefore not sufficient: transport
|
||||
//! settings, stable account headers, and the protocol adapter can all change
|
||||
//! the connection that would receive the next event. Rotating bearer values
|
||||
//! are intentionally excluded because they do not change an already-upgraded
|
||||
//! socket's physical binding.
|
||||
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::fmt;
|
||||
|
||||
use aether_contracts::{ProxySnapshot, ResolvedTransportProfile};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use super::adapter::ResponsesWebSocketProtocolAdapter;
|
||||
use crate::ai_serving::AiExecutionDecision;
|
||||
use crate::handlers::proxy::websocket::transport::{
|
||||
websocket_handshake_headers, websocket_upstream_url,
|
||||
};
|
||||
use crate::orchestration::ResponsesWebSocketAdapter;
|
||||
|
||||
/// Stable, comparable identity for the actual WebSocket connection target.
|
||||
///
|
||||
/// The identity deliberately owns the normalized handshake values rather than
|
||||
/// retaining a reference to the planner decision. A later re-plan can then
|
||||
/// be compared without accidentally ignoring a field that changes the
|
||||
/// physical connection.
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub(super) struct UpstreamBindingIdentity {
|
||||
adapter_kind: ResponsesWebSocketAdapter,
|
||||
provider_id: Option<String>,
|
||||
endpoint_id: Option<String>,
|
||||
key_id: Option<String>,
|
||||
upstream_url: String,
|
||||
handshake_headers: BTreeMap<String, String>,
|
||||
/// Authentication values are not part of a stable key binding when the
|
||||
/// planner has already supplied a key identity. If that identity is
|
||||
/// unavailable, retain only a one-way fingerprint so two accounts cannot
|
||||
/// accidentally share a continuation socket.
|
||||
auth_fingerprint: Option<[u8; 32]>,
|
||||
proxy: Option<ProxySnapshot>,
|
||||
transport_profile: Option<ResolvedTransportProfile>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(super) enum UpstreamBindingIdentityError {
|
||||
MissingUpstreamUrl,
|
||||
InvalidUpstreamUrl,
|
||||
InvalidHandshakeHeaders,
|
||||
}
|
||||
|
||||
impl UpstreamBindingIdentity {
|
||||
/// Builds an identity from the same normalized URL and headers used by
|
||||
/// the WebSocket transport client.
|
||||
pub(super) fn from_decision(
|
||||
adapter: &'static dyn ResponsesWebSocketProtocolAdapter,
|
||||
decision: &AiExecutionDecision,
|
||||
) -> Result<Self, UpstreamBindingIdentityError> {
|
||||
let raw_url = decision
|
||||
.upstream_url
|
||||
.as_deref()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.ok_or(UpstreamBindingIdentityError::MissingUpstreamUrl)?;
|
||||
let upstream_url = websocket_upstream_url(raw_url, "invalid")
|
||||
.map_err(|_| UpstreamBindingIdentityError::InvalidUpstreamUrl)?
|
||||
.to_string();
|
||||
|
||||
let headers = websocket_handshake_headers(&decision.provider_request_headers, "invalid")
|
||||
.map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?;
|
||||
let authentication_header_names = authentication_header_names(decision);
|
||||
let mut handshake_headers = BTreeMap::new();
|
||||
let mut authentication_headers = BTreeMap::new();
|
||||
for (name, value) in &headers {
|
||||
let name = name.as_str().to_ascii_lowercase();
|
||||
let value = value
|
||||
.to_str()
|
||||
.map_err(|_| UpstreamBindingIdentityError::InvalidHandshakeHeaders)?;
|
||||
if authentication_header_names.contains(name.as_str()) {
|
||||
authentication_headers.insert(name, value.to_string());
|
||||
} else {
|
||||
handshake_headers.insert(name, value.to_string());
|
||||
}
|
||||
}
|
||||
let auth_fingerprint = decision
|
||||
.key_id
|
||||
.is_none()
|
||||
.then(|| fingerprint_headers(&authentication_headers))
|
||||
.filter(|_| !authentication_headers.is_empty());
|
||||
|
||||
Ok(Self {
|
||||
adapter_kind: adapter.kind(),
|
||||
provider_id: decision.provider_id.clone(),
|
||||
endpoint_id: decision.endpoint_id.clone(),
|
||||
key_id: decision.key_id.clone(),
|
||||
upstream_url,
|
||||
handshake_headers,
|
||||
auth_fingerprint,
|
||||
proxy: effective_proxy_snapshot(decision.proxy.as_ref()),
|
||||
transport_profile: decision.transport_profile.clone(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Header names that carry credentials in the provider handshake. The
|
||||
/// planner's explicit `auth_header` extends this list for provider-specific
|
||||
/// schemes; unknown headers remain part of the stable handshake identity.
|
||||
fn authentication_header_names(decision: &AiExecutionDecision) -> BTreeSet<String> {
|
||||
let mut names = BTreeSet::from([
|
||||
"authorization".to_string(),
|
||||
"proxy-authorization".to_string(),
|
||||
"x-api-key".to_string(),
|
||||
"api-key".to_string(),
|
||||
"x-goog-api-key".to_string(),
|
||||
"x-azure-api-key".to_string(),
|
||||
]);
|
||||
if let Some(name) = decision
|
||||
.auth_header
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|name| !name.is_empty())
|
||||
{
|
||||
names.insert(name.to_ascii_lowercase());
|
||||
}
|
||||
names
|
||||
}
|
||||
|
||||
fn fingerprint_headers(headers: &BTreeMap<String, String>) -> [u8; 32] {
|
||||
let mut hasher = Sha256::new();
|
||||
for (name, value) in headers {
|
||||
hasher.update((name.len() as u64).to_be_bytes());
|
||||
hasher.update(name.as_bytes());
|
||||
hasher.update((value.len() as u64).to_be_bytes());
|
||||
hasher.update(value.as_bytes());
|
||||
}
|
||||
hasher.finalize().into()
|
||||
}
|
||||
|
||||
/// Normalize only values that are provably direct transport. Keep node/tunnel
|
||||
/// fields even though the current WebSocket builder rejects those proxies: a
|
||||
/// re-plan must not accidentally reuse an already-bound direct socket for a
|
||||
/// decision that selected a different proxy topology.
|
||||
fn effective_proxy_snapshot(proxy: Option<&ProxySnapshot>) -> Option<ProxySnapshot> {
|
||||
let proxy = proxy?;
|
||||
if proxy.enabled == Some(false) {
|
||||
return None;
|
||||
}
|
||||
let mut normalized = proxy.clone();
|
||||
normalized.url = normalized
|
||||
.url
|
||||
.take()
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty());
|
||||
normalized.mode = normalized
|
||||
.mode
|
||||
.take()
|
||||
.map(|value| value.trim().to_ascii_lowercase())
|
||||
.filter(|value| !value.is_empty());
|
||||
normalized.node_id = normalized
|
||||
.node_id
|
||||
.take()
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty());
|
||||
normalized.label = normalized
|
||||
.label
|
||||
.take()
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty());
|
||||
|
||||
let has_effective_proxy = normalized.url.is_some()
|
||||
|| normalized.node_id.is_some()
|
||||
|| normalized.mode.is_some()
|
||||
|| normalized.extra.is_some();
|
||||
has_effective_proxy.then_some(normalized)
|
||||
}
|
||||
|
||||
impl fmt::Debug for UpstreamBindingIdentity {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter
|
||||
.debug_struct("UpstreamBindingIdentity")
|
||||
.field("adapter_kind", &self.adapter_kind)
|
||||
.field("provider_id", &self.provider_id)
|
||||
.field("endpoint_id", &self.endpoint_id)
|
||||
.field("key_id", &self.key_id)
|
||||
.field("upstream_url", &self.upstream_url)
|
||||
.field(
|
||||
"handshake_header_names",
|
||||
&self.handshake_headers.keys().collect::<Vec<_>>(),
|
||||
)
|
||||
.field("proxy_configured", &self.proxy.is_some())
|
||||
.field(
|
||||
"transport_profile_id",
|
||||
&self
|
||||
.transport_profile
|
||||
.as_ref()
|
||||
.map(|profile| profile.profile_id.as_str()),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::json;
|
||||
|
||||
use super::{UpstreamBindingIdentity, UpstreamBindingIdentityError};
|
||||
use crate::ai_serving::AiExecutionDecision;
|
||||
use crate::handlers::proxy::websocket::responses::adapter::resolve_responses_websocket_adapter;
|
||||
use crate::orchestration::ResponsesWebSocketAdapter;
|
||||
|
||||
fn decision() -> AiExecutionDecision {
|
||||
AiExecutionDecision {
|
||||
action: "execute".to_string(),
|
||||
decision_kind: None,
|
||||
execution_strategy: None,
|
||||
conversion_mode: None,
|
||||
request_id: Some("request-1".to_string()),
|
||||
candidate_id: Some("candidate-1".to_string()),
|
||||
provider_name: Some("provider".to_string()),
|
||||
provider_type: Some("openai".to_string()),
|
||||
provider_id: Some("provider-1".to_string()),
|
||||
endpoint_id: Some("endpoint-1".to_string()),
|
||||
key_id: Some("key-1".to_string()),
|
||||
upstream_base_url: Some("https://api.example.test".to_string()),
|
||||
upstream_url: Some("https://api.example.test/v1/responses".to_string()),
|
||||
provider_request_method: Some("POST".to_string()),
|
||||
auth_header: Some("authorization".to_string()),
|
||||
auth_value: Some("Bearer secret".to_string()),
|
||||
provider_api_format: Some("openai:responses".to_string()),
|
||||
client_api_format: Some("openai:responses".to_string()),
|
||||
provider_contract: None,
|
||||
client_contract: None,
|
||||
model_name: Some("gpt-5.6-sol".to_string()),
|
||||
mapped_model: None,
|
||||
prompt_cache_key: None,
|
||||
extra_headers: BTreeMap::new(),
|
||||
provider_request_headers: BTreeMap::from([
|
||||
("Authorization".to_string(), "Bearer secret".to_string()),
|
||||
("X-Client".to_string(), "aether".to_string()),
|
||||
("Connection".to_string(), "keep-alive".to_string()),
|
||||
]),
|
||||
provider_request_body: Some(json!({"model": "gpt-5.6-sol"})),
|
||||
provider_request_body_base64: None,
|
||||
content_type: Some("application/json".to_string()),
|
||||
content_encoding: None,
|
||||
request_gzip: None,
|
||||
proxy: None,
|
||||
transport_profile: None,
|
||||
timeouts: None,
|
||||
upstream_is_stream: true,
|
||||
report_kind: None,
|
||||
report_context: None,
|
||||
auth_context: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn identity_normalizes_url_and_hop_by_hop_headers() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
let identity = UpstreamBindingIdentity::from_decision(adapter, &decision()).unwrap();
|
||||
|
||||
assert_eq!(identity.upstream_url, "wss://api.example.test/v1/responses");
|
||||
assert_eq!(
|
||||
identity.handshake_headers,
|
||||
BTreeMap::from([("x-client".to_string(), "aether".to_string())])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn identity_changes_when_physical_binding_changes() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
let base = decision();
|
||||
let identity = UpstreamBindingIdentity::from_decision(adapter, &base).unwrap();
|
||||
|
||||
let codex_adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Codex);
|
||||
assert_ne!(
|
||||
identity,
|
||||
UpstreamBindingIdentity::from_decision(codex_adapter, &base).unwrap()
|
||||
);
|
||||
|
||||
for mutate in [
|
||||
|decision: &mut AiExecutionDecision| {
|
||||
decision.key_id = Some("key-2".to_string());
|
||||
},
|
||||
|decision: &mut AiExecutionDecision| {
|
||||
decision.upstream_url = Some("https://other.example.test/v1/responses".to_string());
|
||||
},
|
||||
|decision: &mut AiExecutionDecision| {
|
||||
decision
|
||||
.provider_request_headers
|
||||
.insert("X-Client".to_string(), "other".to_string());
|
||||
},
|
||||
|decision: &mut AiExecutionDecision| {
|
||||
decision.proxy = Some(aether_contracts::ProxySnapshot {
|
||||
enabled: Some(true),
|
||||
url: Some("http://proxy.example.test:8080".to_string()),
|
||||
..Default::default()
|
||||
});
|
||||
},
|
||||
|decision: &mut AiExecutionDecision| {
|
||||
decision.transport_profile = Some(aether_contracts::ResolvedTransportProfile {
|
||||
profile_id: "chrome136".to_string(),
|
||||
..Default::default()
|
||||
});
|
||||
},
|
||||
] {
|
||||
let mut changed = base.clone();
|
||||
mutate(&mut changed);
|
||||
let changed_identity =
|
||||
UpstreamBindingIdentity::from_decision(adapter, &changed).unwrap();
|
||||
assert_ne!(identity, changed_identity);
|
||||
}
|
||||
|
||||
let mut rotated = base.clone();
|
||||
rotated
|
||||
.provider_request_headers
|
||||
.insert("Authorization".to_string(), "Bearer rotated".to_string());
|
||||
assert_eq!(
|
||||
identity,
|
||||
UpstreamBindingIdentity::from_decision(adapter, &rotated).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stable_key_identity_ignores_custom_auth_value_rotation() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
let mut base = decision();
|
||||
base.auth_header = Some("X-Provider-Token".to_string());
|
||||
base.provider_request_headers.remove("Authorization");
|
||||
base.provider_request_headers.insert(
|
||||
"X-Provider-Token".to_string(),
|
||||
"provider-token-1".to_string(),
|
||||
);
|
||||
let identity = UpstreamBindingIdentity::from_decision(adapter, &base).unwrap();
|
||||
assert!(!identity.handshake_headers.contains_key("x-provider-token"));
|
||||
assert!(identity.auth_fingerprint.is_none());
|
||||
|
||||
let mut rotated = base;
|
||||
rotated.provider_request_headers.insert(
|
||||
"X-Provider-Token".to_string(),
|
||||
"provider-token-2".to_string(),
|
||||
);
|
||||
assert_eq!(
|
||||
identity,
|
||||
UpstreamBindingIdentity::from_decision(adapter, &rotated).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_key_identity_fingerprints_authentication_values() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
let mut first = decision();
|
||||
first.key_id = None;
|
||||
let first_identity = UpstreamBindingIdentity::from_decision(adapter, &first).unwrap();
|
||||
assert!(first_identity.auth_fingerprint.is_some());
|
||||
|
||||
let mut same_account_rotation = first.clone();
|
||||
same_account_rotation.provider_request_headers.insert(
|
||||
"Authorization".to_string(),
|
||||
"Bearer different-account-or-token".to_string(),
|
||||
);
|
||||
let changed_identity =
|
||||
UpstreamBindingIdentity::from_decision(adapter, &same_account_rotation).unwrap();
|
||||
assert_ne!(first_identity, changed_identity);
|
||||
|
||||
let mut non_auth_change = first;
|
||||
non_auth_change
|
||||
.provider_request_headers
|
||||
.insert("X-Client".to_string(), "other-client".to_string());
|
||||
assert_ne!(
|
||||
first_identity,
|
||||
UpstreamBindingIdentity::from_decision(adapter, &non_auth_change).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn disabled_proxy_is_equivalent_to_direct_transport() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
let direct = decision();
|
||||
let direct_identity = UpstreamBindingIdentity::from_decision(adapter, &direct).unwrap();
|
||||
let mut explicitly_disabled = direct;
|
||||
explicitly_disabled.proxy = Some(aether_contracts::ProxySnapshot {
|
||||
enabled: Some(false),
|
||||
url: Some("http://ignored.example.test:8080".to_string()),
|
||||
..Default::default()
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
direct_identity,
|
||||
UpstreamBindingIdentity::from_decision(adapter, &explicitly_disabled).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn identity_rejects_missing_or_invalid_connection_fields() {
|
||||
let adapter = resolve_responses_websocket_adapter(ResponsesWebSocketAdapter::Standard);
|
||||
let mut missing = decision();
|
||||
missing.upstream_url = None;
|
||||
assert_eq!(
|
||||
UpstreamBindingIdentity::from_decision(adapter, &missing),
|
||||
Err(UpstreamBindingIdentityError::MissingUpstreamUrl)
|
||||
);
|
||||
|
||||
let mut invalid = decision();
|
||||
invalid.upstream_url = Some("file:///tmp/responses".to_string());
|
||||
assert_eq!(
|
||||
UpstreamBindingIdentity::from_decision(adapter, &invalid),
|
||||
Err(UpstreamBindingIdentityError::InvalidUpstreamUrl)
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,770 @@
|
||||
//! Client-side Responses WebSocket event forwarding and follow-up planning.
|
||||
|
||||
use axum::body::Bytes;
|
||||
use axum::extract::ws::{Message as AxumWsMessage, WebSocket};
|
||||
use futures_util::SinkExt;
|
||||
use serde_json::Value;
|
||||
use uuid::Uuid;
|
||||
use wreq::ws::message::Message as WreqWsMessage;
|
||||
|
||||
use super::adapter::{resolve_responses_websocket_adapter, ResponsesWebSocketDrainDirective};
|
||||
use super::lifecycle::{
|
||||
await_pending_turn_finalization, queue_turn_finalization,
|
||||
send_responses_websocket_turn_start_error, ActiveResponsesWebSocketTurn,
|
||||
};
|
||||
use super::quota::{mark_active_response_retry_unsafe, send_previous_response_not_found};
|
||||
use super::request::{
|
||||
build_planning_parts, changed_followup_response_create_model,
|
||||
continuation_requires_same_upstream, normalize_followup_response_create,
|
||||
planned_response_create_event, provider_model_from_decision,
|
||||
response_create_has_previous_response_id, response_create_model_or_current,
|
||||
};
|
||||
use super::state::{ActiveResponsesWebSocketRequest, BoundResponsesConnection};
|
||||
use super::turn::{
|
||||
begin_responses_websocket_turn, prepare_responses_websocket_turn_decision,
|
||||
ResponsesWebSocketTurnObservation, ResponsesWebSocketTurnOutcome,
|
||||
};
|
||||
use super::upstream::{bind_responses_upstream, decision_reuses_bound_upstream};
|
||||
use crate::ai_serving::maybe_build_responses_websocket_decision;
|
||||
use crate::clock::current_unix_secs;
|
||||
use crate::control::{request_model_local_rejection, GatewayControlDecision};
|
||||
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
|
||||
use crate::handlers::proxy::websocket::session::{CLOSE_INTERNAL_ERROR, WEBSOCKET_LOG_TRANSPORT};
|
||||
use crate::handlers::proxy::websocket::transport::{
|
||||
client_close_to_upstream, close_client_socket, close_upstream_socket, send_client_message,
|
||||
send_gateway_error, send_gateway_error_with_status, send_upstream_message,
|
||||
};
|
||||
use crate::orchestration::release_pool_key_lease_from_report_context;
|
||||
use crate::rate_limit::FrontdoorUserRpmOutcome;
|
||||
use crate::AppState;
|
||||
|
||||
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
|
||||
|
||||
macro_rules! debug {
|
||||
($($arg:tt)*) => {
|
||||
tracing::debug!(target: LOG_TARGET, $($arg)*)
|
||||
};
|
||||
}
|
||||
|
||||
macro_rules! warn {
|
||||
($($arg:tt)*) => {
|
||||
tracing::warn!(target: LOG_TARGET, $($arg)*)
|
||||
};
|
||||
}
|
||||
|
||||
pub(super) enum RelayDisposition {
|
||||
Continue,
|
||||
Close,
|
||||
UpstreamError(&'static str),
|
||||
}
|
||||
|
||||
pub(super) fn adapter_drain_ready(
|
||||
pending_adapter_drain: Option<ResponsesWebSocketDrainDirective>,
|
||||
response_in_flight: bool,
|
||||
observation: Option<ResponsesWebSocketTurnObservation>,
|
||||
upstream_closed: bool,
|
||||
) -> bool {
|
||||
pending_adapter_drain.is_some()
|
||||
&& (upstream_closed
|
||||
|| !response_in_flight
|
||||
|| matches!(
|
||||
observation,
|
||||
Some(ResponsesWebSocketTurnObservation::Terminal(_))
|
||||
))
|
||||
}
|
||||
|
||||
pub(super) async fn forward_client_message(
|
||||
client_message: AxumWsMessage,
|
||||
bound: &mut BoundResponsesConnection,
|
||||
client_socket: &mut WebSocket,
|
||||
state: &AppState,
|
||||
context: &WebSocketRequestContext,
|
||||
) -> RelayDisposition {
|
||||
match client_message {
|
||||
AxumWsMessage::Text(text) => {
|
||||
let text = text.to_string();
|
||||
let client_event = serde_json::from_str::<Value>(&text).ok();
|
||||
let is_response_create = client_event
|
||||
.as_ref()
|
||||
.and_then(|event| event.get("type"))
|
||||
.and_then(Value::as_str)
|
||||
== Some("response.create");
|
||||
if !is_response_create {
|
||||
if bound.upstream.is_none() {
|
||||
send_gateway_error(
|
||||
client_socket,
|
||||
"responses_websocket_upstream_rebind_required",
|
||||
"Send a new response.create to select another Provider connection",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
// We cannot reconstruct arbitrary Responses control events on
|
||||
// a replacement socket. A concurrent quota error must be
|
||||
// surfaced rather than replaying only the response.create.
|
||||
mark_active_response_retry_unsafe(bound, "client_control_event");
|
||||
return send_upstream_message(
|
||||
bound
|
||||
.upstream
|
||||
.as_mut()
|
||||
.expect("upstream presence was checked above"),
|
||||
WreqWsMessage::text(text),
|
||||
)
|
||||
.await
|
||||
.map(|()| RelayDisposition::Continue)
|
||||
.unwrap_or(RelayDisposition::UpstreamError(
|
||||
"responses_websocket_send_failed",
|
||||
));
|
||||
}
|
||||
|
||||
if bound.response_in_flight {
|
||||
send_gateway_error(
|
||||
client_socket,
|
||||
"response_already_in_progress",
|
||||
"This connection runs one response at a time",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
|
||||
// A prior terminal turn may still be writing usage/audit and
|
||||
// projecting provider effects. Do not let a new independent turn
|
||||
// plan against stale health, adaptive, or pool state.
|
||||
await_pending_turn_finalization(bound).await;
|
||||
|
||||
match consume_response_create_rate_limit(state, &context.decision, context.rpm_bypassed)
|
||||
.await
|
||||
{
|
||||
Ok(true) => {}
|
||||
Ok(false) => {
|
||||
send_gateway_error_with_status(
|
||||
client_socket,
|
||||
429,
|
||||
"rate_limit_exceeded",
|
||||
"Too many response.create events; retry later",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
Err(()) => {
|
||||
send_gateway_error_with_status(
|
||||
client_socket,
|
||||
503,
|
||||
"gateway_rate_limit_unavailable",
|
||||
"Gateway could not evaluate the response rate limit",
|
||||
)
|
||||
.await;
|
||||
close_client_socket(
|
||||
client_socket,
|
||||
CLOSE_INTERNAL_ERROR,
|
||||
"rate_limit_unavailable",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Close;
|
||||
}
|
||||
}
|
||||
|
||||
let Some(client_event) = client_event else {
|
||||
send_gateway_error(
|
||||
client_socket,
|
||||
"invalid_response_create",
|
||||
"response.create must be valid JSON",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
};
|
||||
if bound.upstream.is_none() {
|
||||
if response_create_has_previous_response_id(&client_event) {
|
||||
send_previous_response_not_found(client_socket).await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
let mut client_event = client_event;
|
||||
let requested_model = match response_create_model_or_current(
|
||||
&mut client_event,
|
||||
&bound.client_model,
|
||||
) {
|
||||
Ok(model) => model,
|
||||
Err(code) => {
|
||||
send_gateway_error(
|
||||
client_socket,
|
||||
code,
|
||||
"response.create.model must be a non-empty string",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
};
|
||||
return forward_replanned_response_create(
|
||||
bound,
|
||||
client_socket,
|
||||
state,
|
||||
context,
|
||||
client_event,
|
||||
requested_model,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
let changed_model =
|
||||
match changed_followup_response_create_model(&client_event, &bound.client_model) {
|
||||
Ok(model) => model,
|
||||
Err(code) => {
|
||||
send_gateway_error(
|
||||
client_socket,
|
||||
code,
|
||||
"response.create.model must be a non-empty string",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
};
|
||||
if let Some(requested_model) = changed_model {
|
||||
return forward_replanned_response_create(
|
||||
bound,
|
||||
client_socket,
|
||||
state,
|
||||
context,
|
||||
client_event,
|
||||
requested_model,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
if !response_create_has_previous_response_id(&client_event) {
|
||||
return forward_replanned_response_create(
|
||||
bound,
|
||||
client_socket,
|
||||
state,
|
||||
context,
|
||||
client_event,
|
||||
bound.client_model.clone(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
let outbound = match normalize_followup_response_create(
|
||||
&client_event,
|
||||
&bound.provider_model,
|
||||
&bound.body_normalization,
|
||||
) {
|
||||
Ok(value) => value,
|
||||
Err(code) => {
|
||||
send_gateway_error(
|
||||
client_socket,
|
||||
code,
|
||||
"Gateway could not prepare the response.create event",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
};
|
||||
let provider_event = match serde_json::from_str::<Value>(&outbound) {
|
||||
Ok(event) => event,
|
||||
Err(_) => {
|
||||
send_gateway_error(
|
||||
client_socket,
|
||||
"response_create_serialization_failed",
|
||||
"Gateway could not prepare the response.create event",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
};
|
||||
let turn_index = bound.next_turn_index;
|
||||
let turn_request_id = Uuid::new_v4().to_string();
|
||||
let logical_turn_id = Uuid::new_v4().to_string();
|
||||
debug!(
|
||||
event_name = "responses_websocket_response_create_forwarding",
|
||||
log_type = "event",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
turn_index,
|
||||
client_model = %bound.client_model,
|
||||
provider_model = %bound.provider_model,
|
||||
model_replanned = false,
|
||||
has_previous_response_id = response_create_has_previous_response_id(&client_event),
|
||||
"gateway is forwarding a Responses response.create"
|
||||
);
|
||||
let turn_decision = prepare_responses_websocket_turn_decision(
|
||||
&bound.decision_template,
|
||||
turn_request_id,
|
||||
false,
|
||||
&client_event,
|
||||
&provider_event,
|
||||
&context.trace_id,
|
||||
turn_index,
|
||||
&logical_turn_id,
|
||||
1,
|
||||
);
|
||||
let planning_parts = build_planning_parts(context);
|
||||
let mut turn = match begin_responses_websocket_turn(
|
||||
state,
|
||||
&planning_parts,
|
||||
&context.decision,
|
||||
turn_decision,
|
||||
&client_event,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(turn) => turn,
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_followup_turn_lifecycle_start_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
error = ?error,
|
||||
"gateway could not start Responses WebSocket follow-up usage/audit lifecycle"
|
||||
);
|
||||
send_responses_websocket_turn_start_error(client_socket, &error).await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
};
|
||||
turn.set_provider_response_headers(bound.upstream_response_headers.clone());
|
||||
bound.active_turn = Some(ActiveResponsesWebSocketTurn::new(state, turn));
|
||||
bound.active_response_create = Some(ActiveResponsesWebSocketRequest::new(
|
||||
client_event.clone(),
|
||||
turn_index,
|
||||
logical_turn_id,
|
||||
));
|
||||
bound.next_turn_index = bound.next_turn_index.saturating_add(1);
|
||||
bound.response_in_flight = true;
|
||||
|
||||
let Some(upstream) = bound.upstream.as_mut() else {
|
||||
return RelayDisposition::UpstreamError("responses_websocket_send_failed");
|
||||
};
|
||||
match send_upstream_message(upstream, WreqWsMessage::text(outbound)).await {
|
||||
Ok(()) => {
|
||||
if let Some(turn) = bound.active_turn.as_mut() {
|
||||
turn.mark_upstream_request_sent();
|
||||
}
|
||||
RelayDisposition::Continue
|
||||
}
|
||||
Err(_) => RelayDisposition::UpstreamError("responses_websocket_send_failed"),
|
||||
}
|
||||
}
|
||||
AxumWsMessage::Binary(data) => {
|
||||
if bound.upstream.is_some() {
|
||||
mark_active_response_retry_unsafe(bound, "client_binary_frame");
|
||||
send_upstream_message(
|
||||
bound
|
||||
.upstream
|
||||
.as_mut()
|
||||
.expect("upstream presence was checked above"),
|
||||
WreqWsMessage::Binary(data),
|
||||
)
|
||||
.await
|
||||
.map(|()| RelayDisposition::Continue)
|
||||
.unwrap_or(RelayDisposition::UpstreamError(
|
||||
"responses_websocket_send_failed",
|
||||
))
|
||||
} else {
|
||||
send_gateway_error(
|
||||
client_socket,
|
||||
"responses_websocket_upstream_rebind_required",
|
||||
"Send a new response.create to select another Provider connection",
|
||||
)
|
||||
.await;
|
||||
RelayDisposition::Continue
|
||||
}
|
||||
}
|
||||
AxumWsMessage::Ping(data) => match bound.upstream.as_mut() {
|
||||
Some(upstream) => send_upstream_message(upstream, WreqWsMessage::Ping(data))
|
||||
.await
|
||||
.map(|()| RelayDisposition::Continue)
|
||||
.unwrap_or(RelayDisposition::UpstreamError(
|
||||
"responses_websocket_send_failed",
|
||||
)),
|
||||
None => send_client_message(client_socket, AxumWsMessage::Pong(data))
|
||||
.await
|
||||
.map(|()| RelayDisposition::Continue)
|
||||
.unwrap_or(RelayDisposition::Close),
|
||||
},
|
||||
AxumWsMessage::Pong(data) => match bound.upstream.as_mut() {
|
||||
Some(upstream) => send_upstream_message(upstream, WreqWsMessage::Pong(data))
|
||||
.await
|
||||
.map(|()| RelayDisposition::Continue)
|
||||
.unwrap_or(RelayDisposition::UpstreamError(
|
||||
"responses_websocket_send_failed",
|
||||
)),
|
||||
None => RelayDisposition::Continue,
|
||||
},
|
||||
AxumWsMessage::Close(frame) => {
|
||||
if let Some(upstream) = bound.upstream.as_mut() {
|
||||
close_upstream_socket(upstream, client_close_to_upstream(frame)).await;
|
||||
}
|
||||
RelayDisposition::Close
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn forward_replanned_response_create(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
client_socket: &mut WebSocket,
|
||||
state: &AppState,
|
||||
context: &WebSocketRequestContext,
|
||||
client_event: Value,
|
||||
requested_model: String,
|
||||
) -> RelayDisposition {
|
||||
let planning_parts = build_planning_parts(context);
|
||||
let client_event_text = match serde_json::to_vec(&client_event) {
|
||||
Ok(value) => Bytes::from(value),
|
||||
Err(_) => {
|
||||
send_gateway_error(
|
||||
client_socket,
|
||||
"invalid_response_create",
|
||||
"response.create must be valid JSON",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
};
|
||||
match request_model_local_rejection(
|
||||
state,
|
||||
Some(&context.decision),
|
||||
&planning_parts.uri,
|
||||
&planning_parts.headers,
|
||||
&client_event_text,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(_)) => {
|
||||
send_gateway_error(
|
||||
client_socket,
|
||||
"model_not_allowed",
|
||||
"The requested model is not available to this API key",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
Ok(None) => {}
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_followup_model_access_check_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
requested_model = %requested_model,
|
||||
error = ?error,
|
||||
"gateway failed to evaluate follow-up WebSocket model access policy"
|
||||
);
|
||||
send_gateway_error(
|
||||
client_socket,
|
||||
"gateway_auth_unavailable",
|
||||
"Gateway could not evaluate request access",
|
||||
)
|
||||
.await;
|
||||
close_client_socket(
|
||||
client_socket,
|
||||
CLOSE_INTERNAL_ERROR,
|
||||
"gateway_auth_unavailable",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Close;
|
||||
}
|
||||
}
|
||||
|
||||
let turn_request_id = Uuid::new_v4().to_string();
|
||||
let logical_turn_id = Uuid::new_v4().to_string();
|
||||
let now_unix_secs = current_unix_secs();
|
||||
let excluded_key_ids = bound.exhausted_exclusions.key_ids(now_unix_secs);
|
||||
let excluded_codex_account_ids = bound.exhausted_exclusions.codex_account_ids(now_unix_secs);
|
||||
let excluded_key_ids = (!excluded_key_ids.is_empty()).then_some(&excluded_key_ids);
|
||||
let excluded_codex_account_ids =
|
||||
(!excluded_codex_account_ids.is_empty()).then_some(&excluded_codex_account_ids);
|
||||
let planned = match maybe_build_responses_websocket_decision(
|
||||
state,
|
||||
&planning_parts,
|
||||
&turn_request_id,
|
||||
&context.decision,
|
||||
&client_event,
|
||||
excluded_key_ids,
|
||||
excluded_codex_account_ids,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(decision)) => decision,
|
||||
Ok(None) => {
|
||||
send_gateway_error_with_status(
|
||||
client_socket,
|
||||
503,
|
||||
"responses_provider_unavailable",
|
||||
"No eligible WebSocket-enabled Responses provider is available for the requested model",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_followup_model_planning_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
requested_model = %requested_model,
|
||||
error = ?error,
|
||||
"gateway failed to re-plan Responses WebSocket follow-up model"
|
||||
);
|
||||
send_gateway_error_with_status(
|
||||
client_socket,
|
||||
503,
|
||||
"responses_provider_unavailable",
|
||||
"Gateway could not prepare the requested model",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
};
|
||||
let adapter = resolve_responses_websocket_adapter(planned.adapter);
|
||||
let normalization = planned.normalization;
|
||||
let decision = planned.execution;
|
||||
let reuses_bound_upstream = decision_reuses_bound_upstream(bound, adapter, &decision);
|
||||
if continuation_requires_same_upstream(&client_event, reuses_bound_upstream) {
|
||||
release_pool_key_lease_from_report_context(state, decision.report_context.as_ref()).await;
|
||||
debug!(
|
||||
event_name = "responses_websocket_continuation_rebind_rejected",
|
||||
log_type = "event",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
requested_model = %requested_model,
|
||||
previous_key_id = ?bound.decision_template.key_id,
|
||||
planned_key_id = ?decision.key_id,
|
||||
error_code = "previous_response_not_found",
|
||||
"gateway refused to move a Responses continuation to a different upstream account or connection"
|
||||
);
|
||||
send_previous_response_not_found(client_socket).await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
let provider_event =
|
||||
match planned_response_create_event(&decision, &client_event).and_then(|event| {
|
||||
serde_json::from_str::<Value>(&event)
|
||||
.map_err(|_| "response_create_serialization_failed")
|
||||
}) {
|
||||
Ok(event) => event,
|
||||
Err(code) => {
|
||||
release_pool_key_lease_from_report_context(state, decision.report_context.as_ref())
|
||||
.await;
|
||||
send_gateway_error(
|
||||
client_socket,
|
||||
code,
|
||||
"Gateway could not prepare the requested model",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
};
|
||||
let turn_index = bound.next_turn_index;
|
||||
let turn_decision = prepare_responses_websocket_turn_decision(
|
||||
&decision,
|
||||
turn_request_id,
|
||||
true,
|
||||
&client_event,
|
||||
&provider_event,
|
||||
&context.trace_id,
|
||||
turn_index,
|
||||
&logical_turn_id,
|
||||
1,
|
||||
);
|
||||
let mut turn = match begin_responses_websocket_turn(
|
||||
state,
|
||||
&planning_parts,
|
||||
&context.decision,
|
||||
turn_decision,
|
||||
&client_event,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(turn) => turn,
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_replanned_turn_lifecycle_start_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
requested_model = %requested_model,
|
||||
error = ?error,
|
||||
"gateway could not start re-planned WebSocket usage/audit lifecycle"
|
||||
);
|
||||
send_responses_websocket_turn_start_error(client_socket, &error).await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
};
|
||||
|
||||
if reuses_bound_upstream {
|
||||
let outbound = match serde_json::to_string(&provider_event) {
|
||||
Ok(outbound) => outbound,
|
||||
Err(_) => {
|
||||
queue_turn_finalization(
|
||||
bound,
|
||||
state,
|
||||
turn,
|
||||
ResponsesWebSocketTurnOutcome::upstream_send_failed(),
|
||||
)
|
||||
.await;
|
||||
send_gateway_error(
|
||||
client_socket,
|
||||
"response_create_serialization_failed",
|
||||
"Gateway could not prepare the requested model",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
};
|
||||
let Some(upstream) = bound.upstream.as_mut() else {
|
||||
queue_turn_finalization(
|
||||
bound,
|
||||
state,
|
||||
turn,
|
||||
ResponsesWebSocketTurnOutcome::upstream_send_failed(),
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::UpstreamError("responses_websocket_send_failed");
|
||||
};
|
||||
if send_upstream_message(upstream, WreqWsMessage::text(outbound))
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
queue_turn_finalization(
|
||||
bound,
|
||||
state,
|
||||
turn,
|
||||
ResponsesWebSocketTurnOutcome::upstream_send_failed(),
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::UpstreamError("responses_websocket_send_failed");
|
||||
}
|
||||
|
||||
turn.mark_upstream_request_sent();
|
||||
turn.set_provider_response_headers(bound.upstream_response_headers.clone());
|
||||
let provider_model =
|
||||
provider_model_from_decision(&decision).unwrap_or_else(|| bound.provider_model.clone());
|
||||
let previous_client_model = std::mem::replace(&mut bound.client_model, requested_model);
|
||||
let previous_provider_model = std::mem::replace(&mut bound.provider_model, provider_model);
|
||||
bound.decision_template = decision;
|
||||
// The re-plan keeps this upstream but resolved a new model, so later
|
||||
// continuations must normalize against the new plan, not the old one.
|
||||
bound.body_normalization = normalization;
|
||||
bound.active_turn = Some(ActiveResponsesWebSocketTurn::new(state, turn));
|
||||
bound.active_response_create = Some(ActiveResponsesWebSocketRequest::new(
|
||||
client_event.clone(),
|
||||
turn_index,
|
||||
logical_turn_id.clone(),
|
||||
));
|
||||
bound.next_turn_index = bound.next_turn_index.saturating_add(1);
|
||||
bound.response_in_flight = true;
|
||||
debug!(
|
||||
event_name = "responses_websocket_followup_model_replanned",
|
||||
log_type = "event",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
turn_index,
|
||||
previous_client_model = %previous_client_model,
|
||||
client_model = %bound.client_model,
|
||||
previous_provider_model = %previous_provider_model,
|
||||
provider_model = %bound.provider_model,
|
||||
upstream_rebound = false,
|
||||
model_replanned = true,
|
||||
"gateway re-planned a Responses WebSocket model on the existing upstream"
|
||||
);
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
|
||||
let mut replacement =
|
||||
match bind_responses_upstream(&decision, normalization, &client_event, adapter).await {
|
||||
Ok(connection) => connection,
|
||||
Err(code) => {
|
||||
queue_turn_finalization(
|
||||
bound,
|
||||
state,
|
||||
turn,
|
||||
ResponsesWebSocketTurnOutcome::upstream_connect_failed(code),
|
||||
)
|
||||
.await;
|
||||
warn!(
|
||||
event_name = "responses_websocket_followup_model_rebind_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
requested_model = %requested_model,
|
||||
error_code = code,
|
||||
"gateway failed to rebind Responses WebSocket follow-up model"
|
||||
);
|
||||
send_gateway_error_with_status(
|
||||
client_socket,
|
||||
502,
|
||||
code,
|
||||
"Gateway could not establish the requested model",
|
||||
)
|
||||
.await;
|
||||
return RelayDisposition::Continue;
|
||||
}
|
||||
};
|
||||
|
||||
turn.mark_upstream_request_sent();
|
||||
turn.set_provider_response_headers(replacement.upstream_response_headers.clone());
|
||||
let previous_client_model = bound.client_model.clone();
|
||||
let previous_provider_model = bound.provider_model.clone();
|
||||
let replacement_upstream = replacement
|
||||
.upstream
|
||||
.take()
|
||||
.expect("newly bound Responses upstream should be present");
|
||||
if let Some(mut previous_upstream) = bound.upstream.replace(replacement_upstream) {
|
||||
close_upstream_socket(&mut previous_upstream, None).await;
|
||||
}
|
||||
bound.adapter = replacement.adapter;
|
||||
bound.client_model = replacement.client_model;
|
||||
bound.provider_model = replacement.provider_model;
|
||||
bound.response_in_flight = true;
|
||||
bound.decision_template = replacement.decision_template;
|
||||
bound.body_normalization = replacement.body_normalization;
|
||||
bound.binding_identity = replacement.binding_identity;
|
||||
bound.active_turn = Some(ActiveResponsesWebSocketTurn::new(state, turn));
|
||||
bound.active_response_create = Some(ActiveResponsesWebSocketRequest::new(
|
||||
client_event,
|
||||
turn_index,
|
||||
logical_turn_id,
|
||||
));
|
||||
bound.next_turn_index = bound.next_turn_index.saturating_add(1);
|
||||
bound.upstream_response_headers = replacement.upstream_response_headers;
|
||||
bound.pending_adapter_drain = replacement.pending_adapter_drain;
|
||||
debug!(
|
||||
event_name = "responses_websocket_followup_model_rebound",
|
||||
log_type = "event",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
turn_index,
|
||||
previous_client_model = %previous_client_model,
|
||||
requested_model = %requested_model,
|
||||
previous_provider_model = %previous_provider_model,
|
||||
provider_model = %bound.provider_model,
|
||||
upstream_rebound = true,
|
||||
model_replanned = true,
|
||||
"gateway rebound Responses WebSocket for a follow-up model"
|
||||
);
|
||||
RelayDisposition::Continue
|
||||
}
|
||||
|
||||
pub(super) async fn consume_response_create_rate_limit(
|
||||
state: &AppState,
|
||||
decision: &GatewayControlDecision,
|
||||
rpm_bypassed: bool,
|
||||
) -> Result<bool, ()> {
|
||||
if rpm_bypassed {
|
||||
return Ok(true);
|
||||
}
|
||||
match state
|
||||
.frontdoor_user_rpm()
|
||||
.check_and_consume(state, Some(decision))
|
||||
.await
|
||||
.map_err(|_| ())?
|
||||
{
|
||||
FrontdoorUserRpmOutcome::Rejected(_) => Ok(false),
|
||||
FrontdoorUserRpmOutcome::Allowed | FrontdoorUserRpmOutcome::NotApplicable => Ok(true),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,575 @@
|
||||
//! 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,
|
||||
ActiveResponsesWebSocketTurn,
|
||||
};
|
||||
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, classify_upstream_frame, fatal_relay_policy, FatalRelaySignal,
|
||||
QuotaRelayAction, QuotaRelayFacts, UpstreamFrameAction, UpstreamFrameKind,
|
||||
};
|
||||
use super::state::BoundResponsesConnection;
|
||||
use super::turn::{
|
||||
ResponsesWebSocketTurn, 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";
|
||||
|
||||
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.active_turn.as_ref().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;
|
||||
bound.active_response_create = None;
|
||||
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;
|
||||
bound.active_response_create = None;
|
||||
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;
|
||||
bound.active_response_create = None;
|
||||
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;
|
||||
bound.active_response_create = None;
|
||||
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;
|
||||
bound.active_response_create = None;
|
||||
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;
|
||||
bound.active_response_create = None;
|
||||
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.active_response_create = None;
|
||||
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.active_response_create = None;
|
||||
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.active_turn.is_some(),
|
||||
"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
|
||||
.active_turn
|
||||
.as_mut()
|
||||
.and_then(|turn| turn.observe_upstream_frame(frame, adapter)),
|
||||
None => {
|
||||
if let Some(turn) = bound.active_turn.as_mut() {
|
||||
turn.observe_invalid_upstream_text(text.as_str())
|
||||
}
|
||||
else {
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
_ => None,
|
||||
};
|
||||
update_response_in_flight(bound, parsed_upstream_frame.as_ref());
|
||||
if matches!(
|
||||
observation,
|
||||
Some(ResponsesWebSocketTurnObservation::Started)
|
||||
| Some(ResponsesWebSocketTurnObservation::Terminal(_))
|
||||
) {
|
||||
if let Some(turn) = bound.active_turn.as_mut() {
|
||||
turn.mark_stream_started(state).await;
|
||||
}
|
||||
}
|
||||
let terminal_outcome = match observation {
|
||||
Some(ResponsesWebSocketTurnObservation::Terminal(outcome)) => Some(outcome),
|
||||
_ => None,
|
||||
};
|
||||
if terminal_outcome.is_some() {
|
||||
bound.response_in_flight = false;
|
||||
}
|
||||
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;
|
||||
bound.active_response_create = None;
|
||||
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.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) {
|
||||
let mut retry_turn = bound.active_turn.take().map(ActiveResponsesWebSocketTurn::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;
|
||||
}
|
||||
bound.active_turn = retry_turn.map(|turn| ActiveResponsesWebSocketTurn::new(state, turn));
|
||||
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.active_turn.take().map(ActiveResponsesWebSocketTurn::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;
|
||||
}
|
||||
bound.active_response_create = None;
|
||||
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;
|
||||
bound.active_response_create = None;
|
||||
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(),
|
||||
"gateway could not relay a provider event to the client"
|
||||
);
|
||||
finalize_active_turn(
|
||||
bound,
|
||||
state,
|
||||
ResponsesWebSocketTurnOutcome::client_disconnected(),
|
||||
).await;
|
||||
bound.active_response_create = None;
|
||||
close_bound_upstream(bound).await;
|
||||
break;
|
||||
}
|
||||
if let (Some(turn), Some(frame)) =
|
||||
(bound.active_turn.as_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;
|
||||
bound.active_response_create = None;
|
||||
} 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");
|
||||
bound.active_response_create = None;
|
||||
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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn update_response_in_flight(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
frame: Option<&ParsedResponsesWebSocketFrame<'_>>,
|
||||
) {
|
||||
let Some(frame) = frame else {
|
||||
return;
|
||||
};
|
||||
let frame_kind = if frame.is_terminal() {
|
||||
UpstreamFrameKind::Terminal
|
||||
} else if frame.is_started() {
|
||||
UpstreamFrameKind::Started
|
||||
} else {
|
||||
UpstreamFrameKind::Other
|
||||
};
|
||||
match classify_upstream_frame(frame_kind) {
|
||||
UpstreamFrameAction::Continue if frame.is_started() => {
|
||||
bound.response_in_flight = true;
|
||||
}
|
||||
UpstreamFrameAction::FinalizeTurn => {
|
||||
bound.response_in_flight = false;
|
||||
}
|
||||
UpstreamFrameAction::Continue | UpstreamFrameAction::FinalizeAndClose => {}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,371 @@
|
||||
//! Parsed OpenAI Responses WebSocket text frames.
|
||||
//!
|
||||
//! A relay frame is parsed once and then shared by the protocol adapter, turn
|
||||
//! accounting, retry safety, and connection lifecycle code. Keeping the raw
|
||||
//! text as a borrow avoids copying the websocket payload while the relay is
|
||||
//! processing it.
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(super) struct ResponsesWebSocketFrameTerminal {
|
||||
pub(super) status_code: u16,
|
||||
pub(super) cancelled: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(super) struct ParsedResponsesWebSocketFrame<'a> {
|
||||
raw_text: &'a str,
|
||||
event: Value,
|
||||
event_type: Option<String>,
|
||||
status: Option<u16>,
|
||||
started: bool,
|
||||
terminal: Option<ResponsesWebSocketFrameTerminal>,
|
||||
terminal_event: Option<Value>,
|
||||
chunked: bool,
|
||||
}
|
||||
|
||||
impl<'a> ParsedResponsesWebSocketFrame<'a> {
|
||||
pub(super) fn parse(raw_text: &'a str) -> serde_json::Result<Self> {
|
||||
let event = serde_json::from_str::<Value>(raw_text)?;
|
||||
let events = protocol_events_of(&event);
|
||||
let started = events.iter().copied().any(event_is_started);
|
||||
// A batch carries at most one terminal in practice. Taking the first
|
||||
// in document order keeps the outcome deterministic if that ever
|
||||
// stops being true.
|
||||
let terminal_entry = events
|
||||
.iter()
|
||||
.copied()
|
||||
.find_map(|candidate| terminal_for_event(candidate).map(|term| (candidate, term)));
|
||||
let terminal = terminal_entry.map(|(_, terminal)| terminal);
|
||||
// The terminal event describes the turn's outcome, so it is the one
|
||||
// worth naming in logs and recording as the terminal error body.
|
||||
let event_type = terminal_entry
|
||||
.map(|(candidate, _)| candidate)
|
||||
.or_else(|| events.last().copied())
|
||||
.and_then(event_type_of)
|
||||
.map(str::to_string);
|
||||
let terminal_event = terminal_entry.map(|(candidate, _)| candidate.clone());
|
||||
let chunked = event.get("chunks").and_then(Value::as_array).is_some();
|
||||
let status = terminal.map(|terminal| terminal.status_code);
|
||||
|
||||
Ok(Self {
|
||||
raw_text,
|
||||
event,
|
||||
event_type,
|
||||
status,
|
||||
started,
|
||||
terminal,
|
||||
terminal_event,
|
||||
chunked,
|
||||
})
|
||||
}
|
||||
|
||||
/// The protocol events this frame carries.
|
||||
///
|
||||
/// Codex batches standard `response.*` events into a `{"chunks":[...]}`
|
||||
/// envelope, so one frame can carry several events — and the terminal one
|
||||
/// may be buried inside the batch. Every consumer that interprets event
|
||||
/// semantics must walk this rather than the envelope, or a batched
|
||||
/// `response.completed` goes unnoticed and wedges the turn.
|
||||
pub(super) fn protocol_events(&self) -> Vec<&Value> {
|
||||
protocol_events_of(&self.event)
|
||||
}
|
||||
|
||||
/// The individual event that ended the turn, unwrapped from its batch.
|
||||
pub(super) fn terminal_event(&self) -> Option<&Value> {
|
||||
self.terminal_event.as_ref()
|
||||
}
|
||||
|
||||
pub(super) fn is_chunked(&self) -> bool {
|
||||
self.chunked
|
||||
}
|
||||
|
||||
pub(super) fn raw_text(&self) -> &'a str {
|
||||
self.raw_text
|
||||
}
|
||||
|
||||
pub(super) fn event(&self) -> &Value {
|
||||
&self.event
|
||||
}
|
||||
|
||||
pub(super) fn event_type(&self) -> Option<&str> {
|
||||
self.event_type.as_deref()
|
||||
}
|
||||
|
||||
pub(super) fn status(&self) -> Option<u16> {
|
||||
self.status
|
||||
}
|
||||
|
||||
pub(super) fn is_started(&self) -> bool {
|
||||
self.started
|
||||
}
|
||||
|
||||
pub(super) fn is_terminal(&self) -> bool {
|
||||
self.terminal.is_some()
|
||||
}
|
||||
|
||||
pub(super) fn terminal(&self) -> Option<ResponsesWebSocketFrameTerminal> {
|
||||
self.terminal
|
||||
}
|
||||
|
||||
/// Return a bounded label suitable for structured logs. Event payloads
|
||||
/// are never inserted directly into a log field.
|
||||
pub(super) fn event_type_for_log(&self) -> String {
|
||||
self.event_type
|
||||
.as_deref()
|
||||
.map(safe_websocket_event_label)
|
||||
.unwrap_or_else(|| "invalid_json".to_string())
|
||||
}
|
||||
}
|
||||
|
||||
/// Flattens a frame into the events it carries. An envelope may name its own
|
||||
/// `type` *and* batch further events under `chunks`; both are protocol events.
|
||||
fn protocol_events_of(event: &Value) -> Vec<&Value> {
|
||||
let mut events = Vec::new();
|
||||
if event_type_of(event).is_some() {
|
||||
events.push(event);
|
||||
}
|
||||
if let Some(chunks) = event.get("chunks").and_then(Value::as_array) {
|
||||
events.extend(chunks.iter().filter(|chunk| event_type_of(chunk).is_some()));
|
||||
}
|
||||
// An unrecognized shape is still relayed and still accounted for, so it
|
||||
// must not vanish from the observer's view of the stream.
|
||||
if events.is_empty() {
|
||||
events.push(event);
|
||||
}
|
||||
events
|
||||
}
|
||||
|
||||
fn event_type_of(event: &Value) -> Option<&str> {
|
||||
event.get("type").and_then(Value::as_str)
|
||||
}
|
||||
|
||||
fn event_is_started(event: &Value) -> bool {
|
||||
matches!(
|
||||
event_type_of(event).unwrap_or_default(),
|
||||
"response.created" | "response.in_progress" | "response.queued"
|
||||
)
|
||||
}
|
||||
|
||||
fn terminal_for_event(event: &Value) -> Option<ResponsesWebSocketFrameTerminal> {
|
||||
match event_type_of(event).unwrap_or_default() {
|
||||
"response.completed" => Some(ResponsesWebSocketFrameTerminal {
|
||||
status_code: websocket_event_status_code(event, 200),
|
||||
cancelled: false,
|
||||
}),
|
||||
"response.incomplete" => Some(ResponsesWebSocketFrameTerminal {
|
||||
status_code: websocket_event_status_code(event, 502),
|
||||
cancelled: false,
|
||||
}),
|
||||
"response.cancelled" => Some(ResponsesWebSocketFrameTerminal {
|
||||
status_code: 499,
|
||||
cancelled: true,
|
||||
}),
|
||||
"response.failed" => Some(ResponsesWebSocketFrameTerminal {
|
||||
status_code: websocket_event_status_code(event, 502),
|
||||
cancelled: false,
|
||||
}),
|
||||
"error" => Some(ResponsesWebSocketFrameTerminal {
|
||||
status_code: websocket_event_status_code(event, 502),
|
||||
cancelled: false,
|
||||
}),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn websocket_event_status_code(event: &Value, default: u16) -> u16 {
|
||||
if let Some(status_code) = event
|
||||
.get("status_code")
|
||||
.or_else(|| event.get("status"))
|
||||
.or_else(|| {
|
||||
event
|
||||
.get("response")
|
||||
.and_then(|response| response.get("status_code"))
|
||||
})
|
||||
.and_then(Value::as_u64)
|
||||
.and_then(|value| u16::try_from(value).ok())
|
||||
.filter(|value| *value > 0)
|
||||
{
|
||||
return status_code;
|
||||
}
|
||||
|
||||
let error_code = [
|
||||
event.pointer("/error/type"),
|
||||
event.pointer("/error/code"),
|
||||
event.pointer("/response/error/type"),
|
||||
event.pointer("/response/error/code"),
|
||||
]
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.filter_map(Value::as_str)
|
||||
.map(str::to_ascii_lowercase)
|
||||
.find(|value| !value.trim().is_empty());
|
||||
match error_code.as_deref() {
|
||||
Some(
|
||||
"usage_limit_reached" | "insufficient_quota" | "rate_limit_exceeded" | "quota_exceeded",
|
||||
) => 429,
|
||||
Some("invalid_api_key" | "authentication_error") => 401,
|
||||
Some("invalid_request_error" | "invalid_request" | "model_not_found") => 400,
|
||||
Some("overloaded" | "server_error" | "service_unavailable") => 503,
|
||||
_ => default,
|
||||
}
|
||||
}
|
||||
|
||||
fn safe_websocket_event_label(value: &str) -> String {
|
||||
let value = value.trim();
|
||||
if value.is_empty()
|
||||
|| value.len() > 80
|
||||
|| !value
|
||||
.bytes()
|
||||
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'.' | b'_' | b'-'))
|
||||
{
|
||||
return "unknown".to_string();
|
||||
}
|
||||
value.to_string()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::ParsedResponsesWebSocketFrame;
|
||||
|
||||
#[test]
|
||||
fn parses_started_frame_once_with_raw_text_and_event_metadata() {
|
||||
let raw = r#"{"type":"response.in_progress","response":{"status":200}}"#;
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(raw).expect("valid frame");
|
||||
|
||||
assert_eq!(frame.raw_text(), raw);
|
||||
assert_eq!(frame.event_type(), Some("response.in_progress"));
|
||||
assert_eq!(frame.status(), None);
|
||||
assert!(frame.is_started());
|
||||
assert!(!frame.is_terminal());
|
||||
assert_eq!(frame.event()["response"]["status"], 200);
|
||||
assert_eq!(frame.event_type_for_log(), "response.in_progress");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_terminal_status_and_cancellation() {
|
||||
let completed = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"type":"response.completed","status_code":201}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
assert_eq!(completed.status(), Some(201));
|
||||
assert_eq!(
|
||||
completed
|
||||
.terminal()
|
||||
.map(|terminal| (terminal.status_code, terminal.cancelled)),
|
||||
Some((201, false))
|
||||
);
|
||||
|
||||
let cancelled = ParsedResponsesWebSocketFrame::parse(r#"{"type":"response.cancelled"}"#)
|
||||
.expect("valid frame");
|
||||
assert_eq!(cancelled.status(), Some(499));
|
||||
assert_eq!(
|
||||
cancelled
|
||||
.terminal()
|
||||
.map(|terminal| (terminal.status_code, terminal.cancelled)),
|
||||
Some((499, true))
|
||||
);
|
||||
|
||||
let error = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"type":"error","status_code":429,"error":{"type":"usage_limit_reached"}}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
assert_eq!(error.status(), Some(429));
|
||||
assert!(error.is_terminal());
|
||||
|
||||
let failed = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded"}}}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
assert_eq!(failed.status(), Some(429));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detects_a_terminal_batched_inside_a_chunks_envelope() {
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"chunks":[{"type":"response.output_text.delta","delta":"hi"},{"type":"response.completed","response":{"usage":{"total_tokens":8}}}]}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
|
||||
assert!(frame.is_chunked());
|
||||
assert!(frame.is_terminal());
|
||||
assert_eq!(frame.status(), Some(200));
|
||||
// The label and the recorded error body must name the event that ended
|
||||
// the turn, not the envelope.
|
||||
assert_eq!(frame.event_type(), Some("response.completed"));
|
||||
assert_eq!(
|
||||
frame.terminal_event().and_then(|event| event
|
||||
.pointer("/response/usage/total_tokens")
|
||||
.and_then(serde_json::Value::as_u64)),
|
||||
Some(8)
|
||||
);
|
||||
assert_eq!(frame.protocol_events().len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn detects_a_start_event_batched_inside_a_chunks_envelope() {
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"chunks":[{"type":"codex.rate_limits"},{"type":"response.created"}]}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
|
||||
assert!(frame.is_started());
|
||||
assert!(!frame.is_terminal());
|
||||
assert_eq!(frame.protocol_events().len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn an_envelope_may_carry_its_own_type_alongside_batched_events() {
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"type":"codex.response.metadata","chunks":[{"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded"}}}]}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
|
||||
assert_eq!(frame.protocol_events().len(), 2);
|
||||
assert!(frame.is_terminal());
|
||||
assert_eq!(frame.status(), Some(429));
|
||||
assert_eq!(frame.event_type(), Some("response.failed"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_batch_without_a_terminal_does_not_end_the_turn() {
|
||||
let frame = ParsedResponsesWebSocketFrame::parse(
|
||||
r#"{"chunks":[{"type":"response.output_text.delta","delta":"a"},{"type":"response.output_text.delta","delta":"b"}]}"#,
|
||||
)
|
||||
.expect("valid frame");
|
||||
|
||||
assert!(!frame.is_terminal());
|
||||
assert!(!frame.is_started());
|
||||
assert!(frame.terminal_event().is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn an_unrecognized_shape_is_still_surfaced_as_one_event() {
|
||||
let frame =
|
||||
ParsedResponsesWebSocketFrame::parse(r#"{"unexpected":true}"#).expect("valid frame");
|
||||
|
||||
assert_eq!(frame.protocol_events().len(), 1);
|
||||
assert!(!frame.is_chunked());
|
||||
assert!(!frame.is_terminal());
|
||||
assert_eq!(frame.event_type(), None);
|
||||
assert_eq!(frame.event_type_for_log(), "invalid_json");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preserves_safe_log_label_boundaries() {
|
||||
let unsafe_label =
|
||||
ParsedResponsesWebSocketFrame::parse(r#"{"type":"not safe / contains spaces"}"#)
|
||||
.expect("valid frame");
|
||||
assert_eq!(unsafe_label.event_type_for_log(), "unknown");
|
||||
|
||||
let missing_label =
|
||||
ParsedResponsesWebSocketFrame::parse(r#"{"message":"ok"}"#).expect("valid frame");
|
||||
assert_eq!(missing_label.event_type_for_log(), "invalid_json");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_json() {
|
||||
assert!(ParsedResponsesWebSocketFrame::parse("not-json").is_err());
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,251 @@
|
||||
//! Turn finalization and terminal error mapping for a Responses WebSocket.
|
||||
//!
|
||||
//! A connection can outlive a turn, so persistence and adapter observation
|
||||
//! handles are joined in order before the next turn is planned.
|
||||
|
||||
use std::time::Duration;
|
||||
|
||||
use axum::extract::ws::WebSocket;
|
||||
use tokio::task::JoinHandle;
|
||||
use tokio::time::timeout;
|
||||
|
||||
use super::state::BoundResponsesConnection;
|
||||
use super::turn::{
|
||||
spawn_responses_websocket_turn_finalization, ResponsesWebSocketTurn,
|
||||
ResponsesWebSocketTurnOutcome,
|
||||
};
|
||||
use crate::handlers::proxy::websocket::session::{
|
||||
CLOSE_INTERNAL_ERROR, CLOSE_POLICY_VIOLATION, CLOSE_TRY_AGAIN, WEBSOCKET_LOG_TRANSPORT,
|
||||
};
|
||||
use crate::handlers::proxy::websocket::transport::send_responses_websocket_error;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
const RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT: Duration = Duration::from_secs(5);
|
||||
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
|
||||
|
||||
macro_rules! warn {
|
||||
($($arg:tt)*) => {
|
||||
tracing::warn!(target: LOG_TARGET, $($arg)*)
|
||||
};
|
||||
}
|
||||
|
||||
/// Owns the in-flight turn so that losing the relay task still finalizes it.
|
||||
///
|
||||
/// Every ordinary exit path takes the turn out of here and finalizes it
|
||||
/// explicitly. This guard only covers the paths that are not exit paths at all
|
||||
/// — a panic in the relay loop, or the task being dropped — where the turn
|
||||
/// would otherwise be discarded with its usage row left `Pending`, its
|
||||
/// candidate row left `Streaming`, and its distributed pool key lease leaked
|
||||
/// until the lease expires. Mirrors the HTTP path's `DirectPassthroughFinalizer`.
|
||||
pub(super) struct ActiveResponsesWebSocketTurn {
|
||||
turn: Option<ResponsesWebSocketTurn>,
|
||||
state: AppState,
|
||||
}
|
||||
|
||||
impl ActiveResponsesWebSocketTurn {
|
||||
pub(super) fn new(state: &AppState, turn: ResponsesWebSocketTurn) -> Self {
|
||||
Self {
|
||||
turn: Some(turn),
|
||||
state: state.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Hands the turn back to a caller that will finalize it explicitly.
|
||||
pub(super) fn disarm(mut self) -> ResponsesWebSocketTurn {
|
||||
self.turn
|
||||
.take()
|
||||
.expect("an armed active turn always holds its turn")
|
||||
}
|
||||
}
|
||||
|
||||
impl std::ops::Deref for ActiveResponsesWebSocketTurn {
|
||||
type Target = ResponsesWebSocketTurn;
|
||||
|
||||
fn deref(&self) -> &Self::Target {
|
||||
self.turn
|
||||
.as_ref()
|
||||
.expect("an armed active turn always holds its turn")
|
||||
}
|
||||
}
|
||||
|
||||
impl std::ops::DerefMut for ActiveResponsesWebSocketTurn {
|
||||
fn deref_mut(&mut self) -> &mut Self::Target {
|
||||
self.turn
|
||||
.as_mut()
|
||||
.expect("an armed active turn always holds its turn")
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for ActiveResponsesWebSocketTurn {
|
||||
fn drop(&mut self) {
|
||||
let Some(turn) = self.turn.take() else {
|
||||
return;
|
||||
};
|
||||
let state = self.state.clone();
|
||||
// No runtime means the process is going down; the spawn could not
|
||||
// complete anyway.
|
||||
if let Ok(handle) = tokio::runtime::Handle::try_current() {
|
||||
warn!(
|
||||
event_name = "responses_websocket_turn_abandoned",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
"gateway finalized a Responses WebSocket turn whose relay task went away"
|
||||
);
|
||||
handle.spawn(async move {
|
||||
turn.finalize_detached(
|
||||
&state,
|
||||
ResponsesWebSocketTurnOutcome::relay_task_abandoned(),
|
||||
)
|
||||
.await;
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn finalize_active_turn(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
state: &AppState,
|
||||
outcome: ResponsesWebSocketTurnOutcome,
|
||||
) {
|
||||
if let Some(turn) = bound.active_turn.take() {
|
||||
queue_turn_finalization(bound, state, turn.disarm(), outcome).await;
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn queue_turn_finalization(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
state: &AppState,
|
||||
turn: ResponsesWebSocketTurn,
|
||||
outcome: ResponsesWebSocketTurnOutcome,
|
||||
) {
|
||||
await_pending_adapter_observation(bound).await;
|
||||
await_pending_turn_finalization(bound).await;
|
||||
bound.pending_turn_finalization =
|
||||
Some(spawn_responses_websocket_turn_finalization(state.clone(), turn, outcome).await);
|
||||
}
|
||||
|
||||
pub(super) async fn await_pending_adapter_observation(bound: &mut BoundResponsesConnection) {
|
||||
if let Some(mut handle) = bound.pending_adapter_observation.take() {
|
||||
match timeout(RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT, &mut handle).await {
|
||||
Ok(Err(error)) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_adapter_observation_join_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
error = ?error,
|
||||
"gateway Responses WebSocket adapter observation task failed"
|
||||
);
|
||||
}
|
||||
Ok(Ok(())) => {}
|
||||
Err(_) => {
|
||||
handle.abort();
|
||||
let _ = handle.await;
|
||||
warn!(
|
||||
event_name = "responses_websocket_adapter_observation_timeout",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
timeout_ms = RESPONSES_WEBSOCKET_ADAPTER_OBSERVATION_TIMEOUT.as_millis() as u64,
|
||||
"gateway stopped waiting for a Responses WebSocket adapter observation"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn finalize_unbound_turn(
|
||||
state: AppState,
|
||||
turn: ResponsesWebSocketTurn,
|
||||
outcome: ResponsesWebSocketTurnOutcome,
|
||||
) -> JoinHandle<()> {
|
||||
spawn_responses_websocket_turn_finalization(state, turn, outcome).await
|
||||
}
|
||||
|
||||
pub(super) async fn await_turn_finalization_handle(handle: JoinHandle<()>) {
|
||||
// Do not abort terminal persistence here. Each I/O stage inside the turn
|
||||
// finalizer is independently bounded, and aborting the owner would skip
|
||||
// pool-lease cleanup and leave usage/candidate state non-terminal.
|
||||
match handle.await {
|
||||
Ok(()) => {}
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_turn_finalization_join_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
error = ?error,
|
||||
"gateway Responses WebSocket turn finalizer task failed"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn await_pending_turn_finalization(bound: &mut BoundResponsesConnection) {
|
||||
if let Some(handle) = bound.pending_turn_finalization.take() {
|
||||
await_turn_finalization_handle(handle).await;
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn send_responses_websocket_turn_start_error(
|
||||
client_socket: &mut WebSocket,
|
||||
error: &GatewayError,
|
||||
) {
|
||||
match error {
|
||||
GatewayError::Client { status, message } => {
|
||||
let (error_type, code) = if status.as_u16() == 429 {
|
||||
("rate_limit_error", "gateway_request_capacity_exceeded")
|
||||
} else {
|
||||
("invalid_request_error", "gateway_request_not_allowed")
|
||||
};
|
||||
send_responses_websocket_error(
|
||||
client_socket,
|
||||
status.as_u16(),
|
||||
error_type,
|
||||
code,
|
||||
message,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
GatewayError::AdmissionTimeout { .. } => {
|
||||
send_responses_websocket_error(
|
||||
client_socket,
|
||||
503,
|
||||
"server_error",
|
||||
"gateway_admission_timeout",
|
||||
"Gateway capacity is busy; retry this response",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
GatewayError::LocalExecutionPlanningTimeout { .. } => {
|
||||
send_responses_websocket_error(
|
||||
client_socket,
|
||||
504,
|
||||
"server_error",
|
||||
"gateway_planning_timeout",
|
||||
"Gateway planning timed out; retry this response",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
_ => {
|
||||
send_responses_websocket_error(
|
||||
client_socket,
|
||||
500,
|
||||
"server_error",
|
||||
"responses_websocket_turn_start_failed",
|
||||
"Gateway could not start this response",
|
||||
)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn responses_websocket_turn_start_close(error: &GatewayError) -> (u16, &'static str) {
|
||||
match error {
|
||||
GatewayError::Client { .. } => (CLOSE_POLICY_VIOLATION, "request_not_allowed"),
|
||||
GatewayError::AdmissionTimeout { .. }
|
||||
| GatewayError::LocalExecutionPlanningTimeout { .. } => (CLOSE_TRY_AGAIN, "gateway_busy"),
|
||||
_ => (CLOSE_INTERNAL_ERROR, "turn_start_failed"),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
//! OpenAI Responses WebSocket protocol entry point, session engine, and adapters.
|
||||
//!
|
||||
//! The route is protocol-oriented. `session` bootstraps the authenticated
|
||||
//! connection, `connection` owns the socket FSM, `client` and `quota` own
|
||||
//! protocol/retry policy, and `lifecycle`/`turn` bridge each turn into the
|
||||
//! existing usage and audit runtime. Adapters contain only provider-specific
|
||||
//! connection and metadata behavior.
|
||||
|
||||
mod adapter;
|
||||
mod adapters;
|
||||
mod admission;
|
||||
mod binding;
|
||||
mod client;
|
||||
mod connection;
|
||||
mod frame;
|
||||
mod lifecycle;
|
||||
mod quota;
|
||||
mod relay_policy;
|
||||
mod request;
|
||||
mod session;
|
||||
mod state;
|
||||
mod turn;
|
||||
mod upstream;
|
||||
|
||||
use std::net::SocketAddr;
|
||||
|
||||
use axum::body::Body;
|
||||
use axum::extract::ws::WebSocketUpgrade;
|
||||
use axum::extract::{ConnectInfo, State};
|
||||
use axum::http::{HeaderMap, Response, Uri};
|
||||
|
||||
use crate::handlers::proxy::websocket::ingress::{
|
||||
upgrade_authenticated_ai_websocket, WebSocketIngressSpec,
|
||||
};
|
||||
use crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
pub(crate) async fn responses_websocket(
|
||||
State(state): State<AppState>,
|
||||
ConnectInfo(remote_addr): ConnectInfo<SocketAddr>,
|
||||
ws: WebSocketUpgrade,
|
||||
headers: HeaderMap,
|
||||
uri: Uri,
|
||||
) -> Result<Response<Body>, GatewayError> {
|
||||
upgrade_authenticated_ai_websocket(
|
||||
state,
|
||||
remote_addr,
|
||||
ws,
|
||||
headers,
|
||||
uri,
|
||||
RESPONSES_WEBSOCKET_SESSION_LIMITS,
|
||||
RESPONSES_WEBSOCKET_INGRESS_SPEC,
|
||||
session::run_responses_websocket,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
const RESPONSES_WEBSOCKET_INGRESS_SPEC: WebSocketIngressSpec = WebSocketIngressSpec {
|
||||
route_unavailable_message: "WebSocket route is unavailable",
|
||||
ip_whitelist_failure_event_name: "responses_websocket_ip_whitelist_check_failed",
|
||||
};
|
||||
@@ -0,0 +1,380 @@
|
||||
//! Quota exhaustion, replay safety, and upstream replacement policy.
|
||||
|
||||
use axum::extract::ws::WebSocket;
|
||||
use futures_util::SinkExt;
|
||||
use serde_json::Value;
|
||||
use uuid::Uuid;
|
||||
use wreq::ws::message::Message as WreqWsMessage;
|
||||
|
||||
use super::adapter::{
|
||||
resolve_responses_websocket_adapter, ResponsesWebSocketDrainDirective,
|
||||
ResponsesWebSocketRebindSafety,
|
||||
};
|
||||
use super::lifecycle::{queue_turn_finalization, ActiveResponsesWebSocketTurn};
|
||||
use super::request::{
|
||||
build_planning_parts, planned_response_create_event, response_create_has_previous_response_id,
|
||||
};
|
||||
use super::state::BoundResponsesConnection;
|
||||
use super::turn::{
|
||||
begin_responses_websocket_turn, prepare_responses_websocket_turn_decision,
|
||||
ResponsesWebSocketTurnOutcome,
|
||||
};
|
||||
use super::upstream::{bind_responses_upstream, close_bound_upstream};
|
||||
use crate::ai_serving::maybe_build_responses_websocket_decision;
|
||||
use crate::clock::current_unix_secs;
|
||||
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
|
||||
use crate::handlers::proxy::websocket::session::WEBSOCKET_LOG_TRANSPORT;
|
||||
use crate::handlers::proxy::websocket::transport::{
|
||||
close_upstream_socket, send_responses_websocket_error,
|
||||
};
|
||||
use crate::orchestration::release_pool_key_lease_from_report_context;
|
||||
use crate::AppState;
|
||||
|
||||
const PREVIOUS_RESPONSE_NOT_FOUND_MESSAGE: &str =
|
||||
"Previous response was not found. Retrying the full request.";
|
||||
const LOG_TARGET: &str = "aether_gateway::handlers::proxy::responses_ws";
|
||||
|
||||
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 detach_exhausted_upstream(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
directive: ResponsesWebSocketDrainDirective,
|
||||
trace_id: &str,
|
||||
) {
|
||||
let exclusion = record_exhausted_bound_key(bound, directive.retry_exclusion_until_unix_secs);
|
||||
close_bound_upstream(bound).await;
|
||||
bound.response_in_flight = false;
|
||||
bound.pending_adapter_drain = None;
|
||||
let now_unix_secs = current_unix_secs();
|
||||
debug!(
|
||||
event_name = "responses_websocket_upstream_detached",
|
||||
log_type = "event",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %trace_id,
|
||||
reason = directive.error_code,
|
||||
exhausted_key_id = ?exclusion.as_ref().map(|(key_id, _)| key_id),
|
||||
retry_exclusion_until_unix_secs = ?exclusion.as_ref().map(|(_, until)| until),
|
||||
exhausted_exclusion_count = bound.exhausted_exclusions.len(now_unix_secs),
|
||||
"gateway detached an exhausted Responses WebSocket upstream while preserving the client socket"
|
||||
);
|
||||
}
|
||||
|
||||
pub(super) fn record_exhausted_bound_key(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
reset_at_unix_secs: Option<u64>,
|
||||
) -> Option<(String, u64)> {
|
||||
let key_id = bound
|
||||
.decision_template
|
||||
.key_id
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|key_id| !key_id.is_empty())?
|
||||
.to_string();
|
||||
let provider_account_id = bound
|
||||
.adapter
|
||||
.exhaustion_exclusion_identity(&bound.decision_template)
|
||||
.and_then(|identity| identity.account_id);
|
||||
let exclusion_until = bound.exhausted_exclusions.exclude(
|
||||
key_id.clone(),
|
||||
provider_account_id,
|
||||
reset_at_unix_secs,
|
||||
current_unix_secs(),
|
||||
);
|
||||
Some((key_id, exclusion_until))
|
||||
}
|
||||
|
||||
pub(super) async fn retry_active_turn_after_quota_exhaustion(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
state: &AppState,
|
||||
context: &WebSocketRequestContext,
|
||||
) -> bool {
|
||||
let Some(active) = bound.active_response_create.as_mut() else {
|
||||
return false;
|
||||
};
|
||||
if let Some(reason) = active.quota_retry_block_reason() {
|
||||
debug!(
|
||||
event_name = "responses_websocket_quota_retry_skipped",
|
||||
log_type = "event",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
turn_index = active.turn_index,
|
||||
logical_turn_id = %active.logical_turn_id,
|
||||
turn_attempt = active.turn_attempt,
|
||||
reason,
|
||||
"gateway will not transparently replay an unsafe Responses WebSocket turn"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
active.retry_attempted = true;
|
||||
active.turn_attempt = active.turn_attempt.saturating_add(1);
|
||||
let client_event = active.client_event.clone();
|
||||
let turn_index = active.turn_index;
|
||||
let logical_turn_id = active.logical_turn_id.clone();
|
||||
let turn_attempt = active.turn_attempt;
|
||||
|
||||
let retry_exclusion_until_unix_secs = bound
|
||||
.pending_adapter_drain
|
||||
.and_then(|directive| directive.retry_exclusion_until_unix_secs);
|
||||
let exhausted_key = record_exhausted_bound_key(bound, retry_exclusion_until_unix_secs);
|
||||
let exhausted_key_id = exhausted_key.as_ref().map(|(key_id, _)| key_id.clone());
|
||||
|
||||
let planning_parts = build_planning_parts(context);
|
||||
let turn_request_id = Uuid::new_v4().to_string();
|
||||
let now_unix_secs = current_unix_secs();
|
||||
let excluded_key_ids = bound.exhausted_exclusions.key_ids(now_unix_secs);
|
||||
let excluded_codex_account_ids = bound.exhausted_exclusions.codex_account_ids(now_unix_secs);
|
||||
let excluded_key_ids = (!excluded_key_ids.is_empty()).then_some(&excluded_key_ids);
|
||||
let excluded_codex_account_ids =
|
||||
(!excluded_codex_account_ids.is_empty()).then_some(&excluded_codex_account_ids);
|
||||
let planned = match maybe_build_responses_websocket_decision(
|
||||
state,
|
||||
&planning_parts,
|
||||
&turn_request_id,
|
||||
&context.decision,
|
||||
&client_event,
|
||||
excluded_key_ids,
|
||||
excluded_codex_account_ids,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(Some(decision)) => decision,
|
||||
Ok(None) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_quota_retry_provider_unavailable",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
exhausted_key_id = ?exhausted_key_id,
|
||||
"gateway could not find an alternate Responses WebSocket provider after quota exhaustion"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_quota_retry_planning_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
exhausted_key_id = ?exhausted_key_id,
|
||||
error = ?error,
|
||||
"gateway could not plan an alternate Responses WebSocket provider after quota exhaustion"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
};
|
||||
let adapter = resolve_responses_websocket_adapter(planned.adapter);
|
||||
let normalization = planned.normalization;
|
||||
let decision = planned.execution;
|
||||
if exhausted_key_id.as_deref() == decision.key_id.as_deref() {
|
||||
release_pool_key_lease_from_report_context(state, decision.report_context.as_ref()).await;
|
||||
warn!(
|
||||
event_name = "responses_websocket_quota_retry_selected_exhausted_key",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
key_id = ?decision.key_id,
|
||||
"gateway rejected an alternate Responses WebSocket plan that reused the exhausted key"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
let provider_event = match planned_response_create_event(&decision, &client_event).and_then(
|
||||
|event| {
|
||||
serde_json::from_str::<Value>(&event)
|
||||
.map_err(|_| "response_create_serialization_failed")
|
||||
},
|
||||
) {
|
||||
Ok(event) => event,
|
||||
Err(code) => {
|
||||
release_pool_key_lease_from_report_context(state, decision.report_context.as_ref())
|
||||
.await;
|
||||
warn!(
|
||||
event_name = "responses_websocket_quota_retry_normalization_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
error_code = code,
|
||||
"gateway could not rebuild a Responses response.create for transparent quota retry"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
};
|
||||
let turn_decision = prepare_responses_websocket_turn_decision(
|
||||
&decision,
|
||||
turn_request_id,
|
||||
true,
|
||||
&client_event,
|
||||
&provider_event,
|
||||
&context.trace_id,
|
||||
turn_index,
|
||||
&logical_turn_id,
|
||||
turn_attempt,
|
||||
);
|
||||
let mut turn = match begin_responses_websocket_turn(
|
||||
state,
|
||||
&planning_parts,
|
||||
&context.decision,
|
||||
turn_decision,
|
||||
&client_event,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(turn) => turn,
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = "responses_websocket_quota_retry_reporting_unavailable",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
error = ?error,
|
||||
"gateway could not start usage and audit tracking for transparent quota retry"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
};
|
||||
let mut replacement = match bind_responses_upstream(
|
||||
&decision,
|
||||
normalization,
|
||||
&client_event,
|
||||
adapter,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(connection) => connection,
|
||||
Err(code) => {
|
||||
queue_turn_finalization(
|
||||
bound,
|
||||
state,
|
||||
turn,
|
||||
ResponsesWebSocketTurnOutcome::upstream_connect_failed(code),
|
||||
)
|
||||
.await;
|
||||
warn!(
|
||||
event_name = "responses_websocket_quota_retry_rebind_failed",
|
||||
log_type = "ops",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
error_code = code,
|
||||
"gateway could not bind an alternate Responses WebSocket provider after quota exhaustion"
|
||||
);
|
||||
return false;
|
||||
}
|
||||
};
|
||||
|
||||
turn.mark_upstream_request_sent();
|
||||
turn.set_provider_response_headers(replacement.upstream_response_headers.clone());
|
||||
let replacement_upstream = replacement
|
||||
.upstream
|
||||
.take()
|
||||
.expect("newly bound Responses upstream should be present");
|
||||
if let Some(mut previous_upstream) = bound.upstream.replace(replacement_upstream) {
|
||||
close_upstream_socket(&mut previous_upstream, None).await;
|
||||
}
|
||||
let previous_key_id = bound.decision_template.key_id.clone();
|
||||
bound.adapter = replacement.adapter;
|
||||
bound.client_model = replacement.client_model;
|
||||
bound.provider_model = replacement.provider_model;
|
||||
bound.response_in_flight = true;
|
||||
bound.decision_template = replacement.decision_template;
|
||||
bound.body_normalization = replacement.body_normalization;
|
||||
bound.binding_identity = replacement.binding_identity;
|
||||
bound.active_turn = Some(ActiveResponsesWebSocketTurn::new(state, turn));
|
||||
bound.upstream_response_headers = replacement.upstream_response_headers;
|
||||
bound.pending_adapter_drain = None;
|
||||
debug!(
|
||||
event_name = "responses_websocket_quota_retry_rebound",
|
||||
log_type = "event",
|
||||
transport = WEBSOCKET_LOG_TRANSPORT,
|
||||
websocket = true,
|
||||
trace_id = %context.trace_id,
|
||||
turn_index,
|
||||
logical_turn_id = %logical_turn_id,
|
||||
turn_attempt,
|
||||
previous_key_id = ?previous_key_id,
|
||||
key_id = ?bound.decision_template.key_id,
|
||||
"gateway transparently rebound a Responses WebSocket turn after quota exhaustion"
|
||||
);
|
||||
true
|
||||
}
|
||||
|
||||
pub(super) fn active_continuation_can_retry_from_full_input(
|
||||
bound: &BoundResponsesConnection,
|
||||
) -> bool {
|
||||
bound.active_response_create.as_ref().is_some_and(|active| {
|
||||
response_create_has_previous_response_id(&active.client_event)
|
||||
&& active.retry_unsafe_reason.is_none()
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn is_usage_limit_error_event(event: &Value) -> bool {
|
||||
let is_error = |value: &Value| {
|
||||
value.get("type").and_then(Value::as_str) == Some("error")
|
||||
&& value.pointer("/error/type").and_then(Value::as_str) == Some("usage_limit_reached")
|
||||
};
|
||||
is_error(event)
|
||||
|| event
|
||||
.get("chunks")
|
||||
.and_then(Value::as_array)
|
||||
.is_some_and(|chunks| chunks.iter().any(is_error))
|
||||
}
|
||||
|
||||
pub(super) fn should_request_full_continuation_retry(
|
||||
bound: &BoundResponsesConnection,
|
||||
retry_current_turn: bool,
|
||||
upstream_event: Option<&Value>,
|
||||
) -> bool {
|
||||
retry_current_turn
|
||||
&& active_continuation_can_retry_from_full_input(bound)
|
||||
&& upstream_event.is_some_and(is_usage_limit_error_event)
|
||||
}
|
||||
|
||||
pub(super) async fn send_previous_response_not_found(client_socket: &mut WebSocket) {
|
||||
send_responses_websocket_error(
|
||||
client_socket,
|
||||
400,
|
||||
"invalid_request_error",
|
||||
"previous_response_not_found",
|
||||
PREVIOUS_RESPONSE_NOT_FOUND_MESSAGE,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
pub(super) fn observe_active_response_rebind_safety(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
event: &Value,
|
||||
) {
|
||||
let ResponsesWebSocketRebindSafety::Unsafe { reason } =
|
||||
bound.adapter.rebind_safety_for_upstream_event(event)
|
||||
else {
|
||||
return;
|
||||
};
|
||||
if let Some(active) = bound.active_response_create.as_mut() {
|
||||
active.mark_retry_unsafe(reason);
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn mark_active_response_retry_unsafe(
|
||||
bound: &mut BoundResponsesConnection,
|
||||
reason: &'static str,
|
||||
) {
|
||||
if let Some(active) = bound.active_response_create.as_mut() {
|
||||
active.mark_retry_unsafe(reason);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,298 @@
|
||||
//! Pure relay policy decisions for the Responses WebSocket session.
|
||||
//!
|
||||
//! The session owns sockets, provider planning, and usage persistence. This
|
||||
//! module deliberately owns none of those resources: it only turns observed
|
||||
//! protocol facts into a bounded action. Keeping this layer dependency-free
|
||||
//! makes the failure paths executable with `rustc --test` without linking the
|
||||
//! full gateway (which is useful on constrained CI/diagnostic hosts).
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum FatalRelaySignal {
|
||||
ConnectionAdmissionLost,
|
||||
InvalidUpstreamText,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct FatalRelayPolicy {
|
||||
pub status_code: u16,
|
||||
pub close_code: u16,
|
||||
pub error_code: &'static str,
|
||||
pub client_message: &'static str,
|
||||
pub close_reason: &'static str,
|
||||
}
|
||||
|
||||
/// Map a local relay failure to the status/event/close tuple sent after the
|
||||
/// HTTP upgrade. In particular, capacity loss is retryable (1013), while a
|
||||
/// malformed provider frame is an internal relay error (1011).
|
||||
pub const fn fatal_relay_policy(signal: FatalRelaySignal) -> FatalRelayPolicy {
|
||||
match signal {
|
||||
FatalRelaySignal::ConnectionAdmissionLost => FatalRelayPolicy {
|
||||
status_code: 503,
|
||||
close_code: 1013,
|
||||
error_code: "gateway_connection_admission_lost",
|
||||
client_message: "Gateway capacity lease was lost; reconnect to continue",
|
||||
close_reason: "connection_admission_lost",
|
||||
},
|
||||
FatalRelaySignal::InvalidUpstreamText => FatalRelayPolicy {
|
||||
status_code: 502,
|
||||
close_code: 1011,
|
||||
error_code: "responses_websocket_invalid_upstream_event",
|
||||
client_message: "Provider returned an invalid WebSocket event",
|
||||
close_reason: "invalid_upstream_event",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum UpstreamFrameKind {
|
||||
Other,
|
||||
Started,
|
||||
Terminal,
|
||||
Close,
|
||||
InvalidText,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum UpstreamFrameAction {
|
||||
Continue,
|
||||
FinalizeTurn,
|
||||
FinalizeAndClose,
|
||||
}
|
||||
|
||||
/// Classify the lifecycle effect of one upstream frame. A malformed text
|
||||
/// frame and a non-terminal close both finalize the active turn before the
|
||||
/// client socket is closed; a valid terminal event finalizes the turn but is
|
||||
/// still eligible for the normal downstream forwarding path.
|
||||
pub const fn classify_upstream_frame(kind: UpstreamFrameKind) -> UpstreamFrameAction {
|
||||
match kind {
|
||||
UpstreamFrameKind::Other | UpstreamFrameKind::Started => UpstreamFrameAction::Continue,
|
||||
UpstreamFrameKind::Terminal => UpstreamFrameAction::FinalizeTurn,
|
||||
UpstreamFrameKind::Close | UpstreamFrameKind::InvalidText => {
|
||||
UpstreamFrameAction::FinalizeAndClose
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct QuotaRelayFacts {
|
||||
/// The adapter has observed a definitive quota signal and is ready to
|
||||
/// drain/rebind the current upstream.
|
||||
pub drain_ready: bool,
|
||||
/// The adapter allows a transparent replay of this turn.
|
||||
pub retry_current_turn: bool,
|
||||
/// The session already attempted the adapter-approved transparent replay
|
||||
/// and could not bind an alternate upstream. A continuation may request
|
||||
/// complete input only after that first recovery path was exhausted.
|
||||
pub transparent_retry_failed: bool,
|
||||
/// The event contains the definitive `usage_limit_reached` error. A
|
||||
/// merely exhausted-looking rate-limit snapshot must not trigger retry.
|
||||
pub usage_limit_error: bool,
|
||||
/// The active request is a continuation that can be retried from complete
|
||||
/// input after the old account is detached.
|
||||
pub continuation_retry_eligible: bool,
|
||||
pub upstream_closed: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum QuotaRelayAction {
|
||||
None,
|
||||
AttemptTransparentRetry,
|
||||
RequestFullContinuationRetry,
|
||||
ForwardQuotaAndDetach,
|
||||
}
|
||||
|
||||
/// Decide the quota branch before any response is forwarded to the client.
|
||||
/// `retry_current_turn` intentionally wins over the continuation branch: the
|
||||
/// session attempts the normal transparent retry first, then calls this again
|
||||
/// with `transparent_retry_failed` after that attempt fails. This preserves
|
||||
/// the Codex recovery order while making each fallback explicit.
|
||||
pub const fn classify_quota_relay(facts: QuotaRelayFacts) -> QuotaRelayAction {
|
||||
if !facts.drain_ready {
|
||||
return QuotaRelayAction::None;
|
||||
}
|
||||
if facts.usage_limit_error && facts.retry_current_turn && !facts.transparent_retry_failed {
|
||||
return QuotaRelayAction::AttemptTransparentRetry;
|
||||
}
|
||||
if facts.usage_limit_error
|
||||
&& facts.continuation_retry_eligible
|
||||
&& (facts.transparent_retry_failed || !facts.retry_current_turn)
|
||||
{
|
||||
return QuotaRelayAction::RequestFullContinuationRetry;
|
||||
}
|
||||
if facts.upstream_closed {
|
||||
return QuotaRelayAction::ForwardQuotaAndDetach;
|
||||
}
|
||||
QuotaRelayAction::None
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
struct MockUpstream {
|
||||
frames: &'static [UpstreamFrameKind],
|
||||
cursor: usize,
|
||||
}
|
||||
|
||||
impl MockUpstream {
|
||||
const fn new(frames: &'static [UpstreamFrameKind]) -> Self {
|
||||
Self { frames, cursor: 0 }
|
||||
}
|
||||
|
||||
fn next(&mut self) -> Option<UpstreamFrameKind> {
|
||||
let frame = self.frames.get(self.cursor).copied()?;
|
||||
self.cursor += 1;
|
||||
Some(frame)
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mock_upstream_terminal_event_finalizes_without_waiting_for_an_extra_frame() {
|
||||
let mut upstream = MockUpstream::new(&[
|
||||
UpstreamFrameKind::Started,
|
||||
UpstreamFrameKind::Other,
|
||||
UpstreamFrameKind::Terminal,
|
||||
]);
|
||||
|
||||
assert_eq!(
|
||||
classify_upstream_frame(upstream.next().unwrap()),
|
||||
UpstreamFrameAction::Continue
|
||||
);
|
||||
assert_eq!(
|
||||
classify_upstream_frame(upstream.next().unwrap()),
|
||||
UpstreamFrameAction::Continue
|
||||
);
|
||||
assert_eq!(
|
||||
classify_upstream_frame(upstream.next().unwrap()),
|
||||
UpstreamFrameAction::FinalizeTurn
|
||||
);
|
||||
assert_eq!(upstream.next(), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mock_upstream_quota_429_attempts_one_transparent_retry_then_can_close() {
|
||||
let first = classify_quota_relay(QuotaRelayFacts {
|
||||
drain_ready: true,
|
||||
retry_current_turn: true,
|
||||
transparent_retry_failed: false,
|
||||
usage_limit_error: true,
|
||||
continuation_retry_eligible: false,
|
||||
upstream_closed: false,
|
||||
});
|
||||
assert_eq!(first, QuotaRelayAction::AttemptTransparentRetry);
|
||||
|
||||
// A failed transparent retry must not loop forever. Once the adapter
|
||||
// no longer permits replay, the terminal quota event is forwarded and
|
||||
// the exhausted upstream is detached.
|
||||
let after_retry_failure = classify_quota_relay(QuotaRelayFacts {
|
||||
drain_ready: true,
|
||||
retry_current_turn: false,
|
||||
transparent_retry_failed: true,
|
||||
usage_limit_error: true,
|
||||
continuation_retry_eligible: false,
|
||||
upstream_closed: true,
|
||||
});
|
||||
assert_eq!(after_retry_failure, QuotaRelayAction::ForwardQuotaAndDetach);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn continuation_quota_can_request_full_input_retry_without_replaying_partial_state() {
|
||||
assert_eq!(
|
||||
classify_quota_relay(QuotaRelayFacts {
|
||||
drain_ready: true,
|
||||
retry_current_turn: false,
|
||||
transparent_retry_failed: false,
|
||||
usage_limit_error: true,
|
||||
continuation_retry_eligible: true,
|
||||
upstream_closed: true,
|
||||
}),
|
||||
QuotaRelayAction::RequestFullContinuationRetry
|
||||
);
|
||||
assert_eq!(
|
||||
classify_quota_relay(QuotaRelayFacts {
|
||||
drain_ready: true,
|
||||
retry_current_turn: false,
|
||||
transparent_retry_failed: true,
|
||||
usage_limit_error: true,
|
||||
continuation_retry_eligible: true,
|
||||
upstream_closed: true,
|
||||
}),
|
||||
QuotaRelayAction::RequestFullContinuationRetry
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn continuation_quota_without_transparent_retry_support_uses_full_input_retry() {
|
||||
assert_eq!(
|
||||
classify_quota_relay(QuotaRelayFacts {
|
||||
drain_ready: true,
|
||||
retry_current_turn: false,
|
||||
transparent_retry_failed: false,
|
||||
usage_limit_error: true,
|
||||
continuation_retry_eligible: true,
|
||||
upstream_closed: false,
|
||||
}),
|
||||
QuotaRelayAction::RequestFullContinuationRetry
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn connection_admission_loss_is_retryable_and_invalid_json_is_terminal() {
|
||||
assert_eq!(
|
||||
fatal_relay_policy(FatalRelaySignal::ConnectionAdmissionLost),
|
||||
FatalRelayPolicy {
|
||||
status_code: 503,
|
||||
close_code: 1013,
|
||||
error_code: "gateway_connection_admission_lost",
|
||||
client_message: "Gateway capacity lease was lost; reconnect to continue",
|
||||
close_reason: "connection_admission_lost",
|
||||
}
|
||||
);
|
||||
assert_eq!(
|
||||
fatal_relay_policy(FatalRelaySignal::InvalidUpstreamText),
|
||||
FatalRelayPolicy {
|
||||
status_code: 502,
|
||||
close_code: 1011,
|
||||
error_code: "responses_websocket_invalid_upstream_event",
|
||||
client_message: "Provider returned an invalid WebSocket event",
|
||||
close_reason: "invalid_upstream_event",
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_json_never_maps_to_a_waiting_state() {
|
||||
let mut upstream = MockUpstream::new(&[UpstreamFrameKind::InvalidText]);
|
||||
let action = classify_upstream_frame(upstream.next().unwrap());
|
||||
assert_eq!(action, UpstreamFrameAction::FinalizeAndClose);
|
||||
assert_eq!(upstream.next(), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn quota_snapshot_without_definitive_error_does_not_trigger_retry() {
|
||||
assert_eq!(
|
||||
classify_quota_relay(QuotaRelayFacts {
|
||||
drain_ready: true,
|
||||
retry_current_turn: true,
|
||||
transparent_retry_failed: false,
|
||||
usage_limit_error: false,
|
||||
continuation_retry_eligible: false,
|
||||
upstream_closed: false,
|
||||
}),
|
||||
QuotaRelayAction::None
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
classify_quota_relay(QuotaRelayFacts {
|
||||
drain_ready: true,
|
||||
retry_current_turn: false,
|
||||
transparent_retry_failed: true,
|
||||
usage_limit_error: false,
|
||||
continuation_retry_eligible: false,
|
||||
upstream_closed: true,
|
||||
}),
|
||||
QuotaRelayAction::ForwardQuotaAndDetach
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,352 @@
|
||||
//! Responses WebSocket request normalization and model-selection helpers.
|
||||
//!
|
||||
//! These functions translate client protocol events into the HTTP-shaped
|
||||
//! planning input and provider `response.create` events. They deliberately do
|
||||
//! not depend on connection state or perform I/O.
|
||||
|
||||
use axum::http::header::{AUTHORIZATION, CONNECTION, CONTENT_TYPE, UPGRADE};
|
||||
use axum::http::Method;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
|
||||
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
|
||||
use crate::headers::request_origin_from_headers_and_remote_addr;
|
||||
|
||||
pub(super) fn build_planning_parts(context: &WebSocketRequestContext) -> http::request::Parts {
|
||||
let mut request = http::Request::builder()
|
||||
.method(Method::POST)
|
||||
.uri(context.uri.clone())
|
||||
.body(())
|
||||
.expect("a validated request URI should build planning request parts");
|
||||
let headers = request.headers_mut();
|
||||
*headers = context.headers.clone();
|
||||
headers.remove(AUTHORIZATION);
|
||||
headers.remove("x-api-key");
|
||||
headers.remove("api-key");
|
||||
headers.remove("x-goog-api-key");
|
||||
headers.remove(CONNECTION);
|
||||
headers.remove(UPGRADE);
|
||||
headers.remove("sec-websocket-key");
|
||||
headers.remove("sec-websocket-version");
|
||||
headers.remove("sec-websocket-protocol");
|
||||
headers.remove("sec-websocket-extensions");
|
||||
headers.insert(
|
||||
CONTENT_TYPE,
|
||||
http::HeaderValue::from_static("application/json"),
|
||||
);
|
||||
request
|
||||
.extensions_mut()
|
||||
.insert(request_origin_from_headers_and_remote_addr(
|
||||
&context.headers,
|
||||
&context.remote_addr,
|
||||
));
|
||||
request.into_parts().0
|
||||
}
|
||||
|
||||
pub(super) fn planned_response_create_event(
|
||||
decision: &AiExecutionDecision,
|
||||
fallback: &Value,
|
||||
) -> Result<String, &'static str> {
|
||||
let event = decision
|
||||
.provider_request_body
|
||||
.clone()
|
||||
.unwrap_or_else(|| fallback.clone());
|
||||
finish_response_create_event(event, fallback)
|
||||
}
|
||||
|
||||
/// Restores the WebSocket protocol framing that provider-body normalization is
|
||||
/// not aware of.
|
||||
///
|
||||
/// `previous_response_id` is on the Codex unsupported-field list and `generate`
|
||||
/// is not an HTTP body option at all, so normalization strips both — yet they
|
||||
/// are the entire point of WebSocket mode. They must be re-grafted from the
|
||||
/// client event afterwards. `stream`/`background` go the other way: the
|
||||
/// normalizer inserts `stream`, and the WebSocket protocol has no use for it.
|
||||
fn finish_response_create_event(
|
||||
mut event: Value,
|
||||
client_event: &Value,
|
||||
) -> Result<String, &'static str> {
|
||||
let object = event
|
||||
.as_object_mut()
|
||||
.ok_or("responses_websocket_request_invalid")?;
|
||||
object.insert(
|
||||
"type".to_string(),
|
||||
Value::String("response.create".to_string()),
|
||||
);
|
||||
for field in ["previous_response_id", "generate"] {
|
||||
if let Some(value) = client_event.get(field) {
|
||||
if value.is_null() {
|
||||
object.remove(field);
|
||||
} else {
|
||||
object.insert(field.to_string(), value.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
object.remove("stream");
|
||||
object.remove("background");
|
||||
serde_json::to_string(&event).map_err(|_| "responses_websocket_request_invalid")
|
||||
}
|
||||
|
||||
pub(super) fn response_create_has_previous_response_id(event: &Value) -> bool {
|
||||
event
|
||||
.get("previous_response_id")
|
||||
.is_some_and(|value| !value.is_null())
|
||||
}
|
||||
|
||||
pub(super) fn continuation_requires_same_upstream(
|
||||
event: &Value,
|
||||
reuses_bound_upstream: bool,
|
||||
) -> bool {
|
||||
response_create_has_previous_response_id(event) && !reuses_bound_upstream
|
||||
}
|
||||
|
||||
pub(super) fn changed_followup_response_create_model(
|
||||
event: &Value,
|
||||
current_client_model: &str,
|
||||
) -> Result<Option<String>, &'static str> {
|
||||
let Some(object) = event.as_object() else {
|
||||
return Err("invalid_response_create");
|
||||
};
|
||||
let Some(model) = object.get("model") else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(model) = model
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|model| !model.is_empty())
|
||||
else {
|
||||
return Err("invalid_response_create_model");
|
||||
};
|
||||
if model.eq_ignore_ascii_case(current_client_model) {
|
||||
Ok(None)
|
||||
} else {
|
||||
Ok(Some(model.to_string()))
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn response_create_model_or_current(
|
||||
event: &mut Value,
|
||||
current_client_model: &str,
|
||||
) -> Result<String, &'static str> {
|
||||
let Some(object) = event.as_object_mut() else {
|
||||
return Err("invalid_response_create");
|
||||
};
|
||||
let Some(model) = object.get("model") else {
|
||||
object.insert(
|
||||
"model".to_string(),
|
||||
Value::String(current_client_model.to_string()),
|
||||
);
|
||||
return Ok(current_client_model.to_string());
|
||||
};
|
||||
let Some(model) = model
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|model| !model.is_empty())
|
||||
else {
|
||||
return Err("invalid_response_create_model");
|
||||
};
|
||||
Ok(model.to_string())
|
||||
}
|
||||
|
||||
pub(super) fn provider_model_from_decision(decision: &AiExecutionDecision) -> Option<String> {
|
||||
decision
|
||||
.provider_request_body
|
||||
.as_ref()
|
||||
.and_then(|body| body.get("model"))
|
||||
.and_then(Value::as_str)
|
||||
.or(decision.mapped_model.as_deref())
|
||||
.map(str::trim)
|
||||
.filter(|model| !model.is_empty())
|
||||
.map(str::to_string)
|
||||
}
|
||||
|
||||
/// Prepares a continuation `response.create` for the already-bound upstream.
|
||||
///
|
||||
/// The turn cannot be re-planned without risking a different provider key, so
|
||||
/// the binding's retained normalizer is replayed instead. That keeps model
|
||||
/// directives, endpoint body rules and the Codex body contract applied on every
|
||||
/// turn rather than only on the one that bound the socket.
|
||||
pub(super) fn normalize_followup_response_create(
|
||||
event: &Value,
|
||||
provider_model: &str,
|
||||
normalization: &ResponsesWebSocketBodyNormalization,
|
||||
) -> Result<String, &'static str> {
|
||||
if event.as_object().is_none() {
|
||||
return Err("invalid_response_create");
|
||||
}
|
||||
if event.get("type").and_then(Value::as_str) != Some("response.create") {
|
||||
return Err("invalid_response_create");
|
||||
}
|
||||
// Normalization is best-effort here: a continuation cannot fall back to
|
||||
// another candidate, so a body the contract rejects is still better sent
|
||||
// than dropped.
|
||||
let mut normalized = normalization
|
||||
.normalize_response_create(event)
|
||||
.unwrap_or_else(|| event.clone());
|
||||
let Some(object) = normalized.as_object_mut() else {
|
||||
return Err("invalid_response_create");
|
||||
};
|
||||
// A continuation must never switch models mid-socket, and normalization is
|
||||
// allowed to rewrite `model` (the Codex image-tool path does).
|
||||
object.insert(
|
||||
"model".to_string(),
|
||||
Value::String(provider_model.to_string()),
|
||||
);
|
||||
finish_response_create_event(normalized, event)
|
||||
.map_err(|_| "response_create_serialization_failed")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{normalize_followup_response_create, response_create_has_previous_response_id};
|
||||
use crate::ai_serving::ResponsesWebSocketBodyNormalization;
|
||||
|
||||
fn normalized_continuation(
|
||||
event: &serde_json::Value,
|
||||
normalization: &ResponsesWebSocketBodyNormalization,
|
||||
) -> serde_json::Value {
|
||||
let outbound = normalize_followup_response_create(event, "provider-model", normalization)
|
||||
.expect("continuation should normalize");
|
||||
serde_json::from_str(&outbound).expect("normalized event should be JSON")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn continuation_keeps_protocol_state_that_provider_normalization_strips() {
|
||||
// `previous_response_id` is on the Codex unsupported-field list, so
|
||||
// normalization removes it — yet it is what continues the chain. If
|
||||
// this regresses, every continuation turn silently starts a new one.
|
||||
let event = json!({
|
||||
"type": "response.create",
|
||||
"model": "public-model",
|
||||
"previous_response_id": "resp_123",
|
||||
"input": [],
|
||||
"stream": true,
|
||||
"background": true,
|
||||
});
|
||||
|
||||
let normalized = normalized_continuation(
|
||||
&event,
|
||||
&ResponsesWebSocketBodyNormalization::for_tests("provider-model")
|
||||
.with_provider_type_for_tests("codex"),
|
||||
);
|
||||
|
||||
assert_eq!(normalized["type"], "response.create");
|
||||
assert_eq!(normalized["previous_response_id"], "resp_123");
|
||||
assert_eq!(normalized["model"], "provider-model");
|
||||
assert!(normalized.get("stream").is_none());
|
||||
assert!(normalized.get("background").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn continuation_strips_fields_the_codex_backend_rejects() {
|
||||
// The point of the fix: before it, turns 2..N reached Codex with the
|
||||
// client's raw body, so a `temperature` that turn 1 had stripped would
|
||||
// be rejected upstream. This also proves normalization really runs
|
||||
// rather than silently falling back to the unmodified event.
|
||||
let event = json!({
|
||||
"type": "response.create",
|
||||
"model": "public-model",
|
||||
"previous_response_id": "resp_123",
|
||||
"temperature": 0.7,
|
||||
"top_p": 0.9,
|
||||
"input": [],
|
||||
});
|
||||
|
||||
let normalized = normalized_continuation(
|
||||
&event,
|
||||
&ResponsesWebSocketBodyNormalization::for_tests("provider-model")
|
||||
.with_provider_type_for_tests("codex"),
|
||||
);
|
||||
|
||||
assert!(normalized.get("temperature").is_none());
|
||||
assert!(normalized.get("top_p").is_none());
|
||||
assert_eq!(normalized["store"], false);
|
||||
// ...and the protocol state survives the same pass.
|
||||
assert_eq!(normalized["previous_response_id"], "resp_123");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn continuation_keeps_a_warmup_generate_flag() {
|
||||
let event = json!({
|
||||
"type": "response.create",
|
||||
"model": "public-model",
|
||||
"previous_response_id": "resp_123",
|
||||
"generate": false,
|
||||
"input": [],
|
||||
});
|
||||
|
||||
let normalized = normalized_continuation(
|
||||
&event,
|
||||
&ResponsesWebSocketBodyNormalization::for_tests("provider-model")
|
||||
.with_provider_type_for_tests("codex"),
|
||||
);
|
||||
|
||||
assert_eq!(normalized["generate"], false);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn continuation_applies_the_model_directive_patch_the_binding_turn_received() {
|
||||
let event = json!({
|
||||
"type": "response.create",
|
||||
"model": "public-model",
|
||||
"previous_response_id": "resp_123",
|
||||
"input": [],
|
||||
});
|
||||
|
||||
let normalized = normalized_continuation(
|
||||
&event,
|
||||
&ResponsesWebSocketBodyNormalization::for_tests("provider-model")
|
||||
.with_model_directive_patch_for_tests(json!({"reasoning": {"effort": "high"}})),
|
||||
);
|
||||
|
||||
assert_eq!(normalized["reasoning"]["effort"], "high");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn continuation_still_forces_the_bound_provider_model() {
|
||||
let event = json!({
|
||||
"type": "response.create",
|
||||
"model": "some-other-model",
|
||||
"previous_response_id": "resp_123",
|
||||
"input": [],
|
||||
});
|
||||
|
||||
let normalized = normalized_continuation(
|
||||
&event,
|
||||
&ResponsesWebSocketBodyNormalization::for_tests("provider-model"),
|
||||
);
|
||||
|
||||
assert_eq!(normalized["model"], "provider-model");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_continuation_that_is_not_a_response_create_is_rejected() {
|
||||
let normalization = ResponsesWebSocketBodyNormalization::for_tests("provider-model");
|
||||
|
||||
assert!(normalize_followup_response_create(
|
||||
&json!({"type": "response.cancel"}),
|
||||
"provider-model",
|
||||
&normalization,
|
||||
)
|
||||
.is_err());
|
||||
assert!(normalize_followup_response_create(
|
||||
&json!("not an object"),
|
||||
"provider-model",
|
||||
&normalization,
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn previous_response_id_is_protocol_state_even_when_not_a_string() {
|
||||
assert!(response_create_has_previous_response_id(
|
||||
&json!({"previous_response_id": 42})
|
||||
));
|
||||
assert!(!response_create_has_previous_response_id(
|
||||
&json!({"previous_response_id": null})
|
||||
));
|
||||
assert!(!response_create_has_previous_response_id(&json!({})));
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,141 @@
|
||||
//! Mutable state owned by one Responses WebSocket connection.
|
||||
//!
|
||||
//! The session loop is intentionally kept separate from these containers. A
|
||||
//! connection may survive many `response.create` turns, while the turn
|
||||
//! lifecycle and upstream binding are replaced independently.
|
||||
|
||||
use serde_json::Value;
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use tokio::task::JoinHandle;
|
||||
|
||||
use super::adapter::{ResponsesWebSocketDrainDirective, ResponsesWebSocketProtocolAdapter};
|
||||
use super::binding::UpstreamBindingIdentity;
|
||||
use super::lifecycle::ActiveResponsesWebSocketTurn;
|
||||
use super::request::response_create_has_previous_response_id;
|
||||
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
|
||||
|
||||
const EXHAUSTED_KEY_EXCLUSION_FALLBACK_SECONDS: u64 = 300;
|
||||
|
||||
/// All mutable state associated with the physical upstream connection.
|
||||
pub(super) struct BoundResponsesConnection {
|
||||
pub(super) upstream: Option<wreq::ws::WebSocket>,
|
||||
pub(super) adapter: &'static dyn ResponsesWebSocketProtocolAdapter,
|
||||
pub(super) client_model: String,
|
||||
pub(super) provider_model: String,
|
||||
pub(super) response_in_flight: bool,
|
||||
pub(super) decision_template: AiExecutionDecision,
|
||||
/// Reproduces this binding's provider-body normalization for continuation
|
||||
/// turns, which must not re-enter the planner. Replaced whenever the
|
||||
/// binding or its decision is replaced.
|
||||
pub(super) body_normalization: ResponsesWebSocketBodyNormalization,
|
||||
pub(super) binding_identity: UpstreamBindingIdentity,
|
||||
pub(super) active_turn: Option<ActiveResponsesWebSocketTurn>,
|
||||
pub(super) active_response_create: Option<ActiveResponsesWebSocketRequest>,
|
||||
pub(super) next_turn_index: u64,
|
||||
pub(super) upstream_response_headers: BTreeMap<String, String>,
|
||||
pub(super) pending_adapter_drain: Option<ResponsesWebSocketDrainDirective>,
|
||||
pub(super) pending_adapter_observation: Option<JoinHandle<()>>,
|
||||
pub(super) exhausted_exclusions: ExhaustedResponsesWebSocketExclusions,
|
||||
pub(super) pending_turn_finalization: Option<JoinHandle<()>>,
|
||||
}
|
||||
|
||||
/// Connection-local fallback in addition to the distributed account breaker.
|
||||
/// A key and its provider account are excluded until the upstream's reset
|
||||
/// deadline (or a short fallback when the terminal payload lacks one), so an
|
||||
/// unusually long-lived client socket does not keep it unavailable after the
|
||||
/// quota has recovered.
|
||||
#[derive(Debug, Default)]
|
||||
pub(super) struct ExhaustedResponsesWebSocketExclusions {
|
||||
expires_at_by_key: BTreeMap<String, u64>,
|
||||
expires_at_by_codex_account: BTreeMap<String, u64>,
|
||||
}
|
||||
|
||||
impl ExhaustedResponsesWebSocketExclusions {
|
||||
pub(super) fn exclude(
|
||||
&mut self,
|
||||
key_id: String,
|
||||
codex_account_id: Option<String>,
|
||||
reset_at_unix_secs: Option<u64>,
|
||||
now_unix_secs: u64,
|
||||
) -> u64 {
|
||||
self.prune(now_unix_secs);
|
||||
let requested_expiry = reset_at_unix_secs
|
||||
.filter(|reset_at| *reset_at > now_unix_secs)
|
||||
.unwrap_or_else(|| {
|
||||
now_unix_secs.saturating_add(EXHAUSTED_KEY_EXCLUSION_FALLBACK_SECONDS)
|
||||
});
|
||||
let expiry = self
|
||||
.expires_at_by_key
|
||||
.entry(key_id)
|
||||
.and_modify(|existing| *existing = (*existing).max(requested_expiry))
|
||||
.or_insert(requested_expiry);
|
||||
if let Some(account_id) = codex_account_id {
|
||||
self.expires_at_by_codex_account
|
||||
.entry(account_id)
|
||||
.and_modify(|existing| *existing = (*existing).max(requested_expiry))
|
||||
.or_insert(requested_expiry);
|
||||
}
|
||||
*expiry
|
||||
}
|
||||
|
||||
pub(super) fn codex_account_ids(&mut self, now_unix_secs: u64) -> BTreeSet<String> {
|
||||
self.prune(now_unix_secs);
|
||||
self.expires_at_by_codex_account.keys().cloned().collect()
|
||||
}
|
||||
|
||||
pub(super) fn key_ids(&mut self, now_unix_secs: u64) -> BTreeSet<String> {
|
||||
self.prune(now_unix_secs);
|
||||
self.expires_at_by_key.keys().cloned().collect()
|
||||
}
|
||||
|
||||
pub(super) fn len(&mut self, now_unix_secs: u64) -> usize {
|
||||
self.prune(now_unix_secs);
|
||||
self.expires_at_by_key.len() + self.expires_at_by_codex_account.len()
|
||||
}
|
||||
|
||||
fn prune(&mut self, now_unix_secs: u64) {
|
||||
self.expires_at_by_key
|
||||
.retain(|_, expires_at| *expires_at > now_unix_secs);
|
||||
self.expires_at_by_codex_account
|
||||
.retain(|_, expires_at| *expires_at > now_unix_secs);
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(super) struct ActiveResponsesWebSocketRequest {
|
||||
pub(super) client_event: Value,
|
||||
pub(super) turn_index: u64,
|
||||
pub(super) logical_turn_id: String,
|
||||
pub(super) turn_attempt: u32,
|
||||
pub(super) retry_attempted: bool,
|
||||
pub(super) retry_unsafe_reason: Option<&'static str>,
|
||||
}
|
||||
|
||||
impl ActiveResponsesWebSocketRequest {
|
||||
pub(super) fn new(client_event: Value, turn_index: u64, logical_turn_id: String) -> Self {
|
||||
Self {
|
||||
client_event,
|
||||
turn_index,
|
||||
logical_turn_id,
|
||||
turn_attempt: 1,
|
||||
retry_attempted: false,
|
||||
retry_unsafe_reason: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn quota_retry_block_reason(&self) -> Option<&'static str> {
|
||||
if self.retry_attempted {
|
||||
Some("quota_retry_already_attempted")
|
||||
} else if let Some(reason) = self.retry_unsafe_reason {
|
||||
Some(reason)
|
||||
} else if response_create_has_previous_response_id(&self.client_event) {
|
||||
Some("previous_response_id")
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn mark_retry_unsafe(&mut self, reason: &'static str) {
|
||||
self.retry_unsafe_reason.get_or_insert(reason);
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,113 @@
|
||||
//! Physical upstream WebSocket binding and transport helpers.
|
||||
|
||||
use serde_json::Value;
|
||||
use wreq::ws::message::Message as WreqWsMessage;
|
||||
|
||||
use super::adapter::ResponsesWebSocketProtocolAdapter;
|
||||
use super::binding::{UpstreamBindingIdentity, UpstreamBindingIdentityError};
|
||||
use super::request::planned_response_create_event;
|
||||
use super::state::{BoundResponsesConnection, ExhaustedResponsesWebSocketExclusions};
|
||||
use crate::ai_serving::{AiExecutionDecision, ResponsesWebSocketBodyNormalization};
|
||||
use crate::handlers::proxy::websocket::session::RESPONSES_WEBSOCKET_SESSION_LIMITS;
|
||||
use crate::handlers::proxy::websocket::transport::{
|
||||
close_upstream_socket, connect_upstream_websocket, send_upstream_message,
|
||||
};
|
||||
|
||||
pub(super) async fn bind_responses_upstream(
|
||||
decision: &AiExecutionDecision,
|
||||
normalization: ResponsesWebSocketBodyNormalization,
|
||||
initial_event: &Value,
|
||||
adapter: &'static dyn ResponsesWebSocketProtocolAdapter,
|
||||
) -> Result<BoundResponsesConnection, &'static str> {
|
||||
let binding_identity =
|
||||
UpstreamBindingIdentity::from_decision(adapter, decision).map_err(|error| match error {
|
||||
UpstreamBindingIdentityError::MissingUpstreamUrl => {
|
||||
adapter.upstream_errors().upstream_url_missing
|
||||
}
|
||||
UpstreamBindingIdentityError::InvalidUpstreamUrl => {
|
||||
adapter.upstream_errors().upstream_url_invalid
|
||||
}
|
||||
UpstreamBindingIdentityError::InvalidHandshakeHeaders => {
|
||||
adapter.upstream_errors().headers_invalid
|
||||
}
|
||||
})?;
|
||||
let mut upstream = connect_upstream_websocket(
|
||||
decision,
|
||||
RESPONSES_WEBSOCKET_SESSION_LIMITS,
|
||||
adapter.upstream_errors(),
|
||||
)
|
||||
.await?;
|
||||
let first_event = planned_response_create_event(decision, initial_event)?;
|
||||
send_upstream_message(&mut upstream.socket, WreqWsMessage::text(first_event))
|
||||
.await
|
||||
.map_err(|_| "responses_websocket_initial_send_failed")?;
|
||||
|
||||
let client_model = initial_event
|
||||
.get("model")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.ok_or("responses_websocket_model_missing")?
|
||||
.to_string();
|
||||
let provider_model = decision
|
||||
.provider_request_body
|
||||
.as_ref()
|
||||
.and_then(|body| body.get("model"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.or_else(|| {
|
||||
decision
|
||||
.mapped_model
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
})
|
||||
.ok_or("responses_websocket_mapped_model_missing")?
|
||||
.to_string();
|
||||
|
||||
Ok(BoundResponsesConnection {
|
||||
upstream: Some(upstream.socket),
|
||||
adapter,
|
||||
client_model,
|
||||
provider_model,
|
||||
response_in_flight: true,
|
||||
decision_template: decision.clone(),
|
||||
body_normalization: normalization,
|
||||
binding_identity,
|
||||
active_turn: None,
|
||||
active_response_create: None,
|
||||
next_turn_index: 2,
|
||||
upstream_response_headers: upstream.response_headers,
|
||||
pending_adapter_drain: None,
|
||||
pending_adapter_observation: None,
|
||||
exhausted_exclusions: ExhaustedResponsesWebSocketExclusions::default(),
|
||||
pending_turn_finalization: None,
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) async fn receive_optional_upstream(
|
||||
upstream: &mut Option<wreq::ws::WebSocket>,
|
||||
) -> Option<Result<WreqWsMessage, ()>> {
|
||||
match upstream.as_mut() {
|
||||
Some(upstream) => upstream.recv().await.map(|message| message.map_err(|_| ())),
|
||||
None => std::future::pending().await,
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn close_bound_upstream(bound: &mut BoundResponsesConnection) {
|
||||
if let Some(mut upstream) = bound.upstream.take() {
|
||||
close_upstream_socket(&mut upstream, None).await;
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn decision_reuses_bound_upstream(
|
||||
bound: &BoundResponsesConnection,
|
||||
adapter: &'static dyn ResponsesWebSocketProtocolAdapter,
|
||||
decision: &AiExecutionDecision,
|
||||
) -> bool {
|
||||
bound.upstream.is_some()
|
||||
&& UpstreamBindingIdentity::from_decision(adapter, decision)
|
||||
.map(|identity| bound.binding_identity == identity)
|
||||
.unwrap_or(false)
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
//! Connection-scoped limits and primitives shared by AI WebSocket sessions.
|
||||
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
/// The public Responses WebSocket contract is intentionally bounded so a
|
||||
/// single active socket cannot retain gateway resources indefinitely.
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub(crate) struct WebSocketSessionLimits {
|
||||
pub(crate) max_frame_size: usize,
|
||||
pub(crate) max_message_size: usize,
|
||||
pub(crate) initial_message_timeout: Duration,
|
||||
pub(crate) max_connection_duration: Duration,
|
||||
}
|
||||
|
||||
pub(crate) const RESPONSES_WEBSOCKET_SESSION_LIMITS: WebSocketSessionLimits =
|
||||
WebSocketSessionLimits {
|
||||
max_frame_size: 16 << 20,
|
||||
max_message_size: 16 << 20,
|
||||
initial_message_timeout: Duration::from_secs(60),
|
||||
max_connection_duration: Duration::from_secs(60 * 60),
|
||||
};
|
||||
|
||||
/// A peer that stops draining its receive window must not be able to pin the
|
||||
/// relay loop. Session loops await socket writes inside a `tokio::select!`,
|
||||
/// so an unbounded write also suspends the connection and per-turn deadlines
|
||||
/// that would otherwise reclaim the upstream socket and the shared upstream
|
||||
/// admission permits.
|
||||
pub(crate) const RELAY_WRITE_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
|
||||
/// Frames the gateway emits while tearing a session down are best-effort: the
|
||||
/// session is ending either way, so an unresponsive peer must not delay
|
||||
/// releasing the upstream.
|
||||
pub(crate) const TEARDOWN_WRITE_TIMEOUT: Duration = Duration::from_secs(5);
|
||||
|
||||
pub(crate) const CLOSE_POLICY_VIOLATION: u16 = 1008;
|
||||
pub(crate) const CLOSE_INTERNAL_ERROR: u16 = 1011;
|
||||
pub(crate) const CLOSE_TRY_AGAIN: u16 = 1013;
|
||||
pub(crate) const WEBSOCKET_LOG_TRANSPORT: &str = "websocket";
|
||||
|
||||
/// Waits for an optional per-turn deadline without allocating a timer when no
|
||||
/// turn is active. The protocol adapter retains ownership of the deadline's
|
||||
/// meaning and terminal outcome.
|
||||
pub(crate) async fn wait_for_optional_deadline(deadline: Option<Instant>) {
|
||||
match deadline {
|
||||
Some(deadline) => tokio::time::sleep_until(tokio::time::Instant::from_std(deadline)).await,
|
||||
None => std::future::pending::<()>().await,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,406 @@
|
||||
//! Upstream WebSocket handshake and frame conversion utilities.
|
||||
//!
|
||||
//! These helpers intentionally do not parse messages. A protocol adapter is
|
||||
//! responsible for deciding when and what to send, while this module owns the
|
||||
//! HTTP-to-WebSocket transport conversion and provider transport profile.
|
||||
|
||||
use std::collections::BTreeMap;
|
||||
use std::time::Duration;
|
||||
|
||||
use axum::extract::ws::{CloseFrame as AxumCloseFrame, Message as AxumWsMessage, WebSocket};
|
||||
use axum::http::header::{
|
||||
ACCEPT, ACCEPT_ENCODING, CONNECTION, CONTENT_ENCODING, CONTENT_LENGTH, CONTENT_TYPE, HOST,
|
||||
TRANSFER_ENCODING, UPGRADE,
|
||||
};
|
||||
use axum::http::HeaderMap;
|
||||
use futures_util::{SinkExt, TryFutureExt};
|
||||
use serde_json::json;
|
||||
use url::Url;
|
||||
use wreq::ws::message::{CloseFrame as WreqCloseFrame, Message as WreqWsMessage};
|
||||
|
||||
use crate::ai_serving::AiExecutionDecision;
|
||||
use crate::execution_runtime::transport::{
|
||||
build_browser_wreq_client, build_request_headers, ExecutionTransportControls,
|
||||
};
|
||||
use crate::handlers::proxy::websocket::session::{
|
||||
WebSocketSessionLimits, RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
|
||||
};
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
pub(crate) struct UpstreamWebSocketErrorCodes {
|
||||
pub(crate) upstream_url_missing: &'static str,
|
||||
pub(crate) upstream_url_invalid: &'static str,
|
||||
pub(crate) headers_invalid: &'static str,
|
||||
pub(crate) client_build_failed: &'static str,
|
||||
pub(crate) proxy_invalid: &'static str,
|
||||
pub(crate) tunnel_proxy_unsupported: &'static str,
|
||||
pub(crate) handshake_failed: &'static str,
|
||||
pub(crate) upgrade_rejected: &'static str,
|
||||
pub(crate) upgrade_failed: &'static str,
|
||||
}
|
||||
|
||||
pub(crate) struct UpstreamWebSocketConnection {
|
||||
pub(crate) socket: wreq::ws::WebSocket,
|
||||
pub(crate) response_headers: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
pub(crate) async fn connect_upstream_websocket(
|
||||
decision: &AiExecutionDecision,
|
||||
limits: WebSocketSessionLimits,
|
||||
errors: UpstreamWebSocketErrorCodes,
|
||||
) -> Result<UpstreamWebSocketConnection, &'static str> {
|
||||
let upstream_url = decision
|
||||
.upstream_url
|
||||
.as_deref()
|
||||
.ok_or(errors.upstream_url_missing)?;
|
||||
let upstream_url = websocket_upstream_url(upstream_url, errors.upstream_url_invalid)?;
|
||||
let headers =
|
||||
websocket_handshake_headers(&decision.provider_request_headers, errors.headers_invalid)?;
|
||||
let client = build_websocket_client(decision, errors)?;
|
||||
let response = client
|
||||
.websocket(upstream_url.as_str())
|
||||
.headers(headers)
|
||||
.max_frame_size(limits.max_frame_size)
|
||||
.max_message_size(limits.max_message_size)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|_| errors.handshake_failed)?;
|
||||
if response.status().as_u16() != 101 {
|
||||
return Err(errors.upgrade_rejected);
|
||||
}
|
||||
let response_headers = websocket_response_headers(response.headers());
|
||||
let socket = response
|
||||
.into_websocket()
|
||||
.await
|
||||
.map_err(|_| errors.upgrade_failed)?;
|
||||
Ok(UpstreamWebSocketConnection {
|
||||
socket,
|
||||
response_headers,
|
||||
})
|
||||
}
|
||||
|
||||
fn websocket_response_headers(headers: &HeaderMap) -> BTreeMap<String, String> {
|
||||
headers
|
||||
.iter()
|
||||
.filter_map(|(name, value)| {
|
||||
value
|
||||
.to_str()
|
||||
.ok()
|
||||
.map(|value| (name.as_str().to_string(), value.to_string()))
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(crate) fn websocket_upstream_url(
|
||||
raw: &str,
|
||||
invalid_code: &'static str,
|
||||
) -> Result<Url, &'static str> {
|
||||
let mut url = Url::parse(raw).map_err(|_| invalid_code)?;
|
||||
if url.host_str().is_none() || !url.username().is_empty() || url.password().is_some() {
|
||||
return Err(invalid_code);
|
||||
}
|
||||
let websocket_scheme = match url.scheme() {
|
||||
"https" => "wss",
|
||||
"http" => "ws",
|
||||
"wss" | "ws" => return Ok(url),
|
||||
_ => return Err(invalid_code),
|
||||
};
|
||||
url.set_scheme(websocket_scheme).map_err(|_| invalid_code)?;
|
||||
Ok(url)
|
||||
}
|
||||
|
||||
pub(crate) fn websocket_handshake_headers(
|
||||
provider_headers: &BTreeMap<String, String>,
|
||||
invalid_code: &'static str,
|
||||
) -> Result<HeaderMap, &'static str> {
|
||||
let mut headers =
|
||||
build_request_headers(provider_headers, None, false).map_err(|_| invalid_code)?;
|
||||
for header in [
|
||||
ACCEPT,
|
||||
ACCEPT_ENCODING,
|
||||
CONNECTION,
|
||||
CONTENT_ENCODING,
|
||||
CONTENT_LENGTH,
|
||||
CONTENT_TYPE,
|
||||
HOST,
|
||||
TRANSFER_ENCODING,
|
||||
UPGRADE,
|
||||
] {
|
||||
headers.remove(header);
|
||||
}
|
||||
Ok(headers)
|
||||
}
|
||||
|
||||
fn build_websocket_client(
|
||||
decision: &AiExecutionDecision,
|
||||
errors: UpstreamWebSocketErrorCodes,
|
||||
) -> Result<wreq::Client, &'static str> {
|
||||
let timeouts = websocket_timeouts(decision);
|
||||
if let Some(profile) = decision.transport_profile.as_ref() {
|
||||
return build_browser_wreq_client(
|
||||
timeouts.as_ref(),
|
||||
decision.proxy.as_ref(),
|
||||
profile,
|
||||
ExecutionTransportControls::default(),
|
||||
false,
|
||||
)
|
||||
.map_err(|_| errors.client_build_failed);
|
||||
}
|
||||
|
||||
let mut builder = wreq::Client::builder();
|
||||
if let Some(connect_ms) = timeouts.as_ref().and_then(|timeouts| timeouts.connect_ms) {
|
||||
builder = builder.connect_timeout(Duration::from_millis(connect_ms));
|
||||
}
|
||||
if let Some(proxy) = decision
|
||||
.proxy
|
||||
.as_ref()
|
||||
.filter(|proxy| proxy.enabled != Some(false))
|
||||
{
|
||||
if let Some(proxy_url) = proxy
|
||||
.url
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|url| !url.is_empty())
|
||||
{
|
||||
let proxy = wreq::Proxy::all(proxy_url).map_err(|_| errors.proxy_invalid)?;
|
||||
builder = builder.proxy(proxy);
|
||||
} else if proxy.node_id.is_some() || proxy.mode.as_deref() == Some("tunnel") {
|
||||
return Err(errors.tunnel_proxy_unsupported);
|
||||
}
|
||||
}
|
||||
builder.build().map_err(|_| errors.client_build_failed)
|
||||
}
|
||||
|
||||
pub(crate) fn websocket_timeouts(
|
||||
decision: &AiExecutionDecision,
|
||||
) -> Option<aether_contracts::ExecutionTimeouts> {
|
||||
let mut timeouts = decision.timeouts.clone()?;
|
||||
timeouts.read_ms = None;
|
||||
timeouts.first_byte_ms = None;
|
||||
timeouts.total_ms = None;
|
||||
Some(timeouts)
|
||||
}
|
||||
|
||||
/// Why a frame did not reach its peer. A timeout is reported separately from
|
||||
/// a socket error because the two describe different peers: one has gone away,
|
||||
/// the other is still connected but has stopped reading.
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(crate) enum WebSocketWriteError {
|
||||
Failed,
|
||||
TimedOut,
|
||||
}
|
||||
|
||||
impl WebSocketWriteError {
|
||||
pub(crate) const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Failed => "write_failed",
|
||||
Self::TimedOut => "write_timeout",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Relays one frame to the client under [`RELAY_WRITE_TIMEOUT`].
|
||||
pub(crate) async fn send_client_message(
|
||||
client_socket: &mut WebSocket,
|
||||
message: AxumWsMessage,
|
||||
) -> Result<(), WebSocketWriteError> {
|
||||
bounded_send(
|
||||
RELAY_WRITE_TIMEOUT,
|
||||
client_socket.send(message).map_err(|_| ()),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
/// Sends one frame to the upstream under [`RELAY_WRITE_TIMEOUT`].
|
||||
pub(crate) async fn send_upstream_message(
|
||||
upstream: &mut wreq::ws::WebSocket,
|
||||
message: WreqWsMessage,
|
||||
) -> Result<(), WebSocketWriteError> {
|
||||
bounded_send(RELAY_WRITE_TIMEOUT, upstream.send(message).map_err(|_| ())).await
|
||||
}
|
||||
|
||||
/// Best-effort teardown write. The caller is already ending the session, so
|
||||
/// the outcome only matters for keeping the wait bounded.
|
||||
async fn send_teardown_message<F>(write: F)
|
||||
where
|
||||
F: std::future::Future<Output = Result<(), ()>>,
|
||||
{
|
||||
let _ = bounded_send(TEARDOWN_WRITE_TIMEOUT, write).await;
|
||||
}
|
||||
|
||||
async fn bounded_send<F>(budget: Duration, write: F) -> Result<(), WebSocketWriteError>
|
||||
where
|
||||
F: std::future::Future<Output = Result<(), ()>>,
|
||||
{
|
||||
match tokio::time::timeout(budget, write).await {
|
||||
Ok(Ok(())) => Ok(()),
|
||||
Ok(Err(())) => Err(WebSocketWriteError::Failed),
|
||||
Err(_) => Err(WebSocketWriteError::TimedOut),
|
||||
}
|
||||
}
|
||||
|
||||
/// Sends a WebSocket Close frame upstream without waiting on an unresponsive
|
||||
/// provider. The socket is dropped by the caller either way.
|
||||
pub(crate) async fn close_upstream_socket(
|
||||
upstream: &mut wreq::ws::WebSocket,
|
||||
frame: Option<WreqCloseFrame>,
|
||||
) {
|
||||
send_teardown_message(upstream.send(WreqWsMessage::Close(frame)).map_err(|_| ())).await;
|
||||
}
|
||||
|
||||
pub(crate) fn upstream_message_to_client(message: WreqWsMessage) -> AxumWsMessage {
|
||||
match message {
|
||||
WreqWsMessage::Text(text) => AxumWsMessage::Text(text.to_string().into()),
|
||||
WreqWsMessage::Binary(data) => AxumWsMessage::Binary(data),
|
||||
WreqWsMessage::Ping(data) => AxumWsMessage::Ping(data),
|
||||
WreqWsMessage::Pong(data) => AxumWsMessage::Pong(data),
|
||||
WreqWsMessage::Close(frame) => AxumWsMessage::Close(frame.map(|frame| AxumCloseFrame {
|
||||
code: frame.code.into(),
|
||||
reason: frame.reason.to_string().into(),
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn client_close_to_upstream(frame: Option<AxumCloseFrame>) -> Option<WreqCloseFrame> {
|
||||
frame.map(|frame| WreqCloseFrame {
|
||||
code: frame.code.into(),
|
||||
reason: frame.reason.to_string().into(),
|
||||
})
|
||||
}
|
||||
|
||||
/// Builds a Responses WebSocket error event in the shape understood by the
|
||||
/// official client implementations. The status is part of the event body,
|
||||
/// not the WebSocket handshake, because the connection is already upgraded.
|
||||
pub(crate) fn responses_websocket_error_event(
|
||||
status: u16,
|
||||
error_type: &str,
|
||||
code: &str,
|
||||
message: &str,
|
||||
) -> serde_json::Value {
|
||||
json!({
|
||||
"type": "error",
|
||||
"status": status,
|
||||
"error": {
|
||||
"type": error_type,
|
||||
"code": code,
|
||||
"message": message,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) async fn send_responses_websocket_error(
|
||||
client_socket: &mut WebSocket,
|
||||
status: u16,
|
||||
error_type: &str,
|
||||
code: &str,
|
||||
message: &str,
|
||||
) {
|
||||
let event = responses_websocket_error_event(status, error_type, code, message);
|
||||
send_teardown_message(
|
||||
client_socket
|
||||
.send(AxumWsMessage::Text(event.to_string().into()))
|
||||
.map_err(|_| ()),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
pub(crate) async fn send_gateway_error(client_socket: &mut WebSocket, code: &str, message: &str) {
|
||||
send_gateway_error_with_status(client_socket, 400, code, message).await;
|
||||
}
|
||||
|
||||
pub(crate) async fn send_gateway_error_with_status(
|
||||
client_socket: &mut WebSocket,
|
||||
status: u16,
|
||||
code: &str,
|
||||
message: &str,
|
||||
) {
|
||||
send_responses_websocket_error(client_socket, status, "gateway_error", code, message).await;
|
||||
}
|
||||
|
||||
pub(crate) async fn close_client_socket(client_socket: &mut WebSocket, code: u16, reason: &str) {
|
||||
send_teardown_message(
|
||||
client_socket
|
||||
.send(AxumWsMessage::Close(Some(AxumCloseFrame {
|
||||
code,
|
||||
reason: reason.to_string().into(),
|
||||
})))
|
||||
.map_err(|_| ()),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
bounded_send, responses_websocket_error_event, websocket_upstream_url, WebSocketWriteError,
|
||||
RELAY_WRITE_TIMEOUT, TEARDOWN_WRITE_TIMEOUT,
|
||||
};
|
||||
use std::time::Duration;
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_peer_that_never_drains_its_window_times_out_instead_of_pinning_the_relay() {
|
||||
let stalled = std::future::pending::<Result<(), ()>>();
|
||||
|
||||
let outcome = bounded_send(Duration::from_millis(1), stalled).await;
|
||||
|
||||
assert_eq!(outcome, Err(WebSocketWriteError::TimedOut));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_socket_error_is_reported_separately_from_a_stalled_peer() {
|
||||
let outcome = bounded_send(RELAY_WRITE_TIMEOUT, std::future::ready(Err(()))).await;
|
||||
|
||||
assert_eq!(outcome, Err(WebSocketWriteError::Failed));
|
||||
assert_eq!(WebSocketWriteError::Failed.as_str(), "write_failed");
|
||||
assert_eq!(WebSocketWriteError::TimedOut.as_str(), "write_timeout");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_write_that_completes_within_its_budget_succeeds() {
|
||||
let outcome = bounded_send(RELAY_WRITE_TIMEOUT, std::future::ready(Ok::<(), ()>(()))).await;
|
||||
|
||||
assert_eq!(outcome, Ok(()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn teardown_writes_are_given_a_shorter_budget_than_relayed_frames() {
|
||||
assert!(TEARDOWN_WRITE_TIMEOUT < RELAY_WRITE_TIMEOUT);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_a_client_compatible_responses_error_event() {
|
||||
let event = responses_websocket_error_event(
|
||||
400,
|
||||
"invalid_request_error",
|
||||
"previous_response_not_found",
|
||||
"Previous response was not found.",
|
||||
);
|
||||
|
||||
assert_eq!(event["type"], "error");
|
||||
assert_eq!(event["status"], 400);
|
||||
assert_eq!(event["error"]["type"], "invalid_request_error");
|
||||
assert_eq!(event["error"]["code"], "previous_response_not_found");
|
||||
assert_eq!(
|
||||
event["error"]["message"],
|
||||
"Previous response was not found."
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn maps_http_url_to_websocket_url_without_losing_path_or_query() {
|
||||
let url = websocket_upstream_url(
|
||||
"https://example.test/backend-api/codex/responses?x=1",
|
||||
"invalid",
|
||||
)
|
||||
.expect("URL should be converted");
|
||||
assert_eq!(
|
||||
url.as_str(),
|
||||
"wss://example.test/backend-api/codex/responses?x=1"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_upstream_url_with_credentials() {
|
||||
assert!(websocket_upstream_url("https://[email protected]/responses", "invalid").is_err());
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user