1use async_trait::async_trait;
7use dashmap::DashMap;
8use std::collections::{BTreeMap, BTreeSet};
9use std::sync::Arc;
10use std::time::{Duration, Instant};
11use uuid::Uuid;
12
13use crate::core::error::{McpError, McpResult};
14
15#[derive(Debug, Clone, PartialEq, Eq)]
17pub struct Principal {
18 pub id: String,
19 pub roles: BTreeSet<String>,
20 pub attributes: BTreeMap<String, String>,
21 pub authentication_method: Option<String>,
22}
23
24impl Principal {
25 pub fn new(id: impl Into<String>) -> Self {
26 Self {
27 id: id.into(),
28 roles: BTreeSet::new(),
29 attributes: BTreeMap::new(),
30 authentication_method: None,
31 }
32 }
33
34 pub fn anonymous() -> Self {
35 Self::new("anonymous")
36 }
37
38 pub fn with_role(mut self, role: impl Into<String>) -> Self {
39 self.roles.insert(role.into());
40 self
41 }
42
43 pub fn with_attribute(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
44 self.attributes.insert(key.into(), value.into());
45 self
46 }
47
48 pub fn with_authentication_method(mut self, method: impl Into<String>) -> Self {
49 self.authentication_method = Some(method.into());
50 self
51 }
52}
53
54#[derive(Debug, Clone, PartialEq, Eq)]
56pub struct RequestContext {
57 pub request_id: String,
58 pub principal: Principal,
59 pub transport: String,
60 pub peer_address: Option<String>,
61 pub metadata: BTreeMap<String, String>,
62}
63
64impl RequestContext {
65 pub fn new(principal: Principal) -> Self {
66 Self {
67 request_id: Uuid::new_v4().to_string(),
68 principal,
69 transport: "unknown".to_string(),
70 peer_address: None,
71 metadata: BTreeMap::new(),
72 }
73 }
74
75 pub fn anonymous() -> Self {
76 Self::new(Principal::anonymous())
77 }
78
79 pub fn with_transport(mut self, transport: impl Into<String>) -> Self {
80 self.transport = transport.into();
81 self
82 }
83
84 pub fn with_peer_address(mut self, address: impl Into<String>) -> Self {
85 self.peer_address = Some(address.into());
86 self
87 }
88
89 pub fn with_metadata(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
90 self.metadata.insert(key.into(), value.into());
91 self
92 }
93}
94
95#[derive(Debug, Clone, PartialEq, Eq)]
97pub struct RequestTarget {
98 pub method: String,
99 pub resource: Option<String>,
100}
101
102impl RequestTarget {
103 pub fn new(method: impl Into<String>, resource: Option<String>) -> Self {
104 Self {
105 method: method.into(),
106 resource,
107 }
108 }
109}
110
111#[async_trait]
113pub trait Authorizer: Send + Sync {
114 async fn authorize(&self, context: &RequestContext, target: &RequestTarget) -> McpResult<()>;
115}
116
117#[derive(Debug, Default)]
119pub struct AllowAllAuthorizer;
120
121#[async_trait]
122impl Authorizer for AllowAllAuthorizer {
123 async fn authorize(&self, _context: &RequestContext, _target: &RequestTarget) -> McpResult<()> {
124 Ok(())
125 }
126}
127
128#[derive(Debug, Clone, PartialEq, Eq)]
131pub struct Permission {
132 pub role: String,
133 pub method_pattern: String,
134 pub resource_pattern: Option<String>,
135}
136
137impl Permission {
138 pub fn new(role: impl Into<String>, method_pattern: impl Into<String>) -> Self {
139 Self {
140 role: role.into(),
141 method_pattern: method_pattern.into(),
142 resource_pattern: None,
143 }
144 }
145
146 pub fn for_resource(mut self, resource_pattern: impl Into<String>) -> Self {
147 self.resource_pattern = Some(resource_pattern.into());
148 self
149 }
150}
151
152#[derive(Debug, Clone, Default)]
154pub struct RbacAuthorizer {
155 permissions: Vec<Permission>,
156}
157
158impl RbacAuthorizer {
159 pub fn new(permissions: impl IntoIterator<Item = Permission>) -> Self {
160 Self {
161 permissions: permissions.into_iter().collect(),
162 }
163 }
164
165 pub fn allow(mut self, permission: Permission) -> Self {
166 self.permissions.push(permission);
167 self
168 }
169}
170
171fn pattern_matches(pattern: &str, value: &str) -> bool {
172 if pattern == "*" {
173 return true;
174 }
175 pattern
176 .strip_suffix('*')
177 .map_or(pattern == value, |prefix| value.starts_with(prefix))
178}
179
180#[async_trait]
181impl Authorizer for RbacAuthorizer {
182 async fn authorize(&self, context: &RequestContext, target: &RequestTarget) -> McpResult<()> {
183 let allowed = self.permissions.iter().any(|permission| {
184 context.principal.roles.contains(&permission.role)
185 && pattern_matches(&permission.method_pattern, &target.method)
186 && match (&permission.resource_pattern, &target.resource) {
187 (None, _) => true,
188 (Some(pattern), Some(resource)) => pattern_matches(pattern, resource),
189 (Some(_), None) => false,
190 }
191 });
192
193 if allowed {
194 Ok(())
195 } else {
196 Err(McpError::Forbidden(format!(
197 "principal '{}' cannot access method '{}'{}",
198 context.principal.id,
199 target.method,
200 target
201 .resource
202 .as_ref()
203 .map(|resource| format!(" resource '{resource}'"))
204 .unwrap_or_default()
205 )))
206 }
207 }
208}
209
210#[derive(Debug, Clone)]
212pub struct RateLimitConfig {
213 pub burst: u32,
214 pub requests_per_second: f64,
215 pub idle_entry_ttl: Duration,
216}
217
218impl RateLimitConfig {
219 pub fn new(burst: u32, requests_per_second: f64) -> McpResult<Self> {
220 if burst == 0 || !requests_per_second.is_finite() || requests_per_second <= 0.0 {
221 return Err(McpError::Validation(
222 "rate limit requires burst > 0 and requests_per_second > 0".to_string(),
223 ));
224 }
225 Ok(Self {
226 burst,
227 requests_per_second,
228 idle_entry_ttl: Duration::from_secs(600),
229 })
230 }
231}
232
233#[derive(Debug)]
234struct Bucket {
235 tokens: f64,
236 updated_at: Instant,
237}
238
239#[derive(Debug)]
241pub struct RateLimiter {
242 config: RateLimitConfig,
243 buckets: DashMap<String, Bucket>,
244}
245
246impl RateLimiter {
247 pub fn new(config: RateLimitConfig) -> Self {
248 Self {
249 config,
250 buckets: DashMap::new(),
251 }
252 }
253
254 pub fn check(&self, context: &RequestContext, target: &RequestTarget) -> McpResult<()> {
255 let now = Instant::now();
256 let key = format!("{}\u{1f}{}", context.principal.id, target.method);
257 let mut bucket = self.buckets.entry(key).or_insert_with(|| Bucket {
258 tokens: f64::from(self.config.burst),
259 updated_at: now,
260 });
261
262 let elapsed = now.duration_since(bucket.updated_at).as_secs_f64();
263 bucket.tokens = (bucket.tokens + elapsed * self.config.requests_per_second)
264 .min(f64::from(self.config.burst));
265 bucket.updated_at = now;
266
267 if bucket.tokens >= 1.0 {
268 bucket.tokens -= 1.0;
269 Ok(())
270 } else {
271 let retry_after = (1.0 - bucket.tokens) / self.config.requests_per_second;
272 Err(McpError::RateLimited {
273 retry_after_ms: (retry_after * 1000.0).ceil() as u64,
274 })
275 }
276 }
277
278 pub fn prune_idle(&self) {
281 let now = Instant::now();
282 self.buckets
283 .retain(|_, bucket| now.duration_since(bucket.updated_at) < self.config.idle_entry_ttl);
284 }
285}
286
287#[derive(Clone)]
289pub struct RequestPolicy {
290 authorizer: Arc<dyn Authorizer>,
291 rate_limiter: Option<Arc<RateLimiter>>,
292}
293
294impl Default for RequestPolicy {
295 fn default() -> Self {
296 Self {
297 authorizer: Arc::new(AllowAllAuthorizer),
298 rate_limiter: None,
299 }
300 }
301}
302
303impl RequestPolicy {
304 pub fn new(authorizer: impl Authorizer + 'static) -> Self {
305 Self {
306 authorizer: Arc::new(authorizer),
307 rate_limiter: None,
308 }
309 }
310
311 pub fn with_rate_limiter(mut self, limiter: RateLimiter) -> Self {
312 self.rate_limiter = Some(Arc::new(limiter));
313 self
314 }
315
316 pub async fn enforce(&self, context: &RequestContext, target: &RequestTarget) -> McpResult<()> {
317 self.authorizer.authorize(context, target).await?;
318 if let Some(limiter) = &self.rate_limiter {
319 limiter.check(context, target)?;
320 }
321 Ok(())
322 }
323}
324
325#[cfg(test)]
326mod tests {
327 use super::*;
328
329 #[tokio::test]
330 async fn rbac_is_deny_by_default_and_checks_resources() {
331 let policy = RequestPolicy::new(RbacAuthorizer::new([Permission::new(
332 "reader",
333 "resources/read",
334 )
335 .for_resource("urn:public:*")]));
336 let context = RequestContext::new(Principal::new("alice").with_role("reader"));
337
338 assert!(policy
339 .enforce(
340 &context,
341 &RequestTarget::new("resources/read", Some("urn:public:1".to_string()))
342 )
343 .await
344 .is_ok());
345 assert!(matches!(
346 policy
347 .enforce(
348 &context,
349 &RequestTarget::new("resources/read", Some("urn:private:1".to_string()))
350 )
351 .await,
352 Err(McpError::Forbidden(_))
353 ));
354 }
355
356 #[tokio::test]
357 async fn rate_limit_is_enforced_per_principal_and_method() {
358 let limiter = RateLimiter::new(RateLimitConfig::new(1, 0.01).unwrap());
359 let policy = RequestPolicy::default().with_rate_limiter(limiter);
360 let context = RequestContext::new(Principal::new("alice"));
361 let target = RequestTarget::new("tools/list", None);
362
363 assert!(policy.enforce(&context, &target).await.is_ok());
364 assert!(matches!(
365 policy.enforce(&context, &target).await,
366 Err(McpError::RateLimited { .. })
367 ));
368 assert!(policy
369 .enforce(&context, &RequestTarget::new("ping", None))
370 .await
371 .is_ok());
372 }
373}