use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
use x509_cert::Certificate;
use crate::error::{WxPayError, WxPayResult};
pub struct CertManager {
certificates: Arc<RwLock<HashMap<String, Certificate>>>,
cert_data: Arc<RwLock<HashMap<String, Vec<u8>>>>,
}
impl CertManager {
pub fn new() -> Self {
Self {
certificates: Arc::new(RwLock::new(HashMap::new())),
cert_data: Arc::new(RwLock::new(HashMap::new())),
}
}
pub async fn add_certificate(
&self,
serial_number: String,
cert_der: Vec<u8>,
) -> WxPayResult<()> {
use der::Decode;
let cert = Certificate::from_der(&cert_der)
.map_err(|e| WxPayError::CertificateParseError(format!("证书解析失败:{}", e)))?;
let mut certificates = self.certificates.write().await;
let mut cert_data = self.cert_data.write().await;
certificates.insert(serial_number.clone(), cert);
cert_data.insert(serial_number, cert_der);
Ok(())
}
pub async fn get_certificate(&self, serial_number: &str) -> Option<Certificate> {
let certificates = self.certificates.read().await;
certificates.get(serial_number).cloned()
}
pub async fn get_certificate_data(&self, serial_number: &str) -> Option<Vec<u8>> {
let cert_data = self.cert_data.read().await;
cert_data.get(serial_number).cloned()
}
pub async fn get_serial_numbers(&self) -> Vec<String> {
let certificates = self.certificates.read().await;
certificates.keys().cloned().collect()
}
pub async fn remove_certificate(&self, serial_number: &str) -> WxPayResult<()> {
let mut certificates = self.certificates.write().await;
let mut cert_data = self.cert_data.write().await;
certificates.remove(serial_number);
cert_data.remove(serial_number);
Ok(())
}
pub async fn clear(&self) {
let mut certificates = self.certificates.write().await;
let mut cert_data = self.cert_data.write().await;
certificates.clear();
cert_data.clear();
}
pub async fn count(&self) -> usize {
let certificates = self.certificates.read().await;
certificates.len()
}
pub async fn has_certificate(&self, serial_number: &str) -> bool {
let certificates = self.certificates.read().await;
certificates.contains_key(serial_number)
}
}
impl Default for CertManager {
fn default() -> Self {
Self::new()
}
}
impl std::fmt::Debug for CertManager {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CertManager")
.field("certificates", &"<certificates>")
.finish()
}
}
impl Clone for CertManager {
fn clone(&self) -> Self {
Self {
certificates: self.certificates.clone(),
cert_data: self.cert_data.clone(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_cert_manager_new() {
let manager = CertManager::new();
assert_eq!(manager.count().await, 0);
}
#[tokio::test]
async fn test_cert_manager_has_certificate() {
let manager = CertManager::new();
assert!(!manager.has_certificate("CERT123456").await);
}
#[tokio::test]
async fn test_cert_manager_get_serial_numbers() {
let manager = CertManager::new();
let serial_numbers = manager.get_serial_numbers().await;
assert!(serial_numbers.is_empty());
}
}