use crate::oauth::{OAuth2AuditEvent, OAuth2AuditLogger, OAuth2Error, TokenResponse};
use async_trait::async_trait;
use redis::AsyncCommands;
use std::sync::Arc;
#[async_trait]
pub trait OAuth2TokenStore: Send + Sync {
async fn store_token(&self, client_id: &str, token: &TokenResponse) -> Result<(), OAuth2Error>;
async fn get_token(&self, client_id: &str) -> Result<Option<TokenResponse>, OAuth2Error>;
async fn delete_token(&self, client_id: &str) -> Result<(), OAuth2Error>;
}
use parking_lot::Mutex;
use std::collections::HashMap;
#[derive(Default)]
pub struct MemoryOAuth2TokenStore {
tokens: Mutex<HashMap<String, TokenResponse>>,
}
impl MemoryOAuth2TokenStore {
pub fn new() -> Self {
Self::default()
}
}
#[async_trait]
impl OAuth2TokenStore for MemoryOAuth2TokenStore {
async fn store_token(&self, client_id: &str, token: &TokenResponse) -> Result<(), OAuth2Error> {
self.tokens
.lock()
.insert(client_id.to_string(), token.clone());
Ok(())
}
async fn get_token(&self, client_id: &str) -> Result<Option<TokenResponse>, OAuth2Error> {
Ok(self.tokens.lock().get(client_id).cloned())
}
async fn delete_token(&self, client_id: &str) -> Result<(), OAuth2Error> {
self.tokens.lock().remove(client_id);
Ok(())
}
}
pub struct RedisOAuth2TokenStore {
client: redis::Client,
audit_logger: Option<Arc<dyn OAuth2AuditLogger>>,
}
impl RedisOAuth2TokenStore {
pub fn new(client: redis::Client) -> Self {
Self {
client,
audit_logger: None,
}
}
pub fn with_audit_logger(mut self, logger: Arc<dyn OAuth2AuditLogger>) -> Self {
self.audit_logger = Some(logger);
self
}
fn key(client_id: &str) -> String {
format!("oauth2:token:{client_id}")
}
fn log_audit(
&self,
grant_type: &str,
result: &str,
alert_code: Option<&str>,
message: Option<&str>,
) {
if let Some(logger) = &self.audit_logger {
let event = OAuth2AuditEvent {
client_id: String::new(),
grant_type: grant_type.to_string(),
result: result.to_string(),
timestamp: chrono::Utc::now().timestamp(),
alert_code: alert_code.map(|s| s.to_string()),
message: message.map(|s| s.to_string()),
};
logger.log_event(&event);
}
}
}
#[async_trait]
impl OAuth2TokenStore for RedisOAuth2TokenStore {
async fn store_token(&self, client_id: &str, token: &TokenResponse) -> Result<(), OAuth2Error> {
let key = Self::key(client_id);
let value =
serde_json::to_string(token).map_err(|err| OAuth2Error::Serialize(err.to_string()))?;
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|err| {
self.log_audit(
"token_store",
"failure",
Some("OAUTH2_TOKEN_STORE_FAILED"),
Some(&err.to_string()),
);
OAuth2Error::HttpTransport(err.to_string())
})?;
let ttl = token.expires_in.unwrap_or(3600).max(1) as u64;
let result: redis::RedisResult<()> = redis::pipe()
.atomic()
.cmd("SET")
.arg(&key)
.arg(&value)
.arg("NX")
.arg("EX")
.arg(ttl)
.query_async(&mut conn)
.await;
match result {
Ok(()) => {
self.log_audit("token_store", "success", None, None);
Ok(())
}
Err(err) => {
self.log_audit(
"token_store",
"failure",
Some("OAUTH2_TOKEN_STORE_FAILED"),
Some(&err.to_string()),
);
Err(OAuth2Error::HttpTransport(err.to_string()))
}
}
}
async fn get_token(&self, client_id: &str) -> Result<Option<TokenResponse>, OAuth2Error> {
let key = Self::key(client_id);
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
let value: Option<String> = conn
.get(&key)
.await
.map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
match value {
Some(s) => {
let token: TokenResponse = serde_json::from_str(&s)
.map_err(|err| OAuth2Error::Serialize(err.to_string()))?;
Ok(Some(token))
}
None => Ok(None),
}
}
async fn delete_token(&self, client_id: &str) -> Result<(), OAuth2Error> {
let key = Self::key(client_id);
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
let _: () = conn
.del(&key)
.await
.map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_memory_token_store_set_get() {
let store = MemoryOAuth2TokenStore::new();
let token = TokenResponse {
access_token: "token123".into(),
token_type: Some("Bearer".into()),
expires_in: Some(3600),
scope: Some("read".into()),
refresh_token: Some("refresh456".into()),
};
store
.store_token("client1", &token)
.await
.expect("store_token 失败");
let retrieved = store
.get_token("client1")
.await
.expect("get_token 失败")
.expect("应查到 token");
assert_eq!(retrieved.access_token, "token123");
assert_eq!(retrieved.token_type.as_deref(), Some("Bearer"));
assert_eq!(retrieved.expires_in, Some(3600));
}
#[tokio::test]
async fn test_memory_token_store_delete() {
let store = MemoryOAuth2TokenStore::new();
let token = TokenResponse {
access_token: "token123".into(),
token_type: None,
expires_in: Some(3600),
scope: None,
refresh_token: None,
};
store
.store_token("client1", &token)
.await
.expect("store_token 失败");
store
.delete_token("client1")
.await
.expect("delete_token 失败");
let result = store.get_token("client1").await.expect("get_token 失败");
assert!(result.is_none(), "删除后应查不到 token");
}
#[tokio::test]
async fn test_memory_token_store_get_nonexistent() {
let store = MemoryOAuth2TokenStore::new();
let result = store
.get_token("nonexistent")
.await
.expect("get_token 失败");
assert!(result.is_none());
}
}