use crate::connection::{WebSocketConnection, WebSocketError, WebSocketResult};
use async_trait::async_trait;
use std::sync::Arc;
pub type AuthResult<T> = Result<T, AuthError>;
#[derive(Debug, thiserror::Error)]
pub enum AuthError {
#[error("Authentication failed: {0}")]
AuthenticationFailed(String),
#[error("Authorization denied: {0}")]
AuthorizationDenied(String),
#[error("Invalid credentials")]
InvalidCredentials,
#[error("Token expired")]
TokenExpired,
#[error("Missing authentication")]
MissingAuthentication,
}
pub trait AuthUser: Send + Sync + std::fmt::Debug {
fn id(&self) -> &str;
fn username(&self) -> &str;
fn is_authenticated(&self) -> bool;
fn has_permission(&self, permission: &str) -> bool;
}
#[derive(Debug, Clone)]
pub struct SimpleAuthUser {
id: String,
username: String,
permissions: Vec<String>,
}
impl SimpleAuthUser {
pub fn new(id: String, username: String, permissions: Vec<String>) -> Self {
Self {
id,
username,
permissions,
}
}
}
impl AuthUser for SimpleAuthUser {
fn id(&self) -> &str {
&self.id
}
fn username(&self) -> &str {
&self.username
}
fn is_authenticated(&self) -> bool {
!self.id.is_empty()
}
fn has_permission(&self, permission: &str) -> bool {
self.permissions.contains(&permission.to_string())
}
}
#[async_trait]
pub trait WebSocketAuthenticator: Send + Sync {
async fn authenticate(
&self,
connection: &Arc<WebSocketConnection>,
credentials: &str,
) -> AuthResult<Box<dyn AuthUser>>;
}
pub struct TokenAuthenticator {
tokens: std::collections::HashMap<String, SimpleAuthUser>,
}
impl TokenAuthenticator {
pub fn new(tokens: Vec<(String, SimpleAuthUser)>) -> Self {
Self {
tokens: tokens.into_iter().collect(),
}
}
pub fn add_token(&mut self, token: String, user: SimpleAuthUser) {
self.tokens.insert(token, user);
}
pub fn remove_token(&mut self, token: &str) -> Option<SimpleAuthUser> {
self.tokens.remove(token)
}
}
#[async_trait]
impl WebSocketAuthenticator for TokenAuthenticator {
async fn authenticate(
&self,
_connection: &Arc<WebSocketConnection>,
credentials: &str,
) -> AuthResult<Box<dyn AuthUser>> {
self.tokens
.get(credentials)
.map(|user| Box::new(user.clone()) as Box<dyn AuthUser>)
.ok_or(AuthError::InvalidCredentials)
}
}
#[async_trait]
pub trait AuthorizationPolicy: Send + Sync {
async fn authorize(
&self,
user: &dyn AuthUser,
action: &str,
resource: Option<&str>,
) -> AuthResult<()>;
}
pub struct PermissionBasedPolicy {
action_permissions: std::collections::HashMap<String, String>,
}
impl PermissionBasedPolicy {
pub fn new(action_permissions: Vec<(String, String)>) -> Self {
Self {
action_permissions: action_permissions.into_iter().collect(),
}
}
pub fn add_permission(&mut self, action: String, permission: String) {
self.action_permissions.insert(action, permission);
}
}
#[async_trait]
impl AuthorizationPolicy for PermissionBasedPolicy {
async fn authorize(
&self,
user: &dyn AuthUser,
action: &str,
_resource: Option<&str>,
) -> AuthResult<()> {
let required_permission = self
.action_permissions
.get(action)
.ok_or_else(|| AuthError::AuthorizationDenied(format!("Unknown action: {}", action)))?;
if user.has_permission(required_permission) {
Ok(())
} else {
Err(AuthError::AuthorizationDenied(format!(
"Missing permission: {}",
required_permission
)))
}
}
}
pub struct AuthenticatedConnection {
connection: Arc<WebSocketConnection>,
user: Box<dyn AuthUser>,
}
impl AuthenticatedConnection {
pub fn new(connection: Arc<WebSocketConnection>, user: Box<dyn AuthUser>) -> Self {
Self { connection, user }
}
pub fn connection(&self) -> &Arc<WebSocketConnection> {
&self.connection
}
pub fn user(&self) -> &dyn AuthUser {
self.user.as_ref()
}
pub async fn send_with_auth<P: AuthorizationPolicy>(
&self,
message: crate::connection::Message,
policy: &P,
) -> WebSocketResult<()> {
policy
.authorize(self.user.as_ref(), "send_message", None)
.await
.map_err(|_| WebSocketError::Protocol("authorization failed".to_string()))?;
self.connection.send(message).await
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::connection::Message;
use tokio::sync::mpsc;
#[test]
fn test_simple_auth_user() {
let user = SimpleAuthUser::new(
"user_123".to_string(),
"alice".to_string(),
vec!["read".to_string(), "write".to_string()],
);
assert_eq!(user.id(), "user_123");
assert_eq!(user.username(), "alice");
assert!(user.is_authenticated());
assert!(user.has_permission("read"));
assert!(user.has_permission("write"));
assert!(!user.has_permission("admin"));
}
#[tokio::test]
async fn test_token_authenticator_valid() {
let user = SimpleAuthUser::new(
"user_1".to_string(),
"alice".to_string(),
vec!["chat.read".to_string()],
);
let authenticator = TokenAuthenticator::new(vec![("token123".to_string(), user)]);
let (tx, _rx) = mpsc::unbounded_channel();
let conn = Arc::new(WebSocketConnection::new("conn_1".to_string(), tx));
let auth_user = authenticator.authenticate(&conn, "token123").await.unwrap();
assert_eq!(auth_user.username(), "alice");
}
#[tokio::test]
async fn test_token_authenticator_invalid() {
let authenticator = TokenAuthenticator::new(vec![]);
let (tx, _rx) = mpsc::unbounded_channel();
let conn = Arc::new(WebSocketConnection::new("conn_1".to_string(), tx));
let result = authenticator.authenticate(&conn, "invalid_token").await;
assert!(result.is_err());
assert!(matches!(result.unwrap_err(), AuthError::InvalidCredentials));
}
#[tokio::test]
async fn test_permission_based_policy_authorized() {
let policy = PermissionBasedPolicy::new(vec![(
"send_message".to_string(),
"chat.write".to_string(),
)]);
let user = SimpleAuthUser::new(
"user_1".to_string(),
"alice".to_string(),
vec!["chat.write".to_string()],
);
let result = policy.authorize(&user, "send_message", None).await;
assert!(result.is_ok());
}
#[tokio::test]
async fn test_permission_based_policy_denied() {
let policy = PermissionBasedPolicy::new(vec![(
"delete_message".to_string(),
"chat.admin".to_string(),
)]);
let user = SimpleAuthUser::new(
"user_1".to_string(),
"alice".to_string(),
vec!["chat.write".to_string()],
);
let result = policy.authorize(&user, "delete_message", None).await;
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
AuthError::AuthorizationDenied(_)
));
}
#[tokio::test]
async fn test_authenticated_connection_send_with_auth() {
let policy = PermissionBasedPolicy::new(vec![(
"send_message".to_string(),
"chat.write".to_string(),
)]);
let user = SimpleAuthUser::new(
"user_1".to_string(),
"alice".to_string(),
vec!["chat.write".to_string()],
);
let (tx, mut rx) = mpsc::unbounded_channel();
let conn = Arc::new(WebSocketConnection::new("conn_1".to_string(), tx));
let auth_conn = AuthenticatedConnection::new(conn, Box::new(user));
let msg = Message::text("Hello".to_string());
auth_conn.send_with_auth(msg, &policy).await.unwrap();
assert!(matches!(rx.try_recv(), Ok(Message::Text { .. })));
}
#[tokio::test]
async fn test_authenticated_connection_send_with_auth_denied() {
let policy = PermissionBasedPolicy::new(vec![(
"send_message".to_string(),
"chat.admin".to_string(),
)]);
let user = SimpleAuthUser::new(
"user_1".to_string(),
"alice".to_string(),
vec!["chat.write".to_string()],
);
let (tx, _rx) = mpsc::unbounded_channel();
let conn = Arc::new(WebSocketConnection::new("conn_1".to_string(), tx));
let auth_conn = AuthenticatedConnection::new(conn, Box::new(user));
let msg = Message::text("Hello".to_string());
let result = auth_conn.send_with_auth(msg, &policy).await;
assert!(result.is_err());
}
}