Compare commits

..
Author SHA1 Message Date
stefanfeng 434caac056 feat(avatar): honor square interaction authorization 2026-09-08 13:19:50 +08:00
stefanfeng 2a01a9946a Merge pull request 'fix(avatar): 修复知识库重新索引操作' (#19) from codex/avatar-reindex-action-20260908 into main
Reviewed-on: #19
2026-09-08 09:19:32 +08:00
stefanfeng 2ce1079bb6 fix(avatar): repair knowledge reindex action 2026-09-08 09:06:33 +08:00
stefanfeng 6fba6dbaaa Merge pull request 'fix(avatar): 分片上传大文件知识库' (#18) from codex/avatar-chunk-upload-20260907 into main
Reviewed-on: #18
2026-09-07 17:48:15 +08:00
stefanfeng e71267cf86 fix(avatar): upload knowledge files in chunks 2026-09-07 17:45:22 +08:00
stefanfeng 359e558dbe Merge pull request 'feat(avatar): 多文件知识库上传与进度展示' (#17) from codex/avatar-upload-progress-20260904 into main
Reviewed-on: #17
2026-09-04 17:33:05 +08:00
stefanfeng 3edf92c7cc feat(avatar): show multi-file knowledge upload progress 2026-09-04 16:40:16 +08:00
stefanfeng 97c4c73b58 Merge pull request 'fix(avatar): 异步知识库索引并修复大文件上传' (#16) from codex/avatar-knowledge-async-20260904 into main
Reviewed-on: #16
2026-09-04 15:59:11 +08:00
stefanfeng 08c58fe0e6 fix(avatar): cap knowledge files at 50MB 2026-09-04 15:18:39 +08:00
stefanfeng 6b7201e890 fix(avatar): align knowledge upload limit with production 2026-09-04 13:56:29 +08:00
stefanfeng 28553aba15 fix(avatar): allow knowledge uploads up to 20MB 2026-09-04 13:47:46 +08:00
stefanfeng 95f91450d0 Merge pull request 'fix(avatar): index knowledge documents asynchronously' (#15) from codex/avatar-knowledge-async-20260904 into main
Reviewed-on: #15
2026-09-04 11:56:01 +08:00
stefanfeng b98a2b9507 fix(avatar): index knowledge documents asynchronously 2026-09-04 11:53:56 +08:00
stefanfeng 59350fb41d Merge pull request 'fix(avatar): ground replies in recognized images' (#14) from codex/avatar-image-answer-hotfix-20260902 into main
Reviewed-on: #14
2026-09-02 14:49:15 +08:00
stefanfeng 6a4b35c49a fix(avatar): ground replies in recognized images 2026-09-02 14:47:27 +08:00
stefanfeng 207bbd02cf Merge pull request 'fix(avatar): recover BOXIM image replies' (#13) from codex/avatar-boxim-vision-hotfix-20260902 into main
Reviewed-on: #13
2026-09-02 14:13:19 +08:00
stefanfeng 7cac96356d fix(avatar): recover BOXIM image replies 2026-09-02 14:12:04 +08:00
stefanfeng 03c32309a8 Merge pull request 'feat(avatar): understand BOXIM image messages' (#12) from codex/avatar-boxim-vision-20260901 into main
Merge pull request #12: BOXIM image understanding
2026-09-01 15:36:19 +08:00
stefanfeng 0fc43908ae feat(avatar): understand BOXIM image messages 2026-09-01 15:35:41 +08:00
stefanfeng 3d999f9472 Merge pull request 'fix(avatar): prevent BOXIM polling starvation' (#11) from codex/avatar-takeover-poll-20260901 into main 2026-09-01 14:04:18 +08:00
stefanfeng 540edb58c4 fix(avatar): prevent BOXIM polling starvation 2026-09-01 14:03:35 +08:00
stefanfeng a7eb6ac2a5 Merge pull request #10 from codex/avatar-vision-chat-20260831
feat(avatar): 支持图片与病例理解对话
2026-09-01 10:41:51 +08:00
stefanfeng 016bc22c05 feat(avatar): add private vision chat support 2026-08-31 15:10:50 +08:00
stefanfeng 094f8cd40f Merge pull request 'fix(avatar): normalize embedding API endpoint' (#9) from codex/avatar-embedding-endpoint-20260828 into main
Reviewed-on: #9
2026-08-28 17:00:16 +08:00
stefanfeng 6794e88d53 fix(avatar): normalize embedding API endpoint 2026-08-28 16:59:22 +08:00
stefanfeng 7884430b3d Merge pull request 'feat(avatar): 优化会会 H5 嵌入管理流程' (#8) from codex/avatar-h5-embedded-layout-20260827 into main
Reviewed-on: #8
2026-08-27 11:41:21 +08:00
stefanfeng 46d42b7d98 fix(avatar): prevent takeover loops and isolate settings 2026-08-27 11:30:12 +08:00
stefanfeng ef58c5f2d2 feat(avatar): optimize embedded H5 management flow 2026-08-27 09:28:05 +08:00
stefanfeng 6e3fe5a616 Merge pull request 'feat(avatar): 接入会会支付并统一积分展示' (#7) from codex/avatar-token-copy-to-points-20260826 into main
Reviewed-on: #7
2026-08-26 14:53:58 +08:00
stefanfeng 67b6bd1b48 Merge pull request 'fix(avatar): 部署重启后自动恢复 BOXIM 接管' (#6) from codex/avatar-takeover-restart-safe-20260826 into main
Reviewed-on: #6
2026-08-26 14:53:49 +08:00
stefanfeng 0752001d85 Merge pull request 'feat(avatar): 数字分身自动跟随用户语言回答' (#5) from codex/avatar-auto-reply-language-20260826 into main
Reviewed-on: #5
2026-08-26 14:53:40 +08:00
stefanfeng c37294be17 feat(avatar): integrate Huihui payments 2026-08-26 14:41:21 +08:00
stefanfeng 5360cac8ad fix(avatar): normalize legacy plan copy 2026-08-26 14:02:38 +08:00
stefanfeng ea932f27fb feat(avatar): rename Token display to points 2026-08-26 14:01:12 +08:00
deploy bf0183bef4 add wechat mini-program domain verification file 2026-08-26 13:44:56 +08:00
stefanfeng 0c6419f37e fix(avatar): restore BOXIM takeover after restart 2026-08-26 13:25:16 +08:00
stefanfeng f768e7648f fix(avatar): prevent inferred reply scenarios 2026-08-26 11:56:54 +08:00
stefanfeng e30ab2b889 feat(avatar): follow user language in replies 2026-08-26 11:52:32 +08:00
stefanfeng 730f586784 Merge pull request 'docs(avatar): 配置 digital.99hui.com 生产域名' (#4) from codex/avatar-multi-management-integrated-20260825 into main
Reviewed-on: #4
2026-08-26 10:17:40 +08:00
stefanfeng ef2b1c6dd6 chore(avatar): add WeChat verification file 2026-08-26 09:51:13 +08:00
stefanfeng c693899b12 docs(avatar): set production H5 domain 2026-08-26 09:44:22 +08:00
stefanfeng d274ccb5e2 Merge pull request 'feat: 完成数字分身多分身管理与生产 H5 接入' (#3) from codex/avatar-multi-management-integrated-20260825 into main
Reviewed-on: #3
2026-08-26 09:43:52 +08:00
stefanfeng 5d19992f00 fix(avatar): remove SSO token from router state 2026-08-26 09:25:07 +08:00
stefanfeng 81aec1c63a feat(avatar): support production H5 token SSO 2026-08-26 09:17:41 +08:00
stefanfeng 4029c31ed7 fix(avatar): bundle uni bridge and add favicon 2026-08-25 17:19:30 +08:00
stefanfeng e24e89d326 fix(deploy): serialize model migration and pin nginx 2026-08-25 16:59:15 +08:00
stefanfeng dc34a03357 feat(ai): add dedicated digital avatar model config 2026-08-25 16:50:33 +08:00
stefanfeng 3f7ff9329a fix(avatar): align QA cards to the left 2026-08-25 15:36:49 +08:00
stefanfeng 699bbbde57 feat(avatar): add user token accounting 2026-08-25 13:24:02 +08:00
stefanfeng 7a0199e685 feat(avatar): improve multi-avatar management 2026-08-25 11:44:27 +08:00
stefanfeng 672019830d fix(avatar): prevent takeover replies blocking across chats 2026-08-21 16:21:00 +08:00
stefanfeng e720baa21e fix(avatar): batch knowledge embedding requests 2026-08-21 13:21:05 +08:00
stefanfeng 9e86cc64ac Merge pull request 'Codex/avatar integrated 20260819' (#2) from codex/avatar-integrated-20260819 into main
Reviewed-on: #2
2026-08-21 09:31:43 +08:00
stefanfeng e2b928273c fix(avatar): acknowledge BOXIM messages as read 2026-08-20 17:10:54 +08:00
stefanfeng 64b7680ec4 fix(avatar): keep BOXIM status ready during long polls 2026-08-20 15:58:07 +08:00
stefanfeng 89f52963b7 fix(avatar): persist takeover toggle immediately 2026-08-20 15:55:12 +08:00
stefanfeng 76bd22c24b fix(avatar): preserve ownership across Huihui environments 2026-08-20 15:27:29 +08:00
stefanfeng 51a317ccd9 fix(avatar): fail closed on BOXIM connection errors 2026-08-20 14:47:13 +08:00
stefanfeng 08590bf9ea fix(avatar): refresh BOXIM takeover status 2026-08-19 18:03:59 +08:00
stefanfeng 25fb8fbee5 feat(avatar): add BOXIM chat takeover 2026-08-19 17:56:57 +08:00
stefanfeng cfcfe7146e feat(avatar): align authorization page with design 2026-08-19 17:05:22 +08:00
stefanfeng 2e2adeb9e2 feat(avatar): complete authorization management 2026-08-19 16:21:19 +08:00
stefanfeng 25d2494616 fix: surface missing knowledge files 2026-08-19 15:13:22 +08:00
stefanfeng bd5f64d000 fix: align avatar route titles 2026-08-19 15:09:38 +08:00
stefanfeng 350df1d119 build: pin secure frontend transitive dependencies 2026-08-19 15:01:39 +08:00
stefanfeng 64462fac92 build: enforce locked frontend type checks 2026-08-19 14:57:40 +08:00
stefanfeng 6bf446f889 fix: restore reproducible frontend builds 2026-08-19 14:56:39 +08:00
stefanfeng 68ea87e1b2 fix: persist avatar data and uploads 2026-08-19 14:44:58 +08:00
stefanfeng c5bfa47a23 fix: await takeover polling jobs 2026-08-19 14:32:10 +08:00
stefanfeng 4ab7732da9 fix: preserve takeover and share token integrity 2026-08-19 14:23:43 +08:00
stefanfeng 4a2d788e85 feat: complete grounded digital avatar chat experience 2026-08-19 14:21:56 +08:00
stefanfeng 07d4a21379 fix: sync Huihui app user homepage profile 2026-08-18 08:48:20 +08:00
stefanfeng efc419e301 fix: sync virtual user profiles to Huihui public pages 2026-08-18 08:48:20 +08:00
stefanfeng c748950402 fix: accept camelCase authorizationId from frontend 2026-08-07 17:18:34 +08:00
stefanfeng f13d488f82 feat: add takeover config API endpoint
PUT /api/avatar/{avatar_id}/authorizations/takeover to update authorization
takeover settings (enabled, mode, delay_seconds) with validation.
2026-08-07 17:15:12 +08:00
stefanfeng 43f2ad9edb feat: add takeover config UI to authorization management page
- Added updateTakeoverConfig API function (PUT /api/avatar/{id}/authorizations/takeover)
- Added expandable takeover config section per auth card with toggle switch, mode selector (immediate/delayed), and delay seconds input
- Added takeover summary badge showing current config when not editing
- Styles follow existing light theme with indigo (#6366F1) accent
2026-08-07 17:07:21 +08:00
stefanfeng 8e365e63e1 feat: add scheduled message polling for avatar takeover
Add APScheduler BackgroundScheduler to on_startup that polls
takeover_service.poll_and_process_messages every 10 seconds. Added
redis and apscheduler dependencies. Added poll_and_process_messages,
fetch_unread_messages, and process_message methods to TakeoverService.
Scheduler initialization is graceful — Redis/Box IM failures do not
prevent app startup.
2026-08-07 17:01:04 +08:00
stefanfeng dbd668e451 fix: resolve owner through Avatar model in takeover service
- Fix C1: check_takeover_enabled now uses owner_huihui_id to query Avatar first
- Fix C2: execute_takeover resolves owner_id via Avatar model (Authorization has no owner_id)
- Fix I1: add test for None case (no Avatar found)
- Fix I2: add test verifying filter arguments
- Fix I3: implement process_delayed_queue to scan and dispatch expired messages
2026-08-07 16:55:25 +08:00
stefanfeng 190ec48d9c feat: add takeover service core logic
Implement TakeoverService for avatar takeover chat feature:
- check_takeover_enabled: query Authorization for enabled takeover
- generate_reply: call local chat API via httpx to generate responses
- execute_takeover: get owner IM credentials, generate reply, send via BoxIM
- enqueue_delayed_message: Redis delayed queue with graceful degradation
- process_delayed_queue: placeholder for scheduled worker

Service degrades gracefully when Redis is unavailable.
2026-08-07 16:46:57 +08:00
stefanfeng 0cc43a0d8c fix: exclude signature from signing string in BoxIMClient
Prevents stale signature values from leaking into the MD5 signing
calculation, matching the news_service.py pattern. Adds a test that
passes a stale signature in extra params and verifies the returned
signature is freshly computed.
2026-08-07 16:42:00 +08:00
stefanfeng 43bb57eacf feat: add Box IM client for Netease Yunxin integration
Add BoxIMClient with get_credentials and send_p2p_message methods,
using Huihui platform MD5 signing mechanism (sorted keys, 24h timestamp,
12-char alphanumeric nonce). Includes full test coverage.
2026-08-07 16:36:50 +08:00
stefanfeng 7327939b4d feat: add takeover fields to Authorization model
Add takeover_enabled, takeover_mode, and takeover_delay_seconds columns
to the Authorization model with a SQLite migration, and include them
in to_dict(). Test: test_takeover_model.py.
2026-08-07 16:30:29 +08:00
stefanfeng 76f0730b57 docs: add avatar takeover chat implementation plan 2026-08-07 16:22:58 +08:00
stefanfeng 6f6655e791 docs: add avatar takeover chat design spec 2026-08-07 16:12:49 +08:00
stefanfeng e3cce871b0 fix: add avatar photo upload endpoint to digital-avatar-app backend
- Add POST /api/avatar/{id}/photo endpoint
- Save uploaded photos to UPLOAD_DIR/{avatar_id}/
- Update avatar.photo_url in database
- Import UPLOAD_DIR from knowledge router
2026-08-07 15:10:59 +08:00
stefanfeng b7f8ac203b feat: add avatar photo upload via Huihui filecenter
- Add POST /avatars/{id}/upload-photo endpoint
- Use Huihui platform's filecenter API for image upload (same signing mechanism as news_service)
- Add upload button in avatar list table
- Frontend: handleUploadPhoto function with el-upload component
- Updates avatar photo_url in SQLite after successful upload
2026-08-07 14:28:22 +08:00
stefanfeng a5f7c86a4b fix: avatar list table full-width and image fallback
- Changed fixed width columns to min-width for elastic layout
- Add @error handler on el-avatar to hide broken images and show emoji
- Detail dialog avatar also handles image load errors
2026-08-07 14:04:45 +08:00
stefanfeng 6026b1279e fix: avatar list fixes — full-width table, absolute photo URLs, global token balance
- Table width set to 100%
- Photo URLs resolved to absolute paths using AVATAR_BACKEND_URL
- Token balance now reads from global token_account table
- Config: added AVATAR_BACKEND_URL setting
2026-08-07 13:53:02 +08:00
stefanfeng 57af1c7db1 fix: gracefully handle missing avatar SQLite database
- Return empty list instead of crashing when avatar.db doesn't exist
- Add is_available() check before each endpoint
- 503 status when trying to access unavailable database
2026-08-06 17:52:01 +08:00
stefanfeng bbbc97eb4a feat: add digital avatar management to admin dashboard
- Add avatars list page with pagination, search, status filter
- Add avatar detail dialog showing full info and config params
- Add on/off switch (active/inactive status toggle)
- Backend: avatar_service.py connects to digital avatar SQLite
- Backend: /api/avatars endpoints for list, detail, status update
- Sidebar menu item added below dashboard
- docker-compose: mount SQLite db and set AVATAR_DB_PATH env var
2026-08-06 17:43:50 +08:00
stefanfeng 1da5f7f10d fix: open logged-in users on avatar management 2026-07-24 17:54:02 +08:00
stefanfeng cc939b6298 fix: guard avatar entry and hide chat navigation 2026-07-24 17:50:19 +08:00
stefanfeng e3bda469bb feat: complete huihui square avatar workflows 2026-07-24 14:04:21 +08:00
stefanfeng 2ef58a44b8 chore: normalize avatar app whitespace 2026-07-23 17:29:29 +08:00
stefanfeng 2d9f26a6f0 feat: add avatar chat and knowledge workflow 2026-07-23 17:21:49 +08:00
stefanfeng 501f548bcc docs: add avatar chat implementation plan 2026-07-23 16:27:02 +08:00
stefanfeng cc8af846d5 docs: design avatar chat knowledge flow 2026-07-23 16:06:30 +08:00
stefanfeng 79d57da769 fix: 今日文章不足时混入历史文章,避免所有用户扎堆同一篇
问题:今日只有1篇文章时,所有虚拟用户都只互动这同一篇
原因:Phase 1 找到今日文章后直接返回,不管数量多少

修复逻辑:
- 今日有效文章 → 每用户最多取 1 篇(count//3,最少1篇)
- 剩余名额(count - today_quota)从历史文章补充
- 历史文章:按当前小时对应页拉取(与Phase 2相同),随机打散
- 历史文章排除今日文章ID和当天发布的文章,保证内容不重复
- 最终返回:今日N篇 + 历史M篇,总量接近 count

效果:今日1篇文章时 → 用户取1篇今日 + 4篇历史,互动多样性恢复
2026-04-08 11:49:57 +08:00
stefanfeng c944fbb0ea fix: 今日文章配额控制,避免全部虚拟用户集中互动同一篇
问题:今日只有1篇文章时,所有虚拟用户全部互动该文章,历史文章无人问津

修复方案(配额制):
- 新增 count_today_articles():轻量统计今日广场文章数
- 配额规则:每篇今日文章最多吸引3个虚拟用户(可调)
  - 今日1篇 → 最多3人互动今日,其余全走历史
  - 今日5篇 → 最多15人互动今日,其余走历史
  - 今日10篇以上 → 批次内所有人均可互动今日文章
- get_news_list() 新增 force_history 参数,强制走 Phase 2
- 调度器在分发任务前计算配额,超出配额的用户透传 force_history=True

效果:新文章获得合理曝光,历史文章持续被互动,分布更自然
2026-04-08 11:47:36 +08:00
stefanfeng b43ee777fc feat: 互动记录/数据看板自动刷新 + 日志时间格式修复
1. Interactions.vue: 每30秒自动刷新互动记录列表
2. Dashboard.vue: 每30秒自动刷新数据看板
3. Logs.vue: 时间格式修复(T→空格,去掉时区标识)
4. logs.py: created_at 改用 strftime 输出 +08:00 格式(而非 isoformat 的 +00:00)

页面保持浏览时自动获取最新数据,离开页面时自动清除定时器
2026-04-07 14:51:59 +08:00
stefanfeng 7d9f5a358b fix: 使用正确的广场接口 /business/square/list
问题根因:
- 一直用 /business/member/square/list(成员广场接口)
- 该接口加了 isPlatformShow=true 和 isAdmin=false 过滤
- 导致大量文章(包括今日最新文章)被过滤掉,只剩124篇
- 正确接口 /business/square/list?type=1 有2504篇文章,今日文章可见

修改内容:
1. 接口地址改为 /business/square/list(4处)
2. 去掉 isPlatformShow 和 isAdmin 参数,只保留 type=1
3. _is_today 时间字段优先用 publishTime(新接口主要返回此字段)
4. 排序和热度权重计算也统一用 publishTime 优先
2026-04-07 14:32:35 +08:00
stefanfeng b6e094c9c0 fix: 过滤 recordId/title 为 null 的异常广场条目(type!=1) 2026-04-07 14:14:57 +08:00
stefanfeng d6ad1db535 fix: Phase2始终从第1页(最新)开始,按小时递进页码
问题:
- 原实现用 hour % total_pages 决定起始页
- 13点时 13%3=1,start_page=2,直接跳过第1页(最新文章)
- 导致虚拟用户永远不互动最新发布的文章

修复:
- 第1页(最新文章)始终在获取总页数时一并拉取,零额外开销
- hour_page = (hour % max_pages) + 1,每小时推进一页(1→2→3...循环)
- 0点=第1页最新,1点=第2页,依此类推,形成完整的新→旧覆盖
- 若当前时段页为空则顺序回退,最终兜底第1页

Phase 1(今日新文章)逻辑不变
2026-04-07 13:52:59 +08:00
stefanfeng 169e798718 fix: 历史文章改为从新到旧顺序翻页,而非随机
修复原因:
- 需求要求从最新文章开始往旧的方向依次互动
- 原实现随机选页,导致每次跳到不同页,无法体现「从新到旧」
- 列表API返回的文章已按 createTime 降序排列(第1页最新)

新Phase 2逻辑:
- 计算总页数(最多10页)
- 用当前小时 % 总页数 决定起始页(同小时内分散到不同页)
- 若起始页为空,顺序往后再折回,直到找到有文章的页
- 在同一页内随机打散,保证同一时段不同用户不总是抢相同文章
- validate_article 校验,不够则从剩余补充

Phase 1(今日新文章)逻辑不变:
- 仍从第1页抓取当天文章,按 createTime 降序(新→旧)排列后校验返回
2026-04-07 13:24:42 +08:00
stefanfeng 053d22965c feat: 调度优先今日新文章,无新文章时随机历史翻页
新调度规则:

Phase 1 — 今日新文章优先(从新到旧轮询):
- 从第1页开始拉取(接口返回最新优先)
- 只保留今日发布的文章,按 createTime 降序排列(新→旧)
- 最多扫描3页,发现非今日文章立即停止
- 对今日文章逐篇 validate_article 校验后返回

Phase 2 — 历史兜底(仅今日无新文章时触发):
- 随机翻 1~10 页历史
- 热度+新鲜度加权采样(commentNum×3 + praiseNum×2 + readNum)
- validate_article 校验后返回

两阶段均包含:
- 本人发布文章过滤
- 静态+运行时无效ID过滤
- 文章有效性校验(不可开/正文<100字自动加入缓存黑名单)
2026-04-03 11:33:31 +08:00
stefanfeng e18c241bf0 feat: 文章有效性校验,过滤不可开/字数<100的文章
新增 validate_article() 方法:
- 调用 GET /news/{id} 接口验证文章是否存在(code≠0 则无效)
- 去除 HTML 标签后统计正文字数,< 100 字则过滤
- 运行时缓存 _invalid_ids_cache:校验失败的 ID 进程内永久跳过,避免重复 API 调用

静态黑名单更新:
- 新增 1952296583257133058(测试发现的无效文章)
- 静态黑名单与运行时缓存合并使用

get_news_list 流程:
1. 静态黑名单过滤(无 API 开销)
2. 热度+新鲜度加权采样
3. validate_article 逐篇校验
4. 若候选不足,从剩余池补充直到达到 count
2026-04-03 11:18:22 +08:00
stefanfeng f52bc7d147 fix: 点赞/收藏/转发同一篇文章每日只触发一次
问题:同一用户对同一篇文章重复点赞、收藏、转发
原因:去重逻辑只针对评论,未覆盖其他互动类型

修复:
- 查询今日所有互动类型的已完成记录(不只是评论)
- 点赞/收藏/转发执行前检查 today_done,已做过则跳过
- 概率控制依然有效(未做过的文章才进入随机概率判断)
- 评论去重逻辑保持不变
2026-04-03 10:47:34 +08:00
stefanfeng 4ab8f94663 fix: 修复调度器只选前5个用户的bug,改为全量轮转
问题根因:
1. SQL查询加了 .limit(max_concurrent),导致只有数据库前5条用户参与互动
2. 额外的 random.random() < 0.6 过滤进一步减少了执行用户数

修复方案:
- 查询所有已登录用户(去掉 SQL LIMIT)
- 按最后互动时间升序排序,最久未互动的用户优先
- 前1/3名额给最久未互动用户(优先权),其余随机补充
- 每轮最多执行 max_concurrent 个用户,保证公平轮转
2026-04-03 10:28:49 +08:00
stefanfeng 7203f04be6 feat: 评论去重 + 热度/新鲜度加权选文
评论去重逻辑:
- 查询今日已评论的文章ID,选文时已评论的文章权重降为10%
- 若选中已评论文章:改为回复其他用户的评论(虚拟用户互动链)
- 若选中未评论文章:正常发新评论,评论成功后随机回复他人评论

热度+新鲜度加权选文规则:
- 热度分 = commentNum×3 + praiseNum×2 + readNum×1
- 新鲜度 = 72小时内的新文章获得最高3倍加成,随时间线性衰减
- 综合权重 = (热度分+1) × 新鲜度,确保真实用户互动多的新文章优先被虚拟用户关注
2026-04-02 17:33:07 +08:00
stefanfeng 958eaeda8a fix: 多项修复
- main.py: 加 _CNJSONResponse 修复 datetime 序列化时区(+00:00→+08:00)
- schemas/__init__.py: 加 _fmt_dt 函数和 sync_to_platform 字段
- ai_service.py: 评论 max_tokens 从 300 提升到 500 避免截断
- scheduler.py: datetime.utcnow() 全部改为 datetime.now()(北京时间)
- docker-compose.yml: MySQL 容器加 TZ=Asia/Shanghai
- Interactions.vue: 文章标题链接从系统配置读取域名,格式为 {域名}/huihui-h5/#/news/share?id={id}&login=no
2026-04-01 18:07:42 +08:00
stefanfeng fe9110ca3c fix: 修复文章链接域名读取(data是dict格式,用data['key'].value获取) 2026-04-01 18:03:31 +08:00
stefanfeng da4744f333 feat: 文章标题链接使用正确H5 URL格式,域名从系统配置动态读取 2026-04-01 16:58:51 +08:00
stefanfeng 7448fdcba1 feat: 互动记录文章标题支持点击跳转详情页(el-link) 2026-04-01 16:12:22 +08:00
stefanfeng b27bc216e3 fix: 修复时间显示+文章标题可点击
- axios 拦截器把 +00:00 转 +08:00
- MySQL/后端容器加 TZ=Asia/Shanghai
- scheduler.py 全部改用 datetime.now()
- main.py 加 _CNJSONResponse 修复序列化时区
- 互动记录文章标题支持点击跳转详情页
- 执行时间直接显示字符串无需二次转换
2026-04-01 15:12:12 +08:00
stefanfeng 9b87eeb84b fix: 修复前端时间显示错误(+00:00转+08:00) 2026-03-31 13:38:21 +08:00
stefanfeng 3a06224f5d feat: 添加一键登出全部按钮 2026-03-31 13:02:49 +08:00
stefanfeng cd07776914 feat: 多项功能更新
- 日志时间改为北京时间(TZ=Asia/Shanghai)
- 评论达上限后继续执行点赞/收藏/转发
- 用户信息同步改用 PATCH /v2/users/current
- 一键登出全部功能
- 一键登出全部前端按钮
- update.sh 一键更新脚本
2026-03-31 10:29:26 +08:00
stefanfeng 3fbccbc2b1 chore: 清理多余文件 2026-03-31 10:25:00 +08:00
stefanfeng 0cfc9bf9c8 feat: AI虚拟用户新闻互动系统 v1.3.0 初始提交
- 虚拟用户管理(昵称/头像/性别/简介/邮箱同步到目标平台)
- AI互动调度(点赞/收藏/评论/转发)
- 日志时间改为北京时间
- 评论达上限后继续执行点赞收藏转发
- 一键登出全部功能
- 浅色主题UI
2026-03-31 10:20:57 +08:00
204 changed files with 33942 additions and 6424 deletions
+18
View File
@@ -0,0 +1,18 @@
# Python
__pycache__/
*.pyc
*.pyo
*.pyd
.env
backend/logs/
# Node
frontend/node_modules/
frontend/dist/
# macOS
.DS_Store
# IDE
.idea/
.vscode/
-490
View File
@@ -1,490 +0,0 @@
# 会会虚拟用户 AI 互动系统 - 架构设计文档
## 一、系统架构
### 1.1 整体架构图
```
┌─────────────────────────────────────────────────────┐
│ 用户层 │
│ (浏览器 / 移动端) │
└──────────────────┬──────────────────────────────────┘
│
│ HTTP/HTTPS
│
┌──────────────────▼──────────────────────────────────┐
│ Nginx │
│ (反向代理 / 静态资源) │
└──────────────────┬──────────────────────────────────┘
│
┌─────────┴─────────┐
│ │
┌────────▼────────┐ ┌──────▼────────┐
│ Frontend │ │ Backend │
│ Vue3 + │ │ FastAPI │
│ Element Plus │ │ Python │
│ (端口:80) │ │ (端口:8000) │
└─────────────────┘ └──────┬────────┘
│
┌─────────────┼─────────────┐
│ │ │
┌──────▼──────┐ ┌───▼────┐ ┌─────▼─────┐
│ MySQL │ │ Redis │ │ 定时任务 │
│ 数据库 │ │ (可选) │ │ 调度器 │
└─────────────┘ └────────┘ └───────────┘
│
┌─────────────┼─────────────┐
│ │ │
┌──────▼──────┐ ┌───▼────┐ ┌─────▼─────┐
│ 会会 API │ │ AI 模型 │ │ 文件存储 │
│ 接口服务 │ │ 服务 │ │ ( uploads)│
└─────────────┘ └────────┘ └───────────┘
```
### 1.2 技术栈选型
| 层级 | 技术 | 说明 |
|------|------|------|
| **前端** | Vue 3 | 渐进式 JavaScript 框架 |
| | Element Plus | UI 组件库 |
| | Pinia | 状态管理 |
| | Vue Router | 路由管理 |
| | ECharts | 数据可视化 |
| **后端** | Python 3.11 | 编程语言 |
| | FastAPI | 高性能 Web 框架 |
| | SQLAlchemy | ORM 框架 |
| | Pydantic | 数据验证 |
| | APScheduler | 定时任务调度 |
| **数据库** | MySQL 8.0 | 关系型数据库 |
| **AI 对接** | OpenAI SDK | GPT 系列模型 |
| | Zhipu SDK | 智谱 AI 模型 |
| **部署** | Docker | 容器化 |
| | Docker Compose | 编排工具 |
| | Nginx | Web 服务器 |
## 二、模块设计
### 2.1 后端模块划分
```
backend/app/
├── main.py # 应用入口
├── core/ # 核心配置
│ ├── config.py # 系统配置
│ └── __init__.py
├── models/ # 数据库模型
│ ├── base.py # 数据库基础
│ ├── virtual_user.py # 虚拟用户模型
│ ├── interaction.py # 互动记录模型
│ ├── token_usage.py # Token 使用模型
│ ├── system_config.py # 系统配置模型
│ ├── ai_model.py # AI 模型配置模型
│ └── news_cache.py # 新闻缓存模型
├── schemas/ # Pydantic Schema
│ ├── virtual_user.py # 虚拟用户 Schema
│ ├── interaction.py # 互动 Schema
│ ├── dashboard.py # 控制台 Schema
│ ├── ai_model.py # AI 模型 Schema
│ └── system_config.py # 系统配置 Schema
├── services/ # 业务服务
│ ├── huihui_api.py # 会会接口服务
│ ├── ai_service.py # AI 服务
│ ├── virtual_user_service.py # 虚拟用户服务
│ ├── interaction_service.py # 互动服务
│ ├── token_service.py # Token 统计服务
│ └── scheduler_service.py # 定时任务服务
├── api/ # API 路由
│ ├── router.py # 路由注册
│ ├── virtual_user.py # 虚拟用户 API
│ ├── interaction.py # 互动 API
│ ├── ai_model.py # AI 模型 API
│ ├── dashboard.py # 控制台 API
│ └── system_config.py # 系统配置 API
└── utils/ # 工具函数
```
### 2.2 前端模块划分
```
frontend/src/
├── main.js # 应用入口
├── App.vue # 根组件
├── router/ # 路由配置
│ └── index.js
├── api/ # API 接口
│ ├── request.js # Axios 封装
│ └── index.js # API 定义
├── views/ # 页面组件
│ ├── Dashboard.vue # 控制台
│ ├── VirtualUsers.vue # 虚拟用户管理
│ ├── Interactions.vue # 互动记录
│ ├── AIModels.vue # AI 模型配置
│ └── Settings.vue # 系统设置
├── components/ # 公共组件
├── stores/ # Pinia 状态管理
└── utils/ # 工具函数
```
## 三、数据库设计
### 3.1 ER 图
```
┌─────────────────┐ ┌──────────────────┐
│ virtual_users │ │ ai_model_configs │
│─────────────────│ │──────────────────│
│ id (PK) │ │ id (PK) │
│ username │ │ model_name │
│ password │ │ provider │
│ nickname │ │ api_key │
│ avatar_url │ │ is_default │
│ writing_style │ │ is_active │
│ activity_level │ └──────────────────┘
│ status │
│ session_token │ ┌──────────────────┐
│ total_interactions│ │ system_configs │
│ today_comments │ │──────────────────│
│ today_replies │ │ id (PK) │
└────────┬────────┘ │ config_key │
│ │ config_value │
│ 1 │ config_type │
│ └──────────────────┘
│ N
┌────────▼────────┐ ┌──────────────────┐
│interaction_records│ │ token_usages │
│─────────────────│ │──────────────────│
│ id (PK) │ │ id (PK) │
│ virtual_user_id │(FK) │ virtual_user_id │
│ news_id │ │ interaction_id │
│ news_title │ │ tokens_used │
│ interaction_type│ │ ai_model │
│ content │ │ action_type │
│ status │ │ usage_date │
│ retry_count │ └──────────────────┘
│ tokens_used │
│ ai_model_used │ ┌──────────────────┐
└─────────────────┘ │ news_cache │
│──────────────────│
│ id (PK) │
│ news_id │
│ title │
│ content │
│ category │
│ cache_date │
└──────────────────┘
```
### 3.2 核心表结构
#### virtual_users (虚拟用户表)
| 字段 | 类型 | 说明 |
|------|------|------|
| id | INT | 主键 |
| username | VARCHAR(100) | 用户名(唯一) |
| password | VARCHAR(200) | 密码(加密) |
| nickname | VARCHAR(100) | 昵称 |
| avatar_url | VARCHAR(500) | 头像 URL |
| writing_style | VARCHAR(50) | 写作风格 |
| activity_level | ENUM | 活跃度(low/medium/high) |
| persona_description | TEXT | AI 生成的人格描述 |
| status | ENUM | 状态(active/disabled) |
| is_logged_in | BOOLEAN | 是否已登录 |
| session_token | VARCHAR(500) | 会话 Token |
| total_interactions | INT | 总互动次数 |
| today_comments | INT | 今日评论数 |
| today_replies | INT | 今日回复数 |
| created_at | DATETIME | 创建时间 |
| updated_at | DATETIME | 更新时间 |
#### interaction_records (互动记录表)
| 字段 | 类型 | 说明 |
|------|------|------|
| id | INT | 主键 |
| virtual_user_id | INT | 虚拟用户 ID(外键) |
| news_id | VARCHAR(100) | 新闻 ID |
| news_title | VARCHAR(500) | 新闻标题 |
| interaction_type | ENUM | 类型(comment/reply/like/favorite/share) |
| content | TEXT | 互动内容 |
| target_comment_id | VARCHAR(100) | 目标评论 ID(回复时) |
| status | ENUM | 状态(pending/success/failed) |
| retry_count | INT | 重试次数 |
| error_message | TEXT | 错误信息 |
| tokens_used | INT | 消耗 Token 数 |
| execution_time | DATETIME | 执行时间 |
## 四、核心流程设计
### 4.1 虚拟用户生成流程
```
开始
│
▼
接收生成请求
(count, styles, levels)
│
▼
循环 count 次
│
├─► 生成唯一用户名
├─► 随机生成昵称
├─► 生成随机密码
├─► 选择写作风格
├─► 选择活跃度级别
├─► 调用 AI 生成人格描述
├─► 生成头像 URL
├─► 保存到数据库
│
▼
返回生成的用户列表
│
▼
结束
```
### 4.2 自动互动执行流程
```
定时任务触发
│
▼
检查活动时间段
│
├─否─► 跳过本次执行
│
是
│
▼
获取活跃虚拟用户列表
│
▼
随机选择一个用户
│
▼
检查用户活跃度概率
│
├─不通过─► 跳过
│
通过
│
▼
检查今日限额
│
├─超出─► 跳过
│
未超出
│
▼
随机选择新闻
│
▼
确定互动类型
(点赞/收藏/转发/评论/回复)
│
▼
如果是评论/回复
│
├─► 调用 AI 生成内容
├─► 记录 Token 消耗
│
▼
调用会会接口执行互动
│
▼
记录互动结果
│
▼
更新用户统计
│
▼
结束
```
### 4.3 AI 内容生成流程
```
接收生成请求
(新闻内容,写作风格,人格)
│
▼
构建提示词 Prompt
│
▼
获取默认 AI 模型配置
│
▼
根据 provider 选择 SDK
│
├─► OpenAI: 调用 GPT API
├─► Zhipu: 调用 GLM API
├─► Baidu: 调用文心 API
└─► Aliyun: 调用通义 API
│
▼
接收 AI 响应
│
▼
解析内容和 Token 数
│
▼
记录 Token 使用
│
▼
返回生成结果
│
▼
结束
```
## 五、API 设计规范
### 5.1 RESTful API 规范
所有 API 遵循 RESTful 设计风格:
```
GET /api/v1/resource # 获取资源列表
GET /api/v1/resource/{id} # 获取单个资源
POST /api/v1/resource # 创建资源
PUT /api/v1/resource/{id} # 更新资源
DELETE /api/v1/resource/{id} # 删除资源
POST /api/v1/resource/action # 资源操作
```
### 5.2 统一响应格式
成功响应:
```json
{
"data": { ... },
"message": "success"
}
```
列表响应:
```json
{
"total": 100,
"items": [ ... ]
}
```
错误响应:
```json
{
"detail": "错误信息",
"status_code": 400
}
```
### 5.3 分页参数
```
GET /api/v1/virtual-users?page=1&page_size=20
```
响应:
```json
{
"total": 100,
"items": [...],
"page": 1,
"page_size": 20
}
```
## 六、安全设计
### 6.1 认证授权
- JWT Token 认证
- Token 有效期 7 天
- 支持刷新 Token
### 6.2 数据安全
- 密码加密存储(bcrypt)
- API Key 加密存储
- SQL 注入防护(ORM 参数化)
- XSS 防护(前端输入过滤)
### 6.3 限流策略
- API 请求限流
- Token 消耗限额
- 单用户互动频次限制
## 七、性能优化
### 7.1 数据库优化
- 索引优化(username, status, execution_time)
- 分页查询
- 批量操作
### 7.2 缓存策略
- 新闻数据缓存
- 系统配置缓存
- 字典数据缓存
### 7.3 并发控制
- 数据库连接池(20-60)
- 异步 IO(FastAPI + asyncio)
- 定时任务并发控制
## 八、扩展性设计
### 8.1 插件化 AI 模型
支持通过配置添加新的 AI 模型提供商,无需修改代码。
### 8.2 可配置的互动策略
所有互动参数可通过系统配置调整:
- 互动概率
- 活跃度分级
- 时间段控制
### 8.3 模块化设计
各功能模块独立,易于扩展新功能:
- 新增互动类型
- 新增数据源
- 新增统计维度
## 九、监控与日志
### 9.1 日志系统
- 应用日志(INFO/WARNING/ERROR)
- 访问日志
- 慢查询日志
- 日志轮转(保留 30 天)
### 9.2 健康检查
```bash
GET /health
```
响应:
```json
{
"status": "healthy",
"scheduler_running": true
}
```
### 9.3 指标监控
- API 响应时间
- 数据库查询性能
- AI 调用成功率
- Token 消耗速率
---
**文档版本**: v1.0
**更新日期**: 2026-03-23
-375
View File
@@ -1,375 +0,0 @@
# 会会虚拟用户 AI 互动系统 - 部署与使用指南
## 一、项目概述
本项目是一个基于 AI 大模型的虚拟用户自动化互动系统,主要功能包括:
### 核心功能
1. **虚拟用户管理** - 批量生成、Excel 导入、人格配置
2. **AI 内容生成** - 支持 OpenAI、智谱等主流大模型
3. **自动化互动** - 评论、回复、点赞、收藏、转发
4. **定时任务调度** - 随机时间、活跃度控制、限额管理
5. **数据可视化** - 控制台仪表盘、Token 消耗统计
## 二、快速部署
### 方法一:一键启动脚本(推荐)
```bash
# 进入项目目录
cd /Users/yqq/Works/Projects/会会广场机器人
# 执行启动脚本
./start.sh
```
脚本会自动完成:
- 检查 Docker 环境
- 创建必要目录
- 配置文件初始化
- 启动所有服务
### 方法二:手动 Docker Compose 部署
```bash
# 1. 配置环境变量
cd backend
cp .env.example .env
# 编辑 .env 文件,配置 AI 模型 API Key
# 2. 启动服务
cd ..
docker-compose up -d
# 3. 查看日志
docker-compose logs -f backend
# 4. 访问服务
# 前端:http://localhost
# 后端 API: http://localhost:8000
# API 文档:http://localhost:8000/docs
```
## 三、配置说明
### 1. 环境变量配置(backend/.env)
```bash
# 必选:AI 模型配置(至少配置一个)
OPENAI_API_KEY=sk-your-api-key-here
# 或
ZHIPU_API_KEY=your-zhipu-api-key-here
# 可选:调整系统限制
MAX_TOKENS_PER_DAY=10000 # 每日 Token 上限
MAX_COMMENTS_PER_USER_PER_DAY=20 # 单用户日评论上限
TASK_START_HOUR=9 # 活动开始时间
TASK_END_HOUR=22 # 活动结束时间
```
### 2. 数据库配置
默认配置即可,Docker Compose 会自动创建 MySQL 容器:
- 主机:mysql
- 端口:3306
- 用户:huihui_user
- 密码:huihui_password
- 数据库:huihui_ai_bot
### 3. 会会接口配置
默认已配置为:`http://192.168.1.200:63120`
如需修改,在 `.env` 文件中调整:
```bash
HUIHUI_API_BASE=http://your-huihui-api-address
```
## 四、使用流程
### 第一步:AI 模型配置
1. 访问前端界面:http://localhost
2. 进入"AI 模型配置"页面
3. 点击"添加模型"
4. 选择提供商(OpenAI / 智谱等)
5. 填写 API Key 和其他参数
6. 点击"测试"验证配置
7. 设为默认模型
**支持的 AI 模型:**
- OpenAI: gpt-3.5-turbo, gpt-4
- 智谱 AI: glm-4
- 百度文心:ERNIE-Bot(待实现)
- 阿里通义:Qwen(待实现)
### 第二步:虚拟用户管理
#### 方式一:批量生成
1. 进入"虚拟用户管理"页面
2. 点击"批量生成"
3. 设置生成数量(1-100)
4. 选择写作风格(可不选,随机分配)
5. 选择活跃度级别
6. 勾选"AI 人格描述"(推荐)
7. 点击"生成"
#### 方式二:Excel 导入
准备 Excel 文件,包含以下列:
- `username`(必填)- 用户名/账号
- `password`(必填)- 密码
- `nickname`(可选)- 昵称
- `writing_style`(可选)- 写作风格
- `activity_level`(可选)- 活跃度(low/medium/high)
导入步骤:
1. 点击"Excel 导入"
2. 选择 Excel 文件
3. 勾选"生成 AI 人格描述"
4. 确认导入
### 第三步:系统设置
1. 进入"系统设置"页面
2. 配置活动时间段(如 9:00-22:00)
3. 配置互动频率(如 10-30 分钟)
4. 配置限额(Token 上限、评论上限等)
5. 保存配置
### 第四步:启动定时任务
1. 在"系统设置"页面
2. 点击"启动任务"按钮
3. 系统开始自动执行互动任务
### 第五步:监控与调整
1. **控制台** - 查看实时统计数据
- 虚拟用户总数
- 今日互动次数
- Token 消耗情况
- 最近互动记录
2. **互动记录** - 查看详细执行日志
- 成功/失败状态
- Token 消耗
- 错误信息
3. **调整策略**
- 根据 Token 消耗调整互动频率
- 根据效果调整写作风格
- 根据需求调整限额配置
## 五、高级配置
### 1. 自定义写作风格
在生成虚拟用户时,可以指定写作风格:
- 幽默风趣
- 严肃理性
- 文艺清新
- 吐槽犀利
- 感性温暖
- 客观中立
- 激情澎湃
- 冷静分析
- 活泼可爱
- 深沉内敛
也可在 `.env` 中添加自定义风格。
### 2. 活跃度分级
- **高活跃度**:每日互动 5-10 次(80% 概率执行)
- **中活跃度**:每日互动 2-5 次(50% 概率执行)
- **低活跃度**:每日互动 1-2 次(30% 概率执行)
### 3. 互动概率配置
默认概率:
- 点赞:80%
- 收藏:50%
- 转发:30%
- 评论:剩余概率
可在"系统设置"页面调整。
### 4. Token 限额策略
建议配置:
- 开发测试:1000-5000 tokens/天
- 小规模使用:5000-20000 tokens/天
- 大规模使用:20000+ tokens/天
根据实际 AI 模型价格计算成本。
## 六、运维管理
### 查看日志
```bash
# 查看所有服务日志
docker-compose logs -f
# 查看后端日志
docker-compose logs -f backend
# 查看数据库日志
docker-compose logs -f mysql
# 查看最近 100 行
docker-compose logs --tail=100 backend
```
### 备份数据
```bash
# 备份 MySQL 数据库
docker exec huihui_mysql mysqldump -uhuihui_user -phuihui_password huihui_ai_bot > backup_$(date +%Y%m%d).sql
# 恢复数据库
docker exec -i huihui_mysql mysql -uhuihui_user -phuihui_password huihui_ai_bot < backup_20260323.sql
```
### 服务重启
```bash
# 重启所有服务
docker-compose restart
# 重启单个服务
docker-compose restart backend
# 停止所有服务
docker-compose down
# 停止并删除容器(保留数据)
docker-compose down
# 完全清理(删除数据和容器)
docker-compose down -v
```
### 更新升级
```bash
# 拉取最新代码
git pull
# 重新构建并启动
docker-compose build
docker-compose up -d
```
## 七、故障排查
### 1. 后端服务无法启动
```bash
# 查看日志
docker-compose logs backend
# 常见问题:
# - 数据库连接失败:检查 MySQL 是否启动
# - 端口冲突:修改 docker-compose.yml 中的端口映射
# - 环境变量错误:检查 .env 文件格式
```
### 2. 数据库连接失败
```bash
# 检查 MySQL 容器状态
docker-compose ps mysql
# 查看 MySQL 日志
docker-compose logs mysql
# 手动连接测试
docker exec -it huihui_mysql mysql -uroot -proot123456
```
### 3. AI 模型调用失败
- 检查 API Key 是否正确
- 检查网络连接
- 查看 Token 余额
- 检查 API 地址是否正确
- 查看后端日志中的详细错误信息
### 4. 定时任务未执行
- 检查调度器状态(系统设置页面)
- 查看后端日志中的调度信息
- 确认活动时间段配置
- 检查虚拟用户状态(需为启用且已登录)
## 八、性能优化
### 1. 数据库优化
```sql
-- 添加索引(已自动创建)
CREATE INDEX idx_user_status ON virtual_users(status);
CREATE INDEX idx_interaction_time ON interaction_records(execution_time);
```
### 2. 并发控制
修改 `.env` 配置:
```bash
# 增加数据库连接池大小
DATABASE_POOL_SIZE=30
DATABASE_MAX_OVERFLOW=60
```
### 3. 缓存策略
- 新闻数据缓存到数据库
- 避免重复调用会会接口
- 定期清理过期缓存
## 九、安全建议
1. **修改默认密码**
- 数据库 root 密码
- 应用 JWT_SECRET_KEY
2. **API Key 保护**
- 不要将 .env 文件提交到 Git
- 使用环境变量或密钥管理服务
3. **网络隔离**
- 生产环境配置防火墙规则
- 限制数据库访问
4. **定期备份**
- 每日备份数据库
- 备份重要配置文件
## 十、技术支持
### 常见问题
**Q: 如何重置系统配置?**
A: 删除 system_configs 表数据,重启服务后会自动创建默认配置。
**Q: 如何批量删除虚拟用户?**
A: 暂时需要通过 API 或数据库操作,后续版本会添加批量删除功能。
**Q: Token 消耗过快怎么办?**
A: 降低 MAX_TOKENS_PER_DAY,减少虚拟用户数量,或调低互动频率。
**Q: 如何查看会会接口的详细文档?**
A: 访问 http://192.168.1.200:63120/doc.html 查看完整接口文档。
### 版本更新
关注项目更新日志,及时升级到最新版本获取新功能和安全修复。
---
**祝您使用愉快!**
如有问题,请联系开发团队或提交 Issue。
-310
View File
@@ -1,310 +0,0 @@
# 会会虚拟用户 AI 互动系统 - 项目交付总结
## 项目完成情况
### ✅ 已完成的功能模块
#### 1. 后端服务(FastAPI + Python)
- ✅ **数据库模型设计** - 6 个核心数据表
- virtual_users(虚拟用户)
- interaction_records(互动记录)
- token_usages(Token 使用)
- system_configs(系统配置)
- ai_model_configs(AI 模型配置)
- news_cache(新闻缓存)
- ✅ **核心服务层** - 6 个业务服务
- HuihuiAPIService - 会会接口对接
- AIService - AI 大模型对接(支持 OpenAI/智谱等)
- VirtualUserService - 虚拟用户管理
- InteractionService - 互动执行引擎
- TokenService - Token 统计服务
- SchedulerService - 定时任务调度器
- ✅ **API 接口** - 5 个路由模块
- /api/v1/virtual-users - 虚拟用户管理 API
- /api/v1/interactions - 互动管理 API
- /api/v1/ai-models - AI 模型配置 API
- /api/v1/dashboard - 控制台 API
- /api/v1/system - 系统设置 API
#### 2. 前端界面(Vue3 + Element Plus)
- ✅ **控制台页面** - Dashboard.vue
- 核心指标卡片(用户数、互动数、Token 消耗)
- 每日 Token 消耗折线图
- 每月 Token 消耗柱状图
- 最近互动记录表格
- ✅ **虚拟用户管理** - VirtualUsers.vue
- 用户列表展示(分页、搜索、筛选)
- 批量生成(1-100 个)
- Excel 导入功能
- 编辑/删除操作
- AI 人格描述生成
- ✅ **互动记录** - Interactions.vue
- 互动记录列表
- 类型标签(评论/回复/点赞/收藏/转发)
- 状态显示(成功/失败/待执行)
- Token 消耗展示
- ✅ **AI 模型配置** - AIModels.vue
- 模型列表管理
- 添加/编辑模型配置
- 模型测试功能
- 默认模型设置
- ✅ **系统设置** - Settings.vue
- 活动调度配置(时间段、频率)
- 限额配置(Token、评论、回复上限)
- 概率配置(点赞/收藏/转发概率)
- 定时任务启停控制
#### 3. Docker 部署配置
- ✅ **Dockerfile** - 后端服务容器化
- ✅ **docker-compose.yml** - 多容器编排
- MySQL 8.0 容器
- 后端应用容器
- Nginx 前端容器
- ✅ **初始化脚本** - init.sql
- ✅ **Nginx 配置** - nginx.conf
- ✅ **快速启动脚本** - start.sh
#### 4. 项目文档
- ✅ README.md - 项目说明文档
- ✅ DEPLOYMENT.md - 部署与使用指南
- ✅ ARCHITECTURE.md - 架构设计文档
- ✅ .env.example - 环境变量示例
## 核心技术亮点
### 1. AI 大模型集成
- 支持多个主流 AI 模型提供商
- OpenAI (GPT-3.5/GPT-4)
- 智谱 AI (GLM-4)
- 百度文心(框架已预留)
- 阿里通义(框架已预留)
- 统一的 AI 服务接口
- 支持动态切换模型
- Token 消耗精确统计
### 2. 智能互动引擎
- **随机策略**
- 活跃度分级(高/中/低)
- 互动时间随机(10-30 分钟间隔)
- 互动类型随机(基于概率)
- **限额控制**
- 每日 Token 上限
- 单用户评论上限
- 单用户回复上限
- **容错机制**
- 失败重试(3 次)
- 错误日志记录
- 独立任务隔离
### 3. 人格化虚拟用户
- AI 生成独特人格描述
- 写作风格配置(10 种预设)
- 活跃度分级
- 批量生成(支持 100 个)
- Excel 导入支持
### 4. 数据可视化
- ECharts 图表展示
- 实时数据统计
- 多维度分析(日/月)
- 交互式仪表盘
## 项目文件清单
```
会会广场机器人/
├── backend/ # 后端服务
│ ├── app/
│ │ ├── api/ # API 路由(5 个模块)
│ │ ├── core/ # 核心配置
│ │ ├── models/ # 数据库模型(6 个表)
│ │ ├── schemas/ # Pydantic Schema
│ │ ├── services/ # 业务服务(6 个)
│ │ └── main.py # 应用入口
│ ├── requirements.txt # Python 依赖
│ ├── Dockerfile # Docker 镜像
│ └── .env.example # 环境变量示例
│
├── frontend/ # 前端服务
│ ├── src/
│ │ ├── api/ # API 接口封装
│ │ ├── router/ # 路由配置
│ │ ├── views/ # 页面组件(5 个)
│ │ ├── App.vue # 根组件
│ │ └── main.js # 入口文件
│ ├── package.json # 依赖配置
│ └── vite.config.js # Vite 配置
│
├── docker/ # Docker 配置
│ ├── mysql/
│ │ └── init.sql # 数据库初始化
│ └── nginx/
│ └── nginx.conf # Nginx 配置
│
├── data/ # 数据持久化目录
│ ├── mysql/
│ └── logs/
│
├── docker-compose.yml # Docker Compose 配置
├── start.sh # 快速启动脚本
├── README.md # 项目说明
├── DEPLOYMENT.md # 部署指南
└── ARCHITECTURE.md # 架构文档
```
## 使用说明
### 快速开始(3 步部署)
```bash
# 1. 进入项目目录
cd /Users/yqq/Works/Projects/会会广场机器人
# 2. 配置环境变量
cd backend
cp .env.example .env
# 编辑 .env,填入 AI 模型 API Key
# 3. 一键启动
cd ..
./start.sh
```
### 访问地址
- 前端界面:http://localhost
- 后端 API: http://localhost:8000
- API 文档:http://localhost:8000/docs
### 首次使用流程
1. **配置 AI 模型** - 在"AI 模型配置"页面添加 API Key
2. **生成虚拟用户** - 批量生成或 Excel 导入
3. **设置系统参数** - 配置活动时间、限额等
4. **启动定时任务** - 点击"启动任务"开始自动互动
5. **监控运行状态** - 在控制台查看统计数据
## 技术栈总结
| 类别 | 技术 | 版本 |
|------|------|------|
| **后端框架** | FastAPI | 0.109.0 |
| **ORM** | SQLAlchemy | 2.0.25 |
| **数据验证** | Pydantic | 2.5.3 |
| **定时任务** | APScheduler | 3.10.4 |
| **数据库** | MySQL | 8.0 |
| **前端框架** | Vue | 3.4.0 |
| **UI 组件库** | Element Plus | 2.5.0 |
| **状态管理** | Pinia | 2.1.7 |
| **图表库** | ECharts | 5.4.3 |
| **构建工具** | Vite | 5.0.10 |
| **容器化** | Docker | latest |
| **编排** | Docker Compose | 2.0+ |
## 性能指标
### 设计容量
- 虚拟用户数:1000+
- 日互动量:10000+
- Token 处理能力:100,000+/天
- API 响应时间:< 200ms
- 数据库查询:< 50ms
### 资源需求
- CPU: 2 核(推荐 4 核)
- 内存:2GB(推荐 4GB)
- 磁盘:10GB(根据日志量调整)
- 网络:1Mbps 带宽
## 安全特性
1. **数据安全**
- 密码加密存储(bcrypt)
- API Key 加密
- SQL 注入防护
2. **访问控制**
- JWT Token 认证
- 权限分级
3. **限额保护**
- Token 日限额
- 互动频率限制
- 并发控制
## 后续优化建议
### 短期优化(v1.1)
- [ ] 完善百度文心、阿里通义对接
- [ ] 添加批量删除虚拟用户功能
- [ ] 增加更多图表维度
- [ ] 优化移动端适配
### 中期规划(v2.0)
- [ ] 用户登录认证系统
- [ ] RBAC 权限管理
- [ ] WebSocket 实时通知
- [ ] 导出报表功能
- [ ] 更多 AI 模型支持
### 长期规划(v3.0)
- [ ] 分布式部署支持
- [ ] Redis 缓存层
- [ ] 消息队列(异步任务)
- [ ] 机器学习优化互动策略
- [ ] A/B 测试框架
## 已知限制
1. **会会接口依赖**
- 需要确保会会接口服务可用
- 接口变更需同步更新代码
2. **AI 模型成本**
- Token 消耗会产生费用
- 建议合理设置限额
3. **浏览器兼容**
- 仅测试 Chrome/Edge/Safari
- IE 不支持
## 联系与支持
### 技术文档
- README.md - 快速开始
- DEPLOYMENT.md - 详细部署指南
- ARCHITECTURE.md - 架构设计参考
### 常见问题
详见 DEPLOYMENT.md 第七章节
### 版本信息
- 当前版本:v1.0.0
- 发布日期:2026-03-23
- Python 版本:3.11+
- Node.js 版本:18+
---
## 项目交付清单 ✅
- [x] 完整的后端服务代码
- [x] 完整的前端界面代码
- [x] Docker 部署配置
- [x] 数据库设计文档
- [x] 架构设计文档
- [x] 部署使用指南
- [x] 快速启动脚本
- [x] API 接口文档(Swagger)
- [x] 环境变量配置示例
- [x] 项目总结文档
**项目交付完成!🎉**
所有功能已按需求实现,可直接部署使用。
+258 -157
View File
@@ -1,193 +1,294 @@
# 会会虚拟用户 AI 互动系统
# AI虚拟用户新闻互动系统
基于 AI 大模型的虚拟用户自动化互动系统,支持对接会会平台接口,实现虚拟用户的自动登录、评论、回复、点赞、收藏、转发等功能。
> 基于AI驱动的虚拟用户新闻互动自动化平台,支持批量虚拟用户管理、AI人格生成、真实登录新闻平台、自动随机互动。
## 功能特性
---
### 核心功能
- ✅ **虚拟用户管理**:批量生成、Excel 导入、人格配置
- ✅ **AI 内容生成**:支持 OpenAI、智谱、百度文心、阿里通义等大模型
- ✅ **自动化互动**:定时任务、随机策略、限额控制
- ✅ **数据可视化**:控制台仪表盘、Token 消耗统计
- ✅ **Docker 部署**:支持 1Panel 一键部署
## 📁 项目结构
### 互动类型
- 评论(AI 生成)
- 回复(AI 生成)
- 点赞
- 收藏
- 转发
### 技术栈
- **后端**:Python 3.11 + FastAPI
- **前端**:Vue 3 + Element Plus
- **数据库**:MySQL 8.0
- **AI 对接**:OpenAI / 智谱 / 百度 / 阿里
- **部署**:Docker + Docker Compose
## 快速开始
### 1. 环境要求
- Docker 20.10+
- Docker Compose 2.0+
- 1Panel 面板(可选)
### 2. 配置环境变量
```bash
cd backend
cp .env.example .env
```
ai-virtual-news/
├── docker-compose.yml # Docker编排文件
├── docker/
│ └── mysql/
│ └── init.sql # 数据库初始化脚本
├── backend/ # Python FastAPI 后端
│ ├── Dockerfile
│ ├── requirements.txt
│ └── app/
│ ├── main.py # 应用入口
│ ├── api/ # API路由层
│ ├── core/ # 核心配置(DB/Redis/日志)
│ ├── models/ # SQLAlchemy ORM模型
│ ├── schemas/ # Pydantic数据模型
│ ├── services/ # 业务服务层
│ └── utils/ # 工具类(AES加密等)
└── frontend/ # Vue3 前端
├── Dockerfile
├── nginx.conf
├── src/
│ ├── views/ # 页面组件
│ ├── api/ # Axios API封装
│ ├── router/ # Vue Router
│ ├── layouts/ # 布局组件
│ └── styles/ # 全局样式
└── package.json
```
编辑 `.env` 文件,配置必要参数:
- 数据库配置(默认即可)
- AI 模型 API Key(至少配置一个)
- 会会接口地址(默认:http://192.168.1.200:63120)
---
### 3. 启动服务
## 🚀 快速部署(1Panel Docker)
### 前置要求
- 已安装 1Panel 面板
- 已安装 Docker 及 Docker Compose
- 服务器内网可访问新闻平台接口(192.168.1.200:63120)
### 第一步:修改环境配置
编辑 `docker-compose.yml`,修改以下**必须**更改的安全参数:
```yaml
environment:
- SECRET_KEY=your-secret-key-change-in-production # ⚠️ 必须修改
- AES_KEY=your-aes-key-32-chars-change-now! # ⚠️ 必须修改(必须32字符)
- DB_PASSWORD=AiVirtual@2024 # ⚠️ 建议修改
```
同时修改 MySQL 的 `MYSQL_PASSWORD` 与 `DB_PASSWORD` 保持一致。
### 第二步:通过 1Panel 部署
**方式A:1Panel 应用商店(推荐)**
1. 登录 1Panel → 应用商店 → 搜索 "Docker Compose"
2. 上传本项目目录
3. 点击部署
**方式B:SSH 命令行**
```bash
# 启动所有服务
# 1. 上传项目到服务器
scp -r ai-virtual-news/ root@your-server:/opt/
# 2. 进入项目目录
cd /opt/ai-virtual-news
# 3. 启动所有服务
docker-compose up -d
# 查看日志
docker-compose logs -f backend
# 停止服务
docker-compose down
# 4. 查看启动日志
docker-compose logs -f
```
### 4. 访问服务
- 前端界面:http://localhost
- 后端 API:http://localhost:8000
- API 文档:http://localhost:8000/docs
### 第三步:访问系统
## 项目结构
| 服务 | 地址 |
|------|------|
| 前端控制台 | http://服务器IP:9000 |
| 后端API文档 | http://服务器IP:8000/api/docs |
| MySQL | 服务器IP:3306 |
| Redis | 服务器IP:6379 |
```
会会广场机器人/
├── backend/ # 后端服务
│ ├── app/
│ │ ├── api/ # API 路由
│ │ ├── core/ # 核心配置
│ │ ├── models/ # 数据库模型
│ │ ├── schemas/ # Pydantic Schema
│ │ ├── services/ # 业务服务
│ │ └── main.py # 应用入口
│ ├── requirements.txt # Python 依赖
│ └── Dockerfile
├── frontend/ # 前端服务
│ └── src/
├── docker/ # Docker 配置
│ ├── mysql/
│ └── nginx/
├── data/ # 数据持久化
│ ├── mysql/
│ └── logs/
├── docker-compose.yml
└── README.md
---
## ⚙️ 初始配置
### 1. 配置AI模型
访问控制台 → **AI模型配置** → 添加模型:
| 字段 | 说明 | 示例 |
|------|------|------|
| 模型名称 | 自定义名称 | GPT-4生产 |
| 提供商 | 选择对应供应商 | OpenAI |
| API地址 | 留空用默认 | https://api.openai.com/v1 |
| API Key | 对应平台的Key | sk-... |
| 模型版本 | 具体模型名 | gpt-4-turbo |
> 配置完成后点击「设为默认」,系统将使用此模型进行所有AI操作。
> 点击「测试」验证模型可用性。
**支持的国产模型配置:**
| 提供商 | API地址 | 模型版本示例 |
|--------|---------|-------------|
| 智谱GLM | https://open.bigmodel.cn/api/paas/v4 | glm-4 |
| 文心一言 | https://aip.baidubce.com/rpc/2.0/ai_custom/v1/wenxinworkshop/chat | ERNIE-Bot-4 |
| 通义千问 | https://dashscope.aliyuncs.com/compatible-mode/v1 | qwen-turbo |
### 2. 配置新闻平台地址
访问控制台 → **调度设置** → 修改「平台接口地址」为实际地址。
### 3. 创建虚拟用户
**方式A:单个创建**
控制台 → 虚拟用户 → 新增用户 → 填写账号密码 → 系统自动生成AI人格
**方式B:Excel批量导入**
1. 下载导入模板
2. 填写账号/密码/昵称等信息
3. 上传Excel,系统自动校验并为每个用户生成AI人格
### 4. 启动自动互动
1. 确认用户已登录(状态为「已登录」)
2. 调度设置 → 确认互动时间段和概率配置
3. 调度器默认启动,系统将在设定时间段自动执行互动
---
## 🔧 运维管理
### Docker 常用命令
```bash
# 查看所有容器状态
docker compose ps
# 重启后端服务(后端代码更新后执行)
docker compose restart ai-virtual-backend
# 查看后端实时日志
docker compose logs -f ai-virtual-backend
# 停止所有服务
docker compose down
# 启动所有服务
docker compose up -d
# 进入后端容器
docker exec -it ai-virtual-backend bash
# 进入MySQL
docker exec -it ai-virtual-mysql mysql -u aivirtual -p ai_virtual_news
```
## API 接口
### ⚠️ 前端更新(重要:必须用此方式)
> `docker compose build` 存在缓存问题,前端代码修改后**必须**用以下方式重新 build,否则修改不会生效。
```bash
cd /opt/1panel/docker/compose/ai-virtual-news/frontend
# 第一步:清缓存并 build(使用 node 镜像直接 build 宿主机目录)
rm -rf dist node_modules/.vite
docker run --rm -v $(pwd):/app -w /app node:18-alpine sh -c "npm run build"
# 第二步:把 dist 复制到运行中的容器
docker cp dist/. ai-virtual-frontend:/usr/share/nginx/html/
# 第三步:重载 nginx(无需重启容器,立即生效)
docker exec ai-virtual-frontend nginx -s reload
```
### 数据备份(1Panel)
1. 1Panel → 数据库 → MySQL → 定时备份
2. 建议每天凌晨 3 点备份,保留 30 天
3. 或手动备份:
```bash
docker exec ai-virtual-mysql mysqldump -u aivirtual -pAiVirtual@2024 ai_virtual_news > backup_$(date +%Y%m%d).sql
```
### 日志位置
| 日志类型 | 容器路径 | 宿主机路径 |
|----------|----------|------------|
| 应用日志 | /app/logs/app_*.log | ./backend/logs/ |
| 错误日志 | /app/logs/error_*.log | ./backend/logs/ |
| AI调用日志 | /app/logs/ai_*.log | ./backend/logs/ |
---
## 🔒 安全注意事项
1. **AES密钥**:`AES_KEY` 必须修改为32字符随机字符串,用于加密存储账号密码
2. **数据库密码**:生产环境务必修改默认密码
3. **端口暴露**:建议通过 Nginx 反向代理访问,不要直接暴露 8000 端口
4. **防火墙**:MySQL(3306)、Redis(6379) 端口不应对外暴露
5. **互动频率**:合理设置互动间隔,避免触发新闻平台风控
---
## 📊 功能模块说明
### 数据看板
- 实时展示用户总数、在线数、今日互动量
- Token消耗折线图(近30天/7天)
- 近12个月月度消耗柱状图
- 系统运行状态监控
### 虚拟用户管理
- `GET /api/v1/virtual-users` - 获取用户列表
- `POST /api/v1/virtual-users` - 创建用户
- `POST /api/v1/virtual-users/generate` - 批量生成用户
- `POST /api/v1/virtual-users/import` - Excel 导入用户
- `PUT /api/v1/virtual-users/{id}` - 更新用户
- `DELETE /api/v1/virtual-users/{id}` - 删除用户
- 新增/编辑/删除用户,账号密码AES加密存储
- Excel批量导入(含格式校验、去重、错误详情)
- Excel批量导出(不含密码密文)
- AI人格生成:性格/语言风格/兴趣/互动倾向/字数偏好
- 编辑用户资料(昵称/真实姓名/性别/头像/简介/邮箱),支持同步到目标平台
- 头像上传:上传图片到平台 filecenter,自动更新用户头像
- 单个/批量启用、禁用、登出操作
- 手动触发登录/登出
### 互动管理
- `POST /api/v1/interactions/execute` - 执行互动
- `GET /api/v1/interactions` - 获取互动记录
### AI互动模块
- 真实调用新闻平台登录接口获取会话Token
- 会话自动校验(10分钟/次),失效自动重登
- 随机翻页获取文章,按用户兴趣偏好筛选,自动过滤无效新闻
- AI生成贴合人格的评论/回复内容,内容完整不截断,自动过滤敏感词
- 按概率随机触发:评论/回复/点赞/收藏/转发
- 每日互动次数限额控制
- 互动记录支持手动重试、取消
### AI 模型配置
- `GET /api/v1/ai-models` - 获取模型列表
- `POST /api/v1/ai-models` - 创建模型配置
- `POST /api/v1/ai-models/test` - 测试模型
### AI模型配置
- 支持 OpenAI / 智谱GLM / 文心一言 / 通义千问 / 本地模型
- API Key AES加密存储
- 模型测试功能(验证可用性 + Token消耗预览)
- 多模型管理,设置默认模型
### 控制台
- `GET /api/v1/dashboard` - 获取统计数据
- `GET /api/v1/dashboard/token/stats` - Token 统计
- `GET /api/v1/dashboard/token/daily` - 每日 Token 使用
### 调度设置
- 互动时间段配置(北京时间)
- 最小互动间隔控制(秒),防止同一用户频繁互动
- 各互动类型概率独立配置
- 并发用户数上限(0=不限)
- 每日Token配额管控
- 一键暂停/启动调度器
- 立即触发互动(测试用)
## 系统配置
### 日志管理
- 登录日志:登录/登出/失败记录
- 日志文件:应用日志/错误日志实时查看
- 日志下载
### 默认限额
- 每日 Token 上限:10,000
- 单用户日评论上限:20
- 单用户日回复上限:10
---
### 活动时间段
- 开始时间:09:00
- 结束时间:22:00
- 互动间隔:10-30 分钟(随机)
## 🐛 常见问题
### 互动概率
- 点赞:80%
- 收藏:50%
- 转发:30%
**Q: 容器启动失败,提示数据库连接失败?**
A: MySQL 启动需要时间,后端依赖 healthcheck。等待 30-60 秒后重试:`docker compose restart ai-virtual-backend`
## 开发指南
**Q: 用户登录始终失败?**
A: 1) 检查新闻平台接口地址是否正确;2) 检查账号密码是否正确;3) 查看后端日志定位具体错误
### 添加新的 AI 模型
1. 在 `app/services/ai_service.py` 添加对应的调用方法
2. 在 `app/models/ai_model.py` 添加提供商枚举
3. 通过 API 配置新模型
**Q: AI人格生成失败?**
A: 未配置AI模型时系统会随机生成人格作为兜底,这是正常行为。配置有效的AI模型后可重新生成。
### 自定义互动策略
修改 `app/services/interaction_service.py` 中的互动逻辑
**Q: 调度器不执行互动?**
A: 检查:1) 调度器是否启用;2) 是否在设定的互动时间段内(北京时间);3) 是否有已登录状态的用户;4) Token是否已达每日上限;5) 用户最近互动时间是否超过最小间隔
### 调整定时任务
修改 `app/services/scheduler_service.py` 中的调度配置
**Q: 前端修改后没有生效?**
A: 不能用 `docker compose build`,必须用上方「前端更新」中的 node 镜像 build 方式。
## 常见问题
**Q: 互动报"服务器繁忙"?**
A: 通常是 orgId 为空导致。系统已自动从广场文章数据获取 orgId,如仍报错请检查文章是否有效。
### 1. 数据库连接失败
检查 MySQL 服务是否启动:
```bash
docker-compose ps mysql
```
**Q: 评论报敏感词?**
A: AI 提示词已包含安全规则,偶发属正常,系统不重试敏感词失败。
### 2. AI 模型调用失败
- 检查 API Key 是否正确
- 检查网络连接
- 查看后端日志:`docker-compose logs backend`
**Q: 后端 502 错误?**
A: 查看日志定位原因:`docker compose logs --tail=20 ai-virtual-backend | grep -E "Error|Exception"`
### 3. Token 消耗过快
- 调整 `MAX_TOKENS_PER_DAY` 配置
- 降低互动频率
- 减少虚拟用户数量
---
## 1Panel 部署
## 📞 技术支持
### 通过 1Panel 部署 Docker Compose
1. 登录 1Panel 面板
2. 进入"容器管理" -> "Compose"
3. 点击"创建",上传 `docker-compose.yml`
4. 配置环境变量
5. 点击"创建"启动服务
### 开放端口
在 1Panel 防火墙中开放:
- 80 (前端)
- 8000 (后端 API)
- 3306 (数据库,可选)
## 更新日志
### v1.0.0 (2026-03-23)
- 初始版本发布
- 支持虚拟用户生成和管理
- 支持 AI 自动生成评论和回复
- 支持多种互动类型
- 完整的 Docker 部署方案
## License
MIT License
## 联系方式
如有问题,请提交 Issue 或联系开发团队。
- 后端API文档:`http://服务器IP:8000/api/docs`
- 接口健康检查:`http://服务器IP:8000/health`
-55
View File
@@ -1,55 +0,0 @@
# 应用基础配置
APP_NAME=会会虚拟用户 AI 互动系统
APP_VERSION=1.0.0
DEBUG=False
API_PREFIX=/api/v1
# 数据库配置
DATABASE_HOST=mysql
DATABASE_PORT=3306
DATABASE_USER=root
DATABASE_PASSWORD=root123456
DATABASE_NAME=huihui_ai_bot
# 或者使用完整的 DATABASE_URL
# DATABASE_URL=mysql+pymysql://root:root123456@mysql:3306/huihui_ai_bot?charset=utf8mb4
# JWT 配置(生产环境请修改)
JWT_SECRET_KEY=your-secret-key-change-in-production-abc123xyz
JWT_ALGORITHM=HS256
JWT_EXPIRE_MINUTES=10080
# 会会接口配置
HUIHUI_API_BASE=http://192.168.1.200:63120
HUIHUI_DOC_URL=http://192.168.1.200:63120/doc.html
# AI 模型配置
DEFAULT_AI_MODEL=openai
# OpenAI 配置
OPENAI_API_KEY=sk-your-openai-api-key
OPENAI_BASE_URL=https://api.openai.com/v1
OPENAI_MODEL=gpt-3.5-turbo
# 智谱 AI 配置
ZHIPU_API_KEY=your-zhipu-api-key
ZHIPU_MODEL=glm-4
# 系统限制配置
MAX_TOKENS_PER_DAY=10000
MAX_COMMENTS_PER_USER_PER_DAY=20
MAX_REPLIES_PER_USER_PER_DAY=10
# 定时任务配置
TASK_START_HOUR=9
TASK_END_HOUR=22
TASK_INTERVAL_MIN=10
TASK_INTERVAL_MAX=30
# 互动概率配置
LIKE_PROBABILITY=0.8
FAVORITE_PROBABILITY=0.5
SHARE_PROBABILITY=0.3
# 文件存储配置
UPLOAD_DIR=/app/data/uploads
LOG_DIR=/app/data/logs
Regular → Executable
+4 -22
View File
@@ -1,39 +1,21 @@
FROM python:3.11-slim
FROM python:3.10-slim
# 设置工作目录
WORKDIR /app
# 设置环境变量
ENV PYTHONDONTWRITEBYTECODE=1 \
PYTHONUNBUFFERED=1 \
PIP_NO_CACHE_DIR=1 \
PIP_DISABLE_PIP_VERSION_CHECK=1
# 安装系统依赖
# Install system dependencies
RUN apt-get update && apt-get install -y \
gcc \
default-libmysqlclient-dev \
pkg-config \
&& rm -rf /var/lib/apt/lists/*
# 复制依赖文件
COPY requirements.txt .
# 安装 Python 依赖
RUN pip install --no-cache-dir -r requirements.txt
# 复制应用代码
COPY . .
# 创建数据目录
RUN mkdir -p /app/data/uploads /app/data/logs
RUN mkdir -p /app/logs /app/config
# 暴露端口
EXPOSE 8000
# 健康检查
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
CMD python -c "import httpx; httpx.get('http://localhost:8000/health')"
# 启动命令
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000"]
CMD ["uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "8000", "--workers", "2"]
View File
Regular → Executable
+13 -3
View File
@@ -1,3 +1,13 @@
"""
API 路由模块初始化
"""
"""API路由汇总"""
from fastapi import APIRouter
from app.api.endpoints import users, interactions, ai_models, dashboard, system, logs, avatars
router = APIRouter()
router.include_router(users.router, prefix="/users", tags=["虚拟用户管理"])
router.include_router(interactions.router, prefix="/interactions", tags=["互动记录"])
router.include_router(ai_models.router, prefix="/ai-models", tags=["AI模型配置"])
router.include_router(dashboard.router, prefix="/dashboard", tags=["数据看板"])
router.include_router(system.router, prefix="/system", tags=["系统设置"])
router.include_router(logs.router, prefix="/logs", tags=["日志管理"])
router.include_router(avatars.router, prefix="/avatars", tags=["数字分身管理"])
-136
View File
@@ -1,136 +0,0 @@
"""
AI 模型配置 API
"""
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session
from typing import List
from app.models.base import get_db
from app.models.ai_model import AIModelConfig
from app.schemas.ai_model import (
AIModelConfigCreate,
AIModelConfigUpdate,
AIModelConfigResponse,
AIModelTestRequest,
AIModelTestResponse
)
from app.services.ai_service import ai_service
router = APIRouter()
@router.get("", response_model=List[AIModelConfigResponse])
def get_ai_models(
db: Session = Depends(get_db)
):
"""获取所有 AI 模型配置"""
models = db.query(AIModelConfig).all()
return models
@router.get("/{model_id}", response_model=AIModelConfigResponse)
def get_ai_model(
model_id: int,
db: Session = Depends(get_db)
):
"""获取 AI 模型详情"""
model = db.query(AIModelConfig).filter(AIModelConfig.id == model_id).first()
if not model:
raise HTTPException(status_code=404, detail="Model not found")
return model
@router.post("", response_model=AIModelConfigResponse)
def create_ai_model(
model_data: AIModelConfigCreate,
db: Session = Depends(get_db)
):
"""创建 AI 模型配置"""
# 检查是否已存在
existing = db.query(AIModelConfig).filter(
AIModelConfig.model_name == model_data.model_name
).first()
if existing:
raise HTTPException(status_code=400, detail="Model already exists")
model = AIModelConfig(**model_data.model_dump())
# 如果是第一个模型,设为默认
if not db.query(AIModelConfig).filter(AIModelConfig.is_default == True).first():
model.is_default = True
db.add(model)
db.commit()
db.refresh(model)
return model
@router.put("/{model_id}", response_model=AIModelConfigResponse)
def update_ai_model(
model_id: int,
model_data: AIModelConfigUpdate,
db: Session = Depends(get_db)
):
"""更新 AI 模型配置"""
model = db.query(AIModelConfig).filter(AIModelConfig.id == model_id).first()
if not model:
raise HTTPException(status_code=404, detail="Model not found")
update_data = model_data.model_dump(exclude_unset=True)
# 如果设置默认模型,先取消其他模型的默认状态
if update_data.get("is_default"):
db.query(AIModelConfig).update({"is_default": False})
for key, value in update_data.items():
setattr(model, key, value)
db.commit()
db.refresh(model)
return model
@router.delete("/{model_id}")
def delete_ai_model(
model_id: int,
db: Session = Depends(get_db)
):
"""删除 AI 模型配置"""
model = db.query(AIModelConfig).filter(AIModelConfig.id == model_id).first()
if not model:
raise HTTPException(status_code=404, detail="Model not found")
db.delete(model)
db.commit()
return {"message": "Model deleted successfully"}
@router.post("/test", response_model=AIModelTestResponse)
async def test_ai_model(
request: AIModelTestRequest,
db: Session = Depends(get_db)
):
"""测试 AI 模型"""
model = db.query(AIModelConfig).filter(AIModelConfig.id == request.model_id).first()
if not model:
raise HTTPException(status_code=404, detail="Model not found")
model_config = {
"provider": model.provider,
"model_name": model.model_name,
"api_key": model.api_key,
"api_url": model.api_url,
"temperature": model.temperature,
"max_tokens": model.max_tokens,
}
result = await ai_service.test_model(
model_config=model_config,
test_prompt=request.test_prompt
)
return result
-150
View File
@@ -1,150 +0,0 @@
"""
控制台 API
"""
from fastapi import APIRouter, Depends, Query
from sqlalchemy.orm import Session
from typing import List
from datetime import date, timedelta
from app.models.base import get_db
from app.schemas.dashboard import DashboardStats, CoreStats, DailyUsageItem, MonthlyUsageItem
from app.services.token_service import TokenService, get_token_service
from app.services.virtual_user_service import VirtualUserService, get_virtual_user_service
from app.models.virtual_user import VirtualUser, UserStatus
from app.models.interaction import InteractionRecord, InteractionType, InteractionStatus
from app.models.token_usage import TokenUsage
from sqlalchemy import func, and_
router = APIRouter()
@router.get("", response_model=DashboardStats)
def get_dashboard_stats(
db: Session = Depends(get_db),
token_service: TokenService = Depends(get_token_service),
user_service: VirtualUserService = Depends(get_virtual_user_service)
):
"""获取控制台统计数据"""
# 核心指标统计
today = date.today()
yesterday = today - timedelta(days=1)
# 用户统计
total_users = db.query(VirtualUser).count()
active_users = db.query(VirtualUser).filter(VirtualUser.status == UserStatus.ACTIVE).count()
disabled_users = total_users - active_users
# 今日互动统计
today_interactions = db.query(InteractionRecord).filter(
and_(
func.date(InteractionRecord.execution_time) == today,
InteractionRecord.status == InteractionStatus.SUCCESS
)
).all()
today_comments = sum(1 for i in today_interactions if i.interaction_type == InteractionType.COMMENT)
today_replies = sum(1 for i in today_interactions if i.interaction_type == InteractionType.REPLY)
today_likes = sum(1 for i in today_interactions if i.interaction_type == InteractionType.LIKE)
today_favorites = sum(1 for i in today_interactions if i.interaction_type == InteractionType.FAVORITE)
today_shares = sum(1 for i in today_interactions if i.interaction_type == InteractionType.SHARE)
# 昨日互动统计
yesterday_interactions = db.query(InteractionRecord).filter(
and_(
func.date(InteractionRecord.execution_time) == yesterday,
InteractionRecord.status == InteractionStatus.SUCCESS
)
).all()
yesterday_comments = sum(1 for i in yesterday_interactions if i.interaction_type == InteractionType.COMMENT)
yesterday_replies = sum(1 for i in yesterday_interactions if i.interaction_type == InteractionType.REPLY)
# Token 统计
today_tokens = token_service.get_today_usage()
month_tokens = token_service.get_month_usage()
remaining_tokens = token_service.get_remaining_tokens()
core_stats = CoreStats(
total_users=total_users,
active_users=active_users,
disabled_users=disabled_users,
today_comments=today_comments,
today_replies=today_replies,
today_likes=today_likes,
today_favorites=today_favorites,
today_shares=today_shares,
yesterday_comments=yesterday_comments,
yesterday_replies=yesterday_replies,
month_tokens=month_tokens,
today_tokens=today_tokens,
remaining_tokens=remaining_tokens
)
# 每日 Token 使用(近 30 天)
daily_usages = token_service.get_daily_usages(days=30)
daily_items = [DailyUsageItem(date=u["date"], tokens=u["tokens"], comments=0, replies=0) for u in daily_usages]
# 每月 Token 使用(近 12 个月)
monthly_usages = token_service.get_monthly_usages(months=12)
monthly_items = [MonthlyUsageItem(month=u["month"], tokens=u["tokens"]) for u in monthly_usages]
# 最近互动记录
recent_interactions = db.query(InteractionRecord).order_by(
InteractionRecord.execution_time.desc()
).limit(10).all()
return DashboardStats(
core_stats=core_stats,
daily_token_usages=daily_items,
monthly_token_usages=monthly_items,
recent_interactions=[
{
"id": r.id,
"virtual_user_id": r.virtual_user_id,
"interaction_type": r.interaction_type.value,
"status": r.status.value,
"execution_time": r.execution_time
}
for r in recent_interactions
]
)
@router.get("/token/stats")
def get_token_stats(
db: Session = Depends(get_db),
token_service: TokenService = Depends(get_token_service)
):
"""获取 Token 统计"""
today_used = token_service.get_today_usage()
today_limit = 10000 # TODO: 从系统配置读取
return {
"today_used": today_used,
"today_limit": today_limit,
"today_remaining": max(0, today_limit - today_used),
"usage_percentage": round((today_used / today_limit) * 100, 2) if today_limit > 0 else 0
}
@router.get("/token/daily", response_model=List[DailyUsageItem])
def get_daily_token_usage(
days: int = Query(30, ge=1, le=90, description="天数"),
db: Session = Depends(get_db),
token_service: TokenService = Depends(get_token_service)
):
"""获取每日 Token 使用"""
usages = token_service.get_daily_usages(days=days)
return [DailyUsageItem(date=u["date"], tokens=u["tokens"], comments=0, replies=0) for u in usages]
@router.get("/token/monthly", response_model=List[MonthlyUsageItem])
def get_monthly_token_usage(
months: int = Query(12, ge=1, le=24, description="月数"),
db: Session = Depends(get_db),
token_service: TokenService = Depends(get_token_service)
):
"""获取每月 Token 使用"""
usages = token_service.get_monthly_usages(months=months)
return [MonthlyUsageItem(month=u["month"], tokens=u["tokens"]) for u in usages]
View File
+141
View File
@@ -0,0 +1,141 @@
"""AI模型配置接口"""
import secrets
from fastapi import APIRouter, Depends, Header, HTTPException
from sqlalchemy import select, update
from app.core.database import get_db
from app.core.config import settings
from app.schemas import ApiResponse, AIModelCreateRequest, AIModelUpdateRequest, AIModelTestRequest
from app.models import AIModelConfig
from app.utils.crypto import encrypt, decrypt
from app.services.ai_service import ai_service
router = APIRouter()
@router.get("")
async def list_models(db=Depends(get_db)):
result = await db.execute(select(AIModelConfig).order_by(AIModelConfig.created_at.desc()))
models = result.scalars().all()
items = [_format_model(m) for m in models]
return ApiResponse(data=items)
@router.post("")
async def create_model(req: AIModelCreateRequest, db=Depends(get_db)):
if req.is_default:
await db.execute(
update(AIModelConfig)
.where(AIModelConfig.usage_scope == req.usage_scope)
.values(is_default=0)
)
model = AIModelConfig(
model_name=req.model_name,
provider=req.provider,
usage_scope=req.usage_scope,
api_base_url=req.api_base_url,
api_key_enc=encrypt(req.api_key) if req.api_key else None,
model_version=req.model_version,
vision_model_version=req.vision_model_version,
ocr_model_version=req.ocr_model_version,
temperature=req.temperature,
max_tokens=req.max_tokens,
timeout_seconds=req.timeout_seconds,
is_default=req.is_default,
is_enabled=1,
)
db.add(model)
await db.commit()
await db.refresh(model)
return ApiResponse(data=_format_model(model), message="模型添加成功")
@router.put("/{model_id}")
async def update_model(model_id: int, req: AIModelUpdateRequest, db=Depends(get_db)):
result = await db.execute(select(AIModelConfig).where(AIModelConfig.id == model_id))
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="模型不存在")
target_scope = req.usage_scope or model.usage_scope
if req.is_default or (req.usage_scope and model.is_default):
await db.execute(
update(AIModelConfig)
.where(
AIModelConfig.id != model_id,
AIModelConfig.usage_scope == target_scope,
)
.values(is_default=0)
)
for field, val in req.model_dump(exclude_none=True).items():
if field == "api_key":
model.api_key_enc = encrypt(val) if val else None
else:
setattr(model, field, val)
await db.commit()
await db.refresh(model)
return ApiResponse(data=_format_model(model), message="更新成功")
@router.get("/runtime/digital-avatar")
async def get_digital_avatar_runtime_model(
x_avatar_config_token: str | None = Header(default=None),
db=Depends(get_db),
):
expected = settings.AVATAR_MODEL_CONFIG_TOKEN
if not expected:
raise HTTPException(status_code=503, detail="数字分身模型配置服务未启用")
if not x_avatar_config_token or not secrets.compare_digest(x_avatar_config_token, expected):
raise HTTPException(status_code=401, detail="无权读取数字分身模型配置")
result = await db.execute(
select(AIModelConfig).where(
AIModelConfig.usage_scope == "digital_avatar",
AIModelConfig.is_default == 1,
AIModelConfig.is_enabled == 1,
)
)
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="尚未配置启用的数字分身专用模型")
return ApiResponse(data={
"api_base_url": model.api_base_url or "https://api.openai.com/v1",
"api_key": decrypt(model.api_key_enc) if model.api_key_enc else "",
"model": model.model_version or model.model_name,
"vision_model": model.vision_model_version or "qwen3.6-flash",
"ocr_model": model.ocr_model_version or "qwen-vl-ocr",
"temperature": model.temperature,
"max_tokens": model.max_tokens,
"timeout_seconds": model.timeout_seconds,
})
@router.delete("/{model_id}")
async def delete_model(model_id: int, db=Depends(get_db)):
result = await db.execute(select(AIModelConfig).where(AIModelConfig.id == model_id))
model = result.scalar_one_or_none()
if not model:
raise HTTPException(status_code=404, detail="模型不存在")
await db.delete(model)
await db.commit()
return ApiResponse(message="删除成功")
@router.post("/test")
async def test_model(req: AIModelTestRequest, db=Depends(get_db)):
result = await ai_service.test_model(db, req.model_id, req.test_prompt)
return ApiResponse(data=result)
def _format_model(m: AIModelConfig) -> dict:
return {
"id": m.id, "model_name": m.model_name, "provider": m.provider,
"usage_scope": m.usage_scope,
"api_base_url": m.api_base_url, "has_api_key": bool(m.api_key_enc),
"model_version": m.model_version, "temperature": m.temperature,
"vision_model_version": m.vision_model_version,
"ocr_model_version": m.ocr_model_version,
"max_tokens": m.max_tokens, "timeout_seconds": m.timeout_seconds,
"is_default": m.is_default, "is_enabled": m.is_enabled,
"created_at": m.created_at.isoformat(),
}
+201
View File
@@ -0,0 +1,201 @@
"""数字分身管理 API 端点"""
from typing import Optional
import mimetypes
import uuid
from fastapi import APIRouter, Query, HTTPException, UploadFile, File
from fastapi.responses import JSONResponse
from app.schemas import ApiResponse
from app.services.avatar_service import avatar_service, get_session, is_available
from app.core.database import AsyncSessionLocal
from app.core.config import settings
router = APIRouter()
@router.get("")
def list_avatars(
page: int = Query(default=1, ge=1),
page_size: int = Query(default=20, ge=1, le=100),
keyword: str = Query(default=None),
status: str = Query(default=None),
):
"""分页查询所有数字分身"""
if not is_available():
return ApiResponse(data={"total": 0, "page": page, "page_size": page_size, "items": []})
db = get_session()
try:
total, items = avatar_service.list_avatars(
db, page, page_size, keyword, status
)
return ApiResponse(
data={
"total": total,
"page": page,
"page_size": page_size,
"items": items,
}
)
finally:
db.close()
@router.get("/{avatar_id}")
def get_avatar(avatar_id: str):
"""获取单个数字分身详情"""
if not is_available():
raise HTTPException(status_code=503, detail="数字分身数据库尚未初始化")
db = get_session()
try:
data = avatar_service.get_avatar(db, avatar_id)
if not data:
raise HTTPException(status_code=404, detail="数字分身不存在")
return ApiResponse(data=data)
finally:
db.close()
@router.put("/{avatar_id}/status")
def update_avatar_status(avatar_id: str, body: dict):
"""更新数字分身状态(开机/关机)"""
status = body.get("status")
if status not in ("active", "inactive", "training"):
raise HTTPException(
status_code=400,
detail="status 必须是 active、inactive 或 training",
)
if not is_available():
raise HTTPException(status_code=503, detail="数字分身数据库尚未初始化")
db = get_session()
try:
data = avatar_service.update_status(db, avatar_id, status)
if not data:
raise HTTPException(status_code=404, detail="数字分身不存在")
msg = "已开机" if status == "active" else "已关机"
return ApiResponse(data=data, message=msg)
finally:
db.close()
async def _upload_to_filecenter(file_bytes: bytes, filename: str) -> str:
"""利用会会平台 filecenter 上传头像,返回 URL"""
import httpx
import hashlib
import random
from datetime import datetime
async with AsyncSessionLocal() as db:
# 读取平台配置
from sqlalchemy import select
from app.models import SystemConfig
result = await db.execute(select(SystemConfig))
configs = {row.config_key: row.config_value for row in result.scalars().all()}
cfg = {
"appId": configs.get("platform_app_id", ""),
"accessId": configs.get("platform_access_id", ""),
"accessSecret": configs.get("platform_access_secret", ""),
}
biz_url = configs.get("news_platform_base_url", "http://192.168.1.200:63120")
if not cfg["accessSecret"]:
raise RuntimeError("平台 accessSecret 未配置")
# 构建签名参数
nonce = str(random.random())[2:][: random.randint(8, 12)]
timestamp = datetime.now().strftime("%Y%m%d%I%M%S")
sign_params = {
"appId": cfg["appId"],
"accessId": cfg["accessId"],
"timestamp": timestamp,
"nonce": nonce,
"module": "userInfo",
"service": "kccloud",
}
# 计算签名
keys = sorted(sign_params.keys())
sign_parts = []
for k in keys:
v = sign_params.get(k)
if v and v != "" and v != []:
sign_parts.append(f"{k}={v}")
sign_str = "&".join(sign_parts) + f"&accessSecret={cfg['accessSecret']}"
signature = hashlib.md5(sign_str.encode("utf-8")).hexdigest().upper()
# 构建 filecenter URL
filecenter_url = biz_url.replace("/huihuibusiness", "/filecenter")
if "/api/" in filecenter_url:
filecenter_url = filecenter_url.split("/api/", 1)[0] + "/api/filecenter"
else:
filecenter_url = filecenter_url.rstrip("/") + "/filecenter"
# 确定 MIME 类型
mime = mimetypes.guess_type(filename)[0] or "image/jpeg"
files = {"file": (filename, file_bytes, mime)}
# 发送请求
async with httpx.AsyncClient(timeout=30) as client:
r = await client.post(
f"{filecenter_url}/fileUpload",
files=files,
data={**sign_params, "signature": signature},
)
d = r.json()
if d.get("code") in [0, 200]:
url = d.get("data") or d.get("url") or ""
if isinstance(url, dict):
url = url.get("url") or url.get("path") or ""
if not url:
raise RuntimeError(f"filecenter 返回空 URL: {d}")
return url
raise RuntimeError(f"filecenter 上传失败: {d.get('message', '未知错误')}")
@router.post("/{avatar_id}/upload-photo")
async def upload_avatar_photo(
avatar_id: str,
file: UploadFile = File(...),
):
"""上传数字分身头像(通过会会平台 filecenter)"""
if not is_available():
raise HTTPException(status_code=503, detail="数字分身数据库尚未初始化")
# 验证文件
if not file.content_type or not file.content_type.startswith("image/"):
raise HTTPException(status_code=400, detail="仅支持图片文件")
file_bytes = await file.read()
if len(file_bytes) > 5 * 1024 * 1024:
raise HTTPException(status_code=400, detail="头像文件不能超过5MB")
# 确保文件扩展名正确
filename = file.filename or "avatar.jpg"
ext = filename.split(".")[-1].lower() if "." in filename else "jpg"
if ext not in ("jpg", "jpeg", "png", "gif", "webp"):
filename = f"avatar.{ext}"
try:
photo_url = await _upload_to_filecenter(file_bytes, filename)
except Exception as e:
raise HTTPException(status_code=500, detail=f"头像上传失败: {str(e)}")
# 更新数据库
db = get_session()
try:
data = avatar_service.get_avatar(db, avatar_id)
if not data:
raise HTTPException(status_code=404, detail="数字分身不存在")
from sqlalchemy import text
db.execute(
text("UPDATE avatars SET photo_url = :url, updated_at = datetime('now') WHERE id = :id"),
{"url": photo_url, "id": avatar_id},
)
db.commit()
# 刷新数据
updated = avatar_service.get_avatar(db, avatar_id)
return ApiResponse(data=updated, message="头像上传成功")
finally:
db.close()
+25
View File
@@ -0,0 +1,25 @@
"""数据看板接口"""
from fastapi import APIRouter, Depends, Query
from app.core.database import get_db
from app.schemas import ApiResponse
from app.services.stats_service import stats_service
router = APIRouter()
@router.get("")
async def get_dashboard(db=Depends(get_db)):
data = await stats_service.get_dashboard(db)
return ApiResponse(data=data)
@router.get("/token-trend")
async def get_token_trend(days: int = Query(default=30, ge=7, le=90), db=Depends(get_db)):
trend = await stats_service.get_token_trend(db, days)
return ApiResponse(data=trend)
@router.get("/monthly-token-trend")
async def get_monthly_token_trend(db=Depends(get_db)):
trend = await stats_service.get_monthly_token_trend(db)
return ApiResponse(data=trend)
+249
View File
@@ -0,0 +1,249 @@
"""互动记录接口"""
from typing import Optional
from fastapi import APIRouter, Depends, Query, HTTPException
from fastapi.responses import StreamingResponse
import io, pandas as pd
from app.core.database import get_db
from app.schemas import ApiResponse
from app.services.stats_service import stats_service
from app.models import InteractionRecord
from sqlalchemy import select, update
router = APIRouter()
@router.get("")
async def list_interactions(
page: int = Query(default=1, ge=1),
page_size: int = Query(default=20, ge=1, le=100),
user_id: Optional[int] = None,
interact_type: Optional[str] = None,
status: Optional[int] = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
keyword: Optional[str] = None,
db=Depends(get_db)
):
result = await stats_service.get_interaction_records(
db, page, page_size, user_id, interact_type, status, start_date, end_date, keyword
)
return ApiResponse(data=result)
@router.post("/{record_id}/retry")
async def retry_interaction(record_id: int, db=Depends(get_db)):
"""手动重试失败任务"""
result = await db.execute(select(InteractionRecord).where(InteractionRecord.id == record_id))
record = result.scalar_one_or_none()
if not record:
raise HTTPException(status_code=404, detail="记录不存在")
if record.status != 2:
raise HTTPException(status_code=400, detail="只能重试失败的任务")
if record.retry_count >= 3:
raise HTTPException(status_code=400, detail="已超过最大重试次数(3次)")
from app.services.news_service import news_service
from app.services.ai_service import ai_service
from app.models import VirtualUser, UserPersonality
user_result = await db.execute(select(VirtualUser).where(VirtualUser.id == record.user_id))
user = user_result.scalar_one_or_none()
if not user or user.status != 2:
raise HTTPException(status_code=400, detail="用户未登录,无法重试")
success, err, platform_record_id = False, "未知类型", ""
if record.interact_type == "comment" and record.content:
success, err, platform_record_id = await news_service.post_comment_with_record_id(
db, user, record.article_id, record.article_title or "", record.content
)
elif record.interact_type == "like":
success, err = await news_service.like_news(db, user, record.article_id, org_id="", title=record.article_title or "")
elif record.interact_type == "collect":
success, err = await news_service.collect_news(db, user, record.article_id, title=record.article_title or "")
elif record.interact_type == "forward":
success, err = await news_service.forward_news(db, user, record.article_id)
elif record.interact_type == "reply" and record.content:
parent_comment = None
if record.parent_comment_id:
comments = await news_service.get_comments(db, user, record.article_id or "")
parent_comment = next(
(
c for c in comments
if str(c.get("id") or c.get("commentId") or "") == str(record.parent_comment_id)
),
None,
)
if not parent_comment:
success, err, platform_record_id = False, "未找到被回复的原评论,无法重试回复", ""
else:
success, err, platform_record_id = await news_service.post_reply_with_record_id(
db, user,
record.article_id or "",
record.parent_comment_id or "",
record.content,
parent_comment=parent_comment,
article_title=record.article_title or "",
)
await db.execute(
update(InteractionRecord).where(InteractionRecord.id == record_id).values(
status=1 if success else 2,
error_msg=None if success else err,
platform_record_id=platform_record_id or record.platform_record_id,
retry_count=record.retry_count + 1,
)
)
await db.commit()
return ApiResponse(message="重试成功" if success else f"重试失败: {err}")
@router.get("/export")
async def export_interactions(
user_id: Optional[int] = None,
interact_type: Optional[str] = None,
status: Optional[int] = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
db=Depends(get_db)
):
"""导出互动记录"""
data = await stats_service.get_interaction_records(
db, 1, 10000, user_id, interact_type, status, start_date, end_date
)
rows = [{
"ID": r["id"], "用户昵称": r["user_nickname"], "用户账号": r["user_account"],
"文章标题": r["article_title"], "互动类型": r["interact_type_label"],
"内容": r["content"] or "", "Token消耗": r["token_consumed"],
"状态": r["status_label"], "失败原因": r["error_msg"] or "",
"重试次数": r["retry_count"], "执行时间": r["executed_at"],
} for r in data["items"]]
df = pd.DataFrame(rows)
buf = io.BytesIO()
df.to_excel(buf, index=False, sheet_name="互动记录")
buf.seek(0)
return StreamingResponse(
buf,
media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
headers={"Content-Disposition": "attachment; filename=interactions_export.xlsx"}
)
@router.post("/{record_id}/cancel")
async def cancel_interaction(record_id: int, db=Depends(get_db)):
"""取消互动(取消点赞/收藏/删除评论),转发不支持取消"""
from sqlalchemy import select, update
from app.models import InteractionRecord, VirtualUser
from app.services.news_service import news_service
# 查找互动记录
r = await db.execute(select(InteractionRecord).where(InteractionRecord.id == record_id))
record = r.scalar_one_or_none()
if not record:
return ApiResponse(code=404, message="记录不存在")
if record.status != 1:
return ApiResponse(code=400, message="只能取消成功的互动")
if record.interact_type == "forward":
return ApiResponse(code=400, message="转发互动无法取消")
if record.interact_type == "read":
return ApiResponse(code=400, message="阅读记录无法取消")
# 查找对应用户
ur = await db.execute(select(VirtualUser).where(VirtualUser.id == record.user_id))
user = ur.scalar_one_or_none()
if not user:
return ApiResponse(code=404, message="用户不存在")
# 执行取消
ok = False
err = ""
if record.interact_type in ("like",):
ok, err = await news_service.cancel_like(
db, user,
news_id=record.article_id or "",
org_id=record.session_id or "", # session_id 字段暂存 org_id
title=record.article_title or "",
)
elif record.interact_type == "collect":
ok, err = await news_service.cancel_collect(
db, user,
news_id=record.article_id or "",
title=record.article_title or "",
)
elif record.interact_type == "comment":
comment_id = record.platform_record_id or ""
if not comment_id:
comment_id, comment_count = await news_service.find_comment_id(
db, user,
news_id=record.article_id or "",
content=record.content or "",
)
if comment_id:
await db.execute(
update(InteractionRecord).where(InteractionRecord.id == record_id).values(
platform_record_id=comment_id
)
)
await db.commit()
elif comment_count == 0:
await db.execute(
update(InteractionRecord).where(InteractionRecord.id == record_id).values(
status=3,
error_msg="平台未查询到评论,已标记取消",
)
)
await db.commit()
return ApiResponse(message="平台未查询到评论,已标记取消")
else:
return ApiResponse(code=400, message="缺少评论ID,且未能从平台评论列表反查到该评论")
ok, err = await news_service.cancel_comment(
db, user,
news_id=record.article_id or "",
comment_id=comment_id,
)
elif record.interact_type == "reply":
reply_id = record.platform_record_id or ""
if not reply_id:
reply_id, reply_count = await news_service.find_reply_id(
db, user,
news_id=record.article_id or "",
content=record.content or "",
parent_comment_id=record.parent_comment_id or "",
)
if reply_id:
await db.execute(
update(InteractionRecord).where(InteractionRecord.id == record_id).values(
platform_record_id=reply_id
)
)
await db.commit()
elif reply_count == 0:
await db.execute(
update(InteractionRecord).where(InteractionRecord.id == record_id).values(
status=3,
error_msg="平台未查询到回复,已标记取消",
)
)
await db.commit()
return ApiResponse(message="平台未查询到回复,已标记取消")
else:
return ApiResponse(code=400, message="缺少回复ID,且未能从平台回复列表反查到该回复")
ok, err = await news_service.cancel_reply(
db, user,
news_id=record.article_id or "",
reply_id=reply_id,
)
if ok:
# 更新状态为手动取消(status=3)
await db.execute(
update(InteractionRecord).where(InteractionRecord.id == record_id).values(
status=3, error_msg="手动取消"
)
)
await db.commit()
return ApiResponse(message="取消成功")
else:
return ApiResponse(code=500, message=f"取消失败: {err}")
+83
View File
@@ -0,0 +1,83 @@
"""日志管理接口"""
import os
from typing import Optional
from fastapi import APIRouter, Depends, Query, HTTPException
from fastapi.responses import FileResponse
from sqlalchemy import select, func, and_
from app.core.database import get_db
from app.schemas import ApiResponse
from app.models import LoginLog
from app.core.config import settings
router = APIRouter()
@router.get("/login")
async def get_login_logs(
page: int = Query(default=1, ge=1),
page_size: int = Query(default=50, ge=1, le=200),
user_id: Optional[int] = None,
action: Optional[str] = None,
db=Depends(get_db)
):
query = select(LoginLog)
conditions = []
if user_id:
conditions.append(LoginLog.user_id == user_id)
if action:
conditions.append(LoginLog.action == action)
if conditions:
query = query.where(and_(*conditions))
total = (await db.execute(select(func.count()).select_from(query.subquery()))).scalar()
query = query.order_by(LoginLog.created_at.desc()).offset((page - 1) * page_size).limit(page_size)
result = await db.execute(query)
logs = result.scalars().all()
items = [{
"id": l.id, "user_id": l.user_id, "user_account": l.user_account,
"action": l.action, "session_id": l.session_id,
"error_msg": l.error_msg, "created_at": l.created_at.strftime("%Y-%m-%dT%H:%M:%S+08:00") if l.created_at else None
} for l in logs]
return ApiResponse(data={"total": total, "page": page, "page_size": page_size, "items": items})
@router.get("/files")
async def list_log_files():
"""列出日志文件"""
log_dir = settings.LOG_DIR
files = []
if os.path.exists(log_dir):
for fname in sorted(os.listdir(log_dir), reverse=True):
if fname.endswith(".log"):
fpath = os.path.join(log_dir, fname)
size = os.path.getsize(fpath)
files.append({"name": fname, "size": size,
"size_kb": round(size / 1024, 1)})
return ApiResponse(data=files)
@router.get("/files/{filename}/tail")
async def tail_log_file(filename: str, lines: int = Query(default=100, ge=10, le=1000)):
"""读取日志文件末尾"""
# 安全校验
if ".." in filename or "/" in filename:
raise HTTPException(status_code=400, detail="非法文件名")
fpath = os.path.join(settings.LOG_DIR, filename)
if not os.path.exists(fpath):
raise HTTPException(status_code=404, detail="文件不存在")
with open(fpath, "r", encoding="utf-8", errors="replace") as f:
all_lines = f.readlines()
tail = all_lines[-lines:]
return ApiResponse(data={"filename": filename, "lines": tail, "total_lines": len(all_lines)})
@router.get("/files/{filename}/download")
async def download_log_file(filename: str):
"""下载日志文件"""
if ".." in filename or "/" in filename:
raise HTTPException(status_code=400, detail="非法文件名")
fpath = os.path.join(settings.LOG_DIR, filename)
if not os.path.exists(fpath):
raise HTTPException(status_code=404, detail="文件不存在")
return FileResponse(fpath, filename=filename, media_type="text/plain")
+115
View File
@@ -0,0 +1,115 @@
"""系统设置接口"""
from fastapi import APIRouter, Depends
from sqlalchemy import select, update as sql_update
from app.core.database import get_db
from app.schemas import ApiResponse
from app.models import SystemConfig
router = APIRouter()
@router.get("/configs")
async def get_configs(db=Depends(get_db)):
result = await db.execute(select(SystemConfig).order_by(SystemConfig.config_key))
configs = result.scalars().all()
data = {c.config_key: {"value": c.config_value, "type": c.config_type, "desc": c.description}
for c in configs}
return ApiResponse(data=data)
@router.put("/configs")
async def update_configs(body: dict, db=Depends(get_db)):
"""批量更新配置"""
for key, value in body.items():
result = await db.execute(select(SystemConfig).where(SystemConfig.config_key == key))
cfg = result.scalar_one_or_none()
if cfg:
cfg.config_value = str(value)
else:
db.add(SystemConfig(config_key=key, config_value=str(value)))
await db.commit()
return ApiResponse(message="配置已保存")
@router.post("/scheduler/toggle")
async def toggle_scheduler(body: dict, db=Depends(get_db)):
enabled = body.get("enabled", True)
result = await db.execute(select(SystemConfig).where(SystemConfig.config_key == "scheduler_enabled"))
cfg = result.scalar_one_or_none()
if cfg:
cfg.config_value = "true" if enabled else "false"
await db.commit()
return ApiResponse(message=f"调度器已{'启用' if enabled else '暂停'}")
@router.post("/sessions/reset-all")
async def reset_all_sessions(db=Depends(get_db)):
"""重置所有用户会话"""
from app.models import VirtualUser
await db.execute(
sql_update(VirtualUser).values(status=0, session_token=None, session_expires_at=None)
)
await db.commit()
return ApiResponse(message="所有会话已重置")
@router.post("/login/diagnose")
async def diagnose_login(body: dict, db=Depends(get_db)):
"""
诊断登录接口原始响应 - 临时调试用
传入: {"username": "xxx", "password": "xxx"}
"""
import httpx, hashlib, uuid
from datetime import datetime
from app.services.news_service import news_service
cfg = await news_service._client(db)
auth = await news_service._auth_url(db)
username = body.get("username", "")
password = body.get("password", "")
# 构建 formData(与真实登录完全一致)
extra = {
"username": username,
"password": password,
"loginType": "password",
"grantType": "password",
"isRegister": "false",
}
if cfg.get("clientCode"):
extra["clientCode"] = cfg["clientCode"]
form = news_service._build_form(extra, cfg)
try:
async with httpx.AsyncClient(timeout=15) as c:
resp = await c.post(f"{auth}/open/login/token", data=form)
# 返回完整诊断信息
try:
resp_json = resp.json()
except Exception:
resp_json = None
return ApiResponse(data={
"status_code": resp.status_code,
"response_text": resp.text[:2000],
"response_json": resp_json,
"request_url": f"{auth}/open/login/token",
"request_form": {k: v if k not in ("password","accessSecret") else "***" for k, v in form.items()},
"content_type": resp.headers.get("content-type", ""),
})
except Exception as e:
return ApiResponse(code=500, message=str(e), data={"error": str(e)})
@router.post("/interaction/run-now")
async def run_interaction_now(db=Depends(get_db)):
"""立即触发一次互动任务(不受时间段限制)"""
from app.services.scheduler import scheduler_service
try:
result = await scheduler_service.run_once_now(db)
return ApiResponse(data=result, message="互动任务已触发")
except Exception as e:
return ApiResponse(code=500, message=f"触发失败: {e}")
+396
View File
@@ -0,0 +1,396 @@
"""虚拟用户管理接口"""
from typing import Optional
from pathlib import Path
import uuid
import mimetypes
from fastapi import APIRouter, Depends, Query, UploadFile, File, HTTPException
from fastapi.responses import StreamingResponse
import io
from app.core.database import get_db
from app.schemas import ApiResponse, UserCreateRequest, UserUpdateRequest, UserBatchRequest, PersonalityUpdateRequest
from app.services.user_service import user_service
router = APIRouter()
_UPLOADS_DIR = Path(__file__).resolve().parents[2] / "uploads" / "avatars"
_UPLOADS_DIR.mkdir(parents=True, exist_ok=True)
async def _save_local_avatar(file_bytes: bytes, filename: str, content_type: str | None) -> str:
ext = Path(filename or "").suffix.lower()
if ext not in {".jpg", ".jpeg", ".png", ".gif", ".webp"}:
guessed_ext = mimetypes.guess_extension(content_type or "") or ".jpg"
ext = ".jpg" if guessed_ext == ".jpe" else guessed_ext
if ext not in {".jpg", ".jpeg", ".png", ".gif", ".webp"}:
ext = ".jpg"
target = _UPLOADS_DIR / f"{uuid.uuid4().hex}{ext}"
target.write_bytes(file_bytes)
return f"/api/uploads/avatars/{target.name}"
@router.get("")
async def list_users(
page: int = Query(default=1, ge=1),
page_size: int = Query(default=20, ge=1, le=100),
keyword: Optional[str] = Query(default=None),
status: Optional[int] = Query(default=None),
is_enabled: Optional[int] = Query(default=None),
db=Depends(get_db)
):
"""获取虚拟用户列表"""
total, items = await user_service.get_users(db, page, page_size, keyword, status, is_enabled)
return ApiResponse(data={"total": total, "page": page, "page_size": page_size, "items": items})
@router.post("")
async def create_user(req: UserCreateRequest, db=Depends(get_db)):
"""创建虚拟用户"""
user = await user_service.create_user(db, req)
return ApiResponse(data=user, message="用户创建成功")
@router.get("/{user_id}")
async def get_user(user_id: int, db=Depends(get_db)):
"""获取单个用户详情"""
total, items = await user_service.get_users(db, 1, 1)
from sqlalchemy import select
from app.models import VirtualUser, UserPersonality
from app.services.user_service import user_service as svc
result = await db.execute(select(VirtualUser).where(VirtualUser.id == user_id))
user = result.scalar_one_or_none()
if not user:
raise HTTPException(status_code=404, detail="用户不存在")
p_result = await db.execute(select(UserPersonality).where(UserPersonality.user_id == user_id))
personality = p_result.scalar_one_or_none()
return ApiResponse(data=svc._format_user(user, personality))
@router.put("/{user_id}")
async def update_user(user_id: int, req: UserUpdateRequest, db=Depends(get_db)):
"""更新用户信息(sync_to_platform=true 时同步到目标平台)"""
result = await user_service.update_user(db, user_id, req)
if req.sync_to_platform:
from sqlalchemy import select
from app.models import VirtualUser as _VU
from app.services.news_service import news_service
ur = await db.execute(select(_VU).where(_VU.id == user_id))
user = ur.scalar_one_or_none()
if user and user.status == 2:
ok, err = await news_service.update_user_profile(
db, user,
nick_name=req.nickname,
real_name=req.real_name,
sex=req.sex,
description=req.description,
email=req.email,
)
if not ok:
return ApiResponse(data=result,
message=f"本地已保存,同步到平台失败: {err}", code=206)
return ApiResponse(data=result, message="更新成功")
@router.delete("/{user_id}")
async def delete_user(user_id: int, db=Depends(get_db)):
"""删除用户"""
await user_service.delete_user(db, user_id)
return ApiResponse(message="删除成功")
@router.post("/batch/action")
async def batch_action(req: UserBatchRequest, db=Depends(get_db)):
"""批量操作用户"""
result = await user_service.batch_action(db, req.user_ids, req.action)
return ApiResponse(data=result, message="批量操作成功")
@router.post("/{user_id}/login")
async def manual_login(user_id: int, db=Depends(get_db)):
"""手动触发用户登录"""
from app.services.news_service import news_service
from sqlalchemy import select
from app.models import VirtualUser
result = await db.execute(select(VirtualUser).where(VirtualUser.id == user_id))
user = result.scalar_one_or_none()
if not user:
raise HTTPException(status_code=404, detail="用户不存在")
success = await news_service.login(db, user)
if success:
return ApiResponse(message="登录成功")
raise HTTPException(status_code=400, detail="登录失败,请检查账号密码")
@router.post("/{user_id}/logout")
async def manual_logout(user_id: int, db=Depends(get_db)):
"""手动登出"""
from app.services.news_service import news_service
await news_service.logout(db, user_id)
return ApiResponse(message="已登出")
@router.post("/{user_id}/personality/generate")
async def generate_personality(user_id: int, db=Depends(get_db)):
"""重新生成AI人格"""
personality = await user_service.generate_personality(db, user_id)
return ApiResponse(data=personality, message="人格生成成功")
@router.put("/{user_id}/personality")
async def update_personality(user_id: int, req: PersonalityUpdateRequest, db=Depends(get_db)):
"""手动编辑人格属性"""
personality = await user_service.update_personality(db, user_id, req)
return ApiResponse(data=personality, message="人格更新成功")
@router.get("/excel/template")
async def download_template():
"""下载Excel导入模板"""
content = await user_service.get_excel_template()
return StreamingResponse(
io.BytesIO(content),
media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
headers={"Content-Disposition": "attachment; filename=virtual_users_template.xlsx"}
)
@router.post("/excel/import")
async def import_excel(file: UploadFile = File(...), db=Depends(get_db)):
"""Excel批量导入"""
if not file.filename.endswith((".xlsx", ".xls")):
raise HTTPException(status_code=400, detail="仅支持Excel文件(.xlsx/.xls)")
content = await file.read()
result = await user_service.import_from_excel(db, content)
return ApiResponse(data=result, message=f"导入完成:成功{result['success']}条,失败{result['failed']}条")
@router.get("/excel/export")
async def export_excel(db=Depends(get_db)):
"""导出用户数据Excel"""
content = await user_service.export_to_excel(db)
return StreamingResponse(
io.BytesIO(content),
media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
headers={"Content-Disposition": "attachment; filename=virtual_users_export.xlsx"}
)
@router.post("/deduplicate")
async def deduplicate_users(db=Depends(get_db)):
"""删除重复用户(保留最早创建的一条)"""
from sqlalchemy import text
# 找出重复的账号,保留 id 最小的,删除其他的
result = await db.execute(
text("""
DELETE FROM virtual_users
WHERE id NOT IN (
SELECT MIN(id) FROM virtual_users GROUP BY account
)
""")
)
await db.commit()
deleted = result.rowcount
return ApiResponse(data={"deleted": deleted}, message=f"已清理 {deleted} 条重复数据")
@router.post("/clear-all")
async def clear_all_users(db=Depends(get_db)):
"""清空所有用户(慎用)"""
from sqlalchemy import text
from app.core.redis_client import get_redis
await db.execute(text("DELETE FROM user_personalities"))
await db.execute(text("DELETE FROM virtual_users"))
await db.commit()
return ApiResponse(message="已清空所有用户数据")
@router.post("/login-all")
async def batch_login_all(db=Depends(get_db)):
"""一键登录所有未登录/登录失效的用户"""
from sqlalchemy import select
from app.services.news_service import news_service
from app.models import VirtualUser as _VU
from app.core.database import AsyncSessionLocal
import asyncio
# 先用当前 session 查出所有待登录用户 ID
result_r = await db.execute(
select(_VU.id, _VU.account).where(
_VU.is_enabled == 1,
_VU.status.in_([0, 3])
)
)
rows = result_r.all()
if not rows:
return ApiResponse(message="没有需要登录的用户", data={"count": 0})
user_ids = [r[0] for r in rows]
total = len(user_ids)
success = failed = 0
# 每个用户独立 session,避免事务污染
async def login_one(uid: int):
async with AsyncSessionLocal() as s:
try:
ur = await s.execute(select(_VU).where(_VU.id == uid))
u = ur.scalar_one_or_none()
if u:
return await news_service.login(s, u)
except Exception as e:
logger.warning(f"login_one {uid} 异常: {e}")
return False
return False
batch_size = 5
for i in range(0, total, batch_size):
batch_ids = user_ids[i:i+batch_size]
results = await asyncio.gather(*[login_one(uid) for uid in batch_ids], return_exceptions=True)
for r in results:
if r is True: success += 1
else: failed += 1
if i + batch_size < total:
await asyncio.sleep(1) # 批次间隔避免过于集中
return ApiResponse(
message=f"登录完成:成功 {success} 个,失败 {failed} 个",
data={"success": success, "failed": failed, "total": total}
)
@router.post("/sync-all-profiles")
async def sync_all_profiles(db=Depends(get_db)):
"""
同步所有已登录用户的平台信息(昵称/真实姓名/性别/头像)到本系统
从登录 session 中的 token 调用目标平台接口获取最新用户信息
"""
from sqlalchemy import select, update
from app.models import VirtualUser as _VU
from app.core.database import AsyncSessionLocal
from app.core.redis_client import get_session
import httpx, asyncio
# 查出所有已登录用户
result_r = await db.execute(select(_VU).where(_VU.status == 2, _VU.is_enabled == 1))
users = result_r.scalars().all()
if not users:
return ApiResponse(message="没有已登录的用户", data={"synced": 0})
synced = failed = 0
async def sync_one(uid: int):
"""从登录 session 中提取已缓存的用户信息,直接写入数据库,无需调用外部接口"""
async with AsyncSessionLocal() as s:
try:
sess = await get_session(uid)
if not sess:
return False
ur = await s.execute(select(_VU).where(_VU.id == uid))
user = ur.scalar_one_or_none()
if not user:
return False
platform_uid = sess.get("platform_uid", "")
# 登录成功时 session 里已存有用户信息
vals = {}
if platform_uid: vals["platform_uid"] = platform_uid
# session 里的字段(登录时写入)
sync_nickname = sess.get("nickname", "")
sync_real_name = sess.get("real_name", "")
resolved_nickname = news_service._resolve_synced_nickname(user, sync_nickname, sync_real_name)
if resolved_nickname: vals["nickname"] = resolved_nickname
if sync_real_name: vals["real_name"] = sync_real_name
if sess.get("sex"): vals["sex"] = int(sess["sex"])
if sess.get("avatar"): vals["avatar_url"] = sess["avatar"]
if vals:
await s.execute(update(_VU).where(_VU.id == uid).values(**vals))
await s.commit()
return True
except Exception as e:
logger.warning(f"sync_one {uid} 失败: {e}")
return False
results = await asyncio.gather(*[sync_one(u.id) for u in users], return_exceptions=True)
for r in results:
if r is True: synced += 1
else: failed += 1
return ApiResponse(
message=f"同步完成:成功 {synced} 个,失败/跳过 {failed} 个",
data={"synced": synced, "failed": failed, "total": len(users)}
)
@router.post("/{user_id}/upload-avatar")
async def upload_avatar(
user_id: int,
file: UploadFile = File(...),
sync_to_platform: bool = Query(default=True),
db=Depends(get_db)
):
"""上传头像并可选同步到目标平台"""
from sqlalchemy import select, update
from app.models import VirtualUser as _VU
from app.services.news_service import news_service
ur = await db.execute(select(_VU).where(_VU.id == user_id))
user = ur.scalar_one_or_none()
if not user:
return ApiResponse(code=404, message="用户不存在")
# 读取文件内容
file_bytes = await file.read()
if len(file_bytes) > 5 * 1024 * 1024:
return ApiResponse(code=400, message="头像文件不能超过5MB")
avatar_url = None
if sync_to_platform and user.status == 2:
# 上传到目标平台
ok, result = await news_service.upload_avatar(db, user, file_bytes, file.filename)
if ok:
avatar_url = result
else:
return ApiResponse(code=500, message=f"头像上传到平台失败: {result}")
else:
# 未登录用户本地落盘,避免 base64 超过 avatar_url 字段长度
avatar_url = await _save_local_avatar(file_bytes, file.filename or "", file.content_type)
# 已登录用户必须同时写入会会当前资料和“TA 的主页”。
# 任一接口失败都不得返回“头像更新成功”。
if sync_to_platform and user.status == 2 and avatar_url:
ok, err = await news_service.update_user_profile(db, user, avatar=avatar_url)
if not ok:
return ApiResponse(code=502, message=f"头像已上传,但同步到会会失败: {err}")
else:
await db.execute(update(_VU).where(_VU.id == user_id).values(avatar_url=avatar_url))
await db.commit()
return ApiResponse(data={"avatar_url": avatar_url}, message="头像更新成功")
@router.post("/logout-all")
async def batch_logout_all(db=Depends(get_db)):
"""一键登出所有已登录用户"""
from sqlalchemy import select, update
from app.models import VirtualUser as _VU
from app.core.redis_client import delete_session
result_r = await db.execute(
select(_VU.id).where(_VU.status == 2, _VU.is_enabled == 1)
)
rows = result_r.all()
if not rows:
return ApiResponse(message="没有已登录的用户", data={"count": 0})
count = 0
for row in rows:
try:
await delete_session(row[0])
count += 1
except Exception:
pass
# 更新所有用户状态为未登录
await db.execute(
update(_VU).where(_VU.status == 2).values(status=0)
)
await db.commit()
return ApiResponse(message=f"已登出 {count} 个用户", data={"count": count})
-72
View File
@@ -1,72 +0,0 @@
"""
互动管理 API
"""
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.orm import Session
from typing import Optional
from app.models.base import get_db
from app.models.interaction import InteractionType
from app.schemas.interaction import (
InteractionRecordResponse,
InteractionRecordListResponse,
InteractionExecuteRequest
)
from app.services.interaction_service import InteractionService, get_interaction_service
router = APIRouter()
@router.get("", response_model=InteractionRecordListResponse)
def get_interaction_records(
page: int = Query(1, ge=1, description="页码"),
page_size: int = Query(20, ge=1, le=100, description="每页数量"),
virtual_user_id: Optional[int] = Query(None, description="虚拟用户 ID"),
interaction_type: Optional[InteractionType] = Query(None, description="互动类型"),
db: Session = Depends(get_db),
service: InteractionService = Depends(get_interaction_service)
):
"""获取互动记录列表"""
# TODO: 实现筛选和分页查询
return {"total": 0, "items": []}
@router.get("/{record_id}", response_model=InteractionRecordResponse)
def get_interaction_record(
record_id: int,
db: Session = Depends(get_db),
service: InteractionService = Depends(get_interaction_service)
):
"""获取互动记录详情"""
# TODO: 实现详情查询
raise HTTPException(status_code=404, detail="Record not found")
@router.post("/execute", response_model=InteractionRecordResponse)
async def execute_interaction(
request: InteractionExecuteRequest,
db: Session = Depends(get_db),
service: InteractionService = Depends(get_interaction_service)
):
"""执行互动"""
record = await service.execute_interaction(
virtual_user_id=request.virtual_user_id,
interaction_type=request.interaction_type,
news_id=request.news_id
)
if not record:
raise HTTPException(status_code=400, detail="Failed to execute interaction")
return record
@router.post("/retry/{record_id}")
async def retry_interaction(
record_id: int,
db: Session = Depends(get_db),
service: InteractionService = Depends(get_interaction_service)
):
"""重试失败的互动"""
# TODO: 实现重试逻辑
raise HTTPException(status_code=404, detail="Record not found")
-19
View File
@@ -1,19 +0,0 @@
"""
API 路由
"""
from fastapi import APIRouter
from .virtual_user import router as virtual_user_router
from .interaction import router as interaction_router
from .ai_model import router as ai_model_router
from .system_config import router as system_config_router
from .dashboard import router as dashboard_router
api_router = APIRouter()
# 注册各模块路由
api_router.include_router(virtual_user_router, prefix="/virtual-users", tags=["虚拟用户管理"])
api_router.include_router(interaction_router, prefix="/interactions", tags=["互动管理"])
api_router.include_router(ai_model_router, prefix="/ai-models", tags=["AI 模型配置"])
api_router.include_router(system_config_router, prefix="/system", tags=["系统设置"])
api_router.include_router(dashboard_router, prefix="/dashboard", tags=["控制台"])
-126
View File
@@ -1,126 +0,0 @@
"""
系统配置 API
"""
from fastapi import APIRouter, Depends, HTTPException
from sqlalchemy.orm import Session
from typing import List
from app.models.base import get_db
from app.models.system_config import SystemConfig
from app.schemas.system_config import (
SystemConfigResponse,
SystemConfigUpdate,
ScheduleConfig,
LimitConfig,
ProbabilityConfig
)
from app.core.config import settings
router = APIRouter()
@router.get("", response_model=List[SystemConfigResponse])
def get_system_configs(
db: Session = Depends(get_db)
):
"""获取所有系统配置"""
configs = db.query(SystemConfig).all()
return configs
@router.get("/schedule", response_model=ScheduleConfig)
def get_schedule_config(
db: Session = Depends(get_db)
):
"""获取调度配置"""
from app.services.scheduler_service import scheduler_service
return ScheduleConfig(
task_start_hour=settings.TASK_START_HOUR,
task_end_hour=settings.TASK_END_HOUR,
task_interval_min=settings.TASK_INTERVAL_MIN,
task_interval_max=settings.TASK_INTERVAL_MAX,
is_task_running=scheduler_service.is_running
)
@router.get("/limits", response_model=LimitConfig)
def get_limit_config(
db: Session = Depends(get_db)
):
"""获取限额配置"""
return LimitConfig(
max_tokens_per_day=settings.MAX_TOKENS_PER_DAY,
max_comments_per_user_per_day=settings.MAX_COMMENTS_PER_USER_PER_DAY,
max_replies_per_user_per_day=settings.MAX_REPLIES_PER_USER_PER_DAY
)
@router.get("/probabilities", response_model=ProbabilityConfig)
def get_probability_config(
db: Session = Depends(get_db)
):
"""获取概率配置"""
return ProbabilityConfig(
like_probability=settings.LIKE_PROBABILITY,
favorite_probability=settings.FAVORITE_PROBABILITY,
share_probability=settings.SHARE_PROBABILITY
)
@router.put("/schedule")
def update_schedule_config(
config: ScheduleConfig,
db: Session = Depends(get_db)
):
"""更新调度配置"""
# TODO: 更新系统配置表并重新加载
return {"message": "Schedule config updated"}
@router.put("/limits")
def update_limit_config(
config: LimitConfig,
db: Session = Depends(get_db)
):
"""更新限额配置"""
# TODO: 更新系统配置表
return {"message": "Limit config updated"}
@router.post("/scheduler/start")
def start_scheduler(
db: Session = Depends(get_db)
):
"""启动定时任务"""
from app.services.scheduler_service import scheduler_service
scheduler_service.start()
scheduler_service.add_interaction_task()
return {"message": "Scheduler started", "running": scheduler_service.is_running}
@router.post("/scheduler/stop")
def stop_scheduler(
db: Session = Depends(get_db)
):
"""停止定时任务"""
from app.services.scheduler_service import scheduler_service
scheduler_service.stop()
return {"message": "Scheduler stopped", "running": scheduler_service.is_running}
@router.get("/scheduler/status")
def get_scheduler_status(
db: Session = Depends(get_db)
):
"""获取定时任务状态"""
from app.services.scheduler_service import scheduler_service
return {
"is_running": scheduler_service.is_running,
"jobs": [job.id for job in scheduler_service.scheduler.get_jobs()]
}
-162
View File
@@ -1,162 +0,0 @@
"""
虚拟用户管理 API
"""
from fastapi import APIRouter, Depends, HTTPException, Query, UploadFile, File
from sqlalchemy.orm import Session
from typing import List, Optional
import pandas as pd
import io
from app.models.base import get_db
from app.schemas.virtual_user import (
VirtualUserCreate,
VirtualUserUpdate,
VirtualUserResponse,
VirtualUserListResponse,
VirtualUserGenerateRequest,
ActivityLevel,
UserStatus
)
from app.services.virtual_user_service import VirtualUserService, get_virtual_user_service
router = APIRouter()
@router.get("", response_model=VirtualUserListResponse)
def get_virtual_users(
page: int = Query(1, ge=1, description="页码"),
page_size: int = Query(20, ge=1, le=100, description="每页数量"),
status: Optional[UserStatus] = Query(None, description="状态筛选"),
search: Optional[str] = Query(None, description="搜索关键词"),
db: Session = Depends(get_db),
service: VirtualUserService = Depends(get_virtual_user_service)
):
"""获取虚拟用户列表"""
result = service.get_users(page=page, page_size=page_size, status=status, search=search)
return result
@router.get("/{user_id}", response_model=VirtualUserResponse)
def get_virtual_user(
user_id: int,
db: Session = Depends(get_db),
service: VirtualUserService = Depends(get_virtual_user_service)
):
"""获取虚拟用户详情"""
user = service.get_user_by_id(user_id)
if not user:
raise HTTPException(status_code=404, detail="User not found")
return user
@router.post("", response_model=VirtualUserResponse)
def create_virtual_user(
user_data: VirtualUserCreate,
db: Session = Depends(get_db),
service: VirtualUserService = Depends(get_virtual_user_service)
):
"""创建虚拟用户"""
user = service.create_user(
username=user_data.username,
password=user_data.password,
nickname=user_data.nickname,
writing_style=user_data.writing_style,
activity_level=user_data.activity_level,
avatar_url=user_data.avatar_url,
persona_description=user_data.persona_description
)
if not user:
raise HTTPException(status_code=400, detail="Failed to create user (username may exist)")
return user
@router.post("/generate", response_model=VirtualUserListResponse)
def generate_virtual_users(
request: VirtualUserGenerateRequest,
db: Session = Depends(get_db),
service: VirtualUserService = Depends(get_virtual_user_service)
):
"""批量生成虚拟用户"""
users = service.generate_users(
count=request.count,
writing_styles=request.writing_styles,
activity_levels=request.activity_levels,
generate_persona=request.generate_persona
)
return {"total": len(users), "items": users}
@router.put("/{user_id}", response_model=VirtualUserResponse)
def update_virtual_user(
user_id: int,
user_data: VirtualUserUpdate,
db: Session = Depends(get_db),
service: VirtualUserService = Depends(get_virtual_user_service)
):
"""更新虚拟用户"""
update_data = user_data.model_dump(exclude_unset=True)
user = service.update_user(user_id, **update_data)
if not user:
raise HTTPException(status_code=404, detail="User not found")
return user
@router.delete("/{user_id}")
def delete_virtual_user(
user_id: int,
db: Session = Depends(get_db),
service: VirtualUserService = Depends(get_virtual_user_service)
):
"""删除虚拟用户"""
success = service.delete_user(user_id)
if not success:
raise HTTPException(status_code=404, detail="User not found")
return {"message": "User deleted successfully"}
@router.post("/import", response_model=dict)
def import_virtual_users(
file: UploadFile = File(...),
generate_persona: bool = Query(True, description="是否生成 AI 人格描述"),
db: Session = Depends(get_db),
service: VirtualUserService = Depends(get_virtual_user_service)
):
"""从 Excel 导入虚拟用户"""
try:
# 读取 Excel 文件
contents = file.file.read()
df = pd.read_excel(io.BytesIO(contents))
# 转换为字典列表
users_data = df.to_dict('records')
# 导入用户
result = service.import_users_from_excel(
users_data=users_data,
generate_persona=generate_persona
)
return {
"message": "Import completed",
"success_count": result["success_count"],
"failed_count": result["failed_count"]
}
except Exception as e:
raise HTTPException(status_code=400, detail=f"Import failed: {str(e)}")
@router.get("/{user_id}/stats")
def get_user_stats(
user_id: int,
db: Session = Depends(get_db),
service: VirtualUserService = Depends(get_virtual_user_service)
):
"""获取用户统计信息"""
stats = service.get_user_stats(user_id)
return stats
Regular → Executable
+1 -6
View File
@@ -1,6 +1 @@
"""
核心模块初始化
"""
from .config import settings
__all__ = ["settings"]
# app.core package
Regular → Executable
+38 -61
View File
@@ -1,76 +1,53 @@
"""
系统配置模块
"""
from pydantic_settings import BaseSettings
from typing import Optional
"""系统配置"""
import os
from urllib.parse import quote_plus
from pydantic_settings import BaseSettings
class Settings(BaseSettings):
"""应用配置"""
# 数据库
DB_HOST: str = os.getenv("DB_HOST", "localhost")
DB_PORT: int = int(os.getenv("DB_PORT", "3306"))
DB_USER: str = os.getenv("DB_USER", "aivirtual")
DB_PASSWORD: str = os.getenv("DB_PASSWORD", "AiVirtual2024")
DB_NAME: str = os.getenv("DB_NAME", "ai_virtual_news")
# 应用基础配置
APP_NAME: str = "会会虚拟用户 AI 互动系统"
APP_VERSION: str = "1.0.0"
DEBUG: bool = True
API_PREFIX: str = "/api/v1"
# Redis
REDIS_HOST: str = os.getenv("REDIS_HOST", "localhost")
REDIS_PORT: int = int(os.getenv("REDIS_PORT", "6379"))
# 数据库配置
DATABASE_HOST: str = "mysql"
DATABASE_PORT: int = 3306
DATABASE_USER: str = "root"
DATABASE_PASSWORD: str = "root123456"
DATABASE_NAME: str = "huihui_ai_bot"
DATABASE_URL: Optional[str] = None
# 安全
SECRET_KEY: str = os.getenv("SECRET_KEY", "dev-secret-key-change-in-prod")
AES_KEY: str = os.getenv("AES_KEY", "your-aes-key-32-chars-change-now!")
AVATAR_MODEL_CONFIG_TOKEN: str = os.getenv("AVATAR_MODEL_CONFIG_TOKEN", "")
# 新闻平台
NEWS_PLATFORM_BASE_URL: str = os.getenv(
"NEWS_PLATFORM_BASE_URL", "http://192.168.1.200:63120"
)
# 数字分身 SQLite 数据库路径(直连数字分身应用的 SQLite)
AVATAR_DB_PATH: str = os.getenv("AVATAR_DB_PATH", "")
# 数字分身后端服务地址(用于拼接头像等文件 URL)
AVATAR_BACKEND_URL: str = os.getenv("AVATAR_BACKEND_URL", "")
# 日志目录
LOG_DIR: str = "/app/logs"
@property
def get_database_url(self) -> str:
if self.DATABASE_URL:
return self.DATABASE_URL
return f"mysql+pymysql://{self.DATABASE_USER}:{self.DATABASE_PASSWORD}@{self.DATABASE_HOST}:{self.DATABASE_PORT}/{self.DATABASE_NAME}?charset=utf8mb4"
def DATABASE_URL(self) -> str:
# 对密码做 URL 编码,防止 @ # ! 等特殊字符破坏连接字符串
pwd = quote_plus(self.DB_PASSWORD)
return f"mysql+aiomysql://{self.DB_USER}:{pwd}@{self.DB_HOST}:{self.DB_PORT}/{self.DB_NAME}?charset=utf8mb4"
# JWT 配置
JWT_SECRET_KEY: str = "your-secret-key-change-in-production"
JWT_ALGORITHM: str = "HS256"
JWT_EXPIRE_MINUTES: int = 60 * 24 * 7 # 7 天
# 会会接口配置
HUIHUI_API_BASE: str = "http://192.168.1.200:63120"
HUIHUI_DOC_URL: str = "http://192.168.1.200:63120/doc.html"
# AI 模型配置(默认)
DEFAULT_AI_MODEL: str = "openai"
OPENAI_API_KEY: Optional[str] = None
OPENAI_BASE_URL: str = "https://api.openai.com/v1"
OPENAI_MODEL: str = "gpt-3.5-turbo"
ZHIPU_API_KEY: Optional[str] = None
ZHIPU_MODEL: str = "glm-4"
# 系统限制配置
MAX_TOKENS_PER_DAY: int = 10000 # 每日 Token 上限
MAX_COMMENTS_PER_USER_PER_DAY: int = 20 # 单用户每日最大评论数
MAX_REPLIES_PER_USER_PER_DAY: int = 10 # 单用户每日最大回复数
# 定时任务配置
TASK_START_HOUR: int = 9 # 活动开始时间
TASK_END_HOUR: int = 22 # 活动结束时间
TASK_INTERVAL_MIN: int = 10 # 最小间隔(分钟)
TASK_INTERVAL_MAX: int = 30 # 最大间隔(分钟)
# 互动概率配置
LIKE_PROBABILITY: float = 0.8 # 点赞概率
FAVORITE_PROBABILITY: float = 0.5 # 收藏概率
SHARE_PROBABILITY: float = 0.3 # 转发概率
# 文件存储配置
UPLOAD_DIR: str = "/app/data/uploads"
LOG_DIR: str = "/app/data/logs"
@property
def SYNC_DATABASE_URL(self) -> str:
pwd = quote_plus(self.DB_PASSWORD)
return f"mysql+pymysql://{self.DB_USER}:{pwd}@{self.DB_HOST}:{self.DB_PORT}/{self.DB_NAME}?charset=utf8mb4"
class Config:
env_file = ".env"
case_sensitive = True
# 创建全局配置实例
settings = Settings()
+104
View File
@@ -0,0 +1,104 @@
"""数据库连接管理"""
import asyncio
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
from sqlalchemy import text
from sqlalchemy.orm import DeclarativeBase
from app.core.config import settings
from app.core.logger import logger
class Base(DeclarativeBase):
pass
engine = create_async_engine(
settings.DATABASE_URL,
echo=False,
pool_pre_ping=True,
pool_recycle=3600,
pool_size=10,
max_overflow=20,
)
AsyncSessionLocal = async_sessionmaker(
engine, class_=AsyncSession, expire_on_commit=False
)
async def get_db():
"""获取数据库会话"""
async with AsyncSessionLocal() as session:
try:
yield session
await session.commit()
except Exception:
await session.rollback()
raise
finally:
await session.close()
async def wait_for_db(max_retries: int = 30, interval: int = 2):
"""等待 MySQL 就绪,最多重试 max_retries 次"""
for attempt in range(1, max_retries + 1):
try:
async with engine.begin() as conn:
await conn.execute(__import__("sqlalchemy").text("SELECT 1"))
logger.info(f"✅ 数据库连接成功(第 {attempt} 次尝试)")
return
except Exception as e:
if attempt == max_retries:
logger.error(f"数据库连接失败,已重试 {max_retries} 次: {e}")
raise
logger.warning(f"数据库未就绪,{interval}s 后重试({attempt}/{max_retries}): {e}")
await asyncio.sleep(interval)
async def init_db():
"""初始化数据库 - 等待 MySQL 就绪并注册所有模型"""
try:
# 等待 MySQL 容器真正就绪
await wait_for_db(max_retries=30, interval=2)
# 导入所有模型类,确保 SQLAlchemy ORM 元数据注册
from app.models import (
VirtualUser, UserPersonality, InteractionRecord,
PendingReplyTask, TokenStat, AIModelConfig, SystemConfig, LoginLog
)
async with engine.begin() as conn:
await conn.execute(text("SELECT GET_LOCK('ai_model_config_migration', 30)"))
try:
columns = (
(
"usage_scope",
"ALTER TABLE ai_model_configs ADD COLUMN usage_scope "
"VARCHAR(16) NOT NULL DEFAULT 'general' AFTER provider",
),
(
"vision_model_version",
"ALTER TABLE ai_model_configs ADD COLUMN vision_model_version "
"VARCHAR(64) NULL AFTER model_version",
),
(
"ocr_model_version",
"ALTER TABLE ai_model_configs ADD COLUMN ocr_model_version "
"VARCHAR(64) NULL AFTER vision_model_version",
),
)
for column_name, ddl in columns:
result = await conn.execute(text(
"SELECT COUNT(*) FROM information_schema.COLUMNS "
"WHERE TABLE_SCHEMA = DATABASE() AND TABLE_NAME = 'ai_model_configs' "
"AND COLUMN_NAME = :column_name"
), {"column_name": column_name})
if result.scalar_one() == 0:
await conn.execute(text(ddl))
logger.info("AI模型配置表已增加 %s 字段", column_name)
finally:
await conn.execute(text("SELECT RELEASE_LOCK('ai_model_config_migration')"))
logger.info("✅ 数据库模型注册成功")
logger.info("✅ 数据库初始化完成")
except Exception as e:
logger.error(f"数据库初始化失败: {e}")
raise
+48
View File
@@ -0,0 +1,48 @@
"""日志配置"""
import sys
import os
from loguru import logger
LOG_DIR = os.getenv("LOG_DIR", "./logs")
os.makedirs(LOG_DIR, exist_ok=True)
# 移除默认处理器
logger.remove()
# 控制台输出
logger.add(
sys.stdout,
level="INFO",
format="<green>{time:YYYY-MM-DD HH:mm:ss}</green> | <level>{level: <8}</level> | <cyan>{name}</cyan>:<cyan>{function}</cyan>:<cyan>{line}</cyan> - <level>{message}</level>",
)
# 通用日志文件
logger.add(
f"{LOG_DIR}/app_{{time:YYYY-MM-DD}}.log",
rotation="00:00",
retention="30 days",
level="INFO",
encoding="utf-8",
format="{time:YYYY-MM-DD HH:mm:ss} | {level: <8} | {name}:{function}:{line} - {message}",
)
# 错误日志文件
logger.add(
f"{LOG_DIR}/error_{{time:YYYY-MM-DD}}.log",
rotation="00:00",
retention="30 days",
level="ERROR",
encoding="utf-8",
)
# AI调用日志
logger.add(
f"{LOG_DIR}/ai_{{time:YYYY-MM-DD}}.log",
rotation="00:00",
retention="30 days",
level="INFO",
encoding="utf-8",
filter=lambda record: "ai_call" in record["extra"],
)
__all__ = ["logger"]
+81
View File
@@ -0,0 +1,81 @@
"""Redis缓存客户端"""
import json
import redis.asyncio as aioredis
from app.core.config import settings
from app.core.logger import logger
_redis_client = None
async def get_redis() -> aioredis.Redis:
global _redis_client
if _redis_client is None:
_redis_client = aioredis.from_url(
f"redis://{settings.REDIS_HOST}:{settings.REDIS_PORT}",
encoding="utf-8",
decode_responses=True,
)
return _redis_client
# Session键前缀
SESSION_PREFIX = "session:"
LOCK_PREFIX = "lock:"
RATE_PREFIX = "rate:"
async def set_session(user_id: int, session_data: dict, expire: int = 86400):
"""存储用户会话"""
r = await get_redis()
key = f"{SESSION_PREFIX}{user_id}"
await r.setex(key, expire, json.dumps(session_data, ensure_ascii=False))
async def get_session(user_id: int) -> dict | None:
"""获取用户会话"""
r = await get_redis()
key = f"{SESSION_PREFIX}{user_id}"
data = await r.get(key)
if data:
return json.loads(data)
return None
async def delete_session(user_id: int):
"""删除用户会话"""
r = await get_redis()
key = f"{SESSION_PREFIX}{user_id}"
await r.delete(key)
async def acquire_lock(name: str, expire: int = 60) -> bool:
"""获取分布式锁"""
r = await get_redis()
key = f"{LOCK_PREFIX}{name}"
result = await r.set(key, "1", nx=True, ex=expire)
return result is True
async def release_lock(name: str):
"""释放分布式锁"""
r = await get_redis()
key = f"{LOCK_PREFIX}{name}"
await r.delete(key)
async def incr_rate(key: str, expire: int = 86400) -> int:
"""限流计数"""
r = await get_redis()
rate_key = f"{RATE_PREFIX}{key}"
count = await r.incr(rate_key)
if count == 1:
await r.expire(rate_key, expire)
return count
async def get_counter(key: str) -> int:
"""获取计数"""
r = await get_redis()
rate_key = f"{RATE_PREFIX}{key}"
val = await r.get(rate_key)
return int(val) if val else 0
Regular → Executable
+74 -59
View File
@@ -1,95 +1,110 @@
"""
FastAPI 应用主文件
AI虚拟用户新闻互动系统 - 后端主入口
"""
import logging
import asyncio
from contextlib import asynccontextmanager
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from loguru import logger as loguru_logger
from fastapi.responses import JSONResponse
from fastapi.staticfiles import StaticFiles
import re as _re
from pathlib import Path
# 修复 Pydantic v2 把 naive datetime 序列化为 +00:00 的问题
# 数据库存的是北京时间,应标记为 +08:00
from fastapi.responses import Response
from fastapi.encoders import jsonable_encoder
import json as _json
class _CNJSONResponse(JSONResponse):
def render(self, content) -> bytes:
text = _json.dumps(content, ensure_ascii=False, allow_nan=False)
# 把 +00:00 替换为 +08:00
text = _re.sub(r'(\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2})\+00:00', r'\1+08:00', text)
return text.encode('utf-8')
from app.core.config import settings
from app.models.base import init_db
from app.api.router import api_router
from app.services.scheduler_service import scheduler_service
from app.core.database import init_db
from app.core.logger import logger
from app.api import router
from app.services.scheduler import scheduler_service
# 配置日志
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
# 自定义 datetime 序列化:数据库存的是北京时间,输出时标记为 +08:00
from fastapi.encoders import jsonable_encoder
import json as _json
logger = logging.getLogger(__name__)
class ChinaDatetimeEncoder(_json.JSONEncoder):
def default(self, obj):
if isinstance(obj, datetime.datetime):
# 标记为 +08:00 时区
return obj.strftime("%Y-%m-%dT%H:%M:%S+08:00")
return super().default(obj)
@asynccontextmanager
async def lifespan(app: FastAPI):
"""应用生命周期管理"""
# 启动时执行
logger.info("Starting application...")
logger.info("🚀 AI虚拟用户新闻互动系统启动中...")
# 初始化数据库
init_db()
logger.info("Database initialized")
# 启动定时任务
scheduler_service.start()
scheduler_service.add_interaction_task()
scheduler_service.add_login_task(hour=8, minute=0)
scheduler_service.reset_daily_counters(hour=0, minute=1)
logger.info("Scheduler started")
await init_db()
# 启动调度器
await scheduler_service.start()
logger.info("✅ 系统启动完成")
yield
# 关闭时执行
logger.info("Shutting down application...")
scheduler_service.stop()
# 关闭调度器
await scheduler_service.stop()
logger.info("🛑 系统已关闭")
# 创建 FastAPI 应用
app = FastAPI(
title=settings.APP_NAME,
version=settings.APP_VERSION,
description="会会虚拟用户 AI 互动系统后端 API",
lifespan=lifespan
import datetime as _dt
import datetime as _dt
class _DatetimeJSONResponse(JSONResponse):
def render(self, content) -> bytes:
import json
def _default(obj):
if isinstance(obj, _dt.datetime):
return obj.strftime("%Y-%m-%dT%H:%M:%S+08:00")
raise TypeError(repr(obj))
return json.dumps(content, ensure_ascii=False, default=_default).encode('utf-8')
app = FastAPI(default_response_class=_CNJSONResponse,
title="AI虚拟用户新闻互动系统",
description="基于AI驱动的虚拟用户新闻互动自动化平台",
version="1.0.0",
lifespan=lifespan,
docs_url="/api/docs",
redoc_url="/api/redoc",
)
# 配置 CORS
# CORS配置
app.add_middleware(
CORSMiddleware,
allow_origins=["*"], # 生产环境应该配置具体的域名
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# 注册路由
app.include_router(api_router, prefix=settings.API_PREFIX)
app.include_router(router, prefix="/api")
@app.get("/")
async def root():
"""根路径"""
return {
"name": settings.APP_NAME,
"version": settings.APP_VERSION,
"status": "running"
}
uploads_dir = Path(__file__).resolve().parent / "uploads"
uploads_dir.mkdir(parents=True, exist_ok=True)
app.mount("/api/uploads", StaticFiles(directory=uploads_dir), name="uploads")
@app.get("/health")
async def health_check():
"""健康检查"""
return {
"status": "healthy",
"scheduler_running": scheduler_service.is_running
}
return {"status": "ok", "service": "ai-virtual-news-backend"}
if __name__ == "__main__":
import uvicorn
uvicorn.run(
"app.main:app",
host="0.0.0.0",
port=8000,
reload=settings.DEBUG
@app.exception_handler(Exception)
async def global_exception_handler(request, exc):
logger.error(f"全局异常: {exc}")
return JSONResponse(
status_code=500,
content={"code": 500, "message": f"服务器内部错误: {str(exc)}"},
)
Regular → Executable
+161 -24
View File
@@ -1,25 +1,162 @@
"""
数据库模型初始化
"""
from .base import Base, engine, get_db, SessionLocal
from .virtual_user import VirtualUser, VirtualUserPersona
from .interaction import InteractionRecord, InteractionType
from .token_usage import TokenUsage
from .system_config import SystemConfig
from .ai_model import AIModelConfig
from .news_cache import NewsCache
"""SQLAlchemy ORM 模型"""
from datetime import datetime
from sqlalchemy import (
BigInteger, Integer, SmallInteger, String, Text, DateTime,
Boolean, Float, Date, JSON, func
)
from sqlalchemy.orm import Mapped, mapped_column
from app.core.database import Base
__all__ = [
"Base",
"engine",
"get_db",
"SessionLocal",
"VirtualUser",
"VirtualUserPersona",
"InteractionRecord",
"InteractionType",
"TokenUsage",
"SystemConfig",
"AIModelConfig",
"NewsCache",
]
class VirtualUser(Base):
__tablename__ = "virtual_users"
id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
nickname: Mapped[str] = mapped_column(String(64), nullable=False)
account: Mapped[str] = mapped_column(String(128), nullable=False, unique=True)
password_enc: Mapped[str] = mapped_column(String(512), nullable=False)
avatar_url: Mapped[str | None] = mapped_column(String(512))
status: Mapped[int] = mapped_column(SmallInteger, default=0)
activity_level: Mapped[int] = mapped_column(SmallInteger, default=1)
daily_comment_limit: Mapped[int] = mapped_column(Integer, default=10)
daily_like_limit: Mapped[int] = mapped_column(Integer, default=30)
today_comment_count: Mapped[int] = mapped_column(Integer, default=0)
today_like_count: Mapped[int] = mapped_column(Integer, default=0)
total_interactions: Mapped[int] = mapped_column(Integer, default=0)
session_token: Mapped[str | None] = mapped_column(Text)
session_expires_at: Mapped[datetime | None] = mapped_column(DateTime)
last_login_at: Mapped[datetime | None] = mapped_column(DateTime)
last_interact_at: Mapped[datetime | None] = mapped_column(DateTime)
real_name: Mapped[str | None] = mapped_column(String(64)) # 真实姓名(从平台同步)
sex: Mapped[int] = mapped_column(SmallInteger, default=0) # 性别 0未知 1男 2女
platform_uid: Mapped[str | None] = mapped_column(String(64)) # 平台用户ID
remark: Mapped[str | None] = mapped_column(String(256))
is_enabled: Mapped[int] = mapped_column(SmallInteger, default=1)
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now(), onupdate=func.now())
class UserPersonality(Base):
__tablename__ = "user_personalities"
id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
user_id: Mapped[int] = mapped_column(BigInteger, nullable=False, unique=True)
character_type: Mapped[str | None] = mapped_column(String(32))
language_style: Mapped[str | None] = mapped_column(String(32))
interest_tags: Mapped[dict | None] = mapped_column(JSON)
interact_tendency: Mapped[str | None] = mapped_column(String(32))
word_count_min: Mapped[int] = mapped_column(Integer, default=20)
word_count_max: Mapped[int] = mapped_column(Integer, default=100)
personality_desc: Mapped[str | None] = mapped_column(Text)
comment_style_prompt: Mapped[str | None] = mapped_column(Text)
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now(), onupdate=func.now())
class InteractionRecord(Base):
__tablename__ = "interaction_records"
id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
user_id: Mapped[int] = mapped_column(BigInteger, nullable=False, index=True)
user_nickname: Mapped[str | None] = mapped_column(String(64))
user_account: Mapped[str | None] = mapped_column(String(128))
article_id: Mapped[str | None] = mapped_column(String(64))
article_title: Mapped[str | None] = mapped_column(String(256))
interact_type: Mapped[str] = mapped_column(String(16), nullable=False, index=True)
content: Mapped[str | None] = mapped_column(Text)
platform_record_id: Mapped[str | None] = mapped_column(String(64)) # 平台返回的记录ID(用于取消互动)
parent_comment_id: Mapped[str | None] = mapped_column(String(64))
session_id: Mapped[str | None] = mapped_column(String(128))
token_consumed: Mapped[int] = mapped_column(Integer, default=0)
status: Mapped[int] = mapped_column(SmallInteger, default=0)
error_msg: Mapped[str | None] = mapped_column(String(512))
retry_count: Mapped[int] = mapped_column(SmallInteger, default=0)
executed_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now(), index=True)
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
class PendingReplyTask(Base):
__tablename__ = "pending_reply_tasks"
id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
actor_user_id: Mapped[int] = mapped_column(BigInteger, nullable=False, index=True)
next_actor_user_id: Mapped[int | None] = mapped_column(BigInteger, index=True)
news_id: Mapped[str] = mapped_column(String(64), nullable=False, index=True)
news_title: Mapped[str | None] = mapped_column(String(256))
article_org_id: Mapped[str | None] = mapped_column(String(64))
parent_comment_id: Mapped[str] = mapped_column(String(64), nullable=False, index=True)
parent_comment: Mapped[dict | None] = mapped_column(JSON)
reply_to: Mapped[dict | None] = mapped_column(JSON)
root_content: Mapped[str | None] = mapped_column(Text)
context: Mapped[str | None] = mapped_column(Text)
next_probability: Mapped[float] = mapped_column(Float, default=0.3)
next_delay_min_seconds: Mapped[int] = mapped_column(Integer, default=30)
next_delay_max_seconds: Mapped[int] = mapped_column(Integer, default=7200)
status: Mapped[int] = mapped_column(SmallInteger, default=0, index=True) # 0待发送 1发送中 2已发送 3失败
attempts: Mapped[int] = mapped_column(SmallInteger, default=0)
scheduled_at: Mapped[datetime] = mapped_column(DateTime, nullable=False, index=True)
locked_at: Mapped[datetime | None] = mapped_column(DateTime)
sent_at: Mapped[datetime | None] = mapped_column(DateTime)
last_error: Mapped[str | None] = mapped_column(String(512))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now(), onupdate=func.now())
class TokenStat(Base):
__tablename__ = "token_stats"
id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
stat_date: Mapped[datetime] = mapped_column(Date, nullable=False, unique=True)
model_name: Mapped[str | None] = mapped_column(String(64))
total_tokens: Mapped[int] = mapped_column(Integer, default=0)
prompt_tokens: Mapped[int] = mapped_column(Integer, default=0)
completion_tokens: Mapped[int] = mapped_column(Integer, default=0)
call_count: Mapped[int] = mapped_column(Integer, default=0)
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now(), onupdate=func.now())
class AIModelConfig(Base):
__tablename__ = "ai_model_configs"
id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
model_name: Mapped[str] = mapped_column(String(64), nullable=False)
provider: Mapped[str] = mapped_column(String(32), nullable=False)
usage_scope: Mapped[str] = mapped_column(String(16), nullable=False, default="general")
api_base_url: Mapped[str | None] = mapped_column(String(256))
api_key_enc: Mapped[str | None] = mapped_column(String(512))
model_version: Mapped[str | None] = mapped_column(String(64))
vision_model_version: Mapped[str | None] = mapped_column(String(64))
ocr_model_version: Mapped[str | None] = mapped_column(String(64))
temperature: Mapped[float] = mapped_column(Float, default=0.7)
max_tokens: Mapped[int] = mapped_column(Integer, default=1000)
timeout_seconds: Mapped[int] = mapped_column(Integer, default=30)
is_default: Mapped[int] = mapped_column(SmallInteger, default=0)
is_enabled: Mapped[int] = mapped_column(SmallInteger, default=1)
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now(), onupdate=func.now())
class SystemConfig(Base):
__tablename__ = "system_configs"
id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
config_key: Mapped[str] = mapped_column(String(64), nullable=False, unique=True)
config_value: Mapped[str | None] = mapped_column(Text)
config_type: Mapped[str] = mapped_column(String(16), default="string")
description: Mapped[str | None] = mapped_column(String(256))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now(), onupdate=func.now())
class LoginLog(Base):
__tablename__ = "login_logs"
id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True)
user_id: Mapped[int] = mapped_column(BigInteger, nullable=False, index=True)
user_account: Mapped[str | None] = mapped_column(String(128))
action: Mapped[str] = mapped_column(String(16), nullable=False)
session_id: Mapped[str | None] = mapped_column(String(128))
ip_address: Mapped[str | None] = mapped_column(String(64))
error_msg: Mapped[str | None] = mapped_column(String(512))
created_at: Mapped[datetime] = mapped_column(DateTime, server_default=func.now(), index=True)
-43
View File
@@ -1,43 +0,0 @@
"""
AI 模型配置模型
"""
from sqlalchemy import Column, Integer, String, DateTime, Boolean, Text, Float
from sqlalchemy.sql import func
from .base import Base
class AIModelConfig(Base):
"""AI 模型配置表"""
__tablename__ = "ai_model_configs"
id = Column(Integer, primary_key=True, autoincrement=True, comment="配置 ID")
# 模型基本信息
model_name = Column(String(100), unique=True, nullable=False, index=True, comment="模型名称(如 gpt-3.5-turbo)")
provider = Column(String(50), nullable=False, comment="提供商(openai/zhipu/baidu/aliyun)")
display_name = Column(String(200), comment="显示名称")
# API 配置
api_url = Column(String(500), nullable=False, comment="API 地址")
api_key = Column(String(500), nullable=False, comment="API Key(加密存储)")
api_version = Column(String(50), comment="API 版本")
# 模型参数
temperature = Column(Float, default=0.7, comment="温度(0-1)")
max_tokens = Column(Integer, default=1000, comment="最大 Token 数")
top_p = Column(Float, default=1.0, comment="Top P 参数")
# 状态控制
is_default = Column(Boolean, default=False, comment="是否为默认模型")
is_active = Column(Boolean, default=True, comment="是否启用")
# 描述信息
description = Column(Text, comment="模型描述")
# 时间戳
created_at = Column(DateTime, server_default=func.now(), comment="创建时间")
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now(), comment="更新时间")
def __repr__(self):
return f"<AIModelConfig(id={self.id}, name='{self.model_name}', provider='{self.provider}')>"
+19
View File
@@ -0,0 +1,19 @@
# models package - re-export all models
from app.models import (
VirtualUser, UserPersonality, InteractionRecord,
TokenStat, AIModelConfig, SystemConfig, LoginLog
)
# Aliases for import compatibility
virtual_user = VirtualUser
personality = UserPersonality
interaction = InteractionRecord
token_stat = TokenStat
ai_model = AIModelConfig
system_config = SystemConfig
login_log = LoginLog
__all__ = [
"VirtualUser", "UserPersonality", "InteractionRecord",
"TokenStat", "AIModelConfig", "SystemConfig", "LoginLog",
]
-49
View File
@@ -1,49 +0,0 @@
"""
数据库基础配置
"""
from sqlalchemy import create_engine
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import sessionmaker
from contextlib import contextmanager
import logging
from app.core.config import settings
logger = logging.getLogger(__name__)
# 创建数据库引擎
engine = create_engine(
settings.get_database_url,
pool_pre_ping=True,
pool_size=20,
max_overflow=40,
echo=settings.DEBUG,
)
# 创建会话工厂
SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine)
# 创建基类
Base = declarative_base()
@contextmanager
def get_db():
"""获取数据库会话的上下文管理器"""
db = SessionLocal()
try:
yield db
db.commit()
except Exception as e:
db.rollback()
logger.error(f"Database error: {e}")
raise
finally:
db.close()
def init_db():
"""初始化数据库表"""
from . import virtual_user, interaction, token_usage, system_config, ai_model, news_cache
Base.metadata.create_all(bind=engine)
logger.info("Database tables created successfully")
-68
View File
@@ -1,68 +0,0 @@
"""
互动记录模型
"""
from sqlalchemy import Column, Integer, String, DateTime, Enum, Text, ForeignKey, Boolean, Float
from sqlalchemy.sql import func
from sqlalchemy.orm import relationship
import enum
from .base import Base
class InteractionType(str, enum.Enum):
"""互动类型枚举"""
COMMENT = "comment" # 评论
REPLY = "reply" # 回复
LIKE = "like" # 点赞
FAVORITE = "favorite" # 收藏
SHARE = "share" # 转发
class InteractionStatus(str, enum.Enum):
"""互动状态枚举"""
PENDING = "pending" # 待执行
SUCCESS = "success" # 成功
FAILED = "failed" # 失败
RETRYING = "retrying" # 重试中
class InteractionRecord(Base):
"""互动记录表"""
__tablename__ = "interaction_records"
id = Column(Integer, primary_key=True, autoincrement=True, comment="记录 ID")
# 关联信息
virtual_user_id = Column(Integer, ForeignKey("virtual_users.id"), nullable=False, index=True, comment="虚拟用户 ID")
virtual_user = relationship("VirtualUser", backref="interaction_records")
news_id = Column(String(100), nullable=False, index=True, comment="新闻 ID")
news_title = Column(String(500), comment="新闻标题")
# 互动内容
interaction_type = Column(Enum(InteractionType), nullable=False, comment="互动类型")
content = Column(Text, comment="互动内容(评论/回复的文本)")
target_comment_id = Column(String(100), comment="目标评论 ID(回复时使用)")
# 执行状态
status = Column(Enum(InteractionStatus), default=InteractionStatus.PENDING, comment="执行状态")
retry_count = Column(Integer, default=0, comment="重试次数")
error_message = Column(Text, comment="错误信息(失败时)")
# AI 相关信息
ai_model_used = Column(String(100), comment="使用的 AI 模型")
tokens_used = Column(Integer, default=0, comment="消耗的 Token 数")
prompt_content = Column(Text, comment="发送给 AI 的提示词")
ai_response = Column(Text, comment="AI 返回的内容")
# 接口响应
api_response = Column(Text, comment="会会接口返回的原始响应")
api_request_id = Column(String(200), comment="接口请求 ID")
# 时间戳
execution_time = Column(DateTime, server_default=func.now(), comment="执行时间")
created_at = Column(DateTime, server_default=func.now(), comment="创建时间")
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now(), comment="更新时间")
def __repr__(self):
return f"<InteractionRecord(id={self.id}, user_id={self.virtual_user_id}, type='{self.interaction_type}', status='{self.status}')>"
-45
View File
@@ -1,45 +0,0 @@
"""
新闻缓存模型
"""
from sqlalchemy import Column, Integer, String, DateTime, Text, Boolean, Date
from sqlalchemy.sql import func
from .base import Base
class NewsCache(Base):
"""新闻缓存表"""
__tablename__ = "news_cache"
id = Column(Integer, primary_key=True, autoincrement=True, comment="缓存 ID")
# 新闻基本信息
news_id = Column(String(100), unique=True, nullable=False, index=True, comment="新闻 ID(来自会会接口)")
title = Column(String(500), nullable=False, comment="新闻标题")
summary = Column(Text, comment="新闻摘要")
content = Column(Text, comment="新闻内容")
# 来源信息
source = Column(String(200), comment="来源")
author = Column(String(100), comment="作者")
publish_time = Column(DateTime, comment="发布时间")
# 分类标签
category = Column(String(100), comment="分类")
tags = Column(String(500), comment="标签(逗号分隔)")
# 互动统计
view_count = Column(Integer, default=0, comment="阅读数")
comment_count = Column(Integer, default=0, comment="评论数")
like_count = Column(Integer, default=0, comment="点赞数")
# 缓存状态
is_cached = Column(Boolean, default=True, comment="是否已缓存")
cache_date = Column(Date, server_default=func.now(), index=True, comment="缓存日期")
# 时间戳
created_at = Column(DateTime, server_default=func.now(), comment="创建时间")
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now(), comment="更新时间")
def __repr__(self):
return f"<NewsCache(id={self.id}, news_id='{self.news_id}', title='{self.title[:30]}...')>"
-32
View File
@@ -1,32 +0,0 @@
"""
系统配置模型
"""
from sqlalchemy import Column, Integer, String, DateTime, Boolean, Text, JSON
from sqlalchemy.sql import func
from .base import Base
class SystemConfig(Base):
"""系统配置表"""
__tablename__ = "system_configs"
id = Column(Integer, primary_key=True, autoincrement=True, comment="配置 ID")
# 配置键值
config_key = Column(String(100), unique=True, nullable=False, index=True, comment="配置键")
config_value = Column(JSON, nullable=False, comment="配置值(JSON 格式)")
config_type = Column(String(50), comment="配置类型(schedule/limit/probability等)")
# 描述信息
description = Column(Text, comment="配置描述")
# 状态
is_active = Column(Boolean, default=True, comment="是否启用")
# 时间戳
created_at = Column(DateTime, server_default=func.now(), comment="创建时间")
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now(), comment="更新时间")
def __repr__(self):
return f"<SystemConfig(id={self.id}, key='{self.config_key}')>"
-36
View File
@@ -1,36 +0,0 @@
"""
Token 使用记录模型
"""
from sqlalchemy import Column, Integer, String, DateTime, ForeignKey, Float, Date
from sqlalchemy.sql import func
from .base import Base
class TokenUsage(Base):
"""Token 使用记录表"""
__tablename__ = "token_usages"
id = Column(Integer, primary_key=True, autoincrement=True, comment="记录 ID")
# 关联信息
virtual_user_id = Column(Integer, ForeignKey("virtual_users.id"), nullable=True, index=True, comment="虚拟用户 ID(可为空,系统级消耗)")
interaction_id = Column(Integer, ForeignKey("interaction_records.id"), nullable=True, comment="互动记录 ID")
# Token 信息
tokens_used = Column(Integer, nullable=False, comment="使用的 Token 数量")
tokens_prompt = Column(Integer, default=0, comment="提示词 Token 数")
tokens_completion = Column(Integer, default=0, comment="完成响应 Token 数")
# AI 模型信息
ai_model = Column(String(100), nullable=False, comment="使用的 AI 模型")
action_type = Column(String(50), comment="操作类型(generate_comment/generate_reply 等)")
# 日期分区(便于统计)
usage_date = Column(Date, server_default=func.now(), index=True, comment="使用日期")
# 时间戳
created_at = Column(DateTime, server_default=func.now(), comment="创建时间")
def __repr__(self):
return f"<TokenUsage(id={self.id}, tokens={self.tokens_used}, model='{self.ai_model}')>"
-91
View File
@@ -1,91 +0,0 @@
"""
虚拟用户模型
"""
from sqlalchemy import Column, Integer, String, DateTime, Boolean, Enum, Text, JSON, Float
from sqlalchemy.sql import func
from enum import Enum as PyEnum
import enum
from .base import Base
class ActivityLevel(str, enum.Enum):
"""活跃度枚举"""
LOW = "low" # 低:每日 1-2 次
MEDIUM = "medium" # 中:每日 2-5 次
HIGH = "high" # 高:每日 5-10 次
class UserStatus(str, enum.Enum):
"""用户状态枚举"""
ACTIVE = "active" # 启用
DISABLED = "disabled" # 禁用
class VirtualUser(Base):
"""虚拟用户表"""
__tablename__ = "virtual_users"
id = Column(Integer, primary_key=True, autoincrement=True, comment="用户 ID")
# 基本信息
username = Column(String(100), unique=True, nullable=False, index=True, comment="用户名(账号)")
password = Column(String(200), nullable=False, comment="密码(加密存储)")
nickname = Column(String(100), nullable=False, comment="昵称")
avatar_url = Column(String(500), comment="头像 URL")
# 人格特征
writing_style = Column(String(50), comment="写作风格")
activity_level = Column(Enum(ActivityLevel), default=ActivityLevel.MEDIUM, comment="活跃度")
persona_description = Column(Text, comment="人格描述(AI 生成)")
# 状态控制
status = Column(Enum(UserStatus), default=UserStatus.ACTIVE, comment="状态")
is_logged_in = Column(Boolean, default=False, comment="是否已登录")
session_token = Column(String(500), comment="会话 Token(登录后)")
token_expire_time = Column(DateTime, comment="Token 过期时间")
# 互动统计
total_interactions = Column(Integer, default=0, comment="总互动次数")
today_comments = Column(Integer, default=0, comment="今日评论数")
today_replies = Column(Integer, default=0, comment="今日回复数")
last_interaction_time = Column(DateTime, comment="最后互动时间")
# 扩展信息
extra_info = Column(JSON, default=dict, comment="扩展信息")
# 时间戳
created_at = Column(DateTime, server_default=func.now(), comment="创建时间")
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now(), comment="更新时间")
def __repr__(self):
return f"<VirtualUser(id={self.id}, nickname='{self.nickname}', status='{self.status}')>"
class VirtualUserPersona(Base):
"""虚拟用户人格模板表"""
__tablename__ = "virtual_user_personas"
id = Column(Integer, primary_key=True, autoincrement=True, comment="ID")
# 人格特征
name = Column(String(100), unique=True, nullable=False, comment="人格名称")
description = Column(Text, comment="人格描述")
# 风格配置
writing_styles = Column(JSON, comment="写作风格列表")
personality_traits = Column(JSON, comment="性格特征列表")
speech_patterns = Column(JSON, comment="说话模式列表")
# AI 提示词
system_prompt = Column(Text, comment="系统提示词(用于 AI 生成)")
# 状态
is_active = Column(Boolean, default=True, comment="是否启用")
# 时间戳
created_at = Column(DateTime, server_default=func.now(), comment="创建时间")
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now(), comment="更新时间")
def __repr__(self):
return f"<VirtualUserPersona(id={self.id}, name='{self.name}')>"
Regular → Executable
+238 -52
View File
@@ -1,53 +1,239 @@
"""
Pydantic Schema 定义
"""
from .virtual_user import (
VirtualUserCreate,
VirtualUserUpdate,
VirtualUserResponse,
VirtualUserListResponse,
VirtualUserGenerateRequest,
VirtualUserImportRequest,
ActivityLevel,
UserStatus,
)
from .interaction import (
InteractionRecordResponse,
InteractionRecordListResponse,
InteractionType,
InteractionStatus,
)
from .token_usage import TokenUsageResponse, TokenUsageStats
from .system_config import SystemConfigResponse, SystemConfigUpdate
from .ai_model import AIModelConfigCreate, AIModelConfigUpdate, AIModelConfigResponse
from .dashboard import DashboardStats, DashboardTokenStats
"""Pydantic数据模型 - 请求/响应模式"""
from datetime import datetime
from typing import Optional, List, Any
from pydantic import BaseModel, Field
from datetime import timezone, timedelta
__all__ = [
# Virtual User
"VirtualUserCreate",
"VirtualUserUpdate",
"VirtualUserResponse",
"VirtualUserListResponse",
"VirtualUserGenerateRequest",
"VirtualUserImportRequest",
"ActivityLevel",
"UserStatus",
# Interaction
"InteractionRecordResponse",
"InteractionRecordListResponse",
"InteractionType",
"InteractionStatus",
# Token Usage
"TokenUsageResponse",
"TokenUsageStats",
# System Config
"SystemConfigResponse",
"SystemConfigUpdate",
# AI Model
"AIModelConfigCreate",
"AIModelConfigUpdate",
"AIModelConfigResponse",
# Dashboard
"DashboardStats",
"DashboardTokenStats",
]
_CST = timedelta(hours=8)
def _fmt_dt(dt):
if dt is None: return None
if hasattr(dt, "strftime"): return dt.strftime("%Y-%m-%dT%H:%M:%S+08:00")
return dt
# ===== 通用响应 =====
class ApiResponse(BaseModel):
code: int = 200
message: str = "success"
data: Any = None
class PageResult(BaseModel):
total: int
page: int
page_size: int
items: List[Any]
# ===== 虚拟用户 =====
class UserCreateRequest(BaseModel):
# 必填
account: str = Field(..., min_length=1, max_length=128, description="新闻平台账号(必填)")
password: str = Field(..., min_length=6, max_length=64, description="登录密码(必填)")
# 选填
nickname: Optional[str] = Field(None, max_length=64, description="昵称(选填,为空自动生成)")
avatar_url: Optional[str] = None
activity_level: int = Field(default=1, ge=0, le=2)
daily_comment_limit: int = Field(default=10, ge=1, le=100)
daily_like_limit: int = Field(default=30, ge=1, le=200)
remark: Optional[str] = None
class UserUpdateRequest(BaseModel):
nickname: Optional[str] = Field(None, min_length=1, max_length=64)
password: Optional[str] = Field(None, min_length=6, max_length=64)
avatar_url: Optional[str] = None
real_name: Optional[str] = None
sex: Optional[int] = None
description: Optional[str] = None
email: Optional[str] = None
activity_level: Optional[int] = Field(None, ge=0, le=2)
daily_comment_limit: Optional[int] = Field(None, ge=1, le=100)
daily_like_limit: Optional[int] = Field(None, ge=1, le=200)
remark: Optional[str] = None
is_enabled: Optional[int] = None
sync_to_platform: bool = False
class UserResponse(BaseModel):
id: int
nickname: str
account: str
avatar_url: Optional[str]
real_name: Optional[str] = None
sex: int = 0
platform_uid: Optional[str] = None
status: int
status_label: str
activity_level: int
activity_label: str
daily_comment_limit: int
daily_like_limit: int
today_comment_count: int
today_like_count: int
total_interactions: int
last_login_at: Optional[datetime]
last_interact_at: Optional[datetime]
remark: Optional[str]
is_enabled: int
created_at: datetime
personality: Optional[dict] = None
class Config:
from_attributes = True
class UserBatchRequest(BaseModel):
user_ids: List[int]
action: str # enable/disable/logout/delete
# ===== 人格 =====
class PersonalityUpdateRequest(BaseModel):
character_type: Optional[str] = None
language_style: Optional[str] = None
interest_tags: Optional[List[str]] = None
interact_tendency: Optional[str] = None
word_count_min: Optional[int] = Field(None, ge=10, le=500)
word_count_max: Optional[int] = Field(None, ge=10, le=1000)
personality_desc: Optional[str] = None
class PersonalityResponse(BaseModel):
id: int
user_id: int
character_type: Optional[str]
language_style: Optional[str]
interest_tags: Optional[List[str]]
interact_tendency: Optional[str]
word_count_min: int
word_count_max: int
personality_desc: Optional[str]
updated_at: datetime
class Config:
from_attributes = True
# ===== 互动记录 =====
class InteractionQueryParams(BaseModel):
page: int = Field(default=1, ge=1)
page_size: int = Field(default=20, ge=1, le=100)
user_id: Optional[int] = None
interact_type: Optional[str] = None
status: Optional[int] = None
start_date: Optional[str] = None
end_date: Optional[str] = None
keyword: Optional[str] = None
class InteractionResponse(BaseModel):
id: int
user_id: int
user_nickname: Optional[str]
user_account: Optional[str]
article_id: Optional[str]
article_title: Optional[str]
interact_type: str
interact_type_label: str
content: Optional[str]
token_consumed: int
status: int
status_label: str
error_msg: Optional[str]
retry_count: int
executed_at: datetime
class Config:
from_attributes = True
# ===== AI模型配置 =====
class AIModelCreateRequest(BaseModel):
model_name: str = Field(..., min_length=1, max_length=64)
provider: str = Field(..., pattern="^(openai|zhipu|wenxin|qianwen|local)$")
usage_scope: str = Field(default="general", pattern="^(general|digital_avatar)$")
api_base_url: Optional[str] = None
api_key: Optional[str] = None
model_version: Optional[str] = None
vision_model_version: Optional[str] = Field(None, max_length=64)
ocr_model_version: Optional[str] = Field(None, max_length=64)
temperature: float = Field(default=0.7, ge=0.0, le=2.0)
max_tokens: int = Field(default=1000, ge=1, le=32000)
timeout_seconds: int = Field(default=30, ge=5, le=300)
is_default: int = Field(default=0, ge=0, le=1)
class AIModelUpdateRequest(BaseModel):
model_name: Optional[str] = None
provider: Optional[str] = Field(None, pattern="^(openai|zhipu|wenxin|qianwen|local)$")
usage_scope: Optional[str] = Field(None, pattern="^(general|digital_avatar)$")
api_base_url: Optional[str] = None
api_key: Optional[str] = None
model_version: Optional[str] = None
vision_model_version: Optional[str] = Field(None, max_length=64)
ocr_model_version: Optional[str] = Field(None, max_length=64)
temperature: Optional[float] = Field(None, ge=0.0, le=2.0)
max_tokens: Optional[int] = Field(None, ge=1, le=32000)
timeout_seconds: Optional[int] = Field(None, ge=5, le=300)
is_default: Optional[int] = None
is_enabled: Optional[int] = None
class AIModelResponse(BaseModel):
id: int
model_name: str
provider: str
usage_scope: str
api_base_url: Optional[str]
has_api_key: bool
model_version: Optional[str]
vision_model_version: Optional[str]
ocr_model_version: Optional[str]
temperature: float
max_tokens: int
timeout_seconds: int
is_default: int
is_enabled: int
created_at: datetime
class Config:
from_attributes = True
class AIModelTestRequest(BaseModel):
model_id: int
test_prompt: str = "你好,请简单介绍一下自己。"
# ===== 系统配置 =====
class SystemConfigUpdateRequest(BaseModel):
configs: dict
# ===== 数据统计 =====
class DashboardResponse(BaseModel):
user_stats: dict
today_interactions: dict
monthly_stats: dict
token_stats: dict
system_status: dict
online_users: int
# ===== 调度配置 =====
class SchedulerConfigRequest(BaseModel):
interact_time_start: Optional[str] = None
interact_time_end: Optional[str] = None
interact_interval_min: Optional[int] = None
interact_interval_max: Optional[int] = None
max_concurrent_users: Optional[int] = None
daily_token_limit: Optional[int] = None
comment_probability: Optional[float] = None
reply_probability: Optional[float] = None
like_probability: Optional[float] = None
collect_probability: Optional[float] = None
forward_probability: Optional[float] = None
scheduler_enabled: Optional[bool] = None
-64
View File
@@ -1,64 +0,0 @@
"""
AI 模型配置相关 Schema
"""
from pydantic import BaseModel, Field
from typing import Optional, List
from datetime import datetime
class AIModelConfigBase(BaseModel):
"""AI 模型配置基础 Schema"""
model_name: str = Field(..., description="模型名称", max_length=100)
provider: str = Field(..., description="提供商", max_length=50)
display_name: Optional[str] = Field(None, description="显示名称", max_length=200)
api_url: str = Field(..., description="API 地址", max_length=500)
api_key: str = Field(..., description="API Key", max_length=500)
temperature: float = Field(0.7, description="温度", ge=0, le=1)
max_tokens: int = Field(1000, description="最大 Token 数", ge=1)
class AIModelConfigCreate(AIModelConfigBase):
"""创建 AI 模型配置请求"""
description: Optional[str] = Field(None, description="模型描述")
class AIModelConfigUpdate(BaseModel):
"""更新 AI 模型配置请求"""
display_name: Optional[str] = Field(None, description="显示名称", max_length=200)
api_url: Optional[str] = Field(None, description="API 地址", max_length=500)
api_key: Optional[str] = Field(None, description="API Key", max_length=500)
temperature: Optional[float] = Field(None, description="温度", ge=0, le=1)
max_tokens: Optional[int] = Field(None, description="最大 Token 数", ge=1)
is_default: Optional[bool] = Field(None, description="是否为默认模型")
is_active: Optional[bool] = Field(None, description="是否启用")
description: Optional[str] = Field(None, description="模型描述")
class AIModelConfigResponse(AIModelConfigBase):
"""AI 模型配置响应"""
id: int
api_version: Optional[str]
top_p: float
is_default: bool
is_active: bool
description: Optional[str]
created_at: datetime
updated_at: datetime
class Config:
from_attributes = True
class AIModelTestRequest(BaseModel):
"""AI 模型测试请求"""
model_id: int = Field(..., description="模型 ID")
test_prompt: str = Field(..., description="测试提示词", min_length=1, max_length=1000)
class AIModelTestResponse(BaseModel):
"""AI 模型测试响应"""
success: bool
content: Optional[str]
tokens_used: int
cost_time: float
error_message: Optional[str]
-55
View File
@@ -1,55 +0,0 @@
"""
控制台仪表盘相关 Schema
"""
from pydantic import BaseModel, Field
from typing import List, Optional
class CoreStats(BaseModel):
"""核心指标统计"""
total_users: int = Field(0, description="虚拟用户总数")
active_users: int = Field(0, description="已启用用户数")
disabled_users: int = Field(0, description="已禁用用户数")
today_comments: int = Field(0, description="今日评论数")
today_replies: int = Field(0, description="今日回复数")
today_likes: int = Field(0, description="今日点赞数")
today_favorites: int = Field(0, description="今日收藏数")
today_shares: int = Field(0, description="今日转发数")
yesterday_comments: int = Field(0, description="昨日评论数")
yesterday_replies: int = Field(0, description="昨日回复数")
month_tokens: int = Field(0, description="当月 Token 消耗")
today_tokens: int = Field(0, description="今日 Token 消耗")
remaining_tokens: int = Field(0, description="今日剩余 Token")
class DashboardTokenStats(BaseModel):
"""Token 统计"""
today_used: int = Field(0, description="今日已用")
today_limit: int = Field(0, description="今日限额")
today_remaining: int = Field(0, description="今日剩余")
usage_percentage: float = Field(0, description="使用百分比")
class DailyUsageItem(BaseModel):
"""每日使用项"""
date: str
tokens: int
comments: int
replies: int
class MonthlyUsageItem(BaseModel):
"""每月使用项"""
month: str
tokens: int
class DashboardStats(BaseModel):
"""控制台统计数据"""
core_stats: CoreStats
daily_token_usages: List[DailyUsageItem] = Field(default_factory=list)
monthly_token_usages: List[MonthlyUsageItem] = Field(default_factory=list)
recent_interactions: List[dict] = Field(default_factory=list)
-63
View File
@@ -1,63 +0,0 @@
"""
互动记录相关 Schema
"""
from pydantic import BaseModel, Field
from typing import Optional, List
from datetime import datetime
from enum import Enum
class InteractionType(str, Enum):
"""互动类型枚举"""
COMMENT = "comment"
REPLY = "reply"
LIKE = "like"
FAVORITE = "favorite"
SHARE = "share"
class InteractionStatus(str, Enum):
"""互动状态枚举"""
PENDING = "pending"
SUCCESS = "success"
FAILED = "failed"
RETRYING = "retrying"
class InteractionRecordBase(BaseModel):
"""互动记录基础 Schema"""
virtual_user_id: int = Field(..., description="虚拟用户 ID")
news_id: str = Field(..., description="新闻 ID")
interaction_type: InteractionType = Field(..., description="互动类型")
content: Optional[str] = Field(None, description="互动内容")
target_comment_id: Optional[str] = Field(None, description="目标评论 ID")
class InteractionRecordResponse(InteractionRecordBase):
"""互动记录响应"""
id: int
news_title: Optional[str]
status: InteractionStatus
retry_count: int
error_message: Optional[str]
ai_model_used: Optional[str]
tokens_used: int
execution_time: datetime
created_at: datetime
class Config:
from_attributes = True
class InteractionRecordListResponse(BaseModel):
"""互动记录列表响应"""
total: int
items: List[InteractionRecordResponse]
class InteractionExecuteRequest(BaseModel):
"""执行互动请求"""
virtual_user_id: int = Field(..., description="虚拟用户 ID")
news_id: Optional[str] = Field(None, description="新闻 ID(不传则随机选择)")
interaction_type: Optional[InteractionType] = Field(None, description="互动类型(不传则随机)")
force_execute: bool = Field(False, description="是否强制执行(忽略限额)")
-60
View File
@@ -1,60 +0,0 @@
"""
系统配置相关 Schema
"""
from pydantic import BaseModel, Field
from typing import Optional, Dict, Any
from datetime import datetime
class SystemConfigBase(BaseModel):
"""系统配置基础 Schema"""
config_key: str = Field(..., description="配置键", max_length=100)
config_value: Dict[str, Any] = Field(..., description="配置值")
config_type: Optional[str] = Field(None, description="配置类型", max_length=50)
description: Optional[str] = Field(None, description="配置描述")
class SystemConfigCreate(SystemConfigBase):
"""创建系统配置请求"""
pass
class SystemConfigUpdate(BaseModel):
"""更新系统配置请求"""
config_value: Optional[Dict[str, Any]] = Field(None, description="配置值")
description: Optional[str] = Field(None, description="配置描述")
is_active: Optional[bool] = Field(None, description="是否启用")
class SystemConfigResponse(SystemConfigBase):
"""系统配置响应"""
id: int
is_active: bool
created_at: datetime
updated_at: datetime
class Config:
from_attributes = True
class ScheduleConfig(BaseModel):
"""调度配置"""
task_start_hour: int = Field(9, description="活动开始时间", ge=0, le=23)
task_end_hour: int = Field(22, description="活动结束时间", ge=0, le=23)
task_interval_min: int = Field(10, description="最小间隔(分钟)", ge=1)
task_interval_max: int = Field(30, description="最大间隔(分钟)", ge=1)
is_task_running: bool = Field(False, description="任务是否运行中")
class LimitConfig(BaseModel):
"""限额配置"""
max_tokens_per_day: int = Field(10000, description="每日 Token 上限", ge=0)
max_comments_per_user_per_day: int = Field(20, description="单用户每日最大评论数", ge=0)
max_replies_per_user_per_day: int = Field(10, description="单用户每日最大回复数", ge=0)
class ProbabilityConfig(BaseModel):
"""概率配置"""
like_probability: float = Field(0.8, description="点赞概率", ge=0, le=1)
favorite_probability: float = Field(0.5, description="收藏概率", ge=0, le=1)
share_probability: float = Field(0.3, description="转发概率", ge=0, le=1)
-54
View File
@@ -1,54 +0,0 @@
"""
Token 使用相关 Schema
"""
from pydantic import BaseModel, Field
from typing import Optional, List
from datetime import date, datetime
class TokenUsageBase(BaseModel):
"""Token 使用基础 Schema"""
tokens_used: int = Field(..., description="使用的 Token 数量")
ai_model: str = Field(..., description="使用的 AI 模型")
action_type: Optional[str] = Field(None, description="操作类型")
class TokenUsageResponse(TokenUsageBase):
"""Token 使用响应"""
id: int
virtual_user_id: Optional[int]
interaction_id: Optional[int]
tokens_prompt: int
tokens_completion: int
usage_date: date
created_at: datetime
class Config:
from_attributes = True
class TokenUsageStats(BaseModel):
"""Token 使用统计"""
today_tokens: int = Field(0, description="今日 Token 数")
yesterday_tokens: int = Field(0, description="昨日 Token 数")
month_tokens: int = Field(0, description="当月 Token 数")
remaining_tokens: int = Field(0, description="剩余 Token 数")
total_limit: int = Field(0, description="总限额")
class DailyTokenUsage(BaseModel):
"""每日 Token 使用"""
date: str
tokens: int
class MonthlyTokenUsage(BaseModel):
"""每月 Token 使用"""
month: str
tokens: int
class TokenUsageChartResponse(BaseModel):
"""Token 使用图表响应"""
daily_usages: List[DailyTokenUsage]
monthly_usages: List[MonthlyTokenUsage]
-86
View File
@@ -1,86 +0,0 @@
"""
虚拟用户相关 Schema
"""
from pydantic import BaseModel, Field
from typing import Optional, List, Dict, Any
from datetime import datetime
from enum import Enum
class ActivityLevel(str, Enum):
"""活跃度枚举"""
LOW = "low"
MEDIUM = "medium"
HIGH = "high"
class UserStatus(str, Enum):
"""用户状态枚举"""
ACTIVE = "active"
DISABLED = "disabled"
class VirtualUserBase(BaseModel):
"""虚拟用户基础 Schema"""
nickname: str = Field(..., description="昵称", min_length=1, max_length=100)
username: str = Field(..., description="用户名(账号)", min_length=1, max_length=100)
password: str = Field(..., description="密码", min_length=1)
avatar_url: Optional[str] = Field(None, description="头像 URL", max_length=500)
writing_style: Optional[str] = Field(None, description="写作风格", max_length=50)
activity_level: ActivityLevel = Field(default=ActivityLevel.MEDIUM, description="活跃度")
persona_description: Optional[str] = Field(None, description="人格描述")
class VirtualUserCreate(VirtualUserBase):
"""创建虚拟用户请求"""
pass
class VirtualUserUpdate(BaseModel):
"""更新虚拟用户请求"""
nickname: Optional[str] = Field(None, description="昵称", min_length=1, max_length=100)
password: Optional[str] = Field(None, description="密码", min_length=1)
avatar_url: Optional[str] = Field(None, description="头像 URL", max_length=500)
writing_style: Optional[str] = Field(None, description="写作风格", max_length=50)
activity_level: Optional[ActivityLevel] = Field(None, description="活跃度")
persona_description: Optional[str] = Field(None, description="人格描述")
status: Optional[UserStatus] = Field(None, description="状态")
class VirtualUserResponse(VirtualUserBase):
"""虚拟用户响应"""
id: int
status: UserStatus
is_logged_in: bool
total_interactions: int
today_comments: int
today_replies: int
last_interaction_time: Optional[datetime]
created_at: datetime
updated_at: datetime
class Config:
from_attributes = True
class VirtualUserListResponse(BaseModel):
"""虚拟用户列表响应"""
total: int
items: List[VirtualUserResponse]
class VirtualUserGenerateRequest(BaseModel):
"""生成虚拟用户请求"""
count: int = Field(1, description="生成数量", ge=1, le=100)
writing_styles: Optional[List[str]] = Field(None, description="写作风格列表")
activity_levels: Optional[List[ActivityLevel]] = Field(
[ActivityLevel.LOW, ActivityLevel.MEDIUM, ActivityLevel.HIGH],
description="活跃度级别列表"
)
generate_persona: bool = Field(True, description="是否生成 AI 人格描述")
class VirtualUserImportRequest(BaseModel):
"""导入虚拟用户请求"""
users: List[Dict[str, Any]] = Field(..., description="用户数据列表")
generate_persona: bool = Field(True, description="是否为导入的用户生成 AI 人格描述")
Regular → Executable
-18
View File
@@ -1,18 +0,0 @@
"""
服务层模块初始化
"""
from .huihui_api import HuihuiAPIService
from .ai_service import AIService
from .virtual_user_service import VirtualUserService
from .interaction_service import InteractionService
from .token_service import TokenService
from .scheduler_service import SchedulerService
__all__ = [
"HuihuiAPIService",
"AIService",
"VirtualUserService",
"InteractionService",
"TokenService",
"SchedulerService",
]
+257 -272
View File
@@ -1,301 +1,286 @@
"""
AI 大模型对接服务
支持 OpenAI、智谱、百度文心、阿里通义等主流大模型
"""
import logging
from typing import Optional, Dict, Any, List
from datetime import datetime
"""AI服务 - 人格生成、内容创作"""
import json
import random
import re
from typing import Optional
import httpx
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, update
logger = logging.getLogger(__name__)
from app.models import AIModelConfig, TokenStat
from app.utils.crypto import decrypt
from app.core.logger import logger
from datetime import date
class AIService:
"""AI 服务类"""
"""AI大模型服务"""
def __init__(self):
self._client_cache: Dict[str, Any] = {}
# 人格候选池
CHARACTER_TYPES = ["开朗", "内敛", "毒舌", "温和", "理性", "感性", "幽默", "严谨"]
LANGUAGE_STYLES = ["严肃", "幽默", "文艺", "吐槽", "口语化", "学术", "简洁", "丰富"]
INTEREST_TAGS_POOL = [
"科技", "财经", "娱乐", "体育", "政治", "文化", "教育", "医疗",
"汽车", "房产", "旅游", "美食", "军事", "国际", "环保", "农业"
]
INTERACT_TENDENCIES = ["爱评论", "爱点赞", "爱收藏", "潜水", "爱转发", "爱回复"]
async def _get_default_model(self, db: AsyncSession) -> Optional[AIModelConfig]:
result = await db.execute(
select(AIModelConfig).where(
AIModelConfig.usage_scope == "general",
AIModelConfig.is_default == 1,
AIModelConfig.is_enabled == 1,
)
)
return result.scalar_one_or_none()
async def _call_api(
self, db: AsyncSession, prompt: str, system_prompt: str = None,
max_tokens: int = None
) -> tuple[str, int]:
"""调用AI接口,返回(内容, token数)"""
model = await self._get_default_model(db)
if not model:
# 无模型配置时返回随机预设
return "", 0
api_key = decrypt(model.api_key_enc) if model.api_key_enc else ""
base_url = model.api_base_url or "https://api.openai.com/v1"
headers = {"Content-Type": "application/json"}
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
messages = []
if system_prompt:
messages.append({"role": "system", "content": system_prompt})
messages.append({"role": "user", "content": prompt})
payload = {
"model": model.model_version or "gpt-3.5-turbo",
"messages": messages,
"temperature": model.temperature,
"max_tokens": max_tokens or model.max_tokens,
}
import asyncio as _asyncio
last_err = None
for attempt in range(3): # 最多重试3次
try:
async with httpx.AsyncClient(timeout=model.timeout_seconds) as client:
resp = await client.post(
f"{base_url}/chat/completions",
headers=headers,
json=payload,
)
# 429 限流:等待后重试
if resp.status_code == 429:
wait = 30 * (attempt + 1) # 30s, 60s, 90s
logger.warning(f"AI接口限流(429),{wait}s后重试({attempt+1}/3)")
await _asyncio.sleep(wait)
continue
resp.raise_for_status()
data = resp.json()
text = data["choices"][0]["message"]["content"].strip()
tokens = data.get("usage", {}).get("total_tokens", 0)
await self._record_token_usage(db, tokens, data.get("usage", {}), model.model_name)
logger.bind(ai_call=True).info(
f"AI调用成功 model={model.model_name} tokens={tokens}"
)
return text, tokens
except Exception as e:
last_err = e
if attempt < 2:
await _asyncio.sleep(5 * (attempt + 1))
logger.error(f"AI调用失败(已重试3次): {last_err}")
return "", 0
async def generate_personality(self, nickname: str, account: str) -> dict:
"""生成用户人格(含fallback随机生成)"""
# 如无AI配置,使用随机生成
from app.core.database import AsyncSessionLocal
try:
async with AsyncSessionLocal() as db:
model = await self._get_default_model(db)
if not model:
return self._random_personality()
prompt = f"""请为以下虚拟新闻读者生成一个独特的人格档案,要求真实自然、贴合中国用户特征。
用户昵称:{nickname}
请严格以JSON格式返回,不要有其他内容:
{{
"character_type": "从[开朗/内敛/毒舌/温和/理性/感性/幽默/严谨]中选一个",
"language_style": "从[严肃/幽默/文艺/吐槽/口语化/学术/简洁/丰富]中选一个",
"interest_tags": ["兴趣1", "兴趣2", "兴趣3"],
"interact_tendency": "从[爱评论/爱点赞/爱收藏/潜水/爱转发/爱回复]中选一个",
"word_count_min": 最少字数(10-50整数),
"word_count_max": 最多字数(50-200整数),
"personality_desc": "一句话描述此人的性格特点(30字以内)"
}}"""
content, _ = await self._call_api(db, prompt, max_tokens=300)
if content:
try:
# 提取JSON
json_match = re.search(r'\{.*\}', content, re.DOTALL)
if json_match:
return json.loads(json_match.group())
except Exception:
pass
return self._random_personality()
except Exception as e:
logger.error(f"人格生成失败: {e}")
return self._random_personality()
def _random_personality(self) -> dict:
"""随机生成人格(无AI时的备用方案)"""
interests = random.sample(self.INTEREST_TAGS_POOL, random.randint(2, 4))
char = random.choice(self.CHARACTER_TYPES)
style = random.choice(self.LANGUAGE_STYLES)
tendency = random.choice(self.INTERACT_TENDENCIES)
w_min = random.randint(15, 40)
w_max = random.randint(60, 150)
return {
"character_type": char,
"language_style": style,
"interest_tags": interests,
"interact_tendency": tendency,
"word_count_min": w_min,
"word_count_max": w_max,
"personality_desc": f"一个{char}性格、{tendency}的新闻读者",
}
async def generate_comment(
self,
news_content: str,
writing_style: str,
persona_description: Optional[str] = None,
model_config: Optional[Dict[str, Any]] = None
) -> Optional[Dict[str, Any]]:
"""
AI 生成评论
:param news_content: 新闻内容
:param writing_style: 写作风格
:param persona_description: 人格描述
:param model_config: 模型配置
:return: 生成结果(包含 content, tokens_used 等)
"""
prompt = self._build_comment_prompt(
news_content,
writing_style,
persona_description
)
self, db: AsyncSession, article_title: str, article_content: str,
personality_prompt: str, word_min: int = 20, word_max: int = 80
) -> tuple[str, int]:
"""生成文章评论"""
system_prompt = f"""你是一名真实的社区用户,正在阅读新闻后发表评论。{personality_prompt}
return await self._call_ai_api(prompt, model_config)
重要规则:
- 评论必须积极正面、文明友善,绝对不包含任何政治敏感、色情、暴力、侮辱、歧视内容
- 不要提及具体政治人物、党派、政策批评、社会矛盾等敏感话题
- 内容围绕文章本身展开,表达个人感受、分享观点、提出建设性问题
- 语言朴实自然,像普通网友留言,不夸张不煽情"""
prompt = f"""请根据以下新闻文章写一条评论。
文章标题:{article_title}
文章摘要:{article_content[:200] if article_content else '(无摘要)'}
要求:
1. 评论字数 {word_min}~{word_max} 字
2. 内容积极正面,贴近文章主题
3. 语气自然真实,符合普通读者口吻
4. 必须是完整的句子,不能被截断,以句号/感叹号/问号结尾
5. 只输出评论正文,不要加任何前缀或解释
评论:"""
return await self._call_api(db, prompt, system_prompt, max_tokens=500)
async def generate_reply(
self,
original_comment: str,
news_content: str,
writing_style: str,
persona_description: Optional[str] = None,
model_config: Optional[Dict[str, Any]] = None
) -> Optional[Dict[str, Any]]:
"""
AI 生成回复
:param original_comment: 原评论
:param news_content: 新闻内容
:param writing_style: 写作风格
:param persona_description: 人格描述
:param model_config: 模型配置
:return: 生成结果
"""
prompt = self._build_reply_prompt(
original_comment,
news_content,
writing_style,
persona_description
)
self, db: AsyncSession, article_title: str, parent_comment: str,
personality_prompt: str, word_min: int = 15, word_max: int = 60
) -> tuple[str, int]:
"""生成回复"""
system_prompt = f"""你是一名真实的社区用户。{personality_prompt}
return await self._call_ai_api(prompt, model_config)
重要规则:回复必须积极正面、文明友善,不含任何敏感违规内容。"""
prompt = f"""文章:{article_title}
原评论:{parent_comment}
def _build_comment_prompt(
self,
news_content: str,
writing_style: str,
persona_description: Optional[str] = None
) -> str:
"""构建评论提示词"""
base_prompt = f"""你是一位虚拟用户,请根据以下要求写一条简短的评论:
请对上面的评论写一条友善自然的回复,{word_min}~{word_max}字,直接输出回复内容。"""
return await self._call_api(db, prompt, system_prompt, max_tokens=150)
写作风格:{writing_style}
"""
async def generate_thread_reply(
self, db: AsyncSession, article_title: str, root_comment: str,
reply_context: str, personality_prompt: str,
word_min: int = 15, word_max: int = 60
) -> tuple[str, int]:
"""结合文章、原评论和上下文生成评论回复链中的下一句。"""
system_prompt = f"""你是一名真实的社区用户。{personality_prompt}
if persona_description:
base_prompt += f"\n人格特征:{persona_description}\n"
重要规则:
- 回复必须积极正面、文明友善,不含任何敏感违规内容
- 要自然接住对方的话,不要机械复述
- 不要透露自己是AI或虚拟用户"""
prompt = f"""文章标题:{article_title}
原评论:{root_comment}
当前对话上下文:{reply_context}
base_prompt += f"""
新闻内容:
{news_content[:1000]} # 限制长度
请结合文章、原评论和当前对话,写一条自然的后续回复。
要求:
1. 字数 {word_min}~{word_max} 字
2. 语气像真实用户交流,可以认同、补充或追问
3. 必须围绕文章和评论内容,不要跑题
4. 只输出回复正文,不要加任何前缀或解释
请写一条 50-100 字的评论,要符合你的写作风格和人格特征。直接输出评论内容,不要有其他说明。"""
回复:"""
return await self._call_api(db, prompt, system_prompt, max_tokens=180)
return base_prompt
async def test_model(self, db: AsyncSession, model_id: int, test_prompt: str) -> dict:
"""测试模型可用性"""
result = await db.execute(select(AIModelConfig).where(AIModelConfig.id == model_id))
model = result.scalar_one_or_none()
if not model:
return {"success": False, "error": "模型不存在"}
def _build_reply_prompt(
self,
original_comment: str,
news_content: str,
writing_style: str,
persona_description: Optional[str] = None
) -> str:
"""构建回复提示词"""
base_prompt = f"""你是一位虚拟用户,请根据以下要求回复另一条评论:
api_key = decrypt(model.api_key_enc) if model.api_key_enc else ""
base_url = model.api_base_url or "https://api.openai.com/v1"
headers = {"Content-Type": "application/json"}
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
写作风格:{writing_style}
"""
if persona_description:
base_prompt += f"\n人格特征:{persona_description}\n"
base_prompt += f"""
新闻内容:
{news_content[:500]}
原评论:
{original_comment}
请写一条 30-80 字的回复,要符合你的写作风格和人格特征。直接输出回复内容,不要有其他说明。"""
return base_prompt
async def _call_ai_api(
self,
prompt: str,
model_config: Optional[Dict[str, Any]] = None
) -> Optional[Dict[str, Any]]:
"""
调用 AI API(根据 model_config 中的 provider 选择对应模型)
:param prompt: 提示词
:param model_config: 模型配置
:return: 生成结果
"""
if not model_config:
# 使用默认配置(需要从数据库加载)
from app.models.ai_model import AIModelConfig
from app.models.base import get_db
with get_db() as db:
default_model = db.query(AIModelConfig).filter(
AIModelConfig.is_default == True,
AIModelConfig.is_active == True
).first()
if not default_model:
logger.error("No default AI model configured")
return None
model_config = {
"provider": default_model.provider,
"model_name": default_model.model_name,
"api_key": default_model.api_key,
"api_url": default_model.api_url,
"temperature": default_model.temperature,
"max_tokens": default_model.max_tokens,
payload = {
"model": model.model_version or "gpt-3.5-turbo",
"messages": [{"role": "user", "content": test_prompt}],
"max_tokens": 200,
}
provider = model_config.get("provider", "").lower()
try:
if provider == "openai":
return await self._call_openai(prompt, model_config)
elif provider == "zhipu":
return await self._call_zhipu(prompt, model_config)
elif provider in ["baidu", "wenxin"]:
return await self._call_baidu_wenxin(prompt, model_config)
elif provider in ["aliyun", "dashscope"]:
return await self._call_aliyun_dashscope(prompt, model_config)
else:
logger.error(f"Unsupported AI provider: {provider}")
return None
except Exception as e:
logger.error(f"AI API call error: {e}")
return None
async def _call_openai(
self,
prompt: str,
config: Dict[str, Any]
) -> Optional[Dict[str, Any]]:
"""调用 OpenAI API"""
try:
from openai import AsyncOpenAI
client = AsyncOpenAI(
api_key=config["api_key"],
base_url=config.get("api_url")
)
response = await client.chat.completions.create(
model=config.get("model_name", "gpt-3.5-turbo"),
messages=[
{"role": "user", "content": prompt}
],
temperature=config.get("temperature", 0.7),
max_tokens=config.get("max_tokens", 1000)
)
content = response.choices[0].message.content
tokens_used = response.usage.total_tokens if response.usage else 0
logger.info(f"OpenAI generated content, tokens: {tokens_used}")
return {
"content": content,
"tokens_used": tokens_used,
"provider": "openai",
"model": config.get("model_name", "gpt-3.5-turbo")
}
except Exception as e:
logger.error(f"OpenAI API error: {e}")
return None
async def _call_zhipu(
self,
prompt: str,
config: Dict[str, Any]
) -> Optional[Dict[str, Any]]:
"""调用智谱 AI API"""
try:
from zhipuai import ZhipuAI
client = ZhipuAI(api_key=config["api_key"])
response = client.chat.completions.create(
model=config.get("model_name", "glm-4"),
messages=[
{"role": "user", "content": prompt}
],
temperature=config.get("temperature", 0.7),
max_tokens=config.get("max_tokens", 1000)
)
content = response.choices[0].message.content
tokens_used = response.usage.total_tokens if response.usage else 0
logger.info(f"Zhipu AI generated content, tokens: {tokens_used}")
return {
"content": content,
"tokens_used": tokens_used,
"provider": "zhipu",
"model": config.get("model_name", "glm-4")
}
except Exception as e:
logger.error(f"Zhipu AI API error: {e}")
return None
async def _call_baidu_wenxin(
self,
prompt: str,
config: Dict[str, Any]
) -> Optional[Dict[str, Any]]:
"""调用百度文心一言 API"""
# TODO: 实现百度文心一言 API 调用
logger.warning("Baidu Wenxin API not implemented yet")
return None
async def _call_aliyun_dashscope(
self,
prompt: str,
config: Dict[str, Any]
) -> Optional[Dict[str, Any]]:
"""调用阿里云通义千问 API"""
# TODO: 实现阿里云 DashScope API 调用
logger.warning("Aliyun DashScope API not implemented yet")
return None
async def test_model(
self,
model_config: Dict[str, Any],
test_prompt: str = "测试评论"
) -> Dict[str, Any]:
"""
测试模型配置
:param model_config: 模型配置
:param test_prompt: 测试提示词
:return: 测试结果
"""
import time
start_time = time.time()
result = await self._call_ai_api(test_prompt, model_config)
cost_time = time.time() - start_time
if result:
start = time.time()
async with httpx.AsyncClient(timeout=model.timeout_seconds) as client:
resp = await client.post(f"{base_url}/chat/completions", headers=headers, json=payload)
resp.raise_for_status()
data = resp.json()
elapsed = round(time.time() - start, 2)
content = data["choices"][0]["message"]["content"]
tokens = data.get("usage", {}).get("total_tokens", 0)
return {
"success": True,
"content": result.get("content"),
"tokens_used": result.get("tokens_used", 0),
"cost_time": round(cost_time, 2),
"error_message": None
"success": True, "content": content,
"tokens": tokens, "elapsed_seconds": elapsed,
}
except Exception as e:
return {"success": False, "error": str(e)}
async def _record_token_usage(
self, db: AsyncSession, total: int, usage: dict, model_name: str
):
"""记录Token消耗"""
today = date.today()
from sqlalchemy.dialects.mysql import insert as mysql_insert
try:
existing = await db.execute(
select(TokenStat).where(TokenStat.stat_date == today)
)
stat = existing.scalar_one_or_none()
if stat:
stat.total_tokens += total
stat.prompt_tokens += usage.get("prompt_tokens", 0)
stat.completion_tokens += usage.get("completion_tokens", 0)
stat.call_count += 1
else:
return {
"success": False,
"content": None,
"tokens_used": 0,
"cost_time": round(cost_time, 2),
"error_message": "Failed to generate content"
}
stat = TokenStat(
stat_date=today,
model_name=model_name,
total_tokens=total,
prompt_tokens=usage.get("prompt_tokens", 0),
completion_tokens=usage.get("completion_tokens", 0),
call_count=1,
)
db.add(stat)
except Exception as e:
logger.error(f"记录Token消耗失败: {e}")
# 创建全局服务实例
ai_service = AIService()
+399
View File
@@ -0,0 +1,399 @@
"""数字分身管理服务层 — 同步连接数字分身应用的 SQLite 数据库"""
import json
import os
from datetime import datetime, timedelta
from typing import Optional, Tuple
from sqlalchemy import create_engine, select, text
from sqlalchemy.orm import sessionmaker, Session
from app.core.config import settings
from app.core.logger import logger
from app.models import UserPersonality, VirtualUser
_engine = None
_SessionLocal: Optional[sessionmaker] = None
AVATAR_ACCOUNT_PREFIX = "__avatar__:"
SQUARE_INTERACTION_PERMISSION = "interact"
SQUARE_INTERACTION_ACTIONS = frozenset({"like", "collect", "comment", "reply"})
def _get_engine_and_session():
global _engine, _SessionLocal
if _engine is None:
db_path = settings.AVATAR_DB_PATH
if not db_path:
# 默认路径:从 backend/app/core/ 向上三级
base = os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
db_path = os.path.join(base, "digital-avatar-app", "backend", "avatar.db")
if not os.path.isabs(db_path):
db_path = os.path.abspath(db_path)
if not os.path.exists(db_path):
# 数据库不存在时返回 None,由调用方处理
return None, None
_engine = create_engine(
f"sqlite:///{db_path}",
connect_args={"check_same_thread": False},
)
_SessionLocal = sessionmaker(bind=_engine, autoflush=False, expire_on_commit=False)
return _engine, _SessionLocal()
def get_session() -> Optional[Session]:
_, session = _get_engine_and_session()
return session
def is_available() -> bool:
"""检查数字分身数据库是否可用"""
engine, _ = _get_engine_and_session()
return engine is not None
def _resolve_photo_url(photo_url: str) -> str:
"""将相对路径的头像 URL 补全为绝对路径"""
if not photo_url:
return ""
if photo_url.startswith(("http://", "https://")):
return photo_url
base = settings.AVATAR_BACKEND_URL
if base:
base = base.rstrip("/")
return f"{base}{photo_url}"
return photo_url
def _get_global_token_balance(db: Session) -> int:
"""获取全局 token_account 余额(单行表)"""
try:
result = db.execute(text("SELECT balance FROM token_account LIMIT 1")).fetchone()
return result.balance if result and result.balance else 0
except Exception:
return 0
def _decode_config(value) -> dict:
if isinstance(value, dict):
return value
if isinstance(value, str):
try:
decoded = json.loads(value)
return decoded if isinstance(decoded, dict) else {}
except (json.JSONDecodeError, ValueError):
return {}
return {}
def is_delegated_avatar_user(user: VirtualUser | None) -> bool:
return bool(user and (user.account or "").startswith(AVATAR_ACCOUNT_PREFIX))
def delegated_avatar_id(user: VirtualUser | None) -> str:
if not is_delegated_avatar_user(user):
return ""
return (user.account or "")[len(AVATAR_ACCOUNT_PREFIX):]
def _list_square_interaction_authorizations(db: Session) -> list[dict]:
"""读取已明确授权分身参与广场互动的身份与会会令牌。"""
rows = db.execute(text("""
SELECT
a.id AS avatar_id,
a.name AS avatar_name,
a.display_name AS avatar_display_name,
a.description AS avatar_description,
a.photo_url AS avatar_photo_url,
a.config AS avatar_config,
u.huihui_user_id,
u.nickname AS owner_nickname,
u.avatar_url AS owner_avatar_url,
u.huihui_token
FROM avatars a
JOIN users u ON u.huihui_user_id = a.owner_id
WHERE a.status = 'active'
""")).fetchall()
authorized = []
for row in rows:
config = _decode_config(row.avatar_config)
permissions = config.get("authorizationPermissions", [])
if not isinstance(permissions, list) or SQUARE_INTERACTION_PERMISSION not in permissions:
continue
platform_uid = str(row.huihui_user_id or "").strip()
token = str(row.huihui_token or "").strip()
if not platform_uid or not token:
continue
authorized.append({
"avatar_id": str(row.avatar_id),
"avatar_name": row.avatar_display_name or row.avatar_name or row.owner_nickname or "数字分身",
"avatar_description": row.avatar_description or "",
"avatar_url": _resolve_photo_url(row.avatar_photo_url or row.owner_avatar_url or ""),
"config": config,
"platform_uid": platform_uid,
"token": token,
})
return authorized
def get_square_interaction_permissions(avatar_id: str) -> frozenset[str]:
"""实时复核授权;数据库不可用、令牌失效或撤权时一律拒绝执行。"""
avatar_db = get_session()
if avatar_db is None:
return frozenset()
try:
authorized_ids = {
item["avatar_id"] for item in _list_square_interaction_authorizations(avatar_db)
}
return SQUARE_INTERACTION_ACTIONS if avatar_id in authorized_ids else frozenset()
except Exception as exc:
logger.error(f"读取数字分身广场互动授权失败: {exc}")
return frozenset()
finally:
avatar_db.close()
def _word_count_range(config: dict) -> tuple[int, int]:
ranges = {
"short": (10, 35),
"medium": (20, 60),
"long": (30, 80),
}
return ranges.get(str(config.get("responseLength") or "medium"), (20, 60))
async def sync_square_interaction_users(db) -> set[str]:
"""把已授权分身同步为调度器身份,并刷新其会会会话。"""
avatar_db = get_session()
if avatar_db is None:
logger.warning("数字分身数据库不可用,跳过广场互动授权同步")
return set()
try:
authorized = _list_square_interaction_authorizations(avatar_db)
except Exception as exc:
logger.error(f"同步数字分身广场互动授权失败: {exc}")
return set()
finally:
avatar_db.close()
from app.core.redis_client import delete_session, set_session
result = await db.execute(
select(VirtualUser).where(VirtualUser.account.like(f"{AVATAR_ACCOUNT_PREFIX}%"))
)
existing_users = {delegated_avatar_id(user): user for user in result.scalars().all()}
authorized_ids = {item["avatar_id"] for item in authorized}
for avatar_id, user in existing_users.items():
if avatar_id not in authorized_ids:
user.is_enabled = 0
user.status = 0
user.session_token = None
user.session_expires_at = None
await delete_session(user.id)
for item in authorized:
avatar_id = item["avatar_id"]
user = existing_users.get(avatar_id)
if user is None:
user = VirtualUser(
nickname=item["avatar_name"],
account=f"{AVATAR_ACCOUNT_PREFIX}{avatar_id}",
password_enc="",
status=2,
is_enabled=1,
platform_uid=item["platform_uid"],
remark="用户授权的数字分身广场互动身份",
)
db.add(user)
await db.flush()
expires_at = datetime.now() + timedelta(days=1)
user.nickname = item["avatar_name"]
user.real_name = item["avatar_name"]
user.avatar_url = item["avatar_url"]
user.platform_uid = item["platform_uid"]
user.session_token = item["token"]
user.session_expires_at = expires_at
user.last_login_at = datetime.now()
user.status = 2
user.is_enabled = 1
config = item["config"]
personality_result = await db.execute(
select(UserPersonality).where(UserPersonality.user_id == user.id)
)
personality = personality_result.scalar_one_or_none()
word_min, word_max = _word_count_range(config)
prompt_parts = [item["avatar_description"], str(config.get("systemPrompt") or "")]
style_prompt = "\n".join(part.strip() for part in prompt_parts if part and part.strip())
if personality is None:
personality = UserPersonality(user_id=user.id)
db.add(personality)
personality.language_style = str(config.get("replyStyle") or "professional")
personality.personality_desc = item["avatar_description"]
personality.comment_style_prompt = style_prompt
personality.word_count_min = word_min
personality.word_count_max = word_max
await set_session(user.id, {
"token": item["token"],
"session_id": f"avatar:{avatar_id}",
"platform_uid": item["platform_uid"],
"org_id": "",
"login_time": datetime.now().isoformat(),
"nickname": item["avatar_name"],
"real_name": item["avatar_name"],
"avatar": item["avatar_url"],
"delegated_avatar_id": avatar_id,
}, expire=86400)
await db.commit()
return authorized_ids
class AvatarService:
@staticmethod
def list_avatars(
db: Session,
page: int = 1,
page_size: int = 20,
keyword: Optional[str] = None,
status: Optional[str] = None,
) -> Tuple[int, list]:
"""分页查询所有数字分身(含归属用户信息)"""
where_clauses = []
params = {}
if keyword:
where_clauses.append(
"(a.name LIKE :kw OR a.display_name LIKE :kw)"
)
params["kw"] = f"%{keyword}%"
if status:
where_clauses.append("a.status = :status")
params["status"] = status
where_sql = ""
if where_clauses:
where_sql = "WHERE " + " AND ".join(where_clauses)
# 计数
count_sql = f"SELECT COUNT(*) FROM avatars a {where_sql}"
total = db.execute(text(count_sql), params).scalar() or 0
# 分页查询
offset = (page - 1) * page_size
params["limit"] = page_size
params["offset"] = offset
query = text(f"""
SELECT
a.id, a.name, a.display_name, a.description,
a.photo_url, a.emoji, a.status, a.token_balance,
a.config, a.created_at, a.updated_at,
u.nickname AS owner_nickname,
u.phone AS owner_phone
FROM avatars a
LEFT JOIN users u ON a.owner_id = u.huihui_user_id
{where_sql}
ORDER BY a.created_at DESC
LIMIT :limit OFFSET :offset
""")
rows = db.execute(query, params).fetchall()
items = []
global_token = _get_global_token_balance(db)
for row in rows:
config = {}
if row.config:
if isinstance(row.config, str):
import json
try:
config = json.loads(row.config)
except (json.JSONDecodeError, ValueError):
config = {}
elif isinstance(row.config, dict):
config = row.config
items.append({
"id": row.id,
"name": row.name,
"display_name": row.display_name,
"description": row.description or "",
"photo_url": _resolve_photo_url(row.photo_url),
"emoji": row.emoji or "🤖",
"status": row.status or "active",
"token_balance": global_token or (row.token_balance or 0),
"config": config,
"owner_nickname": row.owner_nickname or "",
"owner_phone": row.owner_phone or "",
"created_at": str(row.created_at) if row.created_at else "",
"updated_at": str(row.updated_at) if row.updated_at else "",
})
return total, items
@staticmethod
def get_avatar(db: Session, avatar_id: str) -> Optional[dict]:
"""获取单个数字分身详情"""
query = text("""
SELECT
a.id, a.name, a.display_name, a.description,
a.photo_url, a.emoji, a.status, a.token_balance,
a.config, a.created_at, a.updated_at,
u.nickname AS owner_nickname,
u.phone AS owner_phone
FROM avatars a
LEFT JOIN users u ON a.owner_id = u.huihui_user_id
WHERE a.id = :avatar_id
""")
row = db.execute(query, {"avatar_id": avatar_id}).fetchone()
if not row:
return None
config = {}
if row.config:
if isinstance(row.config, str):
import json
try:
config = json.loads(row.config)
except (json.JSONDecodeError, ValueError):
config = {}
elif isinstance(row.config, dict):
config = row.config
global_token = _get_global_token_balance(db)
return {
"id": row.id,
"name": row.name,
"display_name": row.display_name,
"description": row.description or "",
"photo_url": _resolve_photo_url(row.photo_url),
"emoji": row.emoji or "🤖",
"status": row.status or "active",
"token_balance": global_token or (row.token_balance or 0),
"config": config,
"owner_nickname": row.owner_nickname or "",
"owner_phone": row.owner_phone or "",
"created_at": str(row.created_at) if row.created_at else "",
"updated_at": str(row.updated_at) if row.updated_at else "",
}
@staticmethod
def update_status(db: Session, avatar_id: str, status: str) -> Optional[dict]:
"""更新数字分身状态(开关机)"""
query = text("""
UPDATE avatars SET status = :status, updated_at = datetime('now')
WHERE id = :avatar_id
""")
result = db.execute(query, {"status": status, "avatar_id": avatar_id})
db.commit()
if result.rowcount == 0:
return None
return AvatarService.get_avatar(db, avatar_id)
avatar_service = AvatarService()
-291
View File
@@ -1,291 +0,0 @@
"""
会会接口对接服务
基于 http://192.168.1.200:63120/doc.html 接口文档
"""
import httpx
import logging
from typing import Optional, Dict, Any, List
from datetime import datetime
from app.core.config import settings
logger = logging.getLogger(__name__)
class HuihuiAPIService:
"""会会 API 服务类"""
def __init__(self):
self.base_url = settings.HUIHUI_API_BASE
self.timeout = 30 # 秒
self._session_cache: Dict[str, httpx.AsyncClient] = {}
def _get_client(self, session_token: Optional[str] = None) -> httpx.AsyncClient:
"""获取 HTTP 客户端"""
headers = {
"Content-Type": "application/json",
"Accept": "application/json",
}
if session_token:
headers["Authorization"] = f"Bearer {session_token}"
return httpx.AsyncClient(
base_url=self.base_url,
headers=headers,
timeout=self.timeout,
)
async def login(self, username: str, password: str) -> Optional[Dict[str, Any]]:
"""
用户登录
:param username: 用户名
:param password: 密码
:return: 登录响应(包含 session token)
"""
try:
async with self._get_client() as client:
response = await client.post(
"/api/login", # 实际接口路径需根据 doc.html 调整
json={"username": username, "password": password}
)
if response.status_code == 200:
data = response.json()
logger.info(f"Login success for user: {username}")
return data
else:
logger.error(f"Login failed: {response.status_code} - {response.text}")
return None
except Exception as e:
logger.error(f"Login error: {e}")
return None
async def get_news_list(
self,
page: int = 1,
page_size: int = 20,
category: Optional[str] = None
) -> Optional[List[Dict[str, Any]]]:
"""
获取新闻列表
:param page: 页码
:param page_size: 每页数量
:param category: 分类(可选)
:return: 新闻列表
"""
try:
async with self._get_client() as client:
params = {"page": page, "pageSize": page_size}
if category:
params["category"] = category
response = await client.get(
"/api/news/list", # 实际接口路径需根据 doc.html 调整
params=params
)
if response.status_code == 200:
data = response.json()
return data.get("data", [])
else:
logger.error(f"Get news list failed: {response.status_code}")
return None
except Exception as e:
logger.error(f"Get news list error: {e}")
return None
async def get_news_detail(self, news_id: str) -> Optional[Dict[str, Any]]:
"""
获取新闻详情
:param news_id: 新闻 ID
:return: 新闻详情
"""
try:
async with self._get_client() as client:
response = await client.get(f"/api/news/{news_id}")
if response.status_code == 200:
data = response.json()
return data.get("data")
else:
logger.error(f"Get news detail failed: {response.status_code}")
return None
except Exception as e:
logger.error(f"Get news detail error: {e}")
return None
async def create_comment(
self,
news_id: str,
content: str,
session_token: str
) -> Optional[Dict[str, Any]]:
"""
创建评论
:param news_id: 新闻 ID
:param content: 评论内容
:param session_token: 会话 Token
:return: 评论结果
"""
try:
async with self._get_client(session_token) as client:
response = await client.post(
"/api/comment/create", # 实际接口路径需根据 doc.html 调整
json={"newsId": news_id, "content": content}
)
if response.status_code == 200:
data = response.json()
logger.info(f"Comment created for news: {news_id}")
return data
else:
logger.error(f"Create comment failed: {response.status_code} - {response.text}")
return None
except Exception as e:
logger.error(f"Create comment error: {e}")
return None
async def create_reply(
self,
comment_id: str,
content: str,
session_token: str
) -> Optional[Dict[str, Any]]:
"""
创建回复
:param comment_id: 评论 ID
:param content: 回复内容
:param session_token: 会话 Token
:return: 回复结果
"""
try:
async with self._get_client(session_token) as client:
response = await client.post(
"/api/reply/create", # 实际接口路径需根据 doc.html 调整
json={"commentId": comment_id, "content": content}
)
if response.status_code == 200:
data = response.json()
logger.info(f"Reply created for comment: {comment_id}")
return data
else:
logger.error(f"Create reply failed: {response.status_code}")
return None
except Exception as e:
logger.error(f"Create reply error: {e}")
return None
async def like_comment(
self,
comment_id: str,
session_token: str
) -> Optional[Dict[str, Any]]:
"""
点赞评论
:param comment_id: 评论 ID
:param session_token: 会话 Token
:return: 点赞结果
"""
try:
async with self._get_client(session_token) as client:
response = await client.post(f"/api/comment/{comment_id}/like")
if response.status_code == 200:
data = response.json()
logger.info(f"Comment liked: {comment_id}")
return data
else:
logger.error(f"Like comment failed: {response.status_code}")
return None
except Exception as e:
logger.error(f"Like comment error: {e}")
return None
async def favorite_news(
self,
news_id: str,
session_token: str
) -> Optional[Dict[str, Any]]:
"""
收藏新闻
:param news_id: 新闻 ID
:param session_token: 会话 Token
:return: 收藏结果
"""
try:
async with self._get_client(session_token) as client:
response = await client.post(f"/api/news/{news_id}/favorite")
if response.status_code == 200:
data = response.json()
logger.info(f"News favorited: {news_id}")
return data
else:
logger.error(f"Favorite news failed: {response.status_code}")
return None
except Exception as e:
logger.error(f"Favorite news error: {e}")
return None
async def share_news(
self,
news_id: str,
session_token: str
) -> Optional[Dict[str, Any]]:
"""
转发新闻
:param news_id: 新闻 ID
:param session_token: 会话 Token
:return: 转发结果
"""
try:
async with self._get_client(session_token) as client:
response = await client.post(f"/api/news/{news_id}/share")
if response.status_code == 200:
data = response.json()
logger.info(f"News shared: {news_id}")
return data
else:
logger.error(f"Share news failed: {response.status_code}")
return None
except Exception as e:
logger.error(f"Share news error: {e}")
return None
async def get_comments(
self,
news_id: str,
page: int = 1,
page_size: int = 20
) -> Optional[List[Dict[str, Any]]]:
"""
获取新闻评论列表
:param news_id: 新闻 ID
:param page: 页码
:param page_size: 每页数量
:return: 评论列表
"""
try:
async with self._get_client() as client:
params = {"page": page, "pageSize": page_size}
response = await client.get(
f"/api/news/{news_id}/comments",
params=params
)
if response.status_code == 200:
data = response.json()
return data.get("data", [])
else:
logger.error(f"Get comments failed: {response.status_code}")
return None
except Exception as e:
logger.error(f"Get comments error: {e}")
return None
# 创建全局服务实例
huihui_api_service = HuihuiAPIService()
-408
View File
@@ -1,408 +0,0 @@
"""
互动执行服务
"""
import logging
import random
from typing import Optional, List, Dict, Any
from datetime import datetime, timedelta
from sqlalchemy.orm import Session
from sqlalchemy import and_, func
from app.models.virtual_user import VirtualUser, ActivityLevel, UserStatus
from app.models.interaction import InteractionRecord, InteractionType, InteractionStatus
from app.models.token_usage import TokenUsage
from app.models.news_cache import NewsCache
from app.services.huihui_api_service import huihui_api_service
from app.services.ai_service import ai_service
from app.core.config import settings
logger = logging.getLogger(__name__)
class InteractionService:
"""互动执行服务类"""
def __init__(self, db: Session):
self.db = db
async def execute_interaction(
self,
virtual_user_id: int,
interaction_type: Optional[InteractionType] = None,
news_id: Optional[str] = None
) -> Optional[InteractionRecord]:
"""
执行单次互动
:param virtual_user_id: 虚拟用户 ID
:param interaction_type: 互动类型(不传则随机)
:param news_id: 新闻 ID(不传则随机选择)
:return: 互动记录
"""
# 获取虚拟用户
user = self.db.query(VirtualUser).filter(VirtualUser.id == virtual_user_id).first()
if not user:
logger.error(f"Virtual user not found: {virtual_user_id}")
return None
# 检查用户状态
if user.status != UserStatus.ACTIVE:
logger.warning(f"Virtual user is not active: {virtual_user_id}")
return None
# 检查是否已登录
if not user.is_logged_in or not user.session_token:
logger.warning(f"Virtual user not logged in: {virtual_user_id}")
# TODO: 自动登录
return None
# 检查今日限额
if not self._check_daily_limit(user, interaction_type):
logger.warning(f"Daily limit reached for user {virtual_user_id}")
return None
# 选择新闻
if not news_id:
news_id = await self._select_news(user)
if not news_id:
logger.warning("No news available for interaction")
return None
# 获取新闻详情
news = self.db.query(NewsCache).filter(NewsCache.news_id == news_id).first()
if not news:
# 从 API 获取
news_detail = await huihui_api_service.get_news_detail(news_id)
if news_detail:
news = self._cache_news(news_detail)
if not news:
logger.error(f"Cannot get news detail: {news_id}")
return None
# 确定互动类型
if not interaction_type:
interaction_type = self._random_interaction_type()
# 创建互动记录
record = InteractionRecord(
virtual_user_id=virtual_user_id,
news_id=news_id,
news_title=news.title if news else "",
interaction_type=interaction_type,
status=InteractionStatus.PENDING
)
self.db.add(record)
self.db.commit()
self.db.refresh(record)
try:
# 执行互动
if interaction_type == InteractionType.COMMENT:
result = await self._execute_comment(user, news, record)
elif interaction_type == InteractionType.REPLY:
result = await self._execute_reply(user, news, record)
elif interaction_type == InteractionType.LIKE:
result = await self._execute_like(user, news, record)
elif interaction_type == InteractionType.FAVORITE:
result = await self._execute_favorite(user, news, record)
elif interaction_type == InteractionType.SHARE:
result = await self._execute_share(user, news, record)
else:
logger.error(f"Unknown interaction type: {interaction_type}")
return None
if result:
record.status = InteractionStatus.SUCCESS
record.api_response = str(result)
# 更新用户统计
user.total_interactions += 1
if interaction_type == InteractionType.COMMENT:
user.today_comments += 1
elif interaction_type == InteractionType.REPLY:
user.today_replies += 1
user.last_interaction_time = datetime.now()
self.db.commit()
logger.info(f"Interaction executed successfully: user={user.id}, type={interaction_type}")
return record
else:
record.status = InteractionStatus.FAILED
record.error_message = "API call failed"
self.db.commit()
return None
except Exception as e:
logger.error(f"Execute interaction error: {e}")
record.status = InteractionStatus.FAILED
record.error_message = str(e)
self.db.commit()
return None
async def _execute_comment(
self,
user: VirtualUser,
news: NewsCache,
record: InteractionRecord
) -> Optional[Dict[str, Any]]:
"""执行评论"""
# AI 生成评论内容
ai_result = await ai_service.generate_comment(
news_content=news.content or news.summary or news.title,
writing_style=user.writing_style or "普通",
persona_description=user.persona_description
)
if not ai_result or not ai_result.get("content"):
logger.error("AI generate comment failed")
return None
# 记录 Token 使用
self._record_token_usage(
virtual_user_id=user.id,
interaction_id=record.id,
tokens_used=ai_result.get("tokens_used", 0),
ai_model=ai_result.get("model", "unknown"),
action_type="generate_comment",
tokens_prompt=ai_result.get("tokens_prompt", 0),
tokens_completion=ai_result.get("tokens_completion", 0)
)
# 调用接口提交评论
result = await huihui_api_service.create_comment(
news_id=news.news_id,
content=ai_result["content"],
session_token=user.session_token
)
if result:
record.content = ai_result["content"]
record.tokens_used = ai_result.get("tokens_used", 0)
record.ai_model_used = ai_result.get("model")
return result
async def _execute_reply(
self,
user: VirtualUser,
news: NewsCache,
record: InteractionRecord
) -> Optional[Dict[str, Any]]:
"""执行回复"""
# 获取评论列表
comments = await huihui_api_service.get_comments(news_id=news.news_id)
if not comments or len(comments) == 0:
logger.warning(f"No comments available for news {news.news_id}")
return None
# 随机选择一条评论进行回复
target_comment = random.choice(comments)
record.target_comment_id = target_comment.get("id")
# AI 生成回复内容
ai_result = await ai_service.generate_reply(
original_comment=target_comment.get("content", ""),
news_content=news.content or news.summary or news.title,
writing_style=user.writing_style or "普通",
persona_description=user.persona_description
)
if not ai_result or not ai_result.get("content"):
logger.error("AI generate reply failed")
return None
# 记录 Token 使用
self._record_token_usage(
virtual_user_id=user.id,
interaction_id=record.id,
tokens_used=ai_result.get("tokens_used", 0),
ai_model=ai_result.get("model", "unknown"),
action_type="generate_reply"
)
# 调用接口提交回复
result = await huihui_api_service.create_reply(
comment_id=target_comment.get("id"),
content=ai_result["content"],
session_token=user.session_token
)
if result:
record.content = ai_result["content"]
record.tokens_used = ai_result.get("tokens_used", 0)
return result
async def _execute_like(
self,
user: VirtualUser,
news: NewsCache,
record: InteractionRecord
) -> Optional[Dict[str, Any]]:
"""执行点赞"""
# 获取评论列表
comments = await huihui_api_service.get_comments(news_id=news.news_id)
if not comments or len(comments) == 0:
logger.warning(f"No comments available for like")
return None
# 随机选择一条评论点赞
target_comment = random.choice(comments)
record.target_comment_id = target_comment.get("id")
result = await huihui_api_service.like_comment(
comment_id=target_comment.get("id"),
session_token=user.session_token
)
return result
async def _execute_favorite(
self,
user: VirtualUser,
news: NewsCache,
record: InteractionRecord
) -> Optional[Dict[str, Any]]:
"""执行收藏"""
result = await huihui_api_service.favorite_news(
news_id=news.news_id,
session_token=user.session_token
)
return result
async def _execute_share(
self,
user: VirtualUser,
news: NewsCache,
record: InteractionRecord
) -> Optional[Dict[str, Any]]:
"""执行转发"""
result = await huihui_api_service.share_news(
news_id=news.news_id,
session_token=user.session_token
)
return result
def _check_daily_limit(
self,
user: VirtualUser,
interaction_type: Optional[InteractionType]
) -> bool:
"""检查每日限额"""
today = datetime.now().date()
# 统计今日互动
today_records = self.db.query(InteractionRecord).filter(
and_(
InteractionRecord.virtual_user_id == user.id,
func.date(InteractionRecord.execution_time) == today,
InteractionRecord.status == InteractionStatus.SUCCESS
)
).all()
today_comments = sum(1 for r in today_records if r.interaction_type == InteractionType.COMMENT)
today_replies = sum(1 for r in today_records if r.interaction_type == InteractionType.REPLY)
# 检查评论限额
if interaction_type == InteractionType.COMMENT:
if today_comments >= settings.MAX_COMMENTS_PER_USER_PER_DAY:
return False
# 检查回复限额
if interaction_type == InteractionType.REPLY:
if today_replies >= settings.MAX_REPLIES_PER_USER_PER_DAY:
return False
return True
def _random_interaction_type(self) -> InteractionType:
"""随机选择互动类型"""
rand = random.random()
# 根据概率决定互动类型
if rand < settings.LIKE_PROBABILITY:
return InteractionType.LIKE
elif rand < settings.LIKE_PROBABILITY + settings.FAVORITE_PROBABILITY:
return InteractionType.FAVORITE
elif rand < settings.LIKE_PROBABILITY + settings.FAVORITE_PROBABILITY + settings.SHARE_PROBABILITY:
return InteractionType.SHARE
else:
return InteractionType.COMMENT
async def _select_news(self, user: VirtualUser) -> Optional[str]:
"""选择新闻"""
# 优先选择未互动过的新闻
cached_news = self.db.query(NewsCache).order_by(
NewsCache.created_at.desc()
).limit(50).all()
if not cached_news:
# 从 API 获取
news_list = await huihui_api_service.get_news_list(page=1, page_size=20)
if news_list:
for news_data in news_list:
self._cache_news(news_data)
cached_news = self.db.query(NewsCache).order_by(
NewsCache.created_at.desc()
).limit(50).all()
if not cached_news:
return None
# 随机选择一篇
return random.choice(cached_news).news_id
def _cache_news(self, news_data: Dict[str, Any]) -> Optional[NewsCache]:
"""缓存新闻"""
news = NewsCache(
news_id=str(news_data.get("id")),
title=news_data.get("title", ""),
summary=news_data.get("summary", ""),
content=news_data.get("content", ""),
source=news_data.get("source", ""),
author=news_data.get("author", ""),
category=news_data.get("category", "")
)
self.db.add(news)
self.db.commit()
self.db.refresh(news)
return news
def _record_token_usage(
self,
virtual_user_id: int,
interaction_id: int,
tokens_used: int,
ai_model: str,
action_type: str,
tokens_prompt: int = 0,
tokens_completion: int = 0
):
"""记录 Token 使用"""
usage = TokenUsage(
virtual_user_id=virtual_user_id,
interaction_id=interaction_id,
tokens_used=tokens_used,
tokens_prompt=tokens_prompt,
tokens_completion=tokens_completion,
ai_model=ai_model,
action_type=action_type
)
self.db.add(usage)
self.db.commit()
logger.info(f"Token usage recorded: {tokens_used} tokens for user {virtual_user_id}")
# 工厂函数
def get_interaction_service(db: Session) -> InteractionService:
"""获取互动服务实例"""
return InteractionService(db)
+1530
View File
File diff suppressed because it is too large Load Diff
+1023
View File
File diff suppressed because it is too large Load Diff
-215
View File
@@ -1,215 +0,0 @@
"""
定时任务调度服务
基于 APScheduler 实现
"""
import logging
import random
import asyncio
from typing import Optional, List
from datetime import datetime, time
from apscheduler.schedulers.asyncio import AsyncIOScheduler
from apscheduler.triggers.cron import CronTrigger
from apscheduler.triggers.interval import IntervalTrigger
from sqlalchemy.orm import Session
from app.models.virtual_user import VirtualUser, ActivityLevel, UserStatus
from app.models.base import get_db, SessionLocal
from app.services.interaction_service import InteractionService
from app.core.config import settings
logger = logging.getLogger(__name__)
class SchedulerService:
"""定时任务调度服务类"""
def __init__(self):
self.scheduler = AsyncIOScheduler()
self.is_running = False
self._current_job = None
def start(self):
"""启动调度器"""
if not self.is_running:
self.scheduler.start()
self.is_running = True
logger.info("Scheduler started")
def stop(self):
"""停止调度器"""
if self.is_running:
self.scheduler.shutdown()
self.is_running = False
logger.info("Scheduler stopped")
def add_interaction_task(self):
"""添加互动任务"""
# 在活动时间段内,每隔随机时间执行一次互动
# 由于 APScheduler 不支持随机间隔,我们使用固定间隔但通过概率控制执行
# 每 5 分钟检查一次
trigger = IntervalTrigger(minutes=5)
self.scheduler.add_job(
self._execute_random_interaction,
trigger=trigger,
id="random_interaction",
name="Random Interaction Task",
replace_existing=True
)
logger.info("Interaction task added")
def remove_interaction_task(self):
"""移除互动任务"""
try:
self.scheduler.remove_job("random_interaction")
logger.info("Interaction task removed")
except Exception as e:
logger.warning(f"Remove interaction task error: {e}")
async def _execute_random_interaction(self):
"""执行随机互动任务"""
# 检查是否在活动时间段内
now = datetime.now()
current_hour = now.hour
if current_hour < settings.TASK_START_HOUR or current_hour > settings.TASK_END_HOUR:
logger.debug(f"Outside activity hours: {current_hour}")
return
# 随机决定是否执行(通过随机间隔模拟)
if random.random() > 0.5: # 50% 概率执行
logger.debug("Skip this round")
return
logger.info("Executing random interaction task")
# 获取数据库会话
db = SessionLocal()
try:
# 获取所有活跃的虚拟用户
users = db.query(VirtualUser).filter(
VirtualUser.status == UserStatus.ACTIVE,
VirtualUser.is_logged_in == True
).all()
if not users:
logger.debug("No active logged-in users")
return
# 随机选择一个用户
user = random.choice(users)
# 检查用户活跃度
if not self._should_user_interact(user):
logger.debug(f"User {user.id} should not interact now")
return
# 执行互动
interaction_service = InteractionService(db)
await interaction_service.execute_interaction(virtual_user_id=user.id)
except Exception as e:
logger.error(f"Execute random interaction error: {e}")
finally:
db.close()
def _should_user_interact(self, user: VirtualUser) -> bool:
"""根据活跃度判断用户是否应该互动"""
# 根据活跃度决定互动概率
if user.activity_level == ActivityLevel.HIGH:
# 高活跃度:80% 概率
return random.random() < 0.8
elif user.activity_level == ActivityLevel.MEDIUM:
# 中活跃度:50% 概率
return random.random() < 0.5
else:
# 低活跃度:30% 概率
return random.random() < 0.3
def add_login_task(self, hour: int = 8, minute: int = 0):
"""添加每日登录任务"""
trigger = CronTrigger(hour=hour, minute=minute)
self.scheduler.add_job(
self._auto_login_users,
trigger=trigger,
id="daily_login",
name="Daily Auto Login",
replace_existing=True
)
logger.info(f"Daily login task added at {hour:02d}:{minute:02d}")
async def _auto_login_users(self):
"""自动登录所有活跃用户"""
db = SessionLocal()
try:
from app.services.huihui_api_service import huihui_api_service
users = db.query(VirtualUser).filter(
VirtualUser.status == UserStatus.ACTIVE
).all()
for user in users:
try:
# 调用登录接口
result = await huihui_api_service.login(user.username, user.password)
if result and result.get("token"):
user.is_logged_in = True
user.session_token = result["token"]
# TODO: 设置 token 过期时间
logger.info(f"Auto login success: {user.username}")
else:
logger.warning(f"Auto login failed: {user.username}")
except Exception as e:
logger.error(f"Auto login error for {user.username}: {e}")
db.commit()
except Exception as e:
logger.error(f"Auto login task error: {e}")
db.rollback()
finally:
db.close()
def reset_daily_counters(self, hour: int = 0, minute: int = 1):
"""添加每日计数器重置任务"""
trigger = CronTrigger(hour=hour, minute=minute)
self.scheduler.add_job(
self._reset_daily_counters,
trigger=trigger,
id="reset_daily_counters",
name="Reset Daily Counters",
replace_existing=True
)
logger.info(f"Daily reset task added at {hour:02d}:{minute:02d}")
def _reset_daily_counters(self):
"""重置每日计数器"""
db = SessionLocal()
try:
# 重置所有用户的今日计数
db.query(VirtualUser).update({
VirtualUser.today_comments: 0,
VirtualUser.today_replies: 0
})
db.commit()
logger.info("Daily counters reset")
except Exception as e:
logger.error(f"Reset daily counters error: {e}")
db.rollback()
finally:
db.close()
# 创建全局服务实例
scheduler_service = SchedulerService()
+251
View File
@@ -0,0 +1,251 @@
"""数据统计服务"""
from datetime import datetime, date, timedelta, timezone
def _fmt_dt(dt):
"""统一输出 UTC 时间,带时区标识,让前端正确解析为 +8"""
if dt is None:
return None
if dt.tzinfo is None:
# 数据库存的是 UTC,补上时区信息
dt = dt.replace(tzinfo=timezone.utc)
return dt.isoformat()
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, func, and_
from app.models import VirtualUser, InteractionRecord, TokenStat, SystemConfig
from app.core.logger import logger
class StatsService:
async def get_dashboard(self, db: AsyncSession) -> dict:
"""获取控制台数据"""
today = date.today()
now = datetime.now()
month_start = today.replace(day=1)
# 用户统计
user_stats = await self._get_user_stats(db)
# 今日互动统计
today_stats = await self._get_today_stats(db, today)
# 本月互动统计
monthly_stats = await self._get_monthly_stats(db, month_start, today)
# Token统计
token_stats = await self._get_token_stats(db, today)
# 系统状态
system_status = await self._get_system_status(db, now)
# 在线用户数
online_count_result = await db.execute(
select(func.count()).where(VirtualUser.status == 2)
)
online_count = online_count_result.scalar() or 0
return {
"user_stats": user_stats,
"today_interactions": today_stats,
"monthly_stats": monthly_stats,
"token_stats": token_stats,
"system_status": system_status,
"online_users": online_count,
}
async def _get_user_stats(self, db: AsyncSession) -> dict:
total = await db.execute(select(func.count()).select_from(VirtualUser))
normal = await db.execute(select(func.count()).where(VirtualUser.is_enabled == 1))
banned = await db.execute(select(func.count()).where(VirtualUser.status == 4))
abnormal = await db.execute(select(func.count()).where(VirtualUser.status == 3))
return {
"total": total.scalar() or 0,
"normal": normal.scalar() or 0,
"banned": banned.scalar() or 0,
"abnormal": abnormal.scalar() or 0,
}
async def _get_today_stats(self, db: AsyncSession, today: date) -> dict:
result = await db.execute(
select(
InteractionRecord.interact_type,
func.count().label("cnt"),
).where(
func.date(InteractionRecord.executed_at) == today,
InteractionRecord.status == 1,
).group_by(InteractionRecord.interact_type)
)
rows = result.all()
stats = {"comment": 0, "reply": 0, "like": 0, "collect": 0, "forward": 0, "total": 0}
for row in rows:
if row.interact_type in stats:
stats[row.interact_type] = row.cnt
stats["total"] += row.cnt
return stats
async def _get_monthly_stats(self, db: AsyncSession, month_start: date, today: date) -> dict:
result = await db.execute(
select(func.count()).where(
InteractionRecord.executed_at >= month_start,
InteractionRecord.status == 1,
)
)
return {"total": result.scalar() or 0, "month_start": month_start.isoformat()}
async def _get_token_stats(self, db: AsyncSession, today: date) -> dict:
# 今日
today_stat = await db.execute(select(TokenStat).where(TokenStat.stat_date == today))
today_row = today_stat.scalar_one_or_none()
# 每日限额
limit_cfg = await db.execute(
select(SystemConfig).where(SystemConfig.config_key == "daily_token_limit")
)
limit_row = limit_cfg.scalar_one_or_none()
daily_limit = int(limit_row.config_value) if limit_row else 100000
today_used = today_row.total_tokens if today_row else 0
return {
"today_used": today_used,
"daily_limit": daily_limit,
"remaining": max(0, daily_limit - today_used),
"today_calls": today_row.call_count if today_row else 0,
}
async def _get_system_status(self, db: AsyncSession, now: datetime) -> dict:
start_cfg = await db.execute(
select(SystemConfig).where(SystemConfig.config_key == "system_start_time")
)
start_row = start_cfg.scalar_one_or_none()
uptime = ""
if start_row and start_row.config_value:
try:
start_time = datetime.fromisoformat(start_row.config_value)
delta = now - start_time
hours, rem = divmod(int(delta.total_seconds()), 3600)
mins = rem // 60
uptime = f"{hours}小时{mins}分钟"
except Exception:
uptime = "未知"
scheduler_cfg = await db.execute(
select(SystemConfig).where(SystemConfig.config_key == "scheduler_enabled")
)
scheduler_row = scheduler_cfg.scalar_one_or_none()
return {
"uptime": uptime,
"scheduler_enabled": (scheduler_row.config_value == "true") if scheduler_row else True,
"current_time": now.isoformat(),
}
async def get_token_trend(self, db: AsyncSession, days: int = 30) -> list:
"""Token消耗趋势(近N天)"""
end_date = date.today()
start_date = end_date - timedelta(days=days - 1)
result = await db.execute(
select(TokenStat).where(
TokenStat.stat_date >= start_date,
TokenStat.stat_date <= end_date,
).order_by(TokenStat.stat_date)
)
rows = result.scalars().all()
stat_map = {r.stat_date.isoformat(): r.total_tokens for r in rows}
trend = []
for i in range(days):
d = (start_date + timedelta(days=i)).isoformat()
trend.append({"date": d, "tokens": stat_map.get(d, 0)})
return trend
async def get_monthly_token_trend(self, db: AsyncSession) -> list:
"""近12个月Token消耗"""
today = date.today()
months = []
for i in range(11, -1, -1):
if today.month - i <= 0:
year = today.year - 1
month = today.month - i + 12
else:
year = today.year
month = today.month - i
months.append((year, month))
trend = []
for year, month in months:
start = date(year, month, 1)
if month == 12:
end = date(year + 1, 1, 1) - timedelta(days=1)
else:
end = date(year, month + 1, 1) - timedelta(days=1)
result = await db.execute(
select(func.sum(TokenStat.total_tokens)).where(
TokenStat.stat_date >= start, TokenStat.stat_date <= end
)
)
total = result.scalar() or 0
trend.append({"month": f"{year}-{month:02d}", "tokens": total})
return trend
async def get_interaction_records(
self, db: AsyncSession,
page: int = 1, page_size: int = 20,
user_id: int = None, interact_type: str = None,
status: int = None, start_date: str = None,
end_date: str = None, keyword: str = None
) -> dict:
query = select(InteractionRecord)
conditions = []
if user_id:
conditions.append(InteractionRecord.user_id == user_id)
if interact_type:
conditions.append(InteractionRecord.interact_type == interact_type)
if status is not None:
conditions.append(InteractionRecord.status == status)
if start_date:
conditions.append(InteractionRecord.executed_at >= start_date)
if end_date:
conditions.append(InteractionRecord.executed_at <= end_date + " 23:59:59")
if keyword:
from sqlalchemy import or_
conditions.append(
or_(InteractionRecord.article_title.like(f"%{keyword}%"),
InteractionRecord.content.like(f"%{keyword}%"),
InteractionRecord.user_nickname.like(f"%{keyword}%"))
)
if conditions:
query = query.where(and_(*conditions))
count_q = select(func.count()).select_from(query.subquery())
total = (await db.execute(count_q)).scalar()
query = query.order_by(InteractionRecord.executed_at.desc()).offset(
(page - 1) * page_size
).limit(page_size)
result = await db.execute(query)
records = result.scalars().all()
INTERACT_LABELS = {
"comment": "评论", "reply": "回复", "like": "点赞",
"collect": "收藏", "forward": "转发"
}
STATUS_LABELS = {0: "执行中", 1: "成功", 2: "失败", 3: "已取消"}
items = []
for r in records:
items.append({
"id": r.id, "user_id": r.user_id,
"user_nickname": r.user_nickname, "user_account": r.user_account,
"article_id": r.article_id, "article_title": r.article_title,
"interact_type": r.interact_type,
"interact_type_label": INTERACT_LABELS.get(r.interact_type, r.interact_type),
"content": r.content, "token_consumed": r.token_consumed,
"status": r.status, "status_label": STATUS_LABELS.get(r.status, "未知"),
"error_msg": r.error_msg, "retry_count": r.retry_count,
"executed_at": _fmt_dt(r.executed_at),
})
return {"total": total, "page": page, "page_size": page_size, "items": items}
stats_service = StatsService()
-173
View File
@@ -1,173 +0,0 @@
"""
Token 使用统计服务
"""
import logging
from typing import Dict, Any, List
from datetime import datetime, timedelta, date
from sqlalchemy.orm import Session
from sqlalchemy import func, and_, extract
from app.models.token_usage import TokenUsage
from app.core.config import settings
logger = logging.getLogger(__name__)
class TokenService:
"""Token 统计服务类"""
def __init__(self, db: Session):
self.db = db
def get_today_usage(self) -> int:
"""获取今日 Token 使用量"""
today = date.today()
result = self.db.query(func.sum(TokenUsage.tokens_used)).filter(
func.date(TokenUsage.usage_date) == today
).scalar()
return result or 0
def get_yesterday_usage(self) -> int:
"""获取昨日 Token 使用量"""
yesterday = date.today() - timedelta(days=1)
result = self.db.query(func.sum(TokenUsage.tokens_used)).filter(
func.date(TokenUsage.usage_date) == yesterday
).scalar()
return result or 0
def get_month_usage(self, year: Optional[int] = None, month: Optional[int] = None) -> int:
"""获取当月 Token 使用量"""
if not year or not month:
now = datetime.now()
year = now.year
month = now.month
result = self.db.query(func.sum(TokenUsage.tokens_used)).filter(
and_(
extract('year', TokenUsage.usage_date) == year,
extract('month', TokenUsage.usage_date) == month
)
).scalar()
return result or 0
def get_remaining_tokens(self) -> int:
"""获取今日剩余 Token"""
today_used = self.get_today_usage()
remaining = settings.MAX_TOKENS_PER_DAY - today_used
return max(0, remaining)
def get_daily_usages(self, days: int = 30) -> List[Dict[str, Any]]:
"""
获取每日 Token 使用(用于图表)
:param days: 天数
:return: 每日使用列表
"""
end_date = date.today()
start_date = end_date - timedelta(days=days - 1)
results = self.db.query(
func.date(TokenUsage.usage_date).label('usage_date'),
func.sum(TokenUsage.tokens_used).label('total_tokens')
).filter(
and_(
func.date(TokenUsage.usage_date) >= start_date,
func.date(TokenUsage.usage_date) <= end_date
)
).group_by(
func.date(TokenUsage.usage_date)
).order_by(
func.date(TokenUsage.usage_date)
).all()
# 转换为字典列表
usage_dict = {str(row.usage_date): row.total_tokens for row in results}
# 填充缺失的日期
daily_usages = []
current_date = start_date
while current_date <= end_date:
date_str = str(current_date)
tokens = usage_dict.get(date_str, 0)
daily_usages.append({
"date": date_str,
"tokens": tokens
})
current_date += timedelta(days=1)
return daily_usages
def get_monthly_usages(self, months: int = 12) -> List[Dict[str, Any]]:
"""
获取每月 Token 使用(用于图表)
:param months: 月数
:return: 每月使用列表
"""
now = datetime.now()
results = []
for i in range(months):
# 计算月份
month_offset = months - 1 - i
target_date = now - timedelta(days=30 * month_offset)
year = target_date.year
month = target_date.month
# 查询该月的使用量
usage = self.get_month_usage(year, month)
results.append({
"month": f"{year}-{month:02d}",
"tokens": usage
})
return results
def get_user_token_usage(
self,
user_id: int,
days: int = 30
) -> List[Dict[str, Any]]:
"""
获取指定用户的 Token 使用
:param user_id: 用户 ID
:param days: 天数
:return: 每日使用列表
"""
end_date = date.today()
start_date = end_date - timedelta(days=days - 1)
results = self.db.query(
func.date(TokenUsage.usage_date).label('usage_date'),
func.sum(TokenUsage.tokens_used).label('total_tokens')
).filter(
and_(
TokenUsage.virtual_user_id == user_id,
func.date(TokenUsage.usage_date) >= start_date,
func.date(TokenUsage.usage_date) <= end_date
)
).group_by(
func.date(TokenUsage.usage_date)
).order_by(
func.date(TokenUsage.usage_date)
).all()
return [
{"date": str(row.usage_date), "tokens": row.total_tokens}
for row in results
]
def check_token_limit_exceeded(self) -> bool:
"""检查是否超出 Token 限额"""
today_used = self.get_today_usage()
return today_used >= settings.MAX_TOKENS_PER_DAY
# 工厂函数
def get_token_service(db: Session) -> TokenService:
"""获取 Token 服务实例"""
return TokenService(db)
+358
View File
@@ -0,0 +1,358 @@
"""虚拟用户业务服务"""
import io
import uuid
from datetime import datetime, timezone
def _fmt_dt(dt):
if dt is None: return None
if dt.tzinfo is None: dt = dt.replace(tzinfo=timezone.utc)
return dt.isoformat()
from typing import List, Optional, Tuple
import pandas as pd
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, update, delete, func, and_, or_
from fastapi import HTTPException
from app.models import VirtualUser, UserPersonality
from app.schemas import UserCreateRequest, UserUpdateRequest
from app.utils.crypto import encrypt, decrypt
from app.services.ai_service import ai_service
from app.core.logger import logger
STATUS_LABELS = {0: "未登录", 1: "登录中", 2: "已登录", 3: "登录失效", 4: "封禁"}
ACTIVITY_LABELS = {0: "低", 1: "中", 2: "高"}
ACTIVITY_COMMENT_LIMITS = {0: (3, 5), 1: (8, 15), 2: (20, 30)}
class UserService:
async def get_users(
self, db: AsyncSession,
page: int = 1, page_size: int = 20,
keyword: str = None, status: int = None,
is_enabled: int = None
) -> Tuple[int, List[dict]]:
query = select(VirtualUser)
conditions = []
if keyword:
conditions.append(
or_(VirtualUser.nickname.like(f"%{keyword}%"),
VirtualUser.account.like(f"%{keyword}%"))
)
if status is not None:
conditions.append(VirtualUser.status == status)
if is_enabled is not None:
conditions.append(VirtualUser.is_enabled == is_enabled)
if conditions:
query = query.where(and_(*conditions))
count_result = await db.execute(
select(func.count()).select_from(query.subquery())
)
total = count_result.scalar()
query = query.offset((page - 1) * page_size).limit(page_size).order_by(VirtualUser.created_at.desc())
result = await db.execute(query)
users = result.scalars().all()
items = []
for u in users:
# 获取人格
p_result = await db.execute(select(UserPersonality).where(UserPersonality.user_id == u.id))
personality = p_result.scalar_one_or_none()
items.append(self._format_user(u, personality))
return total, items
async def create_user(self, db: AsyncSession, req: UserCreateRequest) -> dict:
# 检查账号重复
existing = await db.execute(select(VirtualUser).where(VirtualUser.account == req.account))
if existing.scalar_one_or_none():
raise HTTPException(status_code=400, detail="账号已存在")
# 昵称选填:为空则自动生成
nickname = req.nickname or f"用户{req.account[-4:]}"
# 检查昵称重复(自动生成的若冲突则加随机后缀)
existing_nick = await db.execute(select(VirtualUser).where(VirtualUser.nickname == nickname))
if existing_nick.scalar_one_or_none():
import random, string
nickname = nickname + "_" + "".join(random.choices(string.digits, k=4))
user = VirtualUser(
nickname=nickname,
account=req.account,
password_enc=encrypt(req.password),
avatar_url=req.avatar_url,
activity_level=req.activity_level,
daily_comment_limit=req.daily_comment_limit,
daily_like_limit=req.daily_like_limit,
remark=req.remark,
status=0,
is_enabled=1,
)
db.add(user)
await db.flush()
# 自动生成AI人格
try:
await self._generate_personality(db, user)
except Exception as e:
logger.warning(f"人格生成失败,跳过: {e}")
await db.commit()
await db.refresh(user)
p_result = await db.execute(select(UserPersonality).where(UserPersonality.user_id == user.id))
personality = p_result.scalar_one_or_none()
return self._format_user(user, personality)
async def update_user(self, db: AsyncSession, user_id: int, req: UserUpdateRequest) -> dict:
user = await self._get_or_404(db, user_id)
if req.nickname and req.nickname != user.nickname:
existing = await db.execute(
select(VirtualUser).where(VirtualUser.nickname == req.nickname, VirtualUser.id != user_id)
)
if existing.scalar_one_or_none():
raise HTTPException(status_code=400, detail="昵称已被使用")
user.nickname = req.nickname
if req.password:
user.password_enc = encrypt(req.password)
if req.avatar_url is not None:
user.avatar_url = req.avatar_url
if req.activity_level is not None:
user.activity_level = req.activity_level
if req.daily_comment_limit is not None:
user.daily_comment_limit = req.daily_comment_limit
if req.daily_like_limit is not None:
user.daily_like_limit = req.daily_like_limit
if req.remark is not None:
user.remark = req.remark
if req.is_enabled is not None:
user.is_enabled = req.is_enabled
if req.is_enabled == 0:
user.status = 0 # 禁用后重置状态
await db.commit()
await db.refresh(user)
p_result = await db.execute(select(UserPersonality).where(UserPersonality.user_id == user.id))
personality = p_result.scalar_one_or_none()
return self._format_user(user, personality)
async def delete_user(self, db: AsyncSession, user_id: int):
user = await self._get_or_404(db, user_id)
await db.execute(delete(UserPersonality).where(UserPersonality.user_id == user_id))
await db.delete(user)
await db.commit()
async def batch_action(self, db: AsyncSession, user_ids: List[int], action: str):
"""批量操作"""
if action == "enable":
await db.execute(update(VirtualUser).where(VirtualUser.id.in_(user_ids)).values(is_enabled=1))
elif action == "disable":
await db.execute(update(VirtualUser).where(VirtualUser.id.in_(user_ids)).values(is_enabled=0, status=0))
elif action == "logout":
await db.execute(update(VirtualUser).where(VirtualUser.id.in_(user_ids)).values(status=0, session_token=None))
elif action == "delete":
await db.execute(delete(UserPersonality).where(UserPersonality.user_id.in_(user_ids)))
await db.execute(delete(VirtualUser).where(VirtualUser.id.in_(user_ids)))
await db.commit()
return {"affected": len(user_ids)}
async def generate_personality(self, db: AsyncSession, user_id: int) -> dict:
"""为用户生成/重新生成AI人格"""
user = await self._get_or_404(db, user_id)
# 删除旧人格
await db.execute(delete(UserPersonality).where(UserPersonality.user_id == user_id))
personality = await self._generate_personality(db, user)
await db.commit()
return self._format_personality(personality)
async def update_personality(self, db: AsyncSession, user_id: int, req) -> dict:
p_result = await db.execute(select(UserPersonality).where(UserPersonality.user_id == user_id))
personality = p_result.scalar_one_or_none()
if not personality:
raise HTTPException(status_code=404, detail="人格不存在")
for field, val in req.model_dump(exclude_none=True).items():
setattr(personality, field, val)
# 重新生成提示词
personality.comment_style_prompt = self._build_style_prompt(personality)
await db.commit()
await db.refresh(personality)
return self._format_personality(personality)
async def import_from_excel(self, db: AsyncSession, file_content: bytes) -> dict:
"""Excel批量导入 - 每行独立事务,互不影响"""
try:
df = pd.read_excel(io.BytesIO(file_content), engine='openpyxl')
except Exception:
df = pd.read_excel(io.BytesIO(file_content))
df.columns = [str(c).strip() for c in df.columns]
required_cols = {"新闻平台账号", "登录密码", "昵称"}
if not required_cols.issubset(set(df.columns)):
raise HTTPException(
status_code=400,
detail=f"缺少必填列: {required_cols - set(df.columns)},当前列: {list(df.columns)}"
)
success_count = 0
error_list = []
for idx, row in df.iterrows():
row_num = idx + 2
row_account = ""
try:
# 账号可能是数字类型(手机号),统一转为字符串
account = str(row.get("新闻平台账号", "") or "").strip().split(".")[0] # 去掉 .0 后缀
password = str(row.get("登录密码", "") or "").strip()
nickname = str(row.get("昵称", "") or "").strip()
row_account = account
if account.lower() in ("nan", "none", ""):
error_list.append({"row": row_num, "error": "账号为空"}); continue
if password.lower() in ("nan", "none", ""):
error_list.append({"row": row_num, "account": account, "error": "密码为空"}); continue
if len(password) < 6:
error_list.append({"row": row_num, "account": account, "error": "密码不足6位"}); continue
# 昵称选填:为空时自动用账号末4位生成
if nickname.lower() in ("nan", "none", ""):
nickname = f"用户{account[-4:]}"
existing = await db.execute(select(VirtualUser).where(VirtualUser.account == account))
if existing.scalar_one_or_none():
error_list.append({"row": row_num, "account": account, "error": "账号已存在"}); continue
existing_nick = await db.execute(select(VirtualUser).where(VirtualUser.nickname == nickname))
if existing_nick.scalar_one_or_none():
error_list.append({"row": row_num, "account": account, "error": "昵称已被使用"}); continue
avatar = str(row.get("头像链接", "") or "").strip()
remark = str(row.get("备注", "") or "").strip()
user = VirtualUser(
nickname=nickname, account=account,
password_enc=encrypt(password),
avatar_url=avatar if avatar.lower() not in ("nan","none","") else None,
remark=remark if remark.lower() not in ("nan","none","") else None,
status=0, is_enabled=1, activity_level=1,
)
db.add(user)
await db.flush()
await db.commit()
try:
await db.refresh(user)
await self._generate_personality(db, user)
await db.commit()
except Exception as pe:
logger.warning(f"第{row_num}行人格生成跳过: {pe}")
await db.rollback()
success_count += 1
except Exception as e:
await db.rollback()
error_list.append({"row": row_num, "account": row_account, "error": str(e)})
logger.warning(f"导入第{row_num}行失败: {e}")
return {"success": success_count, "failed": len(error_list), "errors": error_list}
async def export_to_excel(self, db: AsyncSession) -> bytes:
"""导出全量用户数据(不含密码)"""
result = await db.execute(select(VirtualUser).order_by(VirtualUser.created_at.desc()))
users = result.scalars().all()
rows = []
for u in users:
p_result = await db.execute(select(UserPersonality).where(UserPersonality.user_id == u.id))
p = p_result.scalar_one_or_none()
rows.append({
"ID": u.id, "昵称": u.nickname, "账号": u.account,
"状态": STATUS_LABELS.get(u.status, "未知"),
"活跃度": ACTIVITY_LABELS.get(u.activity_level, "中"),
"性格": p.character_type if p else "", "语言风格": p.language_style if p else "",
"兴趣偏好": ",".join(p.interest_tags or []) if p else "",
"互动倾向": p.interact_tendency if p else "",
"累计互动": u.total_interactions, "今日评论": u.today_comment_count,
"今日点赞": u.today_like_count, "最后登录": u.last_login_at,
"最后互动": u.last_interact_at, "备注": u.remark,
"是否启用": "是" if u.is_enabled else "否", "创建时间": u.created_at,
})
df = pd.DataFrame(rows)
buf = io.BytesIO()
df.to_excel(buf, index=False, sheet_name="虚拟用户")
buf.seek(0)
return buf.read()
async def get_excel_template(self) -> bytes:
"""获取导入模板(账号+密码必填,其他选填)"""
df = pd.DataFrame(columns=["新闻平台账号", "登录密码", "昵称(选填)", "头像链接(选填)", "备注(选填)"])
df.loc[0] = ["13800138000", "password123", "(留空自动生成)", "", ""]
buf = io.BytesIO()
df.to_excel(buf, index=False, sheet_name="导入模板")
buf.seek(0)
return buf.read()
async def _generate_personality(self, db: AsyncSession, user: VirtualUser) -> UserPersonality:
"""调用AI生成人格"""
result = await ai_service.generate_personality(user.nickname, user.account)
personality = UserPersonality(
user_id=user.id,
character_type=result.get("character_type", "温和"),
language_style=result.get("language_style", "幽默"),
interest_tags=result.get("interest_tags", ["科技"]),
interact_tendency=result.get("interact_tendency", "爱评论"),
word_count_min=result.get("word_count_min", 20),
word_count_max=result.get("word_count_max", 80),
personality_desc=result.get("personality_desc", ""),
)
personality.comment_style_prompt = self._build_style_prompt(personality)
db.add(personality)
await db.flush()
return personality
def _build_style_prompt(self, p: UserPersonality) -> str:
interests = "、".join(p.interest_tags or []) if p.interest_tags else "综合"
return (
f"你是一个{p.character_type}性格、{p.language_style}语言风格的新闻读者,"
f"主要对{interests}类内容感兴趣,互动倾向是{p.interact_tendency}。"
f"评论字数控制在{p.word_count_min}~{p.word_count_max}字。"
f"个人简介:{p.personality_desc}"
)
def _format_user(self, u: VirtualUser, p: Optional[UserPersonality]) -> dict:
return {
"id": u.id, "nickname": u.nickname, "account": u.account,
"avatar_url": u.avatar_url,
"real_name": getattr(u, "real_name", None),
"sex": getattr(u, "sex", 0),
"platform_uid": getattr(u, "platform_uid", None),
"status": u.status,
"status_label": STATUS_LABELS.get(u.status, "未知"),
"activity_level": u.activity_level,
"activity_label": ACTIVITY_LABELS.get(u.activity_level, "中"),
"daily_comment_limit": u.daily_comment_limit,
"daily_like_limit": u.daily_like_limit,
"today_comment_count": u.today_comment_count,
"today_like_count": u.today_like_count,
"total_interactions": u.total_interactions,
"last_login_at": _fmt_dt(u.last_login_at),
"last_interact_at": _fmt_dt(u.last_interact_at),
"remark": u.remark, "is_enabled": u.is_enabled,
"created_at": _fmt_dt(u.created_at),
"personality": self._format_personality(p) if p else None,
}
def _format_personality(self, p: Optional[UserPersonality]) -> Optional[dict]:
if not p:
return None
return {
"id": p.id, "user_id": p.user_id,
"character_type": p.character_type, "language_style": p.language_style,
"interest_tags": p.interest_tags or [], "interact_tendency": p.interact_tendency,
"word_count_min": p.word_count_min, "word_count_max": p.word_count_max,
"personality_desc": p.personality_desc,
"updated_at": _fmt_dt(p.updated_at),
}
async def _get_or_404(self, db: AsyncSession, user_id: int) -> VirtualUser:
result = await db.execute(select(VirtualUser).where(VirtualUser.id == user_id))
user = result.scalar_one_or_none()
if not user:
raise HTTPException(status_code=404, detail="用户不存在")
return user
user_service = UserService()
@@ -1,361 +0,0 @@
"""
虚拟用户管理服务
"""
import logging
from typing import Optional, List, Dict, Any
from datetime import datetime, timedelta
from sqlalchemy.orm import Session
from sqlalchemy import and_, func, Date
from app.models.virtual_user import VirtualUser, VirtualUserPersona, ActivityLevel, UserStatus
from app.models.interaction import InteractionRecord, InteractionType
from app.services.ai_service import ai_service
from app.core.config import settings
logger = logging.getLogger(__name__)
class VirtualUserService:
"""虚拟用户服务类"""
# 预设写作风格库
WRITING_STYLES = [
"幽默风趣",
"严肃理性",
"文艺清新",
"吐槽犀利",
"感性温暖",
"客观中立",
"激情澎湃",
"冷静分析",
"活泼可爱",
"深沉内敛"
]
# 昵称前缀和后缀
NICKNAME_PREFIXES = ["清风", "星辰", "云端", "晨曦", "暮色", "流年", "初心", "远方"]
NICKNAME_SUFFIXES = ["行者", "旅人", "追梦", "时光", "记忆", "印象", "故事", "传奇"]
def __init__(self, db: Session):
self.db = db
def get_user_by_id(self, user_id: int) -> Optional[VirtualUser]:
"""根据 ID 获取用户"""
return self.db.query(VirtualUser).filter(VirtualUser.id == user_id).first()
def get_user_by_username(self, username: str) -> Optional[VirtualUser]:
"""根据用户名获取用户"""
return self.db.query(VirtualUser).filter(VirtualUser.username == username).first()
def get_users(
self,
page: int = 1,
page_size: int = 20,
status: Optional[UserStatus] = None,
search: Optional[str] = None
) -> Dict[str, Any]:
"""
获取用户列表
:param page: 页码
:param page_size: 每页数量
:param status: 状态筛选
:param search: 搜索关键词
:return: 用户列表和总数
"""
query = self.db.query(VirtualUser)
if status:
query = query.filter(VirtualUser.status == status)
if search:
query = query.filter(
or_(
VirtualUser.nickname.like(f"%{search}%"),
VirtualUser.username.like(f"%{search}%")
)
)
total = query.count()
users = query.order_by(VirtualUser.created_at.desc()).offset(
(page - 1) * page_size
).limit(page_size).all()
return {"total": total, "items": users}
def create_user(
self,
username: str,
password: str,
nickname: str,
writing_style: Optional[str] = None,
activity_level: ActivityLevel = ActivityLevel.MEDIUM,
avatar_url: Optional[str] = None,
persona_description: Optional[str] = None
) -> Optional[VirtualUser]:
"""
创建虚拟用户
:param username: 用户名
:param password: 密码
:param nickname: 昵称
:param writing_style: 写作风格
:param activity_level: 活跃度
:param avatar_url: 头像 URL
:param persona_description: 人格描述
:return: 创建的用户
"""
# 检查用户名是否已存在
existing = self.get_user_by_username(username)
if existing:
logger.error(f"Username already exists: {username}")
return None
user = VirtualUser(
username=username,
password=password, # TODO: 加密存储
nickname=nickname,
writing_style=writing_style or self._random_writing_style(),
activity_level=activity_level,
avatar_url=avatar_url or self._generate_avatar_url(),
persona_description=persona_description
)
self.db.add(user)
self.db.commit()
self.db.refresh(user)
logger.info(f"Virtual user created: {username}")
return user
def generate_users(
self,
count: int,
writing_styles: Optional[List[str]] = None,
activity_levels: Optional[List[ActivityLevel]] = None,
generate_persona: bool = True
) -> List[VirtualUser]:
"""
批量生成虚拟用户
:param count: 生成数量
:param writing_styles: 写作风格列表
:param activity_levels: 活跃度级别列表
:param generate_persona: 是否生成 AI 人格描述
:return: 生成的用户列表
"""
import random
styles = writing_styles or self.WRITING_STYLES
levels = activity_levels or [ActivityLevel.LOW, ActivityLevel.MEDIUM, ActivityLevel.HIGH]
created_users = []
for i in range(count):
# 生成唯一用户名
timestamp = datetime.now().strftime("%Y%m%d%H%M%S")
username = f"user_{timestamp}_{i}"
# 随机生成昵称
prefix = random.choice(self.NICKNAME_PREFIXES)
suffix = random.choice(self.NICKNAME_SUFFIXES)
nickname = f"{prefix}{suffix}{random.randint(100, 999)}"
# 随机密码
password = f"pwd_{random.randint(100000, 999999)}"
# 随机写作风格
writing_style = random.choice(styles)
# 随机活跃度
activity_level = random.choice(levels)
# 生成头像
avatar_url = self._generate_avatar_url()
# AI 生成人格描述
persona_description = None
if generate_persona:
persona_description = self._generate_persona_description(
writing_style,
activity_level
)
user = self.create_user(
username=username,
password=password,
nickname=nickname,
writing_style=writing_style,
activity_level=activity_level,
avatar_url=avatar_url,
persona_description=persona_description
)
if user:
created_users.append(user)
logger.info(f"Generated {len(created_users)} virtual users")
return created_users
def _generate_persona_description(
self,
writing_style: str,
activity_level: ActivityLevel
) -> str:
"""AI 生成人格描述"""
import asyncio
prompt = f"""请为一位虚拟用户生成人格描述,要求:
- 写作风格:{writing_style}
- 活跃度:{activity_level.value}
请用 50-100 字描述这个人的性格特点、兴趣爱好、说话方式等。直接输出描述内容。"""
try:
# 同步调用异步方法
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
result = loop.run_until_complete(ai_service._call_ai_api(prompt))
loop.close()
if result and result.get("content"):
return result["content"]
except Exception as e:
logger.error(f"Generate persona description error: {e}")
return f"这是一位{writing_style}的虚拟用户,活跃度{activity_level.value}。"
def _random_writing_style(self) -> str:
"""随机选择写作风格"""
import random
return random.choice(self.WRITING_STYLES)
def _generate_avatar_url(self) -> str:
"""生成随机头像 URL(使用第三方头像 API)"""
import random
# 使用 DiceBear 头像 API
seed = f"avatar_{datetime.now().timestamp()}_{random.randint(1000, 9999)}"
return f"https://api.dicebear.com/7.x/avataaars/svg?seed={seed}"
def update_user(
self,
user_id: int,
**kwargs
) -> Optional[VirtualUser]:
"""更新用户信息"""
user = self.get_user_by_id(user_id)
if not user:
return None
for key, value in kwargs.items():
if hasattr(user, key) and value is not None:
setattr(user, key, value)
self.db.commit()
self.db.refresh(user)
return user
def delete_user(self, user_id: int) -> bool:
"""删除用户"""
user = self.get_user_by_id(user_id)
if not user:
return False
self.db.delete(user)
self.db.commit()
logger.info(f"Virtual user deleted: {user_id}")
return True
def import_users_from_excel(
self,
users_data: List[Dict[str, Any]],
generate_persona: bool = True
) -> Dict[str, Any]:
"""
从 Excel 导入虚拟用户
:param users_data: 用户数据列表
:param generate_persona: 是否生成 AI 人格描述
:return: 导入结果
"""
success_count = 0
failed_count = 0
created_users = []
for user_data in users_data:
try:
username = user_data.get("username")
password = user_data.get("password")
nickname = user_data.get("nickname", "")
if not username or not password:
logger.warning(f"Missing username or password: {user_data}")
failed_count += 1
continue
# 如果昵称为空,生成一个
if not nickname:
nickname = f"用户{username}"
writing_style = user_data.get("writing_style")
activity_level_str = user_data.get("activity_level", "medium")
# 转换活跃度枚举
try:
activity_level = ActivityLevel(activity_level_str.lower())
except ValueError:
activity_level = ActivityLevel.MEDIUM
user = self.create_user(
username=username,
password=password,
nickname=nickname,
writing_style=writing_style,
activity_level=activity_level
)
if user:
success_count += 1
created_users.append(user)
else:
failed_count += 1
except Exception as e:
logger.error(f"Import user error: {e}")
failed_count += 1
return {
"success_count": success_count,
"failed_count": failed_count,
"created_users": created_users
}
def get_user_stats(self, user_id: int) -> Dict[str, Any]:
"""获取用户统计信息"""
user = self.get_user_by_id(user_id)
if not user:
return {}
today = datetime.now().date()
# 统计今日互动
today_interactions = self.db.query(InteractionRecord).filter(
and_(
InteractionRecord.virtual_user_id == user_id,
func.date(InteractionRecord.execution_time) == today
)
).all()
today_comments = sum(1 for i in today_interactions if i.interaction_type == InteractionType.COMMENT)
today_replies = sum(1 for i in today_interactions if i.interaction_type == InteractionType.REPLY)
return {
"user_id": user_id,
"nickname": user.nickname,
"total_interactions": user.total_interactions,
"today_comments": today_comments,
"today_replies": today_replies,
"last_interaction_time": user.last_interaction_time
}
# 工厂函数
def get_virtual_user_service(db: Session) -> VirtualUserService:
"""获取虚拟用户服务实例"""
return VirtualUserService(db)
View File
+49
View File
@@ -0,0 +1,49 @@
"""AES加密工具 - 用于密码和API Key加密存储"""
import base64
import hashlib
from Crypto.Cipher import AES
from Crypto.Util.Padding import pad, unpad
from app.core.config import settings
def _get_key() -> bytes:
"""获取32字节AES密钥"""
key = settings.AES_KEY.encode("utf-8")
return hashlib.sha256(key).digest()
def encrypt(plaintext: str) -> str:
"""AES-CBC加密"""
if not plaintext:
return ""
key = _get_key()
cipher = AES.new(key, AES.MODE_CBC)
ct_bytes = cipher.encrypt(pad(plaintext.encode("utf-8"), AES.block_size))
iv = base64.b64encode(cipher.iv).decode("utf-8")
ct = base64.b64encode(ct_bytes).decode("utf-8")
return f"{iv}:{ct}"
def decrypt(ciphertext: str) -> str:
"""AES-CBC解密"""
if not ciphertext or ":" not in ciphertext:
return ""
try:
iv_str, ct_str = ciphertext.split(":", 1)
key = _get_key()
iv = base64.b64decode(iv_str)
ct = base64.b64decode(ct_str)
cipher = AES.new(key, AES.MODE_CBC, iv)
pt = unpad(cipher.decrypt(ct), AES.block_size)
return pt.decode("utf-8")
except Exception:
return ""
def mask_password(password: str) -> str:
"""密码脱敏显示"""
if not password:
return ""
if len(password) <= 2:
return "*" * len(password)
return password[0] + "*" * (len(password) - 2) + password[-1]
Regular → Executable
+22 -33
View File
@@ -1,35 +1,24 @@
# Web Framework
fastapi==0.109.0
uvicorn[standard]==0.27.0
python-multipart==0.0.6
# Database
sqlalchemy==2.0.25
alembic==1.13.1
pymysql==1.1.0
# AI Models
openai==1.10.0
zhipuai==2.0.1
# Utilities
pydantic==2.5.3
pydantic-settings==2.1.0
python-dotenv==1.0.0
httpx==0.26.0
fastapi==0.115.6
uvicorn[standard]==0.34.0
sqlalchemy==2.0.36
pymysql==1.1.1
cryptography==44.0.0
redis==5.2.1
apscheduler==3.10.4
# Excel Support
openpyxl==3.1.2
pandas==2.1.4
# Security
python-jose[cryptography]==3.3.0
pandas==2.2.3
openpyxl==3.1.5
passlib[bcrypt]==1.7.4
# Logging
loguru==0.7.2
# Testing
pytest==7.4.4
pytest-asyncio==0.23.3
pycryptodome==3.21.0
httpx==0.28.1
python-multipart==0.0.20
python-jose[cryptography]==3.3.0
pydantic==2.10.4
pydantic-settings==2.7.0
openai==1.59.6
langchain==0.3.13
langchain-openai==0.3.0
aiofiles==24.1.0
loguru==0.7.3
alembic==1.14.0
aiomysql==0.2.0
greenlet==3.1.1
@@ -0,0 +1,126 @@
import json
import os
import sqlite3
import tempfile
import unittest
from types import SimpleNamespace
from unittest.mock import patch
from app.services import avatar_service
class AvatarSquareAuthorizationTests(unittest.TestCase):
def setUp(self):
fd, self.db_path = tempfile.mkstemp(suffix=".db")
os.close(fd)
connection = sqlite3.connect(self.db_path)
connection.executescript("""
CREATE TABLE users (
huihui_user_id TEXT,
nickname TEXT,
avatar_url TEXT,
huihui_token TEXT
);
CREATE TABLE avatars (
id TEXT,
owner_id TEXT,
name TEXT,
display_name TEXT,
description TEXT,
photo_url TEXT,
config TEXT,
status TEXT
);
""")
connection.execute(
"INSERT INTO users VALUES (?, ?, ?, ?)",
("huihui-7", "主人", "/owner.jpg", "huihui-token"),
)
connection.commit()
connection.close()
avatar_service._engine = None
avatar_service._SessionLocal = None
self.path_patch = patch.object(avatar_service.settings, "AVATAR_DB_PATH", self.db_path)
self.path_patch.start()
def tearDown(self):
self.path_patch.stop()
if avatar_service._engine is not None:
avatar_service._engine.dispose()
avatar_service._engine = None
avatar_service._SessionLocal = None
os.unlink(self.db_path)
def _insert_avatar(self, permissions, *, status="active", token=None):
connection = sqlite3.connect(self.db_path)
connection.execute(
"INSERT INTO avatars VALUES (?, ?, ?, ?, ?, ?, ?, ?)",
(
"avatar-7",
"huihui-7",
"avatar",
"小会",
"语气友好,表达简洁",
"/avatar.jpg",
json.dumps({
"authorizationPermissions": permissions,
"replyStyle": "warm",
"responseLength": "short",
}),
status,
),
)
if token is not None:
connection.execute(
"UPDATE users SET huihui_token = ? WHERE huihui_user_id = ?",
(token, "huihui-7"),
)
connection.commit()
connection.close()
def test_interact_permission_exposes_only_requested_square_actions(self):
self._insert_avatar(["chat", "interact"])
permissions = avatar_service.get_square_interaction_permissions("avatar-7")
self.assertEqual(
permissions,
frozenset({"like", "collect", "comment", "reply"}),
)
self.assertNotIn("forward", permissions)
def test_missing_permission_inactive_avatar_or_missing_token_denies_execution(self):
scenarios = [
(["chat"], "active", "huihui-token"),
(["interact"], "inactive", "huihui-token"),
(["interact"], "active", ""),
]
for permissions, status, token in scenarios:
with self.subTest(permissions=permissions, status=status, token=token):
connection = sqlite3.connect(self.db_path)
connection.execute("DELETE FROM avatars")
connection.commit()
connection.close()
self._insert_avatar(permissions, status=status, token=token)
self.assertEqual(
avatar_service.get_square_interaction_permissions("avatar-7"),
frozenset(),
)
def test_delegated_avatar_identity_is_recognized_without_matching_normal_users(self):
delegated = SimpleNamespace(account="__avatar__:avatar-7")
normal = SimpleNamespace(account="13800000000")
self.assertTrue(avatar_service.is_delegated_avatar_user(delegated))
self.assertEqual(avatar_service.delegated_avatar_id(delegated), "avatar-7")
self.assertFalse(avatar_service.is_delegated_avatar_user(normal))
self.assertEqual(avatar_service.delegated_avatar_id(normal), "")
def test_response_length_maps_to_scheduler_comment_limits(self):
self.assertEqual(avatar_service._word_count_range({"responseLength": "short"}), (10, 35))
self.assertEqual(avatar_service._word_count_range({"responseLength": "long"}), (30, 80))
self.assertEqual(avatar_service._word_count_range({"responseLength": "unknown"}), (20, 60))
if __name__ == "__main__":
unittest.main()
+143
View File
@@ -0,0 +1,143 @@
import unittest
from types import SimpleNamespace
from unittest.mock import AsyncMock, patch
from app.services.news_service import NewsPlatformService
class _Response:
def __init__(self, payload, status_code=200):
self._payload = payload
self.status_code = status_code
self.text = str(payload)
def json(self):
return self._payload
class _Client:
responses = []
calls = []
def __init__(self, *args, **kwargs):
pass
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, exc, tb):
return False
async def patch(self, url, **kwargs):
self.__class__.calls.append(("PATCH", url, kwargs))
return self.__class__.responses.pop(0)
async def get(self, url, **kwargs):
self.__class__.calls.append(("GET", url, kwargs))
return self.__class__.responses.pop(0)
class HuihuiProfileSyncTests(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self):
self.service = NewsPlatformService()
self.service._auth_url = AsyncMock(return_value="https://99hui.com/api/usercenter")
self.service._biz_url = AsyncMock(return_value="https://99hui.com/api/huihuibusiness")
self.service._client = AsyncMock(return_value={
"appId": "app", "accessId": "access", "accessSecret": "secret",
"clientCode": "", "orgId": "",
})
self.db = SimpleNamespace(execute=AsyncMock(), commit=AsyncMock())
self.user = SimpleNamespace(
id=51, account="13721560046", platform_uid="platform-51",
nickname="黎佳怡", real_name="黎佳怡", sex=0, avatar_url="https://img/avatar.jpg",
)
_Client.calls = []
async def test_updates_all_profiles_and_verifies_app_home(self):
_Client.responses = [
_Response({"code": 0, "data": True}),
_Response({"code": 0, "data": True}),
_Response({"code": 0, "data": {
"id": "extend-51", "userId": "platform-51", "name": "", "avatar": "",
}}),
_Response({"code": 0, "data": None}),
_Response({"code": 0, "data": {
"id": "extend-51", "userId": "platform-51", "name": "黎佳怡",
"avatar": "https://img/avatar.jpg",
}}),
]
session = {"token": "token", "platform_uid": "platform-51", "org_id": "org-1"}
with patch("app.services.news_service.get_session", AsyncMock(return_value=session)), \
patch("app.services.news_service.httpx.AsyncClient", _Client):
ok, err = await self.service.update_user_profile(
self.db, self.user,
nick_name="黎佳怡", real_name="黎佳怡",
avatar="https://img/avatar.jpg",
)
self.assertTrue(ok, err)
self.assertEqual(_Client.calls[0][1], "https://99hui.com/api/usercenter/v2/users/current")
self.assertEqual(_Client.calls[1][1], "https://99hui.com/api/usercenter/users/page/platform-51")
public_params = _Client.calls[1][2]["params"]
self.assertEqual(public_params["nickName"], "黎佳怡")
self.assertEqual(public_params["icon"], "https://img/avatar.jpg")
self.assertEqual(_Client.calls[2][0], "GET")
self.assertEqual(
_Client.calls[2][1],
"https://99hui.com/api/huihuiuserextend/user/home/platform-51",
)
self.assertEqual(_Client.calls[3][0], "PATCH")
self.assertEqual(
_Client.calls[3][1],
"https://99hui.com/api/huihuiuserextend/user",
)
self.assertEqual(_Client.calls[3][2]["json"]["id"], "extend-51")
self.assertEqual(_Client.calls[3][2]["json"]["name"], "黎佳怡")
self.assertEqual(_Client.calls[3][2]["json"]["avatar"], "https://img/avatar.jpg")
self.assertEqual(_Client.calls[4][0], "GET")
self.db.commit.assert_awaited_once()
async def test_public_page_failure_is_not_reported_as_success(self):
_Client.responses = [
_Response({"code": 0, "data": True}),
_Response({"code": 500, "message": "page update failed"}),
]
session = {"token": "token", "platform_uid": "platform-51"}
with patch("app.services.news_service.get_session", AsyncMock(return_value=session)), \
patch("app.services.news_service.httpx.AsyncClient", _Client):
ok, err = await self.service.update_user_profile(
self.db, self.user,
nick_name="黎佳怡", avatar="https://img/avatar.jpg",
)
self.assertFalse(ok)
self.assertIn("TA的主页同步失败", err)
self.db.commit.assert_not_awaited()
async def test_app_home_mismatch_is_not_reported_as_success(self):
_Client.responses = [
_Response({"code": 0, "data": True}),
_Response({"code": 0, "data": True}),
_Response({"code": 0, "data": {
"id": "extend-51", "userId": "platform-51", "name": "", "avatar": "",
}}),
_Response({"code": 0, "data": None}),
_Response({"code": 0, "data": {
"id": "extend-51", "userId": "platform-51", "name": "", "avatar": "",
}}),
]
session = {"token": "token", "platform_uid": "platform-51", "org_id": "org-1"}
with patch("app.services.news_service.get_session", AsyncMock(return_value=session)), \
patch("app.services.news_service.httpx.AsyncClient", _Client):
ok, err = await self.service.update_user_profile(
self.db, self.user,
nick_name="黎佳怡", avatar="https://img/avatar.jpg",
)
self.assertFalse(ok)
self.assertIn("App用户主页同步失败", err)
self.db.commit.assert_not_awaited()
if __name__ == "__main__":
unittest.main()
+6
View File
@@ -0,0 +1,6 @@
node_modules/
dist/
.git/
.env
*.log
backend/
+6
View File
@@ -0,0 +1,6 @@
node_modules
dist
.env
*.log
backend/avatar.db
backend/routers/uploads/
+21
View File
@@ -0,0 +1,21 @@
# 构建阶段:安装依赖并打包 H5
FROM node:18-alpine AS build
WORKDIR /app
COPY package*.json ./
RUN npm ci
COPY . .
RUN npm run build
# 运行阶段:nginx 托管静态资源并反向代理 /api 到后端
# 锁定 1.28-alpine:测试服务器 Docker 的 seccomp 拦截 pwrite 系统调用,
# 新版 nginx(>=1.31) 用 pwrite 写 pid 文件会被拦导致致命退出;1.28 用 write() 可正常启动。
FROM nginx:1.28-alpine
COPY --from=build /app/dist /usr/share/nginx/html
# 覆盖 nginx 默认主配置(含唯一可写的 pid /tmp/nginx.pid,规避受限容器内 /run 不可写导致反复重启)
COPY nginx.conf /etc/nginx/nginx.conf
EXPOSE 80
+6
View File
@@ -0,0 +1,6 @@
__pycache__/
*.pyc
*.db
.env
logs/
*.log
+15
View File
@@ -0,0 +1,15 @@
FROM python:3.10-slim
WORKDIR /app
# 后端依赖(fastapi/uvicorn/sqlalchemy/pypdf/python-docx/openpyxl 等)均为纯 Python wheel,
# 无需 gcc 等编译链,故跳过 apt 安装以加快构建并减小镜像体积。
COPY requirements.txt .
RUN pip install --no-cache-dir --timeout 120 --retries 10 -i https://pypi.tuna.tsinghua.edu.cn/simple -r requirements.txt
COPY . .
# 后端使用 SQLite(avatar.db 落在 /app 内),平铺结构以 `uvicorn main:app` 启动
EXPOSE 8000
CMD ["uvicorn", "main:app", "--host", "0.0.0.0", "--port", "8000", "--workers", "1"]
+106
View File
@@ -0,0 +1,106 @@
import os
from sqlalchemy import create_engine, event
from sqlalchemy.orm import sessionmaker, declarative_base, Session
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
DB_FILE = os.path.join(BASE_DIR, "avatar.db")
DATABASE_URL = os.getenv("DATABASE_URL", f"sqlite:///{DB_FILE}")
IS_SQLITE = DATABASE_URL.startswith("sqlite:")
engine = create_engine(
DATABASE_URL,
connect_args={"check_same_thread": False, "timeout": 30} if IS_SQLITE else {},
)
if IS_SQLITE:
@event.listens_for(engine, "connect")
def _configure_sqlite_connection(dbapi_connection, _connection_record):
cursor = dbapi_connection.cursor()
try:
cursor.execute("PRAGMA synchronous=NORMAL")
cursor.execute("PRAGMA busy_timeout=30000")
finally:
cursor.close()
SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False)
Base = declarative_base()
def get_db():
db = SessionLocal()
try:
yield db
finally:
db.close()
def init_db():
import models
if IS_SQLITE:
with engine.connect() as conn:
conn.exec_driver_sql("PRAGMA journal_mode=WAL")
conn.commit()
Base.metadata.create_all(bind=engine)
# 轻量迁移:为已存在的表补充新列(SQLite 不支持自动 ALTER,逐列尝试)
_try_add_columns(
("qa_pairs", "enabled", "BOOLEAN DEFAULT 1"),
("knowledge_docs", "vectorized", "BOOLEAN DEFAULT 0"),
("knowledge_docs", "embedding_model", "VARCHAR DEFAULT ''"),
("knowledge_docs", "chunk_count", "INTEGER DEFAULT 0"),
("knowledge_docs", "vectorized_at", "TIMESTAMP"),
("knowledge_docs", "error_message", "VARCHAR DEFAULT ''"),
("knowledge_docs", "index_stage", "VARCHAR DEFAULT ''"),
("knowledge_docs", "index_progress", "INTEGER DEFAULT 0"),
("avatars", "owner_id", "VARCHAR DEFAULT ''"),
("authorizations", "takeover_enabled", "BOOLEAN DEFAULT 0"),
("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"),
("authorizations", "takeover_delay_seconds", "INTEGER DEFAULT 180"),
("avatars", "share_token", "VARCHAR DEFAULT NULL"),
("token_account", "user_id", "VARCHAR DEFAULT ''"),
("token_account", "total_granted", "BIGINT DEFAULT 0"),
("token_account", "total_consumed", "BIGINT DEFAULT 0"),
("token_account", "created_at", "TIMESTAMP"),
("token_account", "updated_at", "TIMESTAMP"),
("takeover_messages", "attachment_id", "VARCHAR DEFAULT NULL"),
)
_normalize_optional_unique_values()
_normalize_takeover_delays()
_create_token_indexes()
def _try_add_columns(*cols):
with engine.connect() as conn:
for table, col, ddl in cols:
try:
conn.exec_driver_sql(f"ALTER TABLE {table} ADD COLUMN {col} {ddl}")
conn.commit()
except Exception:
# 列已存在(或全新库由 create_all 建好)则忽略
pass
def _normalize_optional_unique_values():
with engine.begin() as conn:
conn.exec_driver_sql("UPDATE avatars SET share_token = NULL WHERE share_token = ''")
def _normalize_takeover_delays():
with engine.begin() as conn:
# The old 30-second column default was never wired into the scheduler.
conn.exec_driver_sql(
"UPDATE authorizations SET takeover_delay_seconds = 180 "
"WHERE takeover_delay_seconds IS NULL OR takeover_delay_seconds = 30"
)
def _create_token_indexes():
with engine.begin() as conn:
conn.exec_driver_sql(
"CREATE UNIQUE INDEX IF NOT EXISTS ux_token_account_user_id "
"ON token_account(user_id) WHERE user_id <> ''"
)
+160
View File
@@ -0,0 +1,160 @@
"""
向量化服务:对文档/查询文本生成向量。
优先级:
1. 若配置了环境变量 EMBEDDING_API_URL,则调用第三方「OpenAI 兼容」的 /embeddings 接口
(需配置 EMBEDDING_API_KEY、EMBEDDING_MODEL,默认 text-embedding-3-small)。
2. 否则使用本地「哈希 TF 嵌入」兜底,使向量检索在无外部依赖时也能端到端跑通,
且相似文本(共享词汇)会得到更高余弦相似度,便于演示召回效果。
"""
import os
import re
import math
import json
import hashlib
import urllib.request
EMBED_DIM = 256
MODEL = os.getenv("EMBEDDING_MODEL", "mock-hash-embed-v1")
def _embedding_endpoint(api_url):
"""Accept either an OpenAI-compatible base URL or its full endpoint."""
api_url = (api_url or "").strip().rstrip("/")
if not api_url or api_url.endswith("/embeddings"):
return api_url
return f"{api_url}/embeddings"
def _tokenize(text):
text = (text or "").lower()
# 英文/数字按词,CJK 逐字(中文无空格,需拆到字级才能命中子词)
tokens = re.findall(r"[a-z0-9]+", text)
tokens += re.findall(r"[一-鿿]", text)
return tokens
def _hash_embedding(texts, dim=EMBED_DIM):
vecs = []
for text in texts:
vec = [0.0] * dim
tokens = _tokenize(text)
if not tokens:
tokens = list(text or "")
for tok in tokens:
h = int(hashlib.md5(tok.encode("utf-8")).hexdigest(), 16)
vec[h % dim] += 1.0
norm = math.sqrt(sum(v * v for v in vec))
if norm > 0:
vec = [v / norm for v in vec]
vecs.append(vec)
return vecs
def embed(texts, on_progress=None):
"""返回 list[list[float]],与输入顺序一致。"""
if not texts:
return []
api_url = _embedding_endpoint(os.getenv("EMBEDDING_API_URL"))
if api_url:
api_key = os.getenv("EMBEDDING_API_KEY", "")
model = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")
try:
batch_size = max(1, int(os.getenv("EMBEDDING_BATCH_SIZE", "10")))
except ValueError:
batch_size = 10
embeddings = []
total = len(texts)
for start in range(0, len(texts), batch_size):
batch = texts[start:start + batch_size]
payload = json.dumps({"input": batch, "model": model}).encode("utf-8")
req = urllib.request.Request(
api_url,
data=payload,
headers={
"Content-Type": "application/json",
"Authorization": f"Bearer {api_key}" if api_key else "",
},
method="POST",
)
with urllib.request.urlopen(req, timeout=30) as resp:
data = json.loads(resp.read().decode("utf-8"))
items = data["data"]
if items and "index" in items[0]:
items = sorted(items, key=lambda x: x["index"])
if len(items) != len(batch):
raise ValueError("embedding response count does not match request")
embeddings.extend(item["embedding"] for item in items)
if on_progress:
on_progress(len(embeddings), total)
return embeddings
vectors = _hash_embedding(texts)
if on_progress:
on_progress(len(vectors), len(texts))
return vectors
def cosine(a, b):
dot = sum(x * y for x, y in zip(a, b))
na = math.sqrt(sum(x * x for x in a))
nb = math.sqrt(sum(y * y for y in b))
if na == 0 or nb == 0:
return 0.0
return dot / (na * nb)
def chunk_text(text, size=400, overlap=50):
text = (text or "").strip()
if not text:
return []
if len(text) <= size:
return [text]
chunks = []
start = 0
while start < len(text):
end = min(start + size, len(text))
chunks.append(text[start:end])
if end == len(text):
break
start = end - overlap
return chunks
def extract_text(path, ext):
"""抽取文档纯文本;未知格式拒绝,已知格式解析失败时保留占位文本。"""
if ext not in {".txt", ".md", ".docx", ".xlsx", ".pdf", ".doc"}:
raise ValueError(f"unsupported file extension: {ext}")
try:
if ext in {".txt", ".md"}:
with open(path, "r", encoding="utf-8", errors="replace") as f:
return f.read()
if ext == ".docx":
from docx import Document
doc = Document(path)
return "\n".join(p.text for p in doc.paragraphs)
if ext == ".xlsx":
import openpyxl
wb = openpyxl.load_workbook(path, data_only=True, read_only=True)
rows = []
for ws in wb.worksheets:
for row in ws.iter_rows(values_only=True):
cells = [str(c) for c in row if c is not None]
if cells:
rows.append(" ".join(cells))
return "\n".join(rows)
if ext == ".pdf":
try:
from pypdf import PdfReader
except ImportError:
from PyPDF2 import PdfReader
reader = PdfReader(path)
return "\n".join((p.extract_text() or "") for p in reader.pages)
if ext == ".doc":
with open(path, "rb") as f:
raw = f.read().decode("utf-8", errors="ignore")
return re.sub(r"[\x00-\x08\x0b\x0c\x0e-\x1f]+", " ", raw)
except Exception as e: # 解析失败时回退
print("extract_text failed:", e)
return f"文档:{os.path.basename(path)} 类型 {ext}"
+266
View File
@@ -0,0 +1,266 @@
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
import os
import logging
from apscheduler.schedulers.asyncio import AsyncIOScheduler
from apscheduler.triggers.interval import IntervalTrigger
from database import init_db, SessionLocal
from models import Avatar, Authorization, Organization, TokenAccount, TokenPlan, User
from fastapi.staticfiles import StaticFiles
import routers.avatars
import routers.tokens
import routers.authorizations
import routers.organizations
import routers.knowledge
import routers.huihui_auth
import routers.chat
import routers.takeover
from responses import ok
from services.chat_attachment_service import purge_expired_chat_attachments
from services.knowledge_vectorizer import knowledge_vectorizer
from services.token_billing import DEFAULT_TOKEN_GRANT, release_stale_reservations
logger = logging.getLogger(__name__)
takeover_scheduler = None
maintenance_scheduler = None
app = FastAPI(title="会会数字分身 API", version="1.0.0")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=False,
allow_methods=["*"],
allow_headers=["*"],
)
app.include_router(routers.avatars.router, prefix="/api")
app.include_router(routers.tokens.router, prefix="/api")
app.include_router(routers.authorizations.router, prefix="/api")
app.include_router(routers.organizations.router, prefix="/api")
app.include_router(routers.knowledge.router, prefix="/api")
app.include_router(routers.huihui_auth.router, prefix="/api")
app.include_router(routers.chat.router, prefix="/api")
app.include_router(routers.takeover.router, prefix="/api")
UPLOAD_DIR = routers.knowledge.UPLOAD_DIR
os.makedirs(UPLOAD_DIR, exist_ok=True)
app.mount("/api/files", StaticFiles(directory=UPLOAD_DIR), name="knowledge-files")
@app.get("/api/health")
def health():
return ok({"status": "ok"})
def seed():
db = SessionLocal()
try:
plan_specs = [
{"id": "1", "name": "基础套餐", "amount": 2_000_000, "price": 10, "badge": "", "desc": "2M 积分"},
{"id": "2", "name": "标准套餐", "amount": 20_000_000, "price": 100, "badge": "常用", "desc": "20M 积分"},
{"id": "3", "name": "专业套餐", "amount": 250_000_000, "price": 1000, "badge": "加赠25%", "desc": "250M 积分"},
{"id": "4", "name": "企业套餐", "amount": 2_500_000_000, "price": 10000, "badge": "企业推荐", "desc": "2500M 积分"},
]
for spec in plan_specs:
plan = db.query(TokenPlan).filter(TokenPlan.id == spec["id"]).first()
if plan is None:
db.add(TokenPlan(**spec))
else:
for key, value in spec.items():
setattr(plan, key, value)
for user in db.query(User).all():
account = db.query(TokenAccount).filter(TokenAccount.user_id == user.id).first()
if account is None:
db.add(TokenAccount(
user_id=user.id,
balance=DEFAULT_TOKEN_GRANT,
total_granted=DEFAULT_TOKEN_GRANT,
total_consumed=0,
))
if db.query(Avatar).count() == 0:
avatar = Avatar(
name="我的数字分身",
display_name="会会助手",
description="我是您的AI数字分身,可以帮您管理日程、回复消息、处理任务。",
emoji="🤖",
status="active",
token_balance=0,
config={
"replyStyle": "professional",
"creativity": 50,
"rigor": 50,
"humor": 30,
"responseLength": "medium",
"systemPrompt": "",
},
)
db.add(avatar)
db.commit()
db.refresh(avatar)
if db.query(Authorization).count() == 0:
auths = [
Authorization(avatar_id=avatar.id, target_type="application", target_name="微信小程序", permissions=["read", "reply"], status="active"),
Authorization(avatar_id=avatar.id, target_type="user", target_name="张三", permissions=["read"], status="active"),
Authorization(avatar_id=avatar.id, target_type="organization", target_name="产品团队", permissions=["read", "edit"], status="inactive"),
]
db.add_all(auths)
if db.query(Organization).count() == 0:
orgs = [
Organization(name="会会增长团队", description="负责会会产品的增长与运营", emoji="🚀", org_type="team", member_count=12),
Organization(name="AI 实验室", description="探索前沿 AI 能力", emoji="💡", org_type="company", member_count=8),
]
db.add_all(orgs)
db.commit()
release_stale_reservations(db)
finally:
db.close()
@app.on_event("startup")
def on_startup():
global takeover_scheduler
init_db()
seed()
knowledge_vectorizer.start()
# Release stale resources when startup is invoked again by a reload/test.
stop_takeover_scheduler()
stop_maintenance_scheduler()
try:
start_maintenance_scheduler()
except Exception as exc:
stop_maintenance_scheduler()
logger.warning(
"Failed to initialize chat attachment cleanup, app will continue: %s",
exc,
)
# --- Takeover scheduler ---
try:
# BOXIM production endpoints are intentionally separate from the login API.
from services.boxim_client import BoxIMClient
boxim_config = {
"HUIHUI_PLATFORM_BASE_URL": os.getenv(
"HUIHUI_PLATFORM_BASE_URL", "https://open.99hui.com/api"
),
"BOXIM_API_BASE_URL": os.getenv(
"BOXIM_API_BASE_URL", "https://im.99hui.com/api"
),
"HUIHUI_APP_ID": os.getenv("HUIHUI_APP_ID", ""),
"HUIHUI_ACCESS_ID": os.getenv("HUIHUI_ACCESS_ID", ""),
"HUIHUI_ACCESS_SECRET": os.getenv("HUIHUI_ACCESS_SECRET", ""),
"BOXIM_TIMEOUT_SECONDS": os.getenv("BOXIM_TIMEOUT_SECONDS", "20"),
}
boxim_client = BoxIMClient(boxim_config)
from services.takeover_service import TakeoverService
takeover_service = TakeoverService(
SessionLocal,
boxim_client,
poll_concurrency=int(os.getenv("BOXIM_POLL_CONCURRENCY", "8")),
max_message_age_seconds=int(
os.getenv("BOXIM_MAX_MESSAGE_AGE_SECONDS", "600")
),
)
poll_interval = max(0.5, float(os.getenv("BOXIM_POLL_INTERVAL_SECONDS", "1")))
takeover_scheduler = AsyncIOScheduler()
takeover_scheduler.add_job(
takeover_service.poll_messages,
trigger=IntervalTrigger(seconds=poll_interval),
id="takeover_message_poll",
max_instances=1,
coalesce=True,
)
process_interval = max(
0.25, float(os.getenv("TAKEOVER_PROCESS_INTERVAL_SECONDS", "0.5"))
)
takeover_scheduler.add_job(
takeover_service.process_reply_tasks,
trigger=IntervalTrigger(seconds=process_interval),
id="takeover_reply_process",
max_instances=1,
coalesce=True,
)
takeover_scheduler.start()
logger.info(
"BOXIM takeover scheduler started (poll=%ss, process=%ss)",
poll_interval,
process_interval,
)
except Exception as e:
stop_takeover_scheduler()
logger.warning(f"Failed to initialize takeover scheduler, app will continue without it: {e}")
def stop_takeover_scheduler():
global takeover_scheduler
if takeover_scheduler is not None:
try:
if takeover_scheduler.running:
takeover_scheduler.shutdown(wait=False)
except Exception as e:
logger.warning(f"Failed to stop takeover scheduler cleanly: {e}")
finally:
takeover_scheduler = None
def purge_expired_chat_attachments_job():
db = SessionLocal()
try:
count = purge_expired_chat_attachments(db)
if count:
logger.info("Purged %s expired chat image attachment(s)", count)
except Exception as exc:
db.rollback()
logger.warning("Failed to purge expired chat image attachments: %s", exc)
finally:
db.close()
def start_maintenance_scheduler():
global maintenance_scheduler
purge_expired_chat_attachments_job()
interval_minutes = max(
5, min(1440, int(os.getenv("CHAT_ATTACHMENT_CLEANUP_MINUTES", "60")))
)
maintenance_scheduler = AsyncIOScheduler()
maintenance_scheduler.add_job(
purge_expired_chat_attachments_job,
trigger=IntervalTrigger(minutes=interval_minutes),
id="chat_attachment_cleanup",
max_instances=1,
coalesce=True,
)
maintenance_scheduler.start()
def stop_maintenance_scheduler():
global maintenance_scheduler
if maintenance_scheduler is not None:
try:
if maintenance_scheduler.running:
maintenance_scheduler.shutdown(wait=False)
except Exception as exc:
logger.warning("Failed to stop maintenance scheduler cleanly: %s", exc)
finally:
maintenance_scheduler = None
@app.on_event("shutdown")
def on_shutdown():
stop_takeover_scheduler()
stop_maintenance_scheduler()
+424
View File
@@ -0,0 +1,424 @@
import uuid
from sqlalchemy import (
BigInteger,
Boolean,
Column,
DateTime,
Float,
Index,
Integer,
JSON,
String,
Text,
UniqueConstraint,
)
from sqlalchemy.sql import func
from database import Base
def _iso(dt):
return dt.isoformat() if dt else None
class Avatar(Base):
__tablename__ = "avatars"
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
owner_id = Column(String, default="", index=True) # 归属用户(会会 huihui_user_id);空=未归属/种子数据
name = Column(String, nullable=False)
display_name = Column(String, default="")
description = Column(Text, default="")
photo_url = Column(String, default="")
emoji = Column(String, default="🤖")
status = Column(String, default="active") # active | inactive | training
share_token = Column(String, nullable=True, default=None, unique=True, index=True) # 对外分享使用的不可猜测令牌
token_balance = Column(Integer, default=0)
config = Column(JSON, default=dict)
created_at = Column(DateTime, server_default=func.now())
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
def to_dict(self):
return {
"id": self.id,
"ownerId": self.owner_id,
"name": self.name,
"displayName": self.display_name,
"description": self.description,
"photoUrl": self.photo_url,
"emoji": self.emoji,
"status": self.status,
"shareToken": self.share_token or "",
"tokenBalance": self.token_balance,
"config": self.config or {},
"createdAt": _iso(self.created_at),
"updatedAt": _iso(self.updated_at),
}
class Authorization(Base):
__tablename__ = "authorizations"
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
avatar_id = Column(String, nullable=False, default="")
target_type = Column(String, default="user") # user | organization | application
target_id = Column(String, default="")
target_name = Column(String, default="")
permissions = Column(JSON, default=list)
status = Column(String, default="active") # active | inactive
takeover_enabled = Column(Boolean, default=False) # 是否开启分身接管
takeover_mode = Column(String, default="immediate") # immediate | delayed
takeover_delay_seconds = Column(Integer, default=180) # 延迟秒数,默认 3 分钟
created_at = Column(DateTime, server_default=func.now())
def to_dict(self):
return {
"id": self.id,
"avatarId": self.avatar_id,
"targetType": self.target_type,
"targetId": self.target_id,
"targetName": self.target_name,
"permissions": self.permissions or [],
"status": self.status,
"takeoverEnabled": self.takeover_enabled,
"takeoverMode": self.takeover_mode,
"takeoverDelaySeconds": self.takeover_delay_seconds,
"createdAt": _iso(self.created_at),
}
class TakeoverCursor(Base):
"""Durable BOXIM polling cursor for one avatar owner."""
__tablename__ = "takeover_cursors"
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
avatar_id = Column(String, nullable=False, unique=True, index=True)
owner_id = Column(String, nullable=False, default="", index=True)
boxim_owner_id = Column(String, default="")
last_message_id = Column(String, default="0")
initialized = Column(Boolean, default=False)
last_polled_at = Column(DateTime)
last_error = Column(Text, default="")
created_at = Column(DateTime, server_default=func.now())
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
class TakeoverMessage(Base):
"""BOXIM message receipt used for audit, deduplication, and chat context."""
__tablename__ = "takeover_messages"
__table_args__ = (
UniqueConstraint("owner_id", "boxim_message_id", name="uq_takeover_message_owner_boxim"),
Index("ix_takeover_message_conversation", "owner_id", "peer_id", "send_time"),
)
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
avatar_id = Column(String, nullable=False, index=True)
owner_id = Column(String, nullable=False, index=True)
boxim_message_id = Column(String, nullable=False)
boxim_local_id = Column(String, nullable=True)
peer_id = Column(String, nullable=False, index=True)
direction = Column(String, nullable=False) # incoming | outgoing
message_type = Column(Integer, default=0)
content = Column(Text, default="")
attachment_id = Column(String, nullable=True)
is_avatar = Column(Boolean, default=False)
send_time = Column(DateTime, nullable=False)
created_at = Column(DateTime, server_default=func.now())
class TakeoverReplyTask(Base):
"""Restart-safe delayed BOXIM reply task."""
__tablename__ = "takeover_reply_tasks"
__table_args__ = (
UniqueConstraint("owner_id", "trigger_message_id", name="uq_takeover_task_owner_trigger"),
Index("ix_takeover_task_due", "status", "scheduled_at"),
Index("ix_takeover_task_conversation", "owner_id", "peer_id", "status"),
)
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
avatar_id = Column(String, nullable=False, index=True)
owner_id = Column(String, nullable=False, index=True)
peer_id = Column(String, nullable=False, index=True)
trigger_message_id = Column(String, nullable=False)
source_message_ids = Column(JSON, default=list)
prompt = Column(Text, default="")
response_text = Column(Text, default="")
status = Column(String, default="pending")
scheduled_at = Column(DateTime, nullable=False)
locked_at = Column(DateTime)
sent_at = Column(DateTime)
attempts = Column(Integer, default=0)
last_error = Column(Text, default="")
cancel_reason = Column(String, default="")
boxim_local_id = Column(String, nullable=False)
boxim_sent_message_id = Column(String, default="")
created_at = Column(DateTime, server_default=func.now())
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
class Organization(Base):
__tablename__ = "organizations"
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
name = Column(String, nullable=False)
description = Column(Text, default="")
emoji = Column(String, default="🏢")
org_type = Column(String, default="team") # team | company | community
role = Column(String, default="admin") # admin | member | viewer
member_count = Column(Integer, default=1)
created_at = Column(DateTime, server_default=func.now())
def to_dict(self):
return {
"id": self.id,
"name": self.name,
"description": self.description,
"emoji": self.emoji,
"type": self.org_type,
"role": self.role,
"memberCount": self.member_count,
"createdAt": _iso(self.created_at),
}
class KnowledgeDoc(Base):
__tablename__ = "knowledge_docs"
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
avatar_id = Column(String, nullable=False, default="")
filename = Column(String, default="")
file_type = Column(String, default="") # pdf | doc | docx | xlsx
file_size = Column(Integer, default=0)
file_url = Column(String, default="")
status = Column(String, default="uploaded") # uploaded | parsing | ready | failed
error_message = Column(String, default="") # 建立索引失败原因
index_stage = Column(String, default="") # queued | extracting | chunking | embedding | ready | failed
index_progress = Column(Integer, default=0) # 0-100
vectorized = Column(Boolean, default=False) # 是否已向量化
embedding_model = Column(String, default="") # 向量模型标识
chunk_count = Column(Integer, default=0) # 切片数量
vectorized_at = Column(DateTime) # 向量化时间
created_at = Column(DateTime, server_default=func.now())
def to_dict(self):
return {
"id": self.id,
"avatarId": self.avatar_id,
"filename": self.filename,
"fileType": self.file_type,
"fileSize": self.file_size,
"fileUrl": self.file_url,
"status": self.status,
"errorMessage": self.error_message or "",
"indexStage": self.index_stage or "",
"indexProgress": int(self.index_progress or 0),
"vectorized": bool(self.vectorized),
"embeddingModel": self.embedding_model,
"chunkCount": self.chunk_count,
"vectorizedAt": _iso(self.vectorized_at),
"createdAt": _iso(self.created_at),
}
class QAPair(Base):
__tablename__ = "qa_pairs"
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
avatar_id = Column(String, nullable=False, default="")
question = Column(Text, default="")
answer = Column(Text, default="")
enabled = Column(Boolean, default=True) # 是否启用(关闭后不参与作答)
created_at = Column(DateTime, server_default=func.now())
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
def to_dict(self):
return {
"id": self.id,
"avatarId": self.avatar_id,
"question": self.question,
"answer": self.answer,
"enabled": bool(self.enabled),
"createdAt": _iso(self.created_at),
"updatedAt": _iso(self.updated_at),
}
class KnowledgeChunk(Base):
__tablename__ = "knowledge_chunks"
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
doc_id = Column(String, default="") # 关联 KnowledgeDoc.id
avatar_id = Column(String, default="")
content = Column(Text, default="") # 切片文本
vector = Column(Text, default="") # JSON 编码的向量
chunk_index = Column(Integer, default=0)
embedding_model = Column(String, default="")
created_at = Column(DateTime, server_default=func.now())
def to_dict(self):
return {
"id": self.id,
"docId": self.doc_id,
"avatarId": self.avatar_id,
"content": self.content,
"chunkIndex": self.chunk_index,
"embeddingModel": self.embedding_model,
"createdAt": _iso(self.created_at),
}
class ChatAttachment(Base):
"""Private, avatar-scoped result of one chat image analysis."""
__tablename__ = "chat_attachments"
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
avatar_id = Column(String, nullable=False, default="", index=True)
uploader_kind = Column(String, default="owner") # owner | public | boxim
filename = Column(String, default="")
mime_type = Column(String, default="")
file_size = Column(Integer, default=0)
status = Column(String, default="processing") # processing | ready | failed
category = Column(String, default="general_image")
summary = Column(Text, default="")
extracted_text = Column(Text, default="")
structured_data = Column(JSON, default=dict)
warning = Column(Text, default="")
vision_model = Column(String, default="")
ocr_model = Column(String, default="")
used_at = Column(DateTime)
expires_at = Column(DateTime, nullable=False)
created_at = Column(DateTime, server_default=func.now())
def to_dict(self):
return {
"id": self.id,
"avatarId": self.avatar_id,
"filename": self.filename,
"mimeType": self.mime_type,
"fileSize": self.file_size,
"status": self.status,
"category": self.category,
"summary": self.summary,
"warning": self.warning,
"expiresAt": _iso(self.expires_at),
"createdAt": _iso(self.created_at),
}
class TokenAccount(Base):
__tablename__ = "token_account"
id = Column(Integer, primary_key=True)
user_id = Column(String, nullable=False, default="", index=True)
balance = Column(BigInteger, default=1_000_000)
total_granted = Column(BigInteger, default=1_000_000)
total_consumed = Column(BigInteger, default=0)
created_at = Column(DateTime, server_default=func.now())
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
class TokenUsage(Base):
__tablename__ = "token_usage"
__table_args__ = (
Index("ix_token_usage_user_created", "user_id", "created_at"),
Index("ix_token_usage_avatar_created", "avatar_id", "created_at"),
)
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
user_id = Column(String, nullable=False, index=True)
avatar_id = Column(String, nullable=False, default="", index=True)
source = Column(String, nullable=False, default="chat")
model = Column(String, default="")
status = Column(String, nullable=False, default="reserved")
reserved_tokens = Column(BigInteger, default=0)
prompt_tokens = Column(BigInteger, default=0)
completion_tokens = Column(BigInteger, default=0)
total_tokens = Column(BigInteger, default=0)
balance_after = Column(BigInteger, default=0)
failure_reason = Column(String, default="")
created_at = Column(DateTime, server_default=func.now())
settled_at = Column(DateTime)
class TokenPlan(Base):
__tablename__ = "token_plans"
id = Column(String, primary_key=True)
name = Column(String, default="")
amount = Column(BigInteger, default=0)
price = Column(Float, default=0)
badge = Column(String, default="")
desc = Column(String, default="")
def to_dict(self):
return {
"id": self.id,
"name": self.name,
"amount": self.amount,
"price": self.price,
"badge": self.badge,
"desc": self.desc,
}
class TokenPaymentOrder(Base):
__tablename__ = "token_payment_orders"
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
order_no = Column(String, nullable=False, unique=True, index=True)
user_id = Column(String, nullable=False, index=True)
plan_id = Column(String, nullable=False)
payment_method = Column(String, nullable=False)
pay_type = Column(String, nullable=False)
pay_way = Column(String, nullable=False)
points_amount = Column(BigInteger, nullable=False)
price_cents = Column(Integer, nullable=False)
status = Column(String, nullable=False, default="pending", index=True)
provider_order_id = Column(String, default="")
provider_order_no = Column(String, default="")
provider_status = Column(String, default="")
pay_message = Column(Text, default="")
failure_reason = Column(String, default="")
created_at = Column(DateTime, server_default=func.now())
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
paid_at = Column(DateTime)
def to_dict(self):
return {
"id": self.id,
"orderNo": self.order_no,
"planId": self.plan_id,
"paymentMethod": self.payment_method,
"payType": self.pay_type,
"payWay": self.pay_way,
"pointsAmount": self.points_amount,
"price": self.price_cents / 100,
"status": self.status,
"providerStatus": self.provider_status,
"payMessage": self.pay_message,
"failureReason": self.failure_reason,
"createdAt": _iso(self.created_at),
"paidAt": _iso(self.paid_at),
}
class User(Base):
"""会会用户 ↔ 本地用户体系映射(短信验证码登录落库)"""
__tablename__ = "users"
id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex)
huihui_user_id = Column(String, default="", index=True) # 会会 userId(唯一标识)
phone = Column(String, default="", index=True)
nickname = Column(String, default="")
avatar_url = Column(String, default="")
huihui_token = Column(String, default="") # 会会 access_token
app_token = Column(String, default="") # 本系统会话 token
last_login_at = Column(DateTime)
created_at = Column(DateTime, server_default=func.now())
updated_at = Column(DateTime, server_default=func.now(), onupdate=func.now())
def to_dict(self):
return {
"id": self.id,
"huihuiUserId": self.huihui_user_id,
"phone": self.phone,
"nickname": self.nickname,
"avatarUrl": self.avatar_url,
"createdAt": _iso(self.created_at),
"lastLoginAt": _iso(self.last_login_at),
}
@@ -0,0 +1,11 @@
fastapi
uvicorn[standard]
sqlalchemy
pydantic
python-multipart
httpx
pypdf
python-docx
openpyxl
apscheduler>=3.10
Pillow>=10.4
+6
View File
@@ -0,0 +1,6 @@
def ok(data=None, message="success"):
return {"code": 200, "message": message, "data": data}
def fail(message="error", code=400):
return {"code": code, "message": message, "data": None}
@@ -0,0 +1,404 @@
from fastapi import APIRouter, Body, Depends, Header, HTTPException
from sqlalchemy.orm import Session
from database import get_db
from models import Authorization, Avatar, TakeoverCursor, TakeoverReplyTask
from responses import fail, ok
from routers.avatars import _require_owned_avatar
router = APIRouter(tags=["授权"])
TARGET_TYPES = {"user", "organization", "application"}
PERMISSION_ORDER = ("friend", "chat", "publish", "browse", "interact", "takeover")
ALLOWED_PERMISSIONS = set(PERMISSION_ORDER)
AVATAR_PERMISSION_ORDER = PERMISSION_ORDER
AVATAR_PERMISSION_KEY = "authorizationPermissions"
DEFAULT_AVATAR_PERMISSIONS = ["friend", "chat"]
TAKEOVER_DELAY_KEY = "takeoverReplyDelaySeconds"
DEFAULT_TAKEOVER_DELAY_SECONDS = 180
MIN_TAKEOVER_DELAY_SECONDS = 3
MAX_TAKEOVER_DELAY_SECONDS = 86_400
LEGACY_PERMISSION_MAP = {
"read": "browse",
"reply": "chat",
"write": "publish",
"edit": "publish",
}
def _read(payload: dict, camel_key: str, snake_key: str | None = None, default=None):
if camel_key in payload:
return payload[camel_key]
if snake_key and snake_key in payload:
return payload[snake_key]
return default
def _clean_text(value, field_name: str, *, max_length: int) -> str:
text = str(value or "").strip()
if not text:
raise ValueError(f"{field_name}不能为空")
if len(text) > max_length:
raise ValueError(f"{field_name}不能超过 {max_length} 个字符")
return text
def _normalize_permissions(value) -> list[str]:
if not isinstance(value, list):
raise ValueError("权限格式不正确")
normalized = []
for raw in value:
permission = LEGACY_PERMISSION_MAP.get(str(raw).strip(), str(raw).strip())
if permission not in ALLOWED_PERMISSIONS:
raise ValueError(f"不支持的权限:{raw}")
if permission not in normalized:
normalized.append(permission)
if not [item for item in normalized if item != "takeover"]:
raise ValueError("请至少选择一项权限")
return sorted(normalized, key=PERMISSION_ORDER.index)
def _normalize_avatar_permissions(value) -> list[str]:
if not isinstance(value, list):
raise ValueError("权限格式不正确")
normalized = []
for raw in value:
permission = LEGACY_PERMISSION_MAP.get(str(raw).strip(), str(raw).strip())
if permission not in AVATAR_PERMISSION_ORDER:
raise ValueError(f"不支持的权限:{raw}")
if permission not in normalized:
normalized.append(permission)
return sorted(normalized, key=AVATAR_PERMISSION_ORDER.index)
def _stored_avatar_permissions(avatar) -> list[str]:
config = avatar.config or {}
if AVATAR_PERMISSION_KEY not in config:
return list(DEFAULT_AVATAR_PERMISSIONS)
stored = config.get(AVATAR_PERMISSION_KEY)
if not isinstance(stored, list):
return list(DEFAULT_AVATAR_PERMISSIONS)
permissions = []
for raw in stored:
permission = LEGACY_PERMISSION_MAP.get(str(raw).strip(), str(raw).strip())
if permission in AVATAR_PERMISSION_ORDER and permission not in permissions:
permissions.append(permission)
return sorted(permissions, key=AVATAR_PERMISSION_ORDER.index)
def _permission_settings_payload(avatar) -> dict:
return {
"avatarId": avatar.id,
"permissions": _stored_avatar_permissions(avatar),
"takeoverReplyDelaySeconds": _stored_takeover_delay(avatar),
}
def _stored_takeover_delay(avatar) -> int:
raw = (avatar.config or {}).get(TAKEOVER_DELAY_KEY, DEFAULT_TAKEOVER_DELAY_SECONDS)
if isinstance(raw, bool):
return DEFAULT_TAKEOVER_DELAY_SECONDS
try:
delay = int(raw)
except (TypeError, ValueError):
return DEFAULT_TAKEOVER_DELAY_SECONDS
if not MIN_TAKEOVER_DELAY_SECONDS <= delay <= MAX_TAKEOVER_DELAY_SECONDS:
return DEFAULT_TAKEOVER_DELAY_SECONDS
return delay
def _validate_takeover_delay(value) -> int:
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError("自动回复等待时间必须是整数秒")
if not MIN_TAKEOVER_DELAY_SECONDS <= value <= MAX_TAKEOVER_DELAY_SECONDS:
raise ValueError("自动回复等待时间需在 3 秒到 24 小时之间")
return value
def _disable_other_takeovers(db: Session, avatar) -> list[str]:
disabled_ids = []
others = (
db.query(Avatar)
.filter(Avatar.owner_id == avatar.owner_id, Avatar.id != avatar.id)
.all()
)
for other in others:
permissions = _stored_avatar_permissions(other)
if "takeover" not in permissions:
continue
other.config = {
**(other.config or {}),
AVATAR_PERMISSION_KEY: [item for item in permissions if item != "takeover"],
}
disabled_ids.append(other.id)
tasks = (
db.query(TakeoverReplyTask)
.filter(
TakeoverReplyTask.avatar_id == other.id,
TakeoverReplyTask.status.in_(("pending", "generating", "ready", "sending")),
)
.all()
)
for task in tasks:
task.status = "cancelled"
task.cancel_reason = "another_avatar_takeover_enabled"
task.locked_at = None
return disabled_ids
def _require_authorization(db: Session, avatar_id: str, authorization_id: str) -> Authorization:
authorization = (
db.query(Authorization)
.filter(
Authorization.id == authorization_id,
Authorization.avatar_id == avatar_id,
)
.first()
)
if not authorization:
raise HTTPException(status_code=404, detail="授权不存在")
return authorization
def _duplicate_target(
db: Session,
avatar_id: str,
target_type: str,
target_id: str,
*,
exclude_id: str | None = None,
):
query = db.query(Authorization).filter(
Authorization.avatar_id == avatar_id,
Authorization.target_type == target_type,
Authorization.target_id == target_id,
)
if exclude_id:
query = query.filter(Authorization.id != exclude_id)
return query.first()
@router.get("/avatar/{avatar_id}/permission-settings")
def get_permission_settings(
avatar_id: str,
authorization: str = Header(None),
db: Session = Depends(get_db),
):
avatar = _require_owned_avatar(db, avatar_id, authorization)
return ok(_permission_settings_payload(avatar))
@router.put("/avatar/{avatar_id}/permission-settings")
def update_permission_settings(
avatar_id: str,
payload: dict = Body(...),
authorization: str = Header(None),
db: Session = Depends(get_db),
):
avatar = _require_owned_avatar(db, avatar_id, authorization)
if "permissions" not in payload and TAKEOVER_DELAY_KEY not in payload:
return fail("缺少授权设置", 400)
try:
permissions = (
_normalize_avatar_permissions(payload["permissions"])
if "permissions" in payload
else _stored_avatar_permissions(avatar)
)
takeover_delay = (
_validate_takeover_delay(payload[TAKEOVER_DELAY_KEY])
if TAKEOVER_DELAY_KEY in payload
else _stored_takeover_delay(avatar)
)
except ValueError as exc:
return fail(str(exc), 400)
previous_permissions = _stored_avatar_permissions(avatar)
avatar.config = {
**(avatar.config or {}),
AVATAR_PERMISSION_KEY: permissions,
TAKEOVER_DELAY_KEY: takeover_delay,
}
disabled_avatar_ids = _disable_other_takeovers(db, avatar) if "takeover" in permissions else []
cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar.id).first()
if cursor and "takeover" in permissions and "takeover" not in previous_permissions:
cursor.initialized = False
cursor.last_message_id = "0"
cursor.last_error = ""
elif cursor and "takeover" not in permissions:
cursor.last_error = ""
if "takeover" not in permissions:
tasks = (
db.query(TakeoverReplyTask)
.filter(
TakeoverReplyTask.avatar_id == avatar.id,
TakeoverReplyTask.status.in_(("pending", "generating", "ready", "sending")),
)
.all()
)
for task in tasks:
task.status = "cancelled"
task.cancel_reason = "takeover_disabled"
task.locked_at = None
db.commit()
db.refresh(avatar)
response = _permission_settings_payload(avatar)
response["disabledAvatarIds"] = disabled_avatar_ids
return ok(response, "授权设置已保存")
@router.get("/avatar/{avatar_id}/authorizations")
def list_auth(
avatar_id: str,
authorization: str = Header(None),
db: Session = Depends(get_db),
):
_require_owned_avatar(db, avatar_id, authorization)
items = (
db.query(Authorization)
.filter(Authorization.avatar_id == avatar_id)
.order_by(Authorization.created_at.desc())
.all()
)
return ok([item.to_dict() for item in items])
@router.post("/avatar/{avatar_id}/authorizations")
def create_auth(
avatar_id: str,
payload: dict = Body(...),
authorization: str = Header(None),
db: Session = Depends(get_db),
):
_require_owned_avatar(db, avatar_id, authorization)
try:
target_type = _clean_text(
_read(payload, "targetType", "target_type", "user"),
"授权类型",
max_length=24,
)
if target_type not in TARGET_TYPES:
return fail("授权类型不正确", 400)
target_id = _clean_text(
_read(payload, "targetId", "target_id"),
"对象标识",
max_length=120,
)
target_name = _clean_text(
_read(payload, "targetName", "target_name"),
"对象名称",
max_length=50,
)
permissions = _normalize_permissions(payload.get("permissions", []))
except ValueError as exc:
return fail(str(exc), 400)
if _duplicate_target(db, avatar_id, target_type, target_id):
return fail("该对象已在授权列表中,可直接编辑现有授权", 409)
item = Authorization(
avatar_id=avatar_id,
target_type=target_type,
target_id=target_id,
target_name=target_name,
permissions=permissions,
status="active",
takeover_enabled=False,
takeover_mode="immediate",
takeover_delay_seconds=DEFAULT_TAKEOVER_DELAY_SECONDS,
)
db.add(item)
db.commit()
db.refresh(item)
return ok(item.to_dict(), "授权已添加")
@router.put("/avatar/{avatar_id}/authorizations")
def update_auth(
avatar_id: str,
payload: dict = Body(...),
authorization: str = Header(None),
db: Session = Depends(get_db),
):
_require_owned_avatar(db, avatar_id, authorization)
auth_id = payload.get("id") or _read(payload, "authorizationId", "authorization_id")
if not auth_id:
return fail("缺少授权 id", 400)
item = _require_authorization(db, avatar_id, str(auth_id))
try:
target_type = item.target_type
target_id = item.target_id
if "targetType" in payload or "target_type" in payload:
target_type = _clean_text(
_read(payload, "targetType", "target_type"),
"授权类型",
max_length=24,
)
if target_type not in TARGET_TYPES:
return fail("授权类型不正确", 400)
if "targetId" in payload or "target_id" in payload:
target_id = _clean_text(
_read(payload, "targetId", "target_id"),
"对象标识",
max_length=120,
)
if "targetName" in payload or "target_name" in payload:
item.target_name = _clean_text(
_read(payload, "targetName", "target_name"),
"对象名称",
max_length=50,
)
if "permissions" in payload:
item.permissions = _normalize_permissions(payload["permissions"])
except ValueError as exc:
return fail(str(exc), 400)
if _duplicate_target(
db,
avatar_id,
target_type,
target_id,
exclude_id=item.id,
):
return fail("该对象已在授权列表中", 409)
if "status" in payload:
status = str(payload["status"] or "")
if status not in ("active", "inactive"):
return fail("授权状态不正确", 400)
item.status = status
item.target_type = target_type
item.target_id = target_id
permissions = list(item.permissions or [])
chat_allowed = "chat" in permissions or "reply" in permissions
if item.status != "active" or item.target_type != "user" or not chat_allowed:
item.takeover_enabled = False
item.permissions = [permission for permission in permissions if permission != "takeover"]
elif item.takeover_enabled and "takeover" not in permissions:
item.permissions = permissions + ["takeover"]
db.commit()
db.refresh(item)
return ok(item.to_dict(), "授权已更新")
@router.delete("/avatar/{avatar_id}/authorizations/{authorization_id}")
def delete_auth(
avatar_id: str,
authorization_id: str,
authorization: str = Header(None),
db: Session = Depends(get_db),
):
_require_owned_avatar(db, avatar_id, authorization)
item = _require_authorization(db, avatar_id, authorization_id)
db.delete(item)
db.commit()
return ok({"id": authorization_id}, "授权已删除")
@@ -0,0 +1,161 @@
import os
import uuid
from fastapi import APIRouter, Depends, Body, Header, UploadFile, File, HTTPException
from sqlalchemy.orm import Session
from database import get_db
from routers.knowledge import UPLOAD_DIR
from models import (
Authorization,
Avatar,
KnowledgeChunk,
KnowledgeDoc,
QAPair,
TakeoverCursor,
TakeoverMessage,
TakeoverReplyTask,
User,
)
from responses import ok, fail
router = APIRouter(tags=["分身"])
ALLOWED_AVATAR_EXTENSIONS = {".jpg", ".jpeg", ".png", ".webp", ".gif"}
MAX_AVATAR_BYTES = 5 * 1024 * 1024
def _resolve_user(authorization: str | None, db: Session):
"""从 Authorization: Bearer <app_token> 解析当前登录用户"""
if not authorization:
return None
token = authorization.replace("Bearer ", "", 1).replace("bearer ", "", 1).strip()
return db.query(User).filter(User.app_token == token).first()
def _require_owned_avatar(db: Session, avatar_id: str, authorization: str | None):
avatar = db.query(Avatar).filter(Avatar.id == avatar_id).first()
if not avatar:
raise HTTPException(status_code=404, detail="分身不存在")
user = _resolve_user(authorization, db)
if not user:
raise HTTPException(status_code=401, detail="未登录")
if avatar.owner_id and avatar.owner_id != user.huihui_user_id:
raise HTTPException(status_code=403, detail="无权访问该分身")
return avatar
@router.post("/avatar/{avatar_id}/photo")
async def upload_avatar_photo(
avatar_id: str,
file: UploadFile = File(...),
authorization: str = Header(None),
db: Session = Depends(get_db),
):
_require_owned_avatar(db, avatar_id, authorization)
extension = os.path.splitext(file.filename or "")[1].lower()
if extension not in ALLOWED_AVATAR_EXTENSIONS or not (file.content_type or "").startswith("image/"):
return fail("仅支持 JPG、PNG、WebP 或 GIF 图片", code=400)
content = await file.read()
if len(content) > MAX_AVATAR_BYTES:
return fail("头像图片不能超过 5MB", code=400)
avatar_dir = os.path.join(UPLOAD_DIR, avatar_id)
os.makedirs(avatar_dir, exist_ok=True)
stored_name = f"avatar-{uuid.uuid4().hex}{extension}"
with open(os.path.join(avatar_dir, stored_name), "wb") as stream:
stream.write(content)
return ok({"photoUrl": f"/api/files/{avatar_id}/{stored_name}"})
@router.get("/avatar")
def list_avatars(page: int = 1, limit: int = 20, authorization: str = Header(None), db: Session = Depends(get_db)):
# 仅返回当前登录用户自己的分身;未登录返回空,避免看到种子/他人数据
user = _resolve_user(authorization, db)
if not user:
return ok({"data": [], "total": 0})
q = db.query(Avatar).filter(Avatar.owner_id == user.huihui_user_id)
total = q.count()
items = (
q.order_by(Avatar.created_at.desc())
.offset((page - 1) * limit)
.limit(limit)
.all()
)
return ok({"data": [a.to_dict() for a in items], "total": total})
@router.get("/avatar/{avatar_id}")
def get_avatar(
avatar_id: str,
authorization: str = Header(None),
db: Session = Depends(get_db),
):
return ok(_require_owned_avatar(db, avatar_id, authorization).to_dict())
@router.post("/avatar")
def create_avatar(payload: dict = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)):
user = _resolve_user(authorization, db)
if not user:
raise HTTPException(status_code=401, detail="未登录")
a = Avatar(
owner_id=user.huihui_user_id,
name=payload.get("name", "未命名分身"),
display_name=payload.get("displayName", "") or payload.get("display_name", ""),
description=payload.get("description", ""),
photo_url=payload.get("photoUrl", "") or payload.get("photo_url", ""),
emoji=payload.get("emoji", "🤖"),
status=payload.get("status", "active"),
token_balance=payload.get("tokenBalance", 0),
config=payload.get("config", {}) or {},
)
db.add(a)
db.commit()
db.refresh(a)
return ok(a.to_dict())
@router.put("/avatar/{avatar_id}")
def update_avatar(
avatar_id: str,
payload: dict = Body(...),
authorization: str = Header(None),
db: Session = Depends(get_db),
):
a = _require_owned_avatar(db, avatar_id, authorization)
mapping = {
"displayName": "display_name",
"photoUrl": "photo_url",
"tokenBalance": "token_balance",
}
for key in ("name", "displayName", "description", "photoUrl", "emoji", "status", "tokenBalance", "config"):
if key in payload:
col = mapping.get(key, key)
value = payload[key]
if key == "config":
if not isinstance(value, dict):
return fail("分身配置格式不正确", 400)
value = {**(a.config or {}), **value}
setattr(a, col, value)
db.commit()
db.refresh(a)
return ok(a.to_dict())
@router.delete("/avatar/{avatar_id}")
def delete_avatar(
avatar_id: str,
authorization: str = Header(None),
db: Session = Depends(get_db),
):
a = _require_owned_avatar(db, avatar_id, authorization)
# 级联清理关联数据,避免孤儿记录
db.query(KnowledgeDoc).filter(KnowledgeDoc.avatar_id == avatar_id).delete()
db.query(KnowledgeChunk).filter(KnowledgeChunk.avatar_id == avatar_id).delete()
db.query(QAPair).filter(QAPair.avatar_id == avatar_id).delete()
db.query(Authorization).filter(Authorization.avatar_id == avatar_id).delete()
db.query(TakeoverReplyTask).filter(TakeoverReplyTask.avatar_id == avatar_id).delete()
db.query(TakeoverMessage).filter(TakeoverMessage.avatar_id == avatar_id).delete()
db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar_id).delete()
db.delete(a)
db.commit()
return ok({"success": True})
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,449 @@
"""
会会短信验证码登录代理(真实开放平台对接)
──────────────────────────────────────────────
严格按会会开放平台 sign.js 签名范式:
- 公共字段 appId/accessId/timestamp(12小时制hh)/signType/signVersion/accessSecret/nonce 合并业务参数
- 过滤空值 → 字典序排序 → k=v& 拼接 → 末尾追加 accessSecret=secretKey → MD5/SHA256 大写
- 认证服务基址 / appId / accessId / accessSecret / clientCode 走环境变量
真实端点(来自 fat-open 网关 usercenter 的 Swagger):
- 发送验证码:POST {BASE}{HUIHUI_SMS_SEND_PATH} 默认 /open/mobile/sms/code (query 参数)
- 短信登录: POST {BASE}{HUIHUI_SMS_LOGIN_PATH} 默认 /open/login/token,loginType=code
登录成功后建/链本地 users 表(按会会 userId 唯一),签发本系统 app_token 作为会话。
无真实凭证时仍可走 DEV_MOCK 兜底联调。
"""
import os
import uuid
import random
import string
import hashlib
import httpx
from datetime import datetime, timezone, timedelta
from fastapi import APIRouter, Body, Depends, Header
from sqlalchemy.orm import Session
# 会会网关按北京时间(Asia/Shanghai, UTC+8)校验时间戳,容器默认 UTC 会导致签名被拒。
# 用固定 +8 偏移(不依赖 tzdata,slim 镜像缺 IANA 库时会抛 ZoneInfoNotFoundError)。
_CN_TZ = timezone(timedelta(hours=8))
from database import get_db
from models import Avatar, TakeoverCursor, TakeoverMessage, TakeoverReplyTask, User
from responses import ok, fail
from services.boxim_client import BoxIMClient, BoxIMError
router = APIRouter(tags=["会会账号"])
# ── 会会开放平台配置(环境变量)──
AUTH_BASE_URL = os.getenv("HUIHUI_AUTH_BASE_URL", "https://fat-open.99hui.com/api/usercenter")
APP_ID = os.getenv("HUIHUI_APP_ID", "")
ACCESS_ID = os.getenv("HUIHUI_ACCESS_ID", "")
ACCESS_SECRET = os.getenv("HUIHUI_ACCESS_SECRET", "")
CLIENT_CODE = os.getenv("HUIHUI_CLIENT_CODE", "")
SMS_SEND_PATH = os.getenv("HUIHUI_SMS_SEND_PATH", "/open/mobile/sms/code")
SMS_LOGIN_PATH = os.getenv("HUIHUI_SMS_LOGIN_PATH", "/open/login/token")
# 临时开发态:无真实会会凭证时,本地模拟短信收发,便于端到端联调。
DEV_MOCK = os.getenv("HUIHUI_DEV_MOCK", "false").lower() in ("1", "true", "yes")
_mock_codes: dict[str, tuple[str, float]] = {}
_MOCK_TTL = 300
# ── 签名体系(完全对应 sign.js)──
def _get_nonce() -> str:
# 与会会 sign.js 一致:base36 随机串
return "".join(random.choices(string.ascii_lowercase + string.digits, k=12))
def _get_timestamp() -> str:
# 与会会 sign.js 一致:yyyyMMddHHmmss(24 小时制大写 HH)。
# 实测:会会 common.format 走 24 小时制,用 12 小时制(%I)会被网关判"签名验证失败"。
# 必须用北京时间,否则容器(UTC)生成的时间戳与会会网关校验窗口偏差 8h 被拒。
return datetime.now(_CN_TZ).strftime("%Y%m%d%H%M%S")
def _make_sign(params: dict, secret_key: str, sign_type: str = "MD5") -> str:
SIGN_KEY = "signature"
SECRET_KEY = "accessSecret"
keys = sorted(params.keys())
parts = []
for k in keys:
if k in (SIGN_KEY, SECRET_KEY):
continue
v = params.get(k)
if v is None or v == "" or v == []:
continue
if isinstance(v, list):
continue
parts.append(f"{k}={v}")
sign_str = "&".join(parts) + f"&{SECRET_KEY}={secret_key}"
if sign_type.upper() == "SHA256":
return hashlib.sha256(sign_str.encode("utf-8")).hexdigest().upper()
return hashlib.md5(sign_str.encode("utf-8")).hexdigest().upper()
def _build_form(extra: dict) -> dict:
sign_type = "MD5"
sign_version = "1.0"
secret = ACCESS_SECRET
base = {
"appId": APP_ID,
"accessId": ACCESS_ID,
"timestamp": _get_timestamp(),
"signType": sign_type,
"signVersion": sign_version,
"accessSecret": secret,
"nonce": _get_nonce(),
}
base.update(extra)
signature = _make_sign(base, secret, sign_type) if secret else ""
base["signature"] = signature
base.pop("accessSecret", None) # 不发送密钥
return base
def _pick(d: dict, *keys, default=""):
for k in keys:
if d.get(k) not in (None, ""):
return d[k]
return default
def _cfg_ready() -> bool:
return bool(AUTH_BASE_URL and APP_ID and ACCESS_ID and ACCESS_SECRET)
def _create_boxim_client() -> BoxIMClient:
return BoxIMClient({
"HUIHUI_PLATFORM_BASE_URL": os.getenv(
"HUIHUI_PLATFORM_BASE_URL", "https://open.99hui.com/api"
),
"BOXIM_API_BASE_URL": os.getenv("BOXIM_API_BASE_URL", "https://im.99hui.com/api"),
"HUIHUI_APP_ID": APP_ID,
"HUIHUI_ACCESS_ID": ACCESS_ID,
"HUIHUI_ACCESS_SECRET": ACCESS_SECRET,
"BOXIM_TIMEOUT_SECONDS": os.getenv("BOXIM_TIMEOUT_SECONDS", "20"),
})
def _call_huihui(path: str, params: dict, as_query: bool = False):
"""调用会会接口,返回 (ok: bool, payload: dict, http_status: int)"""
url = f"{AUTH_BASE_URL}{path}"
with httpx.Client(timeout=30, follow_redirects=True) as c:
if as_query:
resp = c.post(url, params=params)
else:
resp = c.post(url, data=params)
try:
data = resp.json()
except Exception:
return False, {"message": f"会会返回非JSON: {resp.text[:200]}"}, resp.status_code
# 会会统一包装 {code, message, data}
code = data.get("code")
if resp.status_code == 200 and code in (0, 200, "0", "200"):
return True, data, resp.status_code
return False, data, resp.status_code
@router.post("/huihui/sms/send")
def sms_send(body: dict = Body(...)):
"""请求会会发送短信验证码(真实开放平台 /open/mobile/sms/code)"""
phone = (body.get("phone") or "").strip()
if not phone or not phone.isdigit() or len(phone) != 11:
return fail("请输入正确的 11 位手机号", 400)
# 临时开发态:本地模拟发码
if DEV_MOCK and not _cfg_ready():
code = str(random.randint(100000, 999999))
_mock_codes[phone] = (code, datetime.now().timestamp() + _MOCK_TTL)
return ok({"sent": True, "devCode": code, "dev": True})
if not _cfg_ready():
return fail("会会短信服务未配置(缺少 HUIHUI_APP_ID / HUIHUI_ACCESS_ID / HUIHUI_ACCESS_SECRET)", 500)
# /open/mobile/sms/code:mobile 与公共字段均走 query
form = _build_form({"mobile": phone})
ok_flag, data, status = _call_huihui(SMS_SEND_PATH, form, as_query=True)
if not ok_flag:
return fail(data.get("message") or f"发送失败(HTTP {status})", 502)
return ok({"sent": True})
@router.post("/huihui/sms/login")
def sms_login(body: dict = Body(...), db: Session = Depends(get_db)):
"""短信验证码登录:调会会 /open/login/token(loginType=code) 换取 access_token + userId,落库并签发本系统会话"""
phone = (body.get("phone") or "").strip()
code = (body.get("code") or "").strip()
if not phone or not code:
return fail("手机号或验证码缺失", 400)
# 临时开发态:本地校验模拟码
if DEV_MOCK and not _cfg_ready():
rec = _mock_codes.get(phone)
if not rec or rec[1] < datetime.now().timestamp():
return fail("验证码已失效,请重新获取", 401)
if rec[0] != code:
return fail("验证码错误", 401)
_mock_codes.pop(phone, None)
return _issue_session(db, phone, {
"userId": f"dev_{phone}",
"nickname": f"会会用户{phone[-4:]}",
"avatarUrl": "",
"token": f"dev_token_{phone}",
})
if not _cfg_ready():
return fail("会会登录服务未配置", 500)
extra = {
"username": phone,
"password": code,
"loginType": "code",
"grantType": "password",
"isRegister": "true",
}
if CLIENT_CODE:
extra["clientCode"] = CLIENT_CODE
form = _build_form(extra)
ok_flag, data, status = _call_huihui(SMS_LOGIN_PATH, form, as_query=False)
if not ok_flag:
return fail(data.get("message") or f"登录失败(HTTP {status})", 401)
raw = data.get("data") or {}
user_info = raw.get("userInfo", {}) if isinstance(raw, dict) else {}
huihui_token = (
_pick(raw, "accessToken", "access_token", "token")
or _pick(user_info, "accessToken", "access_token", "token")
)
huihui_user_id = (
_pick(raw, "openid", "userId", "uid", "openId", "id")
or _pick(user_info, "openid", "userId", "uid", "openId", "id")
)
nickname = (
_pick(raw, "nickName", "nickname", "name", "userName")
or _pick(user_info, "nickName", "nickname", "name", "userName")
)
avatar_url = (
_pick(raw, "avatar", "avatarUrl", "headImgUrl", "headimgurl")
or _pick(user_info, "avatar", "avatarUrl", "headImgUrl", "headimgurl")
)
if not huihui_user_id:
return fail("会会未返回用户标识", 502)
return _issue_session(db, phone, {
"userId": huihui_user_id,
"nickname": nickname,
"avatarUrl": avatar_url,
"token": huihui_token,
})
@router.post("/huihui/pwd/login")
def pwd_login(body: dict = Body(...), db: Session = Depends(get_db)):
"""账号密码登录:调会会 /open/login/token(loginType=password) 换取 access_token + userId,落库并签发本系统会话。
与 news_service 既有对接完全一致:username=账号/手机号, password=密码, loginType=password, grantType=password, isRegister=false。
"""
account = (body.get("account") or "").strip()
password = (body.get("password") or "").strip()
if not account or not password:
return fail("账号或密码缺失", 400)
if not _cfg_ready():
return fail("会会登录服务未配置", 500)
extra = {
"username": account,
"password": password,
"loginType": "password",
"grantType": "password",
"isRegister": "false",
}
if CLIENT_CODE:
extra["clientCode"] = CLIENT_CODE
form = _build_form(extra)
ok_flag, data, status = _call_huihui(SMS_LOGIN_PATH, form, as_query=False)
if not ok_flag:
return fail(data.get("message") or f"登录失败(HTTP {status})", 401)
raw = data.get("data") or {}
user_info = raw.get("userInfo", {}) if isinstance(raw, dict) else {}
huihui_token = (
_pick(raw, "accessToken", "access_token", "token")
or _pick(user_info, "accessToken", "access_token", "token")
)
huihui_user_id = (
_pick(raw, "openid", "userId", "uid", "openId", "id")
or _pick(user_info, "openid", "userId", "uid", "openId", "id")
)
nickname = (
_pick(raw, "nickName", "nickname", "name", "userName")
or _pick(user_info, "nickName", "nickname", "name", "userName")
)
avatar_url = (
_pick(raw, "avatar", "avatarUrl", "headImgUrl", "headimgurl")
or _pick(user_info, "avatar", "avatarUrl", "headImgUrl", "headimgurl")
)
if not huihui_user_id:
return fail("会会未返回用户标识", 502)
# 账号即手机号时记录,便于资料展示;非手机号(用户名)则不覆盖已有 phone
phone = account if (account.isdigit() and len(account) == 11) else ""
return _issue_session(db, phone, {
"userId": huihui_user_id,
"nickname": nickname,
"avatarUrl": avatar_url,
"token": huihui_token,
})
@router.post("/huihui/token/login")
async def token_login(body: dict = Body(...), db: Session = Depends(get_db)):
"""Validate a production Huihui token through BOXIM and issue an app session."""
huihui_token = (body.get("token") or "").strip()
if not huihui_token or len(huihui_token) > 8192:
return fail("会会登录凭证无效或已过期", 401)
if not _cfg_ready():
return fail("会会登录服务未配置", 500)
client = _create_boxim_client()
try:
token_data = await client.exchange_access_token(huihui_token)
profile = await client.get_self(token_data["accessToken"])
except BoxIMError as exc:
if exc.auth_error:
return fail("会会登录凭证无效或已过期", 401)
return fail("会会登录服务暂时不可用,请稍后重试", 502)
# BOXIM's id is its internal IM id. Account ownership must use huihuiUserId.
huihui_user_id = str(profile.get("huihuiUserId") or "").strip()
if not huihui_user_id:
return fail("会会未返回用户标识", 502)
phone = str(_pick(profile, "mobile", "phone", default="")).strip()
nickname = str(_pick(profile, "nickName", "nickname", "name", "userName", default="")).strip()
avatar_url = str(
_pick(profile, "headImage", "headImageThumb", "avatar", "avatarUrl", default="")
).strip()
return _issue_session(
db,
phone,
{
"userId": huihui_user_id,
"nickname": nickname,
"avatarUrl": avatar_url,
"token": huihui_token,
},
reuse_existing_session=True,
)
def _transfer_avatar_ownership(db: Session, old_owner_id: str, new_owner_id: str) -> int:
"""Move one user's avatar-owned data to a replacement Huihui identity."""
if not old_owner_id or old_owner_id == new_owner_id:
return 0
avatar_ids = [
avatar_id
for (avatar_id,) in db.query(Avatar.id).filter(Avatar.owner_id == old_owner_id).all()
]
if not avatar_ids:
return 0
db.query(Avatar).filter(Avatar.id.in_(avatar_ids)).update(
{Avatar.owner_id: new_owner_id}, synchronize_session="fetch"
)
for model in (TakeoverCursor, TakeoverMessage, TakeoverReplyTask):
db.query(model).filter(model.avatar_id.in_(avatar_ids)).update(
{model.owner_id: new_owner_id}, synchronize_session="fetch"
)
return len(avatar_ids)
def _find_or_link_user(db: Session, phone: str, huihui_user_id: str) -> User:
"""Resolve an account and safely retain avatars across Huihui environments."""
user = db.query(User).filter(User.huihui_user_id == huihui_user_id).first()
if not phone:
return user or User(huihui_user_id=huihui_user_id)
same_phone_users = db.query(User).filter(User.phone == phone).all()
if user is None:
# A unique verified-phone match is the same person whose upstream ID changed.
if len(same_phone_users) == 1:
user = same_phone_users[0]
old_owner_id = user.huihui_user_id
_transfer_avatar_ownership(db, old_owner_id, huihui_user_id)
user.huihui_user_id = huihui_user_id
return user
return User(huihui_user_id=huihui_user_id)
legacy_users = [candidate for candidate in same_phone_users if candidate.id != user.id]
current_avatar_count = db.query(Avatar).filter(Avatar.owner_id == huihui_user_id).count()
if len(legacy_users) == 1 and current_avatar_count == 0:
legacy_user = legacy_users[0]
_transfer_avatar_ownership(db, legacy_user.huihui_user_id, huihui_user_id)
legacy_user.app_token = ""
legacy_user.huihui_token = ""
db.add(legacy_user)
return user
def _issue_session(
db: Session,
phone: str,
info: dict,
*,
reuse_existing_session: bool = False,
):
"""建/链本地用户并签发本系统会话 token"""
huihui_user_id = info.get("userId", "")
user = _find_or_link_user(db, phone, huihui_user_id)
if phone:
user.phone = phone
if info.get("nickname"):
user.nickname = info["nickname"]
if info.get("avatarUrl"):
user.avatar_url = info["avatarUrl"]
user.huihui_token = info.get("token", "")
if not reuse_existing_session or not user.app_token:
user.app_token = uuid.uuid4().hex
user.last_login_at = datetime.now()
db.add(user)
db.commit()
db.refresh(user)
from services.token_billing import get_or_create_account
get_or_create_account(db, user.id)
return ok({
"token": user.app_token,
"user": user.to_dict(),
"huihui": {
"userId": huihui_user_id,
"nickname": info.get("nickname", ""),
"avatarUrl": info.get("avatarUrl", ""),
},
})
@router.get("/huihui/me")
def me(authorization: str = Header(None), db: Session = Depends(get_db)):
"""当前登录用户信息(Bearer app_token)"""
if not authorization:
return fail("未登录", 401)
token = authorization.replace("Bearer ", "", 1).replace("bearer ", "", 1).strip()
user = db.query(User).filter(User.app_token == token).first()
if not user:
return fail("会话无效或已过期", 401)
return ok(user.to_dict())
@router.post("/huihui/logout")
def logout(authorization: str = Header(None), db: Session = Depends(get_db)):
"""退出登录(作废 app_token)"""
if authorization:
token = authorization.replace("Bearer ", "", 1).replace("bearer ", "", 1).strip()
user = db.query(User).filter(User.app_token == token).first()
if user:
user.app_token = ""
db.commit()
return ok({"success": True})
@@ -0,0 +1,471 @@
import os
import json
import shutil
import time
import uuid
from fastapi import APIRouter, UploadFile, File, Depends, Header, HTTPException
from pydantic import BaseModel
from sqlalchemy.orm import Session
from database import get_db
from models import KnowledgeDoc, QAPair, KnowledgeChunk, Avatar, User
from responses import ok, fail
import embeddings
from services.knowledge_vectorizer import knowledge_vectorizer
router = APIRouter()
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
UPLOAD_DIR = os.path.abspath(os.getenv("UPLOAD_DIR", os.path.join(BASE_DIR, "uploads")))
os.makedirs(UPLOAD_DIR, exist_ok=True)
ALLOWED_EXT = {".md", ".txt", ".pdf", ".doc", ".docx", ".xlsx"}
MAX_UPLOAD_BYTES = 50 * 1024 * 1024
UPLOAD_CHUNK_BYTES = 1024 * 1024
MULTIPART_CHUNK_BYTES = 5 * 1024 * 1024
MULTIPART_ROOT = ".multipart"
MULTIPART_TTL_SECONDS = 24 * 60 * 60
class QAIn(BaseModel):
question: str = ""
answer: str = ""
enabled: bool = True
class EnabledIn(BaseModel):
enabled: bool = True
class MultipartUploadIn(BaseModel):
filename: str
fileSize: int
totalChunks: int
def _validate_document(filename: str, file_size: int):
ext = os.path.splitext(filename or "")[1].lower()
if ext not in ALLOWED_EXT:
return None, f"不支持的文件类型:{ext or '空'},仅支持 md/txt/pdf/doc/docx/xlsx"
if file_size <= 0:
return None, "文件内容不能为空"
if file_size > MAX_UPLOAD_BYTES:
return None, "文件不能超过 50MB"
return ext, ""
def _multipart_dir(avatar_id: str, upload_id: str) -> str:
safe_avatar_id = os.path.basename(avatar_id)
safe_upload_id = os.path.basename(upload_id)
if (
safe_avatar_id != avatar_id
or safe_upload_id != upload_id
or len(upload_id) != 32
or any(character not in "0123456789abcdef" for character in upload_id)
):
raise HTTPException(status_code=400, detail="上传标识无效")
return os.path.join(UPLOAD_DIR, MULTIPART_ROOT, safe_avatar_id, safe_upload_id)
def _purge_stale_multipart_uploads(avatar_id: str):
avatar_upload_root = os.path.join(UPLOAD_DIR, MULTIPART_ROOT, os.path.basename(avatar_id))
if not os.path.isdir(avatar_upload_root):
return
cutoff = time.time() - MULTIPART_TTL_SECONDS
for entry in os.scandir(avatar_upload_root):
if entry.is_dir(follow_symlinks=False) and entry.stat(follow_symlinks=False).st_mtime < cutoff:
shutil.rmtree(entry.path, ignore_errors=True)
def _read_multipart_metadata(avatar_id: str, upload_id: str) -> tuple[str, dict]:
upload_dir = _multipart_dir(avatar_id, upload_id)
metadata_path = os.path.join(upload_dir, "metadata.json")
if not os.path.isfile(metadata_path):
raise HTTPException(status_code=404, detail="上传任务不存在或已过期")
with open(metadata_path, "r", encoding="utf-8") as stream:
return upload_dir, json.load(stream)
def _create_knowledge_doc(db: Session, avatar_id: str, filename: str, ext: str, file_size: int, stored: str):
doc = KnowledgeDoc(
id=uuid.uuid4().hex,
avatar_id=avatar_id,
filename=filename,
file_type=ext.lstrip("."),
file_size=file_size,
file_url=f"/api/files/{avatar_id}/{stored}",
status="parsing",
index_stage="queued",
index_progress=0,
)
db.add(doc)
db.commit()
db.refresh(doc)
knowledge_vectorizer.enqueue(doc.id)
return doc
def _doc_payload(doc: KnowledgeDoc) -> dict:
payload = doc.to_dict()
stored_name = os.path.basename(doc.file_url or "")
stored_path = os.path.join(UPLOAD_DIR, doc.avatar_id, stored_name)
payload["filePresent"] = bool(stored_name and os.path.isfile(stored_path))
return payload
def _resolve_user(authorization: str | None, db: Session):
if not authorization:
return None
token = authorization.replace("Bearer ", "", 1).replace("bearer ", "", 1).strip()
return db.query(User).filter(User.app_token == token).first()
def _require_owned_avatar(db: Session, avatar_id: str, authorization: str | None):
avatar = db.query(Avatar).filter(Avatar.id == avatar_id).first()
if not avatar:
raise HTTPException(status_code=404, detail="avatar not found")
user = _resolve_user(authorization, db)
if not user:
raise HTTPException(status_code=401, detail="未登录")
if avatar.owner_id and avatar.owner_id != user.huihui_user_id:
raise HTTPException(status_code=403, detail="无权访问该分身")
return avatar
# ---------------- Documents ----------------
@router.get("/avatar/{avatar_id}/knowledge/docs")
def list_docs(avatar_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
docs = (
db.query(KnowledgeDoc)
.filter(KnowledgeDoc.avatar_id == avatar_id)
.order_by(KnowledgeDoc.created_at.desc())
.all()
)
return ok([_doc_payload(d) for d in docs])
@router.post("/avatar/{avatar_id}/knowledge/docs")
async def upload_doc(avatar_id: str, file: UploadFile = File(...), authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
ext, validation_error = _validate_document(file.filename or "", 1)
if validation_error:
return fail(validation_error, code=400)
avatar_dir = os.path.join(UPLOAD_DIR, avatar_id)
os.makedirs(avatar_dir, exist_ok=True)
stored = f"{uuid.uuid4().hex}{ext}"
path = os.path.join(avatar_dir, stored)
file_size = 0
try:
# Stream large files to disk so a 100MB upload does not occupy 100MB RAM.
with open(path, "wb") as f:
while chunk := await file.read(UPLOAD_CHUNK_BYTES):
file_size += len(chunk)
if file_size > MAX_UPLOAD_BYTES:
raise ValueError("文件不能超过 50MB")
f.write(chunk)
except ValueError as exc:
if os.path.exists(path):
os.remove(path)
return fail(str(exc), code=400)
if file_size == 0:
if os.path.exists(path):
os.remove(path)
return fail("文件内容不能为空", code=400)
doc = _create_knowledge_doc(db, avatar_id, file.filename or stored, ext, file_size, stored)
return ok(_doc_payload(doc))
@router.post("/avatar/{avatar_id}/knowledge/uploads")
def create_multipart_upload(
avatar_id: str,
body: MultipartUploadIn,
authorization: str = Header(None),
db: Session = Depends(get_db),
):
_require_owned_avatar(db, avatar_id, authorization)
ext, validation_error = _validate_document(body.filename, body.fileSize)
if validation_error:
return fail(validation_error, code=400)
expected_chunks = (body.fileSize + MULTIPART_CHUNK_BYTES - 1) // MULTIPART_CHUNK_BYTES
if body.totalChunks != expected_chunks:
return fail("文件分片数量不正确", code=400)
_purge_stale_multipart_uploads(avatar_id)
upload_id = uuid.uuid4().hex
upload_dir = _multipart_dir(avatar_id, upload_id)
os.makedirs(upload_dir, exist_ok=False)
metadata = {
"filename": body.filename,
"fileSize": body.fileSize,
"totalChunks": body.totalChunks,
"extension": ext,
}
with open(os.path.join(upload_dir, "metadata.json"), "w", encoding="utf-8") as stream:
json.dump(metadata, stream, ensure_ascii=False)
return ok({"uploadId": upload_id, "chunkSize": MULTIPART_CHUNK_BYTES})
@router.post("/avatar/{avatar_id}/knowledge/uploads/{upload_id}/chunks/{chunk_index}")
async def upload_multipart_chunk(
avatar_id: str,
upload_id: str,
chunk_index: int,
file: UploadFile = File(...),
authorization: str = Header(None),
db: Session = Depends(get_db),
):
_require_owned_avatar(db, avatar_id, authorization)
upload_dir, metadata = _read_multipart_metadata(avatar_id, upload_id)
total_chunks = int(metadata["totalChunks"])
if chunk_index < 0 or chunk_index >= total_chunks:
return fail("文件分片序号不正确", code=400)
expected_size = min(
MULTIPART_CHUNK_BYTES,
int(metadata["fileSize"]) - chunk_index * MULTIPART_CHUNK_BYTES,
)
part_path = os.path.join(upload_dir, f"{chunk_index}.part")
temporary_path = f"{part_path}.uploading"
received = 0
try:
with open(temporary_path, "wb") as stream:
while chunk := await file.read(UPLOAD_CHUNK_BYTES):
received += len(chunk)
if received > expected_size:
raise ValueError("文件分片大小不正确")
stream.write(chunk)
if received != expected_size:
raise ValueError("文件分片大小不正确")
os.replace(temporary_path, part_path)
except ValueError as exc:
if os.path.exists(temporary_path):
os.remove(temporary_path)
return fail(str(exc), code=400)
return ok({"chunkIndex": chunk_index, "uploadedBytes": received})
@router.post("/avatar/{avatar_id}/knowledge/uploads/{upload_id}/complete")
def complete_multipart_upload(
avatar_id: str,
upload_id: str,
authorization: str = Header(None),
db: Session = Depends(get_db),
):
_require_owned_avatar(db, avatar_id, authorization)
upload_dir, metadata = _read_multipart_metadata(avatar_id, upload_id)
total_chunks = int(metadata["totalChunks"])
part_paths = [os.path.join(upload_dir, f"{index}.part") for index in range(total_chunks)]
if not all(os.path.isfile(path) for path in part_paths):
return fail("文件分片尚未上传完整", code=400)
if sum(os.path.getsize(path) for path in part_paths) != int(metadata["fileSize"]):
return fail("文件分片总大小不正确", code=400)
avatar_dir = os.path.join(UPLOAD_DIR, avatar_id)
os.makedirs(avatar_dir, exist_ok=True)
stored = f"{uuid.uuid4().hex}{metadata['extension']}"
final_path = os.path.join(avatar_dir, stored)
temporary_path = f"{final_path}.assembling"
try:
with open(temporary_path, "wb") as output:
for part_path in part_paths:
with open(part_path, "rb") as source:
shutil.copyfileobj(source, output, UPLOAD_CHUNK_BYTES)
os.replace(temporary_path, final_path)
doc = _create_knowledge_doc(
db,
avatar_id,
metadata["filename"],
metadata["extension"],
int(metadata["fileSize"]),
stored,
)
except Exception:
if os.path.exists(temporary_path):
os.remove(temporary_path)
raise
shutil.rmtree(upload_dir, ignore_errors=True)
return ok(_doc_payload(doc))
@router.post("/avatar/{avatar_id}/knowledge/docs/{doc_id}/retry")
def retry_doc(avatar_id: str, doc_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
doc = db.query(KnowledgeDoc).filter(
KnowledgeDoc.id == doc_id, KnowledgeDoc.avatar_id == avatar_id
).first()
if not doc:
return fail("文档不存在", code=404)
if doc.vectorized and doc.status == "ready":
return ok(_doc_payload(doc))
stored_name = os.path.basename(doc.file_url or "")
if not stored_name or not os.path.isfile(os.path.join(UPLOAD_DIR, avatar_id, stored_name)):
return fail("原文件不可用,请重新上传", code=400)
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == doc.id).delete()
doc.status = "parsing"
doc.vectorized = False
doc.embedding_model = ""
doc.chunk_count = 0
doc.vectorized_at = None
doc.error_message = ""
doc.index_stage = "queued"
doc.index_progress = 0
db.commit()
db.refresh(doc)
knowledge_vectorizer.enqueue(doc.id)
return ok(_doc_payload(doc))
@router.delete("/avatar/{avatar_id}/knowledge/docs/{doc_id}")
def delete_doc(avatar_id: str, doc_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
doc = (
db.query(KnowledgeDoc)
.filter(KnowledgeDoc.id == doc_id, KnowledgeDoc.avatar_id == avatar_id)
.first()
)
if not doc:
return fail("文档不存在", code=404)
# 级联删除切片
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == doc_id).delete()
try:
fp = os.path.join(UPLOAD_DIR, avatar_id, os.path.basename(doc.file_url))
if os.path.exists(fp):
os.remove(fp)
except Exception:
pass
db.delete(doc)
db.commit()
return ok({"id": doc_id})
# ---------------- 向量检索 ----------------
@router.get("/avatar/{avatar_id}/knowledge/search")
def search_knowledge(avatar_id: str, q: str = "", top_k: int = 5, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
q = (q or "").strip()
if not q:
return ok([])
chunks = (
db.query(KnowledgeChunk)
.filter(KnowledgeChunk.avatar_id == avatar_id)
.all()
)
if not chunks:
return ok([])
qvec = embeddings.embed([q])[0]
scored = []
for c in chunks:
try:
vec = json.loads(c.vector)
except Exception:
continue
scored.append((embeddings.cosine(qvec, vec), c))
scored.sort(key=lambda x: x[0], reverse=True)
results = []
for score, c in scored[: max(1, top_k)]:
doc = db.query(KnowledgeDoc).filter(KnowledgeDoc.id == c.doc_id).first()
snippet = c.content[:120] + ("…" if len(c.content) > 120 else "")
results.append(
{
"docId": c.doc_id,
"filename": doc.filename if doc else "",
"fileType": doc.file_type if doc else "",
"snippet": snippet,
"score": round(score, 4),
}
)
return ok(results)
# ---------------- Standard Q&A pairs ----------------
@router.get("/avatar/{avatar_id}/knowledge/qa")
def list_qa(avatar_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
items = (
db.query(QAPair)
.filter(QAPair.avatar_id == avatar_id)
.order_by(QAPair.created_at.desc())
.all()
)
return ok([q.to_dict() for q in items])
@router.post("/avatar/{avatar_id}/knowledge/qa")
def create_qa(avatar_id: str, body: QAIn, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
q = QAPair(
avatar_id=avatar_id,
question=body.question,
answer=body.answer,
enabled=body.enabled,
)
db.add(q)
db.commit()
db.refresh(q)
return ok(q.to_dict())
@router.put("/avatar/{avatar_id}/knowledge/qa/{qa_id}")
def update_qa(avatar_id: str, qa_id: str, body: QAIn, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
q = (
db.query(QAPair)
.filter(QAPair.id == qa_id, QAPair.avatar_id == avatar_id)
.first()
)
if not q:
return fail("问答对不存在", code=404)
q.question = body.question
q.answer = body.answer
q.enabled = body.enabled
db.commit()
db.refresh(q)
return ok(q.to_dict())
@router.put("/avatar/{avatar_id}/knowledge/qa/{qa_id}/enabled")
def set_qa_enabled(avatar_id: str, qa_id: str, body: EnabledIn, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
q = (
db.query(QAPair)
.filter(QAPair.id == qa_id, QAPair.avatar_id == avatar_id)
.first()
)
if not q:
return fail("问答对不存在", code=404)
q.enabled = bool(body.enabled)
db.commit()
db.refresh(q)
return ok(q.to_dict())
@router.delete("/avatar/{avatar_id}/knowledge/qa/{qa_id}")
def delete_qa(avatar_id: str, qa_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
_require_owned_avatar(db, avatar_id, authorization)
q = (
db.query(QAPair)
.filter(QAPair.id == qa_id, QAPair.avatar_id == avatar_id)
.first()
)
if not q:
return fail("问答对不存在", code=404)
db.delete(q)
db.commit()
return ok({"id": qa_id})
# ---------------- HuiHui user profile (mock; plug real interface via HUIHUI_USER_API) ----------------
@router.get("/user/profile")
def user_profile():
# 接入真实会会接口:设置环境变量 HUIHUI_USER_API 后在此请求并映射字段
api = os.getenv("HUIHUI_USER_API")
if api:
# TODO: 调用会会用户接口,返回 { userId, nickname, avatarUrl }
pass
return ok({
"userId": "hh_10001",
"nickname": "会会用户",
"avatarUrl": "https://api.dicebear.com/7.x/initials/svg?seed=HuiHui&backgroundColor=F97316",
})
@@ -0,0 +1,35 @@
from fastapi import APIRouter, Depends, Body
from sqlalchemy.orm import Session
from database import get_db
from models import Organization
from responses import ok, fail
router = APIRouter(tags=["组织"])
@router.get("/organizations")
def list_orgs(page: int = 1, limit: int = 20, db: Session = Depends(get_db)):
q = db.query(Organization)
total = q.count()
items = (
q.order_by(Organization.created_at.desc())
.offset((page - 1) * limit)
.limit(limit)
.all()
)
return ok({"data": [o.to_dict() for o in items], "total": total})
@router.post("/organizations")
def create_org(payload: dict = Body(...), db: Session = Depends(get_db)):
o = Organization(
name=payload.get("name", "未命名组织"),
description=payload.get("desc", "") or payload.get("description", ""),
emoji=payload.get("emoji", "🏢"),
org_type=payload.get("type", "") or payload.get("orgType", "team"),
)
db.add(o)
db.commit()
db.refresh(o)
return ok(o.to_dict())
@@ -0,0 +1,153 @@
"""数字分身 BOXIM 单聊接管 API。"""
from datetime import datetime, timedelta
from fastapi import APIRouter, Body, Depends, Header
from sqlalchemy.orm import Session
from database import get_db
from models import TakeoverCursor, TakeoverReplyTask, User
from responses import fail, ok
from routers.authorizations import (
DEFAULT_TAKEOVER_DELAY_SECONDS,
MAX_TAKEOVER_DELAY_SECONDS,
MIN_TAKEOVER_DELAY_SECONDS,
_require_authorization,
_stored_takeover_delay,
)
from routers.avatars import _require_owned_avatar
router = APIRouter(tags=["分身接管"])
BOXIM_STATUS_FRESH_SECONDS = 60
def _delay_label(seconds: int) -> str:
if seconds % 60 == 0:
return f"{seconds // 60} 分钟"
return f"{seconds} 秒"
@router.get("/avatar/{avatar_id}/takeover/status")
def get_takeover_status(
avatar_id: str,
authorization: str = Header(None),
db: Session = Depends(get_db),
):
avatar = _require_owned_avatar(db, avatar_id, authorization)
permissions = (avatar.config or {}).get("authorizationPermissions", [])
enabled = isinstance(permissions, list) and "takeover" in permissions
reply_delay_seconds = _stored_takeover_delay(avatar)
user = db.query(User).filter(User.huihui_user_id == avatar.owner_id).first()
cursor = db.query(TakeoverCursor).filter(TakeoverCursor.avatar_id == avatar.id).first()
pending_count = (
db.query(TakeoverReplyTask)
.filter(
TakeoverReplyTask.avatar_id == avatar.id,
TakeoverReplyTask.status.in_(("pending", "generating", "ready", "sending")),
)
.count()
)
if cursor and cursor.last_error:
status, message = "error", cursor.last_error
elif not enabled:
status, message = "disabled", "主动接管未开启"
elif not user or not user.huihui_token:
status, message = "needs_login", "请重新登录会会生产账号以连接 BOXIM"
elif (
cursor
and cursor.initialized
and cursor.last_polled_at
# BOXIM offline-message reads can long-poll for about 20 seconds.
and cursor.last_polled_at
>= datetime.utcnow() - timedelta(seconds=BOXIM_STATUS_FRESH_SECONDS)
):
status, message = (
"ready",
f"BOXIM 已连接,收到私聊消息 {_delay_label(reply_delay_seconds)}后自动回复",
)
else:
status, message = "connecting", "正在连接 BOXIM"
return ok(
{
"enabled": enabled,
"status": status,
"message": message,
"pendingCount": pending_count,
"takeoverReplyDelaySeconds": reply_delay_seconds,
"lastPolledAt": cursor.last_polled_at.isoformat() if cursor and cursor.last_polled_at else None,
}
)
def _has(payload: dict, camel_key: str, snake_key: str) -> bool:
return camel_key in payload or snake_key in payload
def _read(payload: dict, camel_key: str, snake_key: str, default=None):
if camel_key in payload:
return payload[camel_key]
if snake_key in payload:
return payload[snake_key]
return default
@router.put("/avatar/{avatar_id}/authorizations/takeover")
def update_takeover_config(
avatar_id: str,
payload: dict = Body(...),
authorization: str = Header(None),
db: Session = Depends(get_db),
):
_require_owned_avatar(db, avatar_id, authorization)
auth_id = _read(payload, "authorizationId", "authorization_id")
if not auth_id:
return fail("缺少 authorization_id", 400)
auth = _require_authorization(db, avatar_id, str(auth_id))
enabled = bool(auth.takeover_enabled)
mode = auth.takeover_mode or "immediate"
delay = auth.takeover_delay_seconds or DEFAULT_TAKEOVER_DELAY_SECONDS
if _has(payload, "takeoverEnabled", "takeover_enabled"):
raw_enabled = _read(payload, "takeoverEnabled", "takeover_enabled")
if not isinstance(raw_enabled, bool):
return fail("takeover_enabled 必须是布尔值", 400)
enabled = raw_enabled
if _has(payload, "takeoverMode", "takeover_mode"):
mode = _read(payload, "takeoverMode", "takeover_mode")
if mode not in ("immediate", "delayed"):
return fail("takeover_mode 必须是 immediate 或 delayed", 400)
if _has(payload, "takeoverDelaySeconds", "takeover_delay_seconds"):
delay = _read(payload, "takeoverDelaySeconds", "takeover_delay_seconds")
if (
isinstance(delay, bool)
or not isinstance(delay, int)
or not MIN_TAKEOVER_DELAY_SECONDS <= delay <= MAX_TAKEOVER_DELAY_SECONDS
):
return fail("延迟时间需在 3 秒到 24 小时之间", 400)
if enabled and auth.target_type != "user":
return fail("本期仅支持对会会用户开启单聊接管", 400)
if enabled and auth.status != "active":
return fail("请先启用该授权,再开启聊天接管", 400)
permissions = list(auth.permissions or [])
if enabled:
if "chat" not in permissions and "reply" not in permissions:
permissions.append("chat")
if "takeover" not in permissions:
permissions.append("takeover")
else:
permissions = [permission for permission in permissions if permission != "takeover"]
auth.permissions = permissions
auth.takeover_enabled = enabled
auth.takeover_mode = mode
auth.takeover_delay_seconds = delay
db.commit()
db.refresh(auth)
return ok(auth.to_dict(), "接管配置已保存")
@@ -0,0 +1,346 @@
import hashlib
import hmac
import json
import os
import uuid
from datetime import datetime
from decimal import Decimal, InvalidOperation, ROUND_HALF_UP
from urllib.parse import parse_qs
from fastapi import APIRouter, Body, Depends, Header, HTTPException, Request
from sqlalchemy import func
from sqlalchemy.orm import Session
from database import get_db
from models import TokenAccount, TokenPaymentOrder, TokenPlan, TokenUsage, User
from responses import fail, ok
from services.huihui_payment import HuihuiPaymentClient, HuihuiPaymentError
from services.token_billing import DEFAULT_TOKEN_GRANT, get_or_create_account
router = APIRouter(tags=["Token"])
PAYMENT_METHODS = {"wechat": "WECHAT", "alipay": "ALIPAY"}
PAYMENT_SCENES = {"APP", "LITE", "JSAPI"}
SUCCESS_STATUSES = {"SUCCESS", "SUCCEEDED", "PAID", "COMPLETED", "TRADE_SUCCESS"}
FAILED_STATUSES = {"FAIL", "FAILED", "CLOSED", "CANCELLED", "CANCELED", "EXPIRED"}
def _require_user(authorization: str | None, db: Session) -> User:
if not authorization:
raise HTTPException(status_code=401, detail="未登录")
token = authorization.replace("Bearer ", "", 1).replace("bearer ", "", 1).strip()
user = db.query(User).filter(User.app_token == token).first()
if not user:
raise HTTPException(status_code=401, detail="会话无效或已过期")
return user
def _payment_client() -> HuihuiPaymentClient:
return HuihuiPaymentClient({
"HUIHUI_PAYMENT_BASE_URL": os.getenv(
"HUIHUI_PAYMENT_BASE_URL", "https://open.99hui.com/api/payment-v3"
),
"HUIHUI_APP_ID": os.getenv("HUIHUI_APP_ID", ""),
"HUIHUI_ACCESS_ID": os.getenv("HUIHUI_ACCESS_ID", ""),
"HUIHUI_ACCESS_SECRET": os.getenv("HUIHUI_ACCESS_SECRET", ""),
"HUIHUI_PAYMENT_TIMEOUT_SECONDS": os.getenv("HUIHUI_PAYMENT_TIMEOUT_SECONDS", "30"),
})
def _callback_url(order_no: str) -> str:
base = os.getenv(
"HUIHUI_PAYMENT_CALLBACK_BASE_URL", "https://digital.99hui.com"
).rstrip("/")
secret = os.getenv("HUIHUI_PAYMENT_CALLBACK_SECRET", "").strip()
if len(secret) < 16:
raise HuihuiPaymentError("会会支付回调密钥未配置")
signature = hmac.new(secret.encode(), order_no.encode(), hashlib.sha256).hexdigest()
return f"{base}/api/token/payment/callback/{order_no}/{signature}"
def _price_cents(price: float) -> int:
return int(
(Decimal(str(price)) * Decimal("100")).quantize(
Decimal("1"), rounding=ROUND_HALF_UP
)
)
def _payment_payload(order: TokenPaymentOrder, account: TokenAccount) -> dict:
return {**order.to_dict(), "balance": account.balance}
def _nested_payload(value):
if isinstance(value, str):
text = value.strip()
if text[:1] in ("{", "["):
try:
return _nested_payload(json.loads(text))
except (TypeError, ValueError):
return value
return value
if isinstance(value, list):
return [_nested_payload(item) for item in value]
if isinstance(value, dict):
return {key: _nested_payload(item) for key, item in value.items()}
return value
def _find_value(payload, *names):
expected = {name.lower() for name in names}
if isinstance(payload, dict):
for key, value in payload.items():
if key.lower() in expected and value not in (None, ""):
return value
for value in payload.values():
found = _find_value(value, *names)
if found not in (None, ""):
return found
elif isinstance(payload, list):
for value in payload:
found = _find_value(value, *names)
if found not in (None, ""):
return found
return None
def _callback_amount_cents(payload) -> int | None:
value = _find_value(
payload,
"actualAmt",
"payAmt",
"masterOrderAmt",
"orderAmt",
"amount",
"totalAmount",
)
if value in (None, ""):
return None
try:
return int(
(Decimal(str(value)) * Decimal("100")).quantize(
Decimal("1"), rounding=ROUND_HALF_UP
)
)
except (InvalidOperation, TypeError, ValueError):
return None
@router.get("/token/balance")
def balance(authorization: str = Header(None), db: Session = Depends(get_db)):
user = _require_user(authorization, db)
acc = get_or_create_account(db, user.id)
return ok({
"balance": acc.balance,
"totalGranted": acc.total_granted,
"totalConsumed": acc.total_consumed,
})
@router.get("/token/plans")
def plans(authorization: str = Header(None), db: Session = Depends(get_db)):
_require_user(authorization, db)
items = db.query(TokenPlan).order_by(TokenPlan.price.asc()).all()
return ok([p.to_dict() for p in items])
# 积分只会在会会支付回调确认成功后到账。
@router.post("/token/charge")
def charge(payload: dict = Body(...), authorization: str = Header(None), db: Session = Depends(get_db)):
user = _require_user(authorization, db)
plan = db.query(TokenPlan).filter(TokenPlan.id == payload.get("planId")).first()
if not plan:
return fail("套餐不存在", 404)
payment_method = str(payload.get("paymentMethod") or "").lower()
pay_type = PAYMENT_METHODS.get(payment_method)
if not pay_type:
return fail("请选择正确的支付方式", 400)
pay_way = str(payload.get("payScene") or "APP").upper()
if pay_way not in PAYMENT_SCENES:
return fail("当前支付场景不受支持", 400)
cents = _price_cents(plan.price)
order = TokenPaymentOrder(
order_no=f"AV{datetime.utcnow().strftime('%Y%m%d%H%M%S')}{uuid.uuid4().hex[:12].upper()}",
user_id=user.id,
plan_id=plan.id,
payment_method=payment_method,
pay_type=pay_type,
pay_way=pay_way,
points_amount=plan.amount,
price_cents=cents,
status="pending",
)
db.add(order)
db.commit()
try:
callback_url = _callback_url(order.order_no)
except HuihuiPaymentError as exc:
order.status = "failed"
order.failure_reason = str(exc)
db.commit()
return fail(str(exc), 503)
try:
result = _payment_client().create_payment(
huihui_token=user.huihui_token,
huihui_user_id=user.huihui_user_id,
real_name=user.nickname,
order_no=order.order_no,
amount=f"{cents / 100:.2f}",
points_amount=plan.amount,
pay_type=pay_type,
pay_way=pay_way,
callback_url=callback_url,
)
except HuihuiPaymentError as exc:
order.status = "failed"
order.failure_reason = str(exc)[:500]
db.commit()
return fail(str(exc), 502)
db.refresh(order)
if order.status != "paid":
order.provider_order_id = str(result.get("orderId") or "")
order.provider_order_no = str(result.get("orderNo") or "")
order.provider_status = str(result.get("status") or "pending")
message = result.get("payMessage") or ""
order.pay_message = (
json.dumps(message, ensure_ascii=False)
if isinstance(message, (dict, list))
else str(message)
)
if order.provider_status.upper() in FAILED_STATUSES:
order.status = "failed"
order.failure_reason = str(result.get("bankReturnMsg") or "支付下单失败")[:500]
db.commit()
return ok(_payment_payload(order, get_or_create_account(db, user.id)))
@router.get("/token/payment/{order_id}")
def payment_status(order_id: str, authorization: str = Header(None), db: Session = Depends(get_db)):
user = _require_user(authorization, db)
order = db.query(TokenPaymentOrder).filter(
TokenPaymentOrder.id == order_id,
TokenPaymentOrder.user_id == user.id,
).first()
if not order:
return fail("支付订单不存在", 404)
return ok(_payment_payload(order, get_or_create_account(db, user.id)))
@router.post("/token/payment/callback/{order_no}/{callback_signature}")
async def payment_callback(
order_no: str,
callback_signature: str,
request: Request,
db: Session = Depends(get_db),
):
secret = os.getenv("HUIHUI_PAYMENT_CALLBACK_SECRET", "").strip()
expected = hmac.new(secret.encode(), order_no.encode(), hashlib.sha256).hexdigest()
if len(secret) < 16 or not hmac.compare_digest(callback_signature, expected):
raise HTTPException(status_code=404, detail="Not found")
content_type = request.headers.get("content-type", "").lower()
if "application/json" in content_type:
try:
payload = await request.json()
except ValueError:
return fail("支付回调格式不正确", 400)
else:
raw = (await request.body()).decode("utf-8", errors="replace")
payload = {key: values[-1] for key, values in parse_qs(raw).items()}
payload = _nested_payload(payload)
payload_order_no = str(_find_value(
payload,
"masterOrderNo",
"master_order_no",
"orderNo",
"order_no",
"bizOrderNo",
) or "").strip()
if payload_order_no and payload_order_no != order_no:
return fail("支付回调订单号不匹配", 422)
order = db.query(TokenPaymentOrder).filter(TokenPaymentOrder.order_no == order_no).first()
if not order:
return fail("支付订单不存在", 404)
if order.status == "paid":
return ok({"received": True, "duplicate": True})
provider_status = str(_find_value(
payload, "status", "payStatus", "tradeStatus", "paymentStatus"
) or "").upper()
order.provider_status = provider_status
if provider_status not in SUCCESS_STATUSES:
if provider_status in FAILED_STATUSES:
order.status = "failed"
order.failure_reason = str(
_find_value(payload, "message", "errorMsg", "failReason") or "支付失败"
)[:500]
db.commit()
return ok({"received": True, "paid": False})
paid_cents = _callback_amount_cents(payload)
if paid_cents is None or paid_cents != order.price_cents:
order.failure_reason = "支付回调金额不匹配"
db.commit()
return fail("支付金额不匹配", 422)
updated = db.query(TokenPaymentOrder).filter(
TokenPaymentOrder.id == order.id,
TokenPaymentOrder.status != "paid",
).update({
TokenPaymentOrder.status: "paid",
TokenPaymentOrder.provider_status: provider_status,
TokenPaymentOrder.paid_at: datetime.utcnow(),
TokenPaymentOrder.failure_reason: "",
}, synchronize_session=False)
if updated:
account = db.query(TokenAccount).filter(TokenAccount.user_id == order.user_id).first()
if account is None:
account = TokenAccount(
user_id=order.user_id,
balance=DEFAULT_TOKEN_GRANT,
total_granted=DEFAULT_TOKEN_GRANT,
total_consumed=0,
)
db.add(account)
db.flush()
account.balance = int(account.balance or 0) + order.points_amount
account.total_granted = int(account.total_granted or 0) + order.points_amount
db.commit()
return ok({"received": True, "paid": True})
@router.get("/token/usage")
def usage(authorization: str = Header(None), db: Session = Depends(get_db)):
user = _require_user(authorization, db)
rows = (
db.query(
TokenUsage.avatar_id,
TokenUsage.source,
func.sum(TokenUsage.prompt_tokens),
func.sum(TokenUsage.completion_tokens),
func.sum(TokenUsage.total_tokens),
func.count(TokenUsage.id),
)
.filter(TokenUsage.user_id == user.id, TokenUsage.status == "completed")
.group_by(TokenUsage.avatar_id, TokenUsage.source)
.all()
)
return ok([
{
"avatarId": avatar_id,
"source": source,
"promptTokens": int(prompt_tokens or 0),
"completionTokens": int(completion_tokens or 0),
"totalTokens": int(total_tokens or 0),
"requestCount": int(request_count or 0),
}
for avatar_id, source, prompt_tokens, completion_tokens, total_tokens, request_count in rows
])
@@ -0,0 +1,210 @@
"""Client for Huihui's self-hosted BOXIM production APIs."""
import hashlib
import random
import secrets
import string
import time
from datetime import datetime, timedelta, timezone
from typing import Any
import httpx
_CN_TZ = timezone(timedelta(hours=8))
class BoxIMError(RuntimeError):
def __init__(self, message: str, *, code: Any = None, auth_error: bool = False):
super().__init__(message)
self.code = code
self.auth_error = auth_error
class BoxIMClient:
"""Exchange Huihui credentials and call BOXIM's private-message API."""
def __init__(self, config: dict):
self.platform_base_url = config.get(
"HUIHUI_PLATFORM_BASE_URL", "https://open.99hui.com/api"
).rstrip("/")
self.im_base_url = config.get(
"BOXIM_API_BASE_URL", "https://im.99hui.com/api"
).rstrip("/")
self.app_id = config.get("HUIHUI_APP_ID", "")
self.access_id = config.get("HUIHUI_ACCESS_ID", "")
self.access_secret = config.get("HUIHUI_ACCESS_SECRET", "")
self.timeout = float(config.get("BOXIM_TIMEOUT_SECONDS", 20))
def _build_sign_params(self, extra: dict | None = None) -> dict:
"""Build the same signed form used by Huihui's current production app."""
params = {
"appId": self.app_id,
"accessId": self.access_id,
"nonce": "".join(random.choices(string.ascii_lowercase + string.digits, k=12)),
"timestamp": datetime.now(_CN_TZ).strftime("%Y%m%d%H%M%S"),
"signType": "MD5",
"signVersion": "1.0",
**(extra or {}),
}
params.pop("accessSecret", None)
params.pop("signature", None)
sign_parts = []
for key in sorted(params):
value = params[key]
if value in (None, "", []):
continue
if isinstance(value, list):
continue
sign_parts.append(f"{key}={value}")
sign_source = "&".join(sign_parts) + f"&accessSecret={self.access_secret}"
params["signature"] = hashlib.md5(sign_source.encode("utf-8")).hexdigest().upper()
return params
@staticmethod
def _response_payload(response: httpx.Response) -> dict:
try:
payload = response.json()
except ValueError as exc:
raise BoxIMError("BOXIM 返回了无效响应") from exc
if not isinstance(payload, dict):
raise BoxIMError("BOXIM 返回格式不正确")
return payload
async def exchange_access_token(self, huihui_token: str) -> dict:
"""Exchange a production Huihui token for a BOXIM access token."""
if not huihui_token:
raise BoxIMError("缺少会会登录凭证", auth_error=True)
if not (self.app_id and self.access_id and self.access_secret):
raise BoxIMError("会会开放平台凭证未配置", auth_error=True)
headers = {
"Authorization": f"Bearer {huihui_token}",
"appId": self.app_id,
"windowAppId": self.app_id,
}
async with httpx.AsyncClient(timeout=self.timeout, follow_redirects=True) as client:
response = await client.post(
f"{self.platform_base_url}/im/box/netease",
headers=headers,
data=self._build_sign_params(),
)
payload = self._response_payload(response)
data = payload.get("data") or {}
code = payload.get("code")
if response.status_code >= 400 or code not in (0, 200, "0", "200"):
raise BoxIMError(
payload.get("message") or "BOXIM 授权失败",
code=code or response.status_code,
auth_error=response.status_code in (400, 401, 403)
or code in (
400,
401,
40100,
40101,
403,
"400",
"401",
"40100",
"40101",
"403",
),
)
if not data.get("accessToken"):
raise BoxIMError("会会未返回 BOXIM 访问凭证", auth_error=True)
return data
async def _request(
self,
method: str,
path: str,
access_token: str,
*,
params: dict | None = None,
json: dict | None = None,
) -> Any:
headers = {"accessToken": access_token}
async with httpx.AsyncClient(timeout=self.timeout) as client:
response = await client.request(
method,
f"{self.im_base_url}{path}",
headers=headers,
params=params,
json=json,
)
payload = self._response_payload(response)
code = payload.get("code")
if response.status_code >= 400 or code not in (200, "200"):
raise BoxIMError(
payload.get("message") or "BOXIM 请求失败",
code=code or response.status_code,
auth_error=response.status_code in (400, 401, 403)
or code in (400, 401, 40100, 40101, 403, "400", "401", "40100", "40101", "403"),
)
return payload.get("data")
async def get_self(self, access_token: str) -> dict:
data = await self._request("GET", "/user/self", access_token)
if not isinstance(data, dict) or data.get("id") is None:
raise BoxIMError("BOXIM 未返回当前用户信息")
return data
async def fetch_private_messages(self, access_token: str, min_id: str = "0") -> list[dict]:
data = await self._request(
"GET",
"/message/private/loadOfflineMessage",
access_token,
params={"minId": str(min_id or "0")},
)
if data is None:
return []
if not isinstance(data, list):
raise BoxIMError("BOXIM 私聊消息格式不正确")
return [item for item in data if isinstance(item, dict)]
async def mark_private_messages_read(
self,
access_token: str,
friend_id: int | str,
message_id: int | str,
) -> None:
"""Mark one private conversation read through its latest received message."""
friend_id_text = str(friend_id).strip()
message_id_text = str(message_id).strip()
if not friend_id_text.isdigit() or not message_id_text.isdigit():
raise BoxIMError("BOXIM 已读回执参数不正确")
await self._request(
"PUT",
"/message/private/readed",
access_token,
params={
"friendId": int(friend_id_text),
"messageId": int(message_id_text),
},
)
async def send_private_message(
self,
access_token: str,
peer_id: str,
content: str,
*,
local_id: int | str | None = None,
) -> dict:
local_id = int(local_id or (int(time.time() * 1000) * 1000 + secrets.randbelow(1000)))
data = await self._request(
"POST",
"/message/private/send",
access_token,
json={
"localId": local_id,
"recvId": int(peer_id) if str(peer_id).isdigit() else peer_id,
"content": content,
"type": 0,
"receipt": False,
"atUserIds": [],
},
)
if not isinstance(data, dict):
raise BoxIMError("BOXIM 未返回发送结果")
return data
@@ -0,0 +1,151 @@
"""Parse and safely download image payloads from BOXIM private messages."""
import ipaddress
import json
import os
import socket
from dataclasses import dataclass
from pathlib import PurePosixPath
from urllib.parse import unquote, urljoin, urlsplit
import httpx
MAX_REDIRECTS = 3
class BoxIMImageError(RuntimeError):
pass
@dataclass(frozen=True)
class DownloadedBoxIMImage:
content: bytes
filename: str
mime_type: str
source_url: str
def parse_boxim_image_url(content: str, *, base_url: str = "") -> str:
try:
payload = json.loads(content or "")
except (TypeError, ValueError) as exc:
raise BoxIMImageError("BOXIM 图片消息格式无效") from exc
if not isinstance(payload, dict):
raise BoxIMImageError("BOXIM 图片消息格式无效")
value = payload.get("originUrl") or payload.get("thumbUrl") or payload.get("url")
if not isinstance(value, str) or not value.strip():
raise BoxIMImageError("BOXIM 图片消息缺少图片地址")
value = value.strip()
if value.startswith("/"):
if not base_url:
raise BoxIMImageError("BOXIM 图片地址不完整")
value = urljoin(f"{base_url.rstrip('/')}/", value)
return value
def _configured_hosts(name: str) -> set[str]:
return {
value.strip().lower().rstrip(".")
for value in os.getenv(name, "").split(",")
if value.strip()
}
def _host_matches(host: str, configured: set[str]) -> bool:
return any(host == value or host.endswith(f".{value}") for value in configured)
def _resolved_addresses(host: str, port: int) -> set[ipaddress.IPv4Address | ipaddress.IPv6Address]:
try:
return {
ipaddress.ip_address(item[4][0])
for item in socket.getaddrinfo(host, port, type=socket.SOCK_STREAM)
}
except (OSError, ValueError) as exc:
raise BoxIMImageError("BOXIM 图片地址无法解析") from exc
def _is_safe_remote_url(url: str) -> None:
parsed = urlsplit(url)
scheme = parsed.scheme.lower()
allow_http = os.getenv("BOXIM_IMAGE_ALLOW_HTTP", "").lower() in {"1", "true", "yes"}
if scheme not in ({"https", "http"} if allow_http else {"https"}):
raise BoxIMImageError("BOXIM 图片地址必须使用 HTTPS")
if parsed.username or parsed.password or not parsed.hostname:
raise BoxIMImageError("BOXIM 图片地址无效")
host = parsed.hostname.lower().rstrip(".")
allowed_hosts = _configured_hosts("BOXIM_IMAGE_ALLOWED_HOSTS")
if allowed_hosts and not _host_matches(host, allowed_hosts):
raise BoxIMImageError("BOXIM 图片地址不在允许的域名范围内")
private_hosts = _configured_hosts("BOXIM_IMAGE_PRIVATE_HOSTS")
try:
addresses = {ipaddress.ip_address(host)}
except ValueError:
addresses = _resolved_addresses(host, parsed.port or (443 if scheme == "https" else 80))
if not addresses:
raise BoxIMImageError("BOXIM 图片地址无法解析")
if _host_matches(host, private_hosts):
return
if any(not address.is_global for address in addresses):
raise BoxIMImageError("BOXIM 图片地址指向受限网络")
def _filename_from_url(url: str) -> str:
value = unquote(PurePosixPath(urlsplit(url).path).name).strip()
value = value.replace("\x00", "")
return (value or "boxim-image")[:255]
def download_boxim_image(
content: str,
*,
base_url: str = "",
transport: httpx.BaseTransport | None = None,
) -> DownloadedBoxIMImage:
"""Download one BOXIM image without redirects or oversized responses escaping checks."""
url = parse_boxim_image_url(content, base_url=base_url)
max_bytes = max(1024, int(os.getenv("CHAT_IMAGE_MAX_BYTES", str(8 * 1024 * 1024))))
timeout = max(1.0, min(float(os.getenv("BOXIM_IMAGE_TIMEOUT_SECONDS", "15")), 60.0))
with httpx.Client(
timeout=timeout,
follow_redirects=False,
trust_env=False,
transport=transport,
) as client:
for _ in range(MAX_REDIRECTS + 1):
_is_safe_remote_url(url)
try:
with client.stream("GET", url, headers={"Accept": "image/*"}) as response:
if response.status_code in {301, 302, 303, 307, 308}:
location = response.headers.get("location", "").strip()
if not location:
raise BoxIMImageError("BOXIM 图片跳转地址无效")
url = urljoin(url, location)
continue
response.raise_for_status()
raw_length = response.headers.get("content-length", "")
if raw_length.isdigit() and int(raw_length) > max_bytes:
raise BoxIMImageError("BOXIM 图片超过大小限制")
chunks = bytearray()
for chunk in response.iter_bytes():
chunks.extend(chunk)
if len(chunks) > max_bytes:
raise BoxIMImageError("BOXIM 图片超过大小限制")
if not chunks:
raise BoxIMImageError("BOXIM 图片内容为空")
return DownloadedBoxIMImage(
content=bytes(chunks),
filename=_filename_from_url(url),
mime_type=response.headers.get("content-type", "").split(";", 1)[0][:100],
source_url=url,
)
except BoxIMImageError:
raise
except (httpx.HTTPError, OSError) as exc:
raise BoxIMImageError("BOXIM 图片下载失败") from exc
raise BoxIMImageError("BOXIM 图片跳转次数过多")
@@ -0,0 +1,20 @@
from datetime import datetime
from sqlalchemy.orm import Session
from models import ChatAttachment
def purge_expired_chat_attachments(
db: Session,
*,
now: datetime | None = None,
) -> int:
"""Remove expired derived image data; raw image bytes are never persisted."""
count = db.query(ChatAttachment).filter(
ChatAttachment.expires_at < (now or datetime.utcnow())
).delete(synchronize_session=False)
if count:
db.commit()
db.expire_all()
return count
@@ -0,0 +1,115 @@
import logging
import os
import threading
import time
from dataclasses import dataclass
import httpx
logger = logging.getLogger(__name__)
@dataclass(frozen=True)
class ChatModelConfig:
api_base_url: str
api_key: str
model: str
max_tokens: int
timeout_seconds: float
vision_model: str
ocr_model: str
vision_max_tokens: int
vision_timeout_seconds: float
source: str
_cache_lock = threading.Lock()
_cached_config: ChatModelConfig | None = None
_cache_expires_at = 0.0
def _environment_config() -> ChatModelConfig:
return ChatModelConfig(
api_base_url=os.getenv(
"CHAT_API_URL", "https://dashscope.aliyuncs.com/compatible-mode/v1"
).rstrip("/"),
api_key=os.getenv("CHAT_API_KEY", ""),
model=os.getenv("CHAT_MODEL", "qwen-plus"),
max_tokens=max(128, int(os.getenv("CHAT_MAX_OUTPUT_TOKENS", "1024"))),
timeout_seconds=max(5.0, float(os.getenv("CHAT_TIMEOUT_SECONDS", "30"))),
vision_model=os.getenv("VISION_MODEL", "qwen3.6-flash"),
ocr_model=os.getenv("VISION_OCR_MODEL", "qwen-vl-ocr"),
vision_max_tokens=max(256, int(os.getenv("VISION_MAX_OUTPUT_TOKENS", "2048"))),
vision_timeout_seconds=max(10.0, float(os.getenv("VISION_TIMEOUT_SECONDS", "90"))),
source="environment",
)
def _fetch_runtime_config() -> ChatModelConfig | None:
url = os.getenv("CHAT_MODEL_CONFIG_URL", "").strip()
token = os.getenv("AVATAR_MODEL_CONFIG_TOKEN", "").strip()
if not url or not token:
return None
response = httpx.get(
url,
headers={"X-Avatar-Config-Token": token},
timeout=max(2.0, float(os.getenv("CHAT_MODEL_CONFIG_TIMEOUT_SECONDS", "5"))),
)
response.raise_for_status()
payload = response.json().get("data") or {}
api_base_url = str(payload.get("api_base_url") or "").rstrip("/")
api_key = str(payload.get("api_key") or "")
model = str(payload.get("model") or "")
if not api_base_url or not api_key or not model:
raise ValueError("数字分身专用模型配置不完整")
return ChatModelConfig(
api_base_url=api_base_url,
api_key=api_key,
model=model,
max_tokens=max(128, int(payload.get("max_tokens") or 1024)),
timeout_seconds=max(5.0, float(payload.get("timeout_seconds") or 30)),
vision_model=str(
payload.get("vision_model")
or os.getenv("VISION_MODEL", "qwen3.6-flash")
),
ocr_model=str(
payload.get("ocr_model")
or os.getenv("VISION_OCR_MODEL", "qwen-vl-ocr")
),
vision_max_tokens=max(
256, int(os.getenv("VISION_MAX_OUTPUT_TOKENS", "2048"))
),
vision_timeout_seconds=max(
10.0, float(os.getenv("VISION_TIMEOUT_SECONDS", "90"))
),
source="admin",
)
def get_chat_model_config(*, force_refresh: bool = False) -> ChatModelConfig:
global _cached_config, _cache_expires_at
now = time.monotonic()
if not force_refresh and _cached_config is not None and now < _cache_expires_at:
return _cached_config
with _cache_lock:
now = time.monotonic()
if not force_refresh and _cached_config is not None and now < _cache_expires_at:
return _cached_config
try:
config = _fetch_runtime_config() or _environment_config()
except (httpx.HTTPError, ValueError, TypeError) as exc:
logger.warning("读取数字分身专用模型配置失败,暂时使用环境变量配置: %s", exc)
config = _environment_config()
_cached_config = config
ttl = max(5, int(os.getenv("CHAT_MODEL_CONFIG_CACHE_SECONDS", "60")))
_cache_expires_at = now + ttl
return config
def clear_chat_model_config_cache() -> None:
global _cached_config, _cache_expires_at
with _cache_lock:
_cached_config = None
_cache_expires_at = 0.0
@@ -0,0 +1,124 @@
"""Signed client for Huihui's production payment-v3 service."""
import hashlib
import random
import string
from datetime import datetime, timedelta, timezone
from typing import Any
import httpx
_CN_TZ = timezone(timedelta(hours=8))
class HuihuiPaymentError(RuntimeError):
pass
class HuihuiPaymentClient:
def __init__(self, config: dict):
self.base_url = config.get(
"HUIHUI_PAYMENT_BASE_URL", "https://open.99hui.com/api/payment-v3"
).rstrip("/")
self.app_id = config.get("HUIHUI_APP_ID", "")
self.access_id = config.get("HUIHUI_ACCESS_ID", "")
self.access_secret = config.get("HUIHUI_ACCESS_SECRET", "")
self.timeout = float(config.get("HUIHUI_PAYMENT_TIMEOUT_SECONDS", 30))
@property
def configured(self) -> bool:
return bool(self.base_url and self.app_id and self.access_id and self.access_secret)
def _signed_params(self, user_id: str) -> dict:
params = {
"appId": self.app_id,
"accessId": self.access_id,
"nonce": "".join(random.choices(string.ascii_lowercase + string.digits, k=12)),
"timestamp": datetime.now(_CN_TZ).strftime("%Y%m%d%H%M%S"),
"signType": "MD5",
"signVersion": "1.0",
"userId": user_id,
}
source = "&".join(
f"{key}={params[key]}"
for key in sorted(params)
if params[key] not in (None, "", [])
)
source += f"&accessSecret={self.access_secret}"
params["signature"] = hashlib.md5(source.encode("utf-8")).hexdigest().upper()
return params
@staticmethod
def _json(response: httpx.Response) -> dict:
try:
payload = response.json()
except ValueError as exc:
raise HuihuiPaymentError("会会支付返回了无效响应") from exc
if not isinstance(payload, dict):
raise HuihuiPaymentError("会会支付返回格式不正确")
return payload
def create_payment(
self,
*,
huihui_token: str,
huihui_user_id: str,
real_name: str,
order_no: str,
amount: str,
points_amount: int,
pay_type: str,
pay_way: str,
callback_url: str,
) -> dict[str, Any]:
if not self.configured:
raise HuihuiPaymentError("会会支付服务未配置")
if not huihui_token or not huihui_user_id:
raise HuihuiPaymentError("当前会会登录凭证无法发起支付")
now = datetime.now(_CN_TZ)
body = {
"appId": self.app_id,
"callbackUrl": callback_url,
"chargeType": 4,
"currency": "cny",
"description": f"充值 {points_amount} 积分",
"expend": {},
"masterOrderAmt": amount,
"masterOrderNo": order_no,
"memberId": huihui_user_id,
"orderDesc": "数字分身积分充值",
"orderTime": now.isoformat(),
"orderTitle": "数字分身积分充值",
"payAmt": float(amount),
"payType": pay_type,
"payWay": pay_way,
"realName": real_name or "会会用户",
"timeExpire": (now + timedelta(hours=2)).strftime("%Y%m%d%H%M%S"),
}
headers = {
"Authorization": f"Bearer {huihui_token}",
"appId": self.app_id,
"windowAppId": self.app_id,
}
try:
response = httpx.post(
f"{self.base_url}/payment/pay",
headers=headers,
params=self._signed_params(huihui_user_id),
json=body,
timeout=self.timeout,
follow_redirects=True,
)
except httpx.HTTPError as exc:
raise HuihuiPaymentError("会会支付连接失败,请稍后重试") from exc
payload = self._json(response)
code = payload.get("code")
if response.status_code >= 400 or code not in (0, 200, "0", "200"):
raise HuihuiPaymentError(payload.get("message") or "会会支付下单失败")
data = payload.get("data") or {}
if not isinstance(data, dict):
raise HuihuiPaymentError("会会支付未返回订单信息")
return data
@@ -0,0 +1,143 @@
"""Durable, serial knowledge-document indexing for the avatar knowledge base."""
import json
import logging
import os
import queue
import threading
from datetime import datetime, timezone
from database import SessionLocal
from models import KnowledgeChunk, KnowledgeDoc
import embeddings
logger = logging.getLogger(__name__)
BACKEND_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
UPLOAD_DIR = os.path.abspath(
os.getenv("UPLOAD_DIR", os.path.join(BACKEND_DIR, "routers", "uploads"))
)
class KnowledgeVectorizer:
"""Indexes one document at a time so slow providers cannot block uploads."""
def __init__(self):
self._queue: queue.Queue[str] = queue.Queue()
self._queued: set[str] = set()
self._lock = threading.Lock()
self._thread: threading.Thread | None = None
def start(self):
if self._thread and self._thread.is_alive():
return
self._thread = threading.Thread(
target=self._run, name="knowledge-vectorizer", daemon=True
)
self._thread.start()
db = SessionLocal()
try:
# A process restart must not abandon documents already accepted by upload.
for (doc_id,) in db.query(KnowledgeDoc.id).filter(KnowledgeDoc.status == "parsing"):
self.enqueue(doc_id)
finally:
db.close()
def enqueue(self, doc_id: str):
with self._lock:
if doc_id in self._queued:
return
self._queued.add(doc_id)
self._queue.put(doc_id)
def _run(self):
while True:
doc_id = self._queue.get()
try:
self.vectorize_document(doc_id)
except Exception:
logger.exception("Unexpected knowledge vectorizer failure for %s", doc_id)
finally:
with self._lock:
self._queued.discard(doc_id)
self._queue.task_done()
def vectorize_document(self, doc_id: str):
db = SessionLocal()
try:
doc = db.get(KnowledgeDoc, doc_id)
if not doc or doc.status != "parsing":
return
stored_name = os.path.basename(doc.file_url or "")
path = os.path.join(UPLOAD_DIR, doc.avatar_id, stored_name)
if not stored_name or not os.path.isfile(path):
raise FileNotFoundError("原文件不可用,请重新上传")
self._set_progress(db, doc, "extracting", 8)
text = embeddings.extract_text(path, f".{doc.file_type}")
self._set_progress(db, doc, "chunking", 22)
chunks = embeddings.chunk_text(text)
if not chunks:
raise ValueError("文档没有可建立索引的文字内容")
self._set_progress(db, doc, "embedding", 30)
def embedding_progress(done: int, total: int):
percent = 30 + int((done / max(1, total)) * 65)
self._set_progress(db, doc, "embedding", min(percent, 95))
vectors = embeddings.embed(chunks, on_progress=embedding_progress)
if len(vectors) != len(chunks):
raise ValueError("向量服务返回数量与文档分段不一致")
# Commit the document and every chunk together. Chat only sees complete indexes.
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == doc.id).delete()
db.add_all(
[
KnowledgeChunk(
doc_id=doc.id,
avatar_id=doc.avatar_id,
content=chunk,
vector=json.dumps(vector),
chunk_index=index,
embedding_model=embeddings.MODEL,
)
for index, (chunk, vector) in enumerate(zip(chunks, vectors))
]
)
doc.vectorized = True
doc.embedding_model = embeddings.MODEL
doc.chunk_count = len(chunks)
doc.vectorized_at = datetime.now(timezone.utc)
doc.status = "ready"
doc.error_message = ""
doc.index_stage = "ready"
doc.index_progress = 100
db.commit()
logger.info("Knowledge document %s indexed with %s chunks", doc.id, len(chunks))
except Exception as exc:
db.rollback()
failed_doc = db.get(KnowledgeDoc, doc_id)
if failed_doc:
db.query(KnowledgeChunk).filter(KnowledgeChunk.doc_id == failed_doc.id).delete()
failed_doc.status = "failed"
failed_doc.vectorized = False
failed_doc.embedding_model = ""
failed_doc.chunk_count = 0
failed_doc.vectorized_at = None
failed_doc.error_message = str(exc)[:300] or "建立知识索引失败"
failed_doc.index_stage = "failed"
failed_doc.index_progress = 0
db.commit()
logger.exception("Knowledge vectorization failed for %s: %s", doc_id, exc)
finally:
db.close()
@staticmethod
def _set_progress(db, doc, stage: str, progress: int):
doc.index_stage = stage
doc.index_progress = progress
db.commit()
knowledge_vectorizer = KnowledgeVectorizer()
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,203 @@
"""User-scoped token accounting for every avatar model request."""
import math
from dataclasses import dataclass
from datetime import datetime, timedelta
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from models import Avatar, TokenAccount, TokenUsage, User
DEFAULT_TOKEN_GRANT = 1_000_000
class InsufficientTokensError(RuntimeError):
pass
@dataclass(frozen=True)
class TokenReservation:
usage_id: str
user_id: str
reserved_tokens: int
def get_or_create_account(db: Session, user_id: str) -> TokenAccount:
account = db.query(TokenAccount).filter(TokenAccount.user_id == user_id).first()
if account:
return account
account = TokenAccount(
user_id=user_id,
balance=DEFAULT_TOKEN_GRANT,
total_granted=DEFAULT_TOKEN_GRANT,
total_consumed=0,
)
db.add(account)
try:
db.commit()
except IntegrityError:
# A concurrent first request may have created the same user account.
db.rollback()
account = db.query(TokenAccount).filter(TokenAccount.user_id == user_id).first()
if account is None:
raise
db.refresh(account)
return account
def avatar_owner_user(db: Session, avatar: Avatar) -> User | None:
owner_id = (avatar.owner_id or "").strip()
if not owner_id:
return None
return db.query(User).filter(User.huihui_user_id == owner_id).first()
def estimate_request_tokens(messages: list[dict], max_output_tokens: int) -> int:
# UTF-8 bytes / 2 deliberately overestimates mixed Chinese/English prompts;
# the unused reservation is returned after provider usage is received.
content_bytes = sum(
len(str(item.get("content", "")).encode("utf-8"))
for item in messages
)
prompt_reserve = max(1, math.ceil(content_bytes / 2) + len(messages) * 6)
return prompt_reserve + max(1, int(max_output_tokens))
def estimate_fallback_usage(messages: list[dict], output: str) -> int:
content_bytes = sum(
len(str(item.get("content", "")).encode("utf-8"))
for item in messages
) + len((output or "").encode("utf-8"))
return max(1, math.ceil(content_bytes / 3) + len(messages) * 4)
def reserve_avatar_tokens(
db: Session,
avatar: Avatar,
source: str,
model: str,
messages: list[dict],
max_output_tokens: int,
*,
minimum_reserve_tokens: int = 0,
) -> TokenReservation:
user = avatar_owner_user(db, avatar)
if not user:
raise InsufficientTokensError("分身尚未关联有效用户,暂时无法使用积分")
account = get_or_create_account(db, user.id)
reserved = max(
estimate_request_tokens(messages, max_output_tokens),
max(0, int(minimum_reserve_tokens or 0)),
)
updated = (
db.query(TokenAccount)
.filter(TokenAccount.id == account.id, TokenAccount.balance >= reserved)
.update(
{TokenAccount.balance: TokenAccount.balance - reserved},
synchronize_session=False,
)
)
if updated != 1:
db.rollback()
raise InsufficientTokensError("积分余额不足,请充值后继续")
db.refresh(account)
usage = TokenUsage(
user_id=user.id,
avatar_id=avatar.id,
source=source,
model=model,
status="reserved",
reserved_tokens=reserved,
)
db.add(usage)
db.flush()
usage.balance_after = account.balance
db.commit()
return TokenReservation(usage.id, user.id, reserved)
def settle_reservation(
db: Session,
reservation: TokenReservation,
usage: dict | None,
*,
fallback_total: int,
) -> dict:
record = db.query(TokenUsage).filter(TokenUsage.id == reservation.usage_id).first()
if not record or record.status != "reserved":
return {}
provider_usage = usage or {}
prompt_tokens = max(0, int(provider_usage.get("prompt_tokens") or 0))
completion_tokens = max(0, int(provider_usage.get("completion_tokens") or 0))
provider_total = max(
int(provider_usage.get("total_tokens") or 0),
prompt_tokens + completion_tokens,
)
total_tokens = max(1, provider_total or int(fallback_total or 0))
updated = (
db.query(TokenAccount)
.filter(TokenAccount.user_id == reservation.user_id)
.update(
{
TokenAccount.balance: TokenAccount.balance + reservation.reserved_tokens - total_tokens,
TokenAccount.total_consumed: TokenAccount.total_consumed + total_tokens,
},
synchronize_session=False,
)
)
if updated != 1:
raise RuntimeError("积分账户不存在")
db.expire_all()
account = db.query(TokenAccount).filter(TokenAccount.user_id == reservation.user_id).first()
record.prompt_tokens = prompt_tokens
record.completion_tokens = completion_tokens
record.total_tokens = total_tokens
record.balance_after = account.balance
record.status = "completed"
record.settled_at = datetime.utcnow()
db.commit()
return {
"promptTokens": prompt_tokens,
"completionTokens": completion_tokens,
"totalTokens": total_tokens,
"balance": account.balance,
}
def release_reservation(db: Session, reservation: TokenReservation, reason: str = "") -> None:
record = db.query(TokenUsage).filter(TokenUsage.id == reservation.usage_id).first()
if not record or record.status != "reserved":
return
updated = (
db.query(TokenAccount)
.filter(TokenAccount.user_id == reservation.user_id)
.update(
{TokenAccount.balance: TokenAccount.balance + reservation.reserved_tokens},
synchronize_session=False,
)
)
if updated:
db.expire_all()
account = db.query(TokenAccount).filter(TokenAccount.user_id == reservation.user_id).first()
record = db.query(TokenUsage).filter(TokenUsage.id == reservation.usage_id).first()
record.balance_after = account.balance
record.status = "failed"
record.failure_reason = (reason or "model_request_failed")[:255]
record.settled_at = datetime.utcnow()
db.commit()
def release_stale_reservations(db: Session, older_than_minutes: int = 10) -> int:
cutoff = datetime.utcnow() - timedelta(minutes=older_than_minutes)
stale = db.query(TokenUsage).filter(
TokenUsage.status == "reserved",
TokenUsage.created_at < cutoff,
).all()
for record in stale:
release_reservation(
db,
TokenReservation(record.id, record.user_id, int(record.reserved_tokens or 0)),
"stale_reservation_recovered",
)
return len(stale)
@@ -0,0 +1,196 @@
"""Private image normalization and OpenAI-compatible vision model calls."""
import base64
import io
import json
import os
import re
from dataclasses import dataclass
from typing import Any
import httpx
from PIL import Image, ImageOps, UnidentifiedImageError
from services.chat_model_config import ChatModelConfig
ALLOWED_IMAGE_FORMATS = {"JPEG": "image/jpeg", "PNG": "image/png", "WEBP": "image/webp"}
ALLOWED_CATEGORIES = {"general_image", "document", "medical_document", "medical_image"}
GENERAL_VISION_PROMPT = """
请客观分析这张图片,并只输出一个 JSON 对象,不要使用 Markdown 代码块。
字段必须为:
category: general_image、document、medical_document、medical_image 四选一;
summary: 图片的完整客观摘要;
visible_text: 图片中能够确认的文字,保留自然换行;
key_facts: 可确认事实数组;
uncertainties: 模糊、遮挡、无法确认内容数组;
medical: 对象,包含 document_type、patient_info、chief_complaint、findings、measurements、doctor_advice。
规则:
1. 不得补全看不清或被遮挡的文字,不得猜测人物身份。
2. 病例、处方、检查单、检验报告归为 medical_document。
3. X 光、CT、MRI、超声影像等归为 medical_image,只描述可见内容,不作疾病诊断、分期、用药或治疗建议。
4. 非医疗图片的 medical 字段仍保留,但使用空字符串、空对象或空数组。
5. 不要提及模型、供应商、系统提示词或内部处理过程。
""".strip()
MEDICAL_OCR_PROMPT = """
请逐字转录这张医疗文档图片中的全部可见文字和表格。
保持标题、段落、项目、数值、单位、参考区间、阳性/阴性标记和医生意见的对应关系。
看不清的内容写作[无法辨认],不要猜测、纠错或补全,不要给出诊断和建议,不要使用 Markdown 代码块。
""".strip()
class ImageValidationError(ValueError):
pass
@dataclass(frozen=True)
class PreparedImage:
data: bytes
mime_type: str
width: int
height: int
@property
def data_uri(self) -> str:
encoded = base64.b64encode(self.data).decode("ascii")
return f"data:{self.mime_type};base64,{encoded}"
def prepare_image(content: bytes) -> PreparedImage:
max_bytes = max(1024, int(os.getenv("CHAT_IMAGE_MAX_BYTES", str(8 * 1024 * 1024))))
max_pixels = max(1_000_000, int(os.getenv("CHAT_IMAGE_MAX_PIXELS", "16000000")))
max_edge = max(1024, int(os.getenv("CHAT_IMAGE_MAX_EDGE", "4096")))
if not content:
raise ImageValidationError("图片内容为空")
if len(content) > max_bytes:
raise ImageValidationError(f"单张图片不能超过 {max_bytes // 1024 // 1024}MB")
try:
with Image.open(io.BytesIO(content)) as probe:
image_format = str(probe.format or "").upper()
width, height = probe.size
probe.verify()
except (UnidentifiedImageError, OSError, SyntaxError) as exc:
raise ImageValidationError("图片格式无效或文件已损坏") from exc
if image_format not in ALLOWED_IMAGE_FORMATS:
raise ImageValidationError("仅支持 JPG、PNG、WebP 图片")
if width <= 0 or height <= 0 or width * height > max_pixels:
raise ImageValidationError("图片像素过大,请压缩后重新上传")
try:
with Image.open(io.BytesIO(content)) as original:
image = ImageOps.exif_transpose(original)
image.load()
if max(image.size) > max_edge:
image.thumbnail((max_edge, max_edge), Image.Resampling.LANCZOS)
if image.mode in {"RGBA", "LA"}:
canvas = Image.new("RGB", image.size, "white")
alpha = image.getchannel("A")
canvas.paste(image.convert("RGB"), mask=alpha)
image = canvas
elif image.mode != "RGB":
image = image.convert("RGB")
output = io.BytesIO()
image.save(output, format="JPEG", quality=92, optimize=True)
normalized = output.getvalue()
normalized_width, normalized_height = image.size
except (OSError, ValueError) as exc:
raise ImageValidationError("图片解码失败,请重新选择图片") from exc
return PreparedImage(
data=normalized,
mime_type="image/jpeg",
width=normalized_width,
height=normalized_height,
)
def call_vision_model(
prepared: PreparedImage,
model_config: ChatModelConfig,
*,
model: str,
prompt: str,
json_output: bool,
) -> dict:
if not model_config.api_key:
raise RuntimeError("视觉模型服务未配置")
payload: dict[str, Any] = {
"model": model,
"messages": [
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": prepared.data_uri}},
{"type": "text", "text": prompt},
],
}
],
"temperature": 0,
"max_tokens": model_config.vision_max_tokens,
}
if json_output:
payload["response_format"] = {"type": "json_object"}
try:
response = httpx.post(
f"{model_config.api_base_url}/chat/completions",
headers={"Authorization": f"Bearer {model_config.api_key}"},
json=payload,
timeout=model_config.vision_timeout_seconds,
)
response.raise_for_status()
data = response.json()
content = data.get("choices", [{}])[0].get("message", {}).get("content", "")
except (httpx.HTTPError, ValueError, KeyError, IndexError) as exc:
raise RuntimeError("图片识别服务暂时不可用") from exc
if not isinstance(content, str) or not content.strip():
raise RuntimeError("图片识别服务没有返回有效结果")
return {"content": content.strip(), "usage": data.get("usage") or {}}
def parse_vision_analysis(content: str) -> dict:
value = (content or "").strip()
fenced = re.match(r"^```(?:json)?\s*(.*?)\s*```$", value, re.DOTALL | re.IGNORECASE)
if fenced:
value = fenced.group(1).strip()
try:
payload = json.loads(value)
except (TypeError, ValueError) as exc:
raise RuntimeError("图片识别结果格式无效") from exc
if not isinstance(payload, dict):
raise RuntimeError("图片识别结果格式无效")
category = str(payload.get("category") or "general_image").strip().lower()
if category not in ALLOWED_CATEGORIES:
category = "general_image"
medical = payload.get("medical") if isinstance(payload.get("medical"), dict) else {}
return {
"category": category,
"summary": str(payload.get("summary") or "").strip(),
"visible_text": str(payload.get("visible_text") or "").strip(),
"key_facts": _string_list(payload.get("key_facts")),
"uncertainties": _string_list(payload.get("uncertainties")),
"medical": medical,
}
def build_attachment_warning(analysis: dict, *, ocr_failed: bool = False) -> str:
warnings = list(analysis.get("uncertainties") or [])
category = analysis.get("category")
if ocr_failed:
warnings.append("精确文字识别暂时不可用,请人工核对图片原文")
if category == "medical_document":
warnings.append("病例识别结果仅供辅助,不能替代医生诊断,请核对原始文档")
elif category == "medical_image":
warnings.append("医学影像仅作客观描述,不能替代影像报告和医生诊断")
return ";".join(dict.fromkeys(item for item in warnings if item))
def _string_list(value: Any) -> list[str]:
if not isinstance(value, list):
return []
return [str(item).strip() for item in value if str(item).strip()]
@@ -0,0 +1 @@
@@ -0,0 +1,131 @@
import uuid
import pytest
from database import init_db, SessionLocal
from models import (
Authorization,
Avatar,
ChatAttachment,
TakeoverCursor,
TakeoverMessage,
TakeoverReplyTask,
TokenAccount,
TokenPaymentOrder,
TokenUsage,
User,
)
@pytest.fixture(scope="session", autouse=True)
def setup_database():
"""Initialize DB tables and seed an Authorization row."""
init_db()
db = SessionLocal()
try:
existing = db.query(Authorization).first()
if existing is None:
auth = Authorization(
id="test-auth-1",
avatar_id="test-avatar-1",
target_type="user",
target_id="test-user-1",
target_name="Test User",
permissions=["read", "write"],
status="active",
)
db.add(auth)
db.commit()
finally:
db.close()
@pytest.fixture
def authorization_context():
"""Create isolated users, avatars, and one authorization for API tests."""
suffix = uuid.uuid4().hex
owner = User(
id=f"owner-{suffix}",
huihui_user_id=f"huihui-owner-{suffix}",
nickname="授权测试用户",
app_token=f"owner-token-{suffix}",
)
other = User(
id=f"other-{suffix}",
huihui_user_id=f"huihui-other-{suffix}",
nickname="其他用户",
app_token=f"other-token-{suffix}",
)
avatar = Avatar(
id=f"avatar-{suffix}",
owner_id=owner.huihui_user_id,
name="授权测试分身",
status="active",
config={},
)
other_avatar = Avatar(
id=f"other-avatar-{suffix}",
owner_id=other.huihui_user_id,
name="其他分身",
status="active",
config={},
)
authorization = Authorization(
id=f"authorization-{suffix}",
avatar_id=avatar.id,
target_type="user",
target_id=f"contact-{suffix}",
target_name="测试联系人",
permissions=["chat", "browse"],
status="active",
)
db = SessionLocal()
try:
db.add_all([owner, other, avatar, other_avatar, authorization])
db.commit()
yield {
"owner": owner,
"other": other,
"avatar": avatar,
"other_avatar": other_avatar,
"authorization": authorization,
"owner_headers": {"Authorization": f"Bearer {owner.app_token}"},
"other_headers": {"Authorization": f"Bearer {other.app_token}"},
"suffix": suffix,
}
finally:
db.rollback()
avatar_ids = [avatar.id, other_avatar.id]
db.query(ChatAttachment).filter(
ChatAttachment.avatar_id.in_(avatar_ids)
).delete(synchronize_session=False)
db.query(TakeoverReplyTask).filter(
TakeoverReplyTask.avatar_id.in_(avatar_ids)
).delete(synchronize_session=False)
db.query(TakeoverMessage).filter(
TakeoverMessage.avatar_id.in_(avatar_ids)
).delete(synchronize_session=False)
db.query(TakeoverCursor).filter(
TakeoverCursor.avatar_id.in_(avatar_ids)
).delete(synchronize_session=False)
db.query(Authorization).filter(
Authorization.avatar_id.in_(avatar_ids)
).delete(synchronize_session=False)
db.query(Avatar).filter(Avatar.id.in_(avatar_ids)).delete(
synchronize_session=False
)
user_ids = [owner.id, other.id]
db.query(TokenPaymentOrder).filter(TokenPaymentOrder.user_id.in_(user_ids)).delete(
synchronize_session=False
)
db.query(TokenUsage).filter(TokenUsage.user_id.in_(user_ids)).delete(
synchronize_session=False
)
db.query(TokenAccount).filter(TokenAccount.user_id.in_(user_ids)).delete(
synchronize_session=False
)
db.query(User).filter(User.id.in_([owner.id, other.id])).delete(
synchronize_session=False
)
db.commit()
db.close()
@@ -0,0 +1,214 @@
from fastapi.testclient import TestClient
from main import app
client = TestClient(app)
def test_authorization_list_is_scoped_to_owned_avatar(authorization_context):
context = authorization_context
response = client.get(
f"/api/avatar/{context['avatar'].id}/authorizations",
headers=context["owner_headers"],
)
assert response.status_code == 200
payload = response.json()
assert payload["code"] == 200
assert [item["id"] for item in payload["data"]] == [context["authorization"].id]
forbidden = client.get(
f"/api/avatar/{context['other_avatar'].id}/authorizations",
headers=context["owner_headers"],
)
assert forbidden.status_code == 403
def test_create_update_and_delete_authorization(authorization_context):
context = authorization_context
avatar_id = context["avatar"].id
target_id = f"new-contact-{context['suffix']}"
created = client.post(
f"/api/avatar/{avatar_id}/authorizations",
headers=context["owner_headers"],
json={
"targetType": "user",
"targetId": target_id,
"targetName": "新联系人",
"permissions": ["friend", "chat", "browse"],
},
).json()
assert created["code"] == 200
authorization_id = created["data"]["id"]
assert created["data"]["permissions"] == ["friend", "chat", "browse"]
duplicate = client.post(
f"/api/avatar/{avatar_id}/authorizations",
headers=context["owner_headers"],
json={
"targetType": "user",
"targetId": target_id,
"targetName": "重复联系人",
"permissions": ["chat"],
},
).json()
assert duplicate["code"] == 409
updated = client.put(
f"/api/avatar/{avatar_id}/authorizations",
headers=context["owner_headers"],
json={
"id": authorization_id,
"targetName": "联系人新名称",
"permissions": ["interact", "publish"],
},
).json()
assert updated["code"] == 200
assert updated["data"]["targetName"] == "联系人新名称"
assert updated["data"]["permissions"] == ["publish", "interact"]
deleted = client.delete(
f"/api/avatar/{avatar_id}/authorizations/{authorization_id}",
headers=context["owner_headers"],
).json()
assert deleted["code"] == 200
assert deleted["data"]["id"] == authorization_id
def test_authorization_requires_login_and_rejects_unknown_permissions(authorization_context):
context = authorization_context
avatar_id = context["avatar"].id
no_session = client.get(f"/api/avatar/{avatar_id}/authorizations")
assert no_session.status_code == 401
invalid = client.post(
f"/api/avatar/{avatar_id}/authorizations",
headers=context["owner_headers"],
json={
"targetType": "user",
"targetId": "invalid-target",
"targetName": "无效权限",
"permissions": ["admin"],
},
).json()
assert invalid["code"] == 400
def test_avatar_permission_settings_default_and_persist(authorization_context):
context = authorization_context
endpoint = f"/api/avatar/{context['avatar'].id}/permission-settings"
initial = client.get(endpoint, headers=context["owner_headers"]).json()
assert initial["code"] == 200
assert initial["data"] == {
"avatarId": context["avatar"].id,
"permissions": ["friend", "chat"],
"takeoverReplyDelaySeconds": 180,
}
updated = client.put(
endpoint,
headers=context["owner_headers"],
json={"permissions": ["interact", "takeover", "publish", "friend", "friend"]},
).json()
assert updated["code"] == 200
assert updated["data"]["permissions"] == ["friend", "publish", "interact", "takeover"]
reloaded = client.get(endpoint, headers=context["owner_headers"]).json()
assert reloaded["data"]["permissions"] == ["friend", "publish", "interact", "takeover"]
assert reloaded["data"]["takeoverReplyDelaySeconds"] == 180
def test_avatar_permission_settings_allow_all_disabled(authorization_context):
context = authorization_context
endpoint = f"/api/avatar/{context['avatar'].id}/permission-settings"
response = client.put(
endpoint,
headers=context["owner_headers"],
json={"permissions": []},
).json()
assert response["code"] == 200
assert response["data"]["permissions"] == []
def test_avatar_permission_settings_validate_owner_and_permissions(authorization_context):
context = authorization_context
endpoint = f"/api/avatar/{context['avatar'].id}/permission-settings"
invalid = client.put(
endpoint,
headers=context["owner_headers"],
json={"permissions": ["admin"]},
).json()
assert invalid["code"] == 400
missing = client.put(
endpoint,
headers=context["owner_headers"],
json={},
).json()
assert missing["code"] == 400
forbidden = client.get(
f"/api/avatar/{context['other_avatar'].id}/permission-settings",
headers=context["owner_headers"],
)
assert forbidden.status_code == 403
unauthenticated = client.get(endpoint)
assert unauthenticated.status_code == 401
def test_takeover_delay_minimum_and_single_active_avatar_per_owner(authorization_context):
from database import SessionLocal
from models import Avatar
context = authorization_context
endpoint = f"/api/avatar/{context['avatar'].id}/permission-settings"
invalid = client.put(
endpoint,
headers=context["owner_headers"],
json={"permissions": ["chat"], "takeoverReplyDelaySeconds": 2},
).json()
assert invalid["code"] == 400
second_avatar_id = f"second-{context['suffix']}"
db = SessionLocal()
try:
db.add(
Avatar(
id=second_avatar_id,
owner_id=context["owner"].huihui_user_id,
name="第二个分身",
status="active",
config={"authorizationPermissions": ["chat", "takeover"]},
)
)
db.commit()
finally:
db.close()
try:
updated = client.put(
endpoint,
headers=context["owner_headers"],
json={"permissions": ["chat", "takeover"], "takeoverReplyDelaySeconds": 3},
).json()
assert updated["code"] == 200
assert updated["data"]["takeoverReplyDelaySeconds"] == 3
assert updated["data"]["disabledAvatarIds"] == [second_avatar_id]
db = SessionLocal()
try:
second = db.query(Avatar).filter(Avatar.id == second_avatar_id).one()
assert "takeover" not in second.config["authorizationPermissions"]
finally:
db.close()
finally:
db = SessionLocal()
try:
db.query(Avatar).filter(Avatar.id == second_avatar_id).delete()
db.commit()
finally:
db.close()
@@ -0,0 +1,93 @@
"""Ownership and configuration-isolation tests for digital avatars."""
from fastapi.testclient import TestClient
from database import SessionLocal
from main import app
from models import Avatar
client = TestClient(app)
def test_avatar_detail_and_update_require_the_owner(authorization_context):
context = authorization_context
avatar_id = context["avatar"].id
assert client.get(f"/api/avatar/{avatar_id}").status_code == 401
assert client.get(
f"/api/avatar/{avatar_id}", headers=context["other_headers"]
).status_code == 403
updated = client.put(
f"/api/avatar/{avatar_id}",
headers=context["owner_headers"],
json={
"description": "独立描述",
"config": {"replyStyle": "concise"},
},
)
assert updated.status_code == 200
assert updated.json()["data"]["description"] == "独立描述"
forbidden = client.put(
f"/api/avatar/{avatar_id}",
headers=context["other_headers"],
json={"description": "越权修改"},
)
assert forbidden.status_code == 403
def test_avatar_config_updates_do_not_erase_takeover_or_knowledge_scope(authorization_context):
context = authorization_context
avatar_id = context["avatar"].id
db = SessionLocal()
try:
avatar = db.query(Avatar).filter(Avatar.id == avatar_id).one()
avatar.config = {
"authorizationPermissions": ["chat", "takeover"],
"takeoverReplyDelaySeconds": 180,
}
db.commit()
finally:
db.close()
response = client.put(
f"/api/avatar/{avatar_id}",
headers=context["owner_headers"],
json={"config": {"replyStyle": "warm", "creativity": 25}},
).json()
config = response["data"]["config"]
assert config["replyStyle"] == "warm"
assert config["creativity"] == 25
assert config["authorizationPermissions"] == ["chat", "takeover"]
assert config["takeoverReplyDelaySeconds"] == 180
def test_avatar_create_and_delete_require_login_and_ownership(authorization_context):
context = authorization_context
assert client.post("/api/avatar", json={"name": "匿名分身"}).status_code == 401
created = client.post(
"/api/avatar",
headers=context["owner_headers"],
json={"name": "待删除分身"},
)
assert created.status_code == 200
avatar_id = created.json()["data"]["id"]
try:
assert client.delete(
f"/api/avatar/{avatar_id}", headers=context["other_headers"]
).status_code == 403
deleted = client.delete(
f"/api/avatar/{avatar_id}", headers=context["owner_headers"]
).json()
assert deleted["code"] == 200
finally:
db = SessionLocal()
try:
db.query(Avatar).filter(Avatar.id == avatar_id).delete()
db.commit()
finally:
db.close()
@@ -0,0 +1,130 @@
"""Contract tests for the self-hosted BOXIM client."""
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from services.boxim_client import BoxIMClient, BoxIMError
@pytest.fixture
def config():
return {
"HUIHUI_PLATFORM_BASE_URL": "https://open.example/api",
"BOXIM_API_BASE_URL": "https://im.example/api",
"HUIHUI_APP_ID": "test_app",
"HUIHUI_ACCESS_ID": "test_access",
"HUIHUI_ACCESS_SECRET": "test_secret",
}
def _response(payload: dict, status_code: int = 200):
response = MagicMock()
response.status_code = status_code
response.json.return_value = payload
return response
def _client_patch(*, post_payload=None, request_payload=None, status_code=200):
client = AsyncMock()
if post_payload is not None:
client.post.return_value = _response(post_payload, status_code)
if request_payload is not None:
client.request.return_value = _response(request_payload, status_code)
context = AsyncMock()
context.__aenter__.return_value = client
context.__aexit__.return_value = None
return patch("services.boxim_client.httpx.AsyncClient", return_value=context), client
@pytest.mark.asyncio
async def test_exchange_access_token_uses_huihui_bearer_and_signed_form(config):
mocked, client = _client_patch(
post_payload={"code": 0, "data": {"accessToken": "box-token", "accessTokenExpiresIn": 3600}}
)
with mocked:
result = await BoxIMClient(config).exchange_access_token("huihui-token")
assert result["accessToken"] == "box-token"
call = client.post.await_args
assert call.args[0] == "https://open.example/api/im/box/netease"
assert call.kwargs["headers"]["Authorization"] == "Bearer huihui-token"
assert call.kwargs["data"]["appId"] == "test_app"
assert len(call.kwargs["data"]["signature"]) == 32
@pytest.mark.asyncio
async def test_get_self_and_incremental_private_messages_use_boxim_header(config):
client_instance = BoxIMClient(config)
mocked, client = _client_patch(
request_payload={"code": 200, "data": {"id": 42, "nickName": "Owner"}}
)
with mocked:
profile = await client_instance.get_self("box-token")
assert profile["id"] == 42
assert client.request.await_args.kwargs["headers"] == {"accessToken": "box-token"}
mocked, client = _client_patch(
request_payload={"code": 200, "data": [{"id": 101, "sendId": 7, "recvId": 42}]}
)
with mocked:
messages = await client_instance.fetch_private_messages("box-token", "100")
assert messages[0]["id"] == 101
assert client.request.await_args.kwargs["params"] == {"minId": "100"}
@pytest.mark.asyncio
async def test_send_private_message_matches_boxim_payload(config):
mocked, client = _client_patch(
request_payload={"code": 200, "data": {"id": 88, "localId": 12345}}
)
with mocked:
result = await BoxIMClient(config).send_private_message(
"box-token", "77", "你好", local_id="12345"
)
assert result["id"] == 88
call = client.request.await_args
assert call.args[:2] == ("POST", "https://im.example/api/message/private/send")
assert call.kwargs["json"] == {
"localId": 12345,
"recvId": 77,
"content": "你好",
"type": 0,
"receipt": False,
"atUserIds": [],
}
@pytest.mark.asyncio
async def test_mark_private_messages_read_uses_latest_message_id(config):
mocked, client = _client_patch(request_payload={"code": 200, "data": None})
with mocked:
await BoxIMClient(config).mark_private_messages_read("box-token", "77", "101")
call = client.request.await_args
assert call.args[:2] == ("PUT", "https://im.example/api/message/private/readed")
assert call.kwargs["headers"] == {"accessToken": "box-token"}
assert call.kwargs["params"] == {"friendId": 77, "messageId": 101}
@pytest.mark.asyncio
async def test_boxim_auth_error_is_explicit(config):
mocked, _ = _client_patch(
request_payload={"code": 400, "message": "未登录"}, status_code=200
)
with mocked, pytest.raises(BoxIMError) as exc_info:
await BoxIMClient(config).get_self("expired")
assert exc_info.value.auth_error is True
def test_sign_params_include_production_required_fields(config):
params = BoxIMClient(config)._build_sign_params()
assert params["appId"] == "test_app"
assert params["accessId"] == "test_access"
assert params["signType"] == "MD5"
assert params["signVersion"] == "1.0"
assert len(params["nonce"]) == 12
assert len(params["timestamp"]) == 14
assert len(params["signature"]) == 32
assert "accessSecret" not in params
@@ -0,0 +1,71 @@
import ipaddress
import json
import httpx
import pytest
from services.boxim_image_service import (
BoxIMImageError,
download_boxim_image,
parse_boxim_image_url,
)
def test_parse_boxim_image_prefers_origin_and_supports_relative_url():
content = json.dumps({"originUrl": "/files/original.png", "thumbUrl": "/thumb.png"})
assert parse_boxim_image_url(content, base_url="https://im.example/api") == (
"https://im.example/files/original.png"
)
def test_download_boxim_image_streams_public_https(monkeypatch):
monkeypatch.setattr(
"services.boxim_image_service._resolved_addresses",
lambda _host, _port: {ipaddress.ip_address("8.8.8.8")},
)
transport = httpx.MockTransport(
lambda request: httpx.Response(
200,
headers={"content-type": "image/png"},
content=b"png-bytes",
request=request,
)
)
image = download_boxim_image(
json.dumps({"originUrl": "https://cdn.example/case%20photo.png"}),
transport=transport,
)
assert image.content == b"png-bytes"
assert image.filename == "case photo.png"
assert image.mime_type == "image/png"
def test_download_boxim_image_rejects_private_network_url():
with pytest.raises(BoxIMImageError, match="受限网络"):
download_boxim_image(
json.dumps({"originUrl": "https://127.0.0.1/private.png"}),
transport=httpx.MockTransport(lambda request: httpx.Response(200, request=request)),
)
def test_download_boxim_image_stops_oversized_stream(monkeypatch):
monkeypatch.setenv("CHAT_IMAGE_MAX_BYTES", "1024")
monkeypatch.setattr(
"services.boxim_image_service._resolved_addresses",
lambda _host, _port: {ipaddress.ip_address("8.8.8.8")},
)
transport = httpx.MockTransport(
lambda request: httpx.Response(
200,
headers={"content-length": "2048"},
request=request,
)
)
with pytest.raises(BoxIMImageError, match="超过大小限制"):
download_boxim_image(
json.dumps({"originUrl": "https://cdn.example/large.png"}),
transport=transport,
)
@@ -0,0 +1,332 @@
import json
from datetime import datetime, timedelta
from types import SimpleNamespace
from unittest.mock import Mock, patch
import pytest
from fastapi import HTTPException
from fastapi.testclient import TestClient
from database import SessionLocal
from main import app
from models import ChatAttachment
from routers.chat import (
ChatIn,
_answer_denies_available_image,
_attachment_contexts,
_load_chat_attachments,
_resolve_reply,
)
from services.chat_attachment_service import purge_expired_chat_attachments
from services.token_billing import InsufficientTokensError
from services.vision_service import PreparedImage
client = TestClient(app)
GENERAL_RESULT = {
"content": json.dumps({
"category": "general_image",
"summary": "一张包含产品路线图的截图",
"visible_text": "产品路线图",
"key_facts": ["包含三个阶段"],
"uncertainties": [],
"medical": {},
}, ensure_ascii=False),
"usage": {"total_tokens": 120},
}
def test_owner_can_upload_and_cache_image_analysis(authorization_context):
context = authorization_context
prepared = PreparedImage(b"jpeg", "image/jpeg", 100, 80)
with (
patch("routers.chat.prepare_image", return_value=prepared),
patch("routers.chat._run_billed_vision_call", return_value=GENERAL_RESULT),
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/chat/images",
headers=context["owner_headers"],
files={"file": ("roadmap.png", b"image-bytes", "image/png")},
)
assert response.status_code == 200
payload = response.json()["data"]
assert payload["status"] == "ready"
assert payload["category"] == "general_image"
assert payload["summary"] == "一张包含产品路线图的截图"
db = SessionLocal()
try:
stored = db.query(ChatAttachment).filter(ChatAttachment.id == payload["id"]).one()
assert stored.avatar_id == context["avatar"].id
assert stored.extracted_text == "产品路线图"
assert stored.structured_data["key_facts"] == ["包含三个阶段"]
finally:
db.close()
def test_non_owner_cannot_upload_chat_image(authorization_context):
context = authorization_context
response = client.post(
f"/api/avatar/{context['avatar'].id}/chat/images",
headers=context["other_headers"],
files={"file": ("private.png", b"image-bytes", "image/png")},
)
assert response.status_code == 403
def test_image_upload_preserves_insufficient_points_response(authorization_context):
context = authorization_context
with patch(
"routers.chat._analyze_image_bytes",
side_effect=InsufficientTokensError("积分余额不足"),
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/chat/images",
headers=context["owner_headers"],
files={"file": ("private.png", b"image-bytes", "image/png")},
)
assert response.status_code == 402
assert response.json()["detail"] == "积分余额不足"
def test_public_share_can_upload_without_exposing_analysis_details(authorization_context):
context = authorization_context
db = SessionLocal()
try:
avatar = db.get(type(context["avatar"]), context["avatar"].id)
avatar.share_token = f"share-{context['suffix']}"
db.commit()
share_token = avatar.share_token
finally:
db.close()
with (
patch(
"routers.chat.prepare_image",
return_value=PreparedImage(b"jpeg", "image/jpeg", 100, 80),
),
patch("routers.chat._run_billed_vision_call", return_value=GENERAL_RESULT),
):
response = client.post(
f"/api/public/avatar/{share_token}/chat/images",
files={"file": ("visitor.png", b"image-bytes", "image/png")},
)
assert response.status_code == 200
payload = response.json()["data"]
assert payload["status"] == "ready"
assert "structuredData" not in payload
assert "extractedText" not in payload
assert "visionModel" not in payload
assert "ocrModel" not in payload
db = SessionLocal()
try:
stored = db.get(ChatAttachment, payload["id"])
assert stored.uploader_kind == "public"
assert stored.avatar_id == context["avatar"].id
finally:
db.close()
def test_medical_document_uses_ocr_result(authorization_context):
context = authorization_context
general = {
"content": json.dumps({
"category": "medical_document",
"summary": "血常规报告",
"visible_text": "初步文字",
"key_facts": [],
"uncertainties": [],
"medical": {"document_type": "检验报告"},
}, ensure_ascii=False),
"usage": {},
}
ocr = {"content": "白细胞 11.2 x10^9/L", "usage": {}}
with (
patch("routers.chat.prepare_image", return_value=PreparedImage(b"jpeg", "image/jpeg", 100, 80)),
patch("routers.chat._run_billed_vision_call", side_effect=[general, ocr]) as model,
):
response = client.post(
f"/api/avatar/{context['avatar'].id}/chat/images",
headers=context["owner_headers"],
files={"file": ("report.jpg", b"image-bytes", "image/jpeg")},
)
assert response.status_code == 200
attachment_id = response.json()["data"]["id"]
assert model.call_count == 2
assert model.call_args_list[1].kwargs["source"] == "vision_medical_ocr"
db = SessionLocal()
try:
stored = db.query(ChatAttachment).filter(ChatAttachment.id == attachment_id).one()
assert stored.extracted_text == "白细胞 11.2 x10^9/L"
assert stored.ocr_model == "qwen-vl-ocr"
assert "不能替代医生诊断" in stored.warning
finally:
db.close()
def test_attachment_cannot_cross_avatar_boundary(authorization_context):
context = authorization_context
db = SessionLocal()
try:
attachment = ChatAttachment(
avatar_id=context["avatar"].id,
filename="private.jpg",
status="ready",
expires_at=datetime.utcnow() + timedelta(hours=1),
)
db.add(attachment)
db.commit()
body = ChatIn(message="看看图片", attachmentIds=[attachment.id])
with pytest.raises(HTTPException, match="不属于当前分身") as caught:
_load_chat_attachments(db, context["other_avatar"].id, body)
assert caught.value.status_code == 400
finally:
db.close()
def test_expired_attachment_is_removed(authorization_context):
context = authorization_context
db = SessionLocal()
try:
attachment = ChatAttachment(
avatar_id=context["avatar"].id,
filename="expired.jpg",
status="ready",
expires_at=datetime.utcnow() - timedelta(seconds=1),
)
db.add(attachment)
db.commit()
attachment_id = attachment.id
body = ChatIn(message="看看图片", attachmentIds=[attachment_id])
with pytest.raises(HTTPException):
_load_chat_attachments(db, context["avatar"].id, body)
assert db.query(ChatAttachment).filter(ChatAttachment.id == attachment_id).first() is None
finally:
db.close()
def test_cleanup_keeps_unexpired_attachment(authorization_context):
context = authorization_context
now = datetime.utcnow()
db = SessionLocal()
try:
expired = ChatAttachment(
avatar_id=context["avatar"].id,
filename="expired.jpg",
status="ready",
expires_at=now - timedelta(seconds=1),
)
active = ChatAttachment(
avatar_id=context["avatar"].id,
filename="active.jpg",
status="ready",
expires_at=now + timedelta(hours=1),
)
db.add_all([expired, active])
db.commit()
expired_id, active_id = expired.id, active.id
assert purge_expired_chat_attachments(db, now=now) == 1
assert db.get(ChatAttachment, expired_id) is None
assert db.get(ChatAttachment, active_id) is not None
finally:
db.close()
def test_image_context_keeps_standard_answer_authoritative():
avatar = SimpleNamespace(
id="avatar-vision",
name="测试分身",
description="产品顾问",
config={},
)
model = Mock(return_value="标准退款期限是七天;图片显示的是商品包装。")
result = _resolve_reply(
None,
avatar,
"退款期限是多少?",
[],
qa_pairs=[SimpleNamespace(question="退款期限是多少?", answer="七天", enabled=True)],
search_fn=Mock(return_value=[]),
model_client=model,
image_contexts=[{
"id": "attachment",
"filename": "product.jpg",
"category": "general_image",
"summary": "商品包装",
"extractedText": "",
"structuredData": {},
"warning": "",
}],
)
assert result["source"] == "qa"
system = model.call_args.kwargs["messages"][0]["content"]
assert "已确认标准答案" in system
assert "七天" in system
assert "商品包装" in system
assert "标准答题对中的事实优先级高于图片资料" in system
def test_ready_image_context_never_returns_whole_image_access_denial():
avatar = SimpleNamespace(
id="avatar-vision",
name="测试分身",
description="产品顾问",
config={},
)
model = Mock(return_value="抱歉,我无法查看或识别图片,请重新上传。")
result = _resolve_reply(
None,
avatar,
"请看看这张图片",
[],
qa_pairs=[],
search_fn=Mock(return_value=[]),
model_client=model,
image_contexts=[{
"id": "attachment",
"filename": "report.jpg",
"category": "medical_document",
"summary": "一份耳鼻喉科门诊记录",
"extractedText": "主诉:咽痛三天",
"structuredData": {"key_facts": ["主诉为咽痛三天"]},
"warning": "请核对原始资料",
}],
)
assert result["source"] == "vision"
assert "一份耳鼻喉科门诊记录" in result["answer"]
assert "主诉为咽痛三天" in result["answer"]
assert "无法查看" not in result["answer"]
system = model.call_args.kwargs["messages"][0]["content"]
assert "当前会话图片已经成功读取" in system
assert "禁止声称无法查看" in system
def test_image_denial_detector_allows_uncertain_field_in_ready_image():
assert _answer_denies_available_image("我无法查看这张图片") is True
assert _answer_denies_available_image("图片中患者姓名无法辨认,主诉为咽痛三天。") is False
def test_attachment_context_does_not_expose_internal_fields():
row = SimpleNamespace(
id="attachment",
filename="case.jpg",
category="medical_document",
summary="门诊病例",
extracted_text="主诉:咳嗽",
structured_data={"medical": {"chief_complaint": "咳嗽"}},
warning="请核对原文",
)
context = _attachment_contexts([row])[0]
assert context["filename"] == "case.jpg"
assert "avatar_id" not in context
assert "vision_model" not in context
@@ -0,0 +1,87 @@
from unittest.mock import Mock, patch
import httpx
from services.chat_model_config import (
clear_chat_model_config_cache,
get_chat_model_config,
)
def setup_function():
clear_chat_model_config_cache()
def teardown_function():
clear_chat_model_config_cache()
def test_admin_runtime_config_takes_priority(monkeypatch):
monkeypatch.setenv("CHAT_MODEL_CONFIG_URL", "http://config.test/runtime")
monkeypatch.setenv("AVATAR_MODEL_CONFIG_TOKEN", "shared-secret")
response = Mock()
response.raise_for_status.return_value = None
response.json.return_value = {
"data": {
"api_base_url": "https://model.test/v1/",
"api_key": "runtime-key",
"model": "avatar-model",
"vision_model": "avatar-vision-model",
"ocr_model": "avatar-ocr-model",
"max_tokens": 2048,
"timeout_seconds": 42,
}
}
with patch("services.chat_model_config.httpx.get", return_value=response) as request:
config = get_chat_model_config()
assert config.source == "admin"
assert config.api_base_url == "https://model.test/v1"
assert config.model == "avatar-model"
assert config.vision_model == "avatar-vision-model"
assert config.ocr_model == "avatar-ocr-model"
assert config.max_tokens == 2048
request.assert_called_once_with(
"http://config.test/runtime",
headers={"X-Avatar-Config-Token": "shared-secret"},
timeout=5.0,
)
def test_runtime_failure_falls_back_to_environment(monkeypatch):
monkeypatch.setenv("CHAT_MODEL_CONFIG_URL", "http://config.test/runtime")
monkeypatch.setenv("AVATAR_MODEL_CONFIG_TOKEN", "shared-secret")
monkeypatch.setenv("CHAT_API_URL", "https://fallback.test/v1/")
monkeypatch.setenv("CHAT_API_KEY", "fallback-key")
monkeypatch.setenv("CHAT_MODEL", "fallback-model")
monkeypatch.setenv("VISION_MODEL", "fallback-vision")
monkeypatch.setenv("VISION_OCR_MODEL", "fallback-ocr")
monkeypatch.setenv("CHAT_MAX_OUTPUT_TOKENS", "1536")
request = httpx.Request("GET", "http://config.test/runtime")
with patch(
"services.chat_model_config.httpx.get",
side_effect=httpx.ConnectError("offline", request=request),
):
config = get_chat_model_config()
assert config.source == "environment"
assert config.api_base_url == "https://fallback.test/v1"
assert config.api_key == "fallback-key"
assert config.model == "fallback-model"
assert config.vision_model == "fallback-vision"
assert config.ocr_model == "fallback-ocr"
assert config.max_tokens == 1536
def test_runtime_config_is_cached(monkeypatch):
monkeypatch.setenv("CHAT_MODEL_CONFIG_URL", "")
monkeypatch.setenv("CHAT_MODEL", "first-model")
first = get_chat_model_config()
monkeypatch.setenv("CHAT_MODEL", "second-model")
second = get_chat_model_config()
assert first is second
assert second.model == "first-model"

Some files were not shown because too many files have changed in this diff Show More