Skip to main content

systemprompt_api/services/proxy/auth/
access.rs

1//! Access enforcement for proxied MCP and agent requests.
2//!
3//! [`AccessValidator`] resolves whether a service requires OAuth, validates the
4//! caller's bearer token and scopes, and either returns the authenticated user
5//! or converts the failure into an RFC 9728 challenge. For MCP it permits a
6//! session-only fallback when a prior authenticated initialize established the
7//! identity in the proxy cache.
8//!
9//! Copyright (c) systemprompt.io — Business Source License 1.1.
10//! See <https://systemprompt.io> for licensing details.
11
12use axum::http::header::AUTHORIZATION;
13use axum::http::{HeaderMap, StatusCode};
14use std::str::FromStr;
15
16use crate::services::proxy::backend::ProxyError;
17use systemprompt_agent::services::AgentRegistryProviderService;
18use systemprompt_database::ServiceConfig;
19use systemprompt_models::RequestContext;
20use systemprompt_models::auth::{AuthenticatedUser, Permission};
21use systemprompt_models::modules::ApiPaths;
22use systemprompt_oauth::services::AuthService;
23use systemprompt_runtime::AppContext;
24use systemprompt_traits::{AgentRegistryProvider, McpRegistryProvider};
25
26use super::challenge::{AuthValidator, ChallengeRequest, challenge_or_error};
27
28#[derive(Debug)]
29pub struct OAuthRequirement {
30    pub module: String,
31    pub required: bool,
32    pub scopes: Vec<String>,
33    pub audience: String,
34}
35
36#[derive(Debug, Clone, Copy)]
37pub struct AccessValidator;
38
39impl AccessValidator {
40    pub(crate) async fn validate(
41        headers: &HeaderMap,
42        service_name: &str,
43        service: &ServiceConfig,
44        ctx: &AppContext,
45        req_context: Option<&RequestContext>,
46    ) -> Result<Option<AuthenticatedUser>, ProxyError> {
47        let requirement = lookup_oauth_requirement(service, service_name, ctx).await?;
48        Self::validate_with_requirement(headers, service_name, &requirement, ctx, req_context)
49    }
50
51    pub fn validate_with_requirement(
52        headers: &HeaderMap,
53        service_name: &str,
54        requirement: &OAuthRequirement,
55        ctx: &AppContext,
56        req_context: Option<&RequestContext>,
57    ) -> Result<Option<AuthenticatedUser>, ProxyError> {
58        if !requirement.required {
59            return Ok(None);
60        }
61        let resource_path = resource_path_for(&requirement.module, service_name);
62        let has_authorization = headers.get(AUTHORIZATION).is_some();
63        let challenge = |status_code: StatusCode| {
64            challenge_or_error(&ChallengeRequest {
65                service_name,
66                resource_path: &resource_path,
67                headers,
68                ctx,
69                status_code,
70                has_authorization,
71            })
72        };
73        let authenticated_user =
74            match AuthValidator::validate_service_access(headers, service_name, req_context) {
75                Ok(user) => user,
76                Err(status_code) => {
77                    if let Some(outcome) = mcp_session_fallback(
78                        &requirement.module,
79                        service_name,
80                        headers,
81                        status_code,
82                    ) {
83                        return outcome;
84                    }
85                    return Err(challenge(status_code));
86                },
87            };
88        if let Err(status_code) =
89            enforce_required_audience(headers, service_name, &requirement.audience)
90        {
91            return Err(challenge(status_code));
92        }
93        ensure_required_scopes(service_name, &requirement.scopes, &authenticated_user)?;
94        Ok(Some(authenticated_user))
95    }
96}
97
98async fn lookup_oauth_requirement(
99    service: &ServiceConfig,
100    service_name: &str,
101    ctx: &AppContext,
102) -> Result<OAuthRequirement, ProxyError> {
103    if service.module_name == "agent" {
104        let registry =
105            AgentRegistryProviderService::new().map_err(|e| ProxyError::ServiceNotRunning {
106                service: service_name.to_owned(),
107                status: format!("Failed to load agent registry: {e}"),
108            })?;
109        let info =
110            registry
111                .get_agent(service_name)
112                .await
113                .map_err(|e| ProxyError::ServiceNotFound {
114                    service: format!("Agent '{}' not found in registry: {}", service_name, e),
115                })?;
116        Ok(OAuthRequirement {
117            module: "agent".to_owned(),
118            required: info.oauth.required,
119            scopes: info.oauth.scopes,
120            audience: info.oauth.audience,
121        })
122    } else if service.module_name == "mcp" {
123        mcp_oauth_requirement(ctx, service_name).await
124    } else {
125        Ok(OAuthRequirement {
126            module: service.module_name.clone(),
127            required: true,
128            scopes: vec![],
129            audience: String::new(),
130        })
131    }
132}
133
134pub(crate) async fn mcp_oauth_requirement(
135    ctx: &AppContext,
136    service_name: &str,
137) -> Result<OAuthRequirement, ProxyError> {
138    let registry = ctx.mcp_registry();
139    registry
140        .validate()
141        .map_err(|e| ProxyError::ServiceNotRunning {
142            service: service_name.to_owned(),
143            status: format!("Failed to load MCP registry: {e}"),
144        })?;
145    let info = McpRegistryProvider::get_server(registry, service_name)
146        .await
147        .map_err(|e| ProxyError::ServiceNotFound {
148            service: format!("MCP server '{}' not found in registry: {}", service_name, e),
149        })?;
150    Ok(OAuthRequirement {
151        module: "mcp".to_owned(),
152        required: info.oauth.required,
153        scopes: info.oauth.scopes,
154        audience: info.oauth.audience,
155    })
156}
157
158fn enforce_required_audience(
159    headers: &HeaderMap,
160    service_name: &str,
161    audience: &str,
162) -> Result<(), StatusCode> {
163    if audience.is_empty() {
164        return Ok(());
165    }
166    AuthService::authorize_required_audience(headers, audience)
167        .map(|_user| ())
168        .inspect_err(|status| {
169            tracing::warn!(
170                service = %service_name,
171                audience = %audience,
172                status = %status,
173                "Token lacks the service's required audience"
174            );
175        })
176}
177
178fn resource_path_for(module_name: &str, service_name: &str) -> String {
179    match module_name {
180        "mcp" => ApiPaths::mcp_server_endpoint(service_name),
181        "agent" => ApiPaths::agent_endpoint(&systemprompt_identifiers::AgentId::new(service_name)),
182        _ => String::new(),
183    }
184}
185
186fn mcp_session_fallback(
187    module_name: &str,
188    service_name: &str,
189    headers: &HeaderMap,
190    status_code: StatusCode,
191) -> Option<Result<Option<AuthenticatedUser>, ProxyError>> {
192    if module_name != "mcp" || status_code != StatusCode::UNAUTHORIZED {
193        return None;
194    }
195    let has_session = headers
196        .get("mcp-session-id")
197        .and_then(|v| v.to_str().ok())
198        .is_some_and(|v| !v.is_empty());
199    if !has_session {
200        return None;
201    }
202    let has_bearer_token = headers
203        .get("authorization")
204        .and_then(|v| v.to_str().ok())
205        .is_some_and(|v| v.starts_with("Bearer "));
206    if has_bearer_token {
207        tracing::info!(
208            service = %service_name,
209            session_id = ?headers.get("mcp-session-id"),
210            "MCP request has expired/invalid Bearer token — returning 401 for client token refresh"
211        );
212        return None;
213    }
214    tracing::info!(
215        service = %service_name,
216        session_id = ?headers.get("mcp-session-id"),
217        "Allowing MCP request with session ID (session-based auth, identity from proxy cache)"
218    );
219    Some(Ok(None))
220}
221
222fn ensure_required_scopes(
223    service_name: &str,
224    required_scopes: &[String],
225    user: &AuthenticatedUser,
226) -> Result<(), ProxyError> {
227    if required_scopes.is_empty() {
228        return Ok(());
229    }
230    let has_required_scope = required_scopes.iter().any(|required_scope_str| {
231        Permission::from_str(required_scope_str).map_or_else(
232            |_| {
233                user.permissions
234                    .iter()
235                    .any(|p| p.as_str() == required_scope_str)
236            },
237            |required_permission| {
238                user.permissions
239                    .iter()
240                    .any(|p| *p == required_permission || p.implies(&required_permission))
241            },
242        )
243    });
244    if !has_required_scope {
245        return Err(ProxyError::Forbidden {
246            service: format!(
247                "Insufficient permissions for {}. Required: {:?}, User has: {:?}",
248                service_name, required_scopes, user.permissions
249            ),
250        });
251    }
252    Ok(())
253}