fix(iam): redirect_uri dinamis + validasi CSRF state di OAuthServiceImpl

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
asepharyana
2026-07-17 09:08:41 +07:00
co-authored by Claude Sonnet 5
parent 910aa5e071
commit be278f8b1c
4 changed files with 140 additions and 52 deletions
@@ -9,6 +9,10 @@
//! `OAuthServiceImpl` drives the authorization-code + PKCE flow: //! `OAuthServiceImpl` drives the authorization-code + PKCE flow:
//! generating the verifier, building the auth URL, exchanging the code //! generating the verifier, building the auth URL, exchanging the code
//! for a token, and persisting the result via the injected repository. //! for a token, and persisting the result via the injected repository.
//! The CSRF `state` token and PKCE verifier are both persisted to sidecar
//! files next to `token_path` so `start_flow` and `complete_flow` can be
//! two separate calls (the caller — see `zesdex-backend`'s
//! `run_oauth_flow` — binds a real loopback listener in between).
use std::path::PathBuf; use std::path::PathBuf;
use base64::engine::general_purpose::URL_SAFE_NO_PAD; use base64::engine::general_purpose::URL_SAFE_NO_PAD;
@@ -18,10 +22,7 @@ use sha2::{Digest, Sha256};
use crate::domain::oauth::{OAuthConfig, OAuthToken}; use crate::domain::oauth::{OAuthConfig, OAuthToken};
use crate::domain::repository::OAuthRepository; use crate::domain::repository::OAuthRepository;
use crate::domain::service::OAuthService; use crate::domain::service::OAuthService;
use crate::infrastructure::rng::secure_token_hex;
// ---------------------------------------------------------------------------
// PKCE primitives (private to this use-case module)
// ---------------------------------------------------------------------------
const VERIFIER_LENGTH: usize = 64; const VERIFIER_LENGTH: usize = 64;
@@ -30,7 +31,7 @@ struct CodeVerifier(String);
impl CodeVerifier { impl CodeVerifier {
fn new() -> Self { fn new() -> Self {
let bytes: Vec<u8> = (0..VERIFIER_LENGTH).map(|_| rand_byte()).collect(); let bytes = hex::decode(secure_token_hex(VERIFIER_LENGTH)).unwrap_or_default();
CodeVerifier(URL_SAFE_NO_PAD.encode(&bytes)) CodeVerifier(URL_SAFE_NO_PAD.encode(&bytes))
} }
@@ -47,28 +48,11 @@ impl CodeVerifier {
} }
} }
/// Produce one pseudo-random byte from the system clock mixed with a
/// monotonic counter, providing ~64 bits of per-call unpredictability
/// without a full `rand` dependency.
fn rand_byte() -> u8 {
use std::sync::atomic::{AtomicU64, Ordering};
static COUNTER: AtomicU64 = AtomicU64::new(0);
let counter = COUNTER.fetch_add(1, Ordering::Relaxed);
let seed = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_nanos() as u64;
((seed ^ counter) & 0xFF) as u8
}
// ---------------------------------------------------------------------------
// OAuthServiceImpl
// ---------------------------------------------------------------------------
/// Concrete OAuth service backed by a generic token repository. /// Concrete OAuth service backed by a generic token repository.
/// ///
/// The code verifier is stored to a sidecar file (`token_path` with /// The code verifier and CSRF state token are each stored to a sidecar
/// `.verifier` extension) in `start_flow` and consumed in `complete_flow`. /// file (`token_path` with `.verifier`/`.state` extensions respectively)
/// in `start_flow` and consumed + deleted in `complete_flow`.
pub struct OAuthServiceImpl<R: OAuthRepository> { pub struct OAuthServiceImpl<R: OAuthRepository> {
pub token_repo: R, pub token_repo: R,
pub token_path: PathBuf, pub token_path: PathBuf,
@@ -86,36 +70,44 @@ impl<R: OAuthRepository> OAuthServiceImpl<R> {
} }
} }
/// Path to the sidecar file that holds the PKCE verifier between fn sidecar_path(&self, suffix: &str) -> PathBuf {
/// `start_flow` and `complete_flow`.
fn verifier_path(&self) -> PathBuf {
let mut p = self.token_path.clone(); let mut p = self.token_path.clone();
let ext = p let ext = p
.extension() .extension()
.map(|e| format!("{}.verifier", e.to_string_lossy())) .map(|e| format!("{}.{suffix}", e.to_string_lossy()))
.unwrap_or_else(|| "verifier".to_string()); .unwrap_or_else(|| suffix.to_string());
p.set_extension(ext); p.set_extension(ext);
p p
} }
fn verifier_path(&self) -> PathBuf {
self.sidecar_path("verifier")
}
fn state_path(&self) -> PathBuf {
self.sidecar_path("state")
}
} }
impl<R: OAuthRepository> OAuthService for OAuthServiceImpl<R> { impl<R: OAuthRepository> OAuthService for OAuthServiceImpl<R> {
fn start_flow(&self, config: &OAuthConfig) -> anyhow::Result<String> { fn start_flow(
&self,
config: &OAuthConfig,
redirect_uri: &str,
) -> anyhow::Result<(String, String)> {
if config.auth_url.is_empty() { if config.auth_url.is_empty() {
anyhow::bail!("OAuth auth_url is empty"); anyhow::bail!("OAuth auth_url is empty");
} }
let verifier = CodeVerifier::new(); let verifier = CodeVerifier::new();
let challenge = verifier.challenge(); let challenge = verifier.challenge();
let state = secure_token_hex(16);
// Persist the verifier so complete_flow can retrieve it.
if let Some(parent) = self.token_path.parent() { if let Some(parent) = self.token_path.parent() {
std::fs::create_dir_all(parent)?; std::fs::create_dir_all(parent)?;
} }
std::fs::write(self.verifier_path(), verifier.as_str())?; std::fs::write(self.verifier_path(), verifier.as_str())?;
std::fs::write(self.state_path(), &state)?;
// Persist a random state token for CSRF protection.
let state = uuid::Uuid::new_v4().to_string();
let mut url = url::Url::parse(&config.auth_url) let mut url = url::Url::parse(&config.auth_url)
.map_err(|e| anyhow::anyhow!("invalid auth_url '{}': {e}", config.auth_url))?; .map_err(|e| anyhow::anyhow!("invalid auth_url '{}': {e}", config.auth_url))?;
@@ -123,27 +115,38 @@ impl<R: OAuthRepository> OAuthService for OAuthServiceImpl<R> {
url.query_pairs_mut() url.query_pairs_mut()
.append_pair("response_type", "code") .append_pair("response_type", "code")
.append_pair("client_id", &config.client_id) .append_pair("client_id", &config.client_id)
.append_pair("redirect_uri", "http://127.0.0.1:0/callback") .append_pair("redirect_uri", redirect_uri)
.append_pair("scope", &config.scopes.join(" ")) .append_pair("scope", &config.scopes.join(" "))
.append_pair("state", &state) .append_pair("state", &state)
.append_pair("code_challenge_method", "S256") .append_pair("code_challenge_method", "S256")
.append_pair("code_challenge", &challenge); .append_pair("code_challenge", &challenge);
Ok(url.to_string()) Ok((url.to_string(), state))
} }
fn complete_flow(&self, config: &OAuthConfig, code: &str) -> anyhow::Result<OAuthToken> { fn complete_flow(
// Load the verifier that was stored during start_flow. &self,
config: &OAuthConfig,
redirect_uri: &str,
code: &str,
state: &str,
) -> anyhow::Result<OAuthToken> {
let state_path = self.state_path();
let expected_state = std::fs::read_to_string(&state_path)
.map_err(|e| anyhow::anyhow!("failed to read persisted OAuth state: {e}"))?;
if expected_state != state {
anyhow::bail!("OAuth state mismatch \u{2014} possible CSRF attack");
}
let verifier_path = self.verifier_path(); let verifier_path = self.verifier_path();
let verifier = std::fs::read_to_string(&verifier_path) let verifier = std::fs::read_to_string(&verifier_path)
.map_err(|e| anyhow::anyhow!("failed to read PKCE verifier: {e}"))?; .map_err(|e| anyhow::anyhow!("failed to read PKCE verifier: {e}"))?;
// Exchange the authorization code for a token.
let client = reqwest::blocking::Client::new(); let client = reqwest::blocking::Client::new();
let mut params = std::collections::HashMap::new(); let mut params = std::collections::HashMap::new();
params.insert("grant_type", "authorization_code"); params.insert("grant_type", "authorization_code");
params.insert("code", code); params.insert("code", code);
params.insert("redirect_uri", "http://127.0.0.1:0/callback"); params.insert("redirect_uri", redirect_uri);
params.insert("client_id", &config.client_id); params.insert("client_id", &config.client_id);
params.insert("code_verifier", &verifier); params.insert("code_verifier", &verifier);
@@ -186,9 +189,9 @@ impl<R: OAuthRepository> OAuthService for OAuthServiceImpl<R> {
.to_string(), .to_string(),
}; };
// Persist the token and clean up the verifier.
self.token_repo.save_token(&self.token_path, &token)?; self.token_repo.save_token(&self.token_path, &token)?;
let _ = std::fs::remove_file(&verifier_path); let _ = std::fs::remove_file(&verifier_path);
let _ = std::fs::remove_file(&state_path);
Ok(token) Ok(token)
} }
@@ -197,3 +200,69 @@ impl<R: OAuthRepository> OAuthService for OAuthServiceImpl<R> {
self.token_repo.load_token(&self.token_path) self.token_repo.load_token(&self.token_path)
} }
} }
#[cfg(test)]
mod tests {
use super::*;
use crate::domain::repository::OAuthRepository;
use std::cell::RefCell;
use std::path::PathBuf;
#[derive(Default)]
struct FakeOAuthRepo {
saved: RefCell<Option<OAuthToken>>,
}
impl OAuthRepository for FakeOAuthRepo {
fn save_token(&self, _path: &std::path::Path, token: &OAuthToken) -> anyhow::Result<()> {
*self.saved.borrow_mut() = Some(token.clone());
Ok(())
}
fn load_token(&self, _path: &std::path::Path) -> anyhow::Result<Option<OAuthToken>> {
Ok(self.saved.borrow().clone())
}
}
fn tmp_token_path() -> PathBuf {
std::env::temp_dir().join(format!("zesdex-iam-oauth-test-{}", uuid::Uuid::new_v4()))
}
#[test]
fn complete_flow_rejects_mismatched_state() {
let svc = OAuthServiceImpl::new(FakeOAuthRepo::default(), tmp_token_path());
let config = OAuthConfig {
auth_url: "https://example.test/authorize".to_string(),
..OAuthConfig::default()
};
let (_, _real_state) = svc
.start_flow(&config, "http://127.0.0.1:12345/callback")
.expect("start_flow should succeed");
let result = svc.complete_flow(
&config,
"http://127.0.0.1:12345/callback",
"some-code",
"attacker-supplied-state",
);
assert!(
result.is_err(),
"complete_flow must reject a state that doesn't match what start_flow persisted"
);
}
#[test]
fn start_flow_returns_url_containing_the_real_redirect_uri() {
let svc = OAuthServiceImpl::new(FakeOAuthRepo::default(), tmp_token_path());
let config = OAuthConfig {
auth_url: "https://example.test/authorize".to_string(),
..OAuthConfig::default()
};
let (auth_url, state) = svc
.start_flow(&config, "http://127.0.0.1:54321/callback")
.expect("start_flow should succeed");
assert!(
auth_url.contains("127.0.0.1%3A54321") || auth_url.contains("127.0.0.1:54321"),
"auth_url must embed the real dynamic redirect_uri, not a hardcoded port-0 placeholder: {auth_url}"
);
assert!(!state.is_empty());
}
}
+21 -6
View File
@@ -23,13 +23,28 @@ pub trait SessionService {
/// OAuth flow use-case boundary. /// OAuth flow use-case boundary.
pub trait OAuthService { pub trait OAuthService {
/// Start an OAuth authorization-code + PKCE flow. /// Start an OAuth authorization-code + PKCE flow for the given
/// Returns the provider's authorization URL to visit. /// `redirect_uri` (the caller is responsible for actually listening on
fn start_flow(&self, config: &OAuthConfig) -> anyhow::Result<String>; /// it — e.g. a bound `LoopbackServer`). Returns `(auth_url, state)`:
/// the URL to send the user to, and the CSRF state token that must be
/// passed back into `complete_flow` unchanged.
fn start_flow(
&self,
config: &OAuthConfig,
redirect_uri: &str,
) -> anyhow::Result<(String, String)>;
/// Complete the OAuth flow by exchanging an authorization code for a /// Complete the OAuth flow: validates `state` against the value
/// token. /// persisted during `start_flow` (bailing on mismatch — this is the
fn complete_flow(&self, config: &OAuthConfig, code: &str) -> anyhow::Result<OAuthToken>; /// CSRF check), then exchanges `code` for a token using the same
/// `redirect_uri` passed to `start_flow`.
fn complete_flow(
&self,
config: &OAuthConfig,
redirect_uri: &str,
code: &str,
state: &str,
) -> anyhow::Result<OAuthToken>;
/// Retrieve the currently stored OAuth token (if any). /// Retrieve the currently stored OAuth token (if any).
fn get_token(&self) -> anyhow::Result<Option<OAuthToken>>; fn get_token(&self) -> anyhow::Result<Option<OAuthToken>>;
@@ -41,19 +41,23 @@ pub struct SessionListResponse {
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OAuthStartRequest { pub struct OAuthStartRequest {
pub config: OAuthConfig, pub config: OAuthConfig,
pub redirect_uri: String,
} }
/// Response containing the authorization URL for an OAuth flow. /// Response containing the authorization URL and CSRF state for an OAuth flow.
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OAuthStartResponse { pub struct OAuthStartResponse {
pub auth_url: String, pub auth_url: String,
pub state: String,
} }
/// Request body for completing an OAuth flow with an authorization code. /// Request body for completing an OAuth flow with an authorization code.
#[derive(Debug, Clone, Serialize, Deserialize)] #[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OAuthCompleteRequest { pub struct OAuthCompleteRequest {
pub config: OAuthConfig, pub config: OAuthConfig,
pub redirect_uri: String,
pub code: String, pub code: String,
pub state: String,
} }
/// Response containing the acquired OAuth token. /// Response containing the acquired OAuth token.
@@ -49,8 +49,8 @@ pub fn handle_start_oauth<O: OAuthService>(
service: &O, service: &O,
req: OAuthStartRequest, req: OAuthStartRequest,
) -> anyhow::Result<OAuthStartResponse> { ) -> anyhow::Result<OAuthStartResponse> {
let auth_url = service.start_flow(&req.config)?; let (auth_url, state) = service.start_flow(&req.config, &req.redirect_uri)?;
Ok(OAuthStartResponse { auth_url }) Ok(OAuthStartResponse { auth_url, state })
} }
/// Handle a complete-OAuth-flow request. /// Handle a complete-OAuth-flow request.
@@ -58,7 +58,7 @@ pub fn handle_complete_oauth<O: OAuthService>(
service: &O, service: &O,
req: OAuthCompleteRequest, req: OAuthCompleteRequest,
) -> anyhow::Result<OAuthTokenResponse> { ) -> anyhow::Result<OAuthTokenResponse> {
let token = service.complete_flow(&req.config, &req.code)?; let token = service.complete_flow(&req.config, &req.redirect_uri, &req.code, &req.state)?;
Ok(OAuthTokenResponse { token }) Ok(OAuthTokenResponse { token })
} }