feat(ai): add tenant-safe hierarchical expense learning

This commit is contained in:
caoxiaozhu
2026-07-16 14:30:41 +08:00
parent 6bdf65bc24
commit ee88a36baf
65 changed files with 6909 additions and 232 deletions

View File

@@ -173,6 +173,11 @@
- 首个切片使用关系库精确键检索,只允许低敏枚举 `travel_application.transport_mode`;不复用当前缺少完整租户键的 FewShot/Qdrant 或员工风险画像存储个人记忆。 - 首个切片使用关系库精确键检索,只允许低敏枚举 `travel_application.transport_mode`;不复用当前缺少完整租户键的 FewShot/Qdrant 或员工风险画像存储个人记忆。
- 记忆证据只从认证动作事务中的服务端核验纠正产生,关联 Decision、Feedback、Outcome 和 Expense Case草稿、客户端观察、重复请求和同一 Case 多次修改不能重复计票。 - 记忆证据只从认证动作事务中的服务端核验纠正产生,关联 Decision、Feedback、Outcome 和 Expense Case草稿、客户端观察、重复请求和同一 Case 多次修改不能重复计票。
- 个人记忆只填充当前申请的空白软字段;当前明确输入、服务端 HR/组织主数据、企业制度和规则计算结果都优先于个人记忆。每次应用或抑制都返回可解释原因,不把优先级交给 Prompt 或 LLM 自行判断。 - 个人记忆只填充当前申请的空白软字段;当前明确输入、服务端 HR/组织主数据、企业制度和规则计算结果都优先于个人记忆。每次应用或抑制都返回可解释原因,不把优先级交给 Prompt 或 LLM 自行判断。
- 分层记忆首个组织级切片继续只允许 `travel_application.transport_mode` 的“飞机/火车/轮船”。企业记忆由租户管理员显式确认,部门记忆必须绑定稳定的 `organization_unit_id` 并只把部门名称作为展示标签;隐式证据暂只生成个人 Candidate不允许少量个人行为自动升级为部门或企业制度。
- 分层解析顺序固定为“当前显式输入/规则禁用 > 企业记忆 > 部门记忆 > 个人记忆”。不同层值冲突时只应用最高优先级可用项,并返回候选层级、冲突值、未采用原因、有效期和衰减后置信度;不得静默覆盖当前输入,也不得把优先级交给模型自由判断。
- 管理员显式确认的组织记忆可直接激活,但必须记录创建人、确认来源、适用作用域和有效期;隐式个人记忆仍使用 3 个不同 Case、2 次审批通过、证据跨 7 天的最小样本门槛。组织记忆过期、撤销或服务异常时降级到下一层,不阻塞申请预览。
- 已确认历史案例接入报销预审前,`few_shot_samples` 与 Qdrant payload 必须补齐租户、业务场景、制度/规则标识、规则版本和 active 状态;每次向量命中还必须回关系库按租户和样本状态二次校验,样本改判时删除或稳定覆盖旧向量,禁止跨租户或陈旧结论进入提示。
- 历史案例只能输出结构化 `historical_case_evidence`,明确标记 `advisory_only`、人工确认结论、相似度、制度和版本是否匹配;它不得修改确定性预审的 `decision``passed``blocking_count`不得改变预算复核和审批路由。旧版本只可作为降权参考Qdrant 或 Embedding 不可用时返回空证据并继续规则链。
#### 节省与价值 #### 节省与价值
@@ -317,6 +322,7 @@
- LLM 不可用:回退到规则和人工填写。 - LLM 不可用:回退到规则和人工填写。
- OCR 不可用:文件保留并进入待识别队列,用户可手工补录。 - OCR 不可用:文件保留并进入待识别队列,用户可手工补录。
- Qdrant/few-shot 不可用:使用 stable Prompt 和基础规则,不阻塞主流程。 - Qdrant/few-shot 不可用:使用 stable Prompt 和基础规则,不阻塞主流程。
- 历史案例证据只返回固定的脱敏标签与摘要,不返回样本 ID、费用单号、人工评论或历史结论原文它只能辅助复核不参与确定性规则、阻断数量、预算复核和审批路由计算。
- 外部支付/ERP/税务连接器不可用:事件进入 retryable 状态,保留人工处理入口和幂等键。 - 外部支付/ERP/税务连接器不可用:事件进入 retryable 状态,保留人工处理入口和幂等键。
- 记忆冲突或可信度不足:只展示建议,不自动应用。 - 记忆冲突或可信度不足:只展示建议,不自动应用。
- 结构化预览签发超时或不可用:允许继续编辑和重新签发,但保存、提交保持 fail-closed只有不存在结构化预览的历史兼容入口保留原有降级路径且只能记录 `client_observed` 行为证据。 - 结构化预览签发超时或不可用:允许继续编辑和重新签发,但保存、提交保持 fail-closed只有不存在结构化预览的历史兼容入口保留原有降级路径且只能记录 `client_observed` 行为证据。
@@ -534,6 +540,10 @@ docker exec -w /app -e SERVER_VENV_DIR=/tmp/x-financial-server-venv \
## 本轮实现记录 ## 本轮实现记录
- 2026-07-16分层组织记忆`travel_application.transport_mode` 已形成当前输入/规则 > 企业 > 部门 > 个人的解析链。租户管理员可维护 30-365 天有效的企业/部门低敏记忆;同级冲突 fail-closed低级覆盖、过期和撤销均返回可解释但脱敏的状态。创建、更新、撤销使用作用域锁、稳定幂等键和请求指纹响应丢失可安全重放请求内容变化返回冲突。
- 2026-07-16租户化历史案例风险观测、人工反馈、few-shot 关系数据和 Qdrant 向量统一绑定租户、场景、制度和规则版本Hermes 按租户分别构建图与历史,风险规则生成不再回落到默认租户。报销预审与 AI 助手接入只读 `historical_case_evidence`,旧版本显式标 stale检索故障降级为空证据。
- 2026-07-16历史证据隐私边界公开协议仅返回“历史已确认/历史误报,仅供复核”的固定摘要,不暴露样本 ID、费用单号、人工评论和历史结论原文历史命中不改变 review ID、规则 findings、passed、blocking count、预算复核或审批路由。
- 2026-07-16分层学习验证容器内分层与组织记忆 44 项、历史案例/风险/预审/规则生成 90 项、报销接口与费用服务 137 项、Golden 与历史注入 28 项、迁移/模型/租户组合 70 项通过;前端 18 项和 Vite 生产构建通过。一次性 tmpfs PostgreSQL 17 完整迁移循环 9 项通过,临时容器自动清理且开发数据库未修改;变更 Python 文件 Ruff F/I、compileall 与 `git diff --check` 通过。
- 2026-07-16P0 提交前预审握手):新增稳定预审决策契约、输入/规则/动态 findings 指纹和结构化整改动作;申请与报销提交前均由服务端重新计算完整风险上下文,可整改重大风险在预算占用前返回 409握手期间风险变化返回 `PRE_REVIEW_CHANGED`,预算治理和人工判断风险继续进入领导与 P8。 - 2026-07-16P0 提交前预审握手):新增稳定预审决策契约、输入/规则/动态 findings 指纹和结构化整改动作;申请与报销提交前均由服务端重新计算完整风险上下文,可整改重大风险在预算占用前返回 409握手期间风险变化返回 `PRE_REVIEW_CHANGED`,预算治理和人工判断风险继续进入领导与 P8。
- 2026-07-16P0 费用事件续接):新增 `application_pre_review_completed` / `claim_pre_review_completed` 时间线语义,预审与提交共享 correlation/causation票据归集完成后刷新预审申请批准生成空报销草稿时 Expense Case 保持 `approved_to_spend`,首张票据归集后再进入 `claiming` - 2026-07-16P0 费用事件续接):新增 `application_pre_review_completed` / `claim_pre_review_completed` 时间线语义,预审与提交共享 correlation/causation票据归集完成后刷新预审申请批准生成空报销草稿时 Expense Case 保持 `approved_to_spend`,首张票据归集后再进入 `claiming`
- 2026-07-16租户与路由修复AI 草稿、申请预览、Steward 和差旅测算透传真实租户Claim 核心访问和关联草稿后台任务按 CaseLink 与服务端租户隔离;仅未整改的人工判断/预算治理高风险申请进入同部门 P8已解决风险不再重复升级。 - 2026-07-16租户与路由修复AI 草稿、申请预览、Steward 和差旅测算透传真实租户Claim 核心访问和关联草稿后台任务按 CaseLink 与服务端租户隔离;仅未整改的人工判断/预算治理高风险申请进入同部门 P8已解决风险不再重复升级。

View File

@@ -142,7 +142,8 @@
- [ ] [CONCEPT: 风险与预审] 统一风险输出为事实、规则、证据、判断、建议动作和降级原因。 - [ ] [CONCEPT: 风险与预审] 统一风险输出为事实、规则、证据、判断、建议动作和降级原因。
- [x] [CONCEPT: 记忆激活] 为首个个人出行方式切片实现 candidate/active/suppressed/expired/revoked 记忆状态机。 - [x] [CONCEPT: 记忆激活] 为首个个人出行方式切片实现 candidate/active/suppressed/expired/revoked 记忆状态机。
证据Candidate/Active 默认有效期分别为 90/180 天;反向或非白名单纠正抑制已激活记忆,过期后新证据创建新 generation忘记后清值并阻止旧请求复活。 证据Candidate/Active 默认有效期分别为 90/180 天;反向或非白名单纠正抑制已激活记忆,过期后新证据创建新 generation忘记后清值并阻止旧请求复活。
- [ ] [CONCEPT: 记忆激活] 实现用户、部门、企业记忆优先级、冲突解释、时间衰减和最小样本要求。 - [x] [CONCEPT: 记忆激活] 实现用户、部门、企业记忆优先级、冲突解释、时间衰减和最小样本要求。
证据:首个切片限定 `travel_application.transport_mode`,按当前输入/规则 > 企业 > 部门 > 个人解析;组织记忆由管理员显式激活并受 30-365 天有效期约束,个人记忆继续执行 3 Case、2 次审批通过、跨 7 天门槛和时间衰减。同优先级不同值 fail-closed低优先级覆盖关系返回脱敏解释过期、撤销和并发换代均不会继续应用。
- [x] [CONCEPT: 记忆与学习] 新增租户化个人记忆与证据链接迁移,唯一键保证同一 Expense Case 对同一候选最多计一票。 - [x] [CONCEPT: 记忆与学习] 新增租户化个人记忆与证据链接迁移,唯一键保证同一 Expense Case 对同一候选最多计一票。
证据记忆、证据、Decision、Feedback、Outcome、Expense Case 和 Claim 采用租户复合外键;证据同时校验 actor、员工和 Case/单据 owner跨租户或跨主体关系不能计票。 证据记忆、证据、Decision、Feedback、Outcome、Expense Case 和 Claim 采用租户复合外键;证据同时校验 actor、员工和 Case/单据 owner跨租户或跨主体关系不能计票。
- [x] [CONCEPT: 记忆激活] 为常用出行方式实现隐式激活门槛3 个不同 Case 的一致服务端纠正、至少 2 次审批通过、证据跨 7 天草稿、client_observed、accepted 和重试不激活。 - [x] [CONCEPT: 记忆激活] 为常用出行方式实现隐式激活门槛3 个不同 Case 的一致服务端纠正、至少 2 次审批通过、证据跨 7 天草稿、client_observed、accepted 和重试不激活。
@@ -157,7 +158,8 @@
- [x] [CONCEPT: 记忆与学习] 将小财管家、Steward 与通用 Orchestrator 的申请预览切换到认证签发与消费链路,并把结构化预览失败策略收口为 fail-closed。 - [x] [CONCEPT: 记忆与学习] 将小财管家、Steward 与通用 Orchestrator 的申请预览切换到认证签发与消费链路,并把结构化预览失败策略收口为 fail-closed。
证据:`expense_application_preview_workflow.py``orchestrator_expense_application_workflow.py``steward_action_executor.py``useTravelReimbursementApplicationPreviewActions.js`;签发/动作请求 ID 可稳定重试,草稿续签回写服务端会话,字段接受/纠正按 `server_verified` 落账,未签发结构化预览不能保存或提交。 证据:`expense_application_preview_workflow.py``orchestrator_expense_application_workflow.py``steward_action_executor.py``useTravelReimbursementApplicationPreviewActions.js`;签发/动作请求 ID 可稳定重试,草稿续签回写服务端会话,字段接受/纠正按 `server_verified` 落账,未签发结构化预览不能保存或提交。
- [ ] [CONCEPT: 记忆与学习] 从字段接受/修改/拒绝、退回、审批覆盖、付款和审计结果生成记忆证据。 - [ ] [CONCEPT: 记忆与学习] 从字段接受/修改/拒绝、退回、审批覆盖、付款和审计结果生成记忆证据。
- [ ] [CONCEPT: 记忆与学习] 将已确认 few-shot 扩展到报销预审和审批辅助,并按租户、场景、制度版本过滤。 - [x] [CONCEPT: 记忆与学习] 将已确认 few-shot 扩展到报销预审和审批辅助,并按租户、场景、制度版本过滤。
证据:`few_shot_ingestion.py``few_shot_retrieval.py``expense_claim_historical_evidence.py``expense_claim_pre_review.py`;样本关系库和 Qdrant 同时绑定租户、场景、制度与规则版本,检索命中后再由关系库校验,旧版本标记 stale依赖异常返回空证据。公开 `historical_case_evidence` 与 AI 助手只显示“历史已确认/历史误报,仅供复核”的固定脱敏摘要,不暴露 sample ID、单号、人工评论或结论原文也不改变确定性结论、阻断数量、预算复核和审批路由。
- [ ] [CONCEPT: 风险与预审] 完成 golden case、Prompt/规则版本、Canary、回归门禁和自动回滚。 - [ ] [CONCEPT: 风险与预审] 完成 golden case、Prompt/规则版本、Canary、回归门禁和自动回滚。
- [ ] [CONCEPT: 自动化决策] 先上线 shadow再按动作逐项开放 L3L4 必须单独评审。 - [ ] [CONCEPT: 自动化决策] 先上线 shadow再按动作逐项开放 L3L4 必须单独评审。
- [ ] [CONCEPT: 降级策略] 对模型、OCR、Qdrant、连接器和记忆服务实现稳定降级和可观测状态。 - [ ] [CONCEPT: 降级策略] 对模型、OCR、Qdrant、连接器和记忆服务实现稳定降级和可观测状态。
@@ -171,7 +173,9 @@
- [ ] [CONCEPT: AI 记忆与自动化设置] 新增“我的 AI 记忆”,支持来源解释、修改、忘记和关闭个性化。 - [ ] [CONCEPT: AI 记忆与自动化设置] 新增“我的 AI 记忆”,支持来源解释、修改、忘记和关闭个性化。
- [x] [CONCEPT: AI 记忆与自动化设置] 在申请核对表展示常用出行方式的记忆来源、证据数量和“忘记此偏好”,并在保存/提交后区分候选记录与已应用回执。 - [x] [CONCEPT: AI 记忆与自动化设置] 在申请核对表展示常用出行方式的记忆来源、证据数量和“忘记此偏好”,并在保存/提交后区分候选记录与已应用回执。
证据:`TravelReimbursementMemoryPanel.vue` 与独立样式分片承载记忆解释、学习回执和可访问的忘记入口;会话快照可跨刷新恢复已应用记忆。 证据:`TravelReimbursementMemoryPanel.vue` 与独立样式分片承载记忆解释、学习回执和可访问的忘记入口;会话快照可跨刷新恢复已应用记忆。
- [ ] [CONCEPT: AI 记忆与自动化设置] 新增企业记忆、保留策略、敏感等级和自动化上限管理 - [x] [CONCEPT: AI 记忆与自动化设置] 新增企业/部门出行方式记忆管理首个低敏切片
证据:租户管理员可创建、查看、更新和撤销企业/部门记忆,配置 30-365 天有效期;创建、更新、撤销使用请求指纹和稳定幂等键,作用域锁与数据库唯一约束阻止并发静默换代。接口只接受“飞机/火车/轮船”,自由文本、金额、客户、项目、附件和支付信息保持禁用。
- [ ] [CONCEPT: AI 记忆与自动化设置] 扩展可配置保留策略、敏感等级目录和动作级自动化上限管理;在新增字段前先完成制度允许性与数据分级评审。
- [ ] [CONCEPT: 前端] 展示自动化动作、执行依据、撤销入口、抽检状态和版本信息。 - [ ] [CONCEPT: 前端] 展示自动化动作、执行依据、撤销入口、抽检状态和版本信息。
## 8. P2 实现:费用经营与价值证明 ## 8. P2 实现:费用经营与价值证明
@@ -211,7 +215,7 @@
证据:`test_expense_case_endpoints.py` 容器内 8 项通过,覆盖跨租户、无 Case、申请与报销关联摘要以及内部字段递归过滤。 证据:`test_expense_case_endpoints.py` 容器内 8 项通过,覆盖跨租户、无 Case、申请与报销关联摘要以及内部字段递归过滤。
- [ ] [CONCEPT: 测试方案] 为 Alembic baseline、升级、旧数据迁移和回滚边界补充 Postgres 集成测试。 - [ ] [CONCEPT: 测试方案] 为 Alembic baseline、升级、旧数据迁移和回滚边界补充 Postgres 集成测试。
- [x] [CONCEPT: 测试方案] 为当前 migration-owned schema 切片补充一次性 PostgreSQL 集成测试和危险 URL 防误连门禁。 - [x] [CONCEPT: 测试方案] 为当前 migration-owned schema 切片补充一次性 PostgreSQL 集成测试和危险 URL 防误连门禁。
证据:`test_alembic_migrations.py` 默认无显式 URL 时跳过,主机和库名必须带 disposable 标记tmpfs PostgreSQL 17 中 4 项通过,覆盖空库升级、重复升级、关键约束/索引、外键级联、base 降级、legacy 哨兵保留、漂移拒绝和再次升级;持久化开发库未被修改。完整 legacy baseline 仍保留在上一条未完成项中。 证据:`test_alembic_migrations.py` 默认无显式 URL 时跳过,主机和库名必须带 disposable 标记;当前 0009 Head 在 tmpfs PostgreSQL 17 中 9 项通过,覆盖空库升级、重复升级、关键约束/索引、组织 active 脏数据升级前拒绝、版本化 few-shot 无损降级拒绝、外键级联、base 降级、legacy 哨兵保留、漂移拒绝和再次升级;临时容器自动清理,持久化开发库未被修改。完整 legacy baseline 仍保留在上一条未完成项中。
- [ ] [CONCEPT: 测试方案] 为连接器幂等、重试、回执、失败恢复、重复付款和对账补充测试。 - [ ] [CONCEPT: 测试方案] 为连接器幂等、重试、回执、失败恢复、重复付款和对账补充测试。
- [ ] [CONCEPT: 测试方案] 跑通申请 → 票据 → 报销 → 预审 → 审批 → 付款 → 入账 → 归档端到端。 - [ ] [CONCEPT: 测试方案] 跑通申请 → 票据 → 报销 → 预审 → 审批 → 付款 → 入账 → 归档端到端。
当前进度:申请批准 → 自动报销草稿 → 票据归集 → 预审 → 报销提交已在同一 Case 中跑通付款回执、ERP 入账和对账仍未接入。 当前进度:申请批准 → 自动报销草稿 → 票据归集 → 预审 → 报销提交已在同一 Case 中跑通付款回执、ERP 入账和对账仍未接入。

View File

@@ -0,0 +1,7 @@
## 修复记录
- 2026-07-16 13:42:57 CST在分层记忆回归测试中发现AI 新建非 `default` 租户费用申请后会立即调用 `submit_claim()`,但此时尚未建立租户化 Expense Case Link提交服务重新查询单据时按 Case Link 执行租户过滤,因此刚创建的申请返回不可见并报“未找到可提交的申请单”。
- 2026-07-16 13:42:57 CST`UserAgentApplicationSlotMixin._create_expense_application_record()` 完成 Claim `flush` 后、进入草稿或提交分支前,使用当前认证租户同步建立 Expense Case 与 Claim Link使新建记录、租户归属和后续提交保持在同一数据库事务失败时仍由既有外层回滚不留下半成品 Case 或 Link。
- 2026-07-16 13:42:57 CST按项目规范执行 `git fetch --all --prune`、上下游状态与提交差异检查;当前分支相对 `origin/main` ahead 13未发现新的上游提交本地 ahead 包含最近的费用闭环、个人记忆、零录入票据和权威预审阶段提交。
- 2026-07-16 13:42:57 CST在 Docker 容器 `local-x-financial-linux` 内运行原失败用例、分层记忆服务与组织记忆接口测试,结果 `9 passed`;验证非默认租户申请可完成预览提交并生成学习回执,组织记忆生命周期和租户隔离未受影响。
- 影响:修复 AI 助手、Steward 和 Orchestrator 共用申请创建链路在非默认租户下无法直接提交的问题;默认租户与既有草稿事件继续复用同一 Case不改变确定性预审、预算和审批规则。

View File

@@ -0,0 +1,8 @@
## 修复记录
- 2026-07-16 14:29:13 CST组织记忆创建在“检查 active → 新建 generation”之间存在竞态后到请求可能静默压制并发产生的新版本创建、更新的幂等键也没有绑定请求内容撤销操作则无法在响应丢失后安全重试。
- 2026-07-16 14:29:13 CST为企业/部门记忆增加租户与作用域级事务锁、幂等请求锁、请求 payload 指纹和数据库唯一约束;创建、更新、撤销现在都满足“同键同内容重放、同键不同内容 409”新 generation 激活前严格核对预期 active 集合。锁协调职责拆到 `organization_memory_locks.py`,主服务从接近 800 行降至 761 行。
- 2026-07-16 14:29:13 CST加固 `20260716_0009` 迁移active 唯一索引只约束企业/部门作用域DDL 前先检测重复 active 并 fail-fast避免个人学习记忆换代受到影响非 PostgreSQL 在任何结构变更前明确拒绝,避免 SQLite 留下半迁移列。
- 2026-07-16 14:29:13 CST执行 `git fetch --all --prune`、工作区状态和上下游提交差异检查;`origin/main` 没有新增提交,当前分支 ahead 13最近本地检查点为权威提交前预审、持久化零录入票据和个人申请记忆本次未改写这些提交。
- 2026-07-16 14:29:13 CST全部验证在 `local-x-financial-linux` 内执行并设置 60 秒超时。分层与组织记忆回归 44 项、迁移/模型/租户组合 70 项通过;前端记忆与申请链路 18 项通过Vite 生产构建通过,变更 Python 文件 Ruff F/I 与 compileall 通过。一次性 tmpfs PostgreSQL 17 完整迁移循环 9 项通过,临时容器自动清理,开发数据库未修改。
- 影响:管理员并发维护组织记忆时不再发生静默覆盖,网络重试不会产生重复版本或把不同请求伪装成成功;个人记忆状态机保持原有最小样本与换代语义,迁移遇到脏数据时在写 DDL 前停止并给出明确错误。

View File

@@ -0,0 +1,9 @@
## 修复记录
- 2026-07-16 14:29:13 CST历史风险观测、few-shot 样本、Qdrant 向量和规则生成链路原先缺少完整租户、场景及制度版本边界Hermes 全局扫描还可能把不同租户的 Claim 放进同一风险图,历史检索结果若直接公开也可能泄露样本 ID、费用单号或人工评语。
- 2026-07-16 14:29:13 CST风险观测与 few-shot 关系数据改为租户复合唯一,风险接口、反馈、查询和 Hermes 图构建按认证租户隔离Qdrant 写入与检索显式携带 tenant、scene、policy_ref、rule_version使用稳定向量 ID并清理旧向量命中后再由关系库校验租户、状态和版本。风险规则生成与再生成从认证用户透传租户缺失租户时禁用历史注入而不回落到 default。
- 2026-07-16 14:29:13 CST报销预审新增只读 `historical_case_evidence`,检索或向量服务异常时降级为空证据。公开协议和 AI 助手只返回固定的“历史已确认/历史误报,仅供复核”标签与脱敏摘要,不返回 sample ID、claim_no、人工评论或历史结论原文历史证据不参与 review ID、确定性 findings、passed、blocking count、预算复核或审批路由计算。
- 2026-07-16 14:29:13 CST加固 `20260716_0008` 迁移:`RiskObservation.claim_id` 明确改为旧 Claim 表的 view-only 软引用,使空库与旧表采用路径的 Head 结构一致;降级若发现非默认租户或非空 policy/rule 版本数据则在 DDL 前 fail-fast禁止静默丢失隔离和版本信息非 PostgreSQL 在变更前明确拒绝。
- 2026-07-16 14:29:13 CST执行 `git fetch --all --prune`、工作区状态和上下游提交差异检查;`origin/main` 没有新增提交,当前分支 ahead 13最近本地检查点为权威提交前预审、持久化零录入票据和个人申请记忆本次未改写这些提交。
- 2026-07-16 14:29:13 CST全部验证在 `local-x-financial-linux` 内执行并设置 60 秒超时。历史案例、风险观测、风险图、预审和规则生成组合 90 项,报销接口与费用服务 137 项,补充 Golden/规则历史注入 28 项通过;迁移/模型/租户组合 70 项通过Ruff F/I、compileall 和 `git diff --check` 通过。一次性 tmpfs PostgreSQL 17 完整迁移循环 9 项通过,临时容器自动清理,开发数据库未修改。
- 影响:同名用户、相同样本键或相似向量不能跨租户读取、覆盖或进入 Prompt历史案例开始帮助预审与复核但只以脱敏、可降级、非决策证据存在不会扩大 AI 自动放行或审批权限。

View File

@@ -0,0 +1,133 @@
"""add hierarchical expense memory governance fields
Revision ID: 20260716_0007
Revises: 20260716_0006
Create Date: 2026-07-16 12:05:00
"""
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "20260716_0007"
down_revision: str | None = "20260716_0006"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
op.add_column(
"memory_entries",
sa.Column(
"origin_type",
sa.String(length=24),
server_default="learned",
nullable=False,
),
)
op.add_column(
"memory_entries",
sa.Column("managed_by", sa.String(length=255), nullable=True),
)
op.add_column(
"memory_entries",
sa.Column("managed_at", sa.DateTime(timezone=True), nullable=True),
)
op.add_column(
"memory_entries",
sa.Column("management_reason", sa.String(length=255), nullable=True),
)
op.add_column(
"memory_entries",
sa.Column("policy_version", sa.String(length=64), nullable=True),
)
op.drop_constraint(
"ck_memory_entries_scope_type",
"memory_entries",
type_="check",
)
op.create_check_constraint(
"ck_memory_entries_scope_type",
"memory_entries",
"scope_type IN ('user', 'department', 'enterprise')",
)
op.create_check_constraint(
"ck_memory_entries_origin_type",
"memory_entries",
"origin_type IN ('learned', 'admin_managed')",
)
op.create_check_constraint(
"ck_memory_entries_scope_origin",
"memory_entries",
"(scope_type = 'user' AND origin_type = 'learned') OR "
"(scope_type IN ('department', 'enterprise') "
"AND origin_type = 'admin_managed')",
)
op.create_check_constraint(
"ck_memory_entries_enterprise_scope",
"memory_entries",
"scope_type != 'enterprise' OR scope_id = tenant_id",
)
op.create_check_constraint(
"ck_memory_entries_management_audit",
"memory_entries",
"(origin_type = 'learned' AND managed_by IS NULL "
"AND managed_at IS NULL AND management_reason IS NULL) OR "
"(origin_type = 'admin_managed' AND managed_by IS NOT NULL "
"AND length(trim(managed_by)) > 0 AND managed_at IS NOT NULL "
"AND management_reason IS NOT NULL "
"AND length(trim(management_reason)) > 0 "
"AND policy_version IS NOT NULL AND length(trim(policy_version)) > 0)",
)
def downgrade() -> None:
op.drop_constraint(
"ck_memory_entries_management_audit",
"memory_entries",
type_="check",
)
op.drop_constraint(
"ck_memory_entries_enterprise_scope",
"memory_entries",
type_="check",
)
op.drop_constraint(
"ck_memory_entries_scope_origin",
"memory_entries",
type_="check",
)
op.drop_constraint(
"ck_memory_entries_origin_type",
"memory_entries",
type_="check",
)
op.drop_constraint(
"ck_memory_entries_scope_type",
"memory_entries",
type_="check",
)
# 0006 只认识个人记忆。组织记忆无法无损映射回旧结构,因此降级时
# 先解除组织版本间的自引用并删除组织作用域记录,个人记忆完整保留。
op.execute(
"UPDATE memory_entries SET superseded_by_id = NULL "
"WHERE scope_type IN ('department', 'enterprise')"
)
op.execute(
"DELETE FROM memory_entries "
"WHERE scope_type IN ('department', 'enterprise')"
)
op.create_check_constraint(
"ck_memory_entries_scope_type",
"memory_entries",
"scope_type = 'user'",
)
op.drop_column("memory_entries", "policy_version")
op.drop_column("memory_entries", "management_reason")
op.drop_column("memory_entries", "managed_at")
op.drop_column("memory_entries", "managed_by")
op.drop_column("memory_entries", "origin_type")

View File

@@ -0,0 +1,435 @@
"""adopt tenant-safe historical case learning tables
Revision ID: 20260716_0008
Revises: 20260716_0007
Create Date: 2026-07-16 13:20:00
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "20260716_0008"
down_revision: str | None = "20260716_0007"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _require_postgresql() -> None:
dialect_name = op.get_bind().dialect.name
if dialect_name != "postgresql":
raise RuntimeError(
"20260716_0008 only supports PostgreSQL; "
f"refusing to mutate {dialect_name} without transactional constraint DDL"
)
def _inspector() -> sa.Inspector:
return sa.inspect(op.get_bind())
def _table_exists(table_name: str) -> bool:
return table_name in _inspector().get_table_names()
def _columns(table_name: str) -> set[str]:
return {str(item["name"]) for item in _inspector().get_columns(table_name)}
def _indexes(table_name: str) -> dict[str, dict[str, object]]:
return {
str(item["name"]): item for item in _inspector().get_indexes(table_name) if item.get("name")
}
def _unique_constraints(table_name: str) -> dict[str, dict[str, object]]:
return {
str(item["name"]): item
for item in _inspector().get_unique_constraints(table_name)
if item.get("name")
}
def _foreign_keys(table_name: str) -> list[dict[str, object]]:
return list(_inspector().get_foreign_keys(table_name))
def _ensure_index(
table_name: str,
index_name: str,
columns: list[str],
*,
unique: bool = False,
) -> None:
if index_name not in _indexes(table_name):
op.create_index(index_name, table_name, columns, unique=unique)
def _drop_legacy_single_key_uniqueness(table_name: str, key_column: str) -> None:
for name, item in list(_unique_constraints(table_name).items()):
if tuple(item.get("column_names") or ()) == (key_column,):
op.drop_constraint(name, table_name, type_="unique")
# PostgreSQL 会把唯一约束的支撑索引同时返回给 get_indexes先删除约束并
# 重新反射,避免对支撑索引执行 DROP INDEX 导致 dependency error。
for name, item in list(_indexes(table_name).items()):
if tuple(item.get("column_names") or ()) == (key_column,) and bool(item.get("unique")):
op.drop_index(name, table_name=table_name)
def _ensure_composite_unique(
table_name: str,
constraint_name: str,
columns: list[str],
) -> None:
expected = tuple(columns)
if any(
tuple(item.get("column_names") or ()) == expected
for item in _unique_constraints(table_name).values()
):
return
op.create_unique_constraint(constraint_name, table_name, columns)
def _drop_index_if_exists(table_name: str, index_name: str) -> None:
if index_name in _indexes(table_name):
op.drop_index(index_name, table_name=table_name)
def _drop_unique_if_exists(table_name: str, constraint_name: str) -> None:
if constraint_name in _unique_constraints(table_name):
op.drop_constraint(constraint_name, table_name, type_="unique")
def _require_lossless_default_tenant_downgrade(table_name: str) -> None:
non_default_count = int(
op.get_bind().scalar(
sa.text(
f"SELECT COUNT(*) FROM {table_name} "
"WHERE tenant_id IS DISTINCT FROM 'default'"
)
)
or 0
)
if non_default_count:
raise RuntimeError(
f"cannot downgrade {table_name}: non-default tenant data would lose isolation"
)
def _require_lossless_few_shot_downgrade() -> None:
enriched_count = int(
op.get_bind().scalar(
sa.text(
"SELECT COUNT(*) FROM few_shot_samples "
"WHERE COALESCE(TRIM(policy_ref), '') <> '' "
"OR COALESCE(TRIM(rule_version), '') <> ''"
)
)
or 0
)
if enriched_count:
raise RuntimeError(
"cannot downgrade few_shot_samples: "
f"{enriched_count} row(s) contain policy_ref or rule_version data"
)
def _drop_risk_observation_claim_foreign_keys() -> None:
"""统一采用软引用,避免 Head 结构取决于 expense_claims 的创建时机。"""
for item in _foreign_keys("risk_observations"):
if (
tuple(item.get("constrained_columns") or ()) == ("claim_id",)
and item.get("referred_table") == "expense_claims"
and item.get("name")
):
op.drop_constraint(
str(item["name"]),
"risk_observations",
type_="foreignkey",
)
def _create_risk_observations() -> None:
op.create_table(
"risk_observations",
sa.Column("id", sa.String(length=36), nullable=False),
sa.Column("tenant_id", sa.String(length=64), server_default="default", nullable=False),
sa.Column("observation_key", sa.String(length=160), nullable=False),
sa.Column("subject_type", sa.String(length=50), nullable=False),
sa.Column("subject_key", sa.String(length=160), nullable=False),
sa.Column("subject_label", sa.String(length=160), nullable=False, server_default=""),
# expense_claims 仍由 legacy bootstrap 创建,空库迁移到此版本时并不存在。
# claim_id 因此是显式软引用,与 ORM 的 viewonly relationship 保持一致。
sa.Column("claim_id", sa.String(length=36), nullable=True),
sa.Column("claim_no", sa.String(length=80), nullable=False, server_default=""),
sa.Column("run_id", sa.String(length=80), nullable=True),
sa.Column("execution_log_id", sa.String(length=36), nullable=True),
sa.Column("risk_type", sa.String(length=80), nullable=False),
sa.Column("risk_signal", sa.String(length=100), nullable=False),
sa.Column("title", sa.String(length=200), nullable=False, server_default=""),
sa.Column("description", sa.Text(), nullable=False, server_default=""),
sa.Column("risk_score", sa.Integer(), nullable=False, server_default="0"),
sa.Column("risk_level", sa.String(length=20), nullable=False),
sa.Column("confidence_score", sa.Float(), nullable=False, server_default="0"),
sa.Column("control_stage", sa.String(length=50), nullable=False, server_default=""),
sa.Column("control_mode", sa.String(length=50), nullable=False, server_default=""),
sa.Column("automation_mode", sa.String(length=50), nullable=False, server_default=""),
sa.Column("source", sa.String(length=60), nullable=False, server_default=""),
sa.Column("algorithm_version", sa.String(length=80), nullable=False, server_default=""),
sa.Column("status", sa.String(length=30), nullable=False, server_default="pending_review"),
sa.Column(
"feedback_status", sa.String(length=30), nullable=False, server_default="unreviewed"
),
sa.Column("contribution_scores_json", sa.JSON(), nullable=False),
sa.Column("baseline_json", sa.JSON(), nullable=False),
sa.Column("evidence_json", sa.JSON(), nullable=False),
sa.Column("graph_node_keys_json", sa.JSON(), nullable=False),
sa.Column("graph_edge_keys_json", sa.JSON(), nullable=False),
sa.Column("policy_refs_json", sa.JSON(), nullable=False),
sa.Column("similar_case_claim_ids_json", sa.JSON(), nullable=False),
sa.Column("ontology_json", sa.JSON(), nullable=False),
sa.Column("decision_trace_json", sa.JSON(), nullable=False),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.Column(
"updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint("tenant_id", "observation_key", name="uq_risk_observations_tenant_key"),
)
def _adopt_risk_observations() -> None:
_drop_risk_observation_claim_foreign_keys()
if "tenant_id" not in _columns("risk_observations"):
op.add_column(
"risk_observations",
sa.Column("tenant_id", sa.String(length=64), server_default="default", nullable=False),
)
op.execute(
sa.text(
"UPDATE risk_observations SET tenant_id = 'default' "
"WHERE tenant_id IS NULL OR tenant_id = ''"
)
)
_drop_legacy_single_key_uniqueness("risk_observations", "observation_key")
_ensure_composite_unique(
"risk_observations",
"uq_risk_observations_tenant_key",
["tenant_id", "observation_key"],
)
def _create_risk_observation_feedback() -> None:
op.create_table(
"risk_observation_feedback",
sa.Column("id", sa.String(length=36), nullable=False),
sa.Column("observation_id", sa.String(length=36), nullable=False),
sa.Column("feedback_type", sa.String(length=30), nullable=False),
sa.Column("action", sa.String(length=50), nullable=False, server_default=""),
sa.Column("actor", sa.String(length=100), nullable=False, server_default=""),
sa.Column("comment", sa.Text(), nullable=True),
sa.Column("payload_json", sa.JSON(), nullable=False),
sa.Column(
"created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False
),
sa.ForeignKeyConstraint(["observation_id"], ["risk_observations.id"]),
sa.PrimaryKeyConstraint("id"),
)
def _create_few_shot_samples() -> None:
op.create_table(
"few_shot_samples",
sa.Column("id", sa.String(length=36), nullable=False),
sa.Column("tenant_id", sa.String(length=64), server_default="default", nullable=False),
sa.Column("sample_key", sa.String(length=160), nullable=False),
sa.Column("source_observation_id", sa.String(length=36), nullable=True),
sa.Column(
"scene", sa.String(length=50), nullable=False, server_default="risk_rule_generation"
),
sa.Column("policy_ref", sa.String(length=160), nullable=False, server_default=""),
sa.Column("rule_version", sa.String(length=80), nullable=False, server_default=""),
sa.Column("domain", sa.String(length=50), nullable=False, server_default=""),
sa.Column("risk_type", sa.String(length=80), nullable=False, server_default=""),
sa.Column("risk_level", sa.String(length=20), nullable=False, server_default=""),
sa.Column("label", sa.String(length=30), nullable=False, server_default="confirmed"),
sa.Column("case_text", sa.Text(), nullable=False, server_default=""),
sa.Column("conclusion_text", sa.Text(), nullable=False, server_default=""),
sa.Column("payload_json", sa.JSON(), nullable=False),
sa.Column("vector_id", sa.String(length=100), nullable=True),
sa.Column("status", sa.String(length=20), nullable=False, server_default="active"),
sa.Column("created_at", sa.DateTime(), server_default=sa.func.now(), nullable=False),
sa.Column("updated_at", sa.DateTime(), server_default=sa.func.now(), nullable=False),
sa.ForeignKeyConstraint(["source_observation_id"], ["risk_observations.id"]),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint("tenant_id", "sample_key", name="uq_few_shot_samples_tenant_key"),
)
def _adopt_few_shot_samples() -> None:
existing = _columns("few_shot_samples")
for column in (
sa.Column("tenant_id", sa.String(length=64), server_default="default", nullable=False),
sa.Column("policy_ref", sa.String(length=160), server_default="", nullable=False),
sa.Column("rule_version", sa.String(length=80), server_default="", nullable=False),
):
if column.name not in existing:
op.add_column("few_shot_samples", column)
op.execute(
sa.text(
"UPDATE few_shot_samples SET tenant_id = 'default' "
"WHERE tenant_id IS NULL OR tenant_id = ''"
)
)
op.execute(sa.text("UPDATE few_shot_samples SET policy_ref = '' WHERE policy_ref IS NULL"))
op.execute(sa.text("UPDATE few_shot_samples SET rule_version = '' WHERE rule_version IS NULL"))
_drop_legacy_single_key_uniqueness("few_shot_samples", "sample_key")
_ensure_composite_unique(
"few_shot_samples",
"uq_few_shot_samples_tenant_key",
["tenant_id", "sample_key"],
)
def _ensure_risk_indexes() -> None:
definitions = {
"ix_risk_observations_tenant_id": ["tenant_id"],
"ix_risk_observations_observation_key": ["observation_key"],
"ix_risk_observations_subject_type": ["subject_type"],
"ix_risk_observations_subject_key": ["subject_key"],
"ix_risk_observations_claim_id": ["claim_id"],
"ix_risk_observations_claim_no": ["claim_no"],
"ix_risk_observations_run_id": ["run_id"],
"ix_risk_observations_execution_log_id": ["execution_log_id"],
"ix_risk_observations_risk_type": ["risk_type"],
"ix_risk_observations_risk_signal": ["risk_signal"],
"ix_risk_observations_risk_score": ["risk_score"],
"ix_risk_observations_risk_level": ["risk_level"],
"ix_risk_observations_source": ["source"],
"ix_risk_observations_algorithm_version": ["algorithm_version"],
"ix_risk_observations_status": ["status"],
"ix_risk_observations_feedback_status": ["feedback_status"],
"ix_risk_observations_subject": ["subject_type", "subject_key"],
"ix_risk_observations_signal_level": ["risk_signal", "risk_level"],
"ix_risk_observations_status_created": ["status", "created_at"],
"ix_risk_observations_tenant_status": ["tenant_id", "status", "created_at"],
}
for name, columns in definitions.items():
_ensure_index("risk_observations", name, columns)
_ensure_index(
"risk_observation_feedback",
"ix_risk_observation_feedback_observation_id",
["observation_id"],
)
_ensure_index(
"risk_observation_feedback", "ix_risk_observation_feedback_feedback_type", ["feedback_type"]
)
_ensure_index(
"risk_observation_feedback",
"ix_risk_observation_feedback_type_created",
["feedback_type", "created_at"],
)
def _ensure_few_shot_indexes() -> None:
definitions = {
"ix_few_shot_samples_tenant_id": ["tenant_id"],
"ix_few_shot_samples_sample_key": ["sample_key"],
"ix_few_shot_samples_source_observation_id": ["source_observation_id"],
"ix_few_shot_samples_scene": ["scene"],
"ix_few_shot_samples_policy_ref": ["policy_ref"],
"ix_few_shot_samples_rule_version": ["rule_version"],
"ix_few_shot_samples_domain": ["domain"],
"ix_few_shot_samples_risk_type": ["risk_type"],
"ix_few_shot_samples_label": ["label"],
"ix_few_shot_samples_status": ["status"],
"ix_few_shot_samples_scene_label": ["scene", "label"],
"ix_few_shot_samples_domain_risk_type": ["domain", "risk_type"],
"ix_few_shot_samples_tenant_rule_lookup": [
"tenant_id",
"scene",
"policy_ref",
"rule_version",
"status",
],
}
for name, columns in definitions.items():
_ensure_index("few_shot_samples", name, columns)
def upgrade() -> None:
_require_postgresql()
if _table_exists("risk_observations"):
_adopt_risk_observations()
else:
_create_risk_observations()
if not _table_exists("risk_observation_feedback"):
_create_risk_observation_feedback()
if _table_exists("few_shot_samples"):
_adopt_few_shot_samples()
else:
_create_few_shot_samples()
_ensure_risk_indexes()
_ensure_few_shot_indexes()
def downgrade() -> None:
_require_postgresql()
# 0008 之前这三张表由旧 bootstrap 管理。降级必须恢复旧结构而不是删除
# 已确认的风险反馈与 few-shot 数据;若存在非默认租户数据则无法无损合并,
# 直接失败并保持当前版本,禁止用数据丢失换取“成功降级”。
few_shot_exists = _table_exists("few_shot_samples")
risk_observations_exists = _table_exists("risk_observations")
# 所有数据安全检查必须先于 DDL任一检查失败时完整保留 0008 结构。
if few_shot_exists:
_require_lossless_default_tenant_downgrade("few_shot_samples")
_require_lossless_few_shot_downgrade()
if risk_observations_exists:
_require_lossless_default_tenant_downgrade("risk_observations")
if few_shot_exists:
for index_name in (
"ix_few_shot_samples_tenant_rule_lookup",
"ix_few_shot_samples_tenant_id",
"ix_few_shot_samples_policy_ref",
"ix_few_shot_samples_rule_version",
):
_drop_index_if_exists("few_shot_samples", index_name)
_drop_unique_if_exists(
"few_shot_samples",
"uq_few_shot_samples_tenant_key",
)
op.create_unique_constraint(
"uq_few_shot_samples_sample_key",
"few_shot_samples",
["sample_key"],
)
op.drop_column("few_shot_samples", "rule_version")
op.drop_column("few_shot_samples", "policy_ref")
op.drop_column("few_shot_samples", "tenant_id")
if risk_observations_exists:
for index_name in (
"ix_risk_observations_tenant_status",
"ix_risk_observations_tenant_id",
):
_drop_index_if_exists("risk_observations", index_name)
_drop_unique_if_exists(
"risk_observations",
"uq_risk_observations_tenant_key",
)
op.create_unique_constraint(
"uq_risk_observations_observation_key",
"risk_observations",
["observation_key"],
)
op.drop_column("risk_observations", "tenant_id")

View File

@@ -0,0 +1,127 @@
"""serialize organization memory mutations and persist idempotency fingerprints
Revision ID: 20260716_0009
Revises: 20260716_0008
Create Date: 2026-07-16 14:20:00
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "20260716_0009"
down_revision: str | None = "20260716_0008"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def _require_postgresql() -> None:
dialect_name = op.get_bind().dialect.name
if dialect_name != "postgresql":
raise RuntimeError(
"20260716_0009 only supports PostgreSQL; "
f"refusing to mutate {dialect_name} without transactional constraint DDL"
)
def _require_unique_active_organization_scopes() -> None:
duplicate_count = int(
op.get_bind().scalar(
sa.text(
"SELECT COUNT(*) FROM ("
"SELECT 1 FROM memory_entries "
"WHERE status = 'active' "
"AND scope_type IN ('department', 'enterprise') "
"GROUP BY tenant_id, scope_type, scope_id, scene, field_key "
"HAVING COUNT(*) > 1"
") AS duplicate_active_organization_scopes"
)
)
or 0
)
if duplicate_count:
raise RuntimeError(
"cannot upgrade organization memory idempotency: "
f"found {duplicate_count} organization scope(s) with duplicate active memories"
)
def upgrade() -> None:
_require_postgresql()
# 必须在添加列或约束前完成数据预检,避免失败后留下半迁移结构。
_require_unique_active_organization_scopes()
op.add_column(
"memory_entries",
sa.Column("management_request_id", sa.String(length=120), nullable=True),
)
op.add_column(
"memory_entries",
sa.Column("management_payload_fingerprint", sa.String(length=80), nullable=True),
)
op.add_column(
"memory_entries",
sa.Column("revoke_request_id", sa.String(length=120), nullable=True),
)
op.add_column(
"memory_entries",
sa.Column("revoke_payload_fingerprint", sa.String(length=80), nullable=True),
)
op.create_check_constraint(
"ck_memory_entries_management_idempotency_pair",
"memory_entries",
"(management_request_id IS NULL AND management_payload_fingerprint IS NULL) OR "
"(management_request_id IS NOT NULL "
"AND management_payload_fingerprint IS NOT NULL)",
)
op.create_check_constraint(
"ck_memory_entries_revoke_idempotency_pair",
"memory_entries",
"(revoke_request_id IS NULL AND revoke_payload_fingerprint IS NULL) OR "
"(revoke_request_id IS NOT NULL AND revoke_payload_fingerprint IS NOT NULL)",
)
op.create_index(
"uq_memory_entries_management_request",
"memory_entries",
["tenant_id", "management_request_id"],
unique=True,
)
op.create_index(
"uq_memory_entries_revoke_request",
"memory_entries",
["tenant_id", "revoke_request_id"],
unique=True,
)
op.create_index(
"uq_memory_entries_active_scope",
"memory_entries",
["tenant_id", "scope_type", "scope_id", "scene", "field_key"],
unique=True,
postgresql_where=sa.text(
"status = 'active' AND scope_type IN ('department', 'enterprise')"
),
)
def downgrade() -> None:
_require_postgresql()
op.drop_index("uq_memory_entries_active_scope", table_name="memory_entries")
op.drop_index("uq_memory_entries_revoke_request", table_name="memory_entries")
op.drop_index("uq_memory_entries_management_request", table_name="memory_entries")
op.drop_constraint(
"ck_memory_entries_revoke_idempotency_pair",
"memory_entries",
type_="check",
)
op.drop_constraint(
"ck_memory_entries_management_idempotency_pair",
"memory_entries",
type_="check",
)
op.drop_column("memory_entries", "revoke_payload_fingerprint")
op.drop_column("memory_entries", "revoke_request_id")
op.drop_column("memory_entries", "management_payload_fingerprint")
op.drop_column("memory_entries", "management_request_id")

View File

@@ -26,6 +26,7 @@ class CurrentUserContext:
is_admin: bool is_admin: bool
tenant_id: str = "default" tenant_id: str = "default"
department_name: str = "" department_name: str = ""
department_id: str = ""
cost_center: str = "" cost_center: str = ""
position: str = "" position: str = ""
grade: str = "" grade: str = ""
@@ -74,6 +75,7 @@ def _authenticate_bearer_user(db: Session, authorization: str | None) -> Current
is_admin=user.is_admin, is_admin=user.is_admin,
tenant_id=user.tenant_id, tenant_id=user.tenant_id,
department_name=user.department, department_name=user.department,
department_id=user.department_id or "",
cost_center=user.cost_center, cost_center=user.cost_center,
position=user.position, position=user.position,
grade=user.grade, grade=user.grade,

View File

@@ -141,6 +141,7 @@ def regenerate_risk_rule(
AgentAssetRiskRuleRegenerationService(db).regenerate( AgentAssetRiskRuleRegenerationService(db).regenerate(
asset_id, asset_id,
payload, payload,
tenant_id=current_user.tenant_id,
actor=_actor_name(current_user, x_actor), actor=_actor_name(current_user, x_actor),
request_id=x_request_id, request_id=x_request_id,
) )

View File

@@ -83,6 +83,7 @@ def _complete_risk_rule_generation_task(
payload: dict, payload: dict,
actor: str, actor: str,
request_id: str | None, request_id: str | None,
tenant_id: str,
) -> None: ) -> None:
db = get_session_factory()() db = get_session_factory()()
try: try:
@@ -90,6 +91,7 @@ def _complete_risk_rule_generation_task(
RiskRuleGenerationJobService(db).complete_rule_asset_generation( RiskRuleGenerationJobService(db).complete_rule_asset_generation(
asset_id, asset_id,
body, body,
tenant_id=tenant_id,
actor=actor, actor=actor,
request_id=request_id, request_id=request_id,
) )
@@ -334,6 +336,7 @@ def generate_agent_asset_risk_rule(
actor = (x_actor or current_user.name or "system").strip() or "system" actor = (x_actor or current_user.name or "system").strip() or "system"
asset_id = RiskRuleGenerationJobService(db).enqueue_rule_asset_generation( asset_id = RiskRuleGenerationJobService(db).enqueue_rule_asset_generation(
payload, payload,
tenant_id=current_user.tenant_id,
actor=actor, actor=actor,
request_id=x_request_id, request_id=x_request_id,
) )
@@ -343,6 +346,7 @@ def generate_agent_asset_risk_rule(
payload.model_dump(mode="json"), payload.model_dump(mode="json"),
actor, actor,
x_request_id, x_request_id,
current_user.tenant_id,
) )
asset = AgentAssetService(db).get_asset(asset_id) asset = AgentAssetService(db).get_asset(asset_id)
if asset is None: if asset is None:
@@ -941,9 +945,10 @@ def create_golden_case(
_: RuleEditorUser, _: RuleEditorUser,
db: DbSession, db: DbSession,
) -> GoldenCaseRead: ) -> GoldenCaseRead:
from app.models.golden_case import GoldenCase
from sqlalchemy import select from sqlalchemy import select
from app.models.golden_case import GoldenCase
existing = db.scalar(select(GoldenCase).where(GoldenCase.case_key == body.case_key)) existing = db.scalar(select(GoldenCase).where(GoldenCase.case_key == body.case_key))
if existing is not None: if existing is not None:
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="case_key 已存在") raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="case_key 已存在")
@@ -975,9 +980,10 @@ def list_golden_cases(
_: CurrentUser, _: CurrentUser,
db: DbSession, db: DbSession,
) -> list[GoldenCaseRead]: ) -> list[GoldenCaseRead]:
from app.models.golden_case import GoldenCase
from sqlalchemy import select from sqlalchemy import select
from app.models.golden_case import GoldenCase
cases = db.scalars( cases = db.scalars(
select(GoldenCase).where(GoldenCase.rule_code == rule_code).order_by(GoldenCase.created_at) select(GoldenCase).where(GoldenCase.rule_code == rule_code).order_by(GoldenCase.created_at)
).all() ).all()
@@ -1013,7 +1019,6 @@ def run_golden_eval(
rule_code = str(manifest.get("rule_code") or "").strip() rule_code = str(manifest.get("rule_code") or "").strip()
if not rule_code: if not rule_code:
raise ValueError("manifest 缺少 rule_code。") raise ValueError("manifest 缺少 rule_code。")
version = body.version or asset.working_version or ""
report = RiskRuleGoldenEvaluator().evaluate_for_rule(db, manifest, rule_code) report = RiskRuleGoldenEvaluator().evaluate_for_rule(db, manifest, rule_code)
return GoldenEvalRead(**report.to_dict()) return GoldenEvalRead(**report.to_dict())
except Exception as exc: except Exception as exc:

View File

@@ -5,16 +5,34 @@ from typing import Annotated
from fastapi import APIRouter, Depends, HTTPException, status from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.api.deps import CurrentUserContext, get_current_user, get_db from app.api.deps import (
CurrentUserContext,
get_current_user,
get_db,
require_platform_admin_user,
)
from app.schemas.expense_application_memory import ( from app.schemas.expense_application_memory import (
ExpenseApplicationMemoryListRead, ExpenseApplicationMemoryListRead,
ExpenseApplicationMemoryRead,
ExpenseApplicationMemoryRevokedRead, ExpenseApplicationMemoryRevokedRead,
ExpenseApplicationOrganizationMemoryCreate,
ExpenseApplicationOrganizationMemoryRevoke,
ExpenseApplicationOrganizationMemoryUpdate,
) )
from app.services.expense_application_memory import ExpenseApplicationMemoryService from app.services.expense_application_memory import ExpenseApplicationMemoryService
from app.services.expense_application_memory_admin import (
ExpenseApplicationOrganizationMemoryService,
OrganizationMemoryConflictError,
OrganizationMemoryNotFoundError,
)
router = APIRouter(prefix="/expense-application-memories") router = APIRouter(prefix="/expense-application-memories")
DbSession = Annotated[Session, Depends(get_db)] DbSession = Annotated[Session, Depends(get_db)]
CurrentUser = Annotated[CurrentUserContext, Depends(get_current_user)] CurrentUser = Annotated[CurrentUserContext, Depends(get_current_user)]
PlatformAdminUser = Annotated[
CurrentUserContext,
Depends(require_platform_admin_user),
]
@router.get( @router.get(
@@ -29,6 +47,105 @@ def list_my_expense_application_memories(
return ExpenseApplicationMemoryService(db).list_current_user_memories(current_user) return ExpenseApplicationMemoryService(db).list_current_user_memories(current_user)
@router.get(
"/organization",
response_model=ExpenseApplicationMemoryListRead,
summary="读取当前租户的企业与部门费用记忆",
)
def list_organization_expense_application_memories(
db: DbSession,
current_user: PlatformAdminUser,
) -> ExpenseApplicationMemoryListRead:
return ExpenseApplicationOrganizationMemoryService(db).list_organization_memories(
current_user
)
@router.post(
"/organization",
response_model=ExpenseApplicationMemoryRead,
status_code=status.HTTP_201_CREATED,
summary="创建企业或部门费用记忆",
)
def create_organization_expense_application_memory(
payload: ExpenseApplicationOrganizationMemoryCreate,
db: DbSession,
current_user: PlatformAdminUser,
) -> ExpenseApplicationMemoryRead:
try:
return ExpenseApplicationOrganizationMemoryService(
db
).create_organization_memory(payload, current_user)
except OrganizationMemoryConflictError as error:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail=str(error),
) from error
except ValueError as error:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_CONTENT,
detail=str(error),
) from error
@router.put(
"/organization/{memory_id}",
response_model=ExpenseApplicationMemoryRead,
summary="换代更新企业或部门费用记忆",
)
def update_organization_expense_application_memory(
memory_id: str,
payload: ExpenseApplicationOrganizationMemoryUpdate,
db: DbSession,
current_user: PlatformAdminUser,
) -> ExpenseApplicationMemoryRead:
try:
return ExpenseApplicationOrganizationMemoryService(
db
).update_organization_memory(memory_id, payload, current_user)
except OrganizationMemoryNotFoundError as error:
raise _organization_memory_not_found() from error
except OrganizationMemoryConflictError as error:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail=str(error),
) from error
except ValueError as error:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_CONTENT,
detail=str(error),
) from error
@router.post(
"/organization/{memory_id}/revoke",
response_model=ExpenseApplicationMemoryRevokedRead,
summary="撤销企业或部门费用记忆",
)
def revoke_organization_expense_application_memory(
memory_id: str,
payload: ExpenseApplicationOrganizationMemoryRevoke,
db: DbSession,
current_user: PlatformAdminUser,
) -> ExpenseApplicationMemoryRevokedRead:
try:
return ExpenseApplicationOrganizationMemoryService(
db
).revoke_organization_memory(memory_id, payload, current_user)
except OrganizationMemoryNotFoundError as error:
raise _organization_memory_not_found() from error
except OrganizationMemoryConflictError as error:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail=str(error),
) from error
except ValueError as error:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_CONTENT,
detail=str(error),
) from error
@router.delete( @router.delete(
"/{memory_id}", "/{memory_id}",
response_model=ExpenseApplicationMemoryRevokedRead, response_model=ExpenseApplicationMemoryRevokedRead,
@@ -49,3 +166,10 @@ def revoke_my_expense_application_memory(
detail="未找到可撤销的个人费用申请记忆。", detail="未找到可撤销的个人费用申请记忆。",
) )
return result return result
def _organization_memory_not_found() -> HTTPException:
return HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="未找到当前租户内可管理的组织费用记忆。",
)

View File

@@ -5,7 +5,7 @@ from typing import Annotated
from fastapi import APIRouter, Depends, HTTPException, Query, status from fastapi import APIRouter, Depends, HTTPException, Query, status
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.api.deps import get_current_user, get_db from app.api.deps import CurrentUserContext, get_current_user, get_db
from app.schemas.common import ErrorResponse from app.schemas.common import ErrorResponse
from app.schemas.risk_observation import ( from app.schemas.risk_observation import (
RiskObservationDashboardRead, RiskObservationDashboardRead,
@@ -16,8 +16,9 @@ from app.schemas.risk_observation import (
) )
from app.services.risk_observations import RiskObservationService from app.services.risk_observations import RiskObservationService
router = APIRouter(prefix="/risk-observations", dependencies=[Depends(get_current_user)]) router = APIRouter(prefix="/risk-observations")
DbSession = Annotated[Session, Depends(get_db)] DbSession = Annotated[Session, Depends(get_db)]
CurrentUser = Annotated[CurrentUserContext, Depends(get_current_user)]
@router.get( @router.get(
@@ -28,6 +29,7 @@ DbSession = Annotated[Session, Depends(get_db)]
) )
def list_risk_observations( def list_risk_observations(
db: DbSession, db: DbSession,
current_user: CurrentUser,
claim_id: Annotated[str | None, Query(max_length=80)] = None, claim_id: Annotated[str | None, Query(max_length=80)] = None,
run_id: Annotated[str | None, Query(max_length=80)] = None, run_id: Annotated[str | None, Query(max_length=80)] = None,
execution_log_id: Annotated[str | None, Query(max_length=80)] = None, execution_log_id: Annotated[str | None, Query(max_length=80)] = None,
@@ -42,6 +44,7 @@ def list_risk_observations(
offset: Annotated[int, Query(ge=0)] = 0, offset: Annotated[int, Query(ge=0)] = 0,
) -> RiskObservationListRead: ) -> RiskObservationListRead:
items, total = RiskObservationService(db).list_observations( items, total = RiskObservationService(db).list_observations(
tenant_id=current_user.tenant_id,
claim_id=claim_id, claim_id=claim_id,
run_id=run_id, run_id=run_id,
execution_log_id=execution_log_id, execution_log_id=execution_log_id,
@@ -63,10 +66,12 @@ def list_risk_observations(
) )
def summarize_risk_observations( def summarize_risk_observations(
db: DbSession, db: DbSession,
current_user: CurrentUser,
window_days: Annotated[int, Query(ge=1, le=365)] = 30, window_days: Annotated[int, Query(ge=1, le=365)] = 30,
limit: Annotated[int, Query(ge=1, le=2000)] = 500, limit: Annotated[int, Query(ge=1, le=2000)] = 500,
) -> RiskObservationDashboardRead: ) -> RiskObservationDashboardRead:
return RiskObservationService(db).summarize_dashboard( return RiskObservationService(db).summarize_dashboard(
tenant_id=current_user.tenant_id,
window_days=window_days, window_days=window_days,
limit=limit, limit=limit,
) )
@@ -78,8 +83,15 @@ def summarize_risk_observations(
summary="查询单据风险观察", summary="查询单据风险观察",
description="按报销单 ID 返回该单据关联的风险观察,供单据详情证据链使用。", description="按报销单 ID 返回该单据关联的风险观察,供单据详情证据链使用。",
) )
def list_claim_risk_observations(claim_id: str, db: DbSession) -> list[RiskObservationRead]: def list_claim_risk_observations(
return RiskObservationService(db).list_claim_observations(claim_id) claim_id: str,
db: DbSession,
current_user: CurrentUser,
) -> list[RiskObservationRead]:
return RiskObservationService(db).list_claim_observations(
claim_id,
tenant_id=current_user.tenant_id,
)
@router.get( @router.get(
@@ -91,8 +103,12 @@ def list_claim_risk_observations(claim_id: str, db: DbSession) -> list[RiskObser
def list_execution_log_risk_observations( def list_execution_log_risk_observations(
execution_log_id: str, execution_log_id: str,
db: DbSession, db: DbSession,
current_user: CurrentUser,
) -> list[RiskObservationRead]: ) -> list[RiskObservationRead]:
return RiskObservationService(db).list_execution_log_observations(execution_log_id) return RiskObservationService(db).list_execution_log_observations(
execution_log_id,
tenant_id=current_user.tenant_id,
)
@router.get( @router.get(
@@ -110,8 +126,12 @@ def list_execution_log_risk_observations(
def get_risk_observation( def get_risk_observation(
observation_key_or_id: str, observation_key_or_id: str,
db: DbSession, db: DbSession,
current_user: CurrentUser,
) -> RiskObservationRead: ) -> RiskObservationRead:
observation = RiskObservationService(db).get_observation(observation_key_or_id) observation = RiskObservationService(db).get_observation(
observation_key_or_id,
tenant_id=current_user.tenant_id,
)
if observation is None: if observation is None:
raise HTTPException( raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, status_code=status.HTTP_404_NOT_FOUND,
@@ -136,9 +156,15 @@ def create_risk_observation_feedback(
observation_key_or_id: str, observation_key_or_id: str,
payload: RiskObservationFeedbackCreate, payload: RiskObservationFeedbackCreate,
db: DbSession, db: DbSession,
current_user: CurrentUser,
) -> RiskObservationFeedbackRead: ) -> RiskObservationFeedbackRead:
try: try:
return RiskObservationService(db).create_feedback(observation_key_or_id, payload) return RiskObservationService(db).create_feedback(
observation_key_or_id,
payload,
tenant_id=current_user.tenant_id,
actor=current_user.name or current_user.username,
)
except LookupError: except LookupError:
raise HTTPException( raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND, status_code=status.HTTP_404_NOT_FOUND,

View File

@@ -79,10 +79,67 @@ MIGRATION_OWNED_TABLES_BY_REVISION: dict[str, frozenset[str]] = {
"workflow_outcomes", "workflow_outcomes",
} }
), ),
"20260716_0007": frozenset(
{
"expense_cases",
"expense_case_links",
"business_events",
"auth_sessions",
"attachment_association_jobs",
"ai_application_preview_decisions",
"ai_decisions",
"ai_decision_feedback",
"memory_entries",
"memory_evidence_links",
"workflow_outcomes",
}
),
"20260716_0008": frozenset(
{
"expense_cases",
"expense_case_links",
"business_events",
"auth_sessions",
"attachment_association_jobs",
"ai_application_preview_decisions",
"ai_decisions",
"ai_decision_feedback",
"memory_entries",
"memory_evidence_links",
"risk_observations",
"risk_observation_feedback",
"few_shot_samples",
"workflow_outcomes",
}
),
"20260716_0009": frozenset(
{
"expense_cases",
"expense_case_links",
"business_events",
"auth_sessions",
"attachment_association_jobs",
"ai_application_preview_decisions",
"ai_decisions",
"ai_decision_feedback",
"memory_entries",
"memory_evidence_links",
"risk_observations",
"risk_observation_feedback",
"few_shot_samples",
"workflow_outcomes",
}
),
} }
if MIGRATION_OWNED_TABLES_BY_REVISION["20260716_0006"] != MIGRATION_OWNED_TABLES: if MIGRATION_OWNED_TABLES_BY_REVISION["20260716_0009"] != MIGRATION_OWNED_TABLES:
raise RuntimeError("latest Alembic revision must own the centralized migration table set") raise RuntimeError("latest Alembic revision must own the centralized migration table set")
# 0008 之前这三张表由旧 bootstrap / 风险服务按需创建。迁移前置检查允许
# 它们作为完整或部分旧资产存在,由 0008 统一收编;其他未来表仍严格拒绝。
LEGACY_ADOPTABLE_HISTORICAL_CASE_TABLES = frozenset(
{"risk_observations", "risk_observation_feedback", "few_shot_samples"}
)
class MigrationPreflightError(RuntimeError): class MigrationPreflightError(RuntimeError):
"""Raised when the database schema cannot be safely advanced by Alembic.""" """Raised when the database schema cannot be safely advanced by Alembic."""
@@ -103,10 +160,11 @@ def _validate_connection(connection: Connection) -> MigrationPreflightState:
owned_tables = table_names & MIGRATION_OWNED_TABLES owned_tables = table_names & MIGRATION_OWNED_TABLES
if "alembic_version" not in table_names: if "alembic_version" not in table_names:
if owned_tables: unsafe_owned_tables = owned_tables - LEGACY_ADOPTABLE_HISTORICAL_CASE_TABLES
if unsafe_owned_tables:
raise MigrationPreflightError( raise MigrationPreflightError(
"unversioned database contains migration-owned tables " "unversioned database contains migration-owned tables "
f"({_format_tables(owned_tables)}); refusing to guess, stamp, or repair" f"({_format_tables(unsafe_owned_tables)}); refusing to guess, stamp, or repair"
) )
return MigrationPreflightState(revision=None, owned_tables=owned_tables) return MigrationPreflightState(revision=None, owned_tables=owned_tables)
@@ -117,10 +175,11 @@ def _validate_connection(connection: Connection) -> MigrationPreflightState:
).scalars() ).scalars()
) )
if not revisions: if not revisions:
if owned_tables: unsafe_owned_tables = owned_tables - LEGACY_ADOPTABLE_HISTORICAL_CASE_TABLES
if unsafe_owned_tables:
raise MigrationPreflightError( raise MigrationPreflightError(
"alembic_version has no recorded revision but migration-owned tables exist " "alembic_version has no recorded revision but migration-owned tables exist "
f"({_format_tables(owned_tables)}); refusing to guess, stamp, or repair" f"({_format_tables(unsafe_owned_tables)}); refusing to guess, stamp, or repair"
) )
return MigrationPreflightState(revision=None, owned_tables=owned_tables) return MigrationPreflightState(revision=None, owned_tables=owned_tables)
@@ -137,9 +196,14 @@ def _validate_connection(connection: Connection) -> MigrationPreflightState:
f"unknown Alembic revision {revision!r}; refusing to run migrations" f"unknown Alembic revision {revision!r}; refusing to run migrations"
) )
if owned_tables != expected_tables: adoptable_tables = (
missing_tables = expected_tables - owned_tables LEGACY_ADOPTABLE_HISTORICAL_CASE_TABLES
unexpected_tables = owned_tables - expected_tables if revision not in {"20260716_0008", "20260716_0009"}
else frozenset()
)
missing_tables = expected_tables - owned_tables
unexpected_tables = owned_tables - expected_tables - adoptable_tables
if missing_tables or unexpected_tables:
raise MigrationPreflightError( raise MigrationPreflightError(
f"migration-owned table set does not match revision {revision}: " f"migration-owned table set does not match revision {revision}: "
f"missing={_format_tables(missing_tables)}; " f"missing={_format_tables(missing_tables)}; "

View File

@@ -16,6 +16,9 @@ MIGRATION_OWNED_TABLES: frozenset[str] = frozenset(
"business_events", "business_events",
"memory_entries", "memory_entries",
"memory_evidence_links", "memory_evidence_links",
"risk_observations",
"risk_observation_feedback",
"few_shot_samples",
"workflow_outcomes", "workflow_outcomes",
} }
) )

View File

@@ -15,6 +15,7 @@ from sqlalchemy import (
String, String,
UniqueConstraint, UniqueConstraint,
func, func,
text,
) )
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
from sqlalchemy.types import JSON from sqlalchemy.types import JSON
@@ -34,7 +35,7 @@ def _candidate_expires_at() -> datetime:
class MemoryEntry(Base): class MemoryEntry(Base):
"""受证据约束、可撤销的个人费用申请记忆。""" """受证据约束、可审计且可撤销的分层费用申请记忆。"""
__tablename__ = "memory_entries" __tablename__ = "memory_entries"
__table_args__ = ( __table_args__ = (
@@ -59,9 +60,33 @@ class MemoryEntry(Base):
name="fk_memory_entries_tenant_superseded_by", name="fk_memory_entries_tenant_superseded_by",
), ),
CheckConstraint( CheckConstraint(
"scope_type = 'user'", "scope_type IN ('user', 'department', 'enterprise')",
name="ck_memory_entries_scope_type", name="ck_memory_entries_scope_type",
), ),
CheckConstraint(
"origin_type IN ('learned', 'admin_managed')",
name="ck_memory_entries_origin_type",
),
CheckConstraint(
"(scope_type = 'user' AND origin_type = 'learned') OR "
"(scope_type IN ('department', 'enterprise') "
"AND origin_type = 'admin_managed')",
name="ck_memory_entries_scope_origin",
),
CheckConstraint(
"scope_type != 'enterprise' OR scope_id = tenant_id",
name="ck_memory_entries_enterprise_scope",
),
CheckConstraint(
"(origin_type = 'learned' AND managed_by IS NULL "
"AND managed_at IS NULL AND management_reason IS NULL) OR "
"(origin_type = 'admin_managed' AND managed_by IS NOT NULL "
"AND length(trim(managed_by)) > 0 AND managed_at IS NOT NULL "
"AND management_reason IS NOT NULL "
"AND length(trim(management_reason)) > 0 "
"AND policy_version IS NOT NULL AND length(trim(policy_version)) > 0)",
name="ck_memory_entries_management_audit",
),
CheckConstraint( CheckConstraint(
"scene = 'travel_application'", "scene = 'travel_application'",
name="ck_memory_entries_scene", name="ck_memory_entries_scene",
@@ -120,6 +145,17 @@ class MemoryEntry(Base):
"superseded_by_id IS NULL OR superseded_by_id != id", "superseded_by_id IS NULL OR superseded_by_id != id",
name="ck_memory_entries_not_self_superseded", name="ck_memory_entries_not_self_superseded",
), ),
CheckConstraint(
"(management_request_id IS NULL AND management_payload_fingerprint IS NULL) OR "
"(management_request_id IS NOT NULL "
"AND management_payload_fingerprint IS NOT NULL)",
name="ck_memory_entries_management_idempotency_pair",
),
CheckConstraint(
"(revoke_request_id IS NULL AND revoke_payload_fingerprint IS NULL) OR "
"(revoke_request_id IS NOT NULL AND revoke_payload_fingerprint IS NOT NULL)",
name="ck_memory_entries_revoke_idempotency_pair",
),
Index( Index(
"ix_memory_entries_scope_lookup", "ix_memory_entries_scope_lookup",
"tenant_id", "tenant_id",
@@ -136,6 +172,31 @@ class MemoryEntry(Base):
"candidate_expires_at", "candidate_expires_at",
"active_expires_at", "active_expires_at",
), ),
Index(
"uq_memory_entries_management_request",
"tenant_id",
"management_request_id",
unique=True,
),
Index(
"uq_memory_entries_revoke_request",
"tenant_id",
"revoke_request_id",
unique=True,
),
Index(
"uq_memory_entries_active_scope",
"tenant_id",
"scope_type",
"scope_id",
"scene",
"field_key",
unique=True,
postgresql_where=text(
"status = 'active' "
"AND scope_type IN ('department', 'enterprise')"
),
).ddl_if(dialect="postgresql"),
) )
id: Mapped[str] = mapped_column(String(36), primary_key=True, default=_new_id) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=_new_id)
@@ -147,6 +208,16 @@ class MemoryEntry(Base):
server_default="user", server_default="user",
) )
scope_id: Mapped[str] = mapped_column(String(120), nullable=False) scope_id: Mapped[str] = mapped_column(String(120), nullable=False)
origin_type: Mapped[str] = mapped_column(
String(24),
nullable=False,
default="learned",
server_default="learned",
)
managed_by: Mapped[str | None] = mapped_column(String(255), nullable=True)
managed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
management_reason: Mapped[str | None] = mapped_column(String(255), nullable=True)
policy_version: Mapped[str | None] = mapped_column(String(64), nullable=True)
scene: Mapped[str] = mapped_column( scene: Mapped[str] = mapped_column(
String(50), String(50),
nullable=False, nullable=False,
@@ -212,6 +283,16 @@ class MemoryEntry(Base):
revoked_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True) revoked_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
revoked_reason: Mapped[str | None] = mapped_column(String(255), nullable=True) revoked_reason: Mapped[str | None] = mapped_column(String(255), nullable=True)
superseded_by_id: Mapped[str | None] = mapped_column(String(36), nullable=True) superseded_by_id: Mapped[str | None] = mapped_column(String(36), nullable=True)
management_request_id: Mapped[str | None] = mapped_column(String(120), nullable=True)
management_payload_fingerprint: Mapped[str | None] = mapped_column(
String(80),
nullable=True,
)
revoke_request_id: Mapped[str | None] = mapped_column(String(120), nullable=True)
revoke_payload_fingerprint: Mapped[str | None] = mapped_column(
String(80),
nullable=True,
)
created_at: Mapped[datetime] = mapped_column( created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), DateTime(timezone=True),
nullable=False, nullable=False,

View File

@@ -4,7 +4,7 @@ import uuid
from datetime import datetime from datetime import datetime
from typing import Any from typing import Any
from sqlalchemy import DateTime, ForeignKey, Index, String, Text, func from sqlalchemy import DateTime, ForeignKey, Index, String, Text, UniqueConstraint, func
from sqlalchemy.orm import Mapped, mapped_column from sqlalchemy.orm import Mapped, mapped_column
from sqlalchemy.types import JSON from sqlalchemy.types import JSON
@@ -20,12 +20,26 @@ class FewShotSample(Base):
__tablename__ = "few_shot_samples" __tablename__ = "few_shot_samples"
__table_args__ = ( __table_args__ = (
UniqueConstraint(
"tenant_id",
"sample_key",
name="uq_few_shot_samples_tenant_key",
),
Index(
"ix_few_shot_samples_tenant_rule_lookup",
"tenant_id",
"scene",
"policy_ref",
"rule_version",
"status",
),
Index("ix_few_shot_samples_scene_label", "scene", "label"), Index("ix_few_shot_samples_scene_label", "scene", "label"),
Index("ix_few_shot_samples_domain_risk_type", "domain", "risk_type"), Index("ix_few_shot_samples_domain_risk_type", "domain", "risk_type"),
) )
id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
sample_key: Mapped[str] = mapped_column(String(160), unique=True, index=True) tenant_id: Mapped[str] = mapped_column(String(64), default="default", index=True)
sample_key: Mapped[str] = mapped_column(String(160), index=True)
source_observation_id: Mapped[str | None] = mapped_column( source_observation_id: Mapped[str | None] = mapped_column(
ForeignKey("risk_observations.id"), ForeignKey("risk_observations.id"),
nullable=True, nullable=True,
@@ -33,6 +47,8 @@ class FewShotSample(Base):
) )
scene: Mapped[str] = mapped_column(String(50), default="risk_rule_generation", index=True) scene: Mapped[str] = mapped_column(String(50), default="risk_rule_generation", index=True)
policy_ref: Mapped[str] = mapped_column(String(160), default="", index=True)
rule_version: Mapped[str] = mapped_column(String(80), default="", index=True)
domain: Mapped[str] = mapped_column(String(50), default="", index=True) domain: Mapped[str] = mapped_column(String(50), default="", index=True)
risk_type: Mapped[str] = mapped_column(String(80), default="", index=True) risk_type: Mapped[str] = mapped_column(String(80), default="", index=True)
risk_level: Mapped[str] = mapped_column(String(20), default="") risk_level: Mapped[str] = mapped_column(String(20), default="")
@@ -45,7 +61,9 @@ class FewShotSample(Base):
vector_id: Mapped[str | None] = mapped_column(String(100), nullable=True) vector_id: Mapped[str | None] = mapped_column(String(100), nullable=True)
status: Mapped[str] = mapped_column(String(20), default="active", index=True) status: Mapped[str] = mapped_column(String(20), default="active", index=True)
created_at: Mapped[datetime] = mapped_column(DateTime, default=func.now(), server_default=func.now()) created_at: Mapped[datetime] = mapped_column(
DateTime, default=func.now(), server_default=func.now()
)
updated_at: Mapped[datetime] = mapped_column( updated_at: Mapped[datetime] = mapped_column(
DateTime, DateTime,
default=func.now(), default=func.now(),

View File

@@ -4,7 +4,17 @@ import uuid
from datetime import datetime from datetime import datetime
from typing import Any from typing import Any
from sqlalchemy import DateTime, Float, ForeignKey, Index, Integer, String, Text, func from sqlalchemy import (
DateTime,
Float,
ForeignKey,
Index,
Integer,
String,
Text,
UniqueConstraint,
func,
)
from sqlalchemy.orm import Mapped, mapped_column, relationship from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.types import JSON from sqlalchemy.types import JSON
@@ -14,18 +24,25 @@ from app.db.base_class import Base
class RiskObservation(Base): class RiskObservation(Base):
__tablename__ = "risk_observations" __tablename__ = "risk_observations"
__table_args__ = ( __table_args__ = (
UniqueConstraint(
"tenant_id",
"observation_key",
name="uq_risk_observations_tenant_key",
),
Index("ix_risk_observations_tenant_status", "tenant_id", "status", "created_at"),
Index("ix_risk_observations_subject", "subject_type", "subject_key"), Index("ix_risk_observations_subject", "subject_type", "subject_key"),
Index("ix_risk_observations_signal_level", "risk_signal", "risk_level"), Index("ix_risk_observations_signal_level", "risk_signal", "risk_level"),
Index("ix_risk_observations_status_created", "status", "created_at"), Index("ix_risk_observations_status_created", "status", "created_at"),
) )
id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
observation_key: Mapped[str] = mapped_column(String(160), unique=True, index=True) tenant_id: Mapped[str] = mapped_column(String(64), default="default", index=True)
observation_key: Mapped[str] = mapped_column(String(160), index=True)
subject_type: Mapped[str] = mapped_column(String(50), index=True) subject_type: Mapped[str] = mapped_column(String(50), index=True)
subject_key: Mapped[str] = mapped_column(String(160), index=True) subject_key: Mapped[str] = mapped_column(String(160), index=True)
subject_label: Mapped[str] = mapped_column(String(160), default="") subject_label: Mapped[str] = mapped_column(String(160), default="")
claim_id: Mapped[str | None] = mapped_column( claim_id: Mapped[str | None] = mapped_column(
ForeignKey("expense_claims.id"), String(36),
nullable=True, nullable=True,
index=True, index=True,
) )
@@ -66,7 +83,12 @@ class RiskObservation(Base):
onupdate=func.now(), onupdate=func.now(),
) )
claim = relationship("ExpenseClaim", foreign_keys=[claim_id]) claim = relationship(
"ExpenseClaim",
primaryjoin="foreign(RiskObservation.claim_id) == ExpenseClaim.id",
foreign_keys=[claim_id],
viewonly=True,
)
feedback_items = relationship( feedback_items = relationship(
"RiskObservationFeedback", "RiskObservationFeedback",
back_populates="observation", back_populates="observation",

View File

@@ -1,10 +1,18 @@
from __future__ import annotations from __future__ import annotations
from datetime import datetime from datetime import datetime
from typing import Literal
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
class ExpenseApplicationMemoryConflict(BaseModel):
scope_type: Literal["user", "department", "enterprise"]
scope_label: str
priority: int
reason: Literal["same_priority_conflict", "lower_priority_overridden"]
class ExpenseApplicationMemoryApplication(BaseModel): class ExpenseApplicationMemoryApplication(BaseModel):
memory_id: str memory_id: str
field_key: str = "transport_mode" field_key: str = "transport_mode"
@@ -15,7 +23,14 @@ class ExpenseApplicationMemoryApplication(BaseModel):
evidence_count: int = 0 evidence_count: int = 0
approved_evidence_count: int = 0 approved_evidence_count: int = 0
confidence: float = 0.0 confidence: float = 0.0
effective_confidence: float = 0.0
expires_at: datetime | None = None expires_at: datetime | None = None
scope_type: Literal["user", "department", "enterprise"] = "user"
scope_id: str = ""
scope_label: str = "个人偏好"
priority: int = 100
conflicts: list[ExpenseApplicationMemoryConflict] = Field(default_factory=list)
can_revoke: bool = True
message: str = "已按可信历史记忆预填常用出行方式,可继续修改。" message: str = "已按可信历史记忆预填常用出行方式,可继续修改。"
@@ -38,6 +53,12 @@ class ExpenseApplicationMemoryRead(BaseModel):
field_key: str field_key: str
value: str = "" value: str = ""
status: str status: str
scope_type: Literal["user", "department", "enterprise"] = "user"
scope_id: str = ""
scope_label: str = "个人偏好"
source: str = "verified_user_history"
origin_type: Literal["learned", "admin_managed"] = "learned"
generation: int = 1
evidence_count: int = 0 evidence_count: int = 0
approved_evidence_count: int = 0 approved_evidence_count: int = 0
confidence: float = 0.0 confidence: float = 0.0
@@ -50,6 +71,11 @@ class ExpenseApplicationMemoryRead(BaseModel):
suppressed_at: datetime | None = None suppressed_at: datetime | None = None
revoked_at: datetime | None = None revoked_at: datetime | None = None
revoked_reason: str = "" revoked_reason: str = ""
managed_by: str = ""
managed_at: datetime | None = None
management_reason: str = ""
superseded_by_id: str = ""
can_revoke: bool = True
created_at: datetime | None = None created_at: datetime | None = None
updated_at: datetime | None = None updated_at: datetime | None = None
@@ -62,3 +88,26 @@ class ExpenseApplicationMemoryRevokedRead(BaseModel):
memory_id: str memory_id: str
status: str = "revoked" status: str = "revoked"
revoked_at: datetime revoked_at: datetime
class ExpenseApplicationOrganizationMemoryCreate(BaseModel):
scope_type: Literal["department", "enterprise"]
scope_id: str | None = Field(default=None, max_length=120)
value: Literal["飞机", "火车", "轮船"]
expires_in_days: int = Field(default=180, ge=30, le=365)
reason: str = Field(min_length=1, max_length=255)
request_id: str = Field(min_length=8, max_length=120)
class ExpenseApplicationOrganizationMemoryUpdate(BaseModel):
value: Literal["飞机", "火车", "轮船"] | None = None
expires_in_days: int | None = Field(default=None, ge=30, le=365)
expected_generation: int = Field(ge=1)
reason: str = Field(min_length=1, max_length=255)
request_id: str = Field(min_length=8, max_length=120)
class ExpenseApplicationOrganizationMemoryRevoke(BaseModel):
expected_generation: int = Field(ge=1)
reason: str = Field(min_length=1, max_length=255)
request_id: str = Field(min_length=8, max_length=120)

View File

@@ -161,6 +161,19 @@ class ExpenseClaimPreReviewFindingRead(BaseModel):
remediation: ExpenseClaimPreReviewRemediationRead remediation: ExpenseClaimPreReviewRemediationRead
class ExpenseClaimHistoricalCaseEvidenceRead(BaseModel):
label: Literal["confirmed", "false_positive"]
label_text: str
advisory_only: Literal[True] = True
score: float = 0.0
scene_code: str = ""
policy_ref: str = ""
rule_version: str = ""
version_status: Literal["matched", "stale"] = "matched"
stale: bool = False
summary: str
class ExpenseClaimPreReviewRead(BaseModel): class ExpenseClaimPreReviewRead(BaseModel):
review_id: str review_id: str
input_fingerprint: str input_fingerprint: str
@@ -173,6 +186,9 @@ class ExpenseClaimPreReviewRead(BaseModel):
blocking_count: int = 0 blocking_count: int = 0
message: str message: str
findings: list[ExpenseClaimPreReviewFindingRead] = Field(default_factory=list) findings: list[ExpenseClaimPreReviewFindingRead] = Field(default_factory=list)
historical_case_evidence: list[ExpenseClaimHistoricalCaseEvidenceRead] = Field(
default_factory=list
)
class ExpenseClaimSubmitPayload(BaseModel): class ExpenseClaimSubmitPayload(BaseModel):

View File

@@ -44,6 +44,7 @@ class RiskObservationRead(BaseModel):
model_config = ConfigDict(from_attributes=True) model_config = ConfigDict(from_attributes=True)
id: str id: str
tenant_id: str
observation_key: str observation_key: str
subject_type: str subject_type: str
subject_key: str subject_key: str
@@ -100,7 +101,11 @@ class RiskObservationListRead(BaseModel):
class RiskObservationFeedbackCreate(BaseModel): class RiskObservationFeedbackCreate(BaseModel):
feedback_type: RiskObservationFeedbackType feedback_type: RiskObservationFeedbackType
action: str | None = Field(default=None, max_length=50) action: str | None = Field(default=None, max_length=50)
actor: str | None = Field(default=None, max_length=100) actor: str | None = Field(
default=None,
max_length=100,
description="兼容字段;服务端始终以当前认证用户覆盖该值。",
)
comment: str | None = Field(default=None, max_length=1000) comment: str | None = Field(default=None, max_length=1000)
payload_json: dict[str, Any] = Field(default_factory=dict) payload_json: dict[str, Any] = Field(default_factory=dict)

View File

@@ -15,6 +15,7 @@ from app.schemas.agent_asset import (
from app.services.agent_asset_rule_library import AgentAssetRuleLibraryManager from app.services.agent_asset_rule_library import AgentAssetRuleLibraryManager
from app.services.agent_asset_spreadsheet import RISK_RULES_LIBRARY from app.services.agent_asset_spreadsheet import RISK_RULES_LIBRARY
from app.services.audit import AuditLogService from app.services.audit import AuditLogService
from app.services.risk_rule_dsl_validator import validate_risk_rule_draft
from app.services.risk_rule_generation import ( from app.services.risk_rule_generation import (
BUSINESS_DOMAIN_LABELS, BUSINESS_DOMAIN_LABELS,
EXPENSE_BUSINESS_STAGE_LABELS, EXPENSE_BUSINESS_STAGE_LABELS,
@@ -22,7 +23,6 @@ from app.services.risk_rule_generation import (
RiskRuleGenerationService, RiskRuleGenerationService,
) )
from app.services.risk_rule_generation_markdown import build_risk_rule_version_markdown from app.services.risk_rule_generation_markdown import build_risk_rule_version_markdown
from app.services.risk_rule_dsl_validator import validate_risk_rule_draft
from app.services.risk_rule_scoring import apply_risk_score_to_draft, calculate_risk_rule_score from app.services.risk_rule_scoring import apply_risk_score_to_draft, calculate_risk_rule_score
from app.services.runtime_chat import RuntimeChatService from app.services.runtime_chat import RuntimeChatService
@@ -52,6 +52,7 @@ class AgentAssetRiskRuleRegenerationService:
asset_id: str, asset_id: str,
body: AgentAssetRiskRuleRegenerateRequest, body: AgentAssetRiskRuleRegenerateRequest,
*, *,
tenant_id: str | None = None,
actor: str, actor: str,
request_id: str | None = None, request_id: str | None = None,
) -> AgentAsset: ) -> AgentAsset:
@@ -60,12 +61,14 @@ class AgentAssetRiskRuleRegenerationService:
return self._regenerate_revision_draft( return self._regenerate_revision_draft(
asset, asset,
body, body,
tenant_id=tenant_id,
actor=actor, actor=actor,
request_id=request_id, request_id=request_id,
) )
return self._regenerate_unpublished_draft( return self._regenerate_unpublished_draft(
asset, asset,
body, body,
tenant_id=tenant_id,
actor=actor, actor=actor,
request_id=request_id, request_id=request_id,
) )
@@ -75,6 +78,7 @@ class AgentAssetRiskRuleRegenerationService:
asset: AgentAsset, asset: AgentAsset,
body: AgentAssetRiskRuleRegenerateRequest, body: AgentAssetRiskRuleRegenerateRequest,
*, *,
tenant_id: str | None,
actor: str, actor: str,
request_id: str | None, request_id: str | None,
) -> AgentAsset: ) -> AgentAsset:
@@ -84,7 +88,12 @@ class AgentAssetRiskRuleRegenerationService:
before = self._snapshot(asset) before = self._snapshot(asset)
config = dict(asset.config_json or {}) config = dict(asset.config_json or {})
request = self._build_generation_request(asset, config, body.model_dump(exclude_unset=True)) request = self._build_generation_request(asset, config, body.model_dump(exclude_unset=True))
payload, risk_score = self._compile_payload(request, actor=actor, created_at=asset.created_at) payload, risk_score = self._compile_payload(
request,
tenant_id=tenant_id,
actor=actor,
created_at=asset.created_at,
)
rule_code = self._stable_rule_code(asset, payload) rule_code = self._stable_rule_code(asset, payload)
payload["rule_code"] = rule_code payload["rule_code"] = rule_code
file_name = f"{rule_code}.json" file_name = f"{rule_code}.json"
@@ -104,6 +113,7 @@ class AgentAssetRiskRuleRegenerationService:
actor=actor, actor=actor,
) )
config.update(self._config_from_payload(payload, risk_score=risk_score, request=request)) config.update(self._config_from_payload(payload, risk_score=risk_score, request=request))
config["tenant_id"] = str(tenant_id or "").strip()
config.update( config.update(
{ {
"generation_status": "completed", "generation_status": "completed",
@@ -138,6 +148,7 @@ class AgentAssetRiskRuleRegenerationService:
asset: AgentAsset, asset: AgentAsset,
body: AgentAssetRiskRuleRegenerateRequest, body: AgentAssetRiskRuleRegenerateRequest,
*, *,
tenant_id: str | None,
actor: str, actor: str,
request_id: str | None, request_id: str | None,
) -> AgentAsset: ) -> AgentAsset:
@@ -151,7 +162,12 @@ class AgentAssetRiskRuleRegenerationService:
body.model_dump(exclude_unset=True), body.model_dump(exclude_unset=True),
base=revision.get("generation_request") if isinstance(revision.get("generation_request"), dict) else {}, base=revision.get("generation_request") if isinstance(revision.get("generation_request"), dict) else {},
) )
payload, risk_score = self._compile_payload(request, actor=actor, created_at=datetime.now(UTC)) payload, risk_score = self._compile_payload(
request,
tenant_id=tenant_id,
actor=actor,
created_at=datetime.now(UTC),
)
payload["rule_code"] = str(asset.code or payload["rule_code"]).strip() payload["rule_code"] = str(asset.code or payload["rule_code"]).strip()
payload["enabled"] = False payload["enabled"] = False
payload.setdefault("metadata", {})["revision_version"] = revision_version payload.setdefault("metadata", {})["revision_version"] = revision_version
@@ -184,6 +200,7 @@ class AgentAssetRiskRuleRegenerationService:
} }
) )
config["revision_draft"] = revision config["revision_draft"] = revision
config["tenant_id"] = str(tenant_id or "").strip()
config["last_operation"] = { config["last_operation"] = {
"action": "regenerate_revision", "action": "regenerate_revision",
"actor": actor, "actor": actor,
@@ -216,6 +233,7 @@ class AgentAssetRiskRuleRegenerationService:
self, self,
request: dict[str, Any], request: dict[str, Any],
*, *,
tenant_id: str | None,
actor: str, actor: str,
created_at: datetime | None, created_at: datetime | None,
) -> tuple[dict[str, Any], dict[str, Any]]: ) -> tuple[dict[str, Any], dict[str, Any]]:
@@ -230,6 +248,7 @@ class AgentAssetRiskRuleRegenerationService:
expense_category_label = EXPENSE_RISK_CATEGORY_LABELS.get(expense_category or "", "") expense_category_label = EXPENSE_RISK_CATEGORY_LABELS.get(expense_category or "", "")
fields = self.generator._resolve_fields(natural_language, domain=domain) fields = self.generator._resolve_fields(natural_language, domain=domain)
draft = self.generator._compile_with_model( draft = self.generator._compile_with_model(
tenant_id=tenant_id,
natural_language=natural_language, natural_language=natural_language,
domain=domain, domain=domain,
business_stage=business_stage, business_stage=business_stage,

View File

@@ -53,6 +53,7 @@ class AuthenticatedUser:
avatar: str avatar: str
is_admin: bool = False is_admin: bool = False
employee_id: str | None = None employee_id: str | None = None
department_id: str | None = None
tenant_id: str = "default" tenant_id: str = "default"
@@ -116,7 +117,7 @@ class AuthService:
} }
if auth_session.username.strip().casefold() not in allowed_identifiers: if auth_session.username.strip().casefold() not in allowed_identifiers:
return None return None
return self._build_admin_user(record) return self._restore_session_scope(self._build_admin_user(record), auth_session)
if auth_session.principal_type != "employee": if auth_session.principal_type != "employee":
return None return None
@@ -133,7 +134,17 @@ class AuthService:
employee = self.db.execute(stmt).scalars().first() employee = self.db.execute(stmt).scalars().first()
if employee is None or employee.employment_status == "停用": if employee is None or employee.employment_status == "停用":
return None return None
return self._build_employee_user(employee) return self._restore_session_scope(self._build_employee_user(employee), auth_session)
@staticmethod
def _restore_session_scope(
user: AuthenticatedUser,
auth_session: AuthSession,
) -> AuthenticatedUser:
"""会话恢复时以签发并认证过的会话租户为准,禁止回落到默认租户。"""
user.tenant_id = str(auth_session.tenant_id or "default").strip() or "default"
return user
def get_user_snapshot(self, identifier: str) -> AuthUserRead | None: def get_user_snapshot(self, identifier: str) -> AuthUserRead | None:
normalized = identifier.strip() normalized = identifier.strip()
@@ -249,6 +260,7 @@ class AuthService:
avatar=(employee.name or "?")[:1].upper(), avatar=(employee.name or "?")[:1].upper(),
is_admin=False, is_admin=False,
employee_id=employee.id, employee_id=employee.id,
department_id=employee.organization_unit_id,
) )
@staticmethod @staticmethod

View File

@@ -24,6 +24,12 @@ from app.schemas.expense_application_memory import (
from app.services.expense_application_memory_evidence import ( from app.services.expense_application_memory_evidence import (
ExpenseApplicationMemoryEvidenceValidator, ExpenseApplicationMemoryEvidenceValidator,
) )
from app.services.expense_application_memory_resolution import (
ExpenseApplicationMemoryResolution,
ExpenseApplicationMemoryResolver,
memory_scope_label,
memory_source,
)
from app.services.expense_application_snapshot import hmac_fingerprint from app.services.expense_application_snapshot import hmac_fingerprint
from app.services.expense_cases import ExpenseCaseService from app.services.expense_cases import ExpenseCaseService
@@ -191,16 +197,16 @@ class ExpenseApplicationMemoryService:
try: try:
with self.db.begin_nested(): with self.db.begin_nested():
entry = self._resolve_active_entry(current_user) resolution = self._resolve_active_entry(current_user)
if entry is None: if resolution is None:
return [] return []
value = self._entry_value(entry) application = resolution.to_application(current_user)
if not value: if resolution.winner is None:
return [] return [application]
facts[MEMORY_FIELD_KEY] = value facts[MEMORY_FIELD_KEY] = resolution.value
return [self._build_application(entry, value)] return [application]
except Exception: except Exception:
logger.warning("个人出行方式记忆读取失败,本轮预览不应用记忆。", exc_info=True) logger.warning("出行方式分层记忆读取失败,本轮预览不应用记忆。", exc_info=True)
return [] return []
def learning_receipts_for_preview_decision( def learning_receipts_for_preview_decision(
@@ -255,7 +261,7 @@ class ExpenseApplicationMemoryService:
self._refresh_entry_metrics(entry, now=now, allow_activation=False) self._refresh_entry_metrics(entry, now=now, allow_activation=False)
self.db.commit() self.db.commit()
return ExpenseApplicationMemoryListRead( return ExpenseApplicationMemoryListRead(
items=[self._serialize_entry(entry) for entry in entries] items=[self._serialize_entry(entry, current_user) for entry in entries]
) )
def revoke_current_user_memory( def revoke_current_user_memory(
@@ -293,38 +299,12 @@ class ExpenseApplicationMemoryService:
def _resolve_active_entry( def _resolve_active_entry(
self, self,
current_user: CurrentUserContext, current_user: CurrentUserContext,
) -> MemoryEntry | None: ) -> ExpenseApplicationMemoryResolution | None:
tenant_id = ExpenseCaseService.normalize_tenant_id(current_user.tenant_id) return ExpenseApplicationMemoryResolver(self.db).resolve(
scope_id = self._scope_id(current_user) current_user,
now = datetime.now(UTC) user_scope_id=self._scope_id(current_user),
entries = list( refresh_user_entry=self._refresh_entry_metrics,
self.db.scalars(
select(MemoryEntry)
.where(
MemoryEntry.tenant_id == tenant_id,
MemoryEntry.scope_type == MEMORY_SCOPE_TYPE,
MemoryEntry.scope_id == scope_id,
MemoryEntry.scene == MEMORY_SCENE,
MemoryEntry.field_key == MEMORY_FIELD_KEY,
MemoryEntry.status.in_(["candidate", "active"]),
)
.order_by(MemoryEntry.last_evidence_at.desc(), MemoryEntry.generation.desc())
.with_for_update()
).all()
) )
active_entry = next((entry for entry in entries if entry.status == "active"), None)
if active_entry is not None:
self._refresh_entry_metrics(active_entry, now=now)
if active_entry.status == "active":
return active_entry
for entry in entries:
if entry.status != "candidate":
continue
self._refresh_entry_metrics(entry, now=now)
if entry.status == "active":
return entry
return None
def _refresh_entry_metrics( def _refresh_entry_metrics(
self, self,
@@ -522,6 +502,7 @@ class ExpenseApplicationMemoryService:
scope_id=scope_id, scope_id=scope_id,
scene=MEMORY_SCENE, scene=MEMORY_SCENE,
field_key=MEMORY_FIELD_KEY, field_key=MEMORY_FIELD_KEY,
origin_type="learned",
generation=generation, generation=generation,
value_json={"value": value}, value_json={"value": value},
value_fingerprint=value_fingerprint, value_fingerprint=value_fingerprint,
@@ -531,6 +512,7 @@ class ExpenseApplicationMemoryService:
confidence=Decimal("0"), confidence=Decimal("0"),
candidate_expires_at=now + MEMORY_CANDIDATE_TTL, candidate_expires_at=now + MEMORY_CANDIDATE_TTL,
last_evidence_at=now, last_evidence_at=now,
policy_version=MEMORY_POLICY_VERSION,
) )
def _approved_case_ids( def _approved_case_ids(
@@ -623,20 +605,6 @@ class ExpenseApplicationMemoryService:
normalized = expires_at if expires_at.tzinfo is not None else expires_at.replace(tzinfo=UTC) normalized = expires_at if expires_at.tzinfo is not None else expires_at.replace(tzinfo=UTC)
return normalized <= now return normalized <= now
@staticmethod
def _build_application(
entry: MemoryEntry,
value: str,
) -> ExpenseApplicationMemoryApplication:
return ExpenseApplicationMemoryApplication(
memory_id=entry.id,
value=value,
evidence_count=int(entry.evidence_count or 0),
approved_evidence_count=int(entry.approved_evidence_count or 0),
confidence=float(entry.confidence or 0),
expires_at=entry.active_expires_at,
)
@staticmethod @staticmethod
def _build_receipt( def _build_receipt(
entry: MemoryEntry, entry: MemoryEntry,
@@ -666,7 +634,11 @@ class ExpenseApplicationMemoryService:
) )
@staticmethod @staticmethod
def _serialize_entry(entry: MemoryEntry) -> ExpenseApplicationMemoryRead: def _serialize_entry(
entry: MemoryEntry,
current_user: CurrentUserContext,
) -> ExpenseApplicationMemoryRead:
scope_type = str(entry.scope_type or "user")
value = ( value = (
"" ""
if entry.status == "revoked" if entry.status == "revoked"
@@ -678,11 +650,20 @@ class ExpenseApplicationMemoryService:
field_key=entry.field_key, field_key=entry.field_key,
value=value, value=value,
status=entry.status, status=entry.status,
scope_type=scope_type,
scope_id=str(entry.scope_id or ""),
scope_label=memory_scope_label(entry, current_user),
source=memory_source(scope_type),
origin_type=str(getattr(entry, "origin_type", "learned") or "learned"),
generation=int(entry.generation or 1),
evidence_count=int(entry.evidence_count or 0), evidence_count=int(entry.evidence_count or 0),
approved_evidence_count=int(entry.approved_evidence_count or 0), approved_evidence_count=int(entry.approved_evidence_count or 0),
confidence=float(entry.confidence or 0), confidence=float(entry.confidence or 0),
activation_threshold=MEMORY_ACTIVATION_THRESHOLD, activation_threshold=MEMORY_ACTIVATION_THRESHOLD,
policy_version=MEMORY_POLICY_VERSION, policy_version=str(
getattr(entry, "policy_version", MEMORY_POLICY_VERSION)
or MEMORY_POLICY_VERSION
),
valid_from=entry.activated_at or entry.created_at, valid_from=entry.activated_at or entry.created_at,
expires_at=( expires_at=(
entry.active_expires_at entry.active_expires_at
@@ -694,6 +675,11 @@ class ExpenseApplicationMemoryService:
suppressed_at=entry.suppressed_at, suppressed_at=entry.suppressed_at,
revoked_at=entry.revoked_at, revoked_at=entry.revoked_at,
revoked_reason=str(entry.revoked_reason or ""), revoked_reason=str(entry.revoked_reason or ""),
managed_by=str(getattr(entry, "managed_by", "") or ""),
managed_at=getattr(entry, "managed_at", None),
management_reason=str(getattr(entry, "management_reason", "") or ""),
superseded_by_id=str(entry.superseded_by_id or ""),
can_revoke=scope_type == "user",
created_at=entry.created_at, created_at=entry.created_at,
updated_at=entry.updated_at, updated_at=entry.updated_at,
) )

View File

@@ -0,0 +1,761 @@
from __future__ import annotations
import uuid
from datetime import UTC, datetime, timedelta
from decimal import Decimal
from sqlalchemy import func, select
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from app.api.deps import CurrentUserContext
from app.models.ai_memory import MemoryEntry
from app.models.organization import OrganizationUnit
from app.schemas.expense_application_memory import (
ExpenseApplicationMemoryListRead,
ExpenseApplicationMemoryRead,
ExpenseApplicationMemoryRevokedRead,
ExpenseApplicationOrganizationMemoryCreate,
ExpenseApplicationOrganizationMemoryRevoke,
ExpenseApplicationOrganizationMemoryUpdate,
)
from app.services.expense_application_memory_resolution import memory_source
from app.services.expense_application_snapshot import hmac_fingerprint
from app.services.expense_cases import ExpenseCaseService
from app.services.organization_memory_locks import (
organization_memory_operation_locks,
organization_memory_request_lock_key,
organization_memory_scope_lock_key,
)
MEMORY_SCENE = "travel_application"
MEMORY_FIELD_KEY = "transport_mode"
ORGANIZATION_SCOPE_TYPES = {"department", "enterprise"}
ORGANIZATION_MEMORY_POLICY_VERSION = "expense_application_transport_org_memory.v1"
SUPPORTED_TRANSPORT_VALUES = {"飞机", "火车", "轮船"}
class OrganizationMemoryNotFoundError(LookupError):
pass
class OrganizationMemoryConflictError(ValueError):
pass
class ExpenseApplicationOrganizationMemoryService:
"""由平台管理员显式维护企业/部门出行方式记忆。"""
def __init__(self, db: Session) -> None:
self.db = db
def list_organization_memories(
self,
current_user: CurrentUserContext,
) -> ExpenseApplicationMemoryListRead:
self._require_admin(current_user)
tenant_id = ExpenseCaseService.normalize_tenant_id(current_user.tenant_id)
now = datetime.now(UTC)
entries = list(
self.db.scalars(
select(MemoryEntry)
.where(
MemoryEntry.tenant_id == tenant_id,
MemoryEntry.scope_type.in_(ORGANIZATION_SCOPE_TYPES),
MemoryEntry.scene == MEMORY_SCENE,
MemoryEntry.field_key == MEMORY_FIELD_KEY,
)
.order_by(
MemoryEntry.scope_type.asc(),
MemoryEntry.scope_id.asc(),
MemoryEntry.generation.desc(),
)
).all()
)
for entry in entries:
if entry.status == "active" and self._is_expired(entry, now):
entry.status = "expired"
entry.expired_at = now
department_names = self._department_names(entries)
self.db.commit()
return ExpenseApplicationMemoryListRead(
items=[self._serialize_entry(entry, department_names) for entry in entries]
)
def create_organization_memory(
self,
payload: ExpenseApplicationOrganizationMemoryCreate,
current_user: CurrentUserContext,
) -> ExpenseApplicationMemoryRead:
self._require_admin(current_user)
tenant_id = ExpenseCaseService.normalize_tenant_id(current_user.tenant_id)
scope_id = self._validate_scope(
tenant_id=tenant_id,
scope_type=payload.scope_type,
scope_id=payload.scope_id,
)
request_id = self._normalize_request_id(payload.request_id)
reason = self._normalize_reason(payload.reason)
request_fingerprint = self._request_fingerprint(
operation="create",
target=f"{payload.scope_type}:{scope_id}",
payload={
"scope_type": payload.scope_type,
"scope_id": scope_id,
"value": payload.value,
"expires_in_days": payload.expires_in_days,
"reason": reason,
},
)
now = datetime.now(UTC)
entry_id = self._request_entry_id(
tenant_id=tenant_id,
operation="create",
target=f"{payload.scope_type}:{scope_id}",
request_id=request_id,
)
expires_at = now + timedelta(days=payload.expires_in_days)
try:
with organization_memory_operation_locks(
self.db,
self._scope_lock_key(
tenant_id=tenant_id,
scope_type=payload.scope_type,
scope_id=scope_id,
),
self._request_lock_key(tenant_id, request_id),
):
replay = self._management_request_replay(
tenant_id=tenant_id,
request_id=request_id,
request_fingerprint=request_fingerprint,
)
if replay is not None:
result = self._serialize_entry(
replay,
self._department_names([replay]),
)
self.db.commit()
return result
active = self._get_scope_active_entry(
tenant_id=tenant_id,
scope_type=payload.scope_type,
scope_id=scope_id,
)
if active is not None:
if self._is_expired(active, now):
active.status = "expired"
active.expired_at = now
self.db.flush()
else:
raise OrganizationMemoryConflictError(
"该组织范围已有生效记忆,请使用更新操作换代。"
)
entry = self._new_active_generation(
entry_id=entry_id,
tenant_id=tenant_id,
scope_type=payload.scope_type,
scope_id=scope_id,
value=payload.value,
expires_at=expires_at,
reason=reason,
request_id=request_id,
request_fingerprint=request_fingerprint,
expected_active_ids=set(),
current_user=current_user,
now=now,
)
self.db.commit()
except OrganizationMemoryConflictError:
self.db.rollback()
raise
except IntegrityError as error:
self.db.rollback()
replay = self._management_request_replay(
tenant_id=tenant_id,
request_id=request_id,
request_fingerprint=request_fingerprint,
)
if replay is None:
raise OrganizationMemoryConflictError(
"该组织记忆已被其他管理员创建,请刷新后重试。"
) from error
entry = replay
self.db.refresh(entry)
return self._serialize_entry(entry, self._department_names([entry]))
def update_organization_memory(
self,
memory_id: str,
payload: ExpenseApplicationOrganizationMemoryUpdate,
current_user: CurrentUserContext,
) -> ExpenseApplicationMemoryRead:
self._require_admin(current_user)
tenant_id = ExpenseCaseService.normalize_tenant_id(current_user.tenant_id)
previous_snapshot = self._get_entry(memory_id, tenant_id, for_update=False)
request_id = self._normalize_request_id(payload.request_id)
reason = self._normalize_reason(payload.reason)
request_fingerprint = self._request_fingerprint(
operation="update",
target=previous_snapshot.id,
payload={
"value": payload.value,
"expires_in_days": payload.expires_in_days,
"expected_generation": payload.expected_generation,
"reason": reason,
},
)
entry_id = self._request_entry_id(
tenant_id=tenant_id,
operation="update",
target=previous_snapshot.id,
request_id=request_id,
)
now = datetime.now(UTC)
try:
with organization_memory_operation_locks(
self.db,
self._scope_lock_key_for_entry(tenant_id, previous_snapshot),
self._request_lock_key(tenant_id, request_id),
):
replay = self._management_request_replay(
tenant_id=tenant_id,
request_id=request_id,
request_fingerprint=request_fingerprint,
)
if replay is not None:
result = self._serialize_entry(
replay,
self._department_names([replay]),
)
self.db.commit()
return result
previous = self._get_entry(memory_id, tenant_id)
self._require_active(previous)
self._validate_expected_generation(previous, payload.expected_generation)
if self._is_expired(previous, now):
previous.status = "expired"
previous.expired_at = now
self.db.commit()
raise OrganizationMemoryConflictError(
"组织记忆已过期,不能继续更新。"
)
expires_at = (
now + timedelta(days=payload.expires_in_days)
if payload.expires_in_days is not None
else self._aware(previous.active_expires_at)
)
entry = self._new_active_generation(
entry_id=entry_id,
tenant_id=tenant_id,
scope_type=str(previous.scope_type),
scope_id=str(previous.scope_id),
value=payload.value or self._entry_value(previous),
expires_at=expires_at,
reason=reason,
request_id=request_id,
request_fingerprint=request_fingerprint,
expected_active_ids={previous.id},
current_user=current_user,
now=now,
)
self.db.commit()
except OrganizationMemoryConflictError:
self.db.rollback()
raise
except IntegrityError as error:
self.db.rollback()
replay = self._management_request_replay(
tenant_id=tenant_id,
request_id=request_id,
request_fingerprint=request_fingerprint,
)
if replay is None:
raise OrganizationMemoryConflictError(
"组织记忆已被其他管理员更新,请刷新后重试。"
) from error
entry = replay
self.db.refresh(entry)
return self._serialize_entry(entry, self._department_names([entry]))
def revoke_organization_memory(
self,
memory_id: str,
payload: ExpenseApplicationOrganizationMemoryRevoke,
current_user: CurrentUserContext,
) -> ExpenseApplicationMemoryRevokedRead:
self._require_admin(current_user)
tenant_id = ExpenseCaseService.normalize_tenant_id(current_user.tenant_id)
entry_snapshot = self._get_entry(memory_id, tenant_id, for_update=False)
request_id = self._normalize_request_id(payload.request_id)
reason = self._normalize_reason(payload.reason)
request_fingerprint = self._request_fingerprint(
operation="revoke",
target=entry_snapshot.id,
payload={
"expected_generation": payload.expected_generation,
"reason": reason,
},
)
now = datetime.now(UTC)
try:
with organization_memory_operation_locks(
self.db,
self._scope_lock_key_for_entry(tenant_id, entry_snapshot),
self._request_lock_key(tenant_id, request_id),
):
replay = self._revoke_request_replay(
tenant_id=tenant_id,
request_id=request_id,
request_fingerprint=request_fingerprint,
)
if replay is not None:
result = self._revoked_response(replay)
self.db.commit()
return result
entry = self._get_entry(memory_id, tenant_id)
self._require_active(entry)
self._validate_expected_generation(entry, payload.expected_generation)
if self._is_expired(entry, now):
entry.status = "expired"
entry.expired_at = now
self.db.commit()
raise OrganizationMemoryConflictError(
"组织记忆已过期,不能继续撤销。"
)
entry.status = "revoked"
entry.revoked_at = now
entry.revoked_reason = reason
entry.managed_by = self._actor_id(current_user)
entry.managed_at = now
entry.management_reason = reason
entry.revoke_request_id = request_id
entry.revoke_payload_fingerprint = request_fingerprint
self.db.commit()
except OrganizationMemoryConflictError:
self.db.rollback()
raise
except IntegrityError as error:
self.db.rollback()
replay = self._revoke_request_replay(
tenant_id=tenant_id,
request_id=request_id,
request_fingerprint=request_fingerprint,
)
if replay is None:
raise OrganizationMemoryConflictError(
"组织记忆已被其他管理员撤销,请刷新后重试。"
) from error
entry = replay
return self._revoked_response(entry)
def _new_active_generation(
self,
*,
entry_id: str,
tenant_id: str,
scope_type: str,
scope_id: str,
value: str,
expires_at: datetime,
reason: str,
request_id: str,
request_fingerprint: str,
expected_active_ids: set[str],
current_user: CurrentUserContext,
now: datetime,
) -> MemoryEntry:
if value not in SUPPORTED_TRANSPORT_VALUES:
raise ValueError("组织记忆只允许飞机、火车或轮船三种低敏枚举值。")
existing = list(
self.db.scalars(
select(MemoryEntry)
.where(
MemoryEntry.tenant_id == tenant_id,
MemoryEntry.scope_type == scope_type,
MemoryEntry.scope_id == scope_id,
MemoryEntry.scene == MEMORY_SCENE,
MemoryEntry.field_key == MEMORY_FIELD_KEY,
MemoryEntry.status == "active",
)
.with_for_update()
).all()
)
actual_active_ids = {entry.id for entry in existing}
if actual_active_ids != expected_active_ids:
raise OrganizationMemoryConflictError(
"组织记忆生效版本已变化,请刷新后重试。"
)
generation = int(
self.db.scalar(
select(func.coalesce(func.max(MemoryEntry.generation), 0)).where(
MemoryEntry.tenant_id == tenant_id,
MemoryEntry.scope_type == scope_type,
MemoryEntry.scope_id == scope_id,
MemoryEntry.scene == MEMORY_SCENE,
MemoryEntry.field_key == MEMORY_FIELD_KEY,
)
)
or 0
) + 1
for old_entry in existing:
old_entry.status = "suppressed"
old_entry.suppressed_at = now
if existing:
self.db.flush()
entry = MemoryEntry(
id=entry_id,
tenant_id=tenant_id,
scope_type=scope_type,
scope_id=scope_id,
scene=MEMORY_SCENE,
field_key=MEMORY_FIELD_KEY,
origin_type="admin_managed",
generation=generation,
value_json={"value": value},
value_fingerprint=hmac_fingerprint(
{"field_key": MEMORY_FIELD_KEY, "value": value}
),
status="active",
evidence_count=0,
approved_evidence_count=0,
confidence=Decimal("1.0000"),
last_evidence_at=now,
candidate_expires_at=expires_at,
activated_at=now,
active_expires_at=expires_at,
managed_by=self._actor_id(current_user),
managed_at=now,
management_reason=reason,
policy_version=ORGANIZATION_MEMORY_POLICY_VERSION,
management_request_id=request_id,
management_payload_fingerprint=request_fingerprint,
)
self.db.add(entry)
self.db.flush()
for old_entry in existing:
old_entry.superseded_by_id = entry.id
self.db.flush()
return entry
@staticmethod
def _scope_lock_key(
*,
tenant_id: str,
scope_type: str,
scope_id: str,
) -> str:
return organization_memory_scope_lock_key(
tenant_id=tenant_id,
scope_type=scope_type,
scope_id=scope_id,
scene=MEMORY_SCENE,
field_key=MEMORY_FIELD_KEY,
)
@classmethod
def _scope_lock_key_for_entry(
cls,
tenant_id: str,
entry: MemoryEntry,
) -> str:
return cls._scope_lock_key(
tenant_id=tenant_id,
scope_type=str(entry.scope_type),
scope_id=str(entry.scope_id),
)
@staticmethod
def _request_lock_key(tenant_id: str, request_id: str) -> str:
return organization_memory_request_lock_key(tenant_id, request_id)
def _management_request_replay(
self,
*,
tenant_id: str,
request_id: str,
request_fingerprint: str,
) -> MemoryEntry | None:
if self._get_revoke_request_entry(tenant_id, request_id) is not None:
raise OrganizationMemoryConflictError(
"该幂等请求标识已用于其他组织记忆操作。"
)
entry = self.db.scalar(
self._organization_entry_query().where(
MemoryEntry.tenant_id == tenant_id,
MemoryEntry.management_request_id == request_id,
)
)
if entry is None:
return None
if entry.management_payload_fingerprint != request_fingerprint:
raise OrganizationMemoryConflictError(
"同一幂等请求标识对应的请求内容不一致。"
)
return entry
def _revoke_request_replay(
self,
*,
tenant_id: str,
request_id: str,
request_fingerprint: str,
) -> MemoryEntry | None:
management_entry = self.db.scalar(
self._organization_entry_query().where(
MemoryEntry.tenant_id == tenant_id,
MemoryEntry.management_request_id == request_id,
)
)
if management_entry is not None:
raise OrganizationMemoryConflictError(
"该幂等请求标识已用于其他组织记忆操作。"
)
entry = self._get_revoke_request_entry(tenant_id, request_id)
if entry is None:
return None
if entry.revoke_payload_fingerprint != request_fingerprint:
raise OrganizationMemoryConflictError(
"同一幂等请求标识对应的请求内容不一致。"
)
if entry.status != "revoked" or entry.revoked_at is None:
raise OrganizationMemoryConflictError("撤销操作审计状态不完整,请人工复核。")
return entry
def _get_revoke_request_entry(
self,
tenant_id: str,
request_id: str,
) -> MemoryEntry | None:
return self.db.scalar(
self._organization_entry_query().where(
MemoryEntry.tenant_id == tenant_id,
MemoryEntry.revoke_request_id == request_id,
)
)
@staticmethod
def _request_fingerprint(
*,
operation: str,
target: str,
payload: dict[str, object],
) -> str:
return hmac_fingerprint(
{
"protocol": "organization_memory_mutation.v1",
"operation": operation,
"target": target,
"payload": payload,
}
)
@staticmethod
def _normalize_request_id(value: str) -> str:
normalized = str(value or "").strip()
if len(normalized) < 8:
raise ValueError("组织记忆管理操作缺少有效的幂等请求标识。")
return normalized[:120]
@staticmethod
def _revoked_response(entry: MemoryEntry) -> ExpenseApplicationMemoryRevokedRead:
if entry.revoked_at is None:
raise OrganizationMemoryConflictError("撤销操作缺少审计时间,请人工复核。")
return ExpenseApplicationMemoryRevokedRead(
memory_id=entry.id,
revoked_at=entry.revoked_at,
)
@staticmethod
def _organization_entry_query():
return select(MemoryEntry).where(
MemoryEntry.scope_type.in_(ORGANIZATION_SCOPE_TYPES),
MemoryEntry.scene == MEMORY_SCENE,
MemoryEntry.field_key == MEMORY_FIELD_KEY,
)
def _get_entry(
self,
memory_id: str,
tenant_id: str,
*,
for_update: bool = True,
) -> MemoryEntry:
statement = self._organization_entry_query().where(
MemoryEntry.id == str(memory_id or "").strip(),
MemoryEntry.tenant_id == tenant_id,
)
if for_update:
statement = statement.with_for_update()
entry = self.db.scalar(statement)
if entry is None:
raise OrganizationMemoryNotFoundError("未找到当前租户内可管理的组织记忆。")
return entry
def _get_scope_active_entry(
self,
*,
tenant_id: str,
scope_type: str,
scope_id: str,
) -> MemoryEntry | None:
return self.db.scalar(
select(MemoryEntry)
.where(
MemoryEntry.tenant_id == tenant_id,
MemoryEntry.scope_type == scope_type,
MemoryEntry.scope_id == scope_id,
MemoryEntry.scene == MEMORY_SCENE,
MemoryEntry.field_key == MEMORY_FIELD_KEY,
MemoryEntry.status == "active",
)
.with_for_update()
)
@staticmethod
def _require_active(entry: MemoryEntry) -> None:
if entry.status != "active":
raise OrganizationMemoryConflictError(
"该组织记忆已被换代、撤销或失效,请刷新后重试。"
)
@staticmethod
def _request_entry_id(
*,
tenant_id: str,
operation: str,
target: str,
request_id: str,
) -> str:
normalized_request_id = str(request_id or "").strip()
if not normalized_request_id:
raise ValueError("组织记忆管理操作缺少幂等请求标识。")
material = "|".join(
(tenant_id, operation, target, normalized_request_id)
)
return str(uuid.uuid5(uuid.NAMESPACE_URL, f"x-financial:memory:{material}"))
def _validate_scope(
self,
*,
tenant_id: str,
scope_type: str,
scope_id: str | None,
) -> str:
normalized_scope_id = str(scope_id or "").strip()
if scope_type == "enterprise":
if normalized_scope_id and normalized_scope_id != tenant_id:
raise ValueError("企业记忆的 scope_id 必须等于当前租户 ID。")
return tenant_id
if scope_type != "department" or not normalized_scope_id:
raise ValueError("部门记忆必须提供稳定的 OrganizationUnit.id。")
department = self.db.get(OrganizationUnit, normalized_scope_id)
if department is None or str(department.unit_type or "") != "department":
raise ValueError("部门记忆只能绑定已存在的 department 类型组织单元。")
return normalized_scope_id
@staticmethod
def _require_admin(current_user: CurrentUserContext) -> None:
if not current_user.is_admin:
raise PermissionError("只有平台管理员可以维护企业或部门记忆。")
@staticmethod
def _validate_expected_generation(entry: MemoryEntry, expected: int) -> None:
if int(entry.generation or 0) != expected:
raise OrganizationMemoryConflictError(
"组织记忆已被其他管理员更新,请刷新后重试。"
)
def _department_names(self, entries: list[MemoryEntry]) -> dict[str, str]:
department_ids = {
str(entry.scope_id)
for entry in entries
if entry.scope_type == "department" and str(entry.scope_id or "")
}
if not department_ids:
return {}
return {
unit.id: str(unit.name or "").strip()
for unit in self.db.scalars(
select(OrganizationUnit).where(OrganizationUnit.id.in_(department_ids))
).all()
}
@staticmethod
def _serialize_entry(
entry: MemoryEntry,
department_names: dict[str, str],
) -> ExpenseApplicationMemoryRead:
scope_type = str(entry.scope_type)
if scope_type == "enterprise":
scope_label = "企业统一规则"
else:
name = department_names.get(str(entry.scope_id), "")
scope_label = f"部门规则({name}" if name else "部门规则"
return ExpenseApplicationMemoryRead(
id=entry.id,
scene=entry.scene,
field_key=entry.field_key,
value=ExpenseApplicationOrganizationMemoryService._entry_value(entry),
status=entry.status,
scope_type=scope_type,
scope_id=str(entry.scope_id),
scope_label=scope_label,
source=memory_source(scope_type),
origin_type=str(getattr(entry, "origin_type", "admin_managed")),
generation=int(entry.generation or 1),
evidence_count=int(entry.evidence_count or 0),
approved_evidence_count=int(entry.approved_evidence_count or 0),
confidence=float(entry.confidence or 0),
activation_threshold=0,
policy_version=str(
getattr(entry, "policy_version", ORGANIZATION_MEMORY_POLICY_VERSION)
),
valid_from=entry.activated_at or entry.created_at,
expires_at=entry.active_expires_at,
last_evidence_at=entry.last_evidence_at,
activated_at=entry.activated_at,
suppressed_at=entry.suppressed_at,
revoked_at=entry.revoked_at,
revoked_reason=str(entry.revoked_reason or ""),
managed_by=str(getattr(entry, "managed_by", "") or ""),
managed_at=getattr(entry, "managed_at", None),
management_reason=str(getattr(entry, "management_reason", "") or ""),
superseded_by_id=str(entry.superseded_by_id or ""),
can_revoke=entry.status == "active",
created_at=entry.created_at,
updated_at=entry.updated_at,
)
@staticmethod
def _entry_value(entry: MemoryEntry) -> str:
value_json = entry.value_json if isinstance(entry.value_json, dict) else {}
return str(value_json.get("value") or "").strip()
@staticmethod
def _actor_id(current_user: CurrentUserContext) -> str:
value = str(current_user.employee_id or current_user.username or "").strip()
if not value:
raise ValueError("当前管理员缺少稳定的操作人标识。")
return value[:120]
@staticmethod
def _normalize_reason(value: str) -> str:
normalized = str(value or "").strip()
if not normalized:
raise ValueError("组织记忆管理操作必须填写原因。")
return normalized[:255]
@staticmethod
def _is_expired(entry: MemoryEntry, now: datetime) -> bool:
return (
entry.active_expires_at is not None
and ExpenseApplicationOrganizationMemoryService._aware(
entry.active_expires_at
)
<= now
)
@staticmethod
def _aware(value: datetime | None) -> datetime:
if value is None:
raise ValueError("组织记忆缺少有效期。")
return value if value.tzinfo is not None else value.replace(tzinfo=UTC)

View File

@@ -0,0 +1,264 @@
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from datetime import UTC, datetime
from sqlalchemy import select
from sqlalchemy.orm import Session
from app.api.deps import CurrentUserContext
from app.models.ai_memory import MemoryEntry
from app.schemas.expense_application_memory import (
ExpenseApplicationMemoryApplication,
ExpenseApplicationMemoryConflict,
)
from app.services.expense_cases import ExpenseCaseService
MEMORY_SCENE = "travel_application"
MEMORY_FIELD_KEY = "transport_mode"
SUPPORTED_TRANSPORT_VALUES = {"飞机", "火车", "轮船"}
SCOPE_PRIORITIES = {"enterprise": 300, "department": 200, "user": 100}
RefreshUserEntry = Callable[..., None]
@dataclass(frozen=True, slots=True)
class ExpenseApplicationMemoryResolution:
winner: MemoryEntry | None
value: str
conflicts: tuple[ExpenseApplicationMemoryConflict, ...]
effective_confidence: float
status: str = "applied"
def to_application(
self,
current_user: CurrentUserContext,
) -> ExpenseApplicationMemoryApplication:
if self.winner is None:
first = self.conflicts[0]
return ExpenseApplicationMemoryApplication(
memory_id="",
value="",
source="memory_conflict",
status="conflict",
confidence=0.0,
effective_confidence=0.0,
scope_type=first.scope_type,
scope_id="",
scope_label=first.scope_label,
priority=first.priority,
conflicts=list(self.conflicts),
can_revoke=False,
message="发现同一层级存在相互冲突的有效记忆,本次未自动填充,请联系管理员处理。",
)
entry = self.winner
scope_type = _scope_type(entry)
return ExpenseApplicationMemoryApplication(
memory_id=entry.id,
value=self.value,
source=memory_source(scope_type),
status=self.status,
evidence_count=int(entry.evidence_count or 0),
approved_evidence_count=int(entry.approved_evidence_count or 0),
confidence=float(entry.confidence or 0),
effective_confidence=self.effective_confidence,
expires_at=entry.active_expires_at,
scope_type=scope_type,
scope_id=str(entry.scope_id or ""),
scope_label=memory_scope_label(entry, current_user),
priority=SCOPE_PRIORITIES[scope_type],
conflicts=list(self.conflicts),
can_revoke=scope_type == "user",
message=_application_message(scope_type),
)
class ExpenseApplicationMemoryResolver:
"""解析当前用户可见的企业、部门、个人记忆,并保守处理冲突。"""
def __init__(self, db: Session) -> None:
self.db = db
def resolve(
self,
current_user: CurrentUserContext,
*,
user_scope_id: str,
refresh_user_entry: RefreshUserEntry,
) -> ExpenseApplicationMemoryResolution | None:
tenant_id = ExpenseCaseService.normalize_tenant_id(current_user.tenant_id)
relevant_scopes = [("enterprise", tenant_id)]
department_id = str(getattr(current_user, "department_id", "") or "").strip()
if department_id:
relevant_scopes.append(("department", department_id))
relevant_scopes.append(("user", user_scope_id))
now = datetime.now(UTC)
entries: list[MemoryEntry] = []
for scope_type, scope_id in relevant_scopes:
entries.extend(
self.db.scalars(
select(MemoryEntry)
.where(
MemoryEntry.tenant_id == tenant_id,
MemoryEntry.scope_type == scope_type,
MemoryEntry.scope_id == scope_id,
MemoryEntry.scene == MEMORY_SCENE,
MemoryEntry.field_key == MEMORY_FIELD_KEY,
MemoryEntry.status.in_(["candidate", "active"]),
)
.order_by(
MemoryEntry.generation.desc(),
MemoryEntry.last_evidence_at.desc(),
)
.with_for_update()
).all()
)
user_entries = [entry for entry in entries if _scope_type(entry) == "user"]
active_user_found = False
for entry in user_entries:
if entry.status != "active":
continue
refresh_user_entry(entry, now=now)
active_user_found = active_user_found or entry.status == "active"
for entry in user_entries:
if entry.status != "candidate":
continue
refresh_user_entry(
entry,
now=now,
allow_activation=not active_user_found,
)
active_user_found = active_user_found or entry.status == "active"
eligible: list[MemoryEntry] = []
for entry in entries:
if (
_scope_type(entry) != "user"
and entry.status == "active"
and _is_expired(entry, now)
):
entry.status = "expired"
entry.expired_at = now
if entry.status != "active" or _is_expired(entry, now):
continue
if _entry_value(entry):
eligible.append(entry)
if not eligible:
return None
highest_priority = max(SCOPE_PRIORITIES[_scope_type(entry)] for entry in eligible)
highest = [
entry
for entry in eligible
if SCOPE_PRIORITIES[_scope_type(entry)] == highest_priority
]
top_values = {_entry_value(entry) for entry in highest}
if len(top_values) > 1:
conflicts = tuple(
_conflict(entry, current_user, reason="same_priority_conflict")
for entry in highest
)
return ExpenseApplicationMemoryResolution(
winner=None,
value="",
conflicts=conflicts,
effective_confidence=0.0,
status="conflict",
)
winner = highest[0]
value = _entry_value(winner)
conflicts = tuple(
_conflict(entry, current_user, reason="lower_priority_overridden")
for entry in eligible
if SCOPE_PRIORITIES[_scope_type(entry)] < highest_priority
and _entry_value(entry) != value
)
return ExpenseApplicationMemoryResolution(
winner=winner,
value=value,
conflicts=conflicts,
effective_confidence=_effective_confidence(winner, now),
)
def _entry_value(entry: MemoryEntry) -> str:
value_json = entry.value_json if isinstance(entry.value_json, dict) else {}
value = str(value_json.get("value") or "").strip()
return value if value in SUPPORTED_TRANSPORT_VALUES else ""
def _scope_type(entry: MemoryEntry) -> str:
value = str(entry.scope_type or "user")
return value if value in SCOPE_PRIORITIES else "user"
def memory_source(scope_type: str) -> str:
return {
"enterprise": "enterprise_policy_memory",
"department": "department_policy_memory",
"user": "verified_user_history",
}[scope_type]
def memory_scope_label(entry: MemoryEntry, current_user: CurrentUserContext) -> str:
scope_type = _scope_type(entry)
if scope_type == "enterprise":
return "企业统一规则"
if scope_type == "department":
department_name = str(current_user.department_name or "").strip()
return f"部门规则({department_name}" if department_name else "部门规则"
return "个人偏好"
def _application_message(scope_type: str) -> str:
return {
"enterprise": "已按企业统一规则预填出行方式,可在规则允许范围内调整。",
"department": "已按当前部门规则预填出行方式,可继续修改。",
"user": "已按可信历史记忆预填常用出行方式,可继续修改。",
}[scope_type]
def _conflict(
entry: MemoryEntry,
current_user: CurrentUserContext,
*,
reason: str,
) -> ExpenseApplicationMemoryConflict:
scope_type = _scope_type(entry)
return ExpenseApplicationMemoryConflict(
scope_type=scope_type,
scope_label=memory_scope_label(entry, current_user),
priority=SCOPE_PRIORITIES[scope_type],
reason=reason,
)
def _effective_confidence(entry: MemoryEntry, now: datetime) -> float:
base = max(0.0, min(1.0, float(entry.confidence or 0)))
if base == 0.0 and str(getattr(entry, "origin_type", "learned")) == "admin_managed":
base = 1.0
started_at = entry.activated_at or entry.created_at
expires_at = entry.active_expires_at
if started_at is None or expires_at is None:
return round(base, 4)
started = _aware(started_at)
expires = _aware(expires_at)
total_seconds = max(1.0, (expires - started).total_seconds())
remaining_ratio = max(0.0, min(1.0, (expires - now).total_seconds() / total_seconds))
# 有效期内仅降低解释性置信度,不会因衰减而提前停止应用。
return round(base * max(0.5, remaining_ratio), 4)
def _is_expired(entry: MemoryEntry, now: datetime) -> bool:
expires_at = entry.active_expires_at if entry.status == "active" else entry.candidate_expires_at
return expires_at is not None and _aware(expires_at) <= now
def _aware(value: datetime) -> datetime:
return value if value.tzinfo is not None else value.replace(tzinfo=UTC)

View File

@@ -0,0 +1,202 @@
from __future__ import annotations
from typing import Any
from sqlalchemy.orm import Session
from app.core.logging import get_logger
from app.core.secret_box import decrypt_secret
from app.models.financial_record import ExpenseClaim
from app.models.system_model_setting import SystemModelSetting
from app.services.embedding_provider import EmbeddingProvider
from app.services.few_shot_retrieval import FewShotRetriever
from app.services.few_shot_store import FewShotStore
from app.services.knowledge_rag_runtime import RuntimeModelConfig
logger = get_logger("app.services.expense_claim_historical_evidence")
_SCENE_BY_BUSINESS_STAGE = {
"expense_application": "expense_application",
"reimbursement": "expense_reimbursement",
}
_LABEL_TEXT = {
"confirmed": "历史已确认,仅供复核",
"false_positive": "历史误报,仅供复核",
}
_SUMMARY_BY_LABEL = {
"confirmed": "历史相似案例经人工复核确认风险成立。",
"false_positive": "历史相似案例经人工复核判定为误报。",
}
_MAX_RULE_CONTEXTS = 3
_MAX_EVIDENCE = 3
class ExpenseClaimHistoricalEvidenceService:
"""为费用预审补充历史案例证据,但不参与确定性决策。"""
def __init__(self, db: Session) -> None:
self.db = db
def retrieve(
self,
claim: ExpenseClaim,
*,
tenant_id: str,
business_stage: str,
findings: list[dict[str, Any]],
) -> list[dict[str, Any]]:
tenant = _text(tenant_id)
scene_code = _SCENE_BY_BUSINESS_STAGE.get(_text(business_stage), "")
if not tenant or not scene_code:
return []
try:
retriever = self._build_retriever()
if retriever is None:
return []
query = self._build_query(claim, findings=findings)
evidence: list[dict[str, Any]] = []
emitted_sample_ids: set[str] = set()
for policy_ref, rule_version in self._rule_contexts(findings):
hits = retriever.retrieve_for_expense_case(
tenant_id=tenant,
scene=scene_code,
policy_ref=policy_ref,
rule_version=rule_version,
query=query,
top_k=_MAX_EVIDENCE,
)
for hit in hits:
sample_id = _text(hit.get("sample_id"))
if not sample_id or sample_id in emitted_sample_ids:
continue
public_item = self._to_public_evidence(hit, scene_code=scene_code)
if not public_item:
continue
evidence.append(public_item)
emitted_sample_ids.add(sample_id)
if len(evidence) >= _MAX_EVIDENCE:
return evidence
return evidence
except Exception:
# 历史案例只作复核参考。检索、向量库或配置异常不得阻断预审。
logger.warning(
"费用预审历史案例检索失败 tenant_id=%s claim_id=%s",
tenant,
_text(claim.id),
exc_info=True,
)
return []
def _build_retriever(self) -> FewShotRetriever | None:
"""只读加载 embedding 配置,禁止在费用事务中触发配置初始化提交。"""
model_row = self.db.get(SystemModelSetting, "embedding")
if model_row is None or not model_row.enabled:
return None
encrypted_api_key = _text(model_row.api_key_encrypted)
try:
api_key = decrypt_secret(encrypted_api_key) if encrypted_api_key else ""
except ValueError:
logger.warning("embedding 配置密钥无法解密,历史案例检索已跳过")
return None
provider = EmbeddingProvider(
RuntimeModelConfig(
slot="embedding",
provider=_text(model_row.provider),
model=_text(model_row.model_name),
endpoint=_text(model_row.endpoint),
api_key=api_key,
capability=_text(model_row.capability) or "embedding",
)
)
return FewShotRetriever(FewShotStore(provider), self.db)
@staticmethod
def _rule_contexts(
findings: list[dict[str, Any]],
) -> list[tuple[str, str]]:
contexts: list[tuple[str, str]] = []
for finding in findings:
context = (
_text(finding.get("rule_code")),
_text(finding.get("rule_version")),
)
if context != ("", "") and context not in contexts:
contexts.append(context)
if len(contexts) >= _MAX_RULE_CONTEXTS:
break
# 无规则命中时仍按租户和业务场景检索,但显式传递空的规则标识。
return contexts or [("", "")]
@staticmethod
def _build_query(
claim: ExpenseClaim,
*,
findings: list[dict[str, Any]],
) -> str:
parts = [
_text(claim.expense_type),
_text(claim.reason),
_text(claim.location),
*[
_text(finding.get("message"))
for finding in findings
if _text(finding.get("message"))
],
]
return "\n".join(part for part in parts if part).strip()
@staticmethod
def _to_public_evidence(
hit: dict[str, Any],
*,
scene_code: str,
) -> dict[str, Any]:
label = _text(hit.get("label")).lower()
label_text = _LABEL_TEXT.get(label)
summary = _SUMMARY_BY_LABEL.get(label)
if not label_text or not summary:
return {}
return {
"label": label,
"label_text": label_text,
"advisory_only": True,
"score": round(float(hit.get("score") or 0.0), 4),
"scene_code": _text(hit.get("scene")) or scene_code,
"policy_ref": _text(hit.get("policy_ref")),
"rule_version": _text(hit.get("rule_version")),
"version_status": (
"stale" if bool(hit.get("stale")) else "matched"
),
"stale": bool(hit.get("stale")),
"summary": summary,
}
def _text(value: Any) -> str:
return str(value or "").strip()
def build_user_agent_historical_evidence_notice(claim: ExpenseClaim) -> str:
"""生成脱敏的 User Agent 提示,不读取或回显历史案例原文。"""
flags = claim.risk_flags_json
if isinstance(flags, dict):
flags = [flags]
if not isinstance(flags, list):
return ""
labels: list[str] = []
for flag in flags:
if (
not isinstance(flag, dict)
or _text(flag.get("source")) != "ai_pre_review"
):
continue
for item in list(flag.get("historical_case_evidence") or []):
if not isinstance(item, dict):
continue
label_text = _LABEL_TEXT.get(_text(item.get("label")).lower(), "")
if label_text and label_text not in labels:
labels.append(label_text)
return "历史案例参考:" + "".join(labels) if labels else ""

View File

@@ -6,6 +6,9 @@ from typing import Any
from app.api.deps import CurrentUserContext from app.api.deps import CurrentUserContext
from app.models.financial_record import ExpenseClaim from app.models.financial_record import ExpenseClaim
from app.services.expense_claim_errors import ExpenseClaimSubmissionBlockedError from app.services.expense_claim_errors import ExpenseClaimSubmissionBlockedError
from app.services.expense_claim_historical_evidence import (
ExpenseClaimHistoricalEvidenceService,
)
from app.services.expense_claim_pre_review_decision import build_pre_review_decision from app.services.expense_claim_pre_review_decision import build_pre_review_decision
from app.services.expense_claim_risk_flags import dedupe_claim_risk_flags from app.services.expense_claim_risk_flags import dedupe_claim_risk_flags
from app.services.expense_claim_risk_stage import ( from app.services.expense_claim_risk_stage import (
@@ -45,6 +48,7 @@ class ExpenseClaimPreReviewMixin:
claim, claim,
is_application_claim=is_application_claim, is_application_claim=is_application_claim,
reviewed_at=datetime.now(UTC), reviewed_at=datetime.now(UTC),
tenant_id=current_user.tenant_id,
) )
if pre_review_flag is None: if pre_review_flag is None:
raise RuntimeError("无法生成费用预审结果。") raise RuntimeError("无法生成费用预审结果。")
@@ -83,6 +87,7 @@ class ExpenseClaimPreReviewMixin:
*, *,
decision_payload: dict[str, Any], decision_payload: dict[str, Any],
business_stage: str, business_stage: str,
historical_case_evidence: list[dict[str, Any]] | None = None,
) -> dict[str, Any]: ) -> dict[str, Any]:
decision = str(decision_payload.get("decision") or "ready_with_review") decision = str(decision_payload.get("decision") or "ready_with_review")
passed = decision != "needs_fix" passed = decision != "needs_fix"
@@ -101,6 +106,7 @@ class ExpenseClaimPreReviewMixin:
"passed": passed, "passed": passed,
"blocking_risk_count": blocking_count, "blocking_risk_count": blocking_count,
**decision_payload, **decision_payload,
"historical_case_evidence": list(historical_case_evidence or []),
"next_action": "next_step" if passed else "risk_explanation_required", "next_action": "next_step" if passed else "risk_explanation_required",
"created_at": str(decision_payload.get("reviewed_at") or ""), "created_at": str(decision_payload.get("reviewed_at") or ""),
}, },
@@ -128,12 +134,14 @@ class ExpenseClaimPreReviewMixin:
*, *,
is_application_claim: bool | None = None, is_application_claim: bool | None = None,
reviewed_at: datetime | None = None, reviewed_at: datetime | None = None,
tenant_id: str = "",
) -> dict[str, Any] | None: ) -> dict[str, Any] | None:
"""业务变更事务内刷新预审快照,不提交、不单独写事件。""" """业务变更事务内刷新预审快照,不提交、不单独写事件。"""
return self._refresh_claim_pre_review_flags( return self._refresh_claim_pre_review_flags(
claim, claim,
is_application_claim=is_application_claim, is_application_claim=is_application_claim,
reviewed_at=reviewed_at, reviewed_at=reviewed_at,
tenant_id=tenant_id,
) )
def _refresh_claim_pre_review_flags( def _refresh_claim_pre_review_flags(
@@ -142,6 +150,7 @@ class ExpenseClaimPreReviewMixin:
*, *,
is_application_claim: bool | None = None, is_application_claim: bool | None = None,
reviewed_at: datetime | None = None, reviewed_at: datetime | None = None,
tenant_id: str = "",
) -> dict[str, Any] | None: ) -> dict[str, Any] | None:
if claim is None: if claim is None:
return None return None
@@ -187,9 +196,18 @@ class ExpenseClaimPreReviewMixin:
platform_rule_set_fingerprint=platform_rule_set_fingerprint, platform_rule_set_fingerprint=platform_rule_set_fingerprint,
reviewed_at=reviewed_at, reviewed_at=reviewed_at,
) )
historical_case_evidence = ExpenseClaimHistoricalEvidenceService(
self.db
).retrieve(
claim,
tenant_id=tenant_id,
business_stage=business_stage,
findings=list(decision_payload.get("findings") or []),
)
pre_review_flag = self._build_ai_pre_review_flag( pre_review_flag = self._build_ai_pre_review_flag(
decision_payload=decision_payload, decision_payload=decision_payload,
business_stage=business_stage, business_stage=business_stage,
historical_case_evidence=historical_case_evidence,
) )
claim.risk_flags_json = self._replace_ai_pre_review_flag( claim.risk_flags_json = self._replace_ai_pre_review_flag(
review_flags, review_flags,

View File

@@ -210,6 +210,48 @@ def pre_review_public_payload(flag: dict[str, Any] | None) -> dict[str, Any] | N
for item in list(flag.get("findings") or []) for item in list(flag.get("findings") or [])
if isinstance(item, dict) if isinstance(item, dict)
], ],
"historical_case_evidence": _public_historical_evidence(flag),
}
def _public_historical_evidence(flag: dict[str, Any]) -> list[dict[str, Any]]:
result: list[dict[str, Any]] = []
for item in list(flag.get("historical_case_evidence") or []):
if not isinstance(item, dict):
continue
public_item = _historical_evidence_public_payload(item)
if public_item is not None:
result.append(public_item)
return result
def _historical_evidence_public_payload(
item: dict[str, Any],
) -> dict[str, Any] | None:
label = _text(item.get("label")).lower()
if label not in {"confirmed", "false_positive"}:
return None
return {
"label": label,
"label_text": (
"历史已确认,仅供复核"
if label == "confirmed"
else "历史误报,仅供复核"
),
"advisory_only": True,
"score": round(float(item.get("score") or 0.0), 4),
"scene_code": _text(item.get("scene_code")),
"policy_ref": _text(item.get("policy_ref")),
"rule_version": _text(item.get("rule_version")),
"version_status": (
"stale" if bool(item.get("stale")) else "matched"
),
"stale": bool(item.get("stale")),
"summary": (
"历史相似案例经人工复核确认风险成立。"
if label == "confirmed"
else "历史相似案例经人工复核判定为误报。"
),
} }

View File

@@ -260,6 +260,7 @@ class ExpenseClaimItemActionMixin:
pre_review_flag = self.refresh_claim_pre_review_state( pre_review_flag = self.refresh_claim_pre_review_state(
claim, claim,
is_application_claim=is_application_claim, is_application_claim=is_application_claim,
tenant_id=current_user.tenant_id,
) )
if pre_review_flag is None: if pre_review_flag is None:
raise RuntimeError("无法生成提交前预审结果。") raise RuntimeError("无法生成提交前预审结果。")

View File

@@ -31,6 +31,16 @@ LABEL_CONCLUSION_FALLBACK = {
"false_positive": "经人工复核判定为误报,相似情形不应触发该风险规则。", "false_positive": "经人工复核判定为误报,相似情形不应触发该风险规则。",
} }
CONTROL_STAGE_SCENES = {
"application": "expense_application",
"expense_application": "expense_application",
"pre_application": "expense_application",
"pre_reimbursement": "expense_reimbursement",
"reimbursement": "expense_reimbursement",
"claim": "expense_reimbursement",
"post_payment": "expense_post_payment",
}
class FewShotIngestionService: class FewShotIngestionService:
"""把已确认的风险观测沉淀为 few-shot 样本。""" """把已确认的风险观测沉淀为 few-shot 样本。"""
@@ -48,22 +58,34 @@ class FewShotIngestionService:
label = observation.feedback_status label = observation.feedback_status
if label not in CONFIRMED_LABELS: if label not in CONFIRMED_LABELS:
return None return None
tenant_id = str(observation.tenant_id or "").strip()
if not tenant_id:
logger.warning("few-shot ingestion 缺少 tenant_id observation_id=%s", observation.id)
return None
sample_key = f"obs:{observation.id}" sample_key = f"obs:{observation.id}"
sample = self.db.scalar( sample = self.db.scalar(
select(FewShotSample).where(FewShotSample.sample_key == sample_key) select(FewShotSample).where(
FewShotSample.tenant_id == tenant_id,
FewShotSample.sample_key == sample_key,
)
) )
domain = self._extract_domain(observation) domain = self._extract_domain(observation)
scene = self._extract_scene(observation)
policy_ref, rule_version = self._extract_rule_identity(observation)
case_text = self._build_case_text(observation) case_text = self._build_case_text(observation)
conclusion_text = self._build_conclusion_text(observation, feedback, label) conclusion_text = self._build_conclusion_text(observation, feedback, label)
payload = self._build_payload(observation, feedback, label) payload = self._build_payload(observation, feedback, label)
if sample is None: if sample is None:
sample = FewShotSample( sample = FewShotSample(
tenant_id=tenant_id,
sample_key=sample_key, sample_key=sample_key,
source_observation_id=observation.id, source_observation_id=observation.id,
scene="risk_rule_generation", scene=scene,
policy_ref=policy_ref,
rule_version=rule_version,
domain=domain, domain=domain,
risk_type=observation.risk_type or "", risk_type=observation.risk_type or "",
risk_level=observation.risk_level or "", risk_level=observation.risk_level or "",
@@ -75,7 +97,11 @@ class FewShotIngestionService:
) )
self.db.add(sample) self.db.add(sample)
else: else:
sample.tenant_id = tenant_id
sample.label = label sample.label = label
sample.scene = scene
sample.policy_ref = policy_ref
sample.rule_version = rule_version
sample.domain = domain sample.domain = domain
sample.risk_type = observation.risk_type or "" sample.risk_type = observation.risk_type or ""
sample.risk_level = observation.risk_level or "" sample.risk_level = observation.risk_level or ""
@@ -83,7 +109,6 @@ class FewShotIngestionService:
sample.conclusion_text = conclusion_text sample.conclusion_text = conclusion_text
sample.payload_json = payload sample.payload_json = payload
sample.status = "active" sample.status = "active"
sample.vector_id = sample.vector_id
try: try:
self.db.commit() self.db.commit()
self.db.refresh(sample) self.db.refresh(sample)
@@ -101,11 +126,17 @@ class FewShotIngestionService:
logger.warning("few-shot vector_id 回写失败 sample_id=%s", sample.id) logger.warning("few-shot vector_id 回写失败 sample_id=%s", sample.id)
return sample return sample
def retract_observation(self, observation_id: str) -> bool: def retract_observation(self, observation_id: str, *, tenant_id: str) -> bool:
"""观测被撤销时删掉对应样本及其向量。""" """观测被撤销时删掉对应样本及其向量。"""
tenant = str(tenant_id or "").strip()
if not tenant:
raise ValueError("tenant_id is required")
sample = self.db.scalar( sample = self.db.scalar(
select(FewShotSample).where(FewShotSample.source_observation_id == observation_id) select(FewShotSample).where(
FewShotSample.tenant_id == tenant,
FewShotSample.source_observation_id == observation_id,
)
) )
if sample is None: if sample is None:
return False return False
@@ -128,6 +159,34 @@ class FewShotIngestionService:
ontology = observation.ontology_json or {} ontology = observation.ontology_json or {}
return str(ontology.get("domain") or "") return str(ontology.get("domain") or "")
def _extract_scene(self, observation: RiskObservation) -> str:
stage = str(observation.control_stage or "").strip().lower()
if stage in CONTROL_STAGE_SCENES:
return CONTROL_STAGE_SCENES[stage]
return stage or "risk_rule_generation"
def _extract_rule_identity(self, observation: RiskObservation) -> tuple[str, str]:
trace = observation.decision_trace_json or {}
policy_ref = _text(
trace.get("policy_ref") or trace.get("rule_code") or trace.get("policy_code")
)
if not policy_ref:
for value in observation.policy_refs_json or []:
if isinstance(value, dict):
policy_ref = _text(
value.get("policy_ref") or value.get("rule_code") or value.get("code")
)
else:
policy_ref = _text(value)
if policy_ref:
break
rule_version = _text(
trace.get("rule_version")
or trace.get("policy_version")
or observation.algorithm_version
)
return policy_ref, rule_version
def _build_case_text(self, observation: RiskObservation) -> str: def _build_case_text(self, observation: RiskObservation) -> str:
parts = [ parts = [
observation.title or "", observation.title or "",
@@ -162,6 +221,8 @@ class FewShotIngestionService:
label: str, label: str,
) -> dict[str, Any]: ) -> dict[str, Any]:
return { return {
"tenant_id": observation.tenant_id,
"scene": self._extract_scene(observation),
"label": label, "label": label,
"risk_type": observation.risk_type, "risk_type": observation.risk_type,
"risk_signal": observation.risk_signal, "risk_signal": observation.risk_signal,
@@ -171,7 +232,13 @@ class FewShotIngestionService:
"feedback_actor": feedback.actor or "", "feedback_actor": feedback.actor or "",
"ontology": observation.ontology_json or {}, "ontology": observation.ontology_json or {},
"policy_refs": observation.policy_refs_json or [], "policy_refs": observation.policy_refs_json or [],
"policy_ref": self._extract_rule_identity(observation)[0],
"rule_version": self._extract_rule_identity(observation)[1],
"evidence": observation.evidence_json or [], "evidence": observation.evidence_json or [],
"subject_label": observation.subject_label or "", "subject_label": observation.subject_label or "",
"claim_no": observation.claim_no or "", "claim_no": observation.claim_no or "",
} }
def _text(value: Any) -> str:
return str(value or "").strip()

View File

@@ -19,9 +19,11 @@ from __future__ import annotations
from typing import Any from typing import Any
from sqlalchemy import select
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.core.logging import get_logger from app.core.logging import get_logger
from app.models.few_shot_sample import FewShotSample
from app.services.embedding_provider import EmbeddingProvider from app.services.embedding_provider import EmbeddingProvider
from app.services.few_shot_store import FewShotStore from app.services.few_shot_store import FewShotStore
@@ -38,17 +40,19 @@ MAX_HISTORICAL_SAMPLES = 3
class FewShotRetriever: class FewShotRetriever:
"""按 case 特征检索已确认样本,返回 prompt 可直接消费的结构。""" """按 case 特征检索已确认样本,返回 prompt 可直接消费的结构。"""
def __init__(self, store: FewShotStore) -> None: def __init__(self, store: FewShotStore, session: Session | None = None) -> None:
self._store = store self._store = store
self._session = session
@classmethod @classmethod
def from_session(cls, session: Session) -> "FewShotRetriever": def from_session(cls, session: Session) -> FewShotRetriever:
provider = EmbeddingProvider.from_settings(session) provider = EmbeddingProvider.from_settings(session)
return cls(FewShotStore(provider)) return cls(FewShotStore(provider), session)
def retrieve_for_risk_rule_generation( def retrieve_for_risk_rule_generation(
self, self,
*, *,
tenant_id: str,
domain: str = "", domain: str = "",
risk_type: str = "", risk_type: str = "",
natural_language: str, natural_language: str,
@@ -65,12 +69,125 @@ class FewShotRetriever:
return [] return []
hits = self._store.search( hits = self._store.search(
case_text, case_text,
tenant_id=tenant_id,
scene="risk_rule_generation", scene="risk_rule_generation",
labels=["confirmed", "false_positive"], labels=["confirmed", "false_positive"],
top_k=top_k, top_k=top_k,
) )
return self._hits_to_injection_blocks(hits) return self._hits_to_injection_blocks(hits)
def retrieve_for_expense_case(
self,
*,
tenant_id: str,
scene: str,
policy_ref: str,
rule_version: str,
query: str,
top_k: int = MAX_HISTORICAL_SAMPLES,
) -> list[dict[str, Any]]:
"""返回仅作建议的历史案例,并以关系库状态做二次授权校验。"""
tenant = str(tenant_id or "").strip()
normalized_scene = str(scene or "").strip()
normalized_policy = str(policy_ref or "").strip()
normalized_version = str(rule_version or "").strip()
if not tenant:
raise ValueError("tenant_id is required for historical case retrieval")
if not query or not normalized_scene or self._session is None:
return []
hits = self._store.search(
query,
tenant_id=tenant,
scene=normalized_scene,
policy_ref=normalized_policy or None,
rule_version=normalized_version or None,
labels=["confirmed", "false_positive"],
top_k=top_k,
)
# 精确版本不足时补检同规则旧版本,但输出会显式标记为 stale绝不自动执行。
if normalized_version and len(hits) < top_k:
fallback_hits = self._store.search(
query,
tenant_id=tenant,
scene=normalized_scene,
policy_ref=normalized_policy or None,
labels=["confirmed", "false_positive"],
top_k=top_k * 2,
)
seen = {str(item.get("sample_id") or "") for item in hits}
hits.extend(
item for item in fallback_hits if str(item.get("sample_id") or "") not in seen
)
return self._validated_expense_case_hits(
hits[: top_k * 2],
tenant_id=tenant,
scene=normalized_scene,
policy_ref=normalized_policy,
rule_version=normalized_version,
top_k=top_k,
)
def _validated_expense_case_hits(
self,
hits: list[dict[str, Any]],
*,
tenant_id: str,
scene: str,
policy_ref: str,
rule_version: str,
top_k: int,
) -> list[dict[str, Any]]:
sample_ids = [str(hit.get("sample_id") or "") for hit in hits]
sample_ids = [sample_id for sample_id in sample_ids if sample_id]
if not sample_ids or self._session is None:
return []
conditions = [
FewShotSample.id.in_(sample_ids),
FewShotSample.tenant_id == tenant_id,
FewShotSample.scene == scene,
FewShotSample.status == "active",
]
if policy_ref:
conditions.append(FewShotSample.policy_ref == policy_ref)
samples = {
item.id: item
for item in self._session.scalars(select(FewShotSample).where(*conditions)).all()
}
result: list[dict[str, Any]] = []
emitted_sample_ids: set[str] = set()
for hit in hits:
sample_id = str(hit.get("sample_id") or "")
sample = samples.get(sample_id)
if sample is None or sample_id in emitted_sample_ids:
continue
version_matches = not rule_version or sample.rule_version == rule_version
result.append(
{
"source": "historical_case",
"advisory_only": True,
"sample_id": sample.id,
"label": sample.label,
"score": round(float(hit.get("score") or 0.0), 4),
"scene": sample.scene,
"policy_ref": sample.policy_ref,
"rule_version": sample.rule_version,
"version_status": "matched" if version_matches else "stale",
"stale": not version_matches,
"conclusion": sample.conclusion_text[:SINGLE_SAMPLE_MAX_CHARS],
"evidence": {
"risk_type": sample.risk_type,
"risk_level": sample.risk_level,
"payload": sample.payload_json or {},
},
}
)
emitted_sample_ids.add(sample_id)
if len(result) >= top_k:
break
return result
def _build_case_text( def _build_case_text(
self, self,
*, *,

View File

@@ -20,6 +20,17 @@ from app.services.knowledge_rag import _resolve_default_qdrant_url
logger = get_logger("app.services.few_shot_store") logger = get_logger("app.services.few_shot_store")
FEW_SHOT_COLLECTION = "few_shot_samples" FEW_SHOT_COLLECTION = "few_shot_samples"
FEW_SHOT_VECTOR_NAMESPACE = uuid.UUID("0ecb5868-cf75-47c1-b3bf-81e4a6ce1df3")
def stable_vector_id(*, tenant_id: str, sample_id: str) -> str:
"""同一租户样本始终映射到同一个 Qdrant point。"""
tenant = str(tenant_id or "").strip()
sample = str(sample_id or "").strip()
if not tenant or not sample:
raise ValueError("tenant_id and sample_id are required")
return str(uuid.uuid5(FEW_SHOT_VECTOR_NAMESPACE, f"{tenant}:{sample}"))
def _resolve_qdrant_config() -> tuple[str, str]: def _resolve_qdrant_config() -> tuple[str, str]:
@@ -73,26 +84,29 @@ class FewShotStore:
try: try:
client.get_collection(FEW_SHOT_COLLECTION) client.get_collection(FEW_SHOT_COLLECTION)
self._ensured = True
return True
except UnexpectedResponse as exc: except UnexpectedResponse as exc:
if exc.status_code != 404: if exc.status_code != 404:
raise raise
# collection 不存在则创建 dim = self._embedding_provider.dimension()
dim = self._embedding_provider.dimension() from qdrant_client.http.models import Distance, VectorParams
client.create_collection(
collection_name=FEW_SHOT_COLLECTION,
vectors_config=VectorParams(size=dim, distance=Distance.COSINE),
)
logger.info("few-shot collection 创建成功 dim=%s", dim)
# 老 collection 也要补齐过滤索引,不能只在首次建表时创建。
from qdrant_client.http.models import ( from qdrant_client.http.models import (
Distance,
VectorParams,
PayloadSchemaType, PayloadSchemaType,
) )
client.create_collection(
collection_name=FEW_SHOT_COLLECTION,
vectors_config=VectorParams(size=dim, distance=Distance.COSINE),
)
for field, field_type in [ for field, field_type in [
("sample_id", PayloadSchemaType.KEYWORD), ("sample_id", PayloadSchemaType.KEYWORD),
("tenant_id", PayloadSchemaType.KEYWORD),
("scene", PayloadSchemaType.KEYWORD), ("scene", PayloadSchemaType.KEYWORD),
("policy_ref", PayloadSchemaType.KEYWORD),
("rule_version", PayloadSchemaType.KEYWORD),
("label", PayloadSchemaType.KEYWORD), ("label", PayloadSchemaType.KEYWORD),
("domain", PayloadSchemaType.KEYWORD), ("domain", PayloadSchemaType.KEYWORD),
("risk_type", PayloadSchemaType.KEYWORD), ("risk_type", PayloadSchemaType.KEYWORD),
@@ -107,7 +121,6 @@ class FewShotStore:
except Exception: except Exception:
logger.debug("payload index 创建跳过 field=%s", field, exc_info=True) logger.debug("payload index 创建跳过 field=%s", field, exc_info=True)
self._ensured = True self._ensured = True
logger.info("few-shot collection 创建成功 dim=%s", dim)
return True return True
except Exception: except Exception:
logger.warning("few-shot collection 初始化失败,本轮操作跳过", exc_info=True) logger.warning("few-shot collection 初始化失败,本轮操作跳过", exc_info=True)
@@ -116,18 +129,31 @@ class FewShotStore:
def upsert(self, sample: Any) -> str | None: def upsert(self, sample: Any) -> str | None:
"""把一条样本向量化并写入 Qdrant返回 vector_id失败返回 None。""" """把一条样本向量化并写入 Qdrant返回 vector_id失败返回 None。"""
tenant_id = str(getattr(sample, "tenant_id", "") or "").strip()
sample_id = str(getattr(sample, "id", "") or "").strip()
if not tenant_id or not sample_id:
logger.warning("few-shot upsert 缺少 tenant_id/sample_id已拒绝")
return None
if not self._ensure_collection(): if not self._ensure_collection():
return None return None
client = self._client client = self._client
try: try:
vector = self._embedding_provider.embed([sample.case_text])[0] vector = self._embedding_provider.embed([sample.case_text])[0]
except Exception: except Exception:
logger.warning("few-shot embedding 失败 sample_key=%s", getattr(sample, "sample_key", ""), exc_info=True) logger.warning(
"few-shot embedding 失败 sample_key=%s",
getattr(sample, "sample_key", ""),
exc_info=True,
)
return None return None
vector_id = uuid.uuid4().hex vector_id = stable_vector_id(tenant_id=tenant_id, sample_id=sample_id)
previous_vector_id = str(getattr(sample, "vector_id", "") or "").strip()
payload = { payload = {
"sample_id": sample.id, "sample_id": sample_id,
"tenant_id": tenant_id,
"scene": sample.scene, "scene": sample.scene,
"policy_ref": getattr(sample, "policy_ref", "") or "",
"rule_version": getattr(sample, "rule_version", "") or "",
"label": sample.label, "label": sample.label,
"domain": sample.domain, "domain": sample.domain,
"risk_type": sample.risk_type, "risk_type": sample.risk_type,
@@ -137,25 +163,51 @@ class FewShotStore:
"payload_json": sample.payload_json, "payload_json": sample.payload_json,
} }
try: try:
points = [{"id": vector_id, "vector": vector, "payload": payload}]
if previous_vector_id and previous_vector_id != vector_id:
# 先用最新判定覆盖旧 point避免后续删除短暂失败时暴露陈旧标签。
points.append({"id": previous_vector_id, "vector": vector, "payload": payload})
client.upsert( client.upsert(
collection_name=FEW_SHOT_COLLECTION, collection_name=FEW_SHOT_COLLECTION,
points=[{"id": vector_id, "vector": vector, "payload": payload}], points=points,
) )
if previous_vector_id and previous_vector_id != vector_id:
try:
client.delete(
collection_name=FEW_SHOT_COLLECTION,
points_selector=[previous_vector_id],
)
except Exception:
logger.warning(
"few-shot 旧向量清理失败,已保留同内容副本 vector_id=%s",
previous_vector_id,
exc_info=True,
)
return vector_id return vector_id
except Exception: except Exception:
logger.warning("few-shot upsert 失败 sample_key=%s", getattr(sample, "sample_key", ""), exc_info=True) logger.warning(
"few-shot upsert 失败 sample_key=%s",
getattr(sample, "sample_key", ""),
exc_info=True,
)
return None return None
def search( def search(
self, self,
case_text: str, case_text: str,
*, *,
tenant_id: str,
scene: str | None = None, scene: str | None = None,
policy_ref: str | None = None,
rule_version: str | None = None,
labels: list[str] | None = None, labels: list[str] | None = None,
top_k: int = 3, top_k: int = 3,
) -> list[dict[str, Any]]: ) -> list[dict[str, Any]]:
"""按 case_text 检索相似样本,可按 scene/label 过滤。失败返回空列表。""" """按 case_text 检索相似样本,可按 scene/label 过滤。失败返回空列表。"""
tenant = str(tenant_id or "").strip()
if not tenant:
raise ValueError("tenant_id is required for few-shot search")
if not case_text or not self._ensure_collection(): if not case_text or not self._ensure_collection():
return [] return []
client = self._client client = self._client
@@ -164,9 +216,16 @@ class FewShotStore:
except Exception: except Exception:
logger.warning("few-shot 检索 embedding 失败", exc_info=True) logger.warning("few-shot 检索 embedding 失败", exc_info=True)
return [] return []
must: list[dict[str, Any]] = [{"key": "status", "match": {"value": "active"}}] must: list[dict[str, Any]] = [
{"key": "tenant_id", "match": {"value": tenant}},
{"key": "status", "match": {"value": "active"}},
]
if scene: if scene:
must.append({"key": "scene", "match": {"value": scene}}) must.append({"key": "scene", "match": {"value": scene}})
if policy_ref:
must.append({"key": "policy_ref", "match": {"value": policy_ref}})
if rule_version:
must.append({"key": "rule_version", "match": {"value": rule_version}})
if labels: if labels:
must.append({"key": "label", "match": {"any": labels}}) must.append({"key": "label", "match": {"any": labels}})
try: try:
@@ -188,6 +247,10 @@ class FewShotStore:
hits.append( hits.append(
{ {
"sample_id": payload.get("sample_id"), "sample_id": payload.get("sample_id"),
"tenant_id": payload.get("tenant_id"),
"scene": payload.get("scene"),
"policy_ref": payload.get("policy_ref") or "",
"rule_version": payload.get("rule_version") or "",
"score": float(getattr(point, "score", 0.0)), "score": float(getattr(point, "score", 0.0)),
"label": payload.get("label"), "label": payload.get("label"),
"domain": payload.get("domain"), "domain": payload.get("domain"),

View File

@@ -14,6 +14,7 @@ from app.core.logging import get_logger
from app.models.financial_record import ExpenseClaim from app.models.financial_record import ExpenseClaim
from app.models.hermes_report import HermesRiskReport from app.models.hermes_report import HermesRiskReport
from app.services.expense_claim_risk_stage import with_risk_business_stage from app.services.expense_claim_risk_stage import with_risk_business_stage
from app.services.expense_claim_tenant_scope import ExpenseClaimTenantScopeMixin
from app.services.risk_observations import RiskObservationService from app.services.risk_observations import RiskObservationService
logger = get_logger("app.services.hermes_risk_scanner") logger = get_logger("app.services.hermes_risk_scanner")
@@ -38,43 +39,60 @@ class HermesRiskScannerService:
logger.info(f"Fetched {len(claims)} claims to analyze.") logger.info(f"Fetched {len(claims)} claims to analyze.")
observation_service = RiskObservationService(self.db) observation_service = RiskObservationService(self.db)
result = evaluate_financial_risk_graph( observation_count = 0
RiskGraphEvaluationContext( graph_node_count = 0
claims=[RiskGraphClaimSnapshot.from_orm(claim) for claim in claims], graph_edge_count = 0
target_claim_ids={claim.id for claim in claims}, for tenant_id, tenant_claims in self._group_claims_by_tenant(claims).items():
history_stats=observation_service.build_history_stats( result = evaluate_financial_risk_graph(
expense_types={str(claim.expense_type or "") for claim in claims}, RiskGraphEvaluationContext(
), claims=[
) RiskGraphClaimSnapshot.from_orm(claim)
) for claim in tenant_claims
claims_by_id = {claim.id: claim for claim in claims} ],
target_claim_ids={claim.id for claim in tenant_claims},
for observation in result.observations: history_stats=observation_service.build_history_stats(
claim = claims_by_id.get(observation.claim_id) tenant_id=tenant_id,
if claim is None: expense_types={
continue str(claim.expense_type or "") for claim in tenant_claims
observation_service.upsert_observation( },
observation, ),
run_id=run_id,
execution_log_id=log_id,
)
claim.hermes_risk_flag = True
claim.risk_flags_json = self._append_algorithm_flag(claim, observation.as_dict())
if log_id:
self.db.add(
HermesRiskReport(
claim_id=observation.claim_id,
execution_log_id=log_id,
risk_level=observation.risk_level,
risk_type=observation.risk_signal,
risk_description=observation.description,
related_claim_ids=[
observation.claim_id,
*observation.similar_case_claim_ids,
],
)
) )
)
claims_by_id = {claim.id: claim for claim in tenant_claims}
observation_count += len(result.observations)
graph_node_count += len(result.nodes)
graph_edge_count += len(result.edges)
for observation in result.observations:
claim = claims_by_id.get(observation.claim_id)
if claim is None:
continue
observation_service.upsert_observation(
observation,
tenant_id=tenant_id,
run_id=run_id,
execution_log_id=log_id,
)
claim.hermes_risk_flag = True
claim.risk_flags_json = self._append_algorithm_flag(
claim,
observation.as_dict(),
)
if log_id:
self.db.add(
HermesRiskReport(
claim_id=observation.claim_id,
execution_log_id=log_id,
risk_level=observation.risk_level,
risk_type=observation.risk_signal,
risk_description=observation.description,
related_claim_ids=[
observation.claim_id,
*observation.similar_case_claim_ids,
],
)
)
now = datetime.now(timezone.utc) now = datetime.now(timezone.utc)
for claim in claims: for claim in claims:
@@ -83,15 +101,28 @@ class HermesRiskScannerService:
self.db.commit() self.db.commit()
logger.info( logger.info(
"Hermes risk graph scan completed. Found %s observations.", "Hermes risk graph scan completed. Found %s observations.",
len(result.observations), observation_count,
) )
return { return {
"scanned_claim_count": len(claims), "scanned_claim_count": len(claims),
"risk_observation_count": len(result.observations), "risk_observation_count": observation_count,
"graph_node_count": len(result.nodes), "graph_node_count": graph_node_count,
"graph_edge_count": len(result.edges), "graph_edge_count": graph_edge_count,
} }
def _group_claims_by_tenant(
self,
claims: list[ExpenseClaim],
) -> dict[str, list[ExpenseClaim]]:
grouped: dict[str, list[ExpenseClaim]] = {}
for claim in claims:
tenant_id = ExpenseClaimTenantScopeMixin.resolve_claim_tenant_id(
self.db,
claim.id,
)
grouped.setdefault(tenant_id, []).append(claim)
return grouped
def _fetch_unscanned_claims(self) -> list[ExpenseClaim]: def _fetch_unscanned_claims(self) -> list[ExpenseClaim]:
stmt = ( stmt = (
select(ExpenseClaim) select(ExpenseClaim)

View File

@@ -0,0 +1,82 @@
from __future__ import annotations
from collections.abc import Iterator
from contextlib import contextmanager
from hashlib import sha256
from threading import Lock, RLock
from sqlalchemy import func, select
from sqlalchemy.orm import Session
_FALLBACK_LOCKS_GUARD = Lock()
_FALLBACK_LOCKS: dict[str, RLock] = {}
@contextmanager
def organization_memory_operation_locks(
db: Session,
*lock_keys: str,
) -> Iterator[None]:
"""串行化组织记忆作用域和幂等键,锁的生命周期覆盖当前事务。"""
normalized_keys = sorted({str(key) for key in lock_keys if str(key)})
bind = db.get_bind()
if bind.dialect.name == "postgresql":
for lock_key in normalized_keys:
db.execute(
select(
func.pg_advisory_xact_lock(
organization_memory_advisory_lock_id(lock_key)
)
)
)
yield
return
# SQLite 等方言没有事务级 advisory lock。进程内锁配合数据库唯一索引
# 提供安全退化,保证测试与单进程部署不会出现空集合竞态。
fallback_locks = [_fallback_lock(lock_key) for lock_key in normalized_keys]
for fallback_lock in fallback_locks:
fallback_lock.acquire()
try:
yield
finally:
for fallback_lock in reversed(fallback_locks):
fallback_lock.release()
def organization_memory_scope_lock_key(
*,
tenant_id: str,
scope_type: str,
scope_id: str,
scene: str,
field_key: str,
) -> str:
return "|".join(
(
"organization-memory-scope",
tenant_id,
scope_type,
scope_id,
scene,
field_key,
)
)
def organization_memory_request_lock_key(tenant_id: str, request_id: str) -> str:
return f"organization-memory-request|{tenant_id}|{request_id}"
def _fallback_lock(lock_key: str) -> RLock:
with _FALLBACK_LOCKS_GUARD:
return _FALLBACK_LOCKS.setdefault(lock_key, RLock())
def organization_memory_advisory_lock_id(lock_key: str) -> int:
return int.from_bytes(
sha256(lock_key.encode("utf-8")).digest()[:8],
byteorder="big",
signed=True,
)

View File

@@ -11,6 +11,7 @@ from sqlalchemy.orm import Session, joinedload
from app.algorithem.risk_graph import RiskHistoryStats, RiskObservationDraft from app.algorithem.risk_graph import RiskHistoryStats, RiskObservationDraft
from app.core.logging import get_logger from app.core.logging import get_logger
from app.db.base import Base from app.db.base import Base
from app.models.expense_case import ExpenseCaseLink
from app.models.financial_record import ExpenseClaim from app.models.financial_record import ExpenseClaim
from app.models.risk_observation import RiskObservation, RiskObservationFeedback from app.models.risk_observation import RiskObservation, RiskObservationFeedback
from app.schemas.risk_observation import ( from app.schemas.risk_observation import (
@@ -34,6 +35,7 @@ FEEDBACK_STATUS_MAP = {
"ignore": ("ignored", "ignored"), "ignore": ("ignored", "ignored"),
"resolve": ("resolved", "resolved"), "resolve": ("resolved", "resolved"),
} }
DEFAULT_TENANT_ID = "default"
class RiskObservationService: class RiskObservationService:
@@ -61,6 +63,7 @@ class RiskObservationService:
self, self,
observation: RiskObservationDraft | dict[str, Any], observation: RiskObservationDraft | dict[str, Any],
*, *,
tenant_id: str | None = None,
run_id: str | None = None, run_id: str | None = None,
execution_log_id: str | None = None, execution_log_id: str | None = None,
) -> RiskObservation: ) -> RiskObservation:
@@ -73,12 +76,22 @@ class RiskObservationService:
observation_key = str(payload.get("observation_key") or "").strip() observation_key = str(payload.get("observation_key") or "").strip()
if not observation_key: if not observation_key:
raise ValueError("Risk observation requires observation_key.") raise ValueError("Risk observation requires observation_key.")
normalized_tenant_id = self._resolve_tenant_id(
tenant_id=tenant_id or _optional_text(payload.get("tenant_id")),
claim_id=_optional_text(payload.get("claim_id")),
)
item = self.db.scalar( item = self.db.scalar(
select(RiskObservation).where(RiskObservation.observation_key == observation_key) select(RiskObservation).where(
RiskObservation.tenant_id == normalized_tenant_id,
RiskObservation.observation_key == observation_key,
)
) )
if item is None: if item is None:
item = RiskObservation(observation_key=observation_key) item = RiskObservation(
tenant_id=normalized_tenant_id,
observation_key=observation_key,
)
self.db.add(item) self.db.add(item)
item.subject_type = _text(payload.get("subject_type")) item.subject_type = _text(payload.get("subject_type"))
@@ -118,9 +131,14 @@ class RiskObservationService:
claim: ExpenseClaim, claim: ExpenseClaim,
flags: list[dict[str, Any]], flags: list[dict[str, Any]],
*, *,
tenant_id: str | None = None,
run_id: str | None = None, run_id: str | None = None,
execution_log_id: str | None = None, execution_log_id: str | None = None,
) -> list[RiskObservation]: ) -> list[RiskObservation]:
normalized_tenant_id = self._resolve_tenant_id(
tenant_id=tenant_id,
claim_id=claim.id,
)
observations: list[RiskObservation] = [] observations: list[RiskObservation] = []
for flag in flags: for flag in flags:
if not isinstance(flag, dict): if not isinstance(flag, dict):
@@ -187,6 +205,7 @@ class RiskObservationService:
"action": _text(flag.get("action")), "action": _text(flag.get("action")),
}, },
}, },
tenant_id=normalized_tenant_id,
run_id=run_id, run_id=run_id,
execution_log_id=execution_log_id, execution_log_id=execution_log_id,
) )
@@ -196,6 +215,7 @@ class RiskObservationService:
def build_history_stats( def build_history_stats(
self, self,
*, *,
tenant_id: str | None = None,
risk_signals: set[str] | None = None, risk_signals: set[str] | None = None,
expense_types: set[str] | None = None, expense_types: set[str] | None = None,
limit: int = 2000, limit: int = 2000,
@@ -204,6 +224,7 @@ class RiskObservationService:
stmt = ( stmt = (
select(RiskObservation, ExpenseClaim.expense_type) select(RiskObservation, ExpenseClaim.expense_type)
.outerjoin(ExpenseClaim, RiskObservation.claim_id == ExpenseClaim.id) .outerjoin(ExpenseClaim, RiskObservation.claim_id == ExpenseClaim.id)
.where(RiskObservation.tenant_id == _normalize_tenant_id(tenant_id))
.order_by(RiskObservation.created_at.desc()) .order_by(RiskObservation.created_at.desc())
.limit(limit) .limit(limit)
) )
@@ -238,6 +259,7 @@ class RiskObservationService:
def list_observations( def list_observations(
self, self,
*, *,
tenant_id: str | None = None,
claim_id: str | None = None, claim_id: str | None = None,
run_id: str | None = None, run_id: str | None = None,
execution_log_id: str | None = None, execution_log_id: str | None = None,
@@ -249,7 +271,7 @@ class RiskObservationService:
offset: int = 0, offset: int = 0,
) -> tuple[list[RiskObservation], int]: ) -> tuple[list[RiskObservation], int]:
self.ensure_storage_ready() self.ensure_storage_ready()
conditions = [] conditions = [RiskObservation.tenant_id == _normalize_tenant_id(tenant_id)]
if claim_id: if claim_id:
conditions.append(RiskObservation.claim_id == claim_id) conditions.append(RiskObservation.claim_id == claim_id)
if run_id: if run_id:
@@ -270,31 +292,52 @@ class RiskObservationService:
RiskObservation.risk_score.desc(), RiskObservation.risk_score.desc(),
RiskObservation.created_at.desc(), RiskObservation.created_at.desc(),
) )
if conditions: count_stmt = count_stmt.where(*conditions)
count_stmt = count_stmt.where(*conditions) stmt = stmt.where(*conditions)
stmt = stmt.where(*conditions)
total = int(self.db.scalar(count_stmt) or 0) total = int(self.db.scalar(count_stmt) or 0)
items = list(self.db.scalars(stmt.offset(offset).limit(limit)).all()) items = list(self.db.scalars(stmt.offset(offset).limit(limit)).all())
return items, total return items, total
def get_observation(self, observation_key_or_id: str) -> RiskObservation | None: def get_observation(
self,
observation_key_or_id: str,
*,
tenant_id: str | None = None,
) -> RiskObservation | None:
self.ensure_storage_ready() self.ensure_storage_ready()
value = str(observation_key_or_id or "").strip() value = str(observation_key_or_id or "").strip()
if not value: if not value:
return None return None
return self.db.scalar( return self.db.scalar(
select(RiskObservation).where( select(RiskObservation).where(
(RiskObservation.observation_key == value) | (RiskObservation.id == value) RiskObservation.tenant_id == _normalize_tenant_id(tenant_id),
(RiskObservation.observation_key == value) | (RiskObservation.id == value),
) )
) )
def list_claim_observations(self, claim_id: str) -> list[RiskObservation]: def list_claim_observations(
items, _ = self.list_observations(claim_id=claim_id, limit=100, offset=0) self,
claim_id: str,
*,
tenant_id: str | None = None,
) -> list[RiskObservation]:
items, _ = self.list_observations(
tenant_id=tenant_id,
claim_id=claim_id,
limit=100,
offset=0,
)
return items return items
def list_execution_log_observations(self, execution_log_id: str) -> list[RiskObservation]: def list_execution_log_observations(
self,
execution_log_id: str,
*,
tenant_id: str | None = None,
) -> list[RiskObservation]:
items, _ = self.list_observations( items, _ = self.list_observations(
tenant_id=tenant_id,
execution_log_id=execution_log_id, execution_log_id=execution_log_id,
limit=200, limit=200,
offset=0, offset=0,
@@ -305,9 +348,15 @@ class RiskObservationService:
self, self,
observation_key_or_id: str, observation_key_or_id: str,
payload: RiskObservationFeedbackCreate, payload: RiskObservationFeedbackCreate,
*,
tenant_id: str | None = None,
actor: str | None = None,
) -> RiskObservationFeedback: ) -> RiskObservationFeedback:
self.ensure_storage_ready() self.ensure_storage_ready()
observation = self.get_observation(observation_key_or_id) observation = self.get_observation(
observation_key_or_id,
tenant_id=tenant_id,
)
if observation is None: if observation is None:
raise LookupError("Risk observation not found.") raise LookupError("Risk observation not found.")
@@ -315,7 +364,7 @@ class RiskObservationService:
observation_id=observation.id, observation_id=observation.id,
feedback_type=payload.feedback_type, feedback_type=payload.feedback_type,
action=payload.action or "", action=payload.action or "",
actor=payload.actor or "", actor=_text(actor) or "system",
comment=payload.comment, comment=payload.comment,
payload_json=payload.payload_json, payload_json=payload.payload_json,
) )
@@ -336,7 +385,8 @@ class RiskObservationService:
) -> None: ) -> None:
"""人工确认/误报后把样本沉淀进 few-shot 池,任何失败都不影响主流程。""" """人工确认/误报后把样本沉淀进 few-shot 池,任何失败都不影响主流程。"""
if os.environ.get("FEW_SHOT_INJECTION_ENABLED", "true").strip().lower() in {"0", "false", "no"}: few_shot_enabled = os.environ.get("FEW_SHOT_INJECTION_ENABLED", "true")
if few_shot_enabled.strip().lower() in {"0", "false", "no"}:
return return
if observation.feedback_status not in {"confirmed", "false_positive"}: if observation.feedback_status not in {"confirmed", "false_positive"}:
return return
@@ -350,15 +400,20 @@ class RiskObservationService:
def summarize_dashboard( def summarize_dashboard(
self, self,
*, *,
tenant_id: str | None = None,
window_days: int = 30, window_days: int = 30,
limit: int = 500, limit: int = 500,
) -> RiskObservationDashboardRead: ) -> RiskObservationDashboardRead:
self.ensure_storage_ready() self.ensure_storage_ready()
normalized_tenant_id = _normalize_tenant_id(tenant_id)
since = datetime.now(UTC) - timedelta(days=window_days) since = datetime.now(UTC) - timedelta(days=window_days)
stmt = ( stmt = (
select(RiskObservation) select(RiskObservation)
.options(joinedload(RiskObservation.claim)) .options(joinedload(RiskObservation.claim))
.where(RiskObservation.created_at >= since) .where(
RiskObservation.tenant_id == normalized_tenant_id,
RiskObservation.created_at >= since,
)
.order_by(RiskObservation.created_at.desc()) .order_by(RiskObservation.created_at.desc())
.limit(limit) .limit(limit)
) )
@@ -371,7 +426,14 @@ class RiskObservationService:
self.db.scalar( self.db.scalar(
select(func.count()) select(func.count())
.select_from(RiskObservationFeedback) .select_from(RiskObservationFeedback)
.where(RiskObservationFeedback.created_at >= since) .join(
RiskObservation,
RiskObservation.id == RiskObservationFeedback.observation_id,
)
.where(
RiskObservation.tenant_id == normalized_tenant_id,
RiskObservationFeedback.created_at >= since,
)
) )
or 0 or 0
) )
@@ -435,6 +497,28 @@ class RiskObservationService:
][:10], ][:10],
) )
def _resolve_tenant_id(
self,
*,
tenant_id: str | None,
claim_id: str | None,
) -> str:
explicit_tenant_id = str(tenant_id or "").strip()
normalized_claim_id = str(claim_id or "").strip()
linked_tenant_id = self.db.scalar(
select(ExpenseCaseLink.tenant_id).where(
ExpenseCaseLink.resource_type == "expense_claim",
ExpenseCaseLink.resource_id == normalized_claim_id,
)
) if normalized_claim_id else None
claim_tenant_id = _normalize_tenant_id(linked_tenant_id)
if explicit_tenant_id:
normalized_tenant_id = _normalize_tenant_id(explicit_tenant_id)
if linked_tenant_id and claim_tenant_id != normalized_tenant_id:
raise PermissionError("Risk observation tenant does not match claim tenant.")
return normalized_tenant_id
return claim_tenant_id
def _count_by(items: list[RiskObservation], field: str) -> dict[str, int]: def _count_by(items: list[RiskObservation], field: str) -> dict[str, int]:
counts: dict[str, int] = {} counts: dict[str, int] = {}
@@ -537,7 +621,8 @@ def _supplier_names(item: RiskObservation) -> list[str]:
names.append(text.split(":", 1)[1] or text) names.append(text.split(":", 1)[1] or text)
for evidence in item.evidence_json or []: for evidence in item.evidence_json or []:
if isinstance(evidence, dict): if isinstance(evidence, dict):
metadata = evidence.get("metadata") if isinstance(evidence.get("metadata"), dict) else {} metadata_value = evidence.get("metadata")
metadata = metadata_value if isinstance(metadata_value, dict) else {}
for key in ("supplier_name", "vendor_name", "merchant_name", "supplier", "vendor"): for key in ("supplier_name", "vendor_name", "merchant_name", "supplier", "vendor"):
name = _text(evidence.get(key)) or _text(metadata.get(key)) name = _text(evidence.get(key)) or _text(metadata.get(key))
if name: if name:
@@ -603,6 +688,10 @@ def _text(value: Any) -> str:
return str(value or "").strip() return str(value or "").strip()
def _normalize_tenant_id(value: Any) -> str:
return _text(value) or DEFAULT_TENANT_ID
def _canonical_key(value: Any) -> str: def _canonical_key(value: Any) -> str:
return "_".join(_text(value).lower().split()) return "_".join(_text(value).lower().split())

View File

@@ -14,7 +14,10 @@ from app.services.agent_asset_rule_library import AgentAssetRuleLibraryManager
from app.services.agent_asset_spreadsheet import RISK_RULES_LIBRARY from app.services.agent_asset_spreadsheet import RISK_RULES_LIBRARY
from app.services.audit import AuditLogService from app.services.audit import AuditLogService
from app.services.expense_claim_risk_stage import infer_risk_domain from app.services.expense_claim_risk_stage import infer_risk_domain
from app.services.risk_rule_dsl_validator import validate_risk_rule_draft
from app.services.risk_rule_explainability import build_risk_rule_explainability_artifacts from app.services.risk_rule_explainability import build_risk_rule_explainability_artifacts
from app.services.risk_rule_generation_interpreter import COMPOSITE_RULE_TEMPLATE_KEY
from app.services.risk_rule_generation_markdown import build_risk_rule_version_markdown
from app.services.risk_rule_generation_ontology import ( from app.services.risk_rule_generation_ontology import (
BUSINESS_DOMAIN_LABELS, BUSINESS_DOMAIN_LABELS,
DOMAIN_FIELD_PREFIXES, DOMAIN_FIELD_PREFIXES,
@@ -26,16 +29,13 @@ from app.services.risk_rule_generation_ontology import (
RiskRuleField, RiskRuleField,
) )
from app.services.risk_rule_generation_prompt import build_risk_rule_compiler_messages from app.services.risk_rule_generation_prompt import build_risk_rule_compiler_messages
from app.services.risk_rule_generation_interpreter import COMPOSITE_RULE_TEMPLATE_KEY from app.services.risk_rule_generation_semantic_plan import unwrap_semantic_plan_payload
from app.services.risk_rule_generation_markdown import build_risk_rule_version_markdown
from app.services.risk_rule_generation_semantics import ( from app.services.risk_rule_generation_semantics import (
CITY_CONSISTENCY_SEMANTIC_TYPE, CITY_CONSISTENCY_SEMANTIC_TYPE,
CITY_CONSISTENCY_SEMANTIC_TYPES, CITY_CONSISTENCY_SEMANTIC_TYPES,
build_city_consistency_draft, build_city_consistency_draft,
build_city_consistency_params, build_city_consistency_params,
) )
from app.services.risk_rule_generation_semantic_plan import unwrap_semantic_plan_payload
from app.services.risk_rule_dsl_validator import validate_risk_rule_draft
from app.services.risk_rule_scoring import apply_risk_score_to_draft, calculate_risk_rule_score from app.services.risk_rule_scoring import apply_risk_score_to_draft, calculate_risk_rule_score
from app.services.runtime_chat import RuntimeChatService from app.services.runtime_chat import RuntimeChatService
@@ -57,6 +57,7 @@ class RiskRuleGenerationService:
self, self,
body: AgentAssetRiskRuleGenerateRequest, body: AgentAssetRiskRuleGenerateRequest,
*, *,
tenant_id: str | None = None,
actor: str, actor: str,
request_id: str | None = None, request_id: str | None = None,
) -> str: ) -> str:
@@ -81,6 +82,7 @@ class RiskRuleGenerationService:
created_at = datetime.now(UTC) created_at = datetime.now(UTC)
fields = self._resolve_fields(natural_language, domain=domain) fields = self._resolve_fields(natural_language, domain=domain)
draft = self._compile_with_model( draft = self._compile_with_model(
tenant_id=tenant_id,
natural_language=natural_language, natural_language=natural_language,
domain=domain, domain=domain,
business_stage=business_stage, business_stage=business_stage,
@@ -174,6 +176,7 @@ class RiskRuleGenerationService:
"ontology_signal": payload.get("ontology_signal"), "ontology_signal": payload.get("ontology_signal"),
"evaluator": payload.get("evaluator"), "evaluator": payload.get("evaluator"),
"generated_by": "natural_language", "generated_by": "natural_language",
"tenant_id": str(tenant_id or "").strip(),
"source_ref": "自然语言风险规则", "source_ref": "自然语言风险规则",
"last_operation": { "last_operation": {
"action": "create", "action": "create",
@@ -217,6 +220,7 @@ class RiskRuleGenerationService:
def _compile_with_model( def _compile_with_model(
self, self,
*, *,
tenant_id: str | None,
natural_language: str, natural_language: str,
domain: str, domain: str,
business_stage: str, business_stage: str,
@@ -235,6 +239,7 @@ class RiskRuleGenerationService:
for item in fields for item in fields
] ]
few_shot_samples = self._retrieve_few_shot_samples( few_shot_samples = self._retrieve_few_shot_samples(
tenant_id=tenant_id,
domain=domain, domain=domain,
natural_language=natural_language, natural_language=natural_language,
) )
@@ -271,6 +276,7 @@ class RiskRuleGenerationService:
def _retrieve_few_shot_samples( def _retrieve_few_shot_samples(
self, self,
*, *,
tenant_id: str | None,
domain: str, domain: str,
natural_language: str, natural_language: str,
) -> list[dict[str, Any]]: ) -> list[dict[str, Any]]:
@@ -280,11 +286,15 @@ class RiskRuleGenerationService:
if os.environ.get("FEW_SHOT_INJECTION_ENABLED", "true").strip().lower() in {"0", "false", "no"}: if os.environ.get("FEW_SHOT_INJECTION_ENABLED", "true").strip().lower() in {"0", "false", "no"}:
return [] return []
normalized_tenant_id = str(tenant_id or "").strip()
if not normalized_tenant_id:
return []
try: try:
from app.services.few_shot_retrieval import FewShotRetriever from app.services.few_shot_retrieval import FewShotRetriever
retriever = FewShotRetriever.from_session(self.db) retriever = FewShotRetriever.from_session(self.db)
return retriever.retrieve_for_risk_rule_generation( return retriever.retrieve_for_risk_rule_generation(
tenant_id=normalized_tenant_id,
domain=domain, domain=domain,
natural_language=natural_language, natural_language=natural_language,
) )

View File

@@ -43,6 +43,7 @@ class RiskRuleGenerationJobService:
self, self,
body: AgentAssetRiskRuleGenerateRequest, body: AgentAssetRiskRuleGenerateRequest,
*, *,
tenant_id: str | None = None,
actor: str, actor: str,
request_id: str | None = None, request_id: str | None = None,
) -> str: ) -> str:
@@ -96,6 +97,7 @@ class RiskRuleGenerationJobService:
"storage_key": f"rules/{RISK_RULES_LIBRARY}/{file_name}", "storage_key": f"rules/{RISK_RULES_LIBRARY}/{file_name}",
}, },
"generated_by": "natural_language", "generated_by": "natural_language",
"tenant_id": str(tenant_id or "").strip(),
"generation_status": AgentAssetStatus.GENERATING.value, "generation_status": AgentAssetStatus.GENERATING.value,
"generation_started_at": created_at.isoformat(), "generation_started_at": created_at.isoformat(),
"generation_request": self._dump_generation_request(body), "generation_request": self._dump_generation_request(body),
@@ -130,6 +132,7 @@ class RiskRuleGenerationJobService:
asset_id: str, asset_id: str,
body: AgentAssetRiskRuleGenerateRequest, body: AgentAssetRiskRuleGenerateRequest,
*, *,
tenant_id: str | None = None,
actor: str, actor: str,
request_id: str | None = None, request_id: str | None = None,
) -> None: ) -> None:
@@ -137,7 +140,13 @@ class RiskRuleGenerationJobService:
asset = self.db.get(AgentAsset, asset_id) asset = self.db.get(AgentAsset, asset_id)
if asset is None or asset.status != AgentAssetStatus.GENERATING.value: if asset is None or asset.status != AgentAssetStatus.GENERATING.value:
return return
self._complete_rule_asset(asset, body, actor=actor, request_id=request_id) self._complete_rule_asset(
asset,
body,
tenant_id=tenant_id,
actor=actor,
request_id=request_id,
)
except Exception as exc: # noqa: BLE001 - 后台任务必须把失败写回资产状态 except Exception as exc: # noqa: BLE001 - 后台任务必须把失败写回资产状态
self.mark_generation_failed( self.mark_generation_failed(
asset_id, asset_id,
@@ -190,6 +199,7 @@ class RiskRuleGenerationJobService:
asset: AgentAsset, asset: AgentAsset,
body: AgentAssetRiskRuleGenerateRequest, body: AgentAssetRiskRuleGenerateRequest,
*, *,
tenant_id: str | None,
actor: str, actor: str,
request_id: str | None, request_id: str | None,
) -> None: ) -> None:
@@ -205,6 +215,7 @@ class RiskRuleGenerationJobService:
fields = self.generator._resolve_fields(natural_language, domain=domain) fields = self.generator._resolve_fields(natural_language, domain=domain)
draft = self.generator._compile_with_model( draft = self.generator._compile_with_model(
tenant_id=tenant_id,
natural_language=natural_language, natural_language=natural_language,
domain=domain, domain=domain,
business_stage=business_stage, business_stage=business_stage,
@@ -282,6 +293,7 @@ class RiskRuleGenerationJobService:
"ontology_signal": payload.get("ontology_signal"), "ontology_signal": payload.get("ontology_signal"),
"evaluator": payload.get("evaluator"), "evaluator": payload.get("evaluator"),
"generated_by": "natural_language", "generated_by": "natural_language",
"tenant_id": str(tenant_id or "").strip(),
"source_ref": "自然语言风险规则", "source_ref": "自然语言风险规则",
"generation_status": "completed", "generation_status": "completed",
"generation_completed_at": datetime.now(UTC).isoformat(), "generation_completed_at": datetime.now(UTC).isoformat(),

View File

@@ -31,7 +31,11 @@ from app.services.expense_application_draft_events import (
ExpenseApplicationDraftEventService, ExpenseApplicationDraftEventService,
) )
from app.services.expense_application_learning import ExpenseApplicationLearningService from app.services.expense_application_learning import ExpenseApplicationLearningService
from app.services.expense_cases import ExpenseCaseService
from app.services.expense_claim_access_policy import ExpenseClaimAccessPolicy from app.services.expense_claim_access_policy import ExpenseClaimAccessPolicy
from app.services.expense_claim_historical_evidence import (
build_user_agent_historical_evidence_notice,
)
from app.services.expense_claim_risk_stage import with_risk_business_stage from app.services.expense_claim_risk_stage import with_risk_business_stage
from app.services.travel_reimbursement_calculator import TravelReimbursementCalculatorService from app.services.travel_reimbursement_calculator import TravelReimbursementCalculatorService
from app.services.user_agent_application_dates import ( from app.services.user_agent_application_dates import (
@@ -884,6 +888,12 @@ class UserAgentApplicationPersistenceMixin:
if existing is not None: if existing is not None:
return existing return existing
raise raise
# 非默认租户的费用单必须先建立 tenant-scoped Case Link后续
# submit_claim() 的访问策略才能在同一事务内重新查询到刚创建的记录。
ExpenseCaseService(self.db).ensure_case_for_claim(
claim,
tenant_id=current_user.tenant_id,
)
if not submit: if not submit:
_, draft_event = draft_event_service.record( _, draft_event = draft_event_service.record(
payload, payload,
@@ -1349,6 +1359,9 @@ class UserAgentApplicationMixin(UserAgentApplicationSlotMixin, UserAgentApplicat
facts["application_no"] = application_claim.claim_no facts["application_no"] = application_claim.claim_no
facts["application_claim_id"] = application_claim.id facts["application_claim_id"] = application_claim.id
facts["manager_name"] = self._resolve_application_manager_name(payload, application_claim) facts["manager_name"] = self._resolve_application_manager_name(payload, application_claim)
facts["historical_case_evidence_notice"] = (
build_user_agent_historical_evidence_notice(application_claim)
)
return UserAgentResponse( return UserAgentResponse(
answer=self._build_expense_application_answer(payload, facts=facts, step=step), answer=self._build_expense_application_answer(payload, facts=facts, step=step),
citations=[], citations=[],
@@ -1417,6 +1430,9 @@ class UserAgentApplicationMixin(UserAgentApplicationSlotMixin, UserAgentApplicat
if step == "submitted": if step == "submitted":
application_no = str(facts.get("application_no") or "").strip() or self._build_application_claim_no(payload, facts) application_no = str(facts.get("application_no") or "").strip() or self._build_application_claim_no(payload, facts)
manager_name = str(facts.get("manager_name") or "").strip() or "直属领导" manager_name = str(facts.get("manager_name") or "").strip() or "直属领导"
historical_notice = str(
facts.get("historical_case_evidence_notice") or ""
).strip()
submitted_title = ( submitted_title = (
"申请单据已修改并重新提交,已进入审批流程。" "申请单据已修改并重新提交,已进入审批流程。"
if str(facts.get("application_edit_mode") or "").strip().lower() == "true" if str(facts.get("application_edit_mode") or "").strip().lower() == "true"
@@ -1427,6 +1443,7 @@ class UserAgentApplicationMixin(UserAgentApplicationSlotMixin, UserAgentApplicat
submitted_title, submitted_title,
f"系统已推送给 {manager_name} 审核,当前节点:{manager_name}审核中。", f"系统已推送给 {manager_name} 审核,当前节点:{manager_name}审核中。",
f"申请单号:{application_no}", f"申请单号:{application_no}",
*([historical_notice] if historical_notice else []),
"下方是简要单据信息。需要查看完整详情时,请点击快捷方式进入单据详情。", "下方是简要单据信息。需要查看完整详情时,请点击快捷方式进入单据详情。",
] ]
) )

View File

@@ -3,7 +3,10 @@ from __future__ import annotations
import os import os
import re import re
from collections.abc import Iterator from collections.abc import Iterator
from importlib.util import module_from_spec, spec_from_file_location
from pathlib import Path from pathlib import Path
from types import SimpleNamespace
from typing import Any
import pytest import pytest
from alembic.config import Config from alembic.config import Config
@@ -14,16 +17,44 @@ from sqlalchemy.pool import NullPool
from alembic import command from alembic import command
from app.core.config import get_settings from app.core.config import get_settings
from app.db.migration_preflight import MigrationPreflightError, validate_migration_state from app.db.migration_preflight import (
LEGACY_ADOPTABLE_HISTORICAL_CASE_TABLES,
MigrationPreflightError,
validate_migration_state,
)
from app.db.schema_ownership import MIGRATION_OWNED_TABLES, create_legacy_schema from app.db.schema_ownership import MIGRATION_OWNED_TABLES, create_legacy_schema
from app.models.ai_memory import MemoryEntry
from app.models.risk_observation import RiskObservation
MIGRATION_TEST_DATABASE_URL = os.getenv("MIGRATION_TEST_DATABASE_URL", "").strip() MIGRATION_TEST_DATABASE_URL = os.getenv("MIGRATION_TEST_DATABASE_URL", "").strip()
LEGACY_PROBE_TABLE = "legacy_migration_probe_records" LEGACY_PROBE_TABLE = "legacy_migration_probe_records"
HEAD_REVISION = "20260716_0006" HEAD_REVISION = "20260716_0009"
SERVER_DIR = Path(__file__).resolve().parents[1] SERVER_DIR = Path(__file__).resolve().parents[1]
ALEMBIC_INI_PATH = SERVER_DIR / "alembic.ini" ALEMBIC_INI_PATH = SERVER_DIR / "alembic.ini"
class _UnsupportedDialectOperationGuard:
def __init__(self, dialect_name: str = "sqlite") -> None:
self.bind = SimpleNamespace(dialect=SimpleNamespace(name=dialect_name))
self.mutation_calls: list[str] = []
def get_bind(self) -> SimpleNamespace:
return self.bind
def __getattr__(self, name: str) -> Any:
self.mutation_calls.append(name)
raise AssertionError(f"unsupported dialect attempted migration operation: {name}")
def _load_migration_module(filename: str) -> Any:
path = SERVER_DIR / "alembic" / "versions" / filename
spec = spec_from_file_location(f"migration_test_{path.stem}", path)
assert spec is not None and spec.loader is not None
module = module_from_spec(spec)
spec.loader.exec_module(module)
return module
def _normalize_probe_component(value: str) -> str: def _normalize_probe_component(value: str) -> str:
return re.sub(r"[^a-z0-9]+", "-", value.lower()).strip("-") return re.sub(r"[^a-z0-9]+", "-", value.lower()).strip("-")
@@ -53,13 +84,9 @@ def _require_disposable_probe_url(raw_url: str) -> str:
host = _normalize_probe_component(parsed.host or "") host = _normalize_probe_component(parsed.host or "")
database = _normalize_probe_component(parsed.database or "") database = _normalize_probe_component(parsed.database or "")
if not _is_disposable_probe_host(host): if not _is_disposable_probe_host(host):
raise RuntimeError( raise RuntimeError("迁移测试数据库主机名必须使用 migration-probe 或 disposable-probe 前缀")
"迁移测试数据库主机名必须使用 migration-probe 或 disposable-probe 前缀"
)
if not _is_disposable_probe_database(database): if not _is_disposable_probe_database(database):
raise RuntimeError( raise RuntimeError("迁移测试数据库名必须使用 migration-probe 或 disposable-probe 前缀")
"迁移测试数据库名必须使用 migration-probe 或 disposable-probe 前缀"
)
return raw_url return raw_url
@@ -93,13 +120,21 @@ def _alembic_config(database_url: str) -> Config:
def _upgrade_head(database_url: str) -> None: def _upgrade_head(database_url: str) -> None:
_upgrade_revision(database_url, "head")
def _upgrade_revision(database_url: str, revision: str) -> None:
get_settings.cache_clear() get_settings.cache_clear()
command.upgrade(_alembic_config(database_url), "head") command.upgrade(_alembic_config(database_url), revision)
def _downgrade_base(database_url: str) -> None: def _downgrade_base(database_url: str) -> None:
_downgrade_revision(database_url, "base")
def _downgrade_revision(database_url: str, revision: str) -> None:
get_settings.cache_clear() get_settings.cache_clear()
command.downgrade(_alembic_config(database_url), "base") command.downgrade(_alembic_config(database_url), revision)
def _table_names(engine: Engine) -> set[str]: def _table_names(engine: Engine) -> set[str]:
@@ -132,6 +167,28 @@ def _assert_indexes(
assert indexes.get(index_name) == expected_columns assert indexes.get(index_name) == expected_columns
def _assert_postgresql_index_predicate(
engine: Engine,
table_name: str,
index_name: str,
*expected_fragments: str,
) -> None:
with engine.connect() as connection:
index_definition = str(
connection.scalar(
text(
"SELECT indexdef FROM pg_indexes "
"WHERE schemaname = 'public' "
"AND tablename = :table_name AND indexname = :index_name"
),
{"table_name": table_name, "index_name": index_name},
)
or ""
).lower()
for fragment in expected_fragments:
assert fragment.lower() in index_definition
def _assert_check_constraint( def _assert_check_constraint(
engine: Engine, engine: Engine,
table_name: str, table_name: str,
@@ -176,6 +233,20 @@ def _assert_composite_foreign_key(
assert str(matching[0].get("options", {}).get("ondelete", "")).upper() == "RESTRICT" assert str(matching[0].get("options", {}).get("ondelete", "")).upper() == "RESTRICT"
def _assert_no_foreign_key(
engine: Engine,
table_name: str,
constrained_columns: tuple[str, ...],
referred_table: str,
) -> None:
foreign_keys = inspect(engine).get_foreign_keys(table_name, schema="public")
assert not any(
tuple(item["constrained_columns"]) == constrained_columns
and item["referred_table"] == referred_table
for item in foreign_keys
)
def _assert_head_schema(engine: Engine) -> None: def _assert_head_schema(engine: Engine) -> None:
names = _table_names(engine) names = _table_names(engine)
assert MIGRATION_OWNED_TABLES.issubset(names) assert MIGRATION_OWNED_TABLES.issubset(names)
@@ -280,11 +351,48 @@ def _assert_head_schema(engine: Engine) -> None:
"uq_attachment_association_jobs_owner_dedupe", "uq_attachment_association_jobs_owner_dedupe",
("tenant_id", "owner_username", "dedupe_key", "generation"), ("tenant_id", "owner_username", "dedupe_key", "generation"),
) )
_assert_unique_constraint(
engine,
"risk_observations",
"uq_risk_observations_tenant_key",
("tenant_id", "observation_key"),
)
_assert_unique_constraint(
engine,
"few_shot_samples",
"uq_few_shot_samples_tenant_key",
("tenant_id", "sample_key"),
)
_assert_check_constraint( _assert_check_constraint(
engine, engine,
"memory_entries", "memory_entries",
"ck_memory_entries_expired_fields", "ck_memory_entries_expired_fields",
) )
_assert_check_constraint(
engine,
"memory_entries",
"ck_memory_entries_scope_origin",
)
_assert_check_constraint(
engine,
"memory_entries",
"ck_memory_entries_enterprise_scope",
)
_assert_check_constraint(
engine,
"memory_entries",
"ck_memory_entries_management_audit",
)
_assert_check_constraint(
engine,
"memory_entries",
"ck_memory_entries_management_idempotency_pair",
)
_assert_check_constraint(
engine,
"memory_entries",
"ck_memory_entries_revoke_idempotency_pair",
)
_assert_check_constraint( _assert_check_constraint(
engine, engine,
"attachment_association_jobs", "attachment_association_jobs",
@@ -304,6 +412,16 @@ def _assert_head_schema(engine: Engine) -> None:
"ix_expense_cases_tenant_status": ("tenant_id", "status"), "ix_expense_cases_tenant_status": ("tenant_id", "status"),
}, },
) )
_assert_postgresql_index_predicate(
engine,
"memory_entries",
"uq_memory_entries_active_scope",
"status",
"'active'",
"scope_type",
"department",
"enterprise",
)
_assert_indexes( _assert_indexes(
engine, engine,
"expense_case_links", "expense_case_links",
@@ -416,6 +534,21 @@ def _assert_head_schema(engine: Engine) -> None:
"field_key", "field_key",
"status", "status",
), ),
"uq_memory_entries_management_request": (
"tenant_id",
"management_request_id",
),
"uq_memory_entries_revoke_request": (
"tenant_id",
"revoke_request_id",
),
"uq_memory_entries_active_scope": (
"tenant_id",
"scope_type",
"scope_id",
"scene",
"field_key",
),
}, },
) )
_assert_indexes( _assert_indexes(
@@ -429,6 +562,30 @@ def _assert_head_schema(engine: Engine) -> None:
), ),
}, },
) )
_assert_indexes(
engine,
"risk_observations",
{
"ix_risk_observations_tenant_status": (
"tenant_id",
"status",
"created_at",
),
},
)
_assert_indexes(
engine,
"few_shot_samples",
{
"ix_few_shot_samples_tenant_rule_lookup": (
"tenant_id",
"scene",
"policy_ref",
"rule_version",
"status",
),
},
)
_assert_cascade_foreign_key(engine, "expense_case_links") _assert_cascade_foreign_key(engine, "expense_case_links")
_assert_cascade_foreign_key(engine, "business_events") _assert_cascade_foreign_key(engine, "business_events")
_assert_composite_foreign_key( _assert_composite_foreign_key(
@@ -443,6 +600,12 @@ def _assert_head_schema(engine: Engine) -> None:
("tenant_id", "expense_case_id"), ("tenant_id", "expense_case_id"),
"expense_cases", "expense_cases",
) )
_assert_no_foreign_key(
engine,
"risk_observations",
("claim_id",),
"expense_claims",
)
_assert_composite_foreign_key( _assert_composite_foreign_key(
engine, engine,
"ai_decisions", "ai_decisions",
@@ -556,15 +719,19 @@ def _assert_runtime_cascade(engine: Engine) -> None:
""" """
) )
) )
connection.execute( connection.execute(text("DELETE FROM expense_cases WHERE id = 'migration-probe-case'"))
text("DELETE FROM expense_cases WHERE id = 'migration-probe-case'") assert (
connection.scalar(
text("SELECT COUNT(*) FROM expense_case_links WHERE id = 'migration-probe-link'")
)
== 0
)
assert (
connection.scalar(
text("SELECT COUNT(*) FROM business_events WHERE id = 'migration-probe-event'")
)
== 0
) )
assert connection.scalar(
text("SELECT COUNT(*) FROM expense_case_links WHERE id = 'migration-probe-link'")
) == 0
assert connection.scalar(
text("SELECT COUNT(*) FROM business_events WHERE id = 'migration-probe-event'")
) == 0
def _assert_learning_ledger_tenant_boundary(engine: Engine) -> None: def _assert_learning_ledger_tenant_boundary(engine: Engine) -> None:
@@ -777,26 +944,273 @@ def _create_legacy_sentinel(engine: Engine) -> None:
) )
def _create_hierarchical_memory_downgrade_probe(engine: Engine) -> None:
"""验证 0007 在真实组织记忆存在时仍可安全降级。"""
with engine.begin() as connection:
connection.execute(
text(
"""
INSERT INTO memory_entries (
id, tenant_id, scope_type, scope_id, origin_type,
managed_by, managed_at, management_reason, policy_version,
scene, field_key, generation, value_json, value_fingerprint,
status, evidence_count, approved_evidence_count, confidence,
last_evidence_at, candidate_expires_at,
activated_at, active_expires_at
) VALUES (
'hierarchical-memory-downgrade-probe',
'migration-probe', 'enterprise', 'migration-probe',
'admin_managed', 'migration-admin', CURRENT_TIMESTAMP,
'验证组织记忆降级', 'expense_application_transport_org_memory.v1',
'travel_application', 'transport_mode', 1,
CAST('{"value":"火车"}' AS JSON),
'sha256:hierarchical-memory-downgrade-probe',
'active', 0, 0, 1,
CURRENT_TIMESTAMP, CURRENT_TIMESTAMP + INTERVAL '180 days',
CURRENT_TIMESTAMP, CURRENT_TIMESTAMP + INTERVAL '180 days'
)
"""
)
)
assert connection.scalar(
text(
"SELECT COUNT(*) FROM memory_entries "
"WHERE id = 'hierarchical-memory-downgrade-probe'"
)
) == 1
def _create_duplicate_active_organization_memory_probe(engine: Engine) -> None:
with engine.begin() as connection:
connection.execute(
text(
"""
INSERT INTO memory_entries (
id, tenant_id, scope_type, scope_id, origin_type,
managed_by, managed_at, management_reason, policy_version,
scene, field_key, generation, value_json, value_fingerprint,
status, evidence_count, approved_evidence_count, confidence,
last_evidence_at, candidate_expires_at,
activated_at, active_expires_at
) VALUES
(
'duplicate-active-organization-a',
'migration-probe', 'department', 'department-probe',
'admin_managed', 'migration-admin', CURRENT_TIMESTAMP,
'验证组织生效记忆重复预检', 'expense_application_transport_org_memory.v1',
'travel_application', 'transport_mode', 1,
CAST('{"value":"火车"}' AS JSON),
'sha256:duplicate-active-organization-a',
'active', 0, 0, 1,
CURRENT_TIMESTAMP, CURRENT_TIMESTAMP + INTERVAL '180 days',
CURRENT_TIMESTAMP, CURRENT_TIMESTAMP + INTERVAL '180 days'
),
(
'duplicate-active-organization-b',
'migration-probe', 'department', 'department-probe',
'admin_managed', 'migration-admin', CURRENT_TIMESTAMP,
'验证组织生效记忆重复预检', 'expense_application_transport_org_memory.v1',
'travel_application', 'transport_mode', 2,
CAST('{"value":"飞机"}' AS JSON),
'sha256:duplicate-active-organization-b',
'active', 0, 0, 1,
CURRENT_TIMESTAMP, CURRENT_TIMESTAMP + INTERVAL '180 days',
CURRENT_TIMESTAMP, CURRENT_TIMESTAMP + INTERVAL '180 days'
)
"""
)
)
def _delete_duplicate_active_organization_memory_probe(engine: Engine) -> None:
with engine.begin() as connection:
connection.execute(
text(
"DELETE FROM memory_entries "
"WHERE id IN ("
"'duplicate-active-organization-a', "
"'duplicate-active-organization-b'"
")"
)
)
def _create_enriched_few_shot_downgrade_probe(engine: Engine) -> None:
with engine.begin() as connection:
connection.execute(
text(
"""
INSERT INTO few_shot_samples (
id, tenant_id, sample_key, scene,
policy_ref, rule_version, payload_json
) VALUES (
'enriched-few-shot-downgrade-probe', 'default',
'enriched:few-shot:downgrade-probe', 'expense_reimbursement',
'TRAVEL-001', 'v2', CAST('{}' AS JSON)
)
"""
)
)
def _delete_enriched_few_shot_downgrade_probe(engine: Engine) -> None:
with engine.begin() as connection:
connection.execute(
text(
"DELETE FROM few_shot_samples "
"WHERE id = 'enriched-few-shot-downgrade-probe'"
)
)
def _create_historical_case_downgrade_probe(engine: Engine) -> None:
with engine.begin() as connection:
connection.execute(
text(
"""
INSERT INTO risk_observations (
id, tenant_id, observation_key, subject_type, subject_key,
risk_type, risk_signal, risk_level,
contribution_scores_json, baseline_json, evidence_json,
graph_node_keys_json, graph_edge_keys_json, policy_refs_json,
similar_case_claim_ids_json, ontology_json, decision_trace_json
) VALUES (
'historical-downgrade-observation', 'default',
'historical:downgrade:observation', 'expense_claim',
'claim:historical-downgrade', 'duplicate_invoice',
'duplicate_invoice', 'high', CAST('{}' AS JSON), CAST('{}' AS JSON),
CAST('[]' AS JSON), CAST('[]' AS JSON), CAST('[]' AS JSON),
CAST('[]' AS JSON), CAST('[]' AS JSON), CAST('{}' AS JSON),
CAST('{}' AS JSON)
)
"""
)
)
connection.execute(
text(
"""
INSERT INTO risk_observation_feedback (
id, observation_id, feedback_type, payload_json
) VALUES (
'historical-downgrade-feedback',
'historical-downgrade-observation', 'confirm', CAST('{}' AS JSON)
)
"""
)
)
connection.execute(
text(
"""
INSERT INTO few_shot_samples (
id, tenant_id, sample_key, source_observation_id,
scene, policy_ref, rule_version, payload_json
) VALUES (
'historical-downgrade-sample', 'default',
'historical:downgrade:sample', 'historical-downgrade-observation',
'expense_reimbursement', '', '', CAST('{}' AS JSON)
)
"""
)
)
def _assert_historical_case_downgrade_probe(engine: Engine) -> None:
inspector = inspect(engine)
risk_columns = {
str(item["name"]) for item in inspector.get_columns("risk_observations")
}
sample_columns = {
str(item["name"]) for item in inspector.get_columns("few_shot_samples")
}
assert "tenant_id" not in risk_columns
assert {"tenant_id", "policy_ref", "rule_version"}.isdisjoint(sample_columns)
assert not any(
item["constrained_columns"] == ["claim_id"]
and item["referred_table"] == "expense_claims"
and item["referred_columns"] == ["id"]
for item in inspector.get_foreign_keys("risk_observations")
)
with engine.connect() as connection:
assert connection.scalar(
text(
"SELECT COUNT(*) FROM risk_observations "
"WHERE id = 'historical-downgrade-observation'"
)
) == 1
assert connection.scalar(
text(
"SELECT COUNT(*) FROM risk_observation_feedback "
"WHERE id = 'historical-downgrade-feedback'"
)
) == 1
assert connection.scalar(
text(
"SELECT COUNT(*) FROM few_shot_samples "
"WHERE id = 'historical-downgrade-sample'"
)
) == 1
def _assert_legacy_sentinel(engine: Engine) -> None: def _assert_legacy_sentinel(engine: Engine) -> None:
assert LEGACY_PROBE_TABLE in _table_names(engine) assert LEGACY_PROBE_TABLE in _table_names(engine)
with engine.connect() as connection: with engine.connect() as connection:
payload = connection.scalar( payload = connection.scalar(
text( text(f"SELECT payload FROM {LEGACY_PROBE_TABLE} WHERE id = 'legacy-sentinel'")
f"SELECT payload FROM {LEGACY_PROBE_TABLE} "
"WHERE id = 'legacy-sentinel'"
)
) )
assert payload == "must-survive-migration-cycle" assert payload == "must-survive-migration-cycle"
def _assert_base_schema(engine: Engine) -> None: def _assert_base_schema(engine: Engine) -> None:
names = _table_names(engine) names = _table_names(engine)
assert not MIGRATION_OWNED_TABLES.intersection(names) unsafe_owned_tables = (
MIGRATION_OWNED_TABLES - LEGACY_ADOPTABLE_HISTORICAL_CASE_TABLES
).intersection(names)
assert not unsafe_owned_tables
assert LEGACY_ADOPTABLE_HISTORICAL_CASE_TABLES.issubset(names)
assert "alembic_version" in names assert "alembic_version" in names
with engine.connect() as connection: with engine.connect() as connection:
assert connection.scalar(text("SELECT COUNT(*) FROM alembic_version")) == 0 assert connection.scalar(text("SELECT COUNT(*) FROM alembic_version")) == 0
@pytest.mark.parametrize(
("filename", "direction"),
[
("20260716_0008_tenant_safe_historical_cases.py", "upgrade"),
("20260716_0008_tenant_safe_historical_cases.py", "downgrade"),
("20260716_0009_organization_memory_idempotency.py", "upgrade"),
("20260716_0009_organization_memory_idempotency.py", "downgrade"),
],
)
def test_postgresql_only_migrations_reject_other_dialects_before_mutation(
filename: str,
direction: str,
) -> None:
migration = _load_migration_module(filename)
operation_guard = _UnsupportedDialectOperationGuard()
migration.op = operation_guard
with pytest.raises(RuntimeError, match="only supports PostgreSQL"):
getattr(migration, direction)()
assert operation_guard.mutation_calls == []
def test_head_model_declares_soft_claim_reference_and_organization_only_active_index() -> None:
claim_column = RiskObservation.__table__.c.claim_id
assert not claim_column.foreign_keys
claim_relationship = inspect(RiskObservation).relationships["claim"]
assert claim_relationship.viewonly is True
active_scope_index = next(
index
for index in MemoryEntry.__table__.indexes
if index.name == "uq_memory_entries_active_scope"
)
predicate = str(active_scope_index.dialect_options["postgresql"]["where"])
assert "status = 'active'" in predicate
assert "scope_type IN ('department', 'enterprise')" in predicate
@pytest.mark.parametrize( @pytest.mark.parametrize(
"database_url", "database_url",
[ [
@@ -818,6 +1232,23 @@ def test_alembic_migration_cycle_on_disposable_postgres(
assert _table_names(engine) == set(), "迁移测试必须从全新空库开始" assert _table_names(engine) == set(), "迁移测试必须从全新空库开始"
assert validate_migration_state(engine).revision is None assert validate_migration_state(engine).revision is None
_upgrade_revision(migration_database_url, "20260716_0008")
_create_duplicate_active_organization_memory_probe(engine)
with pytest.raises(RuntimeError, match="duplicate active memories"):
_upgrade_head(migration_database_url)
assert validate_migration_state(engine).revision == "20260716_0008"
memory_columns = {
str(item["name"])
for item in inspect(engine).get_columns("memory_entries", schema="public")
}
assert {
"management_request_id",
"management_payload_fingerprint",
"revoke_request_id",
"revoke_payload_fingerprint",
}.isdisjoint(memory_columns)
_delete_duplicate_active_organization_memory_probe(engine)
_upgrade_head(migration_database_url) _upgrade_head(migration_database_url)
_assert_head_schema(engine) _assert_head_schema(engine)
assert validate_migration_state(engine).revision == HEAD_REVISION assert validate_migration_state(engine).revision == HEAD_REVISION
@@ -830,10 +1261,28 @@ def test_alembic_migration_cycle_on_disposable_postgres(
_assert_runtime_cascade(engine) _assert_runtime_cascade(engine)
_create_legacy_sentinel(engine) _create_legacy_sentinel(engine)
_downgrade_revision(migration_database_url, "20260716_0008")
_create_enriched_few_shot_downgrade_probe(engine)
with pytest.raises(RuntimeError, match="contain policy_ref or rule_version data"):
_downgrade_revision(migration_database_url, "20260716_0007")
assert validate_migration_state(engine).revision == "20260716_0008"
with engine.connect() as connection:
assert connection.execute(
text(
"SELECT policy_ref, rule_version FROM few_shot_samples "
"WHERE id = 'enriched-few-shot-downgrade-probe'"
)
).one() == ("TRAVEL-001", "v2")
_delete_enriched_few_shot_downgrade_probe(engine)
_upgrade_head(migration_database_url)
_create_hierarchical_memory_downgrade_probe(engine)
_create_historical_case_downgrade_probe(engine)
_downgrade_base(migration_database_url) _downgrade_base(migration_database_url)
_assert_base_schema(engine) _assert_base_schema(engine)
_assert_legacy_sentinel(engine) _assert_legacy_sentinel(engine)
_assert_historical_case_downgrade_probe(engine)
assert validate_migration_state(engine).revision is None assert validate_migration_state(engine).revision is None
with engine.begin() as connection: with engine.begin() as connection:

View File

@@ -0,0 +1,505 @@
from __future__ import annotations
import uuid
from datetime import UTC, datetime, timedelta
from decimal import Decimal
import pytest
from sqlalchemy import select
from app.api.deps import CurrentUserContext
from app.models.ai_memory import MemoryEntry
from app.models.organization import OrganizationUnit
from app.schemas.expense_application_memory import (
ExpenseApplicationOrganizationMemoryCreate,
ExpenseApplicationOrganizationMemoryRevoke,
ExpenseApplicationOrganizationMemoryUpdate,
)
from app.services.expense_application_memory import ExpenseApplicationMemoryService
from app.services.expense_application_memory_admin import (
ExpenseApplicationOrganizationMemoryService,
OrganizationMemoryConflictError,
OrganizationMemoryNotFoundError,
)
from app.services.organization_memory_locks import (
organization_memory_advisory_lock_id,
organization_memory_operation_locks,
)
from app.test_helpers.db import build_in_memory_session_factory
def _user(
*,
tenant_id: str = "tenant-hierarchy",
department_id: str = "department-delivery",
is_admin: bool = False,
) -> CurrentUserContext:
return CurrentUserContext(
username="memory-admin@example.com" if is_admin else "memory-user@example.com",
name="记忆管理员" if is_admin else "记忆员工",
role_codes=["admin"] if is_admin else ["user"],
is_admin=is_admin,
tenant_id=tenant_id,
employee_id="memory-admin" if is_admin else "memory-user",
employee_no="E-MEMORY",
department_id=department_id,
department_name="交付部",
)
def _active_organization_memory(
*,
tenant_id: str,
scope_type: str,
scope_id: str,
value: str,
generation: int,
activated_at: datetime | None = None,
expires_at: datetime | None = None,
) -> MemoryEntry:
now = datetime.now(UTC)
activated_at = activated_at or now
expires_at = expires_at or now + timedelta(days=180)
return MemoryEntry(
id=str(uuid.uuid4()),
tenant_id=tenant_id,
scope_type=scope_type,
scope_id=scope_id,
origin_type="admin_managed",
managed_by="memory-admin",
managed_at=now,
management_reason="测试组织记忆",
policy_version="expense_application_transport_org_memory.v1",
scene="travel_application",
field_key="transport_mode",
generation=generation,
value_json={"value": value},
value_fingerprint=f"fingerprint-{scope_type}-{generation}-{value}",
status="active",
evidence_count=0,
approved_evidence_count=0,
confidence=Decimal("1.0000"),
last_evidence_at=now,
candidate_expires_at=expires_at,
activated_at=activated_at,
active_expires_at=expires_at,
created_at=activated_at,
updated_at=now,
)
def test_enterprise_precedes_department_and_explicit_input_precedes_all() -> None:
with build_in_memory_session_factory()() as db:
current_user = _user()
db.add_all(
[
_active_organization_memory(
tenant_id=current_user.tenant_id,
scope_type="enterprise",
scope_id=current_user.tenant_id,
value="飞机",
generation=1,
),
_active_organization_memory(
tenant_id=current_user.tenant_id,
scope_type="department",
scope_id=current_user.department_id,
value="火车",
generation=1,
),
]
)
db.flush()
facts: dict[str, object] = {}
applications = ExpenseApplicationMemoryService(db).apply_active_transport_memory(
facts, current_user
)
assert facts == {"transport_mode": "飞机"}
assert applications[0].scope_type == "enterprise"
assert applications[0].priority == 300
assert applications[0].source == "enterprise_policy_memory"
assert applications[0].can_revoke is False
assert applications[0].conflicts[0].reason == "lower_priority_overridden"
assert applications[0].conflicts[0].scope_type == "department"
explicit_facts = {"transport_mode": "轮船"}
assert (
ExpenseApplicationMemoryService(db).apply_active_transport_memory(
explicit_facts, current_user
)
== []
)
assert explicit_facts == {"transport_mode": "轮船"}
def test_same_level_conflict_fails_closed_and_is_explainable() -> None:
with build_in_memory_session_factory()() as db:
current_user = _user()
db.add_all(
[
_active_organization_memory(
tenant_id=current_user.tenant_id,
scope_type="enterprise",
scope_id=current_user.tenant_id,
value="飞机",
generation=1,
),
_active_organization_memory(
tenant_id=current_user.tenant_id,
scope_type="enterprise",
scope_id=current_user.tenant_id,
value="火车",
generation=2,
),
]
)
db.flush()
facts: dict[str, object] = {}
applications = ExpenseApplicationMemoryService(db).apply_active_transport_memory(
facts, current_user
)
assert facts == {}
assert len(applications) == 1
assert applications[0].status == "conflict"
assert applications[0].memory_id == ""
assert applications[0].scope_id == ""
assert applications[0].value == ""
assert len(applications[0].conflicts) == 2
assert {
conflict.reason for conflict in applications[0].conflicts
} == {"same_priority_conflict"}
for conflict in applications[0].model_dump()["conflicts"]:
assert "memory_id" not in conflict
assert "scope_id" not in conflict
assert "value" not in conflict
def test_expiry_and_time_decay_never_apply_cross_tenant_or_expire_early() -> None:
with build_in_memory_session_factory()() as db:
current_user = _user()
now = datetime.now(UTC)
expired = _active_organization_memory(
tenant_id=current_user.tenant_id,
scope_type="department",
scope_id=current_user.department_id,
value="火车",
generation=1,
activated_at=now - timedelta(days=181),
expires_at=now - timedelta(days=1),
)
other_tenant = _active_organization_memory(
tenant_id="tenant-other",
scope_type="enterprise",
scope_id="tenant-other",
value="轮船",
generation=1,
)
decayed = _active_organization_memory(
tenant_id=current_user.tenant_id,
scope_type="enterprise",
scope_id=current_user.tenant_id,
value="飞机",
generation=1,
activated_at=now - timedelta(days=100),
expires_at=now + timedelta(days=100),
)
db.add_all([expired, other_tenant, decayed])
db.flush()
facts: dict[str, object] = {}
applications = ExpenseApplicationMemoryService(db).apply_active_transport_memory(
facts, current_user
)
db.refresh(expired)
assert expired.status == "expired"
assert facts == {"transport_mode": "飞机"}
assert 0.49 <= applications[0].effective_confidence <= 0.51
assert applications[0].memory_id == decayed.id
def test_organization_memory_lifecycle_is_audited_versioned_and_tenant_safe() -> None:
with build_in_memory_session_factory()() as db:
admin = _user(is_admin=True)
service = ExpenseApplicationOrganizationMemoryService(db)
created = service.create_organization_memory(
ExpenseApplicationOrganizationMemoryCreate(
scope_type="enterprise",
value="火车",
expires_in_days=90,
reason="统一优先使用高铁",
request_id="create-enterprise-memory-1",
),
admin,
)
assert created.scope_id == admin.tenant_id
assert created.generation == 1
assert created.origin_type == "admin_managed"
assert created.managed_by == admin.employee_id
assert created.management_reason == "统一优先使用高铁"
replayed_create = service.create_organization_memory(
ExpenseApplicationOrganizationMemoryCreate(
scope_type="enterprise",
value="火车",
expires_in_days=90,
reason="统一优先使用高铁",
request_id="create-enterprise-memory-1",
),
admin,
)
assert replayed_create.id == created.id
with pytest.raises(OrganizationMemoryConflictError, match="已有生效记忆"):
service.create_organization_memory(
ExpenseApplicationOrganizationMemoryCreate(
scope_type="enterprise",
value="火车",
expires_in_days=90,
reason="重复创建企业基线",
request_id="create-enterprise-memory-2",
),
admin,
)
updated = service.update_organization_memory(
created.id,
ExpenseApplicationOrganizationMemoryUpdate(
value="飞机",
expires_in_days=120,
expected_generation=1,
reason="适配跨区域差旅",
request_id="update-enterprise-memory-1",
),
admin,
)
assert updated.id != created.id
assert updated.generation == 2
old_entry = db.get(MemoryEntry, created.id)
assert old_entry is not None
assert old_entry.status == "suppressed"
assert old_entry.superseded_by_id == updated.id
replayed_update = service.update_organization_memory(
created.id,
ExpenseApplicationOrganizationMemoryUpdate(
value="飞机",
expires_in_days=120,
expected_generation=1,
reason="适配跨区域差旅",
request_id="update-enterprise-memory-1",
),
admin,
)
assert replayed_update.id == updated.id
with pytest.raises(OrganizationMemoryConflictError, match="刷新后重试"):
service.update_organization_memory(
created.id,
ExpenseApplicationOrganizationMemoryUpdate(
value="轮船",
expected_generation=1,
reason="使用旧记录地址更新",
request_id="update-enterprise-memory-old-id",
),
admin,
)
with pytest.raises(OrganizationMemoryConflictError, match="刷新后重试"):
service.update_organization_memory(
updated.id,
ExpenseApplicationOrganizationMemoryUpdate(
expected_generation=1,
reason="使用陈旧版本更新",
request_id="update-enterprise-memory-stale",
),
admin,
)
with pytest.raises(OrganizationMemoryNotFoundError):
service.revoke_organization_memory(
updated.id,
ExpenseApplicationOrganizationMemoryRevoke(
expected_generation=2,
reason="越权租户撤销",
request_id="revoke-other-tenant-memory",
),
_user(tenant_id="tenant-other", is_admin=True),
)
revoked = service.revoke_organization_memory(
updated.id,
ExpenseApplicationOrganizationMemoryRevoke(
expected_generation=2,
reason="企业政策已调整",
request_id="revoke-enterprise-memory-1",
),
admin,
)
assert revoked.memory_id == updated.id
revoked_entry = db.get(MemoryEntry, updated.id)
assert revoked_entry is not None
assert revoked_entry.status == "revoked"
assert revoked_entry.value_json == {"value": "飞机"}
assert revoked_entry.revoked_reason == "企业政策已调整"
replayed_revoke = service.revoke_organization_memory(
updated.id,
ExpenseApplicationOrganizationMemoryRevoke(
expected_generation=2,
reason="企业政策已调整",
request_id="revoke-enterprise-memory-1",
),
admin,
)
assert replayed_revoke == revoked
with pytest.raises(OrganizationMemoryConflictError, match="请求内容不一致"):
service.revoke_organization_memory(
updated.id,
ExpenseApplicationOrganizationMemoryRevoke(
expected_generation=2,
reason="同一键却改变请求体",
request_id="revoke-enterprise-memory-1",
),
admin,
)
assert ExpenseApplicationMemoryService(db).revoke_current_user_memory(
updated.id, admin
) is None
def test_department_memory_requires_existing_department_and_admin() -> None:
with build_in_memory_session_factory()() as db:
department = OrganizationUnit(
id="department-delivery",
unit_code="DELIVERY",
name="交付部",
unit_type="department",
)
division = OrganizationUnit(
id="division-east",
unit_code="EAST",
name="华东大区",
unit_type="division",
)
db.add_all([department, division])
db.flush()
service = ExpenseApplicationOrganizationMemoryService(db)
payload = ExpenseApplicationOrganizationMemoryCreate(
scope_type="department",
scope_id=department.id,
value="火车",
reason="部门短途差旅基线",
request_id="create-department-memory-1",
)
with pytest.raises(PermissionError, match="平台管理员"):
service.create_organization_memory(payload, _user())
with pytest.raises(ValueError, match="department 类型"):
service.create_organization_memory(
payload.model_copy(update={"scope_id": division.id}),
_user(is_admin=True),
)
created = service.create_organization_memory(payload, _user(is_admin=True))
assert created.scope_id == department.id
assert created.scope_label == "部门规则(交付部)"
memories = service.list_organization_memories(_user(is_admin=True))
assert [item.id for item in memories.items] == [created.id]
assert db.scalar(
select(MemoryEntry).where(MemoryEntry.id == created.id)
) is not None
def test_expired_organization_memory_revoke_returns_conflict_and_marks_expired() -> None:
with build_in_memory_session_factory()() as db:
admin = _user(is_admin=True)
now = datetime.now(UTC)
entry = _active_organization_memory(
tenant_id=admin.tenant_id,
scope_type="enterprise",
scope_id=admin.tenant_id,
value="火车",
generation=1,
activated_at=now - timedelta(days=100),
expires_at=now - timedelta(days=1),
)
db.add(entry)
db.commit()
with pytest.raises(OrganizationMemoryConflictError, match="已过期"):
ExpenseApplicationOrganizationMemoryService(db).revoke_organization_memory(
entry.id,
ExpenseApplicationOrganizationMemoryRevoke(
expected_generation=1,
reason="撤销已过期规则",
request_id="revoke-expired-memory-1",
),
admin,
)
db.refresh(entry)
assert entry.status == "expired"
assert entry.expired_at is not None
def test_postgres_organization_mutations_use_signed_transaction_advisory_locks() -> None:
executed: list[object] = []
class _Dialect:
name = "postgresql"
class _Bind:
dialect = _Dialect()
class _RecordingSession:
@staticmethod
def get_bind() -> _Bind:
return _Bind()
@staticmethod
def execute(statement: object) -> None:
executed.append(statement)
session = _RecordingSession()
with organization_memory_operation_locks( # type: ignore[arg-type]
session,
"scope-lock",
"request-lock",
):
assert len(executed) == 2
assert all("pg_advisory_xact_lock" in str(statement) for statement in executed)
lock_id = organization_memory_advisory_lock_id("scope-lock")
assert -(2**63) <= lock_id < 2**63
assert lock_id == organization_memory_advisory_lock_id("scope-lock")
def test_new_generation_rejects_unexpected_active_set_before_suppression() -> None:
with build_in_memory_session_factory()() as db:
admin = _user(is_admin=True)
active = _active_organization_memory(
tenant_id=admin.tenant_id,
scope_type="enterprise",
scope_id=admin.tenant_id,
value="火车",
generation=1,
)
db.add(active)
db.commit()
now = datetime.now(UTC)
with pytest.raises(OrganizationMemoryConflictError, match="生效版本已变化"):
ExpenseApplicationOrganizationMemoryService(db)._new_active_generation(
entry_id=str(uuid.uuid4()),
tenant_id=admin.tenant_id,
scope_type="enterprise",
scope_id=admin.tenant_id,
value="飞机",
expires_at=now + timedelta(days=180),
reason="模拟并发创建",
request_id="unexpected-active-generation",
request_fingerprint="hmac-sha256:" + "0" * 64,
expected_active_ids=set(),
current_user=admin,
now=now,
)
db.refresh(active)
assert active.status == "active"
assert active.superseded_by_id is None

View File

@@ -0,0 +1,279 @@
from __future__ import annotations
from collections.abc import Generator
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine, select
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.pool import StaticPool
from app.api.deps import get_db
from app.db.base import Base
from app.main import create_app
from app.models.ai_memory import MemoryEntry
from app.models.organization import OrganizationUnit
def _build_client() -> tuple[TestClient, sessionmaker[Session]]:
engine = create_engine(
"sqlite+pysqlite:///:memory:",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
Base.metadata.create_all(bind=engine)
session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
with session_factory() as db:
db.add(
OrganizationUnit(
id="department-finance",
unit_code="FINANCE",
name="财务部",
unit_type="department",
)
)
db.commit()
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
with session_factory() as db:
yield db
app.dependency_overrides[get_db] = override_db
return TestClient(app), session_factory
def _headers(*, tenant_id: str = "tenant-a", admin: bool = True) -> dict[str, str]:
return {
"X-Auth-Username": "admin@example.com" if admin else "employee@example.com",
"X-Auth-Name": "Tenant Admin" if admin else "Employee",
"X-Auth-Tenant-Id": tenant_id,
"X-Auth-Is-Admin": "true" if admin else "false",
"X-Auth-Role-Codes": "manager" if admin else "user",
}
def test_organization_memory_api_requires_platform_admin() -> None:
client, _ = _build_client()
response = client.post(
"/api/v1/expense-application-memories/organization",
headers=_headers(admin=False),
json={
"scope_type": "enterprise",
"value": "火车",
"expires_in_days": 180,
"reason": "统一差旅基线",
"request_id": "create-enterprise-memory-denied",
},
)
assert response.status_code == 403
def test_organization_memory_api_lifecycle_and_tenant_isolation() -> None:
client, session_factory = _build_client()
created = client.post(
"/api/v1/expense-application-memories/organization",
headers=_headers(),
json={
"scope_type": "enterprise",
"value": "火车",
"expires_in_days": 180,
"reason": "统一差旅基线",
"request_id": "create-enterprise-memory-a",
},
)
assert created.status_code == 201, created.text
first = created.json()
assert first["scope_id"] == "tenant-a"
assert first["generation"] == 1
assert first["can_revoke"] is True
replayed = client.post(
"/api/v1/expense-application-memories/organization",
headers=_headers(),
json={
"scope_type": "enterprise",
"value": "火车",
"expires_in_days": 180,
"reason": "统一差旅基线",
"request_id": "create-enterprise-memory-a",
},
)
assert replayed.status_code == 201
assert replayed.json()["id"] == first["id"]
changed_create_payload = client.post(
"/api/v1/expense-application-memories/organization",
headers=_headers(),
json={
"scope_type": "enterprise",
"value": "飞机",
"expires_in_days": 180,
"reason": "同键不同请求体",
"request_id": "create-enterprise-memory-a",
},
)
assert changed_create_payload.status_code == 409
duplicate = client.post(
"/api/v1/expense-application-memories/organization",
headers=_headers(),
json={
"scope_type": "enterprise",
"value": "轮船",
"expires_in_days": 180,
"reason": "不应静默覆盖",
"request_id": "create-enterprise-memory-b",
},
)
assert duplicate.status_code == 409
updated = client.put(
f"/api/v1/expense-application-memories/organization/{first['id']}",
headers=_headers(),
json={
"value": "飞机",
"expires_in_days": 120,
"expected_generation": 1,
"reason": "制度版本调整",
"request_id": "update-enterprise-memory-a",
},
)
assert updated.status_code == 200, updated.text
second = updated.json()
assert second["generation"] == 2
assert second["value"] == "飞机"
changed_update_payload = client.put(
f"/api/v1/expense-application-memories/organization/{first['id']}",
headers=_headers(),
json={
"value": "轮船",
"expires_in_days": 120,
"expected_generation": 1,
"reason": "同键不同请求体",
"request_id": "update-enterprise-memory-a",
},
)
assert changed_update_payload.status_code == 409
stale_old_id = client.put(
f"/api/v1/expense-application-memories/organization/{first['id']}",
headers=_headers(),
json={
"value": "轮船",
"expected_generation": 1,
"reason": "旧地址更新",
"request_id": "update-enterprise-memory-old-id",
},
)
assert stale_old_id.status_code == 409
stale = client.put(
f"/api/v1/expense-application-memories/organization/{second['id']}",
headers=_headers(),
json={
"value": "轮船",
"expected_generation": 1,
"reason": "陈旧页面更新",
"request_id": "update-enterprise-memory-stale",
},
)
assert stale.status_code == 409
assert client.get(
"/api/v1/expense-application-memories/organization",
headers=_headers(tenant_id="tenant-b"),
).json() == {"items": []}
hidden = client.post(
f"/api/v1/expense-application-memories/organization/{second['id']}/revoke",
headers=_headers(tenant_id="tenant-b"),
json={
"expected_generation": 2,
"reason": "越权尝试",
"request_id": "revoke-other-tenant",
},
)
assert hidden.status_code == 404
revoked = client.post(
f"/api/v1/expense-application-memories/organization/{second['id']}/revoke",
headers=_headers(),
json={
"expected_generation": 2,
"reason": "改用新制度",
"request_id": "revoke-enterprise-memory-a",
},
)
assert revoked.status_code == 200, revoked.text
replayed_revoke = client.post(
f"/api/v1/expense-application-memories/organization/{second['id']}/revoke",
headers=_headers(),
json={
"expected_generation": 2,
"reason": "改用新制度",
"request_id": "revoke-enterprise-memory-a",
},
)
assert replayed_revoke.status_code == 200
assert replayed_revoke.json() == revoked.json()
changed_revoke_payload = client.post(
f"/api/v1/expense-application-memories/organization/{second['id']}/revoke",
headers=_headers(),
json={
"expected_generation": 2,
"reason": "同键不同请求体",
"request_id": "revoke-enterprise-memory-a",
},
)
assert changed_revoke_payload.status_code == 409
with session_factory() as db:
entries = list(
db.scalars(
select(MemoryEntry)
.where(MemoryEntry.tenant_id == "tenant-a")
.order_by(MemoryEntry.generation.asc())
).all()
)
assert [entry.status for entry in entries] == ["suppressed", "revoked"]
assert entries[0].superseded_by_id == entries[1].id
assert entries[0].management_request_id == "create-enterprise-memory-a"
assert entries[0].management_payload_fingerprint.startswith("hmac-sha256:")
assert entries[1].management_request_id == "update-enterprise-memory-a"
assert entries[1].revoke_request_id == "revoke-enterprise-memory-a"
assert entries[1].revoke_payload_fingerprint.startswith("hmac-sha256:")
def test_department_memory_api_requires_stable_department_id() -> None:
client, _ = _build_client()
invalid = client.post(
"/api/v1/expense-application-memories/organization",
headers=_headers(),
json={
"scope_type": "department",
"scope_id": "财务部",
"value": "火车",
"expires_in_days": 90,
"reason": "部门差旅基线",
"request_id": "create-department-invalid",
},
)
assert invalid.status_code == 422
created = client.post(
"/api/v1/expense-application-memories/organization",
headers=_headers(),
json={
"scope_type": "department",
"scope_id": "department-finance",
"value": "火车",
"expires_in_days": 90,
"reason": "部门差旅基线",
"request_id": "create-department-finance",
},
)
assert created.status_code == 201, created.text
assert created.json()["scope_id"] == "department-finance"
assert "财务部" in created.json()["scope_label"]

View File

@@ -0,0 +1,357 @@
from __future__ import annotations
from datetime import UTC, datetime
from decimal import Decimal
from unittest.mock import MagicMock
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.pool import StaticPool
from app.db.base import Base
from app.models.few_shot_sample import FewShotSample
from app.models.financial_record import ExpenseClaim
from app.schemas.ontology import OntologyParseResult
from app.schemas.user_agent import UserAgentRequest
from app.services.expense_claim_historical_evidence import (
ExpenseClaimHistoricalEvidenceService,
build_user_agent_historical_evidence_notice,
)
from app.services.expense_claim_pre_review import ExpenseClaimPreReviewMixin
from app.services.expense_claim_pre_review_decision import pre_review_public_payload
from app.services.few_shot_retrieval import FewShotRetriever
from app.services.few_shot_store import FewShotStore
from app.services.user_agent_application import UserAgentApplicationMixin
def _session() -> Session:
engine = create_engine(
"sqlite+pysqlite:///:memory:",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
Base.metadata.create_all(engine)
return sessionmaker(bind=engine)()
def _claim(*, claim_id: str = "claim-history") -> ExpenseClaim:
return ExpenseClaim(
id=claim_id,
claim_no=f"RE-{claim_id}",
employee_id="employee-history",
employee_name="张三",
department_id="department-history",
department_name="市场部",
project_code="PRJ-HISTORY",
expense_type="travel",
reason="客户现场差旅",
location="上海",
amount=Decimal("888.00"),
currency="CNY",
invoice_count=1,
occurred_at=datetime(2026, 7, 16, 9, 0, tzinfo=UTC),
status="draft",
approval_stage="待提交",
risk_flags_json=[],
)
def _finding() -> dict:
return {
"risk_id": "risk-history",
"rule_code": "TRAVEL-001",
"rule_version": "v2",
"severity": "high",
"disposition": "fix",
"resolution_status": "unresolved",
"actionability": "fixable_by_submitter",
"source": "submission_review",
"business_stage": "reimbursement",
"risk_domain": "policy",
"visibility_scope": "employee",
"item_ids": [],
"message": "住宿金额超过差旅标准。",
"remediation": {
"action": "补充说明",
"target_item_ids": [],
"required_fields": [],
"alternative_action": None,
},
}
def test_runtime_retrieval_explicitly_passes_tenant_scene_and_rule_identity(
monkeypatch,
) -> None:
with _session() as db:
retriever = MagicMock()
retriever.retrieve_for_expense_case.return_value = [
{
"sample_id": "sample-confirmed",
"label": "confirmed",
"score": 0.96,
"scene": "expense_reimbursement",
"policy_ref": "TRAVEL-001",
"rule_version": "v2",
"stale": False,
"conclusion": "该类超标案例经复核确认成立。",
}
]
monkeypatch.setattr(
ExpenseClaimHistoricalEvidenceService,
"_build_retriever",
lambda _service: retriever,
)
evidence = ExpenseClaimHistoricalEvidenceService(db).retrieve(
_claim(),
tenant_id="tenant-a",
business_stage="reimbursement",
findings=[_finding()],
)
retriever.retrieve_for_expense_case.assert_called_once_with(
tenant_id="tenant-a",
scene="expense_reimbursement",
policy_ref="TRAVEL-001",
rule_version="v2",
query="travel\n客户现场差旅\n上海\n住宿金额超过差旅标准。",
top_k=3,
)
assert evidence[0]["label_text"] == "历史已确认,仅供复核"
assert evidence[0]["advisory_only"] is True
assert evidence[0]["summary"] == "历史相似案例经人工复核确认风险成立。"
assert "sample_id" not in evidence[0]
assert "conclusion" not in evidence[0]
assert "该类超标案例经复核确认成立" not in repr(evidence[0])
def test_runtime_historical_evidence_db_recheck_never_exposes_other_tenant(
monkeypatch,
) -> None:
with _session() as db:
db.add_all(
[
_sample(
sample_id="sample-a",
tenant_id="tenant-a",
label="confirmed",
),
_sample(
sample_id="sample-b",
tenant_id="tenant-b",
label="false_positive",
),
]
)
db.commit()
store = MagicMock(spec=FewShotStore)
hits = [
{"sample_id": "sample-b", "score": 0.99},
{"sample_id": "sample-a", "score": 0.91},
]
store.search.side_effect = [hits, hits]
monkeypatch.setattr(
ExpenseClaimHistoricalEvidenceService,
"_build_retriever",
lambda _service: FewShotRetriever(store, db),
)
evidence = ExpenseClaimHistoricalEvidenceService(db).retrieve(
_claim(),
tenant_id="tenant-a",
business_stage="reimbursement",
findings=[_finding()],
)
assert len(evidence) == 1
assert evidence[0]["label"] == "confirmed"
assert "sample_id" not in evidence[0]
assert "conclusion" not in evidence[0]
assert "tenant-a 的历史结论" not in repr(evidence)
assert "tenant-b 的历史结论" not in repr(evidence)
assert all(
call.kwargs["tenant_id"] == "tenant-a"
for call in store.search.call_args_list
)
def test_retrieval_failure_returns_empty_and_hard_decision_is_unchanged(
monkeypatch,
) -> None:
with _session() as db:
monkeypatch.setattr(
ExpenseClaimHistoricalEvidenceService,
"_build_retriever",
MagicMock(side_effect=RuntimeError("qdrant unavailable")),
)
service = _PreReviewHarness(db)
with_evidence = service._refresh_claim_pre_review_flags(
_claim(claim_id="claim-with-evidence"),
is_application_claim=False,
reviewed_at=datetime(2026, 7, 16, 9, 0, tzinfo=UTC),
tenant_id="tenant-a",
)
without_evidence = service._refresh_claim_pre_review_flags(
_claim(claim_id="claim-with-evidence"),
is_application_claim=False,
reviewed_at=datetime(2026, 7, 16, 10, 0, tzinfo=UTC),
tenant_id="",
)
assert with_evidence is not None
assert without_evidence is not None
for key in (
"review_id",
"input_fingerprint",
"rule_set_fingerprint",
"review_context_fingerprint",
"decision",
"passed",
"blocking_count",
"blocking_risk_count",
"findings",
):
assert with_evidence[key] == without_evidence[key]
assert with_evidence["historical_case_evidence"] == []
assert without_evidence["historical_case_evidence"] == []
def test_evidence_changes_only_advisory_payload_not_rule_decision(monkeypatch) -> None:
with _session() as db:
retriever = MagicMock()
retriever.retrieve_for_expense_case.return_value = [
{
"sample_id": "sample-false-positive",
"label": "false_positive",
"score": 0.87,
"scene": "expense_reimbursement",
"policy_ref": "TRAVEL-001",
"rule_version": "v1",
"stale": True,
"conclusion": "历史相似案例经人工复核为误报。",
}
]
monkeypatch.setattr(
ExpenseClaimHistoricalEvidenceService,
"_build_retriever",
lambda _service: retriever,
)
service = _PreReviewHarness(db)
with_evidence = service._refresh_claim_pre_review_flags(
_claim(claim_id="same-claim"),
is_application_claim=False,
reviewed_at=datetime(2026, 7, 16, 9, 0, tzinfo=UTC),
tenant_id="tenant-a",
)
without_evidence = service._refresh_claim_pre_review_flags(
_claim(claim_id="same-claim"),
is_application_claim=False,
reviewed_at=datetime(2026, 7, 16, 10, 0, tzinfo=UTC),
tenant_id="",
)
assert with_evidence is not None
assert without_evidence is not None
for key in (
"review_id",
"decision",
"passed",
"blocking_count",
"blocking_risk_count",
"findings",
):
assert with_evidence[key] == without_evidence[key]
assert with_evidence["historical_case_evidence"][0]["label_text"] == (
"历史误报,仅供复核"
)
public_payload = pre_review_public_payload(with_evidence)
assert public_payload is not None
public_evidence = public_payload["historical_case_evidence"][0]
assert public_evidence["advisory_only"] is True
assert public_evidence["version_status"] == "stale"
assert public_evidence["summary"] == "历史相似案例经人工复核判定为误报。"
assert "sample_id" not in public_evidence
assert "conclusion" not in public_evidence
assert "历史相似案例经人工复核为误报" not in repr(public_evidence)
assert "claim_no" not in repr(public_evidence)
def test_user_agent_displays_only_fixed_historical_labels() -> None:
claim = _claim()
claim.risk_flags_json = [
{
"source": "ai_pre_review",
"historical_case_evidence": [
{
"label": "confirmed",
"sample_id": "sensitive-sample-id",
"conclusion": "客户甲与员工乙的人工复核原文",
"claim_no": "RE-SENSITIVE",
},
{"label": "false_positive"},
],
}
]
service = UserAgentApplicationMixin()
notice = build_user_agent_historical_evidence_notice(claim)
answer = service._build_expense_application_answer(
UserAgentRequest(
run_id="history-labels",
user_id="employee@example.com",
message="提交申请",
ontology=OntologyParseResult(run_id="history-labels"),
),
facts={
"application_no": "AP-20260716-001",
"manager_name": "直属领导",
"historical_case_evidence_notice": notice,
},
step="submitted",
)
assert "历史已确认,仅供复核" in answer
assert "历史误报,仅供复核" in answer
assert "sensitive-sample-id" not in answer
assert "客户甲与员工乙" not in answer
assert "RE-SENSITIVE" not in answer
def _sample(*, sample_id: str, tenant_id: str, label: str) -> FewShotSample:
return FewShotSample(
id=sample_id,
tenant_id=tenant_id,
sample_key=f"key-{sample_id}",
scene="expense_reimbursement",
policy_ref="TRAVEL-001",
rule_version="v2",
domain="expense",
risk_type="travel_limit",
risk_level="high",
label=label,
case_text="住宿金额超过差旅标准",
conclusion_text=f"{tenant_id} 的历史结论",
payload_json={},
status="active",
)
class _PreReviewHarness(ExpenseClaimPreReviewMixin):
def __init__(self, db: Session) -> None:
self.db = db
@staticmethod
def _run_ai_submission_review(_claim: ExpenseClaim) -> dict:
return {
"risk_flags": [
{
"source": "submission_review",
"severity": "high",
"actionability": "fixable_by_submitter",
"rule_code": "TRAVEL-001",
"rule_version": "v2",
"message": "住宿金额超过差旅标准。",
}
],
"rule_set_fingerprint": "rules-v2",
}

View File

@@ -1,9 +1,7 @@
from __future__ import annotations from __future__ import annotations
from collections.abc import Generator
from datetime import datetime from datetime import datetime
from decimal import Decimal from decimal import Decimal
from types import SimpleNamespace
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
import pytest import pytest
@@ -55,6 +53,7 @@ def _observation(db: Session, key: str = "risk:c1:dup") -> RiskObservation:
) )
db.flush() db.flush()
obs = RiskObservation( obs = RiskObservation(
tenant_id="tenant-a",
observation_key=key, observation_key=key,
subject_type="expense_claim", subject_type="expense_claim",
subject_key="claim:c1", subject_key="claim:c1",
@@ -91,6 +90,7 @@ def test_ingest_confirmed_persists_sample_and_calls_store() -> None:
) )
assert sample is not None assert sample is not None
assert sample.label == "confirmed" assert sample.label == "confirmed"
assert sample.tenant_id == "tenant-a"
assert sample.sample_key == f"obs:{obs.id}" assert sample.sample_key == f"obs:{obs.id}"
assert "重复发票" in sample.case_text assert "重复发票" in sample.case_text
assert "确认重复发票" in sample.conclusion_text assert "确认重复发票" in sample.conclusion_text
@@ -98,13 +98,38 @@ def test_ingest_confirmed_persists_sample_and_calls_store() -> None:
fake_store.upsert.assert_called_once() fake_store.upsert.assert_called_once()
def test_ingest_extracts_business_scene_and_rule_identity() -> None:
with _build_session() as db:
obs = _observation(db, key="risk:c1:identity")
obs.feedback_status = "confirmed"
obs.control_stage = "reimbursement"
obs.policy_refs_json = [{"rule_code": "TRAVEL-001"}]
obs.decision_trace_json = {"rule_version": "v2.3"}
service = FewShotIngestionService(db)
with patch.object(
service,
"_store",
return_value=MagicMock(upsert=MagicMock(return_value=None)),
):
sample = service.ingest_observation_feedback(
obs,
MagicMock(feedback_type="confirm", comment="确认", actor="audit"),
)
assert sample is not None
assert sample.scene == "expense_reimbursement"
assert sample.policy_ref == "TRAVEL-001"
assert sample.rule_version == "v2.3"
def test_ingest_false_positive_also_persisted() -> None: def test_ingest_false_positive_also_persisted() -> None:
with _build_session() as db: with _build_session() as db:
obs = _observation(db, key="risk:c2:fp") obs = _observation(db, key="risk:c2:fp")
obs.feedback_status = "false_positive" obs.feedback_status = "false_positive"
db.commit() db.commit()
service = FewShotIngestionService(db) service = FewShotIngestionService(db)
with patch.object(service, "_store", return_value=MagicMock(upsert=MagicMock(return_value=None))): with patch.object(
service, "_store", return_value=MagicMock(upsert=MagicMock(return_value=None))
):
sample = service.ingest_observation_feedback( sample = service.ingest_observation_feedback(
obs, obs,
MagicMock(feedback_type="false_positive", comment="", actor="audit"), MagicMock(feedback_type="false_positive", comment="", actor="audit"),
@@ -164,6 +189,7 @@ def test_create_feedback_hook_triggers_ingestion() -> None:
service.create_feedback( service.create_feedback(
obs.observation_key, obs.observation_key,
RiskObservationFeedbackCreate(feedback_type="confirm", actor="audit"), RiskObservationFeedbackCreate(feedback_type="confirm", actor="audit"),
tenant_id="tenant-a",
) )
assert len(ingest_calls) == 1 assert len(ingest_calls) == 1
assert ingest_calls[0][1] == "confirm" assert ingest_calls[0][1] == "confirm"
@@ -178,7 +204,10 @@ def test_create_feedback_hook_skipped_for_comment_feedback() -> None:
) as mock_ingest: ) as mock_ingest:
service.create_feedback( service.create_feedback(
obs.observation_key, obs.observation_key,
RiskObservationFeedbackCreate(feedback_type="comment", action="note", actor="audit"), RiskObservationFeedbackCreate(
feedback_type="comment", action="note", actor="audit"
),
tenant_id="tenant-a",
) )
mock_ingest.assert_not_called() mock_ingest.assert_not_called()
@@ -195,6 +224,7 @@ def test_create_feedback_hook_swallows_ingestion_failure() -> None:
feedback = service.create_feedback( feedback = service.create_feedback(
obs.observation_key, obs.observation_key,
RiskObservationFeedbackCreate(feedback_type="confirm", actor="audit"), RiskObservationFeedbackCreate(feedback_type="confirm", actor="audit"),
tenant_id="tenant-a",
) )
assert feedback.feedback_type == "confirm" assert feedback.feedback_type == "confirm"
@@ -210,5 +240,6 @@ def test_create_feedback_hook_respects_feature_flag(monkeypatch: pytest.MonkeyPa
service.create_feedback( service.create_feedback(
obs.observation_key, obs.observation_key,
RiskObservationFeedbackCreate(feedback_type="confirm", actor="audit"), RiskObservationFeedbackCreate(feedback_type="confirm", actor="audit"),
tenant_id="tenant-a",
) )
mock_ingest.assert_not_called() mock_ingest.assert_not_called()

View File

@@ -1,12 +1,11 @@
from __future__ import annotations from __future__ import annotations
import json import json
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock
import pytest
from app.services.few_shot_retrieval import FewShotRetriever from app.services.few_shot_retrieval import FewShotRetriever
from app.services.few_shot_store import FewShotStore from app.services.few_shot_store import FewShotStore
from app.services.risk_rule_generation import RiskRuleGenerationService
from app.services.risk_rule_generation_prompt import build_risk_rule_compiler_messages from app.services.risk_rule_generation_prompt import build_risk_rule_compiler_messages
@@ -36,7 +35,9 @@ def test_retrieve_returns_injection_blocks_with_token_budget() -> None:
] ]
retriever = FewShotRetriever(store) retriever = FewShotRetriever(store)
blocks = retriever.retrieve_for_risk_rule_generation( blocks = retriever.retrieve_for_risk_rule_generation(
domain="expense", natural_language="同一发票重复报销" tenant_id="tenant-a",
domain="expense",
natural_language="同一发票重复报销",
) )
assert len(blocks) == 2 assert len(blocks) == 2
assert blocks[0]["score"] == 0.9 assert blocks[0]["score"] == 0.9
@@ -51,7 +52,13 @@ def test_retrieve_returns_injection_blocks_with_token_budget() -> None:
def test_retrieve_empty_case_text_returns_empty() -> None: def test_retrieve_empty_case_text_returns_empty() -> None:
store = MagicMock(spec=FewShotStore) store = MagicMock(spec=FewShotStore)
retriever = FewShotRetriever(store) retriever = FewShotRetriever(store)
assert retriever.retrieve_for_risk_rule_generation(natural_language="") == [] assert (
retriever.retrieve_for_risk_rule_generation(
tenant_id="tenant-a",
natural_language="",
)
== []
)
store.search.assert_not_called() store.search.assert_not_called()
@@ -62,7 +69,10 @@ def test_retrieve_truncates_overlong_conclusion() -> None:
_hit(0.9, "confirmed", long_text), _hit(0.9, "confirmed", long_text),
] ]
retriever = FewShotRetriever(store) retriever = FewShotRetriever(store)
blocks = retriever.retrieve_for_risk_rule_generation(natural_language="x") blocks = retriever.retrieve_for_risk_rule_generation(
tenant_id="tenant-a",
natural_language="x",
)
assert len(blocks) == 1 assert len(blocks) == 1
# 超长结论应被截断到单条上限 # 超长结论应被截断到单条上限
from app.services.few_shot_retrieval import SINGLE_SAMPLE_MAX_CHARS from app.services.few_shot_retrieval import SINGLE_SAMPLE_MAX_CHARS
@@ -70,6 +80,45 @@ def test_retrieve_truncates_overlong_conclusion() -> None:
assert len(blocks[0]["conclusion"]) <= SINGLE_SAMPLE_MAX_CHARS assert len(blocks[0]["conclusion"]) <= SINGLE_SAMPLE_MAX_CHARS
def test_rule_generation_forwards_authenticated_tenant_to_retriever(monkeypatch) -> None:
fake_retriever = MagicMock()
fake_retriever.retrieve_for_risk_rule_generation.return_value = []
monkeypatch.setenv("FEW_SHOT_INJECTION_ENABLED", "true")
monkeypatch.setattr(
FewShotRetriever,
"from_session",
classmethod(lambda _cls, _session: fake_retriever),
)
service = RiskRuleGenerationService(MagicMock())
service._retrieve_few_shot_samples(
tenant_id="tenant-smart-learning",
domain="expense",
natural_language="重复发票风险规则",
)
fake_retriever.retrieve_for_risk_rule_generation.assert_called_once_with(
tenant_id="tenant-smart-learning",
domain="expense",
natural_language="重复发票风险规则",
)
def test_rule_generation_without_tenant_disables_historical_injection(monkeypatch) -> None:
from_session = MagicMock()
monkeypatch.setenv("FEW_SHOT_INJECTION_ENABLED", "true")
monkeypatch.setattr(FewShotRetriever, "from_session", from_session)
result = RiskRuleGenerationService(MagicMock())._retrieve_few_shot_samples(
tenant_id=None,
domain="expense",
natural_language="重复发票风险规则",
)
assert result == []
from_session.assert_not_called()
def test_build_prompt_merges_few_shot_into_examples() -> None: def test_build_prompt_merges_few_shot_into_examples() -> None:
samples = [ samples = [
{ {
@@ -89,7 +138,14 @@ def test_build_prompt_merges_few_shot_into_examples() -> None:
expense_category=None, expense_category=None,
expense_category_label="", expense_category_label="",
natural_language="重复发票规则", natural_language="重复发票规则",
available_fields=[{"key": "attachment.invoice_no", "label": "发票号", "type": "string", "source": "attachment"}], available_fields=[
{
"key": "attachment.invoice_no",
"label": "发票号",
"type": "string",
"source": "attachment",
}
],
few_shot_samples=samples, few_shot_samples=samples,
) )
assert len(messages) == 2 assert len(messages) == 2

View File

@@ -0,0 +1,183 @@
from __future__ import annotations
from datetime import UTC, datetime, timedelta
from decimal import Decimal
import pytest
from sqlalchemy import create_engine
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.pool import StaticPool
from app.api.deps import _authenticate_bearer_user
from app.db.base import Base
from app.models.ai_memory import MemoryEntry
from app.models.employee import Employee
from app.services.auth import AuthService
from app.services.auth_sessions import AuthSessionService
from app.services.employee import EmployeeService
def _build_session() -> Session:
engine = create_engine(
"sqlite+pysqlite:///:memory:",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
Base.metadata.create_all(bind=engine)
return sessionmaker(bind=engine, autoflush=False, autocommit=False)()
def _memory_entry(
*,
scope_type: str,
scope_id: str,
tenant_id: str = "tenant-a",
origin_type: str = "learned",
with_management_audit: bool = False,
) -> MemoryEntry:
now = datetime.now(UTC)
return MemoryEntry(
tenant_id=tenant_id,
scope_type=scope_type,
scope_id=scope_id,
origin_type=origin_type,
managed_by="admin@example.com" if with_management_audit else None,
managed_at=now if with_management_audit else None,
management_reason="统一差旅交通基线" if with_management_audit else None,
policy_version="travel-policy-v1" if with_management_audit else None,
value_json={"value": "火车"},
value_fingerprint=f"{scope_type}:{scope_id}",
status="active",
evidence_count=0,
approved_evidence_count=0,
confidence=Decimal("1.0000"),
last_evidence_at=now,
candidate_expires_at=now + timedelta(days=90),
activated_at=now,
active_expires_at=now + timedelta(days=180),
created_at=now,
updated_at=now,
)
def test_memory_entry_accepts_learned_user_and_audited_organization_scopes() -> None:
with _build_session() as db:
learned_entry = _memory_entry(
scope_type="user",
scope_id="employee-a",
)
learned_entry.policy_version = "expense-application-memory.v1"
db.add(learned_entry)
db.add(
_memory_entry(
scope_type="department",
scope_id="department-a",
origin_type="admin_managed",
with_management_audit=True,
)
)
db.add(
_memory_entry(
scope_type="enterprise",
scope_id="tenant-a",
origin_type="admin_managed",
with_management_audit=True,
)
)
db.commit()
assert db.query(MemoryEntry).count() == 3
@pytest.mark.parametrize(
("entry", "expected_constraint"),
[
(
_memory_entry(
scope_type="department",
scope_id="department-a",
origin_type="learned",
),
"ck_memory_entries_scope_origin",
),
(
_memory_entry(
scope_type="enterprise",
scope_id="another-tenant",
origin_type="admin_managed",
with_management_audit=True,
),
"ck_memory_entries_enterprise_scope",
),
(
_memory_entry(
scope_type="department",
scope_id="department-a",
origin_type="admin_managed",
),
"ck_memory_entries_management_audit",
),
],
)
def test_memory_entry_rejects_invalid_scope_origin_and_management_audit(
entry: MemoryEntry,
expected_constraint: str,
) -> None:
with _build_session() as db:
db.add(entry)
with pytest.raises(IntegrityError, match=expected_constraint):
db.commit()
def test_admin_managed_memory_rejects_partially_missing_audit_fields() -> None:
with _build_session() as db:
entry = _memory_entry(
scope_type="department",
scope_id="department-a",
origin_type="admin_managed",
with_management_audit=True,
)
entry.management_reason = " "
db.add(entry)
with pytest.raises(IntegrityError, match="ck_memory_entries_management_audit"):
db.commit()
def test_authenticated_session_restores_tenant_and_stable_department_id() -> None:
with _build_session() as db:
employee_snapshot = EmployeeService(db).list_employees()[0]
employee = db.get(Employee, employee_snapshot.id)
assert employee is not None
authenticated_at_login = AuthService(db)._build_employee_user(employee)
authenticated_at_login.tenant_id = "tenant-session-a"
access_token, auth_session = AuthSessionService(db).issue(
authenticated_at_login,
metric_session_id="metric-session-a",
)
db.commit()
restored_user = AuthService(db).get_session_user(auth_session)
current_user = _authenticate_bearer_user(db, f"Bearer {access_token}")
assert restored_user is not None
assert restored_user.tenant_id == "tenant-session-a"
assert restored_user.department_id == employee.organization_unit_id
assert current_user.tenant_id == "tenant-session-a"
assert current_user.department_id == employee.organization_unit_id
assert current_user.department_name == employee.organization_unit.name
def test_platform_admin_has_no_department_scope() -> None:
with _build_session() as db:
record = type(
"AdminRecord",
(),
{"account": "admin", "email": "admin@example.com"},
)()
admin_user = AuthService(db)._build_admin_user(record)
assert admin_user.department_id is None

View File

@@ -7,6 +7,7 @@ from sqlalchemy import Column, Integer, MetaData, Table, create_engine, text
from sqlalchemy.engine import Engine from sqlalchemy.engine import Engine
from app.db.migration_preflight import ( from app.db.migration_preflight import (
LEGACY_ADOPTABLE_HISTORICAL_CASE_TABLES,
MIGRATION_OWNED_TABLES_BY_REVISION, MIGRATION_OWNED_TABLES_BY_REVISION,
MigrationPreflightError, MigrationPreflightError,
validate_migration_state, validate_migration_state,
@@ -46,9 +47,29 @@ def test_unversioned_database_without_migration_owned_tables_is_safe(engine: Eng
assert state.owned_tables == frozenset() assert state.owned_tables == frozenset()
def test_unversioned_database_can_adopt_legacy_historical_case_tables(engine: Engine) -> None:
_create_tables(engine, LEGACY_ADOPTABLE_HISTORICAL_CASE_TABLES)
state = validate_migration_state(engine)
assert state.revision is None
assert state.owned_tables == LEGACY_ADOPTABLE_HISTORICAL_CASE_TABLES
def test_revision_0007_accepts_partial_legacy_historical_case_tables(engine: Engine) -> None:
expected = MIGRATION_OWNED_TABLES_BY_REVISION["20260716_0007"]
adopted = frozenset({"risk_observations", "risk_observation_feedback"})
_create_tables(engine, expected | adopted)
_create_version_table(engine, "20260716_0007")
state = validate_migration_state(engine)
assert state.owned_tables == expected | adopted
@pytest.mark.parametrize( @pytest.mark.parametrize(
"owned_table", "owned_table",
sorted(MIGRATION_OWNED_TABLES_BY_REVISION["20260716_0006"]), sorted(MIGRATION_OWNED_TABLES_BY_REVISION["20260716_0007"]),
) )
def test_unversioned_database_with_any_migration_owned_table_is_rejected( def test_unversioned_database_with_any_migration_owned_table_is_rejected(
engine: Engine, engine: Engine,
@@ -102,8 +123,19 @@ def test_known_revision_requires_and_accepts_its_exact_owned_table_set(
), ),
( (
"20260716_0006", "20260716_0006",
MIGRATION_OWNED_TABLES_BY_REVISION["20260716_0006"] MIGRATION_OWNED_TABLES_BY_REVISION["20260716_0006"] - {"attachment_association_jobs"},
- {"attachment_association_jobs"}, ),
(
"20260716_0007",
MIGRATION_OWNED_TABLES_BY_REVISION["20260716_0007"] - {"memory_entries"},
),
(
"20260716_0008",
MIGRATION_OWNED_TABLES_BY_REVISION["20260716_0008"] - {"few_shot_samples"},
),
(
"20260716_0009",
MIGRATION_OWNED_TABLES_BY_REVISION["20260716_0009"] - {"memory_entries"},
), ),
], ],
) )

View File

@@ -3,6 +3,7 @@ from __future__ import annotations
from collections.abc import Generator from collections.abc import Generator
from datetime import UTC, datetime from datetime import UTC, datetime
from decimal import Decimal from decimal import Decimal
from types import SimpleNamespace
import pytest import pytest
from auth_helpers import install_legacy_header_auth_override from auth_helpers import install_legacy_header_auth_override
@@ -17,9 +18,11 @@ from app.api.deps import get_db
from app.api.v1.endpoints.risk_observations import router as risk_observations_router from app.api.v1.endpoints.risk_observations import router as risk_observations_router
from app.db.base import Base from app.db.base import Base
from app.models.employee import Employee from app.models.employee import Employee
from app.models.expense_case import ExpenseCase, ExpenseCaseLink
from app.models.financial_record import ExpenseClaim from app.models.financial_record import ExpenseClaim
from app.models.risk_observation import RiskObservation from app.models.risk_observation import RiskObservation
from app.schemas.risk_observation import RiskObservationFeedbackCreate from app.schemas.risk_observation import RiskObservationFeedbackCreate
from app.services.hermes_risk_scanner import HermesRiskScannerService
from app.services.risk_observations import RiskObservationService from app.services.risk_observations import RiskObservationService
@@ -179,6 +182,7 @@ def test_risk_observation_endpoints_return_list_detail_dashboard_and_feedback()
assert "top_departments" in dashboard_response.json() assert "top_departments" in dashboard_response.json()
assert feedback_response.status_code == 200 assert feedback_response.status_code == 200
assert feedback_response.json()["feedback_type"] == "false_positive" assert feedback_response.json()["feedback_type"] == "false_positive"
assert feedback_response.json()["actor"] == "Test Admin"
updated_detail_response = client.get("/api/v1/risk-observations/risk:c1:duplicate_invoice") updated_detail_response = client.get("/api/v1/risk-observations/risk:c1:duplicate_invoice")
assert updated_detail_response.status_code == 200 assert updated_detail_response.status_code == 200
@@ -192,6 +196,237 @@ def test_risk_observation_endpoints_return_list_detail_dashboard_and_feedback()
assert observation.feedback_status == "false_positive" assert observation.feedback_status == "false_positive"
def test_risk_observation_endpoints_enforce_tenant_scope_and_authenticated_actor() -> None:
client, session_factory = _build_client()
with session_factory() as db:
service = RiskObservationService(db)
tenant_a = service.upsert_observation(
{
**_observation_payload("risk:tenant-a:duplicate_invoice"),
"claim_id": "shared-claim",
},
tenant_id="tenant-a",
execution_log_id="shared-execution-log",
)
tenant_b = service.upsert_observation(
{
**_observation_payload("risk:tenant-b:duplicate_invoice"),
"claim_id": "shared-claim",
},
tenant_id="tenant-b",
execution_log_id="shared-execution-log",
)
tenant_a_id = tenant_a.id
tenant_b_id = tenant_b.id
db.commit()
tenant_a_headers = {
"X-Auth-Username": "auditor-a",
"X-Auth-Name": "Tenant A Auditor",
"X-Auth-Tenant-Id": "tenant-a",
}
tenant_b_headers = {
"X-Auth-Username": "auditor-b",
"X-Auth-Name": "Tenant B Auditor",
"X-Auth-Tenant-Id": "tenant-b",
}
list_response = client.get("/api/v1/risk-observations", headers=tenant_a_headers)
detail_response = client.get(
f"/api/v1/risk-observations/{tenant_a_id}",
headers=tenant_a_headers,
)
foreign_detail_response = client.get(
f"/api/v1/risk-observations/{tenant_b_id}",
headers=tenant_a_headers,
)
claim_response = client.get(
"/api/v1/risk-observations/claim/shared-claim",
headers=tenant_a_headers,
)
execution_log_response = client.get(
"/api/v1/risk-observations/execution-log/shared-execution-log",
headers=tenant_a_headers,
)
dashboard_response = client.get(
"/api/v1/risk-observations/dashboard",
headers=tenant_a_headers,
)
foreign_feedback_response = client.post(
f"/api/v1/risk-observations/{tenant_b_id}/feedback",
headers=tenant_a_headers,
json={"feedback_type": "confirm", "actor": "伪造管理员"},
)
own_feedback_response = client.post(
f"/api/v1/risk-observations/{tenant_a_id}/feedback",
headers=tenant_a_headers,
json={"feedback_type": "confirm", "actor": "伪造管理员"},
)
assert list_response.status_code == 200
assert list_response.json()["total"] == 1
assert list_response.json()["items"][0]["tenant_id"] == "tenant-a"
assert detail_response.status_code == 200
assert detail_response.json()["tenant_id"] == "tenant-a"
assert foreign_detail_response.status_code == 404
assert claim_response.status_code == 200
assert [item["tenant_id"] for item in claim_response.json()] == ["tenant-a"]
assert execution_log_response.status_code == 200
assert [item["tenant_id"] for item in execution_log_response.json()] == ["tenant-a"]
assert dashboard_response.status_code == 200
assert dashboard_response.json()["total_observations"] == 1
assert foreign_feedback_response.status_code == 404
assert own_feedback_response.status_code == 200
assert own_feedback_response.json()["actor"] == "Tenant A Auditor"
tenant_b_detail = client.get(
f"/api/v1/risk-observations/{tenant_b_id}",
headers=tenant_b_headers,
)
assert tenant_b_detail.status_code == 200
assert tenant_b_detail.json()["status"] == "pending_review"
assert tenant_b_detail.json()["feedback_items"] == []
def test_risk_observation_service_scopes_history_and_same_key_upserts_by_tenant() -> None:
with _build_session() as db:
service = RiskObservationService(db)
tenant_a = service.upsert_observation(
_observation_payload("risk:shared:duplicate_invoice"),
tenant_id="tenant-a",
)
tenant_b = service.upsert_observation(
_observation_payload("risk:shared:duplicate_invoice"),
tenant_id="tenant-b",
)
service.create_feedback(
tenant_a.id,
RiskObservationFeedbackCreate(feedback_type="confirm", actor="untrusted"),
tenant_id="tenant-a",
actor="trusted-auditor",
)
tenant_a_items, tenant_a_total = service.list_observations(tenant_id="tenant-a")
tenant_b_items, tenant_b_total = service.list_observations(tenant_id="tenant-b")
tenant_a_history = service.build_history_stats(
tenant_id="tenant-a",
risk_signals={"duplicate_invoice"},
)
tenant_b_history = service.build_history_stats(
tenant_id="tenant-b",
risk_signals={"duplicate_invoice"},
)
assert tenant_a.id != tenant_b.id
assert tenant_a_total == tenant_b_total == 1
assert [item.tenant_id for item in tenant_a_items] == ["tenant-a"]
assert [item.tenant_id for item in tenant_b_items] == ["tenant-b"]
assert tenant_a_history[0].confirmed_count == 1
assert tenant_b_history[0].confirmed_count == 0
assert tenant_a.feedback_items[0].actor == "trusted-auditor"
assert service.get_observation(tenant_b.id, tenant_id="tenant-a") is None
def test_risk_observation_rejects_explicit_tenant_mismatching_claim_link() -> None:
with _build_session() as db:
claim = _claim_orm("claim-tenant-boundary", "BX-TENANT-BOUNDARY")
expense_case = ExpenseCase(
id="case-tenant-boundary",
tenant_id="tenant-a",
case_no="CASE-TENANT-BOUNDARY",
scene_code="reimbursement",
title="租户边界测试",
current_stage="claiming",
status="active",
)
link = ExpenseCaseLink(
id="link-tenant-boundary",
tenant_id="tenant-a",
expense_case_id=expense_case.id,
resource_type="expense_claim",
resource_id=claim.id,
relation_type="claim",
)
db.add_all([claim, expense_case, link])
db.flush()
with pytest.raises(PermissionError, match="tenant does not match"):
RiskObservationService(db).upsert_observation(
{
**_observation_payload("risk:tenant-boundary"),
"claim_id": claim.id,
},
tenant_id="tenant-b",
)
assert db.query(RiskObservation).filter_by(
observation_key="risk:tenant-boundary"
).one_or_none() is None
def test_hermes_global_scan_builds_graphs_inside_each_tenant(
monkeypatch: pytest.MonkeyPatch,
) -> None:
with _build_session() as db:
claims = [
_claim_orm("claim-tenant-a", "BX-TENANT-A"),
_claim_orm("claim-tenant-b", "BX-TENANT-B"),
]
cases = [
ExpenseCase(
id=f"case-tenant-{suffix}",
tenant_id=f"tenant-{suffix}",
case_no=f"CASE-TENANT-{suffix.upper()}",
scene_code="reimbursement",
title="租户隔离图扫描",
current_stage="claiming",
status="active",
)
for suffix in ("a", "b")
]
links = [
ExpenseCaseLink(
id=f"link-tenant-{suffix}",
tenant_id=f"tenant-{suffix}",
expense_case_id=cases[index].id,
resource_type="expense_claim",
resource_id=claims[index].id,
relation_type="claim",
)
for index, suffix in enumerate(("a", "b"))
]
db.add_all([*claims, *cases, *links])
db.flush()
evaluated_claim_sets: list[set[str]] = []
history_tenants: list[str] = []
def fake_evaluate(context):
evaluated_claim_sets.append(set(context.target_claim_ids))
return SimpleNamespace(observations=[], nodes=[], edges=[])
def fake_history(_self, *, tenant_id=None, **_kwargs):
history_tenants.append(str(tenant_id))
return []
scanner = HermesRiskScannerService(db)
monkeypatch.setattr(scanner, "_fetch_unscanned_claims", lambda: claims)
monkeypatch.setattr(
"app.services.hermes_risk_scanner.evaluate_financial_risk_graph",
fake_evaluate,
)
monkeypatch.setattr(RiskObservationService, "build_history_stats", fake_history)
summary = scanner.scan_global_risks()
assert evaluated_claim_sets == [
{"claim-tenant-a"},
{"claim-tenant-b"},
]
assert history_tenants == ["tenant-a", "tenant-b"]
assert summary["scanned_claim_count"] == 2
def test_risk_observation_feedback_pool_fields_and_replay_set_contract() -> None: def test_risk_observation_feedback_pool_fields_and_replay_set_contract() -> None:
with _build_session() as db: with _build_session() as db:
service = RiskObservationService(db) service = RiskObservationService(db)

View File

@@ -119,8 +119,16 @@ def test_regenerate_risk_rule_endpoint_returns_updated_detail(tmp_path, monkeypa
client, session_factory = build_client() client, session_factory = build_client()
asset_id = _create_rule(session_factory, tmp_path) asset_id = _create_rule(session_factory, tmp_path)
def fake_regenerate(self, target_asset_id, body, *, actor, request_id=None): def fake_regenerate(
del body, request_id self,
target_asset_id,
body,
*,
tenant_id=None,
actor,
request_id=None,
):
del body, request_id, tenant_id
asset = self.db.get(AgentAsset, target_asset_id) asset = self.db.get(AgentAsset, target_asset_id)
assert asset is not None assert asset is not None
config = dict(asset.config_json or {}) config = dict(asset.config_json or {})

View File

@@ -26,6 +26,9 @@ def test_create_legacy_schema_never_creates_migration_owned_tables() -> None:
"expense_cases", "expense_cases",
"memory_entries", "memory_entries",
"memory_evidence_links", "memory_evidence_links",
"risk_observations",
"risk_observation_feedback",
"few_shot_samples",
"workflow_outcomes", "workflow_outcomes",
} }
) )

View File

@@ -0,0 +1,231 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from sqlalchemy import create_engine, select
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.pool import StaticPool
from app.db.base import Base
from app.models.few_shot_sample import FewShotSample
from app.models.risk_observation import RiskObservation
from app.services.few_shot_retrieval import FewShotRetriever
from app.services.few_shot_store import FewShotStore, stable_vector_id
def _session() -> Session:
engine = create_engine(
"sqlite+pysqlite:///:memory:",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
Base.metadata.create_all(engine)
return sessionmaker(bind=engine)()
def _observation(*, tenant_id: str, key: str) -> RiskObservation:
return RiskObservation(
tenant_id=tenant_id,
observation_key=key,
subject_type="expense_claim",
subject_key="claim:1",
risk_type="duplicate_invoice",
risk_signal="duplicate_invoice",
risk_level="high",
)
def _sample(
*,
sample_id: str,
tenant_id: str,
sample_key: str,
version: str,
) -> FewShotSample:
return FewShotSample(
id=sample_id,
tenant_id=tenant_id,
sample_key=sample_key,
scene="expense_reimbursement",
policy_ref="TRAVEL-001",
rule_version=version,
label="confirmed",
case_text="同一发票重复报销",
conclusion_text=f"历史结论 {version}",
payload_json={"risk_signal": "duplicate_invoice"},
status="active",
)
def test_keys_are_unique_inside_tenant_but_reusable_across_tenants() -> None:
with _session() as db:
db.add_all(
[
_observation(tenant_id="tenant-a", key="same-key"),
_observation(tenant_id="tenant-b", key="same-key"),
_sample(
sample_id="sample-a",
tenant_id="tenant-a",
sample_key="same-sample",
version="v1",
),
_sample(
sample_id="sample-b",
tenant_id="tenant-b",
sample_key="same-sample",
version="v1",
),
]
)
db.commit()
assert len(db.scalars(select(RiskObservation)).all()) == 2
assert len(db.scalars(select(FewShotSample)).all()) == 2
def test_vector_id_is_stable_and_tenant_scoped() -> None:
first = stable_vector_id(tenant_id="tenant-a", sample_id="sample-1")
repeated = stable_vector_id(tenant_id="tenant-a", sample_id="sample-1")
other_tenant = stable_vector_id(tenant_id="tenant-b", sample_id="sample-1")
assert first == repeated
assert first != other_tenant
def test_store_upsert_replaces_legacy_vector_and_writes_tenant_payload() -> None:
provider = MagicMock()
provider.embed.return_value = [[0.1, 0.2]]
store = FewShotStore(provider)
client = MagicMock()
store._client = client
sample = SimpleNamespace(
id="sample-1",
tenant_id="tenant-a",
sample_key="key",
scene="expense_reimbursement",
policy_ref="TRAVEL-001",
rule_version="v2",
label="false_positive",
domain="expense",
risk_type="duplicate_invoice",
risk_level="high",
status="active",
case_text="案例",
conclusion_text="改判为误报",
payload_json={},
vector_id="legacy-random-vector-id",
)
with patch.object(store, "_ensure_collection", return_value=True):
vector_id = store.upsert(sample)
assert vector_id == stable_vector_id(tenant_id="tenant-a", sample_id="sample-1")
client.delete.assert_called_once()
points = client.upsert.call_args.kwargs["points"]
point = points[0]
assert point["id"] == vector_id
assert point["payload"]["tenant_id"] == "tenant-a"
assert point["payload"]["rule_version"] == "v2"
assert point["payload"]["label"] == "false_positive"
assert points[1]["id"] == "legacy-random-vector-id"
assert points[1]["payload"]["label"] == "false_positive"
def test_existing_collection_receives_required_payload_indexes() -> None:
provider = MagicMock()
store = FewShotStore(provider)
client = MagicMock()
client.get_collection.return_value = SimpleNamespace()
store._client = client
assert store._ensure_collection() is True
fields = {call.kwargs["field_name"] for call in client.create_payload_index.call_args_list}
assert {"tenant_id", "scene", "policy_ref", "rule_version", "status"} <= fields
client.create_collection.assert_not_called()
def test_search_always_filters_tenant_and_supports_rule_identity() -> None:
provider = MagicMock()
provider.embed.return_value = [[0.1, 0.2]]
store = FewShotStore(provider)
client = MagicMock()
client.query_points.return_value = SimpleNamespace(points=[])
store._client = client
with patch.object(store, "_ensure_collection", return_value=True):
assert (
store.search(
"重复发票",
tenant_id="tenant-a",
scene="expense_reimbursement",
policy_ref="TRAVEL-001",
rule_version="v2",
)
== []
)
filter_payload = client.query_points.call_args.kwargs["query_filter"].model_dump()
must = filter_payload["must"]
assert any(item["key"] == "tenant_id" and item["match"]["value"] == "tenant-a" for item in must)
assert any(item["key"] == "scene" for item in must)
assert any(item["key"] == "policy_ref" for item in must)
assert any(item["key"] == "rule_version" for item in must)
def test_expense_case_retrieval_db_rechecks_tenant_status_and_marks_old_version() -> None:
with _session() as db:
current = _sample(
sample_id="current",
tenant_id="tenant-a",
sample_key="current-key",
version="v2",
)
stale = _sample(
sample_id="stale",
tenant_id="tenant-a",
sample_key="stale-key",
version="v1",
)
other_tenant = _sample(
sample_id="other",
tenant_id="tenant-b",
sample_key="other-key",
version="v2",
)
db.add_all([current, stale, other_tenant])
db.commit()
store = MagicMock(spec=FewShotStore)
store.search.side_effect = [
[{"sample_id": "current", "score": 0.95}],
[
{"sample_id": "current", "score": 0.95},
{"sample_id": "stale", "score": 0.8},
{"sample_id": "other", "score": 0.99},
],
]
retriever = FewShotRetriever(store, db)
evidence = retriever.retrieve_for_expense_case(
tenant_id="tenant-a",
scene="expense_reimbursement",
policy_ref="TRAVEL-001",
rule_version="v2",
query="重复发票",
top_k=3,
)
assert [item["sample_id"] for item in evidence] == ["current", "stale"]
assert evidence[0]["version_status"] == "matched"
assert evidence[0]["advisory_only"] is True
assert evidence[1]["version_status"] == "stale"
assert evidence[1]["stale"] is True
def test_qdrant_unavailable_fails_closed() -> None:
store = FewShotStore(MagicMock())
with patch.object(store, "_ensure_collection", return_value=False):
assert store.search("案例", tenant_id="tenant-a") == []
assert store.upsert(SimpleNamespace(tenant_id="tenant-a", id="sample-1")) is None

View File

@@ -159,6 +159,42 @@ def test_untrusted_application_path_does_not_write_learning_ledger() -> None:
assert list(db.scalars(select(WorkflowOutcome)).all()) == [] assert list(db.scalars(select(WorkflowOutcome)).all()) == []
def test_non_default_tenant_direct_submit_creates_same_tenant_case_link() -> None:
session_factory = build_session_factory()
with session_factory() as db:
request = build_request(
run_id="application-direct-submit-tenant",
tenant_id="tenant-direct-submit",
)
service = UserAgentService(db)
submitted = service._create_expense_application_record(
request,
build_facts(),
submit=True,
learning_current_user=service._build_application_current_user(request),
)
link = db.scalar(
select(ExpenseCaseLink).where(
ExpenseCaseLink.resource_type == "expense_claim",
ExpenseCaseLink.resource_id == submitted.id,
)
)
event = db.scalar(
select(BusinessEvent).where(
BusinessEvent.aggregate_type == "expense_claim",
BusinessEvent.aggregate_id == submitted.id,
BusinessEvent.event_type == "application_submitted",
)
)
assert submitted.status == "submitted"
assert link is not None
assert link.tenant_id == "tenant-direct-submit"
assert event is not None
assert event.tenant_id == "tenant-direct-submit"
assert event.expense_case_id == link.expense_case_id
@pytest.mark.parametrize( @pytest.mark.parametrize(
"preview_patch", "preview_patch",
[ [

View File

@@ -3,6 +3,7 @@
} }
.application-memory-panel, .application-memory-panel,
.application-memory-notice-panel,
.application-learning-receipt-panel { .application-learning-receipt-panel {
margin-top: 14px; margin-top: 14px;
padding: 12px; padding: 12px;
@@ -12,6 +13,7 @@
} }
.application-memory-panel-head, .application-memory-panel-head,
.application-memory-notice-panel > header,
.application-learning-receipt-panel > header { .application-learning-receipt-panel > header {
display: flex; display: flex;
align-items: flex-start; align-items: flex-start;
@@ -37,6 +39,7 @@
} }
.application-memory-panel-head strong, .application-memory-panel-head strong,
.application-memory-notice-panel > header strong,
.application-learning-receipt-panel > header strong { .application-learning-receipt-panel > header strong {
color: #17324a; color: #17324a;
font-size: 13px; font-size: 13px;
@@ -103,6 +106,10 @@
color: #708196; color: #708196;
} }
.application-memory-conflict-note {
color: #8a5a12 !important;
}
.application-memory-forget-btn { .application-memory-forget-btn {
min-height: 44px; min-height: 44px;
display: inline-flex; display: inline-flex;
@@ -143,6 +150,43 @@
background: #f8fbff; background: #f8fbff;
} }
.application-memory-notice-panel {
display: grid;
gap: 8px;
border-color: #ead7a2;
background: #fffaf0;
}
.application-memory-notice-panel > header {
align-items: center;
color: #8a5a12;
}
.application-memory-notice-panel > header i {
font-size: 17px;
}
.application-memory-notice {
display: grid;
gap: 3px;
padding: 8px 10px;
border-left: 3px solid #d6a12a;
border-radius: 4px;
background: #ffffff;
}
.application-memory-notice strong {
color: #17324a;
font-size: 12px;
font-weight: 850;
}
.application-memory-notice span {
color: #52677a;
font-size: 11px;
line-height: 1.5;
}
.application-learning-receipt-panel > header { .application-learning-receipt-panel > header {
align-items: center; align-items: center;
color: var(--theme-primary-active, #255b7d); color: var(--theme-primary-active, #255b7d);

View File

@@ -0,0 +1,317 @@
.expense-memory-settings {
display: grid;
gap: 20px;
}
.expense-memory-safety-note {
display: flex;
align-items: flex-start;
gap: 9px;
margin: 18px 0;
padding: 11px 12px;
border: 1px solid #cfe3da;
border-radius: 8px;
background: #f3faf7;
color: #315f52;
font-size: 12px;
line-height: 1.6;
}
.expense-memory-safety-note i {
flex: 0 0 auto;
margin-top: 1px;
font-size: 17px;
}
.expense-memory-scope-switch {
display: grid;
grid-template-columns: repeat(2, minmax(0, 1fr));
gap: 10px;
margin-bottom: 18px;
}
.expense-memory-scope-switch button {
min-width: 0;
display: grid;
gap: 4px;
padding: 13px 14px;
border: 1px solid #d8e2ec;
border-radius: 8px;
background: #ffffff;
color: #52677a;
text-align: left;
cursor: pointer;
}
.expense-memory-scope-switch button:hover,
.expense-memory-scope-switch button.active {
border-color: var(--theme-primary, #3a7ca5);
background: color-mix(in srgb, var(--theme-primary-soft, #eaf4fa) 45%, #ffffff);
}
.expense-memory-scope-switch strong {
color: #17324a;
font-size: 13px;
}
.expense-memory-scope-switch span {
font-size: 11px;
line-height: 1.5;
}
.expense-memory-form-grid {
display: grid;
grid-template-columns: repeat(2, minmax(0, 1fr));
gap: 16px;
}
.expense-memory-form-grid .field-wide {
grid-column: 1 / -1;
}
.expense-memory-form-actions {
display: flex;
justify-content: flex-end;
margin-top: 18px;
}
.expense-memory-primary-button,
.expense-memory-refresh-button,
.expense-memory-item-actions button,
.expense-memory-inline-editor button {
min-height: 38px;
display: inline-flex;
align-items: center;
justify-content: center;
gap: 6px;
padding: 0 14px;
border: 1px solid #d6e1eb;
border-radius: 7px;
background: #ffffff;
color: #31556f;
font-size: 12px;
font-weight: 800;
cursor: pointer;
}
.expense-memory-primary-button {
border-color: var(--theme-primary, #3a7ca5);
background: var(--theme-primary, #3a7ca5);
color: #ffffff;
}
.expense-memory-primary-button:disabled,
.expense-memory-refresh-button:disabled,
.expense-memory-inline-editor button:disabled {
cursor: wait;
opacity: 0.58;
}
.expense-memory-list-head {
align-items: center;
}
.expense-memory-list-head > div:first-child {
min-width: 0;
}
.expense-memory-list-head h4,
.expense-memory-list-head p {
margin: 0;
}
.expense-memory-list-head p {
margin-top: 4px;
color: #64748b;
font-size: 12px;
}
.expense-memory-state {
min-height: 110px;
display: flex;
align-items: center;
justify-content: center;
gap: 8px;
color: #64748b;
font-size: 13px;
}
.expense-memory-state.is-error {
color: #b42318;
}
.expense-memory-list {
display: grid;
gap: 10px;
margin-top: 16px;
}
.expense-memory-item {
padding: 14px;
border: 1px solid #dce6ef;
border-radius: 9px;
background: #ffffff;
}
.expense-memory-item > header {
display: flex;
align-items: center;
justify-content: space-between;
gap: 12px;
}
.expense-memory-item > header > div {
min-width: 0;
display: flex;
align-items: center;
gap: 8px;
}
.expense-memory-item > header strong {
color: #17324a;
font-size: 13px;
}
.expense-memory-scope-badge,
.expense-memory-status {
flex: 0 0 auto;
padding: 3px 7px;
border-radius: 999px;
background: #edf4f8;
color: #31556f;
font-size: 10px;
font-weight: 850;
}
.expense-memory-scope-badge.is-enterprise {
background: #e8f1ff;
color: #245596;
}
.expense-memory-status.is-active {
background: #e8f7ef;
color: #287653;
}
.expense-memory-status.is-revoked,
.expense-memory-status.is-expired,
.expense-memory-status.is-suppressed {
background: #f1f3f5;
color: #65727e;
}
.expense-memory-item dl {
display: grid;
grid-template-columns: repeat(4, minmax(0, 1fr));
gap: 10px;
margin: 13px 0 0;
}
.expense-memory-item dl > div {
min-width: 0;
display: grid;
gap: 3px;
}
.expense-memory-item dt,
.expense-memory-item dd {
margin: 0;
}
.expense-memory-item dt {
color: #8190a0;
font-size: 10px;
}
.expense-memory-item dd {
color: #334155;
font-size: 11px;
font-weight: 750;
overflow-wrap: anywhere;
}
.expense-memory-item-actions {
display: flex;
justify-content: flex-end;
gap: 8px;
margin-top: 13px;
}
.expense-memory-item-actions button.danger,
.expense-memory-inline-editor button.danger {
border-color: #efc7c3;
color: #b42318;
}
.expense-memory-inline-editor {
display: grid;
grid-template-columns: minmax(140px, 0.6fr) minmax(120px, 0.4fr) minmax(220px, 1fr);
gap: 9px;
margin-top: 13px;
padding: 12px;
border: 1px solid #cfe0ee;
border-radius: 8px;
background: #f7fbfe;
}
.expense-memory-inline-editor.is-danger {
grid-template-columns: 1fr;
border-color: #efc7c3;
background: #fff8f7;
}
.expense-memory-inline-editor p {
margin: 0;
color: #7a3b35;
font-size: 11px;
line-height: 1.55;
}
.expense-memory-inline-editor input {
min-width: 0;
min-height: 38px;
padding: 0 10px;
border: 1px solid #d5e0ea;
border-radius: 7px;
color: #26394a;
font-size: 12px;
}
.expense-memory-inline-editor > div {
grid-column: 1 / -1;
display: flex;
justify-content: flex-end;
gap: 8px;
}
@media (max-width: 900px) {
.expense-memory-item dl {
grid-template-columns: repeat(2, minmax(0, 1fr));
}
.expense-memory-inline-editor {
grid-template-columns: 1fr;
}
}
@media (max-width: 640px) {
.expense-memory-scope-switch,
.expense-memory-form-grid,
.expense-memory-item dl {
grid-template-columns: 1fr;
}
.expense-memory-item > header {
align-items: flex-start;
}
.expense-memory-form-actions,
.expense-memory-item-actions {
justify-content: stretch;
}
.expense-memory-primary-button,
.expense-memory-item-actions button {
flex: 1;
}
}

View File

@@ -1,17 +1,20 @@
<template> <template>
<div v-if="appliedMemories.length || learningReceipts.length" class="application-memory-stack"> <div
v-if="appliedMemories.length || memoryNotices.length || learningReceipts.length"
class="application-memory-stack"
>
<section <section
v-if="appliedMemories.length" v-if="appliedMemories.length"
class="application-memory-panel" class="application-memory-panel"
aria-label="本次申请使用的历史偏好" aria-label="本次申请使用的智能记忆"
> >
<header class="application-memory-panel-head"> <header class="application-memory-panel-head">
<span class="application-memory-panel-icon" aria-hidden="true"> <span class="application-memory-panel-icon" aria-hidden="true">
<i class="mdi mdi-history"></i> <i class="mdi mdi-history"></i>
</span> </span>
<span> <span>
<strong>个性化填充</strong> <strong>智能记忆填充</strong>
<small>以下内容来自已确认的历史偏好您仍可直接修改上方字段</small> <small>企业制度部门基线和个人偏好会按优先级应用您仍可直接修改上方字段</small>
</span> </span>
</header> </header>
<div class="application-memory-list"> <div class="application-memory-list">
@@ -23,13 +26,17 @@
<div class="application-memory-copy"> <div class="application-memory-copy">
<span class="application-memory-status"> <span class="application-memory-status">
<i class="mdi mdi-check-circle-outline" aria-hidden="true"></i> <i class="mdi mdi-check-circle-outline" aria-hidden="true"></i>
已按历史偏好填入 {{ resolveMemorySourceTitle(memory) }}
</span> </span>
<strong>{{ memory.fieldLabel }}{{ memory.value }}</strong> <strong>{{ memory.fieldLabel }}{{ memory.value }}</strong>
<small v-if="memory.message">{{ memory.message }}</small> <small v-if="memory.message">{{ memory.message }}</small>
<em>{{ ui.resolveApplicationMemoryEvidenceText(memory) }}</em> <small v-if="memory.conflicts.length" class="application-memory-conflict-note">
已按适用优先级选择当前设置其他层级设置未自动采用
</small>
<em>{{ resolveMemorySupportingText(memory) }}</em>
</div> </div>
<button <button
v-if="canForgetMemory(memory)"
type="button" type="button"
class="application-memory-forget-btn" class="application-memory-forget-btn"
:disabled="ui.isForgettingApplicationMemory(memory)" :disabled="ui.isForgettingApplicationMemory(memory)"
@@ -46,6 +53,27 @@
</div> </div>
</section> </section>
<section
v-if="memoryNotices.length"
class="application-memory-notice-panel"
role="status"
aria-live="polite"
aria-label="本次未自动应用的记忆说明"
>
<header>
<i class="mdi mdi-information-outline" aria-hidden="true"></i>
<strong>需要您确认的设置</strong>
</header>
<article
v-for="memory in memoryNotices"
:key="`${message.id}-memory-notice-${memory.memoryId}`"
class="application-memory-notice"
>
<strong>{{ memory.fieldLabel }}</strong>
<span>{{ resolveMemoryNoticeText(memory) }}</span>
</article>
</section>
<section <section
v-if="learningReceipts.length" v-if="learningReceipts.length"
class="application-learning-receipt-panel" class="application-learning-receipt-panel"
@@ -75,6 +103,16 @@
<script setup> <script setup>
import { computed } from 'vue' import { computed } from 'vue'
import { resolveAiApplicationPreviewMemories } from '../../services/aiApplicationPreviewActions.js'
const MEMORY_NOTICE_STATUSES = new Set([
'conflict',
'conflicted',
'not_applied',
'overridden',
'skipped',
'suppressed'
])
const props = defineProps({ const props = defineProps({
message: { message: {
@@ -88,7 +126,44 @@ const props = defineProps({
}) })
const appliedMemories = computed(() => props.ui.resolveAppliedApplicationPreviewMemories(props.message)) const appliedMemories = computed(() => props.ui.resolveAppliedApplicationPreviewMemories(props.message))
const normalizedMemories = computed(() => (
resolveAiApplicationPreviewMemories(props.message?.applicationPreview || {})
))
const memoryNotices = computed(() => normalizedMemories.value.filter((memory) => (
memory.status !== 'applied'
&& (memory.conflicts.length > 0 || MEMORY_NOTICE_STATUSES.has(memory.status))
)))
const learningReceipts = computed(() => props.ui.resolveVisibleApplicationLearningReceipts(props.message)) const learningReceipts = computed(() => props.ui.resolveVisibleApplicationLearningReceipts(props.message))
function resolveMemorySourceTitle(memory = {}) {
if (memory.scopeType === 'enterprise') return '企业制度填充'
if (memory.scopeType === 'department') return '部门基线填充'
return '个性化填充'
}
function resolveMemorySupportingText(memory = {}) {
if (memory.scopeType === 'enterprise') {
return '由企业管理员维护,仅影响自动填充。'
}
if (memory.scopeType === 'department') {
return '由当前部门的适用设置提供,仅影响自动填充。'
}
return props.ui.resolveApplicationMemoryEvidenceText(memory)
}
function canForgetMemory(memory = {}) {
return memory.scopeType === 'user' && memory.canRevoke === true
}
function resolveMemoryNoticeText(memory = {}) {
if (memory.conflicts.length > 0 || ['conflict', 'conflicted'].includes(memory.status)) {
return '检测到同层级设置冲突,本次未自动填入,请按实际情况确认。'
}
if (['overridden', 'suppressed'].includes(memory.status)) {
return '存在更高优先级的适用设置,本次未使用该项。'
}
return '该项本次未自动填入,请按实际情况确认。'
}
</script> </script>
<style scoped src="../../assets/styles/components/travel-reimbursement-message-memory.css"></style> <style scoped src="../../assets/styles/components/travel-reimbursement-message-memory.css"></style>

View File

@@ -19,6 +19,43 @@ const AI_APPLICATION_MEMORY_FIELD_LABELS = {
transport_mode: '出行方式', transport_mode: '出行方式',
transportMode: '出行方式' transportMode: '出行方式'
} }
const AI_APPLICATION_MEMORY_SCOPE_DEFAULTS = {
enterprise: {
label: '企业制度',
priority: 300,
source: 'admin_managed_org'
},
department: {
label: '部门基线',
priority: 200,
source: 'admin_managed_org'
},
user: {
label: '个人偏好',
priority: 100,
source: 'verified_user_history'
}
}
function normalizeOptionalNumber(value) {
if (value === '' || value === null || value === undefined) return null
const normalized = Number(value)
return Number.isFinite(normalized) ? normalized : null
}
function normalizeBoolean(value, fallback = false) {
if (value === true || value === 1 || value === 'true' || value === '1') return true
if (value === false || value === 0 || value === 'false' || value === '0') return false
return fallback
}
function normalizeMemoryConflicts(value) {
if (!Array.isArray(value)) return []
return value.filter((item) => {
if (item && typeof item === 'object') return true
return Boolean(normalizeText(item))
})
}
export function normalizeAiApplicationMemory(item = {}) { export function normalizeAiApplicationMemory(item = {}) {
if (!item || typeof item !== 'object') return null if (!item || typeof item !== 'object') return null
@@ -26,6 +63,12 @@ export function normalizeAiApplicationMemory(item = {}) {
const fieldKey = normalizeText(item.fieldKey || item.field_key) const fieldKey = normalizeText(item.fieldKey || item.field_key)
const rawStatus = normalizeText(item.status).toLowerCase() const rawStatus = normalizeText(item.status).toLowerCase()
const status = rawStatus === 'active' ? 'applied' : rawStatus const status = rawStatus === 'active' ? 'applied' : rawStatus
const scopeType = normalizeText(item.scopeType || item.scope_type).toLowerCase() || 'user'
const scopeDefaults = AI_APPLICATION_MEMORY_SCOPE_DEFAULTS[scopeType]
|| AI_APPLICATION_MEMORY_SCOPE_DEFAULTS.user
const explicitCanRevoke = Object.prototype.hasOwnProperty.call(item, 'canRevoke')
? item.canRevoke
: item.can_revoke
if (!memoryId || !fieldKey || !status) return null if (!memoryId || !fieldKey || !status) return null
return { return {
memoryId, memoryId,
@@ -39,6 +82,17 @@ export function normalizeAiApplicationMemory(item = {}) {
evidenceCount: Number(item.evidenceCount ?? item.evidence_count ?? 0) || 0, evidenceCount: Number(item.evidenceCount ?? item.evidence_count ?? 0) || 0,
approvedEvidenceCount: approvedEvidenceCount:
Number(item.approvedEvidenceCount ?? item.approved_evidence_count ?? 0) || 0, Number(item.approvedEvidenceCount ?? item.approved_evidence_count ?? 0) || 0,
scopeType,
scopeId: normalizeText(item.scopeId || item.scope_id),
scopeLabel: normalizeText(item.scopeLabel || item.scope_label) || scopeDefaults.label,
source: normalizeText(item.source) || scopeDefaults.source,
priority:
normalizeOptionalNumber(item.priority) ?? scopeDefaults.priority,
effectiveConfidence: normalizeOptionalNumber(
item.effectiveConfidence ?? item.effective_confidence
),
conflicts: normalizeMemoryConflicts(item.conflicts),
canRevoke: normalizeBoolean(explicitCanRevoke, scopeType === 'user'),
message: normalizeText(item.message) message: normalizeText(item.message)
} }
} }

View File

@@ -0,0 +1,167 @@
import { apiRequest } from './api.js'
export const EXPENSE_MEMORY_TRANSPORT_OPTIONS = ['飞机', '火车', '轮船']
export const EXPENSE_MEMORY_MIN_DAYS = 30
export const EXPENSE_MEMORY_MAX_DAYS = 365
function normalizeText(value) {
return String(value || '').trim()
}
function normalizeInteger(value, fallback = 0) {
const normalized = Number.parseInt(String(value ?? ''), 10)
return Number.isFinite(normalized) ? normalized : fallback
}
function requireReason(value) {
const reason = normalizeText(value)
if (!reason) throw new Error('请填写操作理由。')
return reason
}
function requireTransportValue(value) {
const normalized = normalizeText(value)
if (!EXPENSE_MEMORY_TRANSPORT_OPTIONS.includes(normalized)) {
throw new Error('出行方式仅支持飞机、火车或轮船。')
}
return normalized
}
function requireExpiresInDays(value) {
const normalized = normalizeInteger(value)
if (normalized < EXPENSE_MEMORY_MIN_DAYS || normalized > EXPENSE_MEMORY_MAX_DAYS) {
throw new Error('有效期必须在 30 到 365 天之间。')
}
return normalized
}
function requireExpectedGeneration(value) {
const normalized = normalizeInteger(value)
if (normalized < 1) throw new Error('缺少有效的记忆版本。')
return normalized
}
function requireRequestId(value) {
const normalized = normalizeText(value)
if (normalized.length < 8) throw new Error('缺少有效的幂等请求标识。')
return normalized
}
export function normalizeOrganizationExpenseMemory(item = {}) {
if (!item || typeof item !== 'object') return null
const id = normalizeText(item.id)
const scopeType = normalizeText(item.scopeType || item.scope_type).toLowerCase()
if (!id || !['enterprise', 'department'].includes(scopeType)) return null
const originType = normalizeText(item.originType || item.origin_type)
return {
id,
fieldKey: normalizeText(item.fieldKey || item.field_key) || 'transport_mode',
value: normalizeText(item.value),
status: normalizeText(item.status).toLowerCase(),
scopeType,
scopeId: normalizeText(item.scopeId || item.scope_id),
scopeLabel: normalizeText(item.scopeLabel || item.scope_label)
|| (scopeType === 'enterprise' ? '企业统一规则' : '部门规则'),
source: normalizeText(item.source),
originType,
sourceLabel: originType === 'admin_managed' ? '管理员显式维护' : '组织规则',
generation: Math.max(1, normalizeInteger(item.generation, 1)),
policyVersion: normalizeText(item.policyVersion || item.policy_version),
validFrom: normalizeText(item.validFrom || item.valid_from),
expiresAt: normalizeText(item.expiresAt || item.expires_at),
managedAt: normalizeText(item.managedAt || item.managed_at),
canRevoke: Boolean(item.canRevoke ?? item.can_revoke),
createdAt: normalizeText(item.createdAt || item.created_at),
updatedAt: normalizeText(item.updatedAt || item.updated_at)
}
}
export function normalizeOrganizationExpenseMemories(payload = {}) {
const items = Array.isArray(payload) ? payload : payload?.items
return (Array.isArray(items) ? items : [])
.map((item) => normalizeOrganizationExpenseMemory(item))
.filter(Boolean)
}
export function buildOrganizationExpenseMemoryCreatePayload(input = {}) {
const scopeType = normalizeText(input.scopeType || input.scope_type).toLowerCase()
if (!['enterprise', 'department'].includes(scopeType)) {
throw new Error('请选择企业或部门作用域。')
}
const payload = {
scope_type: scopeType,
value: requireTransportValue(input.value),
expires_in_days: requireExpiresInDays(input.expiresInDays ?? input.expires_in_days),
reason: requireReason(input.reason),
request_id: requireRequestId(input.requestId || input.request_id)
}
if (scopeType === 'department') {
const scopeId = normalizeText(input.scopeId || input.scope_id)
if (!scopeId) throw new Error('请选择要适用的部门。')
payload.scope_id = scopeId
}
return payload
}
export function buildOrganizationExpenseMemoryUpdatePayload(input = {}) {
return {
value: requireTransportValue(input.value),
expires_in_days: requireExpiresInDays(input.expiresInDays ?? input.expires_in_days),
expected_generation: requireExpectedGeneration(
input.expectedGeneration ?? input.expected_generation
),
reason: requireReason(input.reason),
request_id: requireRequestId(input.requestId || input.request_id)
}
}
export function buildOrganizationExpenseMemoryRevokePayload(input = {}) {
return {
expected_generation: requireExpectedGeneration(
input.expectedGeneration ?? input.expected_generation
),
reason: requireReason(input.reason),
request_id: requireRequestId(input.requestId || input.request_id)
}
}
export async function fetchOrganizationExpenseMemories(options = {}) {
const payload = await apiRequest('/expense-application-memories/organization', options)
return normalizeOrganizationExpenseMemories(payload)
}
export async function createOrganizationExpenseMemory(input = {}, options = {}) {
const payload = await apiRequest('/expense-application-memories/organization', {
method: 'POST',
body: JSON.stringify(buildOrganizationExpenseMemoryCreatePayload(input)),
...options
})
return normalizeOrganizationExpenseMemory(payload)
}
export async function updateOrganizationExpenseMemory(memoryId, input = {}, options = {}) {
const id = normalizeText(memoryId)
if (!id) throw new Error('缺少要更新的记忆标识。')
const payload = await apiRequest(
`/expense-application-memories/organization/${encodeURIComponent(id)}`,
{
method: 'PUT',
body: JSON.stringify(buildOrganizationExpenseMemoryUpdatePayload(input)),
...options
}
)
return normalizeOrganizationExpenseMemory(payload)
}
export function revokeOrganizationExpenseMemory(memoryId, input = {}, options = {}) {
const id = normalizeText(memoryId)
if (!id) return Promise.reject(new Error('缺少要撤销的记忆标识。'))
return apiRequest(
`/expense-application-memories/organization/${encodeURIComponent(id)}/revoke`,
{
method: 'POST',
body: JSON.stringify(buildOrganizationExpenseMemoryRevokePayload(input)),
...options
}
)
}

View File

@@ -43,6 +43,14 @@ export const SECTION_DEFINITIONS = [
longDesc: '统一配置智能体会话的保留天数。超过保留期的历史会话会在后端清理,避免上下文和管理成本无限增长。', longDesc: '统一配置智能体会话的保留天数。超过保留期的历史会话会在后端清理,避免上下文和管理成本无限增长。',
actionLabel: '保存会话设置' actionLabel: '保存会话设置'
}, },
{
id: 'aiMemory',
label: 'AI 记忆',
title: 'AI 记忆策略',
desc: '企业与部门智能填充',
longDesc: '显式维护企业和部门的费用申请记忆,统一控制作用域、版本与有效期。',
actionLabel: ''
},
{ {
id: 'hermes', id: 'hermes',
label: '数字员工设置', label: '数字员工设置',
@@ -607,6 +615,7 @@ export function computeSectionStatus(state) {
Number(state.sessionForm.conversationRetentionDays) >= 1 && Number(state.sessionForm.conversationRetentionDays) >= 1 &&
Number(state.sessionForm.conversationRetentionDays) <= 10 Number(state.sessionForm.conversationRetentionDays) <= 10
), ),
aiMemory: true,
hermes: isHermesEmployeeSettingsReady(state.hermesForm), hermes: isHermesEmployeeSettingsReady(state.hermesForm),
llm: (() => { llm: (() => {
const rows = normalizeLlmModelRows(state.llmForm.models) const rows = normalizeLlmModelRows(state.llmForm.models)

View File

@@ -0,0 +1,357 @@
<template>
<div class="expense-memory-settings">
<section class="settings-card expense-memory-create-card">
<div class="card-head">
<div class="card-title-with-icon">
<div class="model-icon-box slate">
<i class="mdi mdi-brain"></i>
</div>
<div>
<h4>组织费用申请记忆</h4>
<p>为企业或指定部门配置出行方式基线仅用于申请表单的智能填充</p>
</div>
</div>
</div>
<div class="expense-memory-safety-note">
<i class="mdi mdi-shield-check-outline" aria-hidden="true"></i>
<span>只保存飞机火车轮船三种低敏枚举值不保存金额事由或客户信息</span>
</div>
<div class="expense-memory-scope-switch" aria-label="选择记忆作用域">
<button
v-for="option in scopeOptions"
:key="option.value"
type="button"
:class="{ active: createForm.scopeType === option.value }"
@click="selectScope(option.value)"
>
<strong>{{ option.label }}</strong>
<span>{{ option.desc }}</span>
</button>
</div>
<div class="expense-memory-form-grid">
<label v-if="createForm.scopeType === 'department'" class="field">
<span><em>*</em> 适用部门</span>
<EnterpriseSelect
v-model="createForm.scopeId"
:options="departmentOptions"
filterable
placeholder="选择稳定部门 ID"
/>
<small>部门来自员工组织元数据不使用可变的部门名称作为关联键</small>
</label>
<label class="field">
<span><em>*</em> 默认出行方式</span>
<EnterpriseSelect v-model="createForm.value" :options="transportOptions" />
</label>
<label class="field">
<span><em>*</em> 有效天数</span>
<input v-model.number="createForm.expiresInDays" type="number" min="30" max="365" />
<small>30 365 到期后自动停止填充</small>
</label>
<label class="field field-wide">
<span><em>*</em> 操作理由</span>
<input
v-model="createForm.reason"
type="text"
maxlength="255"
placeholder="例如:统一当前差旅出行基线"
/>
<small>理由用于后端审计不会向普通员工展示</small>
</label>
</div>
<div class="expense-memory-form-actions">
<button
class="expense-memory-primary-button"
type="button"
:disabled="busyKey === 'create'"
@click="createMemory"
>
<i :class="busyKey === 'create' ? 'mdi mdi-loading mdi-spin' : 'mdi mdi-plus-circle-outline'"></i>
<span>{{ busyKey === 'create' ? '保存中...' : '新增设置' }}</span>
</button>
</div>
</section>
<section class="settings-card expense-memory-list-card">
<div class="card-head expense-memory-list-head">
<div>
<h4>当前设置与版本记录</h4>
<p>更新会创建新版本旧版本保留为可追溯记录</p>
</div>
<button class="expense-memory-refresh-button" type="button" :disabled="loading" @click="loadPanel">
<i :class="loading ? 'mdi mdi-loading mdi-spin' : 'mdi mdi-refresh'"></i>
<span>刷新</span>
</button>
</div>
<div v-if="loading" class="expense-memory-state" role="status">
<i class="mdi mdi-loading mdi-spin"></i>
<span>正在读取组织记忆...</span>
</div>
<div v-else-if="loadError" class="expense-memory-state is-error" role="alert">
<i class="mdi mdi-alert-circle-outline"></i>
<span>{{ loadError }}</span>
</div>
<div v-else-if="!memories.length" class="expense-memory-state">
<i class="mdi mdi-brain"></i>
<span>尚未配置企业或部门记忆</span>
</div>
<div v-else class="expense-memory-list">
<article v-for="memory in memories" :key="memory.id" class="expense-memory-item">
<header>
<div>
<span class="expense-memory-scope-badge" :class="`is-${memory.scopeType}`">
{{ memory.scopeType === 'enterprise' ? '企业' : '部门' }}
</span>
<strong>{{ memory.scopeLabel }}</strong>
</div>
<span class="expense-memory-status" :class="`is-${memory.status}`">
{{ resolveStatusLabel(memory.status) }}
</span>
</header>
<dl>
<div><dt>出行基线</dt><dd>{{ memory.value || '未设置' }}</dd></div>
<div><dt>版本</dt><dd> {{ memory.generation }} </dd></div>
<div><dt>有效期</dt><dd>{{ formatDateTime(memory.expiresAt) }}</dd></div>
<div><dt>维护来源</dt><dd>{{ memory.sourceLabel }}</dd></div>
</dl>
<div v-if="memory.status === 'active'" class="expense-memory-item-actions">
<button type="button" @click="startEdit(memory)">更新版本</button>
<button class="danger" type="button" @click="startRevoke(memory)">撤销</button>
</div>
<div v-if="editingId === memory.id" class="expense-memory-inline-editor">
<EnterpriseSelect v-model="editForm.value" :options="transportOptions" />
<input v-model.number="editForm.expiresInDays" type="number" min="30" max="365" aria-label="更新后有效天数" />
<input v-model="editForm.reason" type="text" maxlength="255" placeholder="请填写更新理由" />
<div>
<button type="button" :disabled="busyKey === memory.id" @click="saveUpdate(memory)">保存新版本</button>
<button type="button" @click="closeInlineActions">取消</button>
</div>
</div>
<div v-if="revokingId === memory.id" class="expense-memory-inline-editor is-danger">
<p>撤销后该作用域将停止自动填充已生成的申请不受影响</p>
<input v-model="revokeReason" type="text" maxlength="255" placeholder="请填写撤销理由" />
<div>
<button class="danger" type="button" :disabled="busyKey === memory.id" @click="confirmRevoke(memory)">确认撤销</button>
<button type="button" @click="closeInlineActions">取消</button>
</div>
</div>
</article>
</div>
</section>
</div>
</template>
<script setup>
import { onMounted, reactive, ref, watch } from 'vue'
import EnterpriseSelect from '../components/shared/EnterpriseSelect.vue'
import { useToast } from '../composables/useToast.js'
import { fetchEmployeeMeta } from '../services/employees.js'
import {
EXPENSE_MEMORY_TRANSPORT_OPTIONS,
createOrganizationExpenseMemory,
fetchOrganizationExpenseMemories,
revokeOrganizationExpenseMemory,
updateOrganizationExpenseMemory
} from '../services/expenseApplicationMemories.js'
const { toast } = useToast()
const scopeOptions = [
{ value: 'enterprise', label: '企业统一', desc: '对当前租户内所有员工生效' },
{ value: 'department', label: '指定部门', desc: '仅对选定组织单元生效' }
]
const transportOptions = EXPENSE_MEMORY_TRANSPORT_OPTIONS.map((value) => ({ value, label: value }))
const memories = ref([])
const departmentOptions = ref([])
const loading = ref(false)
const loadError = ref('')
const busyKey = ref('')
const editingId = ref('')
const revokingId = ref('')
const revokeReason = ref('')
const revokeRequestId = ref('')
const createForm = reactive({
scopeType: 'enterprise',
scopeId: '',
value: '火车',
expiresInDays: 180,
reason: '',
requestId: ''
})
const editForm = reactive({ value: '火车', expiresInDays: 180, reason: '', requestId: '' })
function createRequestId(prefix) {
const suffix = globalThis.crypto?.randomUUID?.()
|| `${Date.now().toString(36)}-${Math.random().toString(36).slice(2)}`
return `${prefix}:${suffix}`
}
watch(
() => [
createForm.scopeType,
createForm.scopeId,
createForm.value,
createForm.expiresInDays,
createForm.reason
],
() => { createForm.requestId = '' }
)
watch(
() => [editForm.value, editForm.expiresInDays, editForm.reason],
() => { editForm.requestId = '' }
)
watch(revokeReason, () => { revokeRequestId.value = '' })
function normalizeDepartmentOptions(payload = {}) {
const options = Array.isArray(payload?.organizationOptions) ? payload.organizationOptions : []
return options
.filter((item) => item?.id && item?.unitType === 'department')
.map((item) => ({
value: String(item.id),
label: `${String(item.name || '未命名部门')}${String(item.code || item.id)}`
}))
}
async function loadPanel() {
loading.value = true
loadError.value = ''
const [memoryResult, employeeMetaResult] = await Promise.allSettled([
fetchOrganizationExpenseMemories(),
fetchEmployeeMeta()
])
if (memoryResult.status === 'fulfilled') {
memories.value = memoryResult.value
} else {
memories.value = []
loadError.value = memoryResult.reason?.message || '组织记忆加载失败,请稍后重试。'
}
departmentOptions.value = employeeMetaResult.status === 'fulfilled'
? normalizeDepartmentOptions(employeeMetaResult.value)
: []
loading.value = false
}
async function refreshMemories() {
memories.value = await fetchOrganizationExpenseMemories()
}
function selectScope(scopeType) {
createForm.scopeType = scopeType
if (scopeType === 'enterprise') createForm.scopeId = ''
}
async function createMemory() {
busyKey.value = 'create'
try {
createForm.requestId ||= createRequestId('organization-memory-create')
await createOrganizationExpenseMemory(createForm)
createForm.reason = ''
createForm.requestId = ''
await refreshMemories()
toast('组织记忆已保存。')
} catch (error) {
toast(error?.message || '组织记忆保存失败。')
} finally {
busyKey.value = ''
}
}
function remainingDays(expiresAt) {
const target = new Date(expiresAt).getTime()
const days = Math.ceil((target - Date.now()) / 86400000)
if (!Number.isFinite(days)) return 180
return Math.max(30, Math.min(365, days))
}
function startEdit(memory) {
editingId.value = memory.id
revokingId.value = ''
editForm.value = memory.value
editForm.expiresInDays = remainingDays(memory.expiresAt)
editForm.reason = ''
editForm.requestId = ''
}
function startRevoke(memory) {
revokingId.value = memory.id
editingId.value = ''
revokeReason.value = ''
revokeRequestId.value = ''
editForm.requestId = ''
}
function closeInlineActions() {
editingId.value = ''
revokingId.value = ''
revokeReason.value = ''
revokeRequestId.value = ''
}
async function saveUpdate(memory) {
busyKey.value = memory.id
try {
editForm.requestId ||= createRequestId(`organization-memory-update:${memory.id}`)
await updateOrganizationExpenseMemory(memory.id, {
...editForm,
expectedGeneration: memory.generation
})
closeInlineActions()
await refreshMemories()
toast('组织记忆新版本已生效。')
} catch (error) {
toast(error?.message || '组织记忆更新失败。')
await refreshMemories().catch(() => {})
} finally {
busyKey.value = ''
}
}
async function confirmRevoke(memory) {
busyKey.value = memory.id
try {
revokeRequestId.value ||= createRequestId(`organization-memory-revoke:${memory.id}`)
await revokeOrganizationExpenseMemory(memory.id, {
expectedGeneration: memory.generation,
reason: revokeReason.value,
requestId: revokeRequestId.value
})
closeInlineActions()
await refreshMemories()
toast('组织记忆已撤销。')
} catch (error) {
toast(error?.message || '组织记忆撤销失败。')
await refreshMemories().catch(() => {})
} finally {
busyKey.value = ''
}
}
function resolveStatusLabel(status) {
return { active: '生效中', suppressed: '已被新版替换', revoked: '已撤销', expired: '已过期' }[status] || '已停用'
}
function formatDateTime(value) {
const date = new Date(value)
if (!value || Number.isNaN(date.getTime())) return '未设置'
return new Intl.DateTimeFormat('zh-CN', { dateStyle: 'medium', timeStyle: 'short' }).format(date)
}
onMounted(loadPanel)
</script>
<style scoped src="../assets/styles/views/expense-memory-settings-panel.css"></style>

View File

@@ -324,7 +324,7 @@
</section> </section>
</template> </template>
<template v-else-if="activeSection === 'hermes'"> <template v-else-if="activeSection === 'hermes'">
<HermesEmployeeSettingsPanel <HermesEmployeeSettingsPanel
:hermes-form="pageState.hermesForm" :hermes-form="pageState.hermesForm"
@toggle-master="toggleHermesMaster" @toggle-master="toggleHermesMaster"
@@ -332,9 +332,13 @@
@toggle-task="toggleHermesTask" @toggle-task="toggleHermesTask"
@update-task-time="updateHermesTaskTime" @update-task-time="updateHermesTaskTime"
/> />
</template> </template>
<template v-else-if="activeSection === 'llm'"> <template v-else-if="activeSection === 'aiMemory'">
<ExpenseMemorySettingsPanel />
</template>
<template v-else-if="activeSection === 'llm'">
<LlmSettingsPanel :llm-form="pageState.llmForm" :provider-options="providerOptions" /> <LlmSettingsPanel :llm-form="pageState.llmForm" :provider-options="providerOptions" />
</template> </template>

View File

@@ -1,3 +1,4 @@
import ExpenseMemorySettingsPanel from '../ExpenseMemorySettingsPanel.vue'
import HermesEmployeeSettingsPanel from '../HermesEmployeeSettingsPanel.vue' import HermesEmployeeSettingsPanel from '../HermesEmployeeSettingsPanel.vue'
import LlmSettingsPanel from '../LlmSettingsPanel.vue' import LlmSettingsPanel from '../LlmSettingsPanel.vue'
import LogDetailView from '../LogDetailView.vue' import LogDetailView from '../LogDetailView.vue'
@@ -9,6 +10,7 @@ import { useSettings } from '../../composables/useSettings.js'
export default { export default {
name: 'SettingsView', name: 'SettingsView',
components: { components: {
ExpenseMemorySettingsPanel,
HermesEmployeeSettingsPanel, HermesEmployeeSettingsPanel,
EnterpriseSelect, EnterpriseSelect,
LlmSettingsPanel, LlmSettingsPanel,

View File

@@ -151,6 +151,10 @@ export function useTravelReimbursementApplicationPreviewActions({
async function forgetApplicationPreviewMemory(message, memory = {}) { async function forgetApplicationPreviewMemory(message, memory = {}) {
const normalizedMemory = normalizeAiApplicationMemory(memory) const normalizedMemory = normalizeAiApplicationMemory(memory)
if (!message?.applicationPreview || !normalizedMemory) return false if (!message?.applicationPreview || !normalizedMemory) return false
if (normalizedMemory.scopeType !== 'user' || normalizedMemory.canRevoke !== true) {
toast('企业或部门设置由管理员统一维护,当前申请仍可手动修改。')
return false
}
if (isForgettingApplicationMemory(normalizedMemory)) return false if (isForgettingApplicationMemory(normalizedMemory)) return false
if (!(await requestForgetMemoryConfirmation(normalizedMemory))) return false if (!(await requestForgetMemoryConfirmation(normalizedMemory))) return false

View File

@@ -4,6 +4,7 @@ import {
AI_APPLICATION_ACTION_SAVE_DRAFT, AI_APPLICATION_ACTION_SAVE_DRAFT,
AI_APPLICATION_ACTION_SUBMIT, AI_APPLICATION_ACTION_SUBMIT,
forgetAiApplicationMemory, forgetAiApplicationMemory,
normalizeAiApplicationMemory,
registerAiApplicationPreviewDecision, registerAiApplicationPreviewDecision,
resolveAiApplicationLearningReceipts, resolveAiApplicationLearningReceipts,
runAiApplicationPreviewAction runAiApplicationPreviewAction
@@ -148,6 +149,14 @@ async function testRegistrationUsesServerCanonicalPreview() {
status: 'active', status: 'active',
evidenceCount: 4, evidenceCount: 4,
approvedEvidenceCount: 3, approvedEvidenceCount: 3,
scope_type: 'department',
scope_id: 'department-shanghai-delivery',
scope_label: '上海交付部',
source: 'admin_managed_org',
priority: 200,
effective_confidence: 0.92,
conflicts: [{ scope_type: 'user', reason: 'lower_priority' }],
can_revoke: false,
message: '根据历史申请偏好填入' message: '根据历史申请偏好填入'
}] }]
} }
@@ -188,6 +197,14 @@ async function testRegistrationUsesServerCanonicalPreview() {
status: 'applied', status: 'applied',
evidenceCount: 4, evidenceCount: 4,
approvedEvidenceCount: 3, approvedEvidenceCount: 3,
scopeType: 'department',
scopeId: 'department-shanghai-delivery',
scopeLabel: '上海交付部',
source: 'admin_managed_org',
priority: 200,
effectiveConfidence: 0.92,
conflicts: [{ scope_type: 'user', reason: 'lower_priority' }],
canRevoke: false,
message: '根据历史申请偏好填入' message: '根据历史申请偏好填入'
}]) }])
} }
@@ -215,6 +232,14 @@ async function testLearningReceiptNormalizationAndForgetEndpoint() {
status: 'candidate', status: 'candidate',
evidenceCount: 2, evidenceCount: 2,
approvedEvidenceCount: 1, approvedEvidenceCount: 1,
scopeType: 'user',
scopeId: '',
scopeLabel: '个人偏好',
source: 'verified_user_history',
priority: 100,
effectiveConfidence: null,
conflicts: [],
canRevoke: true,
message: '再确认一次后可形成稳定偏好' message: '再确认一次后可形成稳定偏好'
}]) }])
@@ -238,6 +263,52 @@ async function testLearningReceiptNormalizationAndForgetEndpoint() {
assert.equal(capturedOptions.method, 'DELETE') assert.equal(capturedOptions.method, 'DELETE')
} }
function testHierarchicalMemoryNormalizationKeepsCompatibilityAndPermissions() {
const enterpriseMemory = normalizeAiApplicationMemory({
memory_id: 'memory-enterprise-transport',
field_key: 'transport_mode',
value: '火车',
status: 'applied',
scope_type: 'enterprise',
scope_id: 'tenant-acme',
scope_label: '企业统一制度',
source: 'admin_managed_org',
priority: '300',
effective_confidence: '0.88',
conflicts: [{ scope_type: 'department', reason: 'lower_priority' }],
can_revoke: false
})
assert.deepEqual(enterpriseMemory, {
memoryId: 'memory-enterprise-transport',
fieldKey: 'transport_mode',
fieldLabel: '出行方式',
value: '火车',
status: 'applied',
evidenceCount: 0,
approvedEvidenceCount: 0,
scopeType: 'enterprise',
scopeId: 'tenant-acme',
scopeLabel: '企业统一制度',
source: 'admin_managed_org',
priority: 300,
effectiveConfidence: 0.88,
conflicts: [{ scope_type: 'department', reason: 'lower_priority' }],
canRevoke: false,
message: ''
})
const legacyPersonalMemory = normalizeAiApplicationMemory({
memoryId: 'memory-legacy-user',
fieldKey: 'transport_mode',
value: '飞机',
status: 'active'
})
assert.equal(legacyPersonalMemory.scopeType, 'user')
assert.equal(legacyPersonalMemory.source, 'verified_user_history')
assert.equal(legacyPersonalMemory.canRevoke, true)
assert.equal(legacyPersonalMemory.status, 'applied')
}
async function testEditDraftActionCarriesClaimAndEditableFields() { async function testEditDraftActionCarriesClaimAndEditableFields() {
let capturedOptions = null let capturedOptions = null
@@ -327,6 +398,7 @@ async function run() {
await testSaveDraftActionUsesFastPreviewEndpoint() await testSaveDraftActionUsesFastPreviewEndpoint()
await testRegistrationUsesServerCanonicalPreview() await testRegistrationUsesServerCanonicalPreview()
await testLearningReceiptNormalizationAndForgetEndpoint() await testLearningReceiptNormalizationAndForgetEndpoint()
testHierarchicalMemoryNormalizationKeepsCompatibilityAndPermissions()
await testEditDraftActionCarriesClaimAndEditableFields() await testEditDraftActionCarriesClaimAndEditableFields()
await testApplicationActionSourceCanBeConfiguredWithoutChangingWorkbenchDefaults() await testApplicationActionSourceCanBeConfiguredWithoutChangingWorkbenchDefaults()
console.log('ai-application-preview-actions tests passed') console.log('ai-application-preview-actions tests passed')

View File

@@ -407,13 +407,48 @@ test('忘记偏好失败时不移除 memory也不改当前字段', async () =
assert.match(toasts.at(-1), /偏好服务暂不可用/) assert.match(toasts.at(-1), /偏好服务暂不可用/)
}) })
test('申请预览 UI 解释已应用偏好并提供可访问的忘记口', () => { test('组织记忆不允许普通用户调用忘记口', async () => {
const organizationMemory = {
memoryId: 'memory-enterprise-policy',
fieldKey: 'transport_mode',
fieldLabel: '出行方式',
value: '火车',
status: 'applied',
scopeType: 'enterprise',
canRevoke: false
}
const message = {
id: 'preview-enterprise-memory',
applicationPreview: createPreview({ memoryApplications: [organizationMemory] })
}
const { actions, toasts } = createActions()
let requestCount = 0
global.fetch = async () => {
requestCount += 1
throw new Error('不应发起组织记忆删除请求')
}
const forgotten = await actions.forgetApplicationPreviewMemory(message, organizationMemory)
assert.equal(forgotten, false)
assert.equal(requestCount, 0)
assert.equal(message.applicationPreview.memoryApplications.length, 1)
assert.match(toasts.at(-1), /管理员统一维护/)
})
test('申请预览 UI 解释分层记忆,仅个人记忆提供忘记入口', () => {
assert.match(messageItemTemplate, /<TravelReimbursementMemoryPanel :message="message" :ui="ui" \/>/) assert.match(messageItemTemplate, /<TravelReimbursementMemoryPanel :message="message" :ui="ui" \/>/)
assert.match(memoryPanelTemplate, /resolveAppliedApplicationPreviewMemories\(props\.message\)/) assert.match(memoryPanelTemplate, /resolveAppliedApplicationPreviewMemories\(props\.message\)/)
assert.match(memoryPanelTemplate, /已按历史偏好填入/) assert.match(memoryPanelTemplate, /企业制度填充/)
assert.match(memoryPanelTemplate, /部门基线填充/)
assert.match(memoryPanelTemplate, /个性化填充/)
assert.match(memoryPanelTemplate, /您仍可直接修改上方字段/) assert.match(memoryPanelTemplate, /您仍可直接修改上方字段/)
assert.match(memoryPanelTemplate, /忘记此偏好/) assert.match(memoryPanelTemplate, /忘记此偏好/)
assert.match(memoryPanelTemplate, /v-if="canForgetMemory\(memory\)"/)
assert.match(memoryPanelTemplate, /memory\.scopeType === 'user' && memory\.canRevoke === true/)
assert.match(memoryPanelTemplate, /:aria-busy="ui\.isForgettingApplicationMemory\(memory\)"/) assert.match(memoryPanelTemplate, /:aria-busy="ui\.isForgettingApplicationMemory\(memory\)"/)
assert.match(memoryPanelTemplate, /检测到同层级设置冲突/)
assert.doesNotMatch(memoryPanelTemplate, /JSON\.stringify\(memory\.conflicts/)
assert.match(memoryPanelTemplate, /resolveVisibleApplicationLearningReceipts\(props\.message\)/) assert.match(memoryPanelTemplate, /resolveVisibleApplicationLearningReceipts\(props\.message\)/)
assert.match(memoryPanelTemplate, /ui\.resolveApplicationLearningReceiptTitle\(receipt\)/) assert.match(memoryPanelTemplate, /ui\.resolveApplicationLearningReceiptTitle\(receipt\)/)
assert.doesNotMatch(messageItemTemplate, /已按历史偏好填入/) assert.doesNotMatch(messageItemTemplate, /已按历史偏好填入/)

View File

@@ -0,0 +1,190 @@
import assert from 'node:assert/strict'
import { readFileSync } from 'node:fs'
import test from 'node:test'
import {
buildOrganizationExpenseMemoryCreatePayload,
buildOrganizationExpenseMemoryRevokePayload,
buildOrganizationExpenseMemoryUpdatePayload,
createOrganizationExpenseMemory,
fetchOrganizationExpenseMemories,
normalizeOrganizationExpenseMemory,
revokeOrganizationExpenseMemory,
updateOrganizationExpenseMemory
} from '../src/services/expenseApplicationMemories.js'
const settingsModel = readFileSync(new URL('../src/utils/settingsModelHelper.js', import.meta.url), 'utf8')
const settingsView = readFileSync(new URL('../src/views/SettingsView.vue', import.meta.url), 'utf8')
const settingsScript = readFileSync(new URL('../src/views/scripts/SettingsView.js', import.meta.url), 'utf8')
const memoryPanel = readFileSync(new URL('../src/views/ExpenseMemorySettingsPanel.vue', import.meta.url), 'utf8')
test('组织记忆 payload 限制低敏枚举、有效期和必填理由', () => {
assert.deepEqual(buildOrganizationExpenseMemoryCreatePayload({
scopeType: 'enterprise',
scopeId: 'must-not-be-sent',
value: '火车',
expiresInDays: 180,
reason: '统一企业基线',
requestId: 'create-enterprise-1'
}), {
scope_type: 'enterprise',
value: '火车',
expires_in_days: 180,
reason: '统一企业基线',
request_id: 'create-enterprise-1'
})
assert.deepEqual(buildOrganizationExpenseMemoryCreatePayload({
scopeType: 'department',
scopeId: 'organization-unit-stable-id',
value: '飞机',
expiresInDays: 90,
reason: '部门差旅基线',
requestId: 'create-department-1'
}), {
scope_type: 'department',
scope_id: 'organization-unit-stable-id',
value: '飞机',
expires_in_days: 90,
reason: '部门差旅基线',
request_id: 'create-department-1'
})
assert.deepEqual(buildOrganizationExpenseMemoryUpdatePayload({
value: '轮船',
expiresInDays: 120,
expectedGeneration: 3,
reason: '更新制度',
requestId: 'update-memory-3'
}), {
value: '轮船',
expires_in_days: 120,
expected_generation: 3,
reason: '更新制度',
request_id: 'update-memory-3'
})
assert.deepEqual(buildOrganizationExpenseMemoryRevokePayload({
expectedGeneration: 3,
reason: '制度已停用',
requestId: 'revoke-memory-3'
}), {
expected_generation: 3,
reason: '制度已停用',
request_id: 'revoke-memory-3'
})
assert.throws(
() => buildOrganizationExpenseMemoryCreatePayload({
scopeType: 'enterprise', value: '网约车', expiresInDays: 180, reason: '不应通过', requestId: 'invalid-value-1'
}),
/仅支持飞机、火车或轮船/
)
assert.throws(
() => buildOrganizationExpenseMemoryCreatePayload({
scopeType: 'enterprise', value: '火车', expiresInDays: 29, reason: '不应通过', requestId: 'invalid-days-1'
}),
/30 到 365 天/
)
assert.throws(
() => buildOrganizationExpenseMemoryCreatePayload({
scopeType: 'enterprise', value: '火车', expiresInDays: 180, reason: '', requestId: 'invalid-reason-1'
}),
/操作理由/
)
})
test('组织记忆 service 使用管理员接口并保留作用域和版本解释', async () => {
const requests = []
const originalFetch = globalThis.fetch
const responseMemory = {
id: 'memory/enterprise 1',
field_key: 'transport_mode',
value: '火车',
status: 'active',
scope_type: 'enterprise',
scope_id: 'tenant-1',
scope_label: '企业统一规则',
source: 'enterprise_policy_memory',
origin_type: 'admin_managed',
generation: 2,
policy_version: 'expense_application_transport_org_memory.v1',
expires_at: '2026-12-31T00:00:00Z',
managed_by: 'private-admin-identifier',
can_revoke: true
}
globalThis.fetch = async (url, options = {}) => {
requests.push({ url: String(url), options })
return {
ok: true,
status: 200,
async json() {
return requests.length === 1 ? { items: [responseMemory] } : responseMemory
}
}
}
try {
const items = await fetchOrganizationExpenseMemories()
await createOrganizationExpenseMemory({
scopeType: 'enterprise', value: '火车', expiresInDays: 180, reason: '创建企业记忆', requestId: 'create-memory-1'
})
await updateOrganizationExpenseMemory('memory/enterprise 1', {
value: '轮船', expiresInDays: 120, expectedGeneration: 2, reason: '更新记忆', requestId: 'update-memory-1'
})
await revokeOrganizationExpenseMemory('memory/enterprise 1', {
expectedGeneration: 2, reason: '撤销记忆', requestId: 'revoke-memory-1'
})
assert.equal(items[0].scopeType, 'enterprise')
assert.equal(items[0].generation, 2)
assert.equal(items[0].sourceLabel, '管理员显式维护')
assert.equal('managedBy' in items[0], false)
} finally {
globalThis.fetch = originalFetch
}
assert.equal(requests[0].url, '/api/v1/expense-application-memories/organization')
assert.equal(requests[1].options.method, 'POST')
assert.equal(
requests[2].url,
'/api/v1/expense-application-memories/organization/memory%2Fenterprise%201'
)
assert.equal(requests[2].options.method, 'PUT')
assert.equal(
requests[3].url,
'/api/v1/expense-application-memories/organization/memory%2Fenterprise%201/revoke'
)
assert.equal(requests[3].options.method, 'POST')
})
test('组织记忆规范化不将管理员标识带入展示模型', () => {
const memory = normalizeOrganizationExpenseMemory({
id: 'memory-department-1',
scope_type: 'department',
scope_id: 'organization-unit-1',
scope_label: '部门规则(财务部)',
status: 'active',
value: '飞机',
origin_type: 'admin_managed',
managed_by: 'sensitive-user-id'
})
assert.equal(memory.scopeId, 'organization-unit-1')
assert.equal(memory.scopeLabel, '部门规则(财务部)')
assert.equal(memory.sourceLabel, '管理员显式维护')
assert.equal('managedBy' in memory, false)
})
test('设置页接入无顶部保存按钮的 AI 记忆独立面板', () => {
assert.match(settingsModel, /id:\s*'aiMemory'[\s\S]*label:\s*'AI 记忆'[\s\S]*actionLabel:\s*''/)
assert.match(settingsModel, /aiMemory:\s*true/)
assert.match(settingsScript, /import ExpenseMemorySettingsPanel/)
assert.match(settingsScript, /ExpenseMemorySettingsPanel,/)
assert.match(settingsView, /activeSection === 'aiMemory'[\s\S]*<ExpenseMemorySettingsPanel \/>/)
assert.match(memoryPanel, /fetchEmployeeMeta\(\)/)
assert.match(memoryPanel, /item\?\.id && item\?\.unitType === 'department'/)
assert.match(memoryPanel, /value:\s*String\(item\.id\)/)
assert.match(memoryPanel, /飞机[、\s\S]*火车[、\s\S]*轮船/)
assert.match(memoryPanel, /min="30" max="365"/)
assert.match(memoryPanel, /第 \{\{ memory\.generation \}\} 版/)
assert.match(memoryPanel, /memory\.sourceLabel/)
assert.match(memoryPanel, /organization-memory-revoke:/)
assert.match(memoryPanel, /requestId:\s*revokeRequestId\.value/)
assert.doesNotMatch(memoryPanel, /managedBy|managed_by|managementReason|management_reason/)
})