pmcp/server/auth/
proxy.rs1use 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
10pub type TokenValidatorFn =
12 Box<dyn Fn(String) -> Pin<Box<dyn Future<Output = Result<AuthContext>> + Send>> + Send + Sync>;
13
14#[derive(Clone, Debug)]
16pub struct ProxyProviderConfig {
17 pub upstream_url: String,
19
20 pub introspection_endpoint: Option<String>,
22
23 pub client_id: Option<String>,
25
26 pub client_secret: Option<String>,
28
29 pub enable_cache: bool,
31
32 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
49pub 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 pub fn new(config: ProxyProviderConfig) -> Self {
70 Self {
71 config,
72 token_validator: None,
73 validator: None,
74 }
75 }
76
77 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 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 pub fn with_validator(mut self, validator: Arc<dyn TokenValidator>) -> Self {
98 self.validator = Some(validator);
99 self
100 }
101
102 pub fn introspection_endpoint(mut self, endpoint: impl Into<String>) -> Self {
104 self.config.introspection_endpoint = Some(endpoint.into());
105 self
106 }
107
108 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 pub fn cache(mut self, enable: bool) -> Self {
121 self.config.enable_cache = enable;
122 self
123 }
124
125 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 async fn validate_token_internal(&self, token: String) -> Result<AuthContext> {
134 if let Some(ref validator_fn) = self.token_validator {
136 return validator_fn(token).await;
137 }
138
139 if let Some(ref validator) = self.validator {
141 return validator.validate(&token).await;
142 }
143
144 self.introspect_token(token).await
146 }
147
148 async fn introspect_token(&self, _token: String) -> Result<AuthContext> {
150 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 let Some(token) = Self::extract_bearer_token(authorization_header) else {
175 return Ok(None); };
177
178 match self.validate_token_internal(token).await {
180 Ok(auth_context) => {
181 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#[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 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 }
233}
234
235#[derive(Debug)]
237pub struct OptionalAuthProvider<P: AuthProvider> {
238 inner: P,
239}
240
241impl<P: AuthProvider> OptionalAuthProvider<P> {
242 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 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 }
265}