Skip to main content

yggdrasil/mcp/
auth.rs

1//! MCP 请求鉴权:bearer token → 已认证主体。
2//!
3//! 流程(每请求):
4//! 1. axum `from_fn` 中间件 [`mcp_auth_middleware`] 解析 `Authorization: Bearer ygg_...`;
5//! 2. SHA-256 哈希后在 `mcp_tokens` 表常量查找未撤销、未过期的行;
6//! 3. 命中则把 [`McpPrincipal`] { user_id, scope, token_id } 注入 request.extensions,
7//!    并异步刷新 `last_used_at`;
8//! 4. 未命中/缺失 → 401(不区分原因,避免探测)。
9//!
10//! Origin→403、协议版本头、体积上限由 rmcp 的 `StreamableHttpServerConfig` 内置,
11//! 本模块只管 bearer 鉴权。`McpPrincipal` 经 rmcp 的 `Extension<http::request::Parts>`
12//! 提取器在工具内读取(见 `server.rs`)。
13
14use axum::body::Body;
15use axum::http::{HeaderMap, Request, StatusCode};
16use axum::middleware::Next;
17use axum::response::Response;
18
19use crate::db::pool::get_conn;
20use crate::models::mcp_token::TokenScope;
21
22/// Bearer token 前缀(明文形式以 `ygg_` 起头,便于人眼识别与日志脱敏)。
23pub const TOKEN_PREFIX: &str = "ygg_";
24
25/// 已认证主体:由中间件注入 request.extensions,工具经 Extension 提取器读取。
26#[derive(Clone, Debug)]
27pub struct McpPrincipal {
28    pub user_id: i32,
29    pub scope: TokenScope,
30    /// 令牌 DB id(String,与 model 一致;用于审计/last_used_at 刷新)。
31    pub token_id: String,
32}
33
34/// 从 Authorization 头解析 bearer 明文(去前缀后返回完整 token 字符串)。
35/// 非 bearer、缺前缀、解码失败统一返回 None。
36pub fn extract_bearer(headers: &HeaderMap) -> Option<String> {
37    let raw = headers
38        .get(axum::http::header::AUTHORIZATION)?
39        .to_str()
40        .ok()?;
41    let scheme = raw.strip_prefix("Bearer ")?.trim();
42    if scheme.starts_with(TOKEN_PREFIX) {
43        Some(scheme.to_string())
44    } else {
45        None
46    }
47}
48
49/// 明文 token → SHA-256 hex(与 mcp_tokens.token_hash 列一致,用于 DB 查找)。
50pub fn hash_token(token: &str) -> String {
51    crate::utils::server::sha256_hex(token)
52}
53
54/// MCP 鉴权中间件:无 bearer 或 token 无效 → 401;否则注入 McpPrincipal。
55///
56/// 注意:这是 T1 的最小实现——每请求同步查库 + 同步更新 last_used_at。
57/// T6 会把 last_used_at 刷新改为节流(批量/惰性)以减负,鉴权查询本身保持同步
58/// (这是认证的必要代价,无法乐观)。
59/// MCP 限流:按 token 计数的 governor 桶(与 web 的 IP-keyed 限流隔离)。
60///
61/// key 是 token_id(而非 user_id):同一用户的多个令牌各自有独立配额,
62/// 避免一个泄露的高频令牌耗尽其它令牌的额度。阈值经
63/// RATE_LIMIT_MCP_PER_SEC / RATE_LIMIT_MCP_BURST 可调(默认 10/s, burst 30)。
64static MCP_LIMITER: std::sync::LazyLock<governor::DefaultKeyedRateLimiter<String>> =
65    std::sync::LazyLock::new(|| {
66        governor::RateLimiter::keyed(
67            governor::Quota::per_second(mcp_rate_per_sec()).allow_burst(mcp_rate_burst()),
68        )
69    });
70
71fn mcp_rate_per_sec() -> std::num::NonZeroU32 {
72    let v = std::env::var("RATE_LIMIT_MCP_PER_SEC")
73        .ok()
74        .and_then(|s| s.parse::<u32>().ok())
75        .unwrap_or(10);
76    std::num::NonZeroU32::new(v.max(1)).expect("v.max(1) 保证非零")
77}
78
79fn mcp_rate_burst() -> std::num::NonZeroU32 {
80    let v = std::env::var("RATE_LIMIT_MCP_BURST")
81        .ok()
82        .and_then(|s| s.parse::<u32>().ok())
83        .unwrap_or(30);
84    std::num::NonZeroU32::new(v.max(1)).expect("v.max(1) 保证非零")
85}
86/// MCP 上传专用限流(与 MCP_LIMITER 隔离):上传是重操作(转码 + 落盘),
87/// 单独配额避免一个高频令牌的普通请求耗尽上传额度,反之亦然。
88/// 阈值经 RATE_LIMIT_MCP_UPLOAD_PER_SEC / _BURST 可调(默认 2/s, burst 5)。
89static MCP_UPLOAD_LIMITER: std::sync::LazyLock<governor::DefaultKeyedRateLimiter<String>> =
90    std::sync::LazyLock::new(|| {
91        governor::RateLimiter::keyed(
92            governor::Quota::per_second(mcp_upload_rate_per_sec())
93                .allow_burst(mcp_upload_rate_burst()),
94        )
95    });
96
97fn mcp_upload_rate_per_sec() -> std::num::NonZeroU32 {
98    let v = std::env::var("RATE_LIMIT_MCP_UPLOAD_PER_SEC")
99        .ok()
100        .and_then(|s| s.parse::<u32>().ok())
101        .unwrap_or(2);
102    std::num::NonZeroU32::new(v.max(1)).expect("v.max(1) 保证非零")
103}
104
105fn mcp_upload_rate_burst() -> std::num::NonZeroU32 {
106    let v = std::env::var("RATE_LIMIT_MCP_UPLOAD_BURST")
107        .ok()
108        .and_then(|s| s.parse::<u32>().ok())
109        .unwrap_or(5);
110    std::num::NonZeroU32::new(v.max(1)).expect("v.max(1) 保证非零")
111}
112
113/// bearer 端点专用:按 token_id 检查上传配额。超限返回 Err(脱敏消息)。
114pub(crate) fn check_mcp_upload_limit(token_id: &str) -> Result<(), &'static str> {
115    // governor 的 keyed 限流器要求 &String 键;token_id 是 &str,转一次。
116    let key = token_id.to_string();
117    if MCP_UPLOAD_LIMITER.check_key(&key).is_err() {
118        return Err("上传过于频繁,请稍后再试");
119    }
120    Ok(())
121}
122
123/// `last_used_at` 刷新节流:同一令牌的 UPDATE 至少间隔 60s,避免高频请求
124/// 每次都写库。窗口外才刷新,窗口内跳过(best-effort,失败静默)。
125const LAST_USED_REFRESH_SECS: i64 = 60;
126
127pub async fn mcp_auth_middleware(
128    mut req: Request<Body>,
129    next: Next,
130) -> Result<Response, StatusCode> {
131    let token = match extract_bearer(req.headers()) {
132        Some(t) => t,
133        None => return Err(StatusCode::UNAUTHORIZED),
134    };
135    let principal = resolve_principal(&token)
136        .await
137        .ok_or(StatusCode::UNAUTHORIZED)?;
138
139    // 限流:按 token_id 计数。超限返回 429(Too Many Requests)。
140    if MCP_LIMITER.check_key(&principal.token_id).is_err() {
141        return Err(StatusCode::TOO_MANY_REQUESTS);
142    }
143
144    // 审计:记录每次已认证的 MCP 请求(token_id + scope + user_id)。
145    // MCP 令牌是自主 AI 客户端,记录其活动有安全价值(事后可追溯滥用)。
146    tracing::info!(
147        user_id = principal.user_id,
148        scope = principal.scope.as_str(),
149        token_id = %principal.token_id,
150        "mcp request authenticated"
151    );
152
153    req.extensions_mut().insert(principal);
154    Ok(next.run(req).await)
155}
156
157/// 查库解析 token → 主体;未撤销、未过期才返回 Some。
158pub(crate) async fn resolve_principal(token: &str) -> Option<McpPrincipal> {
159    let hash = hash_token(token);
160    let client = get_conn().await.ok()?;
161    // 一次查询取出 + 校验所有条件,并带出 last_used_at 供节流判断。
162    let row = client
163        .query_opt(
164            "SELECT id, user_id, scope, expires_at, revoked_at, last_used_at
165             FROM mcp_tokens
166             WHERE token_hash = $1
167               AND revoked_at IS NULL
168               AND (expires_at IS NULL OR expires_at > NOW())",
169            &[&hash],
170        )
171        .await
172        .ok()??;
173
174    let token_id: uuid::Uuid = row.get(0);
175    let user_id: i32 = row.get(1);
176    let scope_str: &str = row.get(2);
177    let scope = TokenScope::from_db(scope_str)?;
178    let last_used: Option<chrono::DateTime<chrono::Utc>> = row.get(5);
179
180    // 节流刷新 last_used_at:仅在 NULL 或距上次 ≥ 60s 时写库,避免高频请求每次 UPDATE。
181    let needs_refresh = last_used
182        .map(|t| (chrono::Utc::now() - t).num_seconds() >= LAST_USED_REFRESH_SECS)
183        .unwrap_or(true);
184    if needs_refresh {
185        let _ = client
186            .execute(
187                "UPDATE mcp_tokens SET last_used_at = NOW() WHERE id = $1",
188                &[&token_id],
189            )
190            .await;
191    }
192
193    Some(McpPrincipal {
194        user_id,
195        scope,
196        token_id: token_id.to_string(),
197    })
198}
199/// bearer 端点鉴权:从 Authorization 头解析 bearer → 查库解析主体。
200///
201/// 供 `/api/mcp/upload` 等不在 `/mcp` 中间件链上的端点复用——这些端点
202/// 自行做 bearer 鉴权(不经 rmcp,不走 McpPrincipal→extensions 注入路径)。
203/// 失败统一返回 401(不区分缺失/无效/过期,避免 token 探测)。
204pub(crate) async fn resolve_bearer_principal(
205    headers: &HeaderMap,
206) -> Result<McpPrincipal, StatusCode> {
207    let token = extract_bearer(headers).ok_or(StatusCode::UNAUTHORIZED)?;
208    resolve_principal(&token)
209        .await
210        .ok_or(StatusCode::UNAUTHORIZED)
211}