1use std::sync::Arc;
16use std::time::Duration as StdDuration;
17
18use axum::{
19 body::Body,
20 extract::{Query, State},
21 http::{header, StatusCode},
22 response::{IntoResponse, Json, Redirect, Response},
23};
24use axum_extra::extract::cookie::{Cookie, SameSite};
25use pep::oidc_client::OidcClient;
26use pep::oidc_resource_server::ResourceServerClient;
27use pep::oidc::pkce_cookie::PkceCookieManager;
28use pep::session_manager::WebSessionManager;
29use pep::{DevConfig, JwtClaims, JwtValidationOptions, OidcClientConfig};
30use serde::Deserialize;
31use time::Duration as TimeDuration;
32
33use cedar_policy::{Context, Entities, EntityUid, Request};
34use std::str::FromStr;
35
36#[derive(Debug, Clone)]
42pub struct AuthConfig {
43 pub issuer_url: String,
45 pub client_id: String,
47 pub client_secret: Option<String>,
49 pub redirect_uri: String,
51 pub scope: String,
53 pub cookie_name: String,
55 pub dev_config: DevConfig,
57 pub validation_options: JwtValidationOptions,
59 pub pkce_cookie_secret: String,
61}
62
63impl AuthConfig {
64 pub fn from_toml(config_toml: &str) -> Option<Self> {
69 let table: toml::Table = toml::from_str(config_toml).ok()?;
70
71 let dev_config = table.get("dev").and_then(|d| d.as_table()).map(|d| {
73 DevConfig {
74 local_dev_mode: d.get("local_dev_mode").and_then(|v| v.as_bool()).unwrap_or(false),
75 local_dev_email: d.get("local_dev_email").and_then(|v| v.as_str()).map(String::from),
76 local_dev_name: d.get("local_dev_name").and_then(|v| v.as_str()).map(String::from),
77 local_dev_username: d.get("local_dev_username").and_then(|v| v.as_str()).map(String::from),
78 }
79 });
80
81 if let Some(ref dc) = dev_config {
83 if dc.local_dev_mode {
84 let oidc = Self::parse_oidc_section(&table);
86 return Some(Self {
87 issuer_url: oidc.as_ref().map(|o| o.0.clone()).unwrap_or_else(|| "https://auth.example.com".into()),
88 client_id: oidc.as_ref().map(|o| o.1.clone()).unwrap_or_else(|| "trustee".into()),
89 client_secret: oidc.as_ref().and_then(|o| o.2.clone()),
90 redirect_uri: oidc.as_ref().map(|o| o.3.clone()).unwrap_or_else(|| "http://localhost:3000/auth/callback".into()),
91 scope: oidc.as_ref().map(|o| o.4.clone()).unwrap_or_else(|| "openid profile email".into()),
92 cookie_name: "trustee_token".into(),
93 dev_config: dc.clone(),
94 validation_options: JwtValidationOptions::default(),
95 pkce_cookie_secret: oidc.as_ref().map(|o| o.6.clone()).unwrap_or_else(|| "trustee-default-pkce-secret-change-me".into()),
96 });
97 }
98 }
99
100 let (issuer_url, client_id, client_secret, redirect_uri, scope, validation_options, pkce_secret) =
102 Self::parse_oidc_section(&table)?;
103
104 Some(Self {
105 issuer_url,
106 client_id,
107 client_secret,
108 redirect_uri,
109 scope,
110 cookie_name: "trustee_token".into(),
111 dev_config: dev_config.unwrap_or_default(),
112 validation_options,
113 pkce_cookie_secret: pkce_secret,
114 })
115 }
116
117 fn parse_oidc_section(
120 table: &toml::Table,
121 ) -> Option<(String, String, Option<String>, String, String, JwtValidationOptions, String)> {
122 let oidc = table.get("oidc")?.as_table()?;
123
124 let issuer_url = oidc.get("issuer_url")?.as_str()?.to_string();
125 let client_id = oidc.get("client_id")?.as_str()?.to_string();
126 let client_secret = oidc.get("client_secret").and_then(|v| v.as_str()).map(String::from);
127 let redirect_uri = oidc
128 .get("redirect_uri")
129 .or_else(|| oidc.get("redirect_url")) .and_then(|v| v.as_str())
131 .unwrap_or("http://localhost:3000/auth/callback")
132 .to_string();
133 let scope = oidc
134 .get("scope")
135 .and_then(|v| v.as_str())
136 .unwrap_or("openid profile email")
137 .to_string();
138
139 let mut validation_options = JwtValidationOptions::default();
140 if let Some(skip) = oidc.get("skip_issuer_validation").and_then(|v| v.as_bool()) {
141 validation_options.skip_issuer_validation = skip;
142 }
143 if let Some(skip) = oidc.get("skip_audience_validation").and_then(|v| v.as_bool()) {
144 validation_options.skip_audience_validation = skip;
145 }
146 validation_options.expected_audience = oidc
147 .get("expected_audience")
148 .and_then(|v| v.as_str())
149 .map(String::from);
150
151 let pkce_secret = oidc
152 .get("pkce_cookie_secret")
153 .and_then(|v| v.as_str())
154 .unwrap_or("trustee-default-pkce-secret-change-me")
155 .to_string();
156
157 Some((issuer_url, client_id, client_secret, redirect_uri, scope, validation_options, pkce_secret))
158 }
159
160 pub fn oidc_client_config(&self) -> OidcClientConfig {
162 OidcClientConfig {
163 issuer_url: self.issuer_url.clone(),
164 client_id: self.client_id.clone(),
165 client_secret: self.client_secret.clone(),
166 redirect_uri: self.redirect_uri.clone(),
167 scope: self.scope.clone(),
168 code_challenge_method: "S256".to_string(),
169 }
170 }
171}
172
173#[derive(Clone)]
175pub struct AuthState {
176 pub oidc_client: OidcClient,
178 pub resource_server: ResourceServerClient,
180 pub client_config: OidcClientConfig,
182 pub config: AuthConfig,
184 pub pkce_manager: PkceCookieManager,
186 pub session_manager: Arc<WebSessionManager>,
188 pub cedar_authorizer: Option<Arc<pep::cedar::CedarAuthorizer>>,
190}
191
192impl AuthState {
193 pub fn new(config: AuthConfig) -> Self {
195 Self::with_cedar(config, None)
196 }
197
198 pub fn with_cedar(config: AuthConfig, cedar_authorizer: Option<Arc<pep::cedar::CedarAuthorizer>>) -> Self {
200 let pkce_manager = PkceCookieManager::new(
201 config.pkce_cookie_secret.as_bytes(),
202 "trustee_pkce_state",
203 StdDuration::from_secs(600),
204 );
205
206 let session_manager = Arc::new(WebSessionManager::new(
207 OidcClient::new(),
208 config.issuer_url.clone(),
209 config.client_id.clone(),
210 config.client_secret.clone(),
211 config.scope.clone(),
212 ));
213
214 Self {
215 oidc_client: OidcClient::new(),
216 resource_server: ResourceServerClient::new(),
217 client_config: config.oidc_client_config(),
218 pkce_manager,
219 session_manager,
220 config,
221 cedar_authorizer,
222 }
223 }
224
225 pub fn is_dev_mode(&self) -> bool {
227 self.config.dev_config.local_dev_mode
228 }
229
230 pub async fn validate_token(&self, token: &str) -> anyhow::Result<JwtClaims> {
232 let mut claims = self
233 .resource_server
234 .validate_jwt_with_options(
235 token,
236 &self.config.issuer_url,
237 &self.config.client_id,
238 &self.config.validation_options,
239 )
240 .await
241 .map_err(|e| anyhow::anyhow!("Token validation failed: {}", e))?;
242
243 let _ = self
245 .resource_server
246 .enrich_claims_with_userinfo(&mut claims, token, &self.config.issuer_url, None)
247 .await;
248
249 if claims.name.is_none() || claims.email.is_none() {
252 self.fill_userinfo_fields(&mut claims, token).await;
253 }
254
255 Ok(claims)
256 }
257
258 async fn fill_userinfo_fields(&self, claims: &mut JwtClaims, token: &str) {
261 let userinfo_url = format!("{}/userinfo", self.config.issuer_url.trim_end_matches('/'));
265
266 let client = reqwest::Client::new();
267 let resp = client
268 .get(&userinfo_url)
269 .header("Authorization", format!("Bearer {}", token))
270 .header("Accept", "application/json")
271 .send()
272 .await;
273
274 let Ok(resp) = resp else {
275 tracing::debug!("Userinfo request failed for name/email enrichment");
276 return;
277 };
278
279 if !resp.status().is_success() {
280 tracing::debug!("Userinfo returned {} for name/email enrichment", resp.status());
281 return;
282 }
283
284 let Ok(userinfo): Result<serde_json::Map<String, serde_json::Value>, _> = resp.json().await else {
285 return;
286 };
287
288 tracing::debug!("Userinfo keys: {:?}", userinfo.keys().collect::<Vec<_>>());
289
290 if claims.name.is_none() {
291 if let Some(name) = userinfo.get("name").and_then(|v| v.as_str()) {
292 claims.name = Some(name.to_string());
293 }
294 }
295 if claims.email.is_none() {
296 if let Some(email) = userinfo.get("email").and_then(|v| v.as_str()) {
297 claims.email = Some(email.to_string());
298 }
299 }
300 if claims.preferred_username.is_none() {
301 if let Some(uname) = userinfo.get("preferred_username").and_then(|v| v.as_str()) {
302 claims.preferred_username = Some(uname.to_string());
303 }
304 }
305 }
306
307 fn check_cedar_authorized(&self, claims: &JwtClaims) -> Result<(), ()> {
312 let Some(ref authorizer) = self.cedar_authorizer else {
313 return Ok(()); };
315
316 let principal_entity = match pep::cedar::build_principal_entity(claims) {
318 Ok(e) => e,
319 Err(e) => {
320 tracing::error!("Cedar: failed to build principal entity: {}", e);
321 return Err(());
322 }
323 };
324
325 let mut entities_vec = vec![principal_entity];
327
328 let app_uid = match EntityUid::from_str(r#"TrusteeApp::"default""#) {
330 Ok(uid) => uid,
331 Err(e) => {
332 tracing::error!("Cedar: failed to build TrusteeApp uid: {}", e);
333 return Err(());
334 }
335 };
336 let app_entity = match cedar_policy::Entity::new(
337 app_uid,
338 std::collections::HashMap::new(),
339 std::collections::HashSet::new(),
340 ) {
341 Ok(e) => e,
342 Err(e) => {
343 tracing::error!("Cedar: failed to build TrusteeApp entity: {}", e);
344 return Err(());
345 }
346 };
347 entities_vec.push(app_entity);
348
349 let entities = match Entities::from_entities(entities_vec, None) {
350 Ok(e) => e,
351 Err(e) => {
352 tracing::error!("Cedar: failed to build entities set: {}", e);
353 return Err(());
354 }
355 };
356
357 let principal_uid = match pep::cedar::build_principal_uid(claims) {
359 Ok(uid) => uid,
360 Err(e) => {
361 tracing::error!("Cedar: failed to build principal uid: {}", e);
362 return Err(());
363 }
364 };
365
366 let action_uid = match EntityUid::from_str(r#"Action::"Access""#) {
367 Ok(uid) => uid,
368 Err(e) => {
369 tracing::error!("Cedar: failed to build action uid: {}", e);
370 return Err(());
371 }
372 };
373
374 let resource_uid = match EntityUid::from_str(r#"TrusteeApp::"default""#) {
375 Ok(uid) => uid,
376 Err(e) => {
377 tracing::error!("Cedar: failed to build resource uid: {}", e);
378 return Err(());
379 }
380 };
381
382 let request = match Request::new(principal_uid, action_uid, resource_uid, Context::empty(), None) {
383 Ok(r) => r,
384 Err(e) => {
385 tracing::error!("Cedar: failed to build request: {}", e);
386 return Err(());
387 }
388 };
389
390 let response = authorizer.is_allowed_with_entities(&request, &entities);
391
392 if response.allowed() {
393 tracing::debug!(
394 "Cedar: authorized user {} (sub={})",
395 claims.email.as_deref().unwrap_or("unknown"),
396 claims.sub
397 );
398 Ok(())
399 } else {
400 tracing::warn!(
401 "Cedar: DENIED user {} (sub={}) — matched policies: {:?}, errors: {:?}",
402 claims.email.as_deref().unwrap_or("unknown"),
403 claims.sub,
404 response.matched_policies(),
405 response.errors()
406 );
407 Err(())
408 }
409 }
410}
411
412#[derive(Debug, Clone)]
418pub struct AuthUser {
419 pub sub: String,
420 pub email: Option<String>,
421 pub name: Option<String>,
422 pub username: Option<String>,
423 pub is_dev: bool,
424}
425
426impl From<JwtClaims> for AuthUser {
427 fn from(claims: JwtClaims) -> Self {
428 Self {
429 sub: claims.sub,
430 email: claims.email,
431 name: claims.name,
432 username: claims.preferred_username,
433 is_dev: false,
434 }
435 }
436}
437
438const SESSION_COOKIE_MAX_AGE: StdDuration = StdDuration::from_secs(3600);
440
441fn dev_user_key(token: &str) -> Option<String> {
444 let parts: Vec<&str> = token.splitn(4, ':').collect();
445 if parts.len() >= 4 {
446 Some(format!("dev:{}", parts[1]))
447 } else {
448 None
449 }
450}
451
452pub async fn check_auth(
472 auth: &Option<Arc<AuthState>>,
473 headers: &axum::http::HeaderMap,
474) -> Result<(Option<String>, String), StatusCode> {
475 let Some(auth) = auth.as_ref() else {
476 return Ok((None, "default".to_string())); };
478
479 if let Some(token) = headers
481 .get(header::AUTHORIZATION)
482 .and_then(|v| v.to_str().ok())
483 .and_then(|v| v.strip_prefix("Bearer "))
484 .map(|s| s.to_string())
485 {
486 if token.starts_with("dev:") {
488 if !auth.config.dev_config.local_dev_mode {
489 tracing::warn!("Dev token presented but dev mode is disabled — rejecting");
490 return Err(StatusCode::UNAUTHORIZED);
491 }
492 return match dev_user_key(&token) {
493 Some(key) => Ok((None, key)),
494 None => Err(StatusCode::UNAUTHORIZED),
495 };
496 }
497
498 return match auth.validate_token(&token).await {
499 Ok(claims) => {
500 if auth.check_cedar_authorized(&claims).is_err() {
501 return Err(StatusCode::FORBIDDEN);
502 }
503 Ok((None, claims.sub))
504 }
505 Err(e) => {
506 tracing::warn!("Bearer token validation failed: {}", e);
507 Err(StatusCode::UNAUTHORIZED)
508 }
509 };
510 }
511
512 let session_id = headers
514 .get(header::COOKIE)
515 .and_then(|v| v.to_str().ok())
516 .and_then(|cookies| extract_token_from_cookies(cookies, &auth.config.cookie_name));
517
518 let Some(session_id) = session_id else {
519 tracing::warn!("No auth token found in request");
520 return Err(StatusCode::UNAUTHORIZED);
521 };
522
523 if session_id.starts_with("dev:") {
525 if !auth.config.dev_config.local_dev_mode {
526 tracing::warn!("Dev cookie presented but dev mode is disabled — rejecting");
527 return Err(StatusCode::UNAUTHORIZED);
528 }
529 return match dev_user_key(&session_id) {
530 Some(key) => Ok((None, key)),
531 None => Err(StatusCode::UNAUTHORIZED),
532 };
533 }
534
535 match auth.session_manager.get_token(&session_id).await {
537 Ok(access_token) => match auth.validate_token(&access_token).await {
538 Ok(claims) => {
539 if auth.check_cedar_authorized(&claims).is_err() {
541 return Err(StatusCode::FORBIDDEN);
542 }
543 let secure = auth.client_config.redirect_uri.starts_with("https");
545 let cookie = create_auth_cookie(
546 &auth.config.cookie_name,
547 &session_id,
548 SESSION_COOKIE_MAX_AGE,
549 secure,
550 );
551 Ok((Some(cookie.to_string()), claims.sub))
552 }
553 Err(e) => {
554 tracing::warn!("Session token validation failed: {} — attempting force-refresh", e);
557 match auth.session_manager.force_refresh(&session_id).await {
558 Ok(new_token) => match auth.validate_token(&new_token).await {
559 Ok(claims) => {
560 if auth.check_cedar_authorized(&claims).is_err() {
562 return Err(StatusCode::FORBIDDEN);
563 }
564 let secure = auth.client_config.redirect_uri.starts_with("https");
565 let cookie = create_auth_cookie(
566 &auth.config.cookie_name,
567 &session_id,
568 SESSION_COOKIE_MAX_AGE,
569 secure,
570 );
571 Ok((Some(cookie.to_string()), claims.sub))
572 }
573 Err(e2) => {
574 tracing::warn!("Session token still invalid after force-refresh: {}", e2);
575 Err(StatusCode::UNAUTHORIZED)
576 }
577 },
578 Err(e2) => {
579 tracing::warn!("Force-refresh failed: {}", e2);
580 Err(StatusCode::UNAUTHORIZED)
581 }
582 }
583 }
584 },
585 Err(e) => {
586 tracing::warn!("Session lookup/refresh failed: {}", e);
587 Err(StatusCode::UNAUTHORIZED)
588 }
589 }
590}
591
592async fn resolve_access_token(
598 auth: &AuthState,
599 headers: &axum::http::HeaderMap,
600) -> Result<String, StatusCode> {
601 if let Some(token) = headers
603 .get(header::AUTHORIZATION)
604 .and_then(|v| v.to_str().ok())
605 .and_then(|v| v.strip_prefix("Bearer "))
606 .map(|s| s.to_string())
607 {
608 return Ok(token);
609 }
610
611 let session_id = headers
613 .get(header::COOKIE)
614 .and_then(|v| v.to_str().ok())
615 .and_then(|cookies| extract_token_from_cookies(cookies, &auth.config.cookie_name));
616
617 match session_id {
618 Some(sid) if sid.starts_with("dev:") => {
619 if !auth.config.dev_config.local_dev_mode {
620 tracing::warn!("Dev cookie in resolve_access_token but dev mode is disabled — rejecting");
621 Err(StatusCode::UNAUTHORIZED)
622 } else {
623 Ok(sid)
624 }
625 }
626 Some(sid) => auth.session_manager.get_token(&sid).await.map_err(|e| {
627 tracing::warn!("Failed to resolve session token: {}", e);
628 StatusCode::UNAUTHORIZED
629 }),
630 None => Err(StatusCode::UNAUTHORIZED),
631 }
632}
633
634fn extract_token_from_cookies(cookie_header: &str, cookie_name: &str) -> Option<String> {
636 for cookie in cookie_header.split(';') {
637 let cookie = cookie.trim();
638 if let Some(value) = cookie.strip_prefix(&format!("{}=", cookie_name)) {
639 return Some(value.to_string());
640 }
641 }
642 None
643}
644
645pub fn auth_routes() -> axum::Router<crate::ServerState> {
651 axum::Router::new()
652 .route("/login", axum::routing::get(login_handler))
653 .route("/callback", axum::routing::get(callback_handler))
654 .route("/me", axum::routing::get(me_handler))
655 .route("/logout", axum::routing::post(logout_handler))
656 .route("/mcp/login", axum::routing::get(mcp_login_handler))
657 .route("/mcp/callback", axum::routing::get(mcp_callback_handler))
658 .route("/mcp/status", axum::routing::get(mcp_status_handler))
659 .route("/mcp/logout", axum::routing::post(mcp_logout_handler))
660}
661
662#[derive(Debug, Deserialize)]
664pub struct CallbackQuery {
665 pub code: Option<String>,
666 pub state: Option<String>,
667 pub error: Option<String>,
668 pub error_description: Option<String>,
669}
670
671async fn login_handler(
673 State(state): State<crate::ServerState>,
674) -> Result<Response, AuthError> {
675 let auth = state.auth.as_ref().ok_or(AuthError::AuthNotConfigured)?;
676
677 if auth.is_dev_mode() {
679 tracing::info!("Dev mode: creating dev session");
680 let dev = &auth.config.dev_config;
681 let dev_token = format!(
682 "dev:{}:{}:{}",
683 dev.local_dev_email.as_deref().unwrap_or("dev@localhost"),
684 dev.local_dev_name.as_deref().unwrap_or("Dev User"),
685 dev.local_dev_username.as_deref().unwrap_or("dev")
686 );
687 let cookie = create_auth_cookie(&auth.config.cookie_name, &dev_token, StdDuration::from_secs(86400), false);
688 return Ok(Response::builder()
689 .status(StatusCode::FOUND)
690 .header(header::LOCATION, "/")
691 .header(header::SET_COOKIE, cookie.to_string())
692 .body(Body::empty())
693 .unwrap());
694 }
695
696 let pkce_session = auth.pkce_manager.create();
698 let challenge = OidcClient::generate_code_challenge(&pkce_session.verifier);
699
700 let auth_url = auth
701 .oidc_client
702 .build_authorization_url(&auth.client_config, &pkce_session.state, Some(&challenge))
703 .await
704 .map_err(|e| AuthError::OidcError(e.to_string()))?;
705
706 let secure = auth.client_config.redirect_uri.starts_with("https");
710 let pkce_cookie = Cookie::build((
711 auth.pkce_manager.cookie_name().to_string(),
712 pkce_session.cookie_value,
713 ))
714 .path("/")
715 .http_only(true)
716 .same_site(SameSite::Lax)
717 .secure(secure)
718 .max_age(TimeDuration::seconds(auth.pkce_manager.ttl().as_secs() as i64))
719 .build();
720
721 Ok(Response::builder()
722 .status(StatusCode::TEMPORARY_REDIRECT)
723 .header(header::LOCATION, &auth_url)
724 .header(header::SET_COOKIE, pkce_cookie.to_string())
725 .body(Body::empty())
726 .unwrap())
727}
728
729async fn callback_handler(
731 State(state): State<crate::ServerState>,
732 Query(query): Query<CallbackQuery>,
733 headers: axum::http::HeaderMap,
734) -> Result<Response, AuthError> {
735 let auth = state.auth.as_ref().ok_or(AuthError::AuthNotConfigured)?;
736
737 if let Some(error) = query.error {
739 let desc = query.error_description.unwrap_or_default();
740 tracing::error!("OIDC error: {} - {}", error, desc);
741 return Ok(Redirect::temporary(&format!(
742 "/?error={}&error_description={}",
743 urlencoding::encode(&error),
744 urlencoding::encode(&desc)
745 ))
746 .into_response());
747 }
748
749 let code = query.code.ok_or(AuthError::MissingCode)?;
750 let oauth_state = query.state.ok_or(AuthError::MissingState)?;
751
752 let cookie_header = headers
754 .get(header::COOKIE)
755 .and_then(|v| v.to_str().ok())
756 .unwrap_or("");
757 let pkce_value = extract_token_from_cookies(cookie_header, auth.pkce_manager.cookie_name())
758 .ok_or(AuthError::InvalidState)?;
759
760 let verifier = auth
762 .pkce_manager
763 .verify(&pkce_value, &oauth_state)
764 .ok_or(AuthError::InvalidState)?;
765
766 tracing::info!("Exchanging authorization code for tokens");
768 let token_response = auth
769 .oidc_client
770 .exchange_code_for_tokens(&auth.client_config, &code, Some(&verifier))
771 .await
772 .map_err(|e| AuthError::TokenExchangeFailed(e.to_string()))?;
773
774 let session_id = auth
775 .session_manager
776 .create_session(&token_response)
777 .await
778 .map_err(|e| AuthError::TokenExchangeFailed(format!("Session creation failed: {}", e)))?;
779
780 let max_age = SESSION_COOKIE_MAX_AGE;
783
784 let secure = auth.client_config.redirect_uri.starts_with("https");
786 let cookie = create_auth_cookie(&auth.config.cookie_name, &session_id, max_age, secure);
787
788 let clear_pkce = Cookie::build((auth.pkce_manager.cookie_name().to_string(), ""))
790 .path("/")
791 .http_only(true)
792 .same_site(SameSite::Lax)
793 .max_age(TimeDuration::seconds(-1))
794 .build();
795
796 tracing::info!("Authentication successful, redirecting to /");
797
798 Ok(Response::builder()
799 .status(StatusCode::FOUND)
800 .header(header::LOCATION, "/")
801 .header(header::SET_COOKIE, cookie.to_string())
802 .header(header::SET_COOKIE, clear_pkce.to_string())
803 .body(Body::empty())
804 .unwrap())
805}
806
807async fn me_handler(
809 State(state): State<crate::ServerState>,
810 headers: axum::http::HeaderMap,
811) -> Response {
812 let Some(ref auth) = state.auth else {
813 return axum::Json(serde_json::json!({
815 "authenticated": true,
816 "auth_enabled": false
817 }))
818 .into_response();
819 };
820
821 let cookie_header = headers
822 .get(header::COOKIE)
823 .and_then(|v| v.to_str().ok())
824 .unwrap_or("");
825
826 let bearer = headers
828 .get(header::AUTHORIZATION)
829 .and_then(|v| v.to_str().ok())
830 .and_then(|v| v.strip_prefix("Bearer "))
831 .map(String::from);
832
833 let token = bearer.clone().or_else(|| extract_token_from_cookies(cookie_header, &auth.config.cookie_name));
834
835 let Some(cookie_value) = token else {
836 return axum::Json(serde_json::json!({
837 "authenticated": false,
838 "auth_enabled": true
839 }))
840 .into_response();
841 };
842
843 if cookie_value.starts_with("dev:") && auth.config.dev_config.local_dev_mode {
846 let parts: Vec<&str> = cookie_value.splitn(4, ':').collect();
847 if parts.len() >= 4 {
848 return axum::Json(serde_json::json!({
849 "authenticated": true,
850 "auth_enabled": true,
851 "email": parts[1],
852 "name": parts[2],
853 "username": parts[3],
854 "dev_mode": true
855 }))
856 .into_response();
857 }
858 }
859
860 let access_token = if bearer.is_some() {
862 cookie_value
864 } else {
865 match auth.session_manager.get_token(&cookie_value).await {
867 Ok(token) => token,
868 Err(e) => {
869 tracing::debug!("Session token resolution failed for /auth/me: {}", e);
870 return axum::Json(serde_json::json!({
871 "authenticated": false,
872 "auth_enabled": true
873 }))
874 .into_response();
875 }
876 }
877 };
878
879 match auth.validate_token(&access_token).await {
881 Ok(claims) => axum::Json(serde_json::json!({
882 "authenticated": true,
883 "auth_enabled": true,
884 "sub": claims.sub,
885 "email": claims.email,
886 "name": claims.name,
887 "username": claims.preferred_username,
888 "dev_mode": false
889 }))
890 .into_response(),
891 Err(e) => {
892 tracing::debug!("Token validation failed for /auth/me: {}", e);
893 axum::Json(serde_json::json!({
894 "authenticated": false,
895 "auth_enabled": true
896 }))
897 .into_response()
898 }
899 }
900}
901
902async fn logout_handler(
904 State(state): State<crate::ServerState>,
905 headers: axum::http::HeaderMap,
906) -> Response {
907 let cookie_name = state
908 .auth
909 .as_ref()
910 .map(|a| a.config.cookie_name.as_str())
911 .unwrap_or("trustee_token");
912
913 if let Some(ref auth) = state.auth {
915 if let Some(cookie_header) = headers.get(header::COOKIE).and_then(|v| v.to_str().ok()) {
916 if let Some(session_id) = extract_token_from_cookies(cookie_header, cookie_name) {
917 if !session_id.starts_with("dev:") {
918 let _ = auth.session_manager.destroy_session(&session_id);
919 }
920 }
921 }
922 }
923
924 let cookie = Cookie::build((cookie_name.to_string(), ""))
925 .path("/")
926 .http_only(true)
927 .same_site(SameSite::Lax)
928 .max_age(TimeDuration::seconds(-1))
929 .build();
930
931 Response::builder()
932 .status(StatusCode::FOUND)
933 .header(header::LOCATION, "/")
934 .header(header::SET_COOKIE, cookie.to_string())
935 .body(Body::empty())
936 .unwrap()
937}
938
939#[derive(Debug, Deserialize)]
945pub struct McpLoginQuery {
946 pub cred: String,
947}
948
949#[derive(Debug, Deserialize)]
951pub struct McpCallbackQuery {
952 pub code: Option<String>,
953 pub state: Option<String>,
954 pub error: Option<String>,
955 pub error_description: Option<String>,
956}
957
958async fn mcp_login_handler(
963 State(state): State<crate::ServerState>,
964 Query(query): Query<McpLoginQuery>,
965 headers: axum::http::HeaderMap,
966) -> Result<Response, AuthError> {
967 let (_cookie, _user_key) = crate::auth::check_auth(&state.auth, &headers)
969 .await
970 .map_err(|_| AuthError::AuthNotConfigured)?;
971
972 let auth = state.auth.as_ref().ok_or(AuthError::AuthNotConfigured)?;
973
974 let cred_config = load_mcp_credential(&state, &query.cred).await?;
976
977 let (issuer_url, client_id, client_secret, scope) = match &cred_config {
978 McpCredentialInfo::WebInteractive {
979 issuer_url,
980 client_id,
981 client_secret,
982 scope,
983 } => (issuer_url.clone(), client_id.clone(), client_secret.clone(), scope.clone()),
984 _ => {
985 return Ok(Redirect::temporary(&format!(
986 "/?mcp_error={}",
987 urlencoding::encode(&format!("Credential '{}' is not web-interactive type", query.cred))
988 ))
989 .into_response());
990 }
991 };
992
993 let oidc_client = OidcClient::new();
995 let verifier = OidcClient::generate_code_verifier();
996 let challenge = OidcClient::generate_code_challenge(&verifier);
997 let oauth_state = OidcClient::generate_state();
998
999 let mcp_redirect_uri = format!(
1001 "{}/auth/mcp/callback",
1002 auth.client_config.redirect_uri.trim_end_matches('/').trim_end_matches("/auth/callback")
1003 );
1004
1005 let mcp_client_config = OidcClientConfig {
1006 issuer_url: issuer_url.clone(),
1007 client_id: client_id.clone(),
1008 client_secret: client_secret.clone(),
1009 redirect_uri: mcp_redirect_uri.clone(),
1010 scope: scope.clone(),
1011 code_challenge_method: "S256".to_string(),
1012 };
1013
1014 let auth_url = oidc_client
1016 .build_authorization_url(&mcp_client_config, &oauth_state, Some(&challenge))
1017 .await
1018 .map_err(|e| AuthError::OidcError(e.to_string()))?;
1019
1020 mcp_pkce().insert(oauth_state.clone(), verifier.clone(), query.cred.clone()).await;
1022
1023 tracing::info!(
1024 "Initiating MCP browser login for credential '{}' (issuer={})",
1025 query.cred, issuer_url
1026 );
1027
1028 Ok(Response::builder()
1029 .status(StatusCode::TEMPORARY_REDIRECT)
1030 .header(header::LOCATION, &auth_url)
1031 .body(Body::empty())
1032 .unwrap())
1033}
1034
1035async fn mcp_callback_handler(
1037 State(state): State<crate::ServerState>,
1038 Query(query): Query<McpCallbackQuery>,
1039 headers: axum::http::HeaderMap,
1040) -> Result<Response, AuthError> {
1041 let auth = state.auth.as_ref().ok_or(AuthError::AuthNotConfigured)?;
1042
1043 if let Some(error) = query.error {
1045 let desc = query.error_description.unwrap_or_default();
1046 tracing::error!("MCP OIDC error: {} - {}", error, desc);
1047 return Ok(Redirect::temporary(&format!(
1048 "/?mcp_error={}&error_description={}",
1049 urlencoding::encode(&error),
1050 urlencoding::encode(&desc)
1051 ))
1052 .into_response());
1053 }
1054
1055 let code = query.code.ok_or(AuthError::MissingCode)?;
1056 let oauth_state = query.state.ok_or(AuthError::MissingState)?;
1057
1058 let pkce_data = mcp_pkce().take(&oauth_state).await
1060 .ok_or(AuthError::InvalidState)?;
1061
1062 let verifier = pkce_data.verifier;
1063 let cred_name = &pkce_data.cred_name;
1064
1065 let cred_config = load_mcp_credential(&state, cred_name).await?;
1067
1068 let (issuer_url, client_id, client_secret, scope) = match &cred_config {
1069 McpCredentialInfo::WebInteractive {
1070 issuer_url,
1071 client_id,
1072 client_secret,
1073 scope,
1074 } => (issuer_url.clone(), client_id.clone(), client_secret.clone(), scope.clone()),
1075 _ => {
1076 return Ok(Redirect::temporary(&format!(
1077 "/?mcp_error={}",
1078 urlencoding::encode("Credential is not web-interactive type")
1079 ))
1080 .into_response());
1081 }
1082 };
1083
1084 let mcp_redirect_uri = format!(
1086 "{}/auth/mcp/callback",
1087 auth.client_config.redirect_uri.trim_end_matches('/').trim_end_matches("/auth/callback")
1088 );
1089
1090 let mcp_client_config = OidcClientConfig {
1091 issuer_url: issuer_url.clone(),
1092 client_id: client_id.clone(),
1093 client_secret: client_secret.clone(),
1094 redirect_uri: mcp_redirect_uri,
1095 scope: scope.clone(),
1096 code_challenge_method: "S256".to_string(),
1097 };
1098
1099 tracing::info!("Exchanging MCP authorization code for tokens (credential={})", cred_name);
1101 let oidc_client = OidcClient::new();
1102 let token_response = oidc_client
1103 .exchange_code_for_tokens(&mcp_client_config, &code, Some(&verifier))
1104 .await
1105 .map_err(|e| AuthError::TokenExchangeFailed(e.to_string()))?;
1106
1107 let expires_at = {
1109 let now = std::time::SystemTime::now()
1110 .duration_since(std::time::UNIX_EPOCH)
1111 .unwrap_or_default()
1112 .as_secs();
1113 let expires_epoch = now + token_response.expires_in.unwrap_or(900);
1114 let days = expires_epoch / 86400;
1115 let rem = expires_epoch % 86400;
1116 let h = rem / 3600;
1117 let m = (rem % 3600) / 60;
1118 let s = rem % 60;
1119 let z = days as i64 + 719468;
1120 let era = if z >= 0 { z } else { z - 146096 } / 146097;
1121 let doe = (z - era * 146097) as u64;
1122 let yoe = (doe - doe / 1460 + doe / 36524 - doe / 146096) / 365;
1123 let y = yoe as i64 + era * 400;
1124 let doy = doe - (365 * yoe + yoe / 4 - yoe / 100);
1125 let mp = (5 * doy + 2) / 153;
1126 let d = doy - (153 * mp + 2) / 5 + 1;
1127 let mon = if mp < 10 { mp + 3 } else { mp - 9 };
1128 let yr = if mon <= 2 { y + 1 } else { y };
1129 format!("{:04}-{:02}-{:02}T{:02}:{:02}:{:02}Z", yr, mon, d, h, m, s)
1130 };
1131
1132 use pep::{FileTokenStore, StoredToken, TokenStore};
1134
1135 let stored = StoredToken::new(
1136 &token_response.access_token,
1137 token_response.refresh_token.clone(),
1138 "Bearer",
1139 &expires_at,
1140 token_response.scope.clone(),
1141 );
1142
1143 let agent_name = state.config_toml.as_ref().and_then(|t| {
1144 toml::from_str::<toml::Value>(t).ok()
1145 .and_then(|v| v.get("agent").and_then(|a| a.get("name")).and_then(|n| n.as_str()).map(String::from))
1146 }).unwrap_or_else(|| "trustee".to_string());
1147 let token_store = FileTokenStore::new(&agent_name);
1148
1149 if let Err(e) = token_store.save(cred_name, &stored) {
1150 tracing::error!("Failed to store MCP token: {}", e);
1151 return Ok(Redirect::temporary(&format!(
1152 "/?mcp_error={}",
1153 urlencoding::encode(&format!("Failed to store token: {}", e))
1154 ))
1155 .into_response());
1156 }
1157
1158 tracing::info!(
1159 "MCP authentication successful for credential '{}' (expires {})",
1160 cred_name, expires_at
1161 );
1162
1163 Ok(Response::builder()
1164 .status(StatusCode::FOUND)
1165 .header(header::LOCATION, format!("/?mcp_connected={}", urlencoding::encode(cred_name)))
1166 .body(Body::empty())
1167 .unwrap())
1168}
1169
1170async fn mcp_status_handler(
1172 State(state): State<crate::ServerState>,
1173 headers: axum::http::HeaderMap,
1174) -> Response {
1175 use pep::{FileTokenStore, TokenStore};
1176
1177 let (_cookie, user_key) = match crate::auth::check_auth(&state.auth, &headers).await {
1179 Ok(result) => result,
1180 Err(code) => return (code, Json(serde_json::json!({"error": "Unauthorized"}))).into_response(),
1181 };
1182
1183 let config_toml = {
1185 let (_sid, session_arc, _, _) = state.ensure_active_session(&user_key).await;
1186 let session = session_arc.lock().await;
1187 match &session.config_toml {
1188 Some(t) => t.clone(),
1189 None => return (StatusCode::INTERNAL_SERVER_ERROR, "Config not loaded").into_response(),
1190 }
1191 };
1192
1193 let mcp_config: toml::Value = match toml::from_str(&config_toml) {
1194 Ok(v) => v,
1195 Err(_) => return Json(serde_json::json!([])).into_response(),
1196 };
1197
1198 let agent_name = {
1199 let (_sid, session_arc, _, _) = state.ensure_active_session(&user_key).await;
1200 let session = session_arc.lock().await;
1201 session.agent_name.clone()
1202 };
1203 let token_store = FileTokenStore::new(&agent_name);
1204
1205 let servers = mcp_config
1207 .get("mcp")
1208 .and_then(|m| m.get("servers"))
1209 .and_then(|s| s.as_array());
1210 let credentials = mcp_config
1211 .get("mcp")
1212 .and_then(|m| m.get("credentials"))
1213 .and_then(|c| c.as_table());
1214
1215 let mut cred_servers: std::collections::HashMap<String, Vec<String>> = std::collections::HashMap::new();
1216 if let Some(servers) = servers {
1217 for server in servers {
1218 let name = server.get("name").and_then(|n| n.as_str()).unwrap_or("");
1219 let cred_ref = server.get("credentials").and_then(|c| c.as_str()).unwrap_or("");
1220 if !cred_ref.is_empty() {
1221 cred_servers
1222 .entry(cred_ref.to_string())
1223 .or_default()
1224 .push(name.to_string());
1225 }
1226 }
1227 }
1228
1229 let mut result = Vec::new();
1230
1231 if let Some(creds) = credentials {
1232 for (cred_name, cred_config) in creds {
1233 let cred_type = cred_config.get("type").and_then(|t| t.as_str()).unwrap_or("unknown");
1234 let servers_using = cred_servers.get(cred_name).cloned().unwrap_or_default();
1235
1236 if cred_type == "web-session" {
1237 let connected = state.auth.is_some();
1239 result.push(serde_json::json!({
1240 "credential": cred_name,
1241 "type": cred_type,
1242 "connected": connected,
1243 "servers": servers_using,
1244 }));
1245 } else if cred_type == "web-interactive" || cred_type == "interactive" {
1246 let status = match token_store.load(cred_name) {
1248 Ok(Some(token)) => {
1249 let expired = token.is_expired();
1250 serde_json::json!({
1251 "credential": cred_name,
1252 "type": cred_type,
1253 "connected": !expired,
1254 "expires_at": token.expires_at,
1255 "servers": servers_using,
1256 })
1257 }
1258 _ => serde_json::json!({
1259 "credential": cred_name,
1260 "type": cred_type,
1261 "connected": false,
1262 "servers": servers_using,
1263 }),
1264 };
1265 result.push(status);
1266 }
1267 }
1268 }
1269
1270 Json(serde_json::Value::Array(result)).into_response()
1271}
1272
1273async fn mcp_logout_handler(
1275 State(state): State<crate::ServerState>,
1276 Query(query): Query<McpLoginQuery>,
1277 headers: axum::http::HeaderMap,
1278) -> Response {
1279 use pep::{FileTokenStore, TokenStore};
1280
1281 let (_cookie, user_key) = match crate::auth::check_auth(&state.auth, &headers).await {
1283 Ok(result) => result,
1284 Err(code) => return (code, Json(serde_json::json!({"error": "Unauthorized"}))).into_response(),
1285 };
1286
1287 let agent_name = {
1288 let (_sid, session_arc, _, _) = state.ensure_active_session(&user_key).await;
1289 let session = session_arc.lock().await;
1290 session.agent_name.clone()
1291 };
1292 let token_store = FileTokenStore::new(&agent_name);
1293
1294 match token_store.delete(&query.cred) {
1295 Ok(()) => {
1296 tracing::info!("Removed MCP credentials for '{}'", query.cred);
1297 Json(serde_json::json!({"success": true})).into_response()
1298 }
1299 Err(e) => {
1300 tracing::error!("Failed to remove MCP credentials: {}", e);
1301 (
1302 StatusCode::INTERNAL_SERVER_ERROR,
1303 Json(serde_json::json!({"error": e.to_string()})),
1304 )
1305 .into_response()
1306 }
1307 }
1308}
1309
1310struct McpPkceStore {
1317 entries: tokio::sync::Mutex<std::collections::HashMap<String, McpPkceEntry>>,
1318}
1319
1320struct McpPkceEntry {
1321 verifier: String,
1322 cred_name: String,
1323 created_at: std::time::Instant,
1324}
1325
1326impl McpPkceStore {
1327 fn new() -> Self {
1328 Self {
1329 entries: tokio::sync::Mutex::new(std::collections::HashMap::new()),
1330 }
1331 }
1332
1333 async fn insert(&self, state: String, verifier: String, cred_name: String) {
1335 let mut map = self.entries.lock().await;
1336 let cutoff = std::time::Instant::now() - std::time::Duration::from_secs(600);
1338 map.retain(|_, v| v.created_at > cutoff);
1339 map.insert(state, McpPkceEntry {
1340 verifier,
1341 cred_name,
1342 created_at: std::time::Instant::now(),
1343 });
1344 }
1345
1346 async fn take(&self, state: &str) -> Option<McpPkceEntry> {
1348 let mut map = self.entries.lock().await;
1349 map.remove(state)
1350 }
1351}
1352
1353static MCP_PKCE: std::sync::OnceLock<McpPkceStore> = std::sync::OnceLock::new();
1355
1356fn mcp_pkce() -> &'static McpPkceStore {
1358 MCP_PKCE.get_or_init(McpPkceStore::new)
1359}
1360
1361enum McpCredentialInfo {
1363 WebInteractive {
1364 issuer_url: String,
1365 client_id: String,
1366 client_secret: Option<String>,
1367 scope: String,
1368 },
1369 Other(String),
1370}
1371
1372async fn load_mcp_credential(
1374 state: &crate::ServerState,
1375 cred_name: &str,
1376) -> Result<McpCredentialInfo, AuthError> {
1377 let config_toml = state
1378 .config_toml
1379 .clone()
1380 .ok_or(AuthError::AuthNotConfigured)?;
1381
1382 let config: toml::Value = toml::from_str(&config_toml)
1383 .map_err(|e| AuthError::OidcError(format!("Config parse error: {}", e)))?;
1384
1385 let cred = config
1386 .get("mcp")
1387 .and_then(|m| m.get("credentials"))
1388 .and_then(|c| c.as_table())
1389 .and_then(|c| c.get(cred_name))
1390 .ok_or_else(|| AuthError::OidcError(format!("Credential '{}' not found", cred_name)))?;
1391
1392 let cred_type = cred.get("type").and_then(|t| t.as_str()).unwrap_or("unknown");
1393
1394 match cred_type {
1395 "web-interactive" => {
1396 let issuer_url = cred
1397 .get("issuer_url")
1398 .and_then(|v| v.as_str())
1399 .ok_or_else(|| AuthError::OidcError("Missing issuer_url".into()))?
1400 .to_string();
1401 let client_id = cred
1402 .get("client_id")
1403 .and_then(|v| v.as_str())
1404 .ok_or_else(|| AuthError::OidcError("Missing client_id".into()))?
1405 .to_string();
1406 let client_secret = cred
1407 .get("client_secret")
1408 .and_then(|v| v.as_str())
1409 .map(String::from);
1410 let scope = cred
1411 .get("scope")
1412 .and_then(|v| v.as_str())
1413 .unwrap_or("openid profile email")
1414 .to_string();
1415
1416 Ok(McpCredentialInfo::WebInteractive {
1417 issuer_url,
1418 client_id,
1419 client_secret,
1420 scope,
1421 })
1422 }
1423 other => Ok(McpCredentialInfo::Other(other.to_string())),
1424 }
1425}
1426
1427fn create_auth_cookie(name: &str, value: &str, max_age: StdDuration, secure: bool) -> Cookie<'static> {
1433 Cookie::build((name.to_string(), value.to_string()))
1434 .path("/")
1435 .http_only(true)
1436 .same_site(SameSite::Lax)
1437 .secure(secure)
1438 .max_age(TimeDuration::seconds(max_age.as_secs() as i64))
1439 .build()
1440}
1441
1442#[derive(Debug)]
1448pub enum AuthError {
1449 MissingCode,
1450 MissingState,
1451 InvalidState,
1452 OidcError(String),
1453 TokenExchangeFailed(String),
1454 AuthNotConfigured,
1455}
1456
1457impl IntoResponse for AuthError {
1458 fn into_response(self) -> Response {
1459 let (_status, msg) = match self {
1460 AuthError::MissingCode => (StatusCode::BAD_REQUEST, "Missing authorization code"),
1461 AuthError::MissingState => (StatusCode::BAD_REQUEST, "Missing state parameter"),
1462 AuthError::InvalidState => (StatusCode::BAD_REQUEST, "Invalid or expired state"),
1463 AuthError::OidcError(_) => (StatusCode::SERVICE_UNAVAILABLE, "Authentication service error"),
1464 AuthError::TokenExchangeFailed(_) => (StatusCode::BAD_REQUEST, "Token exchange failed"),
1465 AuthError::AuthNotConfigured => (StatusCode::NOT_IMPLEMENTED, "Authentication not configured"),
1466 };
1467 Redirect::temporary(&format!("/?error={}", urlencoding::encode(msg))).into_response()
1468 }
1469}