diff --git a/.secrets.baseline b/.secrets.baseline index 6a0e6bc8..001dbd0f 100644 --- a/.secrets.baseline +++ b/.secrets.baseline @@ -3,7 +3,7 @@ "files": "(?x)(Cargo\\.lock$|\\.lock$)|^\\.secrets\\.baseline$|^.secrets.baseline$", "lines": null }, - "generated_at": "2026-09-16T09:14:49Z", + "generated_at": "2026-09-23T14:09:14Z", "plugins_used": [ { "name": "AWSKeyDetector" @@ -142,7 +142,7 @@ "hashed_secret": "bfc6000db1195a9522813fc405c666dd4ce669ad", "is_secret": false, "is_verified": false, - "line_number": 98, + "line_number": 106, "type": "Secret Keyword", "verified_result": null }, @@ -150,7 +150,7 @@ "hashed_secret": "4a4645604f0b9e29503be96a87f6f47a6e4a7890", "is_secret": false, "is_verified": false, - "line_number": 105, + "line_number": 113, "type": "Secret Keyword", "verified_result": null } @@ -168,7 +168,7 @@ "hashed_secret": "4a4645604f0b9e29503be96a87f6f47a6e4a7890", "is_secret": false, "is_verified": false, - "line_number": 176, + "line_number": 183, "type": "Secret Keyword", "verified_result": null } @@ -374,7 +374,7 @@ "hashed_secret": "fdda45b7f6d2ead95d9991fc4678640c3bab0d84", "is_secret": false, "is_verified": false, - "line_number": 342, + "line_number": 348, "type": "Secret Keyword", "verified_result": null }, @@ -382,7 +382,7 @@ "hashed_secret": "093d378410a5cfa4bd5088f3fef62fbdb8a95665", "is_secret": false, "is_verified": false, - "line_number": 348, + "line_number": 354, "type": "Secret Keyword", "verified_result": null }, @@ -390,7 +390,7 @@ "hashed_secret": "c3de40d5e3fc71ed62771c2127a8e42585026c97", "is_secret": false, "is_verified": false, - "line_number": 350, + "line_number": 356, "type": "Secret Keyword", "verified_result": null }, @@ -398,7 +398,7 @@ "hashed_secret": "4d4acd9b084d13f5fdb23807d857e1c48a1cfd0f", "is_secret": false, "is_verified": false, - "line_number": 439, + "line_number": 445, "type": "Secret Keyword", "verified_result": null }, @@ -406,7 +406,7 @@ "hashed_secret": "bd0160c2cf35d950843c88f3be2b9412ed71f485", "is_secret": false, "is_verified": false, - "line_number": 474, + "line_number": 480, "type": "Secret Keyword", "verified_result": null }, @@ -414,7 +414,7 @@ "hashed_secret": "293324f6824bb3a6db5c4dc42a60ddd4a9851c99", "is_secret": false, "is_verified": false, - "line_number": 637, + "line_number": 643, "type": "Hex High Entropy String", "verified_result": null } diff --git a/crates/contextforge-data-plane-lib/src/authorization/mod.rs b/crates/contextforge-data-plane-lib/src/authorization/mod.rs index 49c95a67..98ad898f 100644 --- a/crates/contextforge-data-plane-lib/src/authorization/mod.rs +++ b/crates/contextforge-data-plane-lib/src/authorization/mod.rs @@ -15,6 +15,12 @@ pub use principal_extractor::{ AuthorizedPrincipal, CelPrincipalExtractor, DefaultPrincipalExtractor, PrincipalExtractor, }; +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] +pub enum Permission { + Admin, + MCPUser, +} + pub fn get_authorization_service( config: &JwksConfig, ) -> Result, AuthorizationError> { @@ -27,6 +33,12 @@ pub trait AuthorizationService: std::fmt::Debug { async fn authorize(&self, authorization_token: &HeaderValue) -> Option; } +#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)] +pub enum AuthenticationError { + #[error("invalid bearer token")] + InvalidToken, +} + #[derive(Debug, thiserror::Error)] #[allow(dead_code)] pub enum AuthorizationError { @@ -86,6 +98,12 @@ pub struct AuthorizationClaims { value: serde_json::Value, } +impl AuthorizationClaims { + pub fn as_value(&self) -> &serde_json::Value { + &self.value + } +} + impl From for AuthorizationClaims { fn from(value: serde_json::Value) -> Self { Self { value } diff --git a/crates/contextforge-data-plane-lib/src/authorization/principal_extractor/default_principal_extractor.rs b/crates/contextforge-data-plane-lib/src/authorization/principal_extractor/default_principal_extractor.rs index 7cae3c4d..986aa3bb 100644 --- a/crates/contextforge-data-plane-lib/src/authorization/principal_extractor/default_principal_extractor.rs +++ b/crates/contextforge-data-plane-lib/src/authorization/principal_extractor/default_principal_extractor.rs @@ -13,8 +13,10 @@ impl PrincipalExtractor for DefaultPrincipalExtractor { ) -> Result> { let user_id = ["sub", "user_id", "UserId"].into_iter().find_map(|claim| claims.get(claim)).and_then(|v| v.as_str()); - let tenant_id = - ["tenantId", "tenant_id"].into_iter().find_map(|claim| claims.get(claim)).and_then(|v| v.as_str()); + let tenant_id = ["tenantId", "tenant_id", "woTenantId"] + .into_iter() + .find_map(|claim| claims.get(claim)) + .and_then(|v| v.as_str()); match (user_id, tenant_id) { (Some(user_id), Some(tenant_id)) => Ok(AuthorizedPrincipal::builder() .user_id(user_id.to_owned()) @@ -25,3 +27,44 @@ impl PrincipalExtractor for DefaultPrincipalExtractor { } } } + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn supported_identity_aliases_are_preserved() { + for user_claim in ["sub", "user_id", "UserId"] { + for tenant_claim in ["tenantId", "tenant_id", "woTenantId"] { + let claims = json!({user_claim: "user", tenant_claim: "tenant"}); + let principal = DefaultPrincipalExtractor {}.extract(&claims).unwrap(); + assert_eq!(principal.user_id, "user"); + assert_eq!(principal.tenant_id, "tenant"); + assert!(principal.scopes.is_empty()); + } + } + } + + #[test] + fn existing_alias_precedence_is_preserved() { + let claims = json!({ + "sub": "subject", "user_id": "alternate", "UserId": "other", + "tenantId": "tenant", "tenant_id": "alternate", "woTenantId": "other" + }); + let principal = DefaultPrincipalExtractor {}.extract(&claims).unwrap(); + assert_eq!(principal.user_id, "subject"); + assert_eq!(principal.tenant_id, "tenant"); + } + + #[test] + fn missing_identity_is_rejected() { + for claims in [ + json!({"sub": "user"}), + json!({"tenant_id": "tenant"}), + json!({"woUserId": "user", "woTenantId": "tenant"}), + ] { + assert!(DefaultPrincipalExtractor {}.extract(&claims).is_err()); + } + } +} diff --git a/crates/contextforge-data-plane-lib/src/layers/mod.rs b/crates/contextforge-data-plane-lib/src/layers/mod.rs index bb341459..1740d2bd 100644 --- a/crates/contextforge-data-plane-lib/src/layers/mod.rs +++ b/crates/contextforge-data-plane-lib/src/layers/mod.rs @@ -6,3 +6,5 @@ pub mod virtual_host_config; pub mod virtual_host_id; pub use principal_extractor::PrincipalExtractorLayer; + +pub mod permission; diff --git a/crates/contextforge-data-plane-lib/src/layers/permission.rs b/crates/contextforge-data-plane-lib/src/layers/permission.rs new file mode 100644 index 00000000..2e3902bd --- /dev/null +++ b/crates/contextforge-data-plane-lib/src/layers/permission.rs @@ -0,0 +1,128 @@ +use crate::{ + authorization::{AuthenticationError, AuthorizationClaims, AuthorizedPrincipal, Permission}, + errors::{custom_error, unauthorized_response}, +}; +use axum::{ + extract::{Request, State}, + middleware::Next, + response::Response, +}; +use http::StatusCode; + +pub async fn require_permission(State(permission): State, request: Request, next: Next) -> Response { + if request.extensions().get::().is_none() { + return unauthorized_response("Missing verified identity"); + } + let Some(claims) = request.extensions().get::() else { + return unauthorized_response("Missing verified claims"); + }; + match has_permission(claims, permission) { + Ok(true) => next.run(request).await, + Ok(false) => custom_error(StatusCode::FORBIDDEN, "Insufficient permission"), + Err(_) => unauthorized_response("Invalid role claims"), + } +} + +/// Test role mapping; awaiting confirmation from WxO. +fn has_permission(claims: &AuthorizationClaims, permission: Permission) -> Result { + let allows = |role: &str| match role { + "admin" => true, + "builder" | "user" => permission == Permission::MCPUser, + _ => false, + }; + let claims = claims.as_value(); + let mut granted = false; + if let Some(roles) = claims.get("roles") { + for role in roles.as_array().ok_or(AuthenticationError::InvalidToken)? { + granted |= allows(role.as_str().ok_or(AuthenticationError::InvalidToken)?); + } + } + if let Some(role) = claims.get("role") { + granted |= allows(role.as_str().ok_or(AuthenticationError::InvalidToken)?); + } + Ok(granted) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::authorization::{CelPrincipalExtractor, DefaultPrincipalExtractor, PrincipalExtractor}; + use axum::{Router, body::Body, middleware, routing::get}; + use serde_json::json; + use tower::ServiceExt; + + fn app(permission: Permission) -> Router { + Router::new() + .route("/", get(|| async { StatusCode::NO_CONTENT })) + .layer(middleware::from_fn_with_state(permission, require_permission)) + } + + #[tokio::test] + async fn guards_require_both_verified_claims_and_identity() { + let claims = json!({"sub":"user", "tenant_id":"tenant", "roles":["admin"]}); + let principal = DefaultPrincipalExtractor {}.extract(&claims).unwrap(); + for (include_claims, include_principal) in [(false, false), (true, false), (false, true)] { + let mut request = Request::new(Body::empty()); + if include_claims { + request.extensions_mut().insert(AuthorizationClaims::from(claims.clone())); + } + if include_principal { + request.extensions_mut().insert(principal.clone()); + } + let response = app(Permission::MCPUser).oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + } + } + + #[tokio::test] + async fn guards_map_verified_roles_to_the_requested_permission() { + use StatusCode as S; + for (role_claims, admin_status, mcp_status) in [ + (json!({"roles":["admin"]}), S::NO_CONTENT, S::NO_CONTENT), + (json!({"roles":["builder"]}), S::FORBIDDEN, S::NO_CONTENT), + (json!({"roles":["user"]}), S::FORBIDDEN, S::NO_CONTENT), + (json!({"role":"admin"}), S::NO_CONTENT, S::NO_CONTENT), + (json!({"roles":["user"], "role":"admin"}), S::NO_CONTENT, S::NO_CONTENT), + (json!({"roles":["unknown", "user"]}), S::FORBIDDEN, S::NO_CONTENT), + (json!({"roles":["ServiceAdmin"]}), S::FORBIDDEN, S::FORBIDDEN), + (json!({"roles":[]}), S::FORBIDDEN, S::FORBIDDEN), + (json!({}), S::FORBIDDEN, S::FORBIDDEN), + (json!({"roles":"admin"}), S::UNAUTHORIZED, S::UNAUTHORIZED), + (json!({"role":["admin"]}), S::UNAUTHORIZED, S::UNAUTHORIZED), + (json!({"roles":["admin", 42]}), S::UNAUTHORIZED, S::UNAUTHORIZED), + (json!({"roles":["admin"], "role":42}), S::UNAUTHORIZED, S::UNAUTHORIZED), + ] { + let mut claims = json!({"sub":"user", "tenant_id":"tenant"}); + claims.as_object_mut().unwrap().extend(role_claims.as_object().unwrap().clone()); + let principal = DefaultPrincipalExtractor {}.extract(&claims).unwrap(); + for (permission, expected) in [(Permission::Admin, admin_status), (Permission::MCPUser, mcp_status)] { + let mut request = Request::new(Body::empty()); + request.extensions_mut().insert(principal.clone()); + request.extensions_mut().insert(AuthorizationClaims::from(claims.clone())); + let response = app(permission).oneshot(request).await.unwrap(); + assert_eq!(response.status(), expected, "{permission:?}: {role_claims}"); + } + } + } + + #[tokio::test] + async fn cel_identity_mapping_cannot_grant_roles_absent_from_the_token() { + let extractor = CelPrincipalExtractor::from_expression( + r#"{"user_id": claims.sub, "tenant_id": claims.woTenantId, "role": "admin", "scopes": ["Admin"]}"#, + ) + .unwrap(); + for (roles, mcp_status) in [(json!(["user"]), StatusCode::NO_CONTENT), (json!([]), StatusCode::FORBIDDEN)] { + let claims = json!({"sub":"user", "woTenantId":"tenant", "roles":roles}); + let principal = extractor.extract(&claims).unwrap(); + for (permission, expected) in + [(Permission::Admin, StatusCode::FORBIDDEN), (Permission::MCPUser, mcp_status)] + { + let mut request = Request::new(Body::empty()); + request.extensions_mut().insert(principal.clone()); + request.extensions_mut().insert(AuthorizationClaims::from(claims.clone())); + let response = app(permission).oneshot(request).await.unwrap(); + assert_eq!(response.status(), expected); + } + } + } +} diff --git a/crates/contextforge-data-plane-lib/src/lib.rs b/crates/contextforge-data-plane-lib/src/lib.rs index 1eec4bca..c5748ab1 100644 --- a/crates/contextforge-data-plane-lib/src/lib.rs +++ b/crates/contextforge-data-plane-lib/src/lib.rs @@ -42,7 +42,7 @@ pub type Error = Box; pub type Result = std::result::Result; use crate::{ - authorization::{CelPrincipalExtractor, DefaultPrincipalExtractor}, + authorization::{CelPrincipalExtractor, DefaultPrincipalExtractor, Permission}, config_stores::RedisStore, layers::{ claims_id::claims_layer, @@ -53,6 +53,7 @@ use crate::{ }, }; pub use authorization::{AuthorizationClaims, AuthorizationService, get_authorization_service}; +pub use layers::permission::require_permission; #[derive(Clone)] pub enum UserConfigStoreType { @@ -143,7 +144,8 @@ impl Gateway { let app = axum::Router::new() .nest_service("/servers/{virtual_host_name}/mcp", mcp_service) .layer(middleware::from_fn(virtual_host_config_layer)) - .layer(middleware::from_fn_with_state(mcp_gateway_state.clone(), user_config_store_layer)); + .layer(middleware::from_fn_with_state(mcp_gateway_state.clone(), user_config_store_layer)) + .layer(middleware::from_fn_with_state(Permission::MCPUser, require_permission)); let app = if let Some(cel_principal_extractor_path) = config.cel_principal_extractor_path.as_ref() { app.layer(layers::PrincipalExtractorLayer::new(CelPrincipalExtractor::from_file( diff --git a/crates/contextforge-data-plane-lib/src/tools.rs b/crates/contextforge-data-plane-lib/src/tools.rs index e3a8a2f4..4c56d14c 100644 --- a/crates/contextforge-data-plane-lib/src/tools.rs +++ b/crates/contextforge-data-plane-lib/src/tools.rs @@ -98,6 +98,8 @@ pub async fn get_token( "iss": "contexforge-dataplane", "sub": user_id.clone(), "aud": "contexforge-dataplane-audience", + "role": "user", + "woUserId": user_id.clone(), "exp": now + Duration::from_hours(1).as_secs(), "nbf": now - Duration::from_mins(1).as_secs(), "iat": now, diff --git a/crates/contextforge-data-plane-lib/tests/gateway.rs b/crates/contextforge-data-plane-lib/tests/gateway.rs index 0abd8d15..8c9120df 100644 --- a/crates/contextforge-data-plane-lib/tests/gateway.rs +++ b/crates/contextforge-data-plane-lib/tests/gateway.rs @@ -17,3 +17,6 @@ mod resources; mod subscriptions; #[path = "gateway/tools.rs"] mod tools; + +#[path = "gateway/downstream_auth.rs"] +mod downstream_auth; diff --git a/crates/contextforge-data-plane-lib/tests/gateway/downstream_auth.rs b/crates/contextforge-data-plane-lib/tests/gateway/downstream_auth.rs new file mode 100644 index 00000000..ee4cdf9a --- /dev/null +++ b/crates/contextforge-data-plane-lib/tests/gateway/downstream_auth.rs @@ -0,0 +1,187 @@ +use crate::harness::{TestServer, create_default_config}; +use async_trait::async_trait; +use axum::{Json, Router, body::Body, routing::get}; +use contextforge_data_plane_apis::{ + User, + user_store::{UserConfig, VirtualHost}, +}; +use contextforge_data_plane_lib::{ + Config, ConfigStore, ConfigStoreError, Gateway, UserConfigStoreType, get_authorization_service, +}; +use http::{Request, StatusCode}; +use jsonwebtoken::{Algorithm, EncodingKey, Header, encode, jwk::Jwk}; +use rmcp::transport::streamable_http_server::session::local::LocalSessionManager; +use serde_json::{Value, json}; +use std::{ + collections::HashMap, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + time::{SystemTime, UNIX_EPOCH}, +}; +use tower::ServiceExt; + +fn key() -> EncodingKey { + EncodingKey::from_rsa_pem(include_bytes!(concat!(env!("CARGO_MANIFEST_DIR"), "/../../assets/jwt.key"))) + .expect("valid authentication test fixture") +} +fn claims() -> Value { + let now = SystemTime::now().duration_since(UNIX_EPOCH).expect("valid authentication test fixture").as_secs(); + json!({"iss":"mcpgateway", "aud":"mcpgateway-api", "sub":"user", "woTenantId":"tenant", "role":"user", "exp":now+3600, "nbf":now-60}) +} +fn token(claims: &Value) -> String { + let mut header = Header::new(Algorithm::RS256); + header.kid = Some("test".into()); + encode(&header, claims, &key()).expect("valid authentication test fixture") +} + +#[derive(Clone)] +struct CountingStore { + reads: Arc, + virtual_host: VirtualHost, +} +#[async_trait] +impl ConfigStore for CountingStore { + async fn get_config<'a>(&self, user: &'a User) -> Result { + assert_eq!(user.key(), "user"); + self.reads.fetch_add(1, Ordering::SeqCst); + Ok(UserConfig { virtual_hosts: HashMap::from([("test".into(), self.virtual_host.clone())]) }) + } + async fn set_config<'a>(&self, _: &'a User, _: &'a UserConfig) -> Result<(), ConfigStoreError> { + unreachable!() + } +} + +async fn gateway(config: Config, store: CountingStore) -> Router { + Gateway::builder() + .with_authorization_service( + get_authorization_service(&config.jwks_config).expect("valid authentication test fixture"), + ) + .with_config(config) + .with_user_config_store_type(UserConfigStoreType::Test(Arc::new(store))) + .with_session_manager(Arc::new(LocalSessionManager::default())) + .build() + .into_router() + .await + .expect("valid authentication test fixture") +} + +fn request(token: &str) -> Request { + let request = Request::builder() + .method("POST") + .uri("/contextforge-rs/servers/test/mcp") + .header("host", "localhost") + .header("content-type", "application/json") + .header("accept", "application/json, text/event-stream") + .header("MCP-Protocol-Version", "2026-07-28") + .header("Mcp-Method", "tools/call") + .header("Mcp-Name", "sum") + .header("authorization", format!("Bearer {token}")); + let params = json!({ + "name": "sum", + "arguments": {"a": 2, "b": 3}, + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientInfo": {"name": "auth-test", "version": "1"}, + "io.modelcontextprotocol/clientCapabilities": {} + } + }); + request + .body(Body::from(json!({"jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": params}).to_string())) + .expect("valid authentication test fixture") +} + +async fn key_server() -> TestServer { + let mut jwk = Jwk::from_encoding_key(&key(), Algorithm::RS256).expect("valid authentication test fixture"); + jwk.common.key_id = Some("test".into()); + let document = json!({"keys":[jwk]}); + TestServer::start_http(Router::new().route( + "/jwks", + get(move || { + let document = document.clone(); + async move { Json(document) } + }), + )) + .await + .expect("valid authentication test fixture") +} + +#[tokio::test] +async fn permission_denial_prevents_backend_calls_after_a_successful_request() { + use axum::{ + extract::State, + middleware::{self, Next}, + }; + use contextforge_data_plane_apis::user_store::{BackendMCPGateway, ServiceRoute}; + use rmcp::transport::{StreamableHttpServerConfig, StreamableHttpService}; + let hits = Arc::new(AtomicUsize::new(0)); + let service = StreamableHttpService::new( + || Ok(crate::harness::mock_counter::Counter::new()), + LocalSessionManager::default().into(), + StreamableHttpServerConfig::default(), + ); + let backend = + TestServer::start_http(Router::new().route_service("/mcp", service).layer(middleware::from_fn_with_state( + Arc::clone(&hits), + |State(hits): State>, request: axum::extract::Request, next: Next| async move { + if request.headers().get("Mcp-Method").is_some_and(|method| method == "tools/call") { + hits.fetch_add(1, Ordering::SeqCst); + } + next.run(request).await + }, + ))) + .await + .expect("valid authentication test fixture"); + let store = CountingStore { + reads: Arc::new(AtomicUsize::new(0)), + virtual_host: VirtualHost { + backends: HashMap::from([( + "counter".into(), + BackendMCPGateway { + name: "counter".into(), + url: backend.url("/mcp").parse().expect("valid authentication test fixture"), + mcp_protocol_version: rmcp::model::ProtocolVersion::V_2026_07_28, + passthrough_headers: vec![], + add_headers: HashMap::new(), + remove_headers: vec![], + tool_schemas: HashMap::new(), + completion: HashMap::new(), + }, + )]), + tools: HashMap::from([( + "sum".into(), + ServiceRoute { backend_name: "counter".into(), upstream_name: "sum".into() }, + )]), + resources: HashMap::new(), + resource_templates: HashMap::new(), + prompts: HashMap::new(), + }, + }; + let server = key_server().await; + let mut config = create_default_config(); + config.upstream_transport_config.upstream_connection_mode = + Some(contextforge_data_plane_lib::UpstreamConnectionMode::PlainTextOrTls); + config.jwks_config.url = server.url("/jwks").parse().expect("valid authentication test fixture"); + let app = gateway(config, store.clone()).await; + for (role, expected) in [("user", StatusCode::OK), ("unknown", StatusCode::FORBIDDEN)] { + let mut c = claims(); + c["role"] = role.into(); + let token = token(&c); + let request = request(&token); + let response = app.clone().oneshot(request).await.expect("valid authentication test fixture"); + let status = response.status(); + let body = axum::body::to_bytes(response.into_body(), 65_536).await.expect("valid authentication test fixture"); + assert_eq!(status, expected, "{}", String::from_utf8_lossy(&body)); + if role == "user" { + let text = String::from_utf8_lossy(&body); + let data = text.lines().find_map(|line| line.strip_prefix("data: ")).unwrap_or(&text); + let message: Value = serde_json::from_str(data).expect("MCP result JSON"); + assert_eq!(message["result"]["content"][0]["text"], "5", "{message}"); + } + assert_eq!(store.reads.load(Ordering::SeqCst), 1); + assert_eq!(hits.load(Ordering::SeqCst), 1, "denied request must never reach backend"); + } + server.shutdown().await.expect("valid authentication test fixture"); + backend.shutdown().await.expect("valid authentication test fixture"); +} diff --git a/crates/contextforge-data-plane-lib/tests/gateway/harness/auth.rs b/crates/contextforge-data-plane-lib/tests/gateway/harness/auth.rs index 0a1495fc..5389a747 100644 --- a/crates/contextforge-data-plane-lib/tests/gateway/harness/auth.rs +++ b/crates/contextforge-data-plane-lib/tests/gateway/harness/auth.rs @@ -19,6 +19,7 @@ fn default_claims(user_id: &str) -> serde_json::Value { "iss": "mcpgateway", "sub": user_id, "tenant_id": "test_tenant", + "role": "user", "aud": "mcpgateway-api", "exp": now + TEST_TOKEN_TTL_SECS, "iat": now,