use crate::bindings::error::BindingError;
pub const RELAY_STATE_MAX_BYTES: usize = 80;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RelayState(String);
impl RelayState {
pub fn new(value: &str) -> Result<Self, BindingError> {
validate_relay_state(value)?;
Ok(RelayState(value.to_string()))
}
pub fn echo(value: &str) -> Self {
RelayState(value.to_string())
}
pub fn as_str(&self) -> &str {
&self.0
}
pub fn into_string(self) -> String {
self.0
}
}
impl std::fmt::Display for RelayState {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.0)
}
}
pub fn validate_relay_state(value: &str) -> Result<(), BindingError> {
if value.len() > RELAY_STATE_MAX_BYTES {
return Err(BindingError::RelayStateTooLong(value.len()));
}
if value.contains('\0') {
return Err(BindingError::RelayStateUnsafe(
"contains null bytes".to_string(),
));
}
let lower = value.to_ascii_lowercase();
let trimmed = lower.trim();
for scheme in &["javascript:", "data:", "vbscript:"] {
if trimmed.starts_with(scheme) {
return Err(BindingError::RelayStateUnsafe(format!(
"dangerous URI scheme: {}",
scheme
)));
}
}
if contains_html_tags(value) {
return Err(BindingError::RelayStateUnsafe(
"contains HTML tags".to_string(),
));
}
Ok(())
}
fn contains_html_tags(s: &str) -> bool {
let bytes = s.as_bytes();
for i in 0..bytes.len().saturating_sub(1) {
if bytes[i] == b'<' {
let next = bytes[i + 1];
if next.is_ascii_alphabetic() || next == b'/' {
return true;
}
}
}
false
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_valid_relay_state() {
let rs = RelayState::new("abc123").unwrap();
assert_eq!(rs.as_str(), "abc123");
}
#[test]
fn test_relay_state_max_length() {
let long = "a".repeat(80);
assert!(RelayState::new(&long).is_ok());
let too_long = "a".repeat(81);
assert!(matches!(
RelayState::new(&too_long),
Err(BindingError::RelayStateTooLong(81))
));
}
#[test]
fn test_relay_state_javascript_xss() {
assert!(RelayState::new("javascript:alert(1)").is_err());
assert!(RelayState::new("JAVASCRIPT:alert(1)").is_err());
assert!(RelayState::new(" javascript:alert(1)").is_err());
}
#[test]
fn test_relay_state_data_uri() {
assert!(RelayState::new("data:text/html,<script>").is_err());
}
#[test]
fn test_relay_state_vbscript() {
assert!(RelayState::new("vbscript:msgbox").is_err());
}
#[test]
fn test_relay_state_html_tags() {
assert!(RelayState::new("<script>alert(1)</script>").is_err());
assert!(RelayState::new("<img src=x onerror=alert(1)>").is_err());
}
#[test]
fn test_relay_state_null_bytes() {
assert!(RelayState::new("abc\0def").is_err());
}
#[test]
fn test_relay_state_echo_no_validation() {
let rs = RelayState::echo("javascript:alert(1)");
assert_eq!(rs.as_str(), "javascript:alert(1)");
}
#[test]
fn test_relay_state_url_safe_token() {
let rs = RelayState::new("ss:mem:6a25b4c3e2d1f0a9b8c7d6e5f4a3b2c1").unwrap();
assert_eq!(rs.as_str(), "ss:mem:6a25b4c3e2d1f0a9b8c7d6e5f4a3b2c1");
}
#[test]
fn test_relay_state_less_than_not_tag() {
assert!(RelayState::new("a < 5").is_ok());
assert!(RelayState::new("a<5").is_ok());
}
#[test]
fn test_relay_state_display() {
let rs = RelayState::new("token123").unwrap();
assert_eq!(format!("{}", rs), "token123");
}
#[test]
fn test_relay_state_into_string() {
let rs = RelayState::new("token123").unwrap();
assert_eq!(rs.into_string(), "token123");
}
}