Skip to main content

prism_mcp_rs/
security.rs

1//! Request identity, authorization, and rate-limiting primitives.
2//!
3//! These controls live above the transport layer so the same policy is applied
4//! to STDIO, HTTP, WebSocket, and custom transports.
5
6use 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/// Authenticated identity attached to an MCP request.
16#[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/// Context shared by validation, policy, observability, and request handlers.
55#[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/// Normalized target used by authorization policies.
96#[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/// Authorization decision provider.
112#[async_trait]
113pub trait Authorizer: Send + Sync {
114    async fn authorize(&self, context: &RequestContext, target: &RequestTarget) -> McpResult<()>;
115}
116
117/// Backwards-compatible authorizer used unless an application installs RBAC.
118#[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/// One fine-grained role permission. Patterns support exact values or a final
129/// `*`, for example `tools/*` and `urn:customer:*`.
130#[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/// Deny-by-default role-based authorizer.
153#[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/// Token-bucket limit applied independently to each principal and method.
211#[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/// Concurrent in-process token-bucket limiter.
240#[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    /// Removes stale principal/method buckets. Applications may call this from
279    /// an existing maintenance loop; request processing never scans the map.
280    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/// Shared server request policy.
288#[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}