use actix_web::{http::StatusCode, HttpResponse, ResponseError};
use derive_more::{Display, Error};
#[derive(Debug, Display, Error)]
pub enum WebSocketSecurityError {
#[display("Unauthorized: authentication required for WebSocket connection")]
Unauthorized,
#[display("Forbidden: missing Origin header")]
MissingOrigin,
#[display("Forbidden: origin '{origin}' is not allowed")]
InvalidOrigin {
origin: String,
},
#[display("Forbidden: required role '{role}' not found")]
MissingRole {
role: String,
},
#[display("Forbidden: required authority '{authority}' not found")]
MissingAuthority {
authority: String,
},
}
impl ResponseError for WebSocketSecurityError {
fn status_code(&self) -> StatusCode {
match self {
WebSocketSecurityError::Unauthorized => StatusCode::UNAUTHORIZED,
WebSocketSecurityError::MissingOrigin
| WebSocketSecurityError::InvalidOrigin { .. }
| WebSocketSecurityError::MissingRole { .. }
| WebSocketSecurityError::MissingAuthority { .. } => StatusCode::FORBIDDEN,
}
}
fn error_response(&self) -> HttpResponse {
HttpResponse::build(self.status_code()).body(self.to_string())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_unauthorized_status() {
let err = WebSocketSecurityError::Unauthorized;
assert_eq!(err.status_code(), StatusCode::UNAUTHORIZED);
}
#[test]
fn test_invalid_origin_status() {
let err = WebSocketSecurityError::InvalidOrigin {
origin: "https://evil.com".into(),
};
assert_eq!(err.status_code(), StatusCode::FORBIDDEN);
assert!(err.to_string().contains("evil.com"));
}
#[test]
fn test_missing_origin_status() {
let err = WebSocketSecurityError::MissingOrigin;
assert_eq!(err.status_code(), StatusCode::FORBIDDEN);
}
#[test]
fn test_missing_role_status() {
let err = WebSocketSecurityError::MissingRole {
role: "ADMIN".into(),
};
assert_eq!(err.status_code(), StatusCode::FORBIDDEN);
assert!(err.to_string().contains("ADMIN"));
}
#[test]
fn test_missing_authority_status() {
let err = WebSocketSecurityError::MissingAuthority {
authority: "ws:connect".into(),
};
assert_eq!(err.status_code(), StatusCode::FORBIDDEN);
assert!(err.to_string().contains("ws:connect"));
}
}