Skip to main content

sz_rust_auth_facade/
oauth.rs

1//! OAuth2 模块 — 对齐 Laravel Socialite
2//!
3//! 提供 OAuth2 客户端抽象,对齐 Laravel Socialite `ProviderInterface` 的核心 API。
4//! 支持通用 OAuth2 提供商(QQ、微信、GitHub、Google 等)。
5//!
6//! ## Laravel Socialite 对齐
7//!
8//! ### 核心 API 映射
9//!
10//! | Laravel Socialite 方法 | Rust 方法 | 说明 |
11//! |-----------------------|-----------|------|
12//! | `Socialite::driver('qq')->redirect()` | [`OAuth2Provider::redirect_url`] | 生成授权 URL |
13//! | `Socialite::driver('qq')->user()` | [`OAuth2Provider::user_from_token`] | 用授权码换取用户信息 |
14//! | `$user->getId()` | [`SocialiteUser::id`] | 第三方用户 ID |
15//! | `$user->getNickname()` | [`SocialiteUser::nickname`] | 用户昵称 |
16//! | `$user->getName()` | [`SocialiteUser::name`] | 用户姓名 |
17//! | `$user->getEmail()` | [`SocialiteUser::email`] | 邮箱 |
18//! | `$user->getAvatar()` | [`SocialiteUser::avatar`] | 头像 URL |
19//! | `$user->token` | [`SocialiteUser::access_token`] | 访问令牌 |
20//! | `$user->refreshToken` | [`SocialiteUser::refresh_token`] | 刷新令牌 |
21//! | `$user->expiresIn` | [`SocialiteUser::expires_in`] | 令牌过期时间 |
22//!
23//! ### Laravel 行为对齐
24//!
25//! - **配置驱动**:Laravel 通过 `config/services.php` 配置 `client_id`/`client_secret`/`redirect`。
26//!   Rust 通过 [`OAuth2Config`] builder 表达相同配置项。
27//! - **state 参数**:Laravel 自动生成 CSRF 防护的 state 参数。Rust 由调用方传入 state
28//!   (对齐 Laravel `Socialite::driver('qq')->withState($state)->redirect()`)。
29//! - **scopes**:Laravel 支持 `scopes()` 设置多个 scope。Rust 通过 [`OAuth2Config::with_scopes`]
30//!   或 [`OAuth2Config::with_scope`] 累加。
31//! - **额外参数**:Laravel 支持 `with()` 追加查询参数。Rust 通过
32//!   [`OAuth2Config::with_extra_param`] / [`OAuth2Config::with_extra_params`]。
33//!
34//! ## 架构说明
35//!
36//! - **OAuth2Provider trait**:对齐 Laravel `ProviderInterface`,业务方实现具体提供商逻辑
37//! - **GenericOAuth2Provider**:通用 OAuth2 实现,接收 [`OAuth2Config`] 和
38//!   [`OAuth2HttpTransport`],支持任意标准 OAuth2 提供商
39//! - **OAuth2HttpTransport trait**:HTTP 传输抽象,解耦 OAuth2 客户端与具体 HTTP 库。
40//!   与 `notify::HttpTransport` 分离,因为 OAuth2 token 交换需要返回响应体(解析 access_token)
41//! - **MemoryOAuth2HttpTransport**:内存 HTTP 传输实现,支持预置 mock 响应,用于测试
42
43use base64::Engine;
44use parking_lot::Mutex;
45use rand::rngs::OsRng;
46use rand::RngCore;
47use sha2::{Digest, Sha256};
48use std::collections::VecDeque;
49use std::sync::Arc;
50use thiserror::Error;
51
52// ============================================================================
53// 错误类型
54// ============================================================================
55
56/// OAuth2 错误 — 对齐 Laravel Socialite 异常体系
57#[derive(Debug, Error)]
58pub enum OAuth2Error {
59    /// 缺少必填字段(client_id / client_secret / redirect_url / auth_url / token_url 等)
60    #[error("OAuth2 字段缺失: {0}")]
61    MissingField(String),
62    /// 授权失败(授权码为空、state 不匹配等)
63    #[error("OAuth2 授权失败: {0}")]
64    AuthFailed(String),
65    /// token 交换失败(HTTP 错误、响应缺少 access_token 等)
66    #[error("OAuth2 token 交换失败: {0}")]
67    TokenExchangeFailed(String),
68    /// 获取用户信息失败(HTTP 错误、响应解析失败等)
69    #[error("OAuth2 获取用户信息失败: {0}")]
70    UserInfoFailed(String),
71    /// HTTP 传输失败(网络错误、连接超时等)
72    #[error("OAuth2 HTTP 传输失败: {0}")]
73    HttpTransport(String),
74    /// 序列化/反序列化失败
75    #[error("OAuth2 序列化失败: {0}")]
76    Serialize(String),
77}
78
79// ============================================================================
80// PKCE(Proof Key for Code Exchange, RFC 7636)
81// ============================================================================
82
83/// PKCE 方法 — 对齐 RFC 7636
84///
85/// 当前仅支持 `S256`(SHA256 派生 code_challenge),不支持 `plain`。
86#[derive(Debug, Clone, Copy, PartialEq, Eq)]
87pub enum PkceMethod {
88    /// SHA256 派生:`code_challenge = BASE64URL(SHA256(code_verifier))`
89    S256,
90}
91
92impl std::fmt::Display for PkceMethod {
93    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
94        match self {
95            PkceMethod::S256 => write!(f, "S256"),
96        }
97    }
98}
99
100/// PKCE 参数对 — 包含 code_verifier 和派生的 code_challenge
101///
102/// `code_verifier` 不实现明文 Debug(自定义 Debug 只输出长度,防止日志泄露)。
103pub struct PkceParams {
104    /// PKCE code_verifier(43-128 字符,此处固定 64 字符 hex)
105    pub code_verifier: String,
106    /// PKCE code_challenge(base64url(SHA256(code_verifier)))
107    pub code_challenge: String,
108    /// PKCE 方法(固定 S256)
109    pub method: PkceMethod,
110}
111
112impl std::fmt::Debug for PkceParams {
113    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
114        f.debug_struct("PkceParams")
115            .field(
116                "code_verifier",
117                &format!("<redacted, len={}>", self.code_verifier.len()),
118            )
119            .field("code_challenge", &self.code_challenge)
120            .field("method", &self.method)
121            .finish()
122    }
123}
124
125impl Clone for PkceParams {
126    fn clone(&self) -> Self {
127        Self {
128            code_verifier: self.code_verifier.clone(),
129            code_challenge: self.code_challenge.clone(),
130            method: self.method,
131        }
132    }
133}
134
135// ============================================================================
136// OAuth2Config
137// ============================================================================
138
139/// OAuth2 配置 — 对齐 Laravel Socialite `config/services.php`
140///
141/// 通过 builder 模式构建,必填字段在 [`OAuth2Config::new`] 中提供,
142/// 可选字段通过 `with_*` 链式方法追加。
143///
144/// # PHP 对齐
145///
146/// ```php
147/// // config/services.php
148/// 'qq' => [
149///     'client_id' => env('QQ_CLIENT_ID'),
150///     'client_secret' => env('QQ_CLIENT_SECRET'),
151///     'redirect' => env('QQ_REDIRECT_URL'),
152/// ],
153/// ```
154///
155/// # Rust 用法
156///
157/// ```ignore
158/// use sz_rust_auth_facade::oauth::OAuth2Config;
159///
160/// let config = OAuth2Config::new(
161///     "100123456",
162///     "secretabc",
163///     "https://example.com/oauth/qq/callback",
164///     "https://graph.qq.com/oauth2.0/authorize",
165///     "https://graph.qq.com/oauth2.0/token",
166/// )
167/// .with_user_url("https://graph.qq.com/user/get_user_info")
168/// .with_scope("get_user_info");
169/// ```
170#[derive(Clone)]
171pub struct OAuth2Config {
172    /// 客户端 ID(必填)
173    pub client_id: String,
174    /// 客户端密钥(必填,Debug 脱敏输出 "***")
175    pub client_secret: String,
176    /// 回调 URL(必填)
177    pub redirect_url: String,
178    /// 授权服务器 authorize 端点(必填,如 `https://graph.qq.com/oauth2.0/authorize`)
179    pub auth_url: String,
180    /// 授权服务器 token 端点(必填,如 `https://graph.qq.com/oauth2.0/token`)
181    pub token_url: String,
182    /// 资源服务器用户信息端点(可选,如 `https://graph.qq.com/user/get_user_info`)
183    pub user_url: Option<String>,
184    /// 默认 scopes(可选)
185    pub scopes: Vec<String>,
186    /// 额外参数(可选)
187    pub extra_params: Vec<(String, String)>,
188    /// PKCE 是否启用(默认 false)
189    pub pkce_enabled: bool,
190    /// device_code 流程的设备授权端点(可选)
191    pub device_auth_url: Option<String>,
192}
193
194impl std::fmt::Debug for OAuth2Config {
195    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
196        f.debug_struct("OAuth2Config")
197            .field("client_id", &self.client_id)
198            .field("client_secret", &"***")
199            .field("redirect_url", &self.redirect_url)
200            .field("auth_url", &self.auth_url)
201            .field("token_url", &self.token_url)
202            .field("user_url", &self.user_url)
203            .field("scopes", &self.scopes)
204            .field("extra_params", &self.extra_params)
205            .field("pkce_enabled", &self.pkce_enabled)
206            .field("device_auth_url", &self.device_auth_url)
207            .finish()
208    }
209}
210
211impl OAuth2Config {
212    /// 创建 OAuth2 配置
213    ///
214    /// # 参数
215    ///
216    /// - `client_id`: 客户端 ID
217    /// - `client_secret`: 客户端密钥
218    /// - `redirect_url`: 回调 URL
219    /// - `auth_url`: 授权服务器 authorize 端点
220    /// - `token_url`: 授权服务器 token 端点
221    #[allow(clippy::too_many_arguments)]
222    pub fn new(
223        client_id: impl Into<String>,
224        client_secret: impl Into<String>,
225        redirect_url: impl Into<String>,
226        auth_url: impl Into<String>,
227        token_url: impl Into<String>,
228    ) -> Self {
229        Self {
230            client_id: client_id.into(),
231            client_secret: client_secret.into(),
232            redirect_url: redirect_url.into(),
233            auth_url: auth_url.into(),
234            token_url: token_url.into(),
235            user_url: None,
236            scopes: Vec::new(),
237            extra_params: Vec::new(),
238            pkce_enabled: false,
239            device_auth_url: None,
240        }
241    }
242
243    /// 设置用户信息端点
244    pub fn with_user_url(mut self, user_url: impl Into<String>) -> Self {
245        self.user_url = Some(user_url.into());
246        self
247    }
248
249    /// 设置 scopes 列表(覆盖原有)
250    pub fn with_scopes(mut self, scopes: Vec<String>) -> Self {
251        self.scopes = scopes;
252        self
253    }
254
255    /// 追加单个 scope
256    pub fn with_scope(mut self, scope: impl Into<String>) -> Self {
257        self.scopes.push(scope.into());
258        self
259    }
260
261    /// 设置额外参数列表(覆盖原有)
262    pub fn with_extra_params(mut self, params: Vec<(String, String)>) -> Self {
263        self.extra_params = params;
264        self
265    }
266
267    /// 追加单个额外参数
268    pub fn with_extra_param(mut self, key: impl Into<String>, value: impl Into<String>) -> Self {
269        self.extra_params.push((key.into(), value.into()));
270        self
271    }
272
273    /// 启用/禁用 PKCE(Proof Key for Code Exchange, RFC 7636)
274    ///
275    /// 启用后,调用方需通过 [`OAuth2Config::generate_pkce_pair`] 生成 PKCE 参数对,
276    /// 将 `code_challenge` 追加到授权 URL,将 `code_verifier` 在 token 交换时提交。
277    pub fn with_pkce(mut self, enabled: bool) -> Self {
278        self.pkce_enabled = enabled;
279        self
280    }
281
282    /// 设置 device_code 流程的设备授权端点
283    pub fn with_device_auth_url(mut self, url: impl Into<String>) -> Self {
284        self.device_auth_url = Some(url.into());
285        self
286    }
287
288    /// 自动生成 state 参数 — 使用 `rand::rngs::OsRng` 生成 16 字节随机数 → 32 字符 hex
289    ///
290    /// 满足 spec 4.3.3 ≥32 字符密码学安全要求。
291    pub fn generate_state() -> String {
292        let mut bytes = [0u8; 16];
293        OsRng.fill_bytes(&mut bytes);
294        hex::encode(bytes)
295    }
296
297    /// 自动生成 PKCE 参数对 — code_verifier + code_challenge
298    ///
299    /// - code_verifier:OsRng 生成 32 字节随机 → 64 字符 hex(满足 43-128 字符要求)
300    /// - code_challenge:`BASE64URL-NO-PAD(SHA256(code_verifier))`
301    /// - method:固定 `S256`
302    pub fn generate_pkce_pair() -> PkceParams {
303        let mut bytes = [0u8; 32];
304        OsRng.fill_bytes(&mut bytes);
305        let code_verifier = hex::encode(bytes);
306        let mut hasher = Sha256::new();
307        hasher.update(code_verifier.as_bytes());
308        let digest = hasher.finalize();
309        let code_challenge = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest);
310        PkceParams {
311            code_verifier,
312            code_challenge,
313            method: PkceMethod::S256,
314        }
315    }
316
317    /// 校验必填字段
318    ///
319    /// 必填字段:`client_id` / `client_secret` / `redirect_url` / `auth_url` / `token_url`。
320    /// 任一为空字符串则返回 [`OAuth2Error::MissingField`]。
321    pub fn validate(&self) -> Result<(), OAuth2Error> {
322        if self.client_id.is_empty() {
323            return Err(OAuth2Error::MissingField("client_id".into()));
324        }
325        if self.client_secret.is_empty() {
326            return Err(OAuth2Error::MissingField("client_secret".into()));
327        }
328        if self.redirect_url.is_empty() {
329            return Err(OAuth2Error::MissingField("redirect_url".into()));
330        }
331        if self.auth_url.is_empty() {
332            return Err(OAuth2Error::MissingField("auth_url".into()));
333        }
334        if self.token_url.is_empty() {
335            return Err(OAuth2Error::MissingField("token_url".into()));
336        }
337        Ok(())
338    }
339}
340
341// ============================================================================
342// SocialiteUser
343// ============================================================================
344
345/// OAuth2 用户信息 — 对齐 Laravel Socialite `User`
346///
347/// 表示从 OAuth2 提供商获取的用户信息,包含标准字段和原始响应数据。
348///
349/// # PHP 对齐
350///
351/// ```php
352/// $user = Socialite::driver('qq')->user();
353/// $user->getId();        // 第三方用户 ID
354/// $user->getNickname();  // 昵称
355/// $user->getName();      // 姓名
356/// $user->getEmail();     // 邮箱
357/// $user->getAvatar();    // 头像
358/// $user->token;          // access_token
359/// $user->refreshToken;   // refresh_token
360/// $user->expiresIn;      // 过期秒数
361/// $user->user;           // 原始响应(raw)
362/// ```
363#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
364pub struct SocialiteUser {
365    /// 第三方平台用户 ID
366    pub id: String,
367    /// 用户昵称
368    pub nickname: Option<String>,
369    /// 用户姓名
370    pub name: Option<String>,
371    /// 邮箱
372    pub email: Option<String>,
373    /// 头像 URL
374    pub avatar: Option<String>,
375    /// 原始响应数据(JSON)
376    pub raw: serde_json::Value,
377    /// 访问令牌
378    #[serde(skip_serializing)]
379    pub access_token: Option<String>,
380    /// 刷新令牌
381    #[serde(skip_serializing)]
382    pub refresh_token: Option<String>,
383    /// 令牌过期时间(Unix 秒)
384    pub expires_in: Option<i64>,
385}
386
387// ============================================================================
388// TokenResponse — token 交换响应
389// ============================================================================
390
391/// OAuth2 token 交换响应 — 对齐 RFC 6749 Section 5.1
392///
393/// `access_token` 和 `refresh_token` 标注 `#[serde(skip_serializing)]` 脱敏,
394/// 防止令牌通过 API 响应或日志泄漏。
395#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize)]
396pub struct TokenResponse {
397    /// 访问令牌(脱敏,不序列化)
398    #[serde(skip_serializing)]
399    pub access_token: String,
400    /// 令牌类型(如 "Bearer")
401    pub token_type: Option<String>,
402    /// 过期秒数
403    pub expires_in: Option<i64>,
404    /// 授权范围
405    pub scope: Option<String>,
406    /// 刷新令牌(脱敏,不序列化)
407    #[serde(skip_serializing)]
408    pub refresh_token: Option<String>,
409}
410
411// ============================================================================
412// OAuth2AuditLogger — 审计日志
413// ============================================================================
414
415/// OAuth2 审计事件 — 记录 token 交换/刷新/CSRF 等安全事件
416#[derive(Debug, Clone)]
417pub struct OAuth2AuditEvent {
418    /// 客户端 ID
419    pub client_id: String,
420    /// grant 类型(authorization_code / implicit / device_code / refresh_token)
421    pub grant_type: String,
422    /// 结果(success / failure)
423    pub result: String,
424    /// Unix 时间戳(秒)
425    pub timestamp: i64,
426    /// 告警码(可选,如 OAUTH2_IMPLICIT_TOKEN_EXPOSED / OAUTH2_CSRF_STATE_MISMATCH)
427    pub alert_code: Option<String>,
428    /// 附加消息
429    pub message: Option<String>,
430}
431
432/// OAuth2 审计日志 trait — best-effort 记录安全事件
433///
434/// 实现者保证 `Send + Sync`,日志写入失败**不应**影响主流程。
435pub trait OAuth2AuditLogger: Send + Sync {
436    /// 记录审计事件
437    fn log_event(&self, event: &OAuth2AuditEvent);
438}
439
440// ============================================================================
441// OAuth2Provider trait
442// ============================================================================
443
444/// OAuth2 提供商 trait — 对齐 Laravel Socialite `ProviderInterface`
445///
446/// 业务方实现此 trait 以对接具体 OAuth2 提供商(QQ / 微信 / GitHub / Google 等)。
447/// 框架内置 [`GenericOAuth2Provider`] 通用实现,满足标准 OAuth2 协议的提供商可直接使用。
448///
449/// # 线程安全
450///
451/// 实现者必须保证 `Send + Sync`,因为 Provider 通常作为单例在多线程下使用。
452pub trait OAuth2Provider: Send + Sync {
453    /// 生成授权 URL(对齐 `Socialite::driver('qq')->redirect()`)
454    ///
455    /// # 参数
456    ///
457    /// - `state`: CSRF 防护的 state 参数(由调用方生成并暂存到 session)
458    ///
459    /// # 返回
460    ///
461    /// 完整的授权 URL,包含 `client_id` / `redirect_uri` / `response_type=code` /
462    /// `state` / `scope` 等查询参数。
463    fn redirect_url(&self, state: &str) -> String;
464
465    /// 用授权码换取访问令牌并获取用户信息(对齐 `Socialite::driver('qq')->user()`)
466    ///
467    /// # 参数
468    ///
469    /// - `code`: 授权服务器回调时携带的授权码
470    ///
471    /// # 返回
472    ///
473    /// 成功返回 [`SocialiteUser`],失败返回 [`OAuth2Error`]。
474    fn user_from_token(&self, code: &str) -> Result<SocialiteUser, OAuth2Error>;
475}
476
477// ============================================================================
478// OAuth2HttpTransport trait(HTTP 传输抽象)
479// ============================================================================
480
481/// OAuth2 HTTP 传输 trait — 用于解耦 OAuth2 客户端与具体 HTTP 库
482///
483/// 与 `notify::HttpTransport` 分离的原因:OAuth2 token 交换需要返回响应体
484/// (解析 `access_token`),而 `notify::HttpTransport::post_json` 仅返回 `Result<(), _>`。
485///
486/// 业务方实现此 trait 注入 reqwest / hyper / etc.,即可让 [`GenericOAuth2Provider`]
487/// 投入生产。
488///
489/// # 线程安全
490///
491/// 实现者必须保证 `Send + Sync`,因为 Provider 通常作为单例在多线程下使用。
492pub trait OAuth2HttpTransport: Send + Sync {
493    /// 发送 POST 请求,Content-Type: application/json
494    ///
495    /// # 参数
496    ///
497    /// - `url`: 目标 URL
498    /// - `body`: 请求体(JSON 字符串)
499    ///
500    /// # 返回
501    ///
502    /// 成功返回响应体字符串,失败返回 [`OAuth2Error`]。
503    fn post_json(&self, url: &str, body: &str) -> Result<String, OAuth2Error>;
504}
505
506// ============================================================================
507// MemoryOAuth2HttpTransport(测试/开发用 HTTP 传输实现)
508// ============================================================================
509
510/// 内存 HTTP 传输 — 用于测试和开发环境
511///
512/// 不实际发送 HTTP 请求,而是:
513/// - 将请求暂存到内部 Vec,供测试断言使用
514/// - 从预置响应队列中依次返回 mock 响应
515///
516/// # 用法
517///
518/// ```ignore
519/// use sz_rust_auth_facade::oauth::MemoryOAuth2HttpTransport;
520///
521/// let transport = MemoryOAuth2HttpTransport::new();
522/// // 预置 mock 响应(按调用顺序消费)
523/// transport.push_response(r#"{"access_token":"token123"}"#);
524/// transport.push_response(r#"{"id":"1","nickname":"test"}"#);
525///
526/// let resp = transport.post_json("https://example.com/token", "{}").unwrap();
527/// assert_eq!(resp, r#"{"access_token":"token123"}"#);
528/// ```
529#[derive(Debug, Default)]
530pub struct MemoryOAuth2HttpTransport {
531    /// 已"发送"的 HTTP 请求列表(url, body)
532    requests: Mutex<Vec<(String, String)>>,
533    /// 预置的 mock 响应队列(FIFO)
534    responses: Mutex<VecDeque<String>>,
535}
536
537impl MemoryOAuth2HttpTransport {
538    /// 创建新的内存 HTTP 传输
539    pub fn new() -> Self {
540        Self::default()
541    }
542
543    /// 预置 mock 响应(追加到队列末尾,按调用顺序消费)
544    pub fn push_response(&self, response: impl Into<String>) {
545        self.responses.lock().push_back(response.into());
546    }
547
548    /// 获取已发送请求数量
549    pub fn count(&self) -> usize {
550        self.requests.lock().len()
551    }
552
553    /// 获取所有已发送请求(快照)
554    pub fn all(&self) -> Vec<(String, String)> {
555        self.requests.lock().clone()
556    }
557
558    /// 获取最后发送的请求
559    pub fn last(&self) -> Option<(String, String)> {
560        self.requests.lock().last().cloned()
561    }
562
563    /// 清空已发送请求和预置响应
564    pub fn clear(&self) {
565        self.requests.lock().clear();
566        self.responses.lock().clear();
567    }
568}
569
570impl OAuth2HttpTransport for MemoryOAuth2HttpTransport {
571    fn post_json(&self, url: &str, body: &str) -> Result<String, OAuth2Error> {
572        self.requests
573            .lock()
574            .push((url.to_string(), body.to_string()));
575        let mut responses = self.responses.lock();
576        match responses.pop_front() {
577            Some(resp) => Ok(resp),
578            None => Ok(String::new()),
579        }
580    }
581}
582
583// ============================================================================
584// GenericOAuth2Provider
585// ============================================================================
586
587/// 通用 OAuth2 提供商 — 对齐 Laravel Socialite `GenericProvider`
588///
589/// 接收 [`OAuth2Config`] 和 [`OAuth2HttpTransport`],支持任意标准 OAuth2 提供商。
590///
591/// # 工作流程
592///
593/// 1. [`GenericOAuth2Provider::redirect_url`]:构造授权 URL,引导用户跳转到授权服务器
594/// 2. 授权服务器回调,携带 `code` 和 `state`
595/// 3. [`GenericOAuth2Provider::user_from_token`]:
596///    - POST `token_url` 换取 `access_token`(grant_type=authorization_code)
597///    - 如果配置了 `user_url`,POST 用户信息端点获取用户资料
598///    - 返回 [`SocialiteUser`]
599///
600/// # 用法
601///
602/// ```ignore
603/// use std::sync::Arc;
604/// use sz_rust_auth_facade::oauth::{
605///     GenericOAuth2Provider, MemoryOAuth2HttpTransport, OAuth2Config, OAuth2Provider,
606/// };
607///
608/// let config = OAuth2Config::new(
609///     "client_id",
610///     "client_secret",
611///     "https://example.com/callback",
612///     "https://provider.com/oauth2.0/authorize",
613///     "https://provider.com/oauth2.0/token",
614/// )
615/// .with_user_url("https://provider.com/user/info");
616///
617/// let transport = Arc::new(MemoryOAuth2HttpTransport::new());
618/// let provider = GenericOAuth2Provider::new(config, transport);
619///
620/// let url = provider.redirect_url("random_state");
621/// let user = provider.user_from_token("auth_code").unwrap();
622/// ```
623pub struct GenericOAuth2Provider {
624    /// OAuth2 配置
625    config: OAuth2Config,
626    /// HTTP 传输实现
627    transport: Arc<dyn OAuth2HttpTransport>,
628    /// 审计日志(可选,best-effort)
629    audit_logger: Option<Arc<dyn OAuth2AuditLogger>>,
630    /// Token 存储(可选,需 `redis-store` feature,best-effort)
631    #[cfg(feature = "redis-store")]
632    token_store: Option<Arc<dyn crate::oauth_store::OAuth2TokenStore>>,
633}
634
635impl GenericOAuth2Provider {
636    /// 创建通用 OAuth2 提供商
637    ///
638    /// # 参数
639    ///
640    /// - `config`: OAuth2 配置
641    /// - `transport`: HTTP 传输实现(业务方注入 reqwest / hyper / etc.)
642    pub fn new(config: OAuth2Config, transport: Arc<dyn OAuth2HttpTransport>) -> Self {
643        Self {
644            config,
645            transport,
646            audit_logger: None,
647            #[cfg(feature = "redis-store")]
648            token_store: None,
649        }
650    }
651
652    /// 注入审计日志
653    pub fn with_audit_logger(mut self, logger: Arc<dyn OAuth2AuditLogger>) -> Self {
654        self.audit_logger = Some(logger);
655        self
656    }
657
658    /// 注入 Token 存储(需 `redis-store` feature)
659    ///
660    /// token 交换成功后自动存储(best-effort,存储失败记录告警不传播)。
661    #[cfg(feature = "redis-store")]
662    pub fn with_token_store(
663        mut self,
664        store: Arc<dyn crate::oauth_store::OAuth2TokenStore>,
665    ) -> Self {
666        self.token_store = Some(store);
667        self
668    }
669
670    /// 记录审计事件(best-effort,logger 未注入时静默跳过)
671    fn log_audit(
672        &self,
673        grant_type: &str,
674        result: &str,
675        alert_code: Option<&str>,
676        message: Option<&str>,
677    ) {
678        if let Some(logger) = &self.audit_logger {
679            let event = OAuth2AuditEvent {
680                client_id: self.config.client_id.clone(),
681                grant_type: grant_type.to_string(),
682                result: result.to_string(),
683                timestamp: chrono::Utc::now().timestamp(),
684                alert_code: alert_code.map(|s| s.to_string()),
685                message: message.map(|s| s.to_string()),
686            };
687            logger.log_event(&event);
688        }
689    }
690
691    /// 用 refresh_token 换取新的 access_token
692    ///
693    /// POST `token_url`,body 为 JSON:
694    /// ```json
695    /// {
696    ///   "grant_type": "refresh_token",
697    ///   "refresh_token": "<refresh_token>",
698    ///   "client_id": "<client_id>",
699    ///   "client_secret": "<client_secret>"
700    /// }
701    /// ```
702    pub fn refresh_token(&self, refresh_token: &str) -> Result<TokenResponse, OAuth2Error> {
703        if refresh_token.is_empty() {
704            self.log_audit("refresh_token", "failure", None, Some("refresh_token 为空"));
705            return Err(OAuth2Error::AuthFailed("refresh_token 不能为空".into()));
706        }
707
708        let body = serde_json::json!({
709            "grant_type": "refresh_token",
710            "refresh_token": refresh_token,
711            "client_id": self.config.client_id,
712            "client_secret": self.config.client_secret,
713        });
714        let body_str =
715            serde_json::to_string(&body).map_err(|err| OAuth2Error::Serialize(err.to_string()))?;
716
717        let response = self
718            .transport
719            .post_json(&self.config.token_url, &body_str)
720            .map_err(|err| {
721                self.log_audit("refresh_token", "failure", None, Some(&err.to_string()));
722                OAuth2Error::HttpTransport(err.to_string())
723            })?;
724
725        if response.is_empty() {
726            self.log_audit("refresh_token", "failure", None, Some("空响应"));
727            return Err(OAuth2Error::TokenExchangeFailed("token 响应为空".into()));
728        }
729
730        let json: serde_json::Value = serde_json::from_str(&response).map_err(|err| {
731            self.log_audit(
732                "refresh_token",
733                "failure",
734                None,
735                Some(&format!("JSON 解析失败: {err}")),
736            );
737            OAuth2Error::TokenExchangeFailed(format!("解析 token 响应失败: {err}"))
738        })?;
739
740        let access_token = json
741            .get("access_token")
742            .and_then(|v| v.as_str())
743            .ok_or_else(|| {
744                self.log_audit(
745                    "refresh_token",
746                    "failure",
747                    None,
748                    Some("响应缺少 access_token"),
749                );
750                OAuth2Error::TokenExchangeFailed("token 响应缺少 access_token 字段".into())
751            })?
752            .to_string();
753
754        let token_response = TokenResponse {
755            access_token,
756            token_type: json
757                .get("token_type")
758                .and_then(|v| v.as_str())
759                .map(|s| s.to_string()),
760            expires_in: json.get("expires_in").and_then(|v| v.as_i64()),
761            scope: json
762                .get("scope")
763                .and_then(|v| v.as_str())
764                .map(|s| s.to_string()),
765            refresh_token: json
766                .get("refresh_token")
767                .and_then(|v| v.as_str())
768                .map(|s| s.to_string()),
769        };
770
771        self.log_audit("refresh_token", "success", None, None);
772        Ok(token_response)
773    }
774
775    /// 构造授权 URL
776    ///
777    /// 拼接 `auth_url` 与查询参数:
778    /// - `client_id`
779    /// - `redirect_uri`
780    /// - `response_type=code`
781    /// - `state`
782    /// - `scope`(如果配置了 scopes,以空格连接)
783    /// - `code_challenge` + `code_challenge_method=S256`(如果传入 PKCE 参数)
784    /// - 额外参数(如果配置了 extra_params)
785    ///
786    /// 如果 `auth_url` 已包含查询字符串,则用 `&` 追加,否则用 `?` 起始。
787    fn build_redirect_url(&self, state: &str, pkce: Option<&PkceParams>) -> String {
788        let mut params: Vec<(String, String)> = vec![
789            ("client_id".into(), self.config.client_id.clone()),
790            ("redirect_uri".into(), self.config.redirect_url.clone()),
791            ("response_type".into(), "code".into()),
792            ("state".into(), state.to_string()),
793        ];
794
795        if !self.config.scopes.is_empty() {
796            params.push(("scope".into(), self.config.scopes.join(" ")));
797        }
798
799        if let Some(pkce) = pkce {
800            params.push(("code_challenge".into(), pkce.code_challenge.clone()));
801            params.push(("code_challenge_method".into(), pkce.method.to_string()));
802        }
803
804        for (key, value) in &self.config.extra_params {
805            params.push((key.clone(), value.clone()));
806        }
807
808        let query = params
809            .iter()
810            .map(|(key, value)| format!("{}={}", percent_encode(key), percent_encode(value)))
811            .collect::<Vec<_>>()
812            .join("&");
813
814        let separator = if self.config.auth_url.contains('?') {
815            "&"
816        } else {
817            "?"
818        };
819        format!("{}{}{}", self.config.auth_url, separator, query)
820    }
821
822    /// 用授权码换取访问令牌
823    ///
824    /// POST `token_url`,body 为 JSON:
825    /// ```json
826    /// {
827    ///   "grant_type": "authorization_code",
828    ///   "code": "<code>",
829    ///   "client_id": "<client_id>",
830    ///   "client_secret": "<client_secret>",
831    ///   "redirect_uri": "<redirect_url>",
832    ///   "code_verifier": "<code_verifier>"  // 仅当 pkce 启用时
833    /// }
834    /// ```
835    ///
836    /// 成功返回解析后的 JSON(含 `access_token` / `refresh_token` / `expires_in` 等)。
837    fn exchange_token(
838        &self,
839        code: &str,
840        code_verifier: Option<&str>,
841    ) -> Result<serde_json::Value, OAuth2Error> {
842        let mut body = serde_json::json!({
843            "grant_type": "authorization_code",
844            "code": code,
845            "client_id": self.config.client_id,
846            "client_secret": self.config.client_secret,
847            "redirect_uri": self.config.redirect_url,
848        });
849
850        if let Some(verifier) = code_verifier {
851            body["code_verifier"] = serde_json::Value::String(verifier.to_string());
852        }
853
854        let body_str =
855            serde_json::to_string(&body).map_err(|err| OAuth2Error::Serialize(err.to_string()))?;
856
857        let response = self
858            .transport
859            .post_json(&self.config.token_url, &body_str)
860            .map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
861
862        if response.is_empty() {
863            return Ok(serde_json::Value::Null);
864        }
865
866        serde_json::from_str(&response)
867            .map_err(|err| OAuth2Error::TokenExchangeFailed(format!("解析 token 响应失败: {err}")))
868    }
869
870    /// 获取用户信息
871    ///
872    /// POST `user_url`,body 为 JSON(含 `access_token` 和 `openid`)。
873    /// 成功返回解析后的 JSON。
874    fn fetch_user_info(
875        &self,
876        access_token: &str,
877        token_json: &serde_json::Value,
878    ) -> Result<serde_json::Value, OAuth2Error> {
879        let user_url = self
880            .config
881            .user_url
882            .as_ref()
883            .ok_or_else(|| OAuth2Error::UserInfoFailed("user_url 未配置".into()))?;
884
885        let body = serde_json::json!({
886            "access_token": access_token,
887            "openid": token_json.get("openid").cloned().unwrap_or(serde_json::Value::Null),
888        });
889        let body_str =
890            serde_json::to_string(&body).map_err(|err| OAuth2Error::Serialize(err.to_string()))?;
891
892        let response = self
893            .transport
894            .post_json(user_url, &body_str)
895            .map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
896
897        if response.is_empty() {
898            return Ok(serde_json::Value::Null);
899        }
900
901        serde_json::from_str(&response)
902            .map_err(|err| OAuth2Error::UserInfoFailed(format!("解析用户信息响应失败: {err}")))
903    }
904
905    /// 从用户信息 JSON 中提取 [`SocialiteUser`] 的标准字段
906    ///
907    /// 兼容多种字段命名(对齐 Laravel Socialite 的字段映射逻辑):
908    /// - `id` / `openid` / `user_id` → `id`
909    /// - `nickname` / `nick_name` → `nickname`
910    /// - `name` / `username` → `name`
911    /// - `email` → `email`
912    /// - `avatar` / `figureurl_qq_1` / `figureurl` / `headimgurl` → `avatar`
913    fn extract_user_fields(user_json: &serde_json::Value) -> SocialiteUser {
914        let id = user_json
915            .get("id")
916            .or_else(|| user_json.get("openid"))
917            .or_else(|| user_json.get("user_id"))
918            .and_then(extract_string)
919            .unwrap_or_default();
920
921        let nickname = user_json
922            .get("nickname")
923            .or_else(|| user_json.get("nick_name"))
924            .and_then(extract_string);
925
926        let name = user_json
927            .get("name")
928            .or_else(|| user_json.get("username"))
929            .and_then(extract_string);
930
931        let email = user_json.get("email").and_then(extract_string);
932
933        let avatar = user_json
934            .get("avatar")
935            .or_else(|| user_json.get("figureurl_qq_1"))
936            .or_else(|| user_json.get("figureurl"))
937            .or_else(|| user_json.get("headimgurl"))
938            .and_then(extract_string);
939
940        SocialiteUser {
941            id,
942            nickname,
943            name,
944            email,
945            avatar,
946            raw: user_json.clone(),
947            access_token: None,
948            refresh_token: None,
949            expires_in: None,
950        }
951    }
952}
953
954impl OAuth2Provider for GenericOAuth2Provider {
955    fn redirect_url(&self, state: &str) -> String {
956        self.build_redirect_url(state, None)
957    }
958
959    fn user_from_token(&self, code: &str) -> Result<SocialiteUser, OAuth2Error> {
960        self.user_from_token_with_pkce(code, None)
961    }
962}
963
964impl GenericOAuth2Provider {
965    /// 构造授权 URL(带 PKCE 参数)
966    ///
967    /// 当配置了 PKCE 参数时,授权 URL 追加 `code_challenge` 和 `code_challenge_method=S256`。
968    pub fn redirect_url_with_pkce(&self, state: &str, pkce: &PkceParams) -> String {
969        self.build_redirect_url(state, Some(pkce))
970    }
971
972    /// 用授权码换取用户信息(可选 PKCE code_verifier)
973    ///
974    /// 当 `code_verifier` 为 `Some` 时,token 交换 POST body 追加 `code_verifier` 字段。
975    pub fn user_from_token_with_pkce(
976        &self,
977        code: &str,
978        code_verifier: Option<&str>,
979    ) -> Result<SocialiteUser, OAuth2Error> {
980        // 1. 校验配置必填字段
981        self.config.validate()?;
982
983        // 2. 校验授权码非空
984        if code.is_empty() {
985            return Err(OAuth2Error::AuthFailed("授权码不能为空".into()));
986        }
987
988        // 3. 用授权码换取访问令牌
989        let token_json = self.exchange_token(code, code_verifier)?;
990
991        // 4. 提取 access_token
992        let access_token = token_json
993            .get("access_token")
994            .and_then(|value| value.as_str())
995            .ok_or_else(|| {
996                self.log_audit(
997                    "authorization_code",
998                    "failure",
999                    None,
1000                    Some("token 响应缺少 access_token"),
1001                );
1002                OAuth2Error::TokenExchangeFailed(format!(
1003                    "token 响应缺少 access_token 字段: {token_json}"
1004                ))
1005            })?
1006            .to_string();
1007
1008        let refresh_token = token_json
1009            .get("refresh_token")
1010            .and_then(|value| value.as_str())
1011            .map(|value| value.to_string());
1012
1013        let expires_in = token_json
1014            .get("expires_in")
1015            .and_then(|value| value.as_i64());
1016
1017        // 4.1 自动刷新:如果 expires_in <= 0 且有 refresh_token,自动调用 refresh_token
1018        let (access_token, refresh_token, expires_in) = if let Some(exp) = expires_in {
1019            if exp <= 0 {
1020                if let Some(ref rt) = refresh_token {
1021                    match self.refresh_token(rt) {
1022                        Ok(new_token) => {
1023                            let new_access = new_token.access_token;
1024                            let new_refresh = new_token.refresh_token.or(refresh_token.clone());
1025                            let new_exp = new_token.expires_in;
1026                            (new_access, new_refresh, new_exp)
1027                        }
1028                        Err(_) => {
1029                            // 刷新失败,保留原 token
1030                            (access_token, refresh_token, expires_in)
1031                        }
1032                    }
1033                } else {
1034                    (access_token, refresh_token, expires_in)
1035                }
1036            } else {
1037                (access_token, refresh_token, expires_in)
1038            }
1039        } else {
1040            (access_token, refresh_token, expires_in)
1041        };
1042
1043        // 5. 如果配置了 user_url,获取用户信息;否则仅返回 token 信息
1044        let mut user = if self.config.user_url.is_some() {
1045            let user_json = self.fetch_user_info(&access_token, &token_json)?;
1046            Self::extract_user_fields(&user_json)
1047        } else {
1048            SocialiteUser::default()
1049        };
1050
1051        user.access_token = Some(access_token);
1052        user.refresh_token = refresh_token;
1053        user.expires_in = expires_in;
1054
1055        // 6. best-effort 存储 token(需 redis-store feature)
1056        #[cfg(feature = "redis-store")]
1057        if let Some(store) = &self.token_store {
1058            let store = store.clone();
1059            let client_id = self.config.client_id.clone();
1060            let token_to_store = TokenResponse {
1061                access_token: user.access_token.clone().unwrap_or_default(),
1062                token_type: None,
1063                expires_in: user.expires_in,
1064                scope: None,
1065                refresh_token: user.refresh_token.clone(),
1066            };
1067            tokio::task::spawn(async move {
1068                if let Err(err) = store.store_token(&client_id, &token_to_store).await {
1069                    tracing::warn!(
1070                        error = %err,
1071                        client_id = %client_id,
1072                        "OAUTH2_TOKEN_STORE_FAILED: token 存储失败(best-effort,不影响主流程)"
1073                    );
1074                }
1075            });
1076        }
1077
1078        self.log_audit("authorization_code", "success", None, None);
1079        Ok(user)
1080    }
1081}
1082
1083// ============================================================================
1084// ImplicitOAuth2Provider — OAuth2 Implicit 流程
1085// ============================================================================
1086
1087/// OAuth2 Implicit 流程提供商 — 对齐 RFC 6749 Section 4.2
1088///
1089/// Implicit 流程直接在授权 URL 的 fragment 中返回 access_token,
1090/// 不需要后端 token 交换步骤。**安全级别较低**(token 经 URL fragment 暴露)。
1091///
1092/// # 工作流程
1093///
1094/// 1. [`ImplicitOAuth2Provider::redirect_url`]:构造 `response_type=token` 的授权 URL
1095/// 2. 授权服务器回调,在 URL fragment 中携带 `access_token` 和 `state`
1096/// 3. [`ImplicitOAuth2Provider::parse_fragment`]:解析 fragment 提取 token,校验 state
1097pub struct ImplicitOAuth2Provider {
1098    /// OAuth2 配置
1099    config: OAuth2Config,
1100    /// 审计日志(可选)
1101    audit_logger: Option<Arc<dyn OAuth2AuditLogger>>,
1102}
1103
1104impl ImplicitOAuth2Provider {
1105    /// 创建 Implicit OAuth2 提供商
1106    pub fn new(config: OAuth2Config) -> Self {
1107        Self {
1108            config,
1109            audit_logger: None,
1110        }
1111    }
1112
1113    /// 注入审计日志
1114    pub fn with_audit_logger(mut self, logger: Arc<dyn OAuth2AuditLogger>) -> Self {
1115        self.audit_logger = Some(logger);
1116        self
1117    }
1118
1119    /// 记录审计事件(best-effort)
1120    fn log_audit(
1121        &self,
1122        grant_type: &str,
1123        result: &str,
1124        alert_code: Option<&str>,
1125        message: Option<&str>,
1126    ) {
1127        if let Some(logger) = &self.audit_logger {
1128            let event = OAuth2AuditEvent {
1129                client_id: self.config.client_id.clone(),
1130                grant_type: grant_type.to_string(),
1131                result: result.to_string(),
1132                timestamp: chrono::Utc::now().timestamp(),
1133                alert_code: alert_code.map(|s| s.to_string()),
1134                message: message.map(|s| s.to_string()),
1135            };
1136            logger.log_event(&event);
1137        }
1138    }
1139
1140    /// 构造授权 URL(`response_type=token`)
1141    ///
1142    /// 拼接 `auth_url` 与查询参数:
1143    /// - `client_id`
1144    /// - `redirect_uri`
1145    /// - `response_type=token`
1146    /// - `state`
1147    /// - `scope`(如果配置了 scopes)
1148    /// - `code_challenge` + `code_challenge_method=S256`(如果传入 PKCE 参数)
1149    /// - 额外参数
1150    pub fn redirect_url(&self, state: &str) -> String {
1151        self.redirect_url_with_pkce(state, None)
1152    }
1153
1154    /// 构造授权 URL(带可选 PKCE 参数)
1155    pub fn redirect_url_with_pkce(&self, state: &str, pkce: Option<&PkceParams>) -> String {
1156        let mut params: Vec<(String, String)> = vec![
1157            ("client_id".into(), self.config.client_id.clone()),
1158            ("redirect_uri".into(), self.config.redirect_url.clone()),
1159            ("response_type".into(), "token".into()),
1160            ("state".into(), state.to_string()),
1161        ];
1162
1163        if !self.config.scopes.is_empty() {
1164            params.push(("scope".into(), self.config.scopes.join(" ")));
1165        }
1166
1167        if let Some(pkce) = pkce {
1168            params.push(("code_challenge".into(), pkce.code_challenge.clone()));
1169            params.push(("code_challenge_method".into(), pkce.method.to_string()));
1170        }
1171
1172        for (key, value) in &self.config.extra_params {
1173            params.push((key.clone(), value.clone()));
1174        }
1175
1176        let query = params
1177            .iter()
1178            .map(|(key, value)| format!("{}={}", percent_encode(key), percent_encode(value)))
1179            .collect::<Vec<_>>()
1180            .join("&");
1181
1182        let separator = if self.config.auth_url.contains('?') {
1183            "&"
1184        } else {
1185            "?"
1186        };
1187        format!("{}{}{}", self.config.auth_url, separator, query)
1188    }
1189
1190    /// 解析 URI fragment 提取 access_token
1191    ///
1192    /// # 参数
1193    ///
1194    /// - `fragment`: URI fragment(`#` 后的部分),如 `access_token=xxx&state=yyy`
1195    /// - `expected_state`: 预期的 state 值(用于 CSRF 校验)
1196    ///
1197    /// # 返回
1198    ///
1199    /// 成功返回 [`TokenResponse`],其中 `refresh_token` 固定为 `None`(implicit 流程不返回 refresh_token)。
1200    pub fn parse_fragment(
1201        &self,
1202        fragment: &str,
1203        expected_state: &str,
1204    ) -> Result<TokenResponse, OAuth2Error> {
1205        // 告警:implicit 流程 token 经 URL fragment 暴露
1206        self.log_audit(
1207            "implicit",
1208            "success",
1209            Some("OAUTH2_IMPLICIT_TOKEN_EXPOSED"),
1210            Some("implicit 流程 token 经 URL fragment 暴露"),
1211        );
1212
1213        if fragment.is_empty() {
1214            self.log_audit("implicit", "failure", None, Some("fragment 为空"));
1215            return Err(OAuth2Error::TokenExchangeFailed("fragment 为空".into()));
1216        }
1217
1218        // 解析 fragment 中的 key=value 对
1219        let params: std::collections::HashMap<&str, &str> = fragment
1220            .split('&')
1221            .filter_map(|pair| {
1222                let (key, value) = pair.split_once('=')?;
1223                Some((key, value))
1224            })
1225            .collect();
1226
1227        // state 校验(CSRF 防护)
1228        let state = params.get("state").copied().unwrap_or("");
1229        if state != expected_state {
1230            self.log_audit(
1231                "implicit",
1232                "failure",
1233                Some("OAUTH2_CSRF_STATE_MISMATCH"),
1234                Some(&format!(
1235                    "state 不匹配: expected={expected_state}, actual={state}"
1236                )),
1237            );
1238            return Err(OAuth2Error::AuthFailed("CSRF state mismatch".into()));
1239        }
1240
1241        // 提取 access_token
1242        let access_token = params.get("access_token").copied().ok_or_else(|| {
1243            self.log_audit(
1244                "implicit",
1245                "failure",
1246                None,
1247                Some("fragment 无 access_token"),
1248            );
1249            OAuth2Error::TokenExchangeFailed("fragment 中缺少 access_token".into())
1250        })?;
1251
1252        Ok(TokenResponse {
1253            access_token: access_token.to_string(),
1254            token_type: params.get("token_type").map(|s| s.to_string()),
1255            expires_in: params.get("expires_in").and_then(|s| s.parse().ok()),
1256            scope: params.get("scope").map(|s| s.to_string()),
1257            // implicit 流程不返回 refresh_token
1258            refresh_token: None,
1259        })
1260    }
1261}
1262
1263impl OAuth2Provider for ImplicitOAuth2Provider {
1264    fn redirect_url(&self, state: &str) -> String {
1265        self.redirect_url(state)
1266    }
1267
1268    fn user_from_token(&self, code: &str) -> Result<SocialiteUser, OAuth2Error> {
1269        // Implicit 流程没有后端 token 交换,code 实际上是 access_token
1270        // 直接构造 SocialiteUser(无用户信息端点调用)
1271        Ok(SocialiteUser {
1272            access_token: Some(code.to_string()),
1273            ..Default::default()
1274        })
1275    }
1276}
1277
1278// ============================================================================
1279// DeviceCodeOAuth2Provider — OAuth2 Device Code 流程(feature = "device-code")
1280// ============================================================================
1281
1282/// OAuth2 Device Code 流程模块(RFC 8628)— 需要 `device-code` feature
1283#[cfg(feature = "device-code")]
1284pub mod device_code {
1285    use super::*;
1286    use async_trait::async_trait;
1287
1288    /// 异步 HTTP 传输 trait — 用于 device_code 流程的异步轮询
1289    ///
1290    /// device_code 流程需要 `tokio::time::sleep` 异步等待,因此传输层也需异步。
1291    #[async_trait]
1292    pub trait AsyncOAuth2HttpTransport: Send + Sync {
1293        /// 发送 POST 请求,Content-Type: application/x-www-form-urlencoded
1294        ///
1295        /// # 参数
1296        ///
1297        /// - `url`: 目标 URL
1298        /// - `params`: 表单参数列表(key-value 对)
1299        ///
1300        /// # 返回
1301        ///
1302        /// 成功返回响应体字符串,失败返回 [`OAuth2Error`]。
1303        async fn post_form(
1304            &self,
1305            url: &str,
1306            params: &[(&str, &str)],
1307        ) -> Result<String, OAuth2Error>;
1308    }
1309
1310    /// Device Code 响应 — 对齐 RFC 8628 Section 3.2
1311    #[derive(Debug, Clone)]
1312    pub struct DeviceCodeResponse {
1313        /// 设备码(客户端用于轮询 token 端点)
1314        pub device_code: String,
1315        /// 用户码(用户在验证页面输入)
1316        pub user_code: String,
1317        /// 验证 URI(用户访问以完成授权)
1318        pub verification_uri: String,
1319        /// 过期秒数
1320        pub expires_in: i64,
1321        /// 轮询间隔秒数
1322        pub interval: i64,
1323    }
1324
1325    /// Device Code OAuth2 提供商 — 对齐 RFC 8628
1326    pub struct DeviceCodeOAuth2Provider {
1327        /// OAuth2 配置
1328        config: OAuth2Config,
1329        /// 异步 HTTP 传输
1330        transport: Arc<dyn AsyncOAuth2HttpTransport>,
1331        /// 审计日志(可选)
1332        audit_logger: Option<Arc<dyn OAuth2AuditLogger>>,
1333    }
1334
1335    impl DeviceCodeOAuth2Provider {
1336        /// 创建 Device Code OAuth2 提供商
1337        pub fn new(config: OAuth2Config, transport: Arc<dyn AsyncOAuth2HttpTransport>) -> Self {
1338            Self {
1339                config,
1340                transport,
1341                audit_logger: None,
1342            }
1343        }
1344
1345        /// 注入审计日志
1346        pub fn with_audit_logger(mut self, logger: Arc<dyn OAuth2AuditLogger>) -> Self {
1347            self.audit_logger = Some(logger);
1348            self
1349        }
1350
1351        /// 记录审计事件(best-effort)
1352        fn log_audit(
1353            &self,
1354            grant_type: &str,
1355            result: &str,
1356            alert_code: Option<&str>,
1357            message: Option<&str>,
1358        ) {
1359            if let Some(logger) = &self.audit_logger {
1360                let event = OAuth2AuditEvent {
1361                    client_id: self.config.client_id.clone(),
1362                    grant_type: grant_type.to_string(),
1363                    result: result.to_string(),
1364                    timestamp: chrono::Utc::now().timestamp(),
1365                    alert_code: alert_code.map(|s| s.to_string()),
1366                    message: message.map(|s| s.to_string()),
1367                };
1368                logger.log_event(&event);
1369            }
1370        }
1371
1372        /// 请求设备码 — POST device_authorization 端点
1373        ///
1374        /// # 参数
1375        ///
1376        /// - `scope`: 请求的授权范围列表
1377        pub async fn request_device_code(
1378            &self,
1379            scope: &[String],
1380        ) -> Result<DeviceCodeResponse, OAuth2Error> {
1381            let device_auth_url = self.config.device_auth_url.as_ref().ok_or_else(|| {
1382                self.log_audit(
1383                    "device_code",
1384                    "failure",
1385                    None,
1386                    Some("device_auth_url 未配置"),
1387                );
1388                OAuth2Error::MissingField("device_auth_url".into())
1389            })?;
1390
1391            let scope_str = scope.join(" ");
1392            let params: Vec<(&str, &str)> = vec![
1393                ("client_id", self.config.client_id.as_str()),
1394                ("scope", scope_str.as_str()),
1395            ];
1396
1397            let response = self
1398                .transport
1399                .post_form(device_auth_url, &params)
1400                .await
1401                .map_err(|err| {
1402                    self.log_audit("device_code", "failure", None, Some(&err.to_string()));
1403                    OAuth2Error::HttpTransport(err.to_string())
1404                })?;
1405
1406            let json: serde_json::Value = serde_json::from_str(&response).map_err(|err| {
1407                self.log_audit(
1408                    "device_code",
1409                    "failure",
1410                    None,
1411                    Some(&format!("JSON 解析失败: {err}")),
1412                );
1413                OAuth2Error::TokenExchangeFailed(format!("解析 device code 响应失败: {err}"))
1414            })?;
1415
1416            let device_code = json
1417                .get("device_code")
1418                .and_then(|v| v.as_str())
1419                .ok_or_else(|| {
1420                    OAuth2Error::TokenExchangeFailed("device code 响应缺少 device_code 字段".into())
1421                })?
1422                .to_string();
1423
1424            let user_code = json
1425                .get("user_code")
1426                .and_then(|v| v.as_str())
1427                .ok_or_else(|| {
1428                    OAuth2Error::TokenExchangeFailed("device code 响应缺少 user_code 字段".into())
1429                })?
1430                .to_string();
1431
1432            let verification_uri = json
1433                .get("verification_uri")
1434                .and_then(|v| v.as_str())
1435                .ok_or_else(|| {
1436                    OAuth2Error::TokenExchangeFailed(
1437                        "device code 响应缺少 verification_uri 字段".into(),
1438                    )
1439                })?
1440                .to_string();
1441
1442            let expires_in = json
1443                .get("expires_in")
1444                .and_then(|v| v.as_i64())
1445                .unwrap_or(600);
1446            let interval = json.get("interval").and_then(|v| v.as_i64()).unwrap_or(5);
1447
1448            self.log_audit("device_code", "success", None, None);
1449            Ok(DeviceCodeResponse {
1450                device_code,
1451                user_code,
1452                verification_uri,
1453                expires_in,
1454                interval,
1455            })
1456        }
1457
1458        /// 轮询 token 端点直到获取 token 或超时
1459        ///
1460        /// # 参数
1461        ///
1462        /// - `device_code`: 设备码
1463        /// - `interval`: 初始轮询间隔(秒)
1464        /// - `expires_in`: 设备码过期时间(秒)
1465        ///
1466        /// # 退避策略
1467        ///
1468        /// - `authorization_pending` → 按 interval 继续轮询
1469        /// - `slow_down` → interval += 5(上限 60s)后继续
1470        /// - 收到 token → 停止返回
1471        /// - 超过 expires_in → `OAUTH2_DEVICE_CODE_EXPIRED`
1472        /// - `access_denied` → `OAUTH2_ACCESS_DENIED`
1473        pub async fn poll_for_token(
1474            &self,
1475            device_code: &str,
1476            mut interval: i64,
1477            expires_in: i64,
1478        ) -> Result<TokenResponse, OAuth2Error> {
1479            let start = std::time::Instant::now();
1480            let expires_duration = std::time::Duration::from_secs(expires_in.max(0) as u64);
1481
1482            loop {
1483                // 检查是否过期
1484                if start.elapsed() >= expires_duration {
1485                    self.log_audit(
1486                        "device_code",
1487                        "failure",
1488                        Some("OAUTH2_DEVICE_CODE_EXPIRED"),
1489                        Some("设备码已过期"),
1490                    );
1491                    return Err(OAuth2Error::AuthFailed(
1492                        "OAUTH2_DEVICE_CODE_EXPIRED: 设备码已过期".into(),
1493                    ));
1494                }
1495
1496                // 轮询 token 端点
1497                let params: Vec<(&str, &str)> = vec![
1498                    ("grant_type", "device_code"),
1499                    ("device_code", device_code),
1500                    ("client_id", self.config.client_id.as_str()),
1501                ];
1502
1503                let response = self
1504                    .transport
1505                    .post_form(&self.config.token_url, &params)
1506                    .await
1507                    .map_err(|err| OAuth2Error::HttpTransport(err.to_string()))?;
1508
1509                let json: serde_json::Value = serde_json::from_str(&response).map_err(|err| {
1510                    OAuth2Error::TokenExchangeFailed(format!("解析 token 响应失败: {err}"))
1511                })?;
1512
1513                // 检查是否有 access_token(成功)
1514                if let Some(access_token) = json.get("access_token").and_then(|v| v.as_str()) {
1515                    let token_response = TokenResponse {
1516                        access_token: access_token.to_string(),
1517                        token_type: json
1518                            .get("token_type")
1519                            .and_then(|v| v.as_str())
1520                            .map(|s| s.to_string()),
1521                        expires_in: json.get("expires_in").and_then(|v| v.as_i64()),
1522                        scope: json
1523                            .get("scope")
1524                            .and_then(|v| v.as_str())
1525                            .map(|s| s.to_string()),
1526                        refresh_token: json
1527                            .get("refresh_token")
1528                            .and_then(|v| v.as_str())
1529                            .map(|s| s.to_string()),
1530                    };
1531                    self.log_audit("device_code", "success", None, None);
1532                    return Ok(token_response);
1533                }
1534
1535                // 检查错误码
1536                let error = json.get("error").and_then(|v| v.as_str()).unwrap_or("");
1537
1538                match error {
1539                    "authorization_pending" => {
1540                        // 继续轮询,interval 不变
1541                    }
1542                    "slow_down" => {
1543                        // interval += 5,上限 60s
1544                        interval = (interval + 5).min(60);
1545                    }
1546                    "access_denied" => {
1547                        self.log_audit(
1548                            "device_code",
1549                            "failure",
1550                            Some("OAUTH2_ACCESS_DENIED"),
1551                            Some("用户拒绝授权"),
1552                        );
1553                        return Err(OAuth2Error::AuthFailed(
1554                            "OAUTH2_ACCESS_DENIED: 用户拒绝授权".into(),
1555                        ));
1556                    }
1557                    "expired_token" => {
1558                        self.log_audit(
1559                            "device_code",
1560                            "failure",
1561                            Some("OAUTH2_DEVICE_CODE_EXPIRED"),
1562                            Some("设备码已过期"),
1563                        );
1564                        return Err(OAuth2Error::AuthFailed(
1565                            "OAUTH2_DEVICE_CODE_EXPIRED: 设备码已过期".into(),
1566                        ));
1567                    }
1568                    _ => {
1569                        return Err(OAuth2Error::TokenExchangeFailed(format!(
1570                            "未知错误: {error}"
1571                        )));
1572                    }
1573                }
1574
1575                // 异步等待(不阻塞 tokio 线程)
1576                if interval > 0 {
1577                    tokio::time::sleep(std::time::Duration::from_secs(interval as u64)).await;
1578                }
1579            }
1580        }
1581    }
1582
1583    // ------------------------------------------------------------------------
1584    // 内存异步 HTTP 传输(测试用)
1585    // ------------------------------------------------------------------------
1586
1587    /// 内存异步 HTTP 传输 — 用于测试 device_code 流程
1588    #[derive(Default)]
1589    pub struct MemoryAsyncOAuth2HttpTransport {
1590        requests: Mutex<Vec<(String, String)>>,
1591        responses: Mutex<VecDeque<String>>,
1592    }
1593
1594    impl MemoryAsyncOAuth2HttpTransport {
1595        /// 创建新的内存异步 HTTP 传输
1596        pub fn new() -> Self {
1597            Self::default()
1598        }
1599
1600        /// 预置 mock 响应(追加到队列末尾,按调用顺序消费)
1601        pub fn push_response(&self, response: impl Into<String>) {
1602            self.responses.lock().push_back(response.into());
1603        }
1604
1605        /// 获取已发送请求数量
1606        pub fn count(&self) -> usize {
1607            self.requests.lock().len()
1608        }
1609    }
1610
1611    #[async_trait]
1612    impl AsyncOAuth2HttpTransport for MemoryAsyncOAuth2HttpTransport {
1613        async fn post_form(
1614            &self,
1615            url: &str,
1616            params: &[(&str, &str)],
1617        ) -> Result<String, OAuth2Error> {
1618            let body = params
1619                .iter()
1620                .map(|(k, v)| format!("{k}={v}"))
1621                .collect::<Vec<_>>()
1622                .join("&");
1623            self.requests.lock().push((url.to_string(), body));
1624            let mut responses = self.responses.lock();
1625            match responses.pop_front() {
1626                Some(resp) => Ok(resp),
1627                None => Ok(String::new()),
1628            }
1629        }
1630    }
1631
1632    // ------------------------------------------------------------------------
1633    // 测试
1634    // ------------------------------------------------------------------------
1635
1636    #[cfg(test)]
1637    mod tests {
1638        use super::*;
1639
1640        /// 测试 device code 请求
1641        #[tokio::test]
1642        async fn test_device_code_request() {
1643            let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
1644            transport.push_response(
1645                r#"{"device_code":"dc123","user_code":"UC-ABCD","verification_uri":"https://provider.com/device","expires_in":600,"interval":5}"#,
1646            );
1647
1648            let config = OAuth2Config::new(
1649                "client123",
1650                "secret456",
1651                "https://example.com/callback",
1652                "https://provider.com/authorize",
1653                "https://provider.com/token",
1654            )
1655            .with_device_auth_url("https://provider.com/device_authorize");
1656            let provider = DeviceCodeOAuth2Provider::new(config, transport);
1657
1658            let resp = provider
1659                .request_device_code(&["read".into(), "write".into()])
1660                .await
1661                .expect("request_device_code 失败");
1662
1663            assert_eq!(resp.device_code, "dc123");
1664            assert_eq!(resp.user_code, "UC-ABCD");
1665            assert_eq!(resp.verification_uri, "https://provider.com/device");
1666            assert_eq!(resp.expires_in, 600);
1667            assert_eq!(resp.interval, 5);
1668        }
1669
1670        /// 测试 device code 轮询 authorization_pending → 继续轮询 → 成功
1671        #[tokio::test]
1672        async fn test_device_code_poll_pending() {
1673            let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
1674            // 第一次:pending
1675            transport.push_response(r#"{"error":"authorization_pending"}"#);
1676            // 第二次:成功
1677            transport.push_response(
1678                r#"{"access_token":"token123","token_type":"Bearer","expires_in":3600}"#,
1679            );
1680
1681            let config = OAuth2Config::new(
1682                "client123",
1683                "secret456",
1684                "https://example.com/callback",
1685                "https://provider.com/authorize",
1686                "https://provider.com/token",
1687            );
1688            let provider = DeviceCodeOAuth2Provider::new(config, transport);
1689
1690            let token = provider
1691                .poll_for_token("dc123", 0, 600)
1692                .await
1693                .expect("poll_for_token 失败");
1694
1695            assert_eq!(token.access_token, "token123");
1696            assert_eq!(token.token_type.as_deref(), Some("Bearer"));
1697        }
1698
1699        /// 测试 device code 轮询 slow_down → interval += 5
1700        #[tokio::test]
1701        async fn test_device_code_poll_slow_down() {
1702            let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
1703            // 第一次:slow_down
1704            transport.push_response(r#"{"error":"slow_down"}"#);
1705            // 第二次:成功
1706            transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
1707
1708            let config = OAuth2Config::new(
1709                "client123",
1710                "secret456",
1711                "https://example.com/callback",
1712                "https://provider.com/authorize",
1713                "https://provider.com/token",
1714            );
1715            let provider = DeviceCodeOAuth2Provider::new(config, transport);
1716
1717            let token = provider
1718                .poll_for_token("dc123", 0, 600)
1719                .await
1720                .expect("poll_for_token 失败");
1721
1722            assert_eq!(token.access_token, "token123");
1723        }
1724
1725        /// 测试 device code 过期 → OAUTH2_DEVICE_CODE_EXPIRED
1726        #[tokio::test]
1727        async fn test_device_code_expired() {
1728            let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
1729            // 持续 pending,直到过期
1730            transport.push_response(r#"{"error":"authorization_pending"}"#);
1731
1732            let config = OAuth2Config::new(
1733                "client123",
1734                "secret456",
1735                "https://example.com/callback",
1736                "https://provider.com/authorize",
1737                "https://provider.com/token",
1738            );
1739            let provider = DeviceCodeOAuth2Provider::new(config, transport);
1740
1741            let err = provider.poll_for_token("dc123", 0, 0).await.unwrap_err();
1742            assert!(
1743                err.to_string().contains("OAUTH2_DEVICE_CODE_EXPIRED"),
1744                "应返回设备码过期错误: {err}"
1745            );
1746        }
1747
1748        /// 测试 device code access_denied → 停止轮询
1749        #[tokio::test]
1750        async fn test_device_code_access_denied() {
1751            let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
1752            transport.push_response(r#"{"error":"access_denied"}"#);
1753
1754            let config = OAuth2Config::new(
1755                "client123",
1756                "secret456",
1757                "https://example.com/callback",
1758                "https://provider.com/authorize",
1759                "https://provider.com/token",
1760            );
1761            let provider = DeviceCodeOAuth2Provider::new(config, transport);
1762
1763            let err = provider.poll_for_token("dc123", 5, 600).await.unwrap_err();
1764            assert!(
1765                err.to_string().contains("OAUTH2_ACCESS_DENIED"),
1766                "应返回 access_denied 错误: {err}"
1767            );
1768        }
1769
1770        /// 测试 device_auth_url 未设置 → 配置错误
1771        #[tokio::test]
1772        async fn test_device_code_no_auth_url() {
1773            let transport = Arc::new(MemoryAsyncOAuth2HttpTransport::new());
1774            let config = OAuth2Config::new(
1775                "client123",
1776                "secret456",
1777                "https://example.com/callback",
1778                "https://provider.com/authorize",
1779                "https://provider.com/token",
1780            );
1781            // 不设置 device_auth_url
1782            let provider = DeviceCodeOAuth2Provider::new(config, transport);
1783
1784            let err = provider.request_device_code(&[]).await.unwrap_err();
1785            assert!(matches!(err, OAuth2Error::MissingField(field) if field == "device_auth_url"));
1786        }
1787    }
1788}
1789
1790// ============================================================================
1791// 辅助函数
1792// ============================================================================
1793
1794/// 从 JSON 值中提取字符串
1795///
1796/// 支持字符串和整数类型(整数转为十进制字符串),其他类型返回 `None`。
1797fn extract_string(value: &serde_json::Value) -> Option<String> {
1798    match value {
1799        serde_json::Value::String(string) => Some(string.clone()),
1800        serde_json::Value::Number(number) => number.as_i64().map(|number| number.to_string()),
1801        _ => None,
1802    }
1803}
1804
1805/// 简易百分号编码 — 用于 URL 查询参数
1806///
1807/// 对齐 RFC 3986 的 unreserved 字符集(`A-Za-z0-9-._~`)保持原样,
1808/// 其余字符编码为 `%XX` 形式(UTF-8 字节)。
1809fn percent_encode(input: &str) -> String {
1810    let mut output = String::with_capacity(input.len());
1811    for byte in input.as_bytes() {
1812        if matches!(byte, b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~') {
1813            output.push(*byte as char);
1814        } else {
1815            output.push_str(&format!("%{byte:02X}"));
1816        }
1817    }
1818    output
1819}
1820
1821// ============================================================================
1822// 单元测试
1823// ============================================================================
1824
1825#[cfg(test)]
1826mod tests {
1827    use super::*;
1828    use proptest::{prop_assert, prop_assert_eq};
1829
1830    // ------------------------------------------------------------------------
1831    // OAuth2Config 测试
1832    // ------------------------------------------------------------------------
1833
1834    /// 测试 OAuth2Config builder 模式(所有可选字段)
1835    #[test]
1836    fn test_oauth2_config_builder() {
1837        let config = OAuth2Config::new(
1838            "client123",
1839            "secret456",
1840            "https://example.com/callback",
1841            "https://provider.com/authorize",
1842            "https://provider.com/token",
1843        )
1844        .with_user_url("https://provider.com/user/info")
1845        .with_scopes(vec!["scope1".into(), "scope2".into()])
1846        .with_extra_param("foo", "bar");
1847
1848        assert_eq!(config.client_id, "client123");
1849        assert_eq!(config.client_secret, "secret456");
1850        assert_eq!(config.redirect_url, "https://example.com/callback");
1851        assert_eq!(config.auth_url, "https://provider.com/authorize");
1852        assert_eq!(config.token_url, "https://provider.com/token");
1853        assert_eq!(
1854            config.user_url.as_deref(),
1855            Some("https://provider.com/user/info")
1856        );
1857        assert_eq!(config.scopes, vec!["scope1", "scope2"]);
1858        assert_eq!(config.extra_params, vec![("foo".into(), "bar".into())]);
1859    }
1860
1861    /// 测试 OAuth2Config 最小配置(仅必填字段)
1862    #[test]
1863    fn test_oauth2_config_minimal() {
1864        let config = OAuth2Config::new(
1865            "client123",
1866            "secret456",
1867            "https://example.com/callback",
1868            "https://provider.com/authorize",
1869            "https://provider.com/token",
1870        );
1871
1872        assert_eq!(config.client_id, "client123");
1873        assert_eq!(config.client_secret, "secret456");
1874        assert_eq!(config.redirect_url, "https://example.com/callback");
1875        assert_eq!(config.auth_url, "https://provider.com/authorize");
1876        assert_eq!(config.token_url, "https://provider.com/token");
1877        assert!(config.user_url.is_none());
1878        assert!(config.scopes.is_empty());
1879        assert!(config.extra_params.is_empty());
1880
1881        // 最小配置应通过校验
1882        assert!(config.validate().is_ok());
1883    }
1884
1885    /// 测试 OAuth2Config::with_scope 追加多个 scope
1886    #[test]
1887    fn test_oauth2_config_with_scope_chained() {
1888        let config = OAuth2Config::new(
1889            "id",
1890            "secret",
1891            "https://example.com/callback",
1892            "https://provider.com/authorize",
1893            "https://provider.com/token",
1894        )
1895        .with_scope("get_user_info")
1896        .with_scope("get_unionid");
1897
1898        assert_eq!(config.scopes, vec!["get_user_info", "get_unionid"]);
1899    }
1900
1901    /// 测试 OAuth2Config::with_extra_params 覆盖
1902    #[test]
1903    fn test_oauth2_config_with_extra_params() {
1904        let config = OAuth2Config::new(
1905            "id",
1906            "secret",
1907            "https://example.com/callback",
1908            "https://provider.com/authorize",
1909            "https://provider.com/token",
1910        )
1911        .with_extra_param("a", "1")
1912        .with_extra_param("b", "2")
1913        .with_extra_params(vec![("x".into(), "10".into())]);
1914
1915        assert_eq!(config.extra_params, vec![("x".into(), "10".into())]);
1916    }
1917
1918    /// 测试 OAuth2Config::validate 检测空字段
1919    #[test]
1920    fn test_oauth2_config_validate_empty_fields() {
1921        // client_id 为空
1922        let config = OAuth2Config::new(
1923            "",
1924            "secret",
1925            "https://example.com/callback",
1926            "https://provider.com/authorize",
1927            "https://provider.com/token",
1928        );
1929        let err = config.validate().unwrap_err();
1930        assert!(matches!(err, OAuth2Error::MissingField(field) if field == "client_id"));
1931
1932        // client_secret 为空
1933        let config = OAuth2Config::new(
1934            "id",
1935            "",
1936            "https://example.com/callback",
1937            "https://provider.com/authorize",
1938            "https://provider.com/token",
1939        );
1940        let err = config.validate().unwrap_err();
1941        assert!(matches!(err, OAuth2Error::MissingField(field) if field == "client_secret"));
1942
1943        // redirect_url 为空
1944        let config = OAuth2Config::new(
1945            "id",
1946            "secret",
1947            "",
1948            "https://provider.com/authorize",
1949            "https://provider.com/token",
1950        );
1951        let err = config.validate().unwrap_err();
1952        assert!(matches!(err, OAuth2Error::MissingField(field) if field == "redirect_url"));
1953
1954        // auth_url 为空
1955        let config = OAuth2Config::new(
1956            "id",
1957            "secret",
1958            "https://example.com/callback",
1959            "",
1960            "https://provider.com/token",
1961        );
1962        let err = config.validate().unwrap_err();
1963        assert!(matches!(err, OAuth2Error::MissingField(field) if field == "auth_url"));
1964
1965        // token_url 为空
1966        let config = OAuth2Config::new(
1967            "id",
1968            "secret",
1969            "https://example.com/callback",
1970            "https://provider.com/authorize",
1971            "",
1972        );
1973        let err = config.validate().unwrap_err();
1974        assert!(matches!(err, OAuth2Error::MissingField(field) if field == "token_url"));
1975    }
1976
1977    // ------------------------------------------------------------------------
1978    // SocialiteUser 测试
1979    // ------------------------------------------------------------------------
1980
1981    /// 测试 SocialiteUser 默认值
1982    #[test]
1983    fn test_socialite_user_default() {
1984        let user = SocialiteUser::default();
1985        assert!(user.id.is_empty());
1986        assert!(user.nickname.is_none());
1987        assert!(user.name.is_none());
1988        assert!(user.email.is_none());
1989        assert!(user.avatar.is_none());
1990        assert!(user.raw.is_null());
1991        assert!(user.access_token.is_none());
1992        assert!(user.refresh_token.is_none());
1993        assert!(user.expires_in.is_none());
1994    }
1995
1996    /// 测试 SocialiteUser 序列化/反序列化
1997    #[test]
1998    fn test_socialite_user_serialize_deserialize() {
1999        let user = SocialiteUser {
2000            id: "123".into(),
2001            nickname: Some("tester".into()),
2002            name: Some("Test User".into()),
2003            email: Some("test@example.com".into()),
2004            avatar: Some("https://example.com/avatar.png".into()),
2005            raw: serde_json::json!({"key": "value"}),
2006            access_token: Some("token123".into()),
2007            refresh_token: Some("refresh456".into()),
2008            expires_in: Some(3600),
2009        };
2010
2011        let json = serde_json::to_string(&user).expect("序列化失败");
2012
2013        // P0-SEC-01 安全修复:access_token / refresh_token 不应出现在序列化输出中
2014        // 防止令牌通过 API 响应泄漏(对齐 MerchantUser.password 的 skip_serializing 策略)
2015        assert!(
2016            !json.contains("access_token"),
2017            "access_token 不应出现在序列化 JSON 中(安全脱敏要求): {json}"
2018        );
2019        assert!(
2020            !json.contains("refresh_token"),
2021            "refresh_token 不应出现在序列化 JSON 中(安全脱敏要求): {json}"
2022        );
2023
2024        let parsed: SocialiteUser = serde_json::from_str(&json).expect("反序列化失败");
2025
2026        assert_eq!(parsed.id, "123");
2027        assert_eq!(parsed.nickname.as_deref(), Some("tester"));
2028        assert_eq!(parsed.name.as_deref(), Some("Test User"));
2029        assert_eq!(parsed.email.as_deref(), Some("test@example.com"));
2030        assert_eq!(
2031            parsed.avatar.as_deref(),
2032            Some("https://example.com/avatar.png")
2033        );
2034        // 反序列化后 token 字段为 None(序列化时未包含,反序列化用 #[serde(default)])
2035        assert_eq!(parsed.access_token, None);
2036        assert_eq!(parsed.refresh_token, None);
2037        assert_eq!(parsed.expires_in, Some(3600));
2038    }
2039
2040    // ------------------------------------------------------------------------
2041    // redirect_url 测试
2042    // ------------------------------------------------------------------------
2043
2044    /// 测试 redirect_url 包含必填查询参数
2045    #[test]
2046    fn test_redirect_url_contains_required_params() {
2047        let config = OAuth2Config::new(
2048            "client123",
2049            "secret456",
2050            "https://example.com/callback",
2051            "https://provider.com/oauth2.0/authorize",
2052            "https://provider.com/oauth2.0/token",
2053        );
2054        let provider =
2055            GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2056
2057        let url = provider.redirect_url("random_state_abc");
2058
2059        assert!(url.starts_with("https://provider.com/oauth2.0/authorize?"));
2060        assert!(url.contains("client_id=client123"));
2061        assert!(url.contains("redirect_uri=https%3A%2F%2Fexample.com%2Fcallback"));
2062        assert!(url.contains("response_type=code"));
2063        assert!(url.contains("state=random_state_abc"));
2064        // 未配置 scopes 时不应包含 scope
2065        assert!(!url.contains("scope="));
2066    }
2067
2068    /// 测试 redirect_url 包含 scopes(空格连接,百分号编码为 %20)
2069    #[test]
2070    fn test_redirect_url_with_scopes() {
2071        let config = OAuth2Config::new(
2072            "client123",
2073            "secret456",
2074            "https://example.com/callback",
2075            "https://provider.com/oauth2.0/authorize",
2076            "https://provider.com/oauth2.0/token",
2077        )
2078        .with_scopes(vec!["get_user_info".into(), "get_unionid".into()]);
2079        let provider =
2080            GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2081
2082        let url = provider.redirect_url("state123");
2083
2084        // scope 以空格连接,空格编码为 %20
2085        assert!(url.contains("scope=get_user_info%20get_unionid"));
2086    }
2087
2088    /// 测试 redirect_url 包含额外参数
2089    #[test]
2090    fn test_redirect_url_with_extra_params() {
2091        let config = OAuth2Config::new(
2092            "client123",
2093            "secret456",
2094            "https://example.com/callback",
2095            "https://provider.com/oauth2.0/authorize",
2096            "https://provider.com/oauth2.0/token",
2097        )
2098        .with_extra_param("foo", "bar")
2099        .with_extra_param("display", "mobile");
2100        let provider =
2101            GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2102
2103        let url = provider.redirect_url("state123");
2104
2105        assert!(url.contains("foo=bar"));
2106        assert!(url.contains("display=mobile"));
2107    }
2108
2109    /// 测试 redirect_url 在已有查询字符串的 auth_url 上追加参数
2110    #[test]
2111    fn test_redirect_url_with_existing_query() {
2112        let config = OAuth2Config::new(
2113            "client123",
2114            "secret456",
2115            "https://example.com/callback",
2116            "https://provider.com/authorize?foo=bar",
2117            "https://provider.com/token",
2118        );
2119        let provider =
2120            GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2121
2122        let url = provider.redirect_url("state123");
2123
2124        // 已有查询字符串时应使用 & 追加
2125        assert!(url.contains("?foo=bar&"));
2126        assert!(url.contains("client_id=client123"));
2127    }
2128
2129    // ------------------------------------------------------------------------
2130    // MemoryOAuth2HttpTransport 测试
2131    // ------------------------------------------------------------------------
2132
2133    /// 测试 MemoryOAuth2HttpTransport 记录请求并返回预置响应
2134    #[test]
2135    fn test_memory_oauth2_http_transport_post_json() {
2136        let transport = MemoryOAuth2HttpTransport::new();
2137        transport.push_response(r#"{"access_token":"token123"}"#);
2138
2139        let response = transport
2140            .post_json("https://example.com/token", r#"{"code":"abc"}"#)
2141            .expect("post_json 失败");
2142
2143        assert_eq!(response, r#"{"access_token":"token123"}"#);
2144        assert_eq!(transport.count(), 1);
2145
2146        let (url, body) = transport.last().expect("应有请求记录");
2147        assert_eq!(url, "https://example.com/token");
2148        assert_eq!(body, r#"{"code":"abc"}"#);
2149    }
2150
2151    /// 测试 MemoryOAuth2HttpTransport 响应队列按顺序消费
2152    #[test]
2153    fn test_memory_oauth2_http_transport_response_queue() {
2154        let transport = MemoryOAuth2HttpTransport::new();
2155        transport.push_response("resp1");
2156        transport.push_response("resp2");
2157
2158        let resp1 = transport
2159            .post_json("url1", "body1")
2160            .expect("第一次调用失败");
2161        let resp2 = transport
2162            .post_json("url2", "body2")
2163            .expect("第二次调用失败");
2164
2165        assert_eq!(resp1, "resp1");
2166        assert_eq!(resp2, "resp2");
2167        assert_eq!(transport.count(), 2);
2168    }
2169
2170    /// 测试 MemoryOAuth2HttpTransport 响应耗尽后返回空字符串
2171    #[test]
2172    fn test_memory_oauth2_http_transport_empty_response() {
2173        let transport = MemoryOAuth2HttpTransport::new();
2174        // 不预置响应
2175        let response = transport
2176            .post_json("url", "body")
2177            .expect("post_json 不应失败");
2178        assert_eq!(response, "");
2179    }
2180
2181    /// 测试 MemoryOAuth2HttpTransport clear
2182    #[test]
2183    fn test_memory_oauth2_http_transport_clear() {
2184        let transport = MemoryOAuth2HttpTransport::new();
2185        transport.push_response("resp");
2186        transport.post_json("url", "body").expect("调用失败");
2187        assert_eq!(transport.count(), 1);
2188
2189        transport.clear();
2190        assert_eq!(transport.count(), 0);
2191        // clear 后响应队列也清空,返回空字符串
2192        let response = transport
2193            .post_json("url", "body")
2194            .expect("post_json 不应失败");
2195        assert_eq!(response, "");
2196    }
2197
2198    // ------------------------------------------------------------------------
2199    // GenericOAuth2Provider::user_from_token 测试
2200    // ------------------------------------------------------------------------
2201
2202    /// 测试 GenericOAuth2Provider::user_from_token 完整流程
2203    ///
2204    /// 使用 MemoryOAuth2HttpTransport mock token 响应和用户信息响应。
2205    #[test]
2206    fn test_generic_oauth2_provider_user_from_token() {
2207        let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2208        // 预置 mock 响应:token 响应 + 用户信息响应
2209        transport.push_response(r#"{"access_token":"token123","refresh_token":"refresh456","expires_in":3600,"openid":"openid_abc"}"#);
2210        transport.push_response(
2211            r#"{"id":"12345","nickname":"test_user","name":"Test","email":"test@example.com","avatar":"https://example.com/avatar.png"}"#,
2212        );
2213
2214        let config = OAuth2Config::new(
2215            "client123",
2216            "secret456",
2217            "https://example.com/callback",
2218            "https://provider.com/authorize",
2219            "https://provider.com/token",
2220        )
2221        .with_user_url("https://provider.com/user/info");
2222        let provider = GenericOAuth2Provider::new(config, transport.clone());
2223
2224        let user = provider
2225            .user_from_token("auth_code_abc")
2226            .expect("user_from_token 失败");
2227
2228        // 验证 token 字段
2229        assert_eq!(user.access_token.as_deref(), Some("token123"));
2230        assert_eq!(user.refresh_token.as_deref(), Some("refresh456"));
2231        assert_eq!(user.expires_in, Some(3600));
2232
2233        // 验证用户信息字段
2234        assert_eq!(user.id, "12345");
2235        assert_eq!(user.nickname.as_deref(), Some("test_user"));
2236        assert_eq!(user.name.as_deref(), Some("Test"));
2237        assert_eq!(user.email.as_deref(), Some("test@example.com"));
2238        assert_eq!(
2239            user.avatar.as_deref(),
2240            Some("https://example.com/avatar.png")
2241        );
2242
2243        // 验证原始响应数据
2244        assert_eq!(user.raw["id"], "12345");
2245        assert_eq!(user.raw["nickname"], "test_user");
2246
2247        // 验证 HTTP 请求次数(token + user info)
2248        assert_eq!(transport.count(), 2);
2249    }
2250
2251    /// 测试 GenericOAuth2Provider::user_from_token 无 user_url 时仅返回 token
2252    #[test]
2253    fn test_generic_oauth2_provider_user_from_token_no_user_url() {
2254        let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2255        transport.push_response(r#"{"access_token":"token123","expires_in":7200}"#);
2256
2257        let config = OAuth2Config::new(
2258            "client123",
2259            "secret456",
2260            "https://example.com/callback",
2261            "https://provider.com/authorize",
2262            "https://provider.com/token",
2263        );
2264        // 不配置 user_url
2265        let provider = GenericOAuth2Provider::new(config, transport.clone());
2266
2267        let user = provider
2268            .user_from_token("auth_code")
2269            .expect("user_from_token 失败");
2270
2271        assert_eq!(user.access_token.as_deref(), Some("token123"));
2272        assert_eq!(user.expires_in, Some(7200));
2273        assert!(user.refresh_token.is_none());
2274        // 无 user_url 时不获取用户信息,id 为空
2275        assert!(user.id.is_empty());
2276        // 仅一次 HTTP 请求(token 交换)
2277        assert_eq!(transport.count(), 1);
2278    }
2279
2280    /// 测试 GenericOAuth2Provider::user_from_token 授权码为空时返回错误
2281    #[test]
2282    fn test_generic_oauth2_provider_missing_code() {
2283        let config = OAuth2Config::new(
2284            "client123",
2285            "secret456",
2286            "https://example.com/callback",
2287            "https://provider.com/authorize",
2288            "https://provider.com/token",
2289        );
2290        let provider =
2291            GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2292
2293        let err = provider.user_from_token("").unwrap_err();
2294        assert!(matches!(err, OAuth2Error::AuthFailed(msg) if msg.contains("授权码")));
2295    }
2296
2297    /// 测试 GenericOAuth2Provider 配置缺少必填字段时返回错误
2298    #[test]
2299    fn test_oauth2_provider_missing_config_fields() {
2300        // client_id 为空
2301        let config = OAuth2Config::new(
2302            "",
2303            "secret456",
2304            "https://example.com/callback",
2305            "https://provider.com/authorize",
2306            "https://provider.com/token",
2307        );
2308        let provider = GenericOAuth2Provider::new(config, Arc::new(MemoryHttpTransport));
2309
2310        let err = provider.user_from_token("code").unwrap_err();
2311        assert!(matches!(err, OAuth2Error::MissingField(field) if field == "client_id"));
2312
2313        // token_url 为空
2314        let config = OAuth2Config::new(
2315            "client123",
2316            "secret456",
2317            "https://example.com/callback",
2318            "https://provider.com/authorize",
2319            "",
2320        );
2321        let provider = GenericOAuth2Provider::new(config, Arc::new(MemoryHttpTransport));
2322
2323        let err = provider.user_from_token("code").unwrap_err();
2324        assert!(matches!(err, OAuth2Error::MissingField(field) if field == "token_url"));
2325    }
2326
2327    /// 测试 token 响应缺少 access_token 时返回 TokenExchangeFailed
2328    #[test]
2329    fn test_generic_oauth2_provider_token_response_missing_access_token() {
2330        let transport = MemoryOAuth2HttpTransport::new();
2331        transport.push_response(r#"{"error":"invalid_grant"}"#);
2332
2333        let config = OAuth2Config::new(
2334            "client123",
2335            "secret456",
2336            "https://example.com/callback",
2337            "https://provider.com/authorize",
2338            "https://provider.com/token",
2339        );
2340        let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
2341
2342        let err = provider.user_from_token("code").unwrap_err();
2343        assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
2344    }
2345
2346    /// 测试 token 响应非 JSON 时返回 TokenExchangeFailed
2347    #[test]
2348    fn test_generic_oauth2_provider_token_response_invalid_json() {
2349        let transport = MemoryOAuth2HttpTransport::new();
2350        transport.push_response("not a json");
2351
2352        let config = OAuth2Config::new(
2353            "client123",
2354            "secret456",
2355            "https://example.com/callback",
2356            "https://provider.com/authorize",
2357            "https://provider.com/token",
2358        );
2359        let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
2360
2361        let err = provider.user_from_token("code").unwrap_err();
2362        assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
2363    }
2364
2365    /// 测试用户信息字段兼容多种命名(openid / figureurl_qq_1 等)
2366    #[test]
2367    fn test_generic_oauth2_provider_user_info_field_aliases() {
2368        let transport = MemoryOAuth2HttpTransport::new();
2369        transport.push_response(r#"{"access_token":"token123","openid":"openid_abc"}"#);
2370        transport.push_response(
2371            r#"{"openid":"qq_12345","nickname":"qq_user","figureurl_qq_1":"https://qzapp.qlogo.cn/1.png"}"#,
2372        );
2373
2374        let config = OAuth2Config::new(
2375            "client123",
2376            "secret456",
2377            "https://example.com/callback",
2378            "https://provider.com/authorize",
2379            "https://provider.com/token",
2380        )
2381        .with_user_url("https://provider.com/user/info");
2382        let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
2383
2384        let user = provider
2385            .user_from_token("code")
2386            .expect("user_from_token 失败");
2387
2388        // openid 作为 id
2389        assert_eq!(user.id, "qq_12345");
2390        assert_eq!(user.nickname.as_deref(), Some("qq_user"));
2391        // figureurl_qq_1 作为 avatar
2392        assert_eq!(user.avatar.as_deref(), Some("https://qzapp.qlogo.cn/1.png"));
2393    }
2394
2395    /// 测试用户 ID 为整数类型时正确转为字符串
2396    #[test]
2397    fn test_generic_oauth2_provider_user_id_integer() {
2398        let transport = MemoryOAuth2HttpTransport::new();
2399        transport.push_response(r#"{"access_token":"token123"}"#);
2400        transport.push_response(r#"{"id":12345,"nickname":"github_user"}"#);
2401
2402        let config = OAuth2Config::new(
2403            "client123",
2404            "secret456",
2405            "https://example.com/callback",
2406            "https://provider.com/authorize",
2407            "https://provider.com/token",
2408        )
2409        .with_user_url("https://provider.com/user/info");
2410        let provider = GenericOAuth2Provider::new(config, Arc::new(transport));
2411
2412        let user = provider
2413            .user_from_token("code")
2414            .expect("user_from_token 失败");
2415
2416        assert_eq!(user.id, "12345");
2417        assert_eq!(user.nickname.as_deref(), Some("github_user"));
2418    }
2419
2420    /// 测试 HTTP 传输失败时返回 HttpTransport 错误
2421    #[test]
2422    fn test_generic_oauth2_provider_http_transport_failure() {
2423        let config = OAuth2Config::new(
2424            "client123",
2425            "secret456",
2426            "https://example.com/callback",
2427            "https://provider.com/authorize",
2428            "https://provider.com/token",
2429        );
2430        let provider = GenericOAuth2Provider::new(config, Arc::new(FailingTransport));
2431
2432        let err = provider.user_from_token("code").unwrap_err();
2433        assert!(matches!(err, OAuth2Error::HttpTransport(_)));
2434    }
2435
2436    /// 测试 percent_encode 函数
2437    #[test]
2438    fn test_percent_encode() {
2439        // unreserved 字符保持原样
2440        assert_eq!(percent_encode("abcXYZ09-._~"), "abcXYZ09-._~");
2441        // 空格编码为 %20
2442        assert_eq!(percent_encode("a b"), "a%20b");
2443        // 斜杠编码为 %2F
2444        assert_eq!(percent_encode("/"), "%2F");
2445        // 冒号编码为 %3A
2446        assert_eq!(percent_encode(":"), "%3A");
2447        // URL 编码
2448        assert_eq!(
2449            percent_encode("https://example.com/path"),
2450            "https%3A%2F%2Fexample.com%2Fpath"
2451        );
2452        // 中文字符(UTF-8 编码)
2453        assert_eq!(percent_encode("中"), "%E4%B8%AD");
2454    }
2455
2456    /// 测试 extract_string 函数
2457    #[test]
2458    fn test_extract_string() {
2459        // 字符串
2460        assert_eq!(
2461            extract_string(&serde_json::json!("hello")),
2462            Some("hello".into())
2463        );
2464        // 整数
2465        assert_eq!(
2466            extract_string(&serde_json::json!(12345)),
2467            Some("12345".into())
2468        );
2469        // 浮点数(不支持,返回 None)
2470        assert_eq!(extract_string(&serde_json::json!(1.5)), None);
2471        // 布尔值(不支持,返回 None)
2472        assert_eq!(extract_string(&serde_json::json!(true)), None);
2473        // null(不支持,返回 None)
2474        assert_eq!(extract_string(&serde_json::Value::Null), None);
2475        // 对象(不支持,返回 None)
2476        assert_eq!(extract_string(&serde_json::json!({"a": 1})), None);
2477    }
2478
2479    // ------------------------------------------------------------------------
2480    // 测试辅助类型
2481    // ------------------------------------------------------------------------
2482
2483    /// 始终失败的 HTTP 传输(用于测试错误路径)
2484    struct FailingTransport;
2485
2486    impl OAuth2HttpTransport for FailingTransport {
2487        fn post_json(&self, _url: &str, _body: &str) -> Result<String, OAuth2Error> {
2488            Err(OAuth2Error::HttpTransport("connection refused".into()))
2489        }
2490    }
2491
2492    /// 空的 HTTP 传输(用于仅需校验、不实际发送的测试)
2493    struct MemoryHttpTransport;
2494
2495    impl OAuth2HttpTransport for MemoryHttpTransport {
2496        fn post_json(&self, _url: &str, _body: &str) -> Result<String, OAuth2Error> {
2497            Ok(String::new())
2498        }
2499    }
2500
2501    // ------------------------------------------------------------------------
2502    // T1: PKCE 与 state 自动生成测试
2503    // ------------------------------------------------------------------------
2504
2505    /// 测试 state 自动生成:32 字符 hex、两次调用不同
2506    #[test]
2507    fn test_state_auto_generate() {
2508        let state1 = OAuth2Config::generate_state();
2509        let state2 = OAuth2Config::generate_state();
2510
2511        // 32 字符 hex
2512        assert_eq!(state1.len(), 32, "state 应为 32 字符 hex(16 字节)");
2513        assert_eq!(state2.len(), 32, "state 应为 32 字符 hex(16 字节)");
2514
2515        // 全部为 hex 字符
2516        assert!(
2517            state1.chars().all(|c| c.is_ascii_hexdigit()),
2518            "state 应全为 hex 字符: {state1}"
2519        );
2520        assert!(
2521            state2.chars().all(|c| c.is_ascii_hexdigit()),
2522            "state 应全为 hex 字符: {state2}"
2523        );
2524
2525        // 两次调用不同(密码学随机,碰撞概率极低)
2526        assert_ne!(state1, state2, "两次生成的 state 不应相同");
2527    }
2528
2529    /// 测试 PKCE 参数对生成:verifier 43-128 字符、challenge == base64url(SHA256(verifier))
2530    #[test]
2531    fn test_pkce_pair_generate() {
2532        let pkce = OAuth2Config::generate_pkce_pair();
2533
2534        // code_verifier 应为 64 字符 hex(32 字节)
2535        assert!(
2536            pkce.code_verifier.len() >= 43 && pkce.code_verifier.len() <= 128,
2537            "code_verifier 长度应在 43-128 之间,实际: {}",
2538            pkce.code_verifier.len()
2539        );
2540        assert_eq!(
2541            pkce.code_verifier.len(),
2542            64,
2543            "code_verifier 应为 64 字符 hex"
2544        );
2545        assert!(
2546            pkce.code_verifier.chars().all(|c| c.is_ascii_hexdigit()),
2547            "code_verifier 应全为 hex 字符"
2548        );
2549
2550        // code_challenge 应为 base64url(SHA256(code_verifier))
2551        let mut hasher = sha2::Sha256::new();
2552        sha2::Digest::update(&mut hasher, pkce.code_verifier.as_bytes());
2553        let digest = sha2::Digest::finalize(hasher);
2554        let expected_challenge = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest);
2555        assert_eq!(
2556            pkce.code_challenge, expected_challenge,
2557            "code_challenge 应等于 base64url(SHA256(code_verifier))"
2558        );
2559
2560        // method 固定 S256
2561        assert_eq!(pkce.method, PkceMethod::S256);
2562    }
2563
2564    /// 测试 authorization_code 流程带 PKCE:URL 含 code_challenge
2565    #[test]
2566    fn test_authorization_code_with_pkce() {
2567        let config = OAuth2Config::new(
2568            "client123",
2569            "secret456",
2570            "https://example.com/callback",
2571            "https://provider.com/oauth2.0/authorize",
2572            "https://provider.com/oauth2.0/token",
2573        )
2574        .with_pkce(true);
2575        let provider =
2576            GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2577
2578        let pkce = OAuth2Config::generate_pkce_pair();
2579        let url = provider.redirect_url_with_pkce("state123", &pkce);
2580
2581        assert!(url.contains("code_challenge="), "URL 应包含 code_challenge");
2582        assert!(
2583            url.contains("code_challenge_method=S256"),
2584            "URL 应包含 code_challenge_method=S256"
2585        );
2586        assert!(url.contains("response_type=code"));
2587        assert!(url.contains("state=state123"));
2588    }
2589
2590    /// 测试 authorization_code 流程不带 PKCE:URL 不含 code_challenge
2591    #[test]
2592    fn test_authorization_code_without_pkce() {
2593        let config = OAuth2Config::new(
2594            "client123",
2595            "secret456",
2596            "https://example.com/callback",
2597            "https://provider.com/oauth2.0/authorize",
2598            "https://provider.com/oauth2.0/token",
2599        );
2600        // pkce_enabled 默认 false
2601        assert!(!config.pkce_enabled);
2602
2603        let provider =
2604            GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2605        let url = provider.redirect_url("state123");
2606
2607        assert!(
2608            !url.contains("code_challenge"),
2609            "URL 不应包含 code_challenge(PKCE 未启用)"
2610        );
2611    }
2612
2613    /// 测试 client_secret 在 Debug 输出中脱敏
2614    #[test]
2615    fn test_client_secret_not_in_debug() {
2616        let config = OAuth2Config::new(
2617            "client123",
2618            "super_secret_value_456",
2619            "https://example.com/callback",
2620            "https://provider.com/authorize",
2621            "https://provider.com/token",
2622        );
2623
2624        let debug_str = format!("{:?}", config);
2625        assert!(
2626            !debug_str.contains("super_secret_value_456"),
2627            "client_secret 不应出现在 Debug 输出中: {debug_str}"
2628        );
2629        assert!(
2630            debug_str.contains("***"),
2631            "Debug 输出应包含脱敏标记 '***': {debug_str}"
2632        );
2633    }
2634
2635    /// 测试 PkceParams 的 Debug 输出不包含 code_verifier 明文
2636    #[test]
2637    fn test_pkce_params_debug_redacted() {
2638        let pkce = OAuth2Config::generate_pkce_pair();
2639        let debug_str = format!("{:?}", pkce);
2640        assert!(
2641            !debug_str.contains(&pkce.code_verifier),
2642            "code_verifier 明文不应出现在 Debug 输出中: {debug_str}"
2643        );
2644        assert!(
2645            debug_str.contains("redacted"),
2646            "Debug 输出应包含 'redacted' 标记: {debug_str}"
2647        );
2648    }
2649
2650    /// 测试 with_pkce builder 方法
2651    #[test]
2652    fn test_with_pkce_builder() {
2653        let config = OAuth2Config::new(
2654            "id",
2655            "secret",
2656            "https://example.com/callback",
2657            "https://provider.com/authorize",
2658            "https://provider.com/token",
2659        );
2660        assert!(!config.pkce_enabled, "默认 pkce_enabled 应为 false");
2661
2662        let config = config.with_pkce(true);
2663        assert!(
2664            config.pkce_enabled,
2665            "with_pkce(true) 后 pkce_enabled 应为 true"
2666        );
2667
2668        let config = config.with_pkce(false);
2669        assert!(
2670            !config.pkce_enabled,
2671            "with_pkce(false) 后 pkce_enabled 应为 false"
2672        );
2673    }
2674
2675    /// 测试 with_device_auth_url builder 方法
2676    #[test]
2677    fn test_with_device_auth_url_builder() {
2678        let config = OAuth2Config::new(
2679            "id",
2680            "secret",
2681            "https://example.com/callback",
2682            "https://provider.com/authorize",
2683            "https://provider.com/token",
2684        );
2685        assert!(
2686            config.device_auth_url.is_none(),
2687            "默认 device_auth_url 应为 None"
2688        );
2689
2690        let config = config.with_device_auth_url("https://provider.com/device_authorize");
2691        assert_eq!(
2692            config.device_auth_url.as_deref(),
2693            Some("https://provider.com/device_authorize"),
2694        );
2695    }
2696
2697    /// 测试 PKCE 流程的 token 交换:POST body 包含 code_verifier
2698    #[test]
2699    fn test_exchange_token_with_pkce_verifier() {
2700        let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2701        transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
2702
2703        let config = OAuth2Config::new(
2704            "client123",
2705            "secret456",
2706            "https://example.com/callback",
2707            "https://provider.com/authorize",
2708            "https://provider.com/token",
2709        )
2710        .with_pkce(true);
2711        let provider = GenericOAuth2Provider::new(config, transport.clone());
2712
2713        let pkce = OAuth2Config::generate_pkce_pair();
2714        let user = provider
2715            .user_from_token_with_pkce("auth_code", Some(&pkce.code_verifier))
2716            .expect("user_from_token_with_pkce 失败");
2717
2718        assert_eq!(user.access_token.as_deref(), Some("token123"));
2719
2720        // 验证 POST body 包含 code_verifier
2721        let (_url, body) = transport.last().expect("应有请求记录");
2722        assert!(
2723            body.contains("code_verifier"),
2724            "token 交换 body 应包含 code_verifier: {body}"
2725        );
2726        assert!(
2727            body.contains(&pkce.code_verifier),
2728            "token 交换 body 应包含 code_verifier 值"
2729        );
2730    }
2731
2732    /// 测试不带 PKCE 的 token 交换:POST body 不包含 code_verifier
2733    #[test]
2734    fn test_exchange_token_without_pkce_verifier() {
2735        let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2736        transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
2737
2738        let config = OAuth2Config::new(
2739            "client123",
2740            "secret456",
2741            "https://example.com/callback",
2742            "https://provider.com/authorize",
2743            "https://provider.com/token",
2744        );
2745        let provider = GenericOAuth2Provider::new(config, transport.clone());
2746
2747        let _user = provider
2748            .user_from_token("auth_code")
2749            .expect("user_from_token 失败");
2750
2751        let (_url, body) = transport.last().expect("应有请求记录");
2752        assert!(
2753            !body.contains("code_verifier"),
2754            "token 交换 body 不应包含 code_verifier(PKCE 未启用): {body}"
2755        );
2756    }
2757
2758    // ------------------------------------------------------------------------
2759    // T2: refresh_token + audit logger 测试
2760    // ------------------------------------------------------------------------
2761
2762    /// Mock 审计日志收集器(用于测试)
2763    #[derive(Default)]
2764    struct MockAuditLogger {
2765        events: Mutex<Vec<OAuth2AuditEvent>>,
2766    }
2767
2768    impl MockAuditLogger {
2769        fn events(&self) -> Vec<OAuth2AuditEvent> {
2770            self.events.lock().clone()
2771        }
2772    }
2773
2774    impl OAuth2AuditLogger for MockAuditLogger {
2775        fn log_event(&self, event: &OAuth2AuditEvent) {
2776            self.events.lock().push(event.clone());
2777        }
2778    }
2779
2780    /// 测试 refresh_token:mock transport 返回新 token → 成功
2781    #[test]
2782    fn test_refresh_token() {
2783        let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2784        transport.push_response(
2785            r#"{"access_token":"new_token","token_type":"Bearer","expires_in":7200,"scope":"read"}"#,
2786        );
2787
2788        let config = OAuth2Config::new(
2789            "client123",
2790            "secret456",
2791            "https://example.com/callback",
2792            "https://provider.com/authorize",
2793            "https://provider.com/token",
2794        );
2795        let provider = GenericOAuth2Provider::new(config, transport.clone());
2796
2797        let token_resp = provider
2798            .refresh_token("old_refresh_token")
2799            .expect("refresh_token 失败");
2800
2801        assert_eq!(token_resp.access_token, "new_token");
2802        assert_eq!(token_resp.token_type.as_deref(), Some("Bearer"));
2803        assert_eq!(token_resp.expires_in, Some(7200));
2804        assert_eq!(token_resp.scope.as_deref(), Some("read"));
2805    }
2806
2807    /// 测试审计日志:token 交换成功后记录事件
2808    #[test]
2809    fn test_audit_log_on_token_exchange() {
2810        let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2811        transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
2812
2813        let logger = Arc::new(MockAuditLogger::default());
2814        let config = OAuth2Config::new(
2815            "client123",
2816            "secret456",
2817            "https://example.com/callback",
2818            "https://provider.com/authorize",
2819            "https://provider.com/token",
2820        );
2821        let provider =
2822            GenericOAuth2Provider::new(config, transport).with_audit_logger(logger.clone());
2823
2824        let _user = provider
2825            .user_from_token("auth_code")
2826            .expect("user_from_token 失败");
2827
2828        let events = logger.events();
2829        assert!(
2830            events.iter().any(|e| e.grant_type == "authorization_code"
2831                && e.result == "success"
2832                && e.client_id == "client123"),
2833            "应记录 authorization_code success 事件: {events:?}"
2834        );
2835    }
2836
2837    /// 测试自动刷新:过期 token + 有效 refresh → 自动刷新
2838    #[test]
2839    fn test_auto_refresh_on_expired() {
2840        let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2841        // 第一次响应:token 已过期 (expires_in=0) + refresh_token
2842        transport.push_response(
2843            r#"{"access_token":"expired_token","refresh_token":"valid_refresh","expires_in":0}"#,
2844        );
2845        // 第二次响应:refresh_token 返回新 token
2846        transport.push_response(r#"{"access_token":"refreshed_token","expires_in":3600}"#);
2847
2848        let config = OAuth2Config::new(
2849            "client123",
2850            "secret456",
2851            "https://example.com/callback",
2852            "https://provider.com/authorize",
2853            "https://provider.com/token",
2854        );
2855        let provider = GenericOAuth2Provider::new(config, transport.clone());
2856
2857        let user = provider
2858            .user_from_token("auth_code")
2859            .expect("user_from_token 失败");
2860
2861        // 应自动刷新为 refreshed_token
2862        assert_eq!(
2863            user.access_token.as_deref(),
2864            Some("refreshed_token"),
2865            "过期 token 应自动刷新"
2866        );
2867        assert_eq!(user.expires_in, Some(3600));
2868        // 应有 2 次 HTTP 请求(token 交换 + refresh)
2869        assert_eq!(transport.count(), 2);
2870    }
2871
2872    /// 测试 refresh_token 为空 → AuthFailed
2873    #[test]
2874    fn test_refresh_token_empty() {
2875        let config = OAuth2Config::new(
2876            "client123",
2877            "secret456",
2878            "https://example.com/callback",
2879            "https://provider.com/authorize",
2880            "https://provider.com/token",
2881        );
2882        let provider =
2883            GenericOAuth2Provider::new(config, Arc::new(MemoryOAuth2HttpTransport::new()));
2884
2885        let err = provider.refresh_token("").unwrap_err();
2886        assert!(matches!(err, OAuth2Error::AuthFailed(msg) if msg.contains("refresh_token")));
2887    }
2888
2889    /// 测试 refresh_token Provider 返回非 JSON → TokenExchangeFailed
2890    #[test]
2891    fn test_refresh_token_invalid_json() {
2892        let transport = Arc::new(MemoryOAuth2HttpTransport::new());
2893        transport.push_response("not a json");
2894
2895        let config = OAuth2Config::new(
2896            "client123",
2897            "secret456",
2898            "https://example.com/callback",
2899            "https://provider.com/authorize",
2900            "https://provider.com/token",
2901        );
2902        let provider = GenericOAuth2Provider::new(config, transport);
2903
2904        let err = provider.refresh_token("valid_refresh").unwrap_err();
2905        assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
2906    }
2907
2908    /// 测试 TokenResponse 序列化脱敏
2909    #[test]
2910    fn test_token_response_serialize_redacted() {
2911        let token_resp = TokenResponse {
2912            access_token: "secret_access_token".into(),
2913            token_type: Some("Bearer".into()),
2914            expires_in: Some(3600),
2915            scope: Some("read".into()),
2916            refresh_token: Some("secret_refresh_token".into()),
2917        };
2918
2919        let json = serde_json::to_string(&token_resp).expect("序列化失败");
2920        assert!(
2921            !json.contains("secret_access_token"),
2922            "access_token 不应出现在序列化 JSON 中: {json}"
2923        );
2924        assert!(
2925            !json.contains("secret_refresh_token"),
2926            "refresh_token 不应出现在序列化 JSON 中: {json}"
2927        );
2928    }
2929
2930    // ------------------------------------------------------------------------
2931    // T3: ImplicitOAuth2Provider 测试
2932    // ------------------------------------------------------------------------
2933
2934    /// 测试 implicit redirect_url 含 response_type=token
2935    #[test]
2936    fn test_implicit_redirect_url() {
2937        let config = OAuth2Config::new(
2938            "client123",
2939            "secret456",
2940            "https://example.com/callback",
2941            "https://provider.com/oauth2.0/authorize",
2942            "https://provider.com/oauth2.0/token",
2943        )
2944        .with_scope("profile");
2945        let provider = ImplicitOAuth2Provider::new(config);
2946
2947        let url = provider.redirect_url("state_abc");
2948
2949        assert!(
2950            url.contains("response_type=token"),
2951            "URL 应含 response_type=token"
2952        );
2953        assert!(url.contains("client_id=client123"));
2954        assert!(url.contains("state=state_abc"));
2955        assert!(url.contains("scope=profile"));
2956    }
2957
2958    /// 测试 implicit parse_fragment:含 access_token + state → 解析成功
2959    #[test]
2960    fn test_implicit_parse_fragment() {
2961        let config = OAuth2Config::new(
2962            "client123",
2963            "secret456",
2964            "https://example.com/callback",
2965            "https://provider.com/authorize",
2966            "https://provider.com/token",
2967        );
2968        let provider = ImplicitOAuth2Provider::new(config);
2969
2970        let fragment = "access_token=token123&token_type=Bearer&expires_in=3600&state=mystate";
2971        let token_resp = provider
2972            .parse_fragment(fragment, "mystate")
2973            .expect("parse_fragment 失败");
2974
2975        assert_eq!(token_resp.access_token, "token123");
2976        assert_eq!(token_resp.token_type.as_deref(), Some("Bearer"));
2977        assert_eq!(token_resp.expires_in, Some(3600));
2978        // implicit 流程不返回 refresh_token
2979        assert!(
2980            token_resp.refresh_token.is_none(),
2981            "implicit 流程 refresh_token 应为 None"
2982        );
2983    }
2984
2985    /// 测试 implicit state 不匹配 → AuthFailed
2986    #[test]
2987    fn test_implicit_state_mismatch() {
2988        let config = OAuth2Config::new(
2989            "client123",
2990            "secret456",
2991            "https://example.com/callback",
2992            "https://provider.com/authorize",
2993            "https://provider.com/token",
2994        );
2995        let provider = ImplicitOAuth2Provider::new(config);
2996
2997        let fragment = "access_token=token123&state=wrong_state";
2998        let err = provider
2999            .parse_fragment(fragment, "expected_state")
3000            .unwrap_err();
3001        assert!(
3002            matches!(&err, OAuth2Error::AuthFailed(msg) if msg.contains("CSRF state mismatch")),
3003            "state 不匹配应返回 CSRF 错误: {err}"
3004        );
3005    }
3006
3007    /// 测试 implicit 回调不含 refresh_token
3008    #[test]
3009    fn test_implicit_no_refresh_token() {
3010        let config = OAuth2Config::new(
3011            "client123",
3012            "secret456",
3013            "https://example.com/callback",
3014            "https://provider.com/authorize",
3015            "https://provider.com/token",
3016        );
3017        let provider = ImplicitOAuth2Provider::new(config);
3018
3019        // fragment 中即使包含 refresh_token 也应被忽略
3020        let fragment = "access_token=token123&refresh_token=should_be_ignored&state=mystate";
3021        let token_resp = provider
3022            .parse_fragment(fragment, "mystate")
3023            .expect("parse_fragment 失败");
3024
3025        assert!(
3026            token_resp.refresh_token.is_none(),
3027            "implicit 流程 refresh_token 应固定为 None"
3028        );
3029    }
3030
3031    /// 测试 implicit fragment 为空 → TokenExchangeFailed
3032    #[test]
3033    fn test_implicit_empty_fragment() {
3034        let config = OAuth2Config::new(
3035            "client123",
3036            "secret456",
3037            "https://example.com/callback",
3038            "https://provider.com/authorize",
3039            "https://provider.com/token",
3040        );
3041        let provider = ImplicitOAuth2Provider::new(config);
3042
3043        let err = provider.parse_fragment("", "state").unwrap_err();
3044        assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
3045    }
3046
3047    /// 测试 implicit fragment 无 access_token → TokenExchangeFailed
3048    #[test]
3049    fn test_implicit_no_access_token() {
3050        let config = OAuth2Config::new(
3051            "client123",
3052            "secret456",
3053            "https://example.com/callback",
3054            "https://provider.com/authorize",
3055            "https://provider.com/token",
3056        );
3057        let provider = ImplicitOAuth2Provider::new(config);
3058
3059        let fragment = "token_type=Bearer&state=mystate";
3060        let err = provider.parse_fragment(fragment, "mystate").unwrap_err();
3061        assert!(matches!(err, OAuth2Error::TokenExchangeFailed(_)));
3062    }
3063
3064    /// 测试 implicit scope 列表为空 → URL 不含 scope 参数
3065    #[test]
3066    fn test_implicit_empty_scopes() {
3067        let config = OAuth2Config::new(
3068            "client123",
3069            "secret456",
3070            "https://example.com/callback",
3071            "https://provider.com/authorize",
3072            "https://provider.com/token",
3073        );
3074        let provider = ImplicitOAuth2Provider::new(config);
3075
3076        let url = provider.redirect_url("state123");
3077        assert!(
3078            !url.contains("scope="),
3079            "空 scopes 时 URL 不应含 scope 参数"
3080        );
3081    }
3082
3083    /// 测试 implicit 审计日志记录 OAUTH2_IMPLICIT_TOKEN_EXPOSED 告警
3084    #[test]
3085    fn test_implicit_audit_log_token_exposed() {
3086        let config = OAuth2Config::new(
3087            "client123",
3088            "secret456",
3089            "https://example.com/callback",
3090            "https://provider.com/authorize",
3091            "https://provider.com/token",
3092        );
3093        let logger = Arc::new(MockAuditLogger::default());
3094        let provider = ImplicitOAuth2Provider::new(config).with_audit_logger(logger.clone());
3095
3096        let fragment = "access_token=token123&state=mystate";
3097        let _ = provider.parse_fragment(fragment, "mystate");
3098
3099        let events = logger.events();
3100        assert!(
3101            events
3102                .iter()
3103                .any(|e| e.alert_code.as_deref() == Some("OAUTH2_IMPLICIT_TOKEN_EXPOSED")),
3104            "应记录 OAUTH2_IMPLICIT_TOKEN_EXPOSED 告警: {events:?}"
3105        );
3106    }
3107
3108    // ------------------------------------------------------------------------
3109    // T5: Token Store 集成测试(需 redis-store feature)
3110    // ------------------------------------------------------------------------
3111
3112    /// 测试 token store 集成:user_from_token 后 store 收到 token
3113    #[cfg(feature = "redis-store")]
3114    #[tokio::test]
3115    async fn test_token_store_integration() {
3116        use crate::oauth_store::{MemoryOAuth2TokenStore, OAuth2TokenStore};
3117
3118        let transport = Arc::new(MemoryOAuth2HttpTransport::new());
3119        transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
3120
3121        let store = Arc::new(MemoryOAuth2TokenStore::new());
3122        let config = OAuth2Config::new(
3123            "client123",
3124            "secret456",
3125            "https://example.com/callback",
3126            "https://provider.com/authorize",
3127            "https://provider.com/token",
3128        );
3129        let provider =
3130            GenericOAuth2Provider::new(config, transport).with_token_store(store.clone());
3131
3132        let user = provider
3133            .user_from_token("auth_code")
3134            .expect("user_from_token 失败");
3135        assert_eq!(user.access_token.as_deref(), Some("token123"));
3136
3137        // 等待 spawn 的存储任务完成
3138        tokio::time::sleep(std::time::Duration::from_millis(100)).await;
3139
3140        let stored = store
3141            .get_token("client123")
3142            .await
3143            .expect("get_token 失败")
3144            .expect("应查到存储的 token");
3145        assert_eq!(stored.access_token, "token123");
3146    }
3147
3148    /// 测试 token store 失败 best-effort:存储失败不影响 token 发放
3149    #[cfg(feature = "redis-store")]
3150    #[tokio::test]
3151    async fn test_token_store_failure_best_effort() {
3152        use crate::oauth_store::MemoryOAuth2TokenStore;
3153
3154        let transport = Arc::new(MemoryOAuth2HttpTransport::new());
3155        transport.push_response(r#"{"access_token":"token123","expires_in":3600}"#);
3156
3157        // 使用一个总是成功的 MemoryStore(模拟 best-effort)
3158        let store = Arc::new(MemoryOAuth2TokenStore::new());
3159        let config = OAuth2Config::new(
3160            "client123",
3161            "secret456",
3162            "https://example.com/callback",
3163            "https://provider.com/authorize",
3164            "https://provider.com/token",
3165        );
3166        let provider =
3167            GenericOAuth2Provider::new(config, transport).with_token_store(store.clone());
3168
3169        // 即使 store 出错,user_from_token 也应成功
3170        let user = provider
3171            .user_from_token("auth_code")
3172            .expect("user_from_token 应成功(best-effort)");
3173        assert_eq!(user.access_token.as_deref(), Some("token123"));
3174    }
3175
3176    // ------------------------------------------------------------------------
3177    // T1: proptest 属性测试
3178    // ------------------------------------------------------------------------
3179
3180    // proptest: state 不可预测(长度固定 32、全 hex 字符)
3181    proptest::proptest! {
3182        #[test]
3183        fn proptest_state_unpredictable(_n in 0u32..1000) {
3184            let s1 = OAuth2Config::generate_state();
3185            let s2 = OAuth2Config::generate_state();
3186            prop_assert_eq!(s1.len(), 32);
3187            prop_assert_eq!(s2.len(), 32);
3188            prop_assert!(s1.chars().all(|c| c.is_ascii_hexdigit()));
3189            prop_assert!(s2.chars().all(|c| c.is_ascii_hexdigit()));
3190        }
3191    }
3192
3193    // proptest: PKCE code_verifier 长度 ∈ [43, 128]
3194    proptest::proptest! {
3195        #[test]
3196        fn proptest_pkce_verifier_length(_n in 0u32..1000) {
3197            let pkce = OAuth2Config::generate_pkce_pair();
3198            prop_assert!(pkce.code_verifier.len() >= 43);
3199            prop_assert!(pkce.code_verifier.len() <= 128);
3200        }
3201    }
3202
3203    // proptest: PKCE code_challenge == base64url(SHA256(code_verifier))
3204    proptest::proptest! {
3205        #[test]
3206        fn proptest_pkce_challenge_matches_verifier(_n in 0u32..1000) {
3207            let pkce = OAuth2Config::generate_pkce_pair();
3208            let mut hasher = sha2::Sha256::new();
3209            sha2::Digest::update(&mut hasher, pkce.code_verifier.as_bytes());
3210            let digest = sha2::Digest::finalize(hasher);
3211            let expected = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(digest);
3212            prop_assert_eq!(pkce.code_challenge, expected);
3213        }
3214    }
3215}