feat: add FastAPI manager API compatibility baseline

This commit is contained in:
Tyke Chen
2026-07-20 17:00:13 +08:00
parent 7c58fa37b2
commit 804ddb51f2
140 changed files with 47169 additions and 1 deletions
+20 -1
View File
@@ -4,4 +4,23 @@ __pycache__
.env
Dockerfile
tmp/
data/
data/
# Repository-local runtimes and test/build products are several gigabytes and
# are never inputs to either manager-api-fastapi image.
.runtime/
.venv-*/
**/.venv/
**/.test-runtime/
**/.pytest_cache/
**/.mypy_cache/
**/.ruff_cache/
**/node_modules/
**/dist/
**/target/
**/uploadfile/
# Runtime state and local model assets must not enter the Docker build context.
main/xiaozhi-server/mysql/
main/xiaozhi-server/models/
main/xiaozhi-server/data/
+267
View File
@@ -0,0 +1,267 @@
# manager-api FastAPI 兼容性矩阵
> 生成依据:`main/manager-api-fastapi/compatibility/java-routes.json`、
> `main/manager-api-fastapi/compatibility/consumer-routes.json`、`route-surface-results.json`、
> `authenticated-route-results.json`、`contract-results.json` 和当前 Java 源码。接口路径均省略
> 共同前缀 `/xiaozhi`。
## 结论与状态口径
Java 基线共有 **154** 条 Spring MVC 路由;FastAPI 已注册 **154/154100%**,并由
`tests/test_java_route_manifest.py` 对源码清单 freshness、数量和 method/path 注册闭合进行检查。
此外实现 3 条仅由仓库消费者使用、Java Controller 中不存在的兼容路由,因此这 3 条不计入
154 条 Java 覆盖率。三端 188 个调用点均能解析到 FastAPI 路由。
矩阵状态必须按下列含义阅读:
- `结构✓`method/path 已注册且清单闭合;它不等于业务行为逐接口实测。
- `请求面差分✓1`:本行已向隔离 Java/FastAPI 各发送一次缺少鉴权或安全非法输入,精确比较
HTTP status、body 与 Content-Type;最终为 **154/154 通过、0 失败、0 跳过**,且不发送成功写请求。
- `认证业务面差分✓1`:本行已使用有效 DB Token、server-secret 或匿名身份,再向隔离
Java/FastAPI 各发送一次安全业务/校验请求,精确比较 HTTP status、body 与 Content-Type
最终为 **154/154 通过、0 失败、0 跳过**,且不主动发送成功写请求。该状态不等于每条路由的
完整成功生命周期均已差分,完整副作用证据仍以 `差分✓N` 为准。
- `领域✓(x,域级)`:该领域有 service/repository/协议自动测试,但不保证本行每条成功与错误路径
都被直接请求。`领域—` 表示除结构测试外没有可归属的域级直接测试证据。
- `差分✓N`:本行除安全请求面外,还参与了成功、主要错误、协议或数据库副作用的深度对照;
括号说明覆盖面。深度结果为 **49/49 checks 通过、0 失败、0 跳过**,覆盖 **21/154** 条路由。
`差分间接✓` 表示 J125 作为下载链路的 URL 生成步骤被间接覆盖;`差分—` 表示没有深度对照,
不能把 154/154 请求面差分误读成 154 条全部成功路径与副作用都已逐接口对照。
- 所有 `Result<T>` 均表示 `{code,msg,data}` envelope;原 Java 为 HTTP 200 的认证、权限、业务和
参数错误由全局兼容层维持 HTTP 200。二进制/OTA 裸响应在“响应类型”列单独标明。
## 三端消费者闭合
| 消费者 | 调用点 | 唯一结构路由 | 方法分布 |
|---|---:|---:|---|
| `manager-web` | 134 | 130 | DELETE 12、GET 59、POST 40、PUT 23 |
| `manager-mobile` | 46 | 40 | DELETE 3、GET 26、POST 12、PUT 5 |
| `xiaozhi-server` | 8 | 8 | GET 2、POST 6 |
| **合计** | **188** | **140** | — |
### 3 条消费者孤儿兼容路由
| Method/path | 来源 | FastAPI 语义 | 鉴权 | 状态 |
|---|---|---|---|---|
| `GET /api/ping` | manager-mobile 环境设置探活 | `{code:0,msg:"success",data:"pong"}` | 匿名 | 实现✓;consumer resolve✓ |
| `PUT /user/configDevice/{device_id}` | manager-web 遗留设备配置调用 | 按现有设备更新契约处理 body | DB Token | 实现✓;consumer resolve✓ |
| `GET /device/address-book/lookup` | xiaozhi-server 管理客户端 | `callerMac/nickname/answer` 地址簿查询/呼叫兼容别名 | server-secret | 实现✓;consumer resolve✓;device 域测试✓ |
`GET /admin/dict/data/type/FIRMWARE_TYPE` 是动态 Java 路由
`GET /admin/dict/data/type/{dictType}` 的一个字面调用,不是第四条孤儿路由。
## Java 基线静态盘点
- Controller24 个、154 条映射。按 Controller 的路由数为:`AdminController`(5)、`AgentChatHistoryController`(4)、`AgentController`(21)、`AgentMcpAccessPointController`(2)、`AgentSnapshotController`(4)、`AgentTemplateController`(6)、`AgentVoicePrintController`(4)、`ConfigController`(3)、`CorrectWordController`(7)、`DeviceController`(13)、`KnowledgeBaseController`(7)、`KnowledgeFilesController`(8)、`LoginController`(8)、`ModelController`(11)、`ModelProviderController`(5)、`OTAController`(3)、`OTAMagController`(9)、`ServerSideManageController`(2)、`SysDictDataController`(6)、`SysDictTypeController`(5)、`SysParamsController`(5)、`TimbreController`(4)、`VoiceCloneController`(6)、`VoiceResourceController`(6)。
- 数据分层:`entity/` 29 个 Java 文件(28 个 `*Entity.java``BaseEntity`)、`dto/` 58 个、
`vo/` 14 个、`dao/` 29 个、`service/` 树 78 个文件(其中
`service/impl/` 38 个)。FastAPI 对应落在 `schemas/``repositories/`
`services/``routers/``integrations/``jobs/`,没有把跨表事务放进路由。
- MyBatis XML20 个,分别是 `mapper/agent/AgentCorrectWordMappingDao.xml``mapper/agent/AgentDao.xml``mapper/agent/AgentPluginMappingMapper.xml``mapper/agent/AgentSnapshotDao.xml``mapper/agent/AgentTagDao.xml``mapper/agent/AgentTagRelationDao.xml``mapper/agent/AgentTemplateMapper.xml``mapper/agent/AiAgentChatHistoryDao.xml``mapper/correctword/CorrectWordItemDao.xml``mapper/device/DeviceAddressBookDao.xml``mapper/device/DeviceDao.xml``mapper/knowledge/KnowledgeBaseDao.xml``mapper/model/ModelConfigDao.xml``mapper/model/ModelProviderDao.xml``mapper/security/SysUserTokenDao.xml``mapper/sys/SysDictDataDao.xml``mapper/sys/SysDictTypeDao.xml``mapper/sys/SysParamsDao.xml``mapper/sys/SysUserDao.xml``mapper/voiceclone/VoiceCloneDao.xml`
- Liquibase`db.changelog-master.yaml` 含 101 个 `changeSet` 引用,目录中恰有
101 个 SQLPython 部署继续执行这 101 个原始 SQL,不改写历史。
- 定时工作:`DocumentStatusSyncTask` 每次完成后延迟 30 秒,扫描 RAGFlow RUNNING 文档并
回写 SUCCESS/FAIL/CANCEL 与统计;当前 Java 源码另有 `AgentSnapshotRedactionRunner`,启动时
执行一次并在滚动部署期每 15 秒补偿脱敏旧快照。FastAPI 将工作移到独立 jobs 进程,并以
Redis 分布式锁/watchdog 防止多 worker 重复执行。
- 外部集成:RAGFlow dataset/document/chunk/retrieval/upload;阿里云短信;火山语音克隆训练与
音频;声纹 HTTPOpenAI-compatible LLM 摘要/标题;MQTT gateway HTTPMCP/管理动作
WebSocketOTA/WS/MQTT 的 HMAC、Base64、时间戳与下载文件存储。自动测试只访问可重复 mock,
未使用真实付费凭证。
## 154 条 Java→FastAPI 逐接口矩阵
副作用缩写:`DB-R/W`=数据库读/写,`Redis-R/W/DEL`=缓存读/写/失效,`文件-R/W`=文件
读取/写入;外部调用均在 service/integration 层。权限为空时表示只需对应鉴权身份。
| # | Method/path | Java Controller.handler | 请求面 | 响应类型 | 鉴权 / 权限 | DB/Redis/文件/外部副作用 | 实现与测试状态 |
|---:|---|---|---|---|---|---|---|
| J001 | `GET /admin/device/all` | `AdminController.pageDevice` | Query:params:Map<String, Object> | envelope <PageData<UserShowDeviceListVO>> | DB Token / `sys:role:superAdmin` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J002 | `POST /admin/dict/data/delete` | `SysDictDataController.delete` | Body:Long[] | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(dict cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J003 | `GET /admin/dict/data/page` | `SysDictDataController.page` | Query:params:Map<String, Object> | envelope <PageData<SysDictDataVO>> | DB Token / `sys:role:superAdmin` | DB-R; Redis-R/W(dict cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J004 | `POST /admin/dict/data/save` | `SysDictDataController.save` | Body:SysDictDataDTO | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(dict cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J005 | `GET /admin/dict/data/type/{dictType}` | `SysDictDataController.getDictDataByType` | Path:dictType | envelope <List<SysDictDataItem>> | DB Token / `sys:role:normal` | DB-R; Redis-R/W(dict cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J006 | `PUT /admin/dict/data/update` | `SysDictDataController.update` | Body:SysDictDataDTO | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(dict cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J007 | `GET /admin/dict/data/{id}` | `SysDictDataController.get` | Path:id | envelope <SysDictDataVO> | DB Token / `sys:role:superAdmin` | DB-R; Redis-R/W(dict cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J008 | `POST /admin/dict/type/delete` | `SysDictTypeController.delete` | Body:Long[] | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(dict cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J009 | `GET /admin/dict/type/page` | `SysDictTypeController.page` | Query:params:Map<String, Object> | envelope <PageData<SysDictTypeVO>> | DB Token / `sys:role:superAdmin` | DB-R; Redis-R/W(dict cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J010 | `POST /admin/dict/type/save` | `SysDictTypeController.save` | Body:SysDictTypeDTO | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(dict cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J011 | `PUT /admin/dict/type/update` | `SysDictTypeController.update` | Body:SysDictTypeDTO | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(dict cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J012 | `GET /admin/dict/type/{id}` | `SysDictTypeController.get` | Path:id | envelope <SysDictTypeVO> | DB Token / `sys:role:superAdmin` | DB-R; Redis-R/W(dict cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J013 | `POST /admin/params` | `SysParamsController.save` | Body:SysParamsDTO | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-W/DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J014 | `PUT /admin/params` | `SysParamsController.update` | Body:SysParamsDTO | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-W; 外部-配置端点探测(按 paramCode) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J015 | `POST /admin/params/delete` | `SysParamsController.delete` | Body:String[] | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-W/DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J016 | `GET /admin/params/page` | `SysParamsController.page` | Query:params:Map<String, Object> | envelope <PageData<SysParamsDTO>> | DB Token / `sys:role:superAdmin` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J017 | `GET /admin/params/{id}` | `SysParamsController.get` | Path:id | envelope <SysParamsDTO> | DB Token / `sys:role:superAdmin` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J018 | `POST /admin/server/emit-action` | `ServerSideManageController.emitServerAction` | Body:EmitSeverActionDTO | envelope <Boolean> | DB Token / `sys:role:superAdmin` | DB/Redis-R(secret/WS); Redis-W(one-shot); 外部-WebSocket | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J019 | `GET /admin/server/server-list` | `ServerSideManageController.getWsServerList` | — | envelope <List<String>> | DB Token / `sys:role:superAdmin` | DB/Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J020 | `GET /admin/users` | `AdminController.pageUser` | Query:params:Map<String, Object> | envelope <PageData<AdminPageUserVO>> | DB Token / `sys:role:superAdmin` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分✓3(权限/序列化/非法分页) |
| J021 | `PUT /admin/users/changeStatus/{status}` | `AdminController.changeStatus` | Path:status; Body:String[] | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W(user/password/status/token) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J022 | `DELETE /admin/users/{id}` | `AdminController.delete` | Path:id | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W(用户/token/device/agent 级联) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J023 | `PUT /admin/users/{id}` | `AdminController.update` | Path:id | envelope <String> | DB Token / `sys:role:superAdmin` | DB-W(user/password/status/token) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(sys,域级);差分— |
| J024 | `POST /agent` | `AgentController.save` | Body:AgentCreateDTO | envelope <String> | DB Token / `sys:role:normal` | DB-W(含快照/映射/标签事务); Redis-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J025 | `GET /agent/all` | `AgentController.adminAgentList` | Query:params:Map<String, Object> | envelope <PageData<AgentEntity>> | DB Token / `sys:role:superAdmin` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J026 | `POST /agent/audio/{audioId}` | `AgentController.getAudioId` | Path:audioId | envelope <String> | DB Token / `sys:role:normal` | DB-R(audio); Redis-W(one-shot URL) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J027 | `GET /agent/chat-history/download/{uuid}/current` | `AgentChatHistoryController.downloadCurrentSession` | Path:uuid | 流式/二进制 + 原下载 headers | 匿名 / — | DB/Redis-R(one-shot); 文件-R/流式 | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J028 | `GET /agent/chat-history/download/{uuid}/previous` | `AgentChatHistoryController.downloadCurrentSessionWithPrevious` | Path:uuid | 流式/二进制 + 原下载 headers | 匿名 / — | DB/Redis-R(one-shot); 文件-R/流式 | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J029 | `POST /agent/chat-history/getDownloadUrl/{agentId}/{sessionId}` | `AgentChatHistoryController.getDownloadUrl` | Path:agentId,sessionId | envelope <String> | DB Token / — | DB-R(chat/session); Redis-W(download token TTL) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J030 | `POST /agent/chat-history/report` | `AgentChatHistoryController.uploadFile` | Body:AgentChatHistoryReportDTO | envelope <Boolean> | server-secret / — | DB-W(chat/session); server-secret | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J031 | `POST /agent/chat-summary/{sessionId}/save` | `AgentController.generateAndSaveChatSummary` | Path:sessionId | envelope <Void> | server-secret / — | DB-R/W(chat); 外部-OpenAI-compatible LLM | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J032 | `POST /agent/chat-title/{sessionId}/generate` | `AgentController.generateAndSaveChatTitle` | Path:sessionId | envelope <Void> | server-secret / — | DB-R/W(chat); 外部-OpenAI-compatible LLM | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J033 | `GET /agent/list` | `AgentController.getUserAgents` | Query:keyword:String,searchType:String | envelope <List<AgentDTO>> | DB Token / `sys:role:normal` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分✓1 |
| J034 | `GET /agent/mcp/address/{agentId}` | `AgentMcpAccessPointController.getAgentMcpAccessAddress` | Path:agentId | envelope <String> | DB Token / `sys:role:normal` | DB/Redis-R; AES token 生成 | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J035 | `GET /agent/mcp/tools/{agentId}` | `AgentMcpAccessPointController.getAgentMcpToolsList` | Path:agentId | envelope <List<String>> | DB Token / `sys:role:normal` | DB/Redis-R; 外部-WebSocket MCP | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J036 | `GET /agent/play/{uuid}` | `AgentController.playAudio` | Path:uuid | 流式/二进制 + 原下载 headers | 匿名 / — | DB/Redis-R(one-shot); 文件-R/流式 | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J037 | `PUT /agent/saveMemory/{macAddress}` | `AgentController.updateByDeviceId` | Path:macAddress; Body:AgentMemoryDTO | envelope <Void> | DB Token / `sys:role:normal` | DB-W(含快照/映射/标签事务); Redis-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J038 | `POST /agent/tag` | `AgentController.createTag` | Body:Map<String, String> | envelope <AgentTagEntity> | DB Token / `sys:role:normal` | DB-W(含快照/映射/标签事务); Redis-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J039 | `GET /agent/tag/list` | `AgentController.getAllTags` | — | envelope <List<AgentTagDTO>> | DB Token / `sys:role:normal` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J040 | `DELETE /agent/tag/{id}` | `AgentController.deleteTag` | Path:id | envelope <Void> | DB Token / `sys:role:normal` | DB-W(含快照/映射/标签事务); Redis-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J041 | `GET /agent/template` | `AgentController.templateList` | — | envelope <List<AgentTemplateEntity>> | DB Token / `sys:role:normal` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J042 | `POST /agent/template` | `AgentTemplateController.createAgentTemplate` | Body:AgentTemplateEntity | envelope <AgentTemplateEntity> | DB Token / `sys:role:superAdmin` | DB-W(含快照/映射/标签事务); Redis-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J043 | `PUT /agent/template` | `AgentTemplateController.updateAgentTemplate` | Body:AgentTemplateEntity | envelope <AgentTemplateEntity> | DB Token / `sys:role:superAdmin` | DB-W(含快照/映射/标签事务); Redis-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J044 | `POST /agent/template/batch-remove` | `AgentTemplateController.batchRemoveAgentTemplates` | Body:List<String> | envelope <String> | DB Token / `sys:role:superAdmin` | DB-W(含快照/映射/标签事务); Redis-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J045 | `GET /agent/template/page` | `AgentTemplateController.getAgentTemplatesPage` | Query:params:Map<String, Object> | envelope <PageData<AgentTemplateVO>> | DB Token / `sys:role:superAdmin` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J046 | `DELETE /agent/template/{id}` | `AgentTemplateController.deleteAgentTemplate` | Path:id | envelope <String> | DB Token / `sys:role:superAdmin` | DB-W(含快照/映射/标签事务); Redis-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J047 | `GET /agent/template/{id}` | `AgentTemplateController.getAgentTemplateById` | Path:id | envelope <AgentTemplateVO> | DB Token / `sys:role:superAdmin` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J048 | `POST /agent/voice-print` | `AgentVoicePrintController.save` | Body:AgentVoicePrintSaveDTO | envelope <Void> | DB Token / `sys:role:normal` | DB-W; 外部-voiceprint HTTP | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J049 | `PUT /agent/voice-print` | `AgentVoicePrintController.update` | Body:AgentVoicePrintUpdateDTO | envelope <Void> | DB Token / `sys:role:normal` | DB-W; 外部-voiceprint HTTP | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J050 | `GET /agent/voice-print/list/{id}` | `AgentVoicePrintController.list` | Path:id | envelope <List<AgentVoicePrintVO>> | DB Token / `sys:role:normal` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J051 | `DELETE /agent/voice-print/{id}` | `AgentVoicePrintController.delete` | Path:id | envelope <Void> | DB Token / `sys:role:normal` | DB-W; 外部-voiceprint HTTP | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J052 | `GET /agent/{agentId}/snapshots` | `AgentSnapshotController.page` | Path:agentId; Query:params:AgentSnapshotPageDTO | envelope <PageData<AgentSnapshotVO>> | DB Token / `sys:role:normal` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J053 | `DELETE /agent/{agentId}/snapshots/{snapshotId}` | `AgentSnapshotController.deleteSnapshot` | Path:agentId,snapshotId | envelope <Void> | DB Token / `sys:role:normal` | DB-W(含快照/映射/标签事务); Redis-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J054 | `GET /agent/{agentId}/snapshots/{snapshotId}` | `AgentSnapshotController.getSnapshot` | Path:agentId,snapshotId | envelope <AgentSnapshotVO> | DB Token / `sys:role:normal` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J055 | `POST /agent/{agentId}/snapshots/{snapshotId}/restore` | `AgentSnapshotController.restore` | Path:agentId,snapshotId; Body:AgentSnapshotRestoreDTO | envelope <Void> | DB Token / `sys:role:normal` | DB-W(含快照/映射/标签事务); Redis-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J056 | `DELETE /agent/{id}` | `AgentController.delete` | Path:id | envelope <Void> | DB Token / `sys:role:normal` | DB-W(含快照/映射/标签事务); Redis-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J057 | `GET /agent/{id}` | `AgentController.getAgentById` | Path:id | envelope <AgentInfoVO> | DB Token / `sys:role:normal` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J058 | `PUT /agent/{id}` | `AgentController.update` | Path:id; Body:AgentUpdateDTO | envelope <Void> | DB Token / `sys:role:normal` | DB-W(含快照/映射/标签事务); Redis-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J059 | `GET /agent/{id}/chat-history/audio` | `AgentController.getContentByAudioId` | Path:id | envelope <String> | DB Token / `sys:role:normal` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J060 | `GET /agent/{id}/chat-history/user` | `AgentController.getRecentlyFiftyByAgentId` | Path:id | envelope <List<AgentChatHistoryUserVO>> | DB Token / `sys:role:normal` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J061 | `GET /agent/{id}/chat-history/{sessionId}` | `AgentController.getAgentChatHistory` | Path:id,sessionId | envelope <List<AgentChatHistoryDTO>> | DB Token / `sys:role:normal` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J062 | `GET /agent/{id}/sessions` | `AgentController.getAgentSessions` | Path:id; Query:params:Map<String, Object> | envelope <PageData<AgentChatSessionDTO>> | DB Token / `sys:role:normal` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J063 | `GET /agent/{id}/tags` | `AgentController.getAgentTags` | Path:id | envelope <List<AgentTagDTO>> | DB Token / `sys:role:normal` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J064 | `PUT /agent/{id}/tags` | `AgentController.saveAgentTags` | Path:id; Body:Map<String, Object> | envelope <Void> | DB Token / `sys:role:normal` | DB-W(含快照/映射/标签事务); Redis-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(agent,域级);差分— |
| J065 | `POST /config/agent-models` | `ConfigController.getAgentModels` | Body:AgentModelsDTO | envelope <Object> | server-secret / — | DB-R; Redis-R/W(runtime/model/timbre cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(config,域级);差分— |
| J066 | `POST /config/correct-words` | `ConfigController.getCorrectWords` | Body:CorrectWordsDTO | envelope <Object> | server-secret / — | DB-R; Redis-R/W(runtime/model/timbre cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(config,域级);差分— |
| J067 | `POST /config/server-base` | `ConfigController.getConfig` | — | envelope <Object> | server-secret / — | DB-R; Redis-R/W(runtime/model/timbre cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(config,域级);差分✓3(缺失/错误/正确 secret) |
| J068 | `POST /correct-word/file` | `CorrectWordController.createFile` | Body:CorrectWordFileCreateDTO | envelope <CorrectWordFileVO> | DB Token / `sys:role:normal` | DB-W(file/items/mapping 事务) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(correctword,域级);差分✓2(响应/DB |
| J069 | `POST /correct-word/file/batch-delete` | `CorrectWordController.batchDeleteFiles` | Body:List<String> | envelope <Void> | DB Token / `sys:role:normal` | DB-W(file/items/mapping 事务) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(correctword,域级);差分— |
| J070 | `GET /correct-word/file/download/{fileId}` | `CorrectWordController.downloadFile` | Path:fileId | 流式/二进制 + 原下载 headers | DB Token / `sys:role:normal` | DB-R(content); 二进制 | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(correctword,域级);差分✓2(二进制/更新后下载) |
| J071 | `GET /correct-word/file/list` | `CorrectWordController.listFiles` | Query:params:Map<String, Object> | envelope <PageData<CorrectWordFileVO>> | DB Token / `sys:role:normal` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(correctword,域级);差分✓1 |
| J072 | `GET /correct-word/file/select` | `CorrectWordController.listAllFiles` | — | envelope <List<CorrectWordFileVO>> | DB Token / `sys:role:normal` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(correctword,域级);差分— |
| J073 | `DELETE /correct-word/file/{fileId}` | `CorrectWordController.deleteFile` | Path:fileId | envelope <Void> | DB Token / `sys:role:normal` | DB-W(file/items/mapping 事务) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(correctword,域级);差分✓1(级联副作用) |
| J074 | `PUT /correct-word/file/{fileId}` | `CorrectWordController.updateFile` | Path:fileId; Body:CorrectWordFileCreateDTO | envelope <Void> | DB Token / `sys:role:normal` | DB-W(file/items/mapping 事务) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(correctword,域级);差分✓2(响应/DB |
| J075 | `GET /datasets` | `KnowledgeBaseController.getPageList` | Query:name:String,page:Integer,page_size:Integer | envelope <PageData<KnowledgeBaseDTO>> | DB Token / `sys:role:normal` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J076 | `POST /datasets` | `KnowledgeBaseController.save` | Body:KnowledgeBaseDTO | envelope <KnowledgeBaseDTO> | DB Token / `sys:role:normal` | DB-R/W; 外部-RAGFlow HTTP(upload/dataset/document/chunk/retrieval) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J077 | `DELETE /datasets/batch` | `KnowledgeBaseController.deleteBatch` | Query:ids:String | envelope <Void> | DB Token / `sys:role:normal` | DB-R/W; 外部-RAGFlow HTTP(upload/dataset/document/chunk/retrieval) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J078 | `GET /datasets/rag-models` | `KnowledgeBaseController.getRAGModels` | — | envelope <List<ModelConfigEntity>> | DB Token / `sys:role:normal` | DB-R(model config) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J079 | `DELETE /datasets/{dataset_id}` | `KnowledgeBaseController.delete` | Path:dataset_id | envelope <Void> | DB Token / `sys:role:normal` | DB-R/W; 外部-RAGFlow HTTP(upload/dataset/document/chunk/retrieval) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J080 | `GET /datasets/{dataset_id}` | `KnowledgeBaseController.getByDatasetId` | Path:dataset_id | envelope <KnowledgeBaseDTO> | DB Token / `sys:role:normal` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J081 | `PUT /datasets/{dataset_id}` | `KnowledgeBaseController.update` | Path:dataset_id; Body:KnowledgeBaseDTO | envelope <KnowledgeBaseDTO> | DB Token / `sys:role:normal` | DB-R/W; 外部-RAGFlow HTTP(upload/dataset/document/chunk/retrieval) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J082 | `POST /datasets/{dataset_id}/chunks` | `KnowledgeFilesController.parseDocuments` | Path:dataset_id; Body:Map<String, List<String>> | envelope <Void> | DB Token / `sys:role:normal` | DB-R/W; 外部-RAGFlow HTTP(upload/dataset/document/chunk/retrieval) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J083 | `DELETE /datasets/{dataset_id}/documents` | `KnowledgeFilesController.delete` | Path:dataset_id; Body:DocumentDTO.BatchIdReq | envelope <Void> | DB Token / `sys:role:normal` | DB-R/W; 外部-RAGFlow HTTP(upload/dataset/document/chunk/retrieval) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J084 | `GET /datasets/{dataset_id}/documents` | `KnowledgeFilesController.getPageList` | Path:dataset_id; Query:name:String,status:String,page:Integer,page_size:Integer | envelope <PageData<KnowledgeFilesDTO>> | DB Token / `sys:role:normal` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J085 | `POST /datasets/{dataset_id}/documents` | `KnowledgeFilesController.uploadDocument` | Path:dataset_id; Query:name:String,chunkMethod:String,metaFields:String,parserConfig:String; Multipart:file | envelope <KnowledgeFilesDTO> | DB Token / `sys:role:normal` | DB-R/W; 外部-RAGFlow HTTP(upload/dataset/document/chunk/retrieval) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J086 | `GET /datasets/{dataset_id}/documents/status/{status}` | `KnowledgeFilesController.getPageListByStatus` | Path:dataset_id,status; Query:page:Integer,page_size:Integer | envelope <PageData<KnowledgeFilesDTO>> | DB Token / `sys:role:normal` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J087 | `DELETE /datasets/{dataset_id}/documents/{document_id}` | `KnowledgeFilesController.deleteSingle` | Path:dataset_id,document_id | envelope <Void> | DB Token / `sys:role:normal` | DB-R/W; 外部-RAGFlow HTTP(upload/dataset/document/chunk/retrieval) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J088 | `GET /datasets/{dataset_id}/documents/{document_id}/chunks` | `KnowledgeFilesController.listChunks` | Path:dataset_id,document_id; Query:page:Integer,pageSize:Integer,keywords:String,id:String | envelope <ChunkDTO.ListVO> | DB Token / `sys:role:normal` | DB-R/W; 外部-RAGFlow HTTP(upload/dataset/document/chunk/retrieval) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J089 | `POST /datasets/{dataset_id}/retrieval-test` | `KnowledgeFilesController.retrievalTest` | Path:dataset_id; Body:RetrievalDTO.TestReq | envelope <RetrievalDTO.ResultVO> | DB Token / `sys:role:normal` | DB-R/W; 外部-RAGFlow HTTP(upload/dataset/document/chunk/retrieval) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(knowledge,域级);差分— |
| J090 | `PUT /device/address-book/alias` | `DeviceController.updateAlias` | Body:DeviceAddressBookAliasDTO | envelope <Void> | DB Token / `sys:role:normal` | DB-W(device/bind/address-book); Redis-R/W | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J091 | `GET /device/address-book/call` | `DeviceController.callByNickname` | Query:callerMac:String,nickname:String,answer:boolean | envelope <Map<String, Object>> | server-secret / — | DB-R; 外部-MQTT gateway HTTP; server-secret | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J092 | `PUT /device/address-book/permission` | `DeviceController.updatePermission` | Body:DeviceAddressBookPermissionDTO | envelope <Void> | DB Token / `sys:role:normal` | DB-W(device/bind/address-book); Redis-R/W | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J093 | `GET /device/address-book/{macAddress}` | `DeviceController.getAddressBook` | Path:macAddress | envelope <Object> | DB Token / `sys:role:normal` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J094 | `GET /device/bind/{agentId}` | `DeviceController.getUserDevices` | Path:agentId | envelope <List<UserShowDeviceListVO>> | DB Token / `sys:role:normal` | DB-R; Redis-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分✓1 |
| J095 | `POST /device/bind/{agentId}` | `DeviceController.forwardToMqttGateway` | Path:agentId; Body:String | envelope <String> | DB Token / `sys:role:normal` | DB/Redis-R; 外部-MQTT gateway HTTP + daily auth | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J096 | `POST /device/bind/{agentId}/{deviceCode}` | `DeviceController.bindDevice` | Path:agentId,deviceCode | envelope <Void> | DB Token / `sys:role:normal` | DB-W(device/bind/address-book); Redis-R/W | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J097 | `POST /device/manual-add` | `DeviceController.manualAddDevice` | Body:DeviceManualAddDTO | envelope <Void> | DB Token / `sys:role:normal` | DB-W(device/bind/address-book); Redis-R/W | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J098 | `POST /device/register` | `DeviceController.registerDevice` | Body:DeviceRegisterDTO | envelope <String> | DB Token / — | DB-W(device/bind/address-book); Redis-R/W | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J099 | `POST /device/tools/call/{deviceId}` | `DeviceController.callDeviceTool` | Path:deviceId; Body:DeviceToolsCallReqDTO | envelope <Object> | DB Token / `sys:role:normal` | DB/Redis-R; 外部-MQTT gateway HTTP + daily auth | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J100 | `POST /device/tools/list/{deviceId}` | `DeviceController.getDeviceTools` | Path:deviceId | envelope <Object> | DB Token / `sys:role:normal` | DB/Redis-R; 外部-MQTT gateway HTTP + daily auth | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分✓2(响应/外呼格式) |
| J101 | `POST /device/unbind` | `DeviceController.unbindDevice` | Body:DeviceUnBindDTO | envelope <Void> | DB Token / `sys:role:normal` | DB-W(device/bind/address-book); Redis-R/W | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J102 | `PUT /device/update/{id}` | `DeviceController.updateDeviceInfo` | Path:id; Body:DeviceUpdateDTO | envelope <Void> | DB Token / `sys:role:normal` | DB-W(device/bind/address-book); Redis-R/W | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分✓3(上下界/UTF-16 长度) |
| J103 | `PUT /models/default/{id}` | `ModelController.setDefaultModel` | Path:id | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(model/config cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J104 | `PUT /models/enable/{id}/{status}` | `ModelController.enableModelConfig` | Path:id,status | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(model/config cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J105 | `GET /models/list` | `ModelController.getModelConfigList` | Query:modelType:String,modelName:String,page:String,limit:String | envelope <PageData<ModelConfigDTO>> | DB Token / `sys:role:superAdmin` | DB-R; Redis-R/W(model cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J106 | `GET /models/llm/names` | `ModelController.getLlmModelCodeList` | Query:modelName:String | envelope <List<LlmModelBasicInfoDTO>> | DB Token / `sys:role:normal` | DB-R; Redis-R/W(model cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J107 | `GET /models/names` | `ModelController.getModelNames` | Query:modelType:String,modelName:String | envelope <List<ModelBasicInfoDTO>> | DB Token / `sys:role:normal` | DB-R; Redis-R/W(model cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J108 | `GET /models/provider` | `ModelProviderController.getListPage` | Query:modelProviderDTO:ModelProviderDTO,page:String,limit:String | envelope <PageData<ModelProviderDTO>> | DB Token / `sys:role:superAdmin` | DB-R; Redis-R/W(model cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分✓1 |
| J109 | `POST /models/provider` | `ModelProviderController.add` | Body:ModelProviderDTO | envelope <ModelProviderDTO> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(model/config cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分✓1(约束集合) |
| J110 | `PUT /models/provider` | `ModelProviderController.edit` | Body:ModelProviderDTO | envelope <ModelProviderDTO> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(model/config cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J111 | `POST /models/provider/delete` | `ModelProviderController.delete` | Body:List<String> | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(model/config cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J112 | `GET /models/provider/plugin/names` | `ModelProviderController.getPluginNameList` | — | envelope <List<ModelProviderDTO>> | DB Token / — | DB-R; Redis-R/W(model cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J113 | `DELETE /models/{id}` | `ModelController.deleteModelConfig` | Path:id | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(model/config cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J114 | `GET /models/{id}` | `ModelController.getModelConfig` | Path:id | envelope <ModelConfigDTO> | DB Token / `sys:role:superAdmin` | DB-R; Redis-R/W(model cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J115 | `GET /models/{modelId}/voices` | `ModelController.getVoiceList` | Path:modelId; Query:voiceName:String | envelope <List<VoiceDTO>> | DB Token / `sys:role:normal` | DB-R; Redis-R/W(model cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J116 | `GET /models/{modelType}/provideTypes` | `ModelController.getModelProviderList` | Path:modelType | envelope <List<ModelProviderDTO>> | DB Token / `sys:role:superAdmin` | DB-R; Redis-R/W(model cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J117 | `POST /models/{modelType}/{provideCode}` | `ModelController.addModelConfig` | Path:modelType,provideCode; Body:ModelConfigBodyDTO | envelope <ModelConfigDTO> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(model/config cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J118 | `PUT /models/{modelType}/{provideCode}/{id}` | `ModelController.editModelConfig` | Path:modelType,provideCode,id; Body:ModelConfigBodyDTO | envelope <ModelConfigDTO> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(model/config cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(model,域级);差分— |
| J119 | `GET /ota/` | `OTAController.getOTA` | — | 裸 text/plain | 匿名 / — | — | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分✓1MIME/body |
| J120 | `POST /ota/` | `OTAController.checkOTAVersion` | Header:Device-Id,Client-Id; Body:DeviceReportReqDTO | 裸 application/json | 匿名 / — | DB/Redis-R(设备/固件/配置); HMAC/Base64/时间戳凭证 | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分✓4(必填/格式/凭证/密码学) |
| J121 | `POST /ota/activate` | `OTAController.activateDevice` | Header:Device-Id,Client-Id | 裸 application/json | 匿名 / — | DB-R/W(device activation); Redis-R/W(TTL) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分✓3 |
| J122 | `GET /otaMag` | `OTAMagController.page` | Query:params:Map<String, Object> | envelope <PageData<OtaEntity>> | DB Token / `sys:role:superAdmin` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J123 | `POST /otaMag` | `OTAMagController.save` | Body:OtaEntity | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W(OTA metadata) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分✓2(响应/DB |
| J124 | `GET /otaMag/download/{uuid}` | `OTAMagController.downloadFirmware` | Path:uuid | 流式/二进制 + 原下载 headers | 匿名 / — | Redis-R/W(一次性/次数); 文件-R/流式 | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分✓4(次数限制及二进制) |
| J125 | `GET /otaMag/getDownloadUrl/{id}` | `OTAMagController.getDownloadUrl` | Path:id | envelope <String> | DB Token / `sys:role:superAdmin` | DB-R; Redis-W(download token TTL) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分间接✓(供下载链路) |
| J126 | `POST /otaMag/upload` | `OTAMagController.uploadFirmware` | Multipart:file | envelope <String> | DB Token / `sys:role:superAdmin` | 文件-W(MD5/扩展名/大小) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分✓2(上传/扩展名错误) |
| J127 | `POST /otaMag/uploadAssetsBin` | `OTAMagController.uploadAssetsBin` | Multipart:file | envelope <String> | DB Token / `sys:role:normal` | 文件-W(MD5/扩展名/大小) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J128 | `DELETE /otaMag/{id}` | `OTAMagController.delete` | Path:id | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W(OTA metadata); 文件-DEL | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J129 | `GET /otaMag/{id}` | `OTAMagController.get` | Path:id | envelope <OtaEntity> | DB Token / `sys:role:superAdmin` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J130 | `PUT /otaMag/{id}` | `OTAMagController.update` | Path:id; Body:OtaEntity | envelope <?> | DB Token / `sys:role:superAdmin` | DB-W(OTA metadata) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(device,域级);差分— |
| J131 | `GET /ttsVoice` | `TimbreController.page` | Query:params:Map<String, Object> | envelope <PageData<TimbreDetailsVO>> | DB Token / `sys:role:superAdmin` | DB-R; Redis-R/W(timbre cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(timbre,域级);差分— |
| J132 | `POST /ttsVoice` | `TimbreController.save` | Body:TimbreDataDTO | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(timbre/config cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(timbre,域级);差分— |
| J133 | `POST /ttsVoice/delete` | `TimbreController.delete` | Body:String[] | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(timbre/config cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(timbre,域级);差分— |
| J134 | `PUT /ttsVoice/{id}` | `TimbreController.update` | Path:id; Body:TimbreDataDTO | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W; Redis-DEL(timbre/config cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(timbre,域级);差分— |
| J135 | `GET /user/captcha` | `LoginController.captcha` | Query:uuid:String | image/gif 二进制 | 匿名 / — | Redis-W(captcha TTL); GIF | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(security,域级);差分— |
| J136 | `PUT /user/change-password` | `LoginController.changePassword` | Body:PasswordDTO | envelope <?> | DB Token / — | DB-W(user/token); Redis-R/DEL(SMS) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(security,域级);差分— |
| J137 | `GET /user/info` | `LoginController.info` | — | envelope <UserDetail> | DB Token / — | DB-R; Redis-R/W(cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(security,域级);差分✓9(七语言/过期 Token/Long |
| J138 | `POST /user/login` | `LoginController.login` | Body:LoginDTO | envelope <TokenDTO> | 匿名 / — | DB-R/W(token); Redis-R/DEL(captcha) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(security,域级);差分— |
| J139 | `GET /user/pub-config` | `LoginController.pubConfig` | — | envelope <Map<String, Object>> | 匿名 / — | DB-R; Redis-R/W(cache) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(security,域级);差分✓1 |
| J140 | `POST /user/register` | `LoginController.register` | Body:LoginDTO | envelope <Void> | 匿名 / — | DB-W(user/token); Redis-R/DEL(SMS) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(security,域级);差分— |
| J141 | `PUT /user/retrieve-password` | `LoginController.retrievePassword` | Body:RetrievePasswordDTO | envelope <?> | 匿名 / — | DB-W(user/token); Redis-R/DEL(SMS) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(security,域级);差分— |
| J142 | `POST /user/smsVerification` | `LoginController.smsVerification` | Body:SmsVerificationDTO | envelope <Void> | 匿名 / — | Redis-R/W(TTL/频控); 外部-Aliyun SMS | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(security,域级);差分— |
| J143 | `GET /voiceClone` | `VoiceCloneController.page` | Query:params:Map<String, Object> | envelope <PageData<VoiceCloneResponseDTO>> | DB Token / `sys:role:normal` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(voiceclone,域级);差分— |
| J144 | `POST /voiceClone/audio/{id}` | `VoiceCloneController.getAudioId` | Path:id | envelope <String> | DB Token / `sys:role:normal` | DB-R; Redis-W(one-shot URL) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(voiceclone,域级);差分— |
| J145 | `POST /voiceClone/cloneAudio` | `VoiceCloneController.cloneAudio` | Body:Map<String, String> | envelope <String> | DB Token / `sys:role:normal` | DB-R/W(train state); 文件-W; 外部-火山语音克隆 HTTP | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(voiceclone,域级);差分— |
| J146 | `GET /voiceClone/play/{uuid}` | `VoiceCloneController.playVoice` | Path:uuid | 流式/二进制 + 原下载 headers | 匿名 / — | Redis-R/DEL(one-shot); 文件/外部音频-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(voiceclone,域级);差分— |
| J147 | `POST /voiceClone/updateName` | `VoiceCloneController.updateName` | Body:Map<String, String> | envelope <String> | DB Token / `sys:role:normal` | DB-W(train record name) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(voiceclone,域级);差分— |
| J148 | `POST /voiceClone/upload` | `VoiceCloneController.uploadVoice` | Query:id:String; Multipart:voiceFile | envelope <String> | DB Token / `sys:role:normal` | DB-R/W(train state); 文件-W; 外部-火山语音克隆 HTTP | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(voiceclone,域级);差分— |
| J149 | `GET /voiceResource` | `VoiceResourceController.page` | Query:params:Map<String, Object> | envelope <PageData<VoiceCloneResponseDTO>> | DB Token / `sys:role:superAdmin` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(voiceclone,域级);差分— |
| J150 | `POST /voiceResource` | `VoiceResourceController.save` | Body:VoiceCloneDTO | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W(voice resource) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(voiceclone,域级);差分— |
| J151 | `GET /voiceResource/ttsPlatforms` | `VoiceResourceController.getTtsPlatformList` | — | envelope <List<Map<String, Object>>> | DB Token / `sys:role:superAdmin` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(voiceclone,域级);差分— |
| J152 | `GET /voiceResource/user/{userId}` | `VoiceResourceController.getByUserId` | Path:userId | envelope <List<VoiceCloneResponseDTO>> | DB Token / `sys:role:normal` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(voiceclone,域级);差分— |
| J153 | `DELETE /voiceResource/{id}` | `VoiceResourceController.delete` | Path:id | envelope <Void> | DB Token / `sys:role:superAdmin` | DB-W(voice resource) | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(voiceclone,域级);差分— |
| J154 | `GET /voiceResource/{id}` | `VoiceResourceController.get` | Path:id | envelope <VoiceCloneResponseDTO> | DB Token / `sys:role:superAdmin` | DB-R | 结构✓;请求面差分✓1;认证业务面差分✓1;领域✓(voiceclone,域级);差分— |
## 已观测差异与未覆盖面
- 154 条安全请求面差分最终全部一致。首轮曾发现 5 个空 Body 映射差异;修复 FastAPI 对
Spring `HttpMessageNotReadableException` 的 code-500 语义后,重新从零执行才得到 154/154。
- 154 条认证业务面差分最终全部一致;该轮使用有效鉴权与安全业务/校验输入,在不主动成功
写入的前提下逐路由对照。证据是 `authenticated-route-results.json`,渲染器会在结果不是
154/154、存在失败或跳过时硬失败。
- 2026-07-20 的隔离差分报告未在 49 个 checks 中观测到响应/所选 headers/数据库副作用
不一致;证据是 `main/manager-api-fastapi/compatibility/contract-results.json`,不是人工推断。
- Hibernate Validator 的 `ConstraintViolation Set` 首条消息无稳定顺序;模型 provider 必填
用例比较“消息属于 Java 声明约束集合”与相同错误码,而不伪造一个固定顺序。
- OTA 时间戳/token 是动态值,差分先比较归一化结构,再分别校验两端 HMAC/Base64 密码学
有效性;这属于有意的测试归一化,不是声称字节恒等。
- 深度差分未直接命中的 133 条中,J125 是下载链路间接覆盖,另 132 条标为 `差分—`
它们有请求面、认证业务面与所属领域测试,但尚无逐路由成功+主要错误+副作用深度对照,不能据此宣称
每一种业务状态均已逐接口行为等价。
- FastAPI 额外提供上述 3 条消费者兼容路由与 live/ready 健康检查;它们没有 Java
Controller 基线,属于明确、可回退的加法差异。
- Java 把定时任务放在 Spring 进程;FastAPI 使用独立 jobs 进程和 Redis 分布式锁。这是
部署拓扑差异,业务状态和幂等目标保持一致。
- RAGFlow、阿里云短信、火山语音克隆、真实声纹、真实 LLM、真实 MQTT/MCP/WS 均未用
生产凭证联调;自动化只证明 mock 请求格式、超时/错误映射/重试中的已覆盖场景。
## 可复现检查
```bash
cd main/manager-api-fastapi
.venv/bin/python scripts/extract_java_routes.py --output compatibility/java-routes.json
.venv/bin/python scripts/extract_consumer_routes.py > /tmp/consumer-routes.json
.venv/bin/pytest -q tests/test_java_route_manifest.py tests/test_consumer_route_manifest.py tests/test_compatibility_document.py
```
逐接口差分的启动、隔离库、mock 与执行命令见 `docs/manager-api-fastapi-test-report.md`
本文件只陈述已落盘的结果,不把缺少真实密钥的外部联调列为通过。
+290
View File
@@ -0,0 +1,290 @@
# manager-api 到 FastAPI 迁移说明
## 目标与边界
`main/manager-api-fastapi``main/manager-api` 的兼容替代实现。迁移只替换管理 API
进程,不改变 MySQL 表、Liquibase 历史、Redis 业务语义,也不要求 manager-web、
manager-mobile 或 xiaozhi-server 修改现有 URL。Java 实现继续保留,作为行为基线和
回滚实现。
本次迁移不采用双写。灰度期间,每个业务域在任一时刻只有一个写入方;读流量可以按
请求切分,写流量必须按业务域整体切换。
## 架构
```mermaid
flowchart LR
C["Web / Mobile / xiaozhi-server / Device"] --> N["Nginx /xiaozhi"]
N --> A["FastAPI API workers"]
A --> R["Routers"]
R --> S["Services / transaction boundaries"]
S --> P["Repositories / SQLAlchemy async"]
S --> I["External integration clients"]
P --> M[("Existing MySQL schema")]
S --> D[("Redis Java compatibility layer")]
J["Standalone jobs worker"] --> S
L["Original Liquibase changelog"] --> X["Migration runner"]
X --> M
```
代码按职责分为:
- `app/routers`URL、HTTP 方法、Header/Query/Body 绑定和响应形态。
- `app/services`:权限后的业务规则、事务边界、跨表级联和外部调用编排。
- `app/repositories`:现有 MySQL 表上的参数化 SQL、行锁和分页。
- `app/schemas`:兼容现有字段别名及请求模型。
- `app/integrations`LLM、MQTT gateway、MCP、RAGFlow、语音克隆和声纹客户端。
- `app/jobs`:独立的定时任务进程;API worker 不启动定时任务。
- `app/core`:数据库 Token 认证、国际化、Java JSON/Long 兼容、SM2、Redis 编解码、
Snowflake ID 和健康检查。
所有业务路由保留 `/xiaozhi` 前缀。普通 API 返回 `{code,msg,data}`;认证过滤器、
业务异常和请求校验均由兼容处理器转换,不把 FastAPI 默认的 401、403 或 422 直接
暴露给客户端。OTA、播放和文件下载等原本返回裸 JSON、文本或二进制的接口继续保持
其原始响应类型。
## 数据库兼容策略
### Schema 与迁移
原目录 `main/manager-api/src/main/resources/db/changelog/` 仍是唯一 schema source of
truth。不得把既有 changeset 改写成 Alembic,也不得修改已应用 changeset 的 ID、作者或
校验和。
Python 部署必须先运行独立迁移镜像或 `scripts/run-migrations.sh`。迁移 runner 直接打包
原 Java Liquibase 资源,因此与保留的 Java 服务使用相同的 `DATABASECHANGELOG` 历史。
宿主机执行脚本时,`JAVA_RESOURCES_DIR` 必须指向仓库中的
`main/manager-api/src/main/resources`;不要使用只存在于应用镜像内的
`/opt/xiaozhi/java-resources`。本机隔离数据库的完整命令见“配置与启动 / 本地”。
`docker compose``manager-api-fastapi``manager-api-jobs` 都等待
`manager-api-migrate` 成功退出,避免应用在未完成迁移时接流量。
### 事务、锁与 ID
- Service 层在一次事务中完成智能体创建/更新/删除、插件/标签/纠错词映射、设备解绑、
快照恢复及其他多表操作;异常时显式回滚。
- 智能体更新和恢复先对 `ai_agent` 执行 `SELECT ... FOR UPDATE`,序列化同一智能体的
快照版本分配和状态令牌校验。
- 快照版本仍使用 `(agent_id, version_no)` 唯一约束及同一事务内的 `MAX+1` 插入;行锁
防止并发恢复或更新竞争。
- 需要 Long 主键的管理表继续使用与 Java epoch/node/sequence 布局一致的 Snowflake
生成器;原本使用 32 位 UUID 的业务表仍使用无连字符 UUID。
- 不新增数据库外键,也不改表、索引、字符集或 MySQL 类型。
### Java/Python 并存
并存期必须给写请求建立确定的域路由,例如 `agent/*` 全部指向一个实现,不能把同一域
中的增删改请求在 Java 与 Python 之间随机分配。建议域切换顺序为只读配置、系统管理、
模型/音色、设备、智能体、知识库和外部集成。域回切前先停止该域的新写入并等待在途
请求完成。
## Redis 兼容策略
默认继续使用 Java 已有 key 名称和 TTL。`app/core/redis.py` 实现 Spring Data
`RedisSerializer.json()` 使用的 Jackson wire format,包括:
- Map 的 `@class`、List/Set 的 wrapper-array
- Object 槽位中 Long 的 `java.lang.Long` 包装;
- `java.util.Date` 的 epoch 毫秒包装;
- Java DTO/Entity 缓存所需的具体类名和字段类型;
- Hash 缓存写入后的 86400 秒默认 TTL。
该兼容层使 Java 回滚进程可以读取 FastAPI 写入的缓存。应用启动和正常测试不会执行
`FLUSHALL`;隔离测试脚本的 `reset` 只操作其自建 Redis 实例。定时任务使用 Redis
分布式锁和自动续租 watchdog,即使启动多个 jobs 容器,同一个任务也只有一个执行者。
## 安全与协议兼容
- 用户 Token 仍保存在 `sys_user_token`,按数据库过期时间校验;没有替换为 JWT。
- 登录密文继续使用 SM2 C1C3C2,黄金向量由 Java 和 Python 双向解密测试校验。
- `/config/*`、聊天记录上报/摘要/标题及地址簿内部接口继续校验数据库中的
`server.secret` Bearer 值。
- OTA、WebSocket 和 MQTT 保留原 HMAC、Base64、时间戳、Client-Id/Device-Id 及 token
格式。
- 密钥只通过环境变量或原参数表提供;`.env.example` 和部署文档不包含真实凭证。
## 配置与启动
### 本地
`.env.example` 是容器 Compose 模板,其中的 `mysql``redis` 是 Compose 服务名,
`/opt/xiaozhi/java-resources` 是镜像内路径,不能原样复制后用于宿主机进程。下面的流程
显式使用 `127.0.0.1:13316` 上的隔离 MySQL、`127.0.0.1:16379` 上的隔离 Redis,以及
仓库原始 Liquibase/i18n 资源;不会连接或迁移现有开发数据库:
```bash
cd main/manager-api-fastapi
uv sync --locked
./scripts/isolated-env.sh start
LIQUIBASE_URL='jdbc:mysql://127.0.0.1:13316/manager_fastapi_test?useUnicode=true&characterEncoding=UTF-8&serverTimezone=Asia/Shanghai&allowMultiQueries=true' \
LIQUIBASE_USERNAME='xiaozhi_test' \
LIQUIBASE_PASSWORD='isolated-test-only' \
JAVA_RESOURCES_DIR="$PWD/../manager-api/src/main/resources" \
MAVEN_BIN="$PWD/../../.runtime/maven/bin/mvn" \
MAVEN_LOCAL_REPOSITORY="$PWD/../../.runtime/m2" \
JAVA_HOME="$PWD/../../.runtime/jdk" \
./scripts/run-migrations.sh
eval "$(./scripts/isolated-env.sh env)"
export APP_ENVIRONMENT=development
export APP_DATABASE_URL="$TEST_FASTAPI_DATABASE_URL"
export APP_REDIS_URL="$TEST_FASTAPI_REDIS_URL"
export APP_JAVA_RESOURCES_DIR="$PWD/../manager-api/src/main/resources"
export APP_UPLOAD_DIR="$PWD/.test-runtime/local-uploadfile"
mkdir -p "$APP_UPLOAD_DIR"
./scripts/start-api.sh
```
上述 `eval` 会得到明确的 localhost URLFastAPI 测试库和 Redis DB 2)。定时任务必须
作为单独进程启动;新终端需要重复 `eval` 及四个 `APP_*` 路径/URL 导出,不能让 jobs
进程落回 `.env.example` 的 Docker DNS
```bash
cd main/manager-api-fastapi
eval "$(./scripts/isolated-env.sh env)"
export APP_ENVIRONMENT=development
export APP_DATABASE_URL="$TEST_FASTAPI_DATABASE_URL"
export APP_REDIS_URL="$TEST_FASTAPI_REDIS_URL"
export APP_JAVA_RESOURCES_DIR="$PWD/../manager-api/src/main/resources"
export APP_UPLOAD_DIR="$PWD/.test-runtime/local-uploadfile"
./scripts/start-jobs.sh
```
API 默认监听 `0.0.0.0:8002`,兼容根路径为
`http://127.0.0.1:8002/xiaozhi``APP_WORKERS` 可以大于 1;任务不会随 API worker
复制。验证结束后运行 `./scripts/isolated-env.sh stop`。生产环境不得设置
`APP_ALLOW_START_WITHOUT_DEPENDENCIES=true`
仓库统一启动脚本继承当前 shell 的上述环境变量;FastAPI 是默认实现,保留的 Java
实现可直接用于本地回滚:
```bash
cd "$(git rev-parse --show-toplevel)"
scripts/restart-local-services.sh --manager-api fastapi --wait 180
scripts/restart-local-services.sh --manager-api java --wait 180
```
### 容器
下面的 `mysql``redis` 只在容器网络确实提供对应 DNS 名时有效,否则必须替换为该网络
可访问的真实主机名。API 镜像内的 Java 资源路径是 `/opt/xiaozhi/java-resources`,迁移
镜像则把同一仓库资源打包到 `/migration/java-resources`。以下示例中的凭证必须替换,
并应通过部署平台的 secret 注入而不是提交到仓库:
```bash
cd main/manager-api-fastapi
export LIQUIBASE_URL='jdbc:mysql://mysql:3306/xiaozhi_esp32_server?serverTimezone=Asia/Shanghai'
export MYSQL_USER='xiaozhi'
export MYSQL_PASSWORD='replace-me'
export FASTAPI_DATABASE_URL='mysql+asyncmy://xiaozhi:replace-me@mysql:3306/xiaozhi_esp32_server?charset=utf8mb4'
export REDIS_URL='redis://redis:6379/0'
export MANAGER_API_UPSTREAM='manager-api-fastapi:8002'
docker compose build
docker compose up -d
```
`MANAGER_API_UPLOAD_SOURCE` 指向保留 Java 服务的宿主机上传目录,必须在启动前让 Java
运行用户与容器 UID 10001 都具备读写和目录遍历权限;不要盲目 `chown -R` 导致 Java 失去
访问权,应使用部署环境的共享组或 ACL。空 named volume 在 Docker 通常会继承镜像中
`/data/uploads` 的 UID,但并非所有 OCI runtime 都实现相同 copy-up 语义;必须以容器内
UID 10001 做一次写入预检。`/xiaozhi/health/ready` 同时检查上传目录可写性,权限不正确时
返回 HTTP 503 和 `data.uploads=false`,不得绕过该检查接入流量。Apple Container 的一次性
卷初始化实测命令和结果记录在测试报告中。
`MANAGER_API_UPSTREAM` 默认值是 `manager-api-fastapi:8002`。修改变量后必须重建 Nginx
容器才能重新渲染配置;切到 FastAPI 和整服务回滚到 Java 的命令分别为:
```bash
MANAGER_API_UPSTREAM='manager-api-fastapi:8002' \
docker compose up -d --no-deps --force-recreate manager-api-nginx
MANAGER_API_UPSTREAM='<Nginx 容器可访问的 Java 主机名或 IP>:8002' \
docker compose up -d --no-deps --force-recreate manager-api-nginx
```
Java 地址必须能从 Nginx 容器网络解析和访问。内置 Nginx 的这个变量会切换整个
`/xiaozhi`;逐业务域灰度需要在上层网关按路径配置两个 upstream,仍须遵守“同一业务域
只有一个写入方”。
API 镜像(包括 jobs 命令)和迁移镜像都声明 `USER 10001:10001`,以非 root 用户运行。
Compose 还为 API/jobs 设置只读根文件系统和 `/tmp` tmpfs`/data/uploads` 是它们唯一的
持久写目录,并通过 `/app/uploadfile` 符号链接兼容数据库中的 Java 相对路径。迁移容器是
一次性非 root 进程,但 Compose 没有把它标为只读根文件系统。Nginx 镜像没有声明
非 root `USER`,不能将其描述为非 root 镜像;它通过 Compose 的 `read_only: true`
`/var/cache/nginx``/var/run``/tmp` 三个 tmpfs 加固。Nginx 保留 `/xiaozhi/` 路径,
关闭上传请求缓冲并设置 100 MiB 上限。健康检查分为:
- `/xiaozhi/health/live`:进程存活;
- `/xiaozhi/health/ready`MySQL、Redis 和上传目录写权限都可用才返回 HTTP 200,否则
HTTP 503`data.database``data.redis``data.uploads` 可直接定位失败项。
SIGTERM 触发 Uvicorn 优雅关闭;compose 给 API 40 秒清理在途请求和连接池。
### 容器验证边界
本次仓库内的实际容器运行验证使用 Apple Container 1.0.0,覆盖镜像构建、隔离 MySQL
迁移、双 API worker、独立 jobs、Nginx 路由、只读文件系统、上传卷和 SIGTERM 优雅
关闭。`docker-compose.yml` 已通过 YAML 解析和自动化部署断言,但当前验证主机没有执行
`docker compose up`,因此不能把 Apple Container 的运行结果表述为 Docker Compose
端到端通过。镜像摘要、实际命令和运行结果见
`docs/manager-api-fastapi-test-report.md`;在目标 Docker/Compose 环境切流前,仍需按上面
命令执行一次迁移、ready 检查和 Nginx upstream 冒烟测试。
## 灰度切流
1. 备份当前配置并确认 Java 基线健康;不要停止 Java。
2. 对目标数据库执行原 Liquibase runner,确认 changeset 数量和校验和无差异。
3. 启动 FastAPI API 但暂不接写流量,检查 live/ready、日志和外部 mock。Python jobs
只在隔离环境验证后停止,生产中暂不持续运行,避免与 Java 调度器同时写入。
4. 先镜像或回放脱敏的只读请求,比较状态、Body、Header、数据库读结果和缓存读取。
5. 按业务域把只读流量从 1% 提升到 10%、50%、100%,监控错误码、P95、数据库连接、
Redis 命中与外部服务错误。
6. 对一个完整业务域建立维护窗口,停止该域 Java 新写入,等待在途事务结束,然后把该域
写路由切到 FastAPI。记录切换时间和最后写入方。
7. 逐域重复;稳定观察期内保留 Java 镜像、配置和回切路由。
8. 所有域稳定后才把 jobs 所有权切给 Python;Java 的定时任务进程必须同时停用,避免
两套调度器并行。
## 回滚
1. 冻结待回滚业务域的新写入,等待 FastAPI 在途请求和 jobs 当前轮次结束。
2. 停止 Python jobs,确认 Redis 分布式锁已释放;不要清空 Redis。
3. 将该域的 Nginx upstream 切回原 Java `manager-api`,保持 `/xiaozhi` 路径不变。
4. 用 Java 健康检查和代表性读请求确认 Token、缓存、上传文件与数据库数据可读。
5. 恢复 Java 写流量并记录回滚边界;不要让 FastAPI 继续写该域。
6. 若问题来自新 changeset,只能新增一个经评审的 Liquibase 前向修复;不得删除或改写
`DATABASECHANGELOG` 历史。
容器整服务回滚时,先把 `<JAVA_UPSTREAM>` 替换为 Nginx 容器可访问的真实地址,再只
重建代理;`--no-deps` 可避免回滚命令意外重启 FastAPI 或重复执行迁移:
```bash
MANAGER_API_UPSTREAM='<JAVA_UPSTREAM>:8002' \
docker compose up -d --no-deps --force-recreate manager-api-nginx
```
FastAPI 没有改变现有 schema,且 Redis 写入采用 Java 兼容格式,因此正常应用回滚不需要
数据反向迁移。若外部系统已接收不可撤销操作,按对应供应商的业务补偿流程处理,不能用
数据库回滚伪造外部成功或失败。
## 隔离验证
`scripts/isolated-env.sh` 只创建 `manager_java_test``manager_fastapi_test` 两个测试库和
端口 `16379` 上的独立 Redis;其中的测试密码仅用于本机隔离环境。标准流程是:
```bash
cd main/manager-api-fastapi
./scripts/isolated-env.sh start
./scripts/isolated-env.sh reset
./scripts/isolated-env.sh migrate
eval "$(./scripts/isolated-env.sh env)"
.venv/bin/pytest -m integration -q
./scripts/isolated-env.sh stop
```
实际执行结果、差分用例和不能使用真实凭证完成的联调项记录在
`docs/manager-api-fastapi-test-report.md`;逐接口状态记录在
`docs/manager-api-fastapi-compatibility.md`
+582
View File
@@ -0,0 +1,582 @@
# manager-api FastAPI 迁移测试报告
> 执行日期:2026-07-20Asia/Shanghai
>
> 工作目录:`/Users/mie/Desktop/Repo/xiaozhi-esp32-server`
>
> FastAPI 目标:`main/manager-api-fastapi`
> Java 基线:`main/manager-api`
本报告只记录实际执行并有输出或落盘证据的检查。结构路由闭合、154 条未认证/非法请求面
差分、154 条已认证安全业务/校验差分、领域测试和 49 条深度 Java/FastAPI 差分是不同强度
的证据,不互相替代。生产外部服务和真实硬件没有验证的部分,均不会写成通过。
## 1. 结果摘要
| 检查项 | 通过 | 失败/错误 | 跳过 | 结论 |
|---|---:|---:|---:|---|
| Java 基线测试 | 98 | 0 | 0 | 最终复跑 `BUILD SUCCESS`15.602 秒 |
| FastAPI 全量 pytest(最终回归) | 139 | 0 | 0 | 12.75 秒;含上传卷 readiness 与证据脱敏回归 |
| 隔离 MySQL/Redis 集成测试复跑 | 7 | 0 | 0 | 事务、锁、TTL、job 单实例与 watchdog 全绿 |
| Java→FastAPI 未认证/非法请求面差分 | 154 | 0 | 0 | 每条 Java 路由各 1 个无成功写入的缺认证或非法请求,逐项比较 status/body/Content-Type |
| Java→FastAPI 已认证安全业务/校验差分 | 154 | 0 | 0 | 每条 Java 路由各 1 个带正确认证的安全业务或校验请求;runner 有意不执行成功写入 |
| Java→FastAPI 深度差分契约 | 49 | 0 | 0 | 成功、主要错误与数据库副作用;直接覆盖 21/154 条 Java 路由 |
| 简单性能测试 | 480 请求 | 0 请求错误 | 不适用 | 4 场景 × 2 服务 × 60 次计量请求 |
| Java 路由结构清单 | 154/154 | 0 | 0 | FastAPI 注册闭合;结构证据,不等同逐接口行为证据 |
| 三端消费者调用点 | 188/188 | 0 | 0 | Web 134、Mobile 46、xiaozhi-server 8 个调用点均可解析 |
| FastAPI Ruff | 通过 | 0 | 不适用 | `app tests scripts` |
| FastAPI mypy | 70 个源文件 | 0 | 不适用 | strict 配置下通过 |
| FastAPI compileall | 通过 | 0 | 不适用 | `app tests scripts` |
| 锁文件与依赖同步 | 通过 | 0 | 不适用 | locked 环境共 55 packages |
| Python sdist/wheel | 2 个产物 | 0 | 不适用 | 冷环境完整依赖安装后可导入,路由数 163 |
| manager-web | i18n、5 unit、13 snapshot、build 全过 | 0 | 0 | build 有 4 条既有 size/precache warning |
| manager-mobile | type、lint、14 snapshot、mp-weixin build 全过 | 0 | 0 | 使用仓库声明的 pnpm 10.10.0 |
| xiaozhi-server | compileall 通过 | 0 | 不适用 | 除性能脚本外没有自动单测;8 个调用由 consumer 契约检查 |
| 真实付费/生产外部服务 | 0 | 不适用 | 不适用 | 无真实凭证,不声称联调通过 |
| 容器迁移、API、jobs 与 Nginx | 通过 | 0 | 0 | Apple Container 1.0.0 实际 build/runCompose 仅做静态验证,未伪装为 `docker compose up` |
最终可重复执行的全量测试、隔离差分、集成、构建和容器运行验证均为绿色。以下范围限制必须
与绿色测试分开陈述:
1. 全部 154 条 Java 路由均执行了两次差分:一次未认证/非法请求,一次已认证安全业务/校验
请求。第二个 runner 为保护隔离 fixture,有意不执行成功写入;49 个深度 checks 直接命中
21 条路由并覆盖代表性成功、错误和数据库副作用。因此两层全路由差分仍不等同于每条路由
的完整成功写入生命周期和全部错误路径差分。
2. 没有真实凭证、生产网络或硬件的外部集成未被计入通过。
## 2. 验证环境
| 组件 | 实际版本/配置 |
|---|---|
| 主机 | macOS 27.0arm64 |
| Java | Oracle JDK 21.0.11 LTS |
| Maven | 3.9.9,使用仓库 `.runtime/m2` |
| Python | 3.10.20`main/manager-api-fastapi/.venv` |
| MySQL | Community Server 8.0.46;隔离端口 `13316` |
| Redis | 8.8.0;隔离端口 `16379` |
| Node.js | v24.18.0 |
| npm | 11.16.0 |
| manager-mobile pnpm | Corepack 解析的 10.10.0 |
| OCI runtime | Apple Container 1.0.0;本机没有可用 Docker/Podman daemon |
| 容器架构 | linux/arm64 |
| 时区 | Asia/Shanghai |
隔离测试只重置 `manager_java_test``manager_fastapi_test` 两个测试 schema 和端口
`16379` 上的专用 Redis;没有连接、修改或清空开发 MySQL/Redis。Java 差分使用 Redis DB 1
FastAPI 使用 DB 2Java 单元验证显式使用 DB 3。
## 3. 实际执行命令
### 3.1 Java 基线
```bash
cd main/manager-api && \
JAVA_HOME=../../.runtime/jdk \
PATH="../../.runtime/jdk/bin:../../.runtime/maven/bin:$PATH" \
../../.runtime/maven/bin/mvn -o \
-Dmaven.repo.local=../../.runtime/m2 \
-Dspring.datasource.druid.url='jdbc:mysql://127.0.0.1:13316/manager_java_test?useUnicode=true&characterEncoding=UTF-8&serverTimezone=Asia/Shanghai&nullCatalogMeansCurrent=true&allowMultiQueries=true' \
-Dspring.datasource.druid.username=xiaozhi_test \
-Dspring.datasource.druid.password=isolated-test-only \
-Dspring.data.redis.host=127.0.0.1 \
-Dspring.data.redis.port=16379 \
-Dspring.data.redis.database=3 \
-Dspring.data.redis.password= \
-DskipTests=false test
```
最终复跑结果:98 tests0 failures0 errors0 skipped`BUILD SUCCESS`15.602 秒。Surefire
XML 位于 `main/manager-api/target/surefire-reports/`,各 suite 的 tests 合计为 98。
### 3.2 FastAPI 全量测试
```bash
cd main/manager-api-fastapi && \
eval "$(./scripts/isolated-env.sh env)" && \
APP_DATABASE_URL="$TEST_FASTAPI_DATABASE_URL" \
APP_REDIS_URL="$TEST_FASTAPI_REDIS_URL" \
APP_ENVIRONMENT=test \
.venv/bin/pytest -q
```
最终结果:139 passed、0 failed、0 skipped12.75 秒。
### 3.3 隔离集成、两层全路由差分、深度差分和性能测试
```bash
cd main/manager-api-fastapi && ./scripts/run-isolated-contract-tests.sh
```
最终执行结果为 exit 0。脚本实际完成以下阶段:
- 启动并重置隔离 MySQL `13316`、Redis `16379`
- 对 Java 与 FastAPI 两个 schema 分别执行原 Liquibase 101 个 changeSets
- 启动 Java 基线 `18082`、FastAPI `18083`、确定性外部 mock `18084`
- 生成并镜像固定用户、DB Token、Long ID、设备、agent、模型、纠错词和 OTA fixture
- 执行 7 条隔离集成测试;
- 对 154 条 Java 路由各执行一条未认证或非法的安全请求面差分,不产生成功写入;
- 对 154 条 Java 路由各执行一条带正确认证的安全业务或校验差分,仍不产生成功写入;
- 执行 49 条 Java/FastAPI 差分检查并写入 JSON
- 执行 480 次计量性能请求并写入 JSON;
- 对 FastAPI/mock 日志执行 warning、traceback、error 门禁;
- 退出时关闭 Java、FastAPI 和 mock,最终 `18082``18084` 没有监听进程。
最终阶段摘要(两行 154 分别对应未认证/非法和已认证安全业务/校验):
```text
7 passed
{"total": 154, "passed": 154, "failed": 0, "skipped": 0}
{"total": 154, "passed": 154, "failed": 0, "skipped": 0}
{"total": 49, "passed": 49, "failed": 0, "skipped": 0}
{"measurements": 8, "requests_measured": 480, "errors": 0}
Isolated integration, two 154-route surfaces, deep differential, and performance tests passed.
```
机器结果:
- `main/manager-api-fastapi/compatibility/route-surface-results.json`
- 生成时间:`2026-07-20T07:12:28.099864+00:00`
- 154 passed、0 failed、0 skipped
- `main/manager-api-fastapi/compatibility/authenticated-route-results.json`
- 生成时间:`2026-07-20T07:12:30.818073+00:00`
- 154 passed、0 failed、0 skipped
- `main/manager-api-fastapi/compatibility/contract-results.json`
- 生成时间:`2026-07-20T07:12:31.520810+00:00`
- 49 passed、0 failed、0 skipped
- `main/manager-api-fastapi/compatibility/performance-results.json`
- 生成时间:`2026-07-20T07:12:33.158826+00:00`
- 480 requests、0 errors
### 3.4 FastAPI 静态、依赖和构建验证
```bash
cd main/manager-api-fastapi
.venv/bin/ruff check app tests scripts
.venv/bin/mypy app
.venv/bin/python -m compileall -q app tests scripts
uv lock --check && uv sync --locked
uv build --no-cache
```
结果:
- Ruff 通过;
- mypy70 个源文件无问题;
- compileallexit 0
- lock check 与 locked syncexit 0,共 55 packages
- 无缓存构建 sdist 与 wheel 均成功。
构建产物另在临时冷虚拟环境验证。首次执行 `uv pip install --python
<临时环境>/bin/python --no-deps <wheel>` 后直接 import,因刻意没有安装 FastAPI 等运行依赖而
失败;这暴露的是冷 wheel 检查命令不完整,不是把失败隐藏为通过。随后执行带依赖的安装:
```bash
uv pip install --python <临时环境>/bin/python <wheel>
```
共安装 39 个锁定依赖。从仓库外 `/tmp` 导入成功,输出
`xiaozhi-manager-api 163`,证明不是依赖当前工作目录导入源码。
### 3.5 manager-web
```bash
cd main/manager-web && \
npm run check:i18n && \
npm run test:unit && \
npm run test:snapshot && \
npm run build
```
结果:
- i18n6 个 locale,每个 1527 keyskey 结构一致;
- unit5/5
- snapshot13/13
- Vue 生产构建 exit 0hash `71b64d002eb434ee`1069 ms);
- 输出有 4 条既有 bundle size/precache warning,没有将 warning 写成失败,也没有删除或放宽
测试来取得绿色结果;另有 `caniuse-lite` 数据过期 17 个月提示,未擅自更新依赖。
### 3.6 manager-mobile
```bash
cd main/manager-mobile && \
corepack pnpm type-check && \
corepack pnpm lint && \
corepack pnpm test:snapshot && \
corepack pnpm build:mp
```
结果:type-check exit 0、lint exit 0、snapshot 14/14、mp-weixin build exit 0。构建提示
`caniuse-lite` 数据过期 20 个月及 uni-app 有新版本;两项均不影响退出码,未擅自更新依赖。
一次较早的失败尝试直接调用普通 `pnpm`,环境解析到 v11,并因非 TTY 下依赖目录清理提示而
终止;没有把该次尝试写成通过。最终命令显式使用 Corepack,解析到仓库声明的 pnpm 10.10.0。
### 3.7 xiaozhi-server
```bash
cd main/xiaozhi-server && \
../manager-api-fastapi/.venv/bin/python --version && \
PYTHONPYCACHEPREFIX=/tmp/xiaozhi-server-pycache \
../manager-api-fastapi/.venv/bin/python -m compileall -q . && \
rg --files -g '*test*.py' -g '!performance_tester/**'
```
结果:Python 3.10.20compileall exit 0。排除性能测试目录后没有自动单测文件,因此没有虚构
pytest 通过数量;`xiaozhi-server` 的 8 个 manager-api 调用点由 consumer manifest 和 FastAPI
兼容测试验证为可解析。
首次误用不存在的 `../../.runtime/python/bin/python3.10`,结果为 exit 127;改用上面实际存在的
FastAPI Python 3.10.20 后通过。
## 4. 两层全路由差分与深度差分覆盖
### 4.1 154 条未认证/非法安全请求面差分
`tests/compatibility/route_surface_runner.py` 从 Java route manifest 逐条构造不发生成功写入的请求:
133 条 DB Token 路由省略 Token,14 条匿名路由发送安全非法输入,7 条内部路由省略
server-secret。Java 与 FastAPI 精确比较 HTTP status、解析后的 body 和 Content-Type;最终
154 passed、0 failed、0 skipped。它证明每条 Java 路由至少有一个请求路径兼容,不代表每条
路由的成功、全部错误与数据库副作用均已逐项对照。
### 4.2 154 条已认证安全业务/校验差分
`tests/compatibility/authenticated_route_runner.py` 使用与 Java 基线语义一致的 DB Token、
server-secret 或匿名认证方式,对同一份 154 路由清单逐条发送安全业务/校验请求。runner 对
动态管理员密码和 OTA 下载 UUID 先独立验证格式再做最小归一化,并同步会影响响应的固定审计
时间;最终 154 passed、0 failed、0 skipped。为不污染 fixture 或产生不可逆副作用,该 runner
有意选择资源不存在、单一约束失败、幂等空操作等不会成功写入的路径。因此这里的“已认证”
证明请求已经越过认证层并进入业务/校验逻辑,不表示每条写接口都完成了一次成功写入。
### 4.3 49 项深度结果分布
| 类别 | 通过 | 失败 |
|---|---:|---:|
| configuration | 1 | 0 |
| authentication-i18n | 7 | 0 |
| authentication | 1 | 0 |
| authorization | 1 | 0 |
| serialization | 2 | 0 |
| agent | 1 | 0 |
| device | 1 | 0 |
| model | 1 | 0 |
| correct-word | 1 | 0 |
| binary-download | 6 | 0 |
| validation | 5 | 0 |
| server-secret | 3 | 0 |
| ota | 1 | 0 |
| ota-validation | 2 | 0 |
| ota-signing | 2 | 0 |
| activation | 3 | 0 |
| external-mock | 2 | 0 |
| crud | 3 | 0 |
| database-side-effect | 4 | 0 |
| upload | 1 | 0 |
| upload-validation | 1 | 0 |
| **合计** | **49** | **0** |
直接覆盖内容包括:
- 默认语言、`zh-CN``zh-TW``en-US``de-DE``vi-VN``pt-BR` 七种
`Accept-Language` 情形;
- 未登录、DB Token 过期、普通用户访问管理员接口;
- Long ID 字符串、日期、Asia/Shanghai/UTC 兼容、null、别名和分页;
- 缺字段、数字格式、最小/最大边界和 Java UTF-16 长度语义;
- server-secret 缺失、错误和正确三条路径;
- OTA health、缺失/非法 `Device-Id`、激活、WS/MQTT credential
- HMAC-SHA256、URL-safe Base64、MQTT Base64 密码的独立密码学验证;
- 纠错词创建、更新、下载、删除及数据库副作用;
- OTA multipart 上传、扩展名错误、元数据副作用、三次下载与第四次 404;
- 二进制 MIME、`Content-Disposition``Content-Length` 和字节摘要;
- MQTT mock 的请求 body 和按日期生成的 Authorization。
### 4.4 有意的动态值处理
- OTA timestamp、WebSocket token 和生成 UUID 先做最小范围归一化,再分别验证格式与 HMAC;
不是把动态字段全部忽略。
- Hibernate Validator 使用无序 `ConstraintViolation Set`。模型 provider 空 body 用例要求
Java/FastAPI 均返回 HTTP 200、错误码 10034,消息必须属于 Java DTO 声明的五个精确约束,
不强行固定 Java 本身不稳定的首条消息。
- 报告落盘前递归脱敏 `private_key`、server secret、MQTT signature key、Token、password 和
Authorization;比较与密码学验证仍使用未脱敏的内存值。
### 4.5 覆盖边界
Java 清单共有 154 条路由,FastAPI 结构注册为 154/154;另有 3 条消费者兼容路由。全部 154
条均有一次未认证/非法差分和一次已认证安全业务/校验差分;49 个深度 checks 直接命中
21/154 条 Java 路由,并对代表性成功、主要错误和数据库副作用做更完整的生命周期验证。
其余路由虽有两层逐路由差分及相关领域 service/repository/protocol 测试,仍不能宣称其全部
成功写入和错误路径均已逐项深度差分。逐行状态见
`docs/manager-api-fastapi-compatibility.md`
## 5. 隔离数据库、Redis 和 job 测试
7 条集成测试均连接隔离 MySQL/Redis,而不是纯 mock
1. MySQL 事务异常回滚;
2. `SELECT FOR UPDATE` 在并发写入下串行化;
3. Redis key TTL 到期且不删除无关 key
4. Redis 分布式锁只允许一个并发 job 执行;
5. watchdog 在任务超过原始 lease 后续租,仍保持单实例;
6. 实际 knowledge job 函数在 Redis 锁下只执行一次;
7. Java hash 兼容写入的默认 TTL 为 86400 秒。
差分 CRUD 另检查 Java/FastAPI 各自隔离 schema 的行级副作用;没有实施双写,也没有操作开发
数据库。最终 FastAPI 与 external mock 日志没有 warning/error/traceback。两层全路由差分会
让 Java 基线按其既有全局异常处理记录 36 条预期 ERROR,严格门禁将其精确分为 8 类:缺 body
13 条、缺 `Device-Id` 5 条、对象/数组反序列化 8 条、缺 query 4 条、非 multipart 3 条、
`callerMac` 空值 1 条、空消息 1 条及 `null` 消息 1 条。未分类 Java ERROR 为 0;任一类别数量
变化或出现未分类日志都会使脚本失败。
## 6. 性能对比
参数:每个服务每场景先顺序 warmup 10 次,再以并发 6 计量 60 次;4 个场景、2 个服务共
480 次计量请求。结果来自最终 `performance-results.json`
| 场景 | 服务 | p50 ms | p95 ms | 吞吐 req/s | 错误 |
|---|---|---:|---:|---:|---:|
| representative-read | Java | 7.662 | 16.239 | 601.128 | 0 |
| representative-read | FastAPI | 6.749 | 12.552 | 751.538 | 0 |
| representative-crud-update | Java | 9.331 | 15.302 | 543.013 | 0 |
| representative-crud-update | FastAPI | 11.252 | 20.599 | 476.870 | 0 |
| runtime-configuration | Java | 8.746 | 16.127 | 568.863 | 0 |
| runtime-configuration | FastAPI | 6.488 | 13.722 | 765.126 | 0 |
| ota-check-and-signing | Java | 11.301 | 19.854 | 454.655 | 0 |
| ota-check-and-signing | FastAPI | 16.208 | 20.344 | 354.807 | 0 |
FastAPI/Java 比率:读取 p50 `0.881`、p95 `0.773`、吞吐 `1.250`CRUD p50 `1.206`
p95 `1.346`、吞吐 `0.878`;配置 p50 `0.742`、p95 `0.851`、吞吐 `1.345`OTA p50
`1.434`、p95 `1.025`、吞吐 `0.780`。因此不能笼统声称所有 FastAPI 接口都更快:本轮读取和
配置场景更快,CRUD p50/p95 分别较慢约 20.6%/34.6%OTA p50 较慢约 43.4%、p95 接近、
吞吐低约 22.0%。
这是同机、短时、固定 fixture 的简单对比,用于发现数量级回退,不是容量、长稳、生产网络或
多 worker 极限测试。
## 7. 三端兼容验证
- `manager-web`:134 个调用点、130 条唯一结构路由;i18n、unit、snapshot 与 production build
均通过,现有调用无需修改 URL。
- `manager-mobile`:46 个调用点、40 条唯一结构路由;type-check、lint、snapshot 与微信小程序
build 均通过。
- `xiaozhi-server`:8 个调用点、8 条唯一结构路由;compileall 通过,consumer manifest 确认
全部能解析到 FastAPI。该模块没有可执行的一般单测,不能把 compileall 写成运行时集成通过。
三端合计 188 个调用点、140 条唯一结构路由。此结论证明 path/method 解析闭合;全部 Java
路由另有未认证/非法和已认证安全业务/校验两层差分,参与 49 项深度差分或领域测试的调用拥有
更完整的成功、错误或副作用证据。
## 8. 迁移验证过程中发现并修正的问题
测试没有通过删除、跳过或放宽失败用例获得绿色。差分在实现过程中实际暴露并促成修复的兼容
问题包括:分页 `total` 类型、运行时配置字段查询、设备日期时区、认证与 OTA MIME、缺失
`Device-Id` envelope、provider 校验语义、非法分页消息、运行时配置 key 命名。修复后才生成
当前 49/49 深度报告。
154 路由请求面 runner 首次执行只有 149/154 通过,准确暴露 5 条“整个 JSON body 缺失”差异:
`POST /ota/``POST /user/login``POST /user/register`
`POST /user/retrieve-password``POST /user/smsVerification`。Java 对整个 body 缺失返回
HTTP 200、`code=500` 的通用 envelopeFastAPI 当时返回 `code=10034`。全局校验兼容层随后
只对定位恰为 body 根节点的 missing 错误映射 `code=500`,字段级 missing 仍保持 `10034`
新增回归用例并完整重跑后才得到 154/154。
已认证安全业务/校验 runner 首次执行为 111/154,通过实际差分先后修正了根 body 类型、必填
query/multipart 的 Java envelope、knowledge/device/agent/voice resource 的检查顺序、权限与资源
不存在语义,以及 DTO 单约束消息等差异;后续结果依次为 149/154、152/154,最终才达到
154/154。对 Hibernate Validator 无序约束的请求,runner 改用只触发一个约束的定向 payload,
没有跳过路由、忽略响应字段或放宽比较。该 runner 始终保持“已认证但不成功写入”的安全边界。
另有测试基础设施、运行时和构建问题被明确记录:
- Pydantic 2.13 与 FastAPI 0.116 的 TypeAdapter alias 路径产生
`UnsupportedFieldAttributeWarning`;依赖锁定到 Pydantic 2.11.7/core 2.33.2 后,实际 OTA
body alias 验证与最终日志门禁均无 warning。
- 第一版日志门禁把 Logback 初始化文本 `ERROR_FILE` 误判为运行时 ERROR,虽然 7、49、480
阶段均绿,脚本仍按门禁返回 1。正则收紧为带完整日期的 Java 应用 ERROR 后,重新从头执行
完整流程并获得 exit 0;没有直接忽略门禁失败。
- fixture 的 MySQL `VALUES()` upsert 产生 8.0 弃用 warning;改为 row alias 语法并重新执行后
无该 warning。
- 已认证差分报告最初虽未含真实凭证,但 `paramCode=server.secret` / `paramValue` 结构仍写入了
隔离 fixture 值;递归脱敏器补充键值对识别及回归测试后,再次完整执行差分。最终四份 JSON
`contract-server-secret`、测试 Token 和测试数据库密码的扫描均为 0 命中。
- Python 3.10 下 fixed-delay jobs 等待超时抛出 `asyncio.TimeoutError`;原捕获路径导致 worker
首轮后退出。worker 改为捕获该异常并增加 Python 3.10 回归测试;实际 jobs 容器随后观察到
knowledge job 跨 30 秒重复运行,snapshot redaction 多轮执行,SIGTERM 后干净退出。
- API 镜像构建初期遇到 uv/pip registry 传输失败;Dockerfile 固定 uv 版本并增加 timeout、retry
与缓存后完成构建。迁移镜像的 Maven Central 并发下载两次卡住,改为串行 resolver、超时与
retry 后成功。Nginx 初版配置在镜像 build 期校验失败,改用 template + `envsubst` 并在 build
内执行 `nginx -t` 后通过。
- Apple Container 自定义网络没有提供本次验证所需的容器名 DNS,host publish 和单文件挂载也
与 Docker 行为不同;验证改用容器 IP、显式 TCP bridge、named volume 和运行时 upstream 模板。它们
是测试 runtime 限制,不被记作应用通过或失败,也没有据此声称 Docker Compose 已实际启动。
## 9. 外部服务与真实联调状态
所有自动化外部调用只访问本地确定性 mock/fixture,不访问真实付费服务。
| 外部能力 | 自动化证据 | 真实联调状态 |
|---|---|---|
| RAGFlow dataset/document/chunk/retrieval/upload | 请求 JSON/query/header、30 秒 timeout、强 DTO、Long/null、错误映射与补偿路径测试 | 无真实 RAGFlow 凭证/实例,未联调 |
| 阿里云短信 | 配置、错误 envelope 与业务路径测试 | 无真实 AccessKey,不发送短信,未联调 |
| 火山语音克隆/音频 | multipart/JSON、状态及错误映射 mock | 无真实付费凭证,未联调 |
| 声纹 HTTP | Java multipart 形状与错误映射 mock | 无真实声纹服务,未联调 |
| OpenAI-compatible LLM | 请求格式、thinking policy、摘要/标题相关 mock | 无真实模型 key,不访问付费模型,未联调 |
| MQTT gateway HTTP | 差分验证 body、按日期 Authorization 和 401 retry 语义 | 无真实 MQTT broker/gateway,未联调 |
| MCP/管理 WebSocket | token、URL、path/scheme/form 兼容测试 | 无真实远端 MCP/WS,未联调 |
| OTA/WS/MQTT credential | 本地 HMAC/Base64/时间戳和下载行为实测 | 无 ESP32 真机和生产 broker,不属于硬件联调 |
因此,本报告只证明 mock 下已覆盖的请求格式、超时、错误映射、重试和本地密码学行为;不能把
任何一项写成供应商或生产环境端到端通过。
## 10. 已知行为/部署差异
- Java 的 Hibernate Validator 首条约束消息顺序不稳定;FastAPI 保持相同 envelope、错误码和
声明消息集合,而不是伪造固定顺序。
- Java 在 Spring 进程内运行定时任务;FastAPI 把 jobs 分离为独立进程,并用 Redis 锁和
watchdog 防止多 worker 重复执行。集成测试验证单实例和续租语义,但部署拓扑有意不同。
- FastAPI 增加 3 条消费者兼容路由和 live/ready health endpoints;它们没有 Java Controller
基线,属于明确的加法差异。
- 49 项已执行深度差分中没有观测到响应、所选 header 或数据库副作用差异;这句话只适用于
报告中的 49 项,不外推为全部 154 条路由均完成了成功写入和全部错误路径生命周期验证。
## 11. 实际容器与 Nginx 验证
### 11.1 Runtime、镜像与 Compose 口径
本机没有可用的 Docker/Podman daemon,实际 OCI build/run 使用 Apple Container 1.0.0 的
linux/arm64 VM,并显式使用隔离 app/log/install root
```bash
CLI=/Users/mie/.cache/xiaozhi-migration-tools/container-1.0.0-prefix/bin/container
ROOT=/Users/mie/.cache/xiaozhi-migration-tools/container-1.0.0-prefix
"$CLI" system start \
--app-root "$ROOT/runtime-data" \
--install-root "$ROOT" \
--log-root "$ROOT/runtime-logs" \
--disable-kernel-install
"$CLI" builder start
"$CLI" build --tag xiaozhi/manager-api-fastapi:0.1.0 \
--file main/manager-api-fastapi/Dockerfile .
"$CLI" build --tag xiaozhi/manager-api-migrate:fastapi-0.1.0 \
--file main/manager-api-fastapi/Dockerfile.migrations .
"$CLI" build --tag xiaozhi/manager-api-nginx:fastapi-0.1.0 \
--file main/manager-api-fastapi/Dockerfile.nginx .
```
三张镜像均实际构建并运行。迁移镜像 OCI index 为
`sha256:613faace4314b03392e65b64d9b4a9ba7a694cdd751c1a45009824d55f0647f7`,其 arm64
manifest 为 `sha256:6a10850841370d033a3b521fbb1100cb64b5cc6837fac35d00c4257343c0f2f9`
Nginx 镜像 OCI index 为
`sha256:2e6a188ad6d38b62fa4e77329a629ada00c4e773da9ded73c5f3289e40da477a`,其 arm64
manifest 为 `sha256:ff653bc2d11d4a3b1640747626055d6551fb33324500fb2e65b9333142da8526`
API 镜像在上传目录 readiness 最后一处源码变更后重新 build;最终 OCI index 为
`sha256:04ae1a98307b7369368b9665c6caf9f0911c8b2a967f5a91f23c6dde7c7baa16`,其 arm64
manifest 为 `sha256:c3267d307c9898975372539121f118549bfe51012c6c9bfdc3e84f99f3e56214`
config 为 `sha256:7c0b13757da041c0d118d14342e92c5310e4fb2140f029e385e97de9fe21d8cc`
manifest size 为 84,551,252 bytes,镜像配置创建时间为 `2026-07-20T07:04:45Z`
`docker-compose.yml` 已由 `tests/test_deployment_artifacts.py` 静态验证 migration dependency、
read-only root、tmpfs、upload volume、healthcheck、graceful timeout 与可切换 upstreamNginx
镜像 build 内也实际执行 `nginx -t`。由于本机没有 Docker Compose runtime,本报告明确只把
Compose 记为静态通过,不声称执行过 `docker compose up`
### 11.2 Liquibase migration
迁移镜像以 UID 10001 一次性运行,只读取原 Java resources 内的 Liquibase 历史。目标为隔离
schema `manager_container_test`;最终容器回归中再次运行并报告 101 个 changeSets 均
up-to-date。随后实查 `DATABASECHANGELOG` 为 101 条、业务及 Liquibase 表合计 30 张,
`DATABASECHANGELOGLOCK.LOCKED=0`,证明历史完整且锁已释放。没有连接、修改或清空开发数据库。
### 11.3 API、jobs、health、卷与优雅关闭
Apple Container VM 访问 host-only MySQL/Redis 时使用仓库内 TCP bridge,而不是暴露开发服务:
```bash
cd main/manager-api-fastapi
.venv/bin/python -m tests.compatibility.tcp_proxy \
--listen-port 13317 --target-host 127.0.0.1 --target-port 13316
.venv/bin/python -m tests.compatibility.tcp_proxy \
--listen-port 16380 --target-host 127.0.0.1 --target-port 16379
```
API 容器使用 `APP_WORKERS=2`、隔离 schema、Redis DB 4、read-only root、`/tmp` tmpfs 和
named upload volume 启动。实测结果:
- 容器内 UID 为 10001,最终层没有 `/bin/uv``/usr/bin/gcc`,应用路由数为 163
- 日志确认 2 个 Uvicorn worker(容器内 PID 3、4);`/xiaozhi/health/live` 为 HTTP 200
- `Accept-Language: en-US` 的未认证业务请求保持 HTTP 200、英文 `{code:401,...}`
`POST /xiaozhi/user/login` 整个 JSON body 缺失保持 Java 的 HTTP 200、`code=500`
- read-only root 生效。Apple Container 新建空 named volume 首次以其默认 root ownership 挂载,
新增 readiness 检查准确返回 HTTP 503、`database=true``redis=true``uploads=false`,没有让
无法上传的实例接流量;该失败没有伪装为通过。随后用一次性 root 容器仅对卷执行 `chown`
ownership 初始化,ready 变为 HTTP 200 且 `uploads=true`UID 10001 的 API 成功写入,重启后
文件 SHA256 `1ad4cb4f879aa1ddf43a14e1a84cc5dbf8f65e91295165e126b8b08be3cd9a50` 保持不变;
- 发送 SIGTERM 后 worker 完成 lifespan shutdown 并以 exit 0 退出,无 traceback/error。
同一 API 镜像另以 `python -m app.jobs.worker`、read-only root 和 `/tmp` tmpfs 启动。实际等待
超过 31 秒后,knowledge fixed-delay job 执行两次且相隔 30 秒,snapshot redaction 多次执行;
这验证 Python 3.10 timeout 修复与真实调度循环。SIGTERM 后 jobs 也干净退出。API 多 worker
本身不加载 jobs,独立 worker 再由 Redis lock/watchdog 保证单实例。
上述 ownership 初始化是 Apple Container 空 named volume 的实测处理;本机没有 Docker
Compose runtime,因此 Docker Compose 的 named-volume copy-up 行为没有实际验证,不能用
Apple Container 的结果代替。
### 11.4 Nginx 切流与 Java 回滚
Nginx 镜像以 read-only root 和 `/var/cache/nginx``/var/run``/tmp` 三个 tmpfs 运行;其
entrypoint 将 `MANAGER_API_UPSTREAM` 注入模板后 `exec nginx`。Apple Container 自定义网络在
本次环境没有容器名 DNS,因此实测使用 runtime 分配的 API/Java 容器 IP,语义与生产 hostname
upstream 相同。FastAPI upstream 下实际验证:
- `/xiaozhi/health/ready` 为 HTTP 200
- `/xiaozhi` 精确返回 308 到 `/xiaozhi/`
- `Accept-Language: en-US` 的未认证 envelope 由 Nginx 转发后与直连 FastAPI 一致,整个 JSON
body 缺失也保持 HTTP 200、`code=500`
- Nginx、API 均在 SIGTERM 下以 exit 0 干净退出。
随后仅替换 `MANAGER_API_UPSTREAM` 指向保留的 Java 容器并重建 Nginx 运行实例;`/xiaozhi/ota/`
回滚探针的 response body 与直连 Java 按字节完全一致。此步骤证明回滚不需要删除 Java 服务、
改数据库或双写,只需切换 upstream。Nginx 基础镜像未声明非 root USER,因此这里不虚构其
non-root 属性;实际硬化证据是 read-only root、最小 tmpfs 和无持久写入。应用与迁移镜像则
均以 UID 10001 运行。
## 12. 证据文件
- 逐接口矩阵:`docs/manager-api-fastapi-compatibility.md`
- 迁移、切流与回滚说明:`docs/manager-api-fastapi-migration.md`
- Java 路由清单:`main/manager-api-fastapi/compatibility/java-routes.json`
- 三端调用清单:`main/manager-api-fastapi/compatibility/consumer-routes.json`
- 154 路由未认证/非法请求面机器报告:
`main/manager-api-fastapi/compatibility/route-surface-results.json`
- 154 路由已认证安全业务/校验机器报告:
`main/manager-api-fastapi/compatibility/authenticated-route-results.json`
- 深度差分机器报告:`main/manager-api-fastapi/compatibility/contract-results.json`
- 性能机器报告:`main/manager-api-fastapi/compatibility/performance-results.json`
- 一键隔离脚本:`main/manager-api-fastapi/scripts/run-isolated-contract-tests.sh`
- 未认证/非法请求面 runner
`main/manager-api-fastapi/tests/compatibility/route_surface_runner.py`
- 已认证安全业务/校验 runner:
`main/manager-api-fastapi/tests/compatibility/authenticated_route_runner.py`
- 深度差分 runner`main/manager-api-fastapi/tests/compatibility/differential_runner.py`
- 外部 mock`main/manager-api-fastapi/tests/compatibility/external_mock.py`
- 集成测试:`main/manager-api-fastapi/tests/integration/test_isolated_runtime.py`
- 容器静态断言:`main/manager-api-fastapi/tests/test_deployment_artifacts.py`
- 容器网络 bridge`main/manager-api-fastapi/tests/compatibility/tcp_proxy.py`
- API/migration/Nginx 构建定义:`main/manager-api-fastapi/Dockerfile`
`main/manager-api-fastapi/Dockerfile.migrations``main/manager-api-fastapi/Dockerfile.nginx`
- Nginx runtime 配置:`main/manager-api-fastapi/deploy/nginx.conf`
`main/manager-api-fastapi/deploy/nginx-entrypoint.sh`
- Java Surefire`main/manager-api/target/surefire-reports/`
## 13. 当前结论
Java 98、FastAPI 全量 139、隔离集成 7、未认证/非法请求面 154/154、已认证安全业务/校验
154/154、深度差分 49/49、性能 480/0,以及 Web/Mobile 构建、xiaozhi-server compileall 和
实际容器/Nginx 验证均按上述命令完成;各测试集合均为 0 failed、0 errors、0 skipped。原 Java
服务和 Liquibase 历史均未删除。
本地可安全执行的兼容、集成、构建、消费者和容器验证已经通过。每条 Java 路由虽已有两次
全覆盖差分,但已认证 runner 有意不执行成功写入,所以不能将其表述为 154 条全部成功、错误
和副作用生命周期均已深度验证;真实 RAGFlow、短信、语音克隆、声纹、模型、MQTT/MCP/WS
及 ESP32 硬件因没有真实凭证或设备而未联调,也没有在本报告中描述为已通过。
+20
View File
@@ -0,0 +1,20 @@
APP_ENVIRONMENT=development
APP_HOST=0.0.0.0
APP_PORT=8002
APP_CONTEXT_PATH=/xiaozhi
APP_TIMEZONE=Asia/Shanghai
APP_DATABASE_URL=mysql+asyncmy://xiaozhi:replace-me@mysql:3306/xiaozhi_esp32_server?charset=utf8mb4
APP_REDIS_URL=redis://redis:6379/0
# Local default; the container image overrides this with /data/uploads.
APP_UPLOAD_DIR=./uploadfile
# Docker Compose source: use a named volume by default, or set an existing
# Java uploadfile host path while the implementations coexist.
MANAGER_API_UPLOAD_SOURCE=manager-api-uploads
APP_JAVA_RESOURCES_DIR=/opt/xiaozhi/java-resources
APP_EXTERNAL_REQUEST_TIMEOUT_SECONDS=10
APP_TRUSTED_PROXY_COUNT=1
APP_LOG_LEVEL=INFO
APP_GRACEFUL_SHUTDOWN_SECONDS=30
# Test-only escape hatches. Leave both unset in deployments.
# APP_SERVER_SECRET_OVERRIDE=
# APP_ALLOW_START_WITHOUT_DEPENDENCIES=false
+16
View File
@@ -0,0 +1,16 @@
.env
.mypy_cache/
.pytest_cache/
.ruff_cache/
.test-runtime/
.venv/
.coverage
coverage.xml
htmlcov/
__pycache__/
*.py[cod]
data/uploads/
target/
dist/
!compatibility/*.json
!tests/fixtures/*.json
+1
View File
@@ -0,0 +1 @@
3.10
+54
View File
@@ -0,0 +1,54 @@
FROM python:3.10.20-bookworm AS build
ENV UV_COMPILE_BYTECODE=1 \
UV_LINK_MODE=copy \
UV_HTTP_TIMEOUT=120 \
UV_HTTP_RETRIES=10 \
PATH=/app/.venv/bin:$PATH
WORKDIR /app
COPY --from=ghcr.io/astral-sh/uv:0.11.28 /uv /uvx /bin/
COPY main/manager-api-fastapi/pyproject.toml main/manager-api-fastapi/uv.lock main/manager-api-fastapi/README.md ./
COPY main/manager-api-fastapi/app ./app
RUN --mount=type=cache,target=/root/.cache/uv \
uv sync --frozen --no-dev --no-editable
FROM python:3.10.20-slim-bookworm
ENV PYTHONDONTWRITEBYTECODE=1 \
PYTHONUNBUFFERED=1 \
PATH=/app/.venv/bin:$PATH \
APP_JAVA_RESOURCES_DIR=/opt/xiaozhi/java-resources \
APP_UPLOAD_DIR=/data/uploads \
APP_HOST=0.0.0.0 \
APP_PORT=8002 \
APP_TIMEZONE=Asia/Shanghai
WORKDIR /app
RUN groupadd --gid 10001 xiaozhi \
&& useradd --uid 10001 --gid xiaozhi --create-home --shell /usr/sbin/nologin xiaozhi
COPY --from=build --chown=10001:10001 /app/.venv ./.venv
COPY --from=build --chown=10001:10001 /app/app ./app
COPY main/manager-api-fastapi/scripts/container-entrypoint.sh /usr/local/bin/manager-api-entrypoint
COPY main/manager-api/src/main/resources/i18n /opt/xiaozhi/java-resources/i18n
COPY main/manager-api/src/main/resources/db /opt/xiaozhi/java-resources/db
RUN mkdir -p /data/uploads \
&& ln -s /data/uploads /app/uploadfile \
&& chown -R xiaozhi:xiaozhi /data/uploads /opt/xiaozhi \
&& chmod 0555 /usr/local/bin/manager-api-entrypoint
USER 10001:10001
EXPOSE 8002
VOLUME ["/data/uploads"]
STOPSIGNAL SIGTERM
HEALTHCHECK --interval=15s --timeout=3s --start-period=20s --retries=4 \
CMD ["python", "-c", "import urllib.request; urllib.request.urlopen('http://127.0.0.1:8002/xiaozhi/health/live', timeout=2).read()"]
ENTRYPOINT ["/usr/local/bin/manager-api-entrypoint"]
@@ -0,0 +1,35 @@
FROM maven:3.9.9-eclipse-temurin-21 AS build
WORKDIR /migration
COPY main/manager-api-fastapi/migration-pom.xml ./pom.xml
COPY main/manager-api-fastapi/migration-src ./migration-src
COPY main/manager-api/src/main/resources ./java-resources
# Keep the Maven repository outside the committed layer so an interrupted
# registry transfer can resume on the next build. Resolver downloads are
# deliberately serial: Apple Container's BuildKit NAT has proved unreliable
# when several Maven Central responses are multiplexed over one connection.
RUN --mount=type=cache,target=/root/.m2/repository \
mvn -B \
-Dmaven.repo.local=/root/.m2/repository \
-Djava.resources.dir=/migration/java-resources \
-Daether.connector.basic.threads=1 \
-Daether.connector.connectTimeout=15000 \
-Daether.connector.requestTimeout=60000 \
-Daether.connector.http.retryHandler.count=5 \
-Daether.connector.http.retryHandler.interval=1000 \
-Daether.connector.http.retryHandler.intervalMax=5000 \
package
FROM eclipse-temurin:21-jre
WORKDIR /migration
COPY --from=build /migration/target/manager-api-liquibase-runner-1.0.0-all.jar ./runner.jar
COPY main/manager-api-fastapi/scripts/run-migrations.sh /usr/local/bin/run-manager-api-migrations
RUN groupadd --gid 10001 xiaozhi \
&& useradd --uid 10001 --gid xiaozhi --create-home --shell /usr/sbin/nologin xiaozhi \
&& chown -R xiaozhi:xiaozhi /migration \
&& chmod 0555 /usr/local/bin/run-manager-api-migrations
ENV MIGRATION_RUNNER_JAR=/migration/runner.jar \
TZ=Asia/Shanghai
USER 10001:10001
STOPSIGNAL SIGTERM
ENTRYPOINT ["/usr/local/bin/run-manager-api-migrations"]
+14
View File
@@ -0,0 +1,14 @@
FROM nginx:1.28.0-alpine
COPY main/manager-api-fastapi/deploy/nginx.conf /etc/nginx/nginx.conf.template
COPY main/manager-api-fastapi/deploy/nginx-entrypoint.sh /usr/local/bin/manager-api-nginx-entrypoint
RUN MANAGER_API_UPSTREAM=127.0.0.1:8002 \
envsubst '${MANAGER_API_UPSTREAM}' \
< /etc/nginx/nginx.conf.template \
> /tmp/nginx-build-check.conf \
&& nginx -t -c /tmp/nginx-build-check.conf \
&& rm /tmp/nginx-build-check.conf \
&& chmod 0555 /usr/local/bin/manager-api-nginx-entrypoint
ENV MANAGER_API_UPSTREAM=manager-api-fastapi:8002
ENTRYPOINT ["/usr/local/bin/manager-api-nginx-entrypoint"]
+38
View File
@@ -0,0 +1,38 @@
# manager-api-fastapi
`manager-api-fastapi` is the Python/FastAPI implementation of the existing Spring Boot
`main/manager-api`. The Java service remains in the repository as the contract baseline,
Liquibase migration owner, and rollback implementation.
## Local development
The service requires Python 3.10, MySQL 8, and Redis 5 or newer. Never point tests at a
development database: the integration harness creates a dedicated database and Redis
namespace/instance.
```bash
cd main/manager-api-fastapi
cp .env.example .env
uv sync --locked
uv run python -m app
```
The compatible base URL is `http://127.0.0.1:8002/xiaozhi`. OpenAPI is exposed at
`/xiaozhi/v3/api-docs` and the Swagger UI at `/xiaozhi/doc.html`.
Production must set `APP_DATABASE_URL`, `APP_REDIS_URL`, `APP_UPLOAD_DIR`, and
`APP_JAVA_RESOURCES_DIR`. The last path must contain the original Java i18n resources and
Liquibase changelog. `APP_SERVER_SECRET_OVERRIDE` is reserved for isolated tests; leaving it
set in a deployment bypasses the database-backed `server.secret` lookup and is unsupported.
## Commands
```bash
uv run pytest
uv run ruff check app tests scripts
uv run mypy app
uv run python scripts/extract_java_routes.py --output compatibility/java-routes.json
```
Migration, container, differential-contract, and cutover instructions are maintained in the
repository-level migration documents under `docs/manager-api-fastapi-*.md`.
+1
View File
@@ -0,0 +1 @@
"""Xiaozhi manager API FastAPI implementation."""
+7
View File
@@ -0,0 +1,7 @@
import uvicorn
from app.core.config import get_settings
if __name__ == "__main__":
settings = get_settings()
uvicorn.run("app.main:app", host=settings.host, port=settings.port, log_level=settings.log_level.lower())
@@ -0,0 +1 @@
"""Shared compatibility infrastructure."""
@@ -0,0 +1,72 @@
from __future__ import annotations
from functools import lru_cache
from pathlib import Path
from typing import Literal
from pydantic import Field, field_validator
from pydantic_settings import BaseSettings, SettingsConfigDict
def _default_java_resources() -> Path:
return Path(__file__).resolve().parents[3] / "manager-api" / "src" / "main" / "resources"
class Settings(BaseSettings):
model_config = SettingsConfigDict(
env_file=".env",
env_prefix="APP_",
case_sensitive=False,
extra="ignore",
)
environment: Literal["development", "test", "production"] = "development"
host: str = "0.0.0.0" # noqa: S104 - container bind is intentional
port: int = 8002
context_path: str = "/xiaozhi"
timezone: str = "Asia/Shanghai"
database_url: str = "mysql+asyncmy://root:change-me@127.0.0.1:3306/xiaozhi_esp32_server?charset=utf8mb4"
redis_url: str = "redis://127.0.0.1:6379/0"
upload_dir: Path = Path("uploadfile")
java_resources_dir: Path = Field(default_factory=_default_java_resources)
external_request_timeout_seconds: float = 10.0
database_pool_size: int = 20
database_max_overflow: int = 20
trusted_proxy_count: int = 1
log_level: str = "INFO"
server_secret_override: str | None = None
allow_start_without_dependencies: bool = False
job_lock_ttl_seconds: int = 120
graceful_shutdown_seconds: float = 30.0
@field_validator("context_path")
@classmethod
def normalize_context_path(cls, value: str) -> str:
normalized = "/" + value.strip("/")
return "" if normalized == "/" else normalized
@field_validator("database_url")
@classmethod
def require_async_driver(cls, value: str) -> str:
if value.startswith("mysql://"):
return value.replace("mysql://", "mysql+asyncmy://", 1)
if value.startswith("sqlite:///"):
return value.replace("sqlite:///", "sqlite+aiosqlite:///", 1)
return value
@property
def i18n_dir(self) -> Path:
return self.java_resources_dir / "i18n"
@property
def changelog_path(self) -> Path:
return self.java_resources_dir / "db" / "changelog" / "db.changelog-master.yaml"
@lru_cache(maxsize=1)
def get_settings() -> Settings:
return Settings()
def clear_settings_cache() -> None:
get_settings.cache_clear()
@@ -0,0 +1,63 @@
from __future__ import annotations
import hashlib
import secrets
import uuid
import bcrypt
from gmssl import func, sm2 # type: ignore[import-untyped]
def generate_database_token(value: str | None = None) -> str:
source = value if value is not None else str(uuid.uuid4())
return hashlib.md5(source.encode("utf-8"), usedforsecurity=False).hexdigest()
def bcrypt_hash(password: str, rounds: int = 10) -> str:
encoded = bcrypt.hashpw(password.encode("utf-8"), bcrypt.gensalt(rounds=rounds))
# The bundled Java BCryptPasswordEncoder only accepts $2a$ hashes.
return encoded.decode("ascii").replace("$2b$", "$2a$", 1)
def bcrypt_matches(password: str, encoded: str | None) -> bool:
if not encoded or not encoded.startswith(("$2a$", "$2$")):
return False
normalized = encoded.replace("$2$", "$2a$", 1)
try:
return bcrypt.checkpw(password.encode("utf-8"), normalized.encode("ascii"))
except (ValueError, UnicodeEncodeError):
return False
def sm2_generate_keypair() -> tuple[str, str]:
private_key = func.random_hex(64)
helper = sm2.CryptSM2(private_key=private_key, public_key="", mode=1)
public_point = str(helper._kg(int(private_key, 16), sm2.default_ecc_table["g"])) # noqa: SLF001
public_key = "04" + public_point
return public_key, private_key
def sm2_encrypt_c1c3c2(public_key: str, plaintext: str) -> str:
helper = sm2.CryptSM2(private_key="", public_key=public_key, mode=1)
encrypted = helper.encrypt(plaintext.encode("utf-8"))
if encrypted is None:
raise ValueError("SM2 KDF returned an all-zero key")
# BouncyCastle's SM2Engine emits the uncompressed-point marker.
return "04" + bytes(encrypted).hex()
def sm2_decrypt_c1c3c2(private_key: str, ciphertext: str) -> str:
normalized = ciphertext.strip().lower()
if normalized.startswith("04"):
normalized = normalized[2:]
if len(normalized) < 128 + 64 or len(normalized) % 2:
raise ValueError("invalid SM2 C1C3C2 ciphertext")
helper = sm2.CryptSM2(private_key=private_key, public_key="", mode=1)
decrypted = helper.decrypt(bytes.fromhex(normalized))
if decrypted is None:
raise ValueError("SM2 decryption failed")
return bytes(decrypted).decode("utf-8")
def random_hex(length: int) -> str:
return secrets.token_hex((length + 1) // 2)[:length]
@@ -0,0 +1,96 @@
from __future__ import annotations
from collections.abc import AsyncIterator, Mapping, Sequence
from contextlib import asynccontextmanager
from typing import Any
from sqlalchemy import Result, TextClause, text
from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker, create_async_engine
from app.core.config import Settings, get_settings
_engine: AsyncEngine | None = None
_session_factory: async_sessionmaker[AsyncSession] | None = None
def configure_database(settings: Settings | None = None) -> tuple[AsyncEngine, async_sessionmaker[AsyncSession]]:
global _engine, _session_factory
selected = settings or get_settings()
engine_options: dict[str, Any] = {"pool_pre_ping": True}
if not selected.database_url.startswith("sqlite"):
engine_options.update(pool_size=selected.database_pool_size, max_overflow=selected.database_max_overflow)
_engine = create_async_engine(selected.database_url, **engine_options)
_session_factory = async_sessionmaker(_engine, expire_on_commit=False, autoflush=False)
return _engine, _session_factory
def get_engine() -> AsyncEngine:
if _engine is None:
return configure_database()[0]
return _engine
def get_session_factory() -> async_sessionmaker[AsyncSession]:
if _session_factory is None:
return configure_database()[1]
return _session_factory
async def get_db() -> AsyncIterator[AsyncSession]:
async with get_session_factory()() as session:
yield session
@asynccontextmanager
async def transaction() -> AsyncIterator[AsyncSession]:
async with get_session_factory()() as session, session.begin():
yield session
async def dispose_database() -> None:
global _engine, _session_factory
if _engine is not None:
await _engine.dispose()
_engine = None
_session_factory = None
class Repository:
def __init__(self, session: AsyncSession):
self.session = session
@staticmethod
def statement(sql: str | TextClause) -> TextClause:
return text(sql) if isinstance(sql, str) else sql
async def fetch_one(self, sql: str | TextClause, params: Mapping[str, Any] | None = None) -> dict[str, Any] | None:
result = await self.session.execute(self.statement(sql), dict(params or {}))
row = result.mappings().first()
return dict(row) if row is not None else None
async def fetch_all(self, sql: str | TextClause, params: Mapping[str, Any] | None = None) -> list[dict[str, Any]]:
result = await self.session.execute(self.statement(sql), dict(params or {}))
return [dict(row) for row in result.mappings().all()]
async def scalar(self, sql: str | TextClause, params: Mapping[str, Any] | None = None) -> Any:
result = await self.session.execute(self.statement(sql), dict(params or {}))
return result.scalar_one_or_none()
async def execute(self, sql: str | TextClause, params: Mapping[str, Any] | None = None) -> int:
result: Result[Any] = await self.session.execute(self.statement(sql), dict(params or {}))
return int(getattr(result, "rowcount", 0) or 0)
async def execute_many(self, sql: str | TextClause, params: Sequence[Mapping[str, Any]]) -> int:
if not params:
return 0
result: Result[Any] = await self.session.execute(self.statement(sql), [dict(item) for item in params])
return int(getattr(result, "rowcount", 0) or 0)
async def database_ping() -> bool:
try:
async with get_session_factory()() as session:
await session.execute(text("SELECT 1"))
return True
except Exception:
return False
@@ -0,0 +1,44 @@
from __future__ import annotations
from dataclasses import dataclass
class ErrorCode:
INTERNAL_SERVER_ERROR = 500
UNAUTHORIZED = 401
FORBIDDEN = 403
DB_RECORD_EXISTS = 10002
PARAMS_GET_ERROR = 10003
ACCOUNT_PASSWORD_ERROR = 10004
ACCOUNT_DISABLE = 10005
CAPTCHA_ERROR = 10007
PASSWORD_ERROR = 10009
UPLOAD_FILE_EMPTY = 10019
TOKEN_INVALID = 10021
ACCOUNT_LOCK = 10022
INVALID_SYMBOL = 10029
PASSWORD_LENGTH_ERROR = 10030
PASSWORD_WEAK_ERROR = 10031
DEL_MYSELF_ERROR = 10032
DEVICE_CAPTCHA_ERROR = 10033
PARAM_VALUE_NULL = 10034
PARAM_TYPE_NULL = 10035
PARAM_TYPE_INVALID = 10036
PARAM_NUMBER_INVALID = 10037
PARAM_BOOLEAN_INVALID = 10038
PARAM_ARRAY_INVALID = 10039
PARAM_JSON_INVALID = 10040
RESOURCE_NOT_FOUND = 10051
ADD_DATA_FAILED = 10065
UPDATE_DATA_FAILED = 10066
MODEL_TYPE_PROVIDE_CODE_NOT_NULL = 10131
@dataclass(slots=True)
class AppError(Exception):
code: int
message: str | None = None
params: tuple[object, ...] = ()
def __str__(self) -> str:
return self.message or str(self.code)
+90
View File
@@ -0,0 +1,90 @@
from __future__ import annotations
import re
from functools import lru_cache
from pathlib import Path
from app.core.config import get_settings
LANGUAGE_FILES: dict[str, str] = {
"zh-CN": "messages_zh_CN.properties",
"zh-TW": "messages_zh_TW.properties",
"en-US": "messages_en_US.properties",
"de-DE": "messages_de_DE.properties",
"vi-VN": "messages_vi_VN.properties",
"pt-BR": "messages_pt_BR.properties",
}
_UNICODE_ESCAPE = re.compile(r"\\u([0-9a-fA-F]{4})")
def resolve_language(accept_language: str | None) -> str:
if not accept_language:
return "zh-CN"
primary = accept_language.split(",", 1)[0].split(";", 1)[0].strip().replace("_", "-")
exact = {key.lower(): key for key in LANGUAGE_FILES}
if primary.lower() in exact:
return exact[primary.lower()]
prefix = primary.lower().split("-", 1)[0]
return {
"zh": "zh-CN",
"en": "en-US",
"de": "de-DE",
"vi": "vi-VN",
"pt": "pt-BR",
}.get(prefix, "zh-CN")
def _unescape(value: str) -> str:
decoded = _UNICODE_ESCAPE.sub(lambda match: chr(int(match.group(1), 16)), value)
return (
decoded.replace("\\t", "\t")
.replace("\\n", "\n")
.replace("\\r", "\r")
.replace("\\f", "\f")
.replace("\\=", "=")
.replace("\\:", ":")
.replace("\\ ", " ")
.replace("\\\\", "\\")
)
def _load_properties(path: Path) -> dict[str, str]:
messages: dict[str, str] = {}
if not path.exists():
return messages
continuation = ""
for raw_line in path.read_text(encoding="utf-8").splitlines():
line = continuation + raw_line
if line.endswith("\\") and not line.endswith("\\\\"):
continuation = line[:-1]
continue
continuation = ""
stripped = line.strip()
if not stripped or stripped.startswith(("#", "!")):
continue
delimiter = "=" if "=" in line else ":"
if delimiter not in line:
continue
key, value = line.split(delimiter, 1)
messages[key.strip()] = _unescape(value.strip())
return messages
@lru_cache(maxsize=16)
def messages_for(language: str, i18n_dir: str | None = None) -> dict[str, str]:
directory = Path(i18n_dir) if i18n_dir else get_settings().i18n_dir
default_messages = _load_properties(directory / "messages.properties")
default_messages.update(_load_properties(directory / LANGUAGE_FILES.get(language, LANGUAGE_FILES["zh-CN"])))
return default_messages
def message_for(code: int, accept_language: str | None, *params: object) -> str:
language = resolve_language(accept_language)
template = messages_for(language).get(str(code), str(code))
for index, param in enumerate(params):
template = template.replace("{" + str(index) + "}", str(param))
return template
def clear_i18n_cache() -> None:
messages_for.cache_clear()
+56
View File
@@ -0,0 +1,56 @@
from __future__ import annotations
import os
import socket
import threading
import time
class SnowflakeIdGenerator:
"""MyBatis-Plus compatible 41/5/5/12-bit Snowflake identifier generator."""
EPOCH = 1288834974657
SEQUENCE_BITS = 12
WORKER_BITS = 5
DATACENTER_BITS = 5
MAX_SEQUENCE = (1 << SEQUENCE_BITS) - 1
WORKER_SHIFT = SEQUENCE_BITS
DATACENTER_SHIFT = SEQUENCE_BITS + WORKER_BITS
TIMESTAMP_SHIFT = SEQUENCE_BITS + WORKER_BITS + DATACENTER_BITS
def __init__(self, worker_id: int | None = None, datacenter_id: int | None = None):
host_hash = sum(socket.gethostname().encode("utf-8"))
self.worker_id = worker_id if worker_id is not None else (host_hash ^ os.getpid()) & 31
self.datacenter_id = datacenter_id if datacenter_id is not None else host_hash & 31
if not 0 <= self.worker_id <= 31 or not 0 <= self.datacenter_id <= 31:
raise ValueError("worker_id and datacenter_id must be in [0, 31]")
self._sequence = 0
self._last_timestamp = -1
self._lock = threading.Lock()
@staticmethod
def _milliseconds() -> int:
return time.time_ns() // 1_000_000
def next_id(self) -> int:
with self._lock:
timestamp = self._milliseconds()
if timestamp < self._last_timestamp:
raise RuntimeError("clock moved backwards; refusing to generate a duplicate Snowflake ID")
if timestamp == self._last_timestamp:
self._sequence = (self._sequence + 1) & self.MAX_SEQUENCE
if self._sequence == 0:
while timestamp <= self._last_timestamp:
timestamp = self._milliseconds()
else:
self._sequence = 0
self._last_timestamp = timestamp
return (
((timestamp - self.EPOCH) << self.TIMESTAMP_SHIFT)
| (self.datacenter_id << self.DATACENTER_SHIFT)
| (self.worker_id << self.WORKER_SHIFT)
| self._sequence
)
snowflake = SnowflakeIdGenerator()
+226
View File
@@ -0,0 +1,226 @@
from __future__ import annotations
import asyncio
import json
import logging
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager, suppress
from datetime import datetime
from typing import Any, cast
from zoneinfo import ZoneInfo
from redis.asyncio import Redis
from app.core.config import get_settings
_client: Redis | None = None
logger = logging.getLogger(__name__)
class JavaRedisCodec:
"""Wire-compatible subset of Spring Data's ``RedisSerializer.json()``.
Spring enables Jackson default typing for non-final values. Consequently a
plain JSON map/list cannot be read by the retained Java rollback service. A
map carries ``@class`` and a collection uses Jackson's wrapper-array form.
``java_type`` and ``item_java_type`` cover the few caches whose Java readers
cast values to concrete DTO/entity classes.
"""
@staticmethod
def encode(
value: Any,
*,
java_type: str | None = None,
item_java_type: str | None = None,
field_java_types: dict[str, str] | None = None,
) -> bytes:
wire = JavaRedisCodec._encode_value(
value,
java_type=java_type,
item_java_type=item_java_type,
field_java_types=field_java_types,
nested=False,
)
return json.dumps(wire, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
@staticmethod
def decode(value: bytes | str | None) -> Any:
if value is None:
return None
raw = value.decode("utf-8") if isinstance(value, bytes) else value
try:
return JavaRedisCodec._decode_value(json.loads(raw))
except json.JSONDecodeError:
return raw
@staticmethod
def _encode_value(
value: Any,
*,
java_type: str | None = None,
item_java_type: str | None = None,
field_java_types: dict[str, str] | None = None,
nested: bool = True,
) -> Any:
if value is None or isinstance(value, str | bool | float):
return value
if isinstance(value, int):
# Jackson's default typing only adds the Long wrapper when the
# runtime value sits behind an Object-typed container slot. A
# top-level Long, or a field with a declared Long type, is emitted
# as an ordinary JSON number.
if not nested or java_type == "java.lang.Long" or -(2**31) <= value < 2**31:
return value
return ["java.lang.Long", value]
if isinstance(value, datetime):
timezone = ZoneInfo(get_settings().timezone)
localized = value.replace(tzinfo=timezone) if value.tzinfo is None else value.astimezone(timezone)
return ["java.util.Date", int(localized.timestamp() * 1000)]
if isinstance(value, dict):
selected_type = java_type or "java.util.HashMap"
result: dict[str, Any] = {"@class": selected_type}
pojo = selected_type not in {
"java.util.HashMap",
"java.util.LinkedHashMap",
"java.util.TreeMap",
"cn.hutool.json.JSONObject",
}
for raw_key, item in value.items():
if raw_key == "@class":
continue
key = _snake_to_camel(str(raw_key)) if pojo else str(raw_key)
child_type = (field_java_types or {}).get(key) or (field_java_types or {}).get(str(raw_key))
result[key] = JavaRedisCodec._encode_value(item, java_type=child_type, nested=True)
return result
if isinstance(value, set | frozenset):
return [
"java.util.HashSet",
[
JavaRedisCodec._encode_value(item, java_type=item_java_type, nested=True)
for item in value
],
]
if isinstance(value, list | tuple):
return [
"java.util.ArrayList",
[
JavaRedisCodec._encode_value(item, java_type=item_java_type, nested=True)
for item in value
],
]
return value
@staticmethod
def _decode_value(value: Any) -> Any:
if isinstance(value, dict):
return {
str(key): JavaRedisCodec._decode_value(item)
for key, item in value.items()
if key != "@class"
}
if isinstance(value, list):
if len(value) == 2 and isinstance(value[0], str) and value[0].startswith("java."):
type_name, payload = value
if type_name == "java.util.Date":
timezone = ZoneInfo(get_settings().timezone)
return datetime.fromtimestamp(float(payload) / 1000, timezone).replace(tzinfo=None)
if type_name in {
"java.util.ArrayList",
"java.util.LinkedList",
"java.util.HashSet",
"java.util.LinkedHashSet",
} and isinstance(payload, list):
return [JavaRedisCodec._decode_value(item) for item in payload]
return JavaRedisCodec._decode_value(payload)
return [JavaRedisCodec._decode_value(item) for item in value]
return value
def _snake_to_camel(value: str) -> str:
head, *tail = value.split("_")
return head + "".join(part[:1].upper() + part[1:] for part in tail)
def get_redis() -> Redis:
global _client
if _client is None:
_client = Redis.from_url(get_settings().redis_url, decode_responses=False)
return _client
async def close_redis() -> None:
global _client
if _client is not None:
await _client.aclose()
_client = None
async def redis_ping() -> bool:
try:
return bool(await get_redis().ping())
except Exception:
return False
async def java_get(key: str) -> Any:
return JavaRedisCodec.decode(await cast(Any, get_redis().get(key)))
async def java_set(
key: str,
value: Any,
ttl_seconds: int | None = None,
*,
java_type: str | None = None,
item_java_type: str | None = None,
) -> None:
await cast(Any, get_redis().set)(
key,
JavaRedisCodec.encode(value, java_type=java_type, item_java_type=item_java_type),
ex=ttl_seconds,
)
async def java_hget(key: str, field: str) -> Any:
return JavaRedisCodec.decode(await cast(Any, get_redis().hget(key, field)))
async def java_hset(key: str, field: str, value: Any, ttl_seconds: int = 86400) -> None:
redis = get_redis()
await cast(Any, redis.hset)(key, field, JavaRedisCodec.encode(value))
await cast(Any, redis.expire)(key, ttl_seconds)
@asynccontextmanager
async def distributed_lock(name: str, ttl_seconds: int) -> AsyncIterator[bool]:
lock = get_redis().lock(name, timeout=ttl_seconds, blocking_timeout=0)
acquired = bool(await lock.acquire(blocking=False))
renewal: asyncio.Task[None] | None = None
if acquired:
renewal = asyncio.create_task(_renew_lock(lock, ttl_seconds))
try:
yield acquired
finally:
if renewal is not None:
renewal.cancel()
with suppress(asyncio.CancelledError):
await renewal
if acquired:
try:
await lock.release()
except Exception:
logger.warning("Lost ownership of distributed lock %s before release", name, exc_info=True)
async def _renew_lock(lock: Any, ttl_seconds: int) -> None:
"""Keep a held job lock alive until its owner leaves the context."""
interval = max(float(ttl_seconds) / 3, 0.25)
while True:
await asyncio.sleep(interval)
try:
await lock.extend(ttl_seconds, replace_ttl=True)
except Exception:
logger.exception("Unable to renew distributed lock; duplicate execution protection is at risk")
return
@@ -0,0 +1,66 @@
from __future__ import annotations
import json
from typing import Any
from fastapi import Request
from starlette.responses import JSONResponse, Response
from app.core.i18n import message_for
from app.core.serialization import java_compatible
class JavaJSONResponse(JSONResponse):
def render(self, content: Any) -> bytes:
return json.dumps(
java_compatible(content),
ensure_ascii=False,
allow_nan=False,
separators=(",", ":"),
).encode("utf-8")
def envelope(data: Any = None, *, code: int = 0, msg: str = "success") -> dict[str, Any]:
return {"code": code, "msg": msg, "data": data}
def ok(data: Any = None) -> JavaJSONResponse:
return JavaJSONResponse(envelope(data))
def error_response(
request: Request,
code: int,
message: str | None = None,
*,
status_code: int = 200,
params: tuple[object, ...] = (),
media_type: str = "application/json",
) -> JavaJSONResponse:
translated = message or message_for(code, request.headers.get("Accept-Language"), *params)
return JavaJSONResponse(
envelope(None, code=code, msg=translated),
status_code=status_code,
media_type=media_type,
)
def raw_json(content: Any, *, exclude_none: bool = False, status_code: int = 200) -> Response:
normalized = java_compatible(content)
if exclude_none:
normalized = _drop_none(normalized)
body = json.dumps(normalized, ensure_ascii=False, allow_nan=False, separators=(",", ":")).encode("utf-8")
return Response(
body,
status_code=status_code,
media_type="application/json",
headers={"Content-Length": str(len(body))},
)
def _drop_none(value: Any) -> Any:
if isinstance(value, dict):
return {key: _drop_none(item) for key, item in value.items() if item is not None}
if isinstance(value, list):
return [_drop_none(item) for item in value]
return value
@@ -0,0 +1,181 @@
from __future__ import annotations
import fnmatch
import hmac
from dataclasses import dataclass
from datetime import datetime
from typing import Any
from fastapi import Request
from sqlalchemy import text
from starlette.middleware.base import BaseHTTPMiddleware, RequestResponseEndpoint
from starlette.responses import Response
from app.core.config import get_settings
from app.core.database import get_session_factory
from app.core.errors import AppError, ErrorCode
from app.core.responses import error_response
from app.services.system_params import SystemParamService
PUBLIC_PATTERNS = (
"/ota/*",
"/ota",
"/otaMag/download/*",
"/webjars/*",
"/druid/*",
"/v3/api-docs*",
"/doc.html*",
"/favicon.ico",
"/user/captcha",
"/user/smsVerification",
"/user/login",
"/user/pub-config",
"/user/register",
"/user/retrieve-password",
"/api/ping",
"/agent/chat-history/download/*",
"/agent/play/*",
"/voiceClone/play/*",
"/health",
"/health/live",
"/health/ready",
)
SERVER_PATTERNS = (
"/config/*",
"/device/address-book/call",
"/device/address-book/lookup",
"/agent/chat-history/report",
"/agent/chat-summary/*",
"/agent/chat-title/*",
)
@dataclass(slots=True, frozen=True)
class AuthUser:
id: int
username: str
super_admin: int
status: int
token: str
row: dict[str, Any]
@property
def is_super_admin(self) -> bool:
return self.super_admin == 1
def _matches(path: str, patterns: tuple[str, ...]) -> bool:
return any(fnmatch.fnmatchcase(path, pattern) for pattern in patterns)
def _bearer_token(request: Request) -> str | None:
authorization = request.headers.get("Authorization")
if not authorization or not authorization.startswith("Bearer "):
return None
value = authorization[len("Bearer ") :]
return value if value.strip() else None
class AuthenticationMiddleware(BaseHTTPMiddleware):
async def dispatch(self, request: Request, call_next: RequestResponseEndpoint) -> Response:
if request.method == "OPTIONS":
return await call_next(request)
settings = get_settings()
path = request.url.path
if settings.context_path and path.startswith(settings.context_path):
path = path[len(settings.context_path) :] or "/"
if _matches(path, PUBLIC_PATTERNS):
request.state.auth_mode = "anonymous"
return await call_next(request)
if _matches(path, SERVER_PATTERNS):
return await self._server_auth(request, call_next)
return await self._user_auth(request, call_next)
async def _server_auth(self, request: Request, call_next: RequestResponseEndpoint) -> Response:
provided = _bearer_token(request)
if provided is None:
return error_response(
request,
ErrorCode.UNAUTHORIZED,
"服务器密钥不能为空",
media_type="application/json;charset=utf-8",
)
expected = get_settings().server_secret_override
if expected is None:
try:
async with get_session_factory()() as session:
expected = await SystemParamService(session).get_value("server.secret", from_cache=True)
except Exception:
expected = None
if not expected or not hmac.compare_digest(provided, expected):
return error_response(
request,
ErrorCode.UNAUTHORIZED,
"无效的服务器密钥",
media_type="application/json;charset=utf-8",
)
request.state.auth_mode = "server"
return await call_next(request)
async def _user_auth(self, request: Request, call_next: RequestResponseEndpoint) -> Response:
token = _bearer_token(request)
if token is None:
return error_response(
request,
ErrorCode.UNAUTHORIZED,
media_type="application/json;charset=utf-8",
)
try:
async with get_session_factory()() as session:
result = await session.execute(
text(
"SELECT u.* FROM sys_user_token t "
"JOIN sys_user u ON u.id = t.user_id "
"WHERE t.token = :token AND t.expire_date >= CURRENT_TIMESTAMP LIMIT 1"
),
{"token": token},
)
mapping = result.mappings().first()
except Exception:
mapping = None
if mapping is None or mapping.get("status") is None or int(mapping["status"]) != 1:
return error_response(
request,
ErrorCode.UNAUTHORIZED,
media_type="application/json;charset=utf-8",
)
row = dict(mapping)
request.state.user = AuthUser(
id=int(row["id"]),
username=str(row.get("username") or ""),
super_admin=int(row.get("super_admin") or 0),
status=int(row["status"]),
token=token,
row=row,
)
request.state.auth_mode = "user"
return await call_next(request)
def current_user(request: Request) -> AuthUser:
user = getattr(request.state, "user", None)
if not isinstance(user, AuthUser):
raise AppError(ErrorCode.UNAUTHORIZED)
return user
def require_normal(request: Request) -> AuthUser:
return current_user(request)
def require_super_admin(request: Request) -> AuthUser:
user = current_user(request)
if not user.is_super_admin:
raise AppError(ErrorCode.FORBIDDEN)
return user
def shanghai_now_naive() -> datetime:
from zoneinfo import ZoneInfo
return datetime.now(tz=ZoneInfo(get_settings().timezone)).replace(tzinfo=None)
@@ -0,0 +1,105 @@
from __future__ import annotations
import dataclasses
import re
from collections.abc import Mapping, Sequence
from datetime import date, datetime, time
from decimal import Decimal
from enum import Enum
from pathlib import Path
from typing import Any
from zoneinfo import ZoneInfo
from pydantic import BaseModel
from app.core.config import get_settings
_SNAKE_PART = re.compile(r"_([a-zA-Z0-9])")
_LONG_FIELD_NAMES = {
"id",
"userId",
"creator",
"updater",
"createUserId",
"updateUserId",
"createDateTimestamp",
"createTime",
"createTimeFrom",
"createTimeTo",
"fileSize",
"lastConnectedAtTimestamp",
"pid",
"reportTime",
"size",
"timestamp",
"tokenCount",
"tokenNum",
"totalDocCount",
"totalTokenCount",
"updateTime",
}
class JavaMap(dict[str, Any]):
"""Marker for Java ``Map`` payloads whose keys Jackson leaves untouched."""
def preserve_java_map_keys(value: Any) -> Any:
"""Recursively mark a dynamic Java Map/List graph as key-preserving."""
if isinstance(value, Mapping):
return JavaMap({str(key): preserve_java_map_keys(item) for key, item in value.items()})
if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray):
return [preserve_java_map_keys(item) for item in value]
return value
def snake_to_camel(value: str) -> str:
return _SNAKE_PART.sub(lambda match: match.group(1).upper(), value)
def _is_long_field(name: str | None) -> bool:
if not name:
return False
return name in _LONG_FIELD_NAMES or name.endswith("Id") or name.endswith("Ids")
def java_compatible(value: Any, *, field_name: str | None = None) -> Any:
if value is None or isinstance(value, str | bool | float):
return value
if isinstance(value, BaseModel):
return java_compatible(value.model_dump(by_alias=True, exclude_unset=False), field_name=field_name)
if dataclasses.is_dataclass(value) and not isinstance(value, type):
return java_compatible(dataclasses.asdict(value), field_name=field_name)
if isinstance(value, Enum):
return java_compatible(value.value, field_name=field_name)
if isinstance(value, datetime):
timezone = ZoneInfo(get_settings().timezone)
localized = value.astimezone(timezone) if value.tzinfo else value
return localized.strftime("%Y-%m-%d %H:%M:%S")
if isinstance(value, date):
return value.strftime("%Y-%m-%d")
if isinstance(value, time):
return value.strftime("%H:%M:%S")
if isinstance(value, Decimal):
return float(value)
if isinstance(value, int):
return str(value) if _is_long_field(field_name) or not -(2**31) <= value < 2**31 else value
if isinstance(value, bytes):
return value
if isinstance(value, Path):
return str(value)
if isinstance(value, JavaMap):
return {
str(raw_key): java_compatible(item, field_name=snake_to_camel(str(raw_key)))
for raw_key, item in value.items()
}
if isinstance(value, Mapping):
result: dict[str, Any] = {}
for raw_key, item in value.items():
key = snake_to_camel(str(raw_key))
result[key] = java_compatible(item, field_name=key)
return result
if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray):
return [java_compatible(item, field_name=field_name) for item in value]
return value
@@ -0,0 +1 @@
"""Outbound integrations used by the FastAPI manager service."""
@@ -0,0 +1,68 @@
from __future__ import annotations
from typing import Any
import httpx
SUMMARY_PROMPT = """你是一个经验丰富的记忆总结者,擅长将对话内容进行总结摘要,遵循以下规则:
1、总结用户的重要信息,以便在未来的对话中提供更个性化的服务
2、不要重复总结,不要遗忘之前记忆,除非原来的记忆超过了1800字,否则不要遗忘、不要压缩用户的历史记忆
3、用户操控的设备音量、播放音乐、天气、退出、不想对话等和用户本身无关的内容,这些信息不需要加入到总结中
4、聊天内容中的今天的日期时间、今天的天气情况与用户事件无关的数据,这些信息如果当成记忆存储会影响后续对话,这些信息不需要加入到总结中
5、不要把设备操控的成果结果和失败结果加入到总结中,也不要把用户的一些废话加入到总结中
6、不要为了总结而总结,如果用户的聊天没有意义,请返回原来的历史记录也是可以的
7、只需要返回总结摘要,严格控制在1800字内
8、不要包含代码、xml,不需要解释、注释和说明,保存记忆时仅从对话提取信息,不要混入示例内容
9、如果提供了历史记忆,请将新对话内容与历史记忆进行智能合并,保留有价值的历史信息,同时添加新的重要信息
历史记忆:
{history_memory}
新对话内容:
{conversation}"""
TITLE_PROMPT = (
"请根据以下对话内容,生成一个简洁的会话标题(约15字以内),只返回标题,不要包含任何解释或标点符号:\n{conversation}"
)
def _apply_thinking_policy(base_url: str, request: dict[str, Any]) -> None:
if "aliyuncs.com" in base_url:
request["enable_thinking"] = False
elif any(domain in base_url for domain in ("bigmodel.cn", "moonshot.cn", "volces.com")):
request["thinking"] = {"type": "disabled"}
async def openai_completion(
config: dict[str, Any],
prompt: str,
*,
temperature: float,
max_tokens: int,
timeout: float,
) -> str | None:
base_url = str(config.get("base_url") or "")
api_key = str(config.get("api_key") or "")
if not base_url.strip() or not api_key.strip():
return None
api_url = base_url if base_url.endswith("/chat/completions") else f"{base_url.rstrip('/')}/chat/completions"
request: dict[str, Any] = {
"model": config.get("model_name") or "gpt-3.5-turbo",
"messages": [{"role": "user", "content": prompt}],
"temperature": temperature,
"max_tokens": max_tokens,
}
_apply_thinking_policy(base_url, request)
async with httpx.AsyncClient(timeout=timeout) as client:
response = await client.post(
api_url,
json=request,
headers={"Content-Type": "application/json", "Authorization": f"Bearer {api_key}"},
)
response.raise_for_status()
payload = response.json()
choices = payload.get("choices") if isinstance(payload, dict) else None
if not isinstance(choices, list) or not choices:
return None
message = choices[0].get("message") if isinstance(choices[0], dict) else None
content = message.get("content") if isinstance(message, dict) else None
return str(content) if content is not None else None
@@ -0,0 +1,106 @@
from __future__ import annotations
import asyncio
import base64
import hashlib
import json
from typing import Any
from urllib.parse import quote_plus, urlsplit, urlunsplit
from cryptography.hazmat.primitives import padding
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
from websockets.asyncio.client import connect
def _java_aes_key(value: str) -> bytes:
raw = value.encode("utf-8")
if len(raw) in {16, 24, 32}:
return raw
return raw[:32].ljust(32, b"\x00")
def encrypt_agent_token(agent_id: str, key: str) -> str:
digest = hashlib.md5(agent_id.encode("utf-8"), usedforsecurity=False).hexdigest()
plain_text = f'{{"agentId": "{digest}"}}'.encode()
padder = padding.PKCS7(128).padder()
padded = padder.update(plain_text) + padder.finalize()
# ECB is required for byte-for-byte compatibility with Java AES/ECB/PKCS5Padding.
encryptor = Cipher(algorithms.AES(_java_aes_key(key)), modes.ECB()).encryptor() # noqa: S305
encrypted = encryptor.update(padded) + encryptor.finalize()
return base64.b64encode(encrypted).decode("ascii")
def build_agent_mcp_address(endpoint: str | None, agent_id: str) -> str | None:
if endpoint is None or not endpoint.strip() or endpoint == "null":
return None
parsed = urlsplit(endpoint)
if not parsed.scheme or not parsed.netloc:
raise ValueError("mcp的地址存在错误,请进入参数管理修改mcp接入点地址")
marker = "key="
marker_index = parsed.query.find(marker)
# Java takes everything following the first key= marker, including subsequent query text.
key = parsed.query[marker_index + len(marker) :] if marker_index >= 0 else parsed.query[3:]
ws_scheme = "wss" if parsed.scheme == "https" else "ws"
path = parsed.path
parent = path[: path.rfind("/")] if "/" in path else ""
base = urlunsplit((ws_scheme, parsed.netloc, parent, "", "")).rstrip("/")
token = quote_plus(encrypt_agent_token(agent_id, key), safe="")
return f"{base}/mcp/?token={token}"
INITIALIZE_REQUEST = {
"jsonrpc": "2.0",
"method": "initialize",
"params": {
"protocolVersion": "2024-11-05",
"capabilities": {"roots": {"listChanged": False}, "sampling": {}},
"clientInfo": {"name": "xz-mcp-broker", "version": "0.0.1"},
},
"id": 1,
}
INITIALIZED_NOTIFICATION = {"jsonrpc": "2.0", "method": "notifications/initialized"}
TOOLS_REQUEST = {"jsonrpc": "2.0", "method": "tools/list", "params": None, "id": 2}
async def _receive_matching(websocket: Any, request_id: int, timeout: float) -> dict[str, Any] | None:
async def receive() -> dict[str, Any] | None:
async for message in websocket:
try:
value = json.loads(message)
except (TypeError, json.JSONDecodeError):
continue
if isinstance(value, dict) and value.get("id") == request_id:
return value
return None
return await asyncio.wait_for(receive(), timeout=timeout)
async def list_mcp_tools(address: str, *, connect_timeout: float = 8.0, session_timeout: float = 10.0) -> list[str]:
call_address = address.replace("/mcp/", "/call/")
try:
async with connect(
call_address,
open_timeout=connect_timeout,
max_size=1024 * 1024,
close_timeout=1,
) as websocket:
await websocket.send(json.dumps(INITIALIZE_REQUEST, ensure_ascii=False, separators=(",", ":")))
initialized = await _receive_matching(websocket, 1, session_timeout)
if not initialized or "result" not in initialized or "error" in initialized:
return []
await websocket.send(json.dumps(INITIALIZED_NOTIFICATION, separators=(",", ":")))
await websocket.send(json.dumps(TOOLS_REQUEST, separators=(",", ":")))
response = await _receive_matching(websocket, 2, session_timeout)
if not response or "error" in response:
return []
result = response.get("result")
tools = result.get("tools") if isinstance(result, dict) else None
if not isinstance(tools, list):
return []
return sorted(
item["name"] for item in tools if isinstance(item, dict) and isinstance(item.get("name"), str)
)
# Java treats every connect/protocol/parse failure as an empty tool list.
except Exception:
return []
@@ -0,0 +1,63 @@
from __future__ import annotations
import hashlib
import json
from datetime import date, datetime, timedelta, timezone
from typing import Any
import httpx
class MqttGatewayError(RuntimeError):
def __init__(self, message: str, status_code: int | None = None):
super().__init__(message)
self.status_code = status_code
def daily_authorization_tokens(signature_key: str, now: datetime | None = None) -> list[str]:
if not signature_key.strip() or signature_key.strip().lower() == "null":
raise MqttGatewayError("MQTT Gateway signature key is empty")
instant = now or datetime.now(tz=timezone.utc)
utc_date = instant.astimezone(timezone.utc).date()
dates: tuple[date, date, date] = (utc_date, utc_date - timedelta(days=1), utc_date + timedelta(days=1))
return [hashlib.sha256(f"{value.isoformat()}{signature_key}".encode()).hexdigest() for value in dates]
async def post_json(
url: str,
body: Any,
signature_key: str,
*,
timeout_seconds: float,
now: datetime | None = None,
client: httpx.AsyncClient | None = None,
) -> str:
encoded = json.dumps(body, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
owns_client = client is None
selected = client or httpx.AsyncClient()
last_unauthorized: int | None = None
try:
for token in daily_authorization_tokens(signature_key, now):
response = await selected.post(
url,
content=encoded,
headers={"Content-Type": "application/json", "Authorization": f"Bearer {token}"},
timeout=timeout_seconds,
)
if response.status_code == 401:
last_unauthorized = response.status_code
continue
if not 200 <= response.status_code < 300:
raise MqttGatewayError(
f"MQTT Gateway request failed with HTTP status {response.status_code}",
response.status_code,
)
return response.text
finally:
if owns_client:
await selected.aclose()
raise MqttGatewayError(
"MQTT Gateway rejected all daily authorization tokens"
+ ("" if last_unauthorized is None else f" (HTTP {last_unauthorized})"),
last_unauthorized,
)
@@ -0,0 +1,490 @@
from __future__ import annotations
from collections.abc import Mapping
from typing import Any
import httpx
from fastapi import UploadFile
from app.core.errors import AppError
_DOCUMENT_CHUNK_METHODS = {
"naive",
"manual",
"qa",
"table",
"paper",
"book",
"laws",
"presentation",
"picture",
"one",
"knowledge_graph",
"email",
}
_RUN_STATUSES = {"UNSTART", "RUNNING", "CANCEL", "DONE", "FAIL"}
_DOCUMENT_PARSER_FIELDS = (
"chunk_token_num",
"delimiter",
"layout_recognize",
"html4excel",
"auto_keywords",
"auto_questions",
"topn_tags",
"raptor",
"graphrag",
)
_DATASET_PARSER_FIELDS = (
"chunk_token_num",
"delimiter",
"layout_recognize",
"html4excel",
"auto_keywords",
"auto_questions",
)
class RAGFlowClient:
"""Async equivalent of the Java RAGFlow adapter and its wire contract."""
def __init__(self, config: Mapping[str, Any]):
self.config = dict(config)
self.base_url = str(config.get("base_url") or config.get("baseUrl") or "").rstrip("/")
self.api_key = str(config.get("api_key") or config.get("apiKey") or "")
raw_timeout = config.get("timeout")
if raw_timeout is None:
self.timeout = 30.0
else:
try:
self.timeout = float(int(str(raw_timeout)))
except (TypeError, ValueError):
self.timeout = 30.0
self._validate(config)
def _validate(self, config: Mapping[str, Any]) -> None:
if not config:
raise AppError(10164)
if not self.base_url.strip():
raise AppError(10171)
if not self.api_key.strip():
raise AppError(10172)
if "" in self.api_key:
raise AppError(10173)
if not self.base_url.startswith(("http://", "https://")):
raise AppError(10174)
adapter_type = "ragflow" if "type" not in config else str(config.get("type"))
if adapter_type != "ragflow":
raise AppError(10184, params=(f"适配器类型未注册: {adapter_type}",))
async def request(
self,
method: str,
endpoint: str,
*,
params: Mapping[str, Any] | None = None,
json_body: Any = None,
files: Mapping[str, Any] | None = None,
data: Mapping[str, Any] | None = None,
) -> dict[str, Any]:
headers = {"Authorization": f"Bearer {self.api_key}"}
if files is None:
headers["Content-Type"] = "application/json"
headers["Accept-Charset"] = "utf-8"
normalized_params = {
key: self._query_value(value) for key, value in (params or {}).items() if value is not None
}
try:
async with httpx.AsyncClient(timeout=self.timeout) as client:
response = await client.request(
method,
self.base_url + endpoint,
params=normalized_params,
json=json_body,
files=files,
data=data,
headers=headers,
)
response.raise_for_status()
payload = response.json()
except (httpx.HTTPError, ValueError) as exc:
raise AppError(10167, params=(f"Request Failed: {exc}",)) from exc
if not isinstance(payload, dict):
raise AppError(10167, params=("Invalid Response",))
code = payload.get("code")
if code is not None:
if isinstance(code, bool) or not isinstance(code, int):
raise AppError(10167, params=("Request Failed: invalid response code type",))
if code != 0:
message = payload.get("message")
if message is not None and not isinstance(message, str):
raise AppError(10167, params=("Request Failed: invalid response message type",))
raise AppError(10167, params=(message or "Unknown RAGFlow Error",))
return dict(payload)
@staticmethod
def _query_value(value: Any) -> Any:
if isinstance(value, bool):
return str(value).lower()
if isinstance(value, list):
# Java List.toString() is what the baseline URL builder sends.
return "[" + ", ".join(str(item) for item in value) + "]"
return value
async def dataset_info(self, dataset_id: str) -> dict[str, Any] | None:
payload = await self.request(
"GET", "/api/v1/datasets", params={"id": dataset_id, "page": 1, "page_size": 1}
)
data = payload.get("data")
if isinstance(data, list) and data and isinstance(data[0], dict):
return _normalize_dataset_info(data[0])
return None
async def create_dataset(self, body: dict[str, Any]) -> dict[str, Any]:
body = dict(body)
body["permission"] = "me" if _blank(body.get("permission")) else body.get("permission")
body["chunk_method"] = "naive" if _blank(body.get("chunk_method")) else body.get("chunk_method")
if _blank(body.get("embedding_model")):
configured_model = self.config.get("embedding_model", self.config.get("embeddingModel"))
body["embedding_model"] = None if _blank(configured_model) else configured_model
body["avatar"] = body.get("avatar") if not _blank(body.get("avatar")) else (
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8z8BQDwAEhQGAhKmMIQAAAABJRU5ErkJggg=="
)
body["parser_config"] = _normalize_parser_config(
body.get("parser_config"), fields=_DATASET_PARSER_FIELDS
)
payload = await self.request("POST", "/api/v1/datasets", json_body=body)
data = payload.get("data")
if not isinstance(data, dict) or not data.get("id"):
raise AppError(10167, params=("Invalid response from createDataset: missing data object",))
return _normalize_dataset_info(data)
async def update_dataset(self, dataset_id: str, body: dict[str, Any]) -> dict[str, Any] | None:
body = dict(body)
body["parser_config"] = _normalize_parser_config(
body.get("parser_config"), fields=_DATASET_PARSER_FIELDS
)
payload = await self.request("PUT", f"/api/v1/datasets/{dataset_id}", json_body=body)
return _normalize_dataset_info(payload["data"]) if isinstance(payload.get("data"), dict) else None
async def delete_datasets(self, ids: list[str]) -> Any:
return (await self.request("DELETE", "/api/v1/datasets", json_body={"ids": ids})).get("data")
async def documents(
self,
dataset_id: str,
*,
page: int = 1,
page_size: int = 10,
name: str | None = None,
status: str | None = None,
document_id: str | None = None,
) -> tuple[list[dict[str, Any]], int]:
params: dict[str, Any] = {"page": page, "page_size": page_size}
if name:
params["name"] = name
if status:
status_number = int(status) if status.lstrip("-").isdigit() else None
names = {0: "UNSTART", 1: "RUNNING", 2: "CANCEL", 3: "DONE", 4: "FAIL"}
params["run"] = [names[status_number]] if status_number in names else []
if document_id:
params["id"] = document_id
payload = await self.request("GET", f"/api/v1/datasets/{dataset_id}/documents", params=params)
data = payload.get("data")
if not isinstance(data, dict):
return [], 0
docs = data.get("docs")
rows: list[dict[str, Any]] = []
if isinstance(docs, list):
for item in docs:
if not isinstance(item, dict):
continue
try:
rows.append(_normalize_upload_document(item))
except AppError:
# The Java adapter skips an individual document whose
# strong DTO conversion fails and keeps the rest of page.
continue
return rows, int(data.get("total") or 0)
async def upload_document(
self,
dataset_id: str,
file: UploadFile,
content: bytes,
*,
name: str,
meta_fields: dict[str, Any] | None,
chunk_method: str | None,
parser_config: dict[str, Any] | None,
) -> dict[str, Any]:
import json
form: dict[str, Any] = {"name": name}
if meta_fields:
form["meta"] = json.dumps(meta_fields, ensure_ascii=False, separators=(",", ":"))
if not _blank(chunk_method):
normalized_method = str(chunk_method).lower()
if normalized_method in _DOCUMENT_CHUNK_METHODS:
form["chunk_method"] = normalized_method
normalized_parser = (
_normalize_parser_config(
parser_config,
fields=_DOCUMENT_PARSER_FIELDS,
validate_layout=True,
)
if parser_config
else None
)
if normalized_parser is not None:
form["parser_config"] = json.dumps(normalized_parser, ensure_ascii=False, separators=(",", ":"))
payload = await self.request(
"POST",
f"/api/v1/datasets/{dataset_id}/documents",
files={"file": (file.filename or name, content, file.content_type or "application/octet-stream")},
data=form,
)
data = payload.get("data")
if isinstance(data, list) and data and isinstance(data[0], dict):
return _normalize_upload_document(data[0])
if isinstance(data, dict):
return _normalize_upload_document(data)
raise AppError(10167, params=("远程上传成功但未返回有效 DocumentID",))
async def delete_documents(self, dataset_id: str, ids: list[str]) -> None:
await self.request(
"DELETE", f"/api/v1/datasets/{dataset_id}/documents", json_body={"ids": ids}
)
async def parse_documents(self, dataset_id: str, document_ids: list[str]) -> None:
await self.request(
"POST",
f"/api/v1/datasets/{dataset_id}/chunks",
json_body={"document_ids": document_ids},
)
async def chunks(
self, dataset_id: str, document_id: str, params: Mapping[str, Any]
) -> dict[str, Any]:
payload = await self.request(
"GET", f"/api/v1/datasets/{dataset_id}/documents/{document_id}/chunks", params=params
)
data = payload.get("data")
if not isinstance(data, dict):
return {"chunks": [], "doc": None, "total": 0}
try:
return _normalize_chunk_list(data)
except (TypeError, ValueError) as exc:
raise AppError(10167, params=(str(exc),)) from exc
async def retrieval(self, body: dict[str, Any]) -> dict[str, Any]:
payload = await self.request("POST", "/api/v1/retrieval", json_body=body)
data = payload.get("data")
if not isinstance(data, dict):
return {"chunks": [], "doc_aggs": [], "total": 0}
try:
return _normalize_retrieval_result(data)
except (TypeError, ValueError) as exc:
raise AppError(10167, params=(str(exc),)) from exc
def _blank(value: Any) -> bool:
return value is None or (isinstance(value, str) and not value.strip())
def _normalize_parser_config(
value: Any,
*,
fields: tuple[str, ...],
validate_layout: bool = False,
) -> dict[str, Any] | None:
if value is None:
return None
if not isinstance(value, Mapping):
raise ValueError("parser_config must be an object")
result = {key: value.get(key) for key in fields}
if validate_layout and result.get("layout_recognize") not in {None, "DeepDOC", "Simple"}:
raise ValueError("invalid layout_recognize")
if "raptor" in result and result["raptor"] is not None:
nested = result["raptor"]
if not isinstance(nested, Mapping):
raise ValueError("raptor must be an object")
result["raptor"] = {"use_raptor": nested.get("use_raptor")}
if "graphrag" in result and result["graphrag"] is not None:
nested = result["graphrag"]
if not isinstance(nested, Mapping):
raise ValueError("graphrag must be an object")
result["graphrag"] = {"use_graphrag": nested.get("use_graphrag")}
return result
def _normalize_upload_document(value: Mapping[str, Any]) -> dict[str, Any]:
result = dict(value)
try:
if result.get("parser_config") is not None:
result["parser_config"] = _normalize_parser_config(
result["parser_config"], fields=_DOCUMENT_PARSER_FIELDS, validate_layout=True
)
except (TypeError, ValueError) as exc:
raise AppError(10167, params=("远程上传成功但未返回有效 DocumentID",)) from exc
chunk_method = result.get("chunk_method")
if chunk_method is not None:
normalized_method = str(chunk_method).lower()
if normalized_method not in _DOCUMENT_CHUNK_METHODS:
raise AppError(10167, params=("远程上传成功但未返回有效 DocumentID",))
result["chunk_method"] = normalized_method
run = result.get("run")
if run is not None and str(run) not in _RUN_STATUSES:
raise AppError(10167, params=("远程上传成功但未返回有效 DocumentID",))
return result
def _normalize_dataset_info(value: Mapping[str, Any]) -> dict[str, Any]:
"""Apply Jackson's DatasetDTO.InfoVO unknown-field and type boundary."""
fields = (
"id",
"name",
"avatar",
"tenant_id",
"description",
"embedding_model",
"permission",
"chunk_method",
"parser_config",
"chunk_count",
"document_count",
"create_time",
"update_time",
"token_num",
"create_date",
"update_date",
)
try:
result = {field: value.get(field) for field in fields}
result["parser_config"] = _normalize_parser_config(
result.get("parser_config"), fields=_DATASET_PARSER_FIELDS
)
for field in ("chunk_count", "document_count", "create_time", "update_time", "token_num"):
raw_value = result[field]
if raw_value is not None:
result[field] = int(raw_value)
return result
except (TypeError, ValueError) as exc:
raise AppError(10167, params=(str(exc),)) from exc
def _nullable_object(value: Any, fields: tuple[str, ...]) -> dict[str, Any] | None:
if value is None:
return None
if not isinstance(value, Mapping):
raise TypeError("response object has an invalid shape")
return {field: value.get(field) for field in fields}
def _normalize_chunk_list(data: Mapping[str, Any]) -> dict[str, Any]:
chunk_fields = (
"id",
"content",
"document_id",
"docnm_kwd",
"important_keywords",
"questions",
"image_id",
"dataset_id",
"available",
"positions",
"token",
)
raw_chunks = data.get("chunks")
if raw_chunks is None:
chunks: list[dict[str, Any]] = []
elif isinstance(raw_chunks, list):
chunks = []
for item in raw_chunks:
normalized = _nullable_object(item, chunk_fields)
if normalized is not None:
chunks.append(normalized)
else:
raise TypeError("chunks must be an array")
doc_fields = (
"id",
"thumbnail",
"dataset_id",
"chunk_method",
"pipeline_id",
"parser_config",
"source_type",
"type",
"created_by",
"name",
"location",
"size",
"token_count",
"chunk_count",
"progress",
"progress_msg",
"process_begin_at",
"process_duration",
"meta_fields",
"suffix",
"run",
"status",
"create_time",
"create_date",
"update_time",
"update_date",
)
doc = _nullable_object(data.get("doc"), doc_fields)
if doc is not None:
doc["parser_config"] = _normalize_parser_config(
doc.get("parser_config"), fields=_DOCUMENT_PARSER_FIELDS, validate_layout=True
)
if doc.get("chunk_count") is not None:
# DocumentDTO.InfoVO.chunkCount is Long, unlike the Integer field
# on KnowledgeFilesDTO used by the document-list endpoint.
doc["chunk_count"] = str(doc["chunk_count"])
if doc.get("run") is not None and str(doc["run"]) not in _RUN_STATUSES:
raise ValueError("invalid document run status")
# ChunkDTO.ListVO.total is Long and therefore uses the Java global Long
# serializer even for small values (including the adapter's default 0L).
return {"chunks": chunks, "doc": doc, "total": str(int(data.get("total") or 0))}
def _normalize_retrieval_result(data: Mapping[str, Any]) -> dict[str, Any]:
hit_fields = (
"id",
"content",
"document_id",
"dataset_id",
"document_name",
"document_keyword",
"similarity",
"vector_similarity",
"term_similarity",
"index",
"highlight",
"important_keywords",
"questions",
"image_id",
"positions",
)
agg_fields = ("doc_name", "doc_id", "count")
def normalize_list(raw: Any, fields: tuple[str, ...], name: str) -> list[dict[str, Any]]:
if raw is None:
return []
if not isinstance(raw, list):
raise TypeError(f"{name} must be an array")
values: list[dict[str, Any]] = []
for item in raw:
normalized = _nullable_object(item, fields)
if normalized is not None:
values.append(normalized)
return values
return {
"chunks": normalize_list(data.get("chunks"), hit_fields, "chunks"),
"doc_aggs": normalize_list(data.get("doc_aggs"), agg_fields, "doc_aggs"),
# RetrievalDTO.ResultVO.total is also Long.
"total": str(int(data.get("total") or 0)),
}
@@ -0,0 +1,96 @@
from __future__ import annotations
import base64
import json
from dataclasses import dataclass
from typing import Any
import httpx
@dataclass(slots=True)
class VoiceCloneProviderError(Exception):
code: int
message: str
def __str__(self) -> str:
return self.message
class VoiceCloneIntegration:
ENDPOINT = "https://openspeech.bytedance.com/api/v1/mega_tts/audio/upload"
def __init__(
self,
*,
timeout_seconds: float,
client: httpx.AsyncClient | None = None,
endpoint: str | None = None,
):
self.timeout_seconds = timeout_seconds
self.client = client
self.endpoint = endpoint or self.ENDPOINT
async def train_huoshan(
self,
*,
appid: str,
access_token: str,
voice: bytes,
speaker_id: str,
) -> str:
request_body: dict[str, Any] = {
"appid": appid,
"audios": [
{
"audio_bytes": base64.b64encode(voice).decode("ascii"),
"audio_format": "wav",
}
],
"source": 2,
"language": 0,
"model_type": 1,
"speaker_id": speaker_id,
}
owns_client = self.client is None
client = self.client or httpx.AsyncClient()
try:
response = await client.post(
self.endpoint,
content=json.dumps(request_body, ensure_ascii=False, separators=(",", ":")).encode("utf-8"),
headers={
"Content-Type": "application/json",
"Authorization": f"Bearer;{access_token}",
"Resource-Id": "seed-icl-1.0",
},
timeout=self.timeout_seconds,
)
try:
payload = response.json()
except (json.JSONDecodeError, ValueError) as exc:
raise VoiceCloneProviderError(10157, str(exc)) from exc
except httpx.HTTPError as exc:
raise VoiceCloneProviderError(10157, str(exc)) from exc
finally:
if owns_client:
await client.aclose()
if not isinstance(payload, dict):
raise VoiceCloneProviderError(10156, "响应格式错误,缺少BaseResp字段")
base_response = payload.get("BaseResp")
if isinstance(base_response, dict):
raw_status = base_response.get("StatusCode")
try:
status_code = int(raw_status) if raw_status is not None else None
except (TypeError, ValueError):
status_code = None
returned_speaker = payload.get("speaker_id")
if status_code == 0 and isinstance(returned_speaker, str) and returned_speaker.strip():
return returned_speaker
status_message = base_response.get("StatusMessage")
message = str(status_message) if status_message not in (None, "") else "训练失败"
raise VoiceCloneProviderError(500, message)
payload_message = payload.get("message")
if payload_message not in (None, ""):
raise VoiceCloneProviderError(500, str(payload_message))
raise VoiceCloneProviderError(10156, "响应格式错误,缺少BaseResp字段")
@@ -0,0 +1,102 @@
from __future__ import annotations
from collections.abc import Mapping
from dataclasses import dataclass
from urllib.parse import urlsplit
import httpx
class VoicePrintIntegrationError(RuntimeError):
def __init__(self, code: int, message: str | None = None, params: tuple[object, ...] = ()):
super().__init__(message or str(code))
self.code = code
self.message = message
self.params = params
@dataclass(slots=True, frozen=True)
class VoicePrintEndpoint:
base_url: str
authorization: str
@classmethod
def parse(cls, configured_url: str | None) -> VoicePrintEndpoint:
if configured_url is None:
raise VoicePrintIntegrationError(10084)
parsed = urlsplit(configured_url)
if not parsed.scheme or not parsed.hostname:
raise VoicePrintIntegrationError(10084)
marker = "key="
marker_index = parsed.query.find(marker)
key = parsed.query[marker_index + len(marker) :] if marker_index >= 0 else parsed.query[3:]
port = f":{parsed.port}" if parsed.port is not None else ""
return cls(f"{parsed.scheme}://{parsed.hostname}{port}", f"Bearer {key}")
class VoicePrintClient:
def __init__(self, configured_url: str, *, timeout: float = 10.0, client: httpx.AsyncClient | None = None):
self.endpoint = VoicePrintEndpoint.parse(configured_url)
self.timeout = timeout
self._client = client
async def _request(
self,
method: str,
path: str,
*,
data: Mapping[str, str] | None = None,
files: Mapping[str, tuple[str, bytes, str]] | None = None,
) -> httpx.Response:
headers = {"Authorization": self.endpoint.authorization}
if self._client is not None:
return await self._client.request(
method, f"{self.endpoint.base_url}{path}", headers=headers, data=data, files=files
)
async with httpx.AsyncClient(timeout=self.timeout) as client:
return await client.request(
method, f"{self.endpoint.base_url}{path}", headers=headers, data=data, files=files
)
async def identify(self, speaker_ids: list[str], audio: bytes) -> tuple[str | None, float | None] | None:
if not speaker_ids:
return None
response = await self._request(
"POST",
"/voiceprint/identify",
data={"speaker_ids": ",".join(speaker_ids)},
files={"file": ("VoicePrint.WAV", audio, "application/octet-stream")},
)
if response.status_code != 200:
raise VoicePrintIntegrationError(10091)
try:
payload = response.json()
except ValueError as exc:
raise VoicePrintIntegrationError(10091) from exc
if not isinstance(payload, dict):
return None
speaker_id = payload.get("speaker_id")
score = payload.get("score")
return (
str(speaker_id) if speaker_id is not None else None,
float(score) if isinstance(score, int | float) else None,
)
async def register(self, speaker_id: str, audio: bytes) -> None:
response = await self._request(
"POST",
"/voiceprint/register",
data={"speaker_id": speaker_id},
files={"file": ("VoicePrint.WAV", audio, "application/octet-stream")},
)
if response.status_code != 200:
raise VoicePrintIntegrationError(10087)
if "true" not in response.text:
raise VoicePrintIntegrationError(10088)
async def cancel(self, speaker_id: str) -> None:
response = await self._request("DELETE", f"/voiceprint/{speaker_id}")
if response.status_code != 200:
raise VoicePrintIntegrationError(10089)
if "true" not in response.text:
raise VoicePrintIntegrationError(10090)
@@ -0,0 +1 @@
"""Single-instance background jobs for the manager API."""
@@ -0,0 +1,18 @@
from __future__ import annotations
from app.core.config import get_settings
from app.core.database import get_session_factory
from app.core.redis import distributed_lock
from app.repositories.knowledge import KnowledgeRepository
from app.services.knowledge import KnowledgeDocumentService
async def sync_running_knowledge_documents() -> int:
"""Run one document-status pass under a cross-process Redis lock."""
settings = get_settings()
async with distributed_lock("jobs:knowledge-document-status", settings.job_lock_ttl_seconds) as acquired:
if not acquired:
return 0
async with get_session_factory()() as session:
return await KnowledgeDocumentService(KnowledgeRepository(session)).sync_running()
+135
View File
@@ -0,0 +1,135 @@
from __future__ import annotations
import asyncio
import logging
import os
import signal
import time
from collections.abc import Awaitable, Callable
from contextlib import suppress
from app.core.config import get_settings
from app.core.database import configure_database, database_ping, dispose_database
from app.core.redis import close_redis, redis_ping
from app.jobs.tasks import sync_running_knowledge_documents
from app.services.agent import redact_legacy_agent_snapshots
logger = logging.getLogger(__name__)
async def _wait_or_stop(stop: asyncio.Event, seconds: float) -> bool:
try:
await asyncio.wait_for(stop.wait(), timeout=seconds)
except asyncio.TimeoutError:
return False
return True
async def _fixed_delay_loop(
stop: asyncio.Event,
operation: Callable[[], Awaitable[int]],
*,
name: str,
initial_delay: float,
delay: float,
) -> None:
if initial_delay and await _wait_or_stop(stop, initial_delay):
return
while not stop.is_set():
started = time.monotonic()
try:
changed = await operation()
logger.info(
"Background job %s completed changed=%s duration_ms=%d",
name,
changed,
(time.monotonic() - started) * 1000,
)
except Exception:
logger.exception("Background job %s failed; the next fixed-delay pass will retry", name)
if await _wait_or_stop(stop, delay):
return
async def _finish_tasks(tasks: list[asyncio.Task[None]], timeout: float) -> bool:
"""Let active jobs finish, then cancel only those exceeding the shutdown budget."""
_, pending = await asyncio.wait(tasks, timeout=max(timeout, 0.0))
if not pending:
return True
logger.warning(
"Graceful job shutdown timed out after %.1f seconds; cancelling %d task(s)",
timeout,
len(pending),
)
for task in pending:
task.cancel()
for task in pending:
with suppress(asyncio.CancelledError):
await task
return False
async def run_worker(stop: asyncio.Event | None = None) -> None:
settings = get_settings()
os.environ["TZ"] = settings.timezone
if hasattr(time, "tzset"):
time.tzset()
logging.basicConfig(level=settings.log_level, format="%(asctime)s %(levelname)s %(name)s %(message)s")
configure_database(settings)
selected_stop = stop or asyncio.Event()
if not settings.allow_start_without_dependencies:
if not await database_ping():
raise RuntimeError("database readiness check failed")
if not await redis_ping():
raise RuntimeError("Redis readiness check failed")
# The retained Java service performs a blocking startup redaction pass,
# then compensates rolling-deployment writes after 5 seconds and every
# 15 seconds. The standalone worker preserves those timings while the
# Redis lock makes multiple worker replicas safe.
await redact_legacy_agent_snapshots()
tasks = [
asyncio.create_task(
_fixed_delay_loop(
selected_stop,
redact_legacy_agent_snapshots,
name="agent-snapshot-redaction",
initial_delay=5,
delay=15,
)
),
asyncio.create_task(
_fixed_delay_loop(
selected_stop,
sync_running_knowledge_documents,
name="knowledge-document-status",
initial_delay=0,
delay=30,
)
),
]
try:
await selected_stop.wait()
finally:
await _finish_tasks(tasks, settings.graceful_shutdown_seconds)
await close_redis()
await dispose_database()
def main() -> None:
stop = asyncio.Event()
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
for signal_name in (signal.SIGINT, signal.SIGTERM):
with suppress(NotImplementedError):
loop.add_signal_handler(signal_name, stop.set)
try:
loop.run_until_complete(run_worker(stop))
finally:
loop.close()
if __name__ == "__main__":
main()
+217
View File
@@ -0,0 +1,217 @@
from __future__ import annotations
import logging
import os
import time
from collections.abc import AsyncIterator, Mapping, Sequence
from contextlib import asynccontextmanager
from typing import Any
from fastapi import FastAPI, Request
from fastapi.exceptions import RequestValidationError
from sqlalchemy.exc import IntegrityError
from starlette.exceptions import HTTPException as StarletteHTTPException
from starlette.middleware.cors import CORSMiddleware
from app.core.config import get_settings
from app.core.database import configure_database, database_ping, dispose_database
from app.core.errors import AppError, ErrorCode
from app.core.i18n import message_for
from app.core.redis import close_redis, redis_ping
from app.core.responses import JavaJSONResponse, error_response, ok
from app.core.security import AuthenticationMiddleware
from app.routers import application_routers
logger = logging.getLogger(__name__)
settings = get_settings()
_MULTIPART_VALIDATION_PATHS = {
"/datasets/{dataset_id}/documents",
"/otaMag/upload",
"/otaMag/uploadAssetsBin",
"/voiceClone/upload",
}
def _matches_path_template(path: str, template: str) -> bool:
path_parts = path.removeprefix(settings.context_path).strip("/").split("/")
template_parts = template.strip("/").split("/")
return len(path_parts) == len(template_parts) and all(
expected.startswith("{") and expected.endswith("}") or actual == expected
for actual, expected in zip(path_parts, template_parts, strict=True)
)
def _java_required_message(request: Request, errors: Sequence[Mapping[str, Any]]) -> str | None:
path = request.url.path.removeprefix(settings.context_path)
mappings = (
(
"/admin/server/emit-action",
(("action", "操作不能为空"), ("targetWs", "目标ws地址不能为空")),
),
("/agent", (("agentName", "智能体名称不能为空"),)),
(
"/agent/chat-history/report",
tuple((field, "不能为空") for field in ("macAddress", "sessionId", "chatType", "content")),
),
(
"/agent/{agentId}/snapshots/{snapshotId}/restore",
(("currentStateToken", "不能为空"),),
),
(
"/config/agent-models",
(
("macAddress", "设备MAC地址不能为空"),
("clientId", "客户端ID不能为空"),
("selectedModule", "客户端已实例化的模型不能为空"),
),
),
("/config/correct-words", (("macAddress", "设备MAC地址不能为空"),)),
(
"/device/address-book/alias",
(("targetMac", "目标MAC地址不能为空"), ("macAddress", "MAC地址不能为空")),
),
)
missing_fields: set[str] = set()
for error in errors:
location = tuple(error.get("loc", ()))
if error.get("type") == "missing" and location[:1] == ("body",):
missing_fields.add(str(location[-1]))
if not missing_fields:
return None
for template, fields in mappings:
if _matches_path_template(path, template):
return next((message for field, message in fields if field in missing_fields), None)
return None
@asynccontextmanager
async def lifespan(_: FastAPI) -> AsyncIterator[None]:
os.environ["TZ"] = settings.timezone
if hasattr(time, "tzset"):
time.tzset()
settings.upload_dir.mkdir(parents=True, exist_ok=True)
configure_database(settings)
if not settings.i18n_dir.exists():
raise RuntimeError(f"Java i18n resources are missing: {settings.i18n_dir}")
if not settings.changelog_path.exists():
raise RuntimeError(f"Liquibase source of truth is missing: {settings.changelog_path}")
if not settings.allow_start_without_dependencies:
if not await database_ping():
raise RuntimeError("database readiness check failed")
if not await redis_ping():
raise RuntimeError("Redis readiness check failed")
yield
await close_redis()
await dispose_database()
app = FastAPI(
title="xiaozhi-manager-api",
version="0.1.0",
docs_url=f"{settings.context_path}/doc.html",
openapi_url=f"{settings.context_path}/v3/api-docs",
redoc_url=None,
default_response_class=JavaJSONResponse,
lifespan=lifespan,
)
app.add_middleware(AuthenticationMiddleware)
app.add_middleware(
CORSMiddleware,
allow_origins=[],
allow_origin_regex=".*",
allow_credentials=True,
allow_methods=["GET", "POST", "PUT", "DELETE", "OPTIONS"],
allow_headers=["*"],
max_age=3600,
)
for router in application_routers():
app.include_router(router, prefix=settings.context_path)
@app.get(f"{settings.context_path}/health", include_in_schema=False)
async def health() -> JavaJSONResponse:
return ok({"status": "UP"})
@app.get(f"{settings.context_path}/health/live", include_in_schema=False)
async def liveness() -> JavaJSONResponse:
return ok({"status": "UP"})
def upload_storage_ready() -> bool:
"""Report whether the non-root API process can traverse and write its upload mount."""
try:
return settings.upload_dir.is_dir() and os.access(
settings.upload_dir,
os.W_OK | os.X_OK,
)
except OSError:
return False
@app.get(f"{settings.context_path}/health/ready", include_in_schema=False)
async def readiness() -> JavaJSONResponse:
database, redis, uploads = await database_ping(), await redis_ping(), upload_storage_ready()
code = 0 if database and redis and uploads else 503
msg = "success" if code == 0 else "dependencies unavailable"
return JavaJSONResponse(
{
"code": code,
"msg": msg,
"data": {"database": database, "redis": redis, "uploads": uploads},
},
status_code=200 if code == 0 else 503,
)
@app.exception_handler(AppError)
async def app_error_handler(request: Request, exc: AppError) -> JavaJSONResponse:
return error_response(request, exc.code, exc.message, params=exc.params)
@app.exception_handler(RequestValidationError)
async def validation_error_handler(request: Request, exc: RequestValidationError) -> JavaJSONResponse:
errors = exc.errors()
# Spring only maps MethodArgumentNotValidException (a deserialized JSON
# object's @Valid field constraints) to code 10034. Root-body conversion,
# missing query parameters and multipart binding failures reach its generic
# exception handler and therefore keep the HTTP-200/code-500 envelope.
root_body_error = any(tuple(error.get("loc", ())) == ("body",) for error in errors)
missing_query = any(
error.get("type") == "missing" and tuple(error.get("loc", ()))[:1] == ("query",)
for error in errors
)
multipart_binding_error = any(
error.get("type") == "missing"
and tuple(error.get("loc", ()))[:1] == ("body",)
and any(_matches_path_template(request.url.path, path) for path in _MULTIPART_VALIDATION_PATHS)
for error in errors
)
if root_body_error or missing_query or multipart_binding_error:
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR)
first = errors[0] if errors else None
detail = _java_required_message(request, errors) or (str(first.get("msg")) if first else None)
return error_response(request, ErrorCode.PARAM_VALUE_NULL, detail)
@app.exception_handler(IntegrityError)
async def integrity_error_handler(request: Request, _: IntegrityError) -> JavaJSONResponse:
return error_response(request, ErrorCode.DB_RECORD_EXISTS)
@app.exception_handler(StarletteHTTPException)
async def http_error_handler(request: Request, exc: StarletteHTTPException) -> JavaJSONResponse:
if exc.status_code == 404:
not_found = message_for(ErrorCode.RESOURCE_NOT_FOUND, request.headers.get("Accept-Language"))
return error_response(request, 404, not_found)
return error_response(request, exc.status_code, str(exc.detail))
@app.exception_handler(Exception)
async def unhandled_error_handler(request: Request, exc: Exception) -> JavaJSONResponse:
logger.exception("Unhandled manager-api error", exc_info=exc)
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR)
@@ -0,0 +1,665 @@
from __future__ import annotations
# All interpolated SQL fragments are selected from closed column/table allowlists.
# ruff: noqa: S608
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import Any
from sqlalchemy import bindparam, text
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import Repository
AGENT_COLUMNS = (
"id",
"user_id",
"agent_code",
"agent_name",
"asr_model_id",
"vad_model_id",
"llm_model_id",
"slm_model_id",
"vllm_model_id",
"tts_model_id",
"tts_voice_id",
"tts_language",
"tts_volume",
"tts_rate",
"tts_pitch",
"mem_model_id",
"intent_model_id",
"chat_history_conf",
"system_prompt",
"summary_memory",
"lang_code",
"language",
"sort",
"creator",
"created_at",
"updater",
"updated_at",
)
AGENT_MUTABLE_COLUMNS = frozenset(AGENT_COLUMNS) - {"id", "user_id", "creator", "created_at"}
TEMPLATE_COLUMNS = (
"id",
"agent_code",
"agent_name",
"asr_model_id",
"vad_model_id",
"llm_model_id",
"vllm_model_id",
"tts_model_id",
"tts_voice_id",
"tts_language",
"tts_volume",
"tts_rate",
"tts_pitch",
"mem_model_id",
"intent_model_id",
"chat_history_conf",
"system_prompt",
"summary_memory",
"lang_code",
"language",
"sort",
"creator",
"created_at",
"updater",
"updated_at",
)
class AgentRepository(Repository):
def __init__(self, session: AsyncSession):
super().__init__(session)
@property
def is_sqlite(self) -> bool:
bind = self.session.get_bind()
return bool(bind is not None and bind.dialect.name == "sqlite")
async def get_agent(self, agent_id: str, *, for_update: bool = False) -> dict[str, Any] | None:
suffix = "" if self.is_sqlite or not for_update else " FOR UPDATE"
return await self.fetch_one(
f"SELECT {', '.join(AGENT_COLUMNS)} FROM ai_agent WHERE id=:id{suffix}", {"id": agent_id}
)
async def check_agent_owner(self, agent_id: str, user_id: int, *, super_admin: bool) -> bool:
if super_admin:
return bool(await self.scalar("SELECT 1 FROM ai_agent WHERE id=:id LIMIT 1", {"id": agent_id}))
return bool(
await self.scalar(
"SELECT 1 FROM ai_agent WHERE id=:id AND user_id=:user_id LIMIT 1",
{"id": agent_id, "user_id": user_id},
)
)
async def list_user_agents(self, user_id: int, keyword: str | None) -> list[dict[str, Any]]:
params: dict[str, Any] = {"user_id": user_id}
where = "a.user_id=:user_id"
if keyword is not None and keyword.strip():
params["keyword"] = f"%{keyword}%"
where += (
" AND (a.agent_name LIKE :keyword"
" OR EXISTS (SELECT 1 FROM ai_device d0 WHERE d0.agent_id=a.id"
" AND d0.user_id=:user_id AND d0.mac_address LIKE :keyword)"
" OR EXISTS (SELECT 1 FROM ai_agent_tag_relation tr0"
" JOIN ai_agent_tag t0 ON t0.id=tr0.tag_id"
" WHERE tr0.agent_id=a.id AND t0.deleted=0 AND t0.tag_name LIKE :keyword))"
)
return await self.fetch_all(
"SELECT a.*, mt.model_name AS tts_model_name, ml.model_name AS llm_model_name,"
" mv.model_name AS vllm_model_name, COALESCE(tv.name, vc.name) AS tts_voice_name,"
" (SELECT MAX(d.last_connected_at) FROM ai_device d WHERE d.agent_id=a.id) AS last_connected_at,"
" (SELECT COUNT(*) FROM ai_device d WHERE d.agent_id=a.id) AS device_count"
" FROM ai_agent a"
" LEFT JOIN ai_model_config mt ON mt.id=a.tts_model_id"
" LEFT JOIN ai_model_config ml ON ml.id=a.llm_model_id"
" LEFT JOIN ai_model_config mv ON mv.id=a.vllm_model_id"
" LEFT JOIN ai_tts_voice tv ON tv.id=a.tts_voice_id"
" LEFT JOIN ai_voice_clone vc ON vc.id=a.tts_voice_id"
f" WHERE {where} ORDER BY a.created_at DESC",
params,
)
async def list_admin_agents(
self, page: int, limit: int, order_field: str, ascending: bool
) -> tuple[list[dict[str, Any]], int]:
allowed = {"agent_name", "created_at", "updated_at", "sort", "id"}
selected = order_field if order_field in allowed else "agent_name"
direction = "ASC" if ascending else "DESC"
total = int(await self.scalar("SELECT COUNT(*) FROM ai_agent") or 0)
query = (
f"SELECT {', '.join(AGENT_COLUMNS)} FROM ai_agent "
f"ORDER BY {selected} {direction} LIMIT :limit OFFSET :offset"
)
rows = await self.fetch_all(
query,
{"limit": limit, "offset": (page - 1) * limit},
)
return rows, total
async def insert_agent(self, values: Mapping[str, Any]) -> int:
columns = [column for column in AGENT_COLUMNS if column in values]
placeholders = ", ".join(f":{column}" for column in columns)
return await self.execute(
f"INSERT INTO ai_agent ({', '.join(columns)}) VALUES ({placeholders})",
{column: values[column] for column in columns},
)
async def update_agent(self, agent_id: str, values: Mapping[str, Any], *, include_null: bool = False) -> int:
selected = {
key: value
for key, value in values.items()
if key in AGENT_MUTABLE_COLUMNS and (include_null or value is not None)
}
if not selected:
return 0
assignments = ", ".join(f"{column}=:{column}" for column in selected)
return await self.execute(
f"UPDATE ai_agent SET {assignments} WHERE id=:agent_id",
{**selected, "agent_id": agent_id},
)
async def get_agent_plugins(self, agent_id: str) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT m.id,m.agent_id,m.plugin_id,m.param_info,p.provider_code"
" FROM ai_agent_plugin_mapping m LEFT JOIN ai_model_provider p ON p.id=m.plugin_id"
" WHERE m.agent_id=:agent_id ORDER BY m.id ASC",
{"agent_id": agent_id},
)
async def replace_plugins(self, agent_id: str, plugins: Sequence[Mapping[str, Any]]) -> None:
existing = await self.fetch_all(
"SELECT id,plugin_id FROM ai_agent_plugin_mapping WHERE agent_id=:agent_id",
{"agent_id": agent_id},
)
by_plugin = {str(row["plugin_id"]): int(row["id"]) for row in existing}
incoming = {str(item.get("plugin_id") or "") for item in plugins}
remove_ids = [int(row["id"]) for row in existing if str(row["plugin_id"]) not in incoming]
if remove_ids:
statement = text("DELETE FROM ai_agent_plugin_mapping WHERE id IN :ids").bindparams(
bindparam("ids", expanding=True)
)
await self.execute(statement, {"ids": remove_ids})
for item in plugins:
plugin_id = str(item.get("plugin_id") or "")
params = {"agent_id": agent_id, "plugin_id": plugin_id, "param_info": item.get("param_info") or "{}"}
if plugin_id in by_plugin:
await self.execute(
"UPDATE ai_agent_plugin_mapping SET param_info=:param_info WHERE id=:id",
{"id": by_plugin[plugin_id], **params},
)
else:
await self.execute(
"INSERT INTO ai_agent_plugin_mapping (id,agent_id,plugin_id,param_info)"
" VALUES (:id,:agent_id,:plugin_id,:param_info)",
{"id": int(item["id"]), **params},
)
async def delete_plugins(self, agent_id: str) -> int:
return await self.execute("DELETE FROM ai_agent_plugin_mapping WHERE agent_id=:id", {"id": agent_id})
async def get_context_provider(self, agent_id: str) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id,agent_id,context_providers,creator,created_at,updater,updated_at"
" FROM ai_agent_context_provider WHERE agent_id=:id LIMIT 1",
{"id": agent_id},
)
async def upsert_context_provider(self, agent_id: str, encoded: str, new_id: str) -> None:
existing = await self.get_context_provider(agent_id)
if existing:
await self.execute(
"UPDATE ai_agent_context_provider SET context_providers=:value WHERE id=:id",
{"value": encoded, "id": existing["id"]},
)
else:
await self.execute(
"INSERT INTO ai_agent_context_provider (id,agent_id,context_providers)"
" VALUES (:id,:agent_id,:value)",
{"id": new_id, "agent_id": agent_id, "value": encoded},
)
async def get_correct_word_ids(self, agent_id: str) -> list[str]:
rows = await self.fetch_all(
"SELECT file_id FROM ai_agent_correct_word_mapping WHERE agent_id=:id", {"id": agent_id}
)
return [str(row["file_id"]) for row in rows]
async def replace_correct_words(
self, agent_id: str, file_ids: Sequence[str], user_id: int, now: datetime, ids: Sequence[str]
) -> None:
await self.execute("DELETE FROM ai_agent_correct_word_mapping WHERE agent_id=:id", {"id": agent_id})
await self.execute_many(
"INSERT INTO ai_agent_correct_word_mapping"
" (id,agent_id,file_id,creator,created_at,updater,updated_at)"
" VALUES (:id,:agent_id,:file_id,:user_id,:now,:user_id,:now)",
[
{"id": mapping_id, "agent_id": agent_id, "file_id": file_id, "user_id": user_id, "now": now}
for mapping_id, file_id in zip(ids, file_ids, strict=True)
],
)
async def get_agent_tags(self, agent_id: str) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT t.id,t.tag_name,t.sort,r.sort AS relation_sort"
" FROM ai_agent_tag t JOIN ai_agent_tag_relation r ON t.id=r.tag_id"
" WHERE r.agent_id=:id AND t.deleted=0 ORDER BY r.sort ASC,r.created_at ASC",
{"id": agent_id},
)
async def get_tags_for_agents(self, agent_ids: Sequence[str]) -> list[dict[str, Any]]:
if not agent_ids:
return []
statement = text(
"SELECT t.id,t.tag_name,r.agent_id,r.sort AS relation_sort"
" FROM ai_agent_tag t JOIN ai_agent_tag_relation r ON t.id=r.tag_id"
" WHERE r.agent_id IN :ids AND t.deleted=0 ORDER BY r.sort ASC,r.created_at ASC"
).bindparams(bindparam("ids", expanding=True))
return await self.fetch_all(statement, {"ids": list(agent_ids)})
async def list_tags(self) -> list[dict[str, Any]]:
return await self.fetch_all("SELECT id,tag_name,sort FROM ai_agent_tag WHERE deleted=0 ORDER BY sort ASC")
async def get_tag(self, tag_id: str) -> dict[str, Any] | None:
return await self.fetch_one("SELECT * FROM ai_agent_tag WHERE id=:id", {"id": tag_id})
async def find_active_tag_by_name(self, tag_name: str) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT * FROM ai_agent_tag WHERE tag_name=:name AND deleted=0 LIMIT 1", {"name": tag_name}
)
async def find_any_tag_by_name(self, tag_name: str) -> dict[str, Any] | None:
return await self.fetch_one("SELECT * FROM ai_agent_tag WHERE tag_name=:name LIMIT 1", {"name": tag_name})
async def insert_tag(self, values: Mapping[str, Any]) -> int:
return await self.execute(
"INSERT INTO ai_agent_tag"
" (id,tag_name,sort,deleted,creator,created_at,updater,updated_at)"
" VALUES (:id,:tag_name,:sort,:deleted,:creator,:created_at,:updater,:updated_at)",
values,
)
async def soft_delete_tag(self, tag_id: str, now: datetime) -> int:
return await self.execute(
"UPDATE ai_agent_tag SET deleted=1,updated_at=:now WHERE id=:id",
{"id": tag_id, "now": now},
)
async def replace_tag_relations(self, agent_id: str, relations: Sequence[Mapping[str, Any]]) -> None:
await self.execute("DELETE FROM ai_agent_tag_relation WHERE agent_id=:id", {"id": agent_id})
await self.execute_many(
"INSERT INTO ai_agent_tag_relation"
" (id,agent_id,tag_id,sort,creator,created_at,updater,updated_at)"
" VALUES (:id,:agent_id,:tag_id,:sort,:creator,:created_at,:updater,:updated_at)",
relations,
)
async def get_model_config(self, model_id: str | None) -> dict[str, Any] | None:
if not model_id:
return None
return await self.fetch_one("SELECT * FROM ai_model_config WHERE id=:id", {"id": model_id})
async def get_default_llm_config(self) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT * FROM ai_model_config WHERE model_type='LLM' AND is_enabled=1"
" ORDER BY is_default DESC,sort ASC LIMIT 1"
)
async def get_model_provider(self, provider_id: str) -> dict[str, Any] | None:
return await self.fetch_one("SELECT * FROM ai_model_provider WHERE id=:id", {"id": provider_id})
async def get_timbre(self, timbre_id: str | None) -> dict[str, Any] | None:
if not timbre_id:
return None
row = await self.fetch_one("SELECT * FROM ai_tts_voice WHERE id=:id", {"id": timbre_id})
if row is None:
row = await self.fetch_one("SELECT * FROM ai_voice_clone WHERE id=:id", {"id": timbre_id})
return row
async def find_timbre_by_voice_code(self, model_id: str, voice_code: str) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT * FROM ai_tts_voice WHERE tts_model_id=:model_id AND tts_voice=:voice LIMIT 1",
{"model_id": model_id, "voice": voice_code},
)
async def get_device_by_mac(self, mac_address: str) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT * FROM ai_device WHERE mac_address=:mac ORDER BY id DESC LIMIT 1", {"mac": mac_address}
)
async def get_agent_by_device_mac(self, mac_address: str) -> dict[str, Any] | None:
return await self.fetch_one(
f"SELECT {', '.join('a.' + column for column in AGENT_COLUMNS)}"
" FROM ai_device d LEFT JOIN ai_agent a ON d.agent_id=a.id"
" WHERE d.mac_address=:mac ORDER BY d.id DESC LIMIT 1",
{"mac": mac_address},
)
async def update_device_connection(self, device_id: str, now: datetime) -> int:
return await self.execute(
"UPDATE ai_device SET last_connected_at=:now WHERE id=:id", {"id": device_id, "now": now}
)
async def insert_chat_audio(self, audio_id: str, audio: bytes) -> int:
return await self.execute(
"INSERT INTO ai_agent_chat_audio (id,audio) VALUES (:id,:audio)", {"id": audio_id, "audio": audio}
)
async def get_chat_audio(self, audio_id: str) -> bytes | None:
value = await self.scalar("SELECT audio FROM ai_agent_chat_audio WHERE id=:id", {"id": audio_id})
return bytes(value) if value is not None else None
async def insert_chat_history(self, values: Mapping[str, Any]) -> int:
return await self.execute(
"INSERT INTO ai_agent_chat_history"
" (mac_address,agent_id,session_id,chat_type,content,audio_id,created_at)"
" VALUES (:mac_address,:agent_id,:session_id,:chat_type,:content,:audio_id,:created_at)",
values,
)
async def get_session_agent_id(self, session_id: str) -> str | None:
value = await self.scalar(
"SELECT agent_id FROM ai_agent_chat_history WHERE session_id=:id LIMIT 1", {"id": session_id}
)
return str(value) if value is not None else None
async def get_audio_agent_id(self, audio_id: str) -> str | None:
value = await self.scalar(
"SELECT agent_id FROM ai_agent_chat_history WHERE audio_id=:id LIMIT 1", {"id": audio_id}
)
return str(value) if value is not None else None
async def is_audio_owned(self, audio_id: str, agent_id: str) -> bool:
count = await self.scalar(
"SELECT COUNT(*) FROM ai_agent_chat_history WHERE audio_id=:audio_id AND agent_id=:agent_id",
{"audio_id": audio_id, "agent_id": agent_id},
)
return int(count or 0) == 1
async def get_audio_content(self, audio_id: str) -> str | None:
value = await self.scalar(
"SELECT content FROM ai_agent_chat_history WHERE audio_id=:id LIMIT 1", {"id": audio_id}
)
return str(value) if value is not None else None
async def get_chat_history(self, agent_id: str, session_id: str) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT created_at,chat_type,content,audio_id,mac_address"
" FROM ai_agent_chat_history WHERE agent_id=:agent_id AND session_id=:session_id"
" ORDER BY created_at ASC",
{"agent_id": agent_id, "session_id": session_id},
)
async def get_recent_user_history(self, agent_id: str) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT content,audio_id FROM ai_agent_chat_history"
" WHERE agent_id=:id AND chat_type=1 AND audio_id IS NOT NULL ORDER BY id DESC LIMIT 50",
{"id": agent_id},
)
async def list_sessions(self, agent_id: str, page: int, limit: int) -> tuple[list[dict[str, Any]], int]:
total = int(
await self.scalar(
"SELECT COUNT(*) FROM (SELECT session_id FROM ai_agent_chat_history"
" WHERE agent_id=:id GROUP BY session_id) sessions",
{"id": agent_id},
)
or 0
)
rows = await self.fetch_all(
"SELECT h.session_id,MAX(h.created_at) AS created_at,COUNT(*) AS chat_count,"
" (SELECT t.title FROM ai_agent_chat_title t WHERE t.session_id=h.session_id LIMIT 1) AS title"
" FROM ai_agent_chat_history h WHERE h.agent_id=:id GROUP BY h.session_id"
" ORDER BY created_at DESC LIMIT :limit OFFSET :offset",
{"id": agent_id, "limit": limit, "offset": (page - 1) * limit},
)
return rows, total
async def upsert_chat_title(self, session_id: str, title: str, now: datetime, title_id: str) -> None:
existing = await self.fetch_one(
"SELECT id FROM ai_agent_chat_title WHERE session_id=:session_id LIMIT 1", {"session_id": session_id}
)
if existing:
await self.execute(
"UPDATE ai_agent_chat_title SET title=:title,updated_at=:now WHERE id=:id",
{"id": existing["id"], "title": title, "now": now},
)
else:
await self.execute(
"INSERT INTO ai_agent_chat_title (id,session_id,title,created_at,updated_at)"
" VALUES (:id,:session_id,:title,:now,:now)",
{"id": title_id, "session_id": session_id, "title": title, "now": now},
)
async def delete_chat_history(self, agent_id: str, *, delete_audio: bool, delete_text: bool) -> None:
if delete_audio:
ids = await self.fetch_all(
"SELECT DISTINCT audio_id FROM ai_agent_chat_history WHERE agent_id=:id AND audio_id IS NOT NULL",
{"id": agent_id},
)
audio_ids = [str(row["audio_id"]) for row in ids]
for offset in range(0, len(audio_ids), 1000):
batch = audio_ids[offset : offset + 1000]
statement = text("DELETE FROM ai_agent_chat_audio WHERE id IN :ids").bindparams(
bindparam("ids", expanding=True)
)
await self.execute(statement, {"ids": batch})
if delete_audio and not delete_text:
await self.execute("UPDATE ai_agent_chat_history SET audio_id=NULL WHERE agent_id=:id", {"id": agent_id})
if delete_text:
await self.execute("DELETE FROM ai_agent_chat_history WHERE agent_id=:id", {"id": agent_id})
async def delete_agent_cascade(self, agent_id: str) -> None:
devices = await self.fetch_all("SELECT mac_address FROM ai_device WHERE agent_id=:id", {"id": agent_id})
macs = [str(row["mac_address"]) for row in devices if row.get("mac_address") is not None]
await self.execute("DELETE FROM ai_device WHERE agent_id=:id", {"id": agent_id})
if macs:
statement = text(
"DELETE FROM ai_device_address_book WHERE mac_address IN :macs OR target_mac IN :targets"
).bindparams(bindparam("macs", expanding=True), bindparam("targets", expanding=True))
await self.execute(statement, {"macs": macs, "targets": macs})
await self.delete_chat_history(agent_id, delete_audio=True, delete_text=True)
for table in (
"ai_agent_plugin_mapping",
"ai_agent_context_provider",
"ai_agent_correct_word_mapping",
"ai_agent_tag_relation",
"ai_agent_snapshot",
):
await self.execute(f"DELETE FROM {table} WHERE agent_id=:id", {"id": agent_id})
await self.execute("DELETE FROM ai_agent WHERE id=:id", {"id": agent_id})
async def list_templates(
self, *, name: str | None = None, page: int | None = None, limit: int | None = None
) -> tuple[list[dict[str, Any]], int]:
params: dict[str, Any] = {}
where = ""
if name:
where = " WHERE agent_name LIKE :name"
params["name"] = f"%{name}%"
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_agent_template{where}", params) or 0)
paging = ""
if page is not None and limit is not None:
params.update(limit=limit, offset=(page - 1) * limit)
paging = " LIMIT :limit OFFSET :offset"
rows = await self.fetch_all(f"SELECT * FROM ai_agent_template{where} ORDER BY sort ASC{paging}", params)
return rows, total
async def get_template(self, template_id: str) -> dict[str, Any] | None:
return await self.fetch_one("SELECT * FROM ai_agent_template WHERE id=:id", {"id": template_id})
async def get_default_template(self) -> dict[str, Any] | None:
return await self.fetch_one("SELECT * FROM ai_agent_template ORDER BY sort ASC LIMIT 1")
async def next_template_sort(self) -> int:
rows = await self.fetch_all("SELECT sort FROM ai_agent_template WHERE sort IS NOT NULL ORDER BY sort ASC")
expected = 1
for row in rows:
value = int(row["sort"])
if value > expected:
return expected
expected = value + 1
return expected
async def insert_template(self, values: Mapping[str, Any]) -> int:
columns = [column for column in TEMPLATE_COLUMNS if column in values]
return await self.execute(
f"INSERT INTO ai_agent_template ({', '.join(columns)})"
f" VALUES ({', '.join(':' + column for column in columns)})",
{column: values[column] for column in columns},
)
async def update_template(self, template_id: str, values: Mapping[str, Any]) -> int:
selected = {
key: value for key, value in values.items() if key in TEMPLATE_COLUMNS and key != "id" and value is not None
}
if not selected:
return 0
return await self.execute(
f"UPDATE ai_agent_template SET {', '.join(key + '=:' + key for key in selected)} WHERE id=:id",
{**selected, "id": template_id},
)
async def delete_template(self, template_id: str) -> int:
return await self.execute("DELETE FROM ai_agent_template WHERE id=:id", {"id": template_id})
async def reorder_templates(self, deleted_sort: int) -> int:
return await self.execute("UPDATE ai_agent_template SET sort=sort-1 WHERE sort>:sort", {"sort": deleted_sort})
async def delete_templates(self, ids: Sequence[str]) -> int:
if not ids:
return 0
statement = text("DELETE FROM ai_agent_template WHERE id IN :ids").bindparams(bindparam("ids", expanding=True))
return await self.execute(statement, {"ids": list(ids)})
async def list_voiceprints(self, agent_id: str, user_id: int) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT id,audio_id,source_name,introduce,create_date"
" FROM ai_agent_voice_print WHERE agent_id=:agent_id AND creator=:user_id",
{"agent_id": agent_id, "user_id": user_id},
)
async def list_voiceprint_ids(self, agent_id: str) -> list[str]:
rows = await self.fetch_all("SELECT id FROM ai_agent_voice_print WHERE agent_id=:id", {"id": agent_id})
return [str(row["id"]) for row in rows]
async def get_voiceprint(self, voiceprint_id: str, user_id: int | None = None) -> dict[str, Any] | None:
where = "id=:id"
params: dict[str, Any] = {"id": voiceprint_id}
if user_id is not None:
where += " AND creator=:user_id"
params["user_id"] = user_id
return await self.fetch_one(f"SELECT * FROM ai_agent_voice_print WHERE {where} LIMIT 1", params)
async def insert_voiceprint(self, values: Mapping[str, Any]) -> int:
return await self.execute(
"INSERT INTO ai_agent_voice_print"
" (id,agent_id,audio_id,source_name,introduce,creator,create_date,updater,update_date)"
" VALUES (:id,:agent_id,:audio_id,:source_name,:introduce,:creator,:create_date,:updater,:update_date)",
values,
)
async def update_voiceprint(self, voiceprint_id: str, user_id: int, values: Mapping[str, Any]) -> int:
allowed = {"audio_id", "source_name", "introduce", "updater", "update_date"}
selected = {key: value for key, value in values.items() if key in allowed and value is not None}
if not selected:
return 0
return await self.execute(
f"UPDATE ai_agent_voice_print SET {', '.join(key + '=:' + key for key in selected)}"
" WHERE id=:id AND creator=:user_id",
{**selected, "id": voiceprint_id, "user_id": user_id},
)
async def delete_voiceprint(self, voiceprint_id: str, user_id: int) -> int:
return await self.execute(
"DELETE FROM ai_agent_voice_print WHERE id=:id AND creator=:user_id",
{"id": voiceprint_id, "user_id": user_id},
)
async def snapshot_max_version(self, agent_id: str) -> int:
return int(
await self.scalar(
"SELECT COALESCE(MAX(version_no),0) FROM ai_agent_snapshot WHERE agent_id=:id", {"id": agent_id}
)
or 0
)
async def latest_snapshot(self, agent_id: str) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT * FROM ai_agent_snapshot WHERE agent_id=:id ORDER BY version_no DESC LIMIT 1", {"id": agent_id}
)
async def next_snapshot(self, agent_id: str, version_no: int) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT * FROM ai_agent_snapshot WHERE agent_id=:id AND version_no>:version"
" ORDER BY version_no ASC LIMIT 1",
{"id": agent_id, "version": version_no},
)
async def get_snapshot(self, snapshot_id: str) -> dict[str, Any] | None:
return await self.fetch_one("SELECT * FROM ai_agent_snapshot WHERE id=:id", {"id": snapshot_id})
async def list_snapshots(
self, agent_id: str, page: int, limit: int, max_version_no: int | None
) -> tuple[list[dict[str, Any]], int]:
params: dict[str, Any] = {"id": agent_id, "limit": limit, "offset": (page - 1) * limit}
extra = ""
if max_version_no is not None:
extra = " AND version_no<=:max_version"
params["max_version"] = max_version_no
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_agent_snapshot WHERE agent_id=:id{extra}", params) or 0)
rows = await self.fetch_all(
f"SELECT * FROM ai_agent_snapshot WHERE agent_id=:id{extra}"
" ORDER BY version_no DESC LIMIT :limit OFFSET :offset",
params,
)
return rows, total
async def insert_snapshot_next_version(self, values: Mapping[str, Any]) -> int:
return await self.execute(
"INSERT INTO ai_agent_snapshot"
" (id,agent_id,user_id,version_no,snapshot_data,changed_fields,source,"
" restore_from_snapshot_id,restore_from_version_no,creator,created_at,redaction_version)"
" SELECT :id,:agent_id,:user_id,COALESCE(MAX(version_no),0)+1,:snapshot_data,:changed_fields,:source,"
" :restore_from_snapshot_id,:restore_from_version_no,:creator,:created_at,:redaction_version"
" FROM ai_agent_snapshot WHERE agent_id=:agent_id",
values,
)
async def prune_snapshots(self, agent_id: str, keep: int) -> int:
rows = await self.fetch_all(
"SELECT id FROM ai_agent_snapshot WHERE agent_id=:id ORDER BY version_no DESC LIMIT :keep",
{"id": agent_id, "keep": keep},
)
retained = [str(row["id"]) for row in rows]
if not retained:
return 0
statement = text("DELETE FROM ai_agent_snapshot WHERE agent_id=:agent_id AND id NOT IN :retained").bindparams(
bindparam("retained", expanding=True)
)
return await self.execute(statement, {"agent_id": agent_id, "retained": retained})
async def delete_snapshot(self, snapshot_id: str) -> int:
return await self.execute("DELETE FROM ai_agent_snapshot WHERE id=:id", {"id": snapshot_id})
async def list_legacy_snapshots(self, after_id: str | None, limit: int, version: int) -> list[dict[str, Any]]:
params: dict[str, Any] = {"version": version, "limit": limit}
extra = ""
if after_id is not None:
extra = " AND id>:after_id"
params["after_id"] = after_id
return await self.fetch_all(
f"SELECT id,snapshot_data,redaction_version FROM ai_agent_snapshot"
f" WHERE redaction_version<:version{extra} ORDER BY id ASC LIMIT :limit",
params,
)
async def update_redacted_snapshot(self, snapshot_id: str, snapshot_data: str, version: int) -> int:
return await self.execute(
"UPDATE ai_agent_snapshot SET snapshot_data=:data,redaction_version=:version"
" WHERE id=:id AND redaction_version<:version",
{"id": snapshot_id, "data": snapshot_data, "version": version},
)
@@ -0,0 +1,105 @@
from __future__ import annotations
from typing import Any
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import Repository
class ConfigRepository(Repository):
def __init__(self, session: AsyncSession):
super().__init__(session)
async def list_params(self) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT param_code, param_value, value_type FROM sys_params WHERE param_type = 1"
)
async def get_param_value(self, code: str) -> str | None:
value = await self.scalar(
"SELECT param_value FROM sys_params WHERE param_code = :code LIMIT 1",
{"code": code},
)
return None if value is None else str(value)
async def get_default_template(self) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id, agent_code, agent_name, asr_model_id, vad_model_id, llm_model_id, "
"vllm_model_id, tts_model_id, tts_voice_id, tts_language, tts_volume, tts_rate, tts_pitch, "
"mem_model_id, intent_model_id, chat_history_conf, system_prompt, summary_memory, lang_code, "
"language, sort FROM ai_agent_template ORDER BY sort ASC LIMIT 1"
)
async def get_device_by_mac(self, mac_address: str) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id, user_id, mac_address, board, agent_id, app_version, auto_update "
"FROM ai_device WHERE mac_address = :mac_address LIMIT 1",
{"mac_address": mac_address},
)
async def get_agent(self, agent_id: str) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id, user_id, agent_code, agent_name, asr_model_id, vad_model_id, llm_model_id, slm_model_id, "
"vllm_model_id, tts_model_id, tts_voice_id, tts_language, tts_volume, tts_rate, tts_pitch, "
"mem_model_id, intent_model_id, chat_history_conf, system_prompt, summary_memory, lang_code, language "
"FROM ai_agent WHERE id = :id LIMIT 1",
{"id": agent_id},
)
async def get_model(self, model_id: str) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id, model_type, model_code, model_name, is_default, is_enabled, config_json, doc_link, "
"remark, sort, creator, create_date, updater, update_date "
"FROM ai_model_config WHERE id = :id LIMIT 1",
{"id": model_id},
)
async def get_timbre(self, timbre_id: str) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id, languages, name, remark, reference_audio, reference_text, sort, tts_model_id, "
"tts_voice, voice_demo FROM ai_tts_voice WHERE id = :id LIMIT 1",
{"id": timbre_id},
)
async def get_voice_clone(self, clone_id: str) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id, name, model_id, voice_id, languages, user_id, train_status, train_error "
"FROM ai_voice_clone WHERE id = :id LIMIT 1",
{"id": clone_id},
)
async def get_plugin_mappings(self, agent_id: str) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT m.id, m.agent_id, m.plugin_id, m.param_info, "
"(SELECT p.provider_code FROM ai_model_provider p WHERE p.id = m.plugin_id LIMIT 1) AS provider_code "
"FROM ai_agent_plugin_mapping m WHERE m.agent_id = :agent_id",
{"agent_id": agent_id},
)
async def get_dataset(self, dataset_id: str) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id, dataset_id, rag_model_id, name, description, status "
"FROM ai_rag_dataset WHERE id = :id LIMIT 1",
{"id": dataset_id},
)
async def get_context_providers(self, agent_id: str) -> Any:
return await self.scalar(
"SELECT context_providers FROM ai_agent_context_provider WHERE agent_id = :agent_id LIMIT 1",
{"agent_id": agent_id},
)
async def get_voiceprints(self, agent_id: str) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT id, agent_id, source_name, introduce, create_date "
"FROM ai_agent_voice_print WHERE agent_id = :agent_id ORDER BY create_date ASC",
{"agent_id": agent_id},
)
async def get_correct_word_items(self, agent_id: str) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT i.source_word, i.target_word FROM ai_agent_correct_word_mapping m "
"JOIN ai_agent_correct_word_item i ON i.file_id = m.file_id WHERE m.agent_id = :agent_id",
{"agent_id": agent_id},
)
@@ -0,0 +1,115 @@
from __future__ import annotations
import uuid
from collections.abc import Sequence
from typing import Any
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import Repository
class CorrectWordRepository(Repository):
def __init__(self, session: AsyncSession):
super().__init__(session)
async def name_exists(self, user_id: int, file_name: str, exclude_id: str | None = None) -> bool:
return bool(
await self.scalar(
"SELECT COUNT(*) FROM ai_agent_correct_word_file WHERE creator=:creator AND file_name=:file_name "
"AND (:exclude_id IS NULL OR id<>:exclude_id)",
{"creator": user_id, "file_name": file_name, "exclude_id": exclude_id},
)
)
async def insert_file(self, values: dict[str, Any]) -> None:
await self.execute(
"INSERT INTO ai_agent_correct_word_file "
"(id, file_name, word_count, content, creator, created_at) "
"VALUES (:id, :file_name, :word_count, :content, :creator, :now)",
values,
)
async def insert_items(self, values: Sequence[dict[str, Any]]) -> None:
await self.execute_many(
"INSERT INTO ai_agent_correct_word_item (id, file_id, source_word, target_word) "
"VALUES (:id, :file_id, :source_word, :target_word)",
values,
)
async def get_file(self, file_id: str, *, for_update: bool = False) -> dict[str, Any] | None:
suffix = " FOR UPDATE" if for_update and self.session.get_bind().dialect.name != "sqlite" else ""
return await self.fetch_one(
f"SELECT * FROM ai_agent_correct_word_file WHERE id=:id{suffix}", # noqa: S608
{"id": file_id},
)
async def update_file(self, values: dict[str, Any]) -> int:
return await self.execute(
"UPDATE ai_agent_correct_word_file SET file_name=:file_name, word_count=:word_count, "
"content=:content, updater=:updater, updated_at=:now WHERE id=:id",
values,
)
async def list_files(
self, user_id: int, *, offset: int | None = None, limit: int | None = None
) -> tuple[list[dict[str, Any]], int]:
total = int(
await self.scalar(
"SELECT COUNT(*) FROM ai_agent_correct_word_file WHERE creator=:creator", {"creator": user_id}
)
or 0
)
if offset is None or limit is None:
rows = await self.fetch_all(
"SELECT * FROM ai_agent_correct_word_file WHERE creator=:creator ORDER BY created_at DESC",
{"creator": user_id},
)
else:
rows = await self.fetch_all(
"SELECT * FROM ai_agent_correct_word_file WHERE creator=:creator ORDER BY created_at DESC "
"LIMIT :offset, :limit",
{"creator": user_id, "offset": offset, "limit": limit},
)
return rows, total
async def delete_file_graph(self, file_id: str) -> None:
await self.execute("DELETE FROM ai_agent_correct_word_mapping WHERE file_id=:id", {"id": file_id})
await self.execute("DELETE FROM ai_agent_correct_word_item WHERE file_id=:id", {"id": file_id})
await self.execute("DELETE FROM ai_agent_correct_word_file WHERE id=:id", {"id": file_id})
async def delete_items(self, file_id: str) -> None:
await self.execute("DELETE FROM ai_agent_correct_word_item WHERE file_id=:id", {"id": file_id})
async def items_for_agent(self, agent_id: str) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT i.source_word, i.target_word FROM ai_agent_correct_word_item i "
"JOIN ai_agent_correct_word_mapping m ON m.file_id=i.file_id WHERE m.agent_id=:agent_id",
{"agent_id": agent_id},
)
async def file_ids_for_agent(self, agent_id: str) -> list[str]:
rows = await self.fetch_all(
"SELECT file_id FROM ai_agent_correct_word_mapping WHERE agent_id=:agent_id", {"agent_id": agent_id}
)
return [str(row["file_id"]) for row in rows]
async def replace_agent_mappings(
self, agent_id: str, file_ids: Sequence[str], user_id: int, now: Any
) -> None:
await self.execute("DELETE FROM ai_agent_correct_word_mapping WHERE agent_id=:agent_id", {"agent_id": agent_id})
await self.execute_many(
"INSERT INTO ai_agent_correct_word_mapping "
"(id, agent_id, file_id, creator, created_at, updater, updated_at) "
"VALUES (:id, :agent_id, :file_id, :user_id, :now, :user_id, :now)",
[
{
"id": uuid.uuid4().hex,
"agent_id": agent_id,
"file_id": file_id,
"user_id": user_id,
"now": now,
}
for file_id in file_ids
],
)
@@ -0,0 +1,294 @@
from __future__ import annotations
# Every interpolated SQL fragment below is a module constant or a service-side allowlist.
# ruff: noqa: S608
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import Any
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import Repository
DEVICE_COLUMNS = (
"id, user_id, mac_address, last_connected_at, auto_update, board, alias, "
"agent_id, app_version, sort, updater, update_date, creator, create_date"
)
OTA_COLUMNS = (
"id, firmware_name, type, version, size, remark, firmware_path, sort, "
"updater, update_date, creator, create_date"
)
ADDRESS_BOOK_COLUMNS = (
"mac_address, target_mac, alias, has_permission, creator, create_date, updater, update_date"
)
class DeviceRepository(Repository):
def __init__(self, session: AsyncSession):
super().__init__(session)
async def get_device(self, device_id: str) -> dict[str, Any] | None:
return await self.fetch_one(
f"SELECT {DEVICE_COLUMNS} FROM ai_device WHERE id = :id LIMIT 1",
{"id": device_id},
)
async def get_device_by_mac(self, mac_address: str) -> dict[str, Any] | None:
return await self.fetch_one(
f"SELECT {DEVICE_COLUMNS} FROM ai_device WHERE mac_address = :mac_address LIMIT 1",
{"mac_address": mac_address},
)
async def get_user_devices(self, user_id: int, agent_id: str) -> list[dict[str, Any]]:
return await self.fetch_all(
f"SELECT {DEVICE_COLUMNS} FROM ai_device WHERE user_id = :user_id AND agent_id = :agent_id",
{"user_id": user_id, "agent_id": agent_id},
)
async def insert_device(self, values: Mapping[str, Any]) -> None:
await self.execute(
"INSERT INTO ai_device "
"(id, user_id, mac_address, last_connected_at, auto_update, board, alias, agent_id, app_version, "
"sort, updater, update_date, creator, create_date) "
"VALUES (:id, :user_id, :mac_address, :last_connected_at, :auto_update, :board, :alias, :agent_id, "
":app_version, :sort, :updater, :update_date, :creator, :create_date)",
values,
)
async def update_device_info(
self,
device_id: str,
*,
auto_update: int | None,
alias: str | None,
updater: int,
now: datetime,
) -> int:
assignments = ["updater = :updater", "update_date = :now"]
params: dict[str, Any] = {"id": device_id, "updater": updater, "now": now}
if auto_update is not None:
assignments.append("auto_update = :auto_update")
params["auto_update"] = auto_update
if alias is not None:
assignments.append("alias = :alias")
params["alias"] = alias
return await self.execute(
f"UPDATE ai_device SET {', '.join(assignments)} WHERE id = :id",
params,
)
async def update_connection(
self,
device_id: str,
*,
app_version: str | None,
now: datetime,
) -> int:
if app_version is None or not app_version.strip():
return await self.execute(
"UPDATE ai_device SET last_connected_at = :now WHERE id = :id",
{"id": device_id, "now": now},
)
return await self.execute(
"UPDATE ai_device SET last_connected_at = :now, app_version = :app_version WHERE id = :id",
{"id": device_id, "now": now, "app_version": app_version},
)
async def delete_device_for_user(self, device_id: str, user_id: int) -> int:
return await self.execute(
"DELETE FROM ai_device WHERE id = :id AND user_id = :user_id",
{"id": device_id, "user_id": user_id},
)
async def get_address_book(self, mac_address: str) -> list[dict[str, Any]]:
return await self.fetch_all(
f"SELECT {ADDRESS_BOOK_COLUMNS} FROM ai_device_address_book "
"WHERE mac_address = :mac_address ORDER BY update_date DESC",
{"mac_address": mac_address},
)
async def get_all_address_book(self) -> list[dict[str, Any]]:
return await self.fetch_all(f"SELECT {ADDRESS_BOOK_COLUMNS} FROM ai_device_address_book")
async def get_address_book_record(self, mac_address: str, target_mac: str) -> dict[str, Any] | None:
return await self.fetch_one(
f"SELECT {ADDRESS_BOOK_COLUMNS} FROM ai_device_address_book "
"WHERE mac_address = :mac_address AND target_mac = :target_mac LIMIT 1",
{"mac_address": mac_address, "target_mac": target_mac},
)
async def get_aliases(self, mac_address: str) -> list[str]:
rows = await self.fetch_all(
"SELECT alias FROM ai_device_address_book WHERE mac_address = :mac_address",
{"mac_address": mac_address},
)
return [str(row["alias"]) for row in rows if row.get("alias") not in (None, "")]
async def insert_address_book(
self,
*,
mac_address: str,
target_mac: str,
alias: str | None,
has_permission: bool | None,
actor: int,
now: datetime,
) -> None:
await self.execute(
"INSERT INTO ai_device_address_book "
"(mac_address, target_mac, alias, has_permission, creator, create_date, updater, update_date) "
"VALUES (:mac_address, :target_mac, :alias, :has_permission, :actor, :now, :actor, :now)",
{
"mac_address": mac_address,
"target_mac": target_mac,
"alias": alias,
"has_permission": has_permission,
"actor": actor,
"now": now,
},
)
async def update_address_alias(
self,
mac_address: str,
target_mac: str,
alias: str | None,
*,
now: datetime,
) -> int:
return await self.execute(
"UPDATE ai_device_address_book SET alias = :alias, update_date = :now "
"WHERE mac_address = :mac_address AND target_mac = :target_mac",
{
"mac_address": mac_address,
"target_mac": target_mac,
"alias": alias,
"now": now,
},
)
async def update_address_permission(
self,
mac_address: str,
target_mac: str,
has_permission: bool,
*,
now: datetime,
) -> int:
return await self.execute(
"UPDATE ai_device_address_book "
"SET has_permission = :has_permission, update_date = :now "
"WHERE mac_address = :mac_address AND target_mac = :target_mac",
{
"mac_address": mac_address,
"target_mac": target_mac,
"has_permission": has_permission,
"now": now,
},
)
async def delete_address_books_for_macs(self, mac_addresses: Sequence[str]) -> int:
if not mac_addresses:
return 0
placeholders = ", ".join(f":mac_{index}" for index in range(len(mac_addresses)))
params = {f"mac_{index}": mac for index, mac in enumerate(mac_addresses)}
return await self.execute(
f"DELETE FROM ai_device_address_book WHERE mac_address IN ({placeholders}) "
f"OR target_mac IN ({placeholders})",
params,
)
async def count_ota(self, firmware_name: str | None = None) -> int:
where = ""
params: dict[str, Any] = {}
if firmware_name is not None and firmware_name.strip():
where = " WHERE firmware_name LIKE :firmware_name"
params["firmware_name"] = f"%{firmware_name}%"
return int(await self.scalar(f"SELECT COUNT(*) FROM ai_ota{where}", params) or 0)
async def list_ota(
self,
*,
page: int,
limit: int,
firmware_name: str | None,
order_fields: Sequence[str],
ascending: bool,
) -> list[dict[str, Any]]:
where = ""
params: dict[str, Any] = {"limit": limit, "offset": max(page - 1, 0) * limit}
if firmware_name is not None and firmware_name.strip():
where = " WHERE firmware_name LIKE :firmware_name"
params["firmware_name"] = f"%{firmware_name}%"
direction = "ASC" if ascending else "DESC"
order_by = ", ".join(f"{field} {direction}" for field in order_fields)
return await self.fetch_all(
f"SELECT {OTA_COLUMNS} FROM ai_ota{where} ORDER BY {order_by} LIMIT :limit OFFSET :offset",
params,
)
async def get_ota(self, ota_id: str) -> dict[str, Any] | None:
return await self.fetch_one(
f"SELECT {OTA_COLUMNS} FROM ai_ota WHERE id = :id LIMIT 1",
{"id": ota_id},
)
async def get_first_ota_by_type(self, ota_type: str) -> dict[str, Any] | None:
return await self.fetch_one(
f"SELECT {OTA_COLUMNS} FROM ai_ota WHERE type = :type LIMIT 1",
{"type": ota_type},
)
async def get_latest_ota(self, ota_type: str) -> dict[str, Any] | None:
return await self.fetch_one(
f"SELECT {OTA_COLUMNS} FROM ai_ota WHERE type = :type ORDER BY update_date DESC LIMIT 1",
{"type": ota_type},
)
async def count_duplicate_ota(self, *, ota_id: str, ota_type: str | None, version: str | None) -> int:
return int(
await self.scalar(
"SELECT COUNT(*) FROM ai_ota WHERE type = :type AND version = :version AND id <> :id",
{"id": ota_id, "type": ota_type, "version": version},
)
or 0
)
async def insert_ota(self, values: Mapping[str, Any]) -> None:
await self.execute(
"INSERT INTO ai_ota "
"(id, firmware_name, type, version, size, remark, firmware_path, sort, updater, update_date, creator, "
"create_date) VALUES (:id, :firmware_name, :type, :version, :size, :remark, :firmware_path, :sort, "
":updater, :update_date, :creator, :create_date)",
values,
)
async def update_ota(self, ota_id: str, values: Mapping[str, Any]) -> int:
allowed = {
"firmware_name",
"type",
"version",
"size",
"remark",
"firmware_path",
"sort",
"updater",
"update_date",
"creator",
"create_date",
}
selected = {key: value for key, value in values.items() if key in allowed and value is not None}
if not selected:
return 0
assignments = ", ".join(f"{key} = :{key}" for key in selected)
return await self.execute(
f"UPDATE ai_ota SET {assignments} WHERE id = :id",
{"id": ota_id, **selected},
)
async def delete_ota(self, ids: Sequence[str]) -> int:
if not ids:
return 0
placeholders = ", ".join(f":id_{index}" for index in range(len(ids)))
params = {f"id_{index}": value for index, value in enumerate(ids)}
return await self.execute(f"DELETE FROM ai_ota WHERE id IN ({placeholders})", params)
@@ -0,0 +1,344 @@
from __future__ import annotations
import json
import uuid
from collections.abc import Sequence
from datetime import datetime
from typing import Any
from zoneinfo import ZoneInfo
from sqlalchemy import bindparam, text
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.config import get_settings
from app.core.database import Repository
class KnowledgeRepository(Repository):
def __init__(self, session: AsyncSession):
super().__init__(session)
async def dataset_page(
self, user_id: int, name: str | None, offset: int, limit: int
) -> tuple[list[dict[str, Any]], int]:
where = (
"WHERE creator=:creator AND (:name IS NULL OR :name='' OR name LIKE CONCAT('%', :name, '%'))"
)
params = {"creator": user_id, "name": name, "offset": offset, "limit": limit}
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_rag_dataset {where}", params) or 0) # noqa: S608
rows = await self.fetch_all(
f"SELECT * FROM ai_rag_dataset {where} ORDER BY created_at DESC LIMIT :offset, :limit", # noqa: S608
params,
)
return rows, total
async def get_dataset(self, identifier: str, *, for_update: bool = False) -> dict[str, Any] | None:
suffix = " FOR UPDATE" if for_update and self.session.get_bind().dialect.name != "sqlite" else ""
return await self.fetch_one(
f"SELECT * FROM ai_rag_dataset WHERE dataset_id=:id OR id=:id LIMIT 1{suffix}", # noqa: S608
{"id": identifier},
)
async def datasets_by_ids(self, identifiers: Sequence[str]) -> list[dict[str, Any]]:
if not identifiers:
return []
statement = text("SELECT * FROM ai_rag_dataset WHERE dataset_id IN :ids OR id IN :ids").bindparams(
bindparam("ids", expanding=True)
)
result = await self.session.execute(statement, {"ids": list(identifiers)})
return [dict(row) for row in result.mappings().all()]
async def duplicate_dataset_name(self, user_id: int, name: str, exclude_id: str | None = None) -> bool:
return bool(
await self.scalar(
"SELECT COUNT(*) FROM ai_rag_dataset WHERE creator=:creator AND name=:name "
"AND (:exclude_id IS NULL OR id<>:exclude_id)",
{"creator": user_id, "name": name, "exclude_id": exclude_id},
)
)
async def dataset_id_conflict(self, dataset_id: str, exclude_id: str) -> bool:
return bool(
await self.scalar(
"SELECT COUNT(*) FROM ai_rag_dataset WHERE dataset_id=:dataset_id AND id<>:exclude_id",
{"dataset_id": dataset_id, "exclude_id": exclude_id},
)
)
async def insert_dataset(self, values: dict[str, Any]) -> None:
await self.execute(
"INSERT INTO ai_rag_dataset "
"(id,dataset_id,rag_model_id,tenant_id,name,avatar,description,embedding_model,permission,chunk_method,"
"parser_config,chunk_count,document_count,token_num,status,creator,created_at,updater,updated_at) VALUES "
"(:id,:dataset_id,:rag_model_id,:tenant_id,:name,:avatar,:description,:embedding_model,:permission,"
":chunk_method,:parser_config,:chunk_count,:document_count,:token_num,:status,:creator,:created_at,"
":updater,:updated_at)",
values,
)
async def update_dataset(self, values: dict[str, Any]) -> int:
return await self.execute(
"UPDATE ai_rag_dataset SET dataset_id=COALESCE(:dataset_id,dataset_id),"
"rag_model_id=COALESCE(:rag_model_id,rag_model_id),name=COALESCE(:name,name),"
"avatar=COALESCE(:avatar,avatar),description=COALESCE(:description,description),"
"embedding_model=COALESCE(:embedding_model,embedding_model),"
"permission=COALESCE(:permission,permission),chunk_method=COALESCE(:chunk_method,chunk_method),"
"parser_config=COALESCE(:parser_config,parser_config),chunk_count=COALESCE(:chunk_count,chunk_count),"
"token_num=COALESCE(:token_num,token_num),status=COALESCE(:status,status),"
"creator=COALESCE(:creator,creator),created_at=COALESCE(:created_at,created_at),updater=:updater,"
"updated_at=:updated_at WHERE id=:id",
values,
)
async def delete_dataset_local(self, row: dict[str, Any]) -> None:
await self.execute("DELETE FROM ai_agent_plugin_mapping WHERE plugin_id=:id", {"id": row["id"]})
await self.execute("DELETE FROM ai_rag_dataset WHERE id=:id", {"id": row["id"]})
async def rag_models(self) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT id, model_name, config_json FROM ai_model_config WHERE model_type='RAG' AND is_enabled=1 "
"ORDER BY is_default DESC, create_date DESC"
)
async def rag_config(self, model_id: str) -> dict[str, Any]:
row = await self.fetch_one("SELECT config_json FROM ai_model_config WHERE id=:id", {"id": model_id})
if row is None or row.get("config_json") is None:
from app.core.errors import AppError
raise AppError(10164)
raw = row["config_json"]
if isinstance(raw, bytes):
raw = raw.decode("utf-8")
config = dict(raw) if isinstance(raw, dict) else dict(json.loads(str(raw)))
config.setdefault("type", "ragflow")
return config
async def documents_page(
self,
dataset_id: str,
*,
name: str | None,
status: str | None,
offset: int,
limit: int,
) -> tuple[list[dict[str, Any]], int]:
where = (
"WHERE dataset_id=:dataset_id "
"AND (:name IS NULL OR :name='' OR name LIKE CONCAT('%', :name, '%')) "
"AND (:status IS NULL OR :status='' OR status=:status)"
)
params = {"dataset_id": dataset_id, "name": name, "status": status, "offset": offset, "limit": limit}
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_rag_knowledge_document {where}", params) or 0) # noqa: S608
rows = await self.fetch_all(
f"SELECT * FROM ai_rag_knowledge_document {where} " # noqa: S608
"ORDER BY created_at DESC LIMIT :offset, :limit",
params,
)
return rows, total
async def all_documents(self, dataset_id: str) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT * FROM ai_rag_knowledge_document WHERE dataset_id=:dataset_id", {"dataset_id": dataset_id}
)
async def documents_by_remote_ids(self, dataset_id: str, ids: Sequence[str]) -> list[dict[str, Any]]:
if not ids:
return []
statement = text(
"SELECT * FROM ai_rag_knowledge_document WHERE dataset_id=:dataset_id AND document_id IN :ids"
).bindparams(bindparam("ids", expanding=True))
result = await self.session.execute(statement, {"dataset_id": dataset_id, "ids": list(ids)})
return [dict(row) for row in result.mappings().all()]
async def upsert_document(self, dataset_id: str, remote: dict[str, Any], *, creator: int | None = None) -> bool:
document_id = str(remote.get("id") or remote.get("document_id") or "")
existing = await self.fetch_one(
"SELECT id,created_at FROM ai_rag_knowledge_document WHERE document_id=:id", {"id": document_id}
)
name = remote.get("name")
size = remote.get("size")
if size is None:
size = remote.get("file_size")
meta_fields = remote.get("meta_fields")
if meta_fields is None:
meta_fields = remote.get("meta")
error = remote.get("progress_msg")
if error is None:
error = remote.get("error")
synced_at = _shanghai_now_naive()
created_at = remote.get("created_at")
if not isinstance(created_at, datetime):
created_at = _millis_date(remote.get("create_time"))
updated_at = remote.get("updated_at")
if not isinstance(updated_at, datetime):
updated_at = _millis_date(remote.get("update_time"))
values = {
"id": existing["id"] if existing else uuid.uuid4().hex,
"dataset_id": remote.get("dataset_id") or dataset_id,
"document_id": document_id,
"name": name,
"size": size,
"type": _file_type(str(name or "")),
"chunk_method": remote.get("chunk_method"),
"parser_config": _json_dump(remote.get("parser_config")),
"status": _remote_status(remote.get("status")),
"run": remote.get("run"),
"progress": remote.get("progress"),
"thumbnail": remote.get("thumbnail"),
"process_duration": remote.get("process_duration"),
"meta_fields": _json_dump(meta_fields),
"source_type": remote.get("source_type"),
"error": error,
"chunk_count": remote.get("chunk_count") or 0,
"token_count": remote.get("token_count") or 0,
"enabled": 1,
"creator": creator,
"created_at": existing.get("created_at") if existing else (created_at or synced_at),
"updated_at": updated_at or synced_at,
"synced_at": synced_at,
}
if existing:
await self.execute(
"UPDATE ai_rag_knowledge_document SET dataset_id=:dataset_id,document_id=:document_id,name=:name,"
"size=:size,type=:type,chunk_method=:chunk_method,parser_config=:parser_config,status=:status,run=:run,"
"progress=:progress,thumbnail=:thumbnail,process_duration=:process_duration,meta_fields=:meta_fields,"
"source_type=:source_type,error=:error,chunk_count=:chunk_count,token_count=:token_count,enabled=:enabled,"
"updated_at=:updated_at,last_sync_at=:synced_at WHERE id=:id",
values,
)
return False
await self.execute(
"INSERT INTO ai_rag_knowledge_document "
"(id,dataset_id,document_id,name,size,type,chunk_method,parser_config,status,run,progress,thumbnail,"
"process_duration,meta_fields,source_type,error,chunk_count,token_count,enabled,creator,created_at,updated_at,"
"last_sync_at) VALUES (:id,:dataset_id,:document_id,:name,:size,:type,:chunk_method,:parser_config,:status,"
":run,:progress,:thumbnail,:process_duration,:meta_fields,:source_type,:error,:chunk_count,:token_count,"
":enabled,:creator,COALESCE(:created_at,:synced_at),COALESCE(:updated_at,:synced_at),:synced_at)",
values,
)
return True
async def update_stats(self, dataset_id: str, docs: int, chunks: int, tokens: int) -> None:
await self.execute(
"UPDATE ai_rag_dataset SET document_count=document_count+:docs,chunk_count=chunk_count+:chunks,"
"token_num=token_num+:tokens,updated_at=:now WHERE dataset_id=:dataset_id",
{
"dataset_id": dataset_id,
"docs": docs,
"chunks": chunks,
"tokens": tokens,
"now": _shanghai_now_naive(),
},
)
async def delete_document_shadows(self, dataset_id: str, ids: Sequence[str]) -> int:
if not ids:
return 0
statement = text(
"DELETE FROM ai_rag_knowledge_document WHERE dataset_id=:dataset_id AND document_id IN :ids"
).bindparams(bindparam("ids", expanding=True))
result = await self.session.execute(statement, {"dataset_id": dataset_id, "ids": list(ids)})
return int(getattr(result, "rowcount", 0) or 0)
async def mark_documents_running(self, dataset_id: str, ids: Sequence[str], now: datetime) -> int:
if not ids:
return 0
statement = text(
"UPDATE ai_rag_knowledge_document SET run='RUNNING',status='1',updated_at=:now "
"WHERE dataset_id=:dataset_id AND document_id IN :ids"
).bindparams(bindparam("ids", expanding=True))
result = await self.session.execute(statement, {"dataset_id": dataset_id, "ids": list(ids), "now": now})
return int(getattr(result, "rowcount", 0) or 0)
async def mark_document_remote_deleted(self, document_id: str, now: datetime) -> int:
return await self.execute(
"UPDATE ai_rag_knowledge_document SET run='CANCEL',error=:error,updated_at=:now,last_sync_at=:now "
"WHERE document_id=:document_id",
{
"document_id": document_id,
"error": "文档在远程服务中已被删除",
"now": now,
},
)
async def sync_running_document(
self,
dataset_id: str,
document_id: str,
remote: dict[str, Any],
now: datetime,
) -> int:
"""Update exactly the columns touched by Java's status-sync helper."""
updated_at = _millis_date(remote.get("update_time")) or now
meta_fields = remote.get("meta_fields")
assignments = (
"status=:status,run=:run,progress=:progress,chunk_count=:chunk_count,token_count=:token_count,"
"error=:error,process_duration=:process_duration,thumbnail=:thumbnail,updated_at=:updated_at,"
"last_sync_at=:now"
)
if meta_fields is not None:
assignments += ",meta_fields=:meta_fields"
return await self.execute(
f"UPDATE ai_rag_knowledge_document SET {assignments} " # noqa: S608
"WHERE document_id=:document_id AND dataset_id=:dataset_id",
{
"dataset_id": dataset_id,
"document_id": document_id,
"status": remote.get("status"),
"run": remote.get("run"),
"progress": remote.get("progress"),
"chunk_count": remote.get("chunk_count"),
"token_count": remote.get("token_count"),
"error": remote.get("progress_msg") if remote.get("progress_msg") is not None else remote.get("error"),
"process_duration": remote.get("process_duration"),
"thumbnail": remote.get("thumbnail"),
"meta_fields": _json_dump(meta_fields),
"updated_at": updated_at,
"now": now,
},
)
async def running_documents(self) -> list[dict[str, Any]]:
return await self.fetch_all("SELECT * FROM ai_rag_knowledge_document WHERE run='RUNNING' AND status='1'")
def _json_dump(value: Any) -> str | None:
if value is None:
return None
if isinstance(value, str):
return value
return json.dumps(value, ensure_ascii=False, separators=(",", ":"))
def _millis_date(value: Any) -> datetime | None:
if value is None:
return None
try:
timezone = ZoneInfo(get_settings().timezone)
return datetime.fromtimestamp(float(value) / 1000, timezone).replace(tzinfo=None)
except (TypeError, ValueError, OSError):
return None
def _file_type(name: str) -> str:
last_dot = name.rfind(".")
if last_dot <= 0 or last_dot == len(name) - 1:
return "unknown"
extension = name.rsplit(".", 1)[1].lower()
if extension in {"pdf", "doc", "docx", "txt", "md", "mdx"}:
return "document"
if extension in {"csv", "xls", "xlsx"}:
return "spreadsheet"
if extension in {"ppt", "pptx"}:
return "presentation"
return extension
def _remote_status(value: Any) -> str:
if value is None or (isinstance(value, str) and not value.strip()):
return "1"
return str(value)
def _shanghai_now_naive() -> datetime:
return datetime.now(ZoneInfo(get_settings().timezone)).replace(tzinfo=None)
@@ -0,0 +1,243 @@
from __future__ import annotations
import json
from collections.abc import Sequence
from typing import Any
from sqlalchemy import bindparam, text
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import Repository
class ModelRepository(Repository):
def __init__(self, session: AsyncSession):
super().__init__(session)
async def list_model_names(self, model_type: str, model_name: str | None) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT id, model_name FROM ai_model_config "
"WHERE model_type = :model_type AND is_enabled = 1 "
"AND (:model_name IS NULL OR :model_name = '' OR model_name LIKE CONCAT('%', :model_name, '%')) "
"ORDER BY sort ASC",
{"model_type": model_type, "model_name": model_name},
)
async def list_llm_names(self, model_name: str | None) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT id, model_name, config_json FROM ai_model_config "
"WHERE model_type = 'llm' AND is_enabled = 1 "
"AND (:model_name IS NULL OR :model_name = '' OR model_name LIKE CONCAT('%', :model_name, '%'))",
{"model_name": model_name},
)
async def list_providers_by_type(self, model_type: str) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT * FROM ai_model_provider WHERE model_type = :model_type ORDER BY sort ASC",
{"model_type": model_type or ""},
)
async def list_providers(
self,
*,
model_type: str | None,
name: str | None,
offset: int,
limit: int,
) -> tuple[list[dict[str, Any]], int]:
where = (
"WHERE (:model_type IS NULL OR :model_type = '' OR model_type = :model_type) "
"AND (:name IS NULL OR :name = '' OR name LIKE CONCAT('%', :name, '%') "
"OR provider_code LIKE CONCAT('%', :name, '%'))"
)
params = {"model_type": model_type, "name": name, "offset": offset, "limit": limit}
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_model_provider {where}", params) or 0) # noqa: S608
rows = await self.fetch_all(
f"SELECT * FROM ai_model_provider {where} " # noqa: S608
"ORDER BY model_type ASC, sort ASC LIMIT :offset, :limit",
params,
)
return rows, total
async def list_model_configs(
self,
*,
model_type: str,
model_name: str | None,
offset: int,
limit: int,
) -> tuple[list[dict[str, Any]], int]:
where = (
"WHERE model_type = :model_type AND "
"(:model_name IS NULL OR :model_name = '' OR model_name LIKE CONCAT('%', :model_name, '%'))"
)
params = {"model_type": model_type, "model_name": model_name, "offset": offset, "limit": limit}
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_model_config {where}", params) or 0) # noqa: S608
rows = await self.fetch_all(
f"SELECT * FROM ai_model_config {where} " # noqa: S608
"ORDER BY is_enabled DESC, sort ASC LIMIT :offset, :limit",
params,
)
return rows, total
async def get_provider(self, model_type: str, provider_code: str) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT * FROM ai_model_provider WHERE model_type = :model_type AND provider_code = :provider_code LIMIT 1",
{"model_type": model_type or "", "provider_code": provider_code or ""},
)
async def get_model(self, model_id: str, *, for_update: bool = False) -> dict[str, Any] | None:
suffix = " FOR UPDATE" if for_update and self.session.get_bind().dialect.name != "sqlite" else ""
return await self.fetch_one(
f"SELECT * FROM ai_model_config WHERE id = :id LIMIT 1{suffix}", # noqa: S608
{"id": model_id},
)
async def insert_model(self, values: dict[str, Any]) -> None:
await self.execute(
"INSERT INTO ai_model_config "
"(id, model_type, model_code, model_name, is_default, is_enabled, config_json, doc_link, remark, sort) "
"VALUES (:id, :model_type, :model_code, :model_name, :is_default, COALESCE(:is_enabled, 0), "
":config_json, :doc_link, :remark, COALESCE(:sort, 0))",
values,
)
async def update_model(self, values: dict[str, Any]) -> int:
return await self.execute(
"UPDATE ai_model_config SET model_type=:model_type, model_code=:model_code, "
"model_name=COALESCE(:model_name, model_name), is_default=:is_default, "
"is_enabled=COALESCE(:is_enabled, is_enabled), config_json=:config_json, doc_link=:doc_link, "
"remark=COALESCE(:remark, remark), sort=COALESCE(:sort, sort) WHERE id=:id",
values,
)
async def delete_model(self, model_id: str) -> int:
return await self.execute("DELETE FROM ai_model_config WHERE id = :id", {"id": model_id})
async def model_agent_references(self, model_id: str) -> list[str]:
rows = await self.fetch_all(
"SELECT agent_name FROM ai_agent WHERE vad_model_id=:id OR asr_model_id=:id OR llm_model_id=:id "
"OR tts_model_id=:id OR mem_model_id=:id OR vllm_model_id=:id OR intent_model_id=:id",
{"id": model_id},
)
return [str(row.get("agent_name") or "") for row in rows]
async def intent_reference_count(self, model_id: str) -> int:
return int(
await self.scalar(
"SELECT COUNT(*) FROM ai_model_config WHERE model_type='Intent' AND CAST(config_json AS CHAR) LIKE "
"CONCAT('%', :id, '%')",
{"id": model_id},
)
or 0
)
async def set_models_default(self, model_type: str, value: int) -> None:
await self.execute(
"UPDATE ai_model_config SET is_default=:value WHERE model_type=:model_type",
{"value": value, "model_type": model_type},
)
async def set_model_enabled(self, model_id: str, status: int) -> int:
return await self.execute(
"UPDATE ai_model_config SET is_enabled=:status WHERE id=:id",
{"status": status, "id": model_id},
)
async def update_default_template_models(self, model_type: str, model_id: str) -> None:
columns = {
"ASR": ("asr_model_id",),
"VAD": ("vad_model_id",),
"LLM": ("llm_model_id",),
"TTS": ("tts_model_id", "tts_voice_id"),
"VLLM": ("vllm_model_id",),
"MEMORY": ("mem_model_id",),
"INTENT": ("intent_model_id",),
}.get(model_type.upper())
if not columns:
return
if columns == ("tts_model_id", "tts_voice_id"):
await self.execute(
"UPDATE ai_agent_template SET tts_model_id=:id, tts_voice_id=NULL WHERE sort >= 0",
{"id": model_id},
)
else:
column = columns[0]
await self.session.execute(
text(f"UPDATE ai_agent_template SET {column}=:id WHERE sort >= 0"), # noqa: S608
{"id": model_id},
)
async def insert_provider(self, values: dict[str, Any]) -> None:
if self.session.get_bind().dialect.name == "sqlite":
statement = (
"INSERT INTO ai_model_provider "
"(id, model_type, provider_code, name, fields, sort, creator, create_date, updater, update_date) "
"VALUES (:id, :model_type, :provider_code, :name, :fields, :sort, :creator, :now, :updater, :now)"
)
else:
statement = (
"INSERT INTO ai_model_provider "
"(id, model_type, provider_code, name, fields, sort, creator, create_date, updater, update_date) "
"VALUES (:id, :model_type, :provider_code, :name, CAST(:fields AS JSON), :sort, :creator, :now, "
":updater, :now)"
)
await self.execute(statement, values)
async def update_provider(self, values: dict[str, Any]) -> int:
if self.session.get_bind().dialect.name == "sqlite":
statement = (
"UPDATE ai_model_provider SET model_type=:model_type, provider_code=:provider_code, name=:name, "
"fields=:fields, sort=:sort, updater=:updater, update_date=:now WHERE id=:id"
)
else:
statement = (
"UPDATE ai_model_provider SET model_type=:model_type, provider_code=:provider_code, name=:name, "
"fields=CAST(:fields AS JSON), sort=:sort, updater=:updater, update_date=:now WHERE id=:id"
)
return await self.execute(statement, values)
async def delete_providers(self, ids: Sequence[str]) -> int:
if not ids:
return 0
statement = text("DELETE FROM ai_model_provider WHERE id IN :ids").bindparams(
bindparam("ids", expanding=True)
)
result = await self.session.execute(statement, {"ids": list(ids)})
return int(getattr(result, "rowcount", 0) or 0)
async def list_plugins_for_user(self, user_id: int) -> list[dict[str, Any]]:
providers = await self.fetch_all("SELECT * FROM ai_model_provider WHERE model_type='Plugin'")
datasets = await self.fetch_all(
"SELECT id, name, created_at, updated_at FROM ai_rag_dataset WHERE creator=:creator AND status=1",
{"creator": user_id},
)
providers.extend(
{
"id": row["id"],
"model_type": "Rag",
"name": f"[知识库]{row['name']}",
"provider_code": "ragflow",
"fields": "[]",
"sort": 0,
"create_date": row.get("created_at"),
"update_date": row.get("updated_at"),
"creator": 0,
"updater": 0,
}
for row in datasets
)
return providers
def parse_json_object(value: Any) -> dict[str, Any] | None:
if value is None:
return None
if isinstance(value, dict):
return dict(value)
if isinstance(value, bytes):
value = value.decode("utf-8")
if isinstance(value, str):
parsed = json.loads(value)
return dict(parsed) if isinstance(parsed, dict) else None
return None
@@ -0,0 +1,153 @@
from __future__ import annotations
from datetime import datetime
from typing import Any
from sqlalchemy import text
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import Repository
class SecurityRepository(Repository):
def __init__(self, session: AsyncSession):
super().__init__(session)
async def get_param_value(self, code: str) -> str | None:
value = await self.scalar(
"SELECT param_value FROM sys_params WHERE param_code = :code LIMIT 1",
{"code": code},
)
return None if value is None else str(value)
async def get_mobile_area_items(self) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT d.dict_label AS name, d.dict_value AS `key` "
"FROM sys_dict_data d "
"LEFT JOIN sys_dict_type t ON d.dict_type_id = t.id "
"WHERE t.dict_type = :dict_type ORDER BY d.sort ASC",
{"dict_type": "MOBILE_AREA"},
)
async def get_user_by_username(self, username: str | None) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id, username, password, super_admin, status, creator, create_date, updater, update_date "
"FROM sys_user WHERE username = :username LIMIT 1",
{"username": username},
)
async def get_user_by_id(self, user_id: int) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id, username, password, super_admin, status, creator, create_date, updater, update_date "
"FROM sys_user WHERE id = :id LIMIT 1",
{"id": user_id},
)
async def count_users(self) -> int:
return int(await self.scalar("SELECT COUNT(*) FROM sys_user") or 0)
async def insert_user(
self,
*,
user_id: int,
username: str | None,
password: str,
super_admin: int,
now: datetime,
) -> None:
await self.execute(
"INSERT INTO sys_user "
"(id, username, password, super_admin, status, creator, create_date, updater, update_date) "
"VALUES (:id, :username, :password, :super_admin, 1, NULL, :now, NULL, :now)",
{
"id": user_id,
"username": username,
"password": password,
"super_admin": super_admin,
"now": now,
},
)
async def get_token_by_user_id(self, user_id: int, *, for_update: bool = False) -> dict[str, Any] | None:
sql = (
"SELECT id, user_id, token, expire_date, update_date, create_date "
"FROM sys_user_token WHERE user_id = :user_id LIMIT 1 FOR UPDATE"
if for_update and self._supports_for_update()
else "SELECT id, user_id, token, expire_date, update_date, create_date "
"FROM sys_user_token WHERE user_id = :user_id LIMIT 1"
)
return await self.fetch_one(
sql,
{"user_id": user_id},
)
async def insert_token(
self,
*,
token_id: int,
user_id: int,
token: str,
now: datetime,
expire_date: datetime,
) -> None:
await self.execute(
"INSERT INTO sys_user_token (id, user_id, token, expire_date, update_date, create_date) "
"VALUES (:id, :user_id, :token, :expire_date, :now, :now)",
{
"id": token_id,
"user_id": user_id,
"token": token,
"expire_date": expire_date,
"now": now,
},
)
async def update_token(self, *, token_id: int, token: str, now: datetime, expire_date: datetime) -> None:
await self.execute(
"UPDATE sys_user_token SET token = :token, expire_date = :expire_date, update_date = :now "
"WHERE id = :id",
{"id": token_id, "token": token, "expire_date": expire_date, "now": now},
)
async def update_password(
self,
user_id: int,
password_hash: str,
now: datetime,
*,
preserve_audit_fields: bool = False,
) -> int:
return await self.execute(
"UPDATE sys_user SET password = :password, "
"update_date = CASE WHEN :preserve_audit = 1 THEN update_date ELSE :now END WHERE id = :id",
{
"id": user_id,
"password": password_hash,
"now": now,
"preserve_audit": int(preserve_audit_fields),
},
)
async def expire_user_token(self, user_id: int, expire_date: datetime) -> int:
return await self.execute(
"UPDATE sys_user_token SET expire_date = :expire_date WHERE user_id = :user_id",
{"user_id": user_id, "expire_date": expire_date},
)
def _supports_for_update(self) -> bool:
bind = self.session.get_bind()
return bind.dialect.name != "sqlite"
async def raw_user_token(session: AsyncSession, token: str) -> dict[str, Any] | None:
result = await session.execute(
text(
"SELECT t.id AS token_id, t.user_id, t.token, t.expire_date, "
"u.username, u.super_admin, u.status "
"FROM sys_user_token t JOIN sys_user u ON u.id = t.user_id "
"WHERE t.token = :token LIMIT 1"
),
{"token": token},
)
row = result.mappings().first()
return dict(row) if row is not None else None
@@ -0,0 +1,494 @@
from __future__ import annotations
from datetime import datetime
from typing import Any
from sqlalchemy import bindparam, text
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import Repository
class SysRepository(Repository):
def __init__(self, session: AsyncSession):
super().__init__(session)
async def page_users(self, *, mobile: str | None, page: int, limit: int) -> tuple[list[dict[str, Any]], int]:
pattern = f"%{mobile}%" if mobile else None
params = {"mobile": pattern, "offset": (page - 1) * limit, "limit": limit}
total = int(
await self.scalar(
"SELECT COUNT(*) FROM sys_user WHERE (:mobile IS NULL OR username LIKE :mobile)",
params,
)
or 0
)
rows = await self.fetch_all(
"SELECT u.id, u.username, u.status, u.create_date, "
"(SELECT COUNT(*) FROM ai_device d WHERE d.user_id = u.id) AS device_count "
"FROM sys_user u WHERE (:mobile IS NULL OR u.username LIKE :mobile) "
"ORDER BY u.id ASC LIMIT :limit OFFSET :offset",
params,
)
return rows, total
async def reset_user_password(
self,
user_id: int,
password_hash: str,
updater: int,
now: datetime,
) -> int:
return await self.execute(
"UPDATE sys_user SET password = :password, updater = :updater, update_date = :now WHERE id = :id",
{"id": user_id, "password": password_hash, "updater": updater, "now": now},
)
async def change_user_status(self, status: int, user_ids: list[int], updater: int, now: datetime) -> int:
statement = text(
"UPDATE sys_user SET status = :status, updater = :updater, update_date = :now WHERE id IN :ids"
).bindparams(bindparam("ids", expanding=True))
return await self.execute(
statement,
{"status": status, "updater": updater, "now": now, "ids": user_ids},
)
async def delete_user_cascade(self, user_id: int) -> None:
agent_rows = await self.fetch_all("SELECT id FROM ai_agent WHERE user_id = :user_id", {"user_id": user_id})
agent_ids = [str(row["id"]) for row in agent_rows]
await self.execute("DELETE FROM sys_user WHERE id = :id", {"id": user_id})
await self.execute("DELETE FROM ai_device WHERE user_id = :id", {"id": user_id})
for agent_id in agent_ids:
audio_rows = await self.fetch_all(
"SELECT DISTINCT audio_id FROM ai_agent_chat_history "
"WHERE agent_id = :agent_id AND audio_id IS NOT NULL",
{"agent_id": agent_id},
)
audio_ids = [str(row["audio_id"]) for row in audio_rows]
if audio_ids:
statement = text("DELETE FROM ai_agent_chat_audio WHERE id IN :ids").bindparams(
bindparam("ids", expanding=True)
)
await self.execute(statement, {"ids": audio_ids})
for table in (
"ai_agent_chat_history",
"ai_agent_plugin_mapping",
"ai_agent_context_provider",
"ai_agent_correct_word_mapping",
"ai_agent_tag_relation",
"ai_agent_snapshot",
):
# Table names are a closed list mirroring AgentServiceImpl.deleteAgent.
await self.execute(
f"DELETE FROM {table} WHERE agent_id = :agent_id", # noqa: S608 - closed table list above
{"agent_id": agent_id},
)
await self.execute("DELETE FROM ai_device WHERE agent_id = :agent_id", {"agent_id": agent_id})
await self.execute("DELETE FROM ai_agent WHERE id = :agent_id", {"agent_id": agent_id})
async def page_devices(
self,
*,
keywords: str | None,
page: int,
limit: int,
) -> tuple[list[dict[str, Any]], int]:
pattern = f"%{keywords}%" if keywords else None
params = {"keywords": pattern, "offset": (page - 1) * limit, "limit": limit}
total = int(
await self.scalar(
"SELECT COUNT(*) FROM ai_device WHERE (:keywords IS NULL OR alias LIKE :keywords)",
params,
)
or 0
)
rows = await self.fetch_all(
"SELECT d.id, d.user_id, d.mac_address, d.last_connected_at, d.auto_update, d.board, d.alias, "
"d.agent_id, d.app_version, d.sort, d.create_date, d.update_date, u.username AS bind_user_name "
"FROM ai_device d LEFT JOIN sys_user u ON u.id = d.user_id "
"WHERE (:keywords IS NULL OR d.alias LIKE :keywords) "
"ORDER BY d.mac_address ASC LIMIT :limit OFFSET :offset",
params,
)
return rows, total
async def page_params(
self,
*,
param_code: str | None,
page: int,
limit: int,
order_field: str | None,
order: str | None,
) -> tuple[list[dict[str, Any]], int]:
pattern = f"%{param_code}%" if param_code else None
params = {"pattern": pattern, "offset": (page - 1) * limit, "limit": limit}
where = "param_type = 1 AND (:pattern IS NULL OR param_code LIKE :pattern OR remark LIKE :pattern)"
total = int(await self.scalar(f"SELECT COUNT(*) FROM sys_params WHERE {where}", params) or 0) # noqa: S608
allowed = {
"id": "id",
"paramCode": "param_code",
"paramValue": "param_value",
"valueType": "value_type",
"createDate": "create_date",
"updateDate": "update_date",
}
order_column = allowed.get(order_field or "")
order_clause = ""
if order_column is not None:
direction = "ASC" if (order or "").lower() == "asc" else "DESC"
order_clause = f" ORDER BY {order_column} {direction}"
sql = (
"SELECT id, param_code, param_value, value_type, remark, create_date, update_date " # noqa: S608
f"FROM sys_params WHERE {where}{order_clause} LIMIT :limit OFFSET :offset"
)
return await self.fetch_all(sql, params), total # noqa: S608
async def list_config_params(self) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT id, param_code, param_value, value_type, remark, create_date, update_date "
"FROM sys_params WHERE param_type = 1"
)
async def get_param(self, param_id: int) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id, param_code, param_value, value_type, remark, create_date, update_date "
"FROM sys_params WHERE id = :id",
{"id": param_id},
)
async def get_param_value(self, code: str) -> str | None:
value = await self.scalar("SELECT param_value FROM sys_params WHERE param_code = :code", {"code": code})
return None if value is None else str(value)
async def insert_param(
self,
*,
param_id: int,
param_code: str,
param_value: str,
value_type: str,
remark: str | None,
user_id: int,
now: datetime,
) -> None:
await self.execute(
"INSERT INTO sys_params "
"(id, param_code, param_value, value_type, param_type, remark, creator, create_date, updater, update_date) "
"VALUES (:id, :code, :value, :value_type, 1, :remark, :user_id, :now, :user_id, :now)",
{
"id": param_id,
"code": param_code,
"value": param_value,
"value_type": value_type,
"remark": remark,
"user_id": user_id,
"now": now,
},
)
async def update_param(
self,
*,
param_id: int,
param_code: str,
param_value: str,
value_type: str,
remark: str | None,
user_id: int,
now: datetime,
) -> int:
return await self.execute(
"UPDATE sys_params SET param_code = :code, param_value = :value, value_type = :value_type, "
"remark = CASE WHEN :has_remark = 1 THEN :remark ELSE remark END, updater = :user_id, update_date = :now "
"WHERE id = :id",
{
"id": param_id,
"code": param_code,
"value": param_value,
"value_type": value_type,
"has_remark": int(remark is not None),
"remark": remark,
"user_id": user_id,
"now": now,
},
)
async def update_param_value_by_code(self, code: str, value: str, user_id: int, now: datetime) -> int:
return await self.execute(
"UPDATE sys_params SET param_value = :value, updater = :user_id, update_date = :now "
"WHERE param_code = :code",
{"code": code, "value": value, "user_id": user_id, "now": now},
)
async def param_codes_for_ids(self, ids: list[int]) -> list[str]:
statement = text("SELECT param_code FROM sys_params WHERE id IN :ids").bindparams(
bindparam("ids", expanding=True)
)
rows = await self.fetch_all(statement, {"ids": ids})
return [str(row["param_code"]) for row in rows]
async def delete_params(self, ids: list[int]) -> int:
statement = text("DELETE FROM sys_params WHERE id IN :ids").bindparams(bindparam("ids", expanding=True))
return await self.execute(statement, {"ids": ids})
async def delete_plugin_mapping_by_plugin_id(self, plugin_id: str) -> int:
return await self.execute(
"DELETE FROM ai_agent_plugin_mapping WHERE plugin_id = :plugin_id",
{"plugin_id": plugin_id},
)
async def page_dict_types(
self,
*,
dict_type: str | None,
dict_name: str | None,
page: int,
limit: int,
) -> tuple[list[dict[str, Any]], int]:
params = {
"dict_type": f"%{dict_type}%" if dict_type else None,
"dict_name": f"%{dict_name}%" if dict_name else None,
"offset": (page - 1) * limit,
"limit": limit,
}
where = (
"(:dict_type IS NULL OR t.dict_type LIKE :dict_type) "
"AND (:dict_name IS NULL OR t.dict_name LIKE :dict_name)"
)
total = int(await self.scalar(f"SELECT COUNT(*) FROM sys_dict_type t WHERE {where}", params) or 0) # noqa: S608
rows = await self.fetch_all(
"SELECT t.id, t.dict_type, t.dict_name, t.remark, t.sort, t.creator, t.create_date, t.updater, " # noqa: S608
"t.update_date, creator.username AS creator_name, updater.username AS updater_name "
"FROM sys_dict_type t LEFT JOIN sys_user creator ON creator.id = t.creator "
"LEFT JOIN sys_user updater ON updater.id = t.updater "
f"WHERE {where} ORDER BY t.sort ASC LIMIT :limit OFFSET :offset", # noqa: S608
params,
)
return rows, total
async def get_dict_type(self, type_id: int) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id, dict_type, dict_name, remark, sort, creator, create_date, updater, update_date "
"FROM sys_dict_type WHERE id = :id",
{"id": type_id},
)
async def dict_type_exists(self, dict_type: str | None, *, exclude_id: int | None = None) -> bool:
count = await self.scalar(
"SELECT COUNT(*) FROM sys_dict_type WHERE dict_type = :dict_type "
"AND (:exclude_id IS NULL OR id <> :exclude_id)",
{"dict_type": dict_type, "exclude_id": exclude_id},
)
return int(count or 0) > 0
async def insert_dict_type(
self,
*,
type_id: int,
dict_type: str | None,
dict_name: str | None,
remark: str | None,
sort: int | None,
user_id: int,
now: datetime,
) -> None:
await self.execute(
"INSERT INTO sys_dict_type "
"(id, dict_type, dict_name, remark, sort, creator, create_date, updater, update_date) "
"VALUES (:id, :dict_type, :dict_name, :remark, :sort, :user_id, :now, :user_id, :now)",
{
"id": type_id,
"dict_type": dict_type,
"dict_name": dict_name,
"remark": remark,
"sort": sort,
"user_id": user_id,
"now": now,
},
)
async def update_dict_type(
self,
*,
type_id: int | None,
dict_type: str | None,
dict_name: str | None,
remark: str | None,
sort: int | None,
user_id: int,
now: datetime,
) -> int:
return await self.execute(
"UPDATE sys_dict_type SET "
"dict_type = CASE WHEN :has_dict_type = 1 THEN :dict_type ELSE dict_type END, "
"dict_name = CASE WHEN :has_dict_name = 1 THEN :dict_name ELSE dict_name END, "
"remark = CASE WHEN :has_remark = 1 THEN :remark ELSE remark END, "
"sort = CASE WHEN :has_sort = 1 THEN :sort ELSE sort END, updater = :user_id, update_date = :now "
"WHERE id = :id",
{
"id": type_id,
"has_dict_type": int(dict_type is not None),
"dict_type": dict_type,
"has_dict_name": int(dict_name is not None),
"dict_name": dict_name,
"has_remark": int(remark is not None),
"remark": remark,
"has_sort": int(sort is not None),
"sort": sort,
"user_id": user_id,
"now": now,
},
)
async def delete_dict_types(self, ids: list[int]) -> None:
statement_data = text("DELETE FROM sys_dict_data WHERE dict_type_id IN :ids").bindparams(
bindparam("ids", expanding=True)
)
statement_types = text("DELETE FROM sys_dict_type WHERE id IN :ids").bindparams(
bindparam("ids", expanding=True)
)
await self.execute(statement_data, {"ids": ids})
await self.execute(statement_types, {"ids": ids})
async def page_dict_data(
self,
*,
dict_type_id: int | None,
dict_label: str | None,
dict_value: str | None,
page: int,
limit: int,
) -> tuple[list[dict[str, Any]], int]:
params = {
"type_id": dict_type_id,
"dict_label": f"%{dict_label}%" if dict_label else None,
"dict_value": f"%{dict_value}%" if dict_value else None,
"offset": (page - 1) * limit,
"limit": limit,
}
where = (
"d.dict_type_id = :type_id AND (:dict_label IS NULL OR d.dict_label LIKE :dict_label) "
"AND (:dict_value IS NULL OR d.dict_value LIKE :dict_value)"
)
total = int(await self.scalar(f"SELECT COUNT(*) FROM sys_dict_data d WHERE {where}", params) or 0) # noqa: S608
rows = await self.fetch_all(
"SELECT d.id, d.dict_type_id, d.dict_label, d.dict_value, d.remark, d.sort, d.creator, " # noqa: S608
"d.create_date, d.updater, d.update_date, creator.username AS creator_name, "
"updater.username AS updater_name FROM sys_dict_data d "
"LEFT JOIN sys_user creator ON creator.id = d.creator "
"LEFT JOIN sys_user updater ON updater.id = d.updater "
f"WHERE {where} ORDER BY d.sort ASC LIMIT :limit OFFSET :offset", # noqa: S608
params,
)
return rows, total
async def get_dict_data(self, data_id: int) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id, dict_type_id, dict_label, dict_value, remark, sort, creator, create_date, updater, update_date "
"FROM sys_dict_data WHERE id = :id",
{"id": data_id},
)
async def dict_data_label_exists(
self,
dict_type_id: int | None,
compared_label: str | None,
*,
exclude_id: int | None = None,
) -> bool:
count = await self.scalar(
"SELECT COUNT(*) FROM sys_dict_data WHERE dict_type_id = :type_id AND dict_label = :label "
"AND (:exclude_id IS NULL OR id <> :exclude_id)",
{"type_id": dict_type_id, "label": compared_label, "exclude_id": exclude_id},
)
return int(count or 0) > 0
async def dict_type_code(self, type_id: int | None) -> str | None:
value = await self.scalar("SELECT dict_type FROM sys_dict_type WHERE id = :id", {"id": type_id})
return None if value is None else str(value)
async def insert_dict_data(
self,
*,
data_id: int,
dict_type_id: int | None,
dict_label: str | None,
dict_value: str | None,
remark: str | None,
sort: int | None,
user_id: int,
now: datetime,
) -> None:
await self.execute(
"INSERT INTO sys_dict_data "
"(id, dict_type_id, dict_label, dict_value, remark, sort, creator, create_date, updater, update_date) "
"VALUES (:id, :type_id, :label, :value, :remark, :sort, :user_id, :now, :user_id, :now)",
{
"id": data_id,
"type_id": dict_type_id,
"label": dict_label,
"value": dict_value,
"remark": remark,
"sort": sort,
"user_id": user_id,
"now": now,
},
)
async def update_dict_data(
self,
*,
data_id: int | None,
dict_type_id: int | None,
dict_label: str | None,
dict_value: str | None,
remark: str | None,
sort: int | None,
user_id: int,
now: datetime,
) -> int:
return await self.execute(
"UPDATE sys_dict_data SET "
"dict_type_id = CASE WHEN :has_type_id = 1 THEN :type_id ELSE dict_type_id END, "
"dict_label = CASE WHEN :has_label = 1 THEN :label ELSE dict_label END, "
"dict_value = CASE WHEN :has_value = 1 THEN :value ELSE dict_value END, "
"remark = CASE WHEN :has_remark = 1 THEN :remark ELSE remark END, "
"sort = CASE WHEN :has_sort = 1 THEN :sort ELSE sort END, updater = :user_id, update_date = :now "
"WHERE id = :id",
{
"id": data_id,
"has_type_id": int(dict_type_id is not None),
"type_id": dict_type_id,
"has_label": int(dict_label is not None),
"label": dict_label,
"has_value": int(dict_value is not None),
"value": dict_value,
"has_remark": int(remark is not None),
"remark": remark,
"has_sort": int(sort is not None),
"sort": sort,
"user_id": user_id,
"now": now,
},
)
async def dict_type_codes_for_data_ids(self, ids: list[int]) -> list[str]:
statement = text(
"SELECT DISTINCT t.dict_type FROM sys_dict_type t JOIN sys_dict_data d ON d.dict_type_id = t.id "
"WHERE d.id IN :ids"
).bindparams(bindparam("ids", expanding=True))
rows = await self.fetch_all(statement, {"ids": ids})
return [str(row["dict_type"]) for row in rows]
async def delete_dict_data(self, ids: list[int]) -> int:
statement = text("DELETE FROM sys_dict_data WHERE id IN :ids").bindparams(bindparam("ids", expanding=True))
return await self.execute(statement, {"ids": ids})
async def dict_items(self, dict_type: str) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT d.dict_label AS name, d.dict_value AS `key` FROM sys_dict_data d "
"LEFT JOIN sys_dict_type t ON d.dict_type_id = t.id "
"WHERE t.dict_type = :dict_type ORDER BY d.sort ASC",
{"dict_type": dict_type},
)
@@ -0,0 +1,71 @@
from __future__ import annotations
from collections.abc import Sequence
from typing import Any
from sqlalchemy import bindparam, text
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import Repository
class TimbreRepository(Repository):
def __init__(self, session: AsyncSession):
super().__init__(session)
async def page(
self, *, tts_model_id: str, name: str | None, offset: int, limit: int
) -> tuple[list[dict[str, Any]], int]:
where = (
"WHERE tts_model_id=:tts_model_id AND "
"(:name IS NULL OR :name='' OR name LIKE CONCAT('%', :name, '%'))"
)
params = {"tts_model_id": tts_model_id, "name": name, "offset": offset, "limit": limit}
total = int(await self.scalar(f"SELECT COUNT(*) FROM ai_tts_voice {where}", params) or 0) # noqa: S608
rows = await self.fetch_all(
f"SELECT * FROM ai_tts_voice {where} LIMIT :offset, :limit", # noqa: S608
params,
)
return rows, total
async def insert(self, values: dict[str, Any]) -> None:
await self.execute(
"INSERT INTO ai_tts_voice "
"(id, languages, name, remark, reference_audio, reference_text, sort, tts_model_id, tts_voice, "
"voice_demo, creator, create_date) VALUES (:id, :languages, :name, :remark, :reference_audio, "
":reference_text, :sort, :tts_model_id, :tts_voice, :voice_demo, :creator, :now)",
values,
)
async def update(self, values: dict[str, Any]) -> int:
return await self.execute(
"UPDATE ai_tts_voice SET languages=:languages, name=:name, remark=COALESCE(:remark, remark), "
"reference_audio=COALESCE(:reference_audio, reference_audio), "
"reference_text=COALESCE(:reference_text, reference_text), sort=:sort, "
"tts_model_id=:tts_model_id, tts_voice=:tts_voice, "
"voice_demo=COALESCE(:voice_demo, voice_demo), updater=:updater, "
"update_date=:now WHERE id=:id",
values,
)
async def delete(self, ids: Sequence[str]) -> int:
if not ids:
return 0
statement = text("DELETE FROM ai_tts_voice WHERE id IN :ids").bindparams(bindparam("ids", expanding=True))
result = await self.session.execute(statement, {"ids": list(ids)})
return int(getattr(result, "rowcount", 0) or 0)
async def voices(
self, model_id: str, name: str | None, user_id: int
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
normal = await self.fetch_all(
"SELECT id, name, voice_demo, languages FROM ai_tts_voice WHERE tts_model_id=:model_id "
"AND (:name IS NULL OR :name='' OR name LIKE CONCAT('%', :name, '%'))",
{"model_id": model_id or "", "name": name},
)
clones = await self.fetch_all(
"SELECT id, name, voice_id AS voice_demo, languages FROM ai_voice_clone "
"WHERE model_id=:model_id AND user_id=:user_id AND train_status=2",
{"model_id": model_id, "user_id": user_id},
)
return normal, clones
@@ -0,0 +1,170 @@
from __future__ import annotations
# Every interpolated SQL fragment below is a module constant or a service-side allowlist.
# ruff: noqa: S608
from collections.abc import Mapping, Sequence
from typing import Any
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import Repository
VOICE_COLUMNS = (
"id, name, model_id, voice_id, languages, user_id, voice, train_status, train_error, creator, create_date"
)
class VoiceCloneRepository(Repository):
def __init__(self, session: AsyncSession):
super().__init__(session)
async def count(self, *, name: str | None, user_id: str | None) -> int:
where, params = self._filters(name=name, user_id=user_id)
return int(await self.scalar(f"SELECT COUNT(*) FROM ai_voice_clone{where}", params) or 0)
async def page(
self,
*,
page: int,
limit: int,
name: str | None,
user_id: str | None,
order_fields: Sequence[str],
ascending: bool,
) -> list[dict[str, Any]]:
where, params = self._filters(name=name, user_id=user_id)
params.update(limit=limit, offset=max(page - 1, 0) * limit)
direction = "ASC" if ascending else "DESC"
order_by = ", ".join(f"{field} {direction}" for field in order_fields)
return await self.fetch_all(
f"SELECT {VOICE_COLUMNS} FROM ai_voice_clone{where} "
f"ORDER BY {order_by} LIMIT :limit OFFSET :offset",
params,
)
async def get(self, voice_id: str | None) -> dict[str, Any] | None:
if voice_id is None:
return None
return await self.fetch_one(
f"SELECT {VOICE_COLUMNS} FROM ai_voice_clone WHERE id = :id LIMIT 1",
{"id": voice_id},
)
async def list_by_user(self, user_id: int) -> list[dict[str, Any]]:
return await self.fetch_all(
f"SELECT {VOICE_COLUMNS} FROM ai_voice_clone "
"WHERE user_id = :user_id ORDER BY create_date DESC",
{"user_id": user_id},
)
async def voice_id_count(self, *, model_id: str, voice_id: str) -> int:
return int(
await self.scalar(
"SELECT COUNT(*) FROM ai_voice_clone WHERE voice_id = :voice_id AND model_id = :model_id",
{"model_id": model_id, "voice_id": voice_id},
)
or 0
)
async def insert_many(self, values: Sequence[Mapping[str, Any]]) -> int:
return await self.execute_many(
"INSERT INTO ai_voice_clone "
"(id, name, model_id, voice_id, languages, user_id, voice, train_status, train_error, creator, "
"create_date) VALUES (:id, :name, :model_id, :voice_id, :languages, :user_id, :voice, :train_status, "
":train_error, :creator, :create_date)",
values,
)
async def delete_many(self, ids: Sequence[str]) -> int:
if not ids:
return 0
placeholders = ", ".join(f":id_{index}" for index in range(len(ids)))
params = {f"id_{index}": value for index, value in enumerate(ids)}
return await self.execute(f"DELETE FROM ai_voice_clone WHERE id IN ({placeholders})", params)
async def update_voice(self, voice_id: str, data: bytes) -> int:
return await self.execute(
"UPDATE ai_voice_clone SET voice = :voice, train_status = 0 WHERE id = :id",
{"id": voice_id, "voice": data},
)
async def update_name(self, voice_id: str, name: str) -> int:
return await self.execute(
"UPDATE ai_voice_clone SET name = :name WHERE id = :id",
{"id": voice_id, "name": name},
)
async def update_training(
self,
voice_id: str,
*,
train_status: int,
train_error: str | None,
speaker_id: str | None = None,
) -> int:
if speaker_id is None:
return await self.execute(
"UPDATE ai_voice_clone SET train_status = :train_status, train_error = :train_error WHERE id = :id",
{"id": voice_id, "train_status": train_status, "train_error": train_error},
)
return await self.execute(
"UPDATE ai_voice_clone SET train_status = :train_status, train_error = :train_error, "
"voice_id = :speaker_id WHERE id = :id",
{
"id": voice_id,
"train_status": train_status,
"train_error": train_error,
"speaker_id": speaker_id,
},
)
async def get_model_config(self, model_id: str) -> dict[str, Any] | None:
return await self.fetch_one(
"SELECT id, model_name, config_json FROM ai_model_config WHERE id = :id LIMIT 1",
{"id": model_id},
)
async def get_model_name(self, model_id: str) -> str | None:
value = await self.scalar(
"SELECT model_name FROM ai_model_config WHERE id = :id LIMIT 1",
{"id": model_id},
)
return None if value is None else str(value)
async def get_usernames(self, user_ids: Sequence[int]) -> dict[int, str]:
if not user_ids:
return {}
unique_ids = list(dict.fromkeys(user_ids))
placeholders = ", ".join(f":user_{index}" for index in range(len(unique_ids)))
params = {f"user_{index}": value for index, value in enumerate(unique_ids)}
rows = await self.fetch_all(
f"SELECT id, username FROM sys_user WHERE id IN ({placeholders})",
params,
)
return {int(row["id"]): str(row["username"]) for row in rows}
async def get_username(self, user_id: int) -> str | None:
value = await self.scalar(
"SELECT username FROM sys_user WHERE id = :id LIMIT 1",
{"id": user_id},
)
return None if value is None else str(value)
async def get_tts_platforms(self) -> list[dict[str, Any]]:
return await self.fetch_all(
"SELECT id, model_name AS modelName FROM ai_model_config "
"WHERE model_type = 'TTS' AND JSON_EXTRACT(config_json, '$.type') = 'huoshan_double_stream'"
)
@staticmethod
def _filters(*, name: str | None, user_id: str | None) -> tuple[str, dict[str, Any]]:
clauses: list[str] = []
params: dict[str, Any] = {}
if user_id is not None and user_id.strip():
clauses.append("user_id = :user_id")
params["user_id"] = user_id
if name is not None and name.strip():
clauses.append("(name LIKE :name OR voice_id = :exact_name)")
params["name"] = f"%{name}%"
params["exact_name"] = name
return (" WHERE " + " AND ".join(clauses) if clauses else "", params)
@@ -0,0 +1,30 @@
"""HTTP routers grouped by the Java business domains."""
from fastapi import APIRouter
def application_routers() -> list[APIRouter]:
"""Return every migrated business router; imports stay explicit for coverage auditing."""
from app.routers.agent import router as agent_router
from app.routers.config import config_router
from app.routers.correctword import correctword_router
from app.routers.device import device_router
from app.routers.knowledge import knowledge_router
from app.routers.model import model_router
from app.routers.security import security_router
from app.routers.sys import sys_router
from app.routers.timbre import timbre_router
from app.routers.voiceclone import voiceclone_router
return [
security_router,
sys_router,
config_router,
agent_router,
device_router,
voiceclone_router,
model_router,
timbre_router,
correctword_router,
knowledge_router,
]
@@ -0,0 +1,386 @@
from __future__ import annotations
from typing import Annotated
from fastapi import APIRouter, BackgroundTasks, Body, Depends, Query, Request
from sqlalchemy.ext.asyncio import AsyncSession
from starlette.responses import Response
from app.core.database import get_db
from app.core.errors import ErrorCode
from app.core.responses import JavaJSONResponse, error_response, ok
from app.core.security import AuthUser, require_normal, require_super_admin
from app.schemas.agent import (
AgentChatHistoryReport,
AgentCreate,
AgentMemory,
AgentSnapshotPage,
AgentSnapshotRestore,
AgentTagAssignment,
AgentTemplate,
AgentUpdate,
AgentVoicePrintSave,
AgentVoicePrintUpdate,
)
from app.services.agent import AgentService, run_chat_summary_task
router = APIRouter(tags=["agent"])
DbSession = Annotated[AsyncSession, Depends(get_db)]
NormalUser = Annotated[AuthUser, Depends(require_normal)]
SuperUser = Annotated[AuthUser, Depends(require_super_admin)]
def _service(session: AsyncSession, user: AuthUser | None, request: Request) -> AgentService:
return AgentService(session, user, language=request.headers.get("Accept-Language"))
@router.post("/agent/chat-history/report")
async def report_chat_history(report: AgentChatHistoryReport, request: Request, session: DbSession) -> JavaJSONResponse:
return ok(await _service(session, None, request).report_chat(report))
@router.post("/agent/chat-history/getDownloadUrl/{agentId}/{sessionId}")
async def issue_chat_history_download(
agentId: str, sessionId: str, request: Request, session: DbSession, user: NormalUser
) -> JavaJSONResponse:
service = _service(session, user, request)
if not await service.has_agent_permission(agentId):
return error_response(request, 10132)
return ok(await service.issue_history_token(agentId, sessionId))
@router.get("/agent/chat-history/download/{uuid}/current")
async def download_current_chat_history(uuid: str, request: Request, session: DbSession) -> Response:
content = await _service(session, None, request).consume_history_download(uuid, previous=False)
return Response(
content.encode("utf-8"),
media_type="text/plain;charset=UTF-8",
headers={"Content-Disposition": "attachment;filename=history.txt"},
)
@router.get("/agent/chat-history/download/{uuid}/previous")
async def download_previous_chat_history(uuid: str, request: Request, session: DbSession) -> Response:
content = await _service(session, None, request).consume_history_download(uuid, previous=True)
return Response(
content.encode("utf-8"),
media_type="text/plain;charset=UTF-8",
headers={"Content-Disposition": "attachment;filename=history.txt"},
)
# Static paths are deliberately registered before /agent/{id}; Starlette resolves in declaration order.
@router.get("/agent/template/page")
async def template_page(
request: Request,
session: DbSession,
user: SuperUser,
page: int = Query(default=1),
limit: int = Query(default=10),
agentName: str | None = Query(default=None),
) -> JavaJSONResponse:
return ok(await _service(session, user, request).template_page(page, limit, agentName))
@router.post("/agent/template/batch-remove")
async def batch_delete_templates(
ids: list[str], request: Request, session: DbSession, user: SuperUser
) -> JavaJSONResponse:
deleted = await _service(session, user, request).batch_delete_templates(ids)
return (
ok("批量删除成功") if deleted else error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "批量删除模板失败")
)
@router.get("/agent/template/{id}")
async def template_detail(id: str, request: Request, session: DbSession, user: SuperUser) -> JavaJSONResponse:
result = await _service(session, user, request).template_detail(id)
return ok(result) if result is not None else error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "模板不存在")
@router.post("/agent/template")
async def create_template(
template: AgentTemplate, request: Request, session: DbSession, user: SuperUser
) -> JavaJSONResponse:
return ok(await _service(session, user, request).create_template(template))
@router.put("/agent/template")
async def update_template(
template: AgentTemplate, request: Request, session: DbSession, user: SuperUser
) -> JavaJSONResponse:
# MyBatis-Plus raises before returning a boolean when updateById receives
# an entity without its @TableId. Keep Java's generic error envelope for
# that exact input; an unknown but non-empty id still returns the controller's
# explicit "更新模板失败" message below.
if template.id is None:
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR)
updated = await _service(session, user, request).update_template(template)
return ok(template) if updated else error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "更新模板失败")
@router.delete("/agent/template/{id}")
async def delete_template(id: str, request: Request, session: DbSession, user: SuperUser) -> JavaJSONResponse:
service = _service(session, user, request)
if await service.template_detail(id) is None:
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "模板不存在")
return (
ok("删除模板成功")
if await service.delete_template(id)
else error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "删除模板失败")
)
@router.post("/agent/voice-print")
async def create_voiceprint(
dto: AgentVoicePrintSave, request: Request, session: DbSession, user: NormalUser
) -> JavaJSONResponse:
created = await _service(session, user, request).create_voiceprint(dto)
return ok() if created else error_response(request, 10057)
@router.put("/agent/voice-print")
async def update_voiceprint(
dto: AgentVoicePrintUpdate, request: Request, session: DbSession, user: NormalUser
) -> JavaJSONResponse:
updated = await _service(session, user, request).update_voiceprint(dto)
return ok() if updated else error_response(request, 10058)
@router.delete("/agent/voice-print/{id}")
async def delete_voiceprint(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
deleted = await _service(session, user, request).delete_voiceprint(id)
return ok() if deleted else error_response(request, 10059)
@router.get("/agent/voice-print/list/{id}")
async def list_voiceprints(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
return ok(await _service(session, user, request).voiceprint_list(id))
@router.get("/agent/mcp/address/{agentId}")
async def mcp_address(agentId: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
service = _service(session, user, request)
if not await service.has_agent_permission(agentId):
return error_response(request, 10200)
address = await service.mcp_address(agentId)
return ok(address) if address is not None else error_response(request, 10201)
@router.get("/agent/mcp/tools/{agentId}")
async def mcp_tools(agentId: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
service = _service(session, user, request)
if not await service.has_agent_permission(agentId):
return error_response(request, 10202)
return ok(await service.mcp_tools(agentId))
@router.get("/agent/tag/list")
async def all_tags(request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
return ok(await _service(session, user, request).all_tags())
@router.post("/agent/tag")
async def create_tag(
request: Request,
session: DbSession,
user: NormalUser,
params: dict[str, str] = Body(...),
) -> JavaJSONResponse:
tag_name = params.get("tagName")
if tag_name is None or not tag_name.strip():
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "标签名称不能为空")
return ok(await _service(session, user, request).save_tag(tag_name))
@router.delete("/agent/tag/{id}")
async def delete_tag(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
await _service(session, user, request).delete_tag(id)
return ok()
@router.post("/agent/audio/{audioId}")
async def issue_audio_token(audioId: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
token = await _service(session, user, request).issue_audio_token(audioId)
return ok(token) if token is not None else error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "音频不存在")
@router.get("/agent/play/{uuid}")
async def play_agent_audio(uuid: str, request: Request, session: DbSession) -> Response:
audio = await _service(session, None, request).consume_audio_token(uuid)
if audio is None:
return Response(status_code=404)
return Response(
audio,
media_type="application/octet-stream",
headers={"Content-Disposition": 'attachment; filename="play.wav"'},
)
@router.put("/agent/saveMemory/{macAddress}")
async def update_memory(
macAddress: str, dto: AgentMemory, request: Request, session: DbSession, user: NormalUser
) -> JavaJSONResponse:
await _service(session, user, request).update_memory_by_mac(macAddress, dto)
return ok()
@router.post("/agent/chat-summary/{sessionId}/save")
async def save_chat_summary(
sessionId: str, background_tasks: BackgroundTasks, request: Request, session: DbSession
) -> JavaJSONResponse:
await _service(session, None, request).session_agent(sessionId)
background_tasks.add_task(run_chat_summary_task, sessionId)
return ok()
@router.post("/agent/chat-title/{sessionId}/generate")
async def generate_chat_title(sessionId: str, request: Request, session: DbSession) -> JavaJSONResponse:
service = _service(session, None, request)
await service.session_agent(sessionId)
await service.generate_chat_title(sessionId)
return ok()
@router.get("/agent/all")
async def admin_agent_list(
request: Request,
session: DbSession,
user: SuperUser,
page: int = Query(default=1),
limit: int = Query(default=10),
orderField: str | None = Query(default=None),
order: str | None = Query(default=None),
) -> JavaJSONResponse:
return ok(await _service(session, user, request).admin_agents(page, limit, orderField, order))
@router.get("/agent/list")
async def user_agent_list(
request: Request,
session: DbSession,
user: NormalUser,
keyword: str | None = Query(default=None),
searchType: str = Query(default="name"),
) -> JavaJSONResponse:
del searchType # Java accepts the parameter but the consolidated implementation ignores it.
return ok(await _service(session, user, request).user_agents(keyword))
@router.get("/agent/template")
async def template_list(request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
return ok(await _service(session, user, request).templates())
@router.get("/agent/{agentId}/snapshots")
async def snapshot_page(
agentId: str,
request: Request,
session: DbSession,
user: NormalUser,
page: int | None = Query(default=1),
limit: int | None = Query(default=10),
maxVersionNo: int | None = Query(default=None),
) -> JavaJSONResponse:
params = AgentSnapshotPage(page=page, limit=limit, max_version_no=maxVersionNo)
return ok(await _service(session, user, request).snapshot_page(agentId, params))
@router.get("/agent/{agentId}/snapshots/{snapshotId}")
async def snapshot_detail(
agentId: str, snapshotId: str, request: Request, session: DbSession, user: NormalUser
) -> JavaJSONResponse:
return ok(await _service(session, user, request).snapshot_detail(agentId, snapshotId))
@router.post("/agent/{agentId}/snapshots/{snapshotId}/restore")
async def restore_snapshot(
agentId: str,
snapshotId: str,
dto: AgentSnapshotRestore,
request: Request,
session: DbSession,
user: NormalUser,
) -> JavaJSONResponse:
await _service(session, user, request).restore_snapshot(agentId, snapshotId, dto.current_state_token)
return ok()
@router.delete("/agent/{agentId}/snapshots/{snapshotId}")
async def delete_snapshot(
agentId: str, snapshotId: str, request: Request, session: DbSession, user: NormalUser
) -> JavaJSONResponse:
await _service(session, user, request).delete_snapshot(agentId, snapshotId)
return ok()
@router.get("/agent/{id}/sessions")
async def agent_sessions(
id: str,
request: Request,
session: DbSession,
user: NormalUser,
page: str | None = Query(default=None),
limit: str | None = Query(default=None),
) -> JavaJSONResponse:
return ok(await _service(session, user, request).sessions(id, page, limit))
@router.get("/agent/{id}/chat-history/user")
async def recent_agent_history(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
service = _service(session, user, request)
if not await service.has_agent_permission(id):
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "没有权限查看该智能体的聊天记录")
return ok(await service.recent_user_history(id))
@router.get("/agent/{id}/chat-history/audio")
async def agent_audio_content(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
return ok(await _service(session, user, request).audio_content(id))
@router.get("/agent/{id}/chat-history/{sessionId}")
async def agent_history(
id: str, sessionId: str, request: Request, session: DbSession, user: NormalUser
) -> JavaJSONResponse:
service = _service(session, user, request)
if not await service.has_agent_permission(id):
return error_response(request, ErrorCode.INTERNAL_SERVER_ERROR, "没有权限查看该智能体的聊天记录")
return ok(await service.history(id, sessionId))
@router.get("/agent/{id}/tags")
async def agent_tags(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
return ok(await _service(session, user, request).agent_tags(id))
@router.put("/agent/{id}/tags")
async def save_agent_tags(
id: str, dto: AgentTagAssignment, request: Request, session: DbSession, user: NormalUser
) -> JavaJSONResponse:
await _service(session, user, request).save_agent_tags(id, dto.tag_ids, dto.tag_names)
return ok()
@router.post("/agent")
async def create_agent(dto: AgentCreate, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
return ok(await _service(session, user, request).create_agent(dto))
@router.put("/agent/{id}")
async def update_agent(
id: str, dto: AgentUpdate, request: Request, session: DbSession, user: NormalUser
) -> JavaJSONResponse:
await _service(session, user, request).update_agent(id, dto)
return ok()
@router.delete("/agent/{id}")
async def delete_agent(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
await _service(session, user, request).delete_agent(id)
return ok()
@router.get("/agent/{id}")
async def agent_detail(id: str, request: Request, session: DbSession, user: NormalUser) -> JavaJSONResponse:
return ok(await _service(session, user, request).agent_detail(id))
@@ -0,0 +1,34 @@
# ruff: noqa: B008
# FastAPI evaluates dependency marker defaults intentionally when registering routes.
from __future__ import annotations
from fastapi import APIRouter, Depends
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
from app.core.responses import JavaJSONResponse, ok
from app.core.serialization import preserve_java_map_keys
from app.repositories.config import ConfigRepository
from app.schemas.config import AgentModelsRequest, CorrectWordsRequest
from app.services.config import ConfigService
config_router = APIRouter()
def _service(session: AsyncSession) -> ConfigService:
return ConfigService(ConfigRepository(session))
@config_router.post("/config/server-base")
async def server_base(session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
return ok(preserve_java_map_keys(await _service(session).get_config(use_cache=True)))
@config_router.post("/config/agent-models")
async def agent_models(dto: AgentModelsRequest, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
return ok(preserve_java_map_keys(await _service(session).get_agent_models(dto.mac_address, dto.selected_module)))
@config_router.post("/config/correct-words")
async def correct_words(dto: CorrectWordsRequest, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
return ok(await _service(session).get_correct_words(dto.mac_address))
@@ -0,0 +1,97 @@
from __future__ import annotations
from urllib.parse import quote
from fastapi import APIRouter, Depends, Request
from sqlalchemy.ext.asyncio import AsyncSession
from starlette.responses import Response
from app.core.database import get_db
from app.core.responses import JavaJSONResponse, ok
from app.core.security import require_normal
from app.repositories.correctword import CorrectWordRepository
from app.schemas.correctword import CorrectWordFileBody
from app.services.correctword import CorrectWordService
correctword_router = APIRouter()
def _java_urlencode(value: str) -> str:
# java.net.URLEncoder leaves alphanumerics plus .-*_ unescaped, encodes
# spaces as '+', and encodes '~'. The controller then replaces '+' with
# '%20'. urllib always leaves '~', so handle that final difference here.
return quote(value, safe="*.-_").replace("~", "%7E")
def _service(session: AsyncSession) -> CorrectWordService:
return CorrectWordService(CorrectWordRepository(session))
@correctword_router.post("/correct-word/file")
async def create_file(
body: CorrectWordFileBody, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
return ok(await _service(session).create(body, require_normal(request)))
@correctword_router.put("/correct-word/file/{file_id}")
async def update_file(
file_id: str,
body: CorrectWordFileBody,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
await _service(session).update(file_id, body, require_normal(request))
return ok()
@correctword_router.get("/correct-word/file/list")
async def list_files(
request: Request,
page: str | None = None,
limit: str | None = None,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
return ok(await _service(session).page(require_normal(request), page, limit))
@correctword_router.get("/correct-word/file/select")
async def select_files(request: Request, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
return ok(await _service(session).all(require_normal(request)))
@correctword_router.get("/correct-word/file/download/{file_id}")
async def download_file(
file_id: str, request: Request, session: AsyncSession = Depends(get_db)
) -> Response:
require_normal(request)
item = await _service(session).get(file_id)
if item is None or not item["content"]:
return Response(status_code=404)
body = "\n".join(item["content"]).encode("utf-8")
file_name = str(item["fileName"])
ascii_name = "".join(character if ord(character) < 128 else "_" for character in file_name)
disposition = f"attachment; filename=\"{ascii_name}\"; filename*=UTF-8''{_java_urlencode(file_name)}"
return Response(
body,
media_type="application/octet-stream",
headers={"Content-Disposition": disposition, "Content-Length": str(len(body))},
)
@correctword_router.delete("/correct-word/file/{file_id}")
async def delete_file(
file_id: str, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
require_normal(request)
await _service(session).delete([file_id])
return ok()
@correctword_router.post("/correct-word/file/batch-delete")
async def batch_delete_files(
file_ids: list[str], request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
require_normal(request)
await _service(session).delete(file_ids)
return ok()
@@ -0,0 +1,505 @@
from __future__ import annotations
import json
from typing import Annotated, Any
from fastapi import APIRouter, BackgroundTasks, Depends, File, Header, Query, Request, UploadFile
from sqlalchemy.ext.asyncio import AsyncSession
from starlette.responses import Response
from app.core.database import get_db
from app.core.i18n import resolve_language
from app.core.responses import JavaJSONResponse, envelope, error_response, ok
from app.core.security import require_normal, require_super_admin
from app.schemas.device import (
DeviceAddressBookAliasRequest,
DeviceAddressBookPermissionRequest,
DeviceManualAddRequest,
DeviceRegisterRequest,
DeviceReportRequest,
DeviceToolCallRequest,
DeviceUnbindRequest,
DeviceUpdateRequest,
OtaRecord,
)
from app.services.device import MAC_PATTERN, DeviceService, is_blank
device_router = APIRouter()
SessionDep = Annotated[AsyncSession, Depends(get_db)]
FirmwareUpload = Annotated[UploadFile, File()]
CallerMacQuery = Annotated[str, Query(alias="callerMac")]
DeviceIdHeader = Annotated[str | None, Header(alias="Device-Id")]
ClientIdHeader = Annotated[str | None, Header(alias="Client-Id")]
def _query_map(request: Request) -> dict[str, Any]:
result: dict[str, Any] = {}
for key, value in request.query_params.multi_items():
if key in result:
previous = result[key]
result[key] = [*previous, value] if isinstance(previous, list) else [previous, value]
else:
result[key] = value
return result
def _raw_ota(payload: dict[str, Any]) -> Response:
body = json.dumps(payload, ensure_ascii=False, separators=(",", ":")).encode("utf-8")
return Response(
body,
status_code=200,
media_type="application/json",
headers={"Content-Length": str(len(body))},
)
@device_router.post("/device/bind/{agent_id}/{device_code}")
async def bind_device(
agent_id: str,
device_code: str,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
await DeviceService(session).activate_bound_device(agent_id=agent_id, activation_code=device_code, user=user)
return ok()
@device_router.post("/device/register")
async def register_device(
body: DeviceRegisterRequest,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
require_normal(request)
if is_blank(body.mac_address):
return error_response(request, 10175)
return ok(await DeviceService(session).register_device(body.mac_address or ""))
@device_router.get("/device/bind/{agent_id}")
async def get_bound_devices(
agent_id: str,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
return ok(await DeviceService(session).list_user_devices(user.id, agent_id))
@device_router.post("/device/bind/{agent_id}")
async def device_online(
agent_id: str,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
await request.body()
try:
return ok(await DeviceService(session).get_online_data(agent_id, user))
except Exception as exc:
return error_response(request, 500, f"转发请求失败: {exc}")
@device_router.post("/device/unbind")
async def unbind_device(
body: DeviceUnbindRequest,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
# DeviceController does not apply @Valid to DeviceUnBindDTO. An empty
# object reaches the service with a null id and is a successful no-op.
await DeviceService(session).unbind(user_id=user.id, device_id=body.device_id or "")
return ok()
@device_router.put("/device/update/{device_id}")
async def update_device(
device_id: str,
body: DeviceUpdateRequest,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
validation = _validate_device_update(body, request.headers.get("Accept-Language"))
if validation is not None:
return error_response(request, 10034, validation)
if not await DeviceService(session).update_device(device_id=device_id, request=body, user=user):
return error_response(request, 500, "设备不存在")
return ok()
@device_router.put("/user/configDevice/{device_id}")
async def configure_device(
device_id: str,
body: DeviceUpdateRequest,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
validation = _validate_device_update(body, request.headers.get("Accept-Language"))
if validation is not None:
return error_response(request, 10034, validation)
if not await DeviceService(session).update_device(device_id=device_id, request=body, user=user):
return error_response(request, 500, "设备不存在")
return ok()
@device_router.post("/device/manual-add")
async def manual_add_device(
body: DeviceManualAddRequest,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
await DeviceService(session).manual_add(request=body, user=user)
return ok()
@device_router.post("/device/tools/list/{device_id}")
async def list_device_tools(
device_id: str,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
tools = await DeviceService(session).get_tools(device_id=device_id, user=user)
if tools is None:
return error_response(request, 10194)
return ok(tools)
@device_router.post("/device/tools/call/{device_id}")
async def call_device_tool(
device_id: str,
body: DeviceToolCallRequest,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
if is_blank(body.name):
return error_response(request, 10034, "工具名称不能为空")
result = await DeviceService(session).call_tool(
device_id=device_id,
tool_name=body.name or "",
arguments=body.arguments,
user=user,
)
if result is None:
return error_response(request, 10194)
return JavaJSONResponse(envelope(result, msg="Tools called successfully"))
# Static address-book paths deliberately precede /address-book/{mac_address}.
@device_router.get("/device/address-book/call")
async def call_address_book(
request: Request,
session: SessionDep,
caller_mac: CallerMacQuery,
nickname: str,
answer: bool = False,
) -> JavaJSONResponse:
result = await DeviceService(session).call_by_nickname(
caller_mac=caller_mac,
nickname=nickname,
answer=answer,
)
return ok(result)
@device_router.get("/device/address-book/lookup")
async def lookup_address_book(
request: Request,
session: SessionDep,
caller_mac: CallerMacQuery,
nickname: str,
) -> JavaJSONResponse:
result = await DeviceService(session).lookup_address_book(caller_mac=caller_mac, nickname=nickname)
if result is None:
return error_response(request, 500, "未找到对应设备")
return ok(result)
@device_router.put("/device/address-book/alias")
async def update_address_alias(
body: DeviceAddressBookAliasRequest,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
if is_blank(body.target_mac):
return error_response(request, 10034, "目标MAC地址不能为空")
if is_blank(body.mac_address):
return error_response(request, 10034, "MAC地址不能为空")
service = DeviceService(session)
caller = await service.repository.get_device_by_mac(body.mac_address or "")
if caller is None or int(caller.get("user_id") or -1) != user.id:
return error_response(request, 500, "无权限操作该设备")
await service.save_address_book(
mac_address=body.mac_address or "",
target_mac=body.target_mac or "",
alias=body.alias,
has_permission=None,
actor=user.id,
)
return ok()
@device_router.put("/device/address-book/permission")
async def update_address_permission(
body: DeviceAddressBookPermissionRequest,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
if is_blank(body.mac_address):
return error_response(request, 10034, "MAC地址不能为空")
if is_blank(body.target_mac):
return error_response(request, 10034, "目标MAC地址不能为空")
service = DeviceService(session)
caller = await service.repository.get_device_by_mac(body.mac_address or "")
if caller is None or int(caller.get("user_id") or -1) != user.id:
return error_response(request, 500, "无权限操作该设备")
await service.save_address_book(
mac_address=body.mac_address or "",
target_mac=body.target_mac or "",
alias=None,
has_permission=body.has_permission,
actor=user.id,
)
return ok()
@device_router.get("/device/address-book/{mac_address}")
async def get_address_book(
mac_address: str,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
require_normal(request)
return ok(await DeviceService(session).address_book(mac_address))
@device_router.post("/ota/")
async def check_ota_version(
report: DeviceReportRequest,
request: Request,
session: SessionDep,
background_tasks: BackgroundTasks,
device_id: DeviceIdHeader = None,
client_id: ClientIdHeader = None,
) -> Response:
if is_blank(device_id):
# Java's required @RequestHeader fails before the controller's blank
# guard and is translated by its global handler into this envelope.
return error_response(request, 500)
if MAC_PATTERN.fullmatch(device_id or "") is None:
return _raw_ota({"error": "Invalid device ID"})
selected_client = device_id if is_blank(client_id) else client_id
client_ip = request.client.host if request.client is not None else "unknown"
service = DeviceService(session)
def defer_connection_update(device: str, agent: str | None, version: str | None) -> None:
background_tasks.add_task(
DeviceService.persist_connection_update_background,
device,
agent,
version,
)
payload = await service.check_ota(
device_id=device_id or "",
client_id=selected_client or device_id or "",
report=report,
request_url=str(request.url),
client_ip=client_ip,
defer_connection_update=defer_connection_update,
)
return _raw_ota(payload)
@device_router.post("/ota/activate")
async def activate_ota_device(
request: Request,
session: SessionDep,
device_id: DeviceIdHeader = None,
client_id: ClientIdHeader = None,
) -> Response:
del client_id
if is_blank(device_id):
return error_response(request, 500)
if await DeviceService(session).repository.get_device_by_mac(device_id or "") is None:
return Response(status_code=202)
return Response("success", media_type="text/plain;charset=UTF-8")
@device_router.get("/ota/")
async def ota_health(session: SessionDep) -> Response:
return Response(
await DeviceService(session).ota_health_text(),
media_type="text/plain;charset=UTF-8",
)
# Static otaMag paths deliberately precede /otaMag/{id}.
@device_router.get("/otaMag/getDownloadUrl/{ota_id}")
async def get_ota_download_url(
ota_id: str,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await DeviceService(session).create_ota_download_id(ota_id))
@device_router.get("/otaMag/download/{download_id}")
async def download_ota(download_id: str, session: SessionDep) -> Response:
resolved = await DeviceService(session).resolve_ota_download(download_id)
if resolved is None:
return Response(status_code=404)
path, filename = resolved
try:
content = path.read_bytes()
except OSError:
return Response(status_code=500)
return Response(
content,
media_type="application/octet-stream",
headers={
"Content-Disposition": f'attachment; filename="{filename}"',
"Content-Length": str(len(content)),
},
)
@device_router.post("/otaMag/upload")
async def upload_firmware(
request: Request,
file: FirmwareUpload,
session: SessionDep,
) -> JavaJSONResponse:
require_super_admin(request)
service = DeviceService(session)
try:
content = await file.read()
return ok(await service.save_firmware_file(filename=file.filename, content=content))
except ValueError as exc:
return error_response(request, 500, str(exc))
except OSError as exc:
return error_response(request, 500, f"文件上传失败:{exc}")
@device_router.post("/otaMag/uploadAssetsBin")
async def upload_assets_firmware(
request: Request,
file: FirmwareUpload,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
service = DeviceService(session)
try:
content = await file.read()
return ok(await service.save_assets_file(filename=file.filename, content=content, user=user))
except ValueError as exc:
return error_response(request, 500, str(exc))
except OSError as exc:
return error_response(request, 500, f"文件上传失败:{exc}")
@device_router.get("/otaMag")
async def page_ota(
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await DeviceService(session).ota_page(_query_map(request)))
@device_router.get("/otaMag/{ota_id}")
async def get_ota(
ota_id: str,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await DeviceService(session).get_ota_record(ota_id))
@device_router.post("/otaMag")
async def save_ota(
request: Request,
session: SessionDep,
record: OtaRecord | None = None,
) -> JavaJSONResponse:
user = require_super_admin(request)
if record is None:
return error_response(request, 500, "固件信息不能为空")
if is_blank(record.firmware_name):
return error_response(request, 500, "固件名称不能为空")
if is_blank(record.type):
return error_response(request, 500, "固件类型不能为空")
if is_blank(record.version):
return error_response(request, 500, "版本号不能为空")
try:
await DeviceService(session).save_ota(record, user)
return ok()
except RuntimeError as exc:
return error_response(request, 500, str(exc))
@device_router.delete("/otaMag/{ota_id}")
async def delete_ota(
ota_id: str,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
require_super_admin(request)
ids = ota_id.split(",") if ota_id else []
if not ids:
return error_response(request, 500, "删除的固件ID不能为空")
await DeviceService(session).delete_ota(ids)
return ok()
@device_router.put("/otaMag/{ota_id}")
async def update_ota(
ota_id: str,
request: Request,
session: SessionDep,
record: OtaRecord | None = None,
) -> JavaJSONResponse:
user = require_super_admin(request)
if record is None:
return error_response(request, 500, "固件信息不能为空")
try:
await DeviceService(session).update_ota(ota_id, record, user)
return ok()
except RuntimeError as exc:
return error_response(request, 500, str(exc))
def _validate_device_update(body: DeviceUpdateRequest, accept_language: str | None) -> str | None:
language = resolve_language(accept_language)
if body.auto_update is not None and body.auto_update < 0:
return {
"zh-CN": "最小不能小于0",
"zh-TW": "必須大於或等於 0",
"de-DE": "muss größer-gleich 0 sein",
"pt-BR": "deve ser maior que ou igual à 0",
}.get(language, "must be greater than or equal to 0")
if body.auto_update is not None and body.auto_update > 1:
return {
"zh-CN": "最大不能超过1",
"zh-TW": "必須小於或等於 1",
"de-DE": "muss kleiner-gleich 1 sein",
"pt-BR": "deve ser menor que ou igual à 1",
}.get(language, "must be less than or equal to 1")
if body.alias is not None and len(body.alias.encode("utf-16-le")) // 2 > 64:
return {
"zh-CN": "个数必须在0和64之间",
"zh-TW": "大小必須在 0 和 64 之間",
"de-DE": "Größe muss zwischen 0 und 64 sein",
"pt-BR": "tamanho deve ser entre 0 e 64",
}.get(language, "size must be between 0 and 64")
return None
@@ -0,0 +1,267 @@
from __future__ import annotations
import json
from typing import Annotated, Any
from fastapi import APIRouter, Depends, File, Form, Query, Request, UploadFile
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
from app.core.errors import AppError
from app.core.responses import JavaJSONResponse, envelope, ok
from app.core.security import require_normal
from app.repositories.knowledge import KnowledgeRepository
from app.schemas.knowledge import DocumentBatchBody, KnowledgeBaseBody, RetrievalBody
from app.services.knowledge import KnowledgeBaseService, KnowledgeDocumentService, dataset_dto
knowledge_router = APIRouter()
def _base(session: AsyncSession) -> KnowledgeBaseService:
return KnowledgeBaseService(KnowledgeRepository(session))
def _documents(session: AsyncSession) -> KnowledgeDocumentService:
return KnowledgeDocumentService(KnowledgeRepository(session))
@knowledge_router.get("/datasets/rag-models")
async def rag_models(request: Request, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
require_normal(request)
return ok(await _base(session).rag_models())
@knowledge_router.delete("/datasets/batch")
async def delete_datasets_batch(
request: Request, ids: str = Query(), session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
user = require_normal(request)
if not ids.strip():
raise AppError(10003)
await _base(session).batch_delete(
ids.split(","), user, request.headers.get("Accept-Language")
)
return ok()
@knowledge_router.get("/datasets")
async def datasets_page(
request: Request,
name: str | None = None,
page: int = 1,
page_size: int = 10,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
return ok(
await _base(session).page(
require_normal(request),
name,
page,
page_size,
request.headers.get("Accept-Language"),
)
)
@knowledge_router.post("/datasets")
async def create_dataset(
body: KnowledgeBaseBody, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
return ok(await _base(session).create(body, require_normal(request)))
@knowledge_router.get("/datasets/{dataset_id}")
async def get_dataset(
dataset_id: str, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
return ok(dataset_dto(await _base(session).get_owned(dataset_id, require_normal(request))))
@knowledge_router.put("/datasets/{dataset_id}")
async def update_dataset(
dataset_id: str,
body: KnowledgeBaseBody,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
return ok(await _base(session).update(dataset_id, body, require_normal(request)))
@knowledge_router.delete("/datasets/{dataset_id}")
async def delete_dataset(
dataset_id: str, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
await _base(session).delete(
dataset_id, require_normal(request), request.headers.get("Accept-Language")
)
return ok()
@knowledge_router.get("/datasets/{dataset_id}/documents/status/{status}")
async def documents_by_status(
dataset_id: str,
status: str,
request: Request,
page: int = 1,
page_size: int = 10,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
return ok(
await _documents(session).page(
dataset_id,
require_normal(request),
name=None,
status=status,
page=page,
page_size=page_size,
)
)
@knowledge_router.get("/datasets/{dataset_id}/documents")
async def documents_page(
dataset_id: str,
request: Request,
name: str | None = None,
status: str | None = None,
page: int = 1,
page_size: int = 10,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
return ok(
await _documents(session).page(
dataset_id,
require_normal(request),
name=name,
status=status,
page=page,
page_size=page_size,
)
)
@knowledge_router.post("/datasets/{dataset_id}/documents")
async def upload_document(
dataset_id: str,
request: Request,
file: Annotated[UploadFile, File()],
name: Annotated[str | None, Form()] = None,
chunk_method: Annotated[str | None, Form(alias="chunkMethod")] = None,
meta_fields: Annotated[str | None, Form(alias="metaFields")] = None,
parser_config: Annotated[str | None, Form(alias="parserConfig")] = None,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
return ok(
await _documents(session).upload(
dataset_id,
require_normal(request),
file,
name=name,
meta_fields=_parse_form_json(meta_fields),
chunk_method=chunk_method,
parser_config=_parse_form_json(parser_config),
)
)
@knowledge_router.delete("/datasets/{dataset_id}/documents")
async def delete_documents(
dataset_id: str,
body: DocumentBatchBody,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
await _documents(session).delete(
dataset_id,
body.ids,
require_normal(request),
request.headers.get("Accept-Language"),
)
return ok()
@knowledge_router.delete("/datasets/{dataset_id}/documents/{document_id}")
async def delete_document(
dataset_id: str,
document_id: str,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
await _documents(session).delete(
dataset_id,
[document_id],
require_normal(request),
request.headers.get("Accept-Language"),
)
return ok()
@knowledge_router.post("/datasets/{dataset_id}/chunks")
async def parse_documents(
dataset_id: str,
body: dict[str, Any],
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
user = require_normal(request)
# Java validates dataset existence/ownership before it reads document_ids.
# A missing dataset must therefore win over the controller's empty-body
# business error.
await _base(session).get_owned(dataset_id, user)
document_ids = body.get("document_ids")
if document_ids is not None and not isinstance(document_ids, list):
# Spring fails Map<String,List<String>> deserialization before entering
# the controller, which is handled as the generic code-500 envelope.
raise RuntimeError("document_ids must be an array")
if not document_ids:
return JavaJSONResponse(envelope(None, code=500, msg="document_ids参数不能为空"))
success = await _documents(session).parse(dataset_id, document_ids, user)
return ok() if success else JavaJSONResponse(
envelope(None, code=500, msg="文档解析失败,文档可能正在处理中")
)
@knowledge_router.get("/datasets/{dataset_id}/documents/{document_id}/chunks")
async def list_chunks(
dataset_id: str,
document_id: str,
request: Request,
page: int = 1,
page_size: int = 10,
keywords: str | None = None,
id: str | None = None, # noqa: A002 - exact Java query parameter
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
return ok(
await _documents(session).chunks(
dataset_id,
document_id,
require_normal(request),
page=page,
page_size=page_size,
keywords=keywords,
chunk_id=id,
)
)
@knowledge_router.post("/datasets/{dataset_id}/retrieval-test")
async def retrieval_test(
dataset_id: str,
body: RetrievalBody,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
return ok(await _documents(session).retrieval(dataset_id, body, require_normal(request)))
def _parse_form_json(value: str | None) -> dict[str, Any] | None:
if value is None:
return None
try:
result = json.loads(value)
except json.JSONDecodeError as exc:
raise RuntimeError(f"解析JSON字符串失败: {value}") from exc
if not isinstance(result, dict):
raise RuntimeError(f"解析JSON字符串失败: {value}")
return dict(result)
@@ -0,0 +1,178 @@
from __future__ import annotations
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
from app.core.responses import JavaJSONResponse, envelope, ok
from app.core.security import require_normal, require_super_admin
from app.repositories.config import ConfigRepository
from app.repositories.model import ModelRepository
from app.schemas.model import ModelConfigBody, ModelProviderBody
from app.services.config import ConfigService
from app.services.model import ModelProviderService, ModelService
model_router = APIRouter()
def _models(session: AsyncSession) -> ModelService:
return ModelService(ModelRepository(session))
def _providers(session: AsyncSession) -> ModelProviderService:
return ModelProviderService(ModelRepository(session))
async def _refresh_server_config(session: AsyncSession) -> None:
await ConfigService(ConfigRepository(session)).get_config(use_cache=False)
@model_router.get("/models/names")
async def model_names(
request: Request,
model_type: str = Query(alias="modelType"),
model_name: str | None = Query(default=None, alias="modelName"),
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_normal(request)
return ok(await _models(session).names(model_type, model_name))
@model_router.get("/models/llm/names")
async def llm_names(
request: Request,
model_name: str | None = Query(default=None, alias="modelName"),
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_normal(request)
return ok(await _models(session).llm_names(model_name))
@model_router.get("/models/list")
async def model_list(
request: Request,
model_type: str = Query(alias="modelType"),
model_name: str | None = Query(default=None, alias="modelName"),
page: str = "0",
limit: str = "10",
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await _models(session).model_page(model_type, model_name, page, limit))
@model_router.get("/models/provider/plugin/names")
async def plugin_names(request: Request, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
user = require_normal(request)
return ok(await ModelRepository(session).list_plugins_for_user(user.id))
@model_router.get("/models/provider")
async def provider_list(
request: Request,
model_type: str | None = Query(default=None, alias="modelType"),
name: str | None = None,
page: str = "0",
limit: str = "10",
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await _providers(session).page(model_type, name, page, limit))
@model_router.post("/models/provider")
async def provider_add(
body: ModelProviderBody, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
return ok(await _providers(session).add(body, require_super_admin(request)))
@model_router.put("/models/provider")
async def provider_edit(
body: ModelProviderBody, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
return ok(await _providers(session).edit(body, require_super_admin(request)))
@model_router.post("/models/provider/delete")
async def provider_delete(
ids: list[str], request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
require_super_admin(request)
await _providers(session).delete(ids)
return ok()
@model_router.get("/models/{model_type}/provideTypes")
async def provider_types(
model_type: str, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await ModelRepository(session).list_providers_by_type(model_type))
@model_router.post("/models/{model_type}/{provide_code}")
async def model_add(
model_type: str,
provide_code: str,
body: ModelConfigBody,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
result = await _models(session).add(model_type, provide_code, body)
await _refresh_server_config(session)
return ok(result)
@model_router.put("/models/enable/{model_id}/{status}")
async def model_enable(
model_id: str, status: int, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
require_super_admin(request)
message = await _models(session).enable(model_id, status)
return JavaJSONResponse(envelope(None, code=500, msg=message)) if message else ok()
@model_router.put("/models/{model_type}/{provide_code}/{model_id}")
async def model_edit(
model_type: str,
provide_code: str,
model_id: str,
body: ModelConfigBody,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
result = await _models(session).edit(model_type, provide_code, model_id, body)
await _refresh_server_config(session)
return ok(result)
@model_router.put("/models/default/{model_id}")
async def model_default(
model_id: str, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
require_super_admin(request)
message = await _models(session).set_default(model_id)
if message:
return JavaJSONResponse(envelope(None, code=500, msg=message))
await _refresh_server_config(session)
return ok()
@model_router.get("/models/{model_id}")
async def model_get(
model_id: str, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await _models(session).get_model(model_id))
@model_router.delete("/models/{model_id}")
async def model_delete(
model_id: str, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
require_super_admin(request)
await _models(session).delete(model_id)
return ok()
@@ -0,0 +1,111 @@
# ruff: noqa: B008
# FastAPI evaluates dependency marker defaults intentionally when registering routes.
from __future__ import annotations
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy.ext.asyncio import AsyncSession
from starlette.responses import Response
from app.core.database import get_db
from app.core.errors import AppError
from app.core.responses import JavaJSONResponse, ok
from app.core.security import require_normal
from app.repositories.security import SecurityRepository
from app.schemas.security import (
LoginRequest,
PasswordChangeRequest,
RetrievePasswordRequest,
SmsVerificationRequest,
)
from app.services.security import CaptchaService, SecurityService
security_router = APIRouter()
def _service(session: AsyncSession) -> SecurityService:
return SecurityService(SecurityRepository(session))
@security_router.get("/user/captcha")
async def captcha(uuid: str | None = Query(default=None)) -> Response:
if uuid is None or not uuid.strip():
raise AppError(10006)
content = await CaptchaService().create(uuid)
return Response(
content,
media_type="image/gif",
headers={
"Pragma": "No-cache",
"Cache-Control": "no-cache",
"Expires": "Thu, 01 Jan 1970 00:00:00 GMT",
},
)
@security_router.post("/user/smsVerification")
async def sms_verification(dto: SmsVerificationRequest, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
await _service(session).send_sms_verification(dto)
return ok()
@security_router.post("/user/login")
async def login(
dto: LoginRequest,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
return ok(await _service(session).login(dto, request))
@security_router.post("/user/register")
async def register(dto: LoginRequest, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
await _service(session).register(dto)
return ok()
@security_router.get("/user/info")
async def info(request: Request) -> JavaJSONResponse:
user = require_normal(request)
return ok(
{
"id": user.id,
"username": user.username,
"superAdmin": user.super_admin,
"token": user.token,
"status": user.status,
}
)
@security_router.put("/user/change-password")
async def change_password(
dto: PasswordChangeRequest,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
await _service(session).change_password(
require_normal(request),
dto,
request.headers.get("Accept-Language"),
)
return ok()
@security_router.put("/user/retrieve-password")
async def retrieve_password(
dto: RetrievePasswordRequest,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
await _service(session).retrieve_password(dto, request.headers.get("Accept-Language"))
return ok()
@security_router.get("/user/pub-config")
async def public_config(session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
return ok(await _service(session).public_config())
@security_router.get("/api/ping")
async def api_ping() -> JavaJSONResponse:
return ok("pong")
+334
View File
@@ -0,0 +1,334 @@
# ruff: noqa: B008
# FastAPI evaluates dependency and body marker defaults intentionally when registering routes.
from __future__ import annotations
from fastapi import APIRouter, Body, Depends, Query, Request
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
from app.core.errors import AppError
from app.core.responses import JavaJSONResponse, envelope, ok
from app.core.security import require_normal, require_super_admin
from app.repositories.sys import SysRepository
from app.schemas.sys import DictDataPayload, DictTypePayload, EmitServerActionRequest, SysParamPayload
from app.services.sys import AdminService, DictService, ServerActionService, SysParamService
sys_router = APIRouter()
def _repository(session: AsyncSession) -> SysRepository:
return SysRepository(session)
def _admin(session: AsyncSession) -> AdminService:
return AdminService(_repository(session))
def _params(session: AsyncSession) -> SysParamService:
return SysParamService(_repository(session))
def _dict(session: AsyncSession) -> DictService:
return DictService(_repository(session))
async def _refresh_server_config(session: AsyncSession) -> None:
from app.repositories.config import ConfigRepository
from app.services.config import ConfigService
await ConfigService(ConfigRepository(session)).get_config(use_cache=False)
@sys_router.get("/admin/users")
async def page_users(
request: Request,
mobile: str | None = None,
page: str = Query(default="1"),
limit: str = Query(default="10"),
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
try:
current, size = int(page), int(limit)
except ValueError as exc:
# Java parses these Map-backed values inside the service; malformed
# numbers therefore reach its generic code=500 handler rather than
# Bean Validation.
raise AppError(500, "排序值不能小于0") from exc
return ok(await _admin(session).page_users(mobile=mobile, page=current, limit=size))
@sys_router.put("/admin/users/{user_id}")
async def reset_user_password(
user_id: int,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
user = require_super_admin(request)
return ok(await _admin(session).reset_password(user_id, user))
@sys_router.delete("/admin/users/{user_id}")
async def delete_user(
user_id: int,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
await _admin(session).delete_user(user_id)
return ok()
@sys_router.put("/admin/users/changeStatus/{status}")
async def change_user_status(
status: int,
request: Request,
user_ids: list[str] = Body(),
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
user = require_super_admin(request)
await _admin(session).change_status(status, user_ids, user)
return ok()
@sys_router.get("/admin/device/all")
async def page_all_devices(
request: Request,
keywords: str | None = None,
page: int = Query(default=1, ge=0),
limit: int = Query(default=10, ge=0),
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await _admin(session).page_devices(keywords=keywords, page=page, limit=limit))
@sys_router.get("/admin/server/server-list")
async def websocket_server_list(request: Request, session: AsyncSession = Depends(get_db)) -> JavaJSONResponse:
require_super_admin(request)
params = _params(session)
return ok(await ServerActionService(params).server_list())
@sys_router.post("/admin/server/emit-action")
async def emit_server_action(
dto: EmitServerActionRequest,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await ServerActionService(_params(session)).emit(dto))
@sys_router.get("/admin/params/page")
async def page_params(
request: Request,
page: int = Query(default=1, ge=0),
limit: int = Query(default=10, ge=0),
order_field: str | None = Query(default=None, alias="orderField"),
order: str | None = None,
param_code: str | None = Query(default=None, alias="paramCode"),
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
return ok(
await _params(session).page(
param_code=param_code,
page=page,
limit=limit,
order_field=order_field,
order=order,
)
)
@sys_router.get("/admin/params/{param_id}")
async def get_param(
param_id: int,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await _params(session).get(param_id))
@sys_router.post("/admin/params")
async def save_param(
dto: SysParamPayload,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
await _params(session).save(
dto,
require_super_admin(request),
request.headers.get("Accept-Language"),
)
await _refresh_server_config(session)
return ok()
@sys_router.put("/admin/params")
async def update_param(
dto: SysParamPayload,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
await _params(session).update(
dto,
require_super_admin(request),
request.headers.get("Accept-Language"),
)
await _refresh_server_config(session)
return ok()
@sys_router.post("/admin/params/delete")
async def delete_params(
request: Request,
ids: list[str] = Body(),
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
await _params(session).delete(ids)
await _refresh_server_config(session)
return ok()
@sys_router.get("/admin/dict/type/page")
async def page_dict_types(
request: Request,
dict_type: str | None = Query(default=None, alias="dictType"),
dict_name: str | None = Query(default=None, alias="dictName"),
page: int = Query(default=1, ge=0),
limit: int = Query(default=10, ge=0),
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
return ok(
await _dict(session).page_types(
dict_type=dict_type,
dict_name=dict_name,
page=page,
limit=limit,
)
)
@sys_router.get("/admin/dict/type/{type_id}")
async def get_dict_type(
type_id: int,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await _dict(session).get_type(type_id))
@sys_router.post("/admin/dict/type/save")
async def save_dict_type(
dto: DictTypePayload,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
await _dict(session).save_type(dto, require_super_admin(request))
return ok()
@sys_router.put("/admin/dict/type/update")
async def update_dict_type(
dto: DictTypePayload,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
await _dict(session).update_type(dto, require_super_admin(request))
return ok()
@sys_router.post("/admin/dict/type/delete")
async def delete_dict_types(
request: Request,
ids: list[int] = Body(),
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
await _dict(session).delete_types(ids)
return ok()
@sys_router.get("/admin/dict/data/page")
async def page_dict_data(
request: Request,
dict_type_id: str | None = Query(default=None, alias="dictTypeId"),
dict_label: str | None = Query(default=None, alias="dictLabel"),
dict_value: str | None = Query(default=None, alias="dictValue"),
page: int = Query(default=1, ge=0),
limit: int = Query(default=10, ge=0),
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
if dict_type_id is None or not dict_type_id:
return JavaJSONResponse(envelope(None, code=500, msg="dictTypeId不能为空"))
try:
parsed_type_id = int(dict_type_id)
except ValueError as exc:
raise AppError(500) from exc
return ok(
await _dict(session).page_data(
dict_type_id=parsed_type_id,
dict_label=dict_label,
dict_value=dict_value,
page=page,
limit=limit,
)
)
@sys_router.get("/admin/dict/data/type/{dict_type}")
async def dict_items(
dict_type: str,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_normal(request)
return ok(await _dict(session).items(dict_type))
@sys_router.get("/admin/dict/data/{data_id}")
async def get_dict_data(
data_id: int,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await _dict(session).get_data(data_id))
@sys_router.post("/admin/dict/data/save")
async def save_dict_data(
dto: DictDataPayload,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
await _dict(session).save_data(dto, require_super_admin(request))
return ok()
@sys_router.put("/admin/dict/data/update")
async def update_dict_data(
dto: DictDataPayload,
request: Request,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
await _dict(session).update_data(dto, require_super_admin(request))
return ok()
@sys_router.post("/admin/dict/data/delete")
async def delete_dict_data(
request: Request,
ids: list[int] = Body(),
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
await _dict(session).delete_data(ids)
return ok()
@@ -0,0 +1,74 @@
from __future__ import annotations
from fastapi import APIRouter, Depends, Query, Request
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.database import get_db
from app.core.responses import JavaJSONResponse, ok
from app.core.security import require_normal, require_super_admin
from app.repositories.timbre import TimbreRepository
from app.schemas.timbre import TimbreBody
from app.services.timbre import TimbreService
timbre_router = APIRouter()
def _service(session: AsyncSession) -> TimbreService:
return TimbreService(TimbreRepository(session))
@timbre_router.get("/ttsVoice")
async def timbre_page(
request: Request,
tts_model_id: str | None = Query(default=None, alias="ttsModelId"),
name: str | None = None,
page: str | None = None,
limit: str | None = None,
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
require_super_admin(request)
return ok(
await _service(session).page(
tts_model_id, name, page, limit, request.headers.get("Accept-Language")
)
)
@timbre_router.post("/ttsVoice")
async def timbre_save(
body: TimbreBody, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
await _service(session).save(
body, require_super_admin(request), request.headers.get("Accept-Language")
)
return ok()
@timbre_router.put("/ttsVoice/{timbre_id}")
async def timbre_update(
timbre_id: str, body: TimbreBody, request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
await _service(session).update(
timbre_id, body, require_super_admin(request), request.headers.get("Accept-Language")
)
return ok()
@timbre_router.post("/ttsVoice/delete")
async def timbre_delete(
ids: list[str], request: Request, session: AsyncSession = Depends(get_db)
) -> JavaJSONResponse:
require_super_admin(request)
await _service(session).delete(ids)
return ok()
@timbre_router.get("/models/{model_id}/voices")
async def model_voices(
model_id: str,
request: Request,
voice_name: str | None = Query(default=None, alias="voiceName"),
session: AsyncSession = Depends(get_db),
) -> JavaJSONResponse:
user = require_normal(request)
return ok(await _service(session).voices(model_id, voice_name, user, request.headers.get("Accept-Language")))
@@ -0,0 +1,222 @@
from __future__ import annotations
from typing import Annotated, Any
from fastapi import APIRouter, Depends, File, Form, Request, UploadFile
from sqlalchemy.ext.asyncio import AsyncSession
from starlette.responses import Response
from app.core.database import get_db
from app.core.errors import AppError
from app.core.i18n import message_for
from app.core.responses import JavaJSONResponse, error_response, ok
from app.core.security import require_normal, require_super_admin
from app.schemas.voiceclone import VoiceCloneRenameRequest, VoiceCloneTrainRequest, VoiceResourceCreateRequest
from app.services.voiceclone import VoiceCloneService
voiceclone_router = APIRouter()
SessionDep = Annotated[AsyncSession, Depends(get_db)]
VoiceFile = Annotated[UploadFile, File(alias="voiceFile")]
VoiceIdForm = Annotated[str, Form(alias="id")]
def _query_map(request: Request) -> dict[str, Any]:
result: dict[str, Any] = {}
for key, value in request.query_params.multi_items():
if key in result:
previous = result[key]
result[key] = [*previous, value] if isinstance(previous, list) else [previous, value]
else:
result[key] = value
return result
# Static voiceResource paths deliberately precede /voiceResource/{id}.
@voiceclone_router.get("/voiceResource/ttsPlatforms")
async def tts_platforms(
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await VoiceCloneService(session).tts_platforms())
@voiceclone_router.get("/voiceResource/user/{user_id}")
async def voice_resources_by_user(
user_id: int,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
require_normal(request)
return ok(await VoiceCloneService(session).get_by_user(user_id))
@voiceclone_router.get("/voiceResource")
async def page_voice_resources(
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await VoiceCloneService(session).page(_query_map(request)))
@voiceclone_router.get("/voiceResource/{voice_id}")
async def get_voice_resource(
voice_id: str,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
require_super_admin(request)
return ok(await VoiceCloneService(session).get_detail(voice_id))
@voiceclone_router.post("/voiceResource")
async def create_voice_resource(
request: Request,
session: SessionDep,
body: VoiceResourceCreateRequest | None = None,
) -> JavaJSONResponse:
user = require_super_admin(request)
if body is None:
return error_response(request, 10145)
if body.model_id is None or body.model_id == "":
return error_response(request, 10146)
if not body.voice_ids:
return error_response(request, 10147)
if body.user_id is None:
return error_response(request, 10148)
try:
await VoiceCloneService(session).create_resources(body, actor=user)
return ok()
except AppError:
raise
except RuntimeError as exc:
return error_response(request, 10065, str(exc))
@voiceclone_router.delete("/voiceResource/{voice_id}")
async def delete_voice_resource(
voice_id: str,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
require_super_admin(request)
ids = voice_id.split(",") if voice_id else []
if not ids:
return error_response(request, 10149)
await VoiceCloneService(session).delete(ids)
return ok()
@voiceclone_router.get("/voiceClone")
async def page_voice_clones(
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
return ok(await VoiceCloneService(session).page(_query_map(request), user_id=user.id))
@voiceclone_router.post("/voiceClone/upload")
async def upload_voice_clone(
request: Request,
session: SessionDep,
voice_file: VoiceFile,
voice_id: VoiceIdForm = "",
) -> JavaJSONResponse:
user = require_normal(request)
service = VoiceCloneService(session)
try:
content = await voice_file.read()
if not content:
return error_response(request, 10140)
content_type = voice_file.content_type
if content_type is None or not content_type.startswith("audio/"):
return error_response(request, 10141)
filename = voice_file.filename
if filename is None or "." not in filename:
raise RuntimeError("文件名缺少扩展名")
extension = filename[filename.rfind(".") :].lower()
if extension not in {".mp3", ".wav"}:
return error_response(request, 500, "只允许上传.mp3和.wav格式的文件")
if len(content) > 10 * 1024 * 1024:
return error_response(request, 10142)
await service.check_permission(voice_id, user)
await service.upload_voice(voice_id, content)
return ok()
except Exception as exc:
if isinstance(exc, AppError):
message = exc.message or message_for(exc.code, request.headers.get("Accept-Language"))
else:
message = str(exc)
return error_response(request, 10143, message)
@voiceclone_router.post("/voiceClone/updateName")
async def update_voice_clone_name(
body: VoiceCloneRenameRequest,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
if body.id is None or body.id == "":
return error_response(request, 10006)
if body.name is None or body.name == "":
return error_response(request, 10181)
service = VoiceCloneService(session)
try:
await service.check_permission(body.id, user)
await service.rename(body.id or "", body.name or "")
return ok()
except Exception as exc:
if isinstance(exc, AppError):
message = exc.message or message_for(exc.code, request.headers.get("Accept-Language"))
else:
message = str(exc)
return error_response(request, 10066, message)
@voiceclone_router.post("/voiceClone/audio/{voice_id}")
async def get_voice_clone_audio_id(
voice_id: str,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
service = VoiceCloneService(session)
await service.check_permission(voice_id, user)
return ok(await service.create_audio_id(voice_id))
@voiceclone_router.get("/voiceClone/play/{download_id}")
async def play_voice_clone(download_id: str, session: SessionDep) -> Response:
try:
content = await VoiceCloneService(session).consume_audio(download_id)
if content is None:
return Response(status_code=404)
return Response(
content,
media_type="audio/wav",
headers={
"Content-Length": str(len(content)),
"Content-Disposition": "inline; filename=voice.wav",
},
)
except Exception:
return Response(status_code=500)
@voiceclone_router.post("/voiceClone/cloneAudio")
async def train_voice_clone(
body: VoiceCloneTrainRequest,
request: Request,
session: SessionDep,
) -> JavaJSONResponse:
user = require_normal(request)
service = VoiceCloneService(session)
await service.check_permission(body.clone_id, user)
await service.clone_audio(
body.clone_id or "",
accept_language=request.headers.get("Accept-Language"),
)
return ok()
@@ -0,0 +1 @@
"""Pydantic request and response schemas."""
@@ -0,0 +1,238 @@
from __future__ import annotations
import json
from datetime import datetime
from typing import Any
from pydantic import Field, field_validator
from pydantic_core import PydanticCustomError
from app.schemas.common import JavaModel
class AgentCreate(JavaModel):
agent_name: str
@field_validator("agent_name", mode="before")
@classmethod
def require_non_blank_name(cls, value: Any) -> Any:
if value is None or isinstance(value, str) and not value.strip():
raise PydanticCustomError("java_not_blank", "智能体名称不能为空")
return value
class AgentMemory(JavaModel):
summary_memory: str | None = None
class ContextProvider(JavaModel):
url: str | None = None
headers: dict[str, Any] | None = None
class FunctionInfo(JavaModel):
plugin_id: str | None = None
param_info: dict[str, Any] = Field(default_factory=dict)
@field_validator("param_info", mode="before")
@classmethod
def normalize_param_info(cls, value: Any) -> dict[str, Any]:
if value is None or value == "":
return {}
if isinstance(value, str):
parsed = json.loads(value)
if not isinstance(parsed, dict):
raise ValueError("paramInfo must be a JSON object")
return {str(key): item for key, item in parsed.items()}
if isinstance(value, dict):
return {str(key): item for key, item in value.items() if key is not None}
parsed = json.loads(json.dumps(value))
if not isinstance(parsed, dict):
raise ValueError("paramInfo must be an object")
return {str(key): item for key, item in parsed.items()}
class AgentUpdate(JavaModel):
agent_code: str | None = None
agent_name: str | None = None
asr_model_id: str | None = None
vad_model_id: str | None = None
llm_model_id: str | None = None
slm_model_id: str | None = None
vllm_model_id: str | None = None
tts_model_id: str | None = None
tts_voice_id: str | None = None
tts_language: str | None = None
tts_volume: int | None = None
tts_rate: int | None = None
tts_pitch: int | None = None
mem_model_id: str | None = None
intent_model_id: str | None = None
functions: list[FunctionInfo] | None = None
system_prompt: str | None = None
summary_memory: str | None = None
chat_history_conf: int | None = None
lang_code: str | None = None
language: str | None = None
sort: int | None = None
context_providers: list[ContextProvider] | None = None
correct_word_file_ids: list[str] | None = None
tag_names: list[str] | None = None
tag_ids: list[str] | None = None
class AgentChatHistoryReport(JavaModel):
mac_address: str
session_id: str
chat_type: int
content: str
audio_base64: str | None = None
report_time: int | None = None
@field_validator("mac_address", "session_id", "content", mode="before")
@classmethod
def require_non_blank(cls, value: Any) -> Any:
if value is None or isinstance(value, str) and not value.strip():
raise PydanticCustomError("java_not_blank", "不能为空")
return value
@field_validator("chat_type", mode="before")
@classmethod
def require_chat_type(cls, value: Any) -> Any:
if value is None:
raise PydanticCustomError("java_not_null", "不能为空")
return value
class AgentSnapshotPage(JavaModel):
page: int | None = 1
limit: int | None = 10
max_version_no: int | None = None
def page_or_default(self) -> int:
return self.page if self.page is not None and self.page >= 1 else 1
def limit_or_default(self) -> int:
return self.limit if self.limit is not None and self.limit >= 1 else 10
class AgentSnapshotRestore(JavaModel):
current_state_token: str
@field_validator("current_state_token", mode="before")
@classmethod
def require_non_blank_token(cls, value: Any) -> Any:
if value is None or isinstance(value, str) and not value.strip():
raise PydanticCustomError("java_not_blank", "不能为空")
return value
class AgentSnapshotTag(JavaModel):
id: str | None = None
tag_name: str | None = None
sort: int | None = None
class AgentSnapshotData(JavaModel):
agent_code: str | None = None
agent_name: str | None = None
asr_model_id: str | None = None
vad_model_id: str | None = None
llm_model_id: str | None = None
slm_model_id: str | None = None
vllm_model_id: str | None = None
tts_model_id: str | None = None
tts_voice_id: str | None = None
tts_language: str | None = None
tts_volume: int | None = None
tts_rate: int | None = None
tts_pitch: int | None = None
mem_model_id: str | None = None
intent_model_id: str | None = None
chat_history_conf: int | None = None
system_prompt: str | None = None
summary_memory: str | None = None
lang_code: str | None = None
language: str | None = None
sort: int | None = None
functions: list[FunctionInfo] | None = None
context_providers: list[ContextProvider] | None = None
correct_word_file_ids: list[str] | None = None
tag_names: list[str] | None = None
tags: list[AgentSnapshotTag] | None = None
class AgentTemplate(JavaModel):
id: str | None = None
agent_code: str | None = None
agent_name: str | None = None
asr_model_id: str | None = None
vad_model_id: str | None = None
llm_model_id: str | None = None
vllm_model_id: str | None = None
tts_model_id: str | None = None
tts_voice_id: str | None = None
tts_language: str | None = None
tts_volume: int | None = None
tts_rate: int | None = None
tts_pitch: int | None = None
mem_model_id: str | None = None
intent_model_id: str | None = None
chat_history_conf: int | None = None
system_prompt: str | None = None
summary_memory: str | None = None
lang_code: str | None = None
language: str | None = None
sort: int | None = None
creator: int | None = None
created_at: datetime | None = None
updater: int | None = None
updated_at: datetime | None = None
class AgentVoicePrintSave(JavaModel):
agent_id: str | None = None
audio_id: str | None = None
source_name: str | None = None
introduce: str | None = None
class AgentVoicePrintUpdate(JavaModel):
id: str | None = None
audio_id: str | None = None
source_name: str | None = None
introduce: str | None = None
class AgentTagAssignment(JavaModel):
tag_ids: list[str] | None = None
tag_names: list[str] | None = None
SNAPSHOT_FIELD_ORDER = [
"agentCode",
"agentName",
"asrModelId",
"vadModelId",
"llmModelId",
"slmModelId",
"vllmModelId",
"ttsModelId",
"ttsVoiceId",
"ttsLanguage",
"ttsVolume",
"ttsRate",
"ttsPitch",
"memModelId",
"intentModelId",
"chatHistoryConf",
"systemPrompt",
"summaryMemory",
"langCode",
"language",
"sort",
"functions",
"contextProviders",
"correctWordFileIds",
"tagNames",
]
@@ -0,0 +1,57 @@
from __future__ import annotations
from collections.abc import Callable
from typing import Any, Generic, TypeVar
from pydantic import BaseModel, ConfigDict, Field
T = TypeVar("T")
def to_camel(value: str) -> str:
head, *tail = value.split("_")
return head + "".join(part[:1].upper() + part[1:] for part in tail)
class JavaModel(BaseModel):
model_config = ConfigDict(
alias_generator=to_camel,
populate_by_name=True,
extra="ignore",
str_strip_whitespace=False,
serialize_by_alias=True,
)
class PageData(JavaModel, Generic[T]):
total: int
list: list[T]
class PageQuery(JavaModel):
page: int = Field(default=1, ge=1)
limit: int = Field(default=10, ge=1)
order_field: str | list[str] | None = None
order: str | None = None
class DeleteIds(JavaModel):
ids: list[str]
def page_payload(rows: list[Any], total: int) -> dict[str, Any]:
return {"total": int(total), "list": rows}
def safe_order_by(
requested: str | list[str] | None,
*,
allowed: set[str],
default: str,
transform: Callable[[str], str] | None = None,
) -> list[str]:
fields = [requested] if isinstance(requested, str) else list(requested or [])
selected = [field for field in fields if field in allowed]
if not selected:
selected = [default]
return [transform(field) if transform else field for field in selected]
@@ -0,0 +1,25 @@
from __future__ import annotations
from pydantic import field_validator
from app.schemas.common import JavaModel
def _not_blank(value: str) -> str:
if not value or not value.strip():
raise ValueError("must not be blank")
return value
class AgentModelsRequest(JavaModel):
mac_address: str
client_id: str
selected_module: dict[str, str]
_validate_required = field_validator("mac_address", "client_id")(_not_blank)
class CorrectWordsRequest(JavaModel):
mac_address: str
_validate_required = field_validator("mac_address")(_not_blank)
@@ -0,0 +1,9 @@
from __future__ import annotations
from app.schemas.common import JavaModel
class CorrectWordFileBody(JavaModel):
file_name: str | None = None
content: list[str] | None = None
file_size: int | None = None
@@ -0,0 +1,145 @@
from __future__ import annotations
from typing import Any
from pydantic import AliasChoices, Field
from app.schemas.common import JavaModel
class DeviceRegisterRequest(JavaModel):
mac_address: str | None = None
class DeviceUnbindRequest(JavaModel):
device_id: str | None = None
class DeviceUpdateRequest(JavaModel):
auto_update: int | None = None
alias: str | None = None
class DeviceManualAddRequest(JavaModel):
agent_id: str | None = None
board: str | None = None
app_version: str | None = None
mac_address: str | None = None
class DeviceToolCallRequest(JavaModel):
name: str | None = None
arguments: dict[str, Any] | None = None
class DeviceAddressBookAliasRequest(JavaModel):
mac_address: str | None = None
target_mac: str | None = None
alias: str | None = None
class DeviceAddressBookPermissionRequest(JavaModel):
mac_address: str | None = None
target_mac: str | None = None
has_permission: bool | None = None
class ChipInfo(JavaModel):
model: int | None = None
cores: int | None = None
revision: int | None = None
features: int | None = None
class ApplicationInfo(JavaModel):
name: str | None = None
version: str | None = None
compile_time: str | None = Field(
default=None,
validation_alias=AliasChoices("compile_time", "compileTime"),
serialization_alias="compile_time",
)
idf_version: str | None = Field(
default=None,
validation_alias=AliasChoices("idf_version", "idfVersion"),
serialization_alias="idf_version",
)
elf_sha256: str | None = Field(
default=None,
validation_alias=AliasChoices("elf_sha256", "elfSha256"),
serialization_alias="elf_sha256",
)
class PartitionInfo(JavaModel):
label: str | None = None
type: int | None = None
subtype: int | None = None
address: int | None = None
size: int | None = None
class OtaPartitionInfo(JavaModel):
label: str | None = None
class BoardInfo(JavaModel):
type: str | None = None
ssid: str | None = None
rssi: int | None = None
channel: int | None = None
ip: str | None = None
mac: str | None = None
class DeviceReportRequest(JavaModel):
version: int | None = None
flash_size: int | None = Field(
default=None,
validation_alias=AliasChoices("flash_size", "flashSize"),
serialization_alias="flash_size",
)
minimum_free_heap_size: int | None = Field(
default=None,
validation_alias=AliasChoices("minimum_free_heap_size", "minimumFreeHeapSize"),
serialization_alias="minimum_free_heap_size",
)
mac_address: str | None = Field(
default=None,
validation_alias=AliasChoices("mac_address", "macAddress"),
serialization_alias="mac_address",
)
uuid: str | None = None
chip_model_name: str | None = Field(
default=None,
validation_alias=AliasChoices("chip_model_name", "chipModelName"),
serialization_alias="chip_model_name",
)
chip_info: ChipInfo | None = Field(
default=None,
validation_alias=AliasChoices("chip_info", "chipInfo"),
serialization_alias="chip_info",
)
application: ApplicationInfo | None = None
partition_table: list[PartitionInfo] | None = Field(
default=None,
validation_alias=AliasChoices("partition_table", "partitionTable"),
serialization_alias="partition_table",
)
ota: OtaPartitionInfo | None = None
board: BoardInfo | None = None
class OtaRecord(JavaModel):
id: str | None = None
firmware_name: str | None = None
type: str | None = None
version: str | None = None
size: int | None = None
remark: str | None = None
firmware_path: str | None = None
sort: int | None = None
updater: int | None = None
update_date: str | None = None
creator: int | None = None
create_date: str | None = None
@@ -0,0 +1,53 @@
from __future__ import annotations
from datetime import datetime
from typing import Any
from pydantic import AliasChoices, Field
from app.schemas.common import JavaModel
class KnowledgeBaseBody(JavaModel):
id: str | None = None
dataset_id: str | None = None
rag_model_id: str | None = None
name: str | None = None
avatar: str | None = None
description: str | None = None
embedding_model: str | None = None
permission: str | None = None
chunk_method: str | None = None
parser_config: str | None = None
chunk_count: int | None = None
token_num: int | None = None
status: int | None = None
creator: int | None = None
created_at: datetime | None = None
updater: int | None = None
updated_at: datetime | None = None
document_count: int | None = None
error_message: str | None = None
class DocumentBatchBody(JavaModel):
ids: list[str] | None = Field(
default=None,
validation_alias=AliasChoices("ids", "document_ids"),
)
class RetrievalBody(JavaModel):
dataset_ids: list[str] | None = None
document_ids: list[str] | None = None
question: str | None = None
page: int | None = None
page_size: int | None = None
similarity_threshold: float | None = None
vector_similarity_weight: float | None = None
top_k: int | None = None
rerank_id: str | None = None
highlight: bool | None = None
keyword: bool | None = None
cross_languages: list[str] | None = None
metadata_condition: dict[str, Any] | None = None
@@ -0,0 +1,25 @@
from __future__ import annotations
from typing import Any
from app.schemas.common import JavaModel
class ModelConfigBody(JavaModel):
id: str | None = None
model_code: str | None = None
model_name: str | None = None
is_default: int | None = None
is_enabled: int | None = None
config_json: dict[str, Any] | None = None
doc_link: str | None = None
remark: str | None = None
sort: int | None = None
class ModelProviderBody(JavaModel):
id: str | None = None
model_type: str | None = None
provider_code: str | None = None
name: str | None = None
fields: str | None = None
sort: int | None = None
@@ -0,0 +1,59 @@
from __future__ import annotations
from typing import Any
from app.schemas.common import JavaModel
class LoginRequest(JavaModel):
# LoginController does not use @Valid; null/blank values reach its service logic.
username: str | None = None
password: str | None = None
mobile_captcha: str | None = None
captcha_id: str | None = None
class SmsVerificationRequest(JavaModel):
# smsVerification likewise omits @Valid in the Java controller.
phone: str | None = None
captcha: str | None = None
captcha_id: str | None = None
class PasswordChangeRequest(JavaModel):
password: str | None = None
new_password: str | None = None
class RetrievePasswordRequest(JavaModel):
phone: str | None = None
code: str | None = None
password: str | None = None
captcha_id: str | None = None
class TokenData(JavaModel):
token: str
expire: int
client_hash: str | None
class UserDetailData(JavaModel):
id: int
username: str
super_admin: int
token: str
status: int
class PublicConfigData(JavaModel):
enable_mobile_register: bool
version: str
year: str
allow_user_register: bool
mobile_area_list: list[dict[str, Any]]
beian_icp_num: str | None
beian_ga_num: str | None
name: str | None
sm2_public_key: str
system_web_menu: Any | None = None
@@ -0,0 +1,45 @@
from __future__ import annotations
from pydantic import field_validator
from app.schemas.common import JavaModel
def _not_blank(value: str) -> str:
if not value or not value.strip():
raise ValueError("must not be blank")
return value
class SysParamPayload(JavaModel):
id: int | None = None
param_code: str | None = None
param_value: str | None = None
value_type: str | None = None
remark: str | None = None
class DictTypePayload(JavaModel):
# Controller calls ValidatorUtils without the DTO's custom groups, so these constraints don't execute in Java.
id: int | None = None
dict_type: str | None = None
dict_name: str | None = None
remark: str | None = None
sort: int | None = None
class DictDataPayload(JavaModel):
# See DictTypePayload: Add/Update/DefaultGroup annotations are skipped by the Java controller.
id: int | None = None
dict_type_id: int | None = None
dict_label: str | None = None
dict_value: str | None = None
remark: str | None = None
sort: int | None = None
class EmitServerActionRequest(JavaModel):
target_ws: str
action: str | None
_validate_target = field_validator("target_ws")(_not_blank)
@@ -0,0 +1,15 @@
from __future__ import annotations
from app.schemas.common import JavaModel
class TimbreBody(JavaModel):
languages: str | None = None
name: str | None = None
remark: str | None = None
reference_audio: str | None = None
reference_text: str | None = None
sort: int | None = 0
tts_model_id: str | None = None
tts_voice: str | None = None
voice_demo: str | None = None
@@ -0,0 +1,19 @@
from __future__ import annotations
from app.schemas.common import JavaModel
class VoiceResourceCreateRequest(JavaModel):
model_id: str | None = None
voice_ids: list[str] | None = None
user_id: int | None = None
languages: str | None = None
class VoiceCloneRenameRequest(JavaModel):
id: str | None = None
name: str | None = None
class VoiceCloneTrainRequest(JavaModel):
clone_id: str | None = None
@@ -0,0 +1 @@
"""Business services."""
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,521 @@
from __future__ import annotations
import base64
import hashlib
import json
import logging
import math
import urllib.parse
from copy import deepcopy
from typing import Any, cast
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
from redis.asyncio import Redis
from app.core.errors import AppError
from app.core.redis import JavaRedisCodec, get_redis
from app.repositories.config import ConfigRepository
logger = logging.getLogger(__name__)
class ConfigService:
def __init__(self, repository: ConfigRepository, *, redis: Redis | None = None):
self.repository = repository
self.redis = redis or get_redis()
async def get_config(self, *, use_cache: bool) -> dict[str, Any]:
if use_cache:
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)("server:config"))
if isinstance(cached, dict):
return cast(dict[str, Any], cached)
result = self._build_base_config(await self.repository.list_params())
template = await self.repository.get_default_template()
if template is None:
raise AppError(10183)
await self._build_module_config(
result=result,
assistant_name=None,
prompt=None,
summary_memory=None,
voice=None,
reference_audio=None,
reference_text=None,
language=None,
tts_volume=None,
tts_rate=None,
tts_pitch=None,
vad_model_id=self._string(template.get("vad_model_id")),
asr_model_id=self._string(template.get("asr_model_id")),
llm_model_id=None,
vllm_model_id=None,
slm_model_id=None,
tts_model_id=None,
mem_model_id=None,
intent_model_id=None,
rag_model_id=None,
)
await cast(Any, self.redis.set)("server:config", JavaRedisCodec.encode(result), ex=24 * 60 * 60)
return result
async def get_agent_models(
self,
mac_address: str,
selected_module: dict[str, str],
) -> dict[str, Any]:
temporary_key = f"tmp_register_mac:{mac_address}"
temporary = JavaRedisCodec.decode(await cast(Any, self.redis.get)(temporary_key))
if temporary == "true":
await cast(Any, self.redis.delete)(temporary_key)
return await self.get_config(use_cache=True)
device = await self.repository.get_device_by_mac(mac_address)
if device is None:
safe_address = mac_address.replace(":", "_").lower()
activation = JavaRedisCodec.decode(
await cast(Any, self.redis.get)(f"ota:activation:data:{safe_address}")
)
if isinstance(activation, dict) and activation.get("activation_code"):
raise AppError(10042, params=(str(activation["activation_code"]),))
raise AppError(10041)
agent_id = self._string(device.get("agent_id"))
agent = await self.repository.get_agent(agent_id or "") if agent_id else None
if agent is None:
raise AppError(10053)
voice: str | None = None
reference_audio: str | None = None
reference_text: str | None = None
language: str | None = None
voice_id = self._string(agent.get("tts_voice_id"))
timbre = await self._timbre(voice_id) if voice_id else None
if timbre is not None:
voice = self._string(timbre.get("tts_voice"))
reference_audio = self._string(timbre.get("reference_audio"))
reference_text = self._string(timbre.get("reference_text"))
chosen_language = self._string(agent.get("tts_language"))
if chosen_language and chosen_language.strip():
language = chosen_language
else:
languages = self._string(timbre.get("languages"))
if languages and languages.strip():
language = languages.split("", 1)[0].strip()
elif voice_id:
clone = await self.repository.get_voice_clone(voice_id)
if clone is not None:
voice = self._string(clone.get("voice_id"))
chosen_language = self._string(agent.get("tts_language"))
language = chosen_language if chosen_language and chosen_language.strip() else "普通话"
result: dict[str, Any] = {
"device_max_output_size": await self._param("device_max_output_size", from_cache=True)
}
memory_model = self._string(agent.get("mem_model_id"))
chat_history = agent.get("chat_history_conf")
if memory_model == "Memory_nomem":
chat_history = 0
elif memory_model is not None and memory_model != "Memory_nomem" and chat_history is None:
chat_history = 2
result["chat_history_conf"] = chat_history
vad_model_id = self._string(agent.get("vad_model_id"))
asr_model_id = self._string(agent.get("asr_model_id"))
if selected_module.get("VAD") == vad_model_id:
vad_model_id = None
if selected_module.get("ASR") == asr_model_id:
asr_model_id = None
if self._string(agent.get("intent_model_id")) != "Intent_nointent":
plugins = await self._plugins(str(agent["id"]))
if plugins:
result["plugins"] = plugins
mcp_endpoint = await self._mcp_address(str(agent["id"]))
if mcp_endpoint and mcp_endpoint.startswith("ws"):
result["mcp_endpoint"] = mcp_endpoint.replace("/mcp/", "/call/")
context_providers = self._json_value(await self.repository.get_context_providers(str(agent["id"])))
if isinstance(context_providers, list) and context_providers:
result["context_providers"] = context_providers
await self._add_voiceprint(str(agent["id"]), result)
await self._build_module_config(
result=result,
assistant_name=self._string(agent.get("agent_name")),
prompt=self._string(agent.get("system_prompt")),
summary_memory=self._string(agent.get("summary_memory")),
voice=voice,
reference_audio=reference_audio,
reference_text=reference_text,
language=language,
tts_volume=self._integer(agent.get("tts_volume")),
tts_rate=self._integer(agent.get("tts_rate")),
tts_pitch=self._integer(agent.get("tts_pitch")),
vad_model_id=vad_model_id,
asr_model_id=asr_model_id,
llm_model_id=self._string(agent.get("llm_model_id")),
vllm_model_id=self._string(agent.get("vllm_model_id")),
slm_model_id=self._string(agent.get("slm_model_id")),
tts_model_id=self._string(agent.get("tts_model_id")),
mem_model_id=memory_model,
intent_model_id=self._string(agent.get("intent_model_id")),
rag_model_id=None,
)
return result
async def get_correct_words(self, mac_address: str) -> list[str]:
device = await self.repository.get_device_by_mac(mac_address)
if device is None or device.get("agent_id") is None:
return []
rows = await self.repository.get_correct_word_items(str(device["agent_id"]))
return [
f"{self._java_string(row.get('source_word'))}|{self._java_string(row.get('target_word'))}"
for row in rows
]
@staticmethod
def _build_base_config(rows: list[dict[str, Any]]) -> dict[str, Any]:
config: dict[str, Any] = {}
for row in rows:
code = str(row.get("param_code") or "")
keys = code.split(".")
current = config
for key in keys[:-1]:
if key not in current:
current[key] = {}
nested = current[key]
if not isinstance(nested, dict):
raise TypeError(f"configuration path {code} collides with scalar key {key}")
current = nested
value = str(row.get("param_value") or "")
value_type = str(row.get("value_type") or "string").lower()
current[keys[-1]] = ConfigService._typed_param(value, value_type)
return config
@staticmethod
def _typed_param(value: str, value_type: str) -> Any:
if value_type == "number":
try:
number = float(value)
# Java's implementation returns an Integer only when the double
# equals its narrowing conversion to a signed 32-bit int.
if math.isnan(number):
narrowed = 0
elif number >= 2**31 - 1:
narrowed = 2**31 - 1
elif number <= -(2**31):
narrowed = -(2**31)
else:
narrowed = int(number)
return narrowed if number == narrowed else number
except ValueError:
return value
if value_type == "boolean":
return value.lower() == "true"
if value_type == "array":
return [item.strip() for item in value.split(";") if item.strip()]
if value_type == "json":
try:
return json.loads(value)
except json.JSONDecodeError:
return value
return value
async def _build_module_config(
self,
*,
result: dict[str, Any],
assistant_name: str | None,
prompt: str | None,
summary_memory: str | None,
voice: str | None,
reference_audio: str | None,
reference_text: str | None,
language: str | None,
tts_volume: int | None,
tts_rate: int | None,
tts_pitch: int | None,
vad_model_id: str | None,
asr_model_id: str | None,
llm_model_id: str | None,
vllm_model_id: str | None,
slm_model_id: str | None,
tts_model_id: str | None,
mem_model_id: str | None,
intent_model_id: str | None,
rag_model_id: str | None,
) -> None:
selected: dict[str, str] = {}
model_types = ("VAD", "ASR", "TTS", "Memory", "Intent", "LLM", "VLLM", "SLM", "RAG")
model_ids = (
vad_model_id,
asr_model_id,
tts_model_id,
mem_model_id,
intent_model_id,
llm_model_id,
vllm_model_id,
slm_model_id,
rag_model_id,
)
intent_llm_id: str | None = None
memory_llm_id: str | None = None
for model_type, model_id in zip(model_types, model_ids, strict=True):
if model_id is None:
continue
model = await self._model(model_id)
if model is None:
continue
configuration = self._json_value(model.get("config_json"))
type_config: dict[str, Any] = {}
if isinstance(configuration, dict):
configuration = deepcopy(configuration)
type_config[str(model["id"])] = configuration
if model_type == "TTS":
optional_values = {
"private_voice": voice,
"ref_audio": reference_audio,
"ref_text": reference_text,
"language": language,
"ttsVolume": tts_volume,
"ttsRate": tts_rate,
"ttsPitch": tts_pitch,
}
configuration.update({key: value for key, value in optional_values.items() if value is not None})
if configuration.get("type") == "huoshan_double_stream" and voice and voice.startswith("S_"):
configuration["resource_id"] = "seed-icl-1.0"
elif model_type == "Intent":
if configuration.get("type") == "intent_llm":
intent_llm_id = self._string(configuration.get("llm"))
if intent_llm_id == llm_model_id:
intent_llm_id = None
functions = configuration.get("functions")
if isinstance(functions, str) and functions.strip():
configuration["functions"] = functions.split(";")
elif model_type == "Memory" and configuration.get("type") == "mem_local_short":
memory_llm_id = self._string(configuration.get("llm"))
if memory_llm_id == llm_model_id:
memory_llm_id = None
elif model_type == "LLM":
for extra_id in (intent_llm_id, memory_llm_id):
if extra_id and extra_id not in type_config:
extra = await self._model(extra_id)
if extra is not None:
type_config[str(extra["id"])] = deepcopy(self._json_value(extra.get("config_json")))
if slm_model_id and slm_model_id != llm_model_id and slm_model_id not in type_config:
small = await self._model(slm_model_id)
small_config = None if small is None else self._json_value(small.get("config_json"))
if small is not None and small_config is not None:
type_config[str(small["id"])] = deepcopy(small_config)
result[model_type] = type_config
selected[model_type] = str(model["id"])
result["selected_module"] = selected
if prompt and prompt.strip():
replacement = assistant_name if assistant_name and assistant_name.strip() else "小智"
prompt = prompt.replace("{{assistant_name}}", replacement)
result["prompt"] = prompt
result["summaryMemory"] = summary_memory
async def _plugins(self, agent_id: str) -> dict[str, Any]:
mappings = await self.repository.get_plugin_mappings(agent_id)
result: dict[str, Any] = {}
knowledge_groups: dict[str, list[dict[str, Any]]] = {}
knowledge_models: dict[str, dict[str, Any]] = {}
for mapping in mappings:
provider_code = self._string(mapping.get("provider_code"))
if provider_code and provider_code.strip():
value = mapping.get("param_info")
result[provider_code] = (
json.dumps(value, ensure_ascii=False, separators=(",", ":")) if isinstance(value, dict) else value
)
# Java removes knowledge mappings by iterating the original list backwards, which reverses dataset order.
for mapping in reversed(mappings):
provider_code = self._string(mapping.get("provider_code"))
if provider_code and provider_code.strip():
continue
dataset = await self.repository.get_dataset(str(mapping["plugin_id"]))
if dataset is None or dataset.get("rag_model_id") is None:
continue
model = await self._model(str(dataset["rag_model_id"]))
if model is None or not model.get("model_code"):
continue
code = str(model["model_code"])
knowledge_groups.setdefault(code, []).append(dataset)
knowledge_models[code] = model
for code, datasets in knowledge_groups.items():
model_config = self._json_value(knowledge_models[code].get("config_json"))
if not isinstance(model_config, dict):
continue
names = ",".join(self._java_string(dataset.get("name")) for dataset in datasets)
descriptions = ",".join(
self._java_string(dataset.get("description")) for dataset in datasets
)
params = {
"base_url": model_config.get("base_url"),
"api_key": model_config.get("api_key"),
"dataset_ids": [dataset.get("dataset_id") for dataset in datasets],
"description": (
f"如果用户询问与【{names}】涵盖的主体范围相关内容时应调用本方法,"
f"用于查询:{descriptions}"
),
}
result[f"search_from_{code}"] = json.dumps(params, ensure_ascii=False, separators=(",", ":"))
return result
async def _mcp_address(self, agent_id: str) -> str | None:
endpoint = await self._param("server.mcp_endpoint", from_cache=True)
if endpoint is None or not endpoint.strip() or endpoint == "null":
return None
parsed = urllib.parse.urlsplit(endpoint)
query = parsed.query
marker_index = query.find("key=")
key = query[marker_index + len("key=") :]
scheme = "wss" if parsed.scheme == "https" else "ws"
path = parsed.path
prefix_path = path[: path.rfind("/")] if "/" in path else ""
prefix = urllib.parse.urlunsplit((scheme, parsed.netloc, prefix_path, "", ""))
token = self._aes_encrypt(
key,
json.dumps(
{"agentId": hashlib.md5(agent_id.encode(), usedforsecurity=False).hexdigest()},
ensure_ascii=False,
separators=(", ", ": "),
),
)
return f"{prefix}/mcp/?token={urllib.parse.quote_plus(token)}"
@staticmethod
def _aes_encrypt(key: str, plaintext: str) -> str:
key_bytes = key.encode()
if len(key_bytes) not in {16, 24, 32}:
key_bytes = (key_bytes + bytes(32))[:32]
block_size = 16
padding_length = block_size - len(plaintext.encode()) % block_size
padded = plaintext.encode() + bytes([padding_length]) * padding_length
# Java's published MCP token format is AES/ECB/PKCS5Padding; changing modes breaks existing servers.
encryptor = Cipher(algorithms.AES(key_bytes), modes.ECB()).encryptor() # noqa: S305
encrypted = encryptor.update(padded) + encryptor.finalize()
return base64.b64encode(encrypted).decode("ascii")
async def _add_voiceprint(self, agent_id: str, result: dict[str, Any]) -> None:
try:
url = await self._param("server.voice_print", from_cache=True)
if url is None or not url.strip() or url == "null":
return
rows = await self.repository.get_voiceprints(agent_id)
if not rows:
return
speakers = [
(
f"{self._java_string(row.get('id'))},"
f"{self._java_string(row.get('source_name'))},{row.get('introduce') or ''}"
)
for row in rows
]
threshold_value = await self._param("server.voiceprint_similarity_threshold", from_cache=True)
try:
threshold = (
float(threshold_value)
if threshold_value is not None and threshold_value not in ("", "null")
else 0.4
)
except ValueError:
threshold = 0.4
result["voiceprint"] = {"url": url, "speakers": speakers, "similarity_threshold": threshold}
except Exception:
logger.warning("Voiceprint configuration lookup failed", exc_info=True)
async def _param(self, code: str, *, from_cache: bool) -> str | None:
if from_cache:
cached = JavaRedisCodec.decode(await cast(Any, self.redis.hget)("sys:params", code))
if cached is not None:
return str(cached)
value = await self.repository.get_param_value(code)
if value is not None and from_cache:
await cast(Any, self.redis.hset)("sys:params", code, JavaRedisCodec.encode(value))
await cast(Any, self.redis.expire)("sys:params", 24 * 60 * 60)
return value
async def _model(self, model_id: str) -> dict[str, Any] | None:
key = f"model:data:{model_id}"
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)(key))
if isinstance(cached, dict):
return self._normalize_cached(cast(dict[str, Any], cached))
model = await self.repository.get_model(model_id)
if model is not None:
raw_configuration = model.get("config_json")
if isinstance(raw_configuration, str):
parsed_configuration = json.loads(raw_configuration)
if parsed_configuration is not None and not isinstance(parsed_configuration, dict):
raise TypeError("ModelConfigEntity.configJson must be a JSON object")
model["config_json"] = parsed_configuration
await cast(Any, self.redis.set)(
key,
JavaRedisCodec.encode(
model,
java_type="xiaozhi.modules.model.entity.ModelConfigEntity",
field_java_types={
"configJson": "cn.hutool.json.JSONObject",
"creator": "java.lang.Long",
"updater": "java.lang.Long",
},
),
ex=24 * 60 * 60,
)
return model
async def _timbre(self, timbre_id: str) -> dict[str, Any] | None:
key = f"timbre:details:{timbre_id}"
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)(key))
if isinstance(cached, dict):
return self._normalize_cached(cast(dict[str, Any], cached))
timbre = await self.repository.get_timbre(timbre_id)
if timbre is not None:
await cast(Any, self.redis.set)(
key,
JavaRedisCodec.encode(
timbre,
java_type="xiaozhi.modules.timbre.vo.TimbreDetailsVO",
field_java_types={"sort": "java.lang.Long"},
),
ex=24 * 60 * 60,
)
return timbre
@staticmethod
def _normalize_cached(value: dict[str, Any]) -> dict[str, Any]:
aliases = {
"modelType": "model_type",
"modelCode": "model_code",
"modelName": "model_name",
"configJson": "config_json",
"ttsVoice": "tts_voice",
"referenceAudio": "reference_audio",
"referenceText": "reference_text",
"ttsModelId": "tts_model_id",
}
return {aliases.get(key, key): item for key, item in value.items() if key != "@class"}
@staticmethod
def _json_value(value: Any) -> Any:
if isinstance(value, bytes):
value = value.decode()
if isinstance(value, str):
try:
return json.loads(value)
except json.JSONDecodeError:
return value
return value
@staticmethod
def _string(value: Any) -> str | None:
return None if value is None else str(value)
@staticmethod
def _integer(value: Any) -> int | None:
return None if value is None else int(value)
@staticmethod
def _java_string(value: Any) -> str:
return "null" if value is None else str(value)
@@ -0,0 +1,133 @@
from __future__ import annotations
import uuid
from typing import Any
from app.core.errors import AppError
from app.core.security import AuthUser, shanghai_now_naive
from app.repositories.correctword import CorrectWordRepository
from app.schemas.correctword import CorrectWordFileBody
def _parse_lines(lines: list[str]) -> list[tuple[str, str]]:
result: list[tuple[str, str]] = []
for raw in lines:
line = raw.strip()
if not line or "|" not in line:
continue
source, target = line.split("|", 1)
if source.strip() and target.strip():
result.append((source.strip(), target.strip()))
return result
def _content_lines(value: str | None) -> list[str]:
if value is None:
return []
# Java String.split keeps one empty element for the empty source string,
# while still discarding trailing empty elements for non-empty strings.
if value == "":
return [""]
lines = value.split("\n")
while lines and lines[-1] == "":
lines.pop()
return lines
def file_vo(row: dict[str, Any]) -> dict[str, Any]:
return {
"id": row.get("id"),
"fileName": row.get("file_name"),
"wordCount": row.get("word_count"),
"content": _content_lines(row.get("content")),
"createdAt": row.get("created_at"),
"updatedAt": row.get("updated_at"),
}
class CorrectWordService:
def __init__(self, repository: CorrectWordRepository):
self.repository = repository
@staticmethod
def validate(body: CorrectWordFileBody, *, check_size: bool) -> None:
if body.file_name is None or not body.file_name.strip():
raise AppError(10034, "文件名不能为空")
if not body.content:
raise AppError(10034, "替换词内容不能为空")
if check_size and body.file_size is not None and body.file_size > 1024 * 1024:
raise AppError(10204)
async def create(self, body: CorrectWordFileBody, user: AuthUser) -> dict[str, Any]:
self.validate(body, check_size=True)
assert body.file_name is not None
assert body.content is not None
items = _parse_lines(body.content)
file_id, now = uuid.uuid4().hex, shanghai_now_naive()
values = {
"id": file_id,
"file_name": body.file_name,
"word_count": len(items),
"content": "\n".join(body.content),
"creator": user.id,
"now": now,
}
async with self.repository.session.begin():
if await self.repository.name_exists(user.id, body.file_name):
raise AppError(10203)
await self.repository.insert_file(values)
await self.repository.insert_items(
[
{"id": uuid.uuid4().hex, "file_id": file_id, "source_word": source, "target_word": target}
for source, target in items
]
)
return file_vo({**values, "created_at": now, "updated_at": None})
async def update(self, file_id: str, body: CorrectWordFileBody, user: AuthUser) -> None:
self.validate(body, check_size=False)
assert body.file_name is not None
assert body.content is not None
items = _parse_lines(body.content)
async with self.repository.session.begin():
row = await self.repository.get_file(file_id, for_update=True)
if row is None:
return
if await self.repository.name_exists(user.id, body.file_name, file_id):
raise AppError(500, f"文件名已存在:{body.file_name}")
await self.repository.delete_items(file_id)
await self.repository.insert_items(
[
{"id": uuid.uuid4().hex, "file_id": file_id, "source_word": source, "target_word": target}
for source, target in items
]
)
await self.repository.update_file(
{
"id": file_id,
"file_name": body.file_name,
"word_count": len(items),
"content": "\n".join(body.content),
"updater": user.id,
"now": shanghai_now_naive(),
}
)
async def page(self, user: AuthUser, page: str | None, limit: str | None) -> dict[str, Any]:
current, size = max(int(page or "1"), 1), int(limit or "10")
rows, total = await self.repository.list_files(user.id, offset=(current - 1) * size, limit=size)
return {"total": total, "list": [file_vo(row) for row in rows]}
async def all(self, user: AuthUser) -> list[dict[str, Any]]:
rows, _ = await self.repository.list_files(user.id)
return [file_vo(row) for row in rows]
async def get(self, file_id: str) -> dict[str, Any] | None:
row = await self.repository.get_file(file_id)
return file_vo(row) if row else None
async def delete(self, file_ids: list[str]) -> None:
async with self.repository.session.begin():
for file_id in file_ids:
if file_id and file_id.strip():
await self.repository.delete_file_graph(file_id.strip())
@@ -0,0 +1,978 @@
from __future__ import annotations
import base64
import hashlib
import hmac
import json
import logging
import random
import re
import secrets
import uuid
from collections.abc import Callable, Mapping, Sequence
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, cast
from zoneinfo import ZoneInfo
import httpx
from redis.asyncio import Redis
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.config import get_settings
from app.core.database import get_session_factory
from app.core.errors import AppError
from app.core.redis import JavaRedisCodec, get_redis
from app.core.security import AuthUser, shanghai_now_naive
from app.integrations.mqtt_gateway import post_json
from app.repositories.device import DeviceRepository
from app.schemas.device import DeviceManualAddRequest, DeviceReportRequest, DeviceUpdateRequest, OtaRecord
from app.services.system_params import SystemParamService
logger = logging.getLogger(__name__)
DEFAULT_TTL_SECONDS = 24 * 60 * 60
INVALID_FIRMWARE_URL = (
"http://xiaozhi.server.com:8002/xiaozhi/otaMag/download/NOT_ACTIVATED_FIRMWARE_THIS_IS_A_INVALID_URL"
)
MAC_PATTERN = re.compile(r"^([0-9A-Za-z]{2}[:-]){5}([0-9A-Za-z]{2})$")
OTA_ORDER_COLUMNS = {
"id": "id",
"firmwareName": "firmware_name",
"firmware_name": "firmware_name",
"type": "type",
"version": "version",
"size": "size",
"sort": "sort",
"updateDate": "update_date",
"update_date": "update_date",
"createDate": "create_date",
"create_date": "create_date",
}
def is_blank(value: str | None) -> bool:
return value is None or not value.strip()
def _java_semicolon_split(value: str) -> list[str]:
parts = value.split(";")
while parts and parts[-1] == "":
parts.pop()
return parts
def _mapping(value: Any) -> dict[str, Any] | None:
if isinstance(value, dict):
if "@class" in value:
return {str(key): item for key, item in value.items() if key != "@class"}
return {str(key): item for key, item in value.items()}
if isinstance(value, list) and len(value) == 2 and isinstance(value[1], dict):
return {str(key): item for key, item in value[1].items()}
return None
async def redis_get(key: str, client: Redis | None = None) -> Any:
selected = client or get_redis()
raw = await cast(Any, selected.get(key))
return JavaRedisCodec.decode(raw)
async def redis_set(key: str, value: Any, *, ttl: int = DEFAULT_TTL_SECONDS, client: Redis | None = None) -> None:
selected = client or get_redis()
await cast(Any, selected.set(key, JavaRedisCodec.encode(value), ex=ttl))
async def redis_delete(*keys: str, client: Redis | None = None) -> None:
if not keys:
return
selected = client or get_redis()
await cast(Any, selected.delete(*keys))
async def redis_increment(key: str, *, ttl: int = DEFAULT_TTL_SECONDS, client: Redis | None = None) -> int:
selected = client or get_redis()
value = int(await cast(Any, selected.incr(key)))
await cast(Any, selected.expire(key, ttl))
return value
class DeviceService:
def __init__(
self,
session: AsyncSession,
*,
redis_client: Redis | None = None,
http_client: httpx.AsyncClient | None = None,
):
self.session = session
self.repository = DeviceRepository(session)
self.params = SystemParamService(session)
self.redis = redis_client
self.http_client = http_client
async def register_device(self, mac_address: str) -> str:
while True:
code = f"{secrets.randbelow(1_000_000):06d}"
key = f"sys:device:captcha:{code}"
if is_blank(cast(str | None, await redis_get(key, self.redis))):
await redis_set(key, mac_address, client=self.redis)
return code
async def activate_bound_device(self, *, agent_id: str, activation_code: str, user: AuthUser) -> None:
if is_blank(activation_code):
raise AppError(10061)
code_key = f"ota:activation:code:{activation_code}"
device_id_value = await redis_get(code_key, self.redis)
if device_id_value in (None, ""):
raise AppError(10062)
device_id = str(device_id_value)
safe_device_id = device_id.replace(":", "_").lower()
data_key = f"ota:activation:data:{safe_device_id}"
cached = _mapping(await redis_get(data_key, self.redis))
if cached is None or str(cached.get("activation_code") or "") != activation_code:
raise AppError(10062)
if await self.repository.get_device(device_id) is not None:
raise AppError(10063)
now = shanghai_now_naive()
values = {
"id": device_id,
"user_id": user.id,
"mac_address": cached.get("mac_address"),
"last_connected_at": now,
"auto_update": 1,
"board": cached.get("board"),
"alias": None,
"agent_id": agent_id,
"app_version": cached.get("app_version"),
"sort": None,
"updater": user.id,
"update_date": now,
"creator": user.id,
"create_date": now,
}
try:
await self.repository.insert_device(values)
await self.session.commit()
except Exception:
await self.session.rollback()
raise
await redis_delete(data_key, code_key, f"agent:device:count:{agent_id}", client=self.redis)
async def list_user_devices(self, user_id: int, agent_id: str) -> list[dict[str, Any]]:
devices = await self.repository.get_user_devices(user_id, agent_id)
return [self._user_device_view(row) for row in devices]
async def get_online_data(self, agent_id: str, user: AuthUser) -> str:
gateway = await self.params.get_value("server.mqtt_manager_api", from_cache=True)
if is_blank(gateway) or gateway == "null":
return ""
devices = await self.repository.get_user_devices(user.id, agent_id)
client_ids = {
self._mqtt_client_id(
str(device.get("board") or "GID_default"),
str(device.get("mac_address") or "unknown"),
)
for device in devices
}
if not client_ids:
return ""
signature_key = await self.params.get_value("server.mqtt_signature_key", from_cache=False)
return await post_json(
f"http://{gateway}/api/devices/status",
{"clientIds": sorted(client_ids)},
signature_key or "",
timeout_seconds=get_settings().external_request_timeout_seconds,
client=self.http_client,
)
async def unbind(self, *, user_id: int, device_id: str) -> None:
device = await self.repository.get_device(device_id)
if device is None:
return
mac_address = device.get("mac_address")
agent_id = device.get("agent_id")
if not is_blank(None if agent_id is None else str(agent_id)):
await redis_delete(f"agent:device:count:{agent_id}", client=self.redis)
try:
await self.repository.delete_device_for_user(device_id, user_id)
await self.session.commit()
except Exception:
await self.session.rollback()
raise
try:
if mac_address is not None:
await self.repository.delete_address_books_for_macs([str(mac_address)])
await self.session.commit()
except Exception:
await self.session.rollback()
raise
await self.refresh_address_book_cache()
async def update_device(
self,
*,
device_id: str,
request: DeviceUpdateRequest,
user: AuthUser,
) -> bool:
device = await self.repository.get_device(device_id)
if device is None or int(device.get("user_id") or -1) != user.id:
return False
await self.repository.update_device_info(
device_id,
auto_update=request.auto_update,
alias=request.alias,
updater=user.id,
now=shanghai_now_naive(),
)
await self.session.commit()
return True
async def manual_add(self, *, request: DeviceManualAddRequest, user: AuthUser) -> None:
mac_address = request.mac_address
if mac_address is not None and await self.repository.get_device_by_mac(mac_address) is not None:
raise AppError(10161)
now = shanghai_now_naive()
values = {
"id": uuid.uuid4().hex if mac_address in (None, "") else mac_address,
"user_id": user.id,
"mac_address": mac_address,
"last_connected_at": now,
"auto_update": 1,
"board": request.board,
"alias": None,
"agent_id": request.agent_id,
"app_version": request.app_version,
"sort": None,
"updater": user.id,
"update_date": now,
"creator": user.id,
"create_date": now,
}
try:
await self.repository.insert_device(values)
await self.session.commit()
except Exception:
await self.session.rollback()
raise
agent_cache_id = "null" if request.agent_id is None else request.agent_id
await redis_delete(f"agent:device:count:{agent_cache_id}", client=self.redis)
async def get_tools(self, *, device_id: str, user: AuthUser) -> dict[str, Any] | None:
gateway_and_device = await self._gateway_device(device_id, user)
if gateway_and_device is None:
return None
gateway, device = gateway_and_device
client_id = self._mqtt_client_id(
str(device.get("board") or "GID_default"),
str(device.get("mac_address") or "unknown"),
)
url = f"http://{gateway}/api/commands/{client_id}"
all_tools: list[Any] = []
cursor: str | None = None
while True:
params: dict[str, Any] = {"withUserTools": True}
if cursor is not None and cursor.strip():
params["cursor"] = cursor
body = {
"type": "mcp",
"payload": {"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": params},
}
response_body = await self._post_gateway(url, body)
if is_blank(response_body):
break
payload = json.loads(response_body)
if not isinstance(payload, dict) or not bool(payload.get("success", False)):
break
data = payload.get("data")
if not isinstance(data, dict):
break
tools = data.get("tools")
if isinstance(tools, list):
all_tools.extend(tools)
next_cursor = data.get("nextCursor")
if not isinstance(next_cursor, str) or not next_cursor.strip():
break
cursor = next_cursor
return None if not all_tools else {"tools": all_tools}
async def call_tool(
self,
*,
device_id: str,
tool_name: str,
arguments: dict[str, Any] | None,
user: AuthUser,
) -> Any:
gateway_and_device = await self._gateway_device(device_id, user)
if gateway_and_device is None:
return None
gateway, device = gateway_and_device
client_id = self._mqtt_client_id(
str(device.get("board") or "GID_default"),
str(device.get("mac_address") or "unknown"),
)
response_body = await self._post_gateway(
f"http://{gateway}/api/commands/{client_id}",
{
"type": "mcp",
"payload": {
"jsonrpc": "2.0",
"id": 2,
"method": "tools/call",
"params": {"name": tool_name, "arguments": arguments},
},
},
)
if is_blank(response_body):
return None
payload = json.loads(response_body)
if not isinstance(payload, dict) or not bool(payload.get("success", False)):
return None
data = payload.get("data")
content = data.get("content") if isinstance(data, dict) else None
if not isinstance(content, list) or not content or not isinstance(content[0], dict):
return None
first = content[0]
if first.get("type") != "text" or not isinstance(first.get("text"), str):
return None
text = str(first["text"])
if not text.strip():
return None
trimmed = text.strip()
if trimmed.startswith("{") or trimmed.startswith("["):
try:
parsed = json.loads(trimmed)
return parsed if isinstance(parsed, dict) else trimmed
except json.JSONDecodeError:
return trimmed
if trimmed == "true":
return True
if trimmed == "false":
return False
return trimmed
async def address_book(self, mac_address: str) -> list[dict[str, Any]]:
rows = await self.repository.get_address_book(mac_address)
for row in rows:
if row.get("has_permission") is not None:
row["has_permission"] = bool(row["has_permission"])
return rows
async def lookup_address_book(self, *, caller_mac: str, nickname: str) -> dict[str, str | None] | None:
books = await self.all_address_books()
caller_book = books.get(caller_mac.lower())
if caller_book is None:
return None
target_with_permission = caller_book.get(nickname)
if target_with_permission is None:
return None
parts = target_with_permission.split("|")
target_mac = parts[0]
has_permission = len(parts) > 1 and parts[1] == "1"
target_book = books.get(target_mac.lower())
if target_book is None:
return None
caller_nickname = target_book.get(caller_mac.lower())
return {
"targetMac": target_mac,
"callerNickname": caller_nickname,
"hasPermission": "true" if has_permission else "false",
}
async def call_by_nickname(self, *, caller_mac: str, nickname: str, answer: bool) -> dict[str, Any]:
books = await self.all_address_books()
if answer:
return await self._post_call("/api/call/accept", {"mac": caller_mac}, "接听")
caller_book = books.get(caller_mac.lower())
if caller_book is None or nickname not in caller_book:
return {"status": "error", "message": f"未找到备注为'{nickname}'的设备"}
parts = caller_book[nickname].split("|")
target_mac = parts[0]
allowed = len(parts) > 1 and parts[1] == "1"
if not allowed:
return {"status": "error", "message": "呼叫失败,您没有权限呼叫该设备"}
target_book = books.get(target_mac.lower())
caller_nickname = target_book.get(caller_mac.lower()) if target_book is not None else None
if is_blank(caller_nickname):
caller = await self.repository.get_device_by_mac(caller_mac)
if caller is None:
raise RuntimeError("caller device does not exist")
caller_nickname = None if caller.get("alias") is None else str(caller["alias"])
if is_blank(caller_nickname):
caller_nickname = self._mac_device_name(caller_mac)
return await self._post_call(
"/api/call/request",
{"caller_mac": caller_mac, "target_mac": target_mac, "caller_nickname": caller_nickname},
"呼叫",
)
async def save_address_book(
self,
*,
mac_address: str,
target_mac: str,
alias: str | None,
has_permission: bool | None,
actor: int,
) -> None:
record = await self.repository.get_address_book_record(mac_address, target_mac)
now = shanghai_now_naive()
if record is None:
final_alias = alias
if is_blank(final_alias):
target = await self.repository.get_device_by_mac(target_mac)
if target is None:
raise RuntimeError("target device does not exist")
final_alias = None if target.get("alias") is None else str(target["alias"])
final_alias = await self._unique_alias(mac_address, final_alias)
await self.repository.insert_address_book(
mac_address=mac_address,
target_mac=target_mac,
alias=final_alias,
has_permission=has_permission,
actor=actor,
now=now,
)
await self.session.commit()
else:
if alias is not None:
await self.repository.update_address_alias(
mac_address,
target_mac,
await self._unique_alias(mac_address, alias),
now=now,
)
await self.session.commit()
await self.refresh_address_book_cache()
if has_permission is not None:
await self.repository.update_address_permission(
mac_address,
target_mac,
has_permission,
now=now,
)
await self.session.commit()
await self.refresh_address_book_cache()
async def all_address_books(self) -> dict[str, dict[str, str]]:
cached = _mapping(await redis_get("device:address_book:all", self.redis))
if cached is not None:
result: dict[str, dict[str, str]] = {}
for key, value in cached.items():
nested = _mapping(value)
if nested is not None:
result[key] = {str(field): str(item) for field, item in nested.items()}
return result
return await self.refresh_address_book_cache()
async def refresh_address_book_cache(self) -> dict[str, dict[str, str]]:
records = await self.repository.get_all_address_book()
result: dict[str, dict[str, str]] = {}
reverse: dict[str, str] = {}
for record in records:
mac_a = str(record["mac_address"]).lower()
mac_b = str(record["target_mac"]).lower()
alias = record.get("alias")
if alias not in (None, ""):
alias_string = str(alias)
result.setdefault(mac_a, {})[alias_string] = (
f"{mac_b}|{'1' if bool(record.get('has_permission')) else '0'}"
)
reverse[f"{mac_b}:{mac_a}"] = alias_string
for record in records:
mac_a = str(record["mac_address"]).lower()
mac_b = str(record["target_mac"]).lower()
reverse_alias = reverse.get(f"{mac_a}:{mac_b}")
if isinstance(reverse_alias, str) and reverse_alias:
result.setdefault(mac_b, {})[mac_a] = reverse_alias
await redis_set("device:address_book:all", result, client=self.redis)
return result
async def check_ota(
self,
*,
device_id: str,
client_id: str,
report: DeviceReportRequest,
request_url: str,
client_ip: str,
defer_connection_update: Callable[[str, str | None, str | None], None] | None = None,
) -> dict[str, Any]:
now = datetime.now(tz=ZoneInfo(get_settings().timezone))
utc_offset = now.utcoffset()
response: dict[str, Any] = {
"server_time": {
"timestamp": int(now.timestamp() * 1000),
"timeZone": get_settings().timezone,
"timezone_offset": int((utc_offset.total_seconds() if utc_offset is not None else 0) / 60),
},
"activation": None,
"error": None,
"firmware": None,
"websocket": None,
"mqtt": None,
}
device = await self.repository.get_device_by_mac(device_id)
if device is None:
if report.application is None:
raise RuntimeError("application is required")
response["firmware"] = {
"version": report.application.version,
"url": INVALID_FIRMWARE_URL,
}
elif device.get("auto_update") is None:
raise RuntimeError("auto_update is null")
elif int(device["auto_update"]) != 0:
ota_type = report.board.type if report.board is not None else None
current_version = report.application.version if report.application is not None else None
response["firmware"] = await self._firmware_info(ota_type, current_version, request_url)
websocket_url = await self.params.get_value("server.websocket", from_cache=True)
auth_enabled = await self.params.get_value("server.auth.enabled", from_cache=True)
websocket_token = ""
if (auth_enabled or "").lower() == "true":
try:
websocket_token = await self._websocket_token(client_id, device_id)
except Exception:
logger.exception("WebSocket token generation failed")
if is_blank(websocket_url) or websocket_url == "null":
selected_websocket = "ws://xiaozhi.server.com:8000/xiaozhi/v1/"
else:
websocket_urls = _java_semicolon_split(websocket_url or "")
selected_websocket = (
random.choice(websocket_urls) # noqa: S311
if websocket_urls
else "ws://xiaozhi.server.com:8000/xiaozhi/v1/"
)
response["websocket"] = {"url": selected_websocket, "token": websocket_token}
mqtt_endpoint = await self.params.get_value("server.mqtt_gateway", from_cache=True)
if mqtt_endpoint not in (None, "", "null"):
try:
group_id = str(device.get("board") or "GID_default") if device is not None else "GID_default"
mqtt = await self._mqtt_config(device_id, group_id, client_ip)
if mqtt is not None:
mqtt["endpoint"] = mqtt_endpoint
response["mqtt"] = mqtt
except Exception:
logger.exception("MQTT credential generation failed")
if device is None:
response["activation"] = await self._activation(device_id, report)
else:
app_version = report.application.version if report.application is not None else None
agent_id = device.get("agent_id")
normalized_agent_id = None if agent_id is None else str(agent_id)
if defer_connection_update is not None:
defer_connection_update(str(device["id"]), normalized_agent_id, app_version)
else:
try:
await self._persist_connection_update(
str(device["id"]),
normalized_agent_id,
app_version,
)
except Exception:
logger.exception("Asynchronous device connection update failed")
return cast(dict[str, Any], self._drop_none(response))
async def _persist_connection_update(
self,
device_id: str,
agent_id: str | None,
app_version: str | None,
) -> None:
connection_time = shanghai_now_naive()
try:
await self.repository.update_connection(device_id, app_version=app_version, now=connection_time)
await self.session.commit()
except Exception:
await self.session.rollback()
raise
if not is_blank(agent_id):
await redis_set(f"agent:device:lastConnected:{agent_id}", connection_time, client=self.redis)
@staticmethod
async def persist_connection_update_background(
device_id: str,
agent_id: str | None,
app_version: str | None,
) -> None:
try:
async with get_session_factory()() as session:
await DeviceService(session)._persist_connection_update(device_id, agent_id, app_version)
except Exception:
logger.exception("Asynchronous device connection update failed")
async def ota_health_text(self) -> str:
mqtt_gateway = await self.params.get_value("server.mqtt_gateway", from_cache=False)
if is_blank(mqtt_gateway):
return "OTA接口不正常,缺少mqtt_gateway地址,请登录智控台,在参数管理找到【server.mqtt_gateway】配置"
websocket = await self.params.get_value("server.websocket", from_cache=True)
if is_blank(websocket) or websocket == "null":
return "OTA接口不正常,缺少websocket地址,请登录智控台,在参数管理找到【server.websocket】配置"
ota_url = await self.params.get_value("server.ota", from_cache=True)
if is_blank(ota_url) or ota_url == "null":
return "OTA接口不正常,缺少ota地址,请登录智控台,在参数管理找到【server.ota】配置"
return f"OTA接口运行正常,websocket集群数量:{len(_java_semicolon_split(websocket or ''))}"
async def ota_page(self, query: Mapping[str, Any]) -> dict[str, Any]:
page = self._positive_int(query.get("page"), 1)
limit = self._positive_int(query.get("limit"), 10)
requested = query.get("orderField")
requested_fields = [requested] if isinstance(requested, str) else list(requested or [])
fields = [OTA_ORDER_COLUMNS[field] for field in requested_fields if field in OTA_ORDER_COLUMNS]
if not fields:
fields = ["update_date"]
ascending = str(query.get("order") or "").lower() == "asc" if requested_fields else True
firmware_name = query.get("firmwareName")
name = str(firmware_name) if firmware_name is not None else None
rows = await self.repository.list_ota(
page=page,
limit=limit,
firmware_name=name,
order_fields=fields,
ascending=ascending,
)
rows = [self._ota_response_record(row) for row in rows]
return {"total": await self.repository.count_ota(name), "list": rows}
async def get_ota_record(self, ota_id: str) -> dict[str, Any] | None:
row = await self.repository.get_ota(ota_id)
return None if row is None else self._ota_response_record(row)
async def save_ota(self, record: OtaRecord, user: AuthUser) -> None:
values = record.model_dump(by_alias=False)
existing = await self.repository.get_first_ota_by_type(record.type or "")
now = shanghai_now_naive()
if existing is not None:
values["updater"] = record.updater if record.updater is not None else user.id
values["update_date"] = record.update_date if record.update_date is not None else now
await self.repository.update_ota(str(existing["id"]), values)
else:
values["id"] = record.id or uuid.uuid4().hex
values["creator"] = record.creator if record.creator is not None else user.id
values["updater"] = record.updater if record.updater is not None else user.id
values["create_date"] = record.create_date if record.create_date is not None else now
values["update_date"] = record.update_date if record.update_date is not None else now
await self.repository.insert_ota(values)
await self.session.commit()
async def update_ota(self, ota_id: str, record: OtaRecord, user: AuthUser) -> None:
if await self.repository.count_duplicate_ota(
ota_id=ota_id,
ota_type=record.type,
version=record.version,
):
raise RuntimeError("已存在相同类型和版本的固件,请修改后重试")
values = record.model_dump(by_alias=False)
values["updater"] = record.updater if record.updater is not None else user.id
values["update_date"] = shanghai_now_naive()
await self.repository.update_ota(ota_id, values)
await self.session.commit()
async def delete_ota(self, ids: Sequence[str]) -> None:
await self.repository.delete_ota(ids)
await self.session.commit()
async def create_ota_download_id(self, ota_id: str) -> str:
value = str(uuid.uuid4())
await redis_set(f"ota:id:{value}", ota_id, client=self.redis)
return value
async def resolve_ota_download(self, download_id: str) -> tuple[Path, str] | None:
id_key = f"ota:id:{download_id}"
ota_value = await redis_get(id_key, self.redis)
if is_blank(None if ota_value is None else str(ota_value)):
return None
count_key = f"ota:download:count:{download_id}"
count_value = await redis_get(count_key, self.redis)
count = int(count_value or 0)
if count >= 3:
await redis_delete(count_key, id_key, client=self.redis)
return None
await redis_set(count_key, count + 1, client=self.redis)
ota_id = str(ota_value)
if ota_id.startswith("file:"):
firmware_path = ota_id[5:]
ota_type = "assets"
version = "1.0.0"
else:
record = await self.repository.get_ota(ota_id)
firmware_value = None if record is None else record.get("firmware_path")
if record is None or is_blank(None if firmware_value is None else str(firmware_value)):
return None
firmware_path = str(record["firmware_path"])
ota_type = str(record.get("type"))
version = str(record.get("version"))
raw_path = Path(firmware_path)
candidates = [raw_path] if raw_path.is_absolute() else [Path.cwd() / raw_path]
if not raw_path.is_absolute() and raw_path.parts and raw_path.parts[0] == "uploadfile":
candidates.insert(0, get_settings().upload_dir.joinpath(*raw_path.parts[1:]))
candidates.append(Path.cwd() / "firmware" / raw_path.name)
resolved = next((candidate for candidate in candidates if candidate.is_file()), None)
if resolved is None:
return None
original_name = f"{ota_type}_{version}"
dot_index = firmware_path.rfind(".")
if dot_index >= 0:
original_name += firmware_path[dot_index:]
safe_name = re.sub(r"[^a-zA-Z0-9._-]", "_", original_name)
return resolved, safe_name
async def save_firmware_file(self, *, filename: str | None, content: bytes) -> str:
if not content:
raise ValueError("上传文件不能为空")
if filename is None:
raise ValueError("文件名不能为空")
dot_index = filename.rfind(".")
if dot_index < 0:
raise RuntimeError("文件名缺少扩展名")
extension = filename[dot_index:].lower()
if extension not in {".bin", ".apk"}:
raise ValueError("只允许上传.bin和.apk格式的文件")
digest = hashlib.md5(content, usedforsecurity=False).hexdigest()
directory = get_settings().upload_dir
directory.mkdir(parents=True, exist_ok=True)
filename_on_disk = f"{digest}{extension}"
physical_path = directory / filename_on_disk
if not physical_path.exists():
with physical_path.open("xb") as stream:
stream.write(content)
# Keep Java's database/API value stable even when the physical upload
# volume is mounted elsewhere (for example /data/uploads in Docker).
return str(Path("uploadfile") / filename_on_disk)
async def save_assets_file(self, *, filename: str | None, content: bytes, user: AuthUser) -> str:
ota_url = await self.params.get_value("server.ota", from_cache=True)
if is_blank(ota_url) or ota_url == "null":
raise AppError(10102)
if len(content) > 20 * 1024 * 1024:
raise AppError(10142)
if not user.is_super_admin:
count_key = f"ota:upload:count:{user.id}"
current = int(await redis_get(count_key, self.redis) or 0)
if current >= 50:
raise AppError(10195)
await redis_increment(count_key, client=self.redis)
path = await self.save_firmware_file(filename=filename, content=content)
download_id = await self.create_ota_download_id(f"file:{path}")
return (ota_url or "").replace("/ota/", "/otaMag/download/") + download_id
async def _gateway_device(self, device_id: str, user: AuthUser) -> tuple[str, dict[str, Any]] | None:
gateway = await self.params.get_value("server.mqtt_manager_api", from_cache=True)
if is_blank(gateway) or gateway == "null":
return None
device = await self.repository.get_device(device_id)
if device is None or int(device.get("user_id") or -1) != user.id:
return None
return gateway or "", device
async def _post_gateway(self, url: str, body: Any, *, timeout_seconds: float | None = None) -> str:
key = await self.params.get_value("server.mqtt_signature_key", from_cache=False)
return await post_json(
url,
body,
key or "",
timeout_seconds=timeout_seconds or get_settings().external_request_timeout_seconds,
client=self.http_client,
)
async def _post_call(self, path: str, body: dict[str, Any], action: str) -> dict[str, Any]:
gateway = await self.params.get_value("server.mqtt_manager_api", from_cache=True)
key = await self.params.get_value("server.mqtt_signature_key", from_cache=True)
if is_blank(gateway) or gateway == "null" or is_blank(key) or (key or "").strip().lower() == "null":
return {"status": "error", "message": f"{action}失败,网关配置缺失"}
result: dict[str, Any] = {"status": "error"}
try:
text = await post_json(
f"http://{gateway}{path}",
body,
key or "",
timeout_seconds=5.0,
client=self.http_client,
)
if text.strip():
payload = json.loads(text)
if isinstance(payload, dict):
result["status"] = payload.get("status")
result["message"] = payload.get("message")
return result
except Exception:
return {"status": "error", "message": f"{action}失败,请稍后再试"}
async def _firmware_info(
self,
ota_type: str | None,
current_version: str | None,
request_url: str,
) -> dict[str, Any] | None:
if is_blank(ota_type):
return None
selected_version = current_version if not is_blank(current_version) else "0.0.0"
ota = await self.repository.get_latest_ota(ota_type or "")
download_url: str | None = None
if ota is not None and self._compare_versions(ota.get("version"), selected_version) > 0:
ota_url = await self.params.get_value("server.ota", from_cache=True)
if is_blank(ota_url) or ota_url == "null":
ota_url = request_url
download_id = await self.create_ota_download_id(str(ota["id"]))
download_url = (ota_url or "").replace("/ota/", "/otaMag/download/") + download_id
return {
"version": selected_version if ota is None else ota.get("version"),
"url": download_url or INVALID_FIRMWARE_URL,
}
async def _activation(self, device_id: str, report: DeviceReportRequest) -> dict[str, Any]:
safe_device_id = device_id.replace(":", "_").lower()
data_key = f"ota:activation:data:{safe_device_id}"
cached = _mapping(await redis_get(data_key, self.redis))
code = str(cached.get("activation_code")) if cached and cached.get("activation_code") is not None else None
frontend = await self.params.get_value("server.fronted_url", from_cache=True)
if code is None or not code.strip():
code = f"{secrets.randbelow(1_000_000):06d}"
board = (
report.board.type
if report.board is not None and report.board.type is not None
else (report.chip_model_name or "unknown")
)
app_version = report.application.version if report.application is not None else None
await redis_set(
data_key,
{
"id": device_id,
"mac_address": device_id,
"board": board,
"app_version": app_version,
"deviceId": device_id,
"activation_code": code,
},
client=self.redis,
)
await redis_set(f"ota:activation:code:{code}", device_id, client=self.redis)
return {
"code": code,
"message": f"{frontend if frontend is not None else 'null'}\n{code}",
"challenge": device_id,
}
async def _websocket_token(self, client_id: str, username: str) -> str:
secret = await self.params.get_value("server.secret", from_cache=False)
if is_blank(secret):
raise RuntimeError("WebSocket认证密钥未配置(server.secret)")
timestamp = int(datetime.now().timestamp())
message = f"{client_id}|{username}|{timestamp}".encode()
signature = hmac.new((secret or "").encode(), message, hashlib.sha256).digest()
encoded = base64.urlsafe_b64encode(signature).decode().rstrip("=")
return f"{encoded}.{timestamp}"
async def _mqtt_config(self, mac_address: str, group_id: str, client_ip: str) -> dict[str, Any] | None:
key = await self.params.get_value("server.mqtt_signature_key", from_cache=True)
if is_blank(key):
return None
client_id = self._mqtt_client_id(group_id, mac_address)
user_data = json.dumps({"ip": client_ip}, ensure_ascii=False, separators=(",", ":"))
username = base64.b64encode(user_data.encode()).decode()
password = base64.b64encode(
hmac.new((key or "").encode(), f"{client_id}|{username}".encode(), hashlib.sha256).digest()
).decode()
safe_mac = mac_address.replace(":", "_")
return {
"client_id": client_id,
"username": username,
"password": password,
"publish_topic": "device-server",
"subscribe_topic": f"devices/p2p/{safe_mac}",
}
@staticmethod
def _mqtt_client_id(group_id: str, mac_address: str) -> str:
safe_group = group_id.replace(":", "_")
safe_mac = mac_address.replace(":", "_")
return f"{safe_group}@@@{safe_mac}@@@{safe_mac}"
@staticmethod
def _compare_versions(first: Any, second: Any) -> int:
if first is None or second is None:
return 0
first = str(first)
second = str(second)
first_parts = first.split(".")
second_parts = second.split(".")
for index in range(max(len(first_parts), len(second_parts))):
first_value = int(first_parts[index]) if index < len(first_parts) else 0
second_value = int(second_parts[index]) if index < len(second_parts) else 0
if first_value != second_value:
return 1 if first_value > second_value else -1
return 0
async def _unique_alias(self, mac_address: str, alias: str | None) -> str | None:
existing = await self.repository.get_aliases(mac_address)
if alias not in existing:
return alias
suffix = 1
while f"{alias}{suffix}" in existing:
suffix += 1
return f"{alias}{suffix}"
@staticmethod
def _mac_device_name(mac: str) -> str:
return mac if len(mac) < 2 else f"尾号为{mac[-2:]}的设备"
@staticmethod
def _positive_int(value: Any, default: int) -> int:
if value is None:
return default
return int(str(value))
@staticmethod
def _drop_none(value: Any) -> Any:
if isinstance(value, dict):
return {key: DeviceService._drop_none(item) for key, item in value.items() if item is not None}
if isinstance(value, list):
return [DeviceService._drop_none(item) for item in value]
return value
@staticmethod
def _user_device_view(row: Mapping[str, Any]) -> dict[str, Any]:
return {
"app_version": row.get("app_version"),
"bind_user_name": None,
"device_type": row.get("board"),
"board": row.get("board"),
"id": row.get("id"),
"mac_address": row.get("mac_address"),
"alias": row.get("alias"),
"ota_upgrade": None,
"recent_chat_time": None,
"last_connected_at_timestamp": DeviceService._timestamp(row.get("last_connected_at")),
"create_date_timestamp": DeviceService._timestamp(row.get("create_date")),
# UserShowDeviceListVO pins only this field to UTC. The companion
# epoch value still uses the configured Asia/Shanghai instant.
"create_date": DeviceService._utc_datetime(row.get("create_date")),
}
@staticmethod
def _utc_datetime(value: Any) -> Any:
if not isinstance(value, datetime):
return value
localized = value.replace(tzinfo=ZoneInfo(get_settings().timezone)) if value.tzinfo is None else value
return localized.astimezone(timezone.utc).replace(tzinfo=None)
@staticmethod
def _ota_response_record(row: Mapping[str, Any]) -> dict[str, Any]:
result = dict(row)
if result.get("size") is not None:
result["size"] = str(result["size"])
return result
@staticmethod
def _timestamp(value: Any) -> int | None:
if not isinstance(value, datetime):
return None
localized = value.replace(tzinfo=ZoneInfo(get_settings().timezone)) if value.tzinfo is None else value
return int(localized.timestamp() * 1000)
@@ -0,0 +1,21 @@
from __future__ import annotations
from functools import lru_cache
from pathlib import Path
from app.core.config import get_settings
from app.core.i18n import LANGUAGE_FILES, _load_properties, resolve_language
@lru_cache(maxsize=16)
def _validation_messages(language: str, directory: str) -> dict[str, str]:
root = Path(directory)
values = _load_properties(root / "validation.properties")
localized = LANGUAGE_FILES[language].replace("messages_", "validation_")
values.update(_load_properties(root / localized))
return values
def validation_message(key: str, accept_language: str | None) -> str:
language = resolve_language(accept_language)
return _validation_messages(language, str(get_settings().i18n_dir)).get(key, key)
@@ -0,0 +1,723 @@
from __future__ import annotations
import json
from collections import defaultdict
from datetime import datetime
from typing import Any
from zoneinfo import ZoneInfo
from fastapi import UploadFile
from app.core.config import get_settings
from app.core.errors import AppError
from app.core.i18n import message_for
from app.core.redis import get_redis
from app.core.security import AuthUser, shanghai_now_naive
from app.core.serialization import preserve_java_map_keys
from app.integrations.ragflow import RAGFlowClient
from app.repositories.knowledge import KnowledgeRepository
from app.schemas.knowledge import KnowledgeBaseBody, RetrievalBody
def dataset_dto(row: dict[str, Any]) -> dict[str, Any]:
return {
"id": row.get("id"),
"datasetId": row.get("dataset_id"),
"ragModelId": row.get("rag_model_id"),
"name": row.get("name"),
"avatar": row.get("avatar"),
"description": row.get("description"),
"embeddingModel": row.get("embedding_model"),
"permission": row.get("permission"),
"chunkMethod": row.get("chunk_method"),
"parserConfig": row.get("parser_config"),
"chunkCount": None if row.get("chunk_count") is None else str(row["chunk_count"]),
"tokenNum": None if row.get("token_num") is None else str(row["token_num"]),
"status": row.get("status"),
"creator": row.get("creator"),
"createdAt": row.get("created_at"),
"updater": row.get("updater"),
"updatedAt": row.get("updated_at"),
# KnowledgeBaseEntity.documentCount is Long while KnowledgeBaseDTO uses
# Integer. Spring BeanUtils does not coerce that property, so local DTO
# conversion leaves it null; list enrichment fills it from RAGFlow.
"documentCount": None,
"errorMessage": row.get("error_message"),
}
def document_dto(row: dict[str, Any]) -> dict[str, Any]:
return {
"id": row.get("document_id"),
"documentId": row.get("document_id"),
"datasetId": row.get("dataset_id"),
"name": row.get("name"),
# RAGFlowAdapter.mapToKnowledgeFilesDTO does not populate these two
# fields for the immediate upload response.
"fileType": None,
"fileSize": row.get("size"),
"filePath": None,
"progress": row.get("progress"),
"thumbnail": row.get("thumbnail"),
"processDuration": row.get("process_duration"),
"sourceType": row.get("source_type"),
"metaFields": _json_object(row.get("meta_fields")),
"chunkMethod": row.get("chunk_method"),
"parserConfig": _json_object(row.get("parser_config")),
"status": row.get("status"),
"run": row.get("run"),
"creator": row.get("creator"),
"createdAt": row.get("created_at"),
"updater": None,
"updatedAt": row.get("updated_at"),
"chunkCount": row.get("chunk_count"),
"tokenCount": row.get("token_count"),
"error": row.get("error"),
"parseStatusCode": _parse_status(row.get("run")),
}
def remote_document_dto(row: dict[str, Any], dataset_id: str) -> dict[str, Any]:
run = row.get("run")
return {
"id": row.get("id"),
"documentId": row.get("id"),
"datasetId": row.get("dataset_id") or dataset_id,
"name": row.get("name"),
"fileType": row.get("type"),
"fileSize": row.get("size"),
"filePath": None,
"progress": row.get("progress"),
"thumbnail": row.get("thumbnail"),
"processDuration": row.get("process_duration"),
"sourceType": row.get("source_type"),
"metaFields": row.get("meta_fields"),
"chunkMethod": row.get("chunk_method"),
"parserConfig": row.get("parser_config"),
"status": _remote_status(row.get("status")),
"run": run,
"creator": None,
"createdAt": _millis(row.get("create_time")),
"updater": None,
"updatedAt": _millis(row.get("update_time")),
"chunkCount": row.get("chunk_count") or 0,
"tokenCount": row.get("token_count"),
"error": row.get("progress_msg"),
"parseStatusCode": _parse_status(run),
}
def _parse_status(run: Any) -> int:
return {"RUNNING": 1, "CANCEL": 2, "DONE": 3, "FAIL": 4}.get(str(run or "").upper(), 0)
def _json_object(value: Any) -> dict[str, Any] | None:
if value is None:
return None
if isinstance(value, dict):
return dict(value)
try:
parsed = json.loads(value.decode() if isinstance(value, bytes) else str(value))
return dict(parsed) if isinstance(parsed, dict) else None
except (ValueError, TypeError):
return None
def _millis(value: Any) -> Any:
try:
if value is None:
return None
timezone = ZoneInfo(get_settings().timezone)
return datetime.fromtimestamp(float(value) / 1000, timezone).replace(tzinfo=None)
except (TypeError, ValueError, OSError):
return None
def _is_blank(value: str | None) -> bool:
return value is None or not value.strip()
def _remote_status(value: Any) -> str:
if value is None or (isinstance(value, str) and not value.strip()):
return "1"
return str(value)
class KnowledgeBaseService:
def __init__(self, repository: KnowledgeRepository):
self.repository = repository
async def get_owned(self, identifier: str, user: AuthUser) -> dict[str, Any]:
if not identifier.strip():
raise AppError(10003)
row = await self.repository.get_dataset(identifier)
if row is None:
raise AppError(10163)
if row.get("creator") is None or int(row["creator"]) != user.id:
raise AppError(10169)
return row
async def page(
self,
user: AuthUser,
name: str | None,
page: int,
page_size: int,
language: str | None = None,
) -> dict[str, Any]:
rows, total = await self.repository.dataset_page(
user.id, name, (max(page, 1) - 1) * page_size, page_size
)
results: list[dict[str, Any]] = []
changed = False
for row in rows:
dto = dataset_dto(row)
if row.get("dataset_id") and row.get("rag_model_id"):
try:
client = await self._client(str(row["rag_model_id"]))
remote = await client.dataset_info(str(row["dataset_id"]))
if remote is None:
await self.repository.execute(
"DELETE FROM ai_rag_knowledge_document WHERE dataset_id=:dataset_id",
{"dataset_id": row["dataset_id"]},
)
await self.repository.delete_dataset_local(row)
await _delete_cache_ignoring_errors(f"knowledge:base:{row['id']}")
changed = True
continue
remote_name = remote.get("name")
local_name = (
str(remote_name).split("_", 1)[1]
if remote_name and "_" in str(remote_name)
else remote_name
)
updates: dict[str, Any] = {}
if local_name and local_name != row.get("name"):
updates["name"] = local_name
dto["name"] = local_name
if remote.get("description") != row.get("description"):
updates["description"] = remote.get("description")
dto["description"] = remote.get("description")
if updates:
await self.repository.execute(
"UPDATE ai_rag_dataset SET name=COALESCE(:name,name),description=:description WHERE id=:id",
{
"name": updates.get("name"),
"description": updates.get("description", row.get("description")),
"id": row["id"],
},
)
changed = True
if remote.get("document_count") is not None:
dto["documentCount"] = int(remote["document_count"])
except Exception as exc:
dto["documentCount"] = 0
dto["errorMessage"] = (
message_for(exc.code, language, *exc.params)
if isinstance(exc, AppError)
else str(exc)
)
results.append(dto)
if changed:
await self.repository.session.commit()
return {"total": total, "list": results}
async def create(self, body: KnowledgeBaseBody, user: AuthUser) -> dict[str, Any]:
if not _is_blank(body.name) and await self.repository.duplicate_dataset_name(user.id, str(body.name)):
raise AppError(10170)
rag_model_id = body.rag_model_id
if _is_blank(rag_model_id):
models = await self.repository.rag_models()
if not models:
raise AppError(10164, params=("未指定且无可用默认 RAG 模型",))
rag_model_id = str(models[0]["id"])
client = await self._client(str(rag_model_id))
create_body = {
"name": f"{user.username}_{'null' if body.name is None else body.name}",
"avatar": body.avatar,
"description": body.description,
"embedding_model": body.embedding_model,
"permission": body.permission,
"chunk_method": body.chunk_method,
# KnowledgeBaseDTO.parserConfig is a String, while CreateReq uses
# ParserConfig. BeanUtils skips the incompatible property.
"parser_config": None,
}
remote = await client.create_dataset(create_body)
dataset_id = str(remote["id"])
now = shanghai_now_naive()
created_at = body.created_at or now
updated_at = body.updated_at or now
values = {
"id": dataset_id,
"dataset_id": dataset_id,
"rag_model_id": rag_model_id,
"tenant_id": remote.get("tenant_id"),
"name": body.name,
"avatar": remote.get("avatar") if _is_blank(body.avatar) else body.avatar,
"description": body.description,
"embedding_model": remote.get("embedding_model"),
"permission": remote.get("permission"),
"chunk_method": remote.get("chunk_method"),
"parser_config": json.dumps(
remote.get("parser_config"), ensure_ascii=False, separators=(",", ":")
)
if remote.get("parser_config") is not None
else None,
"chunk_count": remote.get("chunk_count") or 0,
"document_count": remote.get("document_count") or 0,
"token_num": remote.get("token_num") or 0,
"status": 1,
"creator": user.id,
"updater": user.id,
"created_at": created_at,
"updated_at": updated_at,
}
try:
await self.repository.insert_dataset(values)
await self.repository.session.commit()
except Exception as exc:
await self.repository.session.rollback()
try:
await client.delete_datasets([dataset_id])
except AppError:
pass
if isinstance(exc, AppError):
raise
raise AppError(10167, params=(f"创建知识库失败: {exc}",)) from exc
return dataset_dto(values)
async def update(
self, identifier: str, body: KnowledgeBaseBody, user: AuthUser
) -> dict[str, Any]:
existing = await self.get_owned(identifier, user)
if not _is_blank(body.name) and await self.repository.duplicate_dataset_name(
user.id, str(body.name), str(existing["id"])
):
raise AppError(10170)
if not _is_blank(identifier) and await self.repository.dataset_id_conflict(
identifier, str(existing["id"])
):
raise AppError(10002)
rag_model_id = body.rag_model_id
effective_permission = body.permission
effective_chunk_method = body.chunk_method
if existing.get("dataset_id") and not _is_blank(rag_model_id):
if _is_blank(effective_permission):
effective_permission = existing.get("permission")
if _is_blank(effective_chunk_method):
effective_chunk_method = existing.get("chunk_method")
client = await self._client(str(rag_model_id))
remote_body = {
"name": f"{user.username}_{body.name}" if not _is_blank(body.name) else None,
"avatar": body.avatar,
"description": body.description,
"embedding_model": body.embedding_model,
"permission": effective_permission,
"chunk_method": effective_chunk_method,
"parser_config": _json_object(body.parser_config),
}
await client.update_dataset(str(existing["dataset_id"]), remote_body)
now = shanghai_now_naive()
updater = body.updater if body.updater is not None else user.id
updated_at = body.updated_at or now
values = {
"id": existing["id"],
# The controller injects the literal path value into datasetId,
# even when a legacy row was found through its local primary key.
"dataset_id": identifier,
"rag_model_id": rag_model_id,
"name": body.name,
"avatar": body.avatar,
"description": body.description,
"embedding_model": body.embedding_model,
"permission": effective_permission,
"chunk_method": effective_chunk_method,
"parser_config": body.parser_config,
"chunk_count": body.chunk_count,
"token_num": body.token_num,
"status": body.status,
"creator": body.creator,
"created_at": body.created_at,
"updater": updater,
"updated_at": updated_at,
}
try:
await self.repository.update_dataset(values)
# Java performs cache eviction inside the database transaction;
# an eviction failure therefore rolls this update back.
await get_redis().delete(f"knowledge:base:{existing['id']}")
await self.repository.session.commit()
except Exception:
await self.repository.session.rollback()
raise
# BeanUtils copies request nulls onto the in-memory entity before
# MyBatis' NOT_NULL update strategy preserves the stored columns. The
# Java response is built from that in-memory entity, so its null fields
# intentionally differ from a subsequent GET of the row.
return dataset_dto(values)
async def delete(self, identifier: str, user: AuthUser, language: str | None = None) -> None:
row = await self.get_owned(identifier, user)
documents = await self.repository.all_documents(str(row["dataset_id"]))
if documents:
# Java's document orchestration necessarily resolves the adapter
# when child records exist.
client = await self._client(str(row.get("rag_model_id") or ""))
ids = [str(item["document_id"]) for item in documents]
if any(item.get("run") == "RUNNING" for item in documents):
raise AppError(10199)
try:
await client.delete_documents(str(row["dataset_id"]), ids)
except Exception as exc:
raise _document_delete_error(exc, language) from exc
await self.repository.delete_document_shadows(str(row["dataset_id"]), ids)
await self.repository.update_stats(
str(row["dataset_id"]),
-len(ids),
-sum(int(item.get("chunk_count") or 0) for item in documents),
-sum(int(item.get("token_count") or 0) for item in documents),
)
# deleteDocuments is NOT_SUPPORTED in Java and its shadow cleanup
# commits before the outer dataset transaction continues.
await self.repository.session.commit()
await _delete_cache_ignoring_errors(f"knowledge:base:{row['dataset_id']}")
if not _is_blank(row.get("rag_model_id")) and not _is_blank(row.get("dataset_id")):
client = await self._client(str(row["rag_model_id"]))
await client.delete_datasets([str(row["dataset_id"])])
await self.repository.delete_dataset_local(row)
try:
await get_redis().delete(f"knowledge:base:{row['id']}")
await self.repository.session.commit()
except Exception:
await self.repository.session.rollback()
raise
async def batch_delete(
self, identifiers: list[str], user: AuthUser, language: str | None = None
) -> None:
rows = await self.repository.datasets_by_ids(identifiers)
for row in rows:
if row.get("creator") is None or int(row["creator"]) != user.id:
raise AppError(10169)
# Preserve Java's sequential external calls and stop-on-first-error semantics.
for row in rows:
await self.delete(str(row["dataset_id"]), user, language)
async def rag_models(self) -> list[dict[str, Any]]:
rows = await self.repository.rag_models()
result: list[dict[str, Any]] = []
for row in rows:
result.append(
{
"id": row.get("id"),
"modelType": None,
"modelCode": None,
"modelName": row.get("model_name"),
"isDefault": None,
"isEnabled": None,
# ModelConfigEntity.configJson is a JSONObject. Jackson
# preserves its dynamic snake_case keys instead of applying
# the DTO property naming strategy recursively.
"configJson": preserve_java_map_keys(_json_object(row.get("config_json"))),
"docLink": None,
"remark": None,
"sort": None,
"updater": None,
"updateDate": None,
"creator": None,
"createDate": None,
}
)
return result
async def _client(self, model_id: str) -> RAGFlowClient:
config = await self.repository.rag_config(model_id)
adapter_type = config.get("type")
if adapter_type != "ragflow":
raise AppError(10184, params=(f"适配器类型未注册: {adapter_type}",))
try:
return RAGFlowClient(config)
except AppError as exc:
# KnowledgeBaseAdapterFactory wraps adapter initialization and
# validateConfig failures as RAG_ADAPTER_CREATION_FAILED.
if exc.code in {10171, 10172, 10173, 10174}:
raise AppError(10186) from exc
raise
class KnowledgeDocumentService:
def __init__(self, repository: KnowledgeRepository):
self.repository = repository
self.datasets = KnowledgeBaseService(repository)
async def page(
self,
dataset_id: str,
user: AuthUser,
*,
name: str | None,
status: str | None,
page: int,
page_size: int,
) -> dict[str, Any]:
await self.datasets.get_owned(dataset_id, user)
try:
await self.reconcile(dataset_id, creator=user.id)
except Exception:
await self.repository.session.rollback()
rows, total = await self.repository.documents_page(
dataset_id,
name=name,
status=status,
offset=(max(page, 1) - 1) * page_size,
limit=page_size,
)
return {"total": total, "list": [document_dto(row) for row in rows]}
async def upload(
self,
dataset_id: str,
user: AuthUser,
file: UploadFile,
*,
name: str | None,
meta_fields: dict[str, Any] | None,
chunk_method: str | None,
parser_config: dict[str, Any] | None,
) -> dict[str, Any]:
await self.datasets.get_owned(dataset_id, user)
content = await file.read()
if not dataset_id.strip() or not content:
raise AppError(10003)
file_name = file.filename if _is_blank(name) else name
if _is_blank(file_name):
raise AppError(10179)
assert file_name is not None
client = await self._client_for_dataset(dataset_id)
remote = await client.upload_document(
dataset_id,
file,
content,
name=file_name,
meta_fields=meta_fields,
chunk_method=chunk_method,
parser_config=parser_config,
)
if not remote.get("id"):
raise AppError(10167, params=("远程上传成功但未返回有效 DocumentID",))
remote.setdefault("dataset_id", dataset_id)
shadow = dict(remote)
if _is_blank(str(shadow.get("name")) if shadow.get("name") is not None else None):
shadow["name"] = file_name
# Java stores the original controller values in the shadow row, even
# when invalid chunk methods were omitted from the RAGFlow request.
shadow["chunk_method"] = chunk_method
shadow["parser_config"] = parser_config
inserted = await self.repository.upsert_document(dataset_id, shadow, creator=user.id)
if inserted:
await self.repository.update_stats(dataset_id, 1, 0, 0)
await self.repository.session.commit()
return remote_document_dto(remote, dataset_id)
async def delete(
self,
dataset_id: str,
ids: list[str] | None,
user: AuthUser,
language: str | None = None,
) -> None:
await self.datasets.get_owned(dataset_id, user)
if not ids:
raise AppError(10178)
rows = await self.repository.documents_by_remote_ids(dataset_id, ids)
if len(rows) != len(ids):
raise AppError(10169)
if any(row.get("run") == "RUNNING" for row in rows):
raise AppError(10199)
chunks = sum(int(row.get("chunk_count") or 0) for row in rows)
tokens = sum(int(row.get("token_count") or 0) for row in rows)
client = await self._client_for_dataset(dataset_id)
try:
await client.delete_documents(dataset_id, ids)
except Exception as exc:
raise _document_delete_error(exc, language) from exc
deleted = await self.repository.delete_document_shadows(dataset_id, ids)
if deleted:
await self.repository.update_stats(dataset_id, -len(ids), -chunks, -tokens)
await self.repository.session.commit()
await _delete_cache_ignoring_errors(f"knowledge:base:{dataset_id}")
async def parse(self, dataset_id: str, ids: list[str], user: AuthUser) -> bool:
await self.datasets.get_owned(dataset_id, user)
if not ids:
raise AppError(10178)
client = await self._client_for_dataset(dataset_id)
await client.parse_documents(dataset_id, ids)
await self.repository.mark_documents_running(dataset_id, ids, shanghai_now_naive())
await self.repository.session.commit()
return True
async def chunks(
self,
dataset_id: str,
document_id: str,
user: AuthUser,
*,
page: int,
page_size: int,
keywords: str | None,
chunk_id: str | None,
) -> dict[str, Any]:
await self.datasets.get_owned(dataset_id, user)
client = await self._client_for_dataset(dataset_id)
return await client.chunks(
dataset_id,
document_id,
{"page": page, "page_size": page_size, "keywords": keywords, "id": chunk_id},
)
async def retrieval(self, dataset_id: str, body: RetrievalBody, user: AuthUser) -> dict[str, Any]:
await self.datasets.get_owned(dataset_id, user)
dataset_ids = body.dataset_ids or [dataset_id]
if not dataset_ids:
raise AppError(500, "未指定召回测试的知识库")
page = body.page if body.page is not None and body.page >= 1 else 1
page_size = body.page_size if body.page_size is not None and body.page_size >= 1 else 100
top_k = body.top_k if body.top_k is None or body.top_k >= 1 else 1024
threshold = body.similarity_threshold
if threshold is not None:
threshold = 0.2 if threshold < 0 else min(threshold, 1.0)
payload: dict[str, Any] = {
"dataset_ids": dataset_ids,
"document_ids": body.document_ids,
"question": body.question,
"page": page,
"page_size": page_size,
"similarity_threshold": threshold,
"vector_similarity_weight": body.vector_similarity_weight,
"top_k": top_k,
"rerank_id": body.rerank_id,
"highlight": body.highlight,
"keyword": body.keyword,
"cross_languages": body.cross_languages,
"metadata_condition": body.metadata_condition,
}
payload = {key: value for key, value in payload.items() if value is not None}
client = await self._client_for_dataset(dataset_ids[0])
return await client.retrieval(payload)
async def reconcile(self, dataset_id: str, *, creator: int | None = None) -> int:
client = await self._client_for_dataset(dataset_id)
remote: list[dict[str, Any]] = []
page, total = 1, 2**63 - 1
while (page - 1) * 100 < total:
rows, total = await client.documents(dataset_id, page=page, page_size=100)
if not rows:
break
remote.extend(rows)
page += 1
local = await self.repository.all_documents(dataset_id)
remote_map = {str(item.get("id")): item for item in remote if item.get("id")}
local_map = {str(item["document_id"]): item for item in local}
new_count = 0
for document_id, item in remote_map.items():
prior = local_map.get(document_id)
inserted = await self.repository.upsert_document(dataset_id, item, creator=creator)
if inserted:
new_count += 1
await self.repository.update_stats(
dataset_id, 1, int(item.get("chunk_count") or 0), int(item.get("token_count") or 0)
)
elif prior:
await self.repository.update_stats(
dataset_id,
0,
int(item.get("chunk_count") or 0) - int(prior.get("chunk_count") or 0),
int(item.get("token_count") or 0) - int(prior.get("token_count") or 0),
)
deleted_ids = [identifier for identifier in local_map if identifier not in remote_map]
if deleted_ids:
deleted_rows = [local_map[identifier] for identifier in deleted_ids]
await self.repository.delete_document_shadows(dataset_id, deleted_ids)
await self.repository.update_stats(
dataset_id,
-len(deleted_ids),
-sum(int(row.get("chunk_count") or 0) for row in deleted_rows),
-sum(int(row.get("token_count") or 0) for row in deleted_rows),
)
await self.repository.session.commit()
return new_count
async def sync_running(self) -> int:
rows = await self.repository.running_documents()
grouped: defaultdict[str, list[dict[str, Any]]] = defaultdict(list)
for row in rows:
grouped[str(row["dataset_id"])].append(row)
updates = 0
for dataset_id, documents in grouped.items():
try:
client = await self._client_for_dataset(dataset_id)
except Exception:
await self.repository.session.rollback()
continue
for local in documents:
try:
remote, _ = await client.documents(
dataset_id, page=1, page_size=1, document_id=str(local["document_id"])
)
if not remote:
await self.repository.mark_document_remote_deleted(
str(local["document_id"]), shanghai_now_naive()
)
await self.repository.session.commit()
updates += 1
continue
remote_status = remote[0].get("status")
remote_run = remote[0].get("run")
status_changed = remote_status is not None and str(remote_status) != str(local.get("status"))
run_changed = remote_run is not None and str(remote_run) != str(local.get("run"))
is_processing = remote_run in {"RUNNING", "UNSTART"}
if not (status_changed or run_changed or is_processing):
await self.repository.session.commit()
continue
before_tokens = int(local.get("token_count") or 0)
await self.repository.sync_running_document(
dataset_id,
str(local["document_id"]),
remote[0],
shanghai_now_naive(),
)
delta = int(remote[0].get("token_count") or 0) - before_tokens
if delta:
await self.repository.update_stats(dataset_id, 0, 0, delta)
await self.repository.session.commit()
updates += 1
except Exception:
await self.repository.session.rollback()
continue
return updates
async def _client_for_dataset(self, dataset_id: str) -> RAGFlowClient:
row = await self.repository.get_dataset(dataset_id)
if row is None or not row.get("rag_model_id"):
raise AppError(10164)
return await self.datasets._client(str(row["rag_model_id"]))
def _document_delete_error(exc: Exception, language: str | None) -> AppError:
"""Match `new RenException(e.getMessage())` in the Java delete flow."""
if isinstance(exc, AppError):
message = exc.message or message_for(exc.code, language, *exc.params)
else:
message = str(exc)
return AppError(500, message)
async def _delete_cache_ignoring_errors(key: str) -> None:
try:
await get_redis().delete(key)
except Exception:
# The Java document cleanup and remote-missing cleanup explicitly log
# and continue when Redis is unavailable.
return
@@ -0,0 +1,306 @@
from __future__ import annotations
import copy
import json
import uuid
from typing import Any
from app.core.errors import AppError
from app.core.redis import get_redis
from app.core.security import AuthUser, shanghai_now_naive
from app.repositories.model import ModelRepository, parse_json_object
from app.schemas.model import ModelConfigBody, ModelProviderBody
SENSITIVE_FIELDS = {
"api_key",
"personal_access_token",
"access_token",
"token",
"secret",
"access_key_secret",
"secret_key",
}
def _mask_middle(value: str) -> str:
if not value.strip() or len(value) == 1:
return value
if len(value) <= 8:
return value[:2] + "****" + value[-2:]
return value[:4] + "*" * (len(value) - 8) + value[-4:]
def mask_sensitive(value: Any) -> Any:
if not isinstance(value, dict):
return value
result: dict[str, Any] = {}
for key, item in value.items():
if key.lower() in SENSITIVE_FIELDS and isinstance(item, str):
result[key] = _mask_middle(item)
elif isinstance(item, dict):
result[key] = mask_sensitive(item)
else:
result[key] = copy.deepcopy(item)
return result
def _merge_config(original: dict[str, Any], updated: dict[str, Any]) -> dict[str, Any]:
result = copy.deepcopy(original)
for key, value in updated.items():
if key.lower() in SENSITIVE_FIELDS:
if isinstance(value, str) and "****" not in value:
result[key] = value
elif isinstance(value, dict):
child = result.get(key)
result[key] = _merge_config(child if isinstance(child, dict) else {}, value)
else:
result[key] = copy.deepcopy(value)
for key in list(result):
if key not in updated and key.lower() not in SENSITIVE_FIELDS:
del result[key]
return result
def _model_dto(row: dict[str, Any], *, masked: bool = True) -> dict[str, Any]:
config = parse_json_object(row.get("config_json"))
return {
"id": row.get("id"),
"modelType": row.get("model_type"),
"modelCode": row.get("model_code"),
"modelName": row.get("model_name"),
"isDefault": row.get("is_default"),
"isEnabled": row.get("is_enabled"),
"configJson": mask_sensitive(config) if masked else config,
"docLink": row.get("doc_link"),
"remark": row.get("remark"),
"sort": row.get("sort"),
}
class ModelService:
def __init__(self, repository: ModelRepository):
self.repository = repository
async def names(self, model_type: str, model_name: str | None) -> list[dict[str, Any]]:
return [
{"id": row.get("id"), "modelName": row.get("model_name")}
for row in await self.repository.list_model_names(model_type, model_name)
]
async def llm_names(self, model_name: str | None) -> list[dict[str, Any]]:
result: list[dict[str, Any]] = []
for row in await self.repository.list_llm_names(model_name):
config = parse_json_object(row.get("config_json")) or {}
result.append(
{"id": row.get("id"), "modelName": row.get("model_name"), "type": str(config.get("type", ""))}
)
return result
async def model_page(self, model_type: str, model_name: str | None, page: str, limit: str) -> dict[str, Any]:
current, size = max(int(page), 1), int(limit)
rows, total = await self.repository.list_model_configs(
model_type=model_type,
model_name=model_name,
offset=(current - 1) * size,
limit=size,
)
return {"total": total, "list": [_model_dto(row) for row in rows]}
async def get_model(self, model_id: str) -> dict[str, Any] | None:
row = await self.repository.get_model(model_id)
return _model_dto(row) if row else None
async def add(self, model_type: str, provider_code: str, body: ModelConfigBody) -> dict[str, Any]:
if not model_type.strip() or not provider_code.strip():
raise AppError(10131)
model_id = body.id or uuid.uuid4().hex
values = {
"id": model_id,
"model_type": model_type,
"model_code": body.model_code,
"model_name": body.model_name,
"is_default": 0,
"is_enabled": body.is_enabled,
"config_json": json.dumps(body.config_json, ensure_ascii=False) if body.config_json is not None else None,
"doc_link": body.doc_link,
"remark": body.remark,
"sort": body.sort,
}
async with self.repository.session.begin():
# Keep the read and write in one transaction. A query before
# ``begin()`` triggers SQLAlchemy autobegin and makes the explicit
# transaction fail with InvalidRequestError.
if await self.repository.get_provider(model_type, provider_code) is None:
raise AppError(10162)
await self.repository.insert_model(values)
return _model_dto(values)
async def edit(
self, model_type: str, provider_code: str, model_id: str, body: ModelConfigBody
) -> dict[str, Any]:
if not model_type.strip() or not provider_code.strip():
raise AppError(10131)
async with self.repository.session.begin():
if await self.repository.get_provider(model_type, provider_code) is None:
raise AppError(10162)
original = await self.repository.get_model(model_id, for_update=True)
if original is None:
raise AppError(10051)
updated_config = body.config_json
if updated_config is not None and "llm" in updated_config:
llm = await self.repository.get_model(str(updated_config["llm"]))
llm_config = parse_json_object(llm.get("config_json")) if llm else None
if llm is None or str(llm.get("model_type") or "").upper() != "LLM":
raise AppError(10092)
if llm_config and "type" in llm_config and llm_config["type"] not in {"openai", "ollama"}:
raise AppError(10049)
original_config = parse_json_object(original.get("config_json"))
merged = (
_merge_config(original_config, updated_config)
if original_config is not None and updated_config is not None
else original_config
)
values = {
"id": model_id,
"model_type": model_type,
"model_code": original.get("model_code"),
"model_name": body.model_name,
"is_default": original.get("is_default"),
"is_enabled": body.is_enabled,
"config_json": json.dumps(merged, ensure_ascii=False) if merged is not None else None,
"doc_link": original.get("doc_link"),
"remark": body.remark,
"sort": body.sort,
}
await self.repository.update_model(values)
await self._clear_cache(model_id)
return _model_dto(values)
async def delete(self, model_id: str) -> None:
if not model_id.strip():
raise AppError(10006)
async with self.repository.session.begin():
model = await self.repository.get_model(model_id, for_update=True)
if model and int(model.get("is_default") or 0) == 1:
raise AppError(10064)
agents = await self.repository.model_agent_references(model_id)
if agents:
raise AppError(10093, params=("".join(agents),))
if model and str(model.get("model_type") or "").upper() == "LLM":
if await self.repository.intent_reference_count(model_id):
raise AppError(10094)
await self.repository.delete_model(model_id)
await self._clear_cache(model_id)
async def enable(self, model_id: str, status: int) -> str | None:
async with self.repository.session.begin():
model = await self.repository.get_model(model_id, for_update=True)
if model is None:
return "模型配置不存在"
if status == 0 and int(model.get("is_default") or 0) > 0:
return "默认模型配置不允许关闭"
await self.repository.set_model_enabled(model_id, status)
await self._clear_cache(model_id)
return None
async def set_default(self, model_id: str) -> str | None:
async with self.repository.session.begin():
model = await self.repository.get_model(model_id, for_update=True)
if model is None:
return "模型配置不存在"
model_type = str(model.get("model_type") or "")
await self.repository.set_models_default(model_type, 0)
await self.repository.execute(
"UPDATE ai_model_config SET is_enabled=1, is_default=1 WHERE id=:id", {"id": model_id}
)
await self.repository.update_default_template_models(model_type, model_id)
await self._clear_type_cache(model_type)
return None
async def _clear_cache(self, model_id: str) -> None:
redis = get_redis()
await redis.delete(f"model:data:{model_id}", f"model:name:{model_id}")
async def _clear_type_cache(self, model_type: str) -> None:
rows = await self.repository.fetch_all(
"SELECT id FROM ai_model_config WHERE model_type=:type", {"type": model_type}
)
if rows:
redis = get_redis()
keys = [key for row in rows for key in (f"model:data:{row['id']}", f"model:name:{row['id']}")]
await redis.delete(*keys)
class ModelProviderService:
def __init__(self, repository: ModelRepository):
self.repository = repository
async def page(
self, model_type: str | None, name: str | None, page: str, limit: str
) -> dict[str, Any]:
current, size = max(int(page), 1), int(limit)
rows, total = await self.repository.list_providers(
model_type=model_type, name=name, offset=(current - 1) * size, limit=size
)
return {"total": total, "list": rows}
@staticmethod
def _validate(body: ModelProviderBody, *, update: bool) -> None:
if update and (body.id is None or not body.id.strip()):
raise AppError(10034, "id不能为空")
for field, message in (
(body.provider_code, "providerCode不能为空"),
(body.model_type, "modelType不能为空"),
(body.name, "name不能为空"),
(body.fields, "fields(JSON格式)不能为空"),
):
if field is None or not field.strip():
raise AppError(10034, message)
if body.sort is None:
raise AppError(10034, "sort不能为空")
async def add(self, body: ModelProviderBody, user: AuthUser) -> dict[str, Any]:
self._validate(body, update=False)
now = shanghai_now_naive()
values = {
"id": body.id or uuid.uuid4().hex,
"model_type": body.model_type,
"provider_code": body.provider_code,
"name": body.name,
"fields": body.fields,
"sort": body.sort,
"creator": user.id,
"updater": user.id,
"now": now,
}
async with self.repository.session.begin():
await self.repository.insert_provider(values)
return {
# The Java service returns the request DTO, not the entity on which
# MyBatis-Plus generated the UUID. Therefore an omitted id remains
# null in the response even though the stored row has an id.
"id": body.id, "modelType": body.model_type, "providerCode": body.provider_code,
"name": body.name, "fields": body.fields, "sort": body.sort, "creator": user.id,
"updater": user.id, "createDate": now, "updateDate": now,
}
async def edit(self, body: ModelProviderBody, user: AuthUser) -> dict[str, Any]:
self._validate(body, update=True)
now = shanghai_now_naive()
values = {
"id": body.id, "model_type": body.model_type, "provider_code": body.provider_code,
"name": body.name, "fields": body.fields, "sort": body.sort, "updater": user.id, "now": now,
}
async with self.repository.session.begin():
if await self.repository.update_provider(values) == 0:
raise AppError(10066)
return {
"id": body.id, "modelType": body.model_type, "providerCode": body.provider_code,
"name": body.name, "fields": body.fields, "sort": body.sort, "updater": user.id,
"updateDate": now, "creator": None, "createDate": None,
}
async def delete(self, ids: list[str]) -> None:
async with self.repository.session.begin():
if await self.repository.delete_providers(ids) == 0:
raise AppError(10043)
@@ -0,0 +1,486 @@
from __future__ import annotations
import base64
import hashlib
import hmac
import io
import json
import logging
import re
import secrets
import string
import time
import urllib.parse
import uuid
from datetime import datetime, timedelta
from typing import Any, Protocol, cast
import httpx
from fastapi import Request
from PIL import Image, ImageDraw, ImageFont
from redis.asyncio import Redis
from app.core.config import get_settings
from app.core.crypto import bcrypt_hash, bcrypt_matches, generate_database_token, sm2_decrypt_c1c3c2
from app.core.errors import AppError, ErrorCode
from app.core.ids import snowflake
from app.core.redis import JavaRedisCodec, get_redis
from app.core.security import AuthUser, shanghai_now_naive
from app.repositories.security import SecurityRepository
from app.schemas.security import (
LoginRequest,
PasswordChangeRequest,
RetrievePasswordRequest,
SmsVerificationRequest,
)
from app.services.java_validation import validation_message
logger = logging.getLogger(__name__)
TOKEN_EXPIRE_SECONDS = 12 * 60 * 60
CAPTCHA_TTL_SECONDS = 5 * 60
CAPTCHA_LENGTH = 5
PHONE_PATTERN = re.compile(r"^\+[1-9]\d{0,3}[1-9]\d{4,14}$")
STRONG_PASSWORD = re.compile(r"^(?=.*[0-9])(?=.*[a-z])(?=.*[A-Z]).+$")
class SmsSender(Protocol):
async def send_verification_code(self, phone: str | None, code: str) -> None: ...
class AliyunSmsSender:
"""Minimal implementation of the Aliyun Dysmsapi RPC request used by the Java SDK."""
def __init__(
self,
repository: SecurityRepository,
*,
redis: Redis | None = None,
client: httpx.AsyncClient | None = None,
endpoint: str = "https://dysmsapi.aliyuncs.com/",
):
self.repository = repository
self.redis = redis or get_redis()
self.client = client
self.endpoint = endpoint
async def send_verification_code(self, phone: str | None, code: str) -> None:
access_key_id = await self._param("aliyun.sms.access_key_id") or ""
access_key_secret = await self._param("aliyun.sms.access_key_secret") or ""
sign_name = await self._param("aliyun.sms.sign_name") or ""
template_code = await self._param("aliyun.sms.sms_code_template_code") or ""
# The Tea SDK constructs its client before the refundable send block;
# blank credentials therefore map to SMS_CONNECTION_FAILED (10056).
if not access_key_id.strip() or not access_key_secret.strip():
raise AppError(10056)
timestamp = time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime())
params: dict[str, str] = {
"AccessKeyId": access_key_id,
"Action": "SendSms",
"Format": "JSON",
"RegionId": "cn-hangzhou",
"SignatureMethod": "HMAC-SHA1",
"SignatureNonce": str(uuid.uuid4()),
"SignatureVersion": "1.0",
"SignName": sign_name,
"TemplateCode": template_code,
"TemplateParam": json.dumps({"code": code}, ensure_ascii=False, separators=(",", ":")),
"Timestamp": timestamp,
"Version": "2017-05-25",
}
if phone is not None:
params["PhoneNumbers"] = phone
params["Signature"] = self._signature(params, access_key_secret)
if self.client is not None:
response = await self.client.post(self.endpoint, data=params)
response.raise_for_status()
return
timeout = get_settings().external_request_timeout_seconds
async with httpx.AsyncClient(timeout=timeout) as client:
response = await client.post(self.endpoint, data=params)
response.raise_for_status()
async def _param(self, code: str) -> str | None:
cached = JavaRedisCodec.decode(await cast(Any, self.redis.hget)("sys:params", code))
if cached is not None:
return str(cached)
value = await self.repository.get_param_value(code)
if value is not None:
await cast(Any, self.redis.hset)("sys:params", code, JavaRedisCodec.encode(value))
await cast(Any, self.redis.expire)("sys:params", 24 * 60 * 60)
return value
@classmethod
def _signature(cls, params: dict[str, str], secret: str) -> str:
canonical = "&".join(
f"{cls._percent_encode(key)}={cls._percent_encode(value)}" for key, value in sorted(params.items())
)
string_to_sign = f"POST&%2F&{cls._percent_encode(canonical)}"
digest = hmac.new(
f"{secret}&".encode(),
string_to_sign.encode(),
digestmod=hashlib.sha1, # noqa: S324 - mandated by Aliyun RPC SignatureMethod
).digest()
return base64.b64encode(digest).decode("ascii")
@staticmethod
def _percent_encode(value: str) -> str:
return urllib.parse.quote(str(value), safe="~")
class CaptchaService:
def __init__(self, redis: Redis | None = None):
self.redis = redis or get_redis()
async def create(self, identifier: str) -> bytes:
code = "".join(secrets.choice(string.ascii_letters + string.digits) for _ in range(CAPTCHA_LENGTH))
await self._set_cache(identifier, code)
return self._render_gif(code)
async def validate(self, identifier: str | None, code: str | None, *, delete: bool) -> bool:
if not code or not code.strip():
return False
key = self._captcha_key(identifier)
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get(key)))
if cached is not None and delete:
await cast(Any, self.redis.delete(key))
return cached is not None and code.casefold() == str(cached).casefold()
async def set_sms_code(self, phone: str | None, code: str) -> None:
await self._set_cache(f"sms:Validate:Code:{phone}", code)
async def validate_sms_code(self, phone: str | None, code: str | None, *, delete: bool = False) -> bool:
return await self.validate(f"sms:Validate:Code:{phone}", code, delete=delete)
async def _set_cache(self, identifier: str, value: str) -> None:
await cast(Any, self.redis.set)(
self._captcha_key(identifier),
JavaRedisCodec.encode(value),
ex=CAPTCHA_TTL_SECONDS,
)
@staticmethod
def _captcha_key(identifier: str | None) -> str:
return f"sys:captcha:{'null' if identifier is None else identifier}"
@staticmethod
def _render_gif(code: str) -> bytes:
image = Image.new("RGB", (150, 40), (248, 248, 248))
draw = ImageDraw.Draw(image)
for _ in range(8):
color = tuple(secrets.randbelow(150) for _ in range(3))
draw.line(
(
secrets.randbelow(150),
secrets.randbelow(40),
secrets.randbelow(150),
secrets.randbelow(40),
),
fill=color,
width=1,
)
font = ImageFont.load_default(size=24)
for index, character in enumerate(code):
color = tuple(secrets.randbelow(120) for _ in range(3))
draw.text((10 + index * 27, 6 + secrets.randbelow(5)), character, font=font, fill=color)
output = io.BytesIO()
image.save(output, format="GIF")
return output.getvalue()
class SecurityService:
def __init__(
self,
repository: SecurityRepository,
*,
redis: Redis | None = None,
captcha: CaptchaService | None = None,
sms_sender: SmsSender | None = None,
):
self.repository = repository
self.redis = redis or get_redis()
self.captcha = captcha or CaptchaService(self.redis)
self.sms_sender = sms_sender or AliyunSmsSender(repository, redis=self.redis)
async def login(self, dto: LoginRequest, request: Request) -> dict[str, Any]:
password = await self._decrypt_and_validate_captcha(dto.password, dto.captcha_id)
user = await self.repository.get_user_by_username(dto.username)
if user is None or not bcrypt_matches(password, cast(str | None, user.get("password"))):
raise AppError(ErrorCode.ACCOUNT_PASSWORD_ERROR)
token = await self._create_token(int(user["id"]))
await self.repository.session.commit()
return {
"token": token,
"expire": TOKEN_EXPIRE_SECONDS,
"clientHash": self._client_hash(request),
}
async def register(self, dto: LoginRequest) -> None:
if not await self.allow_user_register():
raise AppError(10072)
password = await self._decrypt_and_validate_captcha(dto.password, dto.captcha_id)
if await self._mobile_registration_enabled():
if dto.username is None or not PHONE_PATTERN.fullmatch(dto.username):
raise AppError(10069)
if not await self.captcha.validate_sms_code(dto.username, dto.mobile_captcha, delete=False):
raise AppError(10075)
if await self.repository.get_user_by_username(dto.username) is not None:
raise AppError(10070)
if not STRONG_PASSWORD.fullmatch(password):
raise AppError(ErrorCode.PASSWORD_WEAK_ERROR)
now = shanghai_now_naive()
user_count = await self.repository.count_users()
await self.repository.insert_user(
user_id=snowflake.next_id(),
username=dto.username,
password=bcrypt_hash(password),
super_admin=1 if user_count == 0 else 0,
now=now,
)
await self.repository.session.commit()
async def change_password(
self,
user: AuthUser,
dto: PasswordChangeRequest,
accept_language: str | None = None,
) -> None:
self._require_not_blank(dto.password, "sysuser.password.require", accept_language)
self._require_not_blank(dto.new_password, "sysuser.password.require", accept_language)
assert dto.password is not None
assert dto.new_password is not None
row = await self.repository.get_user_by_id(user.id)
if row is None:
raise AppError(ErrorCode.TOKEN_INVALID)
if not bcrypt_matches(dto.password, cast(str | None, row.get("password"))):
raise AppError(10048)
if not STRONG_PASSWORD.fullmatch(dto.new_password):
raise AppError(ErrorCode.PASSWORD_WEAK_ERROR)
now = shanghai_now_naive()
await self.repository.update_password(
user.id,
bcrypt_hash(dto.new_password),
now,
preserve_audit_fields=True,
)
# SysUserService.changePassword commits before the non-transactional token service logs out.
await self.repository.session.commit()
await self.repository.expire_user_token(user.id, now - timedelta(minutes=1))
await self.repository.session.commit()
async def retrieve_password(
self,
dto: RetrievePasswordRequest,
accept_language: str | None = None,
) -> None:
if not await self._mobile_registration_enabled():
raise AppError(10073)
self._require_not_blank(dto.phone, "sysuser.password.require", accept_language)
self._require_not_blank(dto.code, "sysuser.password.require", accept_language)
self._require_not_blank(dto.password, "sysuser.password.require", accept_language)
self._require_not_blank(dto.captcha_id, "sysuser.uuid.require", accept_language)
assert dto.phone is not None
assert dto.code is not None
assert dto.password is not None
assert dto.captcha_id is not None
if not PHONE_PATTERN.fullmatch(dto.phone):
raise AppError(10074)
user = await self.repository.get_user_by_username(dto.phone)
if user is None:
raise AppError(10071)
if not await self.captcha.validate_sms_code(dto.phone, dto.code, delete=False):
raise AppError(10075)
password = await self._decrypt_and_validate_captcha(dto.password, dto.captcha_id)
if not STRONG_PASSWORD.fullmatch(password):
raise AppError(ErrorCode.PASSWORD_WEAK_ERROR)
await self.repository.update_password(int(user["id"]), bcrypt_hash(password), shanghai_now_naive())
await self.repository.session.commit()
async def send_sms_verification(self, dto: SmsVerificationRequest) -> None:
if not await self.captcha.validate(dto.captcha_id, dto.captcha, delete=False):
raise AppError(10067)
if not await self._mobile_registration_enabled():
raise AppError(10068)
phone_key = "null" if dto.phone is None else dto.phone
last_send_key = f"sms:Validate:Code:{phone_key}:last_send_time"
current_ms = int(time.time() * 1000)
created = await cast(Any, self.redis.set)(last_send_key, str(current_ms), ex=60, nx=True)
if not created:
raw_last = await cast(Any, self.redis.get)(last_send_key)
if raw_last is not None:
last_ms = int(raw_last.decode() if isinstance(raw_last, bytes) else raw_last)
difference = current_ms - last_ms
if difference < 60_000:
raise AppError(10060, params=(str(max(0, (60_000 - difference) // 1000)),))
today_key = f"sms:Validate:Code:{phone_key}:today_count"
raw_count = await cast(Any, self.redis.get)(today_key)
decoded_count = JavaRedisCodec.decode(raw_count)
today_count = int(decoded_count or 0)
raw_maximum = await self._get_param("server.sms_max_send_count", from_cache=True)
maximum = int(raw_maximum) if raw_maximum is not None and raw_maximum != "" else 5
if today_count >= maximum:
raise AppError(10047)
code = "".join(secrets.choice(string.digits) for _ in range(6))
await self.captcha.set_sms_code(dto.phone, code)
new_count = await cast(Any, self.redis.incr)(today_key)
if int(new_count) == 1:
await cast(Any, self.redis.expire)(today_key, 24 * 60 * 60)
try:
await self.sms_sender.send_verification_code(dto.phone, code)
except AppError:
# Java raises connection-construction failures before entering its refundable send attempt.
raise
except Exception as exc:
logger.warning("Aliyun SMS request failed", exc_info=exc)
await cast(Any, self.redis.delete)(today_key)
raise AppError(10055) from exc
async def public_config(self) -> dict[str, Any]:
public_key = await self._get_param("server.public_key", from_cache=True)
if public_key is None or not public_key.strip():
raise AppError(10129)
menu_config = await self._get_param("system-web.menu", from_cache=True)
result: dict[str, Any] = {
"enableMobileRegister": await self._mobile_registration_enabled(),
"version": "0.9.5",
"year": f"©{shanghai_now_naive().year}",
"allowUserRegister": await self.allow_user_register(),
"mobileAreaList": await self._dict_data_by_type("MOBILE_AREA"),
"beianIcpNum": await self._get_param("server.beian_icp_num", from_cache=True),
"beianGaNum": await self._get_param("server.beian_ga_num", from_cache=True),
"name": await self._get_param("server.name", from_cache=True),
"sm2PublicKey": public_key,
}
if menu_config is not None and menu_config.strip():
result["systemWebMenu"] = json.loads(menu_config)
return result
async def allow_user_register(self) -> bool:
value = await self._get_param("server.allow_user_register", from_cache=True)
if value == "true":
return True
return await self.repository.count_users() == 0
async def _create_token(self, user_id: int) -> str:
now = shanghai_now_naive()
expire_date = now + timedelta(seconds=TOKEN_EXPIRE_SECONDS)
current = await self.repository.get_token_by_user_id(user_id, for_update=True)
if current is None:
token = generate_database_token()
await self.repository.insert_token(
token_id=snowflake.next_id(),
user_id=user_id,
token=token,
now=now,
expire_date=expire_date,
)
return token
stored_expiry = self._datetime(current.get("expire_date"))
token = str(current["token"])
if stored_expiry is None or stored_expiry < now:
token = generate_database_token()
await self.repository.update_token(
token_id=int(current["id"]),
token=token,
now=now,
expire_date=expire_date,
)
return token
async def _decrypt_and_validate_captcha(
self,
encrypted_password: str | None,
captcha_id: str | None,
) -> str:
private_key = await self._get_param("server.private_key", from_cache=True)
if private_key is None or not private_key.strip():
raise AppError(10129)
try:
if encrypted_password is None:
raise ValueError("encrypted password is null")
content = sm2_decrypt_c1c3c2(private_key, encrypted_password)
except Exception as exc:
raise AppError(10130) from exc
if len(content) > CAPTCHA_LENGTH:
embedded_captcha = content[:CAPTCHA_LENGTH]
if not await self.captcha.validate(captcha_id, embedded_captcha, delete=True):
raise AppError(10067)
return content[CAPTCHA_LENGTH:]
if content:
raise AppError(10067)
raise AppError(10130)
async def _mobile_registration_enabled(self) -> bool:
value = await self._get_param("server.enable_mobile_register", from_cache=True)
if value is None or not value.strip():
return False
try:
parsed = json.loads(value.lower())
except json.JSONDecodeError as exc:
raise AppError(ErrorCode.PARAMS_GET_ERROR) from exc
return bool(parsed)
async def _get_param(self, code: str, *, from_cache: bool) -> str | None:
if from_cache:
cached = JavaRedisCodec.decode(await cast(Any, self.redis.hget)("sys:params", code))
if cached is not None:
return str(cached)
value = await self.repository.get_param_value(code)
if from_cache and value is not None:
await cast(Any, self.redis.hset)("sys:params", code, JavaRedisCodec.encode(value))
await cast(Any, self.redis.expire)("sys:params", 24 * 60 * 60)
return value
async def _dict_data_by_type(self, dict_type: str) -> list[dict[str, Any]]:
key = f"sys:dict:data:{dict_type}"
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)(key))
if isinstance(cached, list):
return cast(list[dict[str, Any]], cached)
values = await self.repository.get_mobile_area_items()
await cast(Any, self.redis.set)(
key,
JavaRedisCodec.encode(
values,
item_java_type="xiaozhi.modules.sys.vo.SysDictDataItem",
),
ex=24 * 60 * 60,
)
return values
@staticmethod
def _client_hash(request: Request) -> str:
user_agent = request.headers.get("User-Agent", "").lower()
forwarded_headers = (
"x-forwarded-for",
"Proxy-Client-IP",
"WL-Proxy-Client-IP",
"HTTP_CLIENT_IP",
"HTTP_X_FORWARDED_FOR",
)
ip_address = next(
(
value
for header in forwarded_headers
if (value := request.headers.get(header)) and value.casefold() != "unknown"
),
request.client.host if request.client else "",
)
date = shanghai_now_naive().strftime("%Y-%m-%d")
return hashlib.md5( # noqa: S324 - Java clientHash compatibility requires MD5
f"{ip_address}{date}{user_agent}".encode(), usedforsecurity=False
).hexdigest()
@staticmethod
def _datetime(value: Any) -> datetime | None:
if value is None or isinstance(value, datetime):
return value
if isinstance(value, str):
return datetime.fromisoformat(value)
raise TypeError(f"Unsupported database datetime value: {type(value).__name__}")
@staticmethod
def _require_not_blank(value: str | None, key: str, accept_language: str | None) -> None:
if value is None or not value.strip():
raise AppError(500, validation_message(key, accept_language))
@@ -0,0 +1,715 @@
from __future__ import annotations
import asyncio
import base64
import hashlib
import hmac
import json
import logging
import re
import secrets
import string
import time
import uuid
from datetime import datetime
from typing import Any, cast
from zoneinfo import ZoneInfo
import httpx
from redis.asyncio import Redis
from websockets.asyncio.client import connect
from app.core.config import get_settings
from app.core.crypto import bcrypt_hash
from app.core.errors import AppError, ErrorCode
from app.core.ids import snowflake
from app.core.redis import JavaRedisCodec, get_redis
from app.core.security import AuthUser, shanghai_now_naive
from app.repositories.sys import SysRepository
from app.schemas.sys import DictDataPayload, DictTypePayload, EmitServerActionRequest, SysParamPayload
from app.services.java_validation import validation_message
logger = logging.getLogger(__name__)
WS_PATTERN = re.compile(r"^wss?://[\w.-]+(?:\.[\w.-]+)*(?::\d+)?(?:/[\w.-]*)*$")
class AdminService:
def __init__(self, repository: SysRepository):
self.repository = repository
async def page_users(self, *, mobile: str | None, page: int, limit: int) -> dict[str, Any]:
rows, total = await self.repository.page_users(
mobile=mobile,
page=max(1, page),
limit=max(0, limit),
)
values = [
{
"deviceCount": str(row.get("device_count") or 0),
"mobile": row.get("username"),
"status": row.get("status"),
"userid": str(row["id"]),
"createDate": row.get("create_date"),
}
for row in rows
]
return {"list": values, "total": total}
async def reset_password(self, user_id: int, user: AuthUser) -> str:
password = self._generate_password()
await self.repository.reset_user_password(user_id, bcrypt_hash(password), user.id, shanghai_now_naive())
await self.repository.session.commit()
return password
async def delete_user(self, user_id: int) -> None:
try:
await self.repository.delete_user_cascade(user_id)
await self.repository.session.commit()
except Exception:
await self.repository.session.rollback()
raise
async def change_status(self, status: int, user_ids: list[str], user: AuthUser) -> None:
# SysUserServiceImpl.changeStatus has an outer Spring transaction: a later
# parse/update failure rolls back every earlier item in the same request.
try:
for value in user_ids:
await self.repository.change_user_status(status, [int(value)], user.id, shanghai_now_naive())
await self.repository.session.commit()
except Exception:
await self.repository.session.rollback()
raise
async def page_devices(self, *, keywords: str | None, page: int, limit: int) -> dict[str, Any]:
rows, total = await self.repository.page_devices(
keywords=keywords,
page=max(1, page),
limit=max(0, limit),
)
result = []
for row in rows:
result.append(
{
"appVersion": row.get("app_version"),
"bindUserName": row.get("bind_user_name"),
"deviceType": row.get("board"),
"board": row.get("board"),
"id": row.get("id"),
"macAddress": row.get("mac_address"),
"alias": row.get("alias"),
"otaUpgrade": None,
"recentChatTime": self._short_time(cast(datetime | str | None, row.get("update_date"))),
"lastConnectedAtTimestamp": self._timestamp_ms(
cast(datetime | str | None, row.get("last_connected_at"))
),
"createDateTimestamp": self._timestamp_ms(
cast(datetime | str | None, row.get("create_date"))
),
"createDate": self._utc_datetime_string(
cast(datetime | str | None, row.get("create_date"))
),
}
)
return {"list": result, "total": total}
@staticmethod
def _generate_password() -> str:
characters = string.ascii_letters + string.digits + "!@#$%^&*()"
values = [
secrets.choice(string.digits),
secrets.choice(string.ascii_lowercase),
secrets.choice(string.ascii_uppercase),
secrets.choice("!@#$%^&*()"),
]
values.extend(secrets.choice(characters) for _ in range(8))
secrets.SystemRandom().shuffle(values)
return "".join(values)
@staticmethod
def _timestamp_ms(value: datetime | str | None) -> int | None:
value = AdminService._database_datetime(value)
if value is None:
return None
timezone = ZoneInfo(get_settings().timezone)
localized = value if value.tzinfo else value.replace(tzinfo=timezone)
return int(localized.timestamp() * 1000)
@staticmethod
def _short_time(value: datetime | str | None) -> str | None:
value = AdminService._database_datetime(value)
if value is None:
return None
now = shanghai_now_naive()
if value.tzinfo:
value = value.astimezone(ZoneInfo(get_settings().timezone)).replace(tzinfo=None)
seconds = int((now - value).total_seconds())
if seconds <= 10:
return "刚刚"
if seconds < 60:
return f"{seconds}秒前"
if seconds < 3600:
return f"{seconds // 60}分钟前"
if seconds < 86400:
return f"{seconds // 3600}小时前"
if seconds < 604800:
return f"{seconds // 86400}天前"
return value.strftime("%Y-%m-%d %H:%M:%S")
@staticmethod
def _utc_datetime_string(value: datetime | str | None) -> str | None:
value = AdminService._database_datetime(value)
if value is None:
return None
timezone = ZoneInfo(get_settings().timezone)
localized = value if value.tzinfo else value.replace(tzinfo=timezone)
return localized.astimezone(ZoneInfo("UTC")).strftime("%Y-%m-%d %H:%M:%S")
@staticmethod
def _database_datetime(value: datetime | str | None) -> datetime | None:
if isinstance(value, str):
return datetime.fromisoformat(value)
return value
class ParamExternalValidator:
def __init__(self, client: httpx.AsyncClient | None = None):
self.client = client
async def validate(self, code: str, value: str) -> None:
if code == "server.websocket":
await self._websockets(value)
elif code == "server.ota":
await self._http_endpoint(value, kind="ota")
elif code == "server.mcp_endpoint":
await self._http_endpoint(value, kind="mcp")
elif code == "server.voice_print":
await self._http_endpoint(value, kind="voiceprint")
elif code == "server.mqtt_signature_key":
self._mqtt_secret(value)
async def _websockets(self, value: str) -> None:
urls = value.split(";")
while urls and urls[-1] == "":
urls.pop()
if not urls:
raise AppError(10098)
for raw_url in urls:
if not raw_url.strip():
continue
if "localhost" in raw_url or "127.0.0.1" in raw_url:
raise AppError(10099)
if not WS_PATTERN.fullmatch(raw_url.strip()):
raise AppError(10100)
try:
async with connect(raw_url, open_timeout=5):
pass
except Exception as exc:
raise AppError(10101) from exc
async def _http_endpoint(self, value: str, *, kind: str) -> None:
if not value.strip() or value == "null":
return
if "localhost" in value or "127.0.0.1" in value:
raise AppError({"ota": 10103, "mcp": 10110, "voiceprint": 10116}[kind])
if kind == "ota":
if not value.lower().startswith("http"):
raise AppError(10104)
if not value.endswith("/ota/"):
raise AppError(10105)
elif kind == "mcp":
if "key" not in value.lower():
raise AppError(10111)
else:
if "key" not in value.lower():
raise AppError(10117)
if not value.lower().startswith("http"):
raise AppError(10118)
final_code = {"ota": 10108, "mcp": 10114, "voiceprint": 10121}[kind]
marker = {"ota": "OTA", "mcp": "success", "voiceprint": "healthy"}[kind]
try:
if self.client is not None:
response = await self.client.get(value)
else:
async with httpx.AsyncClient(timeout=get_settings().external_request_timeout_seconds) as client:
response = await client.get(value)
if response.status_code != 200 or marker not in response.text:
raise ValueError("external endpoint response did not match Java validation")
except Exception as exc:
raise AppError(final_code) from exc
@staticmethod
def _mqtt_secret(secret: str) -> None:
if not secret.strip() or secret == "null": # noqa: S105 - sentinel value from the Java parameter table
raise AppError(10122)
if len(secret) < 8:
raise AppError(10123)
if not re.search(r"[a-z]", secret) or not re.search(r"[A-Z]", secret):
raise AppError(10124)
lowered = secret.lower()
if any(weak in lowered for weak in ("test", "1234", "admin", "password", "qwerty", "xiaozhi")):
raise AppError(10125)
class SysParamService:
def __init__(
self,
repository: SysRepository,
*,
redis: Redis | None = None,
validator: ParamExternalValidator | None = None,
):
self.repository = repository
self.redis = redis or get_redis()
self.validator = validator or ParamExternalValidator()
async def page(
self,
*,
param_code: str | None,
page: int,
limit: int,
order_field: str | None,
order: str | None,
) -> dict[str, Any]:
rows, total = await self.repository.page_params(
param_code=param_code,
page=max(1, page),
limit=max(0, limit),
order_field=order_field,
order=order,
)
return {"list": [self._param_dto(row) for row in rows], "total": total}
async def get(self, param_id: int) -> dict[str, Any] | None:
row = await self.repository.get_param(param_id)
return None if row is None else self._param_dto(row)
async def save(
self,
dto: SysParamPayload,
user: AuthUser,
accept_language: str | None = None,
) -> None:
self._validate_group(dto, update=False, accept_language=accept_language)
self._validate_value(dto)
assert dto.param_code is not None
assert dto.param_value is not None
assert dto.value_type is not None
await self.repository.insert_param(
param_id=snowflake.next_id(),
param_code=dto.param_code,
param_value=dto.param_value,
value_type=dto.value_type,
remark=dto.remark,
user_id=user.id,
now=shanghai_now_naive(),
)
await self._cache_set(dto.param_code, dto.param_value)
await self.repository.session.commit()
async def update(
self,
dto: SysParamPayload,
user: AuthUser,
accept_language: str | None = None,
) -> None:
self._validate_group(dto, update=True, accept_language=accept_language)
assert dto.id is not None
assert dto.param_code is not None
assert dto.param_value is not None
assert dto.value_type is not None
# These checks live in the Java controller and therefore run before
# SysParamsService.update validates the declared value type.
await self.validator.validate(dto.param_code, dto.param_value)
if dto.param_code == "system-web.menu":
await self._update_system_web_menu(dto.param_value, user)
else:
self._validate_value(dto)
await self.repository.update_param(
param_id=dto.id,
param_code=dto.param_code,
param_value=dto.param_value,
value_type=dto.value_type,
remark=dto.remark,
user_id=user.id,
now=shanghai_now_naive(),
)
await self._cache_set(dto.param_code, dto.param_value)
await self.repository.session.commit()
async def delete(self, ids: list[str]) -> None:
if not ids:
raise AppError(10001, "id")
parsed_ids = [int(value) for value in ids]
codes = await self.repository.param_codes_for_ids(parsed_ids)
if codes:
await cast(Any, self.redis.hdel)("sys:params", *codes)
await self.repository.delete_params(parsed_ids)
await self.repository.session.commit()
async def get_value(self, code: str, *, from_cache: bool = True) -> str | None:
if from_cache:
cached = JavaRedisCodec.decode(await cast(Any, self.redis.hget)("sys:params", code))
if cached is not None:
return str(cached)
value = await self.repository.get_param_value(code)
if value is not None and from_cache:
await self._cache_set(code, value)
return value
async def config_rows(self) -> list[dict[str, Any]]:
return await self.repository.list_config_params()
async def _update_system_web_menu(self, config_json: str, user: AuthUser) -> None:
current_config = await self.repository.get_param_value("system-web.menu")
try:
current = json.loads(current_config) if current_config and current_config.strip() else None
updated = json.loads(config_json) if config_json.strip() else None
except json.JSONDecodeError as exc:
raise AppError(ErrorCode.PARAM_JSON_INVALID) from exc
if isinstance(current, dict) and isinstance(updated, dict):
current_features = current.get("features")
updated_features = updated.get("features")
# Java only evaluates addressBook when both feature maps are present.
if isinstance(current_features, dict) and isinstance(updated_features, dict):
current_address = current_features.get("addressBook")
updated_address = updated_features.get("addressBook")
current_enabled = self._java_enabled(current_address)
updated_enabled = self._java_enabled(updated_address)
if current_enabled and not updated_enabled:
await self.repository.delete_plugin_mapping_by_plugin_id("SYSTEM_PLUGIN_CALL_DEVICE")
await self.repository.update_param_value_by_code(
"system-web.menu", config_json, user.id, shanghai_now_naive()
)
await self._cache_set("system-web.menu", config_json)
async def _cache_set(self, code: str, value: str) -> None:
await cast(Any, self.redis.hset)("sys:params", code, JavaRedisCodec.encode(value))
await cast(Any, self.redis.expire)("sys:params", 24 * 60 * 60)
@staticmethod
def _java_enabled(address_book: Any) -> bool:
if not isinstance(address_book, dict):
return False
value = address_book.get("enabled")
if value is None:
return False
if not isinstance(value, bool):
# The Java implementation casts the JSON value to Boolean.
raise TypeError("addressBook.enabled must be a boolean")
return value
@staticmethod
def _validate_value(dto: SysParamPayload) -> None:
assert dto.param_value is not None
assert dto.value_type is not None
if not dto.param_value.strip():
raise AppError(ErrorCode.PARAM_VALUE_NULL)
if not dto.value_type.strip():
raise AppError(ErrorCode.PARAM_TYPE_NULL)
value_type = dto.value_type.lower()
if value_type in {"string", "array"}:
return
if value_type == "number":
try:
float(dto.param_value)
except ValueError as exc:
raise AppError(ErrorCode.PARAM_NUMBER_INVALID) from exc
return
if value_type == "boolean":
if dto.param_value.lower() not in {"true", "false"}:
raise AppError(ErrorCode.PARAM_BOOLEAN_INVALID)
return
if value_type == "json":
stripped = dto.param_value.strip()
if not stripped.startswith("{") or not stripped.endswith("}"):
raise AppError(ErrorCode.PARAM_JSON_INVALID)
try:
json.loads(dto.param_value)
except json.JSONDecodeError as exc:
raise AppError(ErrorCode.PARAM_JSON_INVALID) from exc
return
raise AppError(ErrorCode.PARAM_TYPE_INVALID)
@staticmethod
def _validate_group(
dto: SysParamPayload,
*,
update: bool,
accept_language: str | None,
) -> None:
def fail(key: str) -> None:
raise AppError(500, validation_message(key, accept_language))
if update and dto.id is None:
fail("id.require")
if not update and dto.id is not None:
fail("id.null")
if dto.param_code is None or not dto.param_code.strip():
fail("sysparams.paramcode.require")
if dto.param_value is None or not dto.param_value.strip():
fail("sysparams.paramvalue.require")
if dto.value_type is None or not dto.value_type.strip():
fail("sysparams.valuetype.require")
if dto.value_type not in {"string", "number", "boolean", "array", "json"}:
fail("sysparams.valuetype.pattern")
@staticmethod
def _param_dto(row: dict[str, Any]) -> dict[str, Any]:
return {
"id": row.get("id"),
"paramCode": row.get("param_code"),
"paramValue": row.get("param_value"),
"valueType": row.get("value_type"),
"remark": row.get("remark"),
"createDate": row.get("create_date"),
"updateDate": row.get("update_date"),
}
class DictService:
def __init__(self, repository: SysRepository, *, redis: Redis | None = None):
self.repository = repository
self.redis = redis or get_redis()
async def page_types(
self,
*,
dict_type: str | None,
dict_name: str | None,
page: int,
limit: int,
) -> dict[str, Any]:
rows, total = await self.repository.page_dict_types(
dict_type=dict_type,
dict_name=dict_name,
page=max(1, page),
limit=max(0, limit),
)
return {"list": [self._type_vo(row, include_names=True) for row in rows], "total": total}
async def get_type(self, type_id: int) -> dict[str, Any]:
row = await self.repository.get_dict_type(type_id)
if row is None:
raise AppError(10076)
return self._type_vo(row, include_names=False)
async def save_type(self, dto: DictTypePayload, user: AuthUser) -> None:
if await self.repository.dict_type_exists(dto.dict_type):
raise AppError(10077)
await self.repository.insert_dict_type(
type_id=dto.id if dto.id is not None else snowflake.next_id(),
dict_type=dto.dict_type,
dict_name=dto.dict_name,
remark=dto.remark,
sort=dto.sort,
user_id=user.id,
now=shanghai_now_naive(),
)
await self.repository.session.commit()
async def update_type(self, dto: DictTypePayload, user: AuthUser) -> None:
if await self.repository.dict_type_exists(dto.dict_type, exclude_id=dto.id):
raise AppError(10077)
await self.repository.update_dict_type(
type_id=dto.id,
dict_type=dto.dict_type,
dict_name=dto.dict_name,
remark=dto.remark,
sort=dto.sort,
user_id=user.id,
now=shanghai_now_naive(),
)
await self.repository.session.commit()
async def delete_types(self, ids: list[int]) -> None:
await self.repository.delete_dict_types(ids)
await self.repository.session.commit()
async def page_data(
self,
*,
dict_type_id: int,
dict_label: str | None,
dict_value: str | None,
page: int,
limit: int,
) -> dict[str, Any]:
rows, total = await self.repository.page_dict_data(
dict_type_id=dict_type_id,
dict_label=dict_label,
dict_value=dict_value,
page=max(1, page),
limit=max(0, limit),
)
return {"list": [self._data_vo(row, include_names=True) for row in rows], "total": total}
async def get_data(self, data_id: int) -> dict[str, Any] | None:
row = await self.repository.get_dict_data(data_id)
return None if row is None else self._data_vo(row, include_names=False)
async def save_data(self, dto: DictDataPayload, user: AuthUser) -> None:
# Java compares dict_label against dictValue here; retain that behavior for compatibility.
if await self.repository.dict_data_label_exists(dto.dict_type_id, dto.dict_value):
raise AppError(10128)
await self.repository.insert_dict_data(
data_id=dto.id if dto.id is not None else snowflake.next_id(),
dict_type_id=dto.dict_type_id,
dict_label=dto.dict_label,
dict_value=dto.dict_value,
remark=dto.remark,
sort=dto.sort,
user_id=user.id,
now=shanghai_now_naive(),
)
await self._clear_dict_cache(dto.dict_type_id)
await self.repository.session.commit()
async def update_data(self, dto: DictDataPayload, user: AuthUser) -> None:
if await self.repository.dict_data_label_exists(dto.dict_type_id, dto.dict_value, exclude_id=dto.id):
raise AppError(10128)
await self.repository.update_dict_data(
data_id=dto.id,
dict_type_id=dto.dict_type_id,
dict_label=dto.dict_label,
dict_value=dto.dict_value,
remark=dto.remark,
sort=dto.sort,
user_id=user.id,
now=shanghai_now_naive(),
)
await self._clear_dict_cache(dto.dict_type_id)
await self.repository.session.commit()
async def delete_data(self, ids: list[int]) -> None:
if ids:
codes = await self.repository.dict_type_codes_for_data_ids(ids)
if codes:
await cast(Any, self.redis.delete)(*[f"sys:dict:data:{code}" for code in codes])
await self.repository.delete_dict_data(ids)
await self.repository.session.commit()
async def items(self, dict_type: str) -> list[dict[str, Any]] | None:
if not dict_type.strip():
return None
key = f"sys:dict:data:{dict_type}"
cached = JavaRedisCodec.decode(await cast(Any, self.redis.get)(key))
if isinstance(cached, list):
return cast(list[dict[str, Any]], cached)
rows = await self.repository.dict_items(dict_type)
await cast(Any, self.redis.set)(
key,
JavaRedisCodec.encode(
rows,
item_java_type="xiaozhi.modules.sys.vo.SysDictDataItem",
),
ex=24 * 60 * 60,
)
return rows
async def _clear_dict_cache(self, type_id: int | None) -> None:
dict_type = await self.repository.dict_type_code(type_id)
if dict_type is not None:
await cast(Any, self.redis.delete)(f"sys:dict:data:{dict_type}")
@staticmethod
def _type_vo(row: dict[str, Any], *, include_names: bool) -> dict[str, Any]:
return {
"id": row.get("id"),
"dictType": row.get("dict_type"),
"dictName": row.get("dict_name"),
"remark": row.get("remark"),
"sort": row.get("sort"),
"creator": row.get("creator"),
"creatorName": row.get("creator_name") if include_names else None,
"createDate": row.get("create_date"),
"updater": row.get("updater"),
"updaterName": row.get("updater_name") if include_names else None,
"updateDate": row.get("update_date"),
}
@staticmethod
def _data_vo(row: dict[str, Any], *, include_names: bool) -> dict[str, Any]:
return {
"id": row.get("id"),
"dictTypeId": row.get("dict_type_id"),
"dictLabel": row.get("dict_label"),
"dictValue": row.get("dict_value"),
"remark": row.get("remark"),
"sort": row.get("sort"),
"creator": row.get("creator"),
"creatorName": row.get("creator_name") if include_names else None,
"createDate": row.get("create_date"),
"updater": row.get("updater"),
"updaterName": row.get("updater_name") if include_names else None,
"updateDate": row.get("update_date"),
}
class ServerActionService:
def __init__(self, param_service: SysParamService, *, redis: Redis | None = None):
self.param_service = param_service
self.redis = redis or get_redis()
async def server_list(self) -> list[str]:
value = await self.param_service.get_value("server.websocket", from_cache=True)
if value is None or not value.strip():
return []
values = value.split(";")
while values and values[-1] == "":
values.pop()
return values
async def emit(self, dto: EmitServerActionRequest) -> bool:
action = dto.action.lower() if dto.action is not None else None
if action not in {"restart", "update_config"}:
raise AppError(10095)
websocket_text = await self.param_service.get_value("server.websocket", from_cache=True)
if websocket_text is None or not websocket_text.strip():
raise AppError(10096)
if dto.target_ws not in websocket_text.split(";"):
raise AppError(10097)
payload_secret = await self.param_service.get_value("server.secret", from_cache=True)
device_id = str(uuid.uuid4())
client_id = str(uuid.uuid4())
await cast(Any, self.redis.set)(
f"tmp_register_mac:{device_id}",
JavaRedisCodec.encode("true"),
ex=300,
)
authentication_secret = await self.param_service.get_value("server.secret", from_cache=False)
if authentication_secret is None or not authentication_secret.strip():
raise AppError(10045)
timestamp = int(time.time())
content = f"{client_id}|{device_id}|{timestamp}"
signature = hmac.new(authentication_secret.encode(), content.encode(), digestmod=hashlib.sha256).digest()
token = base64.urlsafe_b64encode(signature).rstrip(b"=").decode() + f".{timestamp}"
headers = {
"device-id": device_id,
"client-id": client_id,
"authorization": f"Bearer {token}",
}
if payload_secret is None:
raise AppError(10045)
payload = {"type": "server", "action": action, "content": {"secret": payload_secret}}
try:
async with connect(dto.target_ws, additional_headers=headers, open_timeout=3) as websocket:
await websocket.send(json.dumps(payload, ensure_ascii=False, separators=(",", ":")))
deadline = time.monotonic() + 120
while True:
remaining = deadline - time.monotonic()
if remaining <= 0:
raise TimeoutError
raw = await asyncio.wait_for(websocket.recv(), timeout=remaining)
response = json.loads(raw)
if (
isinstance(response, dict)
and response.get("status") == "success"
and response.get("type") == "server"
and isinstance(response.get("content"), dict)
and response["content"].get("action") is not None
):
return True
except Exception as exc:
raise AppError(10045) from exc
@@ -0,0 +1,50 @@
from __future__ import annotations
import logging
from sqlalchemy import text
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.redis import java_hget, java_hset
logger = logging.getLogger(__name__)
class SystemParamService:
CACHE_KEY = "sys:params"
def __init__(self, session: AsyncSession):
self.session = session
async def get_value(self, code: str, *, from_cache: bool = True) -> str | None:
if from_cache:
try:
cached = await java_hget(self.CACHE_KEY, code)
if cached is not None:
return str(cached)
except Exception:
logger.warning("Redis parameter cache read failed for %s", code, exc_info=True)
result = await self.session.execute(
text("SELECT param_value FROM sys_params WHERE param_code = :code LIMIT 1"),
{"code": code},
)
value = result.scalar_one_or_none()
if value is not None and from_cache:
try:
await java_hset(self.CACHE_KEY, code, str(value))
except Exception:
logger.warning("Redis parameter cache write failed for %s", code, exc_info=True)
return None if value is None else str(value)
async def set_value(self, code: str, value: str) -> int:
result = await self.session.execute(
text(
"UPDATE sys_params SET param_value = :value, update_date = CURRENT_TIMESTAMP WHERE param_code = :code"
),
{"code": code, "value": value},
)
try:
await java_hset(self.CACHE_KEY, code, value)
except Exception:
logger.warning("Redis parameter cache write failed for %s", code, exc_info=True)
return int(getattr(result, "rowcount", 0) or 0)
@@ -0,0 +1,157 @@
from __future__ import annotations
from functools import lru_cache
from typing import Any
from app.core.config import get_settings
from app.core.i18n import _load_properties, message_for, resolve_language
from app.core.ids import snowflake
from app.core.redis import JavaRedisCodec, get_redis
from app.core.security import AuthUser, shanghai_now_naive
from app.repositories.timbre import TimbreRepository
from app.schemas.timbre import TimbreBody
def _details(row: dict[str, Any]) -> dict[str, Any]:
return {
"id": row.get("id"),
"languages": row.get("languages"),
"name": row.get("name"),
"remark": row.get("remark"),
"referenceAudio": row.get("reference_audio"),
"referenceText": row.get("reference_text"),
# TimbreDetailsVO.sort is primitive long, whose Java serializer always
# emits a string and whose null conversion default is zero.
"sort": str(row.get("sort") if row.get("sort") is not None else 0),
"ttsModelId": row.get("tts_model_id"),
"ttsVoice": row.get("tts_voice"),
"voiceDemo": row.get("voice_demo"),
}
class TimbreService:
def __init__(self, repository: TimbreRepository):
self.repository = repository
@staticmethod
def _validate(body: TimbreBody, language: str | None) -> None:
from app.core.errors import AppError
for value, message in (
(body.languages, "timbre.languages.require"),
(body.name, "timbre.name.require"),
(body.tts_model_id, "timbre.ttsModelId.require"),
(body.tts_voice, "timbre.ttsVoice.require"),
):
if value is None or not value.strip():
# TimbreController invokes ValidatorUtils directly. That
# utility wraps validation text in RenException(String), whose
# response code is 500 rather than the global @Valid code 10034.
raise AppError(500, _validation_message(message, language))
if body.sort is not None and body.sort < 0:
raise AppError(500, _validation_message("sort.number", language))
async def page(
self,
tts_model_id: str | None,
name: str | None,
page: str | None,
limit: str | None,
language: str | None,
) -> dict[str, Any]:
if tts_model_id is None or not tts_model_id.strip():
from app.core.errors import AppError
raise AppError(500, _validation_message("timbre.ttsModelId.require", language))
current, size = max(int(page or "1"), 1), int(limit or "10")
rows, total = await self.repository.page(
tts_model_id=tts_model_id, name=name, offset=(current - 1) * size, limit=size
)
return {"total": total, "list": [_details(row) for row in rows]}
async def save(self, body: TimbreBody, user: AuthUser, language: str | None) -> None:
self._validate(body, language)
values = self._values(body, user, str(snowflake.next_id()))
async with self.repository.session.begin():
await self.repository.insert(values)
async def update(
self, timbre_id: str, body: TimbreBody, user: AuthUser, language: str | None
) -> None:
self._validate(body, language)
values = self._values(body, user, timbre_id)
async with self.repository.session.begin():
await self.repository.update(values)
await get_redis().delete(f"timbre:details:{timbre_id}")
async def delete(self, ids: list[str]) -> None:
async with self.repository.session.begin():
await self.repository.delete(ids)
async def voices(self, model_id: str, voice_name: str | None, user: AuthUser, language: str | None) -> Any:
normal, clones = await self.repository.voices(model_id, voice_name, user.id)
values = [
{
"id": row.get("id"),
"name": row.get("name"),
"voiceDemo": row.get("voice_demo"),
"languages": row.get("languages"),
"isClone": False,
}
for row in normal
]
prefix = message_for(10158, language)
redis = get_redis()
for row in clones:
name = prefix + str(row.get("name") or "")
voice = {
"id": row.get("id"),
"name": name,
"voiceDemo": row.get("voice_demo"),
"languages": row.get("languages"),
"isClone": True,
}
await redis.set(f"timbre:name:{row['id']}", JavaRedisCodec.encode(name))
values.insert(0, voice)
return values or None
@staticmethod
def _values(body: TimbreBody, user: AuthUser, timbre_id: str) -> dict[str, Any]:
assert body.languages is not None
assert body.name is not None
assert body.tts_model_id is not None
assert body.tts_voice is not None
return {
"id": timbre_id,
"languages": body.languages,
"name": body.name,
"remark": body.remark,
"reference_audio": body.reference_audio,
"reference_text": body.reference_text,
"sort": body.sort if body.sort is not None else 0,
"tts_model_id": body.tts_model_id,
"tts_voice": body.tts_voice,
"voice_demo": body.voice_demo,
"creator": user.id,
"updater": user.id,
"now": shanghai_now_naive(),
}
_VALIDATION_FILES = {
"zh-CN": "validation_zh_CN.properties",
"zh-TW": "validation_zh_TW.properties",
"en-US": "validation_en_US.properties",
"de-DE": "validation_de_DE.properties",
"vi-VN": "validation_vi_VN.properties",
"pt-BR": "validation_pt_BR.properties",
}
@lru_cache(maxsize=64)
def _validation_message(key: str, accept_language: str | None) -> str:
language = resolve_language(accept_language)
directory = get_settings().i18n_dir
messages = _load_properties(directory / "validation.properties")
messages.update(_load_properties(directory / _VALIDATION_FILES[language]))
return messages.get(key, key)
@@ -0,0 +1,334 @@
from __future__ import annotations
import json
import uuid
from collections.abc import Mapping, Sequence
from typing import Any
import httpx
from redis.asyncio import Redis
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.config import get_settings
from app.core.errors import AppError
from app.core.i18n import message_for
from app.core.security import AuthUser, shanghai_now_naive
from app.integrations.voice_clone import VoiceCloneIntegration, VoiceCloneProviderError
from app.repositories.voiceclone import VoiceCloneRepository
from app.schemas.voiceclone import VoiceResourceCreateRequest
from app.services.device import is_blank, redis_delete, redis_get, redis_set
VOICE_ORDER_COLUMNS = {
"id": "id",
"name": "name",
"modelId": "model_id",
"model_id": "model_id",
"voiceId": "voice_id",
"voice_id": "voice_id",
"userId": "user_id",
"user_id": "user_id",
"trainStatus": "train_status",
"train_status": "train_status",
"createDate": "create_date",
"create_date": "create_date",
}
class VoiceCloneService:
def __init__(
self,
session: AsyncSession,
*,
redis_client: Redis | None = None,
http_client: httpx.AsyncClient | None = None,
provider: VoiceCloneIntegration | None = None,
):
self.session = session
self.repository = VoiceCloneRepository(session)
self.redis = redis_client
self.provider = provider or VoiceCloneIntegration(
timeout_seconds=get_settings().external_request_timeout_seconds,
client=http_client,
)
async def page(self, query: Mapping[str, Any], *, user_id: int | None = None) -> dict[str, Any]:
page = int(str(query.get("page") or "1"))
limit = int(str(query.get("limit") or "10"))
name_value = query.get("name")
name = None if name_value is None else str(name_value)
effective_user = str(user_id) if user_id is not None else self._optional_string(query.get("userId"))
requested = query.get("orderField")
requested_fields = [requested] if isinstance(requested, str) else list(requested or [])
order_fields = [VOICE_ORDER_COLUMNS[field] for field in requested_fields if field in VOICE_ORDER_COLUMNS]
if not order_fields:
order_fields = ["create_date"]
ascending = str(query.get("order") or "").lower() == "asc" if requested_fields else True
rows = await self.repository.page(
page=page,
limit=limit,
name=name,
user_id=effective_user,
order_fields=order_fields,
ascending=ascending,
)
return {
"total": await self.repository.count(name=name, user_id=effective_user),
"list": await self._response_list(rows),
}
async def get_detail(self, voice_id: str) -> dict[str, Any] | None:
row = await self.repository.get(voice_id)
if row is None:
return None
return await self._response(row, include_has_voice=False)
async def get_by_user(self, user_id: int) -> list[dict[str, Any]]:
del user_id
# VoiceCloneServiceImpl.getByUserId orders ai_voice_clone by the
# nonexistent ``created_at`` column (the schema uses ``create_date``).
# The Java endpoint therefore consistently exposes its generic
# code-500 envelope before result conversion.
raise AppError(500)
async def create_resources(self, request: VoiceResourceCreateRequest, *, actor: AuthUser) -> None:
model_id = request.model_id or ""
config = await self._model_config(model_id)
if config is None:
raise AppError(10152)
provider_type = config.get("type")
if not isinstance(provider_type, str) or not provider_type.strip():
raise AppError(10153)
voice_ids = request.voice_ids or []
for voice_id in voice_ids:
if is_blank(voice_id):
continue
if provider_type == "huoshan_double_stream" and "S_" not in voice_id:
raise AppError(10160)
if await self.repository.voice_id_count(model_id=model_id, voice_id=voice_id):
raise AppError(10159)
now = shanghai_now_naive()
prefix = now.strftime("%m%d%H%M")
values: list[dict[str, Any]] = []
for index, voice_id in enumerate(voice_ids, start=1):
values.append(
{
"id": uuid.uuid4().hex,
"name": f"{prefix}_{index}",
"model_id": model_id,
"voice_id": voice_id,
"languages": request.languages,
"user_id": request.user_id,
"voice": None,
"train_status": 0,
"train_error": None,
"creator": actor.id,
"create_date": now,
}
)
try:
await self.repository.insert_many(values)
await self.session.commit()
except Exception:
await self.session.rollback()
raise
async def delete(self, ids: Sequence[str]) -> None:
await self.repository.delete_many(ids)
await self.session.commit()
async def check_permission(self, voice_id: str | None, user: AuthUser) -> dict[str, Any]:
row = await self.repository.get(voice_id)
if row is None:
raise AppError(10144)
if int(row.get("user_id") or -1) != user.id:
raise AppError(10150)
return row
async def upload_voice(self, voice_id: str, content: bytes) -> None:
if await self.repository.get(voice_id) is None:
raise AppError(10144)
await self.repository.update_voice(voice_id, content)
await self.session.commit()
async def rename(self, voice_id: str, name: str) -> None:
if await self.repository.get(voice_id) is None:
raise AppError(10144)
await self.repository.update_name(voice_id, name)
await self.session.commit()
await redis_delete(f"timbre:name:{voice_id}", client=self.redis)
async def create_audio_id(self, voice_id: str) -> str:
row = await self.repository.get(voice_id)
if row is None or row.get("voice") is None:
raise AppError(10182)
value = str(uuid.uuid4())
await redis_set(f"voiceClone:audio:id:{value}", voice_id, client=self.redis)
return value
async def consume_audio(self, download_id: str) -> bytes | None:
key = f"voiceClone:audio:id:{download_id}"
voice_id = await redis_get(key, self.redis)
await redis_delete(key, client=self.redis)
if is_blank(None if voice_id is None else str(voice_id)):
return None
row = await self.repository.get(str(voice_id))
data = None if row is None else row.get("voice")
if data is None:
return None
result = bytes(data)
return result or None
async def clone_audio(
self,
voice_id: str,
*,
accept_language: str | None,
) -> None:
row = await self.repository.get(voice_id)
if row is None:
raise AppError(10144)
raw_voice = row.get("voice")
if raw_voice is None or len(raw_voice) == 0:
raise AppError(10151)
try:
config = await self._model_config(str(row.get("model_id") or ""))
if config is None:
raise AppError(10152)
provider_type = config.get("type")
if not isinstance(provider_type, str) or not provider_type.strip():
raise AppError(10153)
if provider_type != "huoshan_double_stream":
return
appid = config.get("appid")
access_token = config.get("access_token")
if (
not isinstance(appid, str)
or is_blank(appid)
or not isinstance(access_token, str)
or is_blank(access_token)
):
raise AppError(10155)
speaker_id = await self.provider.train_huoshan(
appid=appid,
access_token=access_token,
voice=bytes(raw_voice),
speaker_id=str(row.get("voice_id") or ""),
)
await self.repository.update_training(
voice_id,
train_status=2,
train_error="",
speaker_id=speaker_id,
)
await self.session.commit()
except AppError as exc:
await self._record_training_failure(voice_id, exc.message or message_for(exc.code, accept_language))
raise
except VoiceCloneProviderError as exc:
if exc.code in {500, 10156}:
await self._record_training_failure(voice_id, exc.message)
raise AppError(exc.code, exc.message) from exc
translated = message_for(10154, accept_language, exc.message)
await self._record_training_failure(voice_id, translated)
raise AppError(10154, translated) from exc
except Exception as exc:
translated = message_for(10154, accept_language, str(exc))
await self._record_training_failure(voice_id, translated)
raise AppError(10154, translated) from exc
async def tts_platforms(self) -> list[dict[str, Any]]:
return await self.repository.get_tts_platforms()
async def _record_training_failure(self, voice_id: str, message: str) -> None:
await self.session.rollback()
await self.repository.update_training(voice_id, train_status=3, train_error=message)
await self.session.commit()
async def _model_config(self, model_id: str) -> dict[str, Any] | None:
if is_blank(model_id):
return None
cached = await redis_get(f"model:data:{model_id}", self.redis)
cached_mapping = self._mapping(cached)
if cached_mapping is not None:
config_value = cached_mapping.get("configJson", cached_mapping.get("config_json"))
parsed = self._json_mapping(config_value)
if parsed is not None:
return parsed
row = await self.repository.get_model_config(model_id)
return None if row is None else self._json_mapping(row.get("config_json"))
async def _model_name(self, model_id: str | None) -> str | None:
if is_blank(model_id):
return None
cache_key = f"model:name:{model_id}"
cached = await redis_get(cache_key, self.redis)
if isinstance(cached, str) and cached.strip():
return cached
value = await self.repository.get_model_name(model_id or "")
if value is not None and value.strip():
await redis_set(cache_key, value, client=self.redis)
return value
async def _response_list(self, rows: Sequence[Mapping[str, Any]]) -> list[dict[str, Any]]:
user_ids = [int(row["user_id"]) for row in rows if row.get("user_id") is not None]
usernames = await self.repository.get_usernames(user_ids)
result: list[dict[str, Any]] = []
for row in rows:
result.append(await self._response(row, usernames=usernames, include_has_voice=True))
return result
async def _response(
self,
row: Mapping[str, Any],
*,
usernames: Mapping[int, str] | None = None,
include_has_voice: bool,
) -> dict[str, Any]:
user_id = None if row.get("user_id") is None else int(row["user_id"])
if user_id is None:
username = None
elif usernames is None:
username = await self.repository.get_username(user_id)
else:
username = usernames.get(user_id)
return {
"id": row.get("id"),
"name": row.get("name"),
"model_id": row.get("model_id"),
"model_name": await self._model_name(self._optional_string(row.get("model_id"))),
"voice_id": row.get("voice_id"),
"languages": row.get("languages"),
"user_id": user_id,
"user_name": username,
"train_status": row.get("train_status"),
"train_error": row.get("train_error"),
"create_date": row.get("create_date"),
"has_voice": row.get("voice") is not None if include_has_voice else None,
}
@staticmethod
def _mapping(value: Any) -> dict[str, Any] | None:
if isinstance(value, dict):
return {str(key): item for key, item in value.items() if key != "@class"}
if isinstance(value, list) and len(value) == 2 and isinstance(value[1], dict):
return {str(key): item for key, item in value[1].items()}
return None
@staticmethod
def _json_mapping(value: Any) -> dict[str, Any] | None:
if isinstance(value, dict):
return {str(key): item for key, item in value.items()}
if isinstance(value, bytes):
value = value.decode("utf-8")
if isinstance(value, str):
try:
parsed = json.loads(value)
except json.JSONDecodeError:
return None
return {str(key): item for key, item in parsed.items()} if isinstance(parsed, dict) else None
return None
@staticmethod
def _optional_string(value: Any) -> str | None:
return None if value is None else str(value)
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,164 @@
{
"schema_version": 1,
"generated_at": "2026-07-20T07:12:33.158826+00:00",
"parameters": {
"requests_per_service_per_scenario": 60,
"concurrency": 6,
"sequential_warmup_requests": 10,
"scenario_order": [
"representative-read",
"representative-crud-update",
"runtime-configuration",
"ota-check-and-signing"
],
"service_order": [
"java",
"fastapi"
]
},
"results": [
{
"service": "java",
"scenario": "representative-read",
"requests": 60,
"concurrency": 6,
"warmup_requests": 10,
"errors": 0,
"elapsed_seconds": 0.099812,
"throughput_requests_per_second": 601.128,
"latency_ms_min": 3.259,
"latency_ms_p50": 7.662,
"latency_ms_p95": 16.239,
"latency_ms_max": 20.352
},
{
"service": "fastapi",
"scenario": "representative-read",
"requests": 60,
"concurrency": 6,
"warmup_requests": 10,
"errors": 0,
"elapsed_seconds": 0.079836,
"throughput_requests_per_second": 751.538,
"latency_ms_min": 4.117,
"latency_ms_p50": 6.749,
"latency_ms_p95": 12.552,
"latency_ms_max": 16.057
},
{
"service": "java",
"scenario": "representative-crud-update",
"requests": 60,
"concurrency": 6,
"warmup_requests": 10,
"errors": 0,
"elapsed_seconds": 0.110495,
"throughput_requests_per_second": 543.013,
"latency_ms_min": 5.684,
"latency_ms_p50": 9.331,
"latency_ms_p95": 15.302,
"latency_ms_max": 19.008
},
{
"service": "fastapi",
"scenario": "representative-crud-update",
"requests": 60,
"concurrency": 6,
"warmup_requests": 10,
"errors": 0,
"elapsed_seconds": 0.12582,
"throughput_requests_per_second": 476.87,
"latency_ms_min": 6.996,
"latency_ms_p50": 11.252,
"latency_ms_p95": 20.599,
"latency_ms_max": 24.284
},
{
"service": "java",
"scenario": "runtime-configuration",
"requests": 60,
"concurrency": 6,
"warmup_requests": 10,
"errors": 0,
"elapsed_seconds": 0.105474,
"throughput_requests_per_second": 568.863,
"latency_ms_min": 3.901,
"latency_ms_p50": 8.746,
"latency_ms_p95": 16.127,
"latency_ms_max": 29.528
},
{
"service": "fastapi",
"scenario": "runtime-configuration",
"requests": 60,
"concurrency": 6,
"warmup_requests": 10,
"errors": 0,
"elapsed_seconds": 0.078418,
"throughput_requests_per_second": 765.126,
"latency_ms_min": 3.868,
"latency_ms_p50": 6.488,
"latency_ms_p95": 13.722,
"latency_ms_max": 18.923
},
{
"service": "java",
"scenario": "ota-check-and-signing",
"requests": 60,
"concurrency": 6,
"warmup_requests": 10,
"errors": 0,
"elapsed_seconds": 0.131968,
"throughput_requests_per_second": 454.655,
"latency_ms_min": 7.543,
"latency_ms_p50": 11.301,
"latency_ms_p95": 19.854,
"latency_ms_max": 21.729
},
{
"service": "fastapi",
"scenario": "ota-check-and-signing",
"requests": 60,
"concurrency": 6,
"warmup_requests": 10,
"errors": 0,
"elapsed_seconds": 0.169106,
"throughput_requests_per_second": 354.807,
"latency_ms_min": 7.066,
"latency_ms_p50": 16.208,
"latency_ms_p95": 20.344,
"latency_ms_max": 28.953
}
],
"comparisons": [
{
"scenario": "representative-read",
"p50_ratio_fastapi_over_java": 0.881,
"p95_ratio_fastapi_over_java": 0.773,
"throughput_ratio_fastapi_over_java": 1.25
},
{
"scenario": "representative-crud-update",
"p50_ratio_fastapi_over_java": 1.206,
"p95_ratio_fastapi_over_java": 1.346,
"throughput_ratio_fastapi_over_java": 0.878
},
{
"scenario": "runtime-configuration",
"p50_ratio_fastapi_over_java": 0.742,
"p95_ratio_fastapi_over_java": 0.851,
"throughput_ratio_fastapi_over_java": 1.345
},
{
"scenario": "ota-check-and-signing",
"p50_ratio_fastapi_over_java": 1.434,
"p95_ratio_fastapi_over_java": 1.025,
"throughput_ratio_fastapi_over_java": 0.78
}
],
"summary": {
"measurements": 8,
"requests_measured": 480,
"errors": 0
}
}
File diff suppressed because it is too large Load Diff
+16
View File
@@ -0,0 +1,16 @@
#!/bin/sh
set -eu
UPSTREAM=${MANAGER_API_UPSTREAM:-manager-api-fastapi:8002}
case "${UPSTREAM}" in
''|*[!A-Za-z0-9._:-]*)
echo "MANAGER_API_UPSTREAM must be a hostname-or-IP and port" >&2
exit 2
;;
esac
export MANAGER_API_UPSTREAM=${UPSTREAM}
envsubst '${MANAGER_API_UPSTREAM}' \
< /etc/nginx/nginx.conf.template \
> /tmp/manager-api-nginx.conf
exec nginx -c /tmp/manager-api-nginx.conf -g 'daemon off;'
@@ -0,0 +1,45 @@
worker_processes auto;
pid /tmp/nginx.pid;
events {
worker_connections 1024;
}
http {
include /etc/nginx/mime.types;
default_type application/octet-stream;
access_log /dev/stdout;
error_log /dev/stderr warn;
sendfile on;
keepalive_timeout 65;
client_max_body_size 100m;
upstream manager_api_fastapi {
server ${MANAGER_API_UPSTREAM};
keepalive 32;
}
server {
listen 8002;
server_name _;
location = /xiaozhi {
return 308 /xiaozhi/;
}
location /xiaozhi/ {
proxy_pass http://manager_api_fastapi;
proxy_http_version 1.1;
proxy_set_header Host $host;
proxy_set_header X-Real-IP $remote_addr;
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
proxy_set_header X-Forwarded-Proto $scheme;
proxy_set_header Connection "";
proxy_connect_timeout 10s;
proxy_send_timeout 130s;
proxy_read_timeout 130s;
proxy_request_buffering off;
proxy_buffering off;
}
}
}
@@ -0,0 +1,93 @@
services:
manager-api-migrate:
image: xiaozhi/manager-api-migrate:fastapi-0.1.0
build:
context: ../..
dockerfile: main/manager-api-fastapi/Dockerfile.migrations
environment:
LIQUIBASE_URL: ${LIQUIBASE_URL:?Set a JDBC MySQL URL}
LIQUIBASE_USERNAME: ${MYSQL_USER:?Set MYSQL_USER}
LIQUIBASE_PASSWORD: ${MYSQL_PASSWORD:?Set MYSQL_PASSWORD}
MIGRATION_POM: /migration/pom.xml
JAVA_RESOURCES_DIR: /migration/java-resources
MAVEN_BIN: mvn
restart: "no"
manager-api-fastapi:
image: xiaozhi/manager-api-fastapi:0.1.0
build:
context: ../..
dockerfile: main/manager-api-fastapi/Dockerfile
init: true
depends_on:
manager-api-migrate:
condition: service_completed_successfully
environment:
APP_ENVIRONMENT: production
APP_DATABASE_URL: ${FASTAPI_DATABASE_URL:?Set an asyncmy MySQL URL}
APP_REDIS_URL: ${REDIS_URL:?Set REDIS_URL}
APP_WORKERS: ${APP_WORKERS:-2}
APP_GRACEFUL_SHUTDOWN_SECONDS: ${APP_GRACEFUL_SHUTDOWN_SECONDS:-30}
APP_FORWARDED_ALLOW_IPS: ${APP_FORWARDED_ALLOW_IPS:-*}
APP_ALLOW_START_WITHOUT_DEPENDENCIES: "false"
expose:
- "8002"
volumes:
# During cutover this source can be the retained Java service's host
# uploadfile directory so both implementations see identical files.
- ${MANAGER_API_UPLOAD_SOURCE:-manager-api-uploads}:/data/uploads
read_only: true
tmpfs:
- /tmp:size=64m,mode=1777
restart: unless-stopped
stop_grace_period: 40s
healthcheck:
test: ["CMD", "python", "-c", "import urllib.request; urllib.request.urlopen('http://127.0.0.1:8002/xiaozhi/health/ready', timeout=2).read()"]
interval: 15s
timeout: 3s
retries: 4
start_period: 20s
manager-api-jobs:
image: xiaozhi/manager-api-fastapi:0.1.0
init: true
depends_on:
manager-api-migrate:
condition: service_completed_successfully
command: ["python", "-m", "app.jobs.worker"]
environment:
APP_ENVIRONMENT: production
APP_DATABASE_URL: ${FASTAPI_DATABASE_URL:?Set an asyncmy MySQL URL}
APP_REDIS_URL: ${REDIS_URL:?Set REDIS_URL}
APP_GRACEFUL_SHUTDOWN_SECONDS: ${APP_GRACEFUL_SHUTDOWN_SECONDS:-30}
APP_ALLOW_START_WITHOUT_DEPENDENCIES: "false"
volumes:
- ${MANAGER_API_UPLOAD_SOURCE:-manager-api-uploads}:/data/uploads
read_only: true
tmpfs:
- /tmp:size=64m,mode=1777
restart: unless-stopped
stop_grace_period: 40s
manager-api-nginx:
image: xiaozhi/manager-api-nginx:fastapi-0.1.0
build:
context: ../..
dockerfile: main/manager-api-fastapi/Dockerfile.nginx
depends_on:
manager-api-fastapi:
condition: service_healthy
ports:
- "8002:8002"
environment:
MANAGER_API_UPSTREAM: ${MANAGER_API_UPSTREAM:-manager-api-fastapi:8002}
read_only: true
tmpfs:
- /var/cache/nginx:size=32m
- /var/run:size=1m
- /tmp:size=16m
restart: unless-stopped
stop_grace_period: 15s
volumes:
manager-api-uploads:
@@ -0,0 +1,84 @@
<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0"
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0 https://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion>
<groupId>xiaozhi</groupId>
<artifactId>manager-api-liquibase-runner</artifactId>
<version>1.0.0</version>
<properties>
<maven.compiler.release>21</maven.compiler.release>
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
<java.resources.dir>${project.basedir}/../manager-api/src/main/resources</java.resources.dir>
<liquibase.version>4.20.0</liquibase.version>
<mysql.version>9.1.0</mysql.version>
<spring.version>6.2.3</spring.version>
</properties>
<dependencies>
<dependency>
<groupId>org.liquibase</groupId>
<artifactId>liquibase-core</artifactId>
<version>${liquibase.version}</version>
</dependency>
<dependency>
<groupId>org.springframework</groupId>
<artifactId>spring-jdbc</artifactId>
<version>${spring.version}</version>
</dependency>
<dependency>
<groupId>org.springframework</groupId>
<artifactId>spring-context</artifactId>
<version>${spring.version}</version>
</dependency>
<dependency>
<groupId>com.mysql</groupId>
<artifactId>mysql-connector-j</artifactId>
<version>${mysql.version}</version>
</dependency>
<dependency>
<groupId>org.slf4j</groupId>
<artifactId>slf4j-simple</artifactId>
<version>2.0.16</version>
</dependency>
</dependencies>
<build>
<sourceDirectory>${project.basedir}/migration-src</sourceDirectory>
<resources>
<resource>
<directory>${java.resources.dir}</directory>
<filtering>false</filtering>
</resource>
</resources>
<plugins>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-compiler-plugin</artifactId>
<version>3.13.0</version>
</plugin>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-shade-plugin</artifactId>
<version>3.6.0</version>
<executions>
<execution>
<phase>package</phase>
<goals>
<goal>shade</goal>
</goals>
<configuration>
<createDependencyReducedPom>false</createDependencyReducedPom>
<shadedArtifactAttached>true</shadedArtifactAttached>
<shadedClassifierName>all</shadedClassifierName>
<transformers>
<transformer implementation="org.apache.maven.plugins.shade.resource.ManifestResourceTransformer">
<mainClass>xiaozhi.migration.LiquibaseMigrationRunner</mainClass>
</transformer>
<transformer implementation="org.apache.maven.plugins.shade.resource.ServicesResourceTransformer"/>
</transformers>
</configuration>
</execution>
</executions>
</plugin>
</plugins>
</build>
</project>
@@ -0,0 +1,57 @@
package xiaozhi.migration;
import java.sql.Connection;
import java.sql.ResultSet;
import java.sql.Statement;
import javax.sql.DataSource;
import liquibase.integration.spring.SpringLiquibase;
import org.springframework.core.io.DefaultResourceLoader;
import org.springframework.jdbc.datasource.DriverManagerDataSource;
/** Runs only the original Spring Liquibase changelog; it never starts manager-api or Redis. */
public final class LiquibaseMigrationRunner {
private LiquibaseMigrationRunner() {
}
public static void main(String[] args) throws Exception {
String url = requiredEnvironment("MIGRATION_JDBC_URL");
String username = requiredEnvironment("MIGRATION_USERNAME");
String password = requiredEnvironment("MIGRATION_PASSWORD");
DriverManagerDataSource dataSource = new DriverManagerDataSource(url, username, password);
dataSource.setDriverClassName("com.mysql.cj.jdbc.Driver");
runLiquibase(dataSource);
System.out.println("Liquibase migration complete; applied changeSets=" + appliedChangeSetCount(dataSource));
}
private static void runLiquibase(DataSource dataSource) throws Exception {
SpringLiquibase liquibase = new SpringLiquibase();
liquibase.setDataSource(dataSource);
liquibase.setChangeLog("classpath:db/changelog/db.changelog-master.yaml");
liquibase.setResourceLoader(new DefaultResourceLoader(Thread.currentThread().getContextClassLoader()));
liquibase.setDropFirst(false);
liquibase.setShouldRun(true);
liquibase.afterPropertiesSet();
}
private static int appliedChangeSetCount(DataSource dataSource) throws Exception {
try (Connection connection = dataSource.getConnection();
Statement statement = connection.createStatement();
ResultSet result = statement.executeQuery("SELECT COUNT(*) FROM DATABASECHANGELOG")) {
if (!result.next()) {
throw new IllegalStateException("DATABASECHANGELOG count query returned no row");
}
return result.getInt(1);
}
}
private static String requiredEnvironment(String name) {
String value = System.getenv(name);
if (value == null || value.isBlank()) {
throw new IllegalArgumentException("Required environment variable is missing: " + name);
}
return value;
}
}
+75
View File
@@ -0,0 +1,75 @@
[project]
name = "xiaozhi-manager-api-fastapi"
version = "0.1.0"
description = "FastAPI-compatible replacement for xiaozhi manager-api"
readme = "README.md"
requires-python = ">=3.10,<3.13"
dependencies = [
"aiosqlite==0.21.0",
"asyncmy==0.2.10",
"bcrypt==4.3.0",
"cryptography==45.0.5",
"fastapi==0.116.1",
"gmssl==3.2.2",
"httpx==0.28.1",
"pillow==11.3.0",
# FastAPI 0.116 reconstructs request fields through TypeAdapter. Pydantic
# 2.13 warns that aliases on that compatibility path are ineffective; pin
# the contemporary 2.11 line so request aliases remain warning-free.
"pydantic==2.11.7",
"pydantic-settings==2.10.1",
"python-multipart==0.0.20",
"pyyaml==6.0.2",
"redis[hiredis]==6.2.0",
"sqlalchemy[asyncio]==2.0.41",
"uvicorn[standard]==0.35.0",
"websockets==15.0.1",
]
[dependency-groups]
dev = [
"mypy==1.17.0",
"pytest==8.4.1",
"pytest-asyncio==1.1.0",
"pytest-cov==6.2.1",
"respx==0.22.0",
"ruff==0.12.4",
]
[tool.uv]
default-groups = ["dev"]
[tool.pytest.ini_options]
addopts = "-ra --strict-config --strict-markers"
asyncio_mode = "auto"
testpaths = ["tests"]
markers = [
"integration: requires isolated MySQL and Redis",
"contract: compares the Java and FastAPI services",
"performance: runs the representative performance comparison",
]
[tool.ruff]
target-version = "py310"
line-length = 120
exclude = ["tests/fixtures/generated"]
[tool.ruff.lint]
select = ["E", "F", "I", "UP", "B", "ASYNC", "S"]
ignore = ["S101"]
[tool.ruff.lint.per-file-ignores]
"app/routers/*.py" = ["B008"]
[tool.mypy]
python_version = "3.10"
strict = true
plugins = ["pydantic.mypy"]
exclude = ["tests/fixtures/generated"]
[build-system]
requires = ["hatchling==1.27.0"]
build-backend = "hatchling.build"
[tool.hatch.build.targets.wheel]
packages = ["app"]
+14
View File
@@ -0,0 +1,14 @@
#!/bin/sh
set -eu
if [ "$#" -gt 0 ]; then
exec "$@"
fi
exec uvicorn app.main:app \
--host "${APP_HOST:-0.0.0.0}" \
--port "${APP_PORT:-8002}" \
--workers "${APP_WORKERS:-2}" \
--timeout-graceful-shutdown "${APP_GRACEFUL_SHUTDOWN_SECONDS:-30}" \
--proxy-headers \
--forwarded-allow-ips "${APP_FORWARDED_ALLOW_IPS:-127.0.0.1}"
@@ -0,0 +1,267 @@
#!/usr/bin/env python3
"""Extract manager-api HTTP call sites from all three in-repository consumers."""
from __future__ import annotations
import json
import re
from collections import Counter
from pathlib import Path
from typing import NamedTuple
TARGET_ROOT = Path(__file__).resolve().parents[1]
REPO_ROOT = TARGET_ROOT.parents[1]
MAIN_ROOT = REPO_ROOT / "main"
WEB_ROOT = MAIN_ROOT / "manager-web" / "src"
MOBILE_ROOT = MAIN_ROOT / "manager-mobile" / "src"
SERVER_ROOT = MAIN_ROOT / "xiaozhi-server"
WEB_CHAIN = re.compile(
r"\.url\(\s*(?P<quote>[`'\"])(?P<url>.*?)(?P=quote)\s*\)"
r"(?:\s*//[^\n]*)?\s*\.method\(\s*['\"](?P<method>[A-Za-z]+)['\"]\s*\)",
re.DOTALL,
)
WEB_CONFIG = re.compile(
r"\burl\s*:\s*(?P<quote>[`'\"])(?P<url>.*?)(?P=quote)\s*,"
r"\s*method\s*:\s*['\"](?P<method>[A-Za-z]+)['\"]",
re.DOTALL,
)
WEB_TEMPLATE_URL = re.compile(
r"(?P<quote>`)(?P<url>\$\{(?:(?:Api|api)\.)?getServiceUrl\(\)\}/.*?)"
r"(?P=quote)"
)
WEB_CONCAT_URL = re.compile(
r"(?:(?:Api|api)\.)?getServiceUrl\(\)\s*\+\s*(?P<quote>`)(?P<url>/.*?)(?P=quote)"
)
MOBILE_HTTP = re.compile(
r"http\.(?P<method>Get|Post|Put|Delete|Patch)(?:<[^\n(]*>)?\(\s*"
r"(?P<quote>[`'\"])(?P<url>.*?)(?P=quote)"
)
MOBILE_UNI = re.compile(
r"uni\.request\(\s*\{(?:(?!\}\s*\)).)*?\burl\s*:\s*"
r"(?P<quote>[`'\"])(?P<url>.*?)(?P=quote)\s*,\s*"
r"method\s*:\s*['\"](?P<method>[A-Za-z]+)['\"]",
re.DOTALL,
)
SERVER_CLIENT = re.compile(
r"\._execute_async_request\(\s*['\"](?P<method>[A-Za-z]+)['\"]\s*,\s*"
r"f?(?P<quote>[`'\"])(?P<url>/.*?)(?P=quote)",
re.DOTALL,
)
SERVER_DIRECT = re.compile(
r"f(?P<quote>[`'\"])\{api_url\}(?P<url>/device/address-book/call)(?P=quote)"
)
JS_EXPRESSION = re.compile(r"\$\{([^{}]+)\}")
PYTHON_EXPRESSION = re.compile(r"\{([^{}]+)\}")
PATH_PARAMETER = re.compile(r"\{[^/{}]+\}")
class CallSite(NamedTuple):
consumer: str
method: str
path: str
source: str
def _source(path: Path, text: str, offset: int) -> str:
line = text.count("\n", 0, offset) + 1
return f"{path.relative_to(REPO_ROOT).as_posix()}:{line}"
def _parameter_name(expression: str) -> str:
identifiers = re.findall(r"[A-Za-z_][A-Za-z0-9_]*", expression)
ignored = {"getServiceUrl", "encodeURIComponent", "toString", "value"}
useful = [item for item in identifiers if item not in ignored]
return useful[-1] if useful else "value"
def normalize_path(raw: str) -> str:
value = raw.strip()
for marker in (
"${getServiceUrl()}",
"${Api.getServiceUrl()}",
"${api.getServiceUrl()}",
"${baseUrlInput.value}",
"${getEnvBaseUrl()}",
"{api_url}",
):
if value.startswith(marker):
value = value[len(marker) :]
break
value = value.split("?", 1)[0].split("#", 1)[0]
value = JS_EXPRESSION.sub(lambda match: "{" + _parameter_name(match.group(1)) + "}", value)
value = PYTHON_EXPRESSION.sub(lambda match: "{" + _parameter_name(match.group(1)) + "}", value)
if not value.startswith("/"):
raise ValueError(f"consumer URL does not resolve to a manager-api path: {raw!r}")
return re.sub(r"/{2,}", "/", value).rstrip("/") or "/"
def _iter_source_files(root: Path, suffixes: set[str]) -> list[Path]:
return sorted(
path
for path in root.rglob("*")
if path.is_file()
and path.suffix in suffixes
and "node_modules" not in path.parts
and "dist" not in path.parts
)
def extract_web() -> list[CallSite]:
calls: list[CallSite] = []
for path in _iter_source_files(WEB_ROOT, {".js", ".mjs", ".vue"}):
text = path.read_text(encoding="utf-8")
covered: list[tuple[int, int]] = []
for pattern in (WEB_CHAIN, WEB_CONFIG):
for match in pattern.finditer(text):
raw_url = match.group("url")
if "getServiceUrl()" not in raw_url:
continue
calls.append(
CallSite(
"manager-web",
match.group("method").upper(),
normalize_path(raw_url),
_source(path, text, match.start("url")),
)
)
covered.append(match.span("url"))
def already_covered(offset: int, spans: list[tuple[int, int]] = covered) -> bool:
return any(start <= offset < end for start, end in spans)
for pattern in (WEB_TEMPLATE_URL, WEB_CONCAT_URL):
for match in pattern.finditer(text):
if already_covered(match.start("url")):
continue
line_start = text.rfind("\n", 0, match.start()) + 1
line_end = text.find("\n", match.end())
line = text[line_start : len(text) if line_end < 0 else line_end]
if "console.log" in line:
continue
calls.append(
CallSite(
"manager-web",
"GET",
normalize_path(match.group("url")),
_source(path, text, match.start("url")),
)
)
literal_builders = len(
re.findall(r"\.url\(\s*[`'\"]\$\{getServiceUrl\(\)\}", text)
)
parsed_builders = sum(
1
for match in WEB_CHAIN.finditer(text)
if "getServiceUrl()" in match.group("url")
)
if literal_builders != parsed_builders:
raise RuntimeError(
f"unparsed manager-web request builder(s) in {path}: "
f"found={literal_builders}, parsed={parsed_builders}"
)
return calls
def extract_mobile() -> list[CallSite]:
calls: list[CallSite] = []
for path in _iter_source_files(MOBILE_ROOT, {".ts", ".vue"}):
text = path.read_text(encoding="utf-8")
parsed_http = list(MOBILE_HTTP.finditer(text))
raw_http_count = len(re.findall(r"\bhttp\.(?:Get|Post|Put|Delete|Patch)(?:<|\()", text))
if raw_http_count != len(parsed_http):
raise RuntimeError(
f"unparsed manager-mobile http call(s) in {path}: "
f"found={raw_http_count}, parsed={len(parsed_http)}"
)
for match in parsed_http:
calls.append(
CallSite(
"manager-mobile",
match.group("method").upper(),
normalize_path(match.group("url")),
_source(path, text, match.start("url")),
)
)
for match in MOBILE_UNI.finditer(text):
raw_url = match.group("url")
if not raw_url.startswith(("${baseUrlInput.value}", "${getEnvBaseUrl()}")):
continue
calls.append(
CallSite(
"manager-mobile",
match.group("method").upper(),
normalize_path(raw_url),
_source(path, text, match.start("url")),
)
)
return calls
def extract_server() -> list[CallSite]:
calls: list[CallSite] = []
path = SERVER_ROOT / "config" / "manage_api_client.py"
text = path.read_text(encoding="utf-8")
matches = list(SERVER_CLIENT.finditer(text))
raw_count = len(re.findall(r"\._execute_async_request\(", text))
if raw_count != len(matches):
raise RuntimeError(
f"unparsed xiaozhi-server manager client call(s): found={raw_count}, parsed={len(matches)}"
)
for match in matches:
calls.append(
CallSite(
"xiaozhi-server",
match.group("method").upper(),
normalize_path(match.group("url")),
_source(path, text, match.start("url")),
)
)
direct_path = SERVER_ROOT / "plugins_func" / "functions" / "call_device.py"
direct_text = direct_path.read_text(encoding="utf-8")
direct_matches = list(SERVER_DIRECT.finditer(direct_text))
if len(direct_matches) != 1:
raise RuntimeError(f"expected one direct manager-api call in {direct_path}, found {len(direct_matches)}")
match = direct_matches[0]
calls.append(
CallSite(
"xiaozhi-server",
"GET",
normalize_path(match.group("url")),
_source(direct_path, direct_text, match.start("url")),
)
)
return calls
def canonical_route(method: str, path: str) -> tuple[str, str]:
return method, PATH_PARAMETER.sub("{}", path)
def build_manifest() -> dict[str, object]:
calls = sorted(
extract_web() + extract_mobile() + extract_server(),
key=lambda item: (item.consumer, item.source, item.method, item.path),
)
consumers: dict[str, dict[str, object]] = {}
for consumer in ("manager-web", "manager-mobile", "xiaozhi-server"):
selected = [item for item in calls if item.consumer == consumer]
consumers[consumer] = {
"callSites": len(selected),
"uniqueRoutes": len({canonical_route(item.method, item.path) for item in selected}),
"methods": dict(sorted(Counter(item.method for item in selected).items())),
}
all_routes = {canonical_route(item.method, item.path) for item in calls}
return {
"count": len(calls),
"uniqueRoutes": len(all_routes),
"consumers": consumers,
"calls": [item._asdict() for item in calls],
}
if __name__ == "__main__":
print(json.dumps(build_manifest(), ensure_ascii=False, indent=2, sort_keys=False))
@@ -0,0 +1,168 @@
#!/usr/bin/env python3
"""Extract the Spring MVC contract without starting the Java application."""
from __future__ import annotations
import argparse
import fnmatch
import json
import re
from dataclasses import asdict, dataclass
from pathlib import Path
MAPPING_RE = re.compile(r"@(Get|Post|Put|Delete|Patch)Mapping(?:\((.*)\))?")
CLASS_MAPPING_RE = re.compile(r"@RequestMapping\(\s*(?:value\s*=\s*)?[\"']([^\"']*)[\"']")
PATH_RE = re.compile(r"[\"']([^\"']*)[\"']")
METHOD_RE = re.compile(r"\bpublic\s+(?:<[^>]+>\s+)?[^=(;]+?\s+(\w+)\s*\(")
PERMISSION_RE = re.compile(r'@RequiresPermissions\(\s*"([^"]+)"')
PUBLIC_PATTERNS = (
"/ota/**",
"/otaMag/download/**",
"/webjars/**",
"/druid/**",
"/v3/api-docs/**",
"/doc.html",
"/favicon.ico",
"/user/captcha",
"/user/smsVerification",
"/user/login",
"/user/pub-config",
"/user/register",
"/user/retrieve-password",
"/agent/chat-history/download/**",
"/agent/play/**",
"/voiceClone/play/**",
)
SERVER_PATTERNS = (
"/config/**",
"/device/address-book/call",
"/agent/chat-history/report",
"/agent/chat-summary/**",
"/agent/chat-title/**",
)
@dataclass(frozen=True, slots=True)
class Route:
method: str
path: str
controller: str
handler: str
auth: str
permission: str | None
source: str
line: int
java_signature: str
def _spring_match(path: str, pattern: str) -> bool:
return fnmatch.fnmatchcase(path, pattern.replace("**", "*"))
def classify_auth(path: str) -> str:
if any(_spring_match(path, pattern) for pattern in PUBLIC_PATTERNS):
return "anonymous"
if any(_spring_match(path, pattern) for pattern in SERVER_PATTERNS):
return "server-secret"
return "database-token"
def join_paths(base: str, child: str) -> str:
if not base:
base = "/"
if not base.startswith("/"):
base = "/" + base
if not child:
return base
if not child.startswith("/"):
child = "/" + child
if base == "/":
return child
return base.rstrip("/") + child
def extract_controller(path: Path, root: Path) -> list[Route]:
source = path.read_text(encoding="utf-8")
class_position = source.find(" class ")
class_header = source[:class_position] if class_position >= 0 else source
class_match = list(CLASS_MAPPING_RE.finditer(class_header))
base = class_match[-1].group(1) if class_match else ""
controller_match = re.search(r"public\s+class\s+(\w+)", source)
controller = controller_match.group(1) if controller_match else path.stem
lines = source.splitlines()
routes: list[Route] = []
for index, line in enumerate(lines):
mapping = MAPPING_RE.search(line)
if not mapping:
continue
method = mapping.group(1).upper()
arguments = mapping.group(2) or ""
path_match = PATH_RE.search(arguments)
child = path_match.group(1) if path_match else ""
decorator_block: list[str] = [line]
signature_lines: list[str] = []
handler = "unknown"
for cursor in range(index + 1, min(index + 40, len(lines))):
candidate = lines[cursor]
if not signature_lines and candidate.lstrip().startswith("@"):
decorator_block.append(candidate)
continue
signature_lines.append(candidate.strip())
signature = " ".join(signature_lines)
handler_match = METHOD_RE.search(signature)
if handler_match:
handler = handler_match.group(1)
break
if "{" in candidate and "public " not in signature:
break
permission_match = PERMISSION_RE.search("\n".join(decorator_block))
route_path = join_paths(base, child)
routes.append(
Route(
method=method,
path=route_path,
controller=controller,
handler=handler,
auth=classify_auth(route_path),
permission=permission_match.group(1) if permission_match else None,
source=str(path.relative_to(root)),
line=index + 1,
java_signature=" ".join(signature_lines),
)
)
return routes
def extract_routes(java_root: Path, repository_root: Path) -> list[Route]:
routes: list[Route] = []
for controller in sorted(java_root.rglob("*Controller.java")):
routes.extend(extract_controller(controller, repository_root))
return sorted(routes, key=lambda route: (route.path, route.method, route.controller, route.handler))
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--repository-root", type=Path, default=Path(__file__).resolve().parents[3])
parser.add_argument("--output", type=Path)
args = parser.parse_args()
repository_root = args.repository_root.resolve()
java_root = repository_root / "main" / "manager-api" / "src" / "main" / "java"
routes = extract_routes(java_root, repository_root)
payload = {
"source": "main/manager-api",
"contextPath": "/xiaozhi",
"count": len(routes),
"routes": [asdict(route) for route in routes],
}
serialized = json.dumps(payload, ensure_ascii=False, indent=2) + "\n"
if args.output:
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(serialized, encoding="utf-8")
else:
print(serialized, end="")
return 0
if __name__ == "__main__":
raise SystemExit(main())
+146
View File
@@ -0,0 +1,146 @@
#!/bin/sh
set -eu
SCRIPT_DIR=$(CDPATH= cd -- "$(dirname -- "$0")" && pwd)
TARGET_DIR=$(CDPATH= cd -- "${SCRIPT_DIR}/.." && pwd)
REPOSITORY_ROOT=$(CDPATH= cd -- "${TARGET_DIR}/../.." && pwd)
RUNTIME="${REPOSITORY_ROOT}/.runtime"
STATE_DIR="${TARGET_DIR}/.test-runtime"
MYSQL_BASE="${RUNTIME}/mysql"
MYSQL_PORT="${TEST_MYSQL_PORT:-13316}"
MYSQL_DATA="${STATE_DIR}/mysql-data"
MYSQL_SOCKET="${STATE_DIR}/mysql.sock"
MYSQL_PID="${STATE_DIR}/mysql.pid"
MYSQL_LOG="${STATE_DIR}/mysql.log"
REDIS_PORT="${TEST_REDIS_PORT:-16379}"
REDIS_DIR="${STATE_DIR}/redis-data"
REDIS_PID="${STATE_DIR}/redis.pid"
REDIS_LOG="${STATE_DIR}/redis.log"
TEST_USER="xiaozhi_test"
TEST_PASSWORD="isolated-test-only"
JAVA_DATABASE="manager_java_test"
FASTAPI_DATABASE="manager_fastapi_test"
require_binaries() {
for binary in "${MYSQL_BASE}/bin/mysqld" "${MYSQL_BASE}/bin/mysql" \
"${RUNTIME}/redis/bin/redis-server" "${RUNTIME}/redis/bin/redis-cli"; do
if [ ! -x "${binary}" ]; then
echo "Missing isolated-test binary: ${binary}" >&2
exit 1
fi
done
}
mysql_ready() {
"${MYSQL_BASE}/bin/mysqladmin" --protocol=SOCKET --socket="${MYSQL_SOCKET}" -uroot ping >/dev/null 2>&1
}
redis_ready() {
"${RUNTIME}/redis/bin/redis-cli" -h 127.0.0.1 -p "${REDIS_PORT}" ping >/dev/null 2>&1
}
wait_until() {
description=$1
shift
attempts=0
until "$@"; do
attempts=$((attempts + 1))
if [ "${attempts}" -ge 100 ]; then
echo "Timed out waiting for ${description}" >&2
exit 1
fi
sleep 0.1
done
}
start_mysql() {
mkdir -p "${STATE_DIR}" "${MYSQL_DATA}"
if [ ! -d "${MYSQL_DATA}/mysql" ]; then
"${MYSQL_BASE}/bin/mysqld" --no-defaults --initialize-insecure \
--basedir="${MYSQL_BASE}" --datadir="${MYSQL_DATA}" --log-error="${MYSQL_LOG}"
fi
if ! mysql_ready; then
"${MYSQL_BASE}/bin/mysqld" --no-defaults --daemonize \
--basedir="${MYSQL_BASE}" --datadir="${MYSQL_DATA}" --port="${MYSQL_PORT}" \
--socket="${MYSQL_SOCKET}" --pid-file="${MYSQL_PID}" --log-error="${MYSQL_LOG}" \
--bind-address=127.0.0.1 --mysqlx=0 --skip-name-resolve \
--character-set-server=utf8mb4 --collation-server=utf8mb4_unicode_ci \
--default-time-zone=+08:00
wait_until "isolated MySQL" mysql_ready
fi
"${MYSQL_BASE}/bin/mysql" --protocol=SOCKET --socket="${MYSQL_SOCKET}" -uroot \
-e "CREATE DATABASE IF NOT EXISTS ${JAVA_DATABASE} CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci; CREATE DATABASE IF NOT EXISTS ${FASTAPI_DATABASE} CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci; CREATE USER IF NOT EXISTS '${TEST_USER}'@'127.0.0.1' IDENTIFIED BY '${TEST_PASSWORD}'; GRANT ALL PRIVILEGES ON ${JAVA_DATABASE}.* TO '${TEST_USER}'@'127.0.0.1'; GRANT ALL PRIVILEGES ON ${FASTAPI_DATABASE}.* TO '${TEST_USER}'@'127.0.0.1'; FLUSH PRIVILEGES;"
}
start_redis() {
mkdir -p "${REDIS_DIR}"
if ! redis_ready; then
"${RUNTIME}/redis/bin/redis-server" --daemonize yes --bind 127.0.0.1 \
--port "${REDIS_PORT}" --pidfile "${REDIS_PID}" --dir "${REDIS_DIR}" \
--logfile "${REDIS_LOG}" --save "" --appendonly no --databases 16
wait_until "isolated Redis" redis_ready
fi
}
start() {
require_binaries
start_mysql
start_redis
echo "Isolated MySQL ${MYSQL_PORT} and Redis ${REDIS_PORT} are ready."
}
stop() {
if redis_ready; then
"${RUNTIME}/redis/bin/redis-cli" -h 127.0.0.1 -p "${REDIS_PORT}" shutdown nosave >/dev/null
fi
if mysql_ready; then
"${MYSQL_BASE}/bin/mysqladmin" --protocol=SOCKET --socket="${MYSQL_SOCKET}" -uroot shutdown
fi
echo "Isolated services stopped."
}
reset() {
start
"${MYSQL_BASE}/bin/mysql" --protocol=SOCKET --socket="${MYSQL_SOCKET}" -uroot \
-e "DROP DATABASE IF EXISTS ${JAVA_DATABASE}; DROP DATABASE IF EXISTS ${FASTAPI_DATABASE}; CREATE DATABASE ${JAVA_DATABASE} CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci; CREATE DATABASE ${FASTAPI_DATABASE} CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci; GRANT ALL PRIVILEGES ON ${JAVA_DATABASE}.* TO '${TEST_USER}'@'127.0.0.1'; GRANT ALL PRIVILEGES ON ${FASTAPI_DATABASE}.* TO '${TEST_USER}'@'127.0.0.1';"
"${RUNTIME}/redis/bin/redis-cli" -h 127.0.0.1 -p "${REDIS_PORT}" flushall >/dev/null
echo "Only the isolated test schemas and isolated Redis instance were reset."
}
migrate_one() {
database=$1
LIQUIBASE_URL="jdbc:mysql://127.0.0.1:${MYSQL_PORT}/${database}?useUnicode=true&characterEncoding=UTF-8&serverTimezone=Asia/Shanghai&allowMultiQueries=true" \
LIQUIBASE_USERNAME="${TEST_USER}" \
LIQUIBASE_PASSWORD="${TEST_PASSWORD}" \
MAVEN_BIN="${RUNTIME}/maven/bin/mvn" \
MAVEN_LOCAL_REPOSITORY="${RUNTIME}/m2" \
JAVA_RESOURCES_DIR="${REPOSITORY_ROOT}/main/manager-api/src/main/resources" \
"${SCRIPT_DIR}/run-migrations.sh"
}
migrate() {
start
migrate_one "${JAVA_DATABASE}"
migrate_one "${FASTAPI_DATABASE}"
}
print_env() {
cat <<EOF
export TEST_MYSQL_PORT='${MYSQL_PORT}'
export TEST_REDIS_PORT='${REDIS_PORT}'
export TEST_JAVA_DATABASE_URL='mysql+asyncmy://${TEST_USER}:${TEST_PASSWORD}@127.0.0.1:${MYSQL_PORT}/${JAVA_DATABASE}?charset=utf8mb4'
export TEST_FASTAPI_DATABASE_URL='mysql+asyncmy://${TEST_USER}:${TEST_PASSWORD}@127.0.0.1:${MYSQL_PORT}/${FASTAPI_DATABASE}?charset=utf8mb4'
export TEST_JAVA_JDBC_URL='jdbc:mysql://127.0.0.1:${MYSQL_PORT}/${JAVA_DATABASE}?useUnicode=true&characterEncoding=UTF-8&serverTimezone=Asia/Shanghai&allowMultiQueries=true'
export TEST_JAVA_REDIS_URL='redis://127.0.0.1:${REDIS_PORT}/1'
export TEST_FASTAPI_REDIS_URL='redis://127.0.0.1:${REDIS_PORT}/2'
EOF
}
case "${1:-}" in
start) start ;;
stop) stop ;;
reset) reset ;;
migrate) migrate ;;
env) print_env ;;
*) echo "usage: $0 {start|stop|reset|migrate|env}" >&2; exit 2 ;;
esac
@@ -0,0 +1,640 @@
#!/usr/bin/env python3
"""Render the auditable Java/FastAPI compatibility matrix from checked-in inventories."""
from __future__ import annotations
import argparse
import json
import re
from collections import Counter
from pathlib import Path
from typing import Any
TARGET_ROOT = Path(__file__).resolve().parents[1]
REPOSITORY_ROOT = TARGET_ROOT.parents[1]
JAVA_ROOT = REPOSITORY_ROOT / "main" / "manager-api"
JAVA_SOURCE_ROOT = JAVA_ROOT / "src" / "main" / "java"
JAVA_RESOURCE_ROOT = JAVA_ROOT / "src" / "main" / "resources"
JAVA_MANIFEST = TARGET_ROOT / "compatibility" / "java-routes.json"
CONSUMER_MANIFEST = TARGET_ROOT / "compatibility" / "consumer-routes.json"
CONTRACT_RESULTS = TARGET_ROOT / "compatibility" / "contract-results.json"
ROUTE_SURFACE_RESULTS = TARGET_ROOT / "compatibility" / "route-surface-results.json"
AUTHENTICATED_ROUTE_RESULTS = TARGET_ROOT / "compatibility" / "authenticated-route-results.json"
PATH_PARAMETER = re.compile(r"\{[^/{}]+\}")
DIFFERENTIAL_CASES: dict[tuple[str, str], str] = {
("GET", "/user/pub-config"): "1",
("GET", "/user/info"): "9(七语言/过期 Token/Long",
("GET", "/admin/users"): "3(权限/序列化/非法分页)",
("GET", "/agent/list"): "1",
("GET", "/device/bind/{}"): "1",
("GET", "/models/provider"): "1",
("GET", "/correct-word/file/list"): "1",
("GET", "/correct-word/file/download/{}"): "2(二进制/更新后下载)",
("PUT", "/device/update/{}"): "3(上下界/UTF-16 长度)",
("POST", "/models/provider"): "1(约束集合)",
("POST", "/config/server-base"): "3(缺失/错误/正确 secret",
("GET", "/ota/"): "1MIME/body",
("POST", "/ota/"): "4(必填/格式/凭证/密码学)",
("POST", "/ota/activate"): "3",
("POST", "/device/tools/list/{}"): "2(响应/外呼格式)",
("POST", "/correct-word/file"): "2(响应/DB",
("PUT", "/correct-word/file/{}"): "2(响应/DB",
("DELETE", "/correct-word/file/{}"): "1(级联副作用)",
("POST", "/otaMag/upload"): "2(上传/扩展名错误)",
("POST", "/otaMag"): "2(响应/DB",
("GET", "/otaMag/download/{}"): "4(次数限制及二进制)",
}
DOMAIN_TESTS = {
"AdminController": "sys",
"SysParamsController": "sys",
"SysDictDataController": "sys",
"SysDictTypeController": "sys",
"ServerSideManageController": "sys",
"LoginController": "security",
"ConfigController": "config",
"AgentController": "agent",
"AgentChatHistoryController": "agent",
"AgentMcpAccessPointController": "agent",
"AgentSnapshotController": "agent",
"AgentTemplateController": "agent",
"AgentVoicePrintController": "agent",
"CorrectWordController": "correctword",
"DeviceController": "device",
"KnowledgeBaseController": "knowledge",
"KnowledgeFilesController": "knowledge",
"ModelController": "model",
"ModelProviderController": "model",
"OTAController": "device",
"OTAMagController": "device",
"TimbreController": "timbre",
"VoiceCloneController": "voiceclone",
"VoiceResourceController": "voiceclone",
}
def _canonical(path: str) -> str:
return PATH_PARAMETER.sub("{}", path)
def _declaration(route: dict[str, Any]) -> str:
source_path = REPOSITORY_ROOT / route["source"]
lines = source_path.read_text(encoding="utf-8").splitlines()
excerpt = "\n".join(lines[int(route["line"]) - 1 : int(route["line"]) + 60])
match = re.search(r"\bpublic\s+", excerpt)
if match is None:
raise ValueError(f"public declaration not found for {route['controller']}.{route['handler']}")
declaration = excerpt[match.start() :]
depth = 0
saw_parenthesis = False
for offset, character in enumerate(declaration):
if character == "(":
depth += 1
saw_parenthesis = True
elif character == ")":
depth -= 1
elif character == "{" and saw_parenthesis and depth == 0:
return " ".join(declaration[:offset].split())
raise ValueError(f"unterminated declaration for {route['controller']}.{route['handler']}")
def _parameter_text(declaration: str, handler: str) -> str:
marker = re.search(rf"\b{re.escape(handler)}\s*\(", declaration)
if marker is None:
return ""
start = marker.end() - 1
depth = 0
for offset in range(start, len(declaration)):
character = declaration[offset]
if character == "(":
depth += 1
elif character == ")":
depth -= 1
if depth == 0:
return declaration[start + 1 : offset]
return ""
def _split_parameters(value: str) -> list[str]:
result: list[str] = []
start = 0
round_depth = 0
angle_depth = 0
in_quote: str | None = None
escaped = False
for offset, character in enumerate(value):
if escaped:
escaped = False
continue
if character == "\\":
escaped = True
continue
if in_quote is not None:
if character == in_quote:
in_quote = None
continue
if character in {'"', "'"}:
in_quote = character
elif character == "(":
round_depth += 1
elif character == ")":
round_depth -= 1
elif character == "<":
angle_depth += 1
elif character == ">":
angle_depth -= 1
elif character == "," and round_depth == 0 and angle_depth == 0:
result.append(value[start:offset].strip())
start = offset + 1
tail = value[start:].strip()
if tail:
result.append(tail)
return result
def _without_annotations(value: str) -> str:
output: list[str] = []
offset = 0
while offset < len(value):
if value[offset] != "@":
output.append(value[offset])
offset += 1
continue
offset += 1
while offset < len(value) and (value[offset].isalnum() or value[offset] in "._$"):
offset += 1
while offset < len(value) and value[offset].isspace():
offset += 1
if offset < len(value) and value[offset] == "(":
depth = 1
offset += 1
in_quote: str | None = None
while offset < len(value) and depth:
character = value[offset]
if in_quote is not None:
if character == in_quote and value[offset - 1] != "\\":
in_quote = None
elif character in {'"', "'"}:
in_quote = character
elif character == "(":
depth += 1
elif character == ")":
depth -= 1
offset += 1
while offset < len(value) and value[offset].isspace():
offset += 1
return " ".join("".join(output).split())
def _type_and_name(parameter: str) -> tuple[str, str]:
cleaned = _without_annotations(parameter).removeprefix("final ").strip()
pieces = cleaned.rsplit(" ", 1)
if len(pieces) != 2:
return cleaned, cleaned
return pieces[0], pieces[1]
def _request_surface(route: dict[str, Any], declaration: str) -> str:
path_names = re.findall(r"\{([^/{}]+)\}", route["path"])
headers: list[str] = []
queries: list[str] = []
bodies: list[str] = []
multipart: list[str] = []
for parameter in _split_parameters(_parameter_text(declaration, route["handler"])):
parameter_type, name = _type_and_name(parameter)
if "HttpServletResponse" in parameter_type or "HttpServletRequest" in parameter_type:
continue
if "@PathVariable" in parameter:
continue
if "@RequestHeader" in parameter:
quoted = re.search(r'@RequestHeader(?:\([^)]*)?["\']([^"\']+)["\']', parameter)
headers.append(quoted.group(1) if quoted else name)
elif "MultipartFile" in parameter_type:
multipart.append(name)
elif "@RequestBody" in parameter:
bodies.append(parameter_type)
elif "@RequestParam" in parameter or "@ParameterObject" in parameter:
queries.append(f"{name}:{parameter_type}" if parameter_type != name else name)
elif route["method"] == "GET" or route["controller"] == "ModelProviderController":
queries.append(f"{name}:{parameter_type}" if parameter_type != name else name)
parts: list[str] = []
if path_names:
parts.append("Path:" + ",".join(path_names))
if headers:
parts.append("Header:" + ",".join(headers))
if queries:
parts.append("Query:" + ",".join(queries))
if bodies:
parts.append("Body:" + ",".join(bodies))
if multipart:
parts.append("Multipart:" + ",".join(multipart))
return "; ".join(parts) if parts else ""
def _response_type(route: dict[str, Any], declaration: str) -> str:
path = route["path"]
if path == "/user/captcha":
return "image/gif 二进制"
if path == "/ota/" and route["method"] == "GET":
return "裸 text/plain"
if path.startswith("/ota/"):
return "裸 application/json"
if path in {
"/agent/play/{uuid}",
"/agent/chat-history/download/{uuid}/current",
"/agent/chat-history/download/{uuid}/previous",
"/correct-word/file/download/{fileId}",
"/otaMag/download/{uuid}",
"/voiceClone/play/{uuid}",
}:
return "流式/二进制 + 原下载 headers"
match = re.search(rf"public\s+(.+?)\s+{re.escape(route['handler'])}\s*\(", declaration)
return_type = match.group(1) if match else "unknown"
if return_type.startswith("Result<"):
return "envelope " + return_type.removeprefix("Result")
return return_type
def _permission(route: dict[str, Any]) -> str | None:
source_path = REPOSITORY_ROOT / route["source"]
lines = source_path.read_text(encoding="utf-8").splitlines()
excerpt = "\n".join(lines[int(route["line"]) - 1 : int(route["line"]) + 60])
declaration_offset = excerpt.find("public ")
decorators = excerpt if declaration_offset < 0 else excerpt[:declaration_offset]
match = re.search(r'@RequiresPermissions\(\s*"([^"]+)"', decorators)
return match.group(1) if match else None
def _side_effect(route: dict[str, Any]) -> str:
method = route["method"]
path = route["path"]
handler = route["handler"]
controller = route["controller"]
if path == "/user/captcha":
return "Redis-W(captcha TTL); GIF"
if path == "/user/login":
return "DB-R/W(token); Redis-R/DEL(captcha)"
if path == "/user/smsVerification":
return "Redis-R/W(TTL/频控); 外部-Aliyun SMS"
if path in {"/user/register", "/user/retrieve-password", "/user/change-password"}:
return "DB-W(user/token); Redis-R/DEL(SMS)"
if path in {"/user/info", "/user/pub-config"}:
return "DB-R; Redis-R/W(cache)"
if path.startswith("/admin/server/"):
return "DB/Redis-R(secret/WS); Redis-W(one-shot); 外部-WebSocket" if method == "POST" else "DB/Redis-R"
if path.startswith("/admin/params"):
if method == "GET":
return "DB-R"
if method == "PUT":
return "DB-W; Redis-W; 外部-配置端点探测(按 paramCode)"
return "DB-W; Redis-W/DEL"
if path.startswith("/admin/dict"):
return "DB-R; Redis-R/W(dict cache)" if method == "GET" else "DB-W; Redis-DEL(dict cache)"
if path == "/admin/device/all" or path == "/admin/users":
return "DB-R"
if path == "/admin/users/{id}" and method == "DELETE":
return "DB-W(用户/token/device/agent 级联)"
if path.startswith("/admin/users"):
return "DB-W(user/password/status/token)"
if path.startswith("/config/"):
return "DB-R; Redis-R/W(runtime/model/timbre cache)"
if path.startswith("/agent/mcp/tools"):
return "DB/Redis-R; 外部-WebSocket MCP"
if path.startswith("/agent/mcp/address"):
return "DB/Redis-R; AES token 生成"
if path.startswith("/agent/voice-print"):
return "DB-R" if method == "GET" else "DB-W; 外部-voiceprint HTTP"
if "/chat-summary/" in path or "/chat-title/" in path:
return "DB-R/W(chat); 外部-OpenAI-compatible LLM"
if path == "/agent/chat-history/report":
return "DB-W(chat/session); server-secret"
if path.startswith("/agent/chat-history/getDownloadUrl/"):
return "DB-R(chat/session); Redis-W(download token TTL)"
if "/chat-history/download/" in path or path.startswith("/agent/play/"):
return "DB/Redis-R(one-shot); 文件-R/流式"
if path.startswith("/agent/audio/"):
return "DB-R(audio); Redis-W(one-shot URL)"
if path.startswith("/agent/") or path == "/agent":
return "DB-R; Redis-R" if method == "GET" else "DB-W(含快照/映射/标签事务); Redis-DEL"
if path.startswith("/correct-word/"):
if "download" in path:
return "DB-R(content); 二进制"
return "DB-R" if method == "GET" else "DB-W(file/items/mapping 事务)"
if path.startswith("/datasets"):
if path == "/datasets/rag-models":
return "DB-R(model config)"
if method == "GET" and not path.endswith("/chunks"):
return "DB-R"
return "DB-R/W; 外部-RAGFlow HTTP(upload/dataset/document/chunk/retrieval)"
if path.startswith("/device/tools/") or (path == "/device/bind/{agentId}" and method == "POST"):
return "DB/Redis-R; 外部-MQTT gateway HTTP + daily auth"
if path in {"/device/address-book/call", "/device/address-book/lookup"}:
return "DB-R; 外部-MQTT gateway HTTP; server-secret"
if path.startswith("/device/"):
return "DB-R; Redis-R" if method == "GET" else "DB-W(device/bind/address-book); Redis-R/W"
if path == "/ota/":
if method == "GET":
return ""
return "DB/Redis-R(设备/固件/配置); HMAC/Base64/时间戳凭证"
if path == "/ota/activate":
return "DB-R/W(device activation); Redis-R/W(TTL)"
if path.startswith("/otaMag/upload"):
return "文件-W(MD5/扩展名/大小)"
if path.startswith("/otaMag/download"):
return "Redis-R/W(一次性/次数); 文件-R/流式"
if path.startswith("/otaMag/getDownloadUrl"):
return "DB-R; Redis-W(download token TTL)"
if path.startswith("/otaMag"):
if method == "GET":
return "DB-R"
if method == "DELETE":
return "DB-W(OTA metadata); 文件-DEL"
return "DB-W(OTA metadata)"
if path.startswith("/models"):
return "DB-R; Redis-R/W(model cache)" if method == "GET" else "DB-W; Redis-DEL(model/config cache)"
if path.startswith("/ttsVoice"):
return "DB-R; Redis-R/W(timbre cache)" if method == "GET" else "DB-W; Redis-DEL(timbre/config cache)"
if path.startswith("/voiceClone"):
if path.startswith("/voiceClone/play"):
return "Redis-R/DEL(one-shot); 文件/外部音频-R"
if handler == "getAudioId":
return "DB-R; Redis-W(one-shot URL)"
if handler == "updateName":
return "DB-W(train record name)"
if method == "GET":
return "DB-R"
return "DB-R/W(train state); 文件-W; 外部-火山语音克隆 HTTP"
if path.startswith("/voiceResource"):
return "DB-R" if method == "GET" else "DB-W(voice resource)"
operation = "DB-R" if method == "GET" else "DB-W"
return f"{operation} ({controller}.{handler})"
def _verification(route: dict[str, Any]) -> str:
domain = DOMAIN_TESTS.get(route["controller"])
domain_status = f"领域✓({domain},域级)" if domain else "领域—"
diff = DIFFERENTIAL_CASES.get((route["method"], _canonical(route["path"])))
diff_status = f"差分✓{diff}" if diff else "差分—"
if route["path"] == "/otaMag/getDownloadUrl/{id}":
diff_status = "差分间接✓(供下载链路)"
return f"结构✓;请求面差分✓1;认证业务面差分✓1;{domain_status}{diff_status}"
def _escape(value: str) -> str:
return value.replace("|", "\\|").replace("\n", " ")
def _inventory_section(java_routes: list[dict[str, Any]]) -> str:
controller_counts = Counter(route["controller"] for route in java_routes)
mapper_files = sorted((JAVA_RESOURCE_ROOT / "mapper").rglob("*.xml"))
changelog_root = JAVA_RESOURCE_ROOT / "db" / "changelog"
sql_files = sorted(changelog_root.rglob("*.sql"))
master = changelog_root / "db.changelog-master.yaml"
changeset_refs = master.read_text(encoding="utf-8").count("changeSet:")
entity_files = list(JAVA_SOURCE_ROOT.rglob("entity/*.java"))
dto_files = [path for path in JAVA_SOURCE_ROOT.rglob("*.java") if "dto" in path.parts]
vo_files = list(JAVA_SOURCE_ROOT.rglob("vo/*.java"))
dao_files = list(JAVA_SOURCE_ROOT.rglob("dao/*.java"))
service_files = list(JAVA_SOURCE_ROOT.rglob("service/**/*.java"))
service_impl_files = list(JAVA_SOURCE_ROOT.rglob("service/impl/*.java"))
assert len(controller_counts) == 24
assert len(mapper_files) == 20
assert len(sql_files) == 101
assert changeset_refs == 101
assert len(entity_files) == 29
assert len(dto_files) == 58
assert len(vo_files) == 14
assert len(dao_files) == 29
mapper_names = "".join(f"`{path.relative_to(JAVA_RESOURCE_ROOT).as_posix()}`" for path in mapper_files)
controller_names = "".join(
f"`{controller}`({controller_counts[controller]})" for controller in sorted(controller_counts)
)
return f"""## Java 基线静态盘点
- Controller24 个、154 条映射。按 Controller 的路由数为:{controller_names}
- 数据分层:`entity/` 29 个 Java 文件(28 个 `*Entity.java` 加 `BaseEntity`)、`dto/` 58 个、
`vo/` 14 个、`dao/` 29 个、`service/` 树 {len(service_files)} 个文件(其中
`service/impl/` {len(service_impl_files)} 个)。FastAPI 对应落在 `schemas/`、`repositories/`、
`services/`、`routers/`、`integrations/` 与 `jobs/`,没有把跨表事务放进路由。
- MyBatis XML20 个,分别是 {mapper_names}
- Liquibase`db.changelog-master.yaml` 含 {changeset_refs} 个 `changeSet` 引用,目录中恰有
{len(sql_files)} 个 SQLPython 部署继续执行这 101 个原始 SQL,不改写历史。
- 定时工作:`DocumentStatusSyncTask` 每次完成后延迟 30 秒,扫描 RAGFlow RUNNING 文档并
回写 SUCCESS/FAIL/CANCEL 与统计;当前 Java 源码另有 `AgentSnapshotRedactionRunner`,启动时
执行一次并在滚动部署期每 15 秒补偿脱敏旧快照。FastAPI 将工作移到独立 jobs 进程,并以
Redis 分布式锁/watchdog 防止多 worker 重复执行。
- 外部集成:RAGFlow dataset/document/chunk/retrieval/upload;阿里云短信;火山语音克隆训练与
音频;声纹 HTTPOpenAI-compatible LLM 摘要/标题;MQTT gateway HTTPMCP/管理动作
WebSocketOTA/WS/MQTT 的 HMAC、Base64、时间戳与下载文件存储。自动测试只访问可重复 mock,
未使用真实付费凭证。
"""
def _require_complete_route_report(
report: dict[str, Any],
routes: list[dict[str, Any]],
*,
request_profile: str,
side_effect_policy: str,
) -> None:
expected_summary = {"total": 154, "passed": 154, "failed": 0, "skipped": 0}
expected_names = [f"{route['method']} {route['path']}" for route in routes]
assert report["summary"] == expected_summary
assert report["coverage"] == {
"java_routes": 154,
"request_profile": request_profile,
"side_effect_policy": side_effect_policy,
}
assert [result["name"] for result in report["results"]] == expected_names
assert all(result["passed"] is True and result["difference"] is None for result in report["results"])
def render() -> str:
java_manifest = json.loads(JAVA_MANIFEST.read_text(encoding="utf-8"))
consumer_manifest = json.loads(CONSUMER_MANIFEST.read_text(encoding="utf-8"))
contract = json.loads(CONTRACT_RESULTS.read_text(encoding="utf-8"))
route_surface = json.loads(ROUTE_SURFACE_RESULTS.read_text(encoding="utf-8"))
authenticated_route = json.loads(AUTHENTICATED_ROUTE_RESULTS.read_text(encoding="utf-8"))
routes: list[dict[str, Any]] = java_manifest["routes"]
assert len(routes) == 154
assert contract["summary"] == {"total": 49, "passed": 49, "failed": 0, "skipped": 0}
_require_complete_route_report(
route_surface,
routes,
request_profile="missing-auth-or-safe-invalid-input",
side_effect_policy="no successful write request is issued",
)
_require_complete_route_report(
authenticated_route,
routes,
request_profile="authenticated-safe-business-or-validation",
side_effect_policy="no intentional successful writes",
)
assert len(DIFFERENTIAL_CASES) == 21
lines = [
"# manager-api FastAPI 兼容性矩阵",
"",
"> 生成依据:`main/manager-api-fastapi/compatibility/java-routes.json`、",
"> `main/manager-api-fastapi/compatibility/consumer-routes.json`、`route-surface-results.json`、",
"> `authenticated-route-results.json`、`contract-results.json` 和当前 Java 源码。接口路径均省略",
"> 共同前缀 `/xiaozhi`。",
"",
"## 结论与状态口径",
"",
"Java 基线共有 **154** 条 Spring MVC 路由;FastAPI 已注册 **154/154100%**,并由",
"`tests/test_java_route_manifest.py` 对源码清单 freshness、数量和 method/path 注册闭合进行检查。",
"此外实现 3 条仅由仓库消费者使用、Java Controller 中不存在的兼容路由,因此这 3 条不计入",
"154 条 Java 覆盖率。三端 188 个调用点均能解析到 FastAPI 路由。",
"",
"矩阵状态必须按下列含义阅读:",
"",
"- `结构✓`method/path 已注册且清单闭合;它不等于业务行为逐接口实测。",
"- `请求面差分✓1`:本行已向隔离 Java/FastAPI 各发送一次缺少鉴权或安全非法输入,精确比较",
" HTTP status、body 与 Content-Type;最终为 **154/154 通过、0 失败、0 跳过**,且不发送成功写请求。",
"- `认证业务面差分✓1`:本行已使用有效 DB Token、server-secret 或匿名身份,再向隔离",
" Java/FastAPI 各发送一次安全业务/校验请求,精确比较 HTTP status、body 与 Content-Type",
" 最终为 **154/154 通过、0 失败、0 跳过**,且不主动发送成功写请求。该状态不等于每条路由的",
" 完整成功生命周期均已差分,完整副作用证据仍以 `差分✓N` 为准。",
"- `领域✓(x,域级)`:该领域有 service/repository/协议自动测试,但不保证本行每条成功与错误路径",
" 都被直接请求。`领域—` 表示除结构测试外没有可归属的域级直接测试证据。",
"- `差分✓N`:本行除安全请求面外,还参与了成功、主要错误、协议或数据库副作用的深度对照;",
" 括号说明覆盖面。深度结果为 **49/49 checks 通过、0 失败、0 跳过**,覆盖 **21/154** 条路由。",
" `差分间接✓` 表示 J125 作为下载链路的 URL 生成步骤被间接覆盖;`差分—` 表示没有深度对照,",
" 不能把 154/154 请求面差分误读成 154 条全部成功路径与副作用都已逐接口对照。",
"- 所有 `Result<T>` 均表示 `{code,msg,data}` envelope;原 Java 为 HTTP 200 的认证、权限、业务和",
" 参数错误由全局兼容层维持 HTTP 200。二进制/OTA 裸响应在“响应类型”列单独标明。",
"",
"## 三端消费者闭合",
"",
"| 消费者 | 调用点 | 唯一结构路由 | 方法分布 |",
"|---|---:|---:|---|",
]
for consumer in ("manager-web", "manager-mobile", "xiaozhi-server"):
item = consumer_manifest["consumers"][consumer]
methods = "".join(f"{method} {count}" for method, count in item["methods"].items())
lines.append(f"| `{consumer}` | {item['callSites']} | {item['uniqueRoutes']} | {methods} |")
lines.extend(
[
f"| **合计** | **{consumer_manifest['count']}** | **{consumer_manifest['uniqueRoutes']}** | — |",
"",
"### 3 条消费者孤儿兼容路由",
"",
"| Method/path | 来源 | FastAPI 语义 | 鉴权 | 状态 |",
"|---|---|---|---|---|",
(
"| `GET /api/ping` | manager-mobile 环境设置探活 | "
"`{code:0,msg:\"success\",data:\"pong\"}` | 匿名 | 实现✓;consumer resolve✓ |"
),
(
"| `PUT /user/configDevice/{device_id}` | manager-web 遗留设备配置调用 | "
"按现有设备更新契约处理 body | DB Token | 实现✓;consumer resolve✓ |"
),
(
"| `GET /device/address-book/lookup` | xiaozhi-server 管理客户端 | "
"`callerMac/nickname/answer` 地址簿查询/呼叫兼容别名 | server-secret | "
"实现✓;consumer resolve✓;device 域测试✓ |"
),
"",
"`GET /admin/dict/data/type/FIRMWARE_TYPE` 是动态 Java 路由",
"`GET /admin/dict/data/type/{dictType}` 的一个字面调用,不是第四条孤儿路由。",
"",
_inventory_section(routes).rstrip(),
"",
"## 154 条 Java→FastAPI 逐接口矩阵",
"",
"副作用缩写:`DB-R/W`=数据库读/写,`Redis-R/W/DEL`=缓存读/写/失效,`文件-R/W`=文件",
"读取/写入;外部调用均在 service/integration 层。权限为空时表示只需对应鉴权身份。",
"",
(
"| # | Method/path | Java Controller.handler | 请求面 | 响应类型 | 鉴权 / 权限 | "
"DB/Redis/文件/外部副作用 | 实现与测试状态 |"
),
"|---:|---|---|---|---|---|---|---|",
]
)
for index, route in enumerate(routes, start=1):
declaration = _declaration(route)
permission = _permission(route) or ""
auth = {
"anonymous": "匿名",
"database-token": "DB Token",
"server-secret": "server-secret",
}[route["auth"]]
values = [
f"J{index:03d}",
f"`{route['method']} {route['path']}`",
f"`{route['controller']}.{route['handler']}`",
_request_surface(route, declaration),
_response_type(route, declaration),
f"{auth} / `{permission}`" if permission != "" else f"{auth} / —",
_side_effect(route),
_verification(route),
]
lines.append("| " + " | ".join(_escape(value) for value in values) + " |")
lines.extend(
[
"",
"## 已观测差异与未覆盖面",
"",
"- 154 条安全请求面差分最终全部一致。首轮曾发现 5 个空 Body 映射差异;修复 FastAPI 对",
" Spring `HttpMessageNotReadableException` 的 code-500 语义后,重新从零执行才得到 154/154。",
"- 154 条认证业务面差分最终全部一致;该轮使用有效鉴权与安全业务/校验输入,在不主动成功",
" 写入的前提下逐路由对照。证据是 `authenticated-route-results.json`,渲染器会在结果不是",
" 154/154、存在失败或跳过时硬失败。",
"- 2026-07-20 的隔离差分报告未在 49 个 checks 中观测到响应/所选 headers/数据库副作用",
" 不一致;证据是 `main/manager-api-fastapi/compatibility/contract-results.json`,不是人工推断。",
"- Hibernate Validator 的 `ConstraintViolation Set` 首条消息无稳定顺序;模型 provider 必填",
" 用例比较“消息属于 Java 声明约束集合”与相同错误码,而不伪造一个固定顺序。",
"- OTA 时间戳/token 是动态值,差分先比较归一化结构,再分别校验两端 HMAC/Base64 密码学",
" 有效性;这属于有意的测试归一化,不是声称字节恒等。",
"- 深度差分未直接命中的 133 条中,J125 是下载链路间接覆盖,另 132 条标为 `差分—`;",
" 它们有请求面、认证业务面与所属领域测试,但尚无逐路由成功+主要错误+副作用深度对照,不能据此宣称",
" 每一种业务状态均已逐接口行为等价。",
"- FastAPI 额外提供上述 3 条消费者兼容路由与 live/ready 健康检查;它们没有 Java",
" Controller 基线,属于明确、可回退的加法差异。",
"- Java 把定时任务放在 Spring 进程;FastAPI 使用独立 jobs 进程和 Redis 分布式锁。这是",
" 部署拓扑差异,业务状态和幂等目标保持一致。",
"- RAGFlow、阿里云短信、火山语音克隆、真实声纹、真实 LLM、真实 MQTT/MCP/WS 均未用",
" 生产凭证联调;自动化只证明 mock 请求格式、超时/错误映射/重试中的已覆盖场景。",
"",
"## 可复现检查",
"",
"```bash",
"cd main/manager-api-fastapi",
".venv/bin/python scripts/extract_java_routes.py --output compatibility/java-routes.json",
".venv/bin/python scripts/extract_consumer_routes.py > /tmp/consumer-routes.json",
(
".venv/bin/pytest -q tests/test_java_route_manifest.py "
"tests/test_consumer_route_manifest.py tests/test_compatibility_document.py"
),
"```",
"",
"逐接口差分的启动、隔离库、mock 与执行命令见 `docs/manager-api-fastapi-test-report.md`",
"本文件只陈述已落盘的结果,不把缺少真实密钥的外部联调列为通过。",
]
)
rendered = "\n".join(lines) + "\n"
if len(re.findall(r"^\| J\d{3} \|", rendered, flags=re.MULTILINE)) != 154:
raise AssertionError("rendered route matrix is not closed")
return rendered
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--output", type=Path)
args = parser.parse_args()
rendered = render()
if args.output is None:
print(rendered, end="")
else:
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(rendered, encoding="utf-8")
return 0
if __name__ == "__main__":
raise SystemExit(main())
@@ -0,0 +1,220 @@
#!/bin/sh
set -eu
SCRIPT_DIR=$(CDPATH= cd -- "$(dirname -- "$0")" && pwd)
TARGET_DIR=$(CDPATH= cd -- "${SCRIPT_DIR}/.." && pwd)
REPOSITORY_ROOT=$(CDPATH= cd -- "${TARGET_DIR}/../.." && pwd)
RUNTIME="${REPOSITORY_ROOT}/.runtime"
JAVA_DIR="${REPOSITORY_ROOT}/main/manager-api"
STATE_DIR="${TARGET_DIR}/.test-runtime/contract"
JAVA_PORT="${CONTRACT_JAVA_PORT:-18082}"
FASTAPI_PORT="${CONTRACT_FASTAPI_PORT:-18083}"
MOCK_PORT="${CONTRACT_MOCK_PORT:-18084}"
JAVA_PID=""
FASTAPI_PID=""
MOCK_PID=""
mkdir -p "${STATE_DIR}"
stop_process() {
pid=$1
if [ -n "${pid}" ] && kill -0 "${pid}" 2>/dev/null; then
children=$(pgrep -P "${pid}" 2>/dev/null || true)
if [ -n "${children}" ]; then
kill ${children} 2>/dev/null || true
fi
kill "${pid}" 2>/dev/null || true
wait "${pid}" 2>/dev/null || true
fi
}
cleanup() {
stop_process "${FASTAPI_PID}"
stop_process "${JAVA_PID}"
stop_process "${MOCK_PID}"
}
trap cleanup EXIT INT TERM
wait_for_url() {
name=$1
url=$2
attempts=0
until curl --fail --silent --show-error "${url}" >/dev/null 2>&1; do
attempts=$((attempts + 1))
if [ "${attempts}" -ge 240 ]; then
echo "Timed out waiting for ${name}; inspect ${STATE_DIR}." >&2
return 1
fi
sleep 0.25
done
}
assert_clean_runtime_logs() {
for log in "${STATE_DIR}/fastapi.log" "${STATE_DIR}/external-mock.log"; do
if rg -n -i 'UnsupportedFieldAttributeWarning|Traceback|\bERROR\b|\bWARNING\b' "${log}"; then
echo "Unexpected warning or error in ${log}." >&2
return 1
fi
done
# Logback's own bootstrap diagnostics contain tokens such as ERROR_FILE.
# Runtime application records start with the configured full ISO date.
java_errors=$(rg -n '^[0-9]{4}-[0-9]{2}-[0-9]{2} .*ERROR.* - ' "${STATE_DIR}/java.log" || true)
if [ -n "${java_errors}" ]; then
java_error_count=$(printf '%s\n' "${java_errors}" | wc -l | tr -d ' ')
missing_body_errors=$(printf '%s\n' "${java_errors}" | rg -F 'Required request body is missing:' | wc -l | tr -d ' ')
missing_device_errors=$(printf '%s\n' "${java_errors}" | rg -F "Required request header 'Device-Id'" | wc -l | tr -d ' ')
object_array_errors=$(printf '%s\n' "${java_errors}" | rg -F 'JSON parse error: Cannot deserialize value of type' | wc -l | tr -d ' ')
missing_query_errors=$(printf '%s\n' "${java_errors}" | rg -e "Required request parameter '(ids|modelType|id)' for method parameter type String is not present" | wc -l | tr -d ' ')
multipart_errors=$(printf '%s\n' "${java_errors}" | rg -F 'Current request is not a multipart request' | wc -l | tr -d ' ')
caller_mac_errors=$(printf '%s\n' "${java_errors}" | rg -F 'Cannot invoke "String.toLowerCase()" because "callerMac" is null' | wc -l | tr -d ' ')
null_message_errors=$(printf '%s\n' "${java_errors}" | rg -e ' - null$' | wc -l | tr -d ' ')
blank_message_errors=$(printf '%s\n' "${java_errors}" | rg -e ' - $' | wc -l | tr -d ' ')
if [ "${java_error_count}" -ne 36 ] \
|| [ "${missing_body_errors}" -ne 13 ] \
|| [ "${missing_device_errors}" -ne 5 ] \
|| [ "${object_array_errors}" -ne 8 ] \
|| [ "${missing_query_errors}" -ne 4 ] \
|| [ "${multipart_errors}" -ne 3 ] \
|| [ "${caller_mac_errors}" -ne 1 ] \
|| [ "${null_message_errors}" -ne 1 ] \
|| [ "${blank_message_errors}" -ne 1 ]; then
printf '%s\n' "${java_errors}" >&2
echo "The two 154-route safe surfaces did not produce the exact expected Java baseline error profile." >&2
return 1
fi
unexpected_java_errors=$(
printf '%s\n' "${java_errors}" |
rg -v -e "Required request header 'Device-Id' for method parameter type String is not present" \
-e 'Required request body is missing:' \
-e 'JSON parse error: Cannot deserialize value of type .* from Object value \(token .*START_OBJECT.*\)' \
-e "Required request parameter '(ids|modelType|id)' for method parameter type String is not present" \
-e 'Current request is not a multipart request' \
-e 'Cannot invoke "String.toLowerCase\(\)" because "callerMac" is null' \
-e ' - null$' \
-e ' - $' || true
)
if [ -n "${unexpected_java_errors}" ]; then
printf '%s\n' "${unexpected_java_errors}" >&2
echo "Unexpected Java baseline error; only the counted safe surface-validation paths are allowed." >&2
return 1
fi
fi
}
start_java() {
(
cd "${JAVA_DIR}"
JAVA_HOME="${RUNTIME}/jdk" \
PATH="${RUNTIME}/jdk/bin:${RUNTIME}/maven/bin:${PATH}" \
exec "${RUNTIME}/maven/bin/mvn" \
-Dmaven.repo.local="${RUNTIME}/m2" \
-DskipTests spring-boot:run \
-Dspring-boot.run.arguments="--server.port=${JAVA_PORT} \
--spring.datasource.druid.url=jdbc:mysql://127.0.0.1:${TEST_MYSQL_PORT}/manager_java_test?useUnicode=true&characterEncoding=UTF-8&serverTimezone=Asia/Shanghai&nullCatalogMeansCurrent=true&allowMultiQueries=true \
--spring.datasource.druid.username=xiaozhi_test \
--spring.datasource.druid.password=isolated-test-only \
--spring.data.redis.host=127.0.0.1 \
--spring.data.redis.port=${TEST_REDIS_PORT} \
--spring.data.redis.database=1 \
--spring.data.redis.password="
) >"${STATE_DIR}/java.log" 2>&1 &
JAVA_PID=$!
wait_for_url "Java baseline" "http://127.0.0.1:${JAVA_PORT}/xiaozhi/ota/"
}
start_mock() {
(
cd "${TARGET_DIR}"
exec .venv/bin/uvicorn tests.compatibility.external_mock:app \
--host 127.0.0.1 --port "${MOCK_PORT}" --log-level warning
) >"${STATE_DIR}/external-mock.log" 2>&1 &
MOCK_PID=$!
wait_for_url "external-service mock" "http://127.0.0.1:${MOCK_PORT}/health"
}
start_fastapi() {
(
cd "${TARGET_DIR}"
APP_ENVIRONMENT=test \
APP_DATABASE_URL="${TEST_FASTAPI_DATABASE_URL}" \
APP_REDIS_URL="${TEST_FASTAPI_REDIS_URL}" \
APP_SERVER_SECRET_OVERRIDE=contract-server-secret \
APP_UPLOAD_DIR="${STATE_DIR}/uploads" \
APP_LOG_LEVEL=WARNING \
exec .venv/bin/uvicorn app.main:app \
--host 127.0.0.1 --port "${FASTAPI_PORT}" --log-level warning
) >"${STATE_DIR}/fastapi.log" 2>&1 &
FASTAPI_PID=$!
wait_for_url "FastAPI target" "http://127.0.0.1:${FASTAPI_PORT}/xiaozhi/health/ready"
}
cd "${TARGET_DIR}"
./scripts/isolated-env.sh reset
./scripts/isolated-env.sh migrate
eval "$(./scripts/isolated-env.sh env)"
start_mock
# The retained Java startup creates the SM2 key pair on an empty schema.
start_java
.venv/bin/python -m tests.compatibility.seed_contract_data \
--mysql-port "${TEST_MYSQL_PORT}" --mock-port "${MOCK_PORT}"
# Reload Java after the deterministic fixture replaced server params. FLUSHALL
# is safe here because this is the dedicated Redis on TEST_REDIS_PORT and the
# FastAPI target has not started yet.
stop_process "${JAVA_PID}"
JAVA_PID=""
"${RUNTIME}/redis/bin/redis-cli" -h 127.0.0.1 -p "${TEST_REDIS_PORT}" FLUSHALL >/dev/null
start_java
start_fastapi
# Restore fixed dates after Java startup and allow no stale async baseline write
# to leak into the first read-only comparison.
.venv/bin/python -m tests.compatibility.seed_contract_data \
--mysql-port "${TEST_MYSQL_PORT}" --mock-port "${MOCK_PORT}"
sleep 1
.venv/bin/python -m tests.compatibility.seed_contract_data \
--mysql-port "${TEST_MYSQL_PORT}" --mock-port "${MOCK_PORT}"
TEST_FASTAPI_DATABASE_URL="${TEST_FASTAPI_DATABASE_URL}" \
TEST_FASTAPI_REDIS_URL="${TEST_FASTAPI_REDIS_URL}" \
APP_DATABASE_URL="${TEST_FASTAPI_DATABASE_URL}" \
APP_REDIS_URL="${TEST_FASTAPI_REDIS_URL}" \
APP_ENVIRONMENT=test \
.venv/bin/pytest -q tests/integration/test_isolated_runtime.py
.venv/bin/python -m tests.compatibility.route_surface_runner \
--java-base "http://127.0.0.1:${JAVA_PORT}/xiaozhi" \
--fastapi-base "http://127.0.0.1:${FASTAPI_PORT}/xiaozhi" \
--mock-base "http://127.0.0.1:${MOCK_PORT}" \
--mysql-port "${TEST_MYSQL_PORT}" \
--output compatibility/route-surface-results.json
.venv/bin/python -m tests.compatibility.authenticated_route_runner \
--java-base "http://127.0.0.1:${JAVA_PORT}/xiaozhi" \
--fastapi-base "http://127.0.0.1:${FASTAPI_PORT}/xiaozhi" \
--mock-base "http://127.0.0.1:${MOCK_PORT}" \
--mysql-port "${TEST_MYSQL_PORT}" \
--output compatibility/authenticated-route-results.json
.venv/bin/python -m tests.compatibility.differential_runner \
--java-base "http://127.0.0.1:${JAVA_PORT}/xiaozhi" \
--fastapi-base "http://127.0.0.1:${FASTAPI_PORT}/xiaozhi" \
--mock-base "http://127.0.0.1:${MOCK_PORT}" \
--mysql-port "${TEST_MYSQL_PORT}" \
--output compatibility/contract-results.json
.venv/bin/python -m tests.compatibility.seed_contract_data \
--mysql-port "${TEST_MYSQL_PORT}" --mock-port "${MOCK_PORT}"
.venv/bin/python -m tests.compatibility.performance_runner \
--java-base "http://127.0.0.1:${JAVA_PORT}/xiaozhi" \
--fastapi-base "http://127.0.0.1:${FASTAPI_PORT}/xiaozhi" \
--requests 60 --concurrency 6 --warmup 10 \
--output compatibility/performance-results.json
assert_clean_runtime_logs
echo "Isolated integration, 154-route unauthenticated and authenticated surfaces, deep differential, and performance tests passed."
+51
View File
@@ -0,0 +1,51 @@
#!/bin/sh
set -eu
: "${LIQUIBASE_URL:?Set LIQUIBASE_URL to the isolated or deployment JDBC URL}"
: "${LIQUIBASE_USERNAME:?Set LIQUIBASE_USERNAME}"
: "${LIQUIBASE_PASSWORD:?Set LIQUIBASE_PASSWORD}"
JDBC_URL=${LIQUIBASE_URL}
JDBC_USERNAME=${LIQUIBASE_USERNAME}
JDBC_PASSWORD=${LIQUIBASE_PASSWORD}
unset LIQUIBASE_URL LIQUIBASE_USERNAME LIQUIBASE_PASSWORD
export MIGRATION_JDBC_URL=${JDBC_URL}
export MIGRATION_USERNAME=${JDBC_USERNAME}
export MIGRATION_PASSWORD=${JDBC_PASSWORD}
SCRIPT_DIR=$(CDPATH= cd -- "$(dirname -- "$0")" && pwd)
PROJECT_DIR=$(CDPATH= cd -- "${SCRIPT_DIR}/.." && pwd)
POM="${MIGRATION_POM:-${PROJECT_DIR}/migration-pom.xml}"
JAVA_RESOURCES="${JAVA_RESOURCES_DIR:-${PROJECT_DIR}/../manager-api/src/main/resources}"
if [ -n "${MIGRATION_RUNNER_JAR:-}" ]; then
exec java -jar "${MIGRATION_RUNNER_JAR}"
fi
RUNTIME_ROOT="${PROJECT_DIR}/../../.runtime"
if [ -n "${MAVEN_BIN:-}" ]; then
MAVEN="${MAVEN_BIN}"
elif [ -x "${RUNTIME_ROOT}/maven/bin/mvn" ]; then
MAVEN="${RUNTIME_ROOT}/maven/bin/mvn"
else
MAVEN="mvn"
fi
MAVEN_REPOSITORY_ARGS=""
if [ -n "${MAVEN_LOCAL_REPOSITORY:-}" ]; then
MAVEN_REPOSITORY_ARGS="-Dmaven.repo.local=${MAVEN_LOCAL_REPOSITORY}"
elif [ -x "${RUNTIME_ROOT}/maven/bin/mvn" ]; then
MAVEN_REPOSITORY_ARGS="-Dmaven.repo.local=${RUNTIME_ROOT}/m2"
fi
"${MAVEN}" -B -f "${POM}" ${MAVEN_REPOSITORY_ARGS} \
-Djava.resources.dir="${JAVA_RESOURCES}" package
if [ -n "${JAVA_BIN:-}" ]; then
JAVA="${JAVA_BIN}"
elif [ -n "${JAVA_HOME:-}" ] && [ -x "${JAVA_HOME}/bin/java" ]; then
JAVA="${JAVA_HOME}/bin/java"
elif [ -x "${RUNTIME_ROOT}/jdk/bin/java" ]; then
JAVA="${RUNTIME_ROOT}/jdk/bin/java"
else
JAVA="java"
fi
exec "${JAVA}" -jar "${PROJECT_DIR}/target/manager-api-liquibase-runner-1.0.0-all.jar"

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