use actix_web::HttpRequest;
use crate::http::security::User;
use super::error::WebSocketSecurityError;
#[derive(Debug, Clone)]
pub struct WebSocketUser(User);
impl WebSocketUser {
pub fn extract(req: &HttpRequest) -> Result<Self, WebSocketSecurityError> {
use actix_web::HttpMessage;
req.extensions()
.get::<User>()
.cloned()
.map(WebSocketUser)
.ok_or(WebSocketSecurityError::Unauthorized)
}
pub fn try_extract(req: &HttpRequest) -> Option<Self> {
use actix_web::HttpMessage;
req.extensions().get::<User>().cloned().map(WebSocketUser)
}
pub fn into_inner(self) -> User {
self.0
}
pub fn as_user(&self) -> &User {
&self.0
}
pub fn has_role(&self, role: &str) -> bool {
self.0.has_role(role)
}
pub fn has_any_role(&self, roles: &[&str]) -> bool {
self.0.has_any_role(roles)
}
pub fn has_authority(&self, authority: &str) -> bool {
self.0.has_authority(authority)
}
pub fn has_any_authority(&self, authorities: &[&str]) -> bool {
self.0.has_any_authority(authorities)
}
pub fn get_username(&self) -> &str {
self.0.get_username()
}
pub fn require_role(self, role: &str) -> Result<Self, WebSocketSecurityError> {
if self.has_role(role) {
Ok(self)
} else {
Err(WebSocketSecurityError::MissingRole {
role: role.to_string(),
})
}
}
pub fn require_any_role(self, roles: &[&str]) -> Result<Self, WebSocketSecurityError> {
if self.has_any_role(roles) {
Ok(self)
} else {
Err(WebSocketSecurityError::MissingRole {
role: roles.join(", "),
})
}
}
pub fn require_authority(self, authority: &str) -> Result<Self, WebSocketSecurityError> {
if self.has_authority(authority) {
Ok(self)
} else {
Err(WebSocketSecurityError::MissingAuthority {
authority: authority.to_string(),
})
}
}
pub fn require_any_authority(
self,
authorities: &[&str],
) -> Result<Self, WebSocketSecurityError> {
if self.has_any_authority(authorities) {
Ok(self)
} else {
Err(WebSocketSecurityError::MissingAuthority {
authority: authorities.join(", "),
})
}
}
}
impl std::ops::Deref for WebSocketUser {
type Target = User;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl From<WebSocketUser> for User {
fn from(ws_user: WebSocketUser) -> Self {
ws_user.0
}
}
#[derive(Debug, Clone)]
pub struct WebSocketUpgrade {
user: Option<User>,
origin: Option<String>,
}
impl WebSocketUpgrade {
pub fn new(user: Option<User>, origin: Option<String>) -> Self {
Self { user, origin }
}
pub fn user(&self) -> Option<&User> {
self.user.as_ref()
}
pub fn into_user(self) -> Option<User> {
self.user
}
pub fn origin(&self) -> Option<&str> {
self.origin.as_deref()
}
pub fn is_authenticated(&self) -> bool {
self.user.is_some()
}
pub fn username(&self) -> Option<&str> {
self.user.as_ref().map(|u| u.get_username())
}
}
#[cfg(test)]
mod tests {
use super::*;
use actix_web::{test::TestRequest, HttpMessage};
#[test]
fn test_websocket_user_extract_success() {
let user = User::new("testuser".into(), "password".into());
let req = TestRequest::default().to_http_request();
req.extensions_mut().insert(user.clone());
let ws_user = WebSocketUser::extract(&req).unwrap();
assert_eq!(ws_user.get_username(), "testuser");
}
#[test]
fn test_websocket_user_extract_unauthorized() {
let req = TestRequest::default().to_http_request();
let result = WebSocketUser::extract(&req);
assert!(matches!(result, Err(WebSocketSecurityError::Unauthorized)));
}
#[test]
fn test_websocket_user_try_extract() {
let req = TestRequest::default().to_http_request();
assert!(WebSocketUser::try_extract(&req).is_none());
let user = User::new("testuser".into(), "password".into());
req.extensions_mut().insert(user);
assert!(WebSocketUser::try_extract(&req).is_some());
}
#[test]
fn test_websocket_user_require_role() {
let user = User::new("admin".into(), "password".into()).roles(&["ADMIN".into()]);
let req = TestRequest::default().to_http_request();
req.extensions_mut().insert(user);
let ws_user = WebSocketUser::extract(&req).unwrap();
assert!(ws_user.clone().require_role("ADMIN").is_ok());
assert!(matches!(
ws_user.require_role("SUPERADMIN"),
Err(WebSocketSecurityError::MissingRole { role }) if role == "SUPERADMIN"
));
}
#[test]
fn test_websocket_user_require_authority() {
let user = User::new("user".into(), "password".into()).authorities(&["ws:connect".into()]);
let req = TestRequest::default().to_http_request();
req.extensions_mut().insert(user);
let ws_user = WebSocketUser::extract(&req).unwrap();
assert!(ws_user.clone().require_authority("ws:connect").is_ok());
assert!(matches!(
ws_user.require_authority("ws:admin"),
Err(WebSocketSecurityError::MissingAuthority { authority }) if authority == "ws:admin"
));
}
#[test]
fn test_websocket_upgrade() {
let user = User::new("testuser".into(), "password".into());
let upgrade = WebSocketUpgrade::new(Some(user), Some("https://myapp.com".into()));
assert!(upgrade.is_authenticated());
assert_eq!(upgrade.username(), Some("testuser"));
assert_eq!(upgrade.origin(), Some("https://myapp.com"));
}
#[test]
fn test_websocket_upgrade_anonymous() {
let upgrade = WebSocketUpgrade::new(None, Some("https://myapp.com".into()));
assert!(!upgrade.is_authenticated());
assert_eq!(upgrade.username(), None);
assert_eq!(upgrade.origin(), Some("https://myapp.com"));
}
}