use turbomcp_protocol::{Error as McpError, Result as McpResult};
pub fn validate_oauth_state(expected_state: &str, received_state: &str) -> McpResult<()> {
use sha2::{Digest, Sha256};
use subtle::ConstantTimeEq;
if expected_state.is_empty() || received_state.is_empty() {
return Err(McpError::invalid_params(
"OAuth state parameter cannot be empty".to_string(),
));
}
let expected_hash = Sha256::digest(expected_state.as_bytes());
let received_hash = Sha256::digest(received_state.as_bytes());
let is_equal = expected_hash.ct_eq(&received_hash);
if bool::from(is_equal) {
Ok(())
} else {
Err(McpError::invalid_params(
"OAuth state parameter mismatch - possible CSRF attack".to_string(),
))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_oauth_state_validation_success() {
let state = "random-csrf-token-123";
assert!(validate_oauth_state(state, state).is_ok());
}
#[test]
fn test_oauth_state_validation_mismatch() {
let expected = "state-abc123";
let received = "state-xyz789";
let result = validate_oauth_state(expected, received);
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("state parameter mismatch")
);
}
#[test]
fn test_oauth_state_validation_empty_expected() {
let result = validate_oauth_state("", "some-state");
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("cannot be empty"));
}
#[test]
fn test_oauth_state_validation_empty_received() {
let result = validate_oauth_state("some-state", "");
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("cannot be empty"));
}
#[test]
fn test_oauth_state_validation_case_sensitive() {
let result = validate_oauth_state("State123", "state123");
assert!(result.is_err());
}
}