use crate::error::SaTokenError;
use crate::event::SaTokenEvent;
use crate::manager::SaTokenManager;
use crate::online::OnlineUser;
use crate::token::TokenValue;
use async_trait::async_trait;
use std::collections::HashMap;
use std::sync::Arc;
#[derive(Debug, Clone)]
pub struct WsAuthInfo {
pub login_id: String,
pub token: String,
pub session_id: String,
pub connect_time: chrono::DateTime<chrono::Utc>,
pub metadata: HashMap<String, String>,
}
#[async_trait]
pub trait WsTokenExtractor: Send + Sync {
async fn extract_token(
&self,
headers: &HashMap<String, String>,
query: &HashMap<String, String>,
) -> Option<String>;
}
pub struct DefaultWsTokenExtractor;
#[async_trait]
impl WsTokenExtractor for DefaultWsTokenExtractor {
async fn extract_token(
&self,
_headers: &HashMap<String, String>,
_query: &HashMap<String, String>,
) -> Option<String> {
None
}
}
impl std::fmt::Debug for DefaultWsTokenExtractor {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("DefaultWsTokenExtractor { .. }")
}
}
pub struct WsAuthManager {
manager: Arc<SaTokenManager>,
extractor: Arc<dyn WsTokenExtractor>,
}
impl std::fmt::Debug for WsAuthManager {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("WsAuthManager { .. }")
}
}
impl WsAuthManager {
pub fn new(manager: Arc<SaTokenManager>) -> Self {
Self {
manager,
extractor: Arc::new(DefaultWsTokenExtractor),
}
}
pub fn with_extractor(
manager: Arc<SaTokenManager>,
extractor: Arc<dyn WsTokenExtractor>,
) -> Self {
Self { manager, extractor }
}
pub async fn authenticate(
&self,
headers: &HashMap<String, String>,
query: &HashMap<String, String>,
) -> Result<WsAuthInfo, SaTokenError> {
let token = match self.extractor.extract_token(headers, query).await {
Some(s) => crate::token_io::apply_token_prefix(
s.trim(),
self.manager.config.token_prefix.as_deref(),
),
None => crate::token_io::read_token_from_maps(headers, query, &self.manager.config),
};
let token_str = token.ok_or(SaTokenError::NotLogin)?;
let token = TokenValue::new(token_str.clone());
let token_info = self.manager.get_token_info(&token).await?;
if let Some(expire_time) = token_info.expire_time
&& chrono::Utc::now() > expire_time
{
return Err(SaTokenError::TokenExpired);
}
let login_id = token_info.login_id.to_string();
let session_id = format!("ws:{}:{}", login_id, uuid::Uuid::new_v4());
let auth_info = WsAuthInfo {
login_id: login_id.clone(),
token: token_str.clone(),
session_id,
connect_time: chrono::Utc::now(),
metadata: HashMap::new(),
};
if let Some(online) = self.manager.online_manager() {
let user = OnlineUser {
login_type: token_info.login_type.to_string(),
login_id: login_id.clone(),
token: token_str.clone(),
device: token_info.device.clone().unwrap_or_else(|| "ws".into()),
connect_time: auth_info.connect_time,
last_activity: chrono::Utc::now(),
metadata: HashMap::new(),
};
if let Err(e) = online.mark_online(user).await {
tracing::warn!(error = %e, "failed to mark websocket presence");
}
}
let event = SaTokenEvent::login(&login_id, &token_str).with_login_type("websocket");
self.manager.event_bus().publish(event).await;
Ok(auth_info)
}
pub async fn verify_token(&self, token: &str) -> Result<String, SaTokenError> {
let token_value = TokenValue::new(token);
let token_info = self.manager.get_token_info(&token_value).await?;
if let Some(expire_time) = token_info.expire_time
&& chrono::Utc::now() > expire_time
{
return Err(SaTokenError::TokenExpired);
}
Ok(token_info.login_id.to_string())
}
pub async fn refresh_ws_session(&self, auth_info: &WsAuthInfo) -> Result<(), SaTokenError> {
let token = TokenValue::new(auth_info.token.clone());
let info = self
.manager
.token_repo()
.load_token_info_no_renew(&token)
.await?;
if self.manager.token_repo().should_auto_renew(&info) {
self.manager
.token_repo()
.apply_auto_renew(auth_info.token.as_str(), info.clone())
.await?;
}
if let Some(online) = self.manager.online_manager() {
let _ = online
.update_activity_with_type(&info.login_type, &auth_info.login_id, &auth_info.token)
.await;
}
Ok(())
}
pub async fn end_ws_session(&self, auth_info: &WsAuthInfo) -> Result<(), SaTokenError> {
if let Some(online) = self.manager.online_manager() {
let info = self
.manager
.get_token_info(&TokenValue::new(auth_info.token.clone()))
.await
.ok();
let login_type = info
.as_ref()
.map(|i| i.login_type.as_ref())
.unwrap_or(crate::keys::LOGIN_TYPE_DEFAULT);
online
.mark_offline_with_type(login_type, &auth_info.login_id, &auth_info.token)
.await?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::SaTokenConfig;
use sa_token_storage_memory::MemoryStorage;
#[tokio::test]
async fn test_ws_auth_manager() {
let config = SaTokenConfig::default();
let storage = Arc::new(MemoryStorage::new());
let manager = Arc::new(SaTokenManager::new(storage, config));
let ws_manager = WsAuthManager::new(manager.clone());
let token = manager.login("user123").await.unwrap();
let mut headers = HashMap::new();
headers.insert(
"Authorization".to_string(),
format!("Bearer {}", token.as_str()),
);
let auth_info = ws_manager
.authenticate(&headers, &HashMap::new())
.await
.unwrap();
assert_eq!(auth_info.login_id, "user123");
}
#[tokio::test]
async fn test_token_extraction_from_query() {
let config = SaTokenConfig::default();
let storage = Arc::new(MemoryStorage::new());
let manager = Arc::new(SaTokenManager::new(storage, config));
let ws_manager = WsAuthManager::new(manager.clone());
let token = manager.login("user456").await.unwrap();
let mut query = HashMap::new();
query.insert("token".to_string(), token.as_str().to_string());
let auth_info = ws_manager
.authenticate(&HashMap::new(), &query)
.await
.unwrap();
assert_eq!(auth_info.login_id, "user456");
}
#[tokio::test]
async fn test_verify_token() {
let config = SaTokenConfig::default();
let storage = Arc::new(MemoryStorage::new());
let manager = Arc::new(SaTokenManager::new(storage, config));
let ws_manager = WsAuthManager::new(manager.clone());
let token = manager.login("user789").await.unwrap();
let login_id = ws_manager.verify_token(token.as_str()).await.unwrap();
assert_eq!(login_id, "user789");
}
}