diff --git a/.gitignore b/.gitignore index 7471e8d..1535f3e 100644 --- a/.gitignore +++ b/.gitignore @@ -3,3 +3,4 @@ # macOS .DS_Store +/data diff --git a/Cargo.lock b/Cargo.lock index e484a0c..cf08457 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -25,7 +25,7 @@ dependencies = [ [[package]] name = "aitokenpool" -version = "0.1.0" +version = "0.2.0" dependencies = [ "anyhow", "argon2", @@ -44,7 +44,8 @@ dependencies = [ "sha2", "thiserror", "tokio", - "tower 0.4.13", + "toml", + "tower", "tower-http 0.5.2", "uuid", ] @@ -177,7 +178,7 @@ dependencies = [ "serde_urlencoded", "sync_wrapper", "tokio", - "tower 0.5.3", + "tower", "tower-layer", "tower-service", "tracing", @@ -1468,7 +1469,7 @@ dependencies = [ "tokio-native-tls", "tokio-rustls", "tokio-util", - "tower 0.5.3", + "tower", "tower-http 0.6.11", "tower-service", "url", @@ -1659,6 +1660,15 @@ dependencies = [ "serde_core", ] +[[package]] +name = "serde_spanned" +version = "0.6.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bf41e0cfaf7226dca15e8197172c295a782857fcb97fad1808a166870dee75a3" +dependencies = [ + "serde", +] + [[package]] name = "serde_urlencoded" version = "0.7.1" @@ -1921,16 +1931,46 @@ dependencies = [ ] [[package]] -name = "tower" -version = "0.4.13" +name = "toml" +version = "0.8.23" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b8fa9be0de6cf49e536ce1851f987bd21a43b771b09473c3549a6c853db37c1c" +checksum = "dc1beb996b9d83529a9e75c17a1686767d148d70663143c7854d8b4a09ced362" dependencies = [ - "tower-layer", - "tower-service", - "tracing", + "serde", + "serde_spanned", + "toml_datetime", + "toml_edit", +] + +[[package]] +name = "toml_datetime" +version = "0.6.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "22cddaf88f4fbc13c51aebbf5f8eceb5c7c5a9da2ac40a13519eb5b0a0e8f11c" +dependencies = [ + "serde", ] +[[package]] +name = "toml_edit" +version = "0.22.27" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41fe8c660ae4257887cf66394862d21dbca4a6ddd26f04a3560410406a2f819a" +dependencies = [ + "indexmap", + "serde", + "serde_spanned", + "toml_datetime", + "toml_write", + "winnow", +] + +[[package]] +name = "toml_write" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5d99f8c9a7727884afe522e9bd5edbfc91a3312b36a77b5fb8926e4c31a41801" + [[package]] name = "tower" version = "0.5.3" @@ -1984,7 +2024,7 @@ dependencies = [ "http", "http-body", "pin-project-lite", - "tower 0.5.3", + "tower", "tower-layer", "tower-service", "url", @@ -2363,6 +2403,15 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" +[[package]] +name = "winnow" +version = "0.7.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df79d97927682d2fd8adb29682d1140b343be4ac0f08fd68b7765d9c059d3945" +dependencies = [ + "memchr", +] + [[package]] name = "writeable" version = "0.6.3" diff --git a/Cargo.toml b/Cargo.toml index 3fd6b31..5fcbc0d 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "aitokenpool" -version = "0.1.0" +version = "0.2.0" edition = "2021" description = "AI Token 共享池 — 企业 key 池 + 公共共享市场" license = "MIT" @@ -11,8 +11,8 @@ rust-version = "1.86" # HTTP 服务器(复用 openlocalrouter 技术栈) axum = "0.7" tokio = { version = "1", features = ["macros", "rt-multi-thread", "time", "signal"] } -tower = "0.4" -tower-http = { version = "0.5", features = ["cors", "fs"] } +tower = { version = "0.5", features = ["util"] } +tower-http = { version = "0.5", features = ["cors", "fs", "trace"] } hyper = { version = "1.0", features = ["full"] } # HTTP 客户端(上游转发) @@ -21,6 +21,7 @@ reqwest = { version = "0.12", features = ["rustls-tls", "json", "stream"] } # 序列化 serde = { version = "1.0", features = ["derive"] } serde_json = { version = "1.0" } +toml = "0.8" # 数据库 rusqlite = { version = "0.31", features = ["bundled"] } diff --git a/README.md b/README.md index f3b218f..b0f4498 100644 --- a/README.md +++ b/README.md @@ -17,7 +17,11 @@ AITokenPool 是一个开源的 **AI Token 共享平台**:企业版(内部 ke ## 状态 -🚧 项目初始化中(2026-08-13) +- ✅ **P0-A(v0.2.0,2026-08-17)**:后端骨架 + 配置加载(`config/config.example.toml`)+ SQLite 数据层(幂等迁移 + demo 种子)+ 认证(argon2 + Bearer API Key)+ API Key 端点。`cargo run` 后: + - `GET /healthz` → `{"status":"ok","version":"0.2.0"}` + - `POST /api/auth/login`(demo@aitokenpool.local / demo1234)→ `{api_key}` + - `POST|GET /api/api-keys`(Bearer 认证,key 脱敏 `atk_live_****xxxx`) +- 🚧 P0 后续(网关路由 / 用量追踪)进行中;UI 原型 v1.20(`ui/` 静态页 + mock 数据)。 ## License diff --git a/config/config.example.toml b/config/config.example.toml index f6016c4..e678f52 100644 --- a/config/config.example.toml +++ b/config/config.example.toml @@ -6,6 +6,13 @@ # 由脚本从 OpenRouter/litellm 同步生成,本文件只放「官方价覆盖」。 # - 分层理由见 docs/plan-api-matrix.md 与架构讨论。 +# ============================================================ +# 0. 服务(监听地址 / 数据库路径)— P0-A 起可配 +# ============================================================ +[server] +addr = "0.0.0.0:8080" +db_path = "data/aitokenpool.db" + # ============================================================ # 1. 点数规则(账本层的锚) # ============================================================ diff --git a/src/auth.rs b/src/auth.rs new file mode 100644 index 0000000..0d6c022 --- /dev/null +++ b/src/auth.rs @@ -0,0 +1,70 @@ +//! 认证:argon2 口令哈希 + API Key 生成/校验 +//! +//! P0-A(rant 2026-08-17T22:21:52): +//! - POST /api/auth/login:email+password → argon2 校验 → 返回该用户有效 API Key(无则生成) +//! - Bearer 认证:查 api_keys 表 → 注入用户身份;无效 401 + +use anyhow::{anyhow, Result}; +use argon2::password_hash::{ + rand_core::OsRng, PasswordHash, PasswordHasher, PasswordVerifier, SaltString, +}; +use argon2::Argon2; +use rand::RngCore; + +/// argon2 口令哈希(OWASP 默认参数:m=19MiB, t=2, p=1) +pub fn hash_password(pw: &str) -> Result { + let salt = SaltString::generate(&mut OsRng); + Ok(Argon2::default() + .hash_password(pw.as_bytes(), &salt) + .map_err(|e| anyhow!("argon2 哈希失败: {e}"))? + .to_string()) +} + +/// 校验口令是否匹配存储哈希 +pub fn verify_password(hash: &str, pw: &str) -> bool { + let Ok(parsed) = PasswordHash::new(hash) else { + return false; + }; + Argon2::default() + .verify_password(pw.as_bytes(), &parsed) + .is_ok() +} + +/// 生成分发 API Key:`atk_live_` + 24 位 hex(12 随机字节),与 UI 原型一致 +pub fn generate_api_key() -> String { + let mut bytes = [0u8; 12]; + rand::rngs::OsRng.fill_bytes(&mut bytes); + format!("atk_live_{}", hex::encode(bytes)) +} + +/// API Key 脱敏展示:atk_live_****xxxx(保留后 4 位) +pub fn mask_api_key(key: &str) -> String { + if key.len() > 8 { + let tail = &key[key.len() - 4..]; + format!("atk_live_****{tail}") + } else { + "****".to_string() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn hash_verify_roundtrip() { + let h = hash_password("demo1234").unwrap(); + assert!(verify_password(&h, "demo1234")); + assert!(!verify_password(&h, "wrong")); + } + + #[test] + fn api_key_format_and_mask() { + let k = generate_api_key(); + assert!(k.starts_with("atk_live_")); + assert_eq!(k.len(), "atk_live_".len() + 24); + let m = mask_api_key(&k); + assert_eq!(m, format!("atk_live_****{}", &k[k.len() - 4..])); + assert!(!m.contains(&k[..k.len() - 4])); + } +} diff --git a/src/config.rs b/src/config.rs new file mode 100644 index 0000000..ccad88f --- /dev/null +++ b/src/config.rs @@ -0,0 +1,176 @@ +//! 配置结构(与 config/config.example.toml 一一对应) +//! +//! 设计约定(见 config/config.example.toml 注释): +//! - providers / plans / 点数规则是「人手工维护」的配置,需可读、可注释; +//! - 模型价格大表在 data/models.json,本文件只放「官方价覆盖」price_overrides。 +//! +//! P0-A(rant 2026-08-17T22:21:52):服务骨架 + 配置加载 + +use serde::Deserialize; + +fn default_addr() -> String { + "0.0.0.0:8080".to_string() +} +fn default_db_path() -> String { + "data/aitokenpool.db".to_string() +} + +/// 服务(监听 / 数据库路径)——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, +} + +impl Default for Server { + fn default() -> Self { + Server { + addr: default_addr(), + db_path: default_db_path(), + } + } +} + +/// 顶层配置 +#[derive(Debug, Clone, Deserialize)] +#[allow(dead_code)] // P0-A 仅用 server;points/providers/plans 由后续 P0 网关/定价阶段消费(parse 测试已校验) +pub struct Config { + #[serde(default)] + pub server: Server, + pub points: Points, + pub providers: Vec, + pub plans: Vec, + #[serde(default)] + pub price_overrides: Vec, +} + +/// 点数规则(账本层的锚) +#[derive(Debug, Clone, Deserialize)] +#[allow(dead_code)] +pub struct Points { + /// 货币锚:USD | CNY + pub anchor_currency: String, + /// 1 个单位锚定货币 = 多少「点」 + pub points_per_unit: u32, + /// 显示名(仅 UI) + pub display_name: String, + /// 符号(仅 UI) + pub symbol: String, +} + +/// 提供商(一家模型厂商) +#[derive(Debug, Clone, Deserialize)] +#[allow(dead_code)] +pub struct Provider { + pub id: String, + pub name: String, + pub country: String, + pub has_plan: bool, +} + +/// Plan 端点(一个可被路由到的上游端点) +#[derive(Debug, Clone, Deserialize)] +#[allow(dead_code)] +pub struct Plan { + pub id: String, + pub provider: String, + /// paygo | token | coding + #[serde(rename = "type")] + pub type_: String, + /// key 前缀约定(错配会 401) + pub key_prefix: String, + #[serde(default)] + pub interactive_only: bool, + pub endpoints: Vec, +} + +#[derive(Debug, Clone, Deserialize)] +#[allow(dead_code)] +pub struct Endpoint { + /// openai_chat | anthropic | responses + pub protocol: String, + pub base_url: String, +} + +/// 官方价覆盖(覆盖 data/models.json 聚合源价格) +#[derive(Debug, Clone, Deserialize)] +#[allow(dead_code)] +pub struct PriceOverride { + pub provider: String, + pub model: String, + pub currency: String, + pub input_per_m: f64, + pub output_per_m: f64, + #[serde(default)] + pub cache_hit_input_per_m: Option, + #[serde(default)] + pub source: Option, +} + +impl Config { + /// 从 TOML 文件加载配置 + pub fn load(path: &str) -> anyhow::Result { + let s = std::fs::read_to_string(path)?; + let cfg: Config = toml::from_str(&s)?; + Ok(cfg) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parse_config_example_ok() { + let cfg = + Config::load("config/config.example.toml").expect("解析 config.example.toml 应成功"); + // 点数 + assert_eq!(cfg.points.anchor_currency, "USD"); + assert_eq!(cfg.points.points_per_unit, 1000); + assert_eq!(cfg.points.display_name, "点数"); + assert_eq!(cfg.points.symbol, "P"); + // providers + assert_eq!(cfg.providers.len(), 6); + assert!(cfg + .providers + .iter() + .any(|p| p.id == "deepseek" && !p.has_plan)); + assert!(cfg + .providers + .iter() + .any(|p| p.id == "zhipu" && p.has_plan && p.country == "CN")); + // plans + assert!(cfg.plans.len() >= 7); + let dp = cfg + .plans + .iter() + .find(|p| p.id == "deepseek-paygo") + .expect("deepseek-paygo 应存在"); + assert_eq!(dp.provider, "deepseek"); + assert_eq!(dp.type_, "paygo"); + assert_eq!(dp.key_prefix, "sk-"); + assert_eq!(dp.endpoints.len(), 3); + assert_eq!(dp.endpoints[0].protocol, "openai_chat"); + assert_eq!(dp.endpoints[0].base_url, "https://api.deepseek.com"); + let al = cfg + .plans + .iter() + .find(|p| p.id == "aliyun-token-plan") + .expect("aliyun-token-plan 应存在"); + assert!(al.interactive_only); + // price_overrides + assert!(cfg.price_overrides.len() >= 2); + let dv = cfg + .price_overrides + .iter() + .find(|o| o.model == "deepseek-v4-pro") + .unwrap(); + assert_eq!(dv.input_per_m, 0.435); + assert!(dv.source.is_some()); + // server 默认值 + assert_eq!(cfg.server.addr, "0.0.0.0:8080"); + assert_eq!(cfg.server.db_path, "data/aitokenpool.db"); + } +} diff --git a/src/dao.rs b/src/dao.rs new file mode 100644 index 0000000..023a3bf --- /dev/null +++ b/src/dao.rs @@ -0,0 +1,99 @@ +//! 数据访问层(rusqlite 直查,简单分层) +//! +//! P0-A(rant 2026-08-17T22:21:52):认证 + API Key 管理所需的最小查询集。 + +use anyhow::{anyhow, Result}; +use rusqlite::Connection; + +use crate::auth; + +/// 按邮箱查用户 → (id, password_hash) +pub fn find_user_by_email(conn: &Connection, email: &str) -> Option<(i64, String)> { + conn.query_row( + "SELECT id, password_hash FROM users WHERE email = ?1", + [email], + |r| Ok((r.get(0)?, r.get(1)?)), + ) + .ok() +} + +/// 取该用户的「有效」API Key;没有则生成一个(登录即返回可用的 key) +pub fn get_or_create_api_key(conn: &Connection, user_id: i64) -> Result { + let existing: Option = conn + .query_row( + "SELECT key_value FROM api_keys WHERE user_id = ?1 AND status = 'active' ORDER BY id LIMIT 1", + [user_id], + |r| r.get(0), + ) + .ok(); + if let Some(k) = existing { + return Ok(k); + } + let key = auth::generate_api_key(); + conn.execute( + "INSERT INTO api_keys (user_id, key_value, name, status) VALUES (?1, ?2, '', 'active')", + rusqlite::params![user_id, key], + )?; + Ok(key) +} + +/// 生成新 API Key(设置页「生成新 Key」) +pub fn create_api_key(conn: &Connection, user_id: i64, name: &str) -> Result { + let key = auth::generate_api_key(); + conn.execute( + "INSERT INTO api_keys (user_id, key_value, name, status) VALUES (?1, ?2, ?3, 'active')", + rusqlite::params![user_id, key, name], + )?; + Ok(key) +} + +/// 列出用户的 API Key(key 值脱敏) +pub fn list_api_keys(conn: &Connection, user_id: i64) -> Result> { + let mut stmt = conn.prepare( + "SELECT id, key_value, name, status, created_at FROM api_keys WHERE user_id = ?1 ORDER BY id DESC", + )?; + let rows = stmt.query_map([user_id], |r| { + let raw: String = r.get(1)?; + Ok(serde_json::json!({ + "id": r.get::<_, i64>(0)?, + "key": auth::mask_api_key(&raw), + "name": r.get::<_, String>(2)?, + "status": r.get::<_, String>(3)?, + "created_at": r.get::<_, String>(4)?, + })) + })?; + let mut out = Vec::new(); + for r in rows { + out.push(r?); + } + Ok(out) +} + +/// Bearer 认证:按 key 查归属用户 → Some(user_id) +pub fn find_user_by_api_key(conn: &Connection, key: &str) -> Option { + conn.query_row( + "SELECT user_id FROM api_keys WHERE key_value = ?1 AND status = 'active'", + [key], + |r| r.get(0), + ) + .ok() +} + +/// 更新最近使用时间 +pub fn touch_api_key(conn: &Connection, key: &str) -> Result<()> { + conn.execute( + "UPDATE api_keys SET last_used = datetime('now') WHERE key_value = ?1", + [key], + )?; + Ok(()) +} + +/// 校验口令(供登录用) +pub fn verify_user_password(conn: &Connection, email: &str, pw: &str) -> Result { + let (id, hash) = find_user_by_email(conn, email).ok_or_else(|| anyhow!("用户不存在"))?; + if auth::verify_password(&hash, pw) { + Ok(id) + } else { + Err(anyhow!("口令错误")) + } +} diff --git a/src/db.rs b/src/db.rs new file mode 100644 index 0000000..936ef21 --- /dev/null +++ b/src/db.rs @@ -0,0 +1,203 @@ +//! SQLite 数据层(rusqlite bundled) +//! +//! 表结构对齐 docs/architecture.md §5 + P0-A rant(2026-08-17T22:21:52): +//! users / keys(上游 key)/ api_keys(分发 key,atk_live_ 前缀)/ models / +//! quotas(点数账户)/ transactions / usage_records + schema_version 迁移表。 +//! +//! 迁移策略:CREATE TABLE IF NOT EXISTS 幂等;schema_version 记录当前版本, +//! 重复启动不报错。 + +use anyhow::{Context, Result}; +use rusqlite::Connection; + +pub const SCHEMA_VERSION: i64 = 1; + +/// 打开(或创建)数据库并执行幂等迁移 + dev 种子 +pub fn open(path: &str) -> Result { + if let Some(dir) = std::path::Path::new(path).parent() { + if !dir.as_os_str().is_empty() { + std::fs::create_dir_all(dir) + .with_context(|| format!("创建数据目录失败: {}", dir.display()))?; + } + } + let conn = Connection::open(path).with_context(|| format!("打开数据库失败: {}", path))?; + migrate(&conn)?; + seed(&conn)?; + Ok(conn) +} + +/// 幂等迁移:建表 + schema_version +pub fn migrate(conn: &Connection) -> Result<()> { + conn.execute_batch( + r#" + CREATE TABLE IF NOT EXISTS schema_version ( + version INTEGER NOT NULL + ); + CREATE TABLE IF NOT EXISTS users ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + email TEXT NOT NULL UNIQUE, + password_hash TEXT NOT NULL, + name TEXT NOT NULL DEFAULT '', + role TEXT NOT NULL DEFAULT 'user', + created_at TEXT NOT NULL DEFAULT (datetime('now')) + ); + CREATE TABLE IF NOT EXISTS 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 REFERENCES users(id), + 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')) + ); + CREATE TABLE IF NOT EXISTS api_keys ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER NOT NULL REFERENCES users(id), + key_value TEXT NOT NULL UNIQUE, + name TEXT NOT NULL DEFAULT '', + status TEXT NOT NULL DEFAULT 'active', + created_at TEXT NOT NULL DEFAULT (datetime('now')), + last_used TEXT + ); + CREATE TABLE IF NOT EXISTS models ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + provider TEXT NOT NULL, + model TEXT NOT NULL, + currency TEXT NOT NULL DEFAULT 'USD', + input_per_m REAL NOT NULL DEFAULT 0, + output_per_m REAL NOT NULL DEFAULT 0, + updated_at TEXT NOT NULL DEFAULT (datetime('now')) + ); + CREATE TABLE IF NOT EXISTS quotas ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER NOT NULL UNIQUE REFERENCES users(id), + balance REAL NOT NULL DEFAULT 0, + updated_at TEXT NOT NULL DEFAULT (datetime('now')) + ); + CREATE TABLE IF NOT EXISTS transactions ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER NOT NULL REFERENCES users(id), + counterpart TEXT NOT NULL DEFAULT '', + key_id INTEGER, + model TEXT NOT NULL DEFAULT '', + tokens REAL NOT NULL DEFAULT 0, + pts REAL NOT NULL DEFAULT 0, + type TEXT NOT NULL DEFAULT '', + status TEXT NOT NULL DEFAULT '成功', + time TEXT NOT NULL DEFAULT (datetime('now')) + ); + CREATE TABLE IF NOT EXISTS usage_records ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id INTEGER NOT NULL REFERENCES users(id), + api_key_id INTEGER, + key_id INTEGER, + model TEXT NOT NULL DEFAULT '', + tokens REAL NOT NULL DEFAULT 0, + cost REAL NOT NULL DEFAULT 0, + time TEXT NOT NULL DEFAULT (datetime('now')) + ); + "#, + )?; + // schema_version:INSERT OR REPLACE 保证幂等 + let v: i64 = conn + .query_row("SELECT version FROM schema_version", [], |r| r.get(0)) + .unwrap_or(0); + if v < SCHEMA_VERSION { + conn.execute( + "INSERT OR REPLACE INTO schema_version (version) VALUES (?1)", + [SCHEMA_VERSION], + )?; + } + Ok(()) +} + +/// dev 种子:demo 用户(demo@aitokenpool.local / demo1234,argon2)+ 点数账户 + 示例上游 key +pub fn seed(conn: &Connection) -> Result<()> { + use crate::auth::hash_password; + + let demo_id: Option = conn + .query_row( + "SELECT id FROM users WHERE email = ?1", + ["demo@aitokenpool.local"], + |r| r.get(0), + ) + .ok(); + let demo_id = match demo_id { + Some(id) => id, + None => { + let hash = hash_password("demo1234")?; + conn.execute( + "INSERT INTO users (email, password_hash, name, role) VALUES (?1, ?2, '阿零', 'user')", + rusqlite::params!["demo@aitokenpool.local", hash], + )?; + conn.last_insert_rowid() + } + }; + // 点数账户(seed 余额 12471,对齐 UI mock D.USER.balance) + conn.execute( + "INSERT OR IGNORE INTO quotas (user_id, balance) VALUES (?1, 12471)", + [demo_id], + )?; + // 示例上游 key(不真实可用的占位:占位密钥 + deepseek paygo plan) + conn.execute( + "INSERT OR IGNORE INTO keys (provider, plan, model, status, owner_id, encrypted_key, quota, used) \ + VALUES ('deepseek', 'deepseek-paygo', 'deepseek-v4-flash', 'on', ?1, 'sk-placeholder-encrypted', 1000, 0)", + [demo_id], + )?; + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn tmp_db(tag: &str) -> (Connection, std::path::PathBuf) { + let p = std::env::temp_dir().join(format!("atp_test_{}_{}.db", std::process::id(), tag)); + let _ = std::fs::remove_file(&p); + let conn = open(p.to_str().unwrap()).expect("open tmp db"); + (conn, p) + } + + #[test] + fn migrate_is_idempotent() { + let (conn, p) = tmp_db("migrate"); + // 二次迁移不报错(重复启动场景) + migrate(&conn).expect("第二次 migrate 应成功"); + migrate(&conn).expect("第三次 migrate 应成功"); + let v: i64 = conn + .query_row("SELECT version FROM schema_version", [], |r| r.get(0)) + .unwrap(); + assert_eq!(v, SCHEMA_VERSION); + drop(conn); + let _ = std::fs::remove_file(p); + } + + #[test] + fn seed_demo_user_and_quota() { + let (conn, p) = tmp_db("seed"); + let u: i64 = conn + .query_row( + "SELECT id FROM users WHERE email = 'demo@aitokenpool.local'", + [], + |r| r.get(0), + ) + .expect("demo 用户已种子"); + let bal: f64 = conn + .query_row("SELECT balance FROM quotas WHERE user_id = ?1", [u], |r| { + r.get(0) + }) + .expect("demo 点数账户已种子"); + assert_eq!(bal, 12471.0); + let n: i64 = conn + .query_row("SELECT COUNT(*) FROM keys WHERE owner_id = ?1", [u], |r| { + r.get(0) + }) + .unwrap(); + assert!(n >= 1, "示例上游 key 已种子"); + drop(conn); + let _ = std::fs::remove_file(p); + } +} diff --git a/src/main.rs b/src/main.rs index 7b8ff13..eb30754 100644 --- a/src/main.rs +++ b/src/main.rs @@ -1,33 +1,63 @@ -//! AITokenPool — AI Token 共享池 +//! AITokenPool — AI Token 共享池(网关 + 账本) //! //! 企业版:内部 key 池 + 员工点数配额 //! 公共版:分享闲置 key 赚点数、消费别人 key //! //! 架构定论见 docs/architecture.md(中心化方案 A:平台托管 key + 平台执行调用) +//! +//! P0-A(rant 2026-08-17T22:21:52):服务骨架 + 配置加载 + SQLite 数据层 + 认证。 +//! 启动:cargo run -- --config config/config.toml(默认 config/config.toml, +//! 不存在时提示复制 config.example.toml)。 -/// 返回项目横幅文本(便于测试与展示) -pub fn banner() -> String { - "AITokenPool — AI Token 共享池\n\ - 企业版:key 池 + 员工点数配额 · 公共版:共享市场\n\ - 架构:中心化(方案 A),Rust + axum\n\ - 状态:项目初始化中(2026-08-13)" - .to_string() -} +mod auth; +mod config; +mod dao; +mod db; +mod routes; + +use std::sync::Arc; + +use clap::Parser; -fn main() { - println!("{}", banner()); +#[derive(Parser)] +#[command( + name = "aitokenpool", + version, + about = "AI Token 共享池 — 企业 key 池 + 公共共享市场" +)] +struct Args { + /// 配置文件路径(默认 config/config.toml) + #[arg(long, default_value = "config/config.toml")] + config: String, } -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn banner_contains_project_identity() { - let b = banner(); - assert!(b.contains("AITokenPool")); - assert!(b.contains("企业版")); - assert!(b.contains("公共版")); - assert!(b.contains("中心化")); - } +#[tokio::main] +async fn main() -> anyhow::Result<()> { + env_logger::init(); + let args = Args::parse(); + + let cfg = match config::Config::load(&args.config) { + Ok(c) => c, + Err(e) => { + eprintln!("加载配置失败: {e}"); + eprintln!("提示: 请先复制示例配置:"); + eprintln!(" cp config/config.example.toml config/config.toml"); + std::process::exit(1); + } + }; + + let addr = cfg.server.addr.clone(); + let db_path = cfg.server.db_path.clone(); + log::info!("打开数据库: {db_path}"); + let conn = db::open(&db_path)?; + + let state = routes::AppState::new(conn, Arc::new(cfg)); + let app = routes::router() + .with_state(state) + .layer(tower_http::trace::TraceLayer::new_for_http()); + + let listener = tokio::net::TcpListener::bind(&addr).await?; + log::info!("AITokenPool 服务已启动: http://{addr}"); + axum::serve(listener, app).await?; + Ok(()) } diff --git a/src/routes/api_keys.rs b/src/routes/api_keys.rs new file mode 100644 index 0000000..3e64739 --- /dev/null +++ b/src/routes/api_keys.rs @@ -0,0 +1,34 @@ +//! API Key 管理端点(Bearer 认证) +//! +//! P0-A(rant 2026-08-17T22:21:52): +//! - POST /api/api-keys:生成(atk_live_ + 24 hex,与 UI 原型一致) +//! - GET /api/api-keys:列表(key 脱敏 atk_live_****xxxx) + +use axum::extract::State; +use axum::Json; + +use crate::auth; +use crate::routes::{internal, ApiErr, AppState, AuthUser}; + +/// POST /api/api-keys +pub async fn create( + State(st): State, + auth: AuthUser, +) -> Result, ApiErr> { + let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; + let key = crate::dao::create_api_key(&conn, auth.user_id, "").map_err(internal)?; + Ok(Json(serde_json::json!({ + "api_key": key, + "masked": auth::mask_api_key(&key), + }))) +} + +/// GET /api/api-keys(脱敏列表) +pub async fn list( + State(st): State, + auth: AuthUser, +) -> Result>, ApiErr> { + let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; + let keys = crate::dao::list_api_keys(&conn, auth.user_id).map_err(internal)?; + Ok(Json(keys)) +} diff --git a/src/routes/mod.rs b/src/routes/mod.rs new file mode 100644 index 0000000..0897a76 --- /dev/null +++ b/src/routes/mod.rs @@ -0,0 +1,274 @@ +//! HTTP 路由装配:healthz + 认证 + API Key 管理 +//! +//! P0-A(rant 2026-08-17T22:21:52): +//! - GET /healthz → {"status":"ok","version":"0.2.0"} +//! - POST /api/auth/login → 200 {api_key} / 401 +//! - POST /api/api-keys / GET /api/api-keys(Bearer 认证) + +pub mod api_keys; + +use std::sync::{Arc, Mutex}; + +use axum::extract::{FromRequestParts, State}; +use axum::http::request::Parts; +use axum::http::StatusCode; +use axum::routing::{get, post}; +use axum::{Json, Router}; +use rusqlite::Connection; +use serde::Deserialize; + +use crate::config::Config; +use crate::dao; + +/// 共享状态:数据库连接 + 配置 +#[derive(Clone)] +pub struct AppState { + pub db: Arc>, + #[allow(dead_code)] // 供后续 P0 网关路由阶段消费 + pub cfg: Arc, +} + +impl AppState { + pub fn new(conn: Connection, cfg: Arc) -> Self { + Self { + db: Arc::new(Mutex::new(conn)), + cfg, + } + } +} + +/// 统一错误响应类型 +pub type ApiErr = (StatusCode, Json); + +fn unauthorized() -> ApiErr { + ( + StatusCode::UNAUTHORIZED, + Json(serde_json::json!({ "error": "unauthorized" })), + ) +} + +fn internal(e: impl std::fmt::Display) -> ApiErr { + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({ "error": format!("{e}") })), + ) +} + +/// 已认证用户(Bearer 提取器):无效 key → 401 +#[derive(Debug, Clone, Copy)] +pub struct AuthUser { + pub user_id: i64, +} + +#[axum::async_trait] +impl FromRequestParts for AuthUser { + type Rejection = ApiErr; + + async fn from_request_parts( + parts: &mut Parts, + state: &AppState, + ) -> Result { + let header = parts + .headers + .get(axum::http::header::AUTHORIZATION) + .and_then(|v| v.to_str().ok()) + .unwrap_or(""); + let key = header.strip_prefix("Bearer ").ok_or_else(unauthorized)?; + let conn = state.db.lock().map_err(|_| internal("db lock poisoned"))?; + match dao::find_user_by_api_key(&conn, key) { + Some(user_id) => { + let _ = dao::touch_api_key(&conn, key); + Ok(AuthUser { user_id }) + } + None => Err(unauthorized()), + } + } +} + +/// GET /healthz +pub async fn healthz() -> Json { + Json(serde_json::json!({ + "status": "ok", + "version": env!("CARGO_PKG_VERSION"), + })) +} + +#[derive(Deserialize)] +pub struct LoginReq { + pub email: String, + pub password: String, +} + +/// POST /api/auth/login:email+password → 200 {api_key} / 401 +pub async fn login( + State(st): State, + Json(req): Json, +) -> Result, ApiErr> { + let conn = st.db.lock().map_err(|_| internal("db lock poisoned"))?; + let user_id = + dao::verify_user_password(&conn, &req.email, &req.password).map_err(|_| unauthorized())?; + let api_key = dao::get_or_create_api_key(&conn, user_id).map_err(internal)?; + Ok(Json(serde_json::json!({ + "api_key": api_key, + "user_id": user_id, + }))) +} + +/// 组装路由 +pub fn router() -> Router { + Router::new() + .route("/healthz", get(healthz)) + .route("/api/auth/login", post(login)) + .route("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/api/api-keys", post(api_keys::create).get(api_keys::list)) +} + +#[cfg(test)] +mod tests { + use super::*; + 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_route_{}_{}.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(); + AppState::new(conn, Arc::new(cfg)) + } + + async fn post( + state: AppState, + uri: &str, + body: &str, + bearer: Option<&str>, + ) -> (StatusCode, String) { + let mut b = Request::builder() + .method("POST") + .uri(uri) + .header("content-type", "application/json"); + if let Some(k) = bearer { + b = b.header("authorization", format!("Bearer {k}")); + } + let resp = router() + .with_state(state) + .oneshot(b.body(Body::from(body.to_string())).unwrap()) + .await + .unwrap(); + let status = resp.status(); + let bytes = axum::body::to_bytes(resp.into_body(), 1024 * 1024) + .await + .unwrap(); + (status, String::from_utf8(bytes.to_vec()).unwrap()) + } + + async fn get(state: AppState, uri: &str, bearer: Option<&str>) -> (StatusCode, String) { + let mut b = Request::builder().method("GET").uri(uri); + if let Some(k) = bearer { + b = b.header("authorization", format!("Bearer {k}")); + } + let resp = router() + .with_state(state) + .oneshot(b.body(Body::empty()).unwrap()) + .await + .unwrap(); + let status = resp.status(); + let bytes = axum::body::to_bytes(resp.into_body(), 1024 * 1024) + .await + .unwrap(); + (status, String::from_utf8(bytes.to_vec()).unwrap()) + } + + #[tokio::test] + async fn healthz_ok() { + let (s, body) = get(test_state("healthz"), "/healthz", None).await; + assert_eq!(s, StatusCode::OK); + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + assert_eq!(v["status"], "ok"); + assert_eq!(v["version"], "0.2.0"); + } + + #[tokio::test] + async fn login_ok_and_returns_key() { + let st = test_state("login"); + let (s, body) = post( + st.clone(), + "/api/auth/login", + r#"{"email":"demo@aitokenpool.local","password":"demo1234"}"#, + None, + ) + .await; + assert_eq!(s, StatusCode::OK, "正确口令应 200: {body}"); + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + let key = v["api_key"].as_str().expect("返回 api_key"); + assert!(key.starts_with("atk_live_")); + // 再次登录返回同一 key(get-or-create) + let (_, body2) = post( + st, + "/api/auth/login", + r#"{"email":"demo@aitokenpool.local","password":"demo1234"}"#, + None, + ) + .await; + let v2: serde_json::Value = serde_json::from_str(&body2).unwrap(); + assert_eq!(v2["api_key"].as_str().unwrap(), key, "重复登录应复用 key"); + } + + #[tokio::test] + async fn login_wrong_password_401() { + let (s, _) = post( + test_state("login401"), + "/api/auth/login", + r#"{"email":"demo@aitokenpool.local","password":"wrong"}"#, + None, + ) + .await; + assert_eq!(s, StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn api_keys_create_and_list_with_bearer() { + let st = test_state("keys"); + // 登录拿 key + let (_, body) = post( + st.clone(), + "/api/auth/login", + r#"{"email":"demo@aitokenpool.local","password":"demo1234"}"#, + None, + ) + .await; + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + let bearer = v["api_key"].as_str().unwrap().to_string(); + // POST 生成 + let (s, body) = post(st.clone(), "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/api/api-keys", "{}", Some(&bearer)).await; + assert_eq!(s, StatusCode::OK, "有效 Bearer 生成 key: {body}"); + let v: serde_json::Value = serde_json::from_str(&body).unwrap(); + let new_key = v["api_key"].as_str().unwrap(); + assert!(new_key.starts_with("atk_live_")); + // GET 列表脱敏 + let (s, body) = get(st.clone(), "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/api/api-keys", Some(&bearer)).await; + assert_eq!(s, StatusCode::OK); + let arr: Vec = serde_json::from_str(&body).unwrap(); + assert!(arr.len() >= 2, "列表含登录 key + 新生成 key"); + assert!( + arr.iter() + .all(|k| k["key"].as_str().unwrap().contains("****")), + "key 全部脱敏" + ); + // 无效 Bearer → 401 + let (s, _) = get( + st, + "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/api/api-keys", + Some("atk_live_000000000000000000000000"), + ) + .await; + assert_eq!(s, StatusCode::UNAUTHORIZED); + } + + #[tokio::test] + async fn api_keys_without_bearer_401() { + let (s, _) = get(test_state("nobearer"), "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/api/api-keys", None).await; + assert_eq!(s, StatusCode::UNAUTHORIZED); + } +}