use actix_web::{HttpMessage, HttpRequest};
use crate::http::security::User;
use super::error::WebSocketSecurityError;
use super::extractor::WebSocketUpgrade;
use super::origin::OriginValidator;
#[derive(Debug, Clone)]
pub struct WebSocketSecurityConfig {
origin_validator: OriginValidator,
require_authentication: bool,
required_roles: Vec<String>,
required_authorities: Vec<String>,
}
impl Default for WebSocketSecurityConfig {
fn default() -> Self {
Self::new()
}
}
impl WebSocketSecurityConfig {
pub fn new() -> Self {
Self {
origin_validator: OriginValidator::allow_any(),
require_authentication: false,
required_roles: Vec::new(),
required_authorities: Vec::new(),
}
}
pub fn allowed_origins(mut self, origins: Vec<String>) -> Self {
let origins_refs: Vec<&str> = origins.iter().map(|s| s.as_str()).collect();
self.origin_validator = OriginValidator::new(&origins_refs);
self
}
pub fn origin_validator(mut self, validator: OriginValidator) -> Self {
self.origin_validator = validator;
self
}
pub fn require_authentication(mut self, require: bool) -> Self {
self.require_authentication = require;
self
}
pub fn required_roles(mut self, roles: Vec<String>) -> Self {
self.required_roles = roles;
if !self.required_roles.is_empty() {
self.require_authentication = true;
}
self
}
pub fn required_authorities(mut self, authorities: Vec<String>) -> Self {
self.required_authorities = authorities;
if !self.required_authorities.is_empty() {
self.require_authentication = true;
}
self
}
pub fn validate_upgrade(
&self,
req: &HttpRequest,
) -> Result<WebSocketUpgrade, WebSocketSecurityError> {
self.origin_validator.validate(req)?;
let user = req.extensions().get::<User>().cloned();
if self.require_authentication && user.is_none() {
return Err(WebSocketSecurityError::Unauthorized);
}
if !self.required_roles.is_empty() {
let roles_refs: Vec<&str> = self.required_roles.iter().map(|s| s.as_str()).collect();
if !user.as_ref().is_some_and(|u| u.has_any_role(&roles_refs)) {
return Err(WebSocketSecurityError::MissingRole {
role: self.required_roles.join(", "),
});
}
}
if !self.required_authorities.is_empty() {
let auth_refs: Vec<&str> = self
.required_authorities
.iter()
.map(|s| s.as_str())
.collect();
if !user
.as_ref()
.is_some_and(|u| u.has_any_authority(&auth_refs))
{
return Err(WebSocketSecurityError::MissingAuthority {
authority: self.required_authorities.join(", "),
});
}
}
let origin = req
.headers()
.get("origin")
.and_then(|h| h.to_str().ok())
.map(|s| s.to_string());
Ok(WebSocketUpgrade::new(user, origin))
}
}
#[derive(Debug, Clone, Default)]
pub struct WebSocketSecurityConfigBuilder {
config: WebSocketSecurityConfig,
}
impl WebSocketSecurityConfigBuilder {
pub fn new() -> Self {
Self {
config: WebSocketSecurityConfig::new(),
}
}
pub fn allowed_origins(mut self, origins: Vec<String>) -> Self {
self.config = self.config.allowed_origins(origins);
self
}
pub fn origin_validator(mut self, validator: OriginValidator) -> Self {
self.config = self.config.origin_validator(validator);
self
}
pub fn require_authentication(mut self) -> Self {
self.config = self.config.require_authentication(true);
self
}
pub fn required_roles(mut self, roles: Vec<String>) -> Self {
self.config = self.config.required_roles(roles);
self
}
pub fn required_authorities(mut self, authorities: Vec<String>) -> Self {
self.config = self.config.required_authorities(authorities);
self
}
pub fn build(self) -> WebSocketSecurityConfig {
self.config
}
}
#[cfg(test)]
mod tests {
use super::*;
use actix_web::test::TestRequest;
fn create_request_with_user(user: User) -> HttpRequest {
let req = TestRequest::default()
.insert_header(("origin", "https://myapp.com"))
.to_http_request();
req.extensions_mut().insert(user);
req
}
#[test]
fn test_default_config_allows_all() {
let config = WebSocketSecurityConfig::new();
let req = TestRequest::default()
.insert_header(("origin", "https://any-origin.com"))
.to_http_request();
assert!(config.validate_upgrade(&req).is_ok());
}
#[test]
fn test_origin_validation() {
let config =
WebSocketSecurityConfig::new().allowed_origins(vec!["https://myapp.com".into()]);
let req = TestRequest::default()
.insert_header(("origin", "https://myapp.com"))
.to_http_request();
assert!(config.validate_upgrade(&req).is_ok());
let req = TestRequest::default()
.insert_header(("origin", "https://evil.com"))
.to_http_request();
assert!(config.validate_upgrade(&req).is_err());
}
#[test]
fn test_authentication_requirement() {
let config = WebSocketSecurityConfig::new()
.origin_validator(OriginValidator::allow_any())
.require_authentication(true);
let req = TestRequest::default().to_http_request();
assert!(matches!(
config.validate_upgrade(&req),
Err(WebSocketSecurityError::Unauthorized)
));
let user = User::new("testuser".into(), "password".into());
let req = create_request_with_user(user);
assert!(config.validate_upgrade(&req).is_ok());
}
#[test]
fn test_role_requirement() {
let config = WebSocketSecurityConfig::new()
.origin_validator(OriginValidator::allow_any())
.required_roles(vec!["ADMIN".into()]);
let user = User::new("user".into(), "password".into()).roles(&["USER".into()]);
let req = create_request_with_user(user);
assert!(matches!(
config.validate_upgrade(&req),
Err(WebSocketSecurityError::MissingRole { .. })
));
let admin = User::new("admin".into(), "password".into()).roles(&["ADMIN".into()]);
let req = create_request_with_user(admin);
assert!(config.validate_upgrade(&req).is_ok());
}
#[test]
fn test_authority_requirement() {
let config = WebSocketSecurityConfig::new()
.origin_validator(OriginValidator::allow_any())
.required_authorities(vec!["ws:connect".into()]);
let user = User::new("user".into(), "password".into());
let req = create_request_with_user(user);
assert!(matches!(
config.validate_upgrade(&req),
Err(WebSocketSecurityError::MissingAuthority { .. })
));
let ws_user =
User::new("user".into(), "password".into()).authorities(&["ws:connect".into()]);
let req = create_request_with_user(ws_user);
assert!(config.validate_upgrade(&req).is_ok());
}
#[test]
fn test_combined_requirements() {
let config = WebSocketSecurityConfig::new()
.allowed_origins(vec!["https://myapp.com".into()])
.required_roles(vec!["USER".into()])
.required_authorities(vec!["ws:connect".into()]);
let user = User::new("testuser".into(), "password".into())
.roles(&["USER".into()])
.authorities(&["ws:connect".into()]);
let req = TestRequest::default()
.insert_header(("origin", "https://myapp.com"))
.to_http_request();
req.extensions_mut().insert(user);
assert!(config.validate_upgrade(&req).is_ok());
}
#[test]
fn test_builder_pattern() {
let config = WebSocketSecurityConfigBuilder::new()
.allowed_origins(vec!["https://myapp.com".into()])
.require_authentication()
.required_roles(vec!["USER".into()])
.build();
let user = User::new("user".into(), "password".into()).roles(&["USER".into()]);
let req = TestRequest::default()
.insert_header(("origin", "https://myapp.com"))
.to_http_request();
req.extensions_mut().insert(user);
assert!(config.validate_upgrade(&req).is_ok());
}
}