systemprompt_api/services/proxy/auth/
access.rs1use 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}