mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-04 16:37:46 +08:00
feat(usage): enrich audit metadata and detail views
This commit is contained in:
@@ -131,6 +131,7 @@ const SSE_KEEPALIVE_INTERVAL: Duration = Duration::from_secs(15);
|
||||
const SSE_KEEPALIVE_BYTES: &[u8] = b": aether-keepalive\n\n";
|
||||
const SSE_CONTROL_FILTER_MAX_BUFFER_BYTES: usize = 1024 * 1024;
|
||||
const SSE_TERMINAL_DETECTOR_MAX_LINE_BYTES: usize = 1024 * 1024;
|
||||
const PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES: usize = SSE_TERMINAL_DETECTOR_MAX_LINE_BYTES;
|
||||
const STREAM_IDLE_LOG_INTERVAL: Duration = Duration::from_secs(60);
|
||||
const STREAM_IDLE_LOG_INTERVAL_MS: u64 = 60_000;
|
||||
const REWRITTEN_STREAM_PREFETCH_TIMEOUT: Duration = Duration::from_millis(750);
|
||||
@@ -140,6 +141,78 @@ const DIRECT_PASSTHROUGH_CHANNEL_CAPACITY_ENV: &str =
|
||||
"AETHER_GATEWAY_DIRECT_PASSTHROUGH_CHANNEL_CAPACITY";
|
||||
const DIRECT_PASSTHROUGH_MODE_ENV: &str = "AETHER_GATEWAY_DIRECT_PASSTHROUGH_MODE";
|
||||
|
||||
/// Retains the incomplete tail needed to recognize provider error events split across transport
|
||||
/// chunks without retaining an unbounded copy of the stream.
|
||||
#[derive(Default)]
|
||||
struct ProviderStreamErrorInspection {
|
||||
buffered: Vec<u8>,
|
||||
}
|
||||
|
||||
impl ProviderStreamErrorInspection {
|
||||
fn observe(&mut self, report_context: Option<&Value>, chunk: &[u8]) -> Option<Value> {
|
||||
if chunk.is_empty() {
|
||||
return None;
|
||||
}
|
||||
if let Some(error_body) = extract_provider_private_stream_error_body(report_context, chunk)
|
||||
{
|
||||
return Some(error_body);
|
||||
}
|
||||
|
||||
self.append_rolling(chunk);
|
||||
let error_body = extract_provider_private_stream_error_body(report_context, &self.buffered);
|
||||
if error_body.is_none() {
|
||||
self.trim_completed_sse_events();
|
||||
}
|
||||
error_body
|
||||
}
|
||||
|
||||
fn append_rolling(&mut self, chunk: &[u8]) {
|
||||
if chunk.len() >= PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES {
|
||||
self.buffered.clear();
|
||||
self.buffered.extend_from_slice(
|
||||
&chunk[chunk.len() - PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES..],
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
let overflow = self
|
||||
.buffered
|
||||
.len()
|
||||
.saturating_add(chunk.len())
|
||||
.saturating_sub(PROVIDER_STREAM_ERROR_INSPECTION_MAX_BYTES);
|
||||
if overflow > 0 {
|
||||
self.buffered.drain(..overflow);
|
||||
}
|
||||
self.buffered.extend_from_slice(chunk);
|
||||
}
|
||||
|
||||
fn trim_completed_sse_events(&mut self) {
|
||||
let Ok(text) = std::str::from_utf8(&self.buffered) else {
|
||||
return;
|
||||
};
|
||||
if !text.lines().any(|line| {
|
||||
let line = line.trim_start();
|
||||
line.starts_with("event:") || line.starts_with("data:") || line.starts_with(':')
|
||||
}) {
|
||||
return;
|
||||
}
|
||||
|
||||
let lf_end = self
|
||||
.buffered
|
||||
.windows(2)
|
||||
.rposition(|window| window == b"\n\n")
|
||||
.map(|index| index + 2);
|
||||
let crlf_end = self
|
||||
.buffered
|
||||
.windows(4)
|
||||
.rposition(|window| window == b"\r\n\r\n")
|
||||
.map(|index| index + 4);
|
||||
if let Some(event_end) = lf_end.into_iter().chain(crlf_end).max() {
|
||||
self.buffered.drain(..event_end);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
struct StageElapsedGuard {
|
||||
stage: &'static str,
|
||||
started_at: Instant,
|
||||
@@ -1192,6 +1265,7 @@ struct DirectPassthroughFinalizerCore {
|
||||
stream_usage_report_context: Option<Value>,
|
||||
stream_usage_observer: Option<StreamingStandardTerminalObserver>,
|
||||
stream_usage_observer_buffered: Vec<u8>,
|
||||
provider_error_inspection: ProviderStreamErrorInspection,
|
||||
provider_buffered_body: Vec<u8>,
|
||||
buffered_body: Vec<u8>,
|
||||
provider_body_truncated: bool,
|
||||
@@ -1318,10 +1392,10 @@ impl DirectPassthroughFinalizer {
|
||||
chunk.as_ref(),
|
||||
);
|
||||
}
|
||||
if let Some(error_body_json) = extract_provider_private_stream_error_body(
|
||||
core.stream_usage_report_context.as_ref(),
|
||||
chunk.as_ref(),
|
||||
) {
|
||||
if let Some(error_body_json) = core
|
||||
.provider_error_inspection
|
||||
.observe(core.stream_usage_report_context.as_ref(), chunk.as_ref())
|
||||
{
|
||||
let error_status_code =
|
||||
resolve_local_sync_error_status_code(core.status_code, &error_body_json);
|
||||
core.terminal_failure = Some(build_stream_failure_from_provider_error_body(
|
||||
@@ -1486,6 +1560,7 @@ impl DirectPassthroughFinalizerCore {
|
||||
stream_usage_report_context,
|
||||
stream_usage_observer: _,
|
||||
stream_usage_observer_buffered: _,
|
||||
provider_error_inspection: _,
|
||||
provider_buffered_body,
|
||||
buffered_body,
|
||||
provider_body_truncated,
|
||||
@@ -2207,6 +2282,7 @@ async fn execute_stream_from_direct_passthrough(
|
||||
stream_usage_report_context,
|
||||
stream_usage_observer,
|
||||
stream_usage_observer_buffered: Vec::new(),
|
||||
provider_error_inspection: ProviderStreamErrorInspection::default(),
|
||||
provider_buffered_body: Vec::new(),
|
||||
buffered_body: Vec::new(),
|
||||
provider_body_truncated: false,
|
||||
@@ -2286,6 +2362,7 @@ async fn execute_stream_from_direct_passthrough(
|
||||
.as_ref()
|
||||
.map(|_| StreamingStandardTerminalObserver::default());
|
||||
let mut stream_usage_observer_buffered = Vec::new();
|
||||
let mut provider_error_inspection = ProviderStreamErrorInspection::default();
|
||||
let mut provider_buffered_body = Vec::new();
|
||||
let mut buffered_body = Vec::new();
|
||||
let mut provider_body_truncated = false;
|
||||
@@ -2458,7 +2535,7 @@ async fn execute_stream_from_direct_passthrough(
|
||||
provider_chunk.as_ref(),
|
||||
);
|
||||
}
|
||||
let provider_private_error_body_json = extract_provider_private_stream_error_body(
|
||||
let provider_private_error_body_json = provider_error_inspection.observe(
|
||||
stream_usage_report_context.as_ref(),
|
||||
provider_chunk.as_ref(),
|
||||
);
|
||||
@@ -5007,6 +5084,7 @@ async fn execute_stream_from_frame_stream(
|
||||
.filter(|_| !sync_json_stream_bridge_active_for_report)
|
||||
.map(|_| StreamingStandardTerminalObserver::default());
|
||||
let mut stream_usage_observer_buffered = Vec::new();
|
||||
let mut provider_error_inspection = ProviderStreamErrorInspection::default();
|
||||
append_stream_capture_bytes(
|
||||
&mut provider_buffered_body,
|
||||
&provider_prefetched_body_for_report,
|
||||
@@ -5168,6 +5246,16 @@ async fn execute_stream_from_frame_stream(
|
||||
let replay_chunk = normalized_prefetched_chunk
|
||||
.as_deref()
|
||||
.unwrap_or(provider_prefetched_body_for_report.as_slice());
|
||||
if let Some(error_body_json) = provider_error_inspection
|
||||
.observe(stream_usage_report_context.as_ref(), replay_chunk)
|
||||
{
|
||||
let error_status_code =
|
||||
resolve_local_sync_error_status_code(status_code, &error_body_json);
|
||||
terminal_failure = Some(build_stream_failure_from_provider_error_body(
|
||||
error_status_code,
|
||||
&error_body_json,
|
||||
));
|
||||
}
|
||||
if let (Some(observer), Some(report_context)) = (
|
||||
stream_usage_observer.as_mut(),
|
||||
stream_usage_report_context.as_ref(),
|
||||
@@ -5343,11 +5431,8 @@ async fn execute_stream_from_frame_stream(
|
||||
} else {
|
||||
chunk
|
||||
};
|
||||
let provider_private_error_body_json =
|
||||
extract_provider_private_stream_error_body(
|
||||
stream_usage_report_context.as_ref(),
|
||||
&normalized_chunk,
|
||||
);
|
||||
let provider_private_error_body_json = provider_error_inspection
|
||||
.observe(stream_usage_report_context.as_ref(), &normalized_chunk);
|
||||
if let (Some(observer), Some(report_context)) = (
|
||||
stream_usage_observer.as_mut(),
|
||||
stream_usage_report_context.as_ref(),
|
||||
@@ -5508,11 +5593,8 @@ async fn execute_stream_from_frame_stream(
|
||||
{
|
||||
match normalizer.finish() {
|
||||
Ok(normalized_chunk) if !normalized_chunk.is_empty() => {
|
||||
let provider_private_error_body_json =
|
||||
extract_provider_private_stream_error_body(
|
||||
stream_usage_report_context.as_ref(),
|
||||
&normalized_chunk,
|
||||
);
|
||||
let provider_private_error_body_json = provider_error_inspection
|
||||
.observe(stream_usage_report_context.as_ref(), &normalized_chunk);
|
||||
if let (Some(observer), Some(report_context)) = (
|
||||
stream_usage_observer.as_mut(),
|
||||
stream_usage_report_context.as_ref(),
|
||||
@@ -6114,7 +6196,7 @@ mod tests {
|
||||
stream_terminal_summary_missing_observed_finish,
|
||||
stream_terminal_summary_missing_observed_finish_with_requirement,
|
||||
stream_terminal_summary_represents_failure_with_requirement,
|
||||
ClientVisibleStreamCompletionTracker, DirectPassthroughMode,
|
||||
ClientVisibleStreamCompletionTracker, DirectPassthroughMode, ProviderStreamErrorInspection,
|
||||
};
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::stage_metrics::RequestStageTrace;
|
||||
@@ -6238,6 +6320,51 @@ mod tests {
|
||||
.observe_chunk(b"data: {\"type\":\"response.completed\",\"response\":{}}\r\n\r\n"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_error_inspection_detects_response_failed_at_every_chunk_boundary() {
|
||||
let body = concat!(
|
||||
"event: response.created\n",
|
||||
"data: {\"type\":\"response.created\",\"response\":{\"status\":\"in_progress\"}}\n\n",
|
||||
"event: response.failed\n",
|
||||
"data: {\"type\":\"response.failed\",\"response\":{\"status\":\"failed\",\"error\":{\"type\":\"invalid_request\",\"message\":\"cyber policy rejected the request\",\"code\":\"cyber_policy_violation\",\"param\":\"input\"}}}\n\n",
|
||||
)
|
||||
.as_bytes();
|
||||
|
||||
for split in 1..body.len() {
|
||||
let mut inspection = ProviderStreamErrorInspection::default();
|
||||
let detected = inspection
|
||||
.observe(None, &body[..split])
|
||||
.or_else(|| inspection.observe(None, &body[split..]))
|
||||
.unwrap_or_else(|| panic!("response.failed was missed at byte split {split}"));
|
||||
|
||||
assert_eq!(
|
||||
detected.pointer("/error/code"),
|
||||
Some(&json!("cyber_policy_violation")),
|
||||
"string provider code changed at byte split {split}"
|
||||
);
|
||||
assert_eq!(
|
||||
detected.pointer("/error/param"),
|
||||
Some(&json!("input")),
|
||||
"provider error fields changed at byte split {split}"
|
||||
);
|
||||
}
|
||||
|
||||
let mut inspection = ProviderStreamErrorInspection::default();
|
||||
let mut detected = None;
|
||||
for byte in body.chunks(1) {
|
||||
if let Some(error_body) = inspection.observe(None, byte) {
|
||||
detected = Some(error_body);
|
||||
break;
|
||||
}
|
||||
}
|
||||
let detected = detected.expect("byte-wise response.failed stream should be detected");
|
||||
assert_eq!(
|
||||
detected.pointer("/error/code"),
|
||||
Some(&json!("cyber_policy_violation"))
|
||||
);
|
||||
assert_eq!(detected.pointer("/error/param"), Some(&json!("input")));
|
||||
}
|
||||
|
||||
fn tunnel_proxy_snapshot(base_url: String) -> aether_contracts::ProxySnapshot {
|
||||
aether_contracts::ProxySnapshot {
|
||||
enabled: Some(true),
|
||||
|
||||
@@ -20,10 +20,9 @@ use crate::execution_runtime::submission::{
|
||||
use crate::log_ids::short_request_id;
|
||||
use crate::orchestration::{
|
||||
apply_local_execution_effect, resolve_local_failover_analysis_for_attempt,
|
||||
trace_upstream_response_body, with_upstream_response_report_context,
|
||||
LocalAdaptiveRateLimitEffect, LocalAttemptFailureEffect, LocalExecutionEffect,
|
||||
LocalExecutionEffectContext, LocalHealthFailureEffect, LocalOAuthInvalidationEffect,
|
||||
LocalPoolErrorEffect,
|
||||
with_upstream_response_report_context, LocalAdaptiveRateLimitEffect, LocalAttemptFailureEffect,
|
||||
LocalExecutionEffect, LocalExecutionEffectContext, LocalHealthFailureEffect,
|
||||
LocalOAuthInvalidationEffect, LocalPoolErrorEffect,
|
||||
};
|
||||
use crate::request_candidate_runtime::record_report_request_candidate_status;
|
||||
use crate::request_diagnostics::attach_current_request_diagnostics_to_report_context;
|
||||
@@ -36,6 +35,7 @@ pub(super) struct StreamFailureReport {
|
||||
pub(super) error_type: String,
|
||||
pub(super) error_message: String,
|
||||
extra_error_fields: Map<String, Value>,
|
||||
provider_body_json: Option<Value>,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
@@ -54,20 +54,28 @@ struct StreamFailureBodyFields<'a> {
|
||||
}
|
||||
|
||||
impl StreamFailureReport {
|
||||
fn into_body_json(self) -> Value {
|
||||
fn into_body_jsons(self) -> (Value, Option<Value>) {
|
||||
let Self {
|
||||
status_code,
|
||||
error_type,
|
||||
error_message,
|
||||
mut extra_error_fields,
|
||||
provider_body_json,
|
||||
} = self;
|
||||
extra_error_fields.insert("type".to_string(), Value::String(error_type));
|
||||
extra_error_fields.insert("message".to_string(), Value::String(error_message));
|
||||
extra_error_fields.insert("code".to_string(), Value::from(status_code));
|
||||
Value::Object(Map::from_iter([(
|
||||
let normalized_body = Value::Object(Map::from_iter([(
|
||||
"error".to_string(),
|
||||
Value::Object(extra_error_fields),
|
||||
)]))
|
||||
)]));
|
||||
match provider_body_json {
|
||||
Some(provider_body) if provider_body != normalized_body => {
|
||||
(provider_body, Some(normalized_body))
|
||||
}
|
||||
Some(provider_body) => (provider_body, None),
|
||||
None => (normalized_body, None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn to_json_string(&self) -> serde_json::Result<String> {
|
||||
@@ -94,6 +102,7 @@ pub(super) fn build_stream_failure_report(
|
||||
error_type,
|
||||
error_message,
|
||||
extra_error_fields: Map::new(),
|
||||
provider_body_json: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -124,6 +133,7 @@ pub(super) fn build_stream_failure_from_execution_error(
|
||||
error_type,
|
||||
error_message,
|
||||
extra_error_fields: error_object,
|
||||
provider_body_json: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -150,6 +160,7 @@ pub(super) fn build_stream_failure_from_provider_error_body(
|
||||
error_type,
|
||||
error_message,
|
||||
extra_error_fields: Map::new(),
|
||||
provider_body_json: Some(body_json.clone()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -185,19 +196,33 @@ fn build_stream_failure_sync_payload(
|
||||
failure: StreamFailureReport,
|
||||
) -> GatewaySyncReportRequest {
|
||||
let status_code = failure.status_code;
|
||||
let body = trace_upstream_response_body(None, provider_buffered_body);
|
||||
let (body, client_body) = failure.into_body_jsons();
|
||||
headers.retain(|name, _| {
|
||||
!name.eq_ignore_ascii_case("content-encoding")
|
||||
&& !name.eq_ignore_ascii_case("content-length")
|
||||
&& !name.eq_ignore_ascii_case("content-type")
|
||||
});
|
||||
headers.insert("content-type".to_string(), "application/json".to_string());
|
||||
let report_context = with_upstream_response_report_context(
|
||||
report_context.as_ref(),
|
||||
status_code,
|
||||
Some(&headers),
|
||||
body.as_ref(),
|
||||
Some(&body),
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.or(report_context);
|
||||
headers.remove("content-encoding");
|
||||
headers.remove("content-length");
|
||||
headers.insert("content-type".to_string(), "application/json".to_string());
|
||||
let report_context = report_context.map(|mut context| {
|
||||
if let Some(object) = context.as_object_mut() {
|
||||
let response_headers = serde_json::to_value(&headers).unwrap_or(Value::Null);
|
||||
object.insert(
|
||||
"provider_response_headers".to_string(),
|
||||
response_headers.clone(),
|
||||
);
|
||||
object.insert("client_response_headers".to_string(), response_headers);
|
||||
}
|
||||
context
|
||||
});
|
||||
|
||||
GatewaySyncReportRequest {
|
||||
trace_id: trace_id.to_string(),
|
||||
@@ -205,8 +230,8 @@ fn build_stream_failure_sync_payload(
|
||||
report_context,
|
||||
status_code,
|
||||
headers,
|
||||
body_json: Some(failure.into_body_json()),
|
||||
client_body_json: None,
|
||||
body_json: Some(body),
|
||||
client_body_json: client_body,
|
||||
body_base64: (!provider_buffered_body.is_empty())
|
||||
.then(|| base64::engine::general_purpose::STANDARD.encode(provider_buffered_body)),
|
||||
telemetry,
|
||||
@@ -218,8 +243,9 @@ fn stream_failure_body_field<'a>(
|
||||
field: &str,
|
||||
) -> Option<&'a str> {
|
||||
payload
|
||||
.body_json
|
||||
.client_body_json
|
||||
.as_ref()
|
||||
.or(payload.body_json.as_ref())
|
||||
.and_then(|body_json| body_json.get("error"))
|
||||
.and_then(|value| value.get(field))
|
||||
.and_then(Value::as_str)
|
||||
@@ -477,3 +503,151 @@ pub(super) async fn submit_midstream_stream_failure(
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use base64::Engine as _;
|
||||
use serde_json::json;
|
||||
|
||||
use super::{build_stream_failure_from_provider_error_body, build_stream_failure_sync_payload};
|
||||
|
||||
#[test]
|
||||
fn midstream_failure_trace_uses_terminal_error_instead_of_buffered_sse() {
|
||||
let provider_buffered_body = concat!(
|
||||
"event: response.created\n",
|
||||
"data: {\"type\":\"response.created\",\"response\":{\"instructions\":\"AGENTS.md secret prompt\",\"tools\":[{\"name\":\"update_plan\"}]}}\n\n",
|
||||
"event: response.failed\n",
|
||||
"data: {\"type\":\"response.failed\",\"response\":{\"status\":\"failed\",\"error\":{\"type\":\"invalid_request\",\"message\":\"This content was flagged for possible cybersecurity risk.\",\"code\":\"cyber_policy_violation\",\"param\":\"input\",\"details\":{\"policy_category\":\"cybersecurity\",\"appeal_allowed\":true}}}}\n\n",
|
||||
)
|
||||
.as_bytes();
|
||||
let terminal_error = aether_ai_formats::api::extract_provider_private_stream_error_body(
|
||||
None,
|
||||
provider_buffered_body,
|
||||
)
|
||||
.expect("raw upstream SSE should expose its terminal provider error JSON");
|
||||
let failure = build_stream_failure_from_provider_error_body(400, &terminal_error);
|
||||
|
||||
let payload = build_stream_failure_sync_payload(
|
||||
"trace-cyber-policy",
|
||||
"openai_responses_sync_error".to_string(),
|
||||
Some(json!({"request_id": "request-cyber-policy"})),
|
||||
BTreeMap::from([
|
||||
("Content-Encoding".to_string(), "gzip".to_string()),
|
||||
("Content-Length".to_string(), "4096".to_string()),
|
||||
("Content-Type".to_string(), "text/event-stream".to_string()),
|
||||
(
|
||||
"x-request-id".to_string(),
|
||||
"req_usage-cyber-risk-demo".to_string(),
|
||||
),
|
||||
]),
|
||||
None,
|
||||
provider_buffered_body,
|
||||
failure,
|
||||
);
|
||||
|
||||
let trace_body = payload
|
||||
.report_context
|
||||
.as_ref()
|
||||
.and_then(|context| context.pointer("/upstream_response/body"))
|
||||
.expect("candidate trace should include the terminal error body");
|
||||
assert_eq!(
|
||||
trace_body,
|
||||
payload.body_json.as_ref().expect("usage error body")
|
||||
);
|
||||
assert_eq!(trace_body, &terminal_error);
|
||||
assert_eq!(trace_body["error"]["type"], json!("invalid_request"));
|
||||
assert_eq!(
|
||||
trace_body["error"]["message"],
|
||||
json!("This content was flagged for possible cybersecurity risk.")
|
||||
);
|
||||
assert_eq!(trace_body["error"]["code"], json!("cyber_policy_violation"));
|
||||
assert_eq!(trace_body["error"]["param"], json!("input"));
|
||||
assert_eq!(
|
||||
trace_body["error"]["details"],
|
||||
json!({
|
||||
"policy_category": "cybersecurity",
|
||||
"appeal_allowed": true
|
||||
})
|
||||
);
|
||||
assert_eq!(
|
||||
payload
|
||||
.report_context
|
||||
.as_ref()
|
||||
.and_then(|context| context.pointer("/upstream_response/headers/content-type")),
|
||||
Some(&json!("application/json"))
|
||||
);
|
||||
assert_eq!(
|
||||
payload
|
||||
.report_context
|
||||
.as_ref()
|
||||
.and_then(|context| context.pointer("/upstream_response/headers/x-request-id")),
|
||||
Some(&json!("req_usage-cyber-risk-demo"))
|
||||
);
|
||||
assert_eq!(
|
||||
payload
|
||||
.report_context
|
||||
.as_ref()
|
||||
.and_then(|context| context.pointer("/provider_response_headers/content-type")),
|
||||
Some(&json!("application/json"))
|
||||
);
|
||||
assert_eq!(
|
||||
payload
|
||||
.report_context
|
||||
.as_ref()
|
||||
.and_then(|context| context.pointer("/client_response_headers/content-type")),
|
||||
Some(&json!("application/json"))
|
||||
);
|
||||
let trace_headers = payload
|
||||
.report_context
|
||||
.as_ref()
|
||||
.and_then(|context| context.pointer("/upstream_response/headers"))
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.expect("candidate trace should include terminal JSON headers");
|
||||
assert!(!trace_headers
|
||||
.keys()
|
||||
.any(|name| name.eq_ignore_ascii_case("content-encoding")));
|
||||
assert!(!trace_headers
|
||||
.keys()
|
||||
.any(|name| name.eq_ignore_ascii_case("content-length")));
|
||||
assert_eq!(
|
||||
payload.headers.get("content-type").map(String::as_str),
|
||||
Some("application/json")
|
||||
);
|
||||
assert!(!payload
|
||||
.headers
|
||||
.keys()
|
||||
.any(|name| name.eq_ignore_ascii_case("content-encoding")));
|
||||
assert!(!payload
|
||||
.headers
|
||||
.keys()
|
||||
.any(|name| name.eq_ignore_ascii_case("content-length")));
|
||||
assert_eq!(
|
||||
payload
|
||||
.client_body_json
|
||||
.as_ref()
|
||||
.and_then(|body| body.pointer("/error/message")),
|
||||
Some(&json!(
|
||||
"This content was flagged for possible cybersecurity risk."
|
||||
))
|
||||
);
|
||||
assert_eq!(
|
||||
payload
|
||||
.client_body_json
|
||||
.as_ref()
|
||||
.and_then(|body| body.pointer("/error/code")),
|
||||
Some(&json!(400))
|
||||
);
|
||||
assert_ne!(payload.client_body_json.as_ref(), Some(&terminal_error));
|
||||
assert!(!trace_body.to_string().contains("AGENTS.md secret prompt"));
|
||||
|
||||
let raw_capture = payload
|
||||
.body_base64
|
||||
.as_deref()
|
||||
.and_then(|body| base64::engine::general_purpose::STANDARD.decode(body).ok())
|
||||
.expect("raw provider stream should remain available for usage auditing");
|
||||
assert_eq!(raw_capture, provider_buffered_body);
|
||||
assert!(String::from_utf8_lossy(&raw_capture).contains("AGENTS.md secret prompt"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -733,6 +733,110 @@ async fn admin_monitoring_trace_request_exposes_failed_candidate_upstream_respon
|
||||
assert!(extra.get("provider_response").is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_monitoring_trace_request_prefers_ref_backed_usage_response_body() {
|
||||
let mut candidate = sample_candidate(
|
||||
"cand-used",
|
||||
"request-ref-body",
|
||||
0,
|
||||
RequestCandidateStatus::Failed,
|
||||
Some(101),
|
||||
Some(33),
|
||||
Some(400),
|
||||
);
|
||||
candidate.extra_data = Some(json!({
|
||||
"upstream_response": {
|
||||
"status_code": 400,
|
||||
"headers": {
|
||||
"content-type": "text/event-stream",
|
||||
"x-request-id": "stale-request-like-body"
|
||||
},
|
||||
"body": {
|
||||
"model": "gpt-5.6-sol",
|
||||
"input": [{"role": "user", "content": "request prompt"}]
|
||||
}
|
||||
}
|
||||
}));
|
||||
|
||||
let request_candidates = Arc::new(InMemoryRequestCandidateRepository::seed(vec![candidate]));
|
||||
let provider_catalog = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider()],
|
||||
vec![sample_endpoint()],
|
||||
vec![sample_key()],
|
||||
));
|
||||
let mut usage = sample_usage(
|
||||
"request-ref-body",
|
||||
"provider-1",
|
||||
"OpenAI",
|
||||
0,
|
||||
0.0,
|
||||
"failed",
|
||||
Some(400),
|
||||
100,
|
||||
);
|
||||
usage.candidate_id = Some("cand-used".to_string());
|
||||
usage.response_headers = Some(json!({
|
||||
"content-type": "application/json",
|
||||
"x-request-id": "req_usage-cyber-risk-demo"
|
||||
}));
|
||||
usage.response_body = Some(json!({
|
||||
"error": {
|
||||
"type": "invalid_request",
|
||||
"message": "This content was flagged for possible cybersecurity risk.",
|
||||
"code": 400
|
||||
}
|
||||
}));
|
||||
usage.response_body_state = Some(UsageBodyCaptureState::Reference);
|
||||
let usage_repository = Arc::new(InMemoryUsageReadRepository::seed_with_detached_bodies(
|
||||
vec![usage],
|
||||
));
|
||||
let data_state =
|
||||
crate::data::GatewayDataState::with_request_candidate_and_usage_repository_for_tests(
|
||||
request_candidates,
|
||||
usage_repository,
|
||||
)
|
||||
.with_provider_catalog_reader(provider_catalog);
|
||||
let state = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let context = request_context(
|
||||
http::Method::GET,
|
||||
"/api/admin/monitoring/trace/request-ref-body",
|
||||
);
|
||||
|
||||
let response = local_monitoring_response(&state, &context)
|
||||
.await
|
||||
.expect("handler should not error")
|
||||
.expect("route should be handled locally");
|
||||
|
||||
assert_eq!(response.status(), http::StatusCode::OK);
|
||||
let body = to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("body should read");
|
||||
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
|
||||
let upstream_response = &payload["candidates"][0]["extra_data"]["upstream_response"];
|
||||
assert_eq!(
|
||||
upstream_response["headers"],
|
||||
json!({
|
||||
"content-type": "application/json",
|
||||
"x-request-id": "req_usage-cyber-risk-demo"
|
||||
})
|
||||
);
|
||||
assert_eq!(
|
||||
upstream_response["body"]["error"],
|
||||
json!({
|
||||
"type": "invalid_request",
|
||||
"message": "This content was flagged for possible cybersecurity risk.",
|
||||
"code": 400
|
||||
})
|
||||
);
|
||||
assert!(upstream_response["body"].get("input").is_none());
|
||||
assert_eq!(
|
||||
upstream_response["body_ref"],
|
||||
json!("usage://request/request-ref-body/response_body")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_monitoring_trace_request_decodes_connect_json_response_body_refs() {
|
||||
let mut candidate = sample_candidate(
|
||||
|
||||
@@ -30,6 +30,25 @@ struct ResolvedAdminMonitoringTrace {
|
||||
usage: Option<StoredRequestUsageAudit>,
|
||||
}
|
||||
|
||||
async fn hydrate_admin_monitoring_trace_response_body(
|
||||
state: &AdminAppState<'_>,
|
||||
mut usage: StoredRequestUsageAudit,
|
||||
) -> Result<StoredRequestUsageAudit, GatewayError> {
|
||||
let is_error_node = !usage.status.eq_ignore_ascii_case("completed")
|
||||
|| usage
|
||||
.status_code
|
||||
.is_some_and(|status| !(200..300).contains(&status));
|
||||
let response_body_ref = if is_error_node && usage.response_body.is_none() {
|
||||
usage.response_body_ref.clone()
|
||||
} else {
|
||||
None
|
||||
};
|
||||
if let Some(body_ref) = response_body_ref.as_deref() {
|
||||
usage.response_body = state.resolve_request_usage_body_ref(body_ref).await?;
|
||||
}
|
||||
Ok(usage)
|
||||
}
|
||||
|
||||
pub(super) async fn build_admin_monitoring_trace_request_response(
|
||||
state: &AdminAppState<'_>,
|
||||
request_context: &AdminRequestContext<'_>,
|
||||
@@ -92,6 +111,10 @@ async fn resolve_admin_monitoring_trace(
|
||||
.read_request_usage_audit_shallow(request_id)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let usage = match usage {
|
||||
Some(usage) => Some(hydrate_admin_monitoring_trace_response_body(state, usage).await?),
|
||||
None => None,
|
||||
};
|
||||
return Ok(Some(ResolvedAdminMonitoringTrace { trace, usage }));
|
||||
}
|
||||
|
||||
@@ -102,9 +125,10 @@ async fn resolve_admin_monitoring_trace(
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?
|
||||
{
|
||||
usage_candidates.push(usage);
|
||||
usage_candidates.push(hydrate_admin_monitoring_trace_response_body(state, usage).await?);
|
||||
}
|
||||
if let Some(usage) = state.find_request_usage_by_id(request_id).await? {
|
||||
let usage = hydrate_admin_monitoring_trace_response_body(state, usage).await?;
|
||||
if !usage_candidates.iter().any(|item| item.id == usage.id) {
|
||||
usage_candidates.push(usage);
|
||||
}
|
||||
|
||||
@@ -35,13 +35,16 @@ pub(crate) async fn build_admin_create_user_api_key_response(
|
||||
let Some(user_id) = admin_user_id_from_api_keys_path(request_context.path()) else {
|
||||
return Ok(build_admin_users_bad_request_response("缺少 user_id"));
|
||||
};
|
||||
if state.find_user_auth_by_id(&user_id).await?.is_none() {
|
||||
let Some(target_user) = state.find_user_auth_by_id(&user_id).await? else {
|
||||
return Ok((
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": "用户不存在" })),
|
||||
)
|
||||
.into_response());
|
||||
}
|
||||
};
|
||||
// The authenticated principal authorizes this admin operation, but ownership and inherited
|
||||
// policies must always come from the user selected in the request path.
|
||||
let target_user_id = target_user.id;
|
||||
|
||||
let Some(request_body) = request_body else {
|
||||
return Ok((
|
||||
@@ -148,7 +151,7 @@ pub(crate) async fn build_admin_create_user_api_key_response(
|
||||
|
||||
let Some(created) = state
|
||||
.create_user_api_key(aether_data::repository::auth::CreateUserApiKeyRecord {
|
||||
user_id: user_id.clone(),
|
||||
user_id: target_user_id.clone(),
|
||||
api_key_id: uuid::Uuid::new_v4().to_string(),
|
||||
key_hash: hash_admin_user_api_key(&plaintext_key),
|
||||
key_encrypted: Some(key_encrypted),
|
||||
@@ -174,7 +177,11 @@ pub(crate) async fn build_admin_create_user_api_key_response(
|
||||
|
||||
let created = if allowed_providers.is_some() {
|
||||
match state
|
||||
.set_user_api_key_allowed_providers(&user_id, &created.api_key_id, allowed_providers)
|
||||
.set_user_api_key_allowed_providers(
|
||||
&target_user_id,
|
||||
&created.api_key_id,
|
||||
allowed_providers,
|
||||
)
|
||||
.await?
|
||||
{
|
||||
Some(updated) => updated,
|
||||
@@ -186,7 +193,7 @@ pub(crate) async fn build_admin_create_user_api_key_response(
|
||||
let created = if feature_settings.is_some() {
|
||||
match state
|
||||
.set_user_api_key_feature_settings(
|
||||
&user_id,
|
||||
&target_user_id,
|
||||
&created.api_key_id,
|
||||
feature_settings.clone(),
|
||||
)
|
||||
|
||||
@@ -506,6 +506,9 @@ fn build_users_me_usage_record_payload(
|
||||
if let Some(reasoning_effort) = item.provider_reasoning_effort() {
|
||||
payload["reasoning_effort"] = json!(reasoning_effort);
|
||||
}
|
||||
if let Some(requested_reasoning_effort) = item.requested_reasoning_effort() {
|
||||
payload["requested_reasoning_effort"] = json!(requested_reasoning_effort);
|
||||
}
|
||||
if let Some(service_tier) = item.provider_service_tier() {
|
||||
payload["service_tier"] = json!(service_tier);
|
||||
}
|
||||
@@ -577,6 +580,9 @@ fn build_users_me_usage_active_payload(item: &StoredRequestUsageAudit) -> serde_
|
||||
if let Some(reasoning_effort) = item.provider_reasoning_effort() {
|
||||
payload["reasoning_effort"] = json!(reasoning_effort);
|
||||
}
|
||||
if let Some(requested_reasoning_effort) = item.requested_reasoning_effort() {
|
||||
payload["requested_reasoning_effort"] = json!(requested_reasoning_effort);
|
||||
}
|
||||
if let Some(service_tier) = item.provider_service_tier() {
|
||||
payload["service_tier"] = json!(service_tier);
|
||||
}
|
||||
@@ -1655,6 +1661,27 @@ mod tests {
|
||||
assert_eq!(payload["cache_creation_ephemeral_1h_input_tokens"], 6);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn user_usage_payloads_expose_requested_and_provider_reasoning_mapping() {
|
||||
let item = StoredRequestUsageAudit {
|
||||
request_body: Some(json!({
|
||||
"reasoning": { "effort": "xhigh" }
|
||||
})),
|
||||
provider_request_body: Some(json!({
|
||||
"reasoning": { "effort": "max" }
|
||||
})),
|
||||
..sample_usage("completed")
|
||||
};
|
||||
|
||||
let record = build_users_me_usage_record_payload(&item, false, &BTreeMap::new(), false);
|
||||
let active = build_users_me_usage_active_payload(&item);
|
||||
|
||||
assert_eq!(record["requested_reasoning_effort"], "xhigh");
|
||||
assert_eq!(active["requested_reasoning_effort"], "xhigh");
|
||||
assert_eq!(record["reasoning_effort"], "max");
|
||||
assert_eq!(active["reasoning_effort"], "max");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn user_usage_active_override_uses_terminal_candidate_latency() {
|
||||
let candidate = sample_candidate(
|
||||
|
||||
@@ -20,7 +20,7 @@ use http::StatusCode;
|
||||
use serde_json::json;
|
||||
|
||||
use super::super::{
|
||||
build_router_with_state, issue_test_admin_access_token, start_server, AppState,
|
||||
build_router_with_state, hash_api_key, issue_test_admin_access_token, start_server, AppState,
|
||||
};
|
||||
use crate::constants::{
|
||||
GATEWAY_HEADER, TRUSTED_ADMIN_SESSION_ID_HEADER, TRUSTED_ADMIN_USER_ID_HEADER,
|
||||
@@ -1672,6 +1672,192 @@ async fn gateway_handles_admin_user_api_key_routes_locally_with_trusted_admin_pr
|
||||
upstream_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_created_user_key_inherits_target_user_policy_not_admin_policy() {
|
||||
let mut admin_snapshot = sample_admin_api_key_snapshot("admin-user", "admin-seed-key");
|
||||
admin_snapshot.user_role = "admin".to_string();
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![
|
||||
(Some(hash_api_key("sk-admin-seed")), admin_snapshot),
|
||||
(
|
||||
Some(hash_api_key("sk-target-seed")),
|
||||
sample_admin_api_key_snapshot("target-user", "target-seed-key"),
|
||||
),
|
||||
]));
|
||||
let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![
|
||||
sample_admin_user_with_role("admin-user", "admin", "[email protected]", "admin"),
|
||||
sample_admin_user_with_role("target-user", "user", "[email protected]", "target"),
|
||||
]));
|
||||
|
||||
let admin_group = user_repository
|
||||
.create_user_group(UpsertUserGroupRecord {
|
||||
name: "Admin OpenAI".to_string(),
|
||||
description: None,
|
||||
priority: 10,
|
||||
allowed_providers: Some(vec!["openai".to_string()]),
|
||||
allowed_providers_mode: "specific".to_string(),
|
||||
allowed_api_formats: Some(vec!["openai:responses".to_string()]),
|
||||
allowed_api_formats_mode: "specific".to_string(),
|
||||
allowed_models: Some(vec!["gpt-5.4".to_string()]),
|
||||
allowed_models_mode: "specific".to_string(),
|
||||
rate_limit: Some(100),
|
||||
rate_limit_mode: "custom".to_string(),
|
||||
})
|
||||
.await
|
||||
.expect("admin group should create")
|
||||
.expect("admin group should exist");
|
||||
user_repository
|
||||
.add_user_to_group(&admin_group.id, "admin-user")
|
||||
.await
|
||||
.expect("admin group membership should create");
|
||||
|
||||
let target_group = user_repository
|
||||
.create_user_group(UpsertUserGroupRecord {
|
||||
name: "Target Claude".to_string(),
|
||||
description: None,
|
||||
priority: 10,
|
||||
allowed_providers: Some(vec!["anthropic".to_string()]),
|
||||
allowed_providers_mode: "specific".to_string(),
|
||||
allowed_api_formats: Some(vec!["claude:messages".to_string()]),
|
||||
allowed_api_formats_mode: "specific".to_string(),
|
||||
allowed_models: Some(vec!["claude-sonnet-4-5".to_string()]),
|
||||
allowed_models_mode: "specific".to_string(),
|
||||
rate_limit: Some(30),
|
||||
rate_limit_mode: "custom".to_string(),
|
||||
})
|
||||
.await
|
||||
.expect("target group should create")
|
||||
.expect("target group should exist");
|
||||
user_repository
|
||||
.add_user_to_group(&target_group.id, "target-user")
|
||||
.await
|
||||
.expect("target group membership should create");
|
||||
user_repository
|
||||
.update_user_feature_settings(
|
||||
"admin-user",
|
||||
Some(json!({"chat_pii_redaction": {"enabled": false}})),
|
||||
)
|
||||
.await
|
||||
.expect("admin feature settings should update");
|
||||
user_repository
|
||||
.update_user_feature_settings(
|
||||
"target-user",
|
||||
Some(json!({"chat_pii_redaction": {"enabled": true}})),
|
||||
)
|
||||
.await
|
||||
.expect("target feature settings should update");
|
||||
|
||||
let data_state = GatewayDataState::with_auth_api_key_repository_for_tests(auth_repository)
|
||||
.with_user_reader(user_repository);
|
||||
let app_state = AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(data_state);
|
||||
let inspection_state = app_state.clone();
|
||||
let gateway = build_router_with_state(app_state);
|
||||
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||
|
||||
let response = reqwest::Client::new()
|
||||
.post(format!(
|
||||
"{gateway_url}/api/admin/users/target-user/api-keys"
|
||||
))
|
||||
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
|
||||
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user")
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "admin-session")
|
||||
.json(&json!({"name": "target-key"}))
|
||||
.send()
|
||||
.await
|
||||
.expect("request should succeed");
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = response.json().await.expect("response should parse");
|
||||
let plaintext_key = payload["key"].as_str().expect("plaintext key should exist");
|
||||
let api_key_id = payload["id"].as_str().expect("key id should exist");
|
||||
|
||||
let resolved = inspection_state
|
||||
.read_cached_auth_api_key_snapshot_by_key_hash(
|
||||
&hash_api_key(plaintext_key),
|
||||
chrono::Utc::now().timestamp().max(0) as u64,
|
||||
)
|
||||
.await
|
||||
.expect("created key snapshot should resolve")
|
||||
.expect("created key snapshot should exist");
|
||||
assert_eq!(resolved.user_id, "target-user");
|
||||
assert_eq!(
|
||||
resolved.effective_allowed_providers(),
|
||||
Some(&["anthropic".to_string()][..])
|
||||
);
|
||||
assert_eq!(
|
||||
resolved.effective_allowed_api_formats(),
|
||||
Some(&["claude:messages".to_string()][..])
|
||||
);
|
||||
assert_eq!(
|
||||
resolved.effective_allowed_models(),
|
||||
Some(&["claude-sonnet-4-5".to_string()][..])
|
||||
);
|
||||
assert_eq!(resolved.user_rate_limit, Some(30));
|
||||
|
||||
inspection_state
|
||||
.update_user_group(
|
||||
&target_group.id,
|
||||
UpsertUserGroupRecord {
|
||||
name: "Target Gemini".to_string(),
|
||||
description: None,
|
||||
priority: 10,
|
||||
allowed_providers: Some(vec!["google".to_string()]),
|
||||
allowed_providers_mode: "specific".to_string(),
|
||||
allowed_api_formats: Some(vec!["gemini:generate-content".to_string()]),
|
||||
allowed_api_formats_mode: "specific".to_string(),
|
||||
allowed_models: Some(vec!["gemini-2.5-pro".to_string()]),
|
||||
allowed_models_mode: "specific".to_string(),
|
||||
rate_limit: Some(15),
|
||||
rate_limit_mode: "custom".to_string(),
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("target group should update")
|
||||
.expect("target group should still exist");
|
||||
let updated = inspection_state
|
||||
.read_cached_auth_api_key_snapshot_by_key_hash(
|
||||
&hash_api_key(plaintext_key),
|
||||
chrono::Utc::now().timestamp().max(0) as u64,
|
||||
)
|
||||
.await
|
||||
.expect("created key snapshot should refresh")
|
||||
.expect("created key snapshot should still exist");
|
||||
assert_eq!(updated.user_id, "target-user");
|
||||
assert_eq!(
|
||||
updated.effective_allowed_providers(),
|
||||
Some(&["google".to_string()][..])
|
||||
);
|
||||
assert_eq!(
|
||||
updated.effective_allowed_api_formats(),
|
||||
Some(&["gemini:generate-content".to_string()][..])
|
||||
);
|
||||
assert_eq!(
|
||||
updated.effective_allowed_models(),
|
||||
Some(&["gemini-2.5-pro".to_string()][..])
|
||||
);
|
||||
assert_eq!(updated.user_rate_limit, Some(15));
|
||||
|
||||
let target_features = inspection_state
|
||||
.read_user_feature_settings("target-user")
|
||||
.await
|
||||
.expect("target feature settings should resolve");
|
||||
let key_features = inspection_state
|
||||
.read_auth_api_key_feature_settings("target-user", api_key_id, false)
|
||||
.await
|
||||
.expect("key feature settings should resolve");
|
||||
assert_eq!(
|
||||
target_features,
|
||||
Some(json!({"chat_pii_redaction": {"enabled": true}}))
|
||||
);
|
||||
assert_eq!(
|
||||
key_features, None,
|
||||
"an omitted key override must inherit target settings"
|
||||
);
|
||||
|
||||
gateway_handle.abort();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn gateway_returns_conflict_for_admin_create_user_api_key_when_writer_unavailable() {
|
||||
let upstream_hits = Arc::new(Mutex::new(0usize));
|
||||
|
||||
Reference in New Issue
Block a user