feat(iam): retry OIDC discovery with issuer URL slash variants (#2360)

Co-authored-by: GatewayJ <8352692332qq.com>
Co-authored-by: 安正超 <anzhengchao@gmail.com>
This commit is contained in:
GatewayJ
2026-04-03 21:10:27 +08:00
committed by GitHub
parent b3f31ad694
commit c244943313
+227 -9
View File
@@ -833,18 +833,38 @@ impl OidcSys {
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.
let issuer_str = normalize_config_url(&config.config_url)?;
let base_issuer = normalize_config_url(&config.config_url)?;
let candidates = issuer_candidates(&base_issuer);
let mut last_errors = Vec::new();
let issuer_url = IssuerUrl::new(issuer_str).map_err(|e| format!("invalid issuer URL: {e}"))?;
for candidate_issuer in candidates.iter() {
let issuer_url = IssuerUrl::new(candidate_issuer.clone()).map_err(|e| format!("invalid issuer URL: {e}"))?;
let metadata = CoreProviderMetadata::discover_async(issuer_url, http_client)
.await
.map_err(|e| format!("discovery failed: {e}"))?;
match CoreProviderMetadata::discover_async(issuer_url, http_client)
.await
.map_err(|e| format!("discovery failed: {e}"))
{
Ok(metadata) => {
return Ok(ProviderState {
metadata,
discovered_at: Instant::now(),
});
}
Err(error) => {
last_errors.push(format!("issuer '{candidate_issuer}': {error}"));
warn!(
"OIDC provider '{}' discovery attempt failed for issuer '{}': {}",
config.id, candidate_issuer, error
);
}
}
}
Ok(ProviderState {
metadata,
discovered_at: Instant::now(),
})
Err(format!(
"discovery failed for all issuer variants {:?}: {}",
candidates,
last_errors.join("; ")
))
}
}
@@ -959,6 +979,21 @@ fn normalize_config_url(config_url: &str) -> Result<String, String> {
Ok(issuer)
}
fn issuer_candidates(base: &str) -> Vec<String> {
let original = base.trim();
let mut variants = Vec::with_capacity(2);
variants.push(original.to_string());
let toggled = if original.ends_with('/') {
original.trim_end_matches('/').to_string()
} else {
format!("{original}/")
};
variants.push(toggled);
variants
}
/// 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();
@@ -1116,6 +1151,189 @@ mod tests {
assert!(normalize_config_url("not-a-url").is_err());
}
#[test]
fn test_issuer_candidates() {
assert_eq!(
issuer_candidates("https://idp.example.com/realm"),
vec![
"https://idp.example.com/realm".to_string(),
"https://idp.example.com/realm/".to_string()
]
);
assert_eq!(
issuer_candidates("https://idp.example.com/realm/"),
vec![
"https://idp.example.com/realm/".to_string(),
"https://idp.example.com/realm".to_string()
]
);
assert_eq!(
issuer_candidates("https://idp.example.com"),
vec!["https://idp.example.com".to_string(), "https://idp.example.com/".to_string()]
);
}
fn build_mocked_oidc_provider_config(id: &str, config_url: &str) -> OidcProviderConfig {
OidcProviderConfig {
id: id.to_string(),
enabled: true,
config_url: config_url.to_string(),
client_id: "rustfs-oidc-test".to_string(),
client_secret: None,
scopes: vec!["openid".to_string()],
redirect_uri: None,
redirect_uri_dynamic: false,
claim_name: "sub".to_string(),
claim_prefix: "oidc".to_string(),
role_policy: String::new(),
display_name: "mock-oidc".to_string(),
groups_claim: "groups".to_string(),
email_claim: "email".to_string(),
username_claim: "username".to_string(),
}
}
fn start_mock_oidc_discovery_server<F>(
build_discovery_issuer: F,
max_requests: usize,
) -> (String, std::thread::JoinHandle<()>)
where
F: Fn(&str) -> String + Send + 'static,
{
use std::io::Read;
use std::io::Write;
use std::net::{Shutdown, TcpListener};
use std::time::{Duration, Instant};
// After the last completed response, exit if no new connection arrives within this window.
const IDLE_SHUTDOWN: Duration = Duration::from_millis(100);
const ABSOLUTE_CAP: Duration = Duration::from_millis(500);
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let discovery_issuer = build_discovery_issuer(&base);
let discovery_body = serde_json::json!({
"issuer": discovery_issuer,
"authorization_endpoint": format!("{base}/authorize"),
"token_endpoint": format!("{base}/token"),
"jwks_uri": format!("{base}/jwks"),
"response_types_supported": ["code"],
"response_modes_supported": ["query"],
"subject_types_supported": ["public"],
"id_token_signing_alg_values_supported": ["RS256"],
})
.to_string();
let jwks_body = r#"{"keys":[]}"#;
let handle = std::thread::spawn(move || {
listener
.set_nonblocking(true)
.expect("failed to set discovery mock listener non-blocking");
let mut seen = 0usize;
let start = Instant::now();
let mut last_completed = Instant::now();
loop {
if seen > 0 && last_completed.elapsed() >= IDLE_SHUTDOWN {
break;
}
if start.elapsed() >= ABSOLUTE_CAP {
break;
}
let mut stream = match listener.accept() {
Ok((stream, _)) => stream,
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => {
std::thread::sleep(Duration::from_millis(5));
continue;
}
Err(_) => break,
};
seen += 1;
let mut request_bytes = Vec::new();
let mut buffer = [0u8; 4096];
loop {
let n = stream.read(&mut buffer).unwrap_or_default();
if n == 0 {
break;
}
request_bytes.extend_from_slice(&buffer[..n]);
if request_bytes.windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
if request_bytes.len() >= 8192 {
break;
}
}
let request = String::from_utf8_lossy(&request_bytes);
let path = request.lines().next().unwrap_or("").split_whitespace().nth(1).unwrap_or("");
let (status, body) = if path.contains("/.well-known/openid-configuration") {
(200, discovery_body.as_str())
} else if path.contains("/jwks") {
(200, jwks_body)
} else {
(404, r#"{"error":"not found"}"#)
};
let response = format!(
"HTTP/1.1 {status} {}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
if status == 200 { "OK" } else { "Not Found" },
body.len()
);
let _ = stream.write_all(response.as_bytes());
let _ = stream.flush();
let _ = stream.shutdown(Shutdown::Both);
last_completed = Instant::now();
if seen >= max_requests {
break;
}
}
});
(base, handle)
}
fn discovery_error_contains_all_variants(err: &str, base: &str) -> bool {
err.contains(base) && err.contains(&format!("{base}/")) && err.contains("discovery failed for all issuer variants")
}
#[tokio::test]
async fn test_validate_oidc_provider_config_retries_with_issuer_candidates() {
// Discovery document must advertise the canonical issuer path. The first candidate has no
// trailing slash; openidconnect rejects issuer mismatch, then the second variant succeeds.
let (base, handle) = start_mock_oidc_discovery_server(|base| format!("{base}/application/o/rustfs/"), 8);
let config_url = format!("{base}/application/o/rustfs");
let config = build_mocked_oidc_provider_config("default", &config_url);
let result = validate_oidc_provider_config(&config).await;
let validation_result = result.expect("OIDC provider validation should succeed");
assert_eq!(validation_result.issuer, format!("{base}/application/o/rustfs/"));
assert!(handle.join().is_ok());
}
#[tokio::test]
async fn test_validate_oidc_provider_config_returns_detailed_errors() {
let (base, handle) = start_mock_oidc_discovery_server(|base| format!("{base}/application/o/other"), 8);
let config_url = format!("{base}/application/o/rustfs");
let config = build_mocked_oidc_provider_config("default", &config_url);
let err = validate_oidc_provider_config(&config)
.await
.expect_err("OIDC provider validation should fail");
assert!(discovery_error_contains_all_variants(&err, &base));
assert!(err.contains("issuer '"));
assert!(err.contains(&format!("issuer '{base}/application/o/rustfs'")));
assert!(err.contains(&format!("issuer '{base}/application/o/rustfs/'")));
assert!(handle.join().is_ok());
}
#[test]
fn test_decode_jwt_payload_invalid() {
assert!(decode_jwt_payload("not-a-jwt").is_empty());