mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-05 00:47:48 +08:00
522 lines
20 KiB
Rust
522 lines
20 KiB
Rust
//! 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;
|
||
use crate::privacy::RedactionSessionSlot;
|
||
|
||
/// Model identifiers are copied into planner diagnostics. Bound them before
|
||
/// any planning/logging so a single 16 MiB WebSocket frame cannot amplify into
|
||
/// repeated multi-megabyte log records.
|
||
pub(super) const MAX_RESPONSES_WEBSOCKET_MODEL_BYTES: usize = 256;
|
||
|
||
pub(super) fn validated_response_create_model(value: &Value) -> Result<&str, &'static str> {
|
||
let Some(model) = value
|
||
.as_str()
|
||
.map(str::trim)
|
||
.filter(|model| !model.is_empty())
|
||
else {
|
||
return Err("invalid_response_create_model");
|
||
};
|
||
if model.len() > MAX_RESPONSES_WEBSOCKET_MODEL_BYTES {
|
||
return Err("invalid_response_create_model");
|
||
}
|
||
Ok(model)
|
||
}
|
||
|
||
/// 把一条 WebSocket turn 还原成 planner 需要的 HTTP 形状请求头部。
|
||
///
|
||
/// 这里必须和 HTTP 前门(`handlers/proxy/mod.rs`)保持同一份 extension 契约:
|
||
/// planner 只在 `parts.extensions` 里拿到 `RedactionSessionSlot` 时才做请求脱敏
|
||
/// (`ai_serving/planner/redaction.rs`),少插这一项等于整条 WS 链路静默绕过
|
||
/// 已启用的 PII 脱敏。
|
||
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,
|
||
));
|
||
// slot 必须每个 turn 新建,不能按连接复用:planner 侧的请求脱敏缓存键是
|
||
// `{format:?}:{body_json 指针地址}`(`ai_serving/planner/redaction.rs:169`),
|
||
// 连接级复用同一个 slot 时,上一轮 client_event 释放后这一轮的 `Value` 很可能
|
||
// 落在同一地址,会命中上一轮缓存,把上一轮的脱敏 body 当成这一轮的发出去。
|
||
// 每个 `response.create` 本身就是独立计费/审计请求,per-turn 也正好对应
|
||
// HTTP 前门「一个请求一个 slot」的语义。
|
||
request
|
||
.extensions_mut()
|
||
.insert(RedactionSessionSlot::default());
|
||
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, Codex HTTP
|
||
/// normalization may force `store`, and `generate` is not an HTTP body option
|
||
/// at all. Those fields are WebSocket protocol state, so an explicitly supplied
|
||
/// value (including `null`) must be re-grafted verbatim from the client event.
|
||
/// `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 ["store", "previous_response_id", "generate"] {
|
||
if let Some(value) = client_event.get(field) {
|
||
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 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 model = validated_response_create_model(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 model = validated_response_create_model(model)?;
|
||
let model = model.to_string();
|
||
object.insert("model".to_string(), Value::String(model.clone()));
|
||
Ok(model)
|
||
}
|
||
|
||
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 std::net::SocketAddr;
|
||
|
||
use axum::http::{HeaderMap, Uri};
|
||
use serde_json::{json, Value};
|
||
|
||
use super::{
|
||
build_planning_parts, normalize_followup_response_create,
|
||
response_create_has_previous_response_id,
|
||
};
|
||
use crate::ai_serving::ResponsesWebSocketBodyNormalization;
|
||
use crate::control::GatewayControlDecision;
|
||
use crate::handlers::proxy::websocket::ingress::WebSocketRequestContext;
|
||
use crate::privacy::RedactionSessionSlot;
|
||
|
||
fn websocket_context() -> WebSocketRequestContext {
|
||
WebSocketRequestContext {
|
||
trace_id: "trace-planning-parts".to_string(),
|
||
headers: HeaderMap::new(),
|
||
uri: Uri::from_static("/v1/responses"),
|
||
remote_addr: "127.0.0.1:65001"
|
||
.parse::<SocketAddr>()
|
||
.expect("remote address should parse"),
|
||
client_ip: "127.0.0.1".parse().expect("client IP should parse"),
|
||
decision: GatewayControlDecision::synthetic(
|
||
"/v1/responses".to_string(),
|
||
Some("ai_public".to_string()),
|
||
Some("openai".to_string()),
|
||
Some("responses_websocket".to_string()),
|
||
Some("openai:responses".to_string()),
|
||
),
|
||
websocket_connection_permit: None,
|
||
}
|
||
}
|
||
|
||
#[test]
|
||
fn planning_parts_carry_a_fresh_redaction_session_slot_per_turn() {
|
||
// 没有这个 extension,planner 会静默跳过已启用的 PII 脱敏
|
||
// (ai_serving/planner/redaction.rs),整条 WS 链路都按原文发上游。
|
||
let context = websocket_context();
|
||
let first = build_planning_parts(&context);
|
||
let second = build_planning_parts(&context);
|
||
|
||
let first_slot = first
|
||
.extensions
|
||
.get::<RedactionSessionSlot>()
|
||
.expect("planning parts must carry a redaction session slot");
|
||
let second_slot = second
|
||
.extensions
|
||
.get::<RedactionSessionSlot>()
|
||
.expect("planning parts must carry a redaction session slot");
|
||
|
||
// 每轮必须是独立 slot:slot 内的请求缓存以 body 指针地址为键,跨轮共享会
|
||
// 命中上一轮缓存。用缓存条目相互不可见来证明两者不是同一个 slot。
|
||
first_slot.put_cached_request_redaction(
|
||
"turn-1",
|
||
crate::privacy::CachedRequestRedaction::unredacted(),
|
||
);
|
||
assert!(first_slot.cached_request_redaction("turn-1").is_some());
|
||
assert!(second_slot.cached_request_redaction("turn-1").is_none());
|
||
}
|
||
|
||
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 explicit_store_and_previous_response_id_are_forwarded_opaquely() {
|
||
let event = json!({
|
||
"type": "response.create",
|
||
"model": "public-model",
|
||
"store": true,
|
||
"previous_response_id": {"future": "opaque"},
|
||
"input": [],
|
||
});
|
||
|
||
let normalized = normalized_continuation(
|
||
&event,
|
||
&ResponsesWebSocketBodyNormalization::for_tests("provider-model")
|
||
.with_provider_type_for_tests("codex"),
|
||
);
|
||
|
||
// Codex HTTP normalization normally forces `store: false` and removes
|
||
// `previous_response_id`. WebSocket framing restores exactly what the
|
||
// client sent so the upstream owns validation and continuation lookup.
|
||
assert_eq!(normalized["store"], true);
|
||
assert_eq!(
|
||
normalized["previous_response_id"],
|
||
json!({"future": "opaque"})
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn explicit_null_websocket_protocol_state_is_not_rewritten() {
|
||
let event = json!({
|
||
"type": "response.create",
|
||
"model": "public-model",
|
||
"store": null,
|
||
"previous_response_id": null,
|
||
"generate": null,
|
||
"input": [],
|
||
});
|
||
|
||
let normalized = normalized_continuation(
|
||
&event,
|
||
&ResponsesWebSocketBodyNormalization::for_tests("provider-model")
|
||
.with_provider_type_for_tests("codex"),
|
||
);
|
||
|
||
assert!(normalized.get("store").is_some_and(Value::is_null));
|
||
assert!(normalized
|
||
.get("previous_response_id")
|
||
.is_some_and(Value::is_null));
|
||
assert!(normalized.get("generate").is_some_and(Value::is_null));
|
||
}
|
||
|
||
#[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 lite_continuations_forward_only_each_turns_incremental_input() {
|
||
let normalization = ResponsesWebSocketBodyNormalization::for_tests("gpt-5.6-sol")
|
||
.with_provider_type_for_tests("codex");
|
||
|
||
for turn in 1..=4 {
|
||
let event = json!({
|
||
"type": "response.create",
|
||
"model": "gpt-5.6-sol",
|
||
"previous_response_id": format!("resp_{turn}"),
|
||
"instructions": "The same large developer instructions.",
|
||
"tools": [{
|
||
"type": "function",
|
||
"name": "shell",
|
||
"parameters": {"type": "object"}
|
||
}],
|
||
"input": [
|
||
{
|
||
"type": "reasoning",
|
||
"id": format!("rs_{turn}"),
|
||
"content": [{
|
||
"type": "reasoning_text",
|
||
"text": format!("reasoning state {turn}")
|
||
}],
|
||
"encrypted_content": format!("opaque-{turn}")
|
||
},
|
||
{
|
||
"type": "function_call_output",
|
||
"call_id": format!("call_{turn}"),
|
||
"output": format!("result {turn}")
|
||
}
|
||
]
|
||
});
|
||
|
||
let normalized = normalized_continuation(&event, &normalization);
|
||
let input = normalized["input"].as_array().expect("incremental input");
|
||
assert_eq!(input.len(), 2);
|
||
assert_eq!(input[0]["type"], "reasoning");
|
||
assert_eq!(input[0]["content"][0]["type"], "reasoning_text");
|
||
assert_eq!(input[1]["type"], "function_call_output");
|
||
assert!(normalized.get("instructions").is_none());
|
||
assert!(normalized.get("tools").is_none());
|
||
assert!(!input.iter().any(|item| item["type"] == "additional_tools"));
|
||
assert!(!input.iter().any(|item| item["role"] == "developer"));
|
||
assert_eq!(normalized["previous_response_id"], format!("resp_{turn}"));
|
||
}
|
||
}
|
||
|
||
#[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!({})));
|
||
}
|
||
}
|