feat(usage): enrich audit metadata and detail views

This commit is contained in:
elky
2026-07-17 19:20:16 +08:00
parent 0be380243b
commit 664c063a06
79 changed files with 6968 additions and 1294 deletions
@@ -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));