95705cbb11
- request_context.rs:client_ip_from_headers 改为优先 nginx 覆盖写入的 x-real-ip,x-forwarded-for 只作回退且取最后一段(nginx 用 $proxy_add_x_forwarded_for 追加的真实对端),不再信任可伪造的首段 - 影响两个调用方:公开游玩上报的匿名身份与 IP+game 限流键、微信支付下单的 payer_client_ip - 更新测试:X-Real-IP 优先、XFF 回退取末段、空 X-Real-IP 回退与回环兜底
206 lines
6.5 KiB
Rust
206 lines
6.5 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-real-ip`:nginx 用 `$remote_addr` 覆盖写入,是真实 TCP 对端,调用方无法伪造。
|
||
/// `x-forwarded-for` 只作回退,并取**最后一段**——nginx 用 `$proxy_add_x_forwarded_for` 会把真实
|
||
/// 对端追加在末尾,前面几段是调用方自带的、可伪造。两者都拿不到时才兜底回环地址。
|
||
pub fn client_ip_from_headers(headers: &HeaderMap) -> String {
|
||
headers
|
||
.get("x-real-ip")
|
||
.and_then(|value| value.to_str().ok())
|
||
.map(str::trim)
|
||
.filter(|value| !value.is_empty())
|
||
.or_else(|| {
|
||
headers
|
||
.get("x-forwarded-for")
|
||
.and_then(|value| value.to_str().ok())
|
||
.and_then(|value| value.rsplit(',').next())
|
||
.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_real_ip_over_forwarded_for() {
|
||
let mut headers = HeaderMap::new();
|
||
headers.insert(
|
||
"x-forwarded-for",
|
||
HeaderValue::from_static("203.0.113.7, 10.0.0.1"),
|
||
);
|
||
headers.insert("x-real-ip", HeaderValue::from_static("198.51.100.9"));
|
||
assert_eq!(client_ip_from_headers(&headers), "198.51.100.9");
|
||
}
|
||
|
||
#[test]
|
||
fn client_ip_forwarded_for_fallback_uses_last_address() {
|
||
// nginx 把真实对端追加在末尾,前面是调用方可伪造的值,只能取最后一段。
|
||
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), "10.0.0.1");
|
||
}
|
||
|
||
#[test]
|
||
fn client_ip_ignores_blank_real_ip_and_falls_back_then_loopback() {
|
||
let mut headers = HeaderMap::new();
|
||
headers.insert("x-real-ip", HeaderValue::from_static(" "));
|
||
headers.insert("x-forwarded-for", HeaderValue::from_static("10.0.0.1"));
|
||
assert_eq!(client_ip_from_headers(&headers), "10.0.0.1");
|
||
assert_eq!(client_ip_from_headers(&HeaderMap::new()), "127.0.0.1");
|
||
}
|
||
}
|