use async_trait::async_trait;
use std::sync::Arc;
use crate::auth::Signer;
use crate::cert::CertManager;
use crate::crypto::Aes256GcmCipher;
use crate::error::{WxPayError, WxPayResult};
use crate::http::HttpClient;
use crate::utils::nonce::generate_nonce;
use crate::utils::timestamp::get_timestamp;
use serde::Deserialize;
use base64::Engine;
#[derive(Debug, Clone, Deserialize)]
#[allow(dead_code)]
struct EncryptedCertificate {
#[serde(default)]
algorithm: String,
#[serde(default)]
associated_data: String,
nonce: String,
ciphertext: String,
}
#[derive(Debug, Clone, Deserialize)]
struct CertificateEntry {
serial_no: String,
#[serde(default)]
_effective_time: Option<String>,
#[serde(default)]
_expire_time: Option<String>,
#[serde(default)]
certificate: Option<String>,
#[serde(default)]
encrypt_certificate: Option<EncryptedCertificate>,
}
#[async_trait]
pub trait CertificateDownloader: Send + Sync {
async fn download(&self) -> WxPayResult<Vec<(String, Vec<u8>)>>;
}
pub struct CertDownloader {
base_url: String,
merchant_id: String,
signer: Arc<dyn Signer>,
http_client: Arc<dyn HttpClient>,
cert_manager: Arc<CertManager>,
api_v3_key: Option<String>,
}
impl CertDownloader {
pub fn new(
base_url: impl Into<String>,
merchant_id: impl Into<String>,
signer: Arc<dyn Signer>,
http_client: Arc<dyn HttpClient>,
cert_manager: Arc<CertManager>,
) -> Self {
Self {
base_url: base_url.into(),
merchant_id: merchant_id.into(),
signer,
http_client,
cert_manager,
api_v3_key: None,
}
}
pub fn with_api_v3_key(mut self, api_v3_key: impl Into<String>) -> Self {
self.api_v3_key = Some(api_v3_key.into());
self
}
fn build_url(&self) -> String {
format!("{}/v3/certificates", self.base_url)
}
async fn build_headers(&self) -> WxPayResult<Vec<(String, String)>> {
let timestamp = get_timestamp();
let nonce = generate_nonce();
let url = "/v3/certificates";
let body = "";
let message = format!("GET\n{}\n{}\n{}\n{}\n", url, timestamp, nonce, body);
let signature = self.signer.sign(&message).await?;
let authorization = format!(
r#"WECHATPAY2-SHA256-RSA2048 mchid="{}",nonce_str="{}",timestamp="{}",serial_no="{}",signature="{}"#,
self.merchant_id,
nonce,
timestamp,
self.signer.cert_serial_number(),
signature
);
Ok(vec![
("Authorization".to_string(), authorization),
("Accept".to_string(), "application/json".to_string()),
("User-Agent".to_string(), "wxpay-rs/0.1.0".to_string()),
])
}
}
#[async_trait]
impl CertificateDownloader for CertDownloader {
async fn download(&self) -> WxPayResult<Vec<(String, Vec<u8>)>> {
let url = self.build_url();
let headers = self.build_headers().await?;
let response = self.http_client.get(&url, headers).await?;
if !response.is_success() {
return Err(WxPayError::CertificateDownloadError(format!(
"下载证书失败,HTTP 状态码: {}",
response.status
)));
}
let body = &response.body;
let response: serde_json::Value = serde_json::from_str(body)?;
let mut items: Vec<CertificateEntry> =
serde_json::from_value(response.clone()).or_else(|_| {
response
.get("data")
.cloned()
.and_then(|v| serde_json::from_value(v).ok())
.ok_or_else(|| {
WxPayError::CertificateParseError("证书响应解析失败".to_string())
})
})?;
if items.is_empty()
&& let Some(certs) = response.get("data").and_then(|v| v.as_array())
{
items = certs
.iter()
.filter_map(|item| serde_json::from_value::<CertificateEntry>(item.clone()).ok())
.collect();
}
let mut result = Vec::new();
for item in items {
let serial = item.serial_no.clone();
let cert_der = if let Some(cert_data) = item.certificate {
decode_certificate_der(&cert_data)?
} else if let Some(encrypted) = item.encrypt_certificate {
let cipher =
Aes256GcmCipher::new(self.api_v3_key.as_deref().ok_or_else(|| {
WxPayError::CertificateParseError("加密证书缺少 API v3 Key".to_string())
})?)?;
let plaintext = cipher.decrypt_notification(
&encrypted.nonce,
&encrypted.ciphertext,
&encrypted.associated_data,
)?;
decode_certificate_der(&plaintext)?
} else {
return Err(WxPayError::CertificateParseError(format!(
"证书 {} 无 certificate/ encrypt_certificate 字段",
item.serial_no
)));
};
self.cert_manager
.add_certificate(serial.to_string(), cert_der.clone())
.await?;
result.push((serial.to_string(), cert_der));
}
Ok(result)
}
}
fn decode_certificate_der(certificate_data: &str) -> WxPayResult<Vec<u8>> {
let trimmed = certificate_data.trim();
if trimmed.contains("BEGIN CERTIFICATE") {
let body = trimmed
.lines()
.filter(|line| !line.starts_with("-----"))
.collect::<String>();
return base64::engine::general_purpose::STANDARD
.decode(body)
.map_err(|e| WxPayError::CertificateParseError(format!("证书 PEM 解码失败:{}", e)));
}
base64::engine::general_purpose::STANDARD
.decode(trimmed)
.map_err(|e| WxPayError::CertificateParseError(format!("证书 Base64 解码失败: {}", e)))
}
impl std::fmt::Debug for CertDownloader {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CertDownloader")
.field("base_url", &self.base_url)
.field("merchant_id", &self.merchant_id)
.finish()
}
}
pub struct CertRefresher {
downloader: Arc<CertDownloader>,
interval: u64,
}
impl CertRefresher {
pub fn new(downloader: Arc<CertDownloader>, interval: u64) -> Self {
Self {
downloader,
interval,
}
}
pub fn start_auto_refresh(&self) {
let downloader = self.downloader.clone();
let interval = self.interval;
tokio::spawn(async move {
loop {
tokio::time::sleep(tokio::time::Duration::from_secs(interval)).await;
match downloader.download().await {
Ok(certificates) => {
tracing::info!("成功刷新 {} 个证书", certificates.len());
}
Err(e) => {
tracing::error!("刷新证书失败: {}", e);
}
}
}
});
}
}
impl std::fmt::Debug for CertRefresher {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("CertRefresher")
.field("interval", &self.interval)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_decode_certificate_der_plain_base64() {
let original = vec![0x30u8, 0x82, 0x01, 0x23, 0xAB, 0xCD];
let b64 = base64::engine::general_purpose::STANDARD.encode(&original);
let decoded = decode_certificate_der(&b64).unwrap();
assert_eq!(decoded, original);
}
#[test]
fn test_decode_certificate_der_pem_form() {
let original = vec![0xAAu8, 0xBB, 0xCC, 0xDD, 0xEE, 0xFF];
let b64 = base64::engine::general_purpose::STANDARD.encode(&original);
let pem = format!("-----BEGIN CERTIFICATE-----\n{b64}\n-----END CERTIFICATE-----");
let decoded = decode_certificate_der(&pem).unwrap();
assert_eq!(decoded, original);
}
#[test]
fn test_decode_certificate_der_rejects_invalid() {
let result = decode_certificate_der("!!!not-base64!!!");
assert!(matches!(result, Err(WxPayError::CertificateParseError(_))));
}
#[test]
fn test_decode_certificate_der_handles_whitespace() {
let original = vec![0x01u8, 0x02, 0x03];
let b64 = base64::engine::general_purpose::STANDARD.encode(&original);
let decoded = decode_certificate_der(&format!(" {b64} ")).unwrap();
assert_eq!(decoded, original);
}
}