diff --git a/Cargo.lock b/Cargo.lock index 1f46d93..59dc3d8 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -60,7 +60,7 @@ dependencies = [ [[package]] name = "aitokenpool" -version = "0.2.2" +version = "0.3.0" dependencies = [ "aes-gcm", "anyhow", diff --git a/Cargo.toml b/Cargo.toml index a599877..703bea2 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "aitokenpool" -version = "0.2.2" +version = "0.3.0" edition = "2021" description = "AI Token 共享池 — 企业 key 池 + 公共共享市场" license = "MIT" diff --git a/src/billing.rs b/src/billing.rs index 21fba83..f0871a9 100644 --- a/src/billing.rs +++ b/src/billing.rs @@ -1,4 +1,4 @@ -//! 计量账本(architecture §4.3/4.4) +//! 计量账本(architecture §4.3/4.4 + P1 点数规则细化) //! //! P0-B(rant 2026-08-18T09:55:57): //! - 成本 = prompt_tokens × input_per_m/1e6 + completion_tokens × output_per_m/1e6 @@ -8,6 +8,10 @@ //! - 分享者(key 属主)得 90%(平台抽成 10%),写 transactions(consume / earn) //! - 写 usage_records;更新 keys.used += tokens //! - 调用+记账事务性处理:上游失败不入账(settle 只在成功响应后调用) +//! +//! P1(rant 2026-08-18T11:03:02): +//! - 可用余额 = gift_balance + balance(预检与 settle 一致) +//! - 扣减顺序:先扣最早到期的赠送点数(gift_grants 按 expires_at ASC),不足再扣永久 balance use anyhow::Result; use rusqlite::Connection; @@ -67,15 +71,17 @@ pub struct SettleParams { pub cost: f64, } -/// 事务性入账:扣消费者 → 加分享者 90% → 两条 transactions → usage_records → keys.used -/// 任一步失败整体回滚(调用方只在成功响应后调用,天然满足「失败不入账」) +/// 事务性入账:扣消费者(先赠送后永久)→ 加分享者 90% → 两条 transactions → +/// usage_records → keys.used。任一步失败整体回滚(调用方只在成功响应后调用, +/// 天然满足「失败不入账」) pub fn settle(conn: &mut Connection, p: &SettleParams) -> Result<()> { let tx = conn.transaction()?; - // 消费者扣 balance(余额允许为负——预检已拦截 ≤0 的请求,负余额由后续充值覆盖) + // 消费者扣减:先扣最早到期的赠送点数,剩余从永久 balance 扣 + let remaining = crate::gift::deduct_gift_first(&tx, p.consumer_id, p.pts)?; tx.execute( "UPDATE quotas SET balance = balance - ?1, updated_at = datetime('now') WHERE user_id = ?2", - rusqlite::params![p.pts, p.consumer_id], + rusqlite::params![remaining, p.consumer_id], )?; // 分享者加 90%(平台抽成 10%) @@ -190,11 +196,11 @@ mod tests { let (mut conn, p) = tmp_db("settle"); // 属主用户(user_id=2)与 key conn.execute( - "INSERT INTO users (id, email, password_hash, name, role) VALUES (2, 'owner@t.local', 'x', '分享者', 'user')", + "INSERT INTO users (id, email, password_hash, name, role) VALUES (100, 'owner@t.local', 'x', '分享者', 'user')", [], ) .unwrap(); - conn.execute("INSERT INTO quotas (user_id, balance) VALUES (2, 0)", []) + conn.execute("INSERT INTO quotas (user_id, balance) VALUES (100, 0)", []) .unwrap(); conn.execute( "INSERT INTO keys (id, provider, plan, model, status, owner_id, encrypted_key, quota, used) \ @@ -214,7 +220,7 @@ mod tests { consumer_id: 1, api_key_id: Some(3), key_id: 9, - owner_id: 2, + owner_id: 100, model: "test-model".into(), tokens: 150.0, pts: 2.0, @@ -231,7 +237,7 @@ mod tests { assert!((bal_c - (12471.0 - 2.0)).abs() < 1e-9); // 分享者加 1.8(90%) let bal_o: f64 = conn - .query_row("SELECT balance FROM quotas WHERE user_id = 2", [], |r| { + .query_row("SELECT balance FROM quotas WHERE user_id = 100", [], |r| { r.get(0) }) .unwrap(); @@ -248,9 +254,11 @@ mod tests { .unwrap(); assert_eq!(t_consume, "consume"); let t_earn: String = conn - .query_row("SELECT type FROM transactions WHERE user_id = 2", [], |r| { - r.get(0) - }) + .query_row( + "SELECT type FROM transactions WHERE user_id = 100", + [], + |r| r.get(0), + ) .unwrap(); assert_eq!(t_earn, "earn"); // usage_records 一条 @@ -297,4 +305,143 @@ mod tests { drop(conn); let _ = std::fs::remove_file(p); } + + #[test] + fn settle_deducts_gift_first_then_permanent() { + let (mut conn, p) = tmp_db("settle_gift"); + // 消费者 user_id=1:赠送 1 点(当天 23:59:59 过期)+ 永久 10 点 + conn.execute( + "INSERT INTO gift_grants (user_id, amount, granted_at, expires_at, status) \ + VALUES (1, 1, '2026-08-18 10:00:00', '2026-08-18 23:59:59', 'active')", + [], + ) + .unwrap(); + conn.execute( + "UPDATE quotas SET gift_balance = 1, balance = 10 WHERE user_id = 1", + [], + ) + .unwrap(); + // 分享者 user_id=2 与 key + conn.execute( + "INSERT INTO users (id, email, password_hash, name, role) VALUES (100, 'owner2@t.local', 'x', '分享者', 'user')", + [], + ) + .unwrap(); + conn.execute("INSERT INTO quotas (user_id, balance) VALUES (100, 0)", []) + .unwrap(); + conn.execute( + "INSERT INTO keys (id, provider, plan, model, status, owner_id, encrypted_key, quota, used) \ + VALUES (8, 'test', 'test-plan', 'test-model', 'on', 2, 'sk-test', 1000, 0)", + [], + ) + .unwrap(); + + let params = SettleParams { + consumer_id: 1, + api_key_id: Some(3), + key_id: 8, + owner_id: 100, + model: "test-model".into(), + tokens: 100.0, + pts: 3.0, + cost: 0.003, + }; + settle(&mut conn, ¶ms).unwrap(); + + // 赠送 1 点全部花掉(used)+ 永久扣 2 点 + let gift: f64 = conn + .query_row( + "SELECT gift_balance FROM quotas WHERE user_id = 1", + [], + |r| r.get(0), + ) + .unwrap(); + assert_eq!(gift, 0.0, "赠送先扣光"); + let g_status: String = conn + .query_row( + "SELECT status FROM gift_grants WHERE user_id = 1", + [], + |r| r.get(0), + ) + .unwrap(); + assert_eq!(g_status, "used"); + let bal: f64 = conn + .query_row("SELECT balance FROM quotas WHERE user_id = 1", [], |r| { + r.get(0) + }) + .unwrap(); + assert!((bal - 8.0).abs() < 1e-9, "永久扣 2 点: {bal}"); + // 分享者照常 90% + let owner: f64 = conn + .query_row("SELECT balance FROM quotas WHERE user_id = 100", [], |r| { + r.get(0) + }) + .unwrap(); + assert!((owner - 2.7).abs() < 1e-9, "分享者 3×0.9=2.7: {owner}"); + // 两条 transactions + let n: i64 = conn + .query_row("SELECT COUNT(*) FROM transactions", [], |r| r.get(0)) + .unwrap(); + assert_eq!(n, 2); + + drop(conn); + let _ = std::fs::remove_file(p); + } + + #[test] + fn settle_expired_gift_not_consumed() { + let (mut conn, p) = tmp_db("settle_expired"); + // 一笔已过期(昨天)的赠送:settle 前惰性清理 → 只扣永久 + conn.execute( + "INSERT INTO gift_grants (user_id, amount, granted_at, expires_at, status) \ + VALUES (1, 1, '2026-08-17 10:00:00', '2026-08-17 23:59:59', 'active')", + [], + ) + .unwrap(); + conn.execute( + "UPDATE quotas SET gift_balance = 1, balance = 10 WHERE user_id = 1", + [], + ) + .unwrap(); + let params = SettleParams { + consumer_id: 1, + api_key_id: None, + key_id: 1, // seed 里的 demo key + owner_id: 1, + model: "m".into(), + tokens: 10.0, + pts: 2.0, + cost: 0.002, + }; + settle(&mut conn, ¶ms).unwrap(); + let gift: f64 = conn + .query_row( + "SELECT gift_balance FROM quotas WHERE user_id = 1", + [], + |r| r.get(0), + ) + .unwrap(); + assert_eq!(gift, 0.0, "过期赠送不参与扣减"); + let g_status: String = conn + .query_row( + "SELECT status FROM gift_grants WHERE user_id = 1", + [], + |r| r.get(0), + ) + .unwrap(); + assert_eq!(g_status, "expired", "惰性标记 expired"); + let bal: f64 = conn + .query_row("SELECT balance FROM quotas WHERE user_id = 1", [], |r| { + r.get(0) + }) + .unwrap(); + // 10 - 2(消费)+ 1.8(同属主 90% 分成)= 9.8 + assert!( + (bal - 9.8).abs() < 1e-9, + "过期赠送不扣,全部从永久扣: {bal}" + ); + + drop(conn); + let _ = std::fs::remove_file(p); + } } diff --git a/src/dao.rs b/src/dao.rs index 7ee8fd2..2a589ff 100644 --- a/src/dao.rs +++ b/src/dao.rs @@ -69,12 +69,14 @@ pub fn list_api_keys(conn: &Connection, user_id: i64) -> Result Option<(i64, i64)> { +/// Bearer 认证:按 key 查归属用户 + api_key id + 角色 → Some((user_id, api_key_id, role)) +pub fn find_api_key_user_and_id(conn: &Connection, key: &str) -> Option<(i64, i64, String)> { conn.query_row( - "SELECT user_id, id FROM api_keys WHERE key_value = ?1 AND status = 'active'", + "SELECT a.user_id, a.id, u.role FROM api_keys a \ + JOIN users u ON u.id = a.user_id \ + WHERE a.key_value = ?1 AND a.status = 'active'", [key], - |r| Ok((r.get(0)?, r.get(1)?)), + |r| Ok((r.get(0)?, r.get(1)?, r.get(2)?)), ) .ok() } @@ -111,14 +113,20 @@ pub fn find_keys_by_model(conn: &Connection, model: &str) -> Result> Ok(out) } -/// 用户点数余额(无账户按 0) -pub fn get_balance(conn: &Connection, user_id: i64) -> f64 { +/// 用户可用余额拆分 → (permanent, gift);可用总额 = 两者之和 +pub fn get_balances(conn: &Connection, user_id: i64) -> (f64, f64) { conn.query_row( - "SELECT balance FROM quotas WHERE user_id = ?1", + "SELECT balance, gift_balance FROM quotas WHERE user_id = ?1", [user_id], - |r| r.get(0), + |r| Ok((r.get(0)?, r.get(1)?)), ) - .unwrap_or(0.0) + .unwrap_or((0.0, 0.0)) +} + +/// 用户可用余额(赠送 + 永久)——网关预检口径 +pub fn get_available_balance(conn: &Connection, user_id: i64) -> f64 { + let (permanent, gift) = get_balances(conn, user_id); + permanent + gift } /// 模型单价(按 provider+model)→ (input_per_m, output_per_m, currency) diff --git a/src/db.rs b/src/db.rs index d2c8208..1be471c 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 = 2; +pub const SCHEMA_VERSION: i64 = 3; /// 打开(或创建)数据库并执行幂等迁移 + dev 种子 pub fn open(path: &str) -> Result { @@ -79,8 +79,17 @@ pub fn migrate(conn: &Connection) -> Result<()> { id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL UNIQUE REFERENCES users(id), balance REAL NOT NULL DEFAULT 0, + gift_balance REAL NOT NULL DEFAULT 0, updated_at TEXT NOT NULL DEFAULT (datetime('now')) ); + CREATE TABLE IF NOT EXISTS gift_grants ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER NOT NULL REFERENCES users(id), + amount REAL NOT NULL DEFAULT 0, + granted_at TEXT NOT NULL DEFAULT (datetime('now')), + expires_at TEXT NOT NULL DEFAULT '', + status TEXT NOT NULL DEFAULT 'active' + ); CREATE TABLE IF NOT EXISTS transactions ( id INTEGER PRIMARY KEY AUTOINCREMENT, user_id INTEGER NOT NULL REFERENCES users(id), @@ -127,6 +136,23 @@ pub fn migrate(conn: &Connection) -> Result<()> { "available_end TEXT NOT NULL DEFAULT ''", )?; ensure_column(conn, "keys", "note", "note TEXT NOT NULL DEFAULT ''")?; + // v3(P1):点数账户拆分——gift_balance(当前有效赠送点数)+ gift_grants 明细表 + ensure_column( + conn, + "quotas", + "gift_balance", + "gift_balance REAL NOT NULL DEFAULT 0", + )?; + conn.execute_batch( + "CREATE TABLE IF NOT EXISTS gift_grants ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER NOT NULL REFERENCES users(id), + amount REAL NOT NULL DEFAULT 0, + granted_at TEXT NOT NULL DEFAULT (datetime('now')), + expires_at TEXT NOT NULL DEFAULT '', + status TEXT NOT NULL DEFAULT 'active' + );", + )?; // schema_version:INSERT OR REPLACE 保证幂等 let v: i64 = conn .query_row("SELECT version FROM schema_version", [], |r| r.get(0)) @@ -212,6 +238,27 @@ pub fn seed(conn: &Connection) -> Result<()> { VALUES ('deepseek', 'deepseek-paygo', 'deepseek-v4-flash', 'on', ?1, 'sk-placeholder-encrypted', 1000, 0)", [demo_id], )?; + + // 管理员账号(P1):admin@aitokenpool.local / admin1234,role=admin + let admin_id: Option = conn + .query_row( + "SELECT id FROM users WHERE email = ?1", + ["admin@aitokenpool.local"], + |r| r.get(0), + ) + .ok(); + if admin_id.is_none() { + let hash = hash_password("admin1234")?; + conn.execute( + "INSERT INTO users (email, password_hash, name, role) VALUES (?1, ?2, '管理员', 'admin')", + rusqlite::params!["admin@aitokenpool.local", hash], + )?; + let id = conn.last_insert_rowid(); + conn.execute( + "INSERT OR IGNORE INTO quotas (user_id, balance) VALUES (?1, 0)", + [id], + )?; + } Ok(()) } @@ -391,7 +438,7 @@ mod tests { let v: i64 = conn .query_row("SELECT version FROM schema_version", [], |r| r.get(0)) .unwrap(); - assert_eq!(v, 2, "旧库迁移后版本应为 2"); + assert_eq!(v, SCHEMA_VERSION, "旧库迁移后版本应为 {SCHEMA_VERSION}"); drop(conn); let _ = std::fs::remove_file(p); } diff --git a/src/gateway.rs b/src/gateway.rs index 369b014..6e606da 100644 --- a/src/gateway.rs +++ b/src/gateway.rs @@ -172,7 +172,9 @@ async fn forward( // 锁作用域严格限定在同步读区内,绝不在 await 期间持有 MutexGuard let (balance, keys) = { let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; - let balance = dao::get_balance(&conn, auth.user_id); + // P1:懒加载当日赠送(赠送也计入可用余额);可用余额 = gift + permanent + let _ = crate::gift::ensure_daily_gift(&conn, auth.user_id); + let balance = dao::get_available_balance(&conn, auth.user_id); let keys = dao::find_keys_by_model(&conn, model).map_err(internal)?; (balance, keys) }; @@ -283,7 +285,9 @@ async fn forward_stream( // 余额预检(与 forward 一致,锁作用域严格块内) let (balance, keys) = { let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; - let balance = dao::get_balance(&conn, auth.user_id); + // P1:懒加载当日赠送(赠送也计入可用余额);可用余额 = gift + permanent + let _ = crate::gift::ensure_daily_gift(&conn, auth.user_id); + let balance = dao::get_available_balance(&conn, auth.user_id); let keys = dao::find_keys_by_model(&conn, model).map_err(internal)?; (balance, keys) }; @@ -552,6 +556,13 @@ mod tests { let p = std::env::temp_dir().join(format!("atp_gw_{}_{}.db", std::process::id(), tag)); let _ = std::fs::remove_file(&p); let conn = crate::db::open(p.to_str().unwrap()).expect("open tmp db"); + // demo 注册时间拨到赠送窗口外(2020 年)→ 网关测试不触发每日赠送, + // 消费扣减断言确定(赠送路径由 gift/routes/billing 测试覆盖) + conn.execute( + "UPDATE users SET created_at = '2020-01-01 00:00:00' WHERE id = 1", + [], + ) + .unwrap(); let mut cfg = crate::config::Config::load("config/config.example.toml").unwrap(); cfg.plans.push(crate::config::Plan { id: plan_id.to_string(), diff --git a/src/gift.rs b/src/gift.rs new file mode 100644 index 0000000..983d56a --- /dev/null +++ b/src/gift.rs @@ -0,0 +1,343 @@ +//! 点数赠送规则(user-stories v1.7 → P1 落地) +//! +//! P1(rant 2026-08-18T11:03:02): +//! - 新人赠送:注册(users.created_at)起连续 10 天内每天 1 点,当日有效 +//! (expires_at = 当天 23:59:59);第 11 天起不再赠送 +//! - 懒加载触发:wallet / dashboard / settle 前调用 ensure_daily_gift +//! - 防薅羊毛:每人每日仅一笔;过期未用自动失效(查询时惰性清理) +//! - 消费扣减顺序:先扣最早到期的赠送点数,不足再扣永久 balance + +use anyhow::Result; +use rusqlite::Connection; + +/// 每日赠送点数 +pub const GIFT_DAILY_AMOUNT: f64 = 1.0; +/// 赠送窗口(注册起连续天数) +pub const GIFT_DAYS: i64 = 10; + +/// 惰性清理:把已过期的 active 赠送标记 expired,并重算 gift_balance +/// (重算 = 自愈:任何路径下 gift_balance 都能与 gift_grants 对齐) +pub fn expire_past_gifts(conn: &Connection, user_id: i64) -> Result<()> { + conn.execute( + "UPDATE gift_grants SET status = 'expired' \ + WHERE user_id = ?1 AND status = 'active' AND expires_at < datetime('now')", + [user_id], + )?; + conn.execute( + "UPDATE quotas SET gift_balance = COALESCE(( + SELECT SUM(amount) FROM gift_grants WHERE user_id = ?1 AND status = 'active' + ), 0) WHERE user_id = ?1", + [user_id], + )?; + Ok(()) +} + +/// 新人每日赠送(懒加载):在 10 天窗口内且今天未赠 → 补 1 点(当天 23:59:59 过期) +pub fn ensure_daily_gift(conn: &Connection, user_id: i64) -> Result { + // 用户必须存在(防薅:只绑定注册用户) + let created_at: Option = conn + .query_row( + "SELECT created_at FROM users WHERE id = ?1", + [user_id], + |r| r.get(0), + ) + .ok(); + let Some(created_at) = created_at else { + return Ok(false); + }; + + // 惰性清理过期赠送 + expire_past_gifts(conn, user_id)?; + + // 10 天窗口判定(注册日 = 第 1 天;julianday 差值 < 10) + let in_window: bool = conn + .query_row( + "SELECT julianday(date('now')) - julianday(date(?1)) < ?2", + rusqlite::params![created_at, GIFT_DAYS as f64], + |r| r.get(0), + ) + .unwrap_or(false); + if !in_window { + return Ok(false); + } + + // 今天已赠 → 跳过 + let today_granted: bool = conn + .query_row( + "SELECT EXISTS(SELECT 1 FROM gift_grants WHERE user_id = ?1 AND date(granted_at) = date('now'))", + [user_id], + |r| r.get(0), + ) + .unwrap_or(false); + if today_granted { + return Ok(false); + } + + // 补发:expires_at = 当天 23:59:59 + conn.execute( + "INSERT INTO gift_grants (user_id, amount, granted_at, expires_at, status) \ + VALUES (?1, ?2, datetime('now'), strftime('%Y-%m-%d 23:59:59', 'now'), 'active')", + rusqlite::params![user_id, GIFT_DAILY_AMOUNT], + )?; + conn.execute( + "INSERT OR IGNORE INTO quotas (user_id, balance, gift_balance) VALUES (?1, 0, 0)", + [user_id], + )?; + conn.execute( + "UPDATE quotas SET gift_balance = gift_balance + ?1 WHERE user_id = ?2", + rusqlite::params![GIFT_DAILY_AMOUNT, user_id], + )?; + Ok(true) +} + +/// 消费扣减:先扣最早到期的赠送点数(expires_at ASC),返回仍需从永久扣的剩余点数 +/// 在调用方事务内执行(tx:&Connection 兼容 rusqlite::Transaction) +pub fn deduct_gift_first(conn: &Connection, user_id: i64, mut pts: f64) -> Result { + // 惰性清理(确保不扣已过期) + expire_past_gifts(conn, user_id)?; + + let mut stmt = conn.prepare( + "SELECT id, amount FROM gift_grants \ + WHERE user_id = ?1 AND status = 'active' \ + ORDER BY expires_at ASC, id ASC", + )?; + let grants: Vec<(i64, f64)> = stmt + .query_map([user_id], |r| Ok((r.get(0)?, r.get(1)?)))? + .collect::>>()?; + drop(stmt); + + for (id, amount) in grants { + if pts <= 0.0 { + break; + } + let take = amount.min(pts); + let new_amount = amount - take; + conn.execute( + "UPDATE gift_grants SET amount = ?1, status = CASE WHEN ?1 <= 0 THEN 'used' ELSE 'active' END WHERE id = ?2", + rusqlite::params![new_amount, id], + )?; + pts -= take; + } + + // 重算 gift_balance(与剩余 active 对齐) + conn.execute( + "UPDATE quotas SET gift_balance = COALESCE(( + SELECT SUM(amount) FROM gift_grants WHERE user_id = ?1 AND status = 'active' + ), 0) WHERE user_id = ?1", + [user_id], + )?; + Ok(pts) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::db; + + fn tmp_db(tag: &str) -> (Connection, std::path::PathBuf) { + let p = std::env::temp_dir().join(format!("atp_gift_{}_{}.db", std::process::id(), tag)); + let _ = std::fs::remove_file(&p); + let conn = db::open(p.to_str().unwrap()).expect("open tmp db"); + (conn, p) + } + + /// 注册一个新用户(created_at 可控)并初始化配额 + fn register(conn: &Connection, email: &str, created: &str) -> i64 { + conn.execute( + "INSERT INTO users (email, password_hash, name, role, created_at) VALUES (?1, 'x', '新用户', 'user', ?2)", + rusqlite::params![email, created], + ) + .unwrap(); + let id = conn.last_insert_rowid(); + conn.execute( + "INSERT OR IGNORE INTO quotas (user_id, balance) VALUES (?1, 0)", + [id], + ) + .unwrap(); + id + } + + #[test] + fn first_gift_granted_with_today_expiry() { + let (conn, p) = tmp_db("g1"); + let uid = register(&conn, "u1@t.local", "2026-08-18 10:00:00"); // 今天注册 + let granted = ensure_daily_gift(&conn, uid).unwrap(); + assert!(granted, "注册当天应补发 1 点"); + let (bal, gift): (f64, f64) = conn + .query_row( + "SELECT balance, gift_balance FROM quotas WHERE user_id = ?1", + [uid], + |r| Ok((r.get(0)?, r.get(1)?)), + ) + .unwrap(); + assert_eq!(gift, 1.0); + assert_eq!(bal, 0.0); + let (expires, status, amount): (String, String, f64) = conn + .query_row( + "SELECT expires_at, status, amount FROM gift_grants WHERE user_id = ?1", + [uid], + |r| Ok((r.get(0)?, r.get(1)?, r.get(2)?)), + ) + .unwrap(); + assert_eq!(status, "active"); + assert_eq!(amount, 1.0); + assert!( + expires.ends_with("23:59:59"), + "当日有效(当天 23:59:59 过期): {expires}" + ); + drop(conn); + let _ = std::fs::remove_file(p); + } + + #[test] + fn same_day_no_duplicate_gift() { + let (conn, p) = tmp_db("g2"); + let uid = register(&conn, "u2@t.local", "2026-08-18 10:00:00"); + assert!(ensure_daily_gift(&conn, uid).unwrap()); + assert!(!ensure_daily_gift(&conn, uid).unwrap(), "同天不重复赠送"); + let n: i64 = conn + .query_row( + "SELECT COUNT(*) FROM gift_grants WHERE user_id = ?1", + [uid], + |r| r.get(0), + ) + .unwrap(); + assert_eq!(n, 1); + drop(conn); + let _ = std::fs::remove_file(p); + } + + #[test] + fn outside_window_no_gift() { + let (conn, p) = tmp_db("g3"); + // 11 天前注册 → 超出 10 天窗口 + let uid = register(&conn, "u3@t.local", "2026-08-07 09:00:00"); + let granted = ensure_daily_gift(&conn, uid).unwrap(); + assert!(!granted, "第 11 天起不再赠送"); + let n: i64 = conn + .query_row( + "SELECT COUNT(*) FROM gift_grants WHERE user_id = ?1", + [uid], + |r| r.get(0), + ) + .unwrap(); + assert_eq!(n, 0); + drop(conn); + let _ = std::fs::remove_file(p); + } + + #[test] + fn expired_gift_cleaned_lazily() { + let (conn, p) = tmp_db("g4"); + let uid = register(&conn, "u4@t.local", "2026-08-18 10:00:00"); + assert!(ensure_daily_gift(&conn, uid).unwrap()); + // 手工把 expires_at 改成过去 → 惰性清理应标记 expired 并扣 gift_balance + conn.execute( + "UPDATE gift_grants SET expires_at = '2026-08-01 00:00:00' WHERE user_id = ?1", + [uid], + ) + .unwrap(); + expire_past_gifts(&conn, uid).unwrap(); + let status: String = conn + .query_row( + "SELECT status FROM gift_grants WHERE user_id = ?1", + [uid], + |r| r.get(0), + ) + .unwrap(); + assert_eq!(status, "expired"); + let gift: f64 = conn + .query_row( + "SELECT gift_balance FROM quotas WHERE user_id = ?1", + [uid], + |r| r.get(0), + ) + .unwrap(); + assert_eq!(gift, 0.0, "过期赠送从 gift_balance 扣减"); + drop(conn); + let _ = std::fs::remove_file(p); + } + + #[test] + fn deduct_gift_first_oldest_expiry() { + let (conn, p) = tmp_db("g5"); + let uid = register(&conn, "u5@t.local", "2026-08-18 10:00:00"); + // 两笔 active 赠送:今天到期 1 点、明天到期 1 点(最早到期先扣) + conn.execute( + "INSERT INTO gift_grants (user_id, amount, granted_at, expires_at, status) \ + VALUES (?1, 1, '2026-08-18 10:00:00', '2026-08-18 23:59:59', 'active')", + [uid], + ) + .unwrap(); + conn.execute( + "INSERT INTO gift_grants (user_id, amount, granted_at, expires_at, status) \ + VALUES (?1, 1, '2026-08-18 10:00:00', '2026-08-19 23:59:59', 'active')", + [uid], + ) + .unwrap(); + conn.execute( + "UPDATE quotas SET gift_balance = 2 WHERE user_id = ?1", + [uid], + ) + .unwrap(); + // 扣 1.5 点:先花今天到期的 1 点(used)再花明天到期的 0.5 + let remaining = deduct_gift_first(&conn, uid, 1.5).unwrap(); + assert!(remaining.abs() < 1e-9, "赠送覆盖 1.5 点,剩余永久应为 0"); + let (st1, amt2): (String, f64) = conn + .query_row( + "SELECT (SELECT status FROM gift_grants WHERE user_id = ?1 ORDER BY expires_at ASC LIMIT 1), \ + (SELECT amount FROM gift_grants WHERE user_id = ?1 ORDER BY expires_at DESC LIMIT 1)", + [uid], + |r| Ok((r.get(0)?, r.get(1)?)), + ) + .unwrap(); + assert_eq!(st1, "used", "最早到期先扣"); + assert!( + (amt2 - 0.5).abs() < 1e-9, + "明天这笔剩 0.5 仍 active: {amt2}" + ); + let gift: f64 = conn + .query_row( + "SELECT gift_balance FROM quotas WHERE user_id = ?1", + [uid], + |r| r.get(0), + ) + .unwrap(); + assert!((gift - 0.5).abs() < 1e-9, "gift_balance 与剩余对齐: {gift}"); + drop(conn); + let _ = std::fs::remove_file(p); + } + + #[test] + fn deduct_overflow_falls_to_permanent() { + let (conn, p) = tmp_db("g6"); + let uid = register(&conn, "u6@t.local", "2026-08-18 10:00:00"); + conn.execute( + "INSERT INTO gift_grants (user_id, amount, granted_at, expires_at, status) \ + VALUES (?1, 1, '2026-08-18 10:00:00', '2026-08-18 23:59:59', 'active')", + [uid], + ) + .unwrap(); + conn.execute( + "UPDATE quotas SET gift_balance = 1, balance = 10 WHERE user_id = ?1", + [uid], + ) + .unwrap(); + // 扣 3 点:gift 1 点 + 永久 2 点 + let remaining = deduct_gift_first(&conn, uid, 3.0).unwrap(); + assert!( + (remaining - 2.0).abs() < 1e-9, + "剩余 2 点从永久扣: {remaining}" + ); + let gift: f64 = conn + .query_row( + "SELECT gift_balance FROM quotas WHERE user_id = ?1", + [uid], + |r| r.get(0), + ) + .unwrap(); + assert_eq!(gift, 0.0); + drop(conn); + let _ = std::fs::remove_file(p); + } +} diff --git a/src/main.rs b/src/main.rs index 1cf622f..fdad655 100644 --- a/src/main.rs +++ b/src/main.rs @@ -17,6 +17,7 @@ mod crypto; mod dao; mod db; mod gateway; +mod gift; mod router; mod routes; diff --git a/src/routes/admin.rs b/src/routes/admin.rs new file mode 100644 index 0000000..1902790 --- /dev/null +++ b/src/routes/admin.rs @@ -0,0 +1,162 @@ +//! 管理员 API(对齐原型管理视图「成员充值/部门」) +//! +//! P1(rant 2026-08-18T11:03:02): +//! - POST /api/admin/credits {user_id, amount, note?}:给成员永久点数充值(role=admin) +//! - GET /api/admin/users:成员列表(id/email/name/balance/gift_balance/role) +//! - GET /api/admin/usage:用量报表(每用户本月 tokens/点数/调用次数) +//! - 权限:require_admin(AuthUser.role == "admin",否则 403) + +use axum::extract::State; +use axum::Json; +use rusqlite::params; +use serde::Deserialize; + +use crate::routes::{internal, ApiErr, AppState, AuthUser}; + +/// 充值请求 +#[derive(Debug, Deserialize)] +pub struct CreditReq { + pub user_id: i64, + pub amount: f64, + #[serde(default)] + pub note: String, +} + +/// 权限中间件判定(处理函数内调用) +fn require_admin(auth: &AuthUser) -> Result<(), ApiErr> { + if auth.role == "admin" { + Ok(()) + } else { + Err(( + axum::http::StatusCode::FORBIDDEN, + Json(serde_json::json!({ "error": "需要管理员权限" })), + )) + } +} + +/// POST /api/admin/credits:给成员永久点数充值 + 写 transactions(type=topup) +pub async fn credits( + State(st): State, + auth: AuthUser, + Json(req): Json, +) -> Result, ApiErr> { + require_admin(&auth)?; + if req.amount <= 0.0 { + return Err(( + axum::http::StatusCode::BAD_REQUEST, + Json(serde_json::json!({ "error": "amount 必须大于 0" })), + )); + } + let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; + // 目标用户存在? + let exists: bool = conn + .query_row( + "SELECT EXISTS(SELECT 1 FROM users WHERE id = ?1)", + [req.user_id], + |r| r.get(0), + ) + .unwrap_or(false); + if !exists { + return Err(( + axum::http::StatusCode::NOT_FOUND, + Json(serde_json::json!({ "error": "用户不存在" })), + )); + } + conn.execute( + "INSERT OR IGNORE INTO quotas (user_id, balance) VALUES (?1, 0)", + [req.user_id], + ) + .map_err(internal)?; + conn.execute( + "UPDATE quotas SET balance = balance + ?1, updated_at = datetime('now') WHERE user_id = ?2", + params![req.amount, req.user_id], + ) + .map_err(internal)?; + conn.execute( + "INSERT INTO transactions (user_id, counterpart, key_id, model, tokens, pts, type, status) \ + VALUES (?1, ?2, NULL, 'recharge', 0, ?3, 'topup', '成功')", + params![req.user_id, auth.user_id.to_string(), req.amount], + ) + .map_err(internal)?; + let balance: f64 = conn + .query_row( + "SELECT balance FROM quotas WHERE user_id = ?1", + [req.user_id], + |r| r.get(0), + ) + .unwrap_or(0.0); + Ok(Json(serde_json::json!({ + "user_id": req.user_id, + "amount": req.amount, + "balance": balance, + "note": req.note, + }))) +} + +/// GET /api/admin/users:成员列表 +pub async fn users( + State(st): State, + auth: AuthUser, +) -> Result>, ApiErr> { + require_admin(&auth)?; + let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; + let mut stmt = conn + .prepare( + "SELECT u.id, u.email, u.name, u.role, COALESCE(q.balance, 0), COALESCE(q.gift_balance, 0) \ + FROM users u LEFT JOIN quotas q ON q.user_id = u.id ORDER BY u.id", + ) + .map_err(internal)?; + let rows = stmt + .query_map([], |r| { + Ok(serde_json::json!({ + "id": r.get::<_, i64>(0)?, + "email": r.get::<_, String>(1)?, + "name": r.get::<_, String>(2)?, + "role": r.get::<_, String>(3)?, + "balance": r.get::<_, f64>(4)?, + "gift_balance": r.get::<_, f64>(5)?, + })) + }) + .map_err(internal)?; + let mut out = Vec::new(); + for r in rows { + out.push(r.map_err(internal)?); + } + Ok(Json(out)) +} + +/// GET /api/admin/usage:用量报表(每用户本月 tokens/点数/调用次数) +pub async fn usage( + State(st): State, + auth: AuthUser, +) -> Result>, ApiErr> { + require_admin(&auth)?; + let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; + let mut stmt = conn + .prepare( + "SELECT u.id, u.email, u.name, \ + COALESCE(SUM(ur.tokens), 0), COALESCE(SUM(ur.cost), 0), COUNT(ur.id) \ + FROM users u \ + LEFT JOIN usage_records ur ON ur.user_id = u.id \ + AND strftime('%Y-%m', ur.time) = strftime('%Y-%m', 'now') \ + GROUP BY u.id ORDER BY u.id", + ) + .map_err(internal)?; + let rows = stmt + .query_map([], |r| { + Ok(serde_json::json!({ + "id": r.get::<_, i64>(0)?, + "email": r.get::<_, String>(1)?, + "name": r.get::<_, String>(2)?, + "month_tokens": r.get::<_, f64>(3)?, + "month_cost": r.get::<_, f64>(4)?, + "month_calls": r.get::<_, i64>(5)?, + })) + }) + .map_err(internal)?; + let mut out = Vec::new(); + for r in rows { + out.push(r.map_err(internal)?); + } + Ok(Json(out)) +} diff --git a/src/routes/mod.rs b/src/routes/mod.rs index ff76acf..17443ab 100644 --- a/src/routes/mod.rs +++ b/src/routes/mod.rs @@ -9,6 +9,7 @@ //! - POST /v1/chat/completions / POST /anthropic/v1/messages(网关) //! - GET /api/models(市场) +pub mod admin; pub mod api_keys; pub mod sharing; pub mod wallet; @@ -72,10 +73,12 @@ pub fn internal(e: impl std::fmt::Display) -> ApiErr { } /// 已认证用户(Bearer 提取器):无效 key → 401 -#[derive(Debug, Clone, Copy)] +#[derive(Debug, Clone)] pub struct AuthUser { pub user_id: i64, pub api_key_id: i64, + /// 用户角色:user | admin(P1 起 require_admin 使用) + pub role: String, } #[axum::async_trait] @@ -94,11 +97,12 @@ impl FromRequestParts for AuthUser { let key = header.strip_prefix("Bearer ").ok_or_else(unauthorized)?; let conn = state.db.lock().map_err(|_| internal("db lock poisoned"))?; match dao::find_api_key_user_and_id(&conn, key) { - Some((user_id, api_key_id)) => { + Some((user_id, api_key_id, role)) => { let _ = dao::touch_api_key(&conn, key); Ok(AuthUser { user_id, api_key_id, + role, }) } None => Err(unauthorized()), @@ -150,6 +154,10 @@ pub fn router() -> Router { .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)) + // P1:管理员(充值 / 成员列表 / 用量报表) + .route("/api/admin/credits", post(admin::credits)) + .route("/api/admin/users", get(admin::users)) + .route("/api/admin/usage", get(admin::usage)) } #[cfg(test)] @@ -303,4 +311,122 @@ mod tests { let (s, _) = get(test_state("nobearer"), "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/api/api-keys", None).await; assert_eq!(s, StatusCode::UNAUTHORIZED); } + + /// 登录并返回 Bearer + async fn login_bearer(st: &AppState, email: &str, password: &str) -> String { + let (s, body) = post( + st.clone(), + "/api/auth/login", + &format!(r#"{{"email":"{email}","password":"{password}"}}"#), + None, + ) + .await; + assert_eq!(s, StatusCode::OK, "登录应成功: {body}"); + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + v["api_key"].as_str().unwrap().to_string() + } + + #[tokio::test] + async fn admin_credits_rejects_non_admin() { + let st = test_state("admin403"); + let demo_bearer = login_bearer(&st, "demo@aitokenpool.local", "demo1234").await; + let (s, body) = post( + st.clone(), + "/api/admin/credits", + r#"{"user_id":2,"amount":50}"#, + Some(&demo_bearer), + ) + .await; + assert_eq!(s, StatusCode::FORBIDDEN, "非 admin 应 403: {body}"); + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + assert!(v["error"].as_str().is_some(), "返回错误信息"); + } + + #[tokio::test] + async fn admin_credits_recharges_permanent_and_writes_topup() { + let st = test_state("admincred"); + let admin_bearer = login_bearer(&st, "admin@aitokenpool.local", "admin1234").await; + // 充值前余额 + let (s, body) = get(st.clone(), "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/api/wallet", Some(&admin_bearer)).await; + assert_eq!(s, StatusCode::OK, "admin 可看钱包: {body}"); + // 给 demo(user_id=1)充 50 点 + let (s, body) = post( + st.clone(), + "/api/admin/credits", + r#"{"user_id":1,"amount":50,"note":"P1 test"}"#, + Some(&admin_bearer), + ) + .await; + assert_eq!(s, StatusCode::OK, "充值应成功: {body}"); + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + assert_eq!(v["amount"], 50.0); + assert_eq!(v["balance"], 12471.0 + 50.0, "demo 余额增加 50"); + // transactions 出现 topup 记录(demo 视角) + let demo_bearer = login_bearer(&st, "demo@aitokenpool.local", "demo1234").await; + let (s, body) = get( + st.clone(), + "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/api/transactions?type=topup", + Some(&demo_bearer), + ) + .await; + assert_eq!(s, StatusCode::OK, "topup 过滤应 200: {body}"); + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + let arr = v["items"].as_array().expect("items 为数组"); + assert_eq!(arr.len(), 1, "一条 topup 记录: {body}"); + assert_eq!(arr[0]["pts"], 50.0); + // 管理员给自己的充值非法 amount → 400 + let (s, body) = post( + st.clone(), + "/api/admin/credits", + r#"{"user_id":1,"amount":-5}"#, + Some(&admin_bearer), + ) + .await; + assert_eq!(s, StatusCode::BAD_REQUEST, "负数金额应 400: {body}"); + } + + #[tokio::test] + async fn admin_users_and_usage_lists() { + let st = test_state("adminlist"); + let admin_bearer = login_bearer(&st, "admin@aitokenpool.local", "admin1234").await; + // users:demo + admin 都在列表 + let (s, body) = get(st.clone(), "/api/admin/users", Some(&admin_bearer)).await; + assert_eq!(s, StatusCode::OK, "users 应 200: {body}"); + let arr: Vec = serde_json::from_str(&body).unwrap(); + assert_eq!(arr.len(), 2, "demo + admin: {body}"); + assert!(arr.iter().any(|u| u["role"] == "admin"), "admin 在列表中"); + assert!( + arr.iter().any(|u| u["email"] == "demo@aitokenpool.local"), + "demo 在列表中" + ); + // usage:每用户本月聚合 + let (s, body) = get(st.clone(), "/api/admin/usage", Some(&admin_bearer)).await; + assert_eq!(s, StatusCode::OK, "usage 应 200: {body}"); + let arr: Vec = serde_json::from_str(&body).unwrap(); + assert_eq!(arr.len(), 2); + assert!( + arr.iter().all(|u| u["month_tokens"] == 0.0), + "无调用时 tokens 为 0: {body}" + ); + // 非 admin 访问 users → 403 + let demo_bearer = login_bearer(&st, "demo@aitokenpool.local", "demo1234").await; + let (s, _) = get(st.clone(), "/api/admin/users", Some(&demo_bearer)).await; + assert_eq!(s, StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn wallet_shows_daily_gift_balance() { + let st = test_state("wallet_gift"); + let demo_bearer = login_bearer(&st, "demo@aitokenpool.local", "demo1234").await; + // demo 今天注册(seed 默认 created_at=now)→ 10 天窗口内 → 首次访问 wallet 补发 1 点 + let (s, body) = get(st.clone(), "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/api/wallet", Some(&demo_bearer)).await; + assert_eq!(s, StatusCode::OK, "wallet 应 200: {body}"); + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + assert_eq!(v["gift_balance"], 1.0, "当日赠送 1 点: {body}"); + assert_eq!(v["balance"], 12471.0, "永久余额不变: {body}"); + // 重复访问不重复赠送 + let (_, body2) = get(st.clone(), "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/api/wallet", Some(&demo_bearer)).await; + let v2: serde_json::Value = serde_json::from_str(&body2).unwrap(); + assert_eq!(v2["gift_balance"], 1.0, "同天不重复: {body2}"); + } } diff --git a/src/routes/wallet.rs b/src/routes/wallet.rs index 1ee765e..6179f17 100644 --- a/src/routes/wallet.rs +++ b/src/routes/wallet.rs @@ -10,6 +10,7 @@ use axum::Json; use rusqlite::params; use serde::Deserialize; +use crate::dao; use crate::routes::{internal, ApiErr, AppState, AuthUser}; /// GET /api/wallet @@ -18,13 +19,9 @@ pub async fn wallet( 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); + // P1:懒加载当日赠送(新人每日 1 点,10 天窗口) + let _ = crate::gift::ensure_daily_gift(&conn, auth.user_id); + let (balance, gift_balance) = dao::get_balances(&conn, auth.user_id); let month_use: f64 = conn .query_row( "SELECT COALESCE(SUM(pts), 0) FROM transactions \ @@ -43,6 +40,8 @@ pub async fn wallet( .unwrap_or(0.0); Ok(Json(serde_json::json!({ "balance": balance, + "gift_balance": gift_balance, + "available": balance + gift_balance, "month_use": month_use, "month_earn": month_earn, }))) @@ -77,11 +76,11 @@ pub async fn transactions( 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()), + t @ ("consume" | "earn" | "topup") => Some(t.to_string()), _ => { return Err(( axum::http::StatusCode::BAD_REQUEST, - Json(serde_json::json!({ "error": "type 必须为 consume / earn / all" })), + Json(serde_json::json!({ "error": "type 必须为 consume / earn / topup / all" })), )) } };