Files
Genarrative/server-rs/crates/api-server/src/request_context.rs
T
k88936 dbf3adeb13 后端:接入游玩计数内存缓冲与 flush worker
- game_play_counter.rs:新增纯内存计数器(30min 身份去重、IP+game 固定窗口限流、按时/按量取增量、requeue、过期清理)与 9 个单测
- game_play_counter_worker.rs:新增 flush worker,Build 失败放回重试、Timeout/ConnectDropped 丢弃并记录丢失量,关停前强制 flush
- main.rs:HTTP 角色注册 worker,并在 finalize_shutdown 内按 outbox 超时强制落库
- state.rs/config.rs:AppState 持有计数器,新增 GENARRATIVE_GAME_PLAY_COUNTER_FLUSH_INTERVAL_MS(默认 5s)
- request_context.rs:抽出 client_ip_from_headers 并补测试,runtime_profile.rs 改为复用
2026-10-03 18:25:49 +08:00

190 lines
5.7 KiB
Rust

use std::time::{Duration, Instant};
use axum::{
extract::Request,
http::{HeaderMap, HeaderValue, Request as HttpRequest, header::HeaderName},
middleware::Next,
response::Response,
};
use shared_contracts::api::API_RESPONSE_ENVELOPE_HEADER;
use uuid::Uuid;
pub use shared_contracts::api::X_REQUEST_ID_HEADER;
tokio::task_local! {
/// 当前请求的上下文:`AppError::into_response` 这类拿不到 extensions 的转换也要按同一份
/// meta 口径给出 requestId / operation,错误 envelope 与成功 envelope 才不会各写一套。
pub static CURRENT_REQUEST_CONTEXT: RequestContext;
}
// 当前阶段先把请求级元信息统一挂到 extensions,后续响应头、envelope 与错误处理中间件继续复用。
#[derive(Clone, Debug)]
pub struct RequestContext {
request_id: String,
operation: String,
request_started_at: Instant,
wants_envelope: bool,
external_call_deadline: Option<Instant>,
}
impl RequestContext {
pub fn new(
request_id: String,
operation: String,
elapsed_seed: Duration,
wants_envelope: bool,
) -> Self {
Self {
request_id,
operation,
request_started_at: Instant::now()
.checked_sub(elapsed_seed)
.unwrap_or_else(Instant::now),
wants_envelope,
external_call_deadline: None,
}
}
pub fn with_external_call_deadline(mut self, deadline: Instant) -> Self {
self.external_call_deadline = Some(deadline);
self
}
pub fn request_id(&self) -> &str {
&self.request_id
}
pub fn operation(&self) -> &str {
&self.operation
}
pub fn wants_envelope(&self) -> bool {
self.wants_envelope
}
pub fn external_call_deadline(&self) -> Option<Instant> {
self.external_call_deadline
}
pub fn elapsed(&self) -> u64 {
self.request_started_at
.elapsed()
.as_millis()
.min(u64::MAX as u128) as u64
}
}
pub async fn attach_request_context(mut request: Request, next: Next) -> Response {
let wants_envelope = wants_api_envelope(&request);
let request_id = request
.headers()
.get(X_REQUEST_ID_HEADER)
.and_then(|value| value.to_str().ok())
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.unwrap_or_else(|| Uuid::new_v4().to_string());
let operation = format!("{} {}", request.method(), request.uri());
let context = RequestContext::new(
request_id.clone(),
operation,
Duration::ZERO,
wants_envelope,
);
let context_for_scope = context.clone();
request.extensions_mut().insert(context);
// 统一把 request_id 写回请求头,方便后续 tracing、响应头与 envelope 层读取同一来源。
if let Ok(header_value) = HeaderValue::from_str(&request_id) {
request
.headers_mut()
.insert(HeaderName::from_static(X_REQUEST_ID_HEADER), header_value);
}
CURRENT_REQUEST_CONTEXT
.scope(context_for_scope, async move { next.run(request).await })
.await
}
/// 从代理头解析客户端 IP:反代固定用 `x-forwarded-for` 的第一个地址,直连(本地开发)
/// 回退 `x-real-ip`,都拿不到时才兜底回环地址。
pub fn client_ip_from_headers(headers: &HeaderMap) -> String {
headers
.get("x-forwarded-for")
.and_then(|value| value.to_str().ok())
.and_then(|value| value.split(',').next())
.map(str::trim)
.filter(|value| !value.is_empty())
.or_else(|| {
headers
.get("x-real-ip")
.and_then(|value| value.to_str().ok())
.map(str::trim)
.filter(|value| !value.is_empty())
})
.unwrap_or("127.0.0.1")
.to_string()
}
pub fn resolve_request_id<B>(request: &HttpRequest<B>) -> Option<String> {
request
.extensions()
.get::<RequestContext>()
.map(|context| context.request_id().to_string())
.or_else(|| {
request
.headers()
.get(X_REQUEST_ID_HEADER)
.and_then(|value| value.to_str().ok())
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
}
fn wants_api_envelope<B>(request: &HttpRequest<B>) -> bool {
request
.headers()
.get(API_RESPONSE_ENVELOPE_HEADER)
.and_then(|value| value.to_str().ok())
.map(str::trim)
.map(str::to_lowercase)
.is_some_and(|value| matches!(value.as_str(), "1" | "true" | "v1" | "envelope"))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn request_context_has_no_external_call_deadline_by_default() {
let context = RequestContext::new(
"request-1".to_string(),
"GET /healthz".to_string(),
Duration::ZERO,
false,
);
assert_eq!(context.external_call_deadline(), None);
}
#[test]
fn client_ip_prefers_first_forwarded_address() {
let mut headers = HeaderMap::new();
headers.insert(
"x-forwarded-for",
HeaderValue::from_static("203.0.113.7, 10.0.0.1"),
);
assert_eq!(client_ip_from_headers(&headers), "203.0.113.7");
}
#[test]
fn client_ip_falls_back_to_real_ip_then_loopback() {
let mut headers = HeaderMap::new();
headers.insert("x-real-ip", HeaderValue::from_static("198.51.100.9"));
assert_eq!(client_ip_from_headers(&headers), "198.51.100.9");
assert_eq!(client_ip_from_headers(&HeaderMap::new()), "127.0.0.1");
}
}