Skip to main content

pmcp/server/auth/
proxy.rs

1//! Proxy authentication provider that delegates to upstream OAuth servers.
2
3use super::traits::{AuthContext, AuthProvider, TokenValidator};
4use crate::error::{Error, ErrorCode, Result};
5use async_trait::async_trait;
6use std::future::Future;
7use std::pin::Pin;
8use std::sync::Arc;
9
10/// Token validation function type.
11pub type TokenValidatorFn =
12    Box<dyn Fn(String) -> Pin<Box<dyn Future<Output = Result<AuthContext>> + Send>> + Send + Sync>;
13
14/// Proxy provider configuration.
15#[derive(Clone, Debug)]
16pub struct ProxyProviderConfig {
17    /// Upstream OAuth server URL for token validation.
18    pub upstream_url: String,
19
20    /// Optional introspection endpoint (defaults to `{upstream_url}/introspect`).
21    pub introspection_endpoint: Option<String>,
22
23    /// Client ID for introspection requests.
24    pub client_id: Option<String>,
25
26    /// Client secret for introspection requests.
27    pub client_secret: Option<String>,
28
29    /// Whether to cache token validation results.
30    pub enable_cache: bool,
31
32    /// Cache TTL in seconds (default 300).
33    pub cache_ttl: u64,
34}
35
36impl Default for ProxyProviderConfig {
37    fn default() -> Self {
38        Self {
39            upstream_url: String::new(),
40            introspection_endpoint: None,
41            client_id: None,
42            client_secret: None,
43            enable_cache: true,
44            cache_ttl: 300,
45        }
46    }
47}
48
49/// Proxy authentication provider that delegates to an upstream OAuth server.
50/// This is similar to the TypeScript SDK's `ProxyProvider` pattern.
51pub struct ProxyProvider {
52    config: ProxyProviderConfig,
53    token_validator: Option<TokenValidatorFn>,
54    validator: Option<Arc<dyn TokenValidator>>,
55}
56
57impl std::fmt::Debug for ProxyProvider {
58    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
59        f.debug_struct("ProxyProvider")
60            .field("config", &self.config)
61            .field("token_validator", &self.token_validator.is_some())
62            .field("validator", &self.validator.is_some())
63            .finish()
64    }
65}
66
67impl ProxyProvider {
68    /// Create a new proxy provider with the given configuration.
69    pub fn new(config: ProxyProviderConfig) -> Self {
70        Self {
71            config,
72            token_validator: None,
73            validator: None,
74        }
75    }
76
77    /// Create a proxy provider with just an upstream URL.
78    pub fn with_upstream(upstream_url: impl Into<String>) -> Self {
79        Self::new(ProxyProviderConfig {
80            upstream_url: upstream_url.into(),
81            ..Default::default()
82        })
83    }
84
85    /// Set a custom token validation function.
86    /// This allows applications to implement custom validation logic.
87    pub fn with_validator_fn<F, Fut>(mut self, validator: F) -> Self
88    where
89        F: Fn(String) -> Fut + Send + Sync + 'static,
90        Fut: Future<Output = Result<AuthContext>> + Send + 'static,
91    {
92        self.token_validator = Some(Box::new(move |token| Box::pin(validator(token))));
93        self
94    }
95
96    /// Set a token validator implementation.
97    pub fn with_validator(mut self, validator: Arc<dyn TokenValidator>) -> Self {
98        self.validator = Some(validator);
99        self
100    }
101
102    /// Set the introspection endpoint.
103    pub fn introspection_endpoint(mut self, endpoint: impl Into<String>) -> Self {
104        self.config.introspection_endpoint = Some(endpoint.into());
105        self
106    }
107
108    /// Set client credentials for introspection.
109    pub fn client_credentials(
110        mut self,
111        client_id: impl Into<String>,
112        client_secret: impl Into<String>,
113    ) -> Self {
114        self.config.client_id = Some(client_id.into());
115        self.config.client_secret = Some(client_secret.into());
116        self
117    }
118
119    /// Enable or disable token caching.
120    pub fn cache(mut self, enable: bool) -> Self {
121        self.config.enable_cache = enable;
122        self
123    }
124
125    /// Extract bearer token from authorization header.
126    fn extract_bearer_token(authorization_header: Option<&str>) -> Option<String> {
127        authorization_header?
128            .strip_prefix("Bearer ")
129            .map(|s| s.to_string())
130    }
131
132    /// Validate token using the configured method.
133    async fn validate_token_internal(&self, token: String) -> Result<AuthContext> {
134        // Use custom validator function if provided
135        if let Some(ref validator_fn) = self.token_validator {
136            return validator_fn(token).await;
137        }
138
139        // Use validator implementation if provided
140        if let Some(ref validator) = self.validator {
141            return validator.validate(&token).await;
142        }
143
144        // Fall back to introspection endpoint
145        self.introspect_token(token).await
146    }
147
148    /// Introspect token using the upstream server.
149    async fn introspect_token(&self, _token: String) -> Result<AuthContext> {
150        // This would make an HTTP request to the introspection endpoint
151        // For now, return a placeholder implementation
152        // Real implementation would use reqwest or similar HTTP client
153        //
154        // The implementation would:
155        // 1. POST to introspection_endpoint with token
156        // 2. Include client credentials if configured
157        // 3. Parse the introspection response
158        // 4. Convert to AuthContext
159
160        Err(Error::protocol(
161            ErrorCode::METHOD_NOT_FOUND,
162            "Token introspection not yet implemented. Please provide a custom validator.",
163        ))
164    }
165}
166
167#[async_trait]
168impl AuthProvider for ProxyProvider {
169    async fn validate_request(
170        &self,
171        authorization_header: Option<&str>,
172    ) -> Result<Option<AuthContext>> {
173        // Extract bearer token from Authorization header
174        let Some(token) = Self::extract_bearer_token(authorization_header) else {
175            return Ok(None); // No auth provided
176        };
177
178        // Validate the token
179        match self.validate_token_internal(token).await {
180            Ok(auth_context) => {
181                // Check if token is expired
182                if auth_context.is_expired() {
183                    return Err(Error::protocol(ErrorCode::INVALID_REQUEST, "Token expired"));
184                }
185                Ok(Some(auth_context))
186            },
187            Err(e) => Err(e),
188        }
189    }
190
191    fn auth_scheme(&self) -> &'static str {
192        "Bearer"
193    }
194}
195
196#[async_trait]
197impl TokenValidator for ProxyProvider {
198    async fn validate(&self, token: &str) -> Result<AuthContext> {
199        self.validate_token_internal(token.to_string()).await
200    }
201}
202
203/// No-op authentication provider for development/testing.
204#[derive(Debug, Clone)]
205pub struct NoOpAuthProvider;
206
207#[async_trait]
208impl AuthProvider for NoOpAuthProvider {
209    async fn validate_request(
210        &self,
211        _authorization_header: Option<&str>,
212    ) -> Result<Option<AuthContext>> {
213        // Always return a valid auth context for development with all common scopes
214        Ok(Some(AuthContext {
215            subject: "dev-user".to_string(),
216            scopes: vec![
217                "read".to_string(),
218                "write".to_string(),
219                "admin".to_string(),
220                "mcp:tools:use".to_string(),
221            ],
222            claims: Default::default(),
223            token: None,
224            client_id: Some("dev-client".to_string()),
225            expires_at: None,
226            authenticated: true,
227        }))
228    }
229
230    fn is_required(&self) -> bool {
231        false // Auth not required in dev mode
232    }
233}
234
235/// Optional authentication provider that makes auth optional.
236#[derive(Debug)]
237pub struct OptionalAuthProvider<P: AuthProvider> {
238    inner: P,
239}
240
241impl<P: AuthProvider> OptionalAuthProvider<P> {
242    /// Wrap an auth provider to make authentication optional.
243    pub fn new(provider: P) -> Self {
244        Self { inner: provider }
245    }
246}
247
248#[async_trait]
249impl<P: AuthProvider> AuthProvider for OptionalAuthProvider<P> {
250    async fn validate_request(
251        &self,
252        authorization_header: Option<&str>,
253    ) -> Result<Option<AuthContext>> {
254        // Try to validate, but don't fail if no auth is provided
255        self.inner.validate_request(authorization_header).await
256    }
257
258    fn auth_scheme(&self) -> &'static str {
259        self.inner.auth_scheme()
260    }
261
262    fn is_required(&self) -> bool {
263        false // Make auth optional
264    }
265}