1use 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
22pub const TOKEN_PREFIX: &str = "ygg_";
24
25#[derive(Clone, Debug)]
27pub struct McpPrincipal {
28 pub user_id: i32,
29 pub scope: TokenScope,
30 pub token_id: String,
32}
33
34pub 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
49pub fn hash_token(token: &str) -> String {
51 crate::utils::server::sha256_hex(token)
52}
53
54static 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}
86static 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
113pub(crate) fn check_mcp_upload_limit(token_id: &str) -> Result<(), &'static str> {
115 let key = token_id.to_string();
117 if MCP_UPLOAD_LIMITER.check_key(&key).is_err() {
118 return Err("上传过于频繁,请稍后再试");
119 }
120 Ok(())
121}
122
123const 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 if MCP_LIMITER.check_key(&principal.token_id).is_err() {
141 return Err(StatusCode::TOO_MANY_REQUESTS);
142 }
143
144 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
157pub(crate) async fn resolve_principal(token: &str) -> Option<McpPrincipal> {
159 let hash = hash_token(token);
160 let client = get_conn().await.ok()?;
161 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 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}
199pub(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}