fix(iam): address PR 1875 review issues for OIDC STS flows (#1969)

This commit is contained in:
安正超
2026-02-26 14:38:57 +08:00
committed by GitHub
parent c35ef84a8c
commit 40692f18ed
6 changed files with 424 additions and 93 deletions
+232 -48
View File
@@ -30,7 +30,12 @@ use std::borrow::Cow;
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use tracing::{error, info};
use std::sync::RwLock;
use std::time::{Duration as StdDuration, Instant};
use tracing::{error, info, warn};
use url::Url;
const OIDC_JWKS_REFRESH_INTERVAL: StdDuration = StdDuration::from_secs(24 * 60 * 60);
// ---- HTTP Client Adapter ----
@@ -139,8 +144,16 @@ pub struct OidcClaims {
/// a `CoreClient` because the crate uses type-state generics that make storing
/// the configured client in a HashMap impractical. The client is reconstructed
/// on-the-fly from metadata when needed.
#[derive(Clone)]
struct ProviderState {
metadata: CoreProviderMetadata,
discovered_at: Instant,
}
impl ProviderState {
fn is_stale(&self) -> bool {
self.discovered_at.elapsed() >= OIDC_JWKS_REFRESH_INTERVAL
}
}
// ---- Core OIDC system ----
@@ -148,7 +161,7 @@ struct ProviderState {
/// Global OIDC manager for all configured providers.
pub struct OidcSys {
configs: HashMap<String, OidcProviderConfig>,
provider_states: HashMap<String, ProviderState>,
provider_states: RwLock<HashMap<String, ProviderState>>,
state_store: OidcStateStore,
http_client: ReqwestHttpClient,
}
@@ -181,7 +194,7 @@ impl OidcSys {
Ok(Self {
configs,
provider_states,
provider_states: RwLock::new(provider_states),
state_store: OidcStateStore::new(),
http_client,
})
@@ -191,7 +204,7 @@ impl OidcSys {
pub fn empty() -> Self {
Self {
configs: HashMap::new(),
provider_states: HashMap::new(),
provider_states: RwLock::new(HashMap::new()),
state_store: OidcStateStore::new(),
http_client: ReqwestHttpClient(reqwest::Client::new()),
}
@@ -224,12 +237,12 @@ impl OidcSys {
.configs
.get(provider_id)
.ok_or_else(|| format!("unknown OIDC provider: {provider_id}"))?;
let state = self
.provider_states
.get(provider_id)
.ok_or_else(|| format!("provider not discovered: {provider_id}"))?;
let state = self.ensure_provider_state(provider_id, config).await?;
let (pkce_challenge, pkce_verifier) = PkceCodeChallenge::new_random_sha256();
let redirect = RedirectUrl::new(redirect_uri.to_string()).map_err(|e| format!("invalid redirect URI: {e}"))?;
// Construct CoreClient on-the-fly (avoids type-state storage issues)
let client = CoreClient::from_provider_metadata(
state.metadata.clone(),
ClientId::new(config.client_id.clone()),
@@ -237,10 +250,6 @@ impl OidcSys {
)
.set_auth_type(AuthType::RequestBody);
let (pkce_challenge, pkce_verifier) = PkceCodeChallenge::new_random_sha256();
let redirect = RedirectUrl::new(redirect_uri.to_string()).map_err(|e| format!("invalid redirect URI: {e}"))?;
let mut auth_req =
client.authorize_url(CoreAuthenticationFlow::AuthorizationCode, CsrfToken::new_random, Nonce::new_random);
auth_req = auth_req.set_redirect_uri(Cow::Owned(redirect));
@@ -287,10 +296,7 @@ impl OidcSys {
.configs
.get(&session.provider_id)
.ok_or_else(|| format!("unknown provider: {}", session.provider_id))?;
let provider_state = self
.provider_states
.get(&session.provider_id)
.ok_or_else(|| format!("provider not discovered: {}", session.provider_id))?;
let provider_state = self.get_provider_state(&session.provider_id)?;
// Construct CoreClient on-the-fly with JWKS from discovery
let client = CoreClient::from_provider_metadata(
@@ -319,16 +325,36 @@ impl OidcSys {
.ok_or_else(|| "no id_token in token response".to_string())?;
let verifier = client.id_token_verifier();
let _verified_claims = id_token
.claims(&verifier, &Nonce::new(session.nonce.clone()))
.map_err(|e| format!("ID token verification failed: {e}"))?;
let verified = id_token.claims(&verifier, &Nonce::new(session.nonce.clone()));
if let Err(e) = verified {
let refreshed_state = self
.refresh_provider_state(&session.provider_id, config)
.await
.map_err(|refresh_err| {
format!("ID token verification failed: {e}; failed to refresh provider metadata: {refresh_err}")
})?;
warn!(
"OIDC provider '{}' JWKS metadata refreshed and verification retried after failure",
session.provider_id
);
let client = CoreClient::from_provider_metadata(
refreshed_state.metadata,
ClientId::new(config.client_id.clone()),
config.client_secret.as_ref().map(|s| ClientSecret::new(s.clone())),
)
.set_auth_type(AuthType::RequestBody);
let verifier = client.id_token_verifier();
id_token
.claims(&verifier, &Nonce::new(session.nonce.clone()))
.map_err(|retry_err| format!("ID token verification failed after JWKS refresh: {retry_err}"))?;
}
// Extract raw claims from the verified JWT for custom claim support
// (the crate verifies signature/expiry/nonce; we decode payload for non-standard claims)
let raw_jwt = serde_json::to_value(id_token)
.ok()
.and_then(|v| v.as_str().map(String::from))
.unwrap_or_default();
let raw_jwt = id_token.to_string();
let raw = decode_jwt_payload(&raw_jwt);
let claims = OidcClaims {
@@ -412,10 +438,12 @@ impl OidcSys {
.ok_or_else(|| "JWT missing 'iss' claim".to_string())?;
// Find matching provider by issuer
let (provider_id, config, state) = self
let (provider_id, config, mut state) = self
.find_provider_by_issuer(issuer)
.ok_or_else(|| format!("no OIDC provider configured for issuer: {issuer}"))?;
state = self.ensure_provider_state_if_stale(&provider_id, &config, &state).await?;
// Reconstruct CoreClient from provider metadata
let client = CoreClient::from_provider_metadata(
state.metadata.clone(),
@@ -432,9 +460,25 @@ impl OidcSys {
// Verify the token (signature, issuer, audience, expiry) — skip nonce
// (nonce is only required for the authorization code flow)
let verifier = client.id_token_verifier();
let _verified = id_token
.claims(&verifier, |_: Option<&Nonce>| Ok(()))
.map_err(|e| format!("ID token verification failed: {e}"))?;
if let Err(e) = id_token.claims(&verifier, |_: Option<&Nonce>| Ok(())) {
state = self
.refresh_provider_state(&provider_id, &config)
.await
.map_err(|refresh_err| {
format!("ID token verification failed: {e}; failed to refresh provider metadata: {refresh_err}")
})?;
let client = CoreClient::from_provider_metadata(
state.metadata,
ClientId::new(config.client_id.clone()),
config.client_secret.as_ref().map(|s| ClientSecret::new(s.clone())),
)
.set_auth_type(AuthType::RequestBody);
let verifier = client.id_token_verifier();
id_token
.claims(&verifier, |_: Option<&Nonce>| Ok(()))
.map_err(|retry_err| format!("ID token verification failed after JWKS refresh: {retry_err}"))?;
}
// Extract claims using the provider's claim configuration
let claims = OidcClaims {
@@ -449,20 +493,85 @@ impl OidcSys {
}
/// Find a provider whose discovered issuer matches the given JWT issuer string.
fn find_provider_by_issuer(&self, issuer: &str) -> Option<(&str, &OidcProviderConfig, &ProviderState)> {
let issuer_normalized = issuer.trim_end_matches('/');
for (id, state) in &self.provider_states {
fn find_provider_by_issuer(&self, issuer: &str) -> Option<(String, OidcProviderConfig, ProviderState)> {
let (issuer_scheme, issuer_host, issuer_port, issuer_path) = normalize_issuer(issuer)?;
let map = self
.provider_states
.read()
.map_err(|e| format!("provider state lock poisoned: {e}"))
.ok()?;
for (id, state) in map.iter() {
let provider_issuer = state.metadata.issuer().as_str();
let provider_normalized = provider_issuer.trim_end_matches('/');
if issuer_normalized == provider_normalized
let Some((provider_scheme, provider_host, provider_port, provider_path)) = normalize_issuer(provider_issuer) else {
continue;
};
if issuer_scheme == provider_scheme
&& issuer_host == provider_host
&& issuer_port == provider_port
&& issuer_path == provider_path
&& let Some(config) = self.configs.get(id)
{
return Some((id, config, state));
return Some((id.clone(), config.clone(), state.clone()));
}
}
None
}
fn get_provider_state(&self, provider_id: &str) -> Result<ProviderState, String> {
self.provider_states
.read()
.map_err(|e| format!("provider state lock poisoned: {e}"))?
.get(provider_id)
.cloned()
.ok_or_else(|| format!("provider not discovered: {provider_id}"))
}
async fn refresh_provider_state(&self, provider_id: &str, config: &OidcProviderConfig) -> Result<ProviderState, String> {
let state = Self::discover_provider(config, &self.http_client).await?;
let mut map = self.provider_states.write().map_err(|e| {
let msg = e.to_string();
format!("provider state lock poisoned: {msg}")
})?;
map.insert(provider_id.to_string(), state.clone());
Ok(state)
}
async fn ensure_provider_state(&self, provider_id: &str, config: &OidcProviderConfig) -> Result<ProviderState, String> {
let state = self.get_provider_state(provider_id)?;
if state.is_stale() {
self.refresh_provider_state(provider_id, config).await.or_else(|refresh_err| {
warn!(
"OIDC provider '{}' JWKS metadata refresh skipped due to transient network issue: {}",
provider_id, refresh_err
);
Ok(state)
})
} else {
Ok(state)
}
}
async fn ensure_provider_state_if_stale(
&self,
provider_id: &str,
config: &OidcProviderConfig,
state: &ProviderState,
) -> Result<ProviderState, String> {
if state.is_stale() {
self.refresh_provider_state(provider_id, config).await.or_else(|refresh_err| {
warn!(
"OIDC provider '{}' JWKS metadata refresh skipped due to transient network issue: {}",
provider_id, refresh_err
);
Ok(state.clone())
})
} else {
Ok(state.clone())
}
}
/// Get the state store (used by HTTP handlers).
pub fn state_store(&self) -> &OidcStateStore {
&self.state_store
@@ -611,18 +720,8 @@ impl OidcSys {
/// `discover_async` fetches the discovery document and JWKS in one step.
async fn discover_provider(config: &OidcProviderConfig, http_client: &ReqwestHttpClient) -> Result<ProviderState, String> {
// The openidconnect crate expects the issuer URL (base), not the
// .well-known/openid-configuration URL. Strip the suffix if present.
let issuer_str = config
.config_url
.strip_suffix("/.well-known/openid-configuration")
.unwrap_or(&config.config_url);
// Ensure trailing slash for correct URL joining in the crate
let issuer_str = if issuer_str.ends_with('/') {
issuer_str.to_string()
} else {
format!("{issuer_str}/")
};
// .well-known/openid-configuration URL.
let issuer_str = normalize_config_url(&config.config_url)?;
let issuer_url = IssuerUrl::new(issuer_str).map_err(|e| format!("invalid issuer URL: {e}"))?;
@@ -630,12 +729,63 @@ impl OidcSys {
.await
.map_err(|e| format!("discovery failed: {e}"))?;
Ok(ProviderState { metadata })
Ok(ProviderState {
metadata,
discovered_at: Instant::now(),
})
}
}
// --- Helper functions ---
fn normalize_issuer(raw: &str) -> Option<(String, String, u16, String)> {
let parsed = Url::parse(raw).ok()?;
if parsed.scheme() != "http" && parsed.scheme() != "https" {
return None;
}
let host = parsed.host_str()?.to_ascii_lowercase();
let port = parsed.port_or_known_default()?;
let normalized_path = {
let path = parsed.path().trim_end_matches('/').to_string();
if path.is_empty() { "/".to_string() } else { path }
};
Some((parsed.scheme().to_string(), host, port, normalized_path))
}
fn normalize_config_url(config_url: &str) -> Result<String, String> {
let config_url = config_url.trim();
let url = Url::parse(config_url).map_err(|e| format!("invalid config_url: {e}"))?;
if url.scheme() != "http" && url.scheme() != "https" {
return Err(format!("invalid config_url scheme: {}", url.scheme()));
}
let host = url.host_str().ok_or_else(|| "config_url missing host".to_string())?;
let path = url.path().trim_end_matches('/');
if path.contains("/.well-known/") && !path.ends_with("/.well-known/openid-configuration") {
return Err("config_url uses an unsupported .well-known discovery URL".into());
}
let normalized_path = path.strip_suffix("/.well-known/openid-configuration").unwrap_or(path);
let mut issuer = format!("{}://{host}", url.scheme());
if let Some(port) = url.port() {
issuer.push(':');
issuer.push_str(&port.to_string());
}
if !normalized_path.is_empty() {
issuer.push_str(normalized_path);
}
if !issuer.ends_with('/') {
issuer.push('/');
}
Ok(issuer)
}
/// Decode the payload section of a JWT without validation (token must already be verified).
pub(crate) fn decode_jwt_payload(token: &str) -> HashMap<String, serde_json::Value> {
let parts: Vec<&str> = token.split('.').collect();
@@ -722,6 +872,40 @@ mod tests {
assert_eq!(claims.get("email").and_then(|v| v.as_str()), Some("user@example.com"));
}
#[test]
fn test_normalize_issuer_matches() {
let lhs = normalize_issuer("https://idp.example.com/.well-known/openid-configuration/").unwrap();
let rhs = normalize_issuer("https://idp.example.com/.well-known/openid-configuration").unwrap();
assert_eq!(lhs, rhs);
assert_eq!(
lhs,
(
"https".to_string(),
"idp.example.com".to_string(),
443,
"/.well-known/openid-configuration".to_string()
)
);
}
#[test]
fn test_normalize_config_url() {
assert_eq!(
normalize_config_url("https://idp.example.com/.well-known/openid-configuration").unwrap(),
"https://idp.example.com/"
);
assert_eq!(
normalize_config_url("https://idp.example.com/.well-known/openid-configuration/").unwrap(),
"https://idp.example.com/"
);
assert_eq!(
normalize_config_url("https://idp.example.com/custom/realm").unwrap(),
"https://idp.example.com/custom/realm/"
);
assert!(normalize_config_url("https://idp.example.com/.well-known/invalid").is_err());
assert!(normalize_config_url("gopher://idp.example.com").is_err());
}
#[test]
fn test_decode_jwt_payload_invalid() {
assert!(decode_jwt_payload("not-a-jwt").is_empty());
@@ -810,7 +994,7 @@ mod tests {
}
OidcSys {
configs: config_map,
provider_states: HashMap::new(),
provider_states: RwLock::new(HashMap::new()),
state_store: OidcStateStore::new(),
http_client: ReqwestHttpClient(reqwest::Client::new()),
}
+46 -3
View File
@@ -13,7 +13,15 @@
// limitations under the License.
use moka::future::Cache;
use std::time::Duration;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use tracing::warn;
const OIDC_STATE_CAPACITY: u64 = 10_000;
const OIDC_STATE_CAPACITY_WARNING: u64 = 9_000;
const OIDC_STATE_CAPACITY_CRITICAL: u64 = 10_000;
const OIDC_STATE_CAPACITY_LOG_INTERVAL_SECS: u64 = 60;
/// Stores the PKCE verifier and nonce for an in-flight OIDC authorization flow.
#[derive(Debug, Clone)]
@@ -29,20 +37,55 @@ pub struct OidcAuthSession {
#[derive(Clone)]
pub struct OidcStateStore {
cache: Cache<String, OidcAuthSession>,
last_capacity_log_at: Arc<AtomicU64>,
}
impl OidcStateStore {
pub fn new() -> Self {
let cache = Cache::builder()
.max_capacity(10_000)
.max_capacity(OIDC_STATE_CAPACITY)
.time_to_live(Duration::from_secs(300)) // 5 minute TTL
.build();
Self { cache }
Self {
cache,
last_capacity_log_at: Arc::new(AtomicU64::new(0)),
}
}
/// Store a new auth session keyed by the OAuth2 `state` parameter.
pub async fn insert(&self, state: String, session: OidcAuthSession) {
self.cache.insert(state, session).await;
let size = self.cache.entry_count();
if !Self::should_log_capacity_warning(&self.last_capacity_log_at) {
return;
}
if size >= OIDC_STATE_CAPACITY_CRITICAL {
self.cache.run_pending_tasks().await;
warn!("OIDC state store reached configured capacity ({size}/{OIDC_STATE_CAPACITY})");
return;
}
if size >= OIDC_STATE_CAPACITY_WARNING {
self.cache.run_pending_tasks().await;
warn!("OIDC state store approaching capacity ({size}/{OIDC_STATE_CAPACITY})");
}
}
fn now_unix_secs() -> u64 {
SystemTime::now().duration_since(UNIX_EPOCH).map_or(0, |d| d.as_secs())
}
fn should_log_capacity_warning(last_log_at: &AtomicU64) -> bool {
let now = Self::now_unix_secs();
let last = last_log_at.load(Ordering::Acquire);
if now.saturating_sub(last) < OIDC_STATE_CAPACITY_LOG_INTERVAL_SECS {
return false;
}
last_log_at
.compare_exchange(last, now, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
}
/// Retrieve and remove an auth session (single-use). Returns None if expired or not found.
+7
View File
@@ -778,6 +778,9 @@ impl<T: Store> IamSys<T> {
use rustfs_policy::policy::default::DEFAULT_POLICIES;
let mut resolved = Vec::new();
for policy_name in claim_policies.split(',').map(|s| s.trim()).filter(|s| !s.is_empty()) {
if !Self::is_safe_claim_policy_name(policy_name) {
continue;
}
for (name, p) in DEFAULT_POLICIES.iter() {
if *name == policy_name {
resolved.push(p.clone());
@@ -827,6 +830,10 @@ impl<T: Store> IamSys<T> {
is_owner || combined_policy.is_allowed(args).await
}
fn is_safe_claim_policy_name(policy: &str) -> bool {
!policy.is_empty() && policy.chars().all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-')
}
pub async fn is_allowed_service_account(&self, args: &Args<'_>, parent_user: &str) -> bool {
let Some(p) = args.claims.get("parent") else {
return false;
+2 -1
View File
@@ -57,6 +57,7 @@ impl Operation for IsAdminHandler {
let is_admin = if is_admin {
true
} else {
let empty_claims = HashMap::new();
let iam_store = rustfs_iam::get().map_err(|_| s3_error!(InternalError, "iam not init"))?;
let conditions = get_condition_values(&req.headers, &cred, None, None, None);
iam_store
@@ -66,7 +67,7 @@ impl Operation for IsAdminHandler {
action: Action::AdminAction(AdminAction::AllAdminActions),
conditions: &conditions,
is_owner: false,
claims: cred.claims.as_ref().unwrap_or(&HashMap::new()),
claims: cred.claims.as_ref().unwrap_or(&empty_claims),
deny_only: false,
bucket: "",
object: "",
+129 -39
View File
@@ -21,6 +21,7 @@ use matchit::Params;
use s3s::{Body, S3Error, S3ErrorCode, S3Request, S3Response, S3Result, s3_error};
use time::OffsetDateTime;
use tracing::{error, info, warn};
use url::Url;
const OIDC_PATH_PREFIX: &str = "/rustfs/admin/v3/oidc";
@@ -78,7 +79,10 @@ impl Operation for ListOidcProvidersHandler {
let json_body = serde_json::to_vec(&providers)
.map_err(|e| S3Error::with_message(S3ErrorCode::InternalError, format!("serialize error: {e}")))?;
Ok(S3Response::new((StatusCode::OK, Body::from(json_body))))
let mut resp = S3Response::new((StatusCode::OK, Body::from(json_body)));
resp.headers
.insert(http::header::CONTENT_TYPE, http::HeaderValue::from_static("application/json"));
Ok(resp)
}
}
@@ -103,7 +107,7 @@ impl Operation for OidcAuthorizeHandler {
let redirect_uri = derive_callback_uri(&req, provider_id)?;
// Optional: redirect_after query parameter (must be a safe relative path)
let redirect_after = extract_query_param(&req.uri, "redirect_after").filter(|p| is_safe_redirect_path(p));
let redirect_after = extract_safe_redirect_after(&req.uri)?;
let auth_url = oidc_sys
.authorize_url(provider_id, &redirect_uri, redirect_after)
@@ -192,7 +196,7 @@ impl Operation for OidcCallbackHandler {
&new_cred.session_token,
new_cred.expiration,
session.redirect_after.as_deref(),
);
)?;
let mut resp = S3Response::new((StatusCode::FOUND, Body::empty()));
resp.headers.insert(
@@ -213,31 +217,25 @@ fn derive_callback_uri(req: &S3Request<Body>, provider_id: &str) -> S3Result<Str
// Use explicitly configured redirect_uri if available
if let Some(oidc_sys) = rustfs_iam::get_oidc()
&& let Some(config) = oidc_sys.get_provider_config(provider_id)
&& let Some(ref uri) = config.redirect_uri
{
return Ok(uri.clone());
}
let scheme = req
.headers
.get("x-forwarded-proto")
.and_then(|v| v.to_str().ok())
.unwrap_or_else(|| req.uri.scheme_str().unwrap_or("http"));
if let Some(ref uri) = config.redirect_uri {
let parsed = Url::parse(uri).map_err(|_| s3_error!(InvalidRequest, "invalid configured redirect_uri"))?;
if !is_valid_scheme(parsed.scheme()) || parsed.host_str().is_none() {
return Err(s3_error!(InvalidRequest, "configured redirect_uri must be absolute http/https URL"));
}
return Ok(uri.clone());
}
if !is_valid_scheme(scheme) {
return Err(s3_error!(InvalidRequest, "invalid scheme in request"));
if !config.redirect_uri_dynamic {
return Err(s3_error!(
InvalidRequest,
"provider requires explicit redirect_uri because redirect_uri_dynamic is disabled"
));
}
}
let host = req
.headers
.get(http::header::HOST)
.and_then(|v| v.to_str().ok())
.or_else(|| req.uri.host())
.ok_or_else(|| s3_error!(InvalidRequest, "cannot determine host for redirect URI"))?;
// Validate host doesn't contain path separators or other injection characters
if host.contains('/') || host.contains('\\') {
return Err(s3_error!(InvalidRequest, "invalid host header"));
}
let scheme = extract_request_scheme(req)?;
let host = extract_request_host(req)?;
Ok(format!("{scheme}://{host}/rustfs/admin/v3/oidc/callback/{provider_id}"))
}
@@ -261,6 +259,15 @@ fn extract_query_param(uri: &http::Uri, key: &str) -> Option<String> {
})
}
fn extract_safe_redirect_after(uri: &http::Uri) -> S3Result<Option<String>> {
let redirect_after = extract_query_param(uri, "redirect_after");
match redirect_after {
Some(value) if !is_safe_redirect_path(&value) => Err(s3_error!(InvalidRequest, "invalid redirect_after")),
Some(value) => Ok(Some(value)),
None => Ok(None),
}
}
/// Build the console redirect URL with STS credentials in the hash fragment.
fn build_console_redirect(
req: &S3Request<Body>,
@@ -269,20 +276,9 @@ fn build_console_redirect(
session_token: &str,
expiration: Option<OffsetDateTime>,
redirect_after: Option<&str>,
) -> String {
let scheme = req
.headers
.get("x-forwarded-proto")
.and_then(|v| v.to_str().ok())
.filter(|s| is_valid_scheme(s))
.unwrap_or_else(|| req.uri.scheme_str().unwrap_or("http"));
let host = req
.headers
.get(http::header::HOST)
.and_then(|v| v.to_str().ok())
.filter(|h| !h.contains('/') && !h.contains('\\'))
.unwrap_or("localhost");
) -> S3Result<String> {
let scheme = extract_request_scheme(req)?;
let host = extract_request_host(req)?;
let console_prefix = "/rustfs/console";
let page = redirect_after.filter(|p| is_safe_redirect_path(p)).unwrap_or("/");
@@ -300,7 +296,63 @@ fn build_console_redirect(
urlencoding::encode(page),
);
format!("{scheme}://{host}{console_prefix}/auth/oidc-callback/#{fragment}")
Ok(format!("{scheme}://{host}{console_prefix}/auth/oidc-callback/#{fragment}"))
}
fn extract_request_scheme(req: &S3Request<Body>) -> S3Result<String> {
let raw_scheme = req
.headers
.get("x-forwarded-proto")
.and_then(|v| v.to_str().ok())
.map(|v| v.split(',').next().map(str::trim).unwrap_or(""))
.filter(|v| !v.is_empty())
.map(ToOwned::to_owned)
.or_else(|| req.uri.scheme_str().map(str::to_owned))
.unwrap_or_else(|| "http".to_owned())
.to_ascii_lowercase();
if !is_valid_scheme(&raw_scheme) {
return Err(s3_error!(InvalidRequest, "invalid scheme in request"));
}
Ok(raw_scheme)
}
fn extract_request_host(req: &S3Request<Body>) -> S3Result<String> {
let host = req
.headers
.get(http::header::HOST)
.and_then(|v| v.to_str().ok())
.and_then(|v| {
let v = v.trim();
if v.is_empty() { None } else { Some(v.to_owned()) }
})
.or_else(|| req.uri.authority().map(|a| a.as_str().to_owned()))
.ok_or_else(|| s3_error!(InvalidRequest, "cannot determine host for redirect URI"))?;
parse_host_authority(&host)
}
fn parse_host_authority(raw_host: &str) -> S3Result<String> {
let host = raw_host.trim();
if host.is_empty() {
return Err(s3_error!(InvalidRequest, "invalid host header"));
}
// Parse as authority to normalize and validate, while rejecting URL-style
// constructions (userinfo, query, fragment, and explicit paths).
let parsed = Url::parse(&format!("http://{host}")).map_err(|_| s3_error!(InvalidRequest, "invalid host header"))?;
if !parsed.username().is_empty() || parsed.password().is_some() {
return Err(s3_error!(InvalidRequest, "invalid host header"));
}
if parsed.query().is_some() || parsed.fragment().is_some() {
return Err(s3_error!(InvalidRequest, "invalid host header"));
}
if parsed.path() != "/" {
return Err(s3_error!(InvalidRequest, "invalid host"));
}
Ok(parsed.authority().to_string())
}
#[cfg(test)]
@@ -336,6 +388,44 @@ mod tests {
assert_eq!(extract_query_param(&uri, "redirect_after"), Some("/dashboard".to_string()));
}
#[test]
fn test_parse_host_authority_rejects_userinfo() {
assert!(parse_host_authority("evil.com@victim.com").is_err());
}
#[test]
fn test_parse_host_authority_rejects_query_fragment() {
assert!(parse_host_authority("example.com?x=y").is_err());
assert!(parse_host_authority("example.com#fragment").is_err());
}
#[test]
fn test_parse_host_authority_rejects_path() {
assert!(parse_host_authority("example.com/path").is_err());
}
#[test]
fn test_parse_host_authority_accepts_valid_host_with_port() {
assert_eq!(
parse_host_authority("example.com:8443").expect("valid host should pass"),
"example.com:8443"
);
}
#[test]
fn test_extract_safe_redirect_after() {
let uri: http::Uri = "http://localhost/callback?redirect_after=%2Fdashboard".parse().unwrap();
assert_eq!(
extract_safe_redirect_after(&uri).expect("valid redirect should pass"),
Some("/dashboard".to_string())
);
let uri: http::Uri = "http://localhost/callback?redirect_after=javascript:alert(1)"
.parse()
.unwrap();
assert!(extract_safe_redirect_after(&uri).is_err());
}
#[test]
fn test_is_valid_provider_id() {
assert!(is_valid_provider_id("AUTHENTIK"));
+8 -2
View File
@@ -19,6 +19,7 @@ use crate::{
server::ADMIN_PREFIX,
};
use http::StatusCode;
use http::header::HeaderValue;
use hyper::Method;
use matchit::Params;
use rustfs_config::MAX_ADMIN_REQUEST_BODY_SIZE;
@@ -184,7 +185,8 @@ async fn handle_assume_role(
};
// getAssumeRoleCredentials
let output = serialize::<AssumeRoleOutput>(&resp).unwrap();
let output = serialize::<AssumeRoleOutput>(&resp)
.map_err(|e| S3Error::with_message(S3ErrorCode::InternalError, format!("serialize assume role output failed: {e}")))?;
Ok(S3Response::new((StatusCode::OK, Body::from(output))))
}
@@ -214,6 +216,10 @@ async fn handle_assume_role_with_web_identity(body: AssumeRoleRequest) -> S3Resu
// Map claims to policies and groups
let (policies, groups) = oidc_sys.map_claims_to_policies(&provider_id, &claims);
if policies.is_empty() && groups.is_empty() {
return Err(s3_error!(InvalidArgument, "no policies are available for this OIDC token"));
}
info!(
"AssumeRoleWithWebIdentity: user='{}', provider='{}', policies={:?}, groups={:?}",
claims.username, provider_id, policies, groups
@@ -278,7 +284,7 @@ async fn handle_assume_role_with_web_identity(body: AssumeRoleRequest) -> S3Resu
let mut resp = S3Response::new((StatusCode::OK, Body::from(xml.into_bytes())));
resp.headers
.insert(http::header::CONTENT_TYPE, "application/xml".parse().unwrap());
.insert(http::header::CONTENT_TYPE, HeaderValue::from_static("application/xml"));
Ok(resp)
}