diff --git a/Cargo.lock b/Cargo.lock index a5dc362..1f46d93 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2,6 +2,41 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "aead" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0" +dependencies = [ + "crypto-common", + "generic-array", +] + +[[package]] +name = "aes" +version = "0.8.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b169f7a6d4742236a0a00c541b845991d0ac43e546831af1249753ab4c3aa3a0" +dependencies = [ + "cfg-if", + "cipher", + "cpufeatures 0.2.17", +] + +[[package]] +name = "aes-gcm" +version = "0.10.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "831010a0f742e1209b3bcea8fab6a8e149051ba6099432c8cb2cc117dec3ead1" +dependencies = [ + "aead", + "aes", + "cipher", + "ctr", + "ghash", + "subtle", +] + [[package]] name = "ahash" version = "0.8.12" @@ -25,14 +60,16 @@ dependencies = [ [[package]] name = "aitokenpool" -version = "0.2.1" +version = "0.2.2" dependencies = [ + "aes-gcm", "anyhow", "argon2", "axum", "chrono", "clap", "env_logger", + "futures-util", "hex", "hyper", "log", @@ -318,6 +355,16 @@ dependencies = [ "windows-link", ] +[[package]] +name = "cipher" +version = "0.4.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" +dependencies = [ + "crypto-common", + "inout", +] + [[package]] name = "clap" version = "4.6.6" @@ -415,9 +462,19 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" dependencies = [ "generic-array", + "rand_core 0.6.4", "typenum", ] +[[package]] +name = "ctr" +version = "0.9.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0369ee1ad671834580515889b80f2ea915f23b8be8d0daa4bbaf2ac5c7590835" +dependencies = [ + "cipher", +] + [[package]] name = "defmt" version = "1.1.1" @@ -670,6 +727,16 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "ghash" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0d8a4362ccb29cb0b265253fb0a2728f592895ee6854fd9bc13f2ffda266ff1" +dependencies = [ + "opaque-debug", + "polyval", +] + [[package]] name = "h2" version = "0.4.15" @@ -992,6 +1059,15 @@ dependencies = [ "hashbrown 0.17.1", ] +[[package]] +name = "inout" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "879f10e63c20629ecabbb64a8010319738c66a5cd0c29b02d63d272b03751d01" +dependencies = [ + "generic-array", +] + [[package]] name = "ipnet" version = "2.12.1" @@ -1175,6 +1251,12 @@ version = "1.70.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" +[[package]] +name = "opaque-debug" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" + [[package]] name = "openssl" version = "0.10.81" @@ -1247,6 +1329,18 @@ version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" +[[package]] +name = "polyval" +version = "0.6.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d1fe60d06143b2430aa532c94cfe9e29783047f06c0d7fd359a9a51b729fa25" +dependencies = [ + "cfg-if", + "cpufeatures 0.2.17", + "opaque-debug", + "universal-hash", +] + [[package]] name = "portable-atomic" version = "1.15.0" @@ -2098,6 +2192,16 @@ version = "1.0.24" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" +[[package]] +name = "universal-hash" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea" +dependencies = [ + "crypto-common", + "subtle", +] + [[package]] name = "untrusted" version = "0.9.0" diff --git a/Cargo.toml b/Cargo.toml index 0602c2d..a599877 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "aitokenpool" -version = "0.2.1" +version = "0.2.2" edition = "2021" description = "AI Token 共享池 — 企业 key 池 + 公共共享市场" license = "MIT" @@ -18,6 +18,12 @@ hyper = { version = "1.0", features = ["full"] } # HTTP 客户端(上游转发) reqwest = { version = "0.12", features = ["rustls-tls", "json", "stream"] } +# SSE 流式(bytes_stream + Body::from_stream) +futures-util = "0.3" + +# 上游 key 加密(AES-256-GCM,RustCrypto) +aes-gcm = "0.10" + # 序列化 serde = { version = "1.0", features = ["derive"] } serde_json = { version = "1.0" } diff --git a/config/config.example.toml b/config/config.example.toml index e678f52..7726421 100644 --- a/config/config.example.toml +++ b/config/config.example.toml @@ -12,6 +12,9 @@ [server] addr = "0.0.0.0:8080" db_path = "data/aitokenpool.db" +# 上游 key 主密钥(P0-C 起):hex 32 字节;env ATP_MASTER_KEY 优先级更高。 +# 生产必须显式配置,否则使用随机 dev 密钥(重启后旧密文不可解)。 +# master_key = "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef" # ============================================================ # 1. 点数规则(账本层的锚) diff --git a/src/config.rs b/src/config.rs index a687037..c16fb94 100644 --- a/src/config.rs +++ b/src/config.rs @@ -15,13 +15,16 @@ fn default_db_path() -> String { "data/aitokenpool.db".to_string() } -/// 服务(监听 / 数据库路径)——config.example.toml 可缺省,走默认值 +/// 服务(监听 / 数据库路径 / 主密钥)——config.example.toml 可缺省,走默认值 #[derive(Debug, Clone, Deserialize)] pub struct Server { #[serde(default = "default_addr")] pub addr: String, #[serde(default = "default_db_path")] pub db_path: String, + /// 上游 key 主密钥(hex 32 字节;P0-C 起生效;env ATP_MASTER_KEY 优先级更高) + #[serde(default)] + pub master_key: String, } impl Default for Server { @@ -29,6 +32,7 @@ impl Default for Server { Server { addr: default_addr(), db_path: default_db_path(), + master_key: String::new(), } } } diff --git a/src/crypto.rs b/src/crypto.rs new file mode 100644 index 0000000..54890c2 --- /dev/null +++ b/src/crypto.rs @@ -0,0 +1,162 @@ +//! 上游 key 加密存储(AES-256-GCM,RustCrypto) +//! +//! P0-C(rant 2026-08-18T10:36:04): +//! - keys.encrypted_key 存 `v1::`(含 12 字节随机 nonce) +//! - 主密钥来源(优先级):env `ATP_MASTER_KEY`(hex 32 字节)→ config `[server] master_key` +//! → dev 模式随机生成并打印警告(生产必须显式配置) +//! - 写路径加密、读路径解密;旧明文占位 key 启动时自动迁移(见 db::migrate_key_encryption) + +use aes_gcm::aead::rand_core::RngCore; +use aes_gcm::aead::{Aead, KeyInit, OsRng}; +use aes_gcm::{Aes256Gcm, Nonce}; +use anyhow::{anyhow, Result}; + +/// 密文格式前缀:`v1:`(旧明文无此前缀 → 迁移/拒绝解密) +pub const PREFIX: &str = "v1:"; +/// nonce 长度(AES-GCM 推荐 12 字节) +const NONCE_LEN: usize = 12; + +/// AES-256-GCM 加密器(内部持有 32 字节主密钥) +#[derive(Clone)] +pub struct Crypto { + key: [u8; 32], +} + +impl Crypto { + /// 直接以 32 字节主密钥构造 + pub fn new(key: [u8; 32]) -> Self { + Self { key } + } + + /// 主密钥解析:env ATP_MASTER_KEY → config master_key(均要求 hex 32 字节) + /// 两者皆空 → dev 模式:随机密钥 + 警告(进程内有效,重启后旧密文不可解) + pub fn from_config(config_master_key: &str) -> Self { + if let Ok(env_key) = std::env::var("ATP_MASTER_KEY") { + match parse_master_key(&env_key) { + Ok(k) => { + log::info!("使用 ATP_MASTER_KEY 主密钥(env)"); + return Self::new(k); + } + Err(e) => { + log::error!("ATP_MASTER_KEY 无效(需 32 字节 hex): {e},回退下一来源"); + } + } + } + if !config_master_key.is_empty() { + match parse_master_key(config_master_key) { + Ok(k) => { + log::info!("使用 config [server].master_key 主密钥"); + return Self::new(k); + } + Err(e) => { + log::error!("config master_key 无效(需 32 字节 hex): {e},回退 dev 模式"); + } + } + } + let mut key = [0u8; 32]; + OsRng.fill_bytes(&mut key); + log::warn!( + "未配置 ATP_MASTER_KEY / [server].master_key —— 使用随机 dev 主密钥(重启后旧密文不可解,生产必须显式配置)" + ); + Self::new(key) + } + + /// 加密:`v1::` + pub fn encrypt(&self, plain: &[u8]) -> Result { + let cipher = + Aes256Gcm::new_from_slice(&self.key).map_err(|_| anyhow!("AES-256-GCM 初始化失败"))?; + let mut nonce = [0u8; NONCE_LEN]; + OsRng.fill_bytes(&mut nonce); + let ct = cipher + .encrypt(Nonce::from_slice(&nonce), plain) + .map_err(|_| anyhow!("加密失败"))?; + Ok(format!( + "{PREFIX}{}:{}", + hex::encode(nonce), + hex::encode(ct) + )) + } + + /// 解密:解析 `v1:` 前缀 + nonce + 密文;非 v1 格式(明文/旧格式)→ 报错 + pub fn decrypt(&self, stored: &str) -> Result> { + let rest = stored + .strip_prefix(PREFIX) + .ok_or_else(|| anyhow!("非 v1: 密文格式(明文占位?)"))?; + let (nonce_hex, ct_hex) = rest + .split_once(':') + .ok_or_else(|| anyhow!("密文格式错误:缺 nonce/密文分隔"))?; + let nonce = hex::decode(nonce_hex)?; + let ct = hex::decode(ct_hex)?; + if nonce.len() != NONCE_LEN { + return Err(anyhow!("nonce 长度异常({})", nonce.len())); + } + let cipher = + Aes256Gcm::new_from_slice(&self.key).map_err(|_| anyhow!("AES-256-GCM 初始化失败"))?; + cipher + .decrypt(Nonce::from_slice(&nonce), ct.as_ref()) + .map_err(|_| anyhow!("解密失败(主密钥不匹配或密文被篡改)")) + } +} + +/// 解析 hex 主密钥(必须 64 hex 字符 = 32 字节) +fn parse_master_key(hex_str: &str) -> Result<[u8; 32]> { + let bytes = hex::decode(hex_str.trim())?; + let arr: [u8; 32] = bytes + .try_into() + .map_err(|_| anyhow!("主密钥必须是 32 字节(64 个 hex 字符)"))?; + Ok(arr) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn test_crypto() -> Crypto { + Crypto::new([7u8; 32]) + } + + #[test] + fn encrypt_decrypt_roundtrip() { + let c = test_crypto(); + let stored = c.encrypt(b"sk-abc123secret").unwrap(); + assert!(stored.starts_with(PREFIX), "格式 v1: 前缀"); + assert_eq!(c.decrypt(&stored).unwrap(), b"sk-abc123secret"); + } + + #[test] + fn same_plaintext_different_nonce() { + let c = test_crypto(); + let a = c.encrypt(b"sk-x").unwrap(); + let b = c.encrypt(b"sk-x").unwrap(); + assert_ne!(a, b, "随机 nonce 保证两次密文不同"); + } + + #[test] + fn wrong_master_key_fails_to_decrypt() { + let c1 = test_crypto(); + let c2 = Crypto::new([8u8; 32]); + let stored = c1.encrypt(b"sk-abc").unwrap(); + assert!(c2.decrypt(&stored).is_err(), "错误主密钥必须解密失败"); + } + + #[test] + fn plaintext_old_format_rejected() { + let c = test_crypto(); + assert!( + c.decrypt("sk-placeholder-encrypted").is_err(), + "非 v1: 前缀(明文占位)应报错 → 触发迁移检测" + ); + assert!(c.decrypt("v1:abcd").is_err(), "截断密文报错"); + } + + #[test] + fn master_key_hex_validation() { + let good = "a".repeat(64); + assert!(parse_master_key(&good).is_ok()); + assert!(parse_master_key("not-hex").is_err()); + assert!( + parse_master_key(&"a".repeat(62)).is_err(), + "长度不足 32 字节" + ); + } +} diff --git a/src/db.rs b/src/db.rs index 1388664..d2c8208 100644 --- a/src/db.rs +++ b/src/db.rs @@ -10,7 +10,7 @@ use anyhow::{Context, Result}; use rusqlite::Connection; -pub const SCHEMA_VERSION: i64 = 1; +pub const SCHEMA_VERSION: i64 = 2; /// 打开(或创建)数据库并执行幂等迁移 + dev 种子 pub fn open(path: &str) -> Result { @@ -51,6 +51,10 @@ pub fn migrate(conn: &Connection) -> Result<()> { encrypted_key TEXT NOT NULL, quota REAL NOT NULL DEFAULT 0, used REAL NOT NULL DEFAULT 0, + available_days TEXT NOT NULL DEFAULT '', + available_start TEXT NOT NULL DEFAULT '', + available_end TEXT NOT NULL DEFAULT '', + note TEXT NOT NULL DEFAULT '', created_at TEXT NOT NULL DEFAULT (datetime('now')) ); CREATE TABLE IF NOT EXISTS api_keys ( @@ -103,6 +107,26 @@ pub fn migrate(conn: &Connection) -> Result<()> { ON models(provider, model); "#, )?; + // v2:为旧库补 available_* 列(新建库已在建表语句里) + ensure_column( + conn, + "keys", + "available_days", + "available_days TEXT NOT NULL DEFAULT ''", + )?; + ensure_column( + conn, + "keys", + "available_start", + "available_start TEXT NOT NULL DEFAULT ''", + )?; + ensure_column( + conn, + "keys", + "available_end", + "available_end TEXT NOT NULL DEFAULT ''", + )?; + ensure_column(conn, "keys", "note", "note TEXT NOT NULL DEFAULT ''")?; // schema_version:INSERT OR REPLACE 保证幂等 let v: i64 = conn .query_row("SELECT version FROM schema_version", [], |r| r.get(0)) @@ -116,6 +140,45 @@ pub fn migrate(conn: &Connection) -> Result<()> { Ok(()) } +/// 幂等补列:列不存在才 ALTER TABLE ADD COLUMN +fn ensure_column(conn: &Connection, table: &str, column: &str, ddl: &str) -> Result<()> { + let sql = format!("SELECT 1 FROM pragma_table_info('{table}') WHERE name = ?1"); + let exists: bool = conn + .prepare(&sql)? + .exists([column]) + .with_context(|| format!("检查列 {table}.{column} 失败"))?; + if !exists { + conn.execute(&format!("ALTER TABLE {table} ADD COLUMN {ddl}"), []) + .with_context(|| format!("为 {table} 添加列 {column} 失败"))?; + } + Ok(()) +} + +/// 上游 key 加密迁移:旧明文占位(非 v1: 前缀)→ 启动时自动加密 +/// 返回迁移条数;已加密 / 已迁移的 key 原样保留 +pub fn migrate_key_encryption(conn: &Connection, crypto: &crate::crypto::Crypto) -> Result { + let mut stmt = conn.prepare("SELECT id, encrypted_key FROM keys")?; + let rows: Vec<(i64, String)> = stmt + .query_map([], |r| Ok((r.get(0)?, r.get(1)?)))? + .collect::>>()?; + let mut n = 0usize; + for (id, stored) in rows { + if stored.starts_with(crate::crypto::PREFIX) { + continue; // 已是 v1 密文 + } + let cipher = crypto.encrypt(stored.as_bytes())?; + conn.execute( + "UPDATE keys SET encrypted_key = ?1 WHERE id = ?2", + rusqlite::params![cipher, id], + )?; + n += 1; + } + if n > 0 { + log::info!("密钥加密迁移完成:{n} 条明文 key 已加密"); + } + Ok(n) +} + /// dev 种子:demo 用户(demo@aitokenpool.local / demo1234,argon2)+ 点数账户 + 示例上游 key pub fn seed(conn: &Connection) -> Result<()> { use crate::auth::hash_password; @@ -291,4 +354,80 @@ mod tests { drop(conn); let _ = std::fs::remove_file(p); } + + #[test] + fn v2_migration_adds_available_and_note_columns() { + // 模拟 v1 旧库:建表后不含 available_* / note 列 → migrate 幂等补列 + let p = std::env::temp_dir().join(format!("atp_v1_{}_{}.db", std::process::id(), "v1")); + let _ = std::fs::remove_file(&p); + let conn = Connection::open(&p).unwrap(); + conn.execute_batch( + "CREATE TABLE keys ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + provider TEXT NOT NULL, + plan TEXT NOT NULL DEFAULT '', + model TEXT NOT NULL DEFAULT '', + status TEXT NOT NULL DEFAULT 'on', + owner_id INTEGER NOT NULL, + encrypted_key TEXT NOT NULL, + quota REAL NOT NULL DEFAULT 0, + used REAL NOT NULL DEFAULT 0, + created_at TEXT NOT NULL DEFAULT (datetime('now')) + ); + INSERT INTO keys (provider, plan, model, status, owner_id, encrypted_key) VALUES ('t','p','m','on',1,'sk-plain');", + ) + .unwrap(); + migrate(&conn).unwrap(); + // 补列成功且旧数据保留 + let row: (String, String, String, String) = conn + .query_row( + "SELECT available_days, available_start, available_end, note FROM keys WHERE id = 1", + [], + |r| Ok((r.get(0)?, r.get(1)?, r.get(2)?, r.get(3)?)), + ) + .unwrap(); + assert_eq!(row.0, ""); + assert_eq!(row.3, ""); + let v: i64 = conn + .query_row("SELECT version FROM schema_version", [], |r| r.get(0)) + .unwrap(); + assert_eq!(v, 2, "旧库迁移后版本应为 2"); + drop(conn); + let _ = std::fs::remove_file(p); + } + + #[test] + fn key_encryption_migration_encrypts_plaintext() { + let crypto = crate::crypto::Crypto::new([11u8; 32]); + let (conn, p) = tmp_db("kenc"); + // 手工插入明文占位 key(模拟 v1 遗留) + conn.execute( + "INSERT INTO keys (provider, plan, model, status, owner_id, encrypted_key) VALUES ('t','p','m','on',1,'sk-placeholder-encrypted')", + [], + ) + .unwrap(); + let n = migrate_key_encryption(&conn, &crypto).expect("迁移成功"); + assert!(n >= 1, "至少迁移一条"); + let stored: String = conn + .query_row( + "SELECT encrypted_key FROM keys WHERE provider = 't'", + [], + |r| r.get(0), + ) + .unwrap(); + assert!( + stored.starts_with(crate::crypto::PREFIX), + "已加密: {stored}" + ); + assert_eq!( + crypto.decrypt(&stored).unwrap(), + b"sk-placeholder-encrypted", + "迁移后可解密还原原文" + ); + // 幂等:二次迁移不再变化 + let n2 = migrate_key_encryption(&conn, &crypto).expect("二次迁移"); + assert_eq!(n2, 0, "已加密的 key 不再重复迁移"); + drop(conn); + let _ = std::fs::remove_file(p); + } } diff --git a/src/gateway.rs b/src/gateway.rs index ffb559c..369b014 100644 --- a/src/gateway.rs +++ b/src/gateway.rs @@ -10,10 +10,12 @@ //! anthropic → {base}/v1/messages)→ 上游响应原样透传 → 解析 usage → 计量入账。 //! 上游 key 当前为明文占位(加密留 P0-C)。 +use axum::body::Body; use axum::extract::State; use axum::http::StatusCode; use axum::response::Response; use axum::Json; +use futures_util::StreamExt; use crate::billing; use crate::config::Config; @@ -84,10 +86,80 @@ fn passthrough(status: StatusCode, body: Vec) -> Response { Response::builder() .status(status) .header("content-type", "application/json") - .body(axum::body::Body::from(body)) + .body(Body::from(body)) .expect("build passthrough response") } +/// 解密上游 key(P0-C:keys.encrypted_key 为 v1: 密文,转发前解密) +fn decrypt_key(st: &AppState, key: &dao::KeyRow) -> Option { + match st.crypto.decrypt(&key.encrypted_key) { + Ok(bytes) => Some(String::from_utf8_lossy(&bytes).into_owned()), + Err(e) => { + log::error!("解密上游 key 失败 key_id={}: {e}", key.id); + None + } + } +} + +/// 计量入账(非流式/流式共用):tokens>0 才入账;失败仅记日志不影响透传 +fn settle_usage( + st: &AppState, + auth: AuthUser, + key: &dao::KeyRow, + model: &str, + input_tokens: f64, + output_tokens: f64, +) { + let tokens = input_tokens + output_tokens; + if tokens <= 0.0 { + return; + } + // 锁作用域严格限定在同步区内(绝不在 await 期间持有 MutexGuard) + let mut conn = match st.db.lock() { + Ok(c) => c, + Err(e) => { + log::error!("计量入账失败(db lock poisoned): {e}"); + return; + } + }; + let price = dao::get_model_price(&conn, &key.provider, model); + let (pts, cost) = match price { + Some((i_per_m, o_per_m, currency)) => { + let pts = billing::calc_points( + input_tokens, + output_tokens, + i_per_m, + o_per_m, + st.cfg.points.points_per_unit, + ¤cy, + &st.cfg.points.anchor_currency, + ); + let cost = billing::to_anchor( + billing::raw_cost(input_tokens, output_tokens, i_per_m, o_per_m), + ¤cy, + &st.cfg.points.anchor_currency, + ); + (pts, cost) + } + None => (0.0, 0.0), + }; + let params = billing::SettleParams { + consumer_id: auth.user_id, + api_key_id: Some(auth.api_key_id), + key_id: key.id, + owner_id: key.owner_id, + model: model.to_string(), + tokens, + pts, + cost, + }; + if let Err(e) = billing::settle(&mut conn, ¶ms) { + log::error!("计量入账失败 key_id={} user={}: {e}", key.id, auth.user_id); + } + drop(conn); + st.router.mark_sticky(auth.user_id, model, key.id); +} + /// 核心转发逻辑(openai_chat / anthropic 共用) async fn forward( st: &AppState, @@ -133,12 +205,17 @@ async fn forward( st.router.mark_unhealthy(key_id); continue; }; + // 解密上游 key;解密失败 → 视为该 key 不可用 + let Some(plain_key) = decrypt_key(st, key) else { + st.router.mark_unhealthy(key_id); + continue; + }; // 转发上游(anthropic 用 x-api-key,openai 用 Bearer) let resp = if protocol == "anthropic" { st.http .post(&url) - .header("x-api-key", &key.encrypted_key) + .header("x-api-key", &plain_key) .header("content-type", "application/json") .header("anthropic-version", "2023-06-01") .body(body.clone()) @@ -147,7 +224,7 @@ async fn forward( } else { st.http .post(&url) - .header("authorization", format!("Bearer {}", key.encrypted_key)) + .header("authorization", format!("Bearer {plain_key}")) .header("content-type", "application/json") .body(body.clone()) .send() @@ -168,47 +245,7 @@ async fn forward( if status.is_success() { // 成功:解析 usage → 计量入账 → 粘性 let (input_tokens, output_tokens) = parse_usage(&bytes, protocol); - let tokens = input_tokens + output_tokens; - if tokens > 0.0 { - let mut conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; - let price = dao::get_model_price(&conn, &key.provider, model); - let (pts, cost) = match price { - Some((i_per_m, o_per_m, currency)) => { - let pts = billing::calc_points( - input_tokens, - output_tokens, - i_per_m, - o_per_m, - st.cfg.points.points_per_unit, - ¤cy, - &st.cfg.points.anchor_currency, - ); - let cost = billing::to_anchor( - billing::raw_cost(input_tokens, output_tokens, i_per_m, o_per_m), - ¤cy, - &st.cfg.points.anchor_currency, - ); - (pts, cost) - } - None => (0.0, 0.0), - }; - let params = billing::SettleParams { - consumer_id: auth.user_id, - api_key_id: Some(auth.api_key_id), - key_id: key.id, - owner_id: key.owner_id, - model: model.to_string(), - tokens, - pts, - cost, - }; - if let Err(e) = billing::settle(&mut conn, ¶ms) { - // 入账失败不影响透传(响应已成功),仅记录日志 - log::error!("计量入账失败 key_id={} user={}: {e}", key.id, auth.user_id); - } - drop(conn); - st.router.mark_sticky(auth.user_id, model, key_id); - } + settle_usage(st, auth, key, model, input_tokens, output_tokens); return Ok(passthrough(status, bytes)); } else if status == StatusCode::UNAUTHORIZED || status == StatusCode::FORBIDDEN @@ -229,6 +266,229 @@ async fn forward( )) } +/// SSE 流式转发(P0-C):请求体带 stream:true 时走此分支。 +/// +/// 流程:余额预检 → 路由选 key(初始连接失败可故障转移,最高 3 次)→ 拿到 200 后 +/// 逐块透传上游响应体(data: 行原样转发,保持 event/comment 原始格式)→ +/// 流尾解析 usage(openai 最后 chunk 的 usage / anthropic message_delta 的 +/// output_tokens + message_start 的 input_tokens)→ 复用 settle_usage 入账 → +/// 记粘性。客户端提前断开 → 响应体被 drop → 上游连接自动中止,不入账。 +async fn forward_stream( + st: &AppState, + auth: AuthUser, + model: &str, + body: String, + protocol: &str, +) -> Result { + // 余额预检(与 forward 一致,锁作用域严格块内) + let (balance, keys) = { + let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; + let balance = dao::get_balance(&conn, auth.user_id); + let keys = dao::find_keys_by_model(&conn, model).map_err(internal)?; + (balance, keys) + }; + if balance <= 0.0 { + return Err(err_json(StatusCode::PAYMENT_REQUIRED, "点数余额不足")); + } + if keys.is_empty() { + return Err(err_json( + StatusCode::SERVICE_UNAVAILABLE, + "该模型暂无可用 key", + )); + } + + // 选 key 并建立上游连接(此阶段失败可切换;连接成功后不再切换) + for _ in 0..crate::router::MAX_SWITCHES { + let Some(key_id) = st.router.pick(&keys, auth.user_id, model) else { + return Err(err_json( + StatusCode::SERVICE_UNAVAILABLE, + "该模型暂无可用 key", + )); + }; + let key = keys + .iter() + .find(|k| k.id == key_id) + .expect("pick 返回的 key 必然在候选集"); + + let Some(url) = resolve_endpoint(&st.cfg, &key.plan, protocol) else { + st.router.mark_unhealthy(key_id); + continue; + }; + let Some(plain_key) = decrypt_key(st, key) else { + st.router.mark_unhealthy(key_id); + continue; + }; + + let resp = if protocol == "anthropic" { + st.http + .post(&url) + .header("x-api-key", &plain_key) + .header("content-type", "application/json") + .header("anthropic-version", "2023-06-01") + .body(body.clone()) + .send() + .await + } else { + st.http + .post(&url) + .header("authorization", format!("Bearer {plain_key}")) + .header("content-type", "application/json") + .body(body.clone()) + .send() + .await + }; + + let resp = match resp { + Ok(r) => r, + Err(_) => { + st.router.mark_unhealthy(key_id); + continue; + } + }; + let status = resp.status(); + if !status.is_success() { + // 连接阶段失败 → 与 P0-B 相同的故障转移判定 + if status == StatusCode::UNAUTHORIZED + || status == StatusCode::FORBIDDEN + || status == StatusCode::TOO_MANY_REQUESTS + || status.is_server_error() + { + st.router.mark_unhealthy(key_id); + continue; + } + // 其它 4xx → 用户请求错误,读 body 透传 + let bytes = resp.bytes().await.unwrap_or_default().to_vec(); + return Ok(passthrough(status, bytes)); + } + + // 连接成功:构建 SSE 透传流(usage 捕获 + 流尾入账) + let key = key.clone(); + let st = st.clone(); + let model = model.to_string(); + let protocol = protocol.to_string(); + let capture = std::sync::Arc::new(std::sync::Mutex::new(UsageCapture::new(&protocol))); + let cap_fwd = std::sync::Arc::clone(&capture); + let fwd = resp.bytes_stream().map(move |item| { + if let Ok(bytes) = &item { + if let Ok(mut cap) = cap_fwd.lock() { + cap.push(bytes); + } + } + item.map_err(|e| -> Box { e.into() }) + }); + // 流尾:解析 usage → 入账(客户端未完整接收时该 future 不会执行 → 不入账) + let finalize = futures_util::stream::once(async move { + let (input, output) = { + let mut cap = capture.lock().expect("usage capture lock"); + cap.finish() + }; + settle_usage(&st, auth, &key, &model, input, output); + Ok::>( + axum::body::Bytes::new(), + ) + }); + let body = Body::from_stream(fwd.chain(finalize)); + return Ok(Response::builder() + .status(StatusCode::OK) + .header("content-type", "text/event-stream") + .header("cache-control", "no-cache") + .header("x-accel-buffering", "no") + .body(body) + .expect("build sse response")); + } + Err(err_json( + StatusCode::SERVICE_UNAVAILABLE, + "该模型暂无可用 key", + )) +} + +/// SSE 流式 usage 捕获:转发时记录尾部数据(≤64KB)用于流尾解析; +/// anthropic 的 input_tokens 在 message_start(头部)单独提前捕获 +struct UsageCapture { + protocol: String, + tail: Vec, + input_tokens: f64, +} + +impl UsageCapture { + fn new(protocol: &str) -> Self { + Self { + protocol: protocol.to_string(), + tail: Vec::new(), + input_tokens: 0.0, + } + } + + fn push(&mut self, chunk: &[u8]) { + // anthropic:message_start 事件携带 input_tokens(在流头部,尾部缓冲会丢) + if self.protocol == "anthropic" && self.input_tokens == 0.0 { + if let Some(v) = parse_anthropic_input(chunk) { + self.input_tokens = v; + } + } + self.tail.extend_from_slice(chunk); + if self.tail.len() > 64 * 1024 { + self.tail.drain(0..(self.tail.len() - 64 * 1024)); + } + } + + /// 流尾解析 usage → (input_tokens, output_tokens) + fn finish(&mut self) -> (f64, f64) { + let text = String::from_utf8_lossy(&self.tail); + let mut input = self.input_tokens; + let mut output = 0.0; + for line in text.lines() { + let Some(data) = line.trim_start().strip_prefix("data:") else { + continue; + }; + let Ok(v) = serde_json::from_str::(data.trim()) else { + continue; + }; + if self.protocol == "anthropic" { + if let Some(u) = v.get("usage") { + // message_delta / message_start 均带 usage 字段 + input = u + .get("input_tokens") + .and_then(|x| x.as_f64()) + .unwrap_or(input); + output = u + .get("output_tokens") + .and_then(|x| x.as_f64()) + .unwrap_or(output); + } + } else if let Some(u) = v.get("usage") { + // openai:最后 chunk 的 usage(stream_options.include_usage 时) + input = u + .get("prompt_tokens") + .and_then(|x| x.as_f64()) + .unwrap_or(0.0); + output = u + .get("completion_tokens") + .and_then(|x| x.as_f64()) + .unwrap_or(0.0); + } + } + (input, output) + } +} + +/// 从 anthropic 流式 chunk 提取 input_tokens(message_start 事件) +fn parse_anthropic_input(chunk: &[u8]) -> Option { + let text = String::from_utf8_lossy(chunk); + for line in text.lines() { + let Some(data) = line.trim_start().strip_prefix("data:") else { + continue; + }; + let v: serde_json::Value = serde_json::from_str(data.trim()).ok()?; + if v.get("type").and_then(|t| t.as_str()) == Some("message_start") { + if let Some(tokens) = v.pointer("/message/usage/input_tokens") { + return tokens.as_f64(); + } + } + } + None +} + /// POST /v1/chat/completions(OpenAI 兼容) #[axum::debug_handler] pub async fn chat_completions( @@ -238,7 +498,11 @@ pub async fn chat_completions( ) -> Result { let model = extract_model(&body) .ok_or_else(|| err_json(StatusCode::BAD_REQUEST, "请求体缺少 model 字段"))?; - forward(&st, auth, &model, body.clone(), "openai_chat").await + if body_streaming(&body) { + forward_stream(&st, auth, &model, body.clone(), "openai_chat").await + } else { + forward(&st, auth, &model, body.clone(), "openai_chat").await + } } /// POST /anthropic/v1/messages(Anthropic 兼容) @@ -249,7 +513,19 @@ pub async fn anthropic_messages( ) -> Result { let model = extract_model(&body) .ok_or_else(|| err_json(StatusCode::BAD_REQUEST, "请求体缺少 model 字段"))?; - forward(&st, auth, &model, body.clone(), "anthropic").await + if body_streaming(&body) { + forward_stream(&st, auth, &model, body.clone(), "anthropic").await + } else { + forward(&st, auth, &model, body.clone(), "anthropic").await + } +} + +/// 请求体是否要求流式(stream:true) +fn body_streaming(body: &str) -> bool { + serde_json::from_str::(body) + .ok() + .and_then(|v| v.get("stream").and_then(|s| s.as_bool())) + .unwrap_or(false) } /// GET /api/models(市场页:models 表 + key 可用性,需认证) @@ -283,22 +559,30 @@ mod tests { type_: "paygo".to_string(), key_prefix: "sk-".to_string(), interactive_only: false, - endpoints: vec![crate::config::Endpoint { - protocol: "openai_chat".to_string(), - base_url: base_url.to_string(), - }], + endpoints: vec![ + crate::config::Endpoint { + protocol: "openai_chat".to_string(), + base_url: base_url.to_string(), + }, + crate::config::Endpoint { + protocol: "anthropic".to_string(), + base_url: base_url.to_string(), + }, + ], }); crate::db::seed_models(&conn, &cfg).expect("seed models"); - AppState::new(conn, Arc::new(cfg)) + let crypto = crate::crypto::Crypto::new([9u8; 32]); + AppState::new(conn, Arc::new(cfg), crypto) } - /// 注入测试 key(属主 user_id,模型 model,plan) + /// 注入测试 key(属主 user_id,模型 model,plan;key 值加密落库) fn insert_key(st: &AppState, id: i64, owner: i64, model: &str, plan: &str) { + let encrypted = st.crypto.encrypt(b"sk-test").expect("encrypt test key"); let conn = st.db.lock().unwrap(); conn.execute( "INSERT OR REPLACE INTO keys (id, provider, plan, model, status, owner_id, encrypted_key, quota, used) \ - VALUES (?1, 'test', ?2, ?3, 'on', ?4, 'sk-test', 1000, 0)", - rusqlite::params![id, plan, model, owner], + VALUES (?1, 'test', ?2, ?3, 'on', ?4, ?5, 1000, 0)", + rusqlite::params![id, plan, model, owner, encrypted], ) .unwrap(); } @@ -585,4 +869,240 @@ mod tests { .unwrap(); assert_eq!(resp2.status(), StatusCode::UNAUTHORIZED); } + + /// 假上游:SSE 流式(openai 风格,尾部带 usage + [DONE]) + async fn fake_sse_upstream(listener: tokio::net::TcpListener) { + let app = axum::Router::new().route( + "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/chat/completions", + axum::routing::post(|_body: String| async { + let body = "data: {\"id\":\"1\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"delta\":{\"content\":\"hel\"}}]}\n\n\ + data: {\"id\":\"1\",\"object\":\"chat.completion.chunk\",\"choices\":[{\"delta\":{\"content\":\"lo\"}}]}\n\n\ + data: {\"id\":\"1\",\"object\":\"chat.completion.chunk\",\"usage\":{\"prompt_tokens\":100,\"completion_tokens\":50},\"choices\":[]}\n\n\ + data: [DONE]\n\n"; + ( + [("content-type", "text/event-stream")], + axum::body::Body::from(body), + ) + }), + ); + axum::serve(listener, app).await.unwrap(); + } + + /// 假上游:SSE 流式(anthropic 风格:message_start 带 input,message_delta 带 output) + async fn fake_anthropic_sse_upstream(listener: tokio::net::TcpListener) { + let app = axum::Router::new().route( + "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/v1/messages", + axum::routing::post(|_body: String| async { + let body = "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"usage\":{\"input_tokens\":80,\"output_tokens\":1}}}\n\n\ + event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"delta\":{\"text\":\"hi\"}}\n\n\ + event: message_delta\ndata: {\"type\":\"message_delta\",\"usage\":{\"output_tokens\":30}}\n\n\ + event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n"; + ( + [("content-type", "text/event-stream")], + axum::body::Body::from(body), + ) + }), + ); + axum::serve(listener, app).await.unwrap(); + } + + #[tokio::test] + async fn sse_openai_stream_passthrough_and_settle() { + let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0)) + .await + .unwrap(); + let port = listener.local_addr().unwrap().port(); + let up = tokio::spawn(fake_sse_upstream(listener)); + let base = format!("http://127.0.0.1:{port}"); + + let st = test_state("sse", "test-sse", &base); + { + let conn = st.db.lock().unwrap(); + models_row(&conn, "test", "test-model", 10.0, 20.0); + conn.execute( + "INSERT OR IGNORE INTO users (id, email, password_hash, name, role) VALUES (2, 'owner2@t.local', 'x', '分享者', 'user')", + [], + ) + .unwrap(); + conn.execute( + "INSERT OR IGNORE INTO quotas (user_id, balance) VALUES (2, 0)", + [], + ) + .unwrap(); + } + insert_key(&st, 300, 2, "test-model", "test-sse"); + let key = login_key(st.clone()).await; + + let (s, body) = post_raw( + st.clone(), + "/v1/chat/completions", + r#"{"model":"test-model","stream":true,"messages":[{"role":"user","content":"hi"}]}"#, + Some(&key), + ) + .await; + assert_eq!( + s, + StatusCode::OK, + "body: {}", + String::from_utf8_lossy(&body) + ); + let text = String::from_utf8_lossy(&body); + assert!(text.contains("data: [DONE]"), "SSE 原文透传含 [DONE]"); + assert!( + text.contains("\"content\":\"hel\"") && text.contains("\"content\":\"lo\""), + "chunk 逐块透传" + ); + // 计量:100×10/1e6 + 50×20/1e6 = 0.002 USD × 1000 = 2.0 点 + let conn = st.db.lock().unwrap(); + let bal: f64 = conn + .query_row("SELECT balance FROM quotas WHERE user_id = 1", [], |r| { + r.get(0) + }) + .unwrap(); + assert!((bal - (12471.0 - 2.0)).abs() < 1e-9, "consumer={bal}"); + let n_ur: i64 = conn + .query_row("SELECT COUNT(*) FROM usage_records", [], |r| r.get(0)) + .unwrap(); + assert_eq!(n_ur, 1, "流尾 usage 入账一次"); + drop(conn); + up.abort(); + } + + #[tokio::test] + async fn sse_anthropic_stream_settle() { + let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0)) + .await + .unwrap(); + let port = listener.local_addr().unwrap().port(); + let up = tokio::spawn(fake_anthropic_sse_upstream(listener)); + let base = format!("http://127.0.0.1:{port}"); + + let st = test_state("ssa", "test-ssa", &base); + { + let conn = st.db.lock().unwrap(); + models_row(&conn, "test", "test-model", 10.0, 20.0); + conn.execute( + "INSERT OR IGNORE INTO users (id, email, password_hash, name, role) VALUES (2, 'owner3@t.local', 'x', '分享者', 'user')", + [], + ) + .unwrap(); + conn.execute( + "INSERT OR IGNORE INTO quotas (user_id, balance) VALUES (2, 0)", + [], + ) + .unwrap(); + } + insert_key(&st, 301, 2, "test-model", "test-ssa"); + let key = login_key(st.clone()).await; + + let (s, body) = post_raw( + st.clone(), + "/anthropic/v1/messages", + r#"{"model":"test-model","stream":true,"messages":[{"role":"user","content":"hi"}]}"#, + Some(&key), + ) + .await; + assert_eq!( + s, + StatusCode::OK, + "body: {}", + String::from_utf8_lossy(&body) + ); + let text = String::from_utf8_lossy(&body); + assert!( + text.contains("event: message_stop"), + "anthropic 事件原样透传" + ); + // 80×10/1e6 + 30×20/1e6 = 0.0014 USD × 1000 = 1.4 点 + let conn = st.db.lock().unwrap(); + let bal: f64 = conn + .query_row("SELECT balance FROM quotas WHERE user_id = 1", [], |r| { + r.get(0) + }) + .unwrap(); + assert!((bal - (12471.0 - 1.4)).abs() < 1e-9, "consumer={bal}"); + drop(conn); + up.abort(); + } + + #[tokio::test] + async fn sse_client_disconnect_skips_settle() { + let listener = tokio::net::TcpListener::bind(("127.0.0.1", 0)) + .await + .unwrap(); + let port = listener.local_addr().unwrap().port(); + let up = tokio::spawn(fake_sse_upstream(listener)); + let base = format!("http://127.0.0.1:{port}"); + + let st = test_state("sse_disc", "test-sse-d", &base); + { + let conn = st.db.lock().unwrap(); + models_row(&conn, "test", "test-model", 10.0, 20.0); + } + insert_key(&st, 302, 1, "test-model", "test-sse-d"); + let key = login_key(st.clone()).await; + + let resp = router() + .with_state(st.clone()) + .oneshot( + Request::builder() + .method("POST") + .uri("/v1/chat/completions") + .header("content-type", "application/json") + .header("authorization", format!("Bearer {key}")) + .body(Body::from( + r#"{"model":"test-model","stream":true,"messages":[{"role":"user","content":"hi"}]}"#.to_string(), + )) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), StatusCode::OK); + // 只读一帧就断开(drop body → 上游中止,流尾 finalize 不执行 → 不入账) + use futures_util::StreamExt; + let mut stream = resp.into_body().into_data_stream(); + let _first = stream.next().await; + drop(stream); + tokio::time::sleep(std::time::Duration::from_millis(300)).await; + let conn = st.db.lock().unwrap(); + let n_ur: i64 = conn + .query_row("SELECT COUNT(*) FROM usage_records", [], |r| r.get(0)) + .unwrap(); + assert_eq!(n_ur, 0, "客户端断开后不应入账"); + drop(conn); + up.abort(); + } + + #[test] + fn usage_capture_parses_openai_and_anthropic() { + // openai:usage 在尾部(SSE 事件以换行分隔) + let mut cap = UsageCapture::new("openai_chat"); + cap.push( + br#"data: {"choices":[{"delta":{"content":"hi"}}]}"# + .to_vec() + .as_slice(), + ); + cap.push(b"\n\n"); + cap.push(br#"data: {"usage":{"prompt_tokens":10,"completion_tokens":5},"choices":[]}"#); + cap.push(b"\n\n"); + cap.push(br#"data: [DONE]"#); + let (i, o) = cap.finish(); + assert_eq!(i, 10.0); + assert_eq!(o, 5.0); + + // anthropic:input 在头部 message_start,output 在尾部 message_delta + let mut cap = UsageCapture::new("anthropic"); + cap.push( + br#"event: message_start +data: {"type":"message_start","message":{"usage":{"input_tokens":80,"output_tokens":1}}}"#, + ); + cap.push(b"\n\n"); + cap.push( + br#"event: message_delta +data: {"type":"message_delta","usage":{"output_tokens":30}}"#, + ); + let (i, o) = cap.finish(); + assert_eq!(i, 80.0, "message_start 的 input_tokens 提前捕获"); + assert_eq!(o, 30.0, "message_delta 的 output_tokens 流尾解析"); + } } diff --git a/src/main.rs b/src/main.rs index 96f60cd..1cf622f 100644 --- a/src/main.rs +++ b/src/main.rs @@ -13,6 +13,7 @@ mod auth; mod billing; mod config; +mod crypto; mod dao; mod db; mod gateway; @@ -56,8 +57,15 @@ async fn main() -> anyhow::Result<()> { let conn = db::open(&db_path)?; db::seed_models(&conn, &cfg)?; + // P0-C:主密钥 + 旧明文 key 加密迁移 + let crypto = crypto::Crypto::from_config(&cfg.server.master_key); + let migrated = db::migrate_key_encryption(&conn, &crypto)?; + if migrated > 0 { + log::info!("已加密迁移 {migrated} 条上游 key"); + } + let cfg = Arc::new(cfg); - let state = routes::AppState::new(conn, cfg); + let state = routes::AppState::new(conn, cfg, crypto); let app = routes::router() .with_state(state) .layer(tower_http::trace::TraceLayer::new_for_http()); diff --git a/src/routes/mod.rs b/src/routes/mod.rs index 263d076..ff76acf 100644 --- a/src/routes/mod.rs +++ b/src/routes/mod.rs @@ -10,6 +10,8 @@ //! - GET /api/models(市场) pub mod api_keys; +pub mod sharing; +pub mod wallet; use std::sync::{Arc, Mutex}; @@ -22,21 +24,23 @@ use rusqlite::Connection; use serde::Deserialize; use crate::config::Config; +use crate::crypto::Crypto; use crate::dao; use crate::gateway; use crate::router::RouterState; -/// 共享状态:数据库连接 + 配置 + 路由状态 + HTTP 客户端 +/// 共享状态:数据库连接 + 配置 + 路由状态 + HTTP 客户端 + 密钥加密器 #[derive(Clone)] pub struct AppState { pub db: Arc>, pub cfg: Arc, pub router: Arc, pub http: reqwest::Client, + pub crypto: Crypto, } impl AppState { - pub fn new(conn: Connection, cfg: Arc) -> Self { + pub fn new(conn: Connection, cfg: Arc, crypto: Crypto) -> Self { Self { db: Arc::new(Mutex::new(conn)), cfg, @@ -45,6 +49,7 @@ impl AppState { .timeout(std::time::Duration::from_secs(120)) .build() .expect("reqwest client 构建失败"), + crypto, } } } @@ -139,6 +144,12 @@ pub fn router() -> Router { .route("/v1/chat/completions", post(gateway::chat_completions)) .route("/anthropic/v1/messages", post(gateway::anthropic_messages)) .route("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/api/models", get(gateway::models)) + // P0-C:共享 / 钱包 / 交易 / 仪表盘 + .route("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/api/sharings", post(sharing::create).get(sharing::list)) + .route("/api/sharings/:id", axum::routing::patch(sharing::patch)) + .route("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/api/wallet", get(wallet::wallet)) + .route("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/api/transactions", get(wallet::transactions)) + .route("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/api/dashboard", get(wallet::dashboard)) } #[cfg(test)] @@ -155,7 +166,8 @@ mod tests { let conn = crate::db::open(p.to_str().unwrap()).expect("open tmp db"); let cfg = crate::config::Config::load("config/config.example.toml").unwrap(); crate::db::seed_models(&conn, &cfg).expect("seed models"); - AppState::new(conn, Arc::new(cfg)) + let crypto = crate::crypto::Crypto::new([9u8; 32]); + AppState::new(conn, Arc::new(cfg), crypto) } async fn post( diff --git a/src/routes/sharing.rs b/src/routes/sharing.rs new file mode 100644 index 0000000..3605811 --- /dev/null +++ b/src/routes/sharing.rs @@ -0,0 +1,450 @@ +//! 共享管理 API(对齐原型共享页 US-8/9) +//! +//! P0-C(rant 2026-08-18T10:36:04): +//! - POST /api/sharings 上架(key 加密落库,DB 无明文) +//! - GET /api/sharings 我的共享列表(key 脱敏 sk-****xxxx) +//! - PATCH /api/sharings/:id 暂停/恢复/删除(status: paused/on/off,软删) +//! - 可用时间段字段:available_days + start/end(先存后展示,生效判定留 P1) + +use axum::extract::{Path, State}; +use axum::Json; +use rusqlite::params; +use serde::Deserialize; + +use crate::routes::{internal, ApiErr, AppState, AuthUser}; + +/// 上架请求 +#[derive(Debug, Deserialize)] +pub struct CreateSharingReq { + pub provider: String, + #[serde(default)] + pub plan: String, + pub model: String, + /// 上游 key(明文,服务端加密后落库) + pub key: String, + #[serde(default)] + pub quota: f64, + #[serde(default)] + pub available: Option, + #[serde(default)] + pub note: String, +} + +/// 可用时间段(先存后展示;生效判定留 P1) +#[derive(Debug, Deserialize)] +pub struct Avail { + /// 星期(1-7) + #[serde(default)] + pub days: Vec, + /// HH:mm + #[serde(default)] + pub start: String, + /// HH:mm + #[serde(default)] + pub end: String, +} + +/// 状态变更请求 +#[derive(Debug, Deserialize)] +pub struct PatchSharingReq { + /// paused / on / off(off = 软删除) + pub status: String, +} + +/// key 脱敏:sk-****xxxx(保留前 2 位前缀 + 后 4 位,与原型一致) +fn mask_upstream_key(key: &str) -> String { + if key.len() > 6 { + let prefix = &key[..2.min(key.len())]; + let tail = &key[key.len() - 4..]; + format!("{prefix}-****{tail}") + } else { + "****".to_string() + } +} + +/// POST /api/sharings:上架共享 key(加密落库) +pub async fn create( + State(st): State, + auth: AuthUser, + Json(req): Json, +) -> Result, ApiErr> { + if req.model.trim().is_empty() || req.key.trim().is_empty() { + return Err(( + axum::http::StatusCode::BAD_REQUEST, + Json(serde_json::json!({ "error": "model 与 key 必填" })), + )); + } + let encrypted = st + .crypto + .encrypt(req.key.trim().as_bytes()) + .map_err(internal)?; + let (days, start, end) = match &req.available { + Some(a) => ( + serde_json::to_string(&a.days).unwrap_or_default(), + a.start.clone(), + a.end.clone(), + ), + None => (String::new(), String::new(), String::new()), + }; + let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; + conn.execute( + "INSERT INTO keys (provider, plan, model, status, owner_id, encrypted_key, quota, available_days, available_start, available_end, note) \ + VALUES (?1, ?2, ?3, 'on', ?4, ?5, ?6, ?7, ?8, ?9, ?10)", + params![ + req.provider, + req.plan, + req.model, + auth.user_id, + encrypted, + req.quota, + days, + start, + end, + req.note + ], + ) + .map_err(internal)?; + let id = conn.last_insert_rowid(); + Ok(Json(serde_json::json!({ + "id": id, + "provider": req.provider, + "model": req.model, + "key": mask_upstream_key(&req.key), + "status": "on", + "available_days": days, + "available_start": start, + "available_end": end, + "note": req.note, + }))) +} + +/// 单条共享(含收益汇总);key 先解密再脱敏展示 +fn sharing_row( + conn: &rusqlite::Connection, + crypto: &crate::crypto::Crypto, + r: &rusqlite::Row, +) -> rusqlite::Result { + let id: i64 = r.get(0)?; + let provider: String = r.get(1)?; + let plan: String = r.get(2)?; + let model: String = r.get(3)?; + let status: String = r.get(4)?; + let encrypted_key: String = r.get(5)?; + let quota: f64 = r.get(6)?; + let used: f64 = r.get(7)?; + let days: String = r.get(8)?; + let start: String = r.get(9)?; + let end: String = r.get(10)?; + let note: String = r.get(11)?; + // 解密 → 脱敏(sk-****xxxx);解密失败展示 **** + let masked = crypto + .decrypt(&encrypted_key) + .ok() + .and_then(|k| String::from_utf8(k).ok()) + .map(|k| mask_upstream_key(&k)) + .unwrap_or_else(|| "****".to_string()); + // 收益:该 key 的 earn 交易累计 + let earn: f64 = conn + .query_row( + "SELECT COALESCE(SUM(pts), 0) FROM transactions WHERE key_id = ?1 AND type = 'earn'", + [id], + |r| r.get(0), + ) + .unwrap_or(0.0); + Ok(serde_json::json!({ + "id": id, + "provider": provider, + "plan": plan, + "model": model, + "status": status, + "key": masked, + "quota": quota, + "used": used, + "earn": earn, + "available_days": days, + "available_start": start, + "available_end": end, + "note": note, + })) +} + +/// GET /api/sharings:我的共享列表(脱敏) +pub async fn list( + State(st): State, + auth: AuthUser, +) -> Result>, ApiErr> { + let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; + let crypto = st.crypto.clone(); + let mut stmt = conn + .prepare( + "SELECT id, provider, plan, model, status, encrypted_key, quota, used, \ + available_days, available_start, available_end, note \ + FROM keys WHERE owner_id = ?1 ORDER BY id DESC", + ) + .map_err(internal)?; + let rows = stmt + .query_map([auth.user_id], |r| sharing_row(&conn, &crypto, r)) + .map_err(internal)?; + let mut out = Vec::new(); + for r in rows { + out.push(r.map_err(internal)?); + } + Ok(Json(out)) +} + +/// PATCH /api/sharings/:id:暂停/恢复/删除(status: paused/on/off) +pub async fn patch( + State(st): State, + auth: AuthUser, + Path(id): Path, + Json(req): Json, +) -> Result, ApiErr> { + if !matches!(req.status.as_str(), "paused" | "on" | "off") { + return Err(( + axum::http::StatusCode::BAD_REQUEST, + Json(serde_json::json!({ "error": "status 必须为 paused / on / off" })), + )); + } + let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; + let n = conn + .execute( + "UPDATE keys SET status = ?1 WHERE id = ?2 AND owner_id = ?3", + params![req.status, id, auth.user_id], + ) + .map_err(internal)?; + if n == 0 { + return Err(( + axum::http::StatusCode::NOT_FOUND, + Json(serde_json::json!({ "error": "共享不存在或不属于当前用户" })), + )); + } + let crypto = st.crypto.clone(); + let row = conn + .query_row( + "SELECT id, provider, plan, model, status, encrypted_key, quota, used, \ + available_days, available_start, available_end, note \ + FROM keys WHERE id = ?1", + [id], + |r| sharing_row(&conn, &crypto, r), + ) + .map_err(internal)?; + Ok(Json(row)) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::routes::router; + use axum::body::Body; + use axum::http::Request; + use std::sync::Arc; + use tower::util::ServiceExt; + + fn test_state(tag: &str) -> AppState { + let p = std::env::temp_dir().join(format!("atp_share_{}_{}.db", std::process::id(), tag)); + let _ = std::fs::remove_file(&p); + let conn = crate::db::open(p.to_str().unwrap()).expect("open tmp db"); + let cfg = crate::config::Config::load("config/config.example.toml").unwrap(); + crate::db::seed_models(&conn, &cfg).expect("seed models"); + let crypto = crate::crypto::Crypto::new([13u8; 32]); + AppState::new(conn, Arc::new(cfg), crypto) + } + + async fn login(st: AppState) -> String { + let resp = router() + .with_state(st) + .oneshot( + Request::builder() + .method("POST") + .uri("/api/auth/login") + .header("content-type", "application/json") + .body(Body::from( + r#"{"email":"demo@aitokenpool.local","password":"demo1234"}"#.to_string(), + )) + .unwrap(), + ) + .await + .unwrap(); + let bytes = axum::body::to_bytes(resp.into_body(), 1024 * 1024) + .await + .unwrap(); + let v: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); + v["api_key"].as_str().unwrap().to_string() + } + + async fn send( + st: AppState, + method: &str, + uri: &str, + payload: Option<&str>, + bearer: &str, + ) -> (axum::http::StatusCode, String) { + let mut b = Request::builder() + .method(method) + .uri(uri) + .header("authorization", format!("Bearer {bearer}")); + let resp = match payload { + Some(body_str) => { + b = b.header("content-type", "application/json"); + router() + .with_state(st) + .oneshot(b.body(Body::from(body_str.to_string())).unwrap()) + .await + .unwrap() + } + None => router() + .with_state(st) + .oneshot(b.body(Body::empty()).unwrap()) + .await + .unwrap(), + }; + let status = resp.status(); + let bytes = axum::body::to_bytes(resp.into_body(), 2 * 1024 * 1024) + .await + .unwrap(); + (status, String::from_utf8_lossy(&bytes).to_string()) + } + + #[tokio::test] + async fn create_encrypts_key_and_list_masks() { + let st = test_state("create"); + let key = login(st.clone()).await; + + let (s, body) = send( + st.clone(), + "POST", + "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/api/sharings", + Some(r#"{"provider":"deepseek","plan":"deepseek-paygo","model":"deepseek-v4-flash","key":"sk-realsecret1234","quota":1000,"available":{"days":[1,2,3,4,5],"start":"09:00","end":"18:00"},"note":"工作日共享"}"#), + &key, + ) + .await; + assert_eq!(s, axum::http::StatusCode::OK, "body: {body}"); + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + let id = v["id"].as_i64().unwrap(); + + // DB 中无明文(加密落库)——锁作用域块内,绝不让 MutexGuard 跨 await + let (_stored, days, start, end, note) = { + let conn = st.db.lock().unwrap(); + let stored: String = conn + .query_row("SELECT encrypted_key FROM keys WHERE id = ?1", [id], |r| { + r.get(0) + }) + .unwrap(); + assert!( + stored.starts_with(crate::crypto::PREFIX), + "密文前缀: {stored}" + ); + assert!(!stored.contains("sk-realsecret1234"), "DB 不得存明文"); + // 可用时间段字段正确 + let (days, start, end, note): (String, String, String, String) = conn + .query_row( + "SELECT available_days, available_start, available_end, note FROM keys WHERE id = ?1", + [id], + |r| Ok((r.get(0)?, r.get(1)?, r.get(2)?, r.get(3)?)), + ) + .unwrap(); + (stored, days, start, end, note) + }; + assert!(days.contains("1") && days.contains("5"), "days={days}"); + assert_eq!(start, "09:00"); + assert_eq!(end, "18:00"); + assert_eq!(note, "工作日共享"); + + // 列表脱敏:sk-****1234,不含真实 key + let (s, body) = send(st.clone(), "GET", "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/api/sharings", None, &key).await; + assert_eq!(s, axum::http::StatusCode::OK, "body: {body}"); + let arr: Vec = serde_json::from_str(&body).unwrap(); + assert!(arr.iter().any(|r| r["id"] == id)); + let row = arr.iter().find(|r| r["id"] == id).unwrap(); + assert_eq!(row["key"], "sk-****1234", "脱敏: {}", row["key"]); + assert!(!body.contains("sk-realsecret1234"), "列表不得泄露明文"); + assert_eq!(row["status"], "on"); + assert_eq!(row["available_end"], "18:00"); + } + + #[tokio::test] + async fn patch_pause_and_off() { + let st = test_state("patch"); + let key = login(st.clone()).await; + let (_, body) = send( + st.clone(), + "POST", + "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/api/sharings", + Some(r#"{"provider":"deepseek","model":"deepseek-v4-flash","key":"sk-patchme9999"}"#), + &key, + ) + .await; + let id: i64 = serde_json::from_str::(&body).unwrap()["id"] + .as_i64() + .unwrap(); + + // 暂停 + let (s, body) = send( + st.clone(), + "PATCH", + &format!("/api/sharings/{id}"), + Some(r#"{"status":"paused"}"#), + &key, + ) + .await; + assert_eq!(s, axum::http::StatusCode::OK, "body: {body}"); + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + assert_eq!(v["status"], "paused"); + + // 软删(off) + let (s, _) = send( + st.clone(), + "PATCH", + &format!("/api/sharings/{id}"), + Some(r#"{"status":"off"}"#), + &key, + ) + .await; + assert_eq!(s, axum::http::StatusCode::OK); + let status: String = { + let conn = st.db.lock().unwrap(); + conn.query_row("SELECT status FROM keys WHERE id = ?1", [id], |r| r.get(0)) + .unwrap() + }; + assert_eq!(status, "off", "软删保留账本引用"); + + // 非法 status → 400 + let (s, _) = send( + st.clone(), + "PATCH", + &format!("/api/sharings/{id}"), + Some(r#"{"status":"deleted"}"#), + &key, + ) + .await; + assert_eq!(s, axum::http::StatusCode::BAD_REQUEST); + + // 他人 id → 404 + let (s, _) = send( + st, + "PATCH", + "/api/sharings/99999", + Some(r#"{"status":"on"}"#), + &key, + ) + .await; + assert_eq!(s, axum::http::StatusCode::NOT_FOUND); + } + + #[tokio::test] + async fn sharing_requires_bearer() { + let st = test_state("nobearer"); + let resp = router() + .with_state(st) + .oneshot( + Request::builder() + .method("GET") + .uri("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/api/sharings") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), axum::http::StatusCode::UNAUTHORIZED); + } +} diff --git a/src/routes/wallet.rs b/src/routes/wallet.rs new file mode 100644 index 0000000..1ee765e --- /dev/null +++ b/src/routes/wallet.rs @@ -0,0 +1,370 @@ +//! 钱包 / 交易 / 仪表盘 API(对齐原型钱包页 + 交易页 + 仪表盘) +//! +//! P0-C(rant 2026-08-18T10:36:04): +//! - GET /api/wallet → {balance, month_use, month_earn} +//! - GET /api/transactions?type=&page=&page_size= → 分页 + type 过滤(consume/earn/all) +//! - GET /api/dashboard → 本月按类型聚合 + 净变化 + 近 N 天序列(sparkline) + +use axum::extract::{Query, State}; +use axum::Json; +use rusqlite::params; +use serde::Deserialize; + +use crate::routes::{internal, ApiErr, AppState, AuthUser}; + +/// GET /api/wallet +pub async fn wallet( + State(st): State, + auth: AuthUser, +) -> Result, ApiErr> { + let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; + let balance: f64 = conn + .query_row( + "SELECT balance FROM quotas WHERE user_id = ?1", + [auth.user_id], + |r| r.get(0), + ) + .unwrap_or(0.0); + let month_use: f64 = conn + .query_row( + "SELECT COALESCE(SUM(pts), 0) FROM transactions \ + WHERE user_id = ?1 AND type = 'consume' AND strftime('%Y-%m', time) = strftime('%Y-%m', 'now')", + [auth.user_id], + |r| r.get(0), + ) + .unwrap_or(0.0); + let month_earn: f64 = conn + .query_row( + "SELECT COALESCE(SUM(pts), 0) FROM transactions \ + WHERE user_id = ?1 AND type = 'earn' AND strftime('%Y-%m', time) = strftime('%Y-%m', 'now')", + [auth.user_id], + |r| r.get(0), + ) + .unwrap_or(0.0); + Ok(Json(serde_json::json!({ + "balance": balance, + "month_use": month_use, + "month_earn": month_earn, + }))) +} + +/// GET /api/transactions 查询参数 +#[derive(Debug, Deserialize)] +pub struct TxQuery { + /// consume / earn / all(缺省 all) + #[serde(default)] + pub r#type: String, + #[serde(default = "default_page")] + pub page: u32, + #[serde(default = "default_page_size")] + pub page_size: u32, +} + +fn default_page() -> u32 { + 1 +} +fn default_page_size() -> u32 { + 20 +} + +/// GET /api/transactions:时间倒序 + type 过滤 + 分页 +pub async fn transactions( + State(st): State, + auth: AuthUser, + Query(q): Query, +) -> Result, ApiErr> { + let page = q.page.max(1); + let page_size = q.page_size.clamp(1, 100); + let type_filter = match q.r#type.as_str() { + "" | "all" => None, + t @ ("consume" | "earn") => Some(t.to_string()), + _ => { + return Err(( + axum::http::StatusCode::BAD_REQUEST, + Json(serde_json::json!({ "error": "type 必须为 consume / earn / all" })), + )) + } + }; + let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; + let total: i64 = match &type_filter { + Some(t) => conn + .query_row( + "SELECT COUNT(*) FROM transactions WHERE user_id = ?1 AND type = ?2", + params![auth.user_id, t], + |r| r.get(0), + ) + .unwrap_or(0), + None => conn + .query_row( + "SELECT COUNT(*) FROM transactions WHERE user_id = ?1", + [auth.user_id], + |r| r.get(0), + ) + .unwrap_or(0), + }; + let offset = (page - 1) * page_size; + let mut stmt = match &type_filter { + Some(_) => conn + .prepare( + "SELECT id, counterpart, key_id, model, tokens, pts, type, status, time \ + FROM transactions WHERE user_id = ?1 AND type = ?2 \ + ORDER BY id DESC LIMIT ?3 OFFSET ?4", + ) + .map_err(internal)?, + None => conn + .prepare( + "SELECT id, counterpart, key_id, model, tokens, pts, type, status, time \ + FROM transactions WHERE user_id = ?1 \ + ORDER BY id DESC LIMIT ?2 OFFSET ?3", + ) + .map_err(internal)?, + }; + let rows: Vec = match &type_filter { + Some(t) => stmt + .query_map(params![auth.user_id, t, page_size, offset], |r| { + Ok(serde_json::json!({ + "id": r.get::<_, i64>(0)?, + "counterpart": r.get::<_, String>(1)?, + "key_id": r.get::<_, Option>(2)?, + "model": r.get::<_, String>(3)?, + "tokens": r.get::<_, f64>(4)?, + "pts": r.get::<_, f64>(5)?, + "type": r.get::<_, String>(6)?, + "status": r.get::<_, String>(7)?, + "time": r.get::<_, String>(8)?, + })) + }) + .map_err(internal)? + .collect::>>() + .map_err(internal)?, + None => stmt + .query_map(params![auth.user_id, page_size, offset], |r| { + Ok(serde_json::json!({ + "id": r.get::<_, i64>(0)?, + "counterpart": r.get::<_, String>(1)?, + "key_id": r.get::<_, Option>(2)?, + "model": r.get::<_, String>(3)?, + "tokens": r.get::<_, f64>(4)?, + "pts": r.get::<_, f64>(5)?, + "type": r.get::<_, String>(6)?, + "status": r.get::<_, String>(7)?, + "time": r.get::<_, String>(8)?, + })) + }) + .map_err(internal)? + .collect::>>() + .map_err(internal)?, + }; + Ok(Json(serde_json::json!({ + "items": rows, + "total": total, + "page": page, + "page_size": page_size, + }))) +} + +/// GET /api/dashboard:本月按类型聚合 + 净变化 + 近 7 天净额序列 +pub async fn dashboard( + State(st): State, + auth: AuthUser, +) -> Result, ApiErr> { + let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; + // 本月按类型聚合 + let mut stmt = conn + .prepare( + "SELECT type, COALESCE(SUM(pts), 0) FROM transactions \ + WHERE user_id = ?1 AND strftime('%Y-%m', time) = strftime('%Y-%m', 'now') \ + GROUP BY type", + ) + .map_err(internal)?; + let month = stmt + .query_map([auth.user_id], |r| { + Ok(serde_json::json!({ + "type": r.get::<_, String>(0)?, + "pts": r.get::<_, f64>(1)?, + })) + }) + .map_err(internal)? + .collect::>>() + .map_err(internal)?; + // 近 7 天净额序列(earn 为正、consume 为负) + let mut stmt = conn + .prepare( + "SELECT date(time), COALESCE(SUM(CASE WHEN type = 'earn' THEN pts ELSE -pts END), 0) \ + FROM transactions \ + WHERE user_id = ?1 AND date(time) >= date('now', '-6 days') \ + GROUP BY date(time) ORDER BY date(time)", + ) + .map_err(internal)?; + let series = stmt + .query_map([auth.user_id], |r| { + Ok(serde_json::json!({ + "date": r.get::<_, String>(0)?, + "pts": r.get::<_, f64>(1)?, + })) + }) + .map_err(internal)? + .collect::>>() + .map_err(internal)?; + let net: f64 = series + .iter() + .map(|s| s["pts"].as_f64().unwrap_or(0.0)) + .sum(); + Ok(Json(serde_json::json!({ + "month": month, + "net": net, + "series": series, + }))) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::routes::router; + use axum::body::Body; + use axum::http::Request; + use std::sync::Arc; + use tower::util::ServiceExt; + + fn test_state(tag: &str) -> AppState { + let p = std::env::temp_dir().join(format!("atp_wallet_{}_{}.db", std::process::id(), tag)); + let _ = std::fs::remove_file(&p); + let conn = crate::db::open(p.to_str().unwrap()).expect("open tmp db"); + let cfg = crate::config::Config::load("config/config.example.toml").unwrap(); + crate::db::seed_models(&conn, &cfg).expect("seed models"); + let crypto = crate::crypto::Crypto::new([17u8; 32]); + AppState::new(conn, Arc::new(cfg), crypto) + } + + async fn get(st: AppState, uri: &str, bearer: &str) -> (axum::http::StatusCode, String) { + let resp = router() + .with_state(st) + .oneshot( + Request::builder() + .method("GET") + .uri(uri) + .header("authorization", format!("Bearer {bearer}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + let status = resp.status(); + let bytes = axum::body::to_bytes(resp.into_body(), 2 * 1024 * 1024) + .await + .unwrap(); + (status, String::from_utf8_lossy(&bytes).to_string()) + } + + async fn login(st: AppState) -> String { + let resp = router() + .with_state(st) + .oneshot( + Request::builder() + .method("POST") + .uri("/api/auth/login") + .header("content-type", "application/json") + .body(Body::from( + r#"{"email":"demo@aitokenpool.local","password":"demo1234"}"#.to_string(), + )) + .unwrap(), + ) + .await + .unwrap(); + let bytes = axum::body::to_bytes(resp.into_body(), 1024 * 1024) + .await + .unwrap(); + let v: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); + v["api_key"].as_str().unwrap().to_string() + } + + #[tokio::test] + async fn wallet_summary_and_dashboard() { + let st = test_state("wallet"); + let key = login(st.clone()).await; + // 种子交易:consume 2.0 + earn 1.8(本月) + { + let conn = st.db.lock().unwrap(); + conn.execute( + "INSERT INTO transactions (user_id, counterpart, key_id, model, tokens, pts, type, status) VALUES (1, '2', 1, 'm', 150, 2.0, 'consume', '成功')", + [], + ) + .unwrap(); + conn.execute( + "INSERT INTO transactions (user_id, counterpart, key_id, model, tokens, pts, type, status) VALUES (1, '3', 2, 'm', 150, 1.8, 'earn', '成功')", + [], + ) + .unwrap(); + } + let (s, body) = get(st.clone(), "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/api/wallet", &key).await; + assert_eq!(s, axum::http::StatusCode::OK, "body: {body}"); + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + assert!((v["balance"].as_f64().unwrap() - 12471.0).abs() < 1e-9); + assert!((v["month_use"].as_f64().unwrap() - 2.0).abs() < 1e-9); + assert!((v["month_earn"].as_f64().unwrap() - 1.8).abs() < 1e-9); + + let (s, body) = get(st.clone(), "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/api/dashboard", &key).await; + assert_eq!(s, axum::http::StatusCode::OK, "body: {body}"); + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + let month = v["month"].as_array().unwrap(); + assert!(month.len() >= 2, "本月 consume+earn 两类聚合: {month:?}"); + let net = v["net"].as_f64().unwrap(); + assert!((net - (1.8 - 2.0)).abs() < 1e-9, "净变化 = earn - consume"); + assert!(!v["series"].as_array().unwrap().is_empty(), "近 7 天序列"); + } + + #[tokio::test] + async fn transactions_filter_and_pagination() { + let st = test_state("tx"); + let key = login(st.clone()).await; + { + let conn = st.db.lock().unwrap(); + for i in 0..5 { + let t = if i % 2 == 0 { "consume" } else { "earn" }; + conn.execute( + "INSERT INTO transactions (user_id, counterpart, key_id, model, tokens, pts, type, status) VALUES (1, 'c', 1, 'm', 1, ?1, ?2, '成功')", + rusqlite::params![i as f64 + 1.0, t], + ) + .unwrap(); + } + } + // 全部 + 分页(page=1 page_size=3 → 3 条;total=5) + let (s, body) = get(st.clone(), "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/api/transactions?page=1&page_size=3", &key).await; + assert_eq!(s, axum::http::StatusCode::OK, "body: {body}"); + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + assert_eq!(v["total"], 5); + assert_eq!(v["items"].as_array().unwrap().len(), 3); + // 时间倒序:最新在前 + let items = v["items"].as_array().unwrap(); + assert!(items[0]["id"].as_i64().unwrap() > items[1]["id"].as_i64().unwrap()); + // type=consume 过滤 → 3 条 + let (_, body) = get(st.clone(), "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/api/transactions?type=consume", &key).await; + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + assert_eq!(v["total"], 3); + assert!(v["items"] + .as_array() + .unwrap() + .iter() + .all(|r| r["type"] == "consume")); + // type=earn → 2 条 + let (_, body) = get(st.clone(), "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/api/transactions?type=earn", &key).await; + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + assert_eq!(v["total"], 2); + // 非法 type → 400 + let (s, _) = get(st.clone(), "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/api/transactions?type=hack", &key).await; + assert_eq!(s, axum::http::StatusCode::BAD_REQUEST); + // 无认证 → 401 + let resp = router() + .with_state(st) + .oneshot( + Request::builder() + .method("GET") + .uri("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/api/wallet") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(resp.status(), axum::http::StatusCode::UNAUTHORIZED); + } +}